autoinference-utils 0.2.6__tar.gz → 0.2.7__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -150,6 +150,7 @@ activemq-data/
150
150
 
151
151
  # Environments
152
152
  .env
153
+ .env.dev
153
154
  .envrc
154
155
  .venv
155
156
  .autoinference/
@@ -241,6 +242,22 @@ Network Trash Folder
241
242
  Temporary Items
242
243
  .apdisk
243
244
 
244
- # Ephemeral files from Experiments
245
+ # Ephemeral files from benchmark runs
245
246
  benchmark_results/
246
247
  src/autoinference/deployments/
248
+
249
+ # Modal Skills
250
+ .agents/skills/modal/
251
+ .agents/skills/.modal-skill-*
252
+
253
+ # cook/etc is an extension point: local tools live there untracked.
254
+ packages/cook/src/cook/etc/*
255
+ # Local tools that existed before /etc are explicitly added back
256
+ !packages/cook/src/cook/etc/__init__.py
257
+ !packages/cook/src/cook/etc/tests/
258
+ !packages/cook/src/cook/etc/inspect_hf_model.py
259
+ !packages/cook/src/cook/etc/download.py
260
+ !packages/cook/src/cook/etc/quantize.py
261
+ !packages/cook/src/cook/etc/graft_mtp.py
262
+ !packages/cook/src/cook/etc/tune_moe.py
263
+ !packages/cook/src/cook/etc/tool_calling.py
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.5
2
2
  Name: autoinference-utils
3
- Version: 0.2.6
3
+ Version: 0.2.7
4
4
  Summary: Shared endpoint abstractions for autoinference deployments
5
5
  Requires-Python: >=3.10
6
6
  Description-Content-Type: text/markdown
@@ -9,6 +9,12 @@ Description-Content-Type: text/markdown
9
9
 
10
10
  Shared endpoint abstractions (`SGLangEndpoint`, `VLLMEndpoint`) for autoinference deployments.
11
11
 
12
+ For SGLang models whose custom Transformers code cannot load from symlinked
13
+ Hugging Face snapshots, set `AUTOINFERENCE_MATERIALIZE_MODEL_PATH=1`. The endpoint
14
+ copies local base/draft model code and metadata into unique temporary directories
15
+ and symlinks top-level safetensors weights. Repository IDs pass through unchanged.
16
+ Leave the setting unset to preserve the original paths.
17
+
12
18
  ## Publishing
13
19
 
14
20
  Bump `version` in `pyproject.toml` and merge to `main` — CI publishes automatically via trusted publishing.
@@ -0,0 +1,15 @@
1
+ # autoinference-utils
2
+
3
+ Shared endpoint abstractions (`SGLangEndpoint`, `VLLMEndpoint`) for autoinference deployments.
4
+
5
+ For SGLang models whose custom Transformers code cannot load from symlinked
6
+ Hugging Face snapshots, set `AUTOINFERENCE_MATERIALIZE_MODEL_PATH=1`. The endpoint
7
+ copies local base/draft model code and metadata into unique temporary directories
8
+ and symlinks top-level safetensors weights. Repository IDs pass through unchanged.
9
+ Leave the setting unset to preserve the original paths.
10
+
11
+ ## Publishing
12
+
13
+ Bump `version` in `pyproject.toml` and merge to `main` — CI publishes automatically via trusted publishing.
14
+
15
+ To publish manually: `cd autoinference_utils && uv build && uv publish`
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "autoinference-utils"
3
- version = "0.2.6"
3
+ version = "0.2.7"
4
4
  description = "Shared endpoint abstractions for autoinference deployments"
5
5
  readme = "README.md"
6
6
  requires-python = ">=3.10"
@@ -10,7 +10,7 @@ requires = ["hatchling"]
10
10
  build-backend = "hatchling.build"
11
11
 
12
12
  [tool.hatch.build.targets.wheel]
13
- packages = ["autoinference_utils"]
13
+ packages = ["src/autoinference_utils"]
14
14
 
15
15
  [tool.hatch.build.targets.sdist]
16
- include = ["autoinference_utils", "README.md", "pyproject.toml"]
16
+ include = ["src", "README.md", "pyproject.toml"]
@@ -23,11 +23,13 @@ import urllib.request
23
23
  from abc import ABC
24
24
  from typing import Any, Callable, Literal, Mapping, Optional, Sequence
25
25
 
26
+ from autoinference_utils.model_path import materialize_model_path
27
+
26
28
  BENCH_MODE_ENV = "ENDPOINT_BENCH_MODE"
27
29
  DUMMY_WEIGHTS_ENV = "ENDPOINT_DUMMY"
28
30
  BENCH_MODE_PORT_OFFSET = 10000
29
31
 
30
- MODAL_FLASH_REQUEST_UUID_HEADER= "modal-flash-request-uuid"
32
+ MODAL_FLASH_REQUEST_UUID_HEADER = "modal-flash-request-uuid"
31
33
  MODAL_SESSION_ID_HEADER = "modal-session-id"
32
34
 
33
35
  ENDPOINTS_REQUIRING_OAI_STREAM_COMPAT: tuple[str, ...] = ("TRTLLMEndpoint",)
@@ -55,9 +57,7 @@ def _flag_name(key: str) -> str:
55
57
  return key.partition("=")[0]
56
58
 
57
59
 
58
- def get_server_arg(
59
- server_args: Mapping[str, str], flag: str
60
- ) -> tuple[str, str] | None:
60
+ def get_server_arg(server_args: Mapping[str, str], flag: str) -> tuple[str, str] | None:
61
61
  """Return the key and value setting a flag under either spelling."""
62
62
  for key, value in server_args.items():
63
63
  if _flag_name(key) == flag:
@@ -153,11 +153,15 @@ class SGLangEndpoint(Endpoint):
153
153
  super().__init__(base_url=f"http://localhost:{sglang_port}")
154
154
 
155
155
  self.worker_port = sglang_port
156
- self.model_path = model_path
156
+ self.model_path = materialize_model_path(model_path)
157
157
  self.tp = tp
158
158
  self.ep = ep
159
159
  self.dp = dp
160
- self.speculative_model_path = speculative_model_path
160
+ self.speculative_model_path = (
161
+ materialize_model_path(speculative_model_path)
162
+ if speculative_model_path is not None
163
+ else None
164
+ )
161
165
  self.load_format = load_format
162
166
  self.nnodes = nnodes
163
167
  self.node_rank = node_rank
@@ -166,9 +170,7 @@ class SGLangEndpoint(Endpoint):
166
170
  self.disaggregation_mode = disaggregation_mode
167
171
  self.prefill_bootstrap_port = prefill_bootstrap_port
168
172
  self.launcher_module = launcher_module
169
- self.extra_server_args = (
170
- dict(extra_server_args) if extra_server_args else {}
171
- )
173
+ self.extra_server_args = dict(extra_server_args) if extra_server_args else {}
172
174
  self.health_timeout = health_timeout
173
175
  self.health_poll_interval = health_poll_interval
174
176
  self.health_request_timeout = health_request_timeout
@@ -183,17 +185,16 @@ class SGLangEndpoint(Endpoint):
183
185
  self._bench_server: Optional[http.server.ThreadingHTTPServer] = None
184
186
 
185
187
  if self.disaggregation_mode not in (None, "prefill", "decode"):
186
- raise ValueError(
187
- "disaggregation_mode must be None, 'prefill', or 'decode'"
188
- )
188
+ raise ValueError("disaggregation_mode must be None, 'prefill', or 'decode'")
189
189
 
190
190
  if os.environ.get(DUMMY_WEIGHTS_ENV) == "1":
191
- if (
192
- self.load_format is None
193
- and not has_server_arg(self.extra_server_args, "--load-format")
191
+ if self.load_format is None and not has_server_arg(
192
+ self.extra_server_args, "--load-format"
194
193
  ):
195
194
  self.load_format = "dummy"
196
- if not has_server_arg(self.extra_server_args, "--model-loader-extra-config"):
195
+ if not has_server_arg(
196
+ self.extra_server_args, "--model-loader-extra-config"
197
+ ):
197
198
  self.extra_server_args["--model-loader-extra-config"] = "{}"
198
199
 
199
200
  # Request logging enabled, log only metadata by default
@@ -204,7 +205,9 @@ class SGLangEndpoint(Endpoint):
204
205
  # Set request logging level
205
206
  level = get_server_arg(self.extra_server_args, "--log-requests-level")
206
207
  if level is None:
207
- self.extra_server_args["--log-requests-level"] = str(self.log_requests_level)
208
+ self.extra_server_args["--log-requests-level"] = str(
209
+ self.log_requests_level
210
+ )
208
211
  else:
209
212
  key, value = level
210
213
  print(
@@ -226,19 +229,21 @@ class SGLangEndpoint(Endpoint):
226
229
  self.log_request_headers = ",".join(sglang_log_request_headers)
227
230
  os.environ["SGLANG_LOG_REQUEST_HEADERS"] = self.log_request_headers
228
231
 
229
-
230
232
  def _build_cmd(self) -> list[str]:
231
233
  cmd = [
232
- "python", "-m", self.launcher_module,
233
- "--host", "0.0.0.0",
234
- "--port", str(self.worker_port),
235
- "--model-path", self.model_path,
234
+ "python",
235
+ "-m",
236
+ self.launcher_module,
237
+ "--host",
238
+ "0.0.0.0",
239
+ "--port",
240
+ str(self.worker_port),
241
+ "--model-path",
242
+ self.model_path,
236
243
  ]
237
244
 
238
245
  if self.speculative_model_path is not None:
239
- cmd.extend(
240
- ["--speculative-draft-model-path", self.speculative_model_path]
241
- )
246
+ cmd.extend(["--speculative-draft-model-path", self.speculative_model_path])
242
247
  if self.load_format is not None:
243
248
  cmd.extend(["--load-format", self.load_format])
244
249
  if self.tp is not None:
@@ -263,14 +268,18 @@ class SGLangEndpoint(Endpoint):
263
268
  raise ValueError("dist_init_host is required when nnodes > 1")
264
269
  cmd.extend(
265
270
  [
266
- "--nnodes", str(self.nnodes),
267
- "--node-rank", str(self.node_rank),
271
+ "--nnodes",
272
+ str(self.nnodes),
273
+ "--node-rank",
274
+ str(self.node_rank),
268
275
  "--dist-init-addr",
269
276
  f"{self.dist_init_host}:{self.dist_init_port}",
270
277
  ]
271
278
  )
272
279
 
273
- merged = _merge_server_args(self.DEFAULT_OPERATIONAL_ARGS, self.extra_server_args)
280
+ merged = _merge_server_args(
281
+ self.DEFAULT_OPERATIONAL_ARGS, self.extra_server_args
282
+ )
274
283
  for key, value in merged.items():
275
284
  cmd.extend(server_arg_tokens(key, value))
276
285
 
@@ -333,9 +342,7 @@ class VLLMEndpoint(Endpoint):
333
342
  super().__init__(base_url=f"http://localhost:{vllm_port}")
334
343
  self.model = model
335
344
  self.worker_port = vllm_port
336
- self.extra_server_args = (
337
- dict(extra_server_args) if extra_server_args else {}
338
- )
345
+ self.extra_server_args = dict(extra_server_args) if extra_server_args else {}
339
346
  self.health_timeout = health_timeout
340
347
  self.health_poll_interval = health_poll_interval
341
348
  self.health_request_timeout = health_request_timeout
@@ -348,10 +355,15 @@ class VLLMEndpoint(Endpoint):
348
355
 
349
356
  def _build_cmd(self) -> list[str]:
350
357
  cmd = [
351
- "python", "-m", "vllm.entrypoints.openai.api_server",
352
- "--host", "0.0.0.0",
353
- "--port", str(self.worker_port),
354
- "--model", self.model,
358
+ "python",
359
+ "-m",
360
+ "vllm.entrypoints.openai.api_server",
361
+ "--host",
362
+ "0.0.0.0",
363
+ "--port",
364
+ str(self.worker_port),
365
+ "--model",
366
+ self.model,
355
367
  ]
356
368
  for key, value in self.extra_server_args.items():
357
369
  cmd.extend(server_arg_tokens(key, value))
@@ -405,9 +417,7 @@ class TRTLLMEndpoint(Endpoint):
405
417
  super().__init__(base_url=f"http://localhost:{worker_port}")
406
418
  self.model = model
407
419
  self.worker_port = worker_port
408
- self.extra_server_args = (
409
- dict(extra_server_args) if extra_server_args else {}
410
- )
420
+ self.extra_server_args = dict(extra_server_args) if extra_server_args else {}
411
421
  self.health_timeout = health_timeout
412
422
  self.health_poll_interval = health_poll_interval
413
423
  self.health_request_timeout = health_request_timeout
@@ -417,8 +427,10 @@ class TRTLLMEndpoint(Endpoint):
417
427
  cmd = [
418
428
  "trtllm-serve",
419
429
  self.model,
420
- "--host", "0.0.0.0",
421
- "--port", str(self.worker_port),
430
+ "--host",
431
+ "0.0.0.0",
432
+ "--port",
433
+ str(self.worker_port),
422
434
  ]
423
435
  for key, value in self.extra_server_args.items():
424
436
  cmd.extend(server_arg_tokens(key, value))
@@ -483,18 +495,30 @@ class RouterEndpoint(Endpoint):
483
495
 
484
496
  def _build_cmd(self) -> list[str]:
485
497
  cmd = [
486
- "python", "-m", "sglang_router.launch_router",
487
- "--host", "0.0.0.0",
488
- "--port", str(self.router_port),
489
- "--prefill-policy", "cache_aware",
490
- "--decode-policy", "round_robin",
491
- "--max-concurrent-requests", "128",
492
- "--rate-limit-tokens-per-second", "0",
493
- "--queue-size", "0",
494
- "--health-check-timeout-secs", "600",
495
- "--log-level", "info",
498
+ "python",
499
+ "-m",
500
+ "sglang_router.launch_router",
501
+ "--host",
502
+ "0.0.0.0",
503
+ "--port",
504
+ str(self.router_port),
505
+ "--prefill-policy",
506
+ "cache_aware",
507
+ "--decode-policy",
508
+ "round_robin",
509
+ "--max-concurrent-requests",
510
+ "128",
511
+ "--rate-limit-tokens-per-second",
512
+ "0",
513
+ "--queue-size",
514
+ "0",
515
+ "--health-check-timeout-secs",
516
+ "600",
517
+ "--log-level",
518
+ "info",
496
519
  "--disable-circuit-breaker",
497
- "--request-timeout-secs", "3600",
520
+ "--request-timeout-secs",
521
+ "3600",
498
522
  ]
499
523
 
500
524
  if self.api_key is not None:
@@ -506,9 +530,7 @@ class RouterEndpoint(Endpoint):
506
530
  for role, node_ip in self.pd_config:
507
531
  node_url = f"http://{node_ip}:{self.worker_port}"
508
532
  if role == "prefill":
509
- cmd.extend(
510
- ["--prefill", node_url, str(self.prefill_bootstrap_port)]
511
- )
533
+ cmd.extend(["--prefill", node_url, str(self.prefill_bootstrap_port)])
512
534
  elif role == "decode":
513
535
  cmd.extend(["--decode", node_url])
514
536
  elif role == "worker":
@@ -552,6 +574,7 @@ class RouterEndpoint(Endpoint):
552
574
  # and this proxy listens on worker_port. POST /bench shells out to a benchmark
553
575
  # task via run_bench(); every other path is forwarded to the upstream server.
554
576
 
577
+
555
578
  def start_bench_proxy(
556
579
  *,
557
580
  listen_port: int,
@@ -562,17 +585,12 @@ def start_bench_proxy(
562
585
  (_BenchProxyHandler,),
563
586
  {"upstream_port": upstream_port},
564
587
  )
565
- server = http.server.ThreadingHTTPServer(
566
- ("0.0.0.0", listen_port), handler_cls
567
- )
588
+ server = http.server.ThreadingHTTPServer(("0.0.0.0", listen_port), handler_cls)
568
589
  thread = threading.Thread(
569
590
  target=server.serve_forever, daemon=True, name="bench-proxy"
570
591
  )
571
592
  thread.start()
572
- print(
573
- f"[bench-proxy] listening on :{listen_port}, "
574
- f"forwarding to :{upstream_port}"
575
- )
593
+ print(f"[bench-proxy] listening on :{listen_port}, forwarding to :{upstream_port}")
576
594
  return server
577
595
 
578
596
 
@@ -601,26 +619,17 @@ class _BenchProxyHandler(http.server.BaseHTTPRequestHandler):
601
619
  try:
602
620
  payload = self._read_json_body()
603
621
  except (ValueError, json.JSONDecodeError) as exc:
604
- return self._send_json(
605
- 400, {"ok": False, "error": f"bad body: {exc}"}
606
- )
622
+ return self._send_json(400, {"ok": False, "error": f"bad body: {exc}"})
607
623
 
608
624
  benchmark = payload.get("benchmark") or ""
609
625
  args = payload.get("args") or []
610
- target = (
611
- payload.get("target")
612
- or f"http://localhost:{self.upstream_port}"
613
- )
626
+ target = payload.get("target") or f"http://localhost:{self.upstream_port}"
614
627
  output_dir = payload.get("output_dir") or "/tmp/bench-output"
615
628
 
616
629
  if not benchmark:
617
- return self._send_json(
618
- 400, {"ok": False, "error": "benchmark is required"}
619
- )
630
+ return self._send_json(400, {"ok": False, "error": "benchmark is required"})
620
631
  if not isinstance(args, list):
621
- return self._send_json(
622
- 400, {"ok": False, "error": "args must be a list"}
623
- )
632
+ return self._send_json(400, {"ok": False, "error": "args must be a list"})
624
633
 
625
634
  rc, body = run_bench(benchmark, list(args), target, output_dir)
626
635
  status = 200 if rc == 0 else 500
@@ -671,9 +680,7 @@ class _BenchProxyHandler(http.server.BaseHTTPRequestHandler):
671
680
  return json.loads(raw)
672
681
 
673
682
  def _send_json(self, status: int, obj: Mapping[str, Any]):
674
- self._send_raw(
675
- status, "application/json", json.dumps(obj).encode()
676
- )
683
+ self._send_raw(status, "application/json", json.dumps(obj).encode())
677
684
 
678
685
  def _send_raw(self, status: int, content_type: str, body: bytes):
679
686
  self.send_response(status)
@@ -698,18 +705,21 @@ def run_bench(
698
705
  return 500, json.dumps(
699
706
  {
700
707
  "ok": False,
701
- "error": (
702
- "autoinference.benchmarks not importable in container"
703
- ),
708
+ "error": ("autoinference.benchmarks not importable in container"),
704
709
  "volume_path": output_dir,
705
710
  }
706
711
  ).encode()
707
712
 
708
713
  cmd = [
709
- sys.executable, "-m", "invoke",
710
- "--search-root", search_root,
711
- "-c", "tasks",
712
- benchmark, *[str(a) for a in args],
714
+ sys.executable,
715
+ "-m",
716
+ "invoke",
717
+ "--search-root",
718
+ search_root,
719
+ "-c",
720
+ "tasks",
721
+ benchmark,
722
+ *[str(a) for a in args],
713
723
  ]
714
724
  env = {
715
725
  **os.environ,
@@ -759,9 +769,7 @@ def wait_ready(
759
769
 
760
770
  while time.time() < deadline:
761
771
  try:
762
- error = _health_check(
763
- url, request_timeout=request_timeout, process=process
764
- )
772
+ error = _health_check(url, request_timeout=request_timeout, process=process)
765
773
  except subprocess.CalledProcessError as exc:
766
774
  print(
767
775
  f"[endpoint] !!! server process exited with code "
@@ -775,13 +783,11 @@ def wait_ready(
775
783
  time.sleep(poll_interval)
776
784
 
777
785
  print(
778
- f"[endpoint] !!! health check timed out after {timeout}s; "
779
- f"last={last_error}",
786
+ f"[endpoint] !!! health check timed out after {timeout}s; last={last_error}",
780
787
  flush=True,
781
788
  )
782
789
  raise TimeoutError(
783
- f"Health check timed out after {timeout}s for {url}. "
784
- f"Last error: {last_error}"
790
+ f"Health check timed out after {timeout}s for {url}. Last error: {last_error}"
785
791
  )
786
792
 
787
793
 
@@ -809,9 +815,7 @@ def warmup_chat_completions(
809
815
  request_timeout=request_timeout,
810
816
  max_attempts=max_attempts_per_request,
811
817
  retry_delay=retry_delay,
812
- description=(
813
- f"warmup request {request_idx + 1}/{successful_requests}"
814
- ),
818
+ description=(f"warmup request {request_idx + 1}/{successful_requests}"),
815
819
  )
816
820
 
817
821
 
@@ -833,9 +837,7 @@ def validate_embeddings_endpoint(
833
837
  1 if all(isinstance(item, int) for item in inputs) else len(inputs)
834
838
  )
835
839
  else:
836
- raise ValueError(
837
- "embedding validation payload must contain non-empty input"
838
- )
840
+ raise ValueError("embedding validation payload must contain non-empty input")
839
841
 
840
842
  request_headers = {"Content-Type": "application/json"}
841
843
  if headers:
@@ -898,10 +900,7 @@ def start_heartbeat_thread(
898
900
  try:
899
901
  error = health_check_fn()
900
902
  except subprocess.CalledProcessError as exc:
901
- print(
902
- f"[heartbeat] server process exited with code "
903
- f"{exc.returncode}"
904
- )
903
+ print(f"[heartbeat] server process exited with code {exc.returncode}")
905
904
  on_failure()
906
905
  return
907
906
  if error is None:
@@ -913,10 +912,7 @@ def start_heartbeat_thread(
913
912
  f"({consecutive_failures}/{max_consecutive_failures})"
914
913
  )
915
914
  if consecutive_failures >= max_consecutive_failures:
916
- print(
917
- "[heartbeat] sustained health-check failure, "
918
- "invoking on_failure"
919
- )
915
+ print("[heartbeat] sustained health-check failure, invoking on_failure")
920
916
  on_failure()
921
917
  return
922
918
 
@@ -1001,8 +997,7 @@ def _wait_ready_url(
1001
997
  last_error = error
1002
998
  time.sleep(poll_interval)
1003
999
  raise TimeoutError(
1004
- f"Timed out after {timeout}s waiting for {url}. "
1005
- f"Last error: {last_error}"
1000
+ f"Timed out after {timeout}s waiting for {url}. Last error: {last_error}"
1006
1001
  )
1007
1002
 
1008
1003
 
@@ -1020,9 +1015,7 @@ def _post_json(
1020
1015
  timeout: float | None = None,
1021
1016
  ) -> bytes:
1022
1017
  body = json.dumps(payload).encode("utf-8")
1023
- req = urllib.request.Request(
1024
- url, data=body, headers=dict(headers), method="POST"
1025
- )
1018
+ req = urllib.request.Request(url, data=body, headers=dict(headers), method="POST")
1026
1019
  with urllib.request.urlopen(req, timeout=timeout) as resp:
1027
1020
  return resp.read()
1028
1021
 
@@ -0,0 +1,38 @@
1
+ """Prepare local model snapshots for Transformers' remote-code loader."""
2
+
3
+ import os
4
+ import shutil
5
+ import tempfile
6
+ from pathlib import Path
7
+
8
+ MATERIALIZE_MODEL_PATH_ENV = "AUTOINFERENCE_MATERIALIZE_MODEL_PATH"
9
+
10
+
11
+ def materialize_model_path(path: str) -> str:
12
+ """Copy snapshot metadata/code into a real directory; retain weight links.
13
+
14
+ Transformers can resolve Python-module symlinks into the HF blob directory,
15
+ losing the relative imports next to the original snapshot entry. This opt-in
16
+ keeps small files together without copying large safetensors shards. The
17
+ unique directory lasts for the container's lifetime.
18
+ """
19
+ source = Path(path)
20
+ if os.environ.get(MATERIALIZE_MODEL_PATH_ENV) != "1" or not source.is_dir():
21
+ return path
22
+
23
+ source = source.resolve()
24
+ destination = Path(tempfile.mkdtemp(prefix="autoinference-model-"))
25
+ try:
26
+ for entry in source.iterdir():
27
+ target = destination / entry.name
28
+ if entry.is_file() and entry.suffix == ".safetensors":
29
+ target.symlink_to(entry.resolve(strict=True))
30
+ elif entry.is_dir():
31
+ shutil.copytree(entry, target, symlinks=False)
32
+ else:
33
+ shutil.copy2(entry, target, follow_symlinks=True)
34
+ except BaseException:
35
+ shutil.rmtree(destination)
36
+ raise
37
+ print(f"Materialized model snapshot {source} -> {destination}")
38
+ return str(destination)
@@ -0,0 +1,84 @@
1
+ import shutil
2
+ from pathlib import Path
3
+
4
+ import pytest
5
+
6
+ from autoinference_utils.endpoint import SGLangEndpoint
7
+ from autoinference_utils.model_path import (
8
+ MATERIALIZE_MODEL_PATH_ENV,
9
+ materialize_model_path,
10
+ )
11
+
12
+
13
+ def snapshot(tmp_path):
14
+ blobs = tmp_path / "blobs"
15
+ source = tmp_path / "snapshots" / "revision"
16
+ blobs.mkdir()
17
+ source.mkdir(parents=True)
18
+ for name, content in {
19
+ "config.json": '{"model_type": "kimi"}',
20
+ "configuration_kimi.py": "from .helper import VALUE\n",
21
+ "helper.py": "VALUE = 42\n",
22
+ "model.safetensors": "weight placeholder",
23
+ }.items():
24
+ blob = blobs / ("hash-" + name)
25
+ blob.write_text(content)
26
+ (source / name).symlink_to(blob)
27
+ return source
28
+
29
+
30
+ def test_opt_in_and_repository_ids(tmp_path, monkeypatch):
31
+ source = snapshot(tmp_path)
32
+ monkeypatch.delenv(MATERIALIZE_MODEL_PATH_ENV, raising=False)
33
+ assert materialize_model_path(str(source)) == str(source)
34
+ monkeypatch.setenv(MATERIALIZE_MODEL_PATH_ENV, "1")
35
+ assert materialize_model_path("moonshotai/Kimi-K3") == "moonshotai/Kimi-K3"
36
+
37
+
38
+ def test_local_code_is_materialized_and_weights_stay_linked(tmp_path, monkeypatch):
39
+ source = snapshot(tmp_path)
40
+ monkeypatch.setenv(MATERIALIZE_MODEL_PATH_ENV, "1")
41
+ monkeypatch.setattr("tempfile.tempdir", str(tmp_path))
42
+ result = Path(materialize_model_path(str(source)))
43
+ assert result != source
44
+ for name in ("config.json", "configuration_kimi.py", "helper.py"):
45
+ assert not (result / name).is_symlink()
46
+ assert (result / name).resolve().parent == result
47
+ assert (result / name).read_bytes() == (source / name).read_bytes()
48
+ assert (result / "model.safetensors").is_symlink()
49
+ assert (result / "model.safetensors").resolve() == (
50
+ source / "model.safetensors"
51
+ ).resolve()
52
+ assert (source / "configuration_kimi.py").is_symlink()
53
+
54
+
55
+ def test_endpoint_uses_materialized_base_and_draft(tmp_path, monkeypatch):
56
+ source = snapshot(tmp_path)
57
+ monkeypatch.setenv(MATERIALIZE_MODEL_PATH_ENV, "1")
58
+ monkeypatch.setattr("tempfile.tempdir", str(tmp_path))
59
+ endpoint = SGLangEndpoint(
60
+ model_path=str(source), speculative_model_path=str(source)
61
+ )
62
+ assert endpoint.model_path != str(source)
63
+ assert endpoint.speculative_model_path != str(source)
64
+ assert endpoint.model_path != endpoint.speculative_model_path
65
+ cmd = endpoint._build_cmd()
66
+ assert cmd[cmd.index("--model-path") + 1] == endpoint.model_path
67
+ assert (
68
+ cmd[cmd.index("--speculative-draft-model-path") + 1]
69
+ == endpoint.speculative_model_path
70
+ )
71
+
72
+
73
+ def test_failed_copy_removes_partial_directory(tmp_path, monkeypatch):
74
+ source = snapshot(tmp_path)
75
+ monkeypatch.setenv(MATERIALIZE_MODEL_PATH_ENV, "1")
76
+ monkeypatch.setattr("tempfile.tempdir", str(tmp_path))
77
+
78
+ def fail(*args, **kwargs):
79
+ raise OSError("copy failed")
80
+
81
+ monkeypatch.setattr(shutil, "copy2", fail)
82
+ with pytest.raises(OSError, match="copy failed"):
83
+ materialize_model_path(str(source))
84
+ assert not list(tmp_path.glob("autoinference-model-*"))
@@ -1,9 +0,0 @@
1
- # autoinference-utils
2
-
3
- Shared endpoint abstractions (`SGLangEndpoint`, `VLLMEndpoint`) for autoinference deployments.
4
-
5
- ## Publishing
6
-
7
- Bump `version` in `pyproject.toml` and merge to `main` — CI publishes automatically via trusted publishing.
8
-
9
- To publish manually: `cd autoinference_utils && uv build && uv publish`