[Refactor] Introduce sock_send/sock_recv wrappers for zmq IPC (#29012)

This commit is contained in:
Lianmin Zheng
2026-06-23 15:54:36 -07:00
committed by GitHub
parent ecab3f322e
commit 34dd9c28ca
22 changed files with 258 additions and 154 deletions
@@ -12,6 +12,7 @@ import zmq
from sglang.srt.entrypoints.http_server import launch_server
from sglang.srt.environ import envs
from sglang.srt.managers.io_struct import sock_recv, sock_send
from sglang.srt.server_args import ServerArgs
from sglang.srt.utils.network import get_free_port, get_zmq_socket_on_host
from sglang.test.scripted_runtime.io_struct import (
@@ -89,7 +90,7 @@ class ScriptedHttpServer:
raise RuntimeError(f"ScriptedHttpServer is dirty: {self._dirty}")
fn_path = f"{script_fn.__module__}:{script_fn.__qualname__}"
self._socket.send_pyobj(RunScript(fn_path=fn_path, args=args))
sock_send(self._socket, RunScript(fn_path=fn_path, args=args))
if not self._socket.poll(int(timeout_s * 1000)):
if not self._server_process.is_alive():
@@ -98,7 +99,7 @@ class ScriptedHttpServer:
self._dirty = f"script {fn_path!r} timed out after {timeout_s}s"
raise TimeoutError(self._dirty)
reply = self._socket.recv_pyobj()
reply = sock_recv(self._socket)
match reply:
case ScriptFailed(traceback=tb):
raise AssertionError(f"scripted-runtime script failed:\n{tb}")
@@ -116,7 +117,7 @@ class ScriptedHttpServer:
fatal_error: Optional[OutOfBandError] = None
try:
try:
self._socket.send_pyobj(Shutdown())
sock_send(self._socket, Shutdown())
except zmq.ZMQError:
pass
@@ -139,7 +140,7 @@ class ScriptedHttpServer:
f"{LISTENER_ACCEPT_TIMEOUT_S}s"
)
ready = self._socket.recv_pyobj()
ready = sock_recv(self._socket)
if not isinstance(ready, HookReady):
raise RuntimeError(
f"ScriptedHttpServer: expected HookReady handshake, got {ready!r}"
@@ -11,6 +11,7 @@ from typing import TYPE_CHECKING, Generator, List, Optional, Tuple
import zmq
from sglang.srt.environ import envs
from sglang.srt.managers.io_struct import sock_recv, sock_send
from sglang.srt.utils.network import get_zmq_socket
from sglang.test.scripted_runtime.background_http_poster import BackgroundHttpPoster
from sglang.test.scripted_runtime.context import ScriptedContext
@@ -152,9 +153,9 @@ class ScriptedSchedulerHook:
socket = get_zmq_socket(ctx_zmq, zmq.PAIR, endpoint, bind=False)
try:
yield from _drive_engine_through_warmup(self._context)
socket.send_pyobj(HookReady())
sock_send(socket, HookReady())
while True:
msg = socket.recv_pyobj()
msg = sock_recv(socket)
match msg:
case Shutdown():
return
@@ -167,11 +168,12 @@ class ScriptedSchedulerHook:
try:
yield from sub_gen
except Exception:
socket.send_pyobj(
ScriptFailed(traceback=traceback.format_exc())
sock_send(
socket,
ScriptFailed(traceback=traceback.format_exc()),
)
else:
socket.send_pyobj(ScriptSucceeded())
sock_send(socket, ScriptSucceeded())
case _:
raise ValueError(f"dispatch loop: unknown command {msg!r}")
finally:
@@ -11,6 +11,7 @@ from sglang.srt.managers.io_struct import (
BatchTokenizedGenerateReqInput,
TokenizedEmbeddingReqInput,
TokenizedGenerateReqInput,
sock_recv,
)
_WORK_REQ_TYPES = (
@@ -64,7 +65,7 @@ class ScriptedTokenizerRecvProxy:
def _drain_underlying(self) -> None:
while True:
try:
req = self._underlying.recv_pyobj(zmq.NOBLOCK)
req = sock_recv(self._underlying, zmq.NOBLOCK)
except zmq.ZMQError:
break
if isinstance(req, _WORK_REQ_TYPES):