[diffusion] feat: make scheduler rpc deadlines explicit (#33965)
Co-authored-by: suoyf <suoyf@nscc-tj.cn> Co-authored-by: Mick <mickjagger19@icloud.com>
This commit is contained in:
@@ -99,6 +99,7 @@ Use `sglang generate --help` and `sglang serve --help` for the full argument lis
|
|||||||
- `--srt-encoder-url {HTTPADDRESS}`: address of SGLang srt server with AR model for GLM-Image like models. See [Models with AR Stage](../models_with_ar).
|
- `--srt-encoder-url {HTTPADDRESS}`: address of SGLang srt server with AR model for GLM-Image like models. See [Models with AR Stage](../models_with_ar).
|
||||||
- `--srt-encoder-timeout {SECONDS}`: Timeout in seconds for HTTP requests to the SGLang encoder server
|
- `--srt-encoder-timeout {SECONDS}`: Timeout in seconds for HTTP requests to the SGLang encoder server
|
||||||
- `--srt-encoder-connection-timeout {SECONDS}`: TCP connection timeout in seconds for SGLang encoder server
|
- `--srt-encoder-connection-timeout {SECONDS}`: TCP connection timeout in seconds for SGLang encoder server
|
||||||
|
- `--scheduler-rpc-timeout {SECONDS}`: optional end-to-end deadline for an internal scheduler RPC, including scheduler queue time. It is unset by default so valid long-running and queued video jobs are not failed by the transport layer. Set it only when the deployment requires a bounded request deadline; caller cancellation and server shutdown remain effective without it.
|
||||||
- `--pe-server-url {HTTPADDRESS}`: url of SGLang server hosting the PE model (e.g., for ERNIE-Image). See [Models with Prompt Enhancement](../models_with_pe).
|
- `--pe-server-url {HTTPADDRESS}`: url of SGLang server hosting the PE model (e.g., for ERNIE-Image). See [Models with Prompt Enhancement](../models_with_pe).
|
||||||
|
|
||||||
### Sampling and output
|
### Sampling and output
|
||||||
|
|||||||
@@ -105,6 +105,9 @@ async def _run_server_warmup_after_http_live(
|
|||||||
|
|
||||||
@asynccontextmanager
|
@asynccontextmanager
|
||||||
async def lifespan(app: FastAPI):
|
async def lifespan(app: FastAPI):
|
||||||
|
from sglang.multimodal_gen.runtime.entrypoints.openai.video_api import (
|
||||||
|
shutdown_video_jobs,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.runtime.scheduler_client import (
|
from sglang.multimodal_gen.runtime.scheduler_client import (
|
||||||
async_scheduler_client,
|
async_scheduler_client,
|
||||||
run_zeromq_broker,
|
run_zeromq_broker,
|
||||||
@@ -136,7 +139,10 @@ async def lifespan(app: FastAPI):
|
|||||||
|
|
||||||
# On shutdown
|
# On shutdown
|
||||||
logger.info("FastAPI app is shutting down...")
|
logger.info("FastAPI app is shutting down...")
|
||||||
|
await shutdown_video_jobs()
|
||||||
broker_task.cancel()
|
broker_task.cancel()
|
||||||
|
with suppress(asyncio.CancelledError):
|
||||||
|
await broker_task
|
||||||
async_scheduler_client.close()
|
async_scheduler_client.close()
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -6,6 +6,8 @@ import os
|
|||||||
import shutil
|
import shutil
|
||||||
import tempfile
|
import tempfile
|
||||||
import time
|
import time
|
||||||
|
from collections.abc import Coroutine
|
||||||
|
from contextlib import suppress
|
||||||
from typing import Any, Dict, Optional
|
from typing import Any, Dict, Optional
|
||||||
|
|
||||||
from fastapi import (
|
from fastapi import (
|
||||||
@@ -61,6 +63,29 @@ _VIDEO_EXTENSIONS = {
|
|||||||
".mpg",
|
".mpg",
|
||||||
".webm",
|
".webm",
|
||||||
}
|
}
|
||||||
|
_VIDEO_JOB_TASKS: dict[str, asyncio.Task[None]] = {}
|
||||||
|
|
||||||
|
|
||||||
|
def _discard_video_job_task(job_id: str, task: asyncio.Task[None]) -> None:
|
||||||
|
if _VIDEO_JOB_TASKS.get(job_id) is task:
|
||||||
|
del _VIDEO_JOB_TASKS[job_id]
|
||||||
|
|
||||||
|
|
||||||
|
def _start_video_job(job_id: str, job: Coroutine[Any, Any, None]) -> asyncio.Task[None]:
|
||||||
|
task = asyncio.create_task(job, name=f"video-job-{job_id}")
|
||||||
|
_VIDEO_JOB_TASKS[job_id] = task
|
||||||
|
task.add_done_callback(lambda completed: _discard_video_job_task(job_id, completed))
|
||||||
|
return task
|
||||||
|
|
||||||
|
|
||||||
|
async def shutdown_video_jobs() -> None:
|
||||||
|
tasks = list(_VIDEO_JOB_TASKS.values())
|
||||||
|
_VIDEO_JOB_TASKS.clear()
|
||||||
|
for task in tasks:
|
||||||
|
task.cancel()
|
||||||
|
for task in tasks:
|
||||||
|
with suppress(asyncio.CancelledError):
|
||||||
|
await task
|
||||||
|
|
||||||
|
|
||||||
def _extra_value(request: VideoGenerationsRequest, name: str) -> Any:
|
def _extra_value(request: VideoGenerationsRequest, name: str) -> Any:
|
||||||
@@ -826,14 +851,15 @@ async def create_video(
|
|||||||
|
|
||||||
assert batch is not None
|
assert batch is not None
|
||||||
# Enqueue the job asynchronously and return immediately
|
# Enqueue the job asynchronously and return immediately
|
||||||
asyncio.create_task(
|
_start_video_job(
|
||||||
|
request_id,
|
||||||
_dispatch_job_async(
|
_dispatch_job_async(
|
||||||
request_id,
|
request_id,
|
||||||
batch,
|
batch,
|
||||||
scheduler_batches=scheduler_batches,
|
scheduler_batches=scheduler_batches,
|
||||||
temp_dirs=temp_dirs or None,
|
temp_dirs=temp_dirs or None,
|
||||||
output_persistent=output_persistent,
|
output_persistent=output_persistent,
|
||||||
)
|
),
|
||||||
)
|
)
|
||||||
return VideoResponse(**job)
|
return VideoResponse(**job)
|
||||||
|
|
||||||
|
|||||||
@@ -25,7 +25,10 @@ from sglang.multimodal_gen.runtime.entrypoints.utils import (
|
|||||||
from sglang.multimodal_gen.runtime.ipc_array import materialize_file_refs
|
from sglang.multimodal_gen.runtime.ipc_array import materialize_file_refs
|
||||||
from sglang.multimodal_gen.runtime.pipelines_core import Req
|
from sglang.multimodal_gen.runtime.pipelines_core import Req
|
||||||
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch
|
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch
|
||||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
from sglang.multimodal_gen.runtime.server_args import (
|
||||||
|
MAX_SCHEDULER_RPC_TIMEOUT_S,
|
||||||
|
ServerArgs,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
from sglang.multimodal_gen.runtime.utils.request_logger import (
|
from sglang.multimodal_gen.runtime.utils.request_logger import (
|
||||||
DiffusionRequestLogger,
|
DiffusionRequestLogger,
|
||||||
@@ -50,6 +53,28 @@ _CONTROL_REQ_TYPES = (
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _configure_recv_timeout(socket: Any, timeout_ms: int | None) -> None:
|
||||||
|
if timeout_ms is None:
|
||||||
|
return
|
||||||
|
max_timeout_ms = MAX_SCHEDULER_RPC_TIMEOUT_S * 1000
|
||||||
|
if (
|
||||||
|
not isinstance(timeout_ms, int)
|
||||||
|
or isinstance(timeout_ms, bool)
|
||||||
|
or not 0 < timeout_ms <= max_timeout_ms
|
||||||
|
):
|
||||||
|
raise ValueError(
|
||||||
|
f"timeout_ms must be None or an integer between 1 and {max_timeout_ms}"
|
||||||
|
)
|
||||||
|
socket.setsockopt(zmq.RCVTIMEO, timeout_ms)
|
||||||
|
|
||||||
|
|
||||||
|
def _resolve_timeout_ms(server_args: ServerArgs, timeout_ms: int | None) -> int | None:
|
||||||
|
if timeout_ms is not None:
|
||||||
|
return timeout_ms
|
||||||
|
timeout_s = server_args.scheduler_rpc_timeout
|
||||||
|
return None if timeout_s is None else timeout_s * 1000
|
||||||
|
|
||||||
|
|
||||||
async def run_zeromq_broker(server_args: ServerArgs):
|
async def run_zeromq_broker(server_args: ServerArgs):
|
||||||
"""
|
"""
|
||||||
This function runs as a background task in the FastAPI process.
|
This function runs as a background task in the FastAPI process.
|
||||||
@@ -61,6 +86,7 @@ async def run_zeromq_broker(server_args: ServerArgs):
|
|||||||
socket.bind(broker_endpoint)
|
socket.bind(broker_endpoint)
|
||||||
logger.info(f"ZMQ Broker is listening for offline jobs on {broker_endpoint}")
|
logger.info(f"ZMQ Broker is listening for offline jobs on {broker_endpoint}")
|
||||||
|
|
||||||
|
try:
|
||||||
while True:
|
while True:
|
||||||
try:
|
try:
|
||||||
# 1. Receive a request from an offline client
|
# 1. Receive a request from an offline client
|
||||||
@@ -78,9 +104,14 @@ async def run_zeromq_broker(server_args: ServerArgs):
|
|||||||
logger.error(f"Error in ZMQ Broker: {e}", exc_info=True)
|
logger.error(f"Error in ZMQ Broker: {e}", exc_info=True)
|
||||||
# A reply must be sent to prevent the client from hanging
|
# A reply must be sent to prevent the client from hanging
|
||||||
try:
|
try:
|
||||||
await socket.send(pickle.dumps({"status": "error", "message": str(e)}))
|
await socket.send(
|
||||||
|
pickle.dumps({"status": "error", "message": str(e)})
|
||||||
|
)
|
||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
|
finally:
|
||||||
|
socket.close(linger=0)
|
||||||
|
ctx.destroy(linger=0)
|
||||||
|
|
||||||
|
|
||||||
def _session_key(batch: Any) -> str | None:
|
def _session_key(batch: Any) -> str | None:
|
||||||
@@ -138,9 +169,10 @@ class SchedulerClient:
|
|||||||
|
|
||||||
def _forward_one(self, endpoint: str, batch: Any, timeout_ms: int | None) -> Any:
|
def _forward_one(self, endpoint: str, batch: Any, timeout_ms: int | None) -> Any:
|
||||||
socket = self.context.socket(zmq.REQ)
|
socket = self.context.socket(zmq.REQ)
|
||||||
socket.setsockopt(zmq.LINGER, 0)
|
|
||||||
socket.setsockopt(zmq.RCVTIMEO, timeout_ms if timeout_ms else 6000000)
|
|
||||||
try:
|
try:
|
||||||
|
socket.setsockopt(zmq.LINGER, 0)
|
||||||
|
effective_timeout = _resolve_timeout_ms(self.server_args, timeout_ms)
|
||||||
|
_configure_recv_timeout(socket, effective_timeout)
|
||||||
socket.connect(endpoint)
|
socket.connect(endpoint)
|
||||||
socket.send_pyobj(batch)
|
socket.send_pyobj(batch)
|
||||||
output_batch = socket.recv_pyobj()
|
output_batch = socket.recv_pyobj()
|
||||||
@@ -220,7 +252,7 @@ class AsyncSchedulerClient:
|
|||||||
self.context = zmq.asyncio.Context()
|
self.context = zmq.asyncio.Context()
|
||||||
logger.debug("AsyncSchedulerClient initialized with zmq.asyncio.Context")
|
logger.debug("AsyncSchedulerClient initialized with zmq.asyncio.Context")
|
||||||
|
|
||||||
async def forward(self, batch: Any) -> Any:
|
async def forward(self, batch: Any, timeout_ms: int | None = None) -> Any:
|
||||||
"""Sends a batch or request to the scheduler and waits for the response."""
|
"""Sends a batch or request to the scheduler and waits for the response."""
|
||||||
self.request_logger.log_received_request(batch)
|
self.request_logger.log_received_request(batch)
|
||||||
if self.context is None:
|
if self.context is None:
|
||||||
@@ -231,23 +263,29 @@ class AsyncSchedulerClient:
|
|||||||
endpoints = self.server_args.scheduler_endpoints
|
endpoints = self.server_args.scheduler_endpoints
|
||||||
if isinstance(batch, _CONTROL_REQ_TYPES):
|
if isinstance(batch, _CONTROL_REQ_TYPES):
|
||||||
# replica state (weights, LoRA, memory) must change everywhere
|
# replica state (weights, LoRA, memory) must change everywhere
|
||||||
results = [await self._forward_one(ep, batch) for ep in endpoints]
|
results = [
|
||||||
|
await self._forward_one(ep, batch, timeout_ms) for ep in endpoints
|
||||||
|
]
|
||||||
output_batch = _merge_fanout_results(results)
|
output_batch = _merge_fanout_results(results)
|
||||||
else:
|
else:
|
||||||
replica = _select_replica(batch, len(endpoints), self._replica_counter)
|
replica = _select_replica(batch, len(endpoints), self._replica_counter)
|
||||||
output_batch = await self._forward_one(endpoints[replica], batch)
|
output_batch = await self._forward_one(
|
||||||
|
endpoints[replica], batch, timeout_ms
|
||||||
|
)
|
||||||
self.request_logger.log_finished_request(batch, output_batch)
|
self.request_logger.log_finished_request(batch, output_batch)
|
||||||
return output_batch
|
return output_batch
|
||||||
|
|
||||||
async def _forward_one(self, endpoint: str, batch: Any) -> Any:
|
async def _forward_one(
|
||||||
|
self, endpoint: str, batch: Any, timeout_ms: int | None
|
||||||
|
) -> Any:
|
||||||
# a temporary REQ socket per request keeps concurrent requests from
|
# a temporary REQ socket per request keeps concurrent requests from
|
||||||
# interleaving on one socket's strict send/recv alternation
|
# interleaving on one socket's strict send/recv alternation
|
||||||
socket = self.context.socket(zmq.REQ)
|
socket = self.context.socket(zmq.REQ)
|
||||||
socket.setsockopt(zmq.LINGER, 0)
|
|
||||||
# 100 minute timeout
|
|
||||||
socket.setsockopt(zmq.RCVTIMEO, 6000000)
|
|
||||||
socket.connect(endpoint)
|
|
||||||
try:
|
try:
|
||||||
|
socket.setsockopt(zmq.LINGER, 0)
|
||||||
|
effective_timeout = _resolve_timeout_ms(self.server_args, timeout_ms)
|
||||||
|
_configure_recv_timeout(socket, effective_timeout)
|
||||||
|
socket.connect(endpoint)
|
||||||
await socket.send(pickle.dumps(batch))
|
await socket.send(pickle.dumps(batch))
|
||||||
payload = await socket.recv()
|
payload = await socket.recv()
|
||||||
output_batch = pickle.loads(payload)
|
output_batch = pickle.loads(payload)
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ from sglang.multimodal_gen.runtime.server_args.server_args import (
|
|||||||
DEFAULT_BCG_TEXT_BUCKETS,
|
DEFAULT_BCG_TEXT_BUCKETS,
|
||||||
LORA_MERGE_MODES,
|
LORA_MERGE_MODES,
|
||||||
LTX2_TWO_STAGE_DEVICE_MODE_CHOICES,
|
LTX2_TWO_STAGE_DEVICE_MODE_CHOICES,
|
||||||
|
MAX_SCHEDULER_RPC_TIMEOUT_S,
|
||||||
Backend,
|
Backend,
|
||||||
PortArgs,
|
PortArgs,
|
||||||
ServerArgs,
|
ServerArgs,
|
||||||
@@ -24,6 +25,7 @@ __all__ = [
|
|||||||
"DEFAULT_BCG_TEXT_BUCKETS",
|
"DEFAULT_BCG_TEXT_BUCKETS",
|
||||||
"LORA_MERGE_MODES",
|
"LORA_MERGE_MODES",
|
||||||
"LTX2_TWO_STAGE_DEVICE_MODE_CHOICES",
|
"LTX2_TWO_STAGE_DEVICE_MODE_CHOICES",
|
||||||
|
"MAX_SCHEDULER_RPC_TIMEOUT_S",
|
||||||
"PortArgs",
|
"PortArgs",
|
||||||
"ServerArgs",
|
"ServerArgs",
|
||||||
"_normalize_ltx2_two_stage_device_mode",
|
"_normalize_ltx2_two_stage_device_mode",
|
||||||
|
|||||||
@@ -71,6 +71,7 @@ LTX2_TWO_STAGE_PIPELINE_NAMES = ("LTX2TwoStagePipeline", "LTX2TwoStageHQPipeline
|
|||||||
# H200-class GPUs (>=130 GiB total) can usually keep both LTX2 DiTs resident.
|
# H200-class GPUs (>=130 GiB total) can usually keep both LTX2 DiTs resident.
|
||||||
LTX2_RESIDENT_AUTO_ENABLE_MEM_GB = 130
|
LTX2_RESIDENT_AUTO_ENABLE_MEM_GB = 130
|
||||||
LORA_MERGE_MODES = ("auto", "merge", "dynamic")
|
LORA_MERGE_MODES = ("auto", "merge", "dynamic")
|
||||||
|
MAX_SCHEDULER_RPC_TIMEOUT_S = 2_147_483
|
||||||
# Mirrors AttentionBackend.supports_ring_rotation; the name-level check
|
# Mirrors AttentionBackend.supports_ring_rotation; the name-level check
|
||||||
# runs before backend classes are importable on every platform.
|
# runs before backend classes are importable on every platform.
|
||||||
RING_CAPABLE_ATTENTION_BACKENDS = ("fa", "sage_attn")
|
RING_CAPABLE_ATTENTION_BACKENDS = ("fa", "sage_attn")
|
||||||
@@ -245,6 +246,7 @@ class ServerArgs(DisaggServerArgsMixin):
|
|||||||
hsdp_replicate_dim: int = 1
|
hsdp_replicate_dim: int = 1
|
||||||
hsdp_shard_dim: Optional[int] = None
|
hsdp_shard_dim: Optional[int] = None
|
||||||
dist_timeout: int | None = 3600 # 1 hour
|
dist_timeout: int | None = 3600 # 1 hour
|
||||||
|
scheduler_rpc_timeout: int | None = None
|
||||||
|
|
||||||
pipeline_config: PipelineConfig = field(default_factory=PipelineConfig, repr=False)
|
pipeline_config: PipelineConfig = field(default_factory=PipelineConfig, repr=False)
|
||||||
|
|
||||||
@@ -484,6 +486,7 @@ class ServerArgs(DisaggServerArgsMixin):
|
|||||||
|
|
||||||
def _validate_parameters(self):
|
def _validate_parameters(self):
|
||||||
"""check consistency and raise errors for invalid configs"""
|
"""check consistency and raise errors for invalid configs"""
|
||||||
|
self._validate_scheduler_rpc_timeout()
|
||||||
self._validate_pipeline()
|
self._validate_pipeline()
|
||||||
self._validate_offload()
|
self._validate_offload()
|
||||||
if not current_platform.is_cpu():
|
if not current_platform.is_cpu():
|
||||||
@@ -493,6 +496,20 @@ class ServerArgs(DisaggServerArgsMixin):
|
|||||||
self._validate_breakable_cuda_graph()
|
self._validate_breakable_cuda_graph()
|
||||||
self.pipeline_config.validate_server_args(self)
|
self.pipeline_config.validate_server_args(self)
|
||||||
|
|
||||||
|
def _validate_scheduler_rpc_timeout(self) -> None:
|
||||||
|
timeout = self.scheduler_rpc_timeout
|
||||||
|
if timeout is None:
|
||||||
|
return
|
||||||
|
if (
|
||||||
|
not isinstance(timeout, int)
|
||||||
|
or isinstance(timeout, bool)
|
||||||
|
or not 0 < timeout <= MAX_SCHEDULER_RPC_TIMEOUT_S
|
||||||
|
):
|
||||||
|
raise ValueError(
|
||||||
|
"scheduler_rpc_timeout must be None or an integer between "
|
||||||
|
f"1 and {MAX_SCHEDULER_RPC_TIMEOUT_S} seconds"
|
||||||
|
)
|
||||||
|
|
||||||
def resolved_bcg_text_buckets(self) -> tuple[int, ...]:
|
def resolved_bcg_text_buckets(self) -> tuple[int, ...]:
|
||||||
"""Sorted, de-duplicated, positive BCG text buckets.
|
"""Sorted, de-duplicated, positive BCG text buckets.
|
||||||
|
|
||||||
@@ -1553,6 +1570,16 @@ class ServerArgs(DisaggServerArgsMixin):
|
|||||||
help="Timeout for torch.distributed operations in seconds. "
|
help="Timeout for torch.distributed operations in seconds. "
|
||||||
"Increase this value if you encounter 'Connection closed by peer' errors after the service is idle. ",
|
"Increase this value if you encounter 'Connection closed by peer' errors after the service is idle. ",
|
||||||
)
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--scheduler-rpc-timeout",
|
||||||
|
type=int,
|
||||||
|
default=ServerArgs.scheduler_rpc_timeout,
|
||||||
|
help=(
|
||||||
|
"Optional end-to-end timeout in seconds for a scheduler RPC, including "
|
||||||
|
"time spent in the scheduler queue. By default no transport-level "
|
||||||
|
"deadline is imposed; callers may still cancel their request."
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
ServerArgs.add_disagg_cli_args(parser)
|
ServerArgs.add_disagg_cli_args(parser)
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,160 @@
|
|||||||
|
import asyncio
|
||||||
|
import pickle
|
||||||
|
import uuid
|
||||||
|
from types import SimpleNamespace
|
||||||
|
from unittest.mock import AsyncMock, MagicMock
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
import zmq
|
||||||
|
import zmq.asyncio
|
||||||
|
|
||||||
|
from sglang.multimodal_gen.runtime.scheduler_client import (
|
||||||
|
AsyncSchedulerClient,
|
||||||
|
SchedulerClient,
|
||||||
|
run_zeromq_broker,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.server_args import MAX_SCHEDULER_RPC_TIMEOUT_S
|
||||||
|
|
||||||
|
|
||||||
|
def test_sync_scheduler_client_converts_configured_seconds_to_milliseconds():
|
||||||
|
response = object()
|
||||||
|
socket = MagicMock()
|
||||||
|
socket.recv_pyobj.return_value = response
|
||||||
|
client = SchedulerClient()
|
||||||
|
client.context = SimpleNamespace(socket=lambda _socket_type: socket)
|
||||||
|
client.server_args = SimpleNamespace(scheduler_rpc_timeout=2)
|
||||||
|
|
||||||
|
assert client._forward_one("tcp://scheduler", object(), None) is response
|
||||||
|
|
||||||
|
socket.setsockopt.assert_any_call(zmq.RCVTIMEO, 2000)
|
||||||
|
socket.close.assert_called_once_with()
|
||||||
|
|
||||||
|
|
||||||
|
def test_async_scheduler_client_has_no_transport_deadline_by_default():
|
||||||
|
response = {"status": "ok"}
|
||||||
|
socket = MagicMock()
|
||||||
|
socket.send = AsyncMock()
|
||||||
|
socket.recv = AsyncMock(return_value=pickle.dumps(response))
|
||||||
|
client = AsyncSchedulerClient()
|
||||||
|
client.context = SimpleNamespace(socket=lambda _socket_type: socket)
|
||||||
|
client.server_args = SimpleNamespace(scheduler_rpc_timeout=None)
|
||||||
|
|
||||||
|
result = asyncio.run(client._forward_one("tcp://scheduler", object(), None))
|
||||||
|
|
||||||
|
assert result == response
|
||||||
|
recv_timeout_calls = [
|
||||||
|
call
|
||||||
|
for call in socket.setsockopt.call_args_list
|
||||||
|
if call.args[0] == zmq.RCVTIMEO
|
||||||
|
]
|
||||||
|
assert recv_timeout_calls == []
|
||||||
|
socket.close.assert_called_once_with()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"invalid_timeout_ms",
|
||||||
|
[0, -1, MAX_SCHEDULER_RPC_TIMEOUT_S * 1000 + 1, True],
|
||||||
|
)
|
||||||
|
def test_scheduler_client_rejects_invalid_override_and_closes_socket(
|
||||||
|
invalid_timeout_ms,
|
||||||
|
):
|
||||||
|
socket = MagicMock()
|
||||||
|
client = SchedulerClient()
|
||||||
|
client.context = SimpleNamespace(socket=lambda _socket_type: socket)
|
||||||
|
client.server_args = SimpleNamespace(scheduler_rpc_timeout=None)
|
||||||
|
|
||||||
|
with pytest.raises(ValueError, match="timeout_ms must be None"):
|
||||||
|
client._forward_one("tcp://scheduler", object(), timeout_ms=invalid_timeout_ms)
|
||||||
|
|
||||||
|
socket.close.assert_called_once_with()
|
||||||
|
|
||||||
|
|
||||||
|
def test_async_scheduler_client_waits_for_delayed_response():
|
||||||
|
async def run_test():
|
||||||
|
context = zmq.asyncio.Context()
|
||||||
|
endpoint = f"inproc://scheduler-{uuid.uuid4().hex}"
|
||||||
|
server = context.socket(zmq.REP)
|
||||||
|
server.bind(endpoint)
|
||||||
|
client = AsyncSchedulerClient()
|
||||||
|
client.context = context
|
||||||
|
client.server_args = SimpleNamespace(scheduler_rpc_timeout=None)
|
||||||
|
|
||||||
|
async def reply():
|
||||||
|
await server.recv()
|
||||||
|
await asyncio.sleep(0.02)
|
||||||
|
await server.send(pickle.dumps({"status": "ok"}))
|
||||||
|
|
||||||
|
reply_task = asyncio.create_task(reply())
|
||||||
|
try:
|
||||||
|
result = await client._forward_one(endpoint, object(), timeout_ms=1000)
|
||||||
|
await reply_task
|
||||||
|
assert result == {"status": "ok"}
|
||||||
|
finally:
|
||||||
|
server.close(linger=0)
|
||||||
|
context.destroy(linger=0)
|
||||||
|
|
||||||
|
asyncio.run(run_test())
|
||||||
|
|
||||||
|
|
||||||
|
def test_async_scheduler_client_honors_explicit_deadline():
|
||||||
|
async def run_test():
|
||||||
|
context = zmq.asyncio.Context()
|
||||||
|
endpoint = f"inproc://scheduler-{uuid.uuid4().hex}"
|
||||||
|
server = context.socket(zmq.REP)
|
||||||
|
server.bind(endpoint)
|
||||||
|
client = AsyncSchedulerClient()
|
||||||
|
client.context = context
|
||||||
|
client.server_args = SimpleNamespace(scheduler_rpc_timeout=None)
|
||||||
|
|
||||||
|
try:
|
||||||
|
with pytest.raises(TimeoutError, match="did not respond"):
|
||||||
|
await client._forward_one(endpoint, object(), timeout_ms=10)
|
||||||
|
finally:
|
||||||
|
server.close(linger=0)
|
||||||
|
context.destroy(linger=0)
|
||||||
|
|
||||||
|
asyncio.run(run_test())
|
||||||
|
|
||||||
|
|
||||||
|
def test_async_scheduler_client_closes_socket_when_cancelled():
|
||||||
|
async def run_test():
|
||||||
|
socket = MagicMock()
|
||||||
|
socket.send = AsyncMock()
|
||||||
|
socket.recv = AsyncMock(side_effect=asyncio.Event().wait)
|
||||||
|
client = AsyncSchedulerClient()
|
||||||
|
client.context = SimpleNamespace(socket=lambda _socket_type: socket)
|
||||||
|
client.server_args = SimpleNamespace(scheduler_rpc_timeout=None)
|
||||||
|
|
||||||
|
task = asyncio.create_task(
|
||||||
|
client._forward_one("tcp://scheduler", object(), timeout_ms=None)
|
||||||
|
)
|
||||||
|
await asyncio.sleep(0)
|
||||||
|
task.cancel()
|
||||||
|
with pytest.raises(asyncio.CancelledError):
|
||||||
|
await task
|
||||||
|
|
||||||
|
socket.close.assert_called_once_with()
|
||||||
|
|
||||||
|
asyncio.run(run_test())
|
||||||
|
|
||||||
|
|
||||||
|
def test_broker_closes_socket_and_context_when_cancelled(monkeypatch):
|
||||||
|
async def run_test():
|
||||||
|
socket = MagicMock()
|
||||||
|
socket.recv = AsyncMock(side_effect=asyncio.Event().wait)
|
||||||
|
context = MagicMock()
|
||||||
|
context.socket.return_value = socket
|
||||||
|
monkeypatch.setattr(zmq.asyncio, "Context", lambda: context)
|
||||||
|
|
||||||
|
task = asyncio.create_task(
|
||||||
|
run_zeromq_broker(SimpleNamespace(broker_port=12345))
|
||||||
|
)
|
||||||
|
await asyncio.sleep(0)
|
||||||
|
task.cancel()
|
||||||
|
with pytest.raises(asyncio.CancelledError):
|
||||||
|
await task
|
||||||
|
|
||||||
|
socket.close.assert_called_once_with(linger=0)
|
||||||
|
context.destroy.assert_called_once_with(linger=0)
|
||||||
|
|
||||||
|
asyncio.run(run_test())
|
||||||
@@ -54,7 +54,10 @@ from sglang.multimodal_gen.runtime.models.dits.qwen_image import (
|
|||||||
from sglang.multimodal_gen.runtime.pipelines.minimax_h3_pipeline import (
|
from sglang.multimodal_gen.runtime.pipelines.minimax_h3_pipeline import (
|
||||||
MiniMaxH3Pipeline,
|
MiniMaxH3Pipeline,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
from sglang.multimodal_gen.runtime.server_args import (
|
||||||
|
MAX_SCHEDULER_RPC_TIMEOUT_S,
|
||||||
|
ServerArgs,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.utils import FlexibleArgumentParser
|
from sglang.multimodal_gen.utils import FlexibleArgumentParser
|
||||||
|
|
||||||
|
|
||||||
@@ -2196,6 +2199,40 @@ class TestDisaggTimeoutArgs(unittest.TestCase):
|
|||||||
self.assertEqual(args.disagg_role, RoleType.DENOISER)
|
self.assertEqual(args.disagg_role, RoleType.DENOISER)
|
||||||
|
|
||||||
|
|
||||||
|
class TestSchedulerRpcTimeoutArgs(unittest.TestCase):
|
||||||
|
def test_scheduler_rpc_timeout_defaults_to_unbounded(self):
|
||||||
|
args = _from_dict_without_model_resolution({"model_path": "/fake"})
|
||||||
|
self.assertIsNone(args.scheduler_rpc_timeout)
|
||||||
|
|
||||||
|
def test_scheduler_rpc_timeout_cli_arg_is_parsed_in_seconds(self):
|
||||||
|
parser = FlexibleArgumentParser()
|
||||||
|
ServerArgs.add_cli_args(parser)
|
||||||
|
argv = [
|
||||||
|
"--model-path",
|
||||||
|
"/fake",
|
||||||
|
"--scheduler-rpc-timeout",
|
||||||
|
"7200",
|
||||||
|
]
|
||||||
|
|
||||||
|
args, _unknown = parser.parse_known_args(argv)
|
||||||
|
self.assertEqual(args.scheduler_rpc_timeout, 7200)
|
||||||
|
|
||||||
|
def test_scheduler_rpc_timeout_rejects_invalid_values(self):
|
||||||
|
invalid_values = (0, -1, MAX_SCHEDULER_RPC_TIMEOUT_S + 1, True, 1.5, "1")
|
||||||
|
|
||||||
|
for invalid_value in invalid_values:
|
||||||
|
with self.subTest(invalid_value=invalid_value):
|
||||||
|
with self.assertRaisesRegex(
|
||||||
|
ValueError, "scheduler_rpc_timeout must be None"
|
||||||
|
):
|
||||||
|
_from_dict_without_model_resolution(
|
||||||
|
{
|
||||||
|
"model_path": "/fake",
|
||||||
|
"scheduler_rpc_timeout": invalid_value,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class TestDisaggTransferBackendArgs(unittest.TestCase):
|
class TestDisaggTransferBackendArgs(unittest.TestCase):
|
||||||
def test_transfer_backend_defaults_to_auto(self):
|
def test_transfer_backend_defaults_to_auto(self):
|
||||||
args = _from_dict_without_model_resolution({"model_path": "/fake"})
|
args = _from_dict_without_model_resolution({"model_path": "/fake"})
|
||||||
|
|||||||
@@ -0,0 +1,47 @@
|
|||||||
|
import asyncio
|
||||||
|
|
||||||
|
from sglang.multimodal_gen.runtime.entrypoints.openai import video_api
|
||||||
|
|
||||||
|
|
||||||
|
def test_video_job_registry_holds_task_until_completion():
|
||||||
|
async def run_test():
|
||||||
|
started = asyncio.Event()
|
||||||
|
finish = asyncio.Event()
|
||||||
|
|
||||||
|
async def job():
|
||||||
|
started.set()
|
||||||
|
await finish.wait()
|
||||||
|
|
||||||
|
task = video_api._start_video_job("job-id", job())
|
||||||
|
await started.wait()
|
||||||
|
assert video_api._VIDEO_JOB_TASKS["job-id"] is task
|
||||||
|
|
||||||
|
finish.set()
|
||||||
|
await task
|
||||||
|
await asyncio.sleep(0)
|
||||||
|
assert "job-id" not in video_api._VIDEO_JOB_TASKS
|
||||||
|
|
||||||
|
asyncio.run(run_test())
|
||||||
|
|
||||||
|
|
||||||
|
def test_shutdown_video_jobs_cancels_and_awaits_cleanup():
|
||||||
|
async def run_test():
|
||||||
|
started = asyncio.Event()
|
||||||
|
cleaned_up = asyncio.Event()
|
||||||
|
|
||||||
|
async def job():
|
||||||
|
started.set()
|
||||||
|
try:
|
||||||
|
await asyncio.Event().wait()
|
||||||
|
finally:
|
||||||
|
cleaned_up.set()
|
||||||
|
|
||||||
|
task = video_api._start_video_job("job-id", job())
|
||||||
|
await started.wait()
|
||||||
|
await video_api.shutdown_video_jobs()
|
||||||
|
|
||||||
|
assert task.cancelled()
|
||||||
|
assert cleaned_up.is_set()
|
||||||
|
assert "job-id" not in video_api._VIDEO_JOB_TASKS
|
||||||
|
|
||||||
|
asyncio.run(run_test())
|
||||||
Reference in New Issue
Block a user