[Refactor] Introduce sock_send/sock_recv wrappers for zmq IPC (#29012)
This commit is contained in:
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user