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],
+18 -4
View File
@@ -31,6 +31,8 @@ from sglang.test.test_utils import (
register_cuda_ci(est_time=134, stage="base-b", runner_config="1-gpu-small")
register_amd_ci(est_time=130, suite="stage-b-test-1-gpu-small-amd")
SERVER_ENV = {"SGLANG_USE_PICKLE_IPC": "0"}
class TestSRTEndpoint(CustomTestCase):
@classmethod
@@ -41,6 +43,7 @@ class TestSRTEndpoint(CustomTestCase):
cls.model,
cls.base_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
env=SERVER_ENV,
other_args=(
"--enable-custom-logit-processor",
"--mem-fraction-static",
@@ -469,6 +472,12 @@ class TestSRTEndpoint(CustomTestCase):
response = requests.post(self.base_url + "/flush_cache")
assert response.status_code == 200
server_info = requests.get(self.base_url + "/server_info").json()
page_size = server_info.get("page_size") or 1
def align_down(num_tokens):
return num_tokens // page_size * page_size
def send_and_check_cached_tokens(input_ids):
response = requests.post(
self.base_url + "/generate",
@@ -483,10 +492,14 @@ class TestSRTEndpoint(CustomTestCase):
return response_json["meta_info"]["cached_tokens"]
self.assertEqual(send_and_check_cached_tokens(range(0, 100)), 0)
self.assertEqual(send_and_check_cached_tokens(range(0, 10000)), 100)
self.assertEqual(send_and_check_cached_tokens(range(0, 10000)), 9999)
self.assertEqual(send_and_check_cached_tokens(range(0, 1000)), 999)
self.assertEqual(send_and_check_cached_tokens(range(0, 11000)), 10000)
self.assertEqual(send_and_check_cached_tokens(range(0, 10000)), align_down(100))
self.assertEqual(
send_and_check_cached_tokens(range(0, 10000)), align_down(9999)
)
self.assertEqual(send_and_check_cached_tokens(range(0, 1000)), align_down(999))
self.assertEqual(
send_and_check_cached_tokens(range(0, 11000)), align_down(10000)
)
def test_get_server_info(self):
response = requests.get(self.base_url + "/server_info")
@@ -648,6 +661,7 @@ class TestTokenizeDetokenize(CustomTestCase):
cls.model,
cls.base_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
env=SERVER_ENV,
)
cls.tokenizer = get_tokenizer(cls.model)
@@ -22,10 +22,12 @@ Current coverage:
import asyncio
import dataclasses
import json
import unittest
from types import SimpleNamespace
from sglang.srt.entrypoints import http_server
from sglang.srt.lora.lora_registry import LoRARef
from sglang.srt.server_args import ServerArgs
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
@@ -33,7 +35,9 @@ from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
def _call_server_info_with(server_args: ServerArgs) -> dict:
def _call_server_info_with(
server_args: ServerArgs, internal_states: list[dict] | None = None
) -> dict:
"""Invoke `http_server.server_info()` against a stub global state.
Bypasses the FastAPI HTTP layer (no TestClient): the handler is an
@@ -44,7 +48,7 @@ def _call_server_info_with(server_args: ServerArgs) -> dict:
"""
async def _fake_internal_state():
return [{"max_req_input_len": 1024}]
return internal_states or [{"max_req_input_len": 1024}]
stub_state = SimpleNamespace(
tokenizer_manager=SimpleNamespace(
@@ -291,6 +295,31 @@ class TestServerInfoExistingFieldsPreserved(CustomTestCase):
self.assertIn("kv_events", info)
self.assertIsNotNone(info["kv_events"])
def test_lora_refs_are_json_serializable_dicts(self):
lora_ref = LoRARef(
lora_id="lora-id",
lora_name="adapter",
lora_path="/tmp/adapter",
pinned=True,
)
args = ServerArgs(model_path="dummy")
args.lora_paths = [lora_ref]
info = _call_server_info_with(
args,
internal_states=[{"lora_paths": [lora_ref]}],
)
expected = {
"lora_id": "lora-id",
"lora_name": "adapter",
"lora_path": "/tmp/adapter",
"pinned": True,
}
self.assertEqual(info["lora_paths"], [expected])
self.assertEqual(info["internal_states"][0]["lora_paths"], [expected])
json.dumps(info)
if __name__ == "__main__":
unittest.main()
@@ -13,10 +13,11 @@ Covers:
"""
import asyncio
import dataclasses
import unittest
from unittest.mock import AsyncMock, MagicMock, Mock
import msgspec
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase, maybe_stub_sgl_kernel
@@ -35,7 +36,7 @@ _NOT_FINISHED = object() # Sentinel: request has not finished yet
# Categorised by value shape so that _make_batch_str_output can assign
# type-appropriate defaults without hardcoding every field name.
# When a field is renamed upstream, the old name simply won't appear in
# dataclasses.fields() and the new name will fall through to the
# msgspec.structs.fields() and the new name will fall through to the
# pattern-matching or safe fallback — no test breakage.
# ---------------------------------------------------------------------------
@@ -145,7 +146,7 @@ def _make_abort_req(rid: str, abort_message: str = "Aborted") -> AbortReq:
def _make_batch_str_output(rid: str, finished_reason=None) -> BatchStrOutput:
"""Create a minimal BatchStrOutput for a single request.
Uses dataclass field introspection so that new or renamed fields in
Uses struct field introspection so that new or renamed fields in
BatchStrOutput don't break this test. Only the fields that matter for
test logic (rids, finished_reasons, output_strs) are set explicitly;
all others receive type-appropriate defaults based on naming patterns.
@@ -159,7 +160,7 @@ def _make_batch_str_output(rid: str, finished_reason=None) -> BatchStrOutput:
fr = finished_reason
kwargs = {}
for f in dataclasses.fields(BatchStrOutput):
for f in msgspec.structs.fields(BatchStrOutput):
if f.name == "rids":
kwargs[f.name] = [rid]
elif f.name == "finished_reasons":
@@ -176,8 +177,8 @@ def _make_batch_str_output(rid: str, finished_reason=None) -> BatchStrOutput:
kwargs[f.name] = [None]
# Fields with class defaults — skip, let the default be used
elif (
f.default is not dataclasses.MISSING
or f.default_factory is not dataclasses.MISSING
f.default is not msgspec.NODEFAULT
or f.default_factory is not msgspec.NODEFAULT
):
continue
# Unknown required field — provide a safe per-request default.