fix(vlm): contain EPD request lifecycle failures (#36944)

Co-authored-by: mickqian <mickqian@users.noreply.github.com>
This commit is contained in:
Mick
2026-09-06 16:05:10 +08:00
committed by GitHub
co-authored by mickqian
parent d61378af77
commit 8ef646a5c6
8 changed files with 2435 additions and 139 deletions
@@ -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()
@@ -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"])
@@ -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,
@@ -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)
@@ -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:
File diff suppressed because it is too large Load Diff
@@ -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"]))
@@ -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"),
[