From 8ef646a5c65bd2f8922483057dddc02e2b0de18c Mon Sep 17 00:00:00 2001 From: Mick Date: Sun, 6 Sep 2026 16:05:10 +0800 Subject: [PATCH] fix(vlm): contain EPD request lifecycle failures (#36944) Co-authored-by: mickqian --- .../srt/disaggregation/encoder/grpc_server.py | 63 +- .../srt/disaggregation/encoder/http_server.py | 114 +- .../srt/disaggregation/encoder/runtime.py | 359 +++- .../srt/disaggregation/encoder/server.py | 359 +++- .../basic/test_openai_completion_rust.py | 4 +- .../unit/disaggregation/test_encode_server.py | 1475 ++++++++++++++++- .../disaggregation/test_encoder_health.py | 89 +- .../disaggregation/test_encoder_scheduler.py | 111 ++ 8 files changed, 2435 insertions(+), 139 deletions(-) diff --git a/python/sglang/srt/disaggregation/encoder/grpc_server.py b/python/sglang/srt/disaggregation/encoder/grpc_server.py index ae2f5aac3..4b41a49ee 100644 --- a/python/sglang/srt/disaggregation/encoder/grpc_server.py +++ b/python/sglang/srt/disaggregation/encoder/grpc_server.py @@ -12,6 +12,7 @@ import logging import multiprocessing as mp import traceback from concurrent import futures +from http import HTTPStatus from typing import List import grpc @@ -21,7 +22,12 @@ from grpc_health.v1 import health_pb2, health_pb2_grpc from grpc_reflection.v1alpha import reflection from smg_grpc_proto import sglang_encoder_pb2, sglang_encoder_pb2_grpc -from sglang.srt.disaggregation.encoder.server import MMEncoder, launch_encoder +from sglang.srt.disaggregation.encoder.runtime import validate_encode_request +from sglang.srt.disaggregation.encoder.server import ( + MMEncoder, + await_task_completion_on_cancel, + launch_encoder, +) from sglang.srt.managers.io_struct import async_sock_send, wrap_as_pickle from sglang.srt.managers.schedule_batch import Modality from sglang.srt.runtime_context import ( @@ -87,6 +93,11 @@ class SGLangEncoderServer(SGLangEncoderServicer): self.send_sockets = send_sockets self.server_args = server_args + async def _dispatch_encode(self, request_dict: dict): + for socket in self.send_sockets: + await async_sock_send(socket, wrap_as_pickle(request_dict)) + return await self.encoder.encode_request(request_dict, Modality.IMAGE) + async def Encode( self, request: sglang_encoder_pb2.EncodeRequest, context ) -> sglang_encoder_pb2.EncodeResponse: @@ -98,21 +109,33 @@ class SGLangEncoderServer(SGLangEncoderServicer): "num_parts": request.num_parts, "part_idx": request.part_idx, } - for socket in self.send_sockets: - await async_sock_send(socket, wrap_as_pickle(request_dict)) + if err := validate_encode_request(request_dict): + context.set_code(grpc.StatusCode.INVALID_ARGUMENT) + context.set_details(err) + return sglang_encoder_pb2.EncodeResponse() # gRPC encode is image-only; the request follows the configured - # cache and transfer backend. - ( - nbytes, - embedding_len, - embedding_dim, - error_msg, - error_code, - ) = await self.encoder.encode_request(request_dict, Modality.IMAGE) + # cache and transfer backend. Keep TP dispatch and rank-0 collective + # launch order identical when gRPC handlers run concurrently. + async with self.encoder.encode_dispatch_lock: + encode_task = asyncio.create_task(self._dispatch_encode(request_dict)) + result = await await_task_completion_on_cancel( + encode_task, f"Encoder request {request.req_id}" + ) + ( + nbytes, + embedding_len, + embedding_dim, + error_msg, + error_code, + ) = result if error_msg is not None: await self.encoder.release_request(request.req_id) - context.set_code(grpc.StatusCode.INTERNAL) + context.set_code( + grpc.StatusCode.INVALID_ARGUMENT + if error_code == HTTPStatus.BAD_REQUEST + else grpc.StatusCode.INTERNAL + ) context.set_details(error_msg) return sglang_encoder_pb2.EncodeResponse() @@ -154,6 +177,14 @@ class SGLangEncoderServer(SGLangEncoderServicer): return sglang_encoder_pb2.EncodeResponse() + except asyncio.CancelledError: + try: + await asyncio.shield(self.encoder.release_request(request.req_id)) + except Exception: + logger.exception( + "Failed to release cancelled encoder request %s", request.req_id + ) + raise except Exception as e: logger.error(f"Encode error: {e}") traceback.print_exc() @@ -178,6 +209,14 @@ class SGLangEncoderServer(SGLangEncoderServicer): await self.encoder.release_request(request.req_id) return sglang_encoder_pb2.SendResponse() + except asyncio.CancelledError: + try: + await asyncio.shield(self.encoder.release_request(request.req_id)) + except Exception: + logger.exception( + "Failed to release cancelled encoder request %s", request.req_id + ) + raise except Exception as e: logger.error(f"Send error: {e}") traceback.print_exc() diff --git a/python/sglang/srt/disaggregation/encoder/http_server.py b/python/sglang/srt/disaggregation/encoder/http_server.py index a1a6771fa..6cf299956 100644 --- a/python/sglang/srt/disaggregation/encoder/http_server.py +++ b/python/sglang/srt/disaggregation/encoder/http_server.py @@ -32,6 +32,8 @@ from sglang.srt.disaggregation.encoder.runtime import ( execute_encode_pipeline, launch_dp_runtime, launch_local_runtime, + send_staged_embedding, + validate_encode_request, ) from sglang.srt.disaggregation.encoder.server import ( EncoderProfiler, @@ -260,8 +262,38 @@ def _summarise_dp_broadcast(results: List[dict]) -> Response: ) +async def _drain_health_encode( + health_encoder: MMEncoder, encode_task: asyncio.Task, req_id: str +): + """Finish a dispatched TP probe before releasing its state and lock.""" + result = None + cleanup_failed = False + try: + result = await asyncio.shield(encode_task) + except Exception: + logger.exception("Encoder health check failed for req_id=%s", req_id) + finally: + try: + await asyncio.shield(health_encoder.release_request(req_id)) + except Exception: + cleanup_failed = True + logger.exception("Encoder health cleanup failed for req_id=%s", req_id) + finally: + health_encoder.encode_dispatch_lock.release() + return None if cleanup_failed else result + + @app.post("/encode") async def handle_encode_request(request: dict): + if err := validate_encode_request(request): + return ORJSONResponse( + status_code=HTTPStatus.BAD_REQUEST, + content={ + "status": "error", + "message": err, + "req_id": request.get("req_id"), + }, + ) req_id = request["req_id"] start_time = time.monotonic() time_stats_json = request.pop("time_stats_json", None) @@ -351,7 +383,6 @@ async def handle_send_request(request: dict): """Mooncake-only: drive the RDMA push of a staged embedding. The zmq backends deliver embeddings inline during /encode and never call /send.""" req_id = request["req_id"] - receive_count = request.get("receive_count") if dp_dispatcher is not None: try: result = await dp_dispatcher.dispatch_send(request) @@ -373,13 +404,23 @@ async def handle_send_request(request: dict): status_code=status_code, ) return ORJSONResponse(content=result.get("content")) - sent = await encoder.send( - req_id=req_id, - prefill_host=request["prefill_host"], - embedding_port=request["embedding_port"], - session_id=request["session_id"], - buffer_address=request["buffer_address"], - ) + try: + sent = await send_staged_embedding( + encoder, + request, + # A pre-refcount decoder may have sibling ranks still to send. + release_without_count=False, + ) + except Exception as error: + logger.error("Mooncake send failed for req_id=%s: %s", req_id, error) + return ORJSONResponse( + status_code=HTTPStatus.INTERNAL_SERVER_ERROR, + content={ + "status": "error", + "message": str(error), + "req_id": req_id, + }, + ) if not sent: # No transfer happened: fail fast rather than 200 + a phantom count. return ORJSONResponse( @@ -390,11 +431,6 @@ async def handle_send_request(request: dict): "req_id": req_id, }, ) - # Sibling ranks share this embedding, so free it only once all have sent. - # No count means a pre-refcount decoder: leave it to the sweep, as when - # some rank never sends at all. - if receive_count: - await server_module.meta_registry.note_send_done(req_id, receive_count) return ORJSONResponse(content=None) @@ -544,10 +580,10 @@ async def health_generate(): # No processor available, fall back to liveness check only return Response(status_code=200) + # uuid keeps rids unique across workers; a bare time.time() can collide. + req_id = f"{HEALTH_CHECK_RID_PREFIX}_{uuid.uuid4().hex}" + owns_dispatch_lock = False try: - # uuid keeps rids unique across workers; a bare time.time() can collide. - req_id = f"{HEALTH_CHECK_RID_PREFIX}_{uuid.uuid4().hex}" - dummy_request = { "mm_items": mm_items, "modality": modality.name, @@ -560,25 +596,36 @@ async def health_generate(): # request. Serialize its broadcast and rank-0 forward with every other # collective dispatch, then recheck whether traffic made the probe # unnecessary while it waited for the lock. - async with encoder.encode_dispatch_lock: - if encoder.has_pending_embeddings(): - return Response(status_code=200) - for socket in send_sockets: - sock_send(socket, wrap_as_pickle(dummy_request)) + await encoder.encode_dispatch_lock.acquire() + owns_dispatch_lock = True + if encoder.has_pending_embeddings(): + return Response(status_code=200) + for socket in send_sockets: + sock_send(socket, wrap_as_pickle(dummy_request)) - _, _, _, error_msg, _ = await asyncio.wait_for( - encoder.encode( - mm_items=mm_items, - modality=modality, - req_id=req_id, - num_parts=1, - part_idx=0, - ), - timeout=HEALTH_CHECK_TIMEOUT, + encode_task = asyncio.create_task( + encoder.encode( + mm_items=mm_items, + modality=modality, + req_id=req_id, + num_parts=1, + part_idx=0, ) + ) + drain_task = asyncio.create_task( + _drain_health_encode(encoder, encode_task, req_id) + ) + # The drain task now owns the lock and request state. A probe timeout or + # client disconnect must not let a later request overtake its TP work. + owns_dispatch_lock = False + result = await asyncio.wait_for( + asyncio.shield(drain_task), + timeout=HEALTH_CHECK_TIMEOUT, + ) - # Clean up stored embedding - await encoder.release_request(req_id) + if result is None: + return Response(status_code=503) + _, _, _, error_msg, _ = result if error_msg: logger.error(f"Encoder health check failed: {error_msg}") @@ -592,6 +639,9 @@ async def health_generate(): except Exception as e: logger.error(f"Encoder health check failed: {e}") return Response(status_code=503) + finally: + if owns_dispatch_lock: + encoder.encode_dispatch_lock.release() @app.api_route("/start_profile", methods=["GET", "POST"]) diff --git a/python/sglang/srt/disaggregation/encoder/runtime.py b/python/sglang/srt/disaggregation/encoder/runtime.py index 67b9b3a1c..9effd06de 100644 --- a/python/sglang/srt/disaggregation/encoder/runtime.py +++ b/python/sglang/srt/disaggregation/encoder/runtime.py @@ -12,6 +12,7 @@ import contextlib import logging import multiprocessing as mp import os +import sys import time import traceback import uuid @@ -31,6 +32,7 @@ from sglang.srt.disaggregation.encoder.server import ( EncoderProfiler, MMEncoder, MMError, + await_task_completion_on_cancel, launch_encoder, ) from sglang.srt.environ import envs @@ -76,6 +78,46 @@ class PendingRequest: # vary per request and can't merge into one HF processor call. _BATCHABLE_MODALITIES = {Modality.IMAGE, Modality.AUDIO} _KIMI_K3_DEFAULT_ENCODER_MAX_BATCH_SIZE = 2 +_DP_RELEASE_AFTER_ENCODE = "release_after_encode" + + +def validate_encode_request(request: dict) -> Optional[str]: + """Return a client-facing error before an encode request is dispatched.""" + if not isinstance(request, dict): + return f"request is not a dict: {type(request).__name__}" + + req_id = request.get("req_id") + if not isinstance(req_id, str) or not req_id: + return "missing or invalid req_id" + + modality = request.get("modality") + if not isinstance(modality, str): + return "missing or invalid modality" + try: + Modality.from_str(modality) + except ValueError: + return f"unsupported modality: {modality}" + + mm_items = request.get("mm_items") + if mm_items is None or (isinstance(mm_items, (list, tuple)) and len(mm_items) == 0): + return "missing or empty mm_items" + + num_parts = request.get("num_parts") + part_idx = request.get("part_idx") + if not isinstance(num_parts, int) or isinstance(num_parts, bool) or num_parts <= 0: + return "num_parts must be a positive integer" + if ( + not isinstance(part_idx, int) + or isinstance(part_idx, bool) + or part_idx < 0 + or part_idx >= num_parts + ): + return f"part_idx must be in [0, {num_parts})" + + hashes = request.get("hashes") + if hashes is not None and not isinstance(hashes, (list, tuple, str, int, bytes)): + return f"hashes must be list/scalar, got {type(hashes).__name__}" + return None def _resolve_encoder_batch_policy( @@ -203,27 +245,12 @@ class EncoderScheduler: if not p.future.done(): p.future.set_exception(e) - @staticmethod - def _validate_request_shape(req: dict) -> Optional[str]: - # Cheap pre-broadcast checks: shape errors that don't require running - # the HF processor. Once a request reaches TP workers they enter - # batch_encode and expect to join its collectives — a malformed batch - # that makes rank-0 bail mid-flight would deadlock the workers. - if not isinstance(req, dict): - return f"request is not a dict: {type(req).__name__}" - if not req.get("req_id"): - return "missing req_id" - if not req.get("mm_items"): - return "missing or empty mm_items" - if "num_parts" not in req or "part_idx" not in req: - return "missing num_parts / part_idx" - h = req.get("hashes") - if h is not None and not isinstance(h, (list, tuple, str, int, bytes)): - return f"hashes must be list/scalar, got {type(h).__name__}" - return None - async def _dispatch_group( - self, group: List[PendingRequest], modality: Modality + self, + group: List[PendingRequest], + modality: Modality, + *, + observe_queue_wait: bool = True, ) -> None: # A request may time out while queued. Never start work that no caller # can observe, or its eventual staged embedding would have no owner. @@ -241,7 +268,7 @@ class EncoderScheduler: # abandoned. valid: List[PendingRequest] = [] for p in group: - err = self._validate_request_shape(p.request) + err = validate_encode_request(p.request) if err is None: valid.append(p) continue @@ -255,7 +282,7 @@ class EncoderScheduler: requests = [p.request for p in group] start = time.time() modality_str = modality.name.lower() - if server_module.encoder_metrics_collector is not None: + if observe_queue_wait and server_module.encoder_metrics_collector is not None: for p in group: server_module.encoder_metrics_collector.observe_queue_wait( max(0.0, start - p.submit_time), modality=modality_str @@ -309,6 +336,22 @@ class EncoderScheduler: p.future.set_exception(err) return + if len(group) > 1 and all( + result[3] is not None + and result[4] is not None + and int(result[4]) == HTTPStatus.BAD_REQUEST + for result in results + ): + logger.warning( + f"Retrying failed {modality.name} batch as {len(group)} " + "individual requests" + ) + for pending in group: + await self._dispatch_group( + [pending], modality, observe_queue_wait=False + ) + return + for p, result in zip(group, results): if not p.future.done(): p.future.set_result(result) @@ -324,6 +367,8 @@ class EncoderScheduler: continue req = p.request try: + if err := validate_encode_request(req): + raise server_module.BadRequestError(err) start = time.time() if server_module.encoder_metrics_collector is not None: server_module.encoder_metrics_collector.observe_queue_wait( @@ -380,6 +425,7 @@ class DPDispatcher: self, dp_size: int, dispatch_sockets: List, + release_sockets: List, result_socket, worker_processes: List[mp.Process], enable_metrics: bool = False, @@ -387,6 +433,7 @@ class DPDispatcher: ): self.dp_size = dp_size self.dispatch_sockets = dispatch_sockets + self.release_sockets = release_sockets self.result_socket = result_socket self.worker_processes = worker_processes # Key = req_id for encode/broadcast, or a per-control-request key for @@ -405,7 +452,7 @@ class DPDispatcher: # Set when _result_listener gives up; makes alive_ranks report empty. self._listener_failed = False # The event loop only keeps weak references to tasks, so the long-lived - # loops started in `start()` need a strong reference to survive GC. + # loops and fire-and-forget notifications need a strong reference. self.background_tasks: Set[asyncio.Task] = set() # Prometheus gauge: pending requests per DP rank. Lives in the main @@ -461,6 +508,28 @@ class DPDispatcher: self.req_id_to_rank.pop(req_id, None) self._update_pending_gauge() + def _release_abandoned_encode(self, rank: int, req_id: str) -> None: + """Tell the owning worker to release an encode that lost its caller.""" + + async def notify_worker() -> None: + try: + await async_sock_send( + self.release_sockets[rank], + wrap_as_pickle( + {"_dp_type": _DP_RELEASE_AFTER_ENCODE, "req_id": req_id} + ), + ) + except Exception: + logger.exception( + "Failed to retire abandoned encoder DP request %s on rank %s", + req_id, + rank, + ) + + task = asyncio.create_task(notify_worker()) + self.background_tasks.add(task) + task.add_done_callback(self.background_tasks.discard) + @staticmethod def _send_req_key(req_id: str, request: dict) -> str: """One in-flight /send future per decoder TP rank, keyed by the rank's @@ -546,6 +615,7 @@ class DPDispatcher: future = asyncio.get_running_loop().create_future() self.pending_futures[rank][req_id] = future self._update_pending_gauge() + dispatched = False logger.info( f"MM-Encoder DP dispatch: req_id={req_id}, " f"modality={request.get('modality', 'image')}, " @@ -563,6 +633,7 @@ class DPDispatcher: await async_sock_send( self.dispatch_sockets[rank], wrap_as_pickle(request) ) + dispatched = True except BaseException: self._drop_pending_and_mapping(rank, req_id) self._mapping_condition.notify_all() @@ -574,6 +645,8 @@ class DPDispatcher: future, timeout=server_module.ENCODER_REQ_TIMEOUT ) except asyncio.TimeoutError: + if dispatched: + self._release_abandoned_encode(rank, req_id) self._drop_pending_and_mapping(rank, req_id) return self._timeout_envelope( req_id, @@ -581,6 +654,8 @@ class DPDispatcher: f"Encoder DP rank={rank} timed out after {server_module.ENCODER_REQ_TIMEOUT}s", ) except BaseException: + if dispatched: + self._release_abandoned_encode(rank, req_id) self._drop_pending_and_mapping(rank, req_id) raise @@ -898,6 +973,12 @@ class DPDispatcher: return await asyncio.sleep(min(0.1 * consecutive_errors, 1.0)) continue + if not isinstance(msg, dict): + logger.error( + "_result_listener received a non-dict envelope (%s); dropping", + type(msg).__name__, + ) + continue req_id = msg.get("req_id", "") dp_type = msg.get("_dp_type", "encode") if dp_type == "send": @@ -1062,6 +1143,49 @@ async def _push_embedding_to_prefill( await enc.release_request(req_id) +async def send_staged_embedding( + enc: MMEncoder, + request: dict, + *, + release_without_count: bool, +) -> bool: + """Send one Mooncake embedding and retire its state on any failure.""" + req_id = request["req_id"] + try: + sent = await enc.send( + req_id=req_id, + prefill_host=request["prefill_host"], + embedding_port=request["embedding_port"], + session_id=request["session_id"], + buffer_address=request["buffer_address"], + ) + if not sent: + return False + + receive_count = request.get("receive_count") + if receive_count: + destination_endpoint = NetworkAddress( + request["prefill_host"], request["embedding_port"] + ).to_host_port_str() + await server_module.meta_registry.note_send_done( + req_id, receive_count, destination_endpoint + ) + elif release_without_count: + await enc.release_request(req_id) + return True + except BaseException as error: + try: + await enc.release_request(req_id) + except Exception as cleanup_error: + if sys.version_info >= (3, 11): + error.add_note(f"Failed to release encoder request: {cleanup_error}") + else: + logger.exception( + "Failed to release encoder request %s after send failure", req_id + ) + raise + + def _record_pipeline_result(modality: Modality, status: str) -> None: if server_module.encoder_metrics_collector is not None: server_module.encoder_metrics_collector.inc_requests_total( @@ -1069,6 +1193,48 @@ def _record_pipeline_result(modality: Modality, status: str) -> None: ) +async def _publish_pipeline_error(req_id: str, error_msg: str) -> bool: + """Report a request error without letting reporting block cleanup.""" + try: + await server_module.meta_registry.publish(req_id, 0, 0, 0, error=error_msg) + return True + except Exception: + logger.exception("Failed to publish encoder error for req_id=%s", req_id) + return False + + +async def _release_failed_request( + enc: MMEncoder, + req_id: str, + *, + preserve_metadata: bool = False, +) -> None: + """Release request resources without hiding the original request error.""" + try: + await enc.release_request(req_id, preserve_metadata=preserve_metadata) + except Exception: + logger.exception("Failed to release encoder resources for req_id=%s", req_id) + + +async def _run_dispatched_encode( + enc: MMEncoder, request: dict, modality: Modality +) -> Tuple: + """Finish TP encode collectives before propagating caller cancellation.""" + encode_task = asyncio.create_task( + enc.encode( + mm_items=request["mm_items"], + modality=modality, + req_id=request["req_id"], + num_parts=request["num_parts"], + part_idx=request["part_idx"], + hashes=request.get("hashes"), + ) + ) + return await await_task_completion_on_cancel( + encode_task, f"Encoder request {request['req_id']}" + ) + + async def execute_encode_pipeline( enc: MMEncoder, sched: Optional[EncoderScheduler], @@ -1082,6 +1248,8 @@ async def execute_encode_pipeline( and keeps the result until follow-up /send calls complete. ZMQ has no early consumer: it waits for encode, sends the embedding, releases it, then returns. """ + if err := validate_encode_request(request): + raise server_module.BadRequestError(err) req_id = request["req_id"] time_stats_json = request.pop("time_stats_json", None) time_stats = EncoderReqTimeStats() @@ -1111,14 +1279,7 @@ async def execute_encode_pipeline( async with enc.encode_dispatch_lock: for socket in send_sockets: sock_send(socket, wrap_as_pickle(request)) - result = await enc.encode( - mm_items=request["mm_items"], - modality=modality, - req_id=request["req_id"], - num_parts=request["num_parts"], - part_idx=request["part_idx"], - hashes=request.get("hashes"), - ) + result = await _run_dispatched_encode(enc, request, modality) else: result = await enc.encode( mm_items=request["mm_items"], @@ -1128,27 +1289,48 @@ async def execute_encode_pipeline( part_idx=request["part_idx"], hashes=request.get("hashes"), ) + except asyncio.CancelledError: + error_msg = "encoder request cancelled" + time_stats.trace_ctx.abort(abort_info={"reason": error_msg}) + try: + await asyncio.shield(enc.release_request(req_id)) + except Exception: + logger.exception("Failed to release cancelled encoder request %s", req_id) + _record_pipeline_result(modality, "error") + raise except asyncio.TimeoutError: error_msg = "encoder batch timed out" time_stats.trace_ctx.abort(abort_info={"reason": error_msg}) - await server_module.meta_registry.publish(req_id, 0, 0, 0, error=error_msg) - await enc.release_request(req_id, preserve_metadata=backend == "mooncake") + error_published = await _publish_pipeline_error(req_id, error_msg) + await _release_failed_request( + enc, + req_id, + preserve_metadata=backend == "mooncake" and error_published, + ) _record_pipeline_result(modality, "error") raise except Exception as e: error_msg = str(e) time_stats.trace_ctx.abort(abort_info={"reason": error_msg}) - await server_module.meta_registry.publish(req_id, 0, 0, 0, error=error_msg) - await enc.release_request(req_id, preserve_metadata=backend == "mooncake") + error_published = await _publish_pipeline_error(req_id, error_msg) + await _release_failed_request( + enc, + req_id, + preserve_metadata=backend == "mooncake" and error_published, + ) _record_pipeline_result(modality, "error") raise nbytes, embedding_len, embedding_dim, error_msg, error_code = result if error_msg: time_stats.trace_ctx.abort(abort_info={"reason": error_msg}) - await server_module.meta_registry.publish(req_id, 0, 0, 0, error=error_msg) + error_published = await _publish_pipeline_error(req_id, error_msg) if backend == "mooncake": - await enc.release_request(req_id, preserve_metadata=True) + await _release_failed_request( + enc, + req_id, + preserve_metadata=error_published, + ) else: try: await _push_embedding_to_prefill( @@ -1161,6 +1343,7 @@ async def execute_encode_pipeline( f"Error-send failed for req_id={req_id}: {send_err}", exc_info=True, ) + await _release_failed_request(enc, req_id) _record_pipeline_result(modality, "error") raise MMError(error_msg, code=error_code or HTTPStatus.INTERNAL_SERVER_ERROR) @@ -1287,12 +1470,10 @@ async def _dp_worker_handle_request( ) from e elif dp_type == "send": req_id = request["req_id"] - sent = await enc.send( - req_id=req_id, - prefill_host=request["prefill_host"], - embedding_port=request["embedding_port"], - session_id=request["session_id"], - buffer_address=request["buffer_address"], + sent = await send_staged_embedding( + enc, + request, + release_without_count=True, ) if not sent: # Error envelope, not 200 + phantom count: the decoder must @@ -1300,13 +1481,6 @@ async def _dp_worker_handle_request( raise MMError( f"no staged embedding for /send req_id={req_id} (already released)" ) - # Releasing on the first /send breaks decoder TP > 1. No count means - # a pre-refcount decoder: stay eager rather than pin until the sweep. - receive_count = request.get("receive_count") - if receive_count: - await server_module.meta_registry.note_send_done(req_id, receive_count) - else: - await enc.release_request(req_id) content = None else: content = await execute_encode_pipeline(enc, sched, request) @@ -1337,7 +1511,12 @@ async def _dp_worker_handle_request( f"req_id={request.get('req_id', '?')}: {e}", exc_info=True, ) - err_code = int(getattr(e, "code", None) or HTTPStatus.INTERNAL_SERVER_ERROR) + # Only MMError carries an HTTP status in this protocol. Third-party + # exceptions may expose a callable ``code`` attribute (for example + # gRPC errors), which must not make error reporting fail a second time. + err_code = int( + e.code if isinstance(e, MMError) else HTTPStatus.INTERNAL_SERVER_ERROR + ) envelope = { "req_id": request.get("req_id", ""), "_dp_type": dp_type, @@ -1368,11 +1547,27 @@ async def _dp_worker_handle_request( ) +async def _retire_abandoned_encode( + enc: MMEncoder, + encode_task: Optional[asyncio.Task], + req_id: str, +) -> None: + """Retire an abandoned request without interrupting its encode work.""" + try: + if encode_task is None or not encode_task.done(): + await enc.abandon_request(req_id) + else: + await enc.release_request(req_id) + except Exception: + logger.exception("Failed to release abandoned encoder DP request %s", req_id) + + async def run_dp_worker( server_args: ServerArgs, dp_rank: int, gpu_id: int, dispatch_path: str, + release_path: str, result_path: str, ): logger.info( @@ -1414,9 +1609,34 @@ async def run_dp_worker( ctx = zmq.asyncio.Context(2) recv_sock = get_zmq_socket(ctx, zmq.PULL, dispatch_path, False) + release_sock = get_zmq_socket(ctx, zmq.PULL, release_path, False) send_sock = get_zmq_socket(ctx, zmq.PUSH, result_path, False) send_lock = asyncio.Lock() inflight: Set[asyncio.Task] = set() + encode_tasks: Dict[str, asyncio.Task] = {} + release_tasks: Set[asyncio.Task] = set() + + async def listen_for_releases() -> None: + # Cleanup must not wait behind the bounded encode queue: under a + # cancellation burst every normal worker slot may already be occupied. + while True: + try: + request = await async_sock_recv(release_sock) + except asyncio.CancelledError: + raise + except Exception: + logger.error(f"DP worker {dp_rank} release recv error", exc_info=True) + continue + if not isinstance(request, dict) or not request.get("req_id"): + logger.error(f"DP worker {dp_rank} received malformed release request") + continue + req_id = request["req_id"] + task = asyncio.create_task( + _retire_abandoned_encode(enc, encode_tasks.get(req_id), req_id) + ) + release_tasks.add(task) + task.add_done_callback(release_tasks.discard) + # Acquire-before-recv → back-pressure propagates to the dispatcher # PUSH buffer. Must be at least max_batch_size or batching degrades. max_inflight = envs.SGLANG_ENCODER_DP_WORKER_MAX_INFLIGHT.get() @@ -1428,6 +1648,7 @@ async def run_dp_worker( ) inflight_sem = asyncio.Semaphore(max_inflight) sched.start() + release_listener_task = asyncio.create_task(listen_for_releases()) logger.info(f"DP worker {dp_rank} ready") try: @@ -1462,12 +1683,30 @@ async def run_dp_worker( spawned = True inflight.add(task) task.add_done_callback(inflight.discard) + if dp_type == "encode": + req_id = request["req_id"] + encode_tasks[req_id] = task + + def forget_encode_task( + completed_task: asyncio.Task, request_id: str = req_id + ) -> None: + if encode_tasks.get(request_id) is completed_task: + encode_tasks.pop(request_id, None) + enc.clear_abandoned_request(request_id) + + task.add_done_callback(forget_encode_task) finally: if not spawned: inflight_sem.release() finally: + release_listener_task.cancel() + with contextlib.suppress(asyncio.CancelledError): + await release_listener_task for task in inflight: task.cancel() + for task in release_tasks: + task.cancel() + await asyncio.gather(*inflight, *release_tasks, return_exceptions=True) ctx.destroy(linger=0) @@ -1476,13 +1715,21 @@ def launch_dp_worker( dp_rank: int, gpu_id: int, dispatch_path: str, + release_path: str, result_path: str, ): publish(server_args, role="encoder") try: configure_logger(server_args, prefix=f" encode_dp_worker[{dp_rank}]") asyncio.run( - run_dp_worker(server_args, dp_rank, gpu_id, dispatch_path, result_path) + run_dp_worker( + server_args, + dp_rank, + gpu_id, + dispatch_path, + release_path, + result_path, + ) ) except KeyboardInterrupt: logger.info(f"DP worker {dp_rank} exiting") @@ -1599,6 +1846,12 @@ def launch_dp_runtime(server_args: ServerArgs) -> DPDispatcher: ) for r in range(dp_size) ] + release_sockets: List[zmq.asyncio.Socket] = [ + get_zmq_socket( + async_zmq_ctx, zmq.PUSH, f"ipc:///tmp/{ipc_prefix}_dp_release_{r}", True + ) + for r in range(dp_size) + ] worker_processes: List[mp.Process] = [] @@ -1628,6 +1881,7 @@ def launch_dp_runtime(server_args: ServerArgs) -> DPDispatcher: dp_rank, gpu_id, f"ipc:///tmp/{ipc_prefix}_dp_dispatch_{dp_rank}", + f"ipc:///tmp/{ipc_prefix}_dp_release_{dp_rank}", result_path, ), daemon=False, @@ -1641,6 +1895,7 @@ def launch_dp_runtime(server_args: ServerArgs) -> DPDispatcher: return DPDispatcher( dp_size, dispatch_sockets, + release_sockets, result_socket, worker_processes, enable_metrics=get_observability().enable_metrics, diff --git a/python/sglang/srt/disaggregation/encoder/server.py b/python/sglang/srt/disaggregation/encoder/server.py index bfc154f94..95f12558a 100644 --- a/python/sglang/srt/disaggregation/encoder/server.py +++ b/python/sglang/srt/disaggregation/encoder/server.py @@ -1,6 +1,7 @@ import asyncio import concurrent.futures import ctypes +import hashlib import logging import os import pickle @@ -87,6 +88,7 @@ rid_to_receive_endpoint: Dict[str, Set[str]] = dict() rid_to_receive_count: Dict[str, int] = dict() cond_dict_lock = asyncio.Lock() rid_to_cond: Dict[str, asyncio.Condition] = {} +encode_state_condition = asyncio.Condition() async def _get_receive_condition(req_id: str) -> asyncio.Condition: @@ -96,6 +98,15 @@ async def _get_receive_condition(req_id: str) -> asyncio.Condition: return rid_to_cond[req_id] +async def _notify_receive_waiters(req_id: str) -> None: + """Wake an existing destination waiter without creating new state.""" + async with cond_dict_lock: + cond = rid_to_cond.get(req_id) + if cond is not None: + async with cond: + cond.notify_all() + + ENCODER_MAX_BATCH_SIZE = envs.SGLANG_ENCODER_MAX_BATCH_SIZE.get() ENCODER_MAX_BATCH_SIZE_EXPLICIT = envs.SGLANG_ENCODER_MAX_BATCH_SIZE.is_set() # Watchdog: max time to wait for a batched /encode result. Bounds HTTP latency @@ -103,6 +114,34 @@ ENCODER_MAX_BATCH_SIZE_EXPLICIT = envs.SGLANG_ENCODER_MAX_BATCH_SIZE.is_set() ENCODER_REQ_TIMEOUT = envs.SGLANG_ENCODER_REQ_TIMEOUT.get() +async def await_task_completion_on_cancel(task: asyncio.Task, operation: str): + """Keep task-owned resources live until cancellation reaches a safe point.""" + try: + return await asyncio.shield(task) + except asyncio.CancelledError: + while not task.done(): + try: + await asyncio.shield(task) + except asyncio.CancelledError: + continue + except Exception: + break + if not task.cancelled() and task.exception() is not None: + logger.error( + "%s failed while draining cancellation", + operation, + exc_info=task.exception(), + ) + raise + + +async def _await_transfer_completion(awaitable, operation: str): + """Do not let cancellation outlive a zero-copy transfer using its buffer.""" + return await await_task_completion_on_cancel( + asyncio.ensure_future(awaitable), operation + ) + + class EncoderMetaRegistry: """Per-part metadata shared by every encoder request lifecycle. @@ -117,9 +156,10 @@ class EncoderMetaRegistry: # Backstop for state whose /send calls never all land. self.sweep_timeout = sweep_timeout self._rid_to_meta: Dict[str, dict] = {} - self._rid_to_send_done: Dict[str, int] = {} + self._rid_to_send_done: Dict[str, Set[str]] = {} self._pending_at: Dict[str, float] = {} self._sweeper_task: Optional[asyncio.Task] = None + self._stale_release_tasks: Dict[str, asyncio.Task] = {} # Set only where the embedding also lives; None in the DP main process. self.on_release: Optional[Callable[[str], Awaitable[None]]] = None @@ -146,9 +186,36 @@ class EncoderMetaRegistry: rid for rid, ts in self._pending_at.items() if now - ts > self.sweep_timeout + and rid not in self._stale_release_tasks ] for rid in stale: - await self._release(rid) + self._schedule_stale_release(rid) + + def _schedule_stale_release(self, req_id: str) -> asyncio.Task: + """Release one stale request without blocking cleanup of other requests.""" + if task := self._stale_release_tasks.get(req_id): + return task + task = asyncio.create_task(self._release_stale(req_id)) + self._stale_release_tasks[req_id] = task + task.add_done_callback( + lambda done, rid=req_id: self._finish_stale_release(rid, done) + ) + return task + + def _finish_stale_release(self, req_id: str, task: asyncio.Task) -> None: + if self._stale_release_tasks.get(req_id) is task: + self._stale_release_tasks.pop(req_id) + + async def _release_stale(self, req_id: str) -> None: + try: + await self._release(req_id) + except Exception: + logger.exception("Failed to release stale encoder request %s", req_id) + # Keep the request eligible for a later sweep without retrying in a + # tight loop. Its metadata and buffer ownership remain intact. + async with rid_lock: + if req_id in self._pending_at: + self._pending_at[req_id] = time.monotonic() async def publish( self, @@ -187,12 +254,15 @@ class EncoderMetaRegistry: ) return self._rid_to_meta.get(req_id) - async def note_send_done(self, req_id: str, receive_count: int) -> None: - """Count one completed ``/send``; release everything at receive_count.""" + async def note_send_done( + self, req_id: str, receive_count: int, destination_endpoint: str + ) -> None: + """Count one destination once; release after every receiver has sent.""" async with rid_lock: - count = self._rid_to_send_done.get(req_id, 0) + 1 - self._rid_to_send_done[req_id] = count - if count >= receive_count: + completed = self._rid_to_send_done.setdefault(req_id, set()) + completed.add(destination_endpoint) + all_done = len(completed) >= receive_count + if all_done: await self._release(req_id) async def _release(self, req_id: str) -> None: @@ -249,6 +319,36 @@ class EncodeContext(msgspec.Struct): is_health_check: bool +def _preprocess_layout_digest(ctx: EncodeContext) -> tuple[int, int]: + """Hash metadata that must agree before TP ranks enter model forward.""" + + def normalize(value): + if isinstance(value, torch.Tensor): + value = value.detach().cpu().numpy() + if isinstance(value, np.ndarray): + return ( + str(value.dtype), + tuple(value.shape), + tuple(value.reshape(-1).tolist()), + ) + if isinstance(value, (list, tuple)): + return tuple(normalize(item) for item in value) + if isinstance(value, np.generic): + return value.item() + return value + + signature = ( + tuple(ctx.items_per_req), + tuple(ctx.preprocess_result.token_counts), + normalize(ctx.preprocess_result.grid_thw), + ) + digest = hashlib.blake2b(pickle.dumps(signature), digest_size=16).digest() + return ( + int.from_bytes(digest[:8], byteorder="little", signed=True), + int.from_bytes(digest[8:], byteorder="little", signed=True), + ) + + @dataclass class ReqState: """The result and in-flight work for one encoder request.""" @@ -582,6 +682,9 @@ class MMEncoder: ) self.req_states: Dict[str, ReqState] = {} + # A DP caller can disappear before its encode creates ReqState. + # Preserve that release intent until _acquire_encode_ref runs. + self.abandoned_req_ids: Set[str] = set() # Need to ensure the NCCL launch order on rank0 matches the dispatch order rank>0 self.encode_dispatch_lock = asyncio.Lock() @@ -641,8 +744,22 @@ class MMEncoder: state = ReqState(req_id) self.req_states[req_id] = state state.active_encodes += 1 + if req_id in self.abandoned_req_ids: + state.release_requested = True + self.abandoned_req_ids.discard(req_id) return state + async def abandon_request(self, req_id: str) -> None: + """Release now, or remember the release until encode state exists.""" + self.abandoned_req_ids.add(req_id) + if req_id in self.req_states: + self.abandoned_req_ids.discard(req_id) + await self.release_request(req_id) + + def clear_abandoned_request(self, req_id: str) -> None: + """Drop an unused release marker after the worker task exits.""" + self.abandoned_req_ids.discard(req_id) + async def _release_encode_ref(self, state: Optional[ReqState]) -> None: if state is None: return @@ -718,8 +835,16 @@ class MMEncoder: async with state.lifecycle_condition: state.release_requested = True state.preserve_metadata_on_release |= preserve_metadata - if state.active_encodes > 0: - return + encode_is_active = state.active_encodes > 0 + + # ``send_with_url`` may be waiting for a destination that will never + # arrive after its HTTP caller disappears. Wake it so the worker slot + # is retired together with the staged embedding. + await _notify_receive_waiters(req_id) + if encode_is_active: + return + + async with state.lifecycle_condition: await state.lifecycle_condition.wait_for(lambda: state.active_sends == 0) if self.req_states.get(req_id) is not state: return @@ -735,21 +860,43 @@ class MMEncoder: expected_destination_count: int, destination_urls: Iterable[str], ) -> None: - async with rid_lock: - if req_id not in rid_to_receive_endpoint: - rid_to_receive_endpoint[req_id] = set() - rid_to_receive_count[req_id] = expected_destination_count - registered_count = rid_to_receive_count[req_id] - if registered_count != expected_destination_count: - raise BadRequestError( - f"Inconsistent receive_count for req_id={req_id}: " - f"registered {registered_count}, got {expected_destination_count}" - ) - rid_to_receive_endpoint[req_id].update(destination_urls) + state = self.req_states.get(req_id) + if state is None: + # registration can beat /encode or its queued batch; only encode creates state + try: + async with encode_state_condition: + await asyncio.wait_for( + encode_state_condition.wait_for( + lambda: req_id in self.req_states + ), + timeout=ENCODER_REQ_TIMEOUT, + ) + except asyncio.TimeoutError as exc: + raise MMError( + f"Timed out waiting for encoder request to start: {req_id}", + code=HTTPStatus.GATEWAY_TIMEOUT, + ) from exc + state = self.req_states.get(req_id) + if state is None: + raise BadRequestError(f"Encoder request is not active: {req_id}") - cond = await _get_receive_condition(req_id) - async with cond: - cond.notify_all() + async with state.lifecycle_condition: + if self.req_states.get(req_id) is not state or state.release_requested: + raise BadRequestError(f"Encoder request is not active: {req_id}") + async with rid_lock: + if req_id not in rid_to_receive_endpoint: + rid_to_receive_endpoint[req_id] = set() + rid_to_receive_count[req_id] = expected_destination_count + registered_count = rid_to_receive_count[req_id] + if registered_count != expected_destination_count: + raise BadRequestError( + f"Inconsistent receive_count for req_id={req_id}: " + f"registered {registered_count}, got {expected_destination_count}" + ) + rid_to_receive_endpoint[req_id].update(destination_urls) + cond = await _get_receive_condition(req_id) + async with cond: + cond.notify_all() def _infer_embedding_dims(self) -> dict: """Infer per-modality embedding dimensions from hf_config at init time.""" @@ -974,10 +1121,14 @@ class MMEncoder: preprocess_result, items_per_req, ) = await self.preprocessor.process_batch_mm_items(requests, modality) + except MMError: + raise except NotImplementedError as e: raise InternalError(f"Not implemented error: {str(e)}") - except Exception as e: + except (TypeError, ValueError) as e: raise BadRequestError(f"Failed to process mm items: {str(e)}") + except Exception as e: + raise InternalError(f"Failed to process mm items: {str(e)}") if len(items_per_req) != len(requests) or any(n <= 0 for n in items_per_req): raise InternalError( @@ -1053,6 +1204,134 @@ class MMEncoder: is_health_check=is_health_check, ) + async def _prepare_encode_context_on_all_ranks( + self, + requests: List[dict], + modality: Modality, + *, + use_global_cache: bool, + is_health_check: bool = False, + ) -> EncodeContext: + """Prepare one context consistently before TP ranks enter model forward.""" + ctx = None + local_error = None + error_phase = 0 + try: + ctx = await self._prepare_encode_context( + requests, + modality, + use_global_cache=use_global_cache, + is_health_check=is_health_check, + ) + except Exception as e: + local_error = e + error_phase = 1 + + if local_error is None: + try: + assert ctx is not None + await self._publish_preprocess_metadata(ctx, requests) + except Exception as e: + local_error = e + error_phase = 2 + + if local_error is None: + assert ctx is not None + layout_digest = _preprocess_layout_digest(ctx) + else: + layout_digest = (0, 0) + statuses = self._sync_tp_prepare_status( + local_error, + error_phase=error_phase, + layout_digest=layout_digest, + ) + + expected_layout = tuple(statuses[0][2:].tolist()) + mismatch_rank = next( + ( + rank + for rank, rank_status in enumerate(statuses[1:], start=1) + if tuple(rank_status[2:].tolist()) != expected_layout + ), + None, + ) + if mismatch_rank is not None: + raise InternalError( + "Encoder preprocessing produced inconsistent layouts across TP " + f"ranks 0 and {mismatch_rank}" + ) + + assert ctx is not None + return ctx + + def _sync_tp_prepare_status( + self, + local_error: Optional[Exception], + *, + error_phase: int, + layout_digest: tuple[int, int], + ) -> List[torch.Tensor]: + """Raise the same preparation error on every TP rank.""" + tp_group = get_tp_group() + error_code = ( + int( + local_error.code + if isinstance(local_error, MMError) + else HTTPStatus.INTERNAL_SERVER_ERROR + ) + if local_error is not None + else 0 + ) + local_status = torch.tensor( + [error_code, error_phase, *layout_digest], dtype=torch.int64 + ) + statuses = [torch.empty_like(local_status) for _ in range(tp_group.world_size)] + if tp_group.world_size > 1: + torch.distributed.all_gather( + statuses, + local_status, + group=tp_group.cpu_group, + ) + else: + statuses[0].copy_(local_status) + + failures = [ + ( + rank, + int(rank_status[0].item()), + int(rank_status[1].item()), + ) + for rank, rank_status in enumerate(statuses) + if rank_status[0].item() != 0 + ] + if not failures: + return statuses + + errors = ( + tp_group.all_gather_object( + str(local_error) if local_error is not None else None + ) + if tp_group.world_size > 1 + else [str(local_error)] + ) + rank, failure_code, failure_phase = next( + ( + (rank, rank_error_code, rank_error_phase) + for rank, rank_error_code, rank_error_phase in failures + if rank_error_code != HTTPStatus.BAD_REQUEST + ), + failures[0], + ) + phase = ( + "Encoder metadata publication" + if failure_phase == 2 + else "Encoder preprocessing" + ) + message = f"{phase} failed on TP rank {rank}: {errors[rank]}" + if failure_code == HTTPStatus.BAD_REQUEST: + raise BadRequestError(message) + raise InternalError(message) + def _broadcast_global_cache_mask(self, mask_tensor: torch.Tensor): if get_parallel().tp_size > 1: torch.distributed.broadcast( @@ -1737,7 +2016,10 @@ class MMEncoder: # Queue sends in order under the lock, then wait for buffer # ownership independently so libzmq can pipeline the connection. try: - await asyncio.to_thread(tracker.wait, self.send_timeout) + await _await_transfer_completion( + asyncio.to_thread(tracker.wait, self.send_timeout), + f"ZMQ transfer for req_id={mm_data.req_id}", + ) except Exception: if self.scheduler_send_sockets.get(endpoint) is sock: self.scheduler_send_sockets.pop(endpoint, None) @@ -1772,7 +2054,10 @@ class MMEncoder: finally: sock.close(linger=5000) - await asyncio.get_event_loop().run_in_executor(self.executor, send_with_socket) + await _await_transfer_completion( + asyncio.get_running_loop().run_in_executor(self.executor, send_with_socket), + f"ZMQ transfer for req_id={mm_data.req_id}", + ) if ( encoder_metrics_collector is not None and get_disagg().encoder_transfer_backend != "mooncake" @@ -1790,23 +2075,16 @@ class MMEncoder: size: int, ) -> int: """Keep the send active until its blocking transfer stops using the MR.""" - transfer_task = asyncio.create_task( + return await _await_transfer_completion( asyncio.to_thread( self.engine.transfer_sync, session_id, source_address, destination_address, size, - ) + ), + f"Mooncake transfer to session={session_id}", ) - try: - return await asyncio.shield(transfer_task) - except asyncio.CancelledError: - try: - await transfer_task - except Exception: - pass - raise def _register_shared_mr(self, mm_data: EmbeddingData, embedding: torch.Tensor): """Register one MR shared by every rank's /send; _send re-registers on failure.""" @@ -1943,13 +2221,15 @@ class MMEncoder: keep_on_gpu = self.use_mooncake and not is_health_check use_global_cache = self.mm_global_cache is not None and not is_health_check try: - ctx = await self._prepare_encode_context( + if self.rank == 0: + async with encode_state_condition: + encode_state_condition.notify_all() + ctx = await self._prepare_encode_context_on_all_ranks( requests, modality, use_global_cache=use_global_cache, is_health_check=is_health_check, ) - await self._publish_preprocess_metadata(ctx, requests) mm_embedding = await self._compute_embedding(ctx, keep_on_gpu=keep_on_gpu) if self.profiler is not None: @@ -2034,6 +2314,9 @@ class MMEncoder: try: while True: + if state.release_requested: + break + async with rid_lock: current_targets = rid_to_receive_endpoint.get(req_id, set()).copy() expected_count = rid_to_receive_count.get(req_id) diff --git a/test/registered/openai_server/basic/test_openai_completion_rust.py b/test/registered/openai_server/basic/test_openai_completion_rust.py index 87c698098..df77a893c 100644 --- a/test/registered/openai_server/basic/test_openai_completion_rust.py +++ b/test/registered/openai_server/basic/test_openai_completion_rust.py @@ -26,8 +26,7 @@ class TestOpenAICompletionRustParity(CustomTestCase): api_key = "sk-123456" def _get_logprobs(self, *, rust_frontend): - # Prefill CUDA graph pads the batch, so numerics follow whichever - # requests share the forward pass; the assertions below need equality. + # compare identical prefill shapes, without graph padding or warmup cache hits process = popen_launch_server( self.model, DEFAULT_URL_FOR_TEST, @@ -38,6 +37,7 @@ class TestOpenAICompletionRustParity(CustomTestCase): "--random-seed", "42", "--disable-prefill-cuda-graph", + "--disable-radix-cache", ], ) try: diff --git a/test/registered/unit/disaggregation/test_encode_server.py b/test/registered/unit/disaggregation/test_encode_server.py index 5994c2394..2283a3da9 100644 --- a/test/registered/unit/disaggregation/test_encode_server.py +++ b/test/registered/unit/disaggregation/test_encode_server.py @@ -2,23 +2,38 @@ import asyncio import pickle import threading import unittest +from http import HTTPStatus from types import SimpleNamespace from unittest.mock import AsyncMock, Mock, patch import numpy as np import torch +import sglang.srt.disaggregation.encoder.server as encoder_server +from sglang.srt.disaggregation.encoder import http_server +from sglang.srt.disaggregation.encoder import runtime as encoder_runtime from sglang.srt.disaggregation.encoder.preprocessor import EncoderPreprocessor from sglang.srt.disaggregation.encoder.receiver import EmbeddingData -from sglang.srt.disaggregation.encoder.runtime import execute_encode_pipeline +from sglang.srt.disaggregation.encoder.runtime import ( + _DP_RELEASE_AFTER_ENCODE, + DPDispatcher, + _retire_abandoned_encode, + execute_encode_pipeline, + send_staged_embedding, +) from sglang.srt.disaggregation.encoder.server import ( + BadRequestError, + EncodeContext, EncoderDelivery, + EncoderMetaRegistry, InternalError, MMEncoder, + MMError, MooncakeDelivery, ReqState, SendDestination, ZmqDelivery, + _await_transfer_completion, meta_registry, rid_to_cond, rid_to_receive_count, @@ -27,6 +42,7 @@ from sglang.srt.disaggregation.encoder.server import ( from sglang.srt.distributed.device_communicators.mooncake_transfer_engine import ( MooncakeTransferEngine, ) +from sglang.srt.managers.io_struct import unwrap_from_pickle from sglang.srt.managers.schedule_batch import Modality from sglang.srt.mem_cache.multimodal_cache import ( EmbeddingResult, @@ -39,6 +55,216 @@ from sglang.test.test_utils import CustomTestCase register_cpu_ci(est_time=2, suite="base-a-test-cpu") +class TestEncoderDPErrorHandling(CustomTestCase): + @staticmethod + async def _run_registration_error(error): + encoder = SimpleNamespace( + register_embedding_destinations=AsyncMock(side_effect=error) + ) + send = AsyncMock() + request = { + "req_id": "req", + "receive_count": 1, + "receive_url": "tcp://127.0.0.1:1", + } + with patch.object(encoder_runtime, "async_sock_send", send): + await encoder_runtime._dp_worker_handle_request( + encoder, + None, + object(), + asyncio.Lock(), + 0, + request, + "register_destinations", + ) + return unwrap_from_pickle(send.await_args.args[1]) + + def test_worker_reports_third_party_exception_with_callable_code(self): + class RpcLikeError(Exception): + def code(self): + return "INTERNAL" + + envelope = asyncio.run( + self._run_registration_error(RpcLikeError("registration failed")) + ) + self.assertEqual(envelope["_error"], "registration failed") + self.assertEqual(envelope["_error_code"], 500) + + def test_worker_preserves_mm_error_status(self): + envelope = asyncio.run( + self._run_registration_error(MMError("bad destination", code=400)) + ) + self.assertEqual(envelope["_error_code"], 400) + + def test_dispatcher_drops_malformed_result_without_stopping_listener(self): + async def run(): + dispatcher = encoder_runtime.DPDispatcher( + dp_size=1, + dispatch_sockets=[object()], + release_sockets=[object()], + result_socket=object(), + worker_processes=[], + ) + future = asyncio.get_running_loop().create_future() + dispatcher.pending_futures[0]["req"] = future + dispatcher.req_id_to_rank["req"] = 0 + valid = {"req_id": "req", "_dp_type": "encode", "content": None} + recv = AsyncMock( + side_effect=[ + ["not", "an", "envelope"], + valid, + asyncio.CancelledError(), + ] + ) + + with patch.object(encoder_runtime, "async_sock_recv", recv): + listener = asyncio.create_task(dispatcher._result_listener()) + await asyncio.wait_for(future, timeout=1) + listener.cancel() + with self.assertRaises(asyncio.CancelledError): + await listener + + self.assertEqual(future.result(), valid) + + asyncio.run(run()) + + +class TestEncoderMetaRegistry(CustomTestCase): + def test_stale_releases_do_not_block_each_other(self): + async def run(): + registry = EncoderMetaRegistry(wait_timeout=1, sweep_timeout=1) + blocked_started = asyncio.Event() + unblock = asyncio.Event() + fast_released = asyncio.Event() + + async def release(req_id): + if req_id == "blocked": + blocked_started.set() + await unblock.wait() + else: + fast_released.set() + + registry.on_release = release + registry._pending_at.update(blocked=0, fast=0) + + blocked_task = registry._schedule_stale_release("blocked") + await asyncio.wait_for(blocked_started.wait(), timeout=1) + fast_task = registry._schedule_stale_release("fast") + await asyncio.wait_for(fast_released.wait(), timeout=1) + await fast_task + + self.assertIn("blocked", registry._pending_at) + self.assertNotIn("fast", registry._pending_at) + + unblock.set() + await blocked_task + self.assertNotIn("blocked", registry._pending_at) + + asyncio.run(run()) + + def test_failed_stale_release_is_retried(self): + async def run(): + registry = EncoderMetaRegistry(wait_timeout=1, sweep_timeout=1) + attempts = 0 + + async def release(_req_id): + nonlocal attempts + attempts += 1 + if attempts == 1: + raise RuntimeError("transient cleanup failure") + + registry.on_release = release + registry._pending_at["req"] = 0 + + await registry._release_stale("req") + self.assertIn("req", registry._pending_at) + retry_at = registry._pending_at["req"] + self.assertGreater(retry_at, 0) + + await registry._release_stale("req") + self.assertNotIn("req", registry._pending_at) + self.assertEqual(attempts, 2) + + asyncio.run(run()) + + def test_send_retries_do_not_release_before_all_destinations_finish(self): + async def run(): + registry = EncoderMetaRegistry(wait_timeout=1, sweep_timeout=1) + released = AsyncMock() + registry.on_release = released + + await registry.note_send_done("req", 2, "10.0.0.1:5000") + await registry.note_send_done("req", 2, "10.0.0.1:5000") + released.assert_not_awaited() + + await registry.note_send_done("req", 2, "10.0.0.2:5000") + released.assert_awaited_once_with("req") + + asyncio.run(run()) + + def test_http_send_counts_the_normalized_destination(self): + async def run(): + send = AsyncMock(return_value=True) + note_send_done = AsyncMock() + request = { + "req_id": "req", + "prefill_host": "127.0.0.1", + "embedding_port": 5000, + "session_id": "session", + "buffer_address": 1234, + "receive_count": 2, + } + with ( + patch.object(http_server, "dp_dispatcher", None), + patch.object(http_server, "encoder", SimpleNamespace(send=send)), + patch.object( + encoder_server.meta_registry, + "note_send_done", + note_send_done, + ), + ): + response = await http_server.handle_send_request(request) + + self.assertEqual(response.status_code, 200) + note_send_done.assert_awaited_once_with("req", 2, "127.0.0.1:5000") + + asyncio.run(run()) + + def test_dp_send_counts_the_normalized_destination(self): + async def run(): + encoder = SimpleNamespace(send=AsyncMock(return_value=True)) + note_send_done = AsyncMock() + request = { + "req_id": "req", + "prefill_host": "127.0.0.1", + "embedding_port": 5000, + "session_id": "session", + "buffer_address": 1234, + "receive_count": 2, + } + with ( + patch.object(encoder_runtime, "async_sock_send", AsyncMock()), + patch.object( + encoder_server.meta_registry, + "note_send_done", + note_send_done, + ), + ): + await encoder_runtime._dp_worker_handle_request( + encoder, + None, + object(), + asyncio.Lock(), + 0, + request, + "send", + ) + + note_send_done.assert_awaited_once_with("req", 2, "127.0.0.1:5000") + + asyncio.run(run()) + + class TestEncoderPreprocessorKimiGrid(CustomTestCase): @staticmethod def _make_preprocessor(model_type="kimi_vl"): @@ -138,6 +364,155 @@ class TestEncoderPreprocessorKimiGrid(CustomTestCase): class TestEncoderDelivery(CustomTestCase): + def test_cancelled_zero_copy_transfer_drains_before_return(self): + async def run(): + transfer_started = threading.Event() + finish_transfer = threading.Event() + + def transfer(): + transfer_started.set() + finish_transfer.wait() + + task = asyncio.create_task( + _await_transfer_completion(asyncio.to_thread(transfer), "test transfer") + ) + while not transfer_started.is_set(): + await asyncio.sleep(0) + + task.cancel() + await asyncio.sleep(0) + self.assertFalse(task.done()) + + task.cancel() + await asyncio.sleep(0) + self.assertFalse(task.done()) + + finish_transfer.set() + with self.assertRaises(asyncio.CancelledError): + await task + + asyncio.run(run()) + + def test_cancelled_mooncake_send_keeps_embedding_until_transfer_stops(self): + async def run(): + transfer_started = threading.Event() + finish_transfer = threading.Event() + + def transfer_sync(*_args): + transfer_started.set() + finish_transfer.wait() + return 0 + + encoder = MMEncoder.__new__(MMEncoder) + encoder.req_states = {} + encoder._element_size = 2 + encoder.transfer_backend = "mooncake" + encoder.engine = SimpleNamespace( + register=unittest.mock.Mock(), + transfer_sync=unittest.mock.Mock(side_effect=transfer_sync), + deregister=unittest.mock.Mock(), + ) + encoder.delivery = MooncakeDelivery(encoder) + + embedding = torch.ones((2, 4), dtype=torch.float16) + state = ReqState( + "cancelled-transfer", + EmbeddingData( + "cancelled-transfer", + 1, + 0, + None, + Modality.IMAGE, + embedding=embedding, + ), + ) + state.embedding_ready.set() + encoder.req_states[state.req_id] = state + + with ( + patch( + "sglang.srt.disaggregation.encoder.server.get_disagg", + return_value=SimpleNamespace(encoder_transfer_backend="mooncake"), + ), + patch.object(meta_registry, "discard", AsyncMock()), + ): + send_task = asyncio.create_task( + encoder.send_to_destination( + state, + SendDestination( + "127.0.0.1:1", session_id="session", buffer_address=1 + ), + ) + ) + while not transfer_started.is_set(): + await asyncio.sleep(0) + + send_task.cancel() + release_task = asyncio.create_task( + encoder.release_request(state.req_id) + ) + await asyncio.sleep(0) + + self.assertFalse(send_task.done()) + self.assertFalse(release_task.done()) + self.assertIs(state.embedding_data.embedding, embedding) + encoder.engine.deregister.assert_not_called() + + finish_transfer.set() + with self.assertRaises(asyncio.CancelledError): + await send_task + await release_task + + encoder.engine.register.assert_called_once_with( + embedding.data_ptr(), embedding.nbytes + ) + encoder.engine.deregister.assert_called_once_with(embedding.data_ptr()) + self.assertIsNone(state.embedding_data) + self.assertNotIn(state.req_id, encoder.req_states) + + asyncio.run(run()) + + def test_failed_mooncake_transfer_releases_per_send_registration(self): + async def run(): + encoder = MMEncoder.__new__(MMEncoder) + encoder._element_size = 2 + encoder.transfer_backend = "mooncake" + encoder.engine = SimpleNamespace( + register=unittest.mock.Mock(), + transfer_sync=unittest.mock.Mock( + side_effect=RuntimeError("transfer failed") + ), + deregister=unittest.mock.Mock(), + ) + embedding = torch.ones((2, 4), dtype=torch.float16) + mm_data = EmbeddingData( + "failed-transfer", + 1, + 0, + None, + Modality.IMAGE, + embedding=embedding, + ) + + with patch( + "sglang.srt.disaggregation.encoder.server.get_disagg", + return_value=SimpleNamespace(encoder_transfer_backend="mooncake"), + ): + with self.assertRaisesRegex(RuntimeError, "transfer failed"): + await encoder._send( + embedding, + mm_data, + session_id="session", + buffer_address=1, + ) + + encoder.engine.register.assert_called_once_with( + embedding.data_ptr(), embedding.nbytes + ) + encoder.engine.deregister.assert_called_once_with(embedding.data_ptr()) + + asyncio.run(run()) + @staticmethod def _global_cache_context(num_items=2): return SimpleNamespace( @@ -271,6 +646,112 @@ class TestEncoderDelivery(CustomTestCase): }, ) + def test_failed_staged_send_releases_request(self): + async def run(): + encoder = SimpleNamespace( + send=AsyncMock(side_effect=RuntimeError("transfer failed")), + release_request=AsyncMock(), + ) + request = { + "req_id": "req", + "prefill_host": "127.0.0.1", + "embedding_port": 1, + "session_id": "session", + "buffer_address": 2, + } + + with self.assertRaisesRegex(RuntimeError, "transfer failed"): + await send_staged_embedding( + encoder, request, release_without_count=False + ) + + encoder.release_request.assert_awaited_once_with("req") + + asyncio.run(run()) + + def test_cleanup_failure_preserves_send_error_on_python_310(self): + async def run(): + encoder = SimpleNamespace( + send=AsyncMock(side_effect=ValueError("transfer failed")), + release_request=AsyncMock(side_effect=RuntimeError("cleanup failed")), + ) + request = { + "req_id": "req", + "prefill_host": "127.0.0.1", + "embedding_port": 1, + "session_id": "session", + "buffer_address": 2, + } + + with ( + patch.object(encoder_runtime.sys, "version_info", (3, 10)), + self.assertLogs(encoder_runtime.logger, level="ERROR"), + self.assertRaisesRegex(ValueError, "transfer failed"), + ): + await send_staged_embedding( + encoder, request, release_without_count=False + ) + + asyncio.run(run()) + + def test_cancelled_staged_send_releases_request(self): + async def run(): + encoder = SimpleNamespace( + send=AsyncMock(side_effect=asyncio.CancelledError()), + release_request=AsyncMock(), + ) + request = { + "req_id": "req", + "prefill_host": "127.0.0.1", + "embedding_port": 1, + "session_id": "session", + "buffer_address": 2, + } + + with self.assertRaises(asyncio.CancelledError): + await send_staged_embedding( + encoder, request, release_without_count=False + ) + + encoder.release_request.assert_awaited_once_with("req") + + asyncio.run(run()) + + def test_staged_send_uses_refcount_or_legacy_release_policy(self): + async def run(): + request = { + "req_id": "req", + "prefill_host": "127.0.0.1", + "embedding_port": 1, + "session_id": "session", + "buffer_address": 2, + "receive_count": 2, + } + encoder = SimpleNamespace( + send=AsyncMock(return_value=True), + release_request=AsyncMock(), + ) + + note_send_done = AsyncMock() + with patch.object(meta_registry, "note_send_done", note_send_done): + self.assertTrue( + await send_staged_embedding( + encoder, request, release_without_count=True + ) + ) + note_send_done.assert_awaited_once_with("req", 2, "127.0.0.1:1") + encoder.release_request.assert_not_awaited() + + request.pop("receive_count") + self.assertTrue( + await send_staged_embedding( + encoder, request, release_without_count=True + ) + ) + encoder.release_request.assert_awaited_once_with("req") + + asyncio.run(run()) + @staticmethod def _make_mooncake_send(engine): embedding = torch.zeros((2, 4), dtype=torch.float32) @@ -370,6 +851,11 @@ class TestEncoderDelivery(CustomTestCase): self.assertFalse(send_task.done()) self.assertNotIn("deregister", events) + send_task.cancel() + await asyncio.sleep(0) + self.assertFalse(send_task.done()) + self.assertNotIn("deregister", events) + finish_transfer.set() with self.assertRaises(asyncio.CancelledError): await send_task @@ -407,6 +893,7 @@ class TestEncoderDelivery(CustomTestCase): encoder = MMEncoder.__new__(MMEncoder) encoder.rank = 0 encoder.req_states = {} + encoder.abandoned_req_ids = set() encoder._embedding_dims = {Modality.IMAGE: 8} encoder._embedding_dtype = torch.float16 encoder._element_size = 2 @@ -455,6 +942,504 @@ class TestEncoderDelivery(CustomTestCase): asyncio.run(run()) + @staticmethod + def _encode_context(): + return EncodeContext( + req_id="req", + modality=Modality.IMAGE, + preprocess_result=SimpleNamespace( + token_counts=[2], + grid_thw=torch.tensor([[1, 2, 4]]), + ), + get_feature_fn=None, + mm_feature=torch.zeros((8, 3)), + num_items=1, + items_per_req=[1], + aux_data={}, + str_mm_hashes=None, + use_global_cache=False, + is_health_check=False, + ) + + @staticmethod + def _load_grpc_server(): + try: + import grpc + from grpc_health.v1 import health_pb2 + from smg_grpc_proto import sglang_encoder_pb2 + + health_pb2.HealthCheckRequest() + except ImportError as e: + raise unittest.SkipTest(f"gRPC test dependencies unavailable: {e}") from e + except Exception as e: + # Generated protobuf modules raise VersionError when the runner's + # protobuf runtime is older than the code generator. + if not ( + type(e).__module__ == "google.protobuf.runtime_version" + and type(e).__name__ == "VersionError" + ): + raise + raise unittest.SkipTest(f"gRPC test dependencies unavailable: {e}") from e + + # Import SGLang outside the dependency guard. Product-code import + # failures are regressions and must fail the test instead of skipping. + from sglang.srt.disaggregation.encoder.grpc_server import ( + SGLangEncoderServer, + ) + + return grpc, sglang_encoder_pb2, SGLangEncoderServer + + def test_remote_preprocess_failure_stops_all_tp_ranks_before_forward(self): + async def run(): + encoder = MMEncoder.__new__(MMEncoder) + encoder._prepare_encode_context = AsyncMock( + return_value=self._encode_context() + ) + encoder._publish_preprocess_metadata = AsyncMock() + + class TPGroup: + world_size = 2 + cpu_group = object() + + @staticmethod + def all_gather_object(local_error): + return [local_error, "bad image"] + + def all_gather(statuses, local_status, group): + self.assertIs(group, TPGroup.cpu_group) + statuses[0].copy_(local_status) + statuses[1].copy_(torch.tensor([400, 1, 0, 0])) + + with ( + patch( + "sglang.srt.disaggregation.encoder.server.get_tp_group", + return_value=TPGroup(), + ), + patch( + "sglang.srt.disaggregation.encoder.server.torch.distributed.all_gather", + side_effect=all_gather, + ), + ): + with self.assertRaisesRegex( + BadRequestError, + "failed on TP rank 1: bad image", + ): + await encoder._prepare_encode_context_on_all_ranks( + [{"req_id": "req"}], + Modality.IMAGE, + use_global_cache=False, + ) + + asyncio.run(run()) + + def test_tp_preprocess_layout_mismatch_fails_before_forward(self): + async def run(): + encoder = MMEncoder.__new__(MMEncoder) + encoder._prepare_encode_context = AsyncMock( + return_value=self._encode_context() + ) + encoder._publish_preprocess_metadata = AsyncMock() + + class TPGroup: + world_size = 2 + cpu_group = object() + + def all_gather(statuses, local_status, group): + self.assertIs(group, TPGroup.cpu_group) + statuses[0].copy_(local_status) + statuses[1].copy_(local_status) + statuses[1][2] += 1 + + with ( + patch( + "sglang.srt.disaggregation.encoder.server.get_tp_group", + return_value=TPGroup(), + ), + patch( + "sglang.srt.disaggregation.encoder.server.torch.distributed.all_gather", + side_effect=all_gather, + ), + ): + with self.assertRaisesRegex( + InternalError, + "inconsistent layouts across TP ranks 0 and 1", + ): + await encoder._prepare_encode_context_on_all_ranks( + [{"req_id": "req"}], + Modality.IMAGE, + use_global_cache=False, + ) + + asyncio.run(run()) + + def test_remote_metadata_failure_stops_tp_peer_before_forward(self): + async def run(): + encoder = MMEncoder.__new__(MMEncoder) + encoder._prepare_encode_context = AsyncMock( + return_value=self._encode_context() + ) + encoder._publish_preprocess_metadata = AsyncMock() + + class TPGroup: + world_size = 2 + cpu_group = object() + + @staticmethod + def all_gather_object(local_error): + self.assertIsNone(local_error) + return ["registry down", None] + + def all_gather(statuses, local_status, group): + self.assertIs(group, TPGroup.cpu_group) + statuses[0].copy_(torch.tensor([500, 2, 0, 0])) + statuses[1].copy_(local_status) + + with ( + patch( + "sglang.srt.disaggregation.encoder.server.get_tp_group", + return_value=TPGroup(), + ), + patch( + "sglang.srt.disaggregation.encoder.server.torch.distributed.all_gather", + side_effect=all_gather, + ), + ): + with self.assertRaisesRegex( + InternalError, + "metadata publication failed on TP rank 0: registry down", + ): + await encoder._prepare_encode_context_on_all_ranks( + [{"req_id": "req"}], + Modality.IMAGE, + use_global_cache=False, + ) + + asyncio.run(run()) + + def test_unexpected_preprocess_failure_is_internal(self): + async def run(): + encoder = MMEncoder.__new__(MMEncoder) + encoder.preprocessor = SimpleNamespace( + process_batch_mm_items=AsyncMock(side_effect=RuntimeError("boom")) + ) + with self.assertRaisesRegex(InternalError, "boom"): + await encoder._prepare_encode_context( + [{"req_id": "req"}], + Modality.IMAGE, + use_global_cache=False, + ) + + asyncio.run(run()) + + def test_grpc_rejects_invalid_request_before_tp_dispatch(self): + async def run(): + grpc, sglang_encoder_pb2, SGLangEncoderServer = self._load_grpc_server() + + context = SimpleNamespace( + set_code=unittest.mock.Mock(), + set_details=unittest.mock.Mock(), + ) + server = SGLangEncoderServer( + encoder=SimpleNamespace(), + send_sockets=[object()], + server_args=SimpleNamespace(), + ) + request = sglang_encoder_pb2.EncodeRequest( + mm_items=["image"], + req_id="invalid", + part_idx=0, + ) + + with patch( + "sglang.srt.disaggregation.encoder.grpc_server.async_sock_send", + new_callable=AsyncMock, + ) as send: + await server.Encode(request, context) + + send.assert_not_awaited() + context.set_code.assert_called_once_with(grpc.StatusCode.INVALID_ARGUMENT) + self.assertIn("num_parts", context.set_details.call_args.args[0]) + + asyncio.run(run()) + + def test_grpc_maps_processor_bad_request_to_invalid_argument(self): + async def run(): + grpc, sglang_encoder_pb2, SGLangEncoderServer = self._load_grpc_server() + + encoder = SimpleNamespace( + encode_dispatch_lock=asyncio.Lock(), + encode_request=AsyncMock( + return_value=( + 0, + 0, + 0, + "invalid image", + HTTPStatus.BAD_REQUEST, + ) + ), + release_request=AsyncMock(), + ) + context = SimpleNamespace( + set_code=unittest.mock.Mock(), + set_details=unittest.mock.Mock(), + ) + server = SGLangEncoderServer( + encoder=encoder, + send_sockets=[], + server_args=SimpleNamespace(), + ) + request = sglang_encoder_pb2.EncodeRequest( + mm_items=["bad-image"], + req_id="bad-image", + num_parts=1, + part_idx=0, + ) + + await server.Encode(request, context) + + context.set_code.assert_called_once_with(grpc.StatusCode.INVALID_ARGUMENT) + context.set_details.assert_called_once_with("invalid image") + encoder.release_request.assert_awaited_once_with("bad-image") + + asyncio.run(run()) + + def test_grpc_serializes_tp_dispatch_with_rank_zero_encode(self): + async def run(): + _, sglang_encoder_pb2, SGLangEncoderServer = self._load_grpc_server() + from sglang.srt.managers.io_struct import unwrap_from_pickle + + first_started = asyncio.Event() + release_first = asyncio.Event() + events = [] + + class Encoder: + def __init__(self): + self.encode_dispatch_lock = asyncio.Lock() + + async def encode_request(self, request, _modality): + req_id = request["req_id"] + events.append(("encode-start", req_id)) + if req_id == "first": + first_started.set() + await release_first.wait() + events.append(("encode-end", req_id)) + return 8, 1, 8, None, None + + async def send(_socket, payload): + request = unwrap_from_pickle(payload) + events.append(("send", request["req_id"])) + + server = SGLangEncoderServer( + encoder=Encoder(), + send_sockets=[object()], + server_args=SimpleNamespace(), + ) + requests = [ + sglang_encoder_pb2.EncodeRequest( + mm_items=["image"], + req_id=req_id, + num_parts=1, + part_idx=0, + ) + for req_id in ("first", "second") + ] + contexts = [ + SimpleNamespace( + set_code=unittest.mock.Mock(), + set_details=unittest.mock.Mock(), + ) + for _ in requests + ] + + with ( + patch( + "sglang.srt.disaggregation.encoder.grpc_server.async_sock_send", + side_effect=send, + ), + patch( + "sglang.srt.disaggregation.encoder.grpc_server.get_disagg", + return_value=SimpleNamespace(encoder_transfer_backend="mooncake"), + ), + ): + first = asyncio.create_task(server.Encode(requests[0], contexts[0])) + await first_started.wait() + second = asyncio.create_task(server.Encode(requests[1], contexts[1])) + await asyncio.sleep(0) + self.assertEqual( + events, + [("send", "first"), ("encode-start", "first")], + ) + release_first.set() + await asyncio.gather(first, second) + + self.assertEqual( + events, + [ + ("send", "first"), + ("encode-start", "first"), + ("encode-end", "first"), + ("send", "second"), + ("encode-start", "second"), + ("encode-end", "second"), + ], + ) + + asyncio.run(run()) + + def test_grpc_encode_cancellation_drains_tp_collective_before_release(self): + async def run(): + _, sglang_encoder_pb2, SGLangEncoderServer = self._load_grpc_server() + + encode_started = asyncio.Event() + finish_encode = asyncio.Event() + + async def encode_request(*_args): + encode_started.set() + await finish_encode.wait() + return 8, 1, 8, None, None + + encoder = SimpleNamespace( + encode_dispatch_lock=asyncio.Lock(), + encode_request=AsyncMock(side_effect=encode_request), + release_request=AsyncMock(), + ) + server = SGLangEncoderServer( + encoder=encoder, + send_sockets=[], + server_args=SimpleNamespace(), + ) + request = sglang_encoder_pb2.EncodeRequest( + mm_items=["image"], + req_id="cancelled-encode", + num_parts=1, + part_idx=0, + ) + context = SimpleNamespace( + set_code=unittest.mock.Mock(), + set_details=unittest.mock.Mock(), + ) + + with patch( + "sglang.srt.disaggregation.encoder.grpc_server.get_disagg", + return_value=SimpleNamespace(encoder_transfer_backend="mooncake"), + ): + task = asyncio.create_task(server.Encode(request, context)) + await encode_started.wait() + task.cancel() + await asyncio.sleep(0) + + # Cancellation cannot interrupt an in-flight TP collective. + self.assertFalse(task.done()) + encoder.release_request.assert_not_awaited() + + task.cancel() + await asyncio.sleep(0) + self.assertFalse(task.done()) + encoder.release_request.assert_not_awaited() + + finish_encode.set() + with self.assertRaises(asyncio.CancelledError): + await task + + encoder.release_request.assert_awaited_once_with("cancelled-encode") + + asyncio.run(run()) + + def test_grpc_encode_cancellation_during_tp_dispatch_completes_encode(self): + async def run(): + _, sglang_encoder_pb2, SGLangEncoderServer = self._load_grpc_server() + + send_started = asyncio.Event() + finish_send = asyncio.Event() + + async def send_to_tp(*_args): + send_started.set() + await finish_send.wait() + + encoder = SimpleNamespace( + encode_dispatch_lock=asyncio.Lock(), + encode_request=AsyncMock(return_value=(8, 1, 8, None, None)), + release_request=AsyncMock(), + ) + server = SGLangEncoderServer( + encoder=encoder, + send_sockets=[object()], + server_args=SimpleNamespace(), + ) + request = sglang_encoder_pb2.EncodeRequest( + mm_items=["image"], + req_id="cancelled-dispatch", + num_parts=1, + part_idx=0, + ) + context = SimpleNamespace( + set_code=unittest.mock.Mock(), + set_details=unittest.mock.Mock(), + ) + + with ( + patch( + "sglang.srt.disaggregation.encoder.grpc_server.async_sock_send", + side_effect=send_to_tp, + ), + patch( + "sglang.srt.disaggregation.encoder.grpc_server.get_disagg", + return_value=SimpleNamespace(encoder_transfer_backend="mooncake"), + ), + ): + task = asyncio.create_task(server.Encode(request, context)) + await send_started.wait() + task.cancel() + await asyncio.sleep(0) + + self.assertFalse(task.done()) + encoder.encode_request.assert_not_awaited() + + finish_send.set() + with self.assertRaises(asyncio.CancelledError): + await task + + encoder.encode_request.assert_awaited_once() + encoder.release_request.assert_awaited_once_with("cancelled-dispatch") + + asyncio.run(run()) + + def test_grpc_send_cancellation_releases_request(self): + async def run(): + _, sglang_encoder_pb2, SGLangEncoderServer = self._load_grpc_server() + + send_started = asyncio.Event() + + async def send(*_args, **_kwargs): + send_started.set() + await asyncio.Event().wait() + + encoder = SimpleNamespace( + send=AsyncMock(side_effect=send), + release_request=AsyncMock(), + ) + server = SGLangEncoderServer( + encoder=encoder, + send_sockets=[], + server_args=SimpleNamespace(), + ) + request = sglang_encoder_pb2.SendRequest( + req_id="cancelled-send", + prefill_host="127.0.0.1", + embedding_port=30001, + ) + context = SimpleNamespace() + + task = asyncio.create_task(server.Send(request, context)) + await send_started.wait() + task.cancel() + with self.assertRaises(asyncio.CancelledError): + await task + + encoder.release_request.assert_awaited_once_with("cancelled-send") + + asyncio.run(run()) + def test_global_cache_lookup_failure_falls_back_to_all_misses(self): async def run(): encoder = MMEncoder.__new__(MMEncoder) @@ -692,6 +1677,7 @@ class TestEncoderDelivery(CustomTestCase): encoder = MMEncoder.__new__(MMEncoder) encoder.rank = 0 encoder.req_states = {} + encoder.abandoned_req_ids = set() encoder.delivery = SimpleNamespace(release=AsyncMock()) state = encoder._acquire_encode_ref("req") @@ -732,6 +1718,38 @@ class TestEncoderDelivery(CustomTestCase): asyncio.run(run()) + def test_abandon_before_encode_is_applied_when_state_is_created(self): + async def run(): + encoder = MMEncoder.__new__(MMEncoder) + encoder.rank = 0 + encoder.req_states = {} + encoder.abandoned_req_ids = set() + encoder.delivery = SimpleNamespace(release=AsyncMock()) + + await encoder.abandon_request("req") + self.assertIn("req", encoder.abandoned_req_ids) + + state = encoder._acquire_encode_ref("req") + self.assertTrue(state.release_requested) + self.assertNotIn("req", encoder.abandoned_req_ids) + encoder._stage_embedding( + EmbeddingData( + "req", + 1, + 0, + None, + Modality.IMAGE, + embedding=torch.ones((1, 1)), + ) + ) + with patch.object(meta_registry, "discard", AsyncMock()): + await encoder._release_encode_ref(state) + + encoder.delivery.release.assert_awaited_once_with(state) + self.assertNotIn("req", encoder.req_states) + + asyncio.run(run()) + def test_error_metadata_survives_buffer_release_for_waiter(self): async def run(): req_id = "test-error-metadata-waiter" @@ -822,11 +1840,123 @@ class TestEncoderDelivery(CustomTestCase): asyncio.run(run()) + def test_cancelled_tp_pipeline_drains_encode_before_release(self): + async def run(): + encode_started = asyncio.Event() + finish_encode = asyncio.Event() + encoder = MMEncoder.__new__(MMEncoder) + encoder.transfer_backend = "zmq_to_tokenizer" + encoder.encode_dispatch_lock = asyncio.Lock() + + async def encode(**_kwargs): + encode_started.set() + await finish_encode.wait() + self.assertTrue(encoder.encode_dispatch_lock.locked()) + return 16, 2, 4, None, None + + encoder.encode = AsyncMock(side_effect=encode) + encoder.release_request = AsyncMock() + request = { + "req_id": "cancelled", + "mm_items": ["item"], + "modality": "video", + "num_parts": 1, + "part_idx": 0, + } + + with patch("sglang.srt.disaggregation.encoder.runtime.sock_send") as send: + task = asyncio.create_task( + execute_encode_pipeline( + encoder, None, request, send_sockets=[object()] + ) + ) + await encode_started.wait() + task.cancel() + await asyncio.sleep(0) + + self.assertFalse(task.done()) + self.assertTrue(encoder.encode_dispatch_lock.locked()) + encoder.release_request.assert_not_awaited() + send.assert_called_once() + + task.cancel() + await asyncio.sleep(0) + self.assertFalse(task.done()) + self.assertTrue(encoder.encode_dispatch_lock.locked()) + + finish_encode.set() + with self.assertRaises(asyncio.CancelledError): + await task + + encoder.release_request.assert_awaited_once_with("cancelled") + self.assertFalse(encoder.encode_dispatch_lock.locked()) + + asyncio.run(run()) + + def test_pipeline_releases_request_when_error_publish_fails(self): + async def run(): + encoder = MMEncoder.__new__(MMEncoder) + encoder.transfer_backend = "mooncake" + encoder.encode = AsyncMock(side_effect=RuntimeError("encode failed")) + encoder.release_request = AsyncMock() + request = { + "req_id": "req", + "mm_items": ["item"], + "modality": "image", + "num_parts": 1, + "part_idx": 0, + } + + with patch.object( + meta_registry, + "publish", + AsyncMock(side_effect=RuntimeError("registry failed")), + ): + with self.assertRaisesRegex(RuntimeError, "encode failed"): + await execute_encode_pipeline(encoder, None, request) + + encoder.release_request.assert_awaited_once_with( + "req", preserve_metadata=False + ) + + asyncio.run(run()) + + def test_pipeline_releases_error_result_when_error_send_fails(self): + async def run(): + encoder = MMEncoder.__new__(MMEncoder) + encoder.transfer_backend = "zmq_to_scheduler" + encoder.encode = AsyncMock(return_value=(0, 0, 0, "bad image", 400)) + encoder.release_request = AsyncMock() + request = { + "req_id": "req", + "mm_items": ["item"], + "modality": "image", + "num_parts": 1, + "part_idx": 0, + } + + with ( + patch.object(meta_registry, "publish", AsyncMock()), + patch( + "sglang.srt.disaggregation.encoder.runtime._push_embedding_to_prefill", + AsyncMock(side_effect=RuntimeError("send failed")), + ), + ): + with self.assertRaisesRegex(MMError, "bad image"): + await execute_encode_pipeline(encoder, None, request) + + encoder.release_request.assert_awaited_once_with( + "req", preserve_metadata=False + ) + + asyncio.run(run()) + def test_send_waits_for_embedding_published_by_encode(self): async def run(): encoder = MMEncoder.__new__(MMEncoder) encoder.rank = 0 encoder.req_states = {} + encoder.abandoned_req_ids = set() state = encoder._acquire_encode_ref("req") state.embedding_data = EmbeddingData( "req", @@ -926,6 +2056,349 @@ class TestEncoderDelivery(CustomTestCase): asyncio.run(run()) + def test_release_wakes_destination_waiter(self): + async def run(): + req_id = "test-release-wakes-destination-waiter" + await meta_registry.discard(req_id) + + encoder = MMEncoder.__new__(MMEncoder) + encoder.req_states = {req_id: ReqState(req_id)} + encoder.send_timeout = 60 + encoder.delivery = ZmqDelivery(encoder, cleanup_receive_state=True) + + send_task = asyncio.create_task(encoder.send_with_url(req_id)) + await asyncio.sleep(0) + self.assertFalse(send_task.done()) + + await asyncio.wait_for(encoder.release_request(req_id), timeout=1) + await asyncio.wait_for(send_task, timeout=1) + + self.assertNotIn(req_id, encoder.req_states) + self.assertNotIn(req_id, rid_to_cond) + self.assertNotIn(req_id, rid_to_receive_endpoint) + self.assertNotIn(req_id, rid_to_receive_count) + + asyncio.run(run()) + + def test_destination_registration_rendezvous_with_encode(self): + async def run(register_first): + req_id = "registration-before-encode" + await meta_registry.discard(req_id) + encoder = MMEncoder.__new__(MMEncoder) + encoder.rank = 0 + encoder.req_states = {} + encoder.abandoned_req_ids = set() + encoder.use_mooncake = False + encoder.mm_global_cache = None + encoder.profiler = None + encoder.send_timeout = 1 + encoder.delivery = ZmqDelivery(encoder, cleanup_receive_state=True) + embedding = torch.ones((1, 4)) + ctx = SimpleNamespace( + req_id=req_id, + modality=Modality.IMAGE, + items_per_req=[1], + preprocess_result=SimpleNamespace( + token_counts=[1], grid_thw=[[1, 1, 1]] + ), + aux_data={}, + use_global_cache=False, + ) + request = { + "req_id": req_id, + "receive_count": 1, + "receive_url": "tcp://127.0.0.1:1", + } + with ( + patch.object( + encoder_server, "encode_state_condition", asyncio.Condition() + ), + patch.object(http_server, "encoder", encoder), + patch.object(http_server, "dp_dispatcher", None), + patch.object( + encoder, + "_prepare_encode_context_on_all_ranks", + AsyncMock(return_value=ctx), + ), + patch.object( + encoder, "_compute_embedding", AsyncMock(return_value=embedding) + ), + patch.object(encoder, "_send", AsyncMock()) as send, + ): + registration = None + if register_first: + registration = asyncio.create_task( + http_server.handle_scheduler_receive_url_request(request) + ) + await asyncio.sleep(0) + self.assertFalse(registration.done()) + self.assertNotIn(req_id, encoder.req_states) + await encoder.encode([], Modality.IMAGE, req_id, 1, 0) + if registration is None: + registration = asyncio.create_task( + http_server.handle_scheduler_receive_url_request(request) + ) + response = await asyncio.wait_for(registration, timeout=1) + self.assertEqual(response.status_code, 200) + self.assertEqual( + rid_to_receive_endpoint[req_id], {request["receive_url"]} + ) + await asyncio.wait_for(encoder.send_with_url(req_id), timeout=1) + send.assert_awaited_once() + torch.testing.assert_close(send.await_args.args[0], embedding) + self.assertNotIn(req_id, encoder.req_states) + self.assertNotIn(req_id, rid_to_receive_endpoint) + self.assertNotIn(req_id, rid_to_receive_count) + self.assertNotIn(req_id, rid_to_cond) + + for register_first in (True, False): + with self.subTest(register_first=register_first): + asyncio.run(run(register_first)) + + def test_destination_registration_timeout_does_not_create_state(self): + async def run(): + req_id = "registration-without-encode" + encoder = MMEncoder.__new__(MMEncoder) + encoder.req_states = {} + with ( + patch.object( + encoder_server, "encode_state_condition", asyncio.Condition() + ), + patch.object(encoder_server, "ENCODER_REQ_TIMEOUT", 0.01), + patch.object(http_server, "encoder", encoder), + patch.object(http_server, "dp_dispatcher", None), + ): + response = await http_server.handle_scheduler_receive_url_request( + { + "req_id": req_id, + "receive_count": 1, + "receive_url": "tcp://127.0.0.1:1", + } + ) + self.assertEqual(response.status_code, 504) + self.assertNotIn(req_id, encoder.req_states) + self.assertNotIn(req_id, rid_to_receive_endpoint) + self.assertNotIn(req_id, rid_to_receive_count) + self.assertNotIn(req_id, rid_to_cond) + + asyncio.run(run()) + + def test_destination_registration_cancellation_does_not_create_state(self): + async def run(): + req_id = "cancelled-registration" + encoder = MMEncoder.__new__(MMEncoder) + encoder.req_states = {} + with patch.object( + encoder_server, "encode_state_condition", asyncio.Condition() + ): + registration = asyncio.create_task( + encoder.register_embedding_destinations( + req_id, 1, ["tcp://127.0.0.1:1"] + ) + ) + await asyncio.sleep(0) + registration.cancel() + with self.assertRaises(asyncio.CancelledError): + await registration + self.assertNotIn(req_id, encoder.req_states) + self.assertNotIn(req_id, rid_to_cond) + + asyncio.run(run()) + + def test_destination_registration_respects_request_lifecycle(self): + async def run(): + req_id = "registration-lifecycle" + await meta_registry.discard(req_id) + + encoder = MMEncoder.__new__(MMEncoder) + state = ReqState(req_id) + state.active_encodes = 1 + encoder.req_states = {req_id: state} + encoder.delivery = ZmqDelivery(encoder, cleanup_receive_state=True) + + await encoder.register_embedding_destinations( + req_id, 1, ["tcp://127.0.0.1:1"] + ) + self.assertIn(req_id, rid_to_receive_endpoint) + + with patch.object(meta_registry, "discard", AsyncMock()): + await encoder.release_request(req_id) + + self.assertTrue(state.release_requested) + self.assertIn(req_id, encoder.req_states) + with self.assertRaisesRegex(BadRequestError, "not active"): + await encoder.register_embedding_destinations( + req_id, 1, ["tcp://127.0.0.1:2"] + ) + self.assertEqual(rid_to_receive_endpoint[req_id], {"tcp://127.0.0.1:1"}) + + with patch.object(meta_registry, "discard", AsyncMock()): + await encoder._release_encode_ref(state) + + self.assertNotIn(req_id, rid_to_receive_endpoint) + self.assertNotIn(req_id, rid_to_receive_count) + self.assertNotIn(req_id, rid_to_cond) + + # a reused ID starts a new request lifecycle + encoder.req_states[req_id] = ReqState(req_id) + await encoder.register_embedding_destinations( + req_id, 1, ["tcp://127.0.0.1:3"] + ) + self.assertEqual(rid_to_receive_endpoint[req_id], {"tcp://127.0.0.1:3"}) + + with patch.object(meta_registry, "discard", AsyncMock()): + await encoder.release_request(req_id) + + asyncio.run(run()) + + +class TestEncoderDPAbandonedRequest(CustomTestCase): + @staticmethod + def _make_dispatcher(): + return DPDispatcher( + dp_size=1, + dispatch_sockets=[object()], + release_sockets=[object()], + result_socket=object(), + worker_processes=[], + ) + + def test_dispatch_timeout_notifies_worker_to_release(self): + async def run(): + dispatcher = self._make_dispatcher() + sent = [] + release_sent = asyncio.Event() + + async def send(socket, payload): + message = unwrap_from_pickle(payload) + sent.append((socket, message)) + if message.get("_dp_type") == _DP_RELEASE_AFTER_ENCODE: + release_sent.set() + + request = {"req_id": "timed-out", "modality": "image"} + with ( + patch( + "sglang.srt.disaggregation.encoder.runtime.async_sock_send", + side_effect=send, + ), + patch( + "sglang.srt.disaggregation.encoder.runtime.server_module.ENCODER_REQ_TIMEOUT", + 0.01, + ), + ): + result = await dispatcher.dispatch(request) + await asyncio.wait_for(release_sent.wait(), timeout=1) + + self.assertEqual(result["_error_type"], "TimeoutError") + self.assertIs(sent[0][0], dispatcher.dispatch_sockets[0]) + self.assertEqual(sent[0][1], request) + self.assertIs(sent[1][0], dispatcher.release_sockets[0]) + self.assertEqual( + sent[1][1], + { + "_dp_type": _DP_RELEASE_AFTER_ENCODE, + "req_id": "timed-out", + }, + ) + self.assertEqual(dispatcher.pending_counts, [0]) + self.assertNotIn("timed-out", dispatcher.req_id_to_rank) + + asyncio.run(run()) + + def test_dispatch_cancellation_notifies_worker_to_release(self): + async def run(): + dispatcher = self._make_dispatcher() + encode_sent = asyncio.Event() + release_sent = asyncio.Event() + + async def send(_socket, payload): + message = unwrap_from_pickle(payload) + if message.get("_dp_type") == _DP_RELEASE_AFTER_ENCODE: + release_sent.set() + else: + encode_sent.set() + + with patch( + "sglang.srt.disaggregation.encoder.runtime.async_sock_send", + side_effect=send, + ): + task = asyncio.create_task( + dispatcher.dispatch({"req_id": "cancelled", "modality": "image"}) + ) + await encode_sent.wait() + task.cancel() + with self.assertRaises(asyncio.CancelledError): + await task + await asyncio.wait_for(release_sent.wait(), timeout=1) + + self.assertEqual(dispatcher.pending_counts, [0]) + self.assertNotIn("cancelled", dispatcher.req_id_to_rank) + + asyncio.run(run()) + + def test_worker_marks_running_encode_abandoned(self): + async def run(): + async def encode(): + await asyncio.Event().wait() + + encode_task = asyncio.create_task(encode()) + encoder = SimpleNamespace( + abandon_request=AsyncMock(), + release_request=AsyncMock(), + ) + await _retire_abandoned_encode(encoder, encode_task, "abandoned") + + encoder.abandon_request.assert_awaited_once_with("abandoned") + encoder.release_request.assert_not_awaited() + encode_task.cancel() + await asyncio.gather(encode_task, return_exceptions=True) + + asyncio.run(run()) + + def test_worker_preserves_release_before_encode_task_exists(self): + async def run(): + encoder = MMEncoder.__new__(MMEncoder) + encoder.rank = 0 + encoder.req_states = {} + encoder.abandoned_req_ids = set() + encoder.delivery = SimpleNamespace(release=AsyncMock()) + + await _retire_abandoned_encode(encoder, None, "abandoned") + self.assertIn("abandoned", encoder.abandoned_req_ids) + + state = encoder._acquire_encode_ref("abandoned") + self.assertTrue(state.release_requested) + with patch.object(meta_registry, "discard", AsyncMock()): + await encoder._release_encode_ref(state) + + encoder.delivery.release.assert_awaited_once_with(state) + self.assertNotIn("abandoned", encoder.req_states) + + asyncio.run(run()) + + def test_worker_release_survives_encode_failure(self): + async def run(): + async def encode(): + raise RuntimeError("bad image") + + encode_task = asyncio.create_task(encode()) + await asyncio.sleep(0) + encoder = SimpleNamespace( + abandon_request=AsyncMock(), + release_request=AsyncMock(), + ) + await _retire_abandoned_encode( + encoder, + encode_task, + "failed", + ) + + encoder.release_request.assert_awaited_once_with("failed") + encoder.abandon_request.assert_not_awaited() + await asyncio.gather(encode_task, return_exceptions=True) + + asyncio.run(run()) + class TestMooncakeRegistration(CustomTestCase): def setUp(self): diff --git a/test/registered/unit/disaggregation/test_encoder_health.py b/test/registered/unit/disaggregation/test_encoder_health.py index c5d3872b7..5a6b64ed8 100644 --- a/test/registered/unit/disaggregation/test_encoder_health.py +++ b/test/registered/unit/disaggregation/test_encoder_health.py @@ -17,6 +17,8 @@ class _FakeEncoder: self.embedding_to_send = {} self.encode_dispatch_lock = asyncio.Lock() self.encode_calls = [] + self.released = [] + self.release_event = asyncio.Event() def has_pending_embeddings(self): return bool(self.embedding_to_send) @@ -28,8 +30,9 @@ class _FakeEncoder: self.encode_calls.append(kwargs) return 1, 1, 1, None, None - async def release_request(self, _req_id): - return None + async def release_request(self, req_id): + self.released.append(req_id) + self.release_event.set() def _install_tp_encoder(monkeypatch, encoder): @@ -84,5 +87,87 @@ def test_health_encode_rechecks_busy_state_after_waiting(monkeypatch): asyncio.run(run_test()) +def test_health_timeout_keeps_dispatch_order_until_encode_drains(monkeypatch): + async def run_test(): + encoder = _FakeEncoder() + _install_tp_encoder(monkeypatch, encoder) + encode_started = asyncio.Event() + finish_encode = asyncio.Event() + + async def encode(**kwargs): + encoder.encode_calls.append(kwargs) + encode_started.set() + await finish_encode.wait() + return 1, 1, 1, None, None + + encoder.encode = encode + monkeypatch.setattr(http_server, "HEALTH_CHECK_TIMEOUT", 0.01) + + response = await http_server.health_generate() + assert response.status_code == 503 + assert encode_started.is_set() + assert encoder.encode_dispatch_lock.locked() + assert encoder.released == [] + + finish_encode.set() + await asyncio.wait_for(encoder.release_event.wait(), timeout=1) + await asyncio.wait_for(encoder.encode_dispatch_lock.acquire(), timeout=1) + encoder.encode_dispatch_lock.release() + assert len(encoder.released) == 1 + + asyncio.run(run_test()) + + +def test_cancelled_health_request_does_not_cancel_dispatched_encode(monkeypatch): + async def run_test(): + encoder = _FakeEncoder() + _install_tp_encoder(monkeypatch, encoder) + encode_started = asyncio.Event() + finish_encode = asyncio.Event() + + async def encode(**kwargs): + encoder.encode_calls.append(kwargs) + encode_started.set() + await finish_encode.wait() + return 1, 1, 1, None, None + + encoder.encode = encode + task = asyncio.create_task(http_server.health_generate()) + await asyncio.wait_for(encode_started.wait(), timeout=1) + + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + assert encoder.encode_dispatch_lock.locked() + assert encoder.released == [] + + finish_encode.set() + await asyncio.wait_for(encoder.release_event.wait(), timeout=1) + await asyncio.wait_for(encoder.encode_dispatch_lock.acquire(), timeout=1) + encoder.encode_dispatch_lock.release() + assert len(encoder.released) == 1 + + asyncio.run(run_test()) + + +def test_health_cleanup_failure_releases_dispatch_lock(monkeypatch): + async def run_test(): + encoder = _FakeEncoder() + _install_tp_encoder(monkeypatch, encoder) + + async def release_request(req_id): + encoder.released.append(req_id) + raise RuntimeError("cleanup failed") + + encoder.release_request = release_request + response = await http_server.health_generate() + + assert response.status_code == 503 + assert not encoder.encode_dispatch_lock.locked() + assert len(encoder.released) == 1 + + asyncio.run(run_test()) + + if __name__ == "__main__": sys.exit(pytest.main([__file__, "-v"])) diff --git a/test/registered/unit/disaggregation/test_encoder_scheduler.py b/test/registered/unit/disaggregation/test_encoder_scheduler.py index 1e31694a4..94c69fc28 100644 --- a/test/registered/unit/disaggregation/test_encoder_scheduler.py +++ b/test/registered/unit/disaggregation/test_encoder_scheduler.py @@ -1,5 +1,7 @@ import asyncio import sys +from types import SimpleNamespace +from unittest.mock import Mock, patch import pytest @@ -7,7 +9,9 @@ from sglang.srt.disaggregation.encoder.runtime import ( EncoderScheduler, PendingRequest, _resolve_encoder_batch_policy, + validate_encode_request, ) +from sglang.srt.managers.schedule_batch import Modality from sglang.test.ci.ci_register import register_cpu_ci register_cpu_ci(est_time=1, suite="base-a-test-cpu") @@ -108,6 +112,113 @@ def test_scheduler_coalesces_concurrent_submissions(): asyncio.run(run_test()) +def test_scheduler_isolates_bad_request_from_failed_fused_batch(): + class FakeEncoder: + def __init__(self): + self.encode_dispatch_lock = asyncio.Lock() + self.batches = [] + + async def batch_encode(self, requests, _modality): + req_ids = [request["req_id"] for request in requests] + self.batches.append(req_ids) + if len(requests) > 1 or req_ids == ["bad"]: + return [(0, 0, 0, "bad image", 400) for _ in requests] + return [(1, 2, 3, None, None)] + + async def run_test(): + encoder = FakeEncoder() + scheduler = EncoderScheduler( + encoder=encoder, + send_sockets=[], + max_batch_size=8, + coalesce_same_turn=True, + ) + collector = SimpleNamespace(observe_queue_wait=Mock()) + with patch( + "sglang.srt.disaggregation.encoder.runtime.server_module.encoder_metrics_collector", + collector, + ): + scheduler.start() + try: + requests = [ + { + "req_id": req_id, + "modality": "image", + "mm_items": [object()], + "num_parts": 1, + "part_idx": 0, + } + for req_id in ("bad", "good") + ] + results = await asyncio.gather( + *(scheduler.submit(request) for request in requests) + ) + finally: + await scheduler.stop() + + assert encoder.batches == [["bad", "good"], ["bad"], ["good"]] + assert results == [(0, 0, 0, "bad image", 400), (1, 2, 3, None, None)] + assert collector.observe_queue_wait.call_count == len(requests) + + asyncio.run(run_test()) + + +@pytest.mark.parametrize( + ("update", "expected"), + [ + ({"req_id": ""}, "missing or invalid req_id"), + ({"modality": "text"}, "unsupported modality"), + ({"mm_items": []}, "missing or empty mm_items"), + ({"num_parts": 0}, "num_parts must be a positive integer"), + ({"part_idx": 1}, "part_idx must be in [0, 1)"), + ], +) +def test_validate_encode_request_rejects_invalid_fields(update, expected): + request = { + "req_id": "request", + "modality": "image", + "mm_items": [object()], + "num_parts": 1, + "part_idx": 0, + } + request.update(update) + + assert expected in validate_encode_request(request) + + +def test_video_request_is_validated_before_tp_broadcast(): + class FakeSocket: + pass + + class FakeEncoder: + async def encode(self, **_kwargs): + raise AssertionError("invalid request must not reach the encoder") + + async def run_test(): + scheduler = EncoderScheduler( + encoder=FakeEncoder(), + send_sockets=[FakeSocket()], + max_batch_size=1, + ) + pending = PendingRequest( + { + "req_id": "bad-video", + "modality": "video", + "mm_items": [object()], + "num_parts": 1, + "part_idx": 1, + }, + asyncio.get_running_loop(), + ) + + await scheduler._dispatch_per_request([pending], Modality.VIDEO) + + with pytest.raises(Exception, match="part_idx must be in"): + pending.future.result() + + asyncio.run(run_test()) + + @pytest.mark.parametrize( ("model_type", "configured", "explicit", "expected"), [