Convert IPC dataclasses to msgspec.Struct with opt-in msgpack transport (#28688)

Co-authored-by: Lianmin Zheng <lianminzheng@gmail.com>
This commit is contained in:
Rain Jiang
2026-06-26 12:04:03 -07:00
committed by GitHub
co-authored by Lianmin Zheng
parent 714011a40f
commit be1930133a
25 changed files with 694 additions and 354 deletions
+9 -7
View File
@@ -22,7 +22,7 @@ import torch
import torch.distributed as dist
import zmq
from sglang.srt.managers.io_struct import sock_recv, sock_send
from sglang.srt.managers.io_struct import sock_recv, sock_send, wrap_as_pickle
# -------------------------------------- config base ------------------------------------------
@@ -1440,7 +1440,7 @@ def _create_zmq_rpc_broadcast(
except Exception as e:
_log(f"[ZmqRpc] error inside handler: {e}")
resp = {"result": None, "error": str(e)}
sock_send(sock, resp)
sock_send(sock, wrap_as_pickle(resp))
thread = threading.Thread(target=serve_loop, daemon=True)
thread.start()
@@ -1479,11 +1479,13 @@ class _ZmqRpcHandle:
def call(*args, **kwargs):
sock_send(
self._socket,
{
"method": method_name,
"args": args,
"kwargs": kwargs,
},
wrap_as_pickle(
{
"method": method_name,
"args": args,
"kwargs": kwargs,
}
),
)
response = sock_recv(self._socket)
if response["error"]:
@@ -26,7 +26,7 @@ from sglang.srt.disaggregation.encode_server import (
handle_scheduler_receive_url_request,
launch_encoder,
)
from sglang.srt.managers.io_struct import async_sock_send
from sglang.srt.managers.io_struct import async_sock_send, wrap_as_pickle
from sglang.srt.managers.schedule_batch import Modality
from sglang.srt.server_args import PortArgs, ServerArgs
from sglang.srt.utils import random_uuid
@@ -96,7 +96,7 @@ class SGLangEncoderServer(SGLangEncoderServicer):
"part_idx": request.part_idx,
}
for socket in self.send_sockets:
await async_sock_send(socket, request_dict)
await async_sock_send(socket, wrap_as_pickle(request_dict))
# gRPC encode is image-only; encoder.encode() requires modality
(
@@ -50,6 +50,7 @@ from sglang.srt.managers.io_struct import (
async_sock_recv,
async_sock_send,
sock_send,
wrap_as_pickle,
)
from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem
from sglang.srt.mem_cache.multimodal_cache import EmbeddingResult, MultiModalStaticCache
@@ -2398,12 +2399,14 @@ class EncoderScheduler:
for sock in self.send_sockets:
sock_send(
sock,
{
"type": "batch_encode",
"modality": modality.name,
"requests": requests,
"enter_time": start,
},
wrap_as_pickle(
{
"type": "batch_encode",
"modality": modality.name,
"requests": requests,
"enter_time": start,
}
),
)
logger.info(f"Dispatching batch of {len(group)} {modality.name} requests")
@@ -2449,7 +2452,7 @@ class EncoderScheduler:
req = p.request
try:
for sock in self.send_sockets:
sock_send(sock, req)
sock_send(sock, wrap_as_pickle(req))
result = await self.encoder.encode_request(req, modality)
if not p.future.done():
p.future.set_result(result)
@@ -2738,7 +2741,7 @@ class DPDispatcher:
)
try:
await async_sock_send(self.dispatch_sockets[rank], request)
await async_sock_send(self.dispatch_sockets[rank], wrap_as_pickle(request))
# An alive-but-stuck worker (NCCL deadlock etc.) wouldn't trip
# the watchdog, so bound the wait explicitly.
return await asyncio.wait_for(future, timeout=ENCODER_REQ_TIMEOUT)
@@ -2787,7 +2790,7 @@ class DPDispatcher:
f"dp_rank={rank}, pending={self.pending_counts}"
)
try:
await async_sock_send(self.dispatch_sockets[rank], request)
await async_sock_send(self.dispatch_sockets[rank], wrap_as_pickle(request))
return await asyncio.wait_for(future, timeout=ENCODER_REQ_TIMEOUT)
except asyncio.TimeoutError:
self.pending_futures[rank].pop(key, None)
@@ -2828,7 +2831,9 @@ class DPDispatcher:
self.req_id_to_rank[req_id] = rank
rank_keys.append((rank, req_id))
request_copy = {**request, "req_id": req_id}
await async_sock_send(self.dispatch_sockets[rank], request_copy)
await async_sock_send(
self.dispatch_sockets[rank], wrap_as_pickle(request_copy)
)
futures.append(future)
# Concurrent wait → total bounded by eff_timeout, not
# dp_size × eff_timeout.
@@ -3056,7 +3061,7 @@ async def _dp_worker_handle_request(
# pyzmq async send isn't safe for concurrent senders.
try:
async with send_lock:
await async_sock_send(send_sock, envelope)
await async_sock_send(send_sock, wrap_as_pickle(envelope))
except Exception:
logger.error(
f"DP worker {dp_rank} failed to send envelope for "
@@ -3603,7 +3608,7 @@ async def handle_encode_request(request: dict):
)
else:
for socket in send_sockets:
sock_send(socket, request)
sock_send(socket, wrap_as_pickle(request))
nbytes, embedding_len, embedding_dim, error_msg, error_code = (
await encoder.encode_request(request, modality)
)
@@ -3824,7 +3829,7 @@ async def health_generate():
# Broadcast to other TP ranks so distributed ops stay in sync
for socket in send_sockets:
sock_send(socket, dummy_request)
sock_send(socket, wrap_as_pickle(dummy_request))
# Run encode on rank 0 with timeout
_, _, _, error_msg, _ = await asyncio.wait_for(
+9 -6
View File
@@ -109,6 +109,7 @@ from sglang.srt.utils import (
set_prometheus_multiproc_dir,
set_ulimit,
)
from sglang.srt.utils.msgspec_utils import msgspec_to_builtins
from sglang.srt.utils.network import get_zmq_socket, is_port_available
from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter
from sglang.srt.utils.watchdog import SubprocessWatchdog
@@ -994,12 +995,14 @@ class Engine(EngineScoreMixin, EngineBase):
internal_states = self.loop.run_until_complete(
self.tokenizer_manager.get_internal_state()
)
return {
**dataclasses.asdict(self.tokenizer_manager.server_args),
**self._scheduler_init_result.scheduler_infos[0],
"internal_states": internal_states,
"version": __version__,
}
return msgspec_to_builtins(
{
**dataclasses.asdict(self.tokenizer_manager.server_args),
**self._scheduler_init_result.scheduler_infos[0],
"internal_states": internal_states,
"version": __version__,
}
)
def init_weights_update_group(
self,
+4 -2
View File
@@ -15,6 +15,8 @@ from typing import Any, Awaitable, Callable, Dict, List, Optional
from pydantic import ValidationError
from sglang.srt.utils.msgspec_utils import msgspec_to_builtins
logger = logging.getLogger(__name__)
@@ -383,9 +385,9 @@ class RuntimeHandle:
return json.dumps(result, default=str)
def get_server_info(self) -> str:
result: Dict[str, Any] = dict(dataclasses.asdict(self.server_args))
result: Dict[str, Any] = dataclasses.asdict(self.server_args)
result.update(self.scheduler_info)
return json.dumps(result, default=str)
return json.dumps(msgspec_to_builtins(result), default=str)
def health_check(self) -> bool:
from sglang.srt.managers.tokenizer_manager import ServerStatus
+16 -13
View File
@@ -173,6 +173,7 @@ from sglang.srt.utils.json_response import (
dumps_json,
orjson_response,
)
from sglang.srt.utils.msgspec_utils import msgspec_to_builtins
from sglang.srt.utils.watchdog import SubprocessWatchdog
from sglang.utils import get_exception_traceback
from sglang.version import __version__
@@ -698,16 +699,18 @@ async def server_info():
server_args = _global_state.tokenizer_manager.server_args
# server_args.model_config is not serializable but should be excluded by asdict.
return {
**dataclasses.asdict(server_args),
**_global_state.scheduler_info,
"internal_states": internal_states,
"version": __version__,
# Structured KV-event publisher descriptor for KV-aware routers.
# `None` when publishing is disabled or misconfigured; see
# `ServerArgs.describe_kv_events_publisher` for the precise contract.
"kv_events": server_args.describe_kv_events_publisher(),
}
return msgspec_to_builtins(
{
**dataclasses.asdict(server_args),
**_global_state.scheduler_info,
"internal_states": internal_states,
"version": __version__,
# Structured KV-event publisher descriptor for KV-aware routers.
# `None` when publishing is disabled or misconfigured; see
# `ServerArgs.describe_kv_events_publisher` for the precise contract.
"kv_events": server_args.describe_kv_events_publisher(),
}
)
@app.get("/get_load")
@@ -1417,7 +1420,7 @@ async def load_lora_adapter(
"""Load a new LoRA adapter without re-launching the server."""
result = await _global_state.tokenizer_manager.load_lora_adapter(obj, request)
status_code = HTTPStatus.OK if result.success else HTTPStatus.BAD_REQUEST
return ORJSONResponse(result, status_code=status_code)
return ORJSONResponse(msgspec_to_builtins(result), status_code=status_code)
@app.api_route("/load_lora_adapter_from_tensors", methods=["POST"])
@@ -1429,7 +1432,7 @@ async def load_lora_adapter_from_tensors(
obj, request
)
status_code = HTTPStatus.OK if result.success else HTTPStatus.BAD_REQUEST
return ORJSONResponse(result, status_code=status_code)
return ORJSONResponse(msgspec_to_builtins(result), status_code=status_code)
@app.api_route("/unload_lora_adapter", methods=["POST"])
@@ -1440,7 +1443,7 @@ async def unload_lora_adapter(
"""Load a new LoRA adapter without re-launching the server."""
result = await _global_state.tokenizer_manager.unload_lora_adapter(obj, request)
status_code = HTTPStatus.OK if result.success else HTTPStatus.BAD_REQUEST
return ORJSONResponse(result, status_code=status_code)
return ORJSONResponse(msgspec_to_builtins(result), status_code=status_code)
@app.api_route("/open_session", methods=["GET", "POST"])
+4
View File
@@ -241,6 +241,10 @@ class Envs:
SGLANG_LOG_SCHEDULER_STATUS_TARGET = EnvStr("")
SGLANG_LOG_SCHEDULER_STATUS_INTERVAL = EnvFloat(60.0)
# IPC
SGLANG_USE_PICKLE_IPC = EnvBool(True)
SGLANG_LOG_PICKLE_IPC_OBJECTS = EnvBool(False)
# SGLang CI
SGLANG_IS_IN_CI = EnvBool(False)
SGLANG_IS_IN_CI_AMD = EnvBool(False)
+5 -4
View File
@@ -15,16 +15,17 @@
import asyncio
from collections import OrderedDict
from dataclasses import dataclass, field, fields
from typing import Dict, List, Optional, Union
from uuid import NAMESPACE_URL, uuid4, uuid5
import msgspec
from msgspec.structs import fields
from sglang.srt.utils import ConcurrentCounter
from sglang.srt.utils.aio_rwlock import RWLock
@dataclass(frozen=True)
class LoRARef:
class LoRARef(msgspec.Struct, frozen=True, array_like=True):
"""
Reference record for a LoRA model.
@@ -33,7 +34,7 @@ class LoRARef:
keys (e.g., radix cache).
"""
lora_id: str = field(default_factory=lambda: uuid4().hex)
lora_id: str = msgspec.field(default_factory=lambda: uuid4().hex)
lora_name: Optional[str] = None
lora_path: Optional[str] = None
pinned: Optional[bool] = None
@@ -38,6 +38,8 @@ from sglang.srt.managers.io_struct import (
TokenizedGenerateReqInput,
sock_recv,
sock_send,
unwrap_from_pickle,
wrap_as_pickle,
)
from sglang.srt.managers.load_snapshot import create_load_snapshot_reader
from sglang.srt.managers.schedule_batch import Req
@@ -238,10 +240,14 @@ class DataParallelController:
if refresh_load_budget and self.refresh_load_budget_on_dispatch:
self.refresh_load_budget()
req.time_stats = DPControllerReqTimeStats.new_from_obj(req.time_stats)
time_stats = DPControllerReqTimeStats.new_from_obj(
unwrap_from_pickle(req.time_stats)
)
req.time_stats.set_dp_dispatch_time()
time_stats.set_dp_dispatch_time()
req.time_stats = wrap_as_pickle(time_stats)
self.dispatching(req)
req.time_stats = time_stats
req.time_stats.set_dp_dispatch_finish_time()
def dispatch_batch_generate(self, batch_req: BatchTokenizedGenerateReqInput):
@@ -382,11 +388,11 @@ class DataParallelController:
connected_clients = 0
while connected_clients < expected_clients:
# Wait for client handshake
client_rank = sock_recv(rep_socket).decode()
client_rank = sock_recv(rep_socket)
logger.debug(f"Received handshake from node {client_rank}")
# Send worker ports to client
sock_send(rep_socket, worker_ports)
sock_send(rep_socket, wrap_as_pickle(worker_ports))
connected_clients += 1
logger.debug(
f"Sent worker ports to {connected_clients}/{expected_clients} nodes"
@@ -411,7 +417,7 @@ class DataParallelController:
while True:
# Wait for client handshake
try:
client_rank = sock_recv(rep_socket).decode()
client_rank = sock_recv(rep_socket)
except Exception:
logger.exception(
"Failed to recv/decode handshake in reply thread; continue"
@@ -420,7 +426,7 @@ class DataParallelController:
logger.debug(f"Received handshake from node {client_rank}")
# Send worker ports to client
sock_send(rep_socket, worker_ports)
sock_send(rep_socket, wrap_as_pickle(worker_ports))
logger.debug(f"Sent worker ports to node {client_rank}")
def _receive_ports_as_client(self, endpoint: str, node_rank: int) -> List[int]:
@@ -433,7 +439,7 @@ class DataParallelController:
try:
# Send handshake with our node rank
sock_send(req_socket, str(node_rank).encode())
sock_send(req_socket, wrap_as_pickle(str(node_rank)))
# Receive worker ports
worker_ports = sock_recv(req_socket)
+3 -4
View File
@@ -12,20 +12,19 @@
# limitations under the License.
# ==============================================================================
"""
Dataclasses for embedding injection.
Structs for embedding injection.
These are placed in a separate module to avoid circular imports between
io_struct.py and schedule_batch.py.
"""
from dataclasses import dataclass
from typing import List
import msgspec
import torch
@dataclass
class PositionalEmbeds:
class PositionalEmbeds(msgspec.Struct, array_like=True):
"""Embeddings to place at specific token positions.
Accepts either a list of [1, hidden_dim] tensors or a pre-stacked [N, hidden_dim] tensor.
File diff suppressed because it is too large Load Diff
@@ -53,6 +53,8 @@ from sglang.srt.managers.io_struct import (
async_sock_send,
sock_recv,
sock_send,
unwrap_from_pickle,
wrap_as_pickle,
)
from sglang.srt.managers.load_snapshot import (
create_load_snapshot_reader,
@@ -122,17 +124,29 @@ def _extract_field_by_index(
if field is None:
return None
should_wrap_result = field_name in ("customized_info", "time_stats")
if should_wrap_result:
field = unwrap_from_pickle(field)
if field is None:
return None
if isinstance(field, dict):
new_field = {}
for k, v in field.items():
new_field[k] = v[index] if len(v) > index else None
if len(v) > index:
new_field[k] = [v[index]] if should_wrap_result else v[index]
else:
new_field[k] = [None] if should_wrap_result else None
if should_wrap_result:
return wrap_as_pickle(new_field) if new_field else None
return new_field
if check_length:
if len(field) <= index:
return None
return [field[index]]
new_field = [field[index]]
return wrap_as_pickle(new_field) if should_wrap_result else new_field
def _handle_output_by_index(output, i):
+3 -2
View File
@@ -266,6 +266,7 @@ from sglang.srt.utils.hf_transformers_utils import (
get_tokenizer,
get_tokenizer_from_processor,
)
from sglang.srt.utils.msgspec_utils import msgspec_to_builtins
from sglang.srt.utils.numa_utils import get_numa_node_if_available, numa_bind_to_node
from sglang.srt.utils.nvtx_utils import scheduler_nvtx_method
from sglang.srt.utils.tensor_bridge import use_mlx
@@ -3752,7 +3753,7 @@ class Scheduler(
# This field is not serializable.
ret.pop("model_config", None)
return GetInternalStateReqOutput(internal_state=ret)
return GetInternalStateReqOutput(internal_state=msgspec_to_builtins(ret))
def set_internal_state(self, recv_req: SetInternalStateReq):
server_args_dict = recv_req.server_args
@@ -3801,7 +3802,7 @@ class Scheduler(
server_args.pop("model_config", None)
return SetInternalStateReqOutput(
updated=if_success,
server_args=server_args,
server_args=msgspec_to_builtins(server_args),
)
def save_remote_model(self, **kwargs):
@@ -1,6 +1,5 @@
from __future__ import annotations
import dataclasses
import time
from dataclasses import dataclass
from typing import (
@@ -10,13 +9,14 @@ from typing import (
Optional,
)
import msgspec
import zmq
from sglang.srt.disaggregation.kv_events import (
EventPublisherFactory,
KVEventBatch,
)
from sglang.srt.managers.io_struct import sock_send
from sglang.srt.managers.io_struct import hook_custom_types, sock_send
if TYPE_CHECKING:
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
@@ -26,8 +26,7 @@ if TYPE_CHECKING:
class SchedulerStats: ... # type: ignore[no-redef]
@dataclasses.dataclass
class KvMetrics:
class KvMetrics(msgspec.Struct, tag=True, kw_only=True, array_like=True):
request_active_slots: int = 0
request_total_slots: int = 0
kv_active_blocks: int = 0
@@ -38,6 +37,9 @@ class KvMetrics:
data_parallel_rank: int = 0
hook_custom_types(KvMetrics)
@dataclass(kw_only=True, slots=True)
class SchedulerKvEventsPublisher:
kv_events_config: Optional[str]
@@ -19,6 +19,7 @@ from sglang.srt.managers.io_struct import (
BatchEmbeddingOutput,
BatchTokenIDOutput,
CachedTokensDetails,
wrap_as_pickle,
)
from sglang.srt.managers.schedule_batch import (
BaseFinishReason,
@@ -229,7 +230,7 @@ class SchedulerOutputStreamer:
BatchEmbeddingOutput(
rids=rids,
http_worker_ipcs=http_worker_ipcs,
time_stats=time_stats,
time_stats=wrap_as_pickle(time_stats),
finished_reasons=finished_reasons,
embeddings=embeddings,
prompt_tokens=prompt_tokens,
@@ -516,7 +517,7 @@ class _GenerationStreamAccumulator:
spec_verify_ct=self.spec_verify_ct,
spec_num_correct_drafts=self.spec_num_correct_drafts,
spec_correct_drafts_histogram=self.spec_correct_drafts_histogram,
time_stats=self.time_stats,
time_stats=wrap_as_pickle(self.time_stats),
finished_reasons=self.finished_reasons,
decoded_texts=self.decoded_texts,
decode_ids=self.decode_ids_list,
@@ -549,7 +550,9 @@ class _GenerationStreamAccumulator:
output_hidden_states=self.output_hidden_states,
routed_experts=self.routed_experts,
indexer_topk=self.indexer_topk,
customized_info=self.customized_info,
customized_info=(
wrap_as_pickle(self.customized_info) if self.customized_info else None
),
placeholder_tokens_idx=None,
placeholder_tokens_val=None,
retraction_counts=self.retraction_counts,
@@ -89,6 +89,9 @@ class SchedulerRequestReceiver:
recv_reqs = self._broadcast_reqs_across_ranks(recv_reqs)
if self.ps.pp_rank == 0:
self.unwrap_pickle_wrapper(recv_reqs)
recv_reqs = self._apply_mm_receiver(recv_reqs)
self._finalize_shm_features(recv_reqs)
@@ -195,6 +198,19 @@ class SchedulerRequestReceiver:
)
return recv_reqs
def unwrap_pickle_wrapper(self, recv_reqs: Optional[List]) -> None:
if not recv_reqs:
return
for req in recv_reqs:
if isinstance(req, (TokenizedGenerateReqInput, TokenizedEmbeddingReqInput)):
req.unwrap_pickle_fields()
elif isinstance(
req, (BatchTokenizedGenerateReqInput, BatchTokenizedEmbeddingReqInput)
):
for sub_req in req:
sub_req.unwrap_pickle_fields()
def _apply_mm_receiver(self, recv_reqs: List) -> List:
# Process MM requests under EPD-disaggregation mode
if (
@@ -80,6 +80,7 @@ from sglang.srt.managers.io_struct import (
async_sock_recv,
async_sock_send,
sock_send,
unwrap_from_pickle,
)
from sglang.srt.managers.load_snapshot import create_load_snapshot_reader
from sglang.srt.managers.mm_utils import TensorTransportMode, wrap_shm_features
@@ -1332,7 +1333,10 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
):
tokenized_obj.time_stats.set_api_server_dispatch_time()
tokenized_obj = wrap_shm_features(tokenized_obj)
time_stats = tokenized_obj.time_stats
tokenized_obj.wrap_pickle_fields()
self._dispatch_to_scheduler(tokenized_obj)
tokenized_obj.time_stats = time_stats
tokenized_obj.time_stats.set_api_server_dispatch_finish_time()
def _send_batch_request(
@@ -1342,13 +1346,19 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
],
):
"""Send a batch of tokenized requests as a single batched request to the scheduler."""
set_time_batch(tokenized_objs, "set_api_server_dispatch_time")
time_stats = [tokenized_obj.time_stats for tokenized_obj in tokenized_objs]
for tokenized_obj in tokenized_objs:
tokenized_obj.wrap_pickle_fields()
if isinstance(tokenized_objs[0], TokenizedGenerateReqInput):
batch_req = BatchTokenizedGenerateReqInput(batch=tokenized_objs)
else:
batch_req = BatchTokenizedEmbeddingReqInput(batch=tokenized_objs)
set_time_batch(tokenized_objs, "set_api_server_dispatch_time")
self._dispatch_to_scheduler(batch_req)
for tokenized_obj, time_stat in zip(tokenized_objs, time_stats):
tokenized_obj.time_stats = time_stat
set_time_batch(tokenized_objs, "set_api_server_dispatch_finish_time")
def _coalesce_streaming_chunks(
@@ -1856,6 +1866,11 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
BatchTokenIDOutput,
],
):
recv_obj.time_stats = unwrap_from_pickle(recv_obj.time_stats)
if isinstance(recv_obj, (BatchStrOutput, BatchTokenIDOutput)):
customized_info = unwrap_from_pickle(recv_obj.customized_info)
else:
customized_info = None
pending_notify: dict[str, ReqState] = {}
batch_notify_size = self.server_args.batch_notify_size
for i, rid in enumerate(recv_obj.rids):
@@ -1911,8 +1926,8 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
meta_info["cached_tokens_details"] = recv_obj.cached_tokens_details[
i
]
if recv_obj.customized_info is not None:
for k, v in recv_obj.customized_info.items():
if customized_info is not None:
for k, v in customized_info.items():
if k not in state.customized_info_accumulated:
state.customized_info_accumulated[k] = []
state.customized_info_accumulated[k].extend(v[i])
+13 -2
View File
@@ -15,7 +15,7 @@
import logging
import math
from typing import Any, Dict, List, Optional, Set, Union
from typing import Dict, List, Optional, Set, Union
import msgspec
@@ -25,6 +25,17 @@ try:
except ImportError:
import sre_parse # Python < 3.11
# JSON-safe value types for custom_params. Must survive msgpack IPC
# without PickleWrapper. After deserialization on the scheduler side,
# Req.__init__ injects "__req__" (a Req object) into the dict in-process;
# that augmented dict is never re-serialized.
_JsonScalar = Union[None, bool, int, float, str]
CustomParamValue = Union[
_JsonScalar,
List[_JsonScalar],
Dict[str, _JsonScalar],
]
_SAMPLING_EPS = 1e-6
TOP_K_ALL = 1 << 30
@@ -96,7 +107,7 @@ class SamplingParams(msgspec.Struct, kw_only=True, omit_defaults=True):
skip_special_tokens: bool = True
spaces_between_special_tokens: bool = True
no_stop_trim: bool = False
custom_params: Optional[Dict[str, Any]] = None
custom_params: Optional[Dict[str, CustomParamValue]] = None
stream_interval: Optional[int] = None
logit_bias: Optional[Dict[str, float]] = None
sampling_seed: Optional[int] = None
+108
View File
@@ -0,0 +1,108 @@
from __future__ import annotations
import base64
import binascii
from typing import Any
import msgspec
from pydantic_core import core_schema
class Base64Bytes:
"""Pydantic marker for HTTP JSON base64-encoded bytes fields."""
def __get_pydantic_core_schema__(self, source_type: Any, handler):
return core_schema.no_info_before_validator_function(
self._decode_value,
handler(source_type),
)
@classmethod
def _decode_value(cls, value: Any) -> Any:
if isinstance(value, str):
try:
return base64.b64decode(value, validate=True)
except binascii.Error as exc:
raise ValueError("Expected base64-encoded bytes") from exc
if isinstance(value, list):
return [cls._decode_value(item) for item in value]
if isinstance(value, tuple):
return tuple(cls._decode_value(item) for item in value)
return value
def msgspec_to_builtins(obj: Any) -> Any:
"""Recursively convert msgspec structs to dict/list Python builtins."""
if isinstance(obj, msgspec.Struct):
return {
field.name: msgspec_to_builtins(getattr(obj, field.name))
for field in msgspec.structs.fields(type(obj))
}
if isinstance(obj, dict):
return {key: msgspec_to_builtins(value) for key, value in obj.items()}
if isinstance(obj, list):
return [msgspec_to_builtins(item) for item in obj]
if isinstance(obj, tuple):
return tuple(msgspec_to_builtins(item) for item in obj)
if isinstance(obj, set):
return [msgspec_to_builtins(item) for item in obj]
return obj
def msgspec_struct_pydantic_core_schema(cls: type[msgspec.Struct], handler):
fields = {}
for struct_field in msgspec.structs.fields(cls):
field_schema = handler.generate_schema(struct_field.type)
required = (
struct_field.default is msgspec.NODEFAULT
and struct_field.default_factory is msgspec.NODEFAULT
)
if struct_field.default is not msgspec.NODEFAULT:
field_schema = core_schema.with_default_schema(
field_schema,
default=struct_field.default,
)
elif struct_field.default_factory is not msgspec.NODEFAULT:
field_schema = core_schema.with_default_schema(
field_schema,
default_factory=struct_field.default_factory,
)
fields[struct_field.name] = core_schema.typed_dict_field(
field_schema,
required=required,
)
typed_dict_schema = core_schema.typed_dict_schema(
fields,
cls_name=cls.__name__,
extra_behavior="ignore",
ref=cls.__name__,
)
def build_struct(value):
return value if isinstance(value, cls) else cls(**value)
dict_to_struct_schema = core_schema.no_info_after_validator_function(
build_struct,
typed_dict_schema,
)
return core_schema.json_or_python_schema(
json_schema=dict_to_struct_schema,
python_schema=core_schema.union_schema(
[
core_schema.is_instance_schema(cls),
dict_to_struct_schema,
],
mode="left_to_right",
),
)
@@ -12,7 +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.managers.io_struct import sock_recv, sock_send, wrap_as_pickle
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 (
@@ -90,7 +90,7 @@ class ScriptedHttpServer:
raise RuntimeError(f"ScriptedHttpServer is dirty: {self._dirty}")
fn_path = f"{script_fn.__module__}:{script_fn.__qualname__}"
sock_send(self._socket, RunScript(fn_path=fn_path, args=args))
sock_send(self._socket, wrap_as_pickle(RunScript(fn_path=fn_path, args=args)))
if not self._socket.poll(int(timeout_s * 1000)):
if not self._server_process.is_alive():
@@ -117,7 +117,7 @@ class ScriptedHttpServer:
fatal_error: Optional[OutOfBandError] = None
try:
try:
sock_send(self._socket, Shutdown())
sock_send(self._socket, wrap_as_pickle(Shutdown()))
except zmq.ZMQError:
pass
@@ -11,7 +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.managers.io_struct import sock_recv, sock_send, wrap_as_pickle
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
@@ -153,7 +153,7 @@ class ScriptedSchedulerHook:
socket = get_zmq_socket(ctx_zmq, zmq.PAIR, endpoint, bind=False)
try:
yield from _drive_engine_through_warmup(self._context)
sock_send(socket, HookReady())
sock_send(socket, wrap_as_pickle(HookReady()))
while True:
msg = sock_recv(socket)
match msg:
@@ -170,10 +170,12 @@ class ScriptedSchedulerHook:
except Exception:
sock_send(
socket,
ScriptFailed(traceback=traceback.format_exc()),
wrap_as_pickle(
ScriptFailed(traceback=traceback.format_exc())
),
)
else:
sock_send(socket, ScriptSucceeded())
sock_send(socket, wrap_as_pickle(ScriptSucceeded()))
case _:
raise ValueError(f"dispatch loop: unknown command {msg!r}")
finally:
@@ -41,6 +41,11 @@ class ScriptedTokenizerRecvProxy:
"ScriptedTokenizerRecvProxy.recv_pyobj: blocking recv is not supported"
)
def recv(self, flags: int = 0) -> bytes:
raise NotImplementedError(
"TODO: support ScriptedTokenizerRecvProxy.recv for msgpack IPC"
)
def wait_until_arrived(
self,
predicate: Callable[[Any], bool],