autoinference-utils 0.2.7__tar.gz → 0.2.8__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.
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.5
2
2
  Name: autoinference-utils
3
- Version: 0.2.7
3
+ Version: 0.2.8
4
4
  Summary: Shared endpoint abstractions for autoinference deployments
5
5
  Requires-Python: >=3.10
6
6
  Description-Content-Type: text/markdown
@@ -9,12 +9,6 @@ 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
-
18
12
  ## Publishing
19
13
 
20
14
  Bump `version` in `pyproject.toml` and merge to `main` — CI publishes automatically via trusted publishing.
@@ -0,0 +1,9 @@
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`
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "autoinference-utils"
3
- version = "0.2.7"
3
+ version = "0.2.8"
4
4
  description = "Shared endpoint abstractions for autoinference deployments"
5
5
  readme = "README.md"
6
6
  requires-python = ">=3.10"
@@ -23,18 +23,24 @@ 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
-
28
26
  BENCH_MODE_ENV = "ENDPOINT_BENCH_MODE"
29
27
  DUMMY_WEIGHTS_ENV = "ENDPOINT_DUMMY"
30
28
  BENCH_MODE_PORT_OFFSET = 10000
31
29
 
32
- MODAL_FLASH_REQUEST_UUID_HEADER = "modal-flash-request-uuid"
30
+ MODAL_FLASH_REQUEST_UUID_HEADER= "modal-flash-request-uuid"
33
31
  MODAL_SESSION_ID_HEADER = "modal-session-id"
34
32
 
35
33
  ENDPOINTS_REQUIRING_OAI_STREAM_COMPAT: tuple[str, ...] = ("TRTLLMEndpoint",)
36
34
 
37
35
 
36
+ def _url(host: str, port: int) -> str:
37
+ return f"http://{'[' + host + ']' if ':' in host else host}:{port}"
38
+
39
+
40
+ def _local_base_url(host: str, port: int) -> str:
41
+ return _url({"0.0.0.0": "127.0.0.1", "::": "::1"}.get(host, host), port)
42
+
43
+
38
44
  def server_arg_tokens(flag: str, value: str) -> list[str]:
39
45
  """Render one server arg as argv tokens.
40
46
 
@@ -57,7 +63,9 @@ def _flag_name(key: str) -> str:
57
63
  return key.partition("=")[0]
58
64
 
59
65
 
60
- def get_server_arg(server_args: Mapping[str, str], flag: str) -> tuple[str, str] | None:
66
+ def get_server_arg(
67
+ server_args: Mapping[str, str], flag: str
68
+ ) -> tuple[str, str] | None:
61
69
  """Return the key and value setting a flag under either spelling."""
62
70
  for key, value in server_args.items():
63
71
  if _flag_name(key) == flag:
@@ -150,18 +158,16 @@ class SGLangEndpoint(Endpoint):
150
158
  sglang_port = (
151
159
  worker_port + BENCH_MODE_PORT_OFFSET if self.bench_mode else worker_port
152
160
  )
153
- super().__init__(base_url=f"http://localhost:{sglang_port}")
161
+ host_arg = get_server_arg(extra_server_args or {}, "--host")
162
+ host = (host_arg[1] or host_arg[0].partition("=")[2]) if host_arg else "0.0.0.0"
163
+ super().__init__(base_url=_local_base_url(host, sglang_port))
154
164
 
155
165
  self.worker_port = sglang_port
156
- self.model_path = materialize_model_path(model_path)
166
+ self.model_path = model_path
157
167
  self.tp = tp
158
168
  self.ep = ep
159
169
  self.dp = dp
160
- self.speculative_model_path = (
161
- materialize_model_path(speculative_model_path)
162
- if speculative_model_path is not None
163
- else None
164
- )
170
+ self.speculative_model_path = speculative_model_path
165
171
  self.load_format = load_format
166
172
  self.nnodes = nnodes
167
173
  self.node_rank = node_rank
@@ -170,7 +176,9 @@ class SGLangEndpoint(Endpoint):
170
176
  self.disaggregation_mode = disaggregation_mode
171
177
  self.prefill_bootstrap_port = prefill_bootstrap_port
172
178
  self.launcher_module = launcher_module
173
- self.extra_server_args = dict(extra_server_args) if extra_server_args else {}
179
+ self.extra_server_args = (
180
+ dict(extra_server_args) if extra_server_args else {}
181
+ )
174
182
  self.health_timeout = health_timeout
175
183
  self.health_poll_interval = health_poll_interval
176
184
  self.health_request_timeout = health_request_timeout
@@ -185,16 +193,17 @@ class SGLangEndpoint(Endpoint):
185
193
  self._bench_server: Optional[http.server.ThreadingHTTPServer] = None
186
194
 
187
195
  if self.disaggregation_mode not in (None, "prefill", "decode"):
188
- raise ValueError("disaggregation_mode must be None, 'prefill', or 'decode'")
196
+ raise ValueError(
197
+ "disaggregation_mode must be None, 'prefill', or 'decode'"
198
+ )
189
199
 
190
200
  if os.environ.get(DUMMY_WEIGHTS_ENV) == "1":
191
- if self.load_format is None and not has_server_arg(
192
- self.extra_server_args, "--load-format"
201
+ if (
202
+ self.load_format is None
203
+ and not has_server_arg(self.extra_server_args, "--load-format")
193
204
  ):
194
205
  self.load_format = "dummy"
195
- if not has_server_arg(
196
- self.extra_server_args, "--model-loader-extra-config"
197
- ):
206
+ if not has_server_arg(self.extra_server_args, "--model-loader-extra-config"):
198
207
  self.extra_server_args["--model-loader-extra-config"] = "{}"
199
208
 
200
209
  # Request logging enabled, log only metadata by default
@@ -205,9 +214,7 @@ class SGLangEndpoint(Endpoint):
205
214
  # Set request logging level
206
215
  level = get_server_arg(self.extra_server_args, "--log-requests-level")
207
216
  if level is None:
208
- self.extra_server_args["--log-requests-level"] = str(
209
- self.log_requests_level
210
- )
217
+ self.extra_server_args["--log-requests-level"] = str(self.log_requests_level)
211
218
  else:
212
219
  key, value = level
213
220
  print(
@@ -229,21 +236,19 @@ class SGLangEndpoint(Endpoint):
229
236
  self.log_request_headers = ",".join(sglang_log_request_headers)
230
237
  os.environ["SGLANG_LOG_REQUEST_HEADERS"] = self.log_request_headers
231
238
 
239
+
232
240
  def _build_cmd(self) -> list[str]:
233
241
  cmd = [
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,
242
+ "python", "-m", self.launcher_module,
243
+ "--host", "0.0.0.0",
244
+ "--port", str(self.worker_port),
245
+ "--model-path", self.model_path,
243
246
  ]
244
247
 
245
248
  if self.speculative_model_path is not None:
246
- cmd.extend(["--speculative-draft-model-path", self.speculative_model_path])
249
+ cmd.extend(
250
+ ["--speculative-draft-model-path", self.speculative_model_path]
251
+ )
247
252
  if self.load_format is not None:
248
253
  cmd.extend(["--load-format", self.load_format])
249
254
  if self.tp is not None:
@@ -268,25 +273,21 @@ class SGLangEndpoint(Endpoint):
268
273
  raise ValueError("dist_init_host is required when nnodes > 1")
269
274
  cmd.extend(
270
275
  [
271
- "--nnodes",
272
- str(self.nnodes),
273
- "--node-rank",
274
- str(self.node_rank),
276
+ "--nnodes", str(self.nnodes),
277
+ "--node-rank", str(self.node_rank),
275
278
  "--dist-init-addr",
276
279
  f"{self.dist_init_host}:{self.dist_init_port}",
277
280
  ]
278
281
  )
279
282
 
280
- merged = _merge_server_args(
281
- self.DEFAULT_OPERATIONAL_ARGS, self.extra_server_args
282
- )
283
+ merged = _merge_server_args(self.DEFAULT_OPERATIONAL_ARGS, self.extra_server_args)
283
284
  for key, value in merged.items():
284
285
  cmd.extend(server_arg_tokens(key, value))
285
286
 
286
287
  return cmd
287
288
 
288
289
  def health_check(self) -> str | None:
289
- url = f"http://127.0.0.1:{self.worker_port}/health"
290
+ url = f"{self.base_url}/health"
290
291
  return _health_check(
291
292
  url,
292
293
  request_timeout=self.health_request_timeout,
@@ -300,6 +301,7 @@ class SGLangEndpoint(Endpoint):
300
301
  wait_ready(
301
302
  self._proc,
302
303
  port=self.worker_port,
304
+ base_url=self.base_url,
303
305
  timeout=self.health_timeout,
304
306
  poll_interval=self.health_poll_interval,
305
307
  request_timeout=self.health_request_timeout,
@@ -308,6 +310,7 @@ class SGLangEndpoint(Endpoint):
308
310
  self._bench_server = start_bench_proxy(
309
311
  listen_port=self.listen_port,
310
312
  upstream_port=self.worker_port,
313
+ upstream_base_url=self.base_url,
311
314
  )
312
315
 
313
316
  def stop(self):
@@ -342,7 +345,9 @@ class VLLMEndpoint(Endpoint):
342
345
  super().__init__(base_url=f"http://localhost:{vllm_port}")
343
346
  self.model = model
344
347
  self.worker_port = vllm_port
345
- self.extra_server_args = dict(extra_server_args) if extra_server_args else {}
348
+ self.extra_server_args = (
349
+ dict(extra_server_args) if extra_server_args else {}
350
+ )
346
351
  self.health_timeout = health_timeout
347
352
  self.health_poll_interval = health_poll_interval
348
353
  self.health_request_timeout = health_request_timeout
@@ -355,15 +360,10 @@ class VLLMEndpoint(Endpoint):
355
360
 
356
361
  def _build_cmd(self) -> list[str]:
357
362
  cmd = [
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,
363
+ "python", "-m", "vllm.entrypoints.openai.api_server",
364
+ "--host", "0.0.0.0",
365
+ "--port", str(self.worker_port),
366
+ "--model", self.model,
367
367
  ]
368
368
  for key, value in self.extra_server_args.items():
369
369
  cmd.extend(server_arg_tokens(key, value))
@@ -417,7 +417,9 @@ class TRTLLMEndpoint(Endpoint):
417
417
  super().__init__(base_url=f"http://localhost:{worker_port}")
418
418
  self.model = model
419
419
  self.worker_port = worker_port
420
- self.extra_server_args = dict(extra_server_args) if extra_server_args else {}
420
+ self.extra_server_args = (
421
+ dict(extra_server_args) if extra_server_args else {}
422
+ )
421
423
  self.health_timeout = health_timeout
422
424
  self.health_poll_interval = health_poll_interval
423
425
  self.health_request_timeout = health_request_timeout
@@ -427,10 +429,8 @@ class TRTLLMEndpoint(Endpoint):
427
429
  cmd = [
428
430
  "trtllm-serve",
429
431
  self.model,
430
- "--host",
431
- "0.0.0.0",
432
- "--port",
433
- str(self.worker_port),
432
+ "--host", "0.0.0.0",
433
+ "--port", str(self.worker_port),
434
434
  ]
435
435
  for key, value in self.extra_server_args.items():
436
436
  cmd.extend(server_arg_tokens(key, value))
@@ -474,13 +474,16 @@ class RouterEndpoint(Endpoint):
474
474
  api_key: Optional[str] = None,
475
475
  health_timeout: float = 10 * 60,
476
476
  health_poll_interval: float = 5.0,
477
+ health_path: str = "/health",
478
+ host: str = "0.0.0.0",
477
479
  ):
478
480
  self.bench_mode = os.environ.get(BENCH_MODE_ENV) == "1"
479
481
  self.listen_port = router_port
480
482
  actual_router_port = (
481
483
  router_port + BENCH_MODE_PORT_OFFSET if self.bench_mode else router_port
482
484
  )
483
- super().__init__(base_url=f"http://localhost:{actual_router_port}")
485
+ self.host = host
486
+ super().__init__(base_url=_local_base_url(host, actual_router_port))
484
487
  self.pd_config = list(pd_config)
485
488
  self.worker_port = (
486
489
  worker_port + BENCH_MODE_PORT_OFFSET if self.bench_mode else worker_port
@@ -490,35 +493,24 @@ class RouterEndpoint(Endpoint):
490
493
  self.api_key = api_key
491
494
  self.health_timeout = health_timeout
492
495
  self.health_poll_interval = health_poll_interval
496
+ self.health_path = health_path
493
497
  self._proc: Optional[subprocess.Popen] = None
494
498
  self._bench_server: Optional[http.server.ThreadingHTTPServer] = None
495
499
 
496
500
  def _build_cmd(self) -> list[str]:
497
501
  cmd = [
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",
502
+ "python", "-m", "sglang_router.launch_router",
503
+ "--host", f"[{self.host}]" if ":" in self.host else self.host,
504
+ "--port", str(self.router_port),
505
+ "--prefill-policy", "cache_aware",
506
+ "--decode-policy", "round_robin",
507
+ "--max-concurrent-requests", "128",
508
+ "--rate-limit-tokens-per-second", "0",
509
+ "--queue-size", "0",
510
+ "--health-check-timeout-secs", "600",
511
+ "--log-level", "info",
519
512
  "--disable-circuit-breaker",
520
- "--request-timeout-secs",
521
- "3600",
513
+ "--request-timeout-secs", "3600",
522
514
  ]
523
515
 
524
516
  if self.api_key is not None:
@@ -528,9 +520,11 @@ class RouterEndpoint(Endpoint):
528
520
  cmd.append("--pd-disaggregation")
529
521
 
530
522
  for role, node_ip in self.pd_config:
531
- node_url = f"http://{node_ip}:{self.worker_port}"
523
+ node_url = _url(node_ip, self.worker_port)
532
524
  if role == "prefill":
533
- cmd.extend(["--prefill", node_url, str(self.prefill_bootstrap_port)])
525
+ cmd.extend(
526
+ ["--prefill", node_url, str(self.prefill_bootstrap_port)]
527
+ )
534
528
  elif role == "decode":
535
529
  cmd.extend(["--decode", node_url])
536
530
  elif role == "worker":
@@ -543,7 +537,7 @@ class RouterEndpoint(Endpoint):
543
537
  def start(self):
544
538
  for _, node_ip in self.pd_config:
545
539
  _wait_ready_url(
546
- f"http://{node_ip}:{self.worker_port}/health",
540
+ f"{_url(node_ip, self.worker_port)}/health",
547
541
  timeout=self.health_timeout,
548
542
  poll_interval=self.health_poll_interval,
549
543
  )
@@ -551,17 +545,25 @@ class RouterEndpoint(Endpoint):
551
545
  cmd = self._build_cmd()
552
546
  print(f"[router] starting: {shlex.join(cmd)}")
553
547
  self._proc = subprocess.Popen(cmd)
554
- _wait_ready_url(
555
- f"http://localhost:{self.router_port}/health",
556
- timeout=self.health_timeout,
557
- poll_interval=self.health_poll_interval,
558
- )
548
+ self.wait_ready()
559
549
  if self.bench_mode:
560
550
  self._bench_server = start_bench_proxy(
561
551
  listen_port=self.listen_port,
562
552
  upstream_port=self.router_port,
553
+ upstream_base_url=self.base_url,
563
554
  )
564
555
 
556
+ def wait_ready(self) -> None:
557
+ assert self._proc is not None
558
+ wait_ready(
559
+ self._proc,
560
+ port=self.router_port,
561
+ base_url=self.base_url,
562
+ health_path=self.health_path,
563
+ timeout=self.health_timeout,
564
+ poll_interval=self.health_poll_interval,
565
+ )
566
+
565
567
  def stop(self):
566
568
  if self._bench_server is not None:
567
569
  self._bench_server.shutdown()
@@ -574,28 +576,34 @@ class RouterEndpoint(Endpoint):
574
576
  # and this proxy listens on worker_port. POST /bench shells out to a benchmark
575
577
  # task via run_bench(); every other path is forwarded to the upstream server.
576
578
 
577
-
578
579
  def start_bench_proxy(
579
580
  *,
580
581
  listen_port: int,
581
582
  upstream_port: int,
583
+ upstream_base_url: str | None = None,
582
584
  ) -> http.server.ThreadingHTTPServer:
583
585
  handler_cls = type(
584
586
  "BenchProxyHandler",
585
587
  (_BenchProxyHandler,),
586
- {"upstream_port": upstream_port},
588
+ {"upstream_port": upstream_port, "upstream_base_url": upstream_base_url},
589
+ )
590
+ server = http.server.ThreadingHTTPServer(
591
+ ("0.0.0.0", listen_port), handler_cls
587
592
  )
588
- server = http.server.ThreadingHTTPServer(("0.0.0.0", listen_port), handler_cls)
589
593
  thread = threading.Thread(
590
594
  target=server.serve_forever, daemon=True, name="bench-proxy"
591
595
  )
592
596
  thread.start()
593
- print(f"[bench-proxy] listening on :{listen_port}, forwarding to :{upstream_port}")
597
+ print(
598
+ f"[bench-proxy] listening on :{listen_port}, "
599
+ f"forwarding to :{upstream_port}"
600
+ )
594
601
  return server
595
602
 
596
603
 
597
604
  class _BenchProxyHandler(http.server.BaseHTTPRequestHandler):
598
605
  upstream_port: int = 8000
606
+ upstream_base_url: str | None = None
599
607
 
600
608
  def log_message(self, format, *args):
601
609
  return
@@ -619,24 +627,35 @@ class _BenchProxyHandler(http.server.BaseHTTPRequestHandler):
619
627
  try:
620
628
  payload = self._read_json_body()
621
629
  except (ValueError, json.JSONDecodeError) as exc:
622
- return self._send_json(400, {"ok": False, "error": f"bad body: {exc}"})
630
+ return self._send_json(
631
+ 400, {"ok": False, "error": f"bad body: {exc}"}
632
+ )
623
633
 
624
634
  benchmark = payload.get("benchmark") or ""
625
635
  args = payload.get("args") or []
626
- target = payload.get("target") or f"http://localhost:{self.upstream_port}"
636
+ target = (
637
+ payload.get("target")
638
+ or self.upstream_base_url
639
+ or f"http://localhost:{self.upstream_port}"
640
+ )
627
641
  output_dir = payload.get("output_dir") or "/tmp/bench-output"
628
642
 
629
643
  if not benchmark:
630
- return self._send_json(400, {"ok": False, "error": "benchmark is required"})
644
+ return self._send_json(
645
+ 400, {"ok": False, "error": "benchmark is required"}
646
+ )
631
647
  if not isinstance(args, list):
632
- return self._send_json(400, {"ok": False, "error": "args must be a list"})
648
+ return self._send_json(
649
+ 400, {"ok": False, "error": "args must be a list"}
650
+ )
633
651
 
634
652
  rc, body = run_bench(benchmark, list(args), target, output_dir)
635
653
  status = 200 if rc == 0 else 500
636
654
  self._send_raw(status, "application/json", body)
637
655
 
638
656
  def _proxy(self, method: str):
639
- url = f"http://localhost:{self.upstream_port}{self.path}"
657
+ base_url = self.upstream_base_url or f"http://localhost:{self.upstream_port}"
658
+ url = f"{base_url}{self.path}"
640
659
  length = int(self.headers.get("Content-Length", 0))
641
660
  body = self.rfile.read(length) if length else None
642
661
  forward_headers = {
@@ -680,7 +699,9 @@ class _BenchProxyHandler(http.server.BaseHTTPRequestHandler):
680
699
  return json.loads(raw)
681
700
 
682
701
  def _send_json(self, status: int, obj: Mapping[str, Any]):
683
- self._send_raw(status, "application/json", json.dumps(obj).encode())
702
+ self._send_raw(
703
+ status, "application/json", json.dumps(obj).encode()
704
+ )
684
705
 
685
706
  def _send_raw(self, status: int, content_type: str, body: bytes):
686
707
  self.send_response(status)
@@ -705,21 +726,18 @@ def run_bench(
705
726
  return 500, json.dumps(
706
727
  {
707
728
  "ok": False,
708
- "error": ("autoinference.benchmarks not importable in container"),
729
+ "error": (
730
+ "autoinference.benchmarks not importable in container"
731
+ ),
709
732
  "volume_path": output_dir,
710
733
  }
711
734
  ).encode()
712
735
 
713
736
  cmd = [
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],
737
+ sys.executable, "-m", "invoke",
738
+ "--search-root", search_root,
739
+ "-c", "tasks",
740
+ benchmark, *[str(a) for a in args],
723
741
  ]
724
742
  env = {
725
743
  **os.environ,
@@ -759,17 +777,20 @@ def wait_ready(
759
777
  port: int,
760
778
  timeout: float,
761
779
  health_path: str = "/health",
780
+ base_url: str | None = None,
762
781
  poll_interval: float = 5.0,
763
782
  request_timeout: float = 5.0,
764
783
  ) -> None:
765
784
  """Poll an HTTP health endpoint until ready, raising if the process dies."""
766
785
  deadline = time.time() + timeout
767
- url = f"http://127.0.0.1:{port}{health_path}"
786
+ url = f"{base_url or _local_base_url('127.0.0.1', port)}{health_path}"
768
787
  last_error = "no response yet"
769
788
 
770
789
  while time.time() < deadline:
771
790
  try:
772
- error = _health_check(url, request_timeout=request_timeout, process=process)
791
+ error = _health_check(
792
+ url, request_timeout=request_timeout, process=process
793
+ )
773
794
  except subprocess.CalledProcessError as exc:
774
795
  print(
775
796
  f"[endpoint] !!! server process exited with code "
@@ -783,17 +804,20 @@ def wait_ready(
783
804
  time.sleep(poll_interval)
784
805
 
785
806
  print(
786
- f"[endpoint] !!! health check timed out after {timeout}s; last={last_error}",
807
+ f"[endpoint] !!! health check timed out after {timeout}s; "
808
+ f"last={last_error}",
787
809
  flush=True,
788
810
  )
789
811
  raise TimeoutError(
790
- f"Health check timed out after {timeout}s for {url}. Last error: {last_error}"
812
+ f"Health check timed out after {timeout}s for {url}. "
813
+ f"Last error: {last_error}"
791
814
  )
792
815
 
793
816
 
794
817
  def warmup_chat_completions(
795
818
  *,
796
819
  port: int,
820
+ base_url: str | None = None,
797
821
  payload: Mapping[str, Any],
798
822
  headers: Mapping[str, str] | None = None,
799
823
  successful_requests: int = 3,
@@ -802,7 +826,7 @@ def warmup_chat_completions(
802
826
  retry_delay: float = 1.0,
803
827
  ) -> None:
804
828
  """Warm the OpenAI chat completions endpoint with strict retries."""
805
- url = f"http://127.0.0.1:{port}/v1/chat/completions"
829
+ url = f"{base_url or _local_base_url('127.0.0.1', port)}/v1/chat/completions"
806
830
  request_headers = {"Content-Type": "application/json"}
807
831
  if headers:
808
832
  request_headers.update(headers)
@@ -815,7 +839,9 @@ def warmup_chat_completions(
815
839
  request_timeout=request_timeout,
816
840
  max_attempts=max_attempts_per_request,
817
841
  retry_delay=retry_delay,
818
- description=(f"warmup request {request_idx + 1}/{successful_requests}"),
842
+ description=(
843
+ f"warmup request {request_idx + 1}/{successful_requests}"
844
+ ),
819
845
  )
820
846
 
821
847
 
@@ -837,7 +863,9 @@ def validate_embeddings_endpoint(
837
863
  1 if all(isinstance(item, int) for item in inputs) else len(inputs)
838
864
  )
839
865
  else:
840
- raise ValueError("embedding validation payload must contain non-empty input")
866
+ raise ValueError(
867
+ "embedding validation payload must contain non-empty input"
868
+ )
841
869
 
842
870
  request_headers = {"Content-Type": "application/json"}
843
871
  if headers:
@@ -900,7 +928,10 @@ def start_heartbeat_thread(
900
928
  try:
901
929
  error = health_check_fn()
902
930
  except subprocess.CalledProcessError as exc:
903
- print(f"[heartbeat] server process exited with code {exc.returncode}")
931
+ print(
932
+ f"[heartbeat] server process exited with code "
933
+ f"{exc.returncode}"
934
+ )
904
935
  on_failure()
905
936
  return
906
937
  if error is None:
@@ -912,7 +943,10 @@ def start_heartbeat_thread(
912
943
  f"({consecutive_failures}/{max_consecutive_failures})"
913
944
  )
914
945
  if consecutive_failures >= max_consecutive_failures:
915
- print("[heartbeat] sustained health-check failure, invoking on_failure")
946
+ print(
947
+ "[heartbeat] sustained health-check failure, "
948
+ "invoking on_failure"
949
+ )
916
950
  on_failure()
917
951
  return
918
952
 
@@ -997,7 +1031,8 @@ def _wait_ready_url(
997
1031
  last_error = error
998
1032
  time.sleep(poll_interval)
999
1033
  raise TimeoutError(
1000
- f"Timed out after {timeout}s waiting for {url}. Last error: {last_error}"
1034
+ f"Timed out after {timeout}s waiting for {url}. "
1035
+ f"Last error: {last_error}"
1001
1036
  )
1002
1037
 
1003
1038
 
@@ -1015,7 +1050,9 @@ def _post_json(
1015
1050
  timeout: float | None = None,
1016
1051
  ) -> bytes:
1017
1052
  body = json.dumps(payload).encode("utf-8")
1018
- req = urllib.request.Request(url, data=body, headers=dict(headers), method="POST")
1053
+ req = urllib.request.Request(
1054
+ url, data=body, headers=dict(headers), method="POST"
1055
+ )
1019
1056
  with urllib.request.urlopen(req, timeout=timeout) as resp:
1020
1057
  return resp.read()
1021
1058
 
@@ -0,0 +1,315 @@
1
+ """SGLang p/d disaggregated serving in a multi-container static cluster."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import asyncio
6
+ import math
7
+ import os
8
+ import subprocess
9
+ import threading
10
+ import time
11
+ from collections.abc import Callable, Mapping
12
+
13
+ from .endpoint import (
14
+ Endpoint,
15
+ RouterEndpoint,
16
+ SGLangEndpoint,
17
+ _flag_name,
18
+ _health_check,
19
+ _local_base_url,
20
+ _merge_server_args,
21
+ _url,
22
+ )
23
+
24
+
25
+ def _configure_fabric() -> None:
26
+ if os.environ.get("FI_PROVIDER") == "efa":
27
+ os.environ["SGLANG_DISAGGREGATION_NIXL_BACKEND"] = "LIBFABRIC"
28
+ elif hca := os.environ.get("NCCL_IB_HCA", ""):
29
+ rails = [name for name in hca.lstrip("=").split(",") if name]
30
+ os.environ["UCX_NET_DEVICES"] = ",".join(f"{name}:1" for name in rails)
31
+
32
+
33
+ async def _stop_container() -> None:
34
+ from modal.client import _Client
35
+ from modal_proto import api_pb2
36
+
37
+ async def stop():
38
+ client = await _Client.from_env()
39
+ await client.stub.ContainerStop(
40
+ api_pb2.ContainerStopRequest(task_id=os.environ["MODAL_TASK_ID"])
41
+ )
42
+
43
+ await asyncio.wait_for(stop(), timeout=3)
44
+
45
+
46
+ def _terminate_container() -> None:
47
+ from modal._utils.async_utils import synchronize_api
48
+
49
+ try:
50
+ synchronize_api(_stop_container)()
51
+ finally:
52
+ os._exit(1)
53
+
54
+
55
+ class _PDRouter(RouterEndpoint):
56
+ def __init__(
57
+ self,
58
+ *,
59
+ max_concurrent_requests,
60
+ drain_timeout,
61
+ health_interval,
62
+ health_request_timeout,
63
+ health_failure_threshold,
64
+ **kwargs,
65
+ ):
66
+ super().__init__(health_path="/readiness", host="::", **kwargs)
67
+ self.max_concurrent_requests = max_concurrent_requests
68
+ self.drain_timeout = drain_timeout
69
+ self.health_interval = health_interval
70
+ self.health_request_timeout = health_request_timeout
71
+ self.health_failure_threshold = health_failure_threshold
72
+
73
+ def _build_cmd(self) -> list[str]:
74
+ cmd = [
75
+ "python",
76
+ "-m",
77
+ "sglang_router.launch_router",
78
+ "--host",
79
+ "[::]",
80
+ "--port",
81
+ str(self.router_port),
82
+ "--pd-disaggregation",
83
+ ]
84
+ for role, host in self.pd_config:
85
+ cmd.extend([f"--{role}", _url(host, self.worker_port)])
86
+ if role == "prefill":
87
+ cmd.append(str(self.prefill_bootstrap_port))
88
+ for flag, value in {
89
+ "--prefill-policy": "cache_aware",
90
+ "--decode-policy": "round_robin",
91
+ "--max-concurrent-requests": self.max_concurrent_requests,
92
+ "--queue-size": 0,
93
+ "--rate-limit-tokens-per-second": 0,
94
+ "--health-check-interval-secs": self.health_interval,
95
+ "--health-check-timeout-secs": self.health_request_timeout,
96
+ "--health-failure-threshold": self.health_failure_threshold,
97
+ "--shutdown-grace-period-secs": self.drain_timeout,
98
+ "--request-timeout-secs": 1800,
99
+ "--log-level": "info",
100
+ }.items():
101
+ cmd.extend([flag, str(value)])
102
+ if self.api_key is not None:
103
+ cmd.extend(["--api-key", self.api_key])
104
+ return cmd
105
+
106
+
107
+ class PDEndpoint(Endpoint):
108
+ """Own a local SGLang engine and the rank-zero router for P/D serving.
109
+
110
+ Call start/stop from Modal enter/exit hooks. Model-specific arguments and
111
+ warmup stay in the recipe. The image must support NIXL and graceful shutdown.
112
+ The ratio is (prefill, decode) counts; prefill ranks precede decode ranks.
113
+ """
114
+
115
+ def __init__(
116
+ self,
117
+ engine: SGLangEndpoint,
118
+ *,
119
+ ratio: tuple[int, int],
120
+ prefill_args: Mapping[str, str] | None = None,
121
+ decode_args: Mapping[str, str] | None = None,
122
+ router_port: int = 9000,
123
+ max_concurrent_requests: int = 96,
124
+ api_key: str | None = None,
125
+ drain_timeout: int = 120,
126
+ health_interval: int = 10,
127
+ health_request_timeout: int = 10,
128
+ router_health_request_timeout: int = 5,
129
+ health_failure_threshold: int = 3,
130
+ health_failure_timeout: float = 30,
131
+ ):
132
+ from modal.experimental import get_cluster_info
133
+
134
+ if not isinstance(engine, SGLangEndpoint):
135
+ raise TypeError("PDEndpoint requires a SGLangEndpoint")
136
+ if engine._proc is not None or engine.bench_mode:
137
+ raise ValueError("PDEndpoint requires an unstarted, non-benchmark engine")
138
+ if (
139
+ min(
140
+ health_interval,
141
+ health_request_timeout,
142
+ router_health_request_timeout,
143
+ health_failure_threshold,
144
+ )
145
+ <= 0
146
+ or drain_timeout < 0
147
+ or not math.isfinite(health_failure_timeout)
148
+ or health_failure_timeout <= 0
149
+ ):
150
+ raise ValueError(
151
+ "health settings must be positive and drain_timeout nonnegative"
152
+ )
153
+ if len(ratio) != 2 or any(
154
+ type(count) is not int or count < 1 for count in ratio
155
+ ):
156
+ raise ValueError(
157
+ "ratio must contain positive integer prefill and decode counts"
158
+ )
159
+ roles = ("prefill",) * ratio[0] + ("decode",) * ratio[1]
160
+ cluster = get_cluster_info()
161
+ hosts = cluster.container_ips
162
+ if len(hosts) != len(roles) or not 0 <= cluster.rank < len(roles):
163
+ raise ValueError(
164
+ f"P/D ratio {ratio} requires a {len(roles)}-container Modal cluster"
165
+ )
166
+ self._host_ip = hosts[cluster.rank]
167
+ self.role = roles[cluster.rank]
168
+ super().__init__(_url(hosts[0], router_port))
169
+ self.engine = engine
170
+ engine.disaggregation_mode = self.role
171
+ engine.health_request_timeout = health_request_timeout
172
+ engine.extra_server_args = _merge_server_args(
173
+ engine.extra_server_args,
174
+ (prefill_args if self.role == "prefill" else decode_args) or {},
175
+ )
176
+ engine.extra_server_args = _merge_server_args(
177
+ engine.extra_server_args,
178
+ {
179
+ "--disaggregation-transfer-backend": "nixl",
180
+ "--host": "::",
181
+ },
182
+ )
183
+ engine.extra_server_args = {
184
+ key: value
185
+ for key, value in engine.extra_server_args.items()
186
+ if _flag_name(key) != "--disaggregation-mode"
187
+ }
188
+ engine.base_url = _local_base_url("::", engine.worker_port)
189
+ self.router = (
190
+ _PDRouter(
191
+ pd_config=list(zip(roles, hosts)),
192
+ worker_port=engine.worker_port,
193
+ router_port=router_port,
194
+ prefill_bootstrap_port=engine.prefill_bootstrap_port,
195
+ health_timeout=engine.health_timeout,
196
+ api_key=api_key,
197
+ max_concurrent_requests=max_concurrent_requests,
198
+ drain_timeout=drain_timeout,
199
+ health_interval=health_interval,
200
+ health_request_timeout=router_health_request_timeout,
201
+ health_failure_threshold=health_failure_threshold,
202
+ )
203
+ if cluster.rank == 0
204
+ else None
205
+ )
206
+ self.drain_timeout = drain_timeout
207
+ self.health_interval = health_interval
208
+ self.health_failure_timeout = health_failure_timeout
209
+ self._stopped = threading.Event()
210
+ self._lock = threading.Lock()
211
+ self._started = False
212
+ self._threads: list[threading.Thread] = []
213
+
214
+ def start(self, *, warmup: Callable[[], None] | None = None) -> None:
215
+ """Run warmup on rank zero before supervising every worker's health."""
216
+ if self._started or self._stopped.is_set():
217
+ raise RuntimeError("PDEndpoint can only be started once")
218
+ self._started = True
219
+ os.environ.setdefault("SGLANG_HOST_IP", self._host_ip)
220
+ _configure_fabric()
221
+ os.environ.update(
222
+ SGLANG_GRACEFUL_SHUTDOWN_TIMEOUT=str(self.drain_timeout),
223
+ SGLANG_ENABLE_HEALTH_ENDPOINT_GENERATION="1",
224
+ )
225
+ try:
226
+ for endpoint in (self.engine, self.router):
227
+ if endpoint is not None:
228
+ endpoint.start()
229
+ self._watch_process(endpoint)
230
+ if self.router is not None:
231
+ if warmup is not None:
232
+ warmup()
233
+ self.router.wait_ready()
234
+ for _, host in self.router.pd_config:
235
+ self._watch_health(_url(host, self.engine.worker_port) + "/health")
236
+ self._watch_health(self.router.base_url + "/health")
237
+ except BaseException:
238
+ self.stop()
239
+ raise
240
+
241
+ def _watch_process(self, endpoint: SGLangEndpoint | RouterEndpoint) -> None:
242
+ process = endpoint._proc
243
+ assert process is not None
244
+
245
+ def wait():
246
+ code = process.wait()
247
+ if not self._stopped.is_set():
248
+ print(f"[pd] Serving process exited with code {code}", flush=True)
249
+ self._fail()
250
+
251
+ thread = threading.Thread(target=wait, daemon=True)
252
+ self._threads.append(thread)
253
+ thread.start()
254
+
255
+ def _watch_health(self, url: str) -> None:
256
+ def poll():
257
+ failed_since = None
258
+ while not self._stopped.wait(self.health_interval):
259
+ started = time.monotonic()
260
+ error = _health_check(
261
+ url, request_timeout=self.engine.health_request_timeout
262
+ )
263
+ if error is None:
264
+ failed_since = None
265
+ continue
266
+ if failed_since is None:
267
+ failed_since = started
268
+ elapsed = time.monotonic() - failed_since
269
+ print(f"[pd] {url}: {error}; unhealthy for {elapsed:.1f}s", flush=True)
270
+ if elapsed >= self.health_failure_timeout:
271
+ self._fail()
272
+ return
273
+
274
+ thread = threading.Thread(target=poll, daemon=True)
275
+ self._threads.append(thread)
276
+ thread.start()
277
+
278
+ def _fail(self) -> None:
279
+ with self._lock:
280
+ if self._stopped.is_set():
281
+ return
282
+ self._stopped.set()
283
+ try:
284
+ for endpoint in (self.router, self.engine):
285
+ if endpoint is not None and endpoint._proc is not None:
286
+ try:
287
+ endpoint._proc.kill()
288
+ except ProcessLookupError:
289
+ pass
290
+ finally:
291
+ _terminate_container()
292
+
293
+ def stop(self) -> None:
294
+ with self._lock:
295
+ if self._stopped.is_set():
296
+ return
297
+ self._stopped.set()
298
+ deadline = time.monotonic() + self.drain_timeout + 30
299
+ try:
300
+ for endpoint in (self.router, self.engine):
301
+ if endpoint is None:
302
+ continue
303
+ process = endpoint._proc
304
+ if process is not None and process.poll() is None:
305
+ process.terminate()
306
+ try:
307
+ process.wait(timeout=max(0, deadline - time.monotonic()))
308
+ except subprocess.TimeoutExpired:
309
+ process.kill()
310
+ endpoint.stop()
311
+ finally:
312
+ for thread in self._threads:
313
+ thread.join(
314
+ timeout=self.engine.health_request_timeout + self.health_interval
315
+ )
@@ -1,15 +0,0 @@
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,38 +0,0 @@
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)
@@ -1,84 +0,0 @@
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-*"))