autoinference-utils 0.2.8__tar.gz → 0.2.10__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.8
3
+ Version: 0.2.10
4
4
  Summary: Shared endpoint abstractions for autoinference deployments
5
5
  Requires-Python: >=3.10
6
6
  Description-Content-Type: text/markdown
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "autoinference-utils"
3
- version = "0.2.8"
3
+ version = "0.2.10"
4
4
  description = "Shared endpoint abstractions for autoinference deployments"
5
5
  readme = "README.md"
6
6
  requires-python = ">=3.10"
@@ -27,7 +27,7 @@ BENCH_MODE_ENV = "ENDPOINT_BENCH_MODE"
27
27
  DUMMY_WEIGHTS_ENV = "ENDPOINT_DUMMY"
28
28
  BENCH_MODE_PORT_OFFSET = 10000
29
29
 
30
- MODAL_FLASH_REQUEST_UUID_HEADER= "modal-flash-request-uuid"
30
+ MODAL_FLASH_REQUEST_UUID_HEADER = "modal-flash-request-uuid"
31
31
  MODAL_SESSION_ID_HEADER = "modal-session-id"
32
32
 
33
33
  ENDPOINTS_REQUIRING_OAI_STREAM_COMPAT: tuple[str, ...] = ("TRTLLMEndpoint",)
@@ -63,9 +63,7 @@ def _flag_name(key: str) -> str:
63
63
  return key.partition("=")[0]
64
64
 
65
65
 
66
- def get_server_arg(
67
- server_args: Mapping[str, str], flag: str
68
- ) -> tuple[str, str] | None:
66
+ def get_server_arg(server_args: Mapping[str, str], flag: str) -> tuple[str, str] | None:
69
67
  """Return the key and value setting a flag under either spelling."""
70
68
  for key, value in server_args.items():
71
69
  if _flag_name(key) == flag:
@@ -176,9 +174,7 @@ class SGLangEndpoint(Endpoint):
176
174
  self.disaggregation_mode = disaggregation_mode
177
175
  self.prefill_bootstrap_port = prefill_bootstrap_port
178
176
  self.launcher_module = launcher_module
179
- self.extra_server_args = (
180
- dict(extra_server_args) if extra_server_args else {}
181
- )
177
+ self.extra_server_args = dict(extra_server_args) if extra_server_args else {}
182
178
  self.health_timeout = health_timeout
183
179
  self.health_poll_interval = health_poll_interval
184
180
  self.health_request_timeout = health_request_timeout
@@ -193,17 +189,16 @@ class SGLangEndpoint(Endpoint):
193
189
  self._bench_server: Optional[http.server.ThreadingHTTPServer] = None
194
190
 
195
191
  if self.disaggregation_mode not in (None, "prefill", "decode"):
196
- raise ValueError(
197
- "disaggregation_mode must be None, 'prefill', or 'decode'"
198
- )
192
+ raise ValueError("disaggregation_mode must be None, 'prefill', or 'decode'")
199
193
 
200
194
  if os.environ.get(DUMMY_WEIGHTS_ENV) == "1":
201
- if (
202
- self.load_format is None
203
- and not has_server_arg(self.extra_server_args, "--load-format")
195
+ if self.load_format is None and not has_server_arg(
196
+ self.extra_server_args, "--load-format"
204
197
  ):
205
198
  self.load_format = "dummy"
206
- if not has_server_arg(self.extra_server_args, "--model-loader-extra-config"):
199
+ if not has_server_arg(
200
+ self.extra_server_args, "--model-loader-extra-config"
201
+ ):
207
202
  self.extra_server_args["--model-loader-extra-config"] = "{}"
208
203
 
209
204
  # Request logging enabled, log only metadata by default
@@ -214,7 +209,9 @@ class SGLangEndpoint(Endpoint):
214
209
  # Set request logging level
215
210
  level = get_server_arg(self.extra_server_args, "--log-requests-level")
216
211
  if level is None:
217
- self.extra_server_args["--log-requests-level"] = str(self.log_requests_level)
212
+ self.extra_server_args["--log-requests-level"] = str(
213
+ self.log_requests_level
214
+ )
218
215
  else:
219
216
  key, value = level
220
217
  print(
@@ -236,19 +233,21 @@ class SGLangEndpoint(Endpoint):
236
233
  self.log_request_headers = ",".join(sglang_log_request_headers)
237
234
  os.environ["SGLANG_LOG_REQUEST_HEADERS"] = self.log_request_headers
238
235
 
239
-
240
236
  def _build_cmd(self) -> list[str]:
241
237
  cmd = [
242
- "python", "-m", self.launcher_module,
243
- "--host", "0.0.0.0",
244
- "--port", str(self.worker_port),
245
- "--model-path", self.model_path,
238
+ "python",
239
+ "-m",
240
+ self.launcher_module,
241
+ "--host",
242
+ "0.0.0.0",
243
+ "--port",
244
+ str(self.worker_port),
245
+ "--model-path",
246
+ self.model_path,
246
247
  ]
247
248
 
248
249
  if self.speculative_model_path is not None:
249
- cmd.extend(
250
- ["--speculative-draft-model-path", self.speculative_model_path]
251
- )
250
+ cmd.extend(["--speculative-draft-model-path", self.speculative_model_path])
252
251
  if self.load_format is not None:
253
252
  cmd.extend(["--load-format", self.load_format])
254
253
  if self.tp is not None:
@@ -273,14 +272,18 @@ class SGLangEndpoint(Endpoint):
273
272
  raise ValueError("dist_init_host is required when nnodes > 1")
274
273
  cmd.extend(
275
274
  [
276
- "--nnodes", str(self.nnodes),
277
- "--node-rank", str(self.node_rank),
275
+ "--nnodes",
276
+ str(self.nnodes),
277
+ "--node-rank",
278
+ str(self.node_rank),
278
279
  "--dist-init-addr",
279
280
  f"{self.dist_init_host}:{self.dist_init_port}",
280
281
  ]
281
282
  )
282
283
 
283
- merged = _merge_server_args(self.DEFAULT_OPERATIONAL_ARGS, self.extra_server_args)
284
+ merged = _merge_server_args(
285
+ self.DEFAULT_OPERATIONAL_ARGS, self.extra_server_args
286
+ )
284
287
  for key, value in merged.items():
285
288
  cmd.extend(server_arg_tokens(key, value))
286
289
 
@@ -345,9 +348,7 @@ class VLLMEndpoint(Endpoint):
345
348
  super().__init__(base_url=f"http://localhost:{vllm_port}")
346
349
  self.model = model
347
350
  self.worker_port = vllm_port
348
- self.extra_server_args = (
349
- dict(extra_server_args) if extra_server_args else {}
350
- )
351
+ self.extra_server_args = dict(extra_server_args) if extra_server_args else {}
351
352
  self.health_timeout = health_timeout
352
353
  self.health_poll_interval = health_poll_interval
353
354
  self.health_request_timeout = health_request_timeout
@@ -360,10 +361,15 @@ class VLLMEndpoint(Endpoint):
360
361
 
361
362
  def _build_cmd(self) -> list[str]:
362
363
  cmd = [
363
- "python", "-m", "vllm.entrypoints.openai.api_server",
364
- "--host", "0.0.0.0",
365
- "--port", str(self.worker_port),
366
- "--model", self.model,
364
+ "python",
365
+ "-m",
366
+ "vllm.entrypoints.openai.api_server",
367
+ "--host",
368
+ "0.0.0.0",
369
+ "--port",
370
+ str(self.worker_port),
371
+ "--model",
372
+ self.model,
367
373
  ]
368
374
  for key, value in self.extra_server_args.items():
369
375
  cmd.extend(server_arg_tokens(key, value))
@@ -417,9 +423,7 @@ class TRTLLMEndpoint(Endpoint):
417
423
  super().__init__(base_url=f"http://localhost:{worker_port}")
418
424
  self.model = model
419
425
  self.worker_port = worker_port
420
- self.extra_server_args = (
421
- dict(extra_server_args) if extra_server_args else {}
422
- )
426
+ self.extra_server_args = dict(extra_server_args) if extra_server_args else {}
423
427
  self.health_timeout = health_timeout
424
428
  self.health_poll_interval = health_poll_interval
425
429
  self.health_request_timeout = health_request_timeout
@@ -429,8 +433,10 @@ class TRTLLMEndpoint(Endpoint):
429
433
  cmd = [
430
434
  "trtllm-serve",
431
435
  self.model,
432
- "--host", "0.0.0.0",
433
- "--port", str(self.worker_port),
436
+ "--host",
437
+ "0.0.0.0",
438
+ "--port",
439
+ str(self.worker_port),
434
440
  ]
435
441
  for key, value in self.extra_server_args.items():
436
442
  cmd.extend(server_arg_tokens(key, value))
@@ -499,18 +505,30 @@ class RouterEndpoint(Endpoint):
499
505
 
500
506
  def _build_cmd(self) -> list[str]:
501
507
  cmd = [
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",
508
+ "python",
509
+ "-m",
510
+ "sglang_router.launch_router",
511
+ "--host",
512
+ f"[{self.host}]" if ":" in self.host else self.host,
513
+ "--port",
514
+ str(self.router_port),
515
+ "--prefill-policy",
516
+ "cache_aware",
517
+ "--decode-policy",
518
+ "round_robin",
519
+ "--max-concurrent-requests",
520
+ "128",
521
+ "--rate-limit-tokens-per-second",
522
+ "0",
523
+ "--queue-size",
524
+ "0",
525
+ "--health-check-timeout-secs",
526
+ "600",
527
+ "--log-level",
528
+ "info",
512
529
  "--disable-circuit-breaker",
513
- "--request-timeout-secs", "3600",
530
+ "--request-timeout-secs",
531
+ "3600",
514
532
  ]
515
533
 
516
534
  if self.api_key is not None:
@@ -522,9 +540,7 @@ class RouterEndpoint(Endpoint):
522
540
  for role, node_ip in self.pd_config:
523
541
  node_url = _url(node_ip, self.worker_port)
524
542
  if role == "prefill":
525
- cmd.extend(
526
- ["--prefill", node_url, str(self.prefill_bootstrap_port)]
527
- )
543
+ cmd.extend(["--prefill", node_url, str(self.prefill_bootstrap_port)])
528
544
  elif role == "decode":
529
545
  cmd.extend(["--decode", node_url])
530
546
  elif role == "worker":
@@ -576,6 +592,7 @@ class RouterEndpoint(Endpoint):
576
592
  # and this proxy listens on worker_port. POST /bench shells out to a benchmark
577
593
  # task via run_bench(); every other path is forwarded to the upstream server.
578
594
 
595
+
579
596
  def start_bench_proxy(
580
597
  *,
581
598
  listen_port: int,
@@ -587,17 +604,12 @@ def start_bench_proxy(
587
604
  (_BenchProxyHandler,),
588
605
  {"upstream_port": upstream_port, "upstream_base_url": upstream_base_url},
589
606
  )
590
- server = http.server.ThreadingHTTPServer(
591
- ("0.0.0.0", listen_port), handler_cls
592
- )
607
+ server = http.server.ThreadingHTTPServer(("0.0.0.0", listen_port), handler_cls)
593
608
  thread = threading.Thread(
594
609
  target=server.serve_forever, daemon=True, name="bench-proxy"
595
610
  )
596
611
  thread.start()
597
- print(
598
- f"[bench-proxy] listening on :{listen_port}, "
599
- f"forwarding to :{upstream_port}"
600
- )
612
+ print(f"[bench-proxy] listening on :{listen_port}, forwarding to :{upstream_port}")
601
613
  return server
602
614
 
603
615
 
@@ -627,9 +639,7 @@ class _BenchProxyHandler(http.server.BaseHTTPRequestHandler):
627
639
  try:
628
640
  payload = self._read_json_body()
629
641
  except (ValueError, json.JSONDecodeError) as exc:
630
- return self._send_json(
631
- 400, {"ok": False, "error": f"bad body: {exc}"}
632
- )
642
+ return self._send_json(400, {"ok": False, "error": f"bad body: {exc}"})
633
643
 
634
644
  benchmark = payload.get("benchmark") or ""
635
645
  args = payload.get("args") or []
@@ -641,13 +651,9 @@ class _BenchProxyHandler(http.server.BaseHTTPRequestHandler):
641
651
  output_dir = payload.get("output_dir") or "/tmp/bench-output"
642
652
 
643
653
  if not benchmark:
644
- return self._send_json(
645
- 400, {"ok": False, "error": "benchmark is required"}
646
- )
654
+ return self._send_json(400, {"ok": False, "error": "benchmark is required"})
647
655
  if not isinstance(args, list):
648
- return self._send_json(
649
- 400, {"ok": False, "error": "args must be a list"}
650
- )
656
+ return self._send_json(400, {"ok": False, "error": "args must be a list"})
651
657
 
652
658
  rc, body = run_bench(benchmark, list(args), target, output_dir)
653
659
  status = 200 if rc == 0 else 500
@@ -699,9 +705,7 @@ class _BenchProxyHandler(http.server.BaseHTTPRequestHandler):
699
705
  return json.loads(raw)
700
706
 
701
707
  def _send_json(self, status: int, obj: Mapping[str, Any]):
702
- self._send_raw(
703
- status, "application/json", json.dumps(obj).encode()
704
- )
708
+ self._send_raw(status, "application/json", json.dumps(obj).encode())
705
709
 
706
710
  def _send_raw(self, status: int, content_type: str, body: bytes):
707
711
  self.send_response(status)
@@ -726,18 +730,21 @@ def run_bench(
726
730
  return 500, json.dumps(
727
731
  {
728
732
  "ok": False,
729
- "error": (
730
- "autoinference.benchmarks not importable in container"
731
- ),
733
+ "error": ("autoinference.benchmarks not importable in container"),
732
734
  "volume_path": output_dir,
733
735
  }
734
736
  ).encode()
735
737
 
736
738
  cmd = [
737
- sys.executable, "-m", "invoke",
738
- "--search-root", search_root,
739
- "-c", "tasks",
740
- benchmark, *[str(a) for a in args],
739
+ sys.executable,
740
+ "-m",
741
+ "invoke",
742
+ "--search-root",
743
+ search_root,
744
+ "-c",
745
+ "tasks",
746
+ benchmark,
747
+ *[str(a) for a in args],
741
748
  ]
742
749
  env = {
743
750
  **os.environ,
@@ -788,9 +795,7 @@ def wait_ready(
788
795
 
789
796
  while time.time() < deadline:
790
797
  try:
791
- error = _health_check(
792
- url, request_timeout=request_timeout, process=process
793
- )
798
+ error = _health_check(url, request_timeout=request_timeout, process=process)
794
799
  except subprocess.CalledProcessError as exc:
795
800
  print(
796
801
  f"[endpoint] !!! server process exited with code "
@@ -804,13 +809,11 @@ def wait_ready(
804
809
  time.sleep(poll_interval)
805
810
 
806
811
  print(
807
- f"[endpoint] !!! health check timed out after {timeout}s; "
808
- f"last={last_error}",
812
+ f"[endpoint] !!! health check timed out after {timeout}s; last={last_error}",
809
813
  flush=True,
810
814
  )
811
815
  raise TimeoutError(
812
- f"Health check timed out after {timeout}s for {url}. "
813
- f"Last error: {last_error}"
816
+ f"Health check timed out after {timeout}s for {url}. Last error: {last_error}"
814
817
  )
815
818
 
816
819
 
@@ -839,9 +842,7 @@ def warmup_chat_completions(
839
842
  request_timeout=request_timeout,
840
843
  max_attempts=max_attempts_per_request,
841
844
  retry_delay=retry_delay,
842
- description=(
843
- f"warmup request {request_idx + 1}/{successful_requests}"
844
- ),
845
+ description=(f"warmup request {request_idx + 1}/{successful_requests}"),
845
846
  )
846
847
 
847
848
 
@@ -863,9 +864,7 @@ def validate_embeddings_endpoint(
863
864
  1 if all(isinstance(item, int) for item in inputs) else len(inputs)
864
865
  )
865
866
  else:
866
- raise ValueError(
867
- "embedding validation payload must contain non-empty input"
868
- )
867
+ raise ValueError("embedding validation payload must contain non-empty input")
869
868
 
870
869
  request_headers = {"Content-Type": "application/json"}
871
870
  if headers:
@@ -928,10 +927,7 @@ def start_heartbeat_thread(
928
927
  try:
929
928
  error = health_check_fn()
930
929
  except subprocess.CalledProcessError as exc:
931
- print(
932
- f"[heartbeat] server process exited with code "
933
- f"{exc.returncode}"
934
- )
930
+ print(f"[heartbeat] server process exited with code {exc.returncode}")
935
931
  on_failure()
936
932
  return
937
933
  if error is None:
@@ -943,10 +939,7 @@ def start_heartbeat_thread(
943
939
  f"({consecutive_failures}/{max_consecutive_failures})"
944
940
  )
945
941
  if consecutive_failures >= max_consecutive_failures:
946
- print(
947
- "[heartbeat] sustained health-check failure, "
948
- "invoking on_failure"
949
- )
942
+ print("[heartbeat] sustained health-check failure, invoking on_failure")
950
943
  on_failure()
951
944
  return
952
945
 
@@ -1031,8 +1024,7 @@ def _wait_ready_url(
1031
1024
  last_error = error
1032
1025
  time.sleep(poll_interval)
1033
1026
  raise TimeoutError(
1034
- f"Timed out after {timeout}s waiting for {url}. "
1035
- f"Last error: {last_error}"
1027
+ f"Timed out after {timeout}s waiting for {url}. Last error: {last_error}"
1036
1028
  )
1037
1029
 
1038
1030
 
@@ -1050,9 +1042,7 @@ def _post_json(
1050
1042
  timeout: float | None = None,
1051
1043
  ) -> bytes:
1052
1044
  body = json.dumps(payload).encode("utf-8")
1053
- req = urllib.request.Request(
1054
- url, data=body, headers=dict(headers), method="POST"
1055
- )
1045
+ req = urllib.request.Request(url, data=body, headers=dict(headers), method="POST")
1056
1046
  with urllib.request.urlopen(req, timeout=timeout) as resp:
1057
1047
  return resp.read()
1058
1048
 
@@ -5,6 +5,8 @@ from __future__ import annotations
5
5
  import asyncio
6
6
  import math
7
7
  import os
8
+ import socket
9
+ import struct
8
10
  import subprocess
9
11
  import threading
10
12
  import time
@@ -19,7 +21,52 @@ from .endpoint import (
19
21
  _local_base_url,
20
22
  _merge_server_args,
21
23
  _url,
24
+ server_arg_tokens,
22
25
  )
26
+ from .router import verify_router
27
+
28
+ # PDEndpoint owns the topology and every flag it mirrors in its own supervision
29
+ # (drain deadline, health probes); recipes tune everything else via router_args.
30
+ _MANAGED_ROUTER_FLAGS = frozenset(
31
+ {
32
+ "--host",
33
+ "--port",
34
+ "--pd-disaggregation",
35
+ "--prefill",
36
+ "--decode",
37
+ "--api-key",
38
+ "--shutdown-grace-period-secs",
39
+ "--health-check-interval-secs",
40
+ "--health-check-timeout-secs",
41
+ "--health-failure-threshold",
42
+ }
43
+ )
44
+
45
+ # Leave headroom inside Modal's 30-second exit-handler deadline.
46
+ _SHUTDOWN_TIMEOUT = 25
47
+
48
+
49
+ def _managed_router_tokens(router_args: Mapping[str, str]) -> list[str]:
50
+ """Rendered override tokens that would reach a managed flag.
51
+
52
+ Checks argv as the router parses it: values split into extra tokens, and
53
+ argparse accepts any unambiguous prefix of a long flag.
54
+ """
55
+ tokens = [
56
+ token
57
+ for key, value in router_args.items()
58
+ for token in server_arg_tokens(key, value)
59
+ ]
60
+ return sorted(
61
+ {
62
+ token
63
+ for token in tokens
64
+ if token.startswith("-")
65
+ and any(
66
+ flag.startswith(_flag_name(token)) for flag in _MANAGED_ROUTER_FLAGS
67
+ )
68
+ }
69
+ )
23
70
 
24
71
 
25
72
  def _configure_fabric() -> None:
@@ -61,6 +108,7 @@ class _PDRouter(RouterEndpoint):
61
108
  health_interval,
62
109
  health_request_timeout,
63
110
  health_failure_threshold,
111
+ router_args,
64
112
  **kwargs,
65
113
  ):
66
114
  super().__init__(health_path="/readiness", host="::", **kwargs)
@@ -69,6 +117,7 @@ class _PDRouter(RouterEndpoint):
69
117
  self.health_interval = health_interval
70
118
  self.health_request_timeout = health_request_timeout
71
119
  self.health_failure_threshold = health_failure_threshold
120
+ self.router_args = dict(router_args)
72
121
 
73
122
  def _build_cmd(self) -> list[str]:
74
123
  cmd = [
@@ -85,20 +134,21 @@ class _PDRouter(RouterEndpoint):
85
134
  cmd.extend([f"--{role}", _url(host, self.worker_port)])
86
135
  if role == "prefill":
87
136
  cmd.append(str(self.prefill_bootstrap_port))
88
- for flag, value in {
137
+ defaults = {
89
138
  "--prefill-policy": "cache_aware",
90
139
  "--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,
140
+ "--max-concurrent-requests": str(self.max_concurrent_requests),
141
+ "--queue-size": "0",
142
+ "--rate-limit-tokens-per-second": "0",
143
+ "--health-check-interval-secs": str(self.health_interval),
144
+ "--health-check-timeout-secs": str(self.health_request_timeout),
145
+ "--health-failure-threshold": str(self.health_failure_threshold),
146
+ "--shutdown-grace-period-secs": str(self.drain_timeout),
147
+ "--request-timeout-secs": "1800",
99
148
  "--log-level": "info",
100
- }.items():
101
- cmd.extend([flag, str(value)])
149
+ }
150
+ for flag, value in _merge_server_args(defaults, self.router_args).items():
151
+ cmd.extend(server_arg_tokens(flag, value))
102
152
  if self.api_key is not None:
103
153
  cmd.extend(["--api-key", self.api_key])
104
154
  return cmd
@@ -110,6 +160,8 @@ class PDEndpoint(Endpoint):
110
160
  Call start/stop from Modal enter/exit hooks. Model-specific arguments and
111
161
  warmup stay in the recipe. The image must support NIXL and graceful shutdown.
112
162
  The ratio is (prefill, decode) counts; prefill ranks precede decode ranks.
163
+ The router is the build pinned in autoinference_utils.router; router_args
164
+ override its default flags, and allow_custom_router skips the pin check.
113
165
  """
114
166
 
115
167
  def __init__(
@@ -128,6 +180,9 @@ class PDEndpoint(Endpoint):
128
180
  router_health_request_timeout: int = 5,
129
181
  health_failure_threshold: int = 3,
130
182
  health_failure_timeout: float = 30,
183
+ router_args: Mapping[str, str] | None = None,
184
+ allow_custom_router: bool = False,
185
+ shutdown_port: int | None = None,
131
186
  ):
132
187
  from modal.experimental import get_cluster_info
133
188
 
@@ -156,6 +211,10 @@ class PDEndpoint(Endpoint):
156
211
  raise ValueError(
157
212
  "ratio must contain positive integer prefill and decode counts"
158
213
  )
214
+ if managed := _managed_router_tokens(router_args or {}):
215
+ raise ValueError(
216
+ f"router_args cannot set PDEndpoint-managed flags: {managed}"
217
+ )
159
218
  roles = ("prefill",) * ratio[0] + ("decode",) * ratio[1]
160
219
  cluster = get_cluster_info()
161
220
  hosts = cluster.container_ips
@@ -164,6 +223,13 @@ class PDEndpoint(Endpoint):
164
223
  f"P/D ratio {ratio} requires a {len(roles)}-container Modal cluster"
165
224
  )
166
225
  self._host_ip = hosts[cluster.rank]
226
+ self._hosts = hosts
227
+ self._rank = cluster.rank
228
+ self._shutdown_port = (
229
+ router_port + 1 if shutdown_port is None else shutdown_port
230
+ )
231
+ if not 1 <= self._shutdown_port <= 65535:
232
+ raise ValueError("shutdown_port must be between 1 and 65535")
167
233
  self.role = roles[cluster.rank]
168
234
  super().__init__(_url(hosts[0], router_port))
169
235
  self.engine = engine
@@ -199,16 +265,22 @@ class PDEndpoint(Endpoint):
199
265
  health_interval=health_interval,
200
266
  health_request_timeout=router_health_request_timeout,
201
267
  health_failure_threshold=health_failure_threshold,
268
+ router_args=router_args or {},
202
269
  )
203
270
  if cluster.rank == 0
204
271
  else None
205
272
  )
273
+ # Fail before any weights load, not after a 20-minute engine boot.
274
+ if self.router is not None and not allow_custom_router:
275
+ verify_router()
206
276
  self.drain_timeout = drain_timeout
207
277
  self.health_interval = health_interval
208
278
  self.health_failure_timeout = health_failure_timeout
209
279
  self._stopped = threading.Event()
210
280
  self._lock = threading.Lock()
211
281
  self._started = False
282
+ self._serving = False
283
+ self._shutdown_listener: socket.socket | None = None
212
284
  self._threads: list[threading.Thread] = []
213
285
 
214
286
  def start(self, *, warmup: Callable[[], None] | None = None) -> None:
@@ -223,6 +295,13 @@ class PDEndpoint(Endpoint):
223
295
  SGLANG_ENABLE_HEALTH_ENDPOINT_GENERATION="1",
224
296
  )
225
297
  try:
298
+ if self._rank == 0:
299
+ self._shutdown_listener = socket.socket(socket.AF_INET6)
300
+ self._shutdown_listener.setsockopt(
301
+ socket.SOL_SOCKET, socket.SO_REUSEADDR, 1
302
+ )
303
+ self._shutdown_listener.bind(("::", self._shutdown_port))
304
+ self._shutdown_listener.listen(len(self._hosts) - 1)
226
305
  for endpoint in (self.engine, self.router):
227
306
  if endpoint is not None:
228
307
  endpoint.start()
@@ -234,6 +313,7 @@ class PDEndpoint(Endpoint):
234
313
  for _, host in self.router.pd_config:
235
314
  self._watch_health(_url(host, self.engine.worker_port) + "/health")
236
315
  self._watch_health(self.router.base_url + "/health")
316
+ self._serving = True
237
317
  except BaseException:
238
318
  self.stop()
239
319
  raise
@@ -295,7 +375,7 @@ class PDEndpoint(Endpoint):
295
375
  if self._stopped.is_set():
296
376
  return
297
377
  self._stopped.set()
298
- deadline = time.monotonic() + self.drain_timeout + 30
378
+ deadline = time.monotonic() + _SHUTDOWN_TIMEOUT
299
379
  try:
300
380
  for endpoint in (self.router, self.engine):
301
381
  if endpoint is None:
@@ -307,9 +387,58 @@ class PDEndpoint(Endpoint):
307
387
  process.wait(timeout=max(0, deadline - time.monotonic()))
308
388
  except subprocess.TimeoutExpired:
309
389
  process.kill()
390
+ process.wait()
310
391
  endpoint.stop()
311
392
  finally:
312
393
  for thread in self._threads:
313
- thread.join(
314
- timeout=self.engine.health_request_timeout + self.health_interval
394
+ thread.join(timeout=max(0, deadline - time.monotonic()))
395
+ if self._serving:
396
+ if self._rank == 0:
397
+ self._wait_for_followers(deadline)
398
+ else:
399
+ self._notify_leader(deadline)
400
+ if self._shutdown_listener is not None:
401
+ self._shutdown_listener.close()
402
+ self._shutdown_listener = None
403
+
404
+ def _notify_leader(self, deadline: float) -> None:
405
+ while True:
406
+ remaining = deadline - time.monotonic()
407
+ if remaining <= 0:
408
+ print("[pd] Failed to notify rank zero of shutdown", flush=True)
409
+ return
410
+ try:
411
+ with socket.create_connection(
412
+ (self._hosts[0], self._shutdown_port), timeout=min(remaining, 1)
413
+ ) as connection:
414
+ connection.sendall(struct.pack("!I", self._rank))
415
+ return
416
+ except OSError:
417
+ time.sleep(min(0.1, max(0, deadline - time.monotonic())))
418
+
419
+ def _wait_for_followers(self, deadline: float) -> None:
420
+ listener = self._shutdown_listener
421
+ if listener is None:
422
+ return
423
+ pending = set(range(1, len(self._hosts)))
424
+ while pending:
425
+ remaining = deadline - time.monotonic()
426
+ if remaining <= 0:
427
+ print(
428
+ f"[pd] Timed out waiting for ranks {sorted(pending)} to stop",
429
+ flush=True,
315
430
  )
431
+ return
432
+ listener.settimeout(min(remaining, 1))
433
+ try:
434
+ connection, _ = listener.accept()
435
+ except TimeoutError:
436
+ continue
437
+ with connection:
438
+ connection.settimeout(min(remaining, 1))
439
+ try:
440
+ payload = connection.recv(4, socket.MSG_WAITALL)
441
+ except TimeoutError:
442
+ continue
443
+ if len(payload) == 4:
444
+ pending.discard(struct.unpack("!I", payload)[0])
@@ -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()