autoinference-utils 0.2.9__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.9
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.9"
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"
@@ -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
@@ -40,6 +42,9 @@ _MANAGED_ROUTER_FLAGS = frozenset(
40
42
  }
41
43
  )
42
44
 
45
+ # Leave headroom inside Modal's 30-second exit-handler deadline.
46
+ _SHUTDOWN_TIMEOUT = 25
47
+
43
48
 
44
49
  def _managed_router_tokens(router_args: Mapping[str, str]) -> list[str]:
45
50
  """Rendered override tokens that would reach a managed flag.
@@ -177,6 +182,7 @@ class PDEndpoint(Endpoint):
177
182
  health_failure_timeout: float = 30,
178
183
  router_args: Mapping[str, str] | None = None,
179
184
  allow_custom_router: bool = False,
185
+ shutdown_port: int | None = None,
180
186
  ):
181
187
  from modal.experimental import get_cluster_info
182
188
 
@@ -217,6 +223,13 @@ class PDEndpoint(Endpoint):
217
223
  f"P/D ratio {ratio} requires a {len(roles)}-container Modal cluster"
218
224
  )
219
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")
220
233
  self.role = roles[cluster.rank]
221
234
  super().__init__(_url(hosts[0], router_port))
222
235
  self.engine = engine
@@ -266,6 +279,8 @@ class PDEndpoint(Endpoint):
266
279
  self._stopped = threading.Event()
267
280
  self._lock = threading.Lock()
268
281
  self._started = False
282
+ self._serving = False
283
+ self._shutdown_listener: socket.socket | None = None
269
284
  self._threads: list[threading.Thread] = []
270
285
 
271
286
  def start(self, *, warmup: Callable[[], None] | None = None) -> None:
@@ -280,6 +295,13 @@ class PDEndpoint(Endpoint):
280
295
  SGLANG_ENABLE_HEALTH_ENDPOINT_GENERATION="1",
281
296
  )
282
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)
283
305
  for endpoint in (self.engine, self.router):
284
306
  if endpoint is not None:
285
307
  endpoint.start()
@@ -291,6 +313,7 @@ class PDEndpoint(Endpoint):
291
313
  for _, host in self.router.pd_config:
292
314
  self._watch_health(_url(host, self.engine.worker_port) + "/health")
293
315
  self._watch_health(self.router.base_url + "/health")
316
+ self._serving = True
294
317
  except BaseException:
295
318
  self.stop()
296
319
  raise
@@ -352,7 +375,7 @@ class PDEndpoint(Endpoint):
352
375
  if self._stopped.is_set():
353
376
  return
354
377
  self._stopped.set()
355
- deadline = time.monotonic() + self.drain_timeout + 30
378
+ deadline = time.monotonic() + _SHUTDOWN_TIMEOUT
356
379
  try:
357
380
  for endpoint in (self.router, self.engine):
358
381
  if endpoint is None:
@@ -364,9 +387,58 @@ class PDEndpoint(Endpoint):
364
387
  process.wait(timeout=max(0, deadline - time.monotonic()))
365
388
  except subprocess.TimeoutExpired:
366
389
  process.kill()
390
+ process.wait()
367
391
  endpoint.stop()
368
392
  finally:
369
393
  for thread in self._threads:
370
- thread.join(
371
- 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,
372
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])