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
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"),
[