[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
+9 -7
View File
@@ -20,6 +20,9 @@ from typing import Any, List, Literal, Optional, Union, get_args, get_type_hints
import torch import torch
import torch.distributed as dist import torch.distributed as dist
import zmq
from sglang.srt.managers.io_struct import sock_recv, sock_send
# -------------------------------------- config base ------------------------------------------ # -------------------------------------- config base ------------------------------------------
@@ -1419,8 +1422,6 @@ def _create_zmq_rpc_broadcast(
handler, timeout_seconds: int = 60 handler, timeout_seconds: int = 60
) -> Optional["_ZmqRpcBroadcast"]: ) -> Optional["_ZmqRpcBroadcast"]:
"""A general-purpose minimal RPC to support broadcasting executions to multi processes""" """A general-purpose minimal RPC to support broadcasting executions to multi processes"""
import zmq
rank = _get_rank() rank = _get_rank()
world_size = dist.get_world_size() if dist.is_initialized() else 1 world_size = dist.get_world_size() if dist.is_initialized() else 1
@@ -1433,13 +1434,13 @@ def _create_zmq_rpc_broadcast(
def serve_loop(): def serve_loop():
while True: while True:
try: try:
req = sock.recv_pyobj() req = sock_recv(sock)
result = getattr(handler, req["method"])(*req["args"], **req["kwargs"]) result = getattr(handler, req["method"])(*req["args"], **req["kwargs"])
resp = {"result": result, "error": None} resp = {"result": result, "error": None}
except Exception as e: except Exception as e:
_log(f"[ZmqRpc] error inside handler: {e}") _log(f"[ZmqRpc] error inside handler: {e}")
resp = {"result": None, "error": str(e)} resp = {"result": None, "error": str(e)}
sock.send_pyobj(resp) sock_send(sock, resp)
thread = threading.Thread(target=serve_loop, daemon=True) thread = threading.Thread(target=serve_loop, daemon=True)
thread.start() thread.start()
@@ -1476,14 +1477,15 @@ class _ZmqRpcHandle:
def __getattr__(self, method_name: str): def __getattr__(self, method_name: str):
def call(*args, **kwargs): def call(*args, **kwargs):
self._socket.send_pyobj( sock_send(
self._socket,
{ {
"method": method_name, "method": method_name,
"args": args, "args": args,
"kwargs": kwargs, "kwargs": kwargs,
} },
) )
response = self._socket.recv_pyobj() response = sock_recv(self._socket)
if response["error"]: if response["error"]:
raise RuntimeError( raise RuntimeError(
f"RPC error on {self._debug_name}: {response['error']}" f"RPC error on {self._debug_name}: {response['error']}"
@@ -26,6 +26,7 @@ from sglang.srt.disaggregation.encode_server import (
handle_scheduler_receive_url_request, handle_scheduler_receive_url_request,
launch_encoder, launch_encoder,
) )
from sglang.srt.managers.io_struct import async_sock_send
from sglang.srt.managers.schedule_batch import Modality from sglang.srt.managers.schedule_batch import Modality
from sglang.srt.server_args import PortArgs, ServerArgs from sglang.srt.server_args import PortArgs, ServerArgs
from sglang.srt.utils import random_uuid from sglang.srt.utils import random_uuid
@@ -95,7 +96,7 @@ class SGLangEncoderServer(SGLangEncoderServicer):
"part_idx": request.part_idx, "part_idx": request.part_idx,
} }
for socket in self.send_sockets: for socket in self.send_sockets:
await socket.send_pyobj(request_dict) await async_sock_send(socket, request_dict)
# gRPC encode is image-only; encoder.encode() requires modality # gRPC encode is image-only; encoder.encode() requires modality
( (
@@ -44,7 +44,14 @@ from sglang.srt.distributed.parallel_state import (
) )
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.layers.dp_attention import initialize_dp_attention from sglang.srt.layers.dp_attention import initialize_dp_attention
from sglang.srt.managers.io_struct import ProfileReq, ProfileReqInput, ProfileReqType from sglang.srt.managers.io_struct import (
ProfileReq,
ProfileReqInput,
ProfileReqType,
async_sock_recv,
async_sock_send,
sock_send,
)
from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem
from sglang.srt.mem_cache.multimodal_cache import EmbeddingResult, MultiModalStaticCache from sglang.srt.mem_cache.multimodal_cache import EmbeddingResult, MultiModalStaticCache
from sglang.srt.model_loader import get_model from sglang.srt.model_loader import get_model
@@ -2390,13 +2397,14 @@ class EncoderScheduler:
requests = [p.request for p in group] requests = [p.request for p in group]
start = time.time() start = time.time()
for sock in self.send_sockets: for sock in self.send_sockets:
sock.send_pyobj( sock_send(
sock,
{ {
"type": "batch_encode", "type": "batch_encode",
"modality": modality.name, "modality": modality.name,
"requests": requests, "requests": requests,
"enter_time": start, "enter_time": start,
} },
) )
logger.info(f"Dispatching batch of {len(group)} {modality.name} requests") logger.info(f"Dispatching batch of {len(group)} {modality.name} requests")
@@ -2442,7 +2450,7 @@ class EncoderScheduler:
req = p.request req = p.request
try: try:
for sock in self.send_sockets: for sock in self.send_sockets:
sock.send_pyobj(req) sock_send(sock, req)
result = await self.encoder.encode_request(req, modality) result = await self.encoder.encode_request(req, modality)
if not p.future.done(): if not p.future.done():
p.future.set_result(result) p.future.set_result(result)
@@ -2731,7 +2739,7 @@ class DPDispatcher:
) )
try: try:
await self.dispatch_sockets[rank].send_pyobj(request) await async_sock_send(self.dispatch_sockets[rank], request)
# An alive-but-stuck worker (NCCL deadlock etc.) wouldn't trip # An alive-but-stuck worker (NCCL deadlock etc.) wouldn't trip
# the watchdog, so bound the wait explicitly. # the watchdog, so bound the wait explicitly.
return await asyncio.wait_for(future, timeout=ENCODER_REQ_TIMEOUT) return await asyncio.wait_for(future, timeout=ENCODER_REQ_TIMEOUT)
@@ -2780,7 +2788,7 @@ class DPDispatcher:
f"dp_rank={rank}, pending={self.pending_counts}" f"dp_rank={rank}, pending={self.pending_counts}"
) )
try: try:
await self.dispatch_sockets[rank].send_pyobj(request) await async_sock_send(self.dispatch_sockets[rank], request)
return await asyncio.wait_for(future, timeout=ENCODER_REQ_TIMEOUT) return await asyncio.wait_for(future, timeout=ENCODER_REQ_TIMEOUT)
except asyncio.TimeoutError: except asyncio.TimeoutError:
self.pending_futures[rank].pop(key, None) self.pending_futures[rank].pop(key, None)
@@ -2821,7 +2829,7 @@ class DPDispatcher:
self.req_id_to_rank[req_id] = rank self.req_id_to_rank[req_id] = rank
rank_keys.append((rank, req_id)) rank_keys.append((rank, req_id))
request_copy = {**request, "req_id": req_id} request_copy = {**request, "req_id": req_id}
await self.dispatch_sockets[rank].send_pyobj(request_copy) await async_sock_send(self.dispatch_sockets[rank], request_copy)
futures.append(future) futures.append(future)
# Concurrent wait → total bounded by eff_timeout, not # Concurrent wait → total bounded by eff_timeout, not
# dp_size × eff_timeout. # dp_size × eff_timeout.
@@ -2901,7 +2909,7 @@ class DPDispatcher:
consecutive_errors = 0 consecutive_errors = 0
while True: while True:
try: try:
msg = await self.result_socket.recv_pyobj() msg = await async_sock_recv(self.result_socket)
consecutive_errors = 0 consecutive_errors = 0
except asyncio.CancelledError: except asyncio.CancelledError:
raise raise
@@ -3051,10 +3059,10 @@ async def _dp_worker_handle_request(
"_error_code": err_code, "_error_code": err_code,
} }
# pyzmq async send_pyobj isn't safe for concurrent senders. # pyzmq async send isn't safe for concurrent senders.
try: try:
async with send_lock: async with send_lock:
await send_sock.send_pyobj(envelope) await async_sock_send(send_sock, envelope)
except Exception: except Exception:
logger.error( logger.error(
f"DP worker {dp_rank} failed to send envelope for " f"DP worker {dp_rank} failed to send envelope for "
@@ -3113,7 +3121,7 @@ async def run_dp_worker(
spawned = False spawned = False
try: try:
try: try:
request = await recv_sock.recv_pyobj() request = await async_sock_recv(recv_sock)
except asyncio.CancelledError: except asyncio.CancelledError:
raise raise
except Exception: except Exception:
@@ -3197,7 +3205,7 @@ async def run_encoder(
): ):
encoder = MMEncoder(server_args, schedule_path, dist_init_method, rank) encoder = MMEncoder(server_args, schedule_path, dist_init_method, rank)
while True: while True:
request = await encoder.schedule_socket.recv_pyobj() request = await async_sock_recv(encoder.schedule_socket)
if isinstance(request, ProfileReq): if isinstance(request, ProfileReq):
if request.type == ProfileReqType.START_PROFILE: if request.type == ProfileReqType.START_PROFILE:
if encoder.profiler is None: if encoder.profiler is None:
@@ -3601,7 +3609,7 @@ async def handle_encode_request(request: dict):
) )
else: else:
for socket in send_sockets: for socket in send_sockets:
socket.send_pyobj(request) sock_send(socket, request)
nbytes, embedding_len, embedding_dim, error_msg, error_code = ( nbytes, embedding_len, embedding_dim, error_msg, error_code = (
await encoder.encode_request(request, modality) await encoder.encode_request(request, modality)
) )
@@ -3822,7 +3830,7 @@ async def health_generate():
# Broadcast to other TP ranks so distributed ops stay in sync # Broadcast to other TP ranks so distributed ops stay in sync
for socket in send_sockets: for socket in send_sockets:
socket.send_pyobj(dummy_request) sock_send(socket, dummy_request)
# Run encode on rank 0 with timeout # Run encode on rank 0 with timeout
_, _, _, error_msg, _ = await asyncio.wait_for( _, _, _, error_msg, _ = await asyncio.wait_for(
@@ -3900,7 +3908,7 @@ async def start_profile_async(obj: Optional[ProfileReqInput] = None):
profile_stages=obj.profile_stages, profile_stages=obj.profile_stages,
) )
for socket in send_sockets: for socket in send_sockets:
socket.send_pyobj(req) sock_send(socket, req)
if encoder.profiler is None: if encoder.profiler is None:
encoder.profiler = EncoderProfiler(encoder.rank) encoder.profiler = EncoderProfiler(encoder.rank)
ok, msg = encoder.profiler.start(req) ok, msg = encoder.profiler.start(req)
@@ -3931,7 +3939,7 @@ async def stop_profile_async():
) )
req = ProfileReq(ProfileReqType.STOP_PROFILE) req = ProfileReq(ProfileReqType.STOP_PROFILE)
for socket in send_sockets: for socket in send_sockets:
socket.send_pyobj(req) sock_send(socket, req)
ok, msg = encoder.profiler.stop() ok, msg = encoder.profiler.stop()
if ok: if ok:
return Response(content="Stop profiling.\n", status_code=200) return Response(content="Stop profiling.\n", status_code=200)
@@ -12,7 +12,7 @@ from sglang.srt.distributed.parallel_state import (
) )
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.eplb.expert_location import get_global_expert_location_metadata from sglang.srt.eplb.expert_location import get_global_expert_location_metadata
from sglang.srt.managers.io_struct import UpdateExpertBackupReq from sglang.srt.managers.io_struct import UpdateExpertBackupReq, sock_recv, sock_send
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.utils.network import get_local_ip_auto from sglang.srt.utils.network import get_local_ip_auto
@@ -65,7 +65,7 @@ class ExpertBackupClient:
self.ready_sockets[i].connect( self.ready_sockets[i].connect(
f"tcp://{all_ips[i * get_world_size() // server_args.nnodes]}:{PORT_BASE + i * 2}" f"tcp://{all_ips[i * get_world_size() // server_args.nnodes]}:{PORT_BASE + i * 2}"
) )
self.ready_sockets[i].send_pyobj(UpdateExpertBackupReq()) sock_send(self.ready_sockets[i], UpdateExpertBackupReq())
self._receive_thread = threading.Thread(target=self._receive_loop, daemon=True) self._receive_thread = threading.Thread(target=self._receive_loop, daemon=True)
self._receive_thread.start() self._receive_thread.start()
@@ -73,7 +73,7 @@ class ExpertBackupClient:
def _receive_loop(self): def _receive_loop(self):
cnt = 0 cnt = 0
while cnt < self.engine_num: while cnt < self.engine_num:
response = self.recv_list[cnt].recv_pyobj() response = sock_recv(self.recv_list[cnt])
self.dram_map_list[response.rank] = response.weight_pointer_map self.dram_map_list[response.rank] = response.weight_pointer_map
self.session_id_list[response.rank] = response.session_id self.session_id_list[response.rank] = response.session_id
self.buffer_size = max(self.buffer_size, response.buffer_size) self.buffer_size = max(self.buffer_size, response.buffer_size)
@@ -9,7 +9,7 @@ import zmq
from sglang.srt.configs.load_config import LoadConfig from sglang.srt.configs.load_config import LoadConfig
from sglang.srt.configs.model_config import ModelConfig from sglang.srt.configs.model_config import ModelConfig
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.managers.io_struct import BackupDramReq from sglang.srt.managers.io_struct import BackupDramReq, sock_recv, sock_send
from sglang.srt.model_loader.loader import DefaultModelLoader, get_model_loader from sglang.srt.model_loader.loader import DefaultModelLoader, get_model_loader
from sglang.srt.model_loader.utils import set_default_torch_dtype from sglang.srt.model_loader.utils import set_default_torch_dtype
from sglang.srt.server_args import ( from sglang.srt.server_args import (
@@ -62,7 +62,7 @@ class ExpertBackupManager:
num_ready_clients = 0 num_ready_clients = 0
while num_ready_clients < server_args.tp_size: while num_ready_clients < server_args.tp_size:
self.recv_from_expert_backup_client.recv_pyobj() sock_recv(self.recv_from_expert_backup_client)
num_ready_clients += 1 num_ready_clients += 1
back_req = BackupDramReq( back_req = BackupDramReq(
@@ -72,7 +72,7 @@ class ExpertBackupManager:
buffer_size=self.continuous_buffer.numel() buffer_size=self.continuous_buffer.numel()
* self.continuous_buffer.element_size(), * self.continuous_buffer.element_size(),
) )
self.send_to_expert_backup_client.send_pyobj(back_req) sock_send(self.send_to_expert_backup_client, back_req)
# Keep the manager subprocess alive until signals # Keep the manager subprocess alive until signals
signal.pause() signal.pause()
+4 -2
View File
@@ -77,6 +77,8 @@ from sglang.srt.managers.io_struct import (
UpdateWeightsFromDistributedReqInput, UpdateWeightsFromDistributedReqInput,
UpdateWeightsFromIPCReqInput, UpdateWeightsFromIPCReqInput,
UpdateWeightsFromTensorReqInput, UpdateWeightsFromTensorReqInput,
sock_recv,
sock_send,
) )
from sglang.srt.managers.multi_tokenizer_mixin import ( from sglang.srt.managers.multi_tokenizer_mixin import (
MultiTokenizerRouter, MultiTokenizerRouter,
@@ -1220,8 +1222,8 @@ class Engine(EngineScoreMixin, EngineBase):
def collective_rpc(self, method: str, **kwargs): def collective_rpc(self, method: str, **kwargs):
obj = RpcReqInput(method=method, parameters=kwargs) obj = RpcReqInput(method=method, parameters=kwargs)
self.send_to_rpc.send_pyobj(obj) sock_send(self.send_to_rpc, obj)
recv_req = self.send_to_rpc.recv_pyobj(zmq.BLOCKY) recv_req = sock_recv(self.send_to_rpc, flags=zmq.BLOCKY)
assert isinstance(recv_req, RpcReqOutput) assert isinstance(recv_req, RpcReqOutput)
assert recv_req.success, recv_req.message assert recv_req.success, recv_req.message
+4 -4
View File
@@ -5,7 +5,7 @@ import copy
from collections import deque from collections import deque
from typing import Deque, Generic, List, Optional, TypeVar from typing import Deque, Generic, List, Optional, TypeVar
import zmq from sglang.srt.managers.io_struct import sock_send
T = TypeVar("T") T = TypeVar("T")
@@ -22,7 +22,7 @@ class FanOutCommunicator(Generic[T]):
Only one request is in-flight at any time in either mode. Only one request is in-flight at any time in either mode.
""" """
def __init__(self, sender: zmq.Socket, fan_out: int, mode="queueing"): def __init__(self, sender, fan_out: int, mode="queueing"):
self._sender = sender self._sender = sender
self._fan_out = fan_out self._fan_out = fan_out
self._mode = mode self._mode = mode
@@ -41,7 +41,7 @@ class FanOutCommunicator(Generic[T]):
assert self._result_values is None assert self._result_values is None
if obj is not None: if obj is not None:
self._sender.send_pyobj(obj) sock_send(self._sender, obj)
self._result_event = asyncio.Event() self._result_event = asyncio.Event()
self._result_values = [] self._result_values = []
@@ -61,7 +61,7 @@ class FanOutCommunicator(Generic[T]):
self._result_event = asyncio.Event() self._result_event = asyncio.Event()
if obj is not None: if obj is not None:
self._sender.send_pyobj(obj) sock_send(self._sender, obj)
# Capture local refs before await -- after event fires, the first # Capture local refs before await -- after event fires, the first
# awakened coroutine clears shared state; later awaiters use local refs. # awakened coroutine clears shared state; later awaiters use local refs.
@@ -36,6 +36,8 @@ from sglang.srt.managers.io_struct import (
ProfileReq, ProfileReq,
TokenizedEmbeddingReqInput, TokenizedEmbeddingReqInput,
TokenizedGenerateReqInput, TokenizedGenerateReqInput,
sock_recv,
sock_send,
) )
from sglang.srt.managers.load_snapshot import create_load_snapshot_reader from sglang.srt.managers.load_snapshot import create_load_snapshot_reader
from sglang.srt.managers.schedule_batch import Req from sglang.srt.managers.schedule_batch import Req
@@ -205,12 +207,12 @@ class DataParallelController:
def send_to_all_workers(self, obj): def send_to_all_workers(self, obj):
for i, worker in enumerate(self.workers): for i, worker in enumerate(self.workers):
if self.status[i]: if self.status[i]:
worker.send_pyobj(obj) sock_send(worker, obj)
def send_control_message(self, obj): def send_control_message(self, obj):
# Send control messages to first worker of tp group # Send control messages to first worker of tp group
for worker in self.workers[:: self.control_message_step]: for worker in self.workers[:: self.control_message_step]:
worker.send_pyobj(obj) sock_send(worker, obj)
def update_active_ranks(self, ranks: ActiveRanksOutput): def update_active_ranks(self, ranks: ActiveRanksOutput):
self.status = ranks.status self.status = ranks.status
@@ -384,7 +386,7 @@ class DataParallelController:
logger.debug(f"Received handshake from node {client_rank}") logger.debug(f"Received handshake from node {client_rank}")
# Send worker ports to client # Send worker ports to client
rep_socket.send_pyobj(worker_ports) sock_send(rep_socket, worker_ports)
connected_clients += 1 connected_clients += 1
logger.debug( logger.debug(
f"Sent worker ports to {connected_clients}/{expected_clients} nodes" f"Sent worker ports to {connected_clients}/{expected_clients} nodes"
@@ -418,7 +420,7 @@ class DataParallelController:
logger.debug(f"Received handshake from node {client_rank}") logger.debug(f"Received handshake from node {client_rank}")
# Send worker ports to client # Send worker ports to client
rep_socket.send_pyobj(worker_ports) sock_send(rep_socket, worker_ports)
logger.debug(f"Sent worker ports to node {client_rank}") logger.debug(f"Sent worker ports to node {client_rank}")
def _receive_ports_as_client(self, endpoint: str, node_rank: int) -> List[int]: def _receive_ports_as_client(self, endpoint: str, node_rank: int) -> List[int]:
@@ -434,7 +436,7 @@ class DataParallelController:
req_socket.send(str(node_rank).encode()) req_socket.send(str(node_rank).encode())
# Receive worker ports # Receive worker ports
worker_ports = req_socket.recv_pyobj() worker_ports = sock_recv(req_socket)
logger.debug(f"Received {len(worker_ports)} worker ports from node 0") logger.debug(f"Received {len(worker_ports)} worker ports from node 0")
return worker_ports return worker_ports
except zmq.Again: except zmq.Again:
@@ -597,7 +599,7 @@ class DataParallelController:
def maybe_external_dp_rank_routing(self, req: Req): def maybe_external_dp_rank_routing(self, req: Req):
if req.routed_dp_rank is not None: if req.routed_dp_rank is not None:
logger.debug(f"Direct routing to DP rank {req.routed_dp_rank}") logger.debug(f"Direct routing to DP rank {req.routed_dp_rank}")
self.workers[req.routed_dp_rank].send_pyobj(req) sock_send(self.workers[req.routed_dp_rank], req)
return True return True
return False return False
@@ -608,7 +610,7 @@ class DataParallelController:
while True: while True:
if self.status[self.round_robin_counter]: if self.status[self.round_robin_counter]:
logger.debug(f"Choose worker {self.round_robin_counter}") logger.debug(f"Choose worker {self.round_robin_counter}")
self.workers[self.round_robin_counter].send_pyobj(req) sock_send(self.workers[self.round_robin_counter], req)
self.round_robin_counter = (self.round_robin_counter + 1) % len( self.round_robin_counter = (self.round_robin_counter + 1) % len(
self.workers self.workers
) )
@@ -626,13 +628,13 @@ class DataParallelController:
"prefill or decode instances; send to the router instead." "prefill or decode instances; send to the router instead."
) )
target_rank = req.bootstrap_room % len(self.workers) target_rank = req.bootstrap_room % len(self.workers)
self.workers[target_rank].send_pyobj(req) sock_send(self.workers[target_rank], req)
def total_requests_scheduler(self, req: Req): def total_requests_scheduler(self, req: Req):
if self.maybe_external_dp_rank_routing(req): if self.maybe_external_dp_rank_routing(req):
return return
target_worker = self.dp_budget.dispatch(LoadBalanceMethod.TOTAL_REQUESTS) target_worker = self.dp_budget.dispatch(LoadBalanceMethod.TOTAL_REQUESTS)
self.workers[target_worker].send_pyobj(req) sock_send(self.workers[target_worker], req)
def total_tokens_scheduler(self, req: Req): def total_tokens_scheduler(self, req: Req):
if self.maybe_external_dp_rank_routing(req): if self.maybe_external_dp_rank_routing(req):
@@ -641,14 +643,14 @@ class DataParallelController:
target_worker = self.dp_budget.dispatch( target_worker = self.dp_budget.dispatch(
LoadBalanceMethod.TOTAL_TOKENS, estimated_tokens=estimated_tokens LoadBalanceMethod.TOTAL_TOKENS, estimated_tokens=estimated_tokens
) )
self.workers[target_worker].send_pyobj(req) sock_send(self.workers[target_worker], req)
def event_loop(self): def event_loop(self):
while True: while True:
while True: while True:
self.soft_watchdog.feed() self.soft_watchdog.feed()
try: try:
recv_req = self.recv_from_tokenizer.recv_pyobj(zmq.NOBLOCK) recv_req = sock_recv(self.recv_from_tokenizer, flags=zmq.NOBLOCK)
except zmq.ZMQError: except zmq.ZMQError:
break break
self._request_dispatcher(recv_req) self._request_dispatcher(recv_req)
@@ -34,6 +34,8 @@ from sglang.srt.managers.io_struct import (
BatchTokenIDOutput, BatchTokenIDOutput,
ConfigureLoggingReq, ConfigureLoggingReq,
FreezeGCReq, FreezeGCReq,
sock_recv,
sock_send,
) )
from sglang.srt.managers.multi_tokenizer_mixin import MultiHttpWorkerDetokenizerMixin from sglang.srt.managers.multi_tokenizer_mixin import MultiHttpWorkerDetokenizerMixin
from sglang.srt.observability.cpu_monitor import start_cpu_monitor_thread from sglang.srt.observability.cpu_monitor import start_cpu_monitor_thread
@@ -160,10 +162,10 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
"""The event loop that handles requests""" """The event loop that handles requests"""
while True: while True:
with self.soft_watchdog.disable(): with self.soft_watchdog.disable():
recv_obj = self.recv_from_scheduler.recv_pyobj() recv_obj = sock_recv(self.recv_from_scheduler)
output = self._request_dispatcher(recv_obj) output = self._request_dispatcher(recv_obj)
if output is not None: if output is not None:
self.send_to_tokenizer.send_pyobj(output) sock_send(self.send_to_tokenizer, output)
self.soft_watchdog.feed() self.soft_watchdog.feed()
def trim_matched_stop( def trim_matched_stop(
+38 -1
View File
@@ -25,9 +25,20 @@ from array import array
from collections import Counter from collections import Counter
from dataclasses import dataclass, field from dataclasses import dataclass, field
from enum import Enum from enum import Enum
from typing import TYPE_CHECKING, Annotated, Any, Dict, List, Literal, Optional, Union from typing import (
TYPE_CHECKING,
Annotated,
Any,
Dict,
List,
Literal,
Optional,
Union,
)
import torch import torch
import zmq
import zmq.asyncio
from pydantic import PlainValidator from pydantic import PlainValidator
from sglang.srt.lora.lora_registry import LoRARef from sglang.srt.lora.lora_registry import LoRARef
@@ -46,6 +57,8 @@ from sglang.srt.utils.field_validators import validate_optional_list_i64_1d_2d
# Handle serialization of Image for pydantic # Handle serialization of Image for pydantic
if TYPE_CHECKING: if TYPE_CHECKING:
from PIL.Image import Image from PIL.Image import Image
from sglang.srt.managers.tokenizer_manager import SenderWrapper
else: else:
Image = Any Image = Any
@@ -2191,6 +2204,30 @@ class DumperControlReqOutput(BaseReq):
error: str = "" error: str = ""
def sock_send(
sender: Union[zmq.Socket, zmq.asyncio.Socket, SenderWrapper],
obj: Any,
flags: int = 0,
) -> None:
sender.send_pyobj(obj, flags=flags)
def sock_recv(socket, flags=0):
return socket.recv_pyobj(flags=flags)
async def async_sock_send(
sender: Union[zmq.asyncio.Socket, SenderWrapper],
obj: Any,
flags: int = 0,
) -> None:
await sender.send_pyobj(obj, flags=flags)
async def async_sock_recv(socket, flags=0):
return await socket.recv_pyobj(flags=flags)
def _check_all_req_types(): def _check_all_req_types():
"""A helper function to check all request types are defined in this file.""" """A helper function to check all request types are defined in this file."""
import inspect import inspect
@@ -49,6 +49,10 @@ from sglang.srt.managers.io_struct import (
PauseContinueBroadcast, PauseContinueBroadcast,
PauseGenerationReqInput, PauseGenerationReqInput,
TokenizerWorkerRegistration, TokenizerWorkerRegistration,
async_sock_recv,
async_sock_send,
sock_recv,
sock_send,
) )
from sglang.srt.managers.load_snapshot import ( from sglang.srt.managers.load_snapshot import (
create_load_snapshot_reader, create_load_snapshot_reader,
@@ -97,7 +101,7 @@ class SocketMapping:
if ipc_name not in self._mapping: if ipc_name not in self._mapping:
self._register_ipc_mapping(ipc_name, is_tokenizer=is_tokenizer) self._register_ipc_mapping(ipc_name, is_tokenizer=is_tokenizer)
self._mapping[ipc_name].send_pyobj(output) sock_send(self._mapping[ipc_name], output)
def _extract_field_by_index( def _extract_field_by_index(
@@ -334,7 +338,7 @@ class MultiHttpWorkerDetokenizerMixin:
"""The event loop that handles requests, for multi multi-http-worker mode""" """The event loop that handles requests, for multi multi-http-worker mode"""
self.socket_mapping = SocketMapping() self.socket_mapping = SocketMapping()
while True: while True:
recv_obj = self.recv_from_scheduler.recv_pyobj() recv_obj = sock_recv(self.recv_from_scheduler)
output = self._request_dispatcher(recv_obj) output = self._request_dispatcher(recv_obj)
if output is None: if output is None:
continue continue
@@ -431,7 +435,7 @@ class MultiTokenizerRouter:
async def router_worker_obj(self): async def router_worker_obj(self):
"""Forward path: workers → scheduler, with pause/continue broadcast.""" """Forward path: workers → scheduler, with pause/continue broadcast."""
while True: while True:
recv_obj = await self.receive_from_worker.recv_pyobj() recv_obj = await async_sock_recv(self.receive_from_worker)
if isinstance(recv_obj, TokenizerWorkerRegistration): if isinstance(recv_obj, TokenizerWorkerRegistration):
if recv_obj.worker_ipc_name not in self.all_worker_ipcs: if recv_obj.worker_ipc_name not in self.all_worker_ipcs:
@@ -456,15 +460,15 @@ class MultiTokenizerRouter:
isinstance(recv_obj, PauseGenerationReqInput) isinstance(recv_obj, PauseGenerationReqInput)
and recv_obj.mode == "abort" and recv_obj.mode == "abort"
): ):
await self.send_to_scheduler.send_pyobj(recv_obj) await async_sock_send(self.send_to_scheduler, recv_obj)
continue continue
await self.send_to_scheduler.send_pyobj(recv_obj) await async_sock_send(self.send_to_scheduler, recv_obj)
async def handle_loop(self): async def handle_loop(self):
"""Backward path: detokenizer → route results to correct worker.""" """Backward path: detokenizer → route results to correct worker."""
while True: while True:
recv_obj = await self.recv_from_detokenizer.recv_pyobj() recv_obj = await async_sock_recv(self.recv_from_detokenizer)
await self._distribute_result_to_workers(recv_obj) await self._distribute_result_to_workers(recv_obj)
async def _distribute_result_to_workers(self, recv_obj): async def _distribute_result_to_workers(self, recv_obj):
@@ -505,7 +509,7 @@ class MultiDetokenizerRouter:
def event_loop(self): def event_loop(self):
while True: while True:
recv_obj = self.recv_from_scheduler.recv_pyobj() recv_obj = sock_recv(self.recv_from_scheduler)
# FreezeGCReq must freeze every detokenizer process. # FreezeGCReq must freeze every detokenizer process.
if isinstance(recv_obj, FreezeGCReq): if isinstance(recv_obj, FreezeGCReq):
@@ -600,7 +604,7 @@ class TokenizerWorker(TokenizerManager):
# Register this worker with the router for pause/continue broadcasting # Register this worker with the router for pause/continue broadcasting
reg = TokenizerWorkerRegistration(worker_ipc_name=self.tokenizer_ipc_name) reg = TokenizerWorkerRegistration(worker_ipc_name=self.tokenizer_ipc_name)
self.send_to_scheduler.send_pyobj(reg) sock_send(self.send_to_scheduler, reg)
# Future for awaiting pause/continue broadcast confirmation # Future for awaiting pause/continue broadcast confirmation
self._pause_continue_future: Optional[asyncio.Future] = None self._pause_continue_future: Optional[asyncio.Future] = None
@@ -618,7 +622,7 @@ class TokenizerWorker(TokenizerManager):
self._pause_continue_future = loop.create_future() self._pause_continue_future = loop.create_future()
# Send to router which will broadcast to all workers # Send to router which will broadcast to all workers
# (router also handles forwarding to scheduler for non-abort modes) # (router also handles forwarding to scheduler for non-abort modes)
self.send_to_scheduler.send_pyobj(obj) sock_send(self.send_to_scheduler, obj)
await self._pause_continue_future await self._pause_continue_future
if obj.mode == "abort": if obj.mode == "abort":
@@ -633,7 +637,7 @@ class TokenizerWorker(TokenizerManager):
async def continue_generation(self, obj: ContinueGenerationReqInput): async def continue_generation(self, obj: ContinueGenerationReqInput):
loop = asyncio.get_event_loop() loop = asyncio.get_event_loop()
self._pause_continue_future = loop.create_future() self._pause_continue_future = loop.create_future()
self.send_to_scheduler.send_pyobj(obj) sock_send(self.send_to_scheduler, obj)
await self._pause_continue_future await self._pause_continue_future
def _handle_pause_continue_broadcast(self, obj: PauseContinueBroadcast): def _handle_pause_continue_broadcast(self, obj: PauseContinueBroadcast):
@@ -734,14 +738,3 @@ def write_data_for_multi_tokenizer(
args_shm.close() args_shm.close()
return args_shm return args_shm
class SenderWrapper:
def __init__(self, port_args: PortArgs, send_to_scheduler: zmq.Socket):
self.port_args = port_args
self.send_to_scheduler = send_to_scheduler
def send_pyobj(self, obj):
if isinstance(obj, BaseReq):
obj.http_worker_ipc = self.port_args.tokenizer_ipc_name
self.send_to_scheduler.send_pyobj(obj)
+2 -1
View File
@@ -147,6 +147,7 @@ from sglang.srt.managers.io_struct import (
UpdateWeightsFromDistributedReqInput, UpdateWeightsFromDistributedReqInput,
UpdateWeightsFromIPCReqInput, UpdateWeightsFromIPCReqInput,
UpdateWeightsFromTensorReqInput, UpdateWeightsFromTensorReqInput,
sock_send,
) )
from sglang.srt.managers.load_snapshot import LoadSnapshot, create_load_snapshot_writer from sglang.srt.managers.load_snapshot import LoadSnapshot, create_load_snapshot_writer
from sglang.srt.managers.multimodal_processor import get_mm_processor, import_processors from sglang.srt.managers.multimodal_processor import get_mm_processor, import_processors
@@ -1645,7 +1646,7 @@ class Scheduler(
self.ipc_channels.send_to_tokenizer.send_output(output, recv_req) self.ipc_channels.send_to_tokenizer.send_output(output, recv_req)
else: else:
if self.ipc_channels.recv_from_rpc is not None: if self.ipc_channels.recv_from_rpc is not None:
self.ipc_channels.recv_from_rpc.send_pyobj(output) sock_send(self.ipc_channels.recv_from_rpc, output)
self.flush_wrapper.check_pending() self.flush_wrapper.check_pending()
if self.external_corpus_manager is not None: if self.external_corpus_manager is not None:
@@ -16,6 +16,7 @@ from sglang.srt.disaggregation.kv_events import (
EventPublisherFactory, EventPublisherFactory,
KVEventBatch, KVEventBatch,
) )
from sglang.srt.managers.io_struct import sock_send
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.distributed.parallel_state_wrapper import ParallelState from sglang.srt.distributed.parallel_state_wrapper import ParallelState
@@ -85,7 +86,7 @@ class SchedulerKvEventsPublisher:
) )
if not self.send_metrics_from_scheduler.closed: if not self.send_metrics_from_scheduler.closed:
self.send_metrics_from_scheduler.send_pyobj(kv_metrics) sock_send(self.send_metrics_from_scheduler, kv_metrics)
def publish_kv_events(self): def publish_kv_events(self):
if not self.enable_kv_cache_events: if not self.enable_kv_cache_events:
@@ -2,7 +2,7 @@ from typing import Optional, Union
import zmq import zmq
from sglang.srt.managers.io_struct import BaseBatchReq, BaseReq from sglang.srt.managers.io_struct import BaseBatchReq, BaseReq, sock_send
class SenderWrapper: class SenderWrapper:
@@ -25,4 +25,4 @@ class SenderWrapper:
# handle communicator reqs for multi-http worker case # handle communicator reqs for multi-http worker case
output.http_worker_ipc = recv_obj.http_worker_ipc output.http_worker_ipc = recv_obj.http_worker_ipc
self.socket.send_pyobj(output) sock_send(self.socket, output)
@@ -20,6 +20,7 @@ from sglang.srt.managers.io_struct import (
BatchTokenizedGenerateReqInput, BatchTokenizedGenerateReqInput,
TokenizedEmbeddingReqInput, TokenizedEmbeddingReqInput,
TokenizedGenerateReqInput, TokenizedGenerateReqInput,
sock_recv,
) )
from sglang.srt.managers.mm_utils import ( from sglang.srt.managers.mm_utils import (
has_shm_features, has_shm_features,
@@ -103,7 +104,7 @@ class SchedulerRequestReceiver:
try: try:
if self.recv_limit_reached(len(recv_reqs)): if self.recv_limit_reached(len(recv_reqs)):
break break
recv_req = self.recv_from_tokenizer.recv_pyobj(zmq.NOBLOCK) recv_req = sock_recv(self.recv_from_tokenizer, zmq.NOBLOCK)
except zmq.ZMQError: except zmq.ZMQError:
break break
recv_reqs.append(recv_req) recv_reqs.append(recv_req)
@@ -112,7 +113,7 @@ class SchedulerRequestReceiver:
try: try:
if self.recv_limit_reached(len(recv_reqs)): if self.recv_limit_reached(len(recv_reqs)):
break break
recv_rpc = self.recv_from_rpc.recv_pyobj(zmq.NOBLOCK) recv_rpc = sock_recv(self.recv_from_rpc, zmq.NOBLOCK)
except zmq.ZMQError: except zmq.ZMQError:
break break
recv_reqs.append(recv_rpc) recv_reqs.append(recv_rpc)
@@ -16,7 +16,7 @@ from contextlib import contextmanager
from enum import Enum, auto from enum import Enum, auto
from typing import Any, List, Optional from typing import Any, List, Optional
from sglang.srt.managers.io_struct import BlockReqInput, BlockReqType from sglang.srt.managers.io_struct import BlockReqInput, BlockReqType, sock_send
from sglang.srt.utils.poll_based_barrier import PollBasedBarrier from sglang.srt.utils.poll_based_barrier import PollBasedBarrier
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -99,8 +99,8 @@ class _State(Enum):
@contextmanager @contextmanager
def input_blocker_guard_region(send_to_scheduler): def input_blocker_guard_region(send_to_scheduler):
send_to_scheduler.send_pyobj(BlockReqInput(BlockReqType.BLOCK)) sock_send(send_to_scheduler, BlockReqInput(BlockReqType.BLOCK))
try: try:
yield yield
finally: finally:
send_to_scheduler.send_pyobj(BlockReqInput(BlockReqType.UNBLOCK)) sock_send(send_to_scheduler, BlockReqInput(BlockReqType.UNBLOCK))
@@ -71,6 +71,8 @@ from sglang.srt.managers.io_struct import (
UpdateWeightsFromIPCReqOutput, UpdateWeightsFromIPCReqOutput,
UpdateWeightsFromTensorReqInput, UpdateWeightsFromTensorReqInput,
UpdateWeightsFromTensorReqOutput, UpdateWeightsFromTensorReqOutput,
async_sock_send,
sock_send,
) )
from sglang.srt.managers.load_snapshot import LoadSnapshot from sglang.srt.managers.load_snapshot import LoadSnapshot
from sglang.srt.server_args import LoRARef, ServerArgs from sglang.srt.server_args import LoRARef, ServerArgs
@@ -857,7 +859,7 @@ class TokenizerControlMixin:
future = asyncio.Future() future = asyncio.Future()
self.session_futures[obj.session_id] = future self.session_futures[obj.session_id] = future
self.send_to_scheduler.send_pyobj(obj) sock_send(self.send_to_scheduler, obj)
try: try:
return await future return await future
@@ -869,7 +871,7 @@ class TokenizerControlMixin:
obj: CloseSessionReqInput, obj: CloseSessionReqInput,
request: Optional[fastapi.Request] = None, request: Optional[fastapi.Request] = None,
): ):
await self.send_to_scheduler.send_pyobj(obj) await async_sock_send(self.send_to_scheduler, obj)
def _update_weight_version_if_provided( def _update_weight_version_if_provided(
self: TokenizerManager, weight_version: Optional[str] self: TokenizerManager, weight_version: Optional[str]
+42 -16
View File
@@ -55,6 +55,7 @@ from sglang.srt.managers.embed_types import PositionalEmbeds
from sglang.srt.managers.io_struct import ( from sglang.srt.managers.io_struct import (
AbortReq, AbortReq,
ActiveRanksOutput, ActiveRanksOutput,
BaseReq,
BatchEmbeddingOutput, BatchEmbeddingOutput,
BatchStrOutput, BatchStrOutput,
BatchTokenIDOutput, BatchTokenIDOutput,
@@ -75,6 +76,9 @@ from sglang.srt.managers.io_struct import (
TokenizedGenerateReqInput, TokenizedGenerateReqInput,
UpdateWeightFromDiskReqInput, UpdateWeightFromDiskReqInput,
UpdateWeightFromDiskReqOutput, UpdateWeightFromDiskReqOutput,
async_sock_recv,
async_sock_send,
sock_send,
) )
from sglang.srt.managers.load_snapshot import create_load_snapshot_reader from sglang.srt.managers.load_snapshot import create_load_snapshot_reader
from sglang.srt.managers.mm_utils import TensorTransportMode, wrap_shm_features from sglang.srt.managers.mm_utils import TensorTransportMode, wrap_shm_features
@@ -379,19 +383,19 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
context, zmq.PULL, port_args.tokenizer_ipc_name, True context, zmq.PULL, port_args.tokenizer_ipc_name, True
) )
if self.server_args.tokenizer_worker_num == 1: if self.server_args.tokenizer_worker_num == 1:
self.send_to_scheduler = get_zmq_socket( send_to_scheduler = get_zmq_socket(
context, zmq.PUSH, port_args.scheduler_input_ipc_name, True context, zmq.PUSH, port_args.scheduler_input_ipc_name, True
) )
self.send_to_scheduler = SenderWrapper(port_args, send_to_scheduler)
else: else:
from sglang.srt.managers.multi_tokenizer_mixin import SenderWrapper
# Use tokenizer_worker_ipc_name in multi-tokenizer mode # Use tokenizer_worker_ipc_name in multi-tokenizer mode
send_to_scheduler = get_zmq_socket( send_to_scheduler = get_zmq_socket(
context, zmq.PUSH, port_args.tokenizer_worker_ipc_name, False context, zmq.PUSH, port_args.tokenizer_worker_ipc_name, False
) )
# Make sure that each request carries the tokenizer_ipc_name for response routing # Make sure that each request carries the tokenizer_ipc_name for response routing
self.send_to_scheduler = SenderWrapper(port_args, send_to_scheduler) self.send_to_scheduler = SenderWrapper(
port_args, send_to_scheduler, attach_multi_http_worker_info=True
)
self.load_snapshot_reader = create_load_snapshot_reader( self.load_snapshot_reader = create_load_snapshot_reader(
self.server_args, self.server_args,
@@ -1322,7 +1326,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
): ):
tokenized_obj.time_stats.set_api_server_dispatch_time() tokenized_obj.time_stats.set_api_server_dispatch_time()
tokenized_obj = wrap_shm_features(tokenized_obj) tokenized_obj = wrap_shm_features(tokenized_obj)
self.send_to_scheduler.send_pyobj(tokenized_obj) sock_send(self.send_to_scheduler, tokenized_obj)
tokenized_obj.time_stats.set_api_server_dispatch_finish_time() tokenized_obj.time_stats.set_api_server_dispatch_finish_time()
def _send_batch_request( def _send_batch_request(
@@ -1338,7 +1342,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
batch_req = BatchTokenizedEmbeddingReqInput(batch=tokenized_objs) batch_req = BatchTokenizedEmbeddingReqInput(batch=tokenized_objs)
set_time_batch(tokenized_objs, "set_api_server_dispatch_time") set_time_batch(tokenized_objs, "set_api_server_dispatch_time")
self.send_to_scheduler.send_pyobj(batch_req) sock_send(self.send_to_scheduler, batch_req)
set_time_batch(tokenized_objs, "set_api_server_dispatch_finish_time") set_time_batch(tokenized_objs, "set_api_server_dispatch_finish_time")
def _coalesce_streaming_chunks( def _coalesce_streaming_chunks(
@@ -1663,7 +1667,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
): ):
return return
req = AbortReq(rid=rid, abort_all=abort_all) req = AbortReq(rid=rid, abort_all=abort_all)
self.send_to_scheduler.send_pyobj(req) sock_send(self.send_to_scheduler, req)
if self.enable_metrics: if self.enable_metrics:
# TODO: also use custom_labels from the request # TODO: also use custom_labels from the request
self.metrics_collector.observe_one_aborted_request( self.metrics_collector.observe_one_aborted_request(
@@ -1674,7 +1678,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
async with self.is_pause_cond: async with self.is_pause_cond:
self.is_pause = True self.is_pause = True
if obj.mode != "abort": if obj.mode != "abort":
await self.send_to_scheduler.send_pyobj(obj) await async_sock_send(self.send_to_scheduler, obj)
else: else:
# we are using the model_update_lock to check if there is still on-going requests. # we are using the model_update_lock to check if there is still on-going requests.
while True: while True:
@@ -1688,7 +1692,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
async def continue_generation(self, obj: ContinueGenerationReqInput): async def continue_generation(self, obj: ContinueGenerationReqInput):
async with self.is_pause_cond: async with self.is_pause_cond:
self.is_pause = False self.is_pause = False
await self.send_to_scheduler.send_pyobj(obj) await async_sock_send(self.send_to_scheduler, obj)
self.is_pause_cond.notify_all() self.is_pause_cond.notify_all()
async def update_weights_from_disk( async def update_weights_from_disk(
@@ -1733,7 +1737,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
async def _wait_for_model_update_from_disk( async def _wait_for_model_update_from_disk(
self, obj: UpdateWeightFromDiskReqInput self, obj: UpdateWeightFromDiskReqInput
) -> Tuple[bool, str]: ) -> Tuple[bool, str]:
self.send_to_scheduler.send_pyobj(obj) sock_send(self.send_to_scheduler, obj)
self.model_update_result = asyncio.Future() self.model_update_result = asyncio.Future()
if self.server_args.dp_size == 1: if self.server_args.dp_size == 1:
result = await self.model_update_result result = await self.model_update_result
@@ -1773,12 +1777,12 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
# Let the exception propagate to the caller. # Let the exception propagate to the caller.
# Only legal requests will be sent to scheduler. # Only legal requests will be sent to scheduler.
logging.getLogger().setLevel(obj.log_level.upper()) logging.getLogger().setLevel(obj.log_level.upper())
self.send_to_scheduler.send_pyobj(obj) sock_send(self.send_to_scheduler, obj)
logging.info(f"Config logging: {obj=}") logging.info(f"Config logging: {obj=}")
async def freeze_gc(self): async def freeze_gc(self):
"""Send a freeze_gc message to the scheduler first, then freeze locally.""" """Send a freeze_gc message to the scheduler first, then freeze locally."""
self.send_to_scheduler.send_pyobj(FreezeGCReq()) sock_send(self.send_to_scheduler, FreezeGCReq())
freeze_gc("Tokenizer Manager") freeze_gc("Tokenizer Manager")
return None return None
@@ -1825,7 +1829,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
"""The event loop that handles requests""" """The event loop that handles requests"""
while True: while True:
with self.soft_watchdog.disable(): with self.soft_watchdog.disable():
recv_obj = await self.recv_from_detokenizer.recv_pyobj() recv_obj = await async_sock_recv(self.recv_from_detokenizer)
if isinstance( if isinstance(
recv_obj, recv_obj,
(BatchStrOutput, BatchEmbeddingOutput, BatchTokenIDOutput), (BatchStrOutput, BatchEmbeddingOutput, BatchTokenIDOutput),
@@ -2648,7 +2652,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
self._subprocess_watchdog.stop() self._subprocess_watchdog.stop()
# Ask schedulers to release resources in userspace and exit (see # Ask schedulers to release resources in userspace and exit (see
# ShutdownReq), then wait for them before hard-killing the rest. # ShutdownReq), then wait for them before hard-killing the rest.
self.send_to_scheduler.send_pyobj(ShutdownReq()) sock_send(self.send_to_scheduler, ShutdownReq())
deadline = time.monotonic() + 15 deadline = time.monotonic() + 15
while time.monotonic() < deadline and collect_scheduler_processes(): while time.monotonic() < deadline and collect_scheduler_processes():
time.sleep(0.1) time.sleep(0.1)
@@ -2717,7 +2721,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
state.event.set() state.event.set()
def update_active_ranks(self, ranks: ActiveRanksOutput): def update_active_ranks(self, ranks: ActiveRanksOutput):
self.send_to_scheduler.send_pyobj(ranks) sock_send(self.send_to_scheduler, ranks)
def _handle_open_session_req_output(self, recv_obj): def _handle_open_session_req_output(self, recv_obj):
future = self.session_futures.get(recv_obj.session_id) future = self.session_futures.get(recv_obj.session_id)
@@ -3115,3 +3119,25 @@ class SignalHandler:
# | http | no | waiting queue | type 1 | type 1 exception | del in _handle_abort_req | # | http | no | waiting queue | type 1 | type 1 exception | del in _handle_abort_req |
# | http | no | running | type 3 | type 3 exception | del in _handle_batch_output | # | http | no | running | type 3 | type 3 exception | del in _handle_batch_output |
# #
class SenderWrapper:
def __init__(
self,
port_args,
send_to_scheduler,
attach_multi_http_worker_info=False,
):
self.port_args = port_args
self.send_to_scheduler = send_to_scheduler
self.attach_multi_http_worker_info = attach_multi_http_worker_info
def _stamp_http_worker_ipc(self, obj):
if not self.attach_multi_http_worker_info:
return
if isinstance(obj, BaseReq):
obj.http_worker_ipc = self.port_args.tokenizer_ipc_name
def send_pyobj(self, obj, flags=0):
self._stamp_http_worker_ipc(obj)
return self.send_to_scheduler.send_pyobj(obj, flags=flags)
@@ -12,6 +12,7 @@ import zmq
from sglang.srt.entrypoints.http_server import launch_server from sglang.srt.entrypoints.http_server import launch_server
from sglang.srt.environ import envs 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.server_args import ServerArgs
from sglang.srt.utils.network import get_free_port, get_zmq_socket_on_host from sglang.srt.utils.network import get_free_port, get_zmq_socket_on_host
from sglang.test.scripted_runtime.io_struct import ( from sglang.test.scripted_runtime.io_struct import (
@@ -89,7 +90,7 @@ class ScriptedHttpServer:
raise RuntimeError(f"ScriptedHttpServer is dirty: {self._dirty}") raise RuntimeError(f"ScriptedHttpServer is dirty: {self._dirty}")
fn_path = f"{script_fn.__module__}:{script_fn.__qualname__}" 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._socket.poll(int(timeout_s * 1000)):
if not self._server_process.is_alive(): 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" self._dirty = f"script {fn_path!r} timed out after {timeout_s}s"
raise TimeoutError(self._dirty) raise TimeoutError(self._dirty)
reply = self._socket.recv_pyobj() reply = sock_recv(self._socket)
match reply: match reply:
case ScriptFailed(traceback=tb): case ScriptFailed(traceback=tb):
raise AssertionError(f"scripted-runtime script failed:\n{tb}") raise AssertionError(f"scripted-runtime script failed:\n{tb}")
@@ -116,7 +117,7 @@ class ScriptedHttpServer:
fatal_error: Optional[OutOfBandError] = None fatal_error: Optional[OutOfBandError] = None
try: try:
try: try:
self._socket.send_pyobj(Shutdown()) sock_send(self._socket, Shutdown())
except zmq.ZMQError: except zmq.ZMQError:
pass pass
@@ -139,7 +140,7 @@ class ScriptedHttpServer:
f"{LISTENER_ACCEPT_TIMEOUT_S}s" f"{LISTENER_ACCEPT_TIMEOUT_S}s"
) )
ready = self._socket.recv_pyobj() ready = sock_recv(self._socket)
if not isinstance(ready, HookReady): if not isinstance(ready, HookReady):
raise RuntimeError( raise RuntimeError(
f"ScriptedHttpServer: expected HookReady handshake, got {ready!r}" f"ScriptedHttpServer: expected HookReady handshake, got {ready!r}"
@@ -11,6 +11,7 @@ from typing import TYPE_CHECKING, Generator, List, Optional, Tuple
import zmq import zmq
from sglang.srt.environ import envs 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.srt.utils.network import get_zmq_socket
from sglang.test.scripted_runtime.background_http_poster import BackgroundHttpPoster from sglang.test.scripted_runtime.background_http_poster import BackgroundHttpPoster
from sglang.test.scripted_runtime.context import ScriptedContext 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) socket = get_zmq_socket(ctx_zmq, zmq.PAIR, endpoint, bind=False)
try: try:
yield from _drive_engine_through_warmup(self._context) yield from _drive_engine_through_warmup(self._context)
socket.send_pyobj(HookReady()) sock_send(socket, HookReady())
while True: while True:
msg = socket.recv_pyobj() msg = sock_recv(socket)
match msg: match msg:
case Shutdown(): case Shutdown():
return return
@@ -167,11 +168,12 @@ class ScriptedSchedulerHook:
try: try:
yield from sub_gen yield from sub_gen
except Exception: except Exception:
socket.send_pyobj( sock_send(
ScriptFailed(traceback=traceback.format_exc()) socket,
ScriptFailed(traceback=traceback.format_exc()),
) )
else: else:
socket.send_pyobj(ScriptSucceeded()) sock_send(socket, ScriptSucceeded())
case _: case _:
raise ValueError(f"dispatch loop: unknown command {msg!r}") raise ValueError(f"dispatch loop: unknown command {msg!r}")
finally: finally:
@@ -11,6 +11,7 @@ from sglang.srt.managers.io_struct import (
BatchTokenizedGenerateReqInput, BatchTokenizedGenerateReqInput,
TokenizedEmbeddingReqInput, TokenizedEmbeddingReqInput,
TokenizedGenerateReqInput, TokenizedGenerateReqInput,
sock_recv,
) )
_WORK_REQ_TYPES = ( _WORK_REQ_TYPES = (
@@ -64,7 +65,7 @@ class ScriptedTokenizerRecvProxy:
def _drain_underlying(self) -> None: def _drain_underlying(self) -> None:
while True: while True:
try: try:
req = self._underlying.recv_pyobj(zmq.NOBLOCK) req = sock_recv(self._underlying, zmq.NOBLOCK)
except zmq.ZMQError: except zmq.ZMQError:
break break
if isinstance(req, _WORK_REQ_TYPES): if isinstance(req, _WORK_REQ_TYPES):
@@ -1,6 +1,10 @@
from __future__ import annotations from __future__ import annotations
import threading
import unittest import unittest
import uuid
import zmq
from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.scripted_runtime.http_server import ScriptedHttpServer from sglang.test.scripted_runtime.http_server import ScriptedHttpServer
@@ -20,23 +24,41 @@ def _sample_script(ctx, *args):
_EXPECTED_FN_PATH = f"{_sample_script.__module__}:{_sample_script.__qualname__}" _EXPECTED_FN_PATH = f"{_sample_script.__module__}:{_sample_script.__qualname__}"
_NO_REPLY = object()
class _FakePairSocket: class _PairSocketHarness:
def __init__(self, *, poll_result: bool, reply: object = None) -> None: def __init__(self, *, reply: object = _NO_REPLY) -> None:
self._poll_result = poll_result self._ctx = zmq.Context()
self.server_socket = self._ctx.socket(zmq.PAIR)
self._peer_socket = self._ctx.socket(zmq.PAIR)
self._endpoint = f"inproc://scripted-http-server-{uuid.uuid4().hex}"
self.server_socket.bind(self._endpoint)
self._peer_socket.connect(self._endpoint)
self._reply = reply self._reply = reply
self.sent: list = [] self.sent: list = []
self._thread = None
if reply is not _NO_REPLY:
self._thread = threading.Thread(target=self._reply_once, daemon=True)
self._thread.start()
def send_pyobj(self, obj: object) -> None: def __enter__(self):
self.sent.append(obj) return self
def poll(self, timeout_ms: int) -> bool: def __exit__(self, exc_type, exc, tb):
return self._poll_result if self._thread is not None:
self._thread.join(timeout=1)
self.server_socket.close(0)
self._peer_socket.close(0)
self._ctx.term()
def recv_pyobj(self) -> object: def _reply_once(self):
return self._reply self.sent.append(self._peer_socket.recv_pyobj())
self._peer_socket.send_pyobj(self._reply)
def assert_no_sent_message(self, test_case: unittest.TestCase) -> None:
test_case.assertFalse(self._peer_socket.poll(50))
class _FakeProcess: class _FakeProcess:
@@ -48,7 +70,7 @@ class _FakeProcess:
return self._alive return self._alive
def _make_server(socket: _FakePairSocket, process: _FakeProcess) -> ScriptedHttpServer: def _make_server(socket: zmq.Socket, process: _FakeProcess) -> ScriptedHttpServer:
server = ScriptedHttpServer.__new__(ScriptedHttpServer) server = ScriptedHttpServer.__new__(ScriptedHttpServer)
server._socket = socket server._socket = socket
server._server_process = process server._server_process = process
@@ -59,69 +81,69 @@ def _make_server(socket: _FakePairSocket, process: _FakeProcess) -> ScriptedHttp
class TestExecuteScriptReplyMatching(CustomTestCase): class TestExecuteScriptReplyMatching(CustomTestCase):
def test_returns_on_script_succeeded(self): def test_returns_on_script_succeeded(self):
socket = _FakePairSocket(poll_result=True, reply=ScriptSucceeded()) with _PairSocketHarness(reply=ScriptSucceeded()) as pair:
server = _make_server(socket, _FakeProcess(alive=True)) server = _make_server(pair.server_socket, _FakeProcess(alive=True))
server.execute_script(_sample_script) server.execute_script(_sample_script)
self.assertEqual(socket.sent, [RunScript(fn_path=_EXPECTED_FN_PATH, args=())]) self.assertEqual(pair.sent, [RunScript(fn_path=_EXPECTED_FN_PATH, args=())])
def test_forwards_args_in_run_script(self): def test_forwards_args_in_run_script(self):
socket = _FakePairSocket(poll_result=True, reply=ScriptSucceeded()) with _PairSocketHarness(reply=ScriptSucceeded()) as pair:
server = _make_server(socket, _FakeProcess(alive=True)) server = _make_server(pair.server_socket, _FakeProcess(alive=True))
server.execute_script(_sample_script, args=(1, "two")) server.execute_script(_sample_script, args=(1, "two"))
self.assertEqual( self.assertEqual(
socket.sent, [RunScript(fn_path=_EXPECTED_FN_PATH, args=(1, "two"))] pair.sent, [RunScript(fn_path=_EXPECTED_FN_PATH, args=(1, "two"))]
) )
def test_script_failed_reply_raises_assertion_with_traceback(self): def test_script_failed_reply_raises_assertion_with_traceback(self):
socket = _FakePairSocket( with _PairSocketHarness(
poll_result=True, reply=ScriptFailed(traceback="REMOTE-TB-MARKER") reply=ScriptFailed(traceback="REMOTE-TB-MARKER")
) ) as pair:
server = _make_server(socket, _FakeProcess(alive=True)) server = _make_server(pair.server_socket, _FakeProcess(alive=True))
with self.assertRaisesRegex(AssertionError, "REMOTE-TB-MARKER"): with self.assertRaisesRegex(AssertionError, "REMOTE-TB-MARKER"):
server.execute_script(_sample_script) server.execute_script(_sample_script)
def test_unexpected_reply_raises_runtime_error(self): def test_unexpected_reply_raises_runtime_error(self):
socket = _FakePairSocket(poll_result=True, reply=HookReady()) with _PairSocketHarness(reply=HookReady()) as pair:
server = _make_server(socket, _FakeProcess(alive=True)) server = _make_server(pair.server_socket, _FakeProcess(alive=True))
with self.assertRaisesRegex(RuntimeError, "unexpected message"): with self.assertRaisesRegex(RuntimeError, "unexpected message"):
server.execute_script(_sample_script) server.execute_script(_sample_script)
class TestExecuteScriptNoReply(CustomTestCase): class TestExecuteScriptNoReply(CustomTestCase):
def test_timeout_when_process_still_alive(self): def test_timeout_when_process_still_alive(self):
socket = _FakePairSocket(poll_result=False) with _PairSocketHarness() as pair:
server = _make_server(socket, _FakeProcess(alive=True)) server = _make_server(pair.server_socket, _FakeProcess(alive=True))
with self.assertRaisesRegex(TimeoutError, "timed out"): with self.assertRaisesRegex(TimeoutError, "timed out"):
server.execute_script(_sample_script, timeout_s=0.01) server.execute_script(_sample_script, timeout_s=0.01)
self.assertIn("timed out", server._dirty) self.assertIn("timed out", server._dirty)
def test_runtime_error_when_process_died(self): def test_runtime_error_when_process_died(self):
socket = _FakePairSocket(poll_result=False) with _PairSocketHarness() as pair:
server = _make_server(socket, _FakeProcess(alive=False)) server = _make_server(pair.server_socket, _FakeProcess(alive=False))
with self.assertRaisesRegex(RuntimeError, "died before responding"): with self.assertRaisesRegex(RuntimeError, "died before responding"):
server.execute_script(_sample_script, timeout_s=0.01) server.execute_script(_sample_script, timeout_s=0.01)
self.assertIn("died before responding", server._dirty) self.assertIn("died before responding", server._dirty)
class TestExecuteScriptDirtyGuard(CustomTestCase): class TestExecuteScriptDirtyGuard(CustomTestCase):
def test_refuses_to_run_when_already_dirty(self): def test_refuses_to_run_when_already_dirty(self):
socket = _FakePairSocket(poll_result=True, reply=ScriptSucceeded()) with _PairSocketHarness() as pair:
server = _make_server(socket, _FakeProcess(alive=True)) server = _make_server(pair.server_socket, _FakeProcess(alive=True))
server._dirty = "prior timeout" server._dirty = "prior timeout"
with self.assertRaisesRegex(RuntimeError, "dirty"): with self.assertRaisesRegex(RuntimeError, "dirty"):
server.execute_script(_sample_script) server.execute_script(_sample_script)
self.assertEqual(socket.sent, []) pair.assert_no_sent_message(self)
if __name__ == "__main__": if __name__ == "__main__":