autoinference-utils 0.2.7__tar.gz → 0.2.9__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.9
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.9"
4
4
  description = "Shared endpoint abstractions for autoinference deployments"
5
5
  readme = "README.md"
6
6
  requires-python = ">=3.10"
@@ -23,8 +23,6 @@ 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
@@ -35,6 +33,14 @@ MODAL_SESSION_ID_HEADER = "modal-session-id"
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
 
@@ -150,18 +156,16 @@ class SGLangEndpoint(Endpoint):
150
156
  sglang_port = (
151
157
  worker_port + BENCH_MODE_PORT_OFFSET if self.bench_mode else worker_port
152
158
  )
153
- super().__init__(base_url=f"http://localhost:{sglang_port}")
159
+ host_arg = get_server_arg(extra_server_args or {}, "--host")
160
+ host = (host_arg[1] or host_arg[0].partition("=")[2]) if host_arg else "0.0.0.0"
161
+ super().__init__(base_url=_local_base_url(host, sglang_port))
154
162
 
155
163
  self.worker_port = sglang_port
156
- self.model_path = materialize_model_path(model_path)
164
+ self.model_path = model_path
157
165
  self.tp = tp
158
166
  self.ep = ep
159
167
  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
- )
168
+ self.speculative_model_path = speculative_model_path
165
169
  self.load_format = load_format
166
170
  self.nnodes = nnodes
167
171
  self.node_rank = node_rank
@@ -286,7 +290,7 @@ class SGLangEndpoint(Endpoint):
286
290
  return cmd
287
291
 
288
292
  def health_check(self) -> str | None:
289
- url = f"http://127.0.0.1:{self.worker_port}/health"
293
+ url = f"{self.base_url}/health"
290
294
  return _health_check(
291
295
  url,
292
296
  request_timeout=self.health_request_timeout,
@@ -300,6 +304,7 @@ class SGLangEndpoint(Endpoint):
300
304
  wait_ready(
301
305
  self._proc,
302
306
  port=self.worker_port,
307
+ base_url=self.base_url,
303
308
  timeout=self.health_timeout,
304
309
  poll_interval=self.health_poll_interval,
305
310
  request_timeout=self.health_request_timeout,
@@ -308,6 +313,7 @@ class SGLangEndpoint(Endpoint):
308
313
  self._bench_server = start_bench_proxy(
309
314
  listen_port=self.listen_port,
310
315
  upstream_port=self.worker_port,
316
+ upstream_base_url=self.base_url,
311
317
  )
312
318
 
313
319
  def stop(self):
@@ -474,13 +480,16 @@ class RouterEndpoint(Endpoint):
474
480
  api_key: Optional[str] = None,
475
481
  health_timeout: float = 10 * 60,
476
482
  health_poll_interval: float = 5.0,
483
+ health_path: str = "/health",
484
+ host: str = "0.0.0.0",
477
485
  ):
478
486
  self.bench_mode = os.environ.get(BENCH_MODE_ENV) == "1"
479
487
  self.listen_port = router_port
480
488
  actual_router_port = (
481
489
  router_port + BENCH_MODE_PORT_OFFSET if self.bench_mode else router_port
482
490
  )
483
- super().__init__(base_url=f"http://localhost:{actual_router_port}")
491
+ self.host = host
492
+ super().__init__(base_url=_local_base_url(host, actual_router_port))
484
493
  self.pd_config = list(pd_config)
485
494
  self.worker_port = (
486
495
  worker_port + BENCH_MODE_PORT_OFFSET if self.bench_mode else worker_port
@@ -490,6 +499,7 @@ class RouterEndpoint(Endpoint):
490
499
  self.api_key = api_key
491
500
  self.health_timeout = health_timeout
492
501
  self.health_poll_interval = health_poll_interval
502
+ self.health_path = health_path
493
503
  self._proc: Optional[subprocess.Popen] = None
494
504
  self._bench_server: Optional[http.server.ThreadingHTTPServer] = None
495
505
 
@@ -499,7 +509,7 @@ class RouterEndpoint(Endpoint):
499
509
  "-m",
500
510
  "sglang_router.launch_router",
501
511
  "--host",
502
- "0.0.0.0",
512
+ f"[{self.host}]" if ":" in self.host else self.host,
503
513
  "--port",
504
514
  str(self.router_port),
505
515
  "--prefill-policy",
@@ -528,7 +538,7 @@ class RouterEndpoint(Endpoint):
528
538
  cmd.append("--pd-disaggregation")
529
539
 
530
540
  for role, node_ip in self.pd_config:
531
- node_url = f"http://{node_ip}:{self.worker_port}"
541
+ node_url = _url(node_ip, self.worker_port)
532
542
  if role == "prefill":
533
543
  cmd.extend(["--prefill", node_url, str(self.prefill_bootstrap_port)])
534
544
  elif role == "decode":
@@ -543,7 +553,7 @@ class RouterEndpoint(Endpoint):
543
553
  def start(self):
544
554
  for _, node_ip in self.pd_config:
545
555
  _wait_ready_url(
546
- f"http://{node_ip}:{self.worker_port}/health",
556
+ f"{_url(node_ip, self.worker_port)}/health",
547
557
  timeout=self.health_timeout,
548
558
  poll_interval=self.health_poll_interval,
549
559
  )
@@ -551,17 +561,25 @@ class RouterEndpoint(Endpoint):
551
561
  cmd = self._build_cmd()
552
562
  print(f"[router] starting: {shlex.join(cmd)}")
553
563
  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
- )
564
+ self.wait_ready()
559
565
  if self.bench_mode:
560
566
  self._bench_server = start_bench_proxy(
561
567
  listen_port=self.listen_port,
562
568
  upstream_port=self.router_port,
569
+ upstream_base_url=self.base_url,
563
570
  )
564
571
 
572
+ def wait_ready(self) -> None:
573
+ assert self._proc is not None
574
+ wait_ready(
575
+ self._proc,
576
+ port=self.router_port,
577
+ base_url=self.base_url,
578
+ health_path=self.health_path,
579
+ timeout=self.health_timeout,
580
+ poll_interval=self.health_poll_interval,
581
+ )
582
+
565
583
  def stop(self):
566
584
  if self._bench_server is not None:
567
585
  self._bench_server.shutdown()
@@ -579,11 +597,12 @@ def start_bench_proxy(
579
597
  *,
580
598
  listen_port: int,
581
599
  upstream_port: int,
600
+ upstream_base_url: str | None = None,
582
601
  ) -> http.server.ThreadingHTTPServer:
583
602
  handler_cls = type(
584
603
  "BenchProxyHandler",
585
604
  (_BenchProxyHandler,),
586
- {"upstream_port": upstream_port},
605
+ {"upstream_port": upstream_port, "upstream_base_url": upstream_base_url},
587
606
  )
588
607
  server = http.server.ThreadingHTTPServer(("0.0.0.0", listen_port), handler_cls)
589
608
  thread = threading.Thread(
@@ -596,6 +615,7 @@ def start_bench_proxy(
596
615
 
597
616
  class _BenchProxyHandler(http.server.BaseHTTPRequestHandler):
598
617
  upstream_port: int = 8000
618
+ upstream_base_url: str | None = None
599
619
 
600
620
  def log_message(self, format, *args):
601
621
  return
@@ -623,7 +643,11 @@ class _BenchProxyHandler(http.server.BaseHTTPRequestHandler):
623
643
 
624
644
  benchmark = payload.get("benchmark") or ""
625
645
  args = payload.get("args") or []
626
- target = payload.get("target") or f"http://localhost:{self.upstream_port}"
646
+ target = (
647
+ payload.get("target")
648
+ or self.upstream_base_url
649
+ or f"http://localhost:{self.upstream_port}"
650
+ )
627
651
  output_dir = payload.get("output_dir") or "/tmp/bench-output"
628
652
 
629
653
  if not benchmark:
@@ -636,7 +660,8 @@ class _BenchProxyHandler(http.server.BaseHTTPRequestHandler):
636
660
  self._send_raw(status, "application/json", body)
637
661
 
638
662
  def _proxy(self, method: str):
639
- url = f"http://localhost:{self.upstream_port}{self.path}"
663
+ base_url = self.upstream_base_url or f"http://localhost:{self.upstream_port}"
664
+ url = f"{base_url}{self.path}"
640
665
  length = int(self.headers.get("Content-Length", 0))
641
666
  body = self.rfile.read(length) if length else None
642
667
  forward_headers = {
@@ -759,12 +784,13 @@ def wait_ready(
759
784
  port: int,
760
785
  timeout: float,
761
786
  health_path: str = "/health",
787
+ base_url: str | None = None,
762
788
  poll_interval: float = 5.0,
763
789
  request_timeout: float = 5.0,
764
790
  ) -> None:
765
791
  """Poll an HTTP health endpoint until ready, raising if the process dies."""
766
792
  deadline = time.time() + timeout
767
- url = f"http://127.0.0.1:{port}{health_path}"
793
+ url = f"{base_url or _local_base_url('127.0.0.1', port)}{health_path}"
768
794
  last_error = "no response yet"
769
795
 
770
796
  while time.time() < deadline:
@@ -794,6 +820,7 @@ def wait_ready(
794
820
  def warmup_chat_completions(
795
821
  *,
796
822
  port: int,
823
+ base_url: str | None = None,
797
824
  payload: Mapping[str, Any],
798
825
  headers: Mapping[str, str] | None = None,
799
826
  successful_requests: int = 3,
@@ -802,7 +829,7 @@ def warmup_chat_completions(
802
829
  retry_delay: float = 1.0,
803
830
  ) -> None:
804
831
  """Warm the OpenAI chat completions endpoint with strict retries."""
805
- url = f"http://127.0.0.1:{port}/v1/chat/completions"
832
+ url = f"{base_url or _local_base_url('127.0.0.1', port)}/v1/chat/completions"
806
833
  request_headers = {"Content-Type": "application/json"}
807
834
  if headers:
808
835
  request_headers.update(headers)
@@ -0,0 +1,372 @@
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
+ server_arg_tokens,
23
+ )
24
+ from .router import verify_router
25
+
26
+ # PDEndpoint owns the topology and every flag it mirrors in its own supervision
27
+ # (drain deadline, health probes); recipes tune everything else via router_args.
28
+ _MANAGED_ROUTER_FLAGS = frozenset(
29
+ {
30
+ "--host",
31
+ "--port",
32
+ "--pd-disaggregation",
33
+ "--prefill",
34
+ "--decode",
35
+ "--api-key",
36
+ "--shutdown-grace-period-secs",
37
+ "--health-check-interval-secs",
38
+ "--health-check-timeout-secs",
39
+ "--health-failure-threshold",
40
+ }
41
+ )
42
+
43
+
44
+ def _managed_router_tokens(router_args: Mapping[str, str]) -> list[str]:
45
+ """Rendered override tokens that would reach a managed flag.
46
+
47
+ Checks argv as the router parses it: values split into extra tokens, and
48
+ argparse accepts any unambiguous prefix of a long flag.
49
+ """
50
+ tokens = [
51
+ token
52
+ for key, value in router_args.items()
53
+ for token in server_arg_tokens(key, value)
54
+ ]
55
+ return sorted(
56
+ {
57
+ token
58
+ for token in tokens
59
+ if token.startswith("-")
60
+ and any(
61
+ flag.startswith(_flag_name(token)) for flag in _MANAGED_ROUTER_FLAGS
62
+ )
63
+ }
64
+ )
65
+
66
+
67
+ def _configure_fabric() -> None:
68
+ if os.environ.get("FI_PROVIDER") == "efa":
69
+ os.environ["SGLANG_DISAGGREGATION_NIXL_BACKEND"] = "LIBFABRIC"
70
+ elif hca := os.environ.get("NCCL_IB_HCA", ""):
71
+ rails = [name for name in hca.lstrip("=").split(",") if name]
72
+ os.environ["UCX_NET_DEVICES"] = ",".join(f"{name}:1" for name in rails)
73
+
74
+
75
+ async def _stop_container() -> None:
76
+ from modal.client import _Client
77
+ from modal_proto import api_pb2
78
+
79
+ async def stop():
80
+ client = await _Client.from_env()
81
+ await client.stub.ContainerStop(
82
+ api_pb2.ContainerStopRequest(task_id=os.environ["MODAL_TASK_ID"])
83
+ )
84
+
85
+ await asyncio.wait_for(stop(), timeout=3)
86
+
87
+
88
+ def _terminate_container() -> None:
89
+ from modal._utils.async_utils import synchronize_api
90
+
91
+ try:
92
+ synchronize_api(_stop_container)()
93
+ finally:
94
+ os._exit(1)
95
+
96
+
97
+ class _PDRouter(RouterEndpoint):
98
+ def __init__(
99
+ self,
100
+ *,
101
+ max_concurrent_requests,
102
+ drain_timeout,
103
+ health_interval,
104
+ health_request_timeout,
105
+ health_failure_threshold,
106
+ router_args,
107
+ **kwargs,
108
+ ):
109
+ super().__init__(health_path="/readiness", host="::", **kwargs)
110
+ self.max_concurrent_requests = max_concurrent_requests
111
+ self.drain_timeout = drain_timeout
112
+ self.health_interval = health_interval
113
+ self.health_request_timeout = health_request_timeout
114
+ self.health_failure_threshold = health_failure_threshold
115
+ self.router_args = dict(router_args)
116
+
117
+ def _build_cmd(self) -> list[str]:
118
+ cmd = [
119
+ "python",
120
+ "-m",
121
+ "sglang_router.launch_router",
122
+ "--host",
123
+ "[::]",
124
+ "--port",
125
+ str(self.router_port),
126
+ "--pd-disaggregation",
127
+ ]
128
+ for role, host in self.pd_config:
129
+ cmd.extend([f"--{role}", _url(host, self.worker_port)])
130
+ if role == "prefill":
131
+ cmd.append(str(self.prefill_bootstrap_port))
132
+ defaults = {
133
+ "--prefill-policy": "cache_aware",
134
+ "--decode-policy": "round_robin",
135
+ "--max-concurrent-requests": str(self.max_concurrent_requests),
136
+ "--queue-size": "0",
137
+ "--rate-limit-tokens-per-second": "0",
138
+ "--health-check-interval-secs": str(self.health_interval),
139
+ "--health-check-timeout-secs": str(self.health_request_timeout),
140
+ "--health-failure-threshold": str(self.health_failure_threshold),
141
+ "--shutdown-grace-period-secs": str(self.drain_timeout),
142
+ "--request-timeout-secs": "1800",
143
+ "--log-level": "info",
144
+ }
145
+ for flag, value in _merge_server_args(defaults, self.router_args).items():
146
+ cmd.extend(server_arg_tokens(flag, value))
147
+ if self.api_key is not None:
148
+ cmd.extend(["--api-key", self.api_key])
149
+ return cmd
150
+
151
+
152
+ class PDEndpoint(Endpoint):
153
+ """Own a local SGLang engine and the rank-zero router for P/D serving.
154
+
155
+ Call start/stop from Modal enter/exit hooks. Model-specific arguments and
156
+ warmup stay in the recipe. The image must support NIXL and graceful shutdown.
157
+ The ratio is (prefill, decode) counts; prefill ranks precede decode ranks.
158
+ The router is the build pinned in autoinference_utils.router; router_args
159
+ override its default flags, and allow_custom_router skips the pin check.
160
+ """
161
+
162
+ def __init__(
163
+ self,
164
+ engine: SGLangEndpoint,
165
+ *,
166
+ ratio: tuple[int, int],
167
+ prefill_args: Mapping[str, str] | None = None,
168
+ decode_args: Mapping[str, str] | None = None,
169
+ router_port: int = 9000,
170
+ max_concurrent_requests: int = 96,
171
+ api_key: str | None = None,
172
+ drain_timeout: int = 120,
173
+ health_interval: int = 10,
174
+ health_request_timeout: int = 10,
175
+ router_health_request_timeout: int = 5,
176
+ health_failure_threshold: int = 3,
177
+ health_failure_timeout: float = 30,
178
+ router_args: Mapping[str, str] | None = None,
179
+ allow_custom_router: bool = False,
180
+ ):
181
+ from modal.experimental import get_cluster_info
182
+
183
+ if not isinstance(engine, SGLangEndpoint):
184
+ raise TypeError("PDEndpoint requires a SGLangEndpoint")
185
+ if engine._proc is not None or engine.bench_mode:
186
+ raise ValueError("PDEndpoint requires an unstarted, non-benchmark engine")
187
+ if (
188
+ min(
189
+ health_interval,
190
+ health_request_timeout,
191
+ router_health_request_timeout,
192
+ health_failure_threshold,
193
+ )
194
+ <= 0
195
+ or drain_timeout < 0
196
+ or not math.isfinite(health_failure_timeout)
197
+ or health_failure_timeout <= 0
198
+ ):
199
+ raise ValueError(
200
+ "health settings must be positive and drain_timeout nonnegative"
201
+ )
202
+ if len(ratio) != 2 or any(
203
+ type(count) is not int or count < 1 for count in ratio
204
+ ):
205
+ raise ValueError(
206
+ "ratio must contain positive integer prefill and decode counts"
207
+ )
208
+ if managed := _managed_router_tokens(router_args or {}):
209
+ raise ValueError(
210
+ f"router_args cannot set PDEndpoint-managed flags: {managed}"
211
+ )
212
+ roles = ("prefill",) * ratio[0] + ("decode",) * ratio[1]
213
+ cluster = get_cluster_info()
214
+ hosts = cluster.container_ips
215
+ if len(hosts) != len(roles) or not 0 <= cluster.rank < len(roles):
216
+ raise ValueError(
217
+ f"P/D ratio {ratio} requires a {len(roles)}-container Modal cluster"
218
+ )
219
+ self._host_ip = hosts[cluster.rank]
220
+ self.role = roles[cluster.rank]
221
+ super().__init__(_url(hosts[0], router_port))
222
+ self.engine = engine
223
+ engine.disaggregation_mode = self.role
224
+ engine.health_request_timeout = health_request_timeout
225
+ engine.extra_server_args = _merge_server_args(
226
+ engine.extra_server_args,
227
+ (prefill_args if self.role == "prefill" else decode_args) or {},
228
+ )
229
+ engine.extra_server_args = _merge_server_args(
230
+ engine.extra_server_args,
231
+ {
232
+ "--disaggregation-transfer-backend": "nixl",
233
+ "--host": "::",
234
+ },
235
+ )
236
+ engine.extra_server_args = {
237
+ key: value
238
+ for key, value in engine.extra_server_args.items()
239
+ if _flag_name(key) != "--disaggregation-mode"
240
+ }
241
+ engine.base_url = _local_base_url("::", engine.worker_port)
242
+ self.router = (
243
+ _PDRouter(
244
+ pd_config=list(zip(roles, hosts)),
245
+ worker_port=engine.worker_port,
246
+ router_port=router_port,
247
+ prefill_bootstrap_port=engine.prefill_bootstrap_port,
248
+ health_timeout=engine.health_timeout,
249
+ api_key=api_key,
250
+ max_concurrent_requests=max_concurrent_requests,
251
+ drain_timeout=drain_timeout,
252
+ health_interval=health_interval,
253
+ health_request_timeout=router_health_request_timeout,
254
+ health_failure_threshold=health_failure_threshold,
255
+ router_args=router_args or {},
256
+ )
257
+ if cluster.rank == 0
258
+ else None
259
+ )
260
+ # Fail before any weights load, not after a 20-minute engine boot.
261
+ if self.router is not None and not allow_custom_router:
262
+ verify_router()
263
+ self.drain_timeout = drain_timeout
264
+ self.health_interval = health_interval
265
+ self.health_failure_timeout = health_failure_timeout
266
+ self._stopped = threading.Event()
267
+ self._lock = threading.Lock()
268
+ self._started = False
269
+ self._threads: list[threading.Thread] = []
270
+
271
+ def start(self, *, warmup: Callable[[], None] | None = None) -> None:
272
+ """Run warmup on rank zero before supervising every worker's health."""
273
+ if self._started or self._stopped.is_set():
274
+ raise RuntimeError("PDEndpoint can only be started once")
275
+ self._started = True
276
+ os.environ.setdefault("SGLANG_HOST_IP", self._host_ip)
277
+ _configure_fabric()
278
+ os.environ.update(
279
+ SGLANG_GRACEFUL_SHUTDOWN_TIMEOUT=str(self.drain_timeout),
280
+ SGLANG_ENABLE_HEALTH_ENDPOINT_GENERATION="1",
281
+ )
282
+ try:
283
+ for endpoint in (self.engine, self.router):
284
+ if endpoint is not None:
285
+ endpoint.start()
286
+ self._watch_process(endpoint)
287
+ if self.router is not None:
288
+ if warmup is not None:
289
+ warmup()
290
+ self.router.wait_ready()
291
+ for _, host in self.router.pd_config:
292
+ self._watch_health(_url(host, self.engine.worker_port) + "/health")
293
+ self._watch_health(self.router.base_url + "/health")
294
+ except BaseException:
295
+ self.stop()
296
+ raise
297
+
298
+ def _watch_process(self, endpoint: SGLangEndpoint | RouterEndpoint) -> None:
299
+ process = endpoint._proc
300
+ assert process is not None
301
+
302
+ def wait():
303
+ code = process.wait()
304
+ if not self._stopped.is_set():
305
+ print(f"[pd] Serving process exited with code {code}", flush=True)
306
+ self._fail()
307
+
308
+ thread = threading.Thread(target=wait, daemon=True)
309
+ self._threads.append(thread)
310
+ thread.start()
311
+
312
+ def _watch_health(self, url: str) -> None:
313
+ def poll():
314
+ failed_since = None
315
+ while not self._stopped.wait(self.health_interval):
316
+ started = time.monotonic()
317
+ error = _health_check(
318
+ url, request_timeout=self.engine.health_request_timeout
319
+ )
320
+ if error is None:
321
+ failed_since = None
322
+ continue
323
+ if failed_since is None:
324
+ failed_since = started
325
+ elapsed = time.monotonic() - failed_since
326
+ print(f"[pd] {url}: {error}; unhealthy for {elapsed:.1f}s", flush=True)
327
+ if elapsed >= self.health_failure_timeout:
328
+ self._fail()
329
+ return
330
+
331
+ thread = threading.Thread(target=poll, daemon=True)
332
+ self._threads.append(thread)
333
+ thread.start()
334
+
335
+ def _fail(self) -> None:
336
+ with self._lock:
337
+ if self._stopped.is_set():
338
+ return
339
+ self._stopped.set()
340
+ try:
341
+ for endpoint in (self.router, self.engine):
342
+ if endpoint is not None and endpoint._proc is not None:
343
+ try:
344
+ endpoint._proc.kill()
345
+ except ProcessLookupError:
346
+ pass
347
+ finally:
348
+ _terminate_container()
349
+
350
+ def stop(self) -> None:
351
+ with self._lock:
352
+ if self._stopped.is_set():
353
+ return
354
+ self._stopped.set()
355
+ deadline = time.monotonic() + self.drain_timeout + 30
356
+ try:
357
+ for endpoint in (self.router, self.engine):
358
+ if endpoint is None:
359
+ continue
360
+ process = endpoint._proc
361
+ if process is not None and process.poll() is None:
362
+ process.terminate()
363
+ try:
364
+ process.wait(timeout=max(0, deadline - time.monotonic()))
365
+ except subprocess.TimeoutExpired:
366
+ process.kill()
367
+ endpoint.stop()
368
+ finally:
369
+ for thread in self._threads:
370
+ thread.join(
371
+ timeout=self.engine.health_request_timeout + self.health_interval
372
+ )
@@ -0,0 +1,103 @@
1
+ """The pinned SGLang model gateway (SMG) wheel that PDEndpoint routes through.
2
+
3
+ This module is the source of truth for the P/D router. The wheel is built once
4
+ from the fork commit below and published as a GitHub release asset; bumping
5
+ the pin here and releasing autoinference-utils moves every recipe that takes
6
+ the new release. Install it inside an image after autoinference-utils:
7
+
8
+ image.run_commands("python -m autoinference_utils.router install")
9
+ """
10
+
11
+ from __future__ import annotations
12
+
13
+ import json
14
+ import subprocess
15
+ import sys
16
+ from importlib.metadata import Distribution, PackageNotFoundError, distribution
17
+
18
+ ROUTER_REPOSITORY = "https://github.com/modal-labs/sglang"
19
+ # Kimi K3 production router: HRRN decode preallocation (SMG_PD_DECODE_HRRN)
20
+ # and brace-free metrics HELP text (kimi-k3-sglang#158). The wheel is the one
21
+ # extracted from the Kimi K3 production image built at this commit.
22
+ ROUTER_COMMIT = "e3e602e3ae46fa4da7ec105d5dc3efcfc85c226b"
23
+ # The wheel is x86-64 Linux and links against glibc 2.39 (Ubuntu 24.04, which
24
+ # the lmsysorg/sglang and vllm base images use); uv refuses it on older images.
25
+ ROUTER_WHEEL_URL = (
26
+ f"{ROUTER_REPOSITORY}/releases/download/smg-router-{ROUTER_COMMIT[:8]}/"
27
+ "sglang_router-0.3.2-cp38-abi3-manylinux_2_39_x86_64.whl"
28
+ )
29
+ ROUTER_WHEEL_SHA256 = "5b867c6c308a609395a6f5c146023274f570b7187b3a7b92ebb11658da52a6c7"
30
+
31
+
32
+ def router_requirement() -> str:
33
+ """The hash-pinned direct requirement for the router wheel."""
34
+ return f"sglang-router @ {ROUTER_WHEEL_URL}#sha256={ROUTER_WHEEL_SHA256}"
35
+
36
+
37
+ def verify_router(dist: Distribution | None = None) -> None:
38
+ """Raise unless the installed sglang-router came from the pinned wheel.
39
+
40
+ Checks the installer's PEP 610 record: the URL must match, and when the
41
+ installer recorded a digest it must match too. uv leaves archive_info empty
42
+ and enforces the digest at install time through the URL's sha256 fragment.
43
+ """
44
+ try:
45
+ dist = dist or distribution("sglang-router")
46
+ origin = json.loads(dist.read_text("direct_url.json") or "{}")
47
+ except PackageNotFoundError:
48
+ origin = None
49
+ except ValueError:
50
+ origin = {}
51
+ if origin is None:
52
+ found = "no sglang-router is installed"
53
+ else:
54
+ recorded = ((origin.get("archive_info") or {}).get("hashes") or {}).get(
55
+ "sha256"
56
+ )
57
+ if origin.get("url") == ROUTER_WHEEL_URL and recorded in (
58
+ None,
59
+ ROUTER_WHEEL_SHA256,
60
+ ):
61
+ return
62
+ found = f"installed sglang-router is {origin.get('url') or 'not a direct URL install'}"
63
+ if recorded is not None:
64
+ found += f" (sha256 {recorded})"
65
+ raise RuntimeError(
66
+ f"{found}; expected {ROUTER_WHEEL_URL} (sha256 {ROUTER_WHEEL_SHA256}). "
67
+ "Run `python -m autoinference_utils.router install` in the image after "
68
+ "autoinference-utils, or pass allow_custom_router=True."
69
+ )
70
+
71
+
72
+ def install() -> None:
73
+ """Install the pinned wheel over whatever router the image ships."""
74
+ # --no-deps keeps the engine image's own web stack untouched.
75
+ subprocess.run(
76
+ [
77
+ "uv",
78
+ "pip",
79
+ "install",
80
+ "--python",
81
+ sys.executable,
82
+ "--no-cache",
83
+ "--no-deps",
84
+ "--reinstall-package",
85
+ "sglang-router",
86
+ router_requirement(),
87
+ ],
88
+ check=True,
89
+ )
90
+ verify_router()
91
+ # --no-deps means a pin whose dependencies the image lacks must fail here,
92
+ # at image build, not at container start.
93
+ subprocess.run(
94
+ [sys.executable, "-m", "sglang_router.launch_router", "--help"],
95
+ check=True,
96
+ stdout=subprocess.DEVNULL,
97
+ )
98
+
99
+
100
+ if __name__ == "__main__":
101
+ if sys.argv[1:] != ["install"]:
102
+ raise SystemExit("usage: python -m autoinference_utils.router install")
103
+ install()
@@ -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-*"))