From 5df60a21cd09785d6995e87e8d2558d419606397 Mon Sep 17 00:00:00 2001 From: Mick Date: Sat, 5 Sep 2026 21:22:37 +0800 Subject: [PATCH] fix(vlm): harden EPD receiver validation and liveness (#36945) --- .../srt/disaggregation/encoder/receiver.py | 477 ++++++++++++++---- python/sglang/srt/managers/io_struct.py | 7 + .../sglang/srt/managers/tokenizer_manager.py | 44 +- .../multimodal/processors/base_processor.py | 63 +++ .../srt/multimodal/processors/interns1pro.py | 2 +- .../disaggregation/test_encode_receiver.py | 342 ++++++++++++- .../test_kimi_k3_encoder_mode.py | 359 ++++++++++++- .../test_tokenizer_manager_rid_cleanup.py | 1 + .../unit/models/test_interns1pro_processor.py | 39 ++ test/registered/unit/models/test_kimi_k25.py | 2 +- .../multimodal/test_gpu_feature_transport.py | 3 + .../test_precomputed_embedding_validation.py | 90 ++++ 12 files changed, 1320 insertions(+), 109 deletions(-) create mode 100644 test/registered/unit/models/test_interns1pro_processor.py create mode 100644 test/registered/unit/multimodal/test_precomputed_embedding_validation.py diff --git a/python/sglang/srt/disaggregation/encoder/receiver.py b/python/sglang/srt/disaggregation/encoder/receiver.py index a1ef1d6ce..257d7c724 100644 --- a/python/sglang/srt/disaggregation/encoder/receiver.py +++ b/python/sglang/srt/disaggregation/encoder/receiver.py @@ -30,7 +30,11 @@ from sglang.srt.distributed.parallel_state import ( get_mooncake_transfer_engine, ) from sglang.srt.environ import envs -from sglang.srt.managers.io_struct import GenerateReqInput, TokenizedGenerateReqInput +from sglang.srt.managers.io_struct import ( + EncoderDispatchErrorReq, + GenerateReqInput, + TokenizedGenerateReqInput, +) from sglang.srt.managers.multimodal_processor import get_mm_processor, import_processors from sglang.srt.managers.schedule_batch import Modality, Req from sglang.srt.multimodal.cache import media_preprocess_kwargs @@ -62,6 +66,22 @@ if TYPE_CHECKING: from sglang.srt.managers.scheduler import Scheduler +class _ReceiveRegistrationRunner: + """Run encoder receive-URL registration off the scheduler thread.""" + + def __init__(self, name: str): + self.loop = asyncio.new_event_loop() + self.thread = threading.Thread(target=self._run, daemon=True, name=name) + self.thread.start() + + def _run(self) -> None: + asyncio.set_event_loop(self.loop) + self.loop.run_forever() + + def submit(self, coroutine): + return asyncio.run_coroutine_threadsafe(coroutine, self.loop) + + def _mark_keep_device_embedding(mm_inputs) -> None: """Tell general_mm_embed_routine not to copy embeddings back to CPU.""" if mm_inputs is None: @@ -356,15 +376,15 @@ def _normalize_embedding_ports(embedding_port): return [embedding_port] -def _grpc_scheduler_receive_url(target, req_id, receive_url, receive_count): +async def _grpc_scheduler_receive_url(target, req_id, receive_url, receive_count): import grpc from smg_grpc_proto import sglang_encoder_pb2, sglang_encoder_pb2_grpc timeout_secs = envs.SGLANG_ENCODER_GRPC_TIMEOUT_SECS.get() - channel = grpc.insecure_channel(target) + channel = grpc.aio.insecure_channel(target) stub = sglang_encoder_pb2_grpc.SglangEncoderStub(channel) try: - stub.SchedulerReceiveUrl( + await stub.SchedulerReceiveUrl( sglang_encoder_pb2.SchedulerReceiveUrlRequest( req_id=req_id, receive_url=receive_url, @@ -373,7 +393,7 @@ def _grpc_scheduler_receive_url(target, req_id, receive_url, receive_count): timeout=timeout_secs, ) finally: - channel.close() + await channel.close() def _grpc_encode_request(target, encode_request): @@ -402,6 +422,24 @@ def _grpc_encode_request(target, encode_request): channel.close() +async def _gather_blocking_grpc_calls(calls): + """Wait for synchronous gRPC threads to stop before propagating cancellation.""" + future = asyncio.gather(*calls, return_exceptions=True) + try: + results = await asyncio.shield(future) + except asyncio.CancelledError: + results = await asyncio.shield(future) + for result in results: + if isinstance(result, Exception): + logger.error("gRPC call failed while draining cancellation: %s", result) + raise + + for result in results: + if isinstance(result, Exception): + raise result + return results + + class EmbeddingData: def __init__( self, @@ -610,6 +648,7 @@ class MultiModalEmbeddingData(EmbeddingData): model_type: Optional[str] = None, ): """Create MultiModalEmbeddingData from an EmbeddingData instance.""" + _validate_embedding_part(embedding_data) # Only forward known optional attrs (e.g. video metadata) so they land on the instance extra = {} for attr in video_meta_attrs_for(model_type): @@ -677,13 +716,7 @@ class MultiModalEmbeddingData(EmbeddingData): return kwargs def add(self, embedding_data: EmbeddingData): - if self.req_id != embedding_data.req_id: - logger.warning( - f"Dropping embedding data with mismatched req_id: " - f"expected {self.req_id}, got {embedding_data.req_id}" - ) - return False - assert not self.ready_list[embedding_data.part_idx] + _validate_embedding_part(embedding_data, current=self) pid = embedding_data.part_idx self.ready_list[pid] = True self.modality_list[pid] = embedding_data.modality @@ -696,6 +729,44 @@ class MultiModalEmbeddingData(EmbeddingData): self._set_image_meta_for_part(pid, embedding_data) +def _validate_embedding_part( + embedding_data: EmbeddingData, + current: Optional[MultiModalEmbeddingData] = None, +) -> None: + """Reject malformed part metadata before indexing aggregation buffers.""" + if not isinstance(embedding_data, EmbeddingData): + raise ValueError(f"expected EmbeddingData, got {type(embedding_data).__name__}") + if ( + not isinstance(embedding_data.num_parts, int) + or isinstance(embedding_data.num_parts, bool) + or embedding_data.num_parts <= 0 + ): + raise ValueError("num_parts must be a positive integer") + if ( + not isinstance(embedding_data.part_idx, int) + or isinstance(embedding_data.part_idx, bool) + or embedding_data.part_idx < 0 + or embedding_data.part_idx >= embedding_data.num_parts + ): + raise ValueError( + f"part_idx must be in [0, {embedding_data.num_parts}), " + f"got {embedding_data.part_idx}" + ) + if current is None: + return + if current.req_id != embedding_data.req_id: + raise ValueError( + f"embedding req_id mismatch: expected {current.req_id}, " + f"got {embedding_data.req_id}" + ) + if current.num_parts != embedding_data.num_parts: + raise ValueError( + f"num_parts changed from {current.num_parts} to {embedding_data.num_parts}" + ) + if current.ready_list[embedding_data.part_idx]: + raise ValueError(f"duplicate embedding part {embedding_data.part_idx}") + + def _aggregate_embedding_part(current, recv_obj, model_type): """Fold one received part into the aggregate (the first part creates it).""" if current is None: @@ -732,6 +803,41 @@ def extract_original_req_id(part_req_id: str) -> str: return part_req_id +def _resolve_embedding_part_request_id( + embedding_data: object, expected_req_id: Optional[str] = None +) -> Optional[str]: + """Validate and normalize an embedding part ID for safe routing.""" + expected = ( + f" for expected rid={expected_req_id}" if expected_req_id is not None else "" + ) + if not isinstance(embedding_data, EmbeddingData): + logger.warning("Dropping non-embedding data%s", expected) + return None + if not isinstance(embedding_data.req_id, str): + logger.warning("Dropping embedding data with a non-string req_id%s", expected) + return None + original_req_id = extract_original_req_id(embedding_data.req_id) + if expected_req_id is not None and original_req_id != expected_req_id: + logger.warning( + "Dropping stale embedding data: expected rid=%s, got rid=%s " + "(likely from ZMQ port reuse)", + expected_req_id, + embedding_data.req_id, + ) + return None + embedding_data.req_id = original_req_id + return original_req_id + + +def _embedding_part_matches_request( + embedding_data: object, expected_req_id: str +) -> bool: + """Normalize a matching part ID; reject stale data from a reused socket.""" + return ( + _resolve_embedding_part_request_id(embedding_data, expected_req_id) is not None + ) + + def _encoder_media_item(mm_item: dict): """Keep per-media options aligned while preserving the legacy URL shape.""" item = { @@ -784,6 +890,7 @@ class WaitingMMRequestBase(ABC): embedding_pool: Optional["EmbeddingPool"] = None, zmq_context=None, embedding_port=None, + registration_runner: Optional[_ReceiveRegistrationRunner] = None, ): self.rid = rid self.recv_req = recv_req @@ -820,6 +927,10 @@ class WaitingMMRequestBase(ABC): # Success-path finalizer handle so abort can release the slot early. self._mm_finalizer: Optional[weakref.finalize] = None self._pool_full_warned = False + self.registration_runner = registration_runner + self.registration_future = None + self.registration_error = None + self.registration_lock = threading.Lock() @abstractmethod def send_encode_request(self) -> None: @@ -829,6 +940,15 @@ class WaitingMMRequestBase(ABC): if self.status != WaitingMMRequestStatus.PENDING: return + with self.registration_lock: + registration_error, self.registration_error = ( + self.registration_error, + None, + ) + if registration_error is not None: + self._fail_and_release(*registration_error) + return + # A complete request can remain pending while the GPU pool is full. # Retry assembly on every scheduler tick, including shared-socket mode. if self.recv_embedding_data is not None and self.recv_embedding_data.ready: @@ -860,6 +980,8 @@ class WaitingMMRequestBase(ABC): try: recv_obj: EmbeddingData = safe_pickle_loads(parts[0]) + if not self._is_valid_embedding_part(recv_obj): + return if getattr(recv_obj, "error_msg", None) is not None: logger.warning( f"Received error signal from encoder for {self.rid}: " @@ -867,8 +989,6 @@ class WaitingMMRequestBase(ABC): ) self._fail_and_release(recv_obj.error_msg, recv_obj.error_code) return - if not self._is_valid_embedding_part(recv_obj): - return # ZMQ materializes frame 1; RDMA already wrote the registered buffer. self._extract_embedding_from_buffer(recv_obj, parts) self.recv_embedding_data = _aggregate_embedding_part( @@ -896,29 +1016,12 @@ class WaitingMMRequestBase(ABC): self.error_msg = error_msg self.error_code = error_code self.status = WaitingMMRequestStatus.FAIL - self._cleanup_gpu_buffer() + self.release_resources() self.close_recv_socket() - async def _check_encoder_responses(self, responses, endpoint: str) -> bool: - """Validate gathered encoder responses; on the first error, FAIL the - request and release its resources. Returns True if all succeeded.""" - msg = await _extract_encoder_error(responses, endpoint, f"rid={self.rid}") - if msg is None: - return True - self._fail_and_release(msg) - return False - def _is_valid_embedding_part(self, recv_obj) -> bool: """Check for and drop stale or out-of-sync payloads; normalize the part req_id to the original rid.""" - original_req_id = extract_original_req_id(recv_obj.req_id) - if original_req_id != self.recv_req.rid: - logger.warning( - f"Dropping stale embedding data: expected rid={self.recv_req.rid}, " - f"got rid={recv_obj.req_id} (likely from ZMQ port reuse)" - ) - return False - recv_obj.req_id = original_req_id - return True + return _embedding_part_matches_request(recv_obj, self.recv_req.rid) @abstractmethod def _extract_embedding_from_buffer(self, recv_obj, parts) -> None: @@ -964,8 +1067,8 @@ class WaitingMMRequestBase(ABC): return True def _finish_assemble(self, recv_embedding) -> None: - """get_mm_data → bind pool slot → publish onto recv_req → SUCCESS.""" - mm_inputs = self.mm_processor.get_mm_data( + """Build validated mm data, bind its pool slot, then publish it.""" + mm_inputs = self.mm_processor.get_validated_mm_data( _select_mm_processor_prompt(self.recv_req, self.mm_processor), recv_embedding, **self.recv_embedding_data.get_mm_extra_meta(), @@ -1005,6 +1108,12 @@ class WaitingMMRequestBase(ABC): def release_resources(self): """Free pool/GPU resources on abort/fail/timeout. Idempotent.""" + registration_future, self.registration_future = ( + self.registration_future, + None, + ) + if registration_future is not None and not registration_future.done(): + registration_future.cancel() self._cleanup_gpu_buffer() finalizer, self._mm_finalizer = self._mm_finalizer, None if finalizer is not None: @@ -1014,8 +1123,37 @@ class WaitingMMRequestBase(ABC): # For zmq_to_scheduler: embedding parts arrive as ZMQ payload frames and # are optionally staged into the GPU EmbeddingPool. class WaitingZmqRequest(WaitingMMRequestBase): - def send_encode_request(self): + def _start_registration(self, coroutine) -> None: + if self.registration_runner is None: + coroutine.close() + self._fail_and_release( + "Encoder receive registration runner is unavailable", + int(HTTPStatus.INTERNAL_SERVER_ERROR), + ) + return + self.registration_future = self.registration_runner.submit(coroutine) + self.registration_future.add_done_callback(self._on_registration_done) + + def _on_registration_done(self, future) -> None: + if future.cancelled(): + return + error = future.exception() + if error is None: + return + logger.error( + "Failed to register encoder receive URL for rid=%s: %s", + self.rid, + error, + exc_info=error, + ) + with self.registration_lock: + self.registration_error = ( + f"Failed to register receive URL with encoder: {error}", + int(HTTPStatus.BAD_GATEWAY), + ) + + def send_encode_request(self): async def _send_single_request(session, url, payload): try: async with session.post(url, json=payload) as response: @@ -1053,7 +1191,7 @@ class WaitingZmqRequest(WaitingMMRequestBase): encoder_url = self.encoder_urls[idx] target_url = f"{encoder_url}/scheduler_receive_url" payload = { - "req_id": part_req_id, # use part_req_id to match encode request + "req_id": part_req_id, "receive_count": receive_count, "receive_url": NetworkAddress( host_name, embedding_port @@ -1091,15 +1229,9 @@ class WaitingZmqRequest(WaitingMMRequestBase): logger.debug(f"Request {i} succeeded.") failed = [r for r in results if isinstance(r, BaseException)] if failed: - # A rank without a registered receive URL can never be - # pushed to; fail via the normal completion path now - # instead of pending until the embedding wait times out. - self._fail_and_release( - f"Failed to register receive URL with encoder: {failed[0]!r}", - int(HTTPStatus.BAD_GATEWAY), - ) + raise failed[0] - asyncio.run( + self._start_registration( send_embedding_port( self.recv_req.rid, self.receive_count, @@ -1180,8 +1312,7 @@ class WaitingZmqRequestGrpc(WaitingZmqRequest): target_url = f"{encoder_url}/SchedulerReceiveUrl" logger.info(f"Preparing to send to {target_url}") tasks.append( - asyncio.to_thread( - _grpc_scheduler_receive_url, + _grpc_scheduler_receive_url( _grpc_target(encoder_url), req_id, receive_url, @@ -1200,8 +1331,11 @@ class WaitingZmqRequestGrpc(WaitingZmqRequest): logger.error(f"Request {i} failed: {result}") else: logger.debug(f"Request {i} succeeded.") + failed = [r for r in results if isinstance(r, BaseException)] + if failed: + raise failed[0] - asyncio.run( + self._start_registration( send_embedding_port( self.recv_req.rid, self.receive_count, @@ -1249,6 +1383,8 @@ class WaitingRDMARequest(WaitingMMRequestBase): self._buffer_lock = threading.Lock() self._terminal = False self._receive_running = False + self._receive_error = None + self._receive_error_lock = threading.Lock() def send_encode_request(self): # Base-class hook. The tokenizer owns /encode, so this rank only pulls @@ -1261,13 +1397,34 @@ class WaitingRDMARequest(WaitingMMRequestBase): asyncio.run(self._pull_meta_and_receive_embedding()) except Exception as e: logger.error(f"RDMA receive failed for rid={self.rid}: {e}") - self._fail_and_release(str(e)) + self._record_receive_error(str(e)) finally: with self._buffer_lock: self._receive_running = False if self._terminal: self._release_buffer_locked() + def _record_receive_error(self, error_msg, error_code=None) -> None: + """Pass a worker-thread failure to the scheduler thread.""" + with self._receive_error_lock: + self._receive_error = (error_msg, error_code) + + def _try_recv_mm_data(self): + with self._receive_error_lock: + receive_error, self._receive_error = self._receive_error, None + if receive_error is not None: + self._fail_and_release(*receive_error) + return + super()._try_recv_mm_data() + + async def _check_encoder_responses(self, responses, endpoint: str) -> bool: + """Record network failures for the scheduler thread to consume.""" + error = await _extract_encoder_error(responses, endpoint, f"rid={self.rid}") + if error is None: + return True + self._record_receive_error(*error) + return False + async def _pull_meta_and_receive_embedding(self): """Pull per-part sizes, allocate the landing buffer, then drive /send. @@ -1334,7 +1491,7 @@ class WaitingRDMARequest(WaitingMMRequestBase): ) if alloc_result is None: # Oversize or alloc timeout — fatal for this request. - self._fail_and_release( + self._record_receive_error( f"EmbeddingPool could not allocate " f"{total_bytes // (1024 * 1024)}MB (oversize or " f"timeout). Raise SGLANG_EMBEDDING_POOL_SIZE_MB." @@ -1449,7 +1606,7 @@ class WaitingRDMARequest(WaitingMMRequestBase): async def _extract_encoder_error(responses, endpoint, context, encode_requests=None): - """Return the first error among gathered encoder responses, or None. + """Return the first ``(message, status)`` error, or None. Pure check — logs each error but has no other side effects; the caller decides how to react. ``encode_requests`` optionally enriches each log @@ -1467,13 +1624,16 @@ async def _extract_encoder_error(responses, endpoint, context, encode_requests=N logger.error( f"Encoder {endpoint} timeout ({timeout_val}s) for {ctx} (request {i})" ) - return f"Encoder {endpoint} timeout ({timeout_val}s)" + return ( + f"Encoder {endpoint} timeout ({timeout_val}s)", + int(HTTPStatus.GATEWAY_TIMEOUT), + ) if isinstance(resp, Exception): logger.error( f"Encoder {endpoint} failed for {ctx} (request {i}): {resp}", exc_info=resp, ) - return str(resp) + return str(resp), int(HTTPStatus.BAD_GATEWAY) if resp.status != 200: try: err = await resp.json() @@ -1481,7 +1641,7 @@ async def _extract_encoder_error(responses, endpoint, context, encode_requests=N except Exception: msg = await resp.text() logger.error(f"Encoder {endpoint} returned error {resp.status}: {msg}") - return msg + return msg, int(resp.status) return None @@ -1739,12 +1899,16 @@ class MMReceiverBase(ABC): self.hostname = get_local_ip_auto() self.waiting_list: List[WaitingMMRequestBase] = [] self.waiting_by_rid: Dict[str, WaitingMMRequestBase] = {} + self.registration_runner = None self.scheduler_embedding_port = None self.scheduler_recv_socket = None if ( self.encoder_transfer_backend == "zmq_to_scheduler" and scheduler is not None ): + self.registration_runner = _ReceiveRegistrationRunner( + f"encoder-receive-registration-{tp_rank}" + ) ( self.scheduler_embedding_port, self.scheduler_recv_socket, @@ -1887,6 +2051,10 @@ class MMReceiverBase(ABC): self, request_obj, mm_processor, prompt, need_wait_for_mm_inputs=True ): req_id = None + recv_socket = None + encode_task = None + recv_task = None + send_time = time.monotonic() try: # ``self.encode_urls`` is shared by reference with the bootstrap # server (when running) so it always reflects the current set. @@ -1931,13 +2099,13 @@ class MMReceiverBase(ABC): done and recv_task not in done and ( - encode_task.exception() is not None or encode_task.result() is False + encode_task.exception() is not None + or encode_task.result() is not None ) ): logger.warning( f"[{req_id}] Encoder dispatch failed; skipping embedding wait" ) - recv_task.cancel() return None result = await asyncio.wait_for( recv_task, @@ -1950,6 +2118,15 @@ class MMReceiverBase(ABC): elapsed = time.monotonic() - send_time logger.warning(f"[{req_id}] Embedding recv timeout after {elapsed:.3f}s") return None + finally: + tasks = [task for task in (encode_task, recv_task) if task is not None] + for task in tasks: + if not task.done(): + task.cancel() + if tasks: + await asyncio.gather(*tasks, return_exceptions=True) + if recv_socket is not None: + recv_socket.close(linger=0) async def _recv_mm_data(self, req_id, recv_socket, mm_processor, prompt): """zmq_to_tokenizer receive: embedding parts arrive as 2-frame ZMQ @@ -1966,6 +2143,8 @@ class MMReceiverBase(ABC): if not parts: continue recv_obj: EmbeddingData = safe_pickle_loads(parts[0]) + if not _embedding_part_matches_request(recv_obj, req_id): + continue if getattr(recv_obj, "error_msg", None) is not None: logger.warning( f"Encoder error for req_id={req_id}: {recv_obj.error_msg} " @@ -1973,8 +2152,6 @@ class MMReceiverBase(ABC): ) return None logger.debug("recv_obj=%s", recv_obj) - # Normalize the part req_id to the original for aggregation. - recv_obj.req_id = extract_original_req_id(recv_obj.req_id) if len(parts) < 2: logger.error( "zmq_to_tokenizer expected 2-part message, got %d parts", @@ -1993,18 +2170,31 @@ class MMReceiverBase(ABC): ) recv_embedding = recv_embedding_data.get_embedding(is_concat=True) - return mm_processor.get_mm_data( + return mm_processor.get_validated_mm_data( prompt, recv_embedding, **recv_embedding_data.get_mm_extra_meta(), ) + except Exception: + logger.exception( + "Failed to receive encoder embeddings for req_id=%s", req_id + ) + return None finally: recv_socket.close() - def send_encode_request(self, obj, time_stats_json=None): - self._send_encode_request(obj, time_stats_json=time_stats_json) + def send_encode_request( + self, obj, time_stats_json=None, on_dispatch_error=None + ) -> Optional[threading.Event]: + return self._send_encode_request( + obj, + time_stats_json=time_stats_json, + on_dispatch_error=on_dispatch_error, + ) - def _send_encode_request(self, obj, time_stats_json=None): + def _send_encode_request( + self, obj, time_stats_json=None, on_dispatch_error=None + ) -> Optional[threading.Event]: mm_data = self._extract_url_data(obj) if obj.rid is None: obj.rid = uuid.uuid4().hex @@ -2028,6 +2218,7 @@ class MMReceiverBase(ABC): # Freeze the encoder URL snapshot onto obj so the scheduler # subprocess uses the same list when indexing encoder_idx. obj.encoder_urls = encode_urls + scheduler_dispatch_ready = threading.Event() encode_thread = threading.Thread( target=self._run_encode_in_thread, @@ -2038,10 +2229,13 @@ class MMReceiverBase(ABC): num_items_assigned, encode_urls, time_stats_json, + scheduler_dispatch_ready, + on_dispatch_error, ), daemon=True, ) encode_thread.start() + return scheduler_dispatch_ready else: # No encoder URLs available (bootstrap may not have any registered yet); # reset the flag so the scheduler does not wait for embeddings that will @@ -2053,6 +2247,7 @@ class MMReceiverBase(ABC): "processing without encoder disaggregation." ) obj.need_wait_for_mm_inputs = False + return None def _sync_fail_info_across_tp(self, waiting_req: WaitingMMRequestBase) -> None: """Share encoder error fields across TP ranks before abort. @@ -2092,8 +2287,18 @@ class MMReceiverBase(ABC): except zmq.Again: return - recv_obj: EmbeddingData = safe_pickle_loads(parts[0]) - rid = extract_original_req_id(recv_obj.req_id) + try: + recv_obj: EmbeddingData = safe_pickle_loads(parts[0]) + except Exception as error: + logger.warning( + "Dropping malformed embedding data from the shared " + "scheduler socket: %s", + error, + ) + continue + rid = _resolve_embedding_part_request_id(recv_obj) + if rid is None: + continue waiting_req = self.waiting_by_rid.get(rid) if waiting_req is None: logger.warning( @@ -2104,7 +2309,21 @@ class MMReceiverBase(ABC): def _process_waiting_requests(self, recv_reqs, waiting_cls, **extra_kwargs): new_recv_reqs = [] + abort_reqs = [] for recv_req in recv_reqs: + if isinstance(recv_req, EncoderDispatchErrorReq): + waiting_req = self.waiting_by_rid.get(recv_req.rid) + if waiting_req is None: + logger.debug( + "Ignoring encoder dispatch error for inactive request %s", + recv_req.rid, + ) + else: + waiting_req._fail_and_release( + recv_req.error_msg, recv_req.error_code + ) + continue + if ( isinstance(recv_req, TokenizedGenerateReqInput) and recv_req.need_wait_for_mm_inputs is True @@ -2116,31 +2335,81 @@ class MMReceiverBase(ABC): # tokenizer never set encoder_urls (legacy / static path). encode_urls = recv_req.encoder_urls or list(self.encode_urls) - waiting_req = waiting_cls( - rid=recv_req.rid, - recv_req=recv_req, - mm_processor=self.mm_processor, - encoder_urls=encode_urls, - model_type=self.model_type, - host_name=self.hostname, - receive_count=self.tp_size, - zmq_context=( - None - if self.scheduler_recv_socket is not None - else self.scheduler_context - ), - embedding_port=self.scheduler_embedding_port, - **extra_kwargs, - ) - if self.scheduler_recv_socket is not None: + waiting_req = None + local_error = None + try: + waiting_req = waiting_cls( + rid=recv_req.rid, + recv_req=recv_req, + mm_processor=self.mm_processor, + encoder_urls=encode_urls, + model_type=self.model_type, + host_name=self.hostname, + receive_count=self.tp_size, + zmq_context=( + None + if self.scheduler_recv_socket is not None + else self.scheduler_context + ), + embedding_port=self.scheduler_embedding_port, + **extra_kwargs, + ) self.waiting_by_rid[waiting_req.rid] = waiting_req - waiting_req.send_encode_request() + waiting_req.send_encode_request() + except Exception as error: + local_error = f"{type(error).__name__}: {error}" + logger.exception( + "Failed to start multimodal receive for rid=%s", recv_req.rid + ) + + # The status all-reduce below requires every TP rank to append + # exactly the same requests. Agree on startup before appending. + rank_errors = ( + [local_error] + if self.tp_size <= 1 + else self.tp_group.all_gather_object(local_error) + ) + failed_ranks = [ + rank for rank, error in enumerate(rank_errors) if error is not None + ] + if failed_ranks: + details = "; ".join( + f"rank {rank}: {rank_errors[rank]}" for rank in failed_ranks + ) + error_msg = f"Failed to start multimodal receive ({details})" + logger.error(error_msg) + if waiting_req is not None: + try: + waiting_req.release_resources() + except Exception: + logger.exception( + "Failed to release multimodal receive resources " + "for rid=%s", + waiting_req.rid, + ) + try: + waiting_req.close_recv_socket() + except Exception: + logger.exception( + "Failed to close multimodal receive socket for rid=%s", + waiting_req.rid, + ) + self.waiting_by_rid.pop(waiting_req.rid, None) + abort_reqs.append( + ( + self.create_req(recv_req), + error_msg, + HTTPStatus.INTERNAL_SERVER_ERROR, + ) + ) + continue + self.waiting_list.append(waiting_req) else: new_recv_reqs.append(recv_req) if len(self.waiting_list) == 0: - return new_recv_reqs, [] + return new_recv_reqs, abort_reqs self._drain_scheduler_embeddings() current_time = time.time() @@ -2164,7 +2433,6 @@ class MMReceiverBase(ABC): ) new_waiting = [] - abort_reqs = [] for i, waiting_req in enumerate(self.waiting_list): status_value = local_status[i].item() if status_value == WaitingMMRequestStatus.SUCCESS: @@ -2199,6 +2467,7 @@ class MMReceiverBase(ABC): else: # status_value == WaitingMMRequestStatus.PENDING new_waiting.append(waiting_req) continue + waiting_req.close_recv_socket() self.waiting_by_rid.pop(waiting_req.rid, None) self.waiting_list = new_waiting @@ -2212,12 +2481,14 @@ class MMReceiverBase(ABC): num_items_assigned, encode_urls=None, time_stats_json=None, + scheduler_dispatch_ready=None, + on_dispatch_error=None, ): # ``embedding_port`` is always None on this path: zmq_to_scheduler / # mooncake ranks register their receive ports with the encoder later # via /scheduler_receive_url, so the dispatch itself carries no port. try: - asyncio.run( + dispatch_error = asyncio.run( self.encode( req_id=req_id, mm_data=mm_data, @@ -2230,6 +2501,15 @@ class MMReceiverBase(ABC): ) except Exception as e: logger.error(f"Encode failed for request {req_id}: {e}", exc_info=True) + dispatch_error = EncoderDispatchErrorReq( + rid=req_id, + error_msg=str(e), + error_code=int(HTTPStatus.BAD_GATEWAY), + ) + + if dispatch_error is not None and on_dispatch_error is not None: + scheduler_dispatch_ready.wait() + on_dispatch_error(dispatch_error) def create_req(self, recv_req: TokenizedGenerateReqInput): req = Req( @@ -2422,7 +2702,10 @@ class MMReceiverHTTP(MMReceiverBase): embedding_pool=self.embedding_pool, ) return self._process_waiting_requests( - recv_reqs, WaitingZmqRequest, embedding_pool=self.embedding_pool + recv_reqs, + WaitingZmqRequest, + embedding_pool=self.embedding_pool, + registration_runner=self.registration_runner, ) async def encode( @@ -2511,11 +2794,15 @@ class MMReceiverHTTP(MMReceiverBase): # zmq_to_tokenizer is pushed to our PULL socket during /encode, # zmq_to_scheduler to the ports its ranks registered, and mooncake # by RDMA once those ranks have pulled sizes and driven /send. - return ( - await _extract_encoder_error( - responses, "HTTP request", f"req_id={req_id}", encode_requests - ) - is None + error = await _extract_encoder_error( + responses, "HTTP request", f"req_id={req_id}", encode_requests + ) + if error is None: + return None + return EncoderDispatchErrorReq( + rid=req_id, + error_msg=error[0], + error_code=error[1], ) @@ -2551,7 +2838,11 @@ class MMReceiverGrpc(MMReceiverBase): # For zmq_to_scheduler def process_waiting_requests(self, recv_reqs): - return self._process_waiting_requests(recv_reqs, WaitingZmqRequestGrpc) + return self._process_waiting_requests( + recv_reqs, + WaitingZmqRequestGrpc, + registration_runner=self.registration_runner, + ) async def encode( self, @@ -2621,7 +2912,7 @@ class MMReceiverGrpc(MMReceiverBase): ) for encode_request in encode_requests ] - await asyncio.gather(*grpc_tasks) + await _gather_blocking_grpc_calls(grpc_tasks) def _validate_transport_mode(transport_mode: str, encoder_urls): diff --git a/python/sglang/srt/managers/io_struct.py b/python/sglang/srt/managers/io_struct.py index d29ff305b..764205f08 100644 --- a/python/sglang/srt/managers/io_struct.py +++ b/python/sglang/srt/managers/io_struct.py @@ -2053,6 +2053,13 @@ class AbortReq(BaseReq, kw_only=True): self.rid = "" +class EncoderDispatchErrorReq(BaseReq, kw_only=True): + """Tokenizer-to-scheduler failure for one EPD encoder dispatch.""" + + error_msg: str + error_code: int + + class ActiveRanksOutput(BaseReq, kw_only=True): status: List[bool] diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index e9d98ecb6..0dc1727a8 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -72,6 +72,7 @@ from sglang.srt.managers.io_struct import ( ContinueGenerationReqInput, ElasticScaleUpdateReq, EmbeddingReqInput, + EncoderDispatchErrorReq, FreezeGCReq, GenerateReqInput, HealthCheckOutput, @@ -588,6 +589,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): def init_running_status(self): # Request states self.rid_to_state: Dict[str, ReqState] = {} + self.encoder_dispatch_ready: Dict[str, threading.Event] = {} self.event_loop = None self.asyncio_tasks = set() @@ -1579,6 +1581,9 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): self._dispatch_to_scheduler(tokenized_obj) self._mark_state_dispatched(tokenized_obj.rid) dispatched = True + dispatch_ready = self.encoder_dispatch_ready.pop(tokenized_obj.rid, None) + if dispatch_ready is not None: + dispatch_ready.set() tokenized_obj.time_stats = time_stats tokenized_obj.time_stats.set_api_server_dispatch_finish_time() finally: @@ -3495,15 +3500,28 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): """ for rid in rids: state = self.rid_to_state.get(rid) - if state is None: - continue - if state.dispatched: - try: - self.abort_request(rid) - except Exception: - logger.exception("Failed to abort request %s during cleanup", rid) - else: - del self.rid_to_state[rid] + if state is not None: + if state.dispatched: + try: + self.abort_request(rid) + except Exception: + logger.exception( + "Failed to abort request %s during cleanup", rid + ) + else: + del self.rid_to_state[rid] + dispatch_ready = self.encoder_dispatch_ready.pop(rid, None) + if dispatch_ready is not None: + dispatch_ready.set() + + def _forward_encoder_dispatch_error(self, error: EncoderDispatchErrorReq) -> None: + if error.rid in self.rid_to_state: + self._dispatch_to_scheduler(error) + + def _schedule_encoder_dispatch_error(self, error: EncoderDispatchErrorReq) -> None: + self.event_loop.call_soon_threadsafe( + self._forward_encoder_dispatch_error, error + ) def _should_dispatch_to_encoder( self, obj: Union[GenerateReqInput, EmbeddingReqInput] @@ -3554,9 +3572,13 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): if state is not None: time_stats_json = state.time_stats.encode_json() - self.mm_receiver.send_encode_request( - obj, time_stats_json=time_stats_json + dispatch_ready = self.mm_receiver.send_encode_request( + obj, + time_stats_json=time_stats_json, + on_dispatch_error=self._schedule_encoder_dispatch_error, ) + if dispatch_ready is not None: + self.encoder_dispatch_ready[obj.rid] = dispatch_ready else: obj.need_wait_for_mm_inputs = False diff --git a/python/sglang/srt/multimodal/processors/base_processor.py b/python/sglang/srt/multimodal/processors/base_processor.py index 60a9e94f6..0deffbfe9 100644 --- a/python/sglang/srt/multimodal/processors/base_processor.py +++ b/python/sglang/srt/multimodal/processors/base_processor.py @@ -647,6 +647,69 @@ class BaseMultimodalProcessor(ABC): video_token_id=getattr(self, "VIDEO_TOKEN_ID", None), ) + def get_validated_mm_data( + self, + prompt, + embeddings: Dict[Modality, torch.Tensor], + **kwargs, + ) -> MultimodalProcessorOutput: + """Build EPD multimodal inputs and validate the embedding layout. + + Model processors may override ``get_mm_data`` to rebuild their prompt + layout. This shared wrapper ensures every override consumes exactly the + encoder rows it received before the result reaches the scheduler. + """ + output = self.get_mm_data(prompt, embeddings, **kwargs) + self._validate_precomputed_embedding_layout(output, embeddings) + return output + + @staticmethod + def _validate_precomputed_embedding_layout( + output: MultimodalProcessorOutput, + embeddings: Dict[Modality, torch.Tensor], + ) -> None: + consumed_per_modality = {modality: 0 for modality in embeddings} + + for item in output.mm_items: + embedding = item.precomputed_embeddings + if not isinstance(embedding, torch.Tensor): + raise RuntimeError( + "EPD multimodal items must contain tensor embeddings; " + f"got {type(embedding).__name__} for " + f"{item.modality.name.lower()}" + ) + + num_rows = embedding.shape[0] + if item.offsets is not None: + expected_rows = sum(end - start + 1 for start, end in item.offsets) + if num_rows != expected_rows: + raise RuntimeError( + "Precomputed multimodal embedding length mismatch for " + f"{item.modality.name.lower()}: expected {expected_rows} " + f"rows from prompt offsets, got {num_rows}" + ) + + if item.modality not in consumed_per_modality: + raise RuntimeError( + "EPD processor returned an unexpected embedding modality: " + f"{item.modality.name.lower()}" + ) + consumed_per_modality[item.modality] += num_rows + + for modality, embedding in embeddings.items(): + if not isinstance(embedding, torch.Tensor): + raise RuntimeError( + "EPD encoder output must contain tensor embeddings; " + f"got {type(embedding).__name__} for {modality.name.lower()}" + ) + consumed_rows = consumed_per_modality[modality] + if consumed_rows != embedding.shape[0]: + raise RuntimeError( + "Precomputed multimodal embedding consumption mismatch for " + f"{modality.name.lower()}: received {embedding.shape[0]} rows, " + f"consumed {consumed_rows}" + ) + def _resolve_processor(self, processor=None): if processor is None: return self._processor, self._tokenizer diff --git a/python/sglang/srt/multimodal/processors/interns1pro.py b/python/sglang/srt/multimodal/processors/interns1pro.py index 8c9f3c4cc..d143daf4d 100644 --- a/python/sglang/srt/multimodal/processors/interns1pro.py +++ b/python/sglang/srt/multimodal/processors/interns1pro.py @@ -26,7 +26,7 @@ class InternS1_1ImageProcessor(QwenVLImageProcessor): MultimodalDataItem( modality=Modality.IMAGE, offsets=offsets, - precomputed_embeddings=embeddings, + precomputed_embeddings=embeddings[Modality.IMAGE], ) ] diff --git a/test/registered/unit/disaggregation/test_encode_receiver.py b/test/registered/unit/disaggregation/test_encode_receiver.py index 55515fddd..b5f31c5bf 100644 --- a/test/registered/unit/disaggregation/test_encode_receiver.py +++ b/test/registered/unit/disaggregation/test_encode_receiver.py @@ -1,11 +1,25 @@ -"""Unit tests for request construction in the encode-disaggregation path.""" +"""Unit tests for the encode-disaggregation receiver.""" +import asyncio +import threading +import time import unittest from array import array +from http import HTTPStatus from types import SimpleNamespace +from unittest.mock import patch -from sglang.srt.disaggregation.encoder.receiver import MMReceiverBase +from sglang.srt.disaggregation.encoder.receiver import ( + MMReceiverBase, + WaitingMMRequestStatus, + WaitingRDMARequest, + WaitingZmqRequest, + WaitingZmqRequestGrpc, + _ReceiveRegistrationRunner, +) from sglang.srt.disaggregation.utils import DisaggregationMode +from sglang.srt.managers.io_struct import EncoderDispatchErrorReq +from sglang.srt.managers.schedule_batch import Modality from sglang.srt.sampling.sampling_params import SamplingParams from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.test_utils import CustomTestCase @@ -13,7 +27,239 @@ from sglang.test.test_utils import CustomTestCase register_cpu_ci(est_time=2, suite="base-a-test-cpu") +def _make_registration_request(request_cls): + request = request_cls.__new__(request_cls) + request.rid = "registration-test" + request.registration_runner = _ReceiveRegistrationRunner( + "test-encoder-receive-registration" + ) + request.registration_future = None + request.registration_error = None + request.registration_lock = threading.Lock() + request.status = WaitingMMRequestStatus.PENDING + request.error_msg = None + request.error_code = None + request.embedding_pool = None + request.embeddings_buffer = None + request.recv_embedding_data = None + request._pool_slot_id = None + request._mm_finalizer = None + request.recv_socket = None + request.recv_req = SimpleNamespace(rid=request.rid) + request.num_items_assigned = {Modality.IMAGE: [1]} + request.encoder_urls = ["http://encoder"] + request.host_name = "127.0.0.1" + request.receive_count = 1 + request.embedding_port = 12345 + return request + + +def _cancel_registration(request): + future = request.registration_future + request.release_resources() + deadline = time.monotonic() + 1 + while not future.done() and time.monotonic() < deadline: + time.sleep(0.01) + assert future.cancelled() + + +class BlockingResponse: + def __init__(self, started): + self.started = started + + async def __aenter__(self): + self.started.set() + await asyncio.Event().wait() + + async def __aexit__(self, *args): + return False + + +class BlockingSession: + started = None + + def __init__(self, *args, **kwargs): + pass + + async def __aenter__(self): + return self + + async def __aexit__(self, *args): + return False + + def post(self, *args, **kwargs): + return BlockingResponse(self.started) + + +class FailingResponse: + async def __aenter__(self): + raise ConnectionError("encoder unavailable") + + async def __aexit__(self, *args): + return False + + +class FailingSession(BlockingSession): + def post(self, *args, **kwargs): + return FailingResponse() + + +class TestReceiveRegistration(CustomTestCase): + def test_http_registration_does_not_block_scheduler(self): + started = threading.Event() + BlockingSession.started = started + request = _make_registration_request(WaitingZmqRequest) + with patch( + "sglang.srt.disaggregation.encoder.receiver.aiohttp.ClientSession", + BlockingSession, + ): + scheduler_call = threading.Thread( + target=request.send_encode_request, daemon=True + ) + scheduler_call.start() + self.assertTrue(started.wait(timeout=1)) + scheduler_call.join(timeout=0.1) + + self.assertFalse(scheduler_call.is_alive()) + self.assertEqual(request.status, WaitingMMRequestStatus.PENDING) + _cancel_registration(request) + + def test_grpc_registration_does_not_block_scheduler(self): + started = threading.Event() + + async def blocking_registration(*args, **kwargs): + started.set() + await asyncio.Event().wait() + + request = _make_registration_request(WaitingZmqRequestGrpc) + with patch( + "sglang.srt.disaggregation.encoder.receiver._grpc_scheduler_receive_url", + blocking_registration, + ): + scheduler_call = threading.Thread( + target=request.send_encode_request, daemon=True + ) + scheduler_call.start() + self.assertTrue(started.wait(timeout=1)) + scheduler_call.join(timeout=0.1) + + self.assertFalse(scheduler_call.is_alive()) + self.assertEqual(request.status, WaitingMMRequestStatus.PENDING) + _cancel_registration(request) + + def test_failure_is_request_local(self): + request = _make_registration_request(WaitingZmqRequest) + with patch( + "sglang.srt.disaggregation.encoder.receiver.aiohttp.ClientSession", + FailingSession, + ): + request.send_encode_request() + deadline = time.monotonic() + 1 + while request.status == WaitingMMRequestStatus.PENDING: + self.assertLess(time.monotonic(), deadline) + request._try_recv_mm_data() + time.sleep(0.01) + + self.assertEqual(request.status, WaitingMMRequestStatus.FAIL) + self.assertEqual(request.error_code, HTTPStatus.BAD_GATEWAY) + self.assertIn("encoder unavailable", request.error_msg) + + class TestEncodeReceiverRequestConstruction(CustomTestCase): + def test_early_dispatch_error_waits_for_scheduler_request(self): + encode_finished = threading.Event() + scheduler_dispatch_ready = threading.Event() + reported = [] + failure = EncoderDispatchErrorReq( + rid="request-1", + error_msg="encoder unavailable", + error_code=HTTPStatus.BAD_GATEWAY, + ) + + async def fail_encode(**kwargs): + encode_finished.set() + return failure + + receiver = SimpleNamespace(encode=fail_encode) + worker = threading.Thread( + target=MMReceiverBase._run_encode_in_thread, + args=( + receiver, + failure.rid, + [], + "encode", + {}, + [], + None, + scheduler_dispatch_ready, + reported.append, + ), + ) + worker.start() + + self.assertTrue(encode_finished.wait(timeout=1)) + worker.join(timeout=0.05) + self.assertTrue(worker.is_alive()) + self.assertEqual(reported, []) + + scheduler_dispatch_ready.set() + worker.join(timeout=1) + self.assertFalse(worker.is_alive()) + self.assertEqual(reported, [failure]) + + def test_dispatch_error_fails_only_owning_wait(self): + class WaitingRequest: + def __init__(self, rid): + self.rid = rid + self.recv_req = SimpleNamespace(rid=rid) + self.status = WaitingMMRequestStatus.PENDING + self.error_msg = None + self.error_code = None + self.start_time = 0 + + def _try_recv_mm_data(self): + pass + + def _fail_and_release(self, error_msg, error_code=None): + self.error_msg = error_msg + self.error_code = error_code + self.status = WaitingMMRequestStatus.FAIL + + def release_resources(self): + pass + + def close_recv_socket(self): + pass + + owner = WaitingRequest("request-1") + other = WaitingRequest("request-2") + receiver = SimpleNamespace( + waiting_list=[owner, other], + waiting_by_rid={owner.rid: owner, other.rid: other}, + scheduler_recv_socket=None, + wait_timeout=float("inf"), + tp_group=SimpleNamespace(cpu_group=object()), + _drain_scheduler_embeddings=lambda: None, + _sync_fail_info_across_tp=lambda request: None, + create_req=lambda request: request, + ) + dispatch_error = EncoderDispatchErrorReq( + rid=owner.rid, + error_msg="bad media", + error_code=HTTPStatus.UNPROCESSABLE_ENTITY, + ) + + with patch("torch.distributed.all_reduce"): + _, abort_reqs = MMReceiverBase._process_waiting_requests( + receiver, [dispatch_error], waiting_cls=None + ) + + self.assertEqual(owner.status, WaitingMMRequestStatus.FAIL) + self.assertEqual(owner.error_msg, dispatch_error.error_msg) + self.assertEqual(owner.error_code, dispatch_error.error_code) + self.assertEqual(other.status, WaitingMMRequestStatus.PENDING) + self.assertEqual([req.rid for req, _, _ in abort_reqs], [owner.rid]) + def test_extra_key_and_cache_salt_are_forwarded(self): scheduler = SimpleNamespace( model_config=SimpleNamespace(hf_eos_token_id={2}, vocab_size=128), @@ -56,6 +302,98 @@ class TestEncodeReceiverRequestConstruction(CustomTestCase): self.assertEqual(req.extra_key, "classification") self.assertEqual(req.cache_salt, "tenant-a") + def test_rdma_worker_error_is_released_on_scheduler_thread(self): + scheduler_thread = threading.get_ident() + + class ThreadCheckedSocket: + closed_by = None + + def close(self): + self.closed_by = threading.get_ident() + + recv_socket = ThreadCheckedSocket() + request = WaitingRDMARequest.__new__(WaitingRDMARequest) + request.rid = "request-1" + request.status = WaitingMMRequestStatus.PENDING + request.error_msg = None + request.error_code = None + request.recv_socket = recv_socket + request._receive_error = None + request._receive_error_lock = threading.Lock() + request._buffer_lock = threading.Lock() + request._terminal = False + request._receive_running = False + request.registration_future = None + request.embeddings_buffer = None + request._pool_slot_id = None + request.embedding_pool = None + request._mm_finalizer = None + + worker = threading.Thread( + target=lambda: asyncio.run( + request._check_encoder_responses( + [ConnectionError("encoder unavailable")], "/send" + ) + ) + ) + worker.start() + worker.join(timeout=1) + + self.assertFalse(worker.is_alive()) + self.assertEqual(request.status, WaitingMMRequestStatus.PENDING) + self.assertIsNone(request.recv_socket.closed_by) + + request._try_recv_mm_data() + + self.assertEqual(request.status, WaitingMMRequestStatus.FAIL) + self.assertIsNone(request.recv_socket) + self.assertTrue(request._terminal) + self.assertEqual(recv_socket.closed_by, scheduler_thread) + + def test_tp_peer_failure_closes_local_receive_socket(self): + class WaitingRequest: + rid = "request-1" + recv_req = SimpleNamespace(rid=rid) + status = WaitingMMRequestStatus.PENDING + error_msg = "peer failed" + error_code = None + start_time = 0 + released = False + closed = False + + def _try_recv_mm_data(self): + pass + + def release_resources(self): + self.released = True + + def close_recv_socket(self): + self.closed = True + + waiting_req = WaitingRequest() + receiver = SimpleNamespace( + waiting_list=[waiting_req], + waiting_by_rid={waiting_req.rid: waiting_req}, + scheduler_recv_socket=None, + wait_timeout=float("inf"), + tp_group=SimpleNamespace(cpu_group=object()), + _drain_scheduler_embeddings=lambda: None, + _sync_fail_info_across_tp=lambda request: None, + create_req=lambda request: request, + ) + + def force_peer_failure(status, **kwargs): + status.fill_(WaitingMMRequestStatus.FAIL) + + with patch("torch.distributed.all_reduce", force_peer_failure): + _, abort_reqs = MMReceiverBase._process_waiting_requests( + receiver, [], waiting_cls=None + ) + + self.assertTrue(waiting_req.released) + self.assertTrue(waiting_req.closed) + self.assertEqual(len(abort_reqs), 1) + if __name__ == "__main__": unittest.main() diff --git a/test/registered/unit/disaggregation/test_kimi_k3_encoder_mode.py b/test/registered/unit/disaggregation/test_kimi_k3_encoder_mode.py index aa4d9229c..d31223f90 100644 --- a/test/registered/unit/disaggregation/test_kimi_k3_encoder_mode.py +++ b/test/registered/unit/disaggregation/test_kimi_k3_encoder_mode.py @@ -7,7 +7,7 @@ import time from array import array from concurrent.futures import ThreadPoolExecutor from types import SimpleNamespace -from unittest.mock import AsyncMock, patch +from unittest.mock import AsyncMock, MagicMock, Mock, patch import pytest import torch @@ -23,8 +23,11 @@ from sglang.srt.disaggregation.encoder.preprocessor import ( ) from sglang.srt.disaggregation.encoder.receiver import ( EmbeddingData, + MMReceiverGrpc, MMReceiverHTTP, MultiModalEmbeddingData, + WaitingMMRequestStatus, + WaitingZmqRequest, _encoder_media_item, _select_mm_processor_prompt, ) @@ -577,6 +580,102 @@ def test_epd_receiver_keeps_content_hash_aligned_with_image(): } +def test_epd_tokenizer_receiver_timeout_cancels_tasks_and_closes_socket(): + async def run(): + receiver = MMReceiverHTTP.__new__(MMReceiverHTTP) + receiver.encode_urls = ["http://encoder"] + receiver.context = object() + receiver.host = "127.0.0.1" + receiver.recv_timeout = 0.01 + receiver._extract_url_data = Mock(return_value=[{"modality": Modality.IMAGE}]) + encode_cancelled = asyncio.Event() + recv_cancelled = asyncio.Event() + + async def wait_until_cancelled(event, *_args, **_kwargs): + try: + await asyncio.Event().wait() + finally: + event.set() + + receiver.encode = lambda *args, **kwargs: wait_until_cancelled( + encode_cancelled, *args, **kwargs + ) + receiver._recv_mm_data = lambda *args, **kwargs: wait_until_cancelled( + recv_cancelled, *args, **kwargs + ) + recv_socket = SimpleNamespace(close=Mock()) + + with patch( + "sglang.srt.disaggregation.encoder.receiver.get_zmq_socket_on_host", + return_value=(12345, recv_socket), + ): + result = await receiver.recv_mm_data( + SimpleNamespace(), + mm_processor=object(), + prompt="prompt", + ) + + assert result is None + assert encode_cancelled.is_set() + assert recv_cancelled.is_set() + recv_socket.close.assert_called_once_with(linger=0) + + asyncio.run(run()) + + +def test_grpc_dispatch_cancellation_waits_for_blocking_calls(): + async def run(): + receiver = MMReceiverGrpc.__new__(MMReceiverGrpc) + receiver.host = "127.0.0.1" + calls_started = 0 + calls_finished = 0 + calls_lock = threading.Lock() + unblock = threading.Event() + + def blocking_encode(_target, _request): + nonlocal calls_started, calls_finished + with calls_lock: + calls_started += 1 + unblock.wait(timeout=2) + with calls_lock: + calls_finished += 1 + + with patch( + "sglang.srt.disaggregation.encoder.receiver._grpc_encode_request", + side_effect=blocking_encode, + ): + task = asyncio.create_task( + receiver.encode( + req_id="req", + mm_data=[ + {"modality": Modality.IMAGE, "url": "image-0"}, + {"modality": Modality.IMAGE, "url": "image-1"}, + ], + embedding_port=1234, + endpoint_encode="encode", + num_items_assigned=[1, 1], + encode_urls=["grpc://encoder-0", "grpc://encoder-1"], + ) + ) + for _ in range(100): + with calls_lock: + if calls_started == 2: + break + await asyncio.sleep(0.01) + assert calls_started == 2 + + task.cancel() + await asyncio.sleep(0) + assert not task.done() + + unblock.set() + with pytest.raises(asyncio.CancelledError): + await task + assert calls_finished == 2 + + asyncio.run(run()) + + def test_kimi_k3_epd_aggregates_original_image_sizes_in_part_order(): first = EmbeddingData( req_id="request", @@ -607,6 +706,109 @@ def test_kimi_k3_epd_aggregates_original_image_sizes_in_part_order(): ] +@pytest.mark.parametrize( + ("num_parts", "part_idx", "error"), + [ + (0, 0, "num_parts must be a positive integer"), + (2, -1, "part_idx must be in"), + (2, 2, "part_idx must be in"), + ], +) +def test_epd_embedding_aggregation_rejects_invalid_part_metadata( + num_parts, part_idx, error +): + part = EmbeddingData( + req_id="request", + num_parts=num_parts, + part_idx=part_idx, + grid_dim=torch.tensor([[1, 2, 2]]), + modality=Modality.IMAGE, + embedding=torch.ones(1, 2), + ) + + with pytest.raises(ValueError, match=error): + MultiModalEmbeddingData.from_embedding_data(part) + + +def test_epd_embedding_aggregation_rejects_duplicate_and_inconsistent_parts(): + def make_part(num_parts, part_idx): + return EmbeddingData( + req_id="request", + num_parts=num_parts, + part_idx=part_idx, + grid_dim=torch.tensor([[1, 2, 2]]), + modality=Modality.IMAGE, + embedding=torch.ones(1, 2), + ) + + combined = MultiModalEmbeddingData.from_embedding_data(make_part(2, 0)) + with pytest.raises(ValueError, match="duplicate embedding part 0"): + combined.add(make_part(2, 0)) + with pytest.raises(ValueError, match="num_parts changed from 2 to 3"): + combined.add(make_part(3, 1)) + + +def test_epd_scheduler_contains_invalid_embedding_part_metadata(): + waiting = WaitingZmqRequest.__new__(WaitingZmqRequest) + waiting.rid = "request" + waiting.recv_req = SimpleNamespace(rid="request") + waiting.status = WaitingMMRequestStatus.PENDING + waiting.recv_embedding_data = None + waiting.model_type = None + waiting._fail_and_release = Mock() + invalid = EmbeddingData( + req_id="request_local_part_2", + num_parts=2, + part_idx=2, + grid_dim=None, + modality=Modality.IMAGE, + embedding=torch.ones(1, 2), + ) + + waiting.consume_parts( + [pickle.dumps(invalid.copy_without_embedding()), invalid.embedding.numpy()] + ) + + waiting._fail_and_release.assert_called_once() + + +def test_epd_tokenizer_contains_duplicate_embedding_part(): + class FakeSocket: + def __init__(self, messages): + self.messages = messages + self.closed = False + + async def recv_multipart(self, copy=False): + return self.messages.pop(0) + + def close(self): + self.closed = True + + async def run_test(): + embedding = torch.tensor([[1.0, 2.0]]) + part = EmbeddingData( + req_id="request_local_part_0", + num_parts=2, + part_idx=0, + grid_dim=torch.tensor([[1, 2, 2]]), + modality=Modality.IMAGE, + embedding=embedding, + ) + frame = [pickle.dumps(part.copy_without_embedding()), embedding.numpy()] + socket = FakeSocket([frame, frame]) + receiver = MMReceiverHTTP.__new__(MMReceiverHTTP) + receiver.model_type = None + + result = await receiver._recv_mm_data( + "request", socket, SimpleNamespace(), "prompt" + ) + + assert result is None + assert socket.closed + + asyncio.run(run_test()) + + def test_kimi_k3_encoder_prefers_grid_thws_and_uses_temporal_pool_length(): grid_thws = torch.tensor([[3, 8, 12]]) stale_grid = torch.tensor([[1, 2, 2]]) @@ -746,6 +948,81 @@ def test_epd_scheduler_uses_token_ids_for_tokenized_mm_processors(): ) +def test_epd_scheduler_ignores_foreign_error_part(): + waiting = WaitingZmqRequest.__new__(WaitingZmqRequest) + waiting.rid = "current" + waiting.recv_req = SimpleNamespace(rid="current") + waiting.status = WaitingMMRequestStatus.PENDING + waiting._fail_and_release = Mock() + stale_error = EmbeddingData( + req_id="stale_local_part_0", + num_parts=1, + part_idx=0, + grid_dim=None, + modality=Modality.IMAGE, + error_msg="stale failure", + error_code=500, + ) + + waiting.consume_parts([pickle.dumps("not embedding data")]) + waiting.consume_parts([pickle.dumps(stale_error)]) + + assert waiting.status == WaitingMMRequestStatus.PENDING + waiting._fail_and_release.assert_not_called() + + +def test_epd_tokenizer_ignores_foreign_part_before_current_embedding(): + class FakeSocket: + def __init__(self, messages): + self.messages = list(messages) + self.closed = False + + async def recv_multipart(self, copy=False): + return self.messages.pop(0) + + def close(self): + self.closed = True + + async def run_test(): + stale_error = EmbeddingData( + req_id="stale_local_part_0", + num_parts=1, + part_idx=0, + grid_dim=None, + modality=Modality.IMAGE, + error_msg="stale failure", + error_code=500, + ) + embedding = torch.tensor([[1.0, 2.0]]) + current = EmbeddingData( + req_id="current_local_part_0", + num_parts=1, + part_idx=0, + grid_dim=None, + modality=Modality.IMAGE, + embedding=embedding, + ) + socket = FakeSocket( + [ + [pickle.dumps(stale_error)], + [pickle.dumps(current.copy_without_embedding()), embedding.numpy()], + ] + ) + receiver = MMReceiverHTTP.__new__(MMReceiverHTTP) + receiver.model_type = None + processor = SimpleNamespace( + get_mm_data=lambda _prompt, embeddings, **_kwargs: embeddings, + get_validated_mm_data=lambda _prompt, embeddings, **_kwargs: embeddings, + ) + + result = await receiver._recv_mm_data("current", socket, processor, "prompt") + + torch.testing.assert_close(result[Modality.IMAGE], embedding) + assert socket.closed + + asyncio.run(run_test()) + + def test_epd_scheduler_routes_many_requests_over_one_receive_socket(): context = zmq.Context() receiver = MMReceiverHTTP.__new__(MMReceiverHTTP) @@ -761,6 +1038,8 @@ def test_epd_scheduler_routes_many_requests_over_one_receive_socket(): sender = context.socket(zmq.PUSH) try: sender.connect(f"tcp://127.0.0.1:{port}") + sender.send_multipart([b"not a pickle"]) + sender.send_multipart([pickle.dumps("not embedding data")]) for i in range(32): mm_data = EmbeddingData( req_id=f"rid-{i}_local_part_0", @@ -784,6 +1063,84 @@ def test_epd_scheduler_routes_many_requests_over_one_receive_socket(): context.term() +def _receiver_for_startup_failure(rank_errors): + receiver = MMReceiverHTTP.__new__(MMReceiverHTTP) + receiver.mm_processor = object() + receiver.model_type = "kimi_k3" + receiver.hostname = "127.0.0.1" + receiver.tp_size = 2 + receiver.tp_group = MagicMock() + receiver.tp_group.all_gather_object.side_effect = rank_errors + receiver.scheduler_recv_socket = object() + receiver.scheduler_context = object() + receiver.scheduler_embedding_port = 1234 + receiver.encode_urls = ["http://encoder"] + receiver.waiting_by_rid = {} + receiver.waiting_list = [] + receiver.create_req = MagicMock(return_value=object()) + return receiver + + +def test_epd_receiver_startup_rejects_remote_rank_failure(): + receiver = _receiver_for_startup_failure( + lambda local_error: [local_error, "RuntimeError: bind failed"] + ) + waiting_req = MagicMock() + waiting_req.rid = "request-id" + waiting_cls = MagicMock(return_value=waiting_req) + + class TokenizedRequest: + rid = "request-id" + need_wait_for_mm_inputs = True + encoder_urls = ["http://encoder"] + + with patch( + "sglang.srt.disaggregation.encoder.receiver.TokenizedGenerateReqInput", + TokenizedRequest, + ): + ready, aborts = receiver._process_waiting_requests( + [TokenizedRequest()], waiting_cls + ) + + assert ready == [] + assert len(aborts) == 1 + assert "rank 1: RuntimeError: bind failed" in aborts[0][1] + assert aborts[0][2] == 500 + waiting_req.send_encode_request.assert_called_once_with() + waiting_req.release_resources.assert_called_once_with() + waiting_req.close_recv_socket.assert_called_once_with() + assert receiver.waiting_list == [] + assert receiver.waiting_by_rid == {} + + +def test_epd_receiver_startup_shares_local_constructor_failure(): + def gather_local_error(local_error): + assert "RuntimeError: socket failed" in local_error + return [local_error, None] + + receiver = _receiver_for_startup_failure(gather_local_error) + waiting_cls = MagicMock(side_effect=RuntimeError("socket failed")) + + class TokenizedRequest: + rid = "request-id" + need_wait_for_mm_inputs = True + encoder_urls = ["http://encoder"] + + with patch( + "sglang.srt.disaggregation.encoder.receiver.TokenizedGenerateReqInput", + TokenizedRequest, + ): + ready, aborts = receiver._process_waiting_requests( + [TokenizedRequest()], waiting_cls + ) + + assert ready == [] + assert len(aborts) == 1 + assert "rank 0: RuntimeError: socket failed" in aborts[0][1] + assert aborts[0][2] == 500 + assert receiver.waiting_list == [] + + def test_epd_encoder_reuses_scheduler_zmq_peer(): async def send_twice(): context = zmq.asyncio.Context() diff --git a/test/registered/unit/managers/test_tokenizer_manager_rid_cleanup.py b/test/registered/unit/managers/test_tokenizer_manager_rid_cleanup.py index 1fab0d801..b57389b81 100644 --- a/test/registered/unit/managers/test_tokenizer_manager_rid_cleanup.py +++ b/test/registered/unit/managers/test_tokenizer_manager_rid_cleanup.py @@ -129,6 +129,7 @@ def _make_tokenizer_manager(case) -> TokenizerManager: tm.server_args.dp_size = 1 tm.disaggregation_mode = "none" tm.rid_to_state = {} + tm.encoder_dispatch_ready = {} tm.enable_metrics = False tm.enable_trace = False tm.enable_lora = False diff --git a/test/registered/unit/models/test_interns1pro_processor.py b/test/registered/unit/models/test_interns1pro_processor.py new file mode 100644 index 000000000..36009bc9e --- /dev/null +++ b/test/registered/unit/models/test_interns1pro_processor.py @@ -0,0 +1,39 @@ +"""CPU tests for InternS1-Pro multimodal processor behavior.""" + +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=5, suite="base-a-test-cpu") + +from types import SimpleNamespace +from unittest.mock import Mock + +import pytest +import torch + +from sglang.srt.managers.schedule_batch import Modality +from sglang.srt.multimodal.processors.interns1pro import InternS1_1ImageProcessor + + +def test_epd_stores_the_image_tensor_in_the_mm_item(): + processor = object.__new__(InternS1_1ImageProcessor) + processor.build_input_ids = Mock(return_value=([1, 2, 3], [(1, 2)])) + processor.IM_START_TOKEN_ID = 10 + processor.IM_END_TOKEN_ID = 11 + processor.mm_tokens = SimpleNamespace( + image_token_id=12, + video_token_id=13, + audio_token_id=14, + ) + image_embedding = torch.zeros(2, 4) + + output = processor.get_validated_mm_data( + [1, 2, 3], + {Modality.IMAGE: image_embedding}, + img_grid_thw=torch.tensor([[1, 2, 2]]), + ) + + assert output.mm_items[0].precomputed_embeddings is image_embedding + + +if __name__ == "__main__": + raise SystemExit(pytest.main([__file__, "-v"])) diff --git a/test/registered/unit/models/test_kimi_k25.py b/test/registered/unit/models/test_kimi_k25.py index 0c368c29b..9fe089120 100644 --- a/test/registered/unit/models/test_kimi_k25.py +++ b/test/registered/unit/models/test_kimi_k25.py @@ -727,7 +727,7 @@ def test_kimi_k3_epd_rebuild_uses_the_same_media_contract(): processor._tokenizer = _Tokenizer() embeddings = {Modality.IMAGE: torch.arange(20, dtype=torch.float32).reshape(5, 4)} - output = processor.get_mm_data( + output = processor.get_validated_mm_data( [1, 99, 2, 99, 3], embeddings, img_grid_thw=torch.tensor([[1, 2, 6], [1, 2, 4]]), diff --git a/test/registered/unit/multimodal/test_gpu_feature_transport.py b/test/registered/unit/multimodal/test_gpu_feature_transport.py index 590f8b2c4..ded81ae54 100644 --- a/test/registered/unit/multimodal/test_gpu_feature_transport.py +++ b/test/registered/unit/multimodal/test_gpu_feature_transport.py @@ -377,6 +377,7 @@ class TestCudaVmmFeatureTransport(unittest.TestCase): manager = object.__new__(TokenizerManager) manager.rid_to_state = {} + manager.encoder_dispatch_ready = {} transport = MagicMock() transport.prepare_for_dispatch_async = AsyncMock(return_value=[]) manager.cuda_vmm_feature_transport = transport @@ -405,6 +406,7 @@ class TestCudaVmmFeatureTransport(unittest.TestCase): manager = object.__new__(tokenizer_manager.TokenizerManager) manager.rid_to_state = {} + manager.encoder_dispatch_ready = {} transport = MagicMock() manager._dispatch_to_scheduler = MagicMock( side_effect=RuntimeError("send failed") @@ -440,6 +442,7 @@ class TestCudaVmmFeatureTransport(unittest.TestCase): manager = object.__new__(tokenizer_manager.TokenizerManager) manager.rid_to_state = {} + manager.encoder_dispatch_ready = {} transport = MagicMock() manager._dispatch_to_scheduler = MagicMock() time_stats = MagicMock() diff --git a/test/registered/unit/multimodal/test_precomputed_embedding_validation.py b/test/registered/unit/multimodal/test_precomputed_embedding_validation.py new file mode 100644 index 000000000..a10c7147e --- /dev/null +++ b/test/registered/unit/multimodal/test_precomputed_embedding_validation.py @@ -0,0 +1,90 @@ +"""Tests for the common EPD precomputed-embedding boundary.""" + +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=5, suite="base-a-test-cpu") + +import unittest + +import torch + +from sglang.srt.managers.schedule_batch import ( + Modality, + MultimodalDataItem, + MultimodalProcessorOutput, +) +from sglang.srt.multimodal.processors.base_processor import BaseMultimodalProcessor + + +class _StubProcessor(BaseMultimodalProcessor): + async def process_mm_data_async(self, *args, **kwargs): + raise NotImplementedError + + def get_mm_data(self, prompt, embeddings, **kwargs): + return self.output + + +def _item(modality, rows, offsets): + return MultimodalDataItem( + modality=modality, + offsets=offsets, + precomputed_embeddings=torch.zeros(rows, 4), + ) + + +class TestPrecomputedEmbeddingValidation(unittest.TestCase): + def setUp(self): + self.processor = object.__new__(_StubProcessor) + + def _validate(self, items, embeddings): + self.processor.output = MultimodalProcessorOutput( + input_ids=[1, 2, 3], + mm_items=items, + ) + return self.processor.get_validated_mm_data([], embeddings) + + def test_accepts_exact_multi_item_layout(self): + image_embedding = torch.zeros(5, 4) + audio_embedding = torch.zeros(2, 4) + output = self._validate( + [ + _item(Modality.IMAGE, 2, [(1, 2)]), + _item(Modality.IMAGE, 3, [(4, 6)]), + _item(Modality.AUDIO, 2, [(8, 9)]), + ], + { + Modality.IMAGE: image_embedding, + Modality.AUDIO: audio_embedding, + }, + ) + + self.assertEqual(len(output.mm_items), 3) + + def test_rejects_item_shorter_than_prompt_offsets(self): + with self.assertRaisesRegex(RuntimeError, "expected 3 rows.*got 2"): + self._validate( + [_item(Modality.IMAGE, 2, [(1, 3)])], + {Modality.IMAGE: torch.zeros(2, 4)}, + ) + + def test_rejects_unconsumed_trailing_rows(self): + with self.assertRaisesRegex(RuntimeError, "received 3 rows, consumed 2"): + self._validate( + [_item(Modality.IMAGE, 2, [(1, 2)])], + {Modality.IMAGE: torch.zeros(3, 4)}, + ) + + def test_rejects_missing_modality(self): + with self.assertRaisesRegex(RuntimeError, "received 2 rows, consumed 0"): + self._validate([], {Modality.VIDEO: torch.zeros(2, 4)}) + + def test_rejects_unexpected_modality(self): + with self.assertRaisesRegex(RuntimeError, "unexpected embedding modality"): + self._validate( + [_item(Modality.VIDEO, 2, [(1, 2)])], + {Modality.IMAGE: torch.zeros(2, 4)}, + ) + + +if __name__ == "__main__": + unittest.main()