From bc148dfdc8478bcc438e5200eb5c12ad05d0c9bb Mon Sep 17 00:00:00 2001 From: Yifei Suo <42202928+Windswithyou@users.noreply.github.com> Date: Fri, 7 Aug 2026 21:33:23 +0800 Subject: [PATCH] [diffusion] feat: make scheduler rpc deadlines explicit (#33965) Co-authored-by: suoyf Co-authored-by: Mick --- docs/docs/sglang-diffusion/api/cli.mdx | 1 + .../runtime/entrypoints/http_server.py | 6 + .../runtime/entrypoints/openai/video_api.py | 30 +++- .../runtime/scheduler_client.py | 98 +++++++---- .../runtime/server_args/__init__.py | 2 + .../runtime/server_args/server_args.py | 27 +++ .../test/unit/test_scheduler_client.py | 160 ++++++++++++++++++ .../test/unit/test_server_args.py | 39 ++++- .../test/unit/test_video_job_lifecycle.py | 47 +++++ 9 files changed, 377 insertions(+), 33 deletions(-) create mode 100644 python/sglang/multimodal_gen/test/unit/test_scheduler_client.py create mode 100644 python/sglang/multimodal_gen/test/unit/test_video_job_lifecycle.py diff --git a/docs/docs/sglang-diffusion/api/cli.mdx b/docs/docs/sglang-diffusion/api/cli.mdx index bae72c1f9..88018e466 100644 --- a/docs/docs/sglang-diffusion/api/cli.mdx +++ b/docs/docs/sglang-diffusion/api/cli.mdx @@ -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-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 +- `--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). ### Sampling and output diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/http_server.py b/python/sglang/multimodal_gen/runtime/entrypoints/http_server.py index 84fc52682..e29123a2b 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/http_server.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/http_server.py @@ -105,6 +105,9 @@ async def _run_server_warmup_after_http_live( @asynccontextmanager 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 ( async_scheduler_client, run_zeromq_broker, @@ -136,7 +139,10 @@ async def lifespan(app: FastAPI): # On shutdown logger.info("FastAPI app is shutting down...") + await shutdown_video_jobs() broker_task.cancel() + with suppress(asyncio.CancelledError): + await broker_task async_scheduler_client.close() diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/openai/video_api.py b/python/sglang/multimodal_gen/runtime/entrypoints/openai/video_api.py index cb773f77d..d96df3ae6 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/openai/video_api.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/openai/video_api.py @@ -6,6 +6,8 @@ import os import shutil import tempfile import time +from collections.abc import Coroutine +from contextlib import suppress from typing import Any, Dict, Optional from fastapi import ( @@ -61,6 +63,29 @@ _VIDEO_EXTENSIONS = { ".mpg", ".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: @@ -826,14 +851,15 @@ async def create_video( assert batch is not None # Enqueue the job asynchronously and return immediately - asyncio.create_task( + _start_video_job( + request_id, _dispatch_job_async( request_id, batch, scheduler_batches=scheduler_batches, temp_dirs=temp_dirs or None, output_persistent=output_persistent, - ) + ), ) return VideoResponse(**job) diff --git a/python/sglang/multimodal_gen/runtime/scheduler_client.py b/python/sglang/multimodal_gen/runtime/scheduler_client.py index df5dcaad1..aa0ef7c1e 100644 --- a/python/sglang/multimodal_gen/runtime/scheduler_client.py +++ b/python/sglang/multimodal_gen/runtime/scheduler_client.py @@ -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.pipelines_core import Req 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.request_logger import ( 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): """ This function runs as a background task in the FastAPI process. @@ -61,26 +86,32 @@ async def run_zeromq_broker(server_args: ServerArgs): socket.bind(broker_endpoint) logger.info(f"ZMQ Broker is listening for offline jobs on {broker_endpoint}") - while True: - try: - # 1. Receive a request from an offline client - payload = await socket.recv() - request_batch = pickle.loads(payload) - logger.debug("Broker received an offline job from a client.") - - # 2. Forward the request to the main Scheduler via the shared client - response_batch = await async_scheduler_client.forward(request_batch) - - # 3. Send the Scheduler's reply back to the offline client - await socket.send(pickle.dumps(response_batch)) - - except Exception as e: - logger.error(f"Error in ZMQ Broker: {e}", exc_info=True) - # A reply must be sent to prevent the client from hanging + try: + while True: try: - await socket.send(pickle.dumps({"status": "error", "message": str(e)})) - except Exception: - pass + # 1. Receive a request from an offline client + payload = await socket.recv() + request_batch = pickle.loads(payload) + logger.debug("Broker received an offline job from a client.") + + # 2. Forward the request to the main Scheduler via the shared client + response_batch = await async_scheduler_client.forward(request_batch) + + # 3. Send the Scheduler's reply back to the offline client + await socket.send(pickle.dumps(response_batch)) + + except Exception as e: + logger.error(f"Error in ZMQ Broker: {e}", exc_info=True) + # A reply must be sent to prevent the client from hanging + try: + await socket.send( + pickle.dumps({"status": "error", "message": str(e)}) + ) + except Exception: + pass + finally: + socket.close(linger=0) + ctx.destroy(linger=0) 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: socket = self.context.socket(zmq.REQ) - socket.setsockopt(zmq.LINGER, 0) - socket.setsockopt(zmq.RCVTIMEO, timeout_ms if timeout_ms else 6000000) 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.send_pyobj(batch) output_batch = socket.recv_pyobj() @@ -220,7 +252,7 @@ class AsyncSchedulerClient: self.context = 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.""" self.request_logger.log_received_request(batch) if self.context is None: @@ -231,23 +263,29 @@ class AsyncSchedulerClient: endpoints = self.server_args.scheduler_endpoints if isinstance(batch, _CONTROL_REQ_TYPES): # 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) else: 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) 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 # interleaving on one socket's strict send/recv alternation socket = self.context.socket(zmq.REQ) - socket.setsockopt(zmq.LINGER, 0) - # 100 minute timeout - socket.setsockopt(zmq.RCVTIMEO, 6000000) - socket.connect(endpoint) 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)) payload = await socket.recv() output_batch = pickle.loads(payload) diff --git a/python/sglang/multimodal_gen/runtime/server_args/__init__.py b/python/sglang/multimodal_gen/runtime/server_args/__init__.py index cdcf73a7a..f31f13d3e 100644 --- a/python/sglang/multimodal_gen/runtime/server_args/__init__.py +++ b/python/sglang/multimodal_gen/runtime/server_args/__init__.py @@ -7,6 +7,7 @@ from sglang.multimodal_gen.runtime.server_args.server_args import ( DEFAULT_BCG_TEXT_BUCKETS, LORA_MERGE_MODES, LTX2_TWO_STAGE_DEVICE_MODE_CHOICES, + MAX_SCHEDULER_RPC_TIMEOUT_S, Backend, PortArgs, ServerArgs, @@ -24,6 +25,7 @@ __all__ = [ "DEFAULT_BCG_TEXT_BUCKETS", "LORA_MERGE_MODES", "LTX2_TWO_STAGE_DEVICE_MODE_CHOICES", + "MAX_SCHEDULER_RPC_TIMEOUT_S", "PortArgs", "ServerArgs", "_normalize_ltx2_two_stage_device_mode", diff --git a/python/sglang/multimodal_gen/runtime/server_args/server_args.py b/python/sglang/multimodal_gen/runtime/server_args/server_args.py index e8d6cb055..3a8d7df3d 100644 --- a/python/sglang/multimodal_gen/runtime/server_args/server_args.py +++ b/python/sglang/multimodal_gen/runtime/server_args/server_args.py @@ -71,6 +71,7 @@ LTX2_TWO_STAGE_PIPELINE_NAMES = ("LTX2TwoStagePipeline", "LTX2TwoStageHQPipeline # H200-class GPUs (>=130 GiB total) can usually keep both LTX2 DiTs resident. LTX2_RESIDENT_AUTO_ENABLE_MEM_GB = 130 LORA_MERGE_MODES = ("auto", "merge", "dynamic") +MAX_SCHEDULER_RPC_TIMEOUT_S = 2_147_483 # Mirrors AttentionBackend.supports_ring_rotation; the name-level check # runs before backend classes are importable on every platform. RING_CAPABLE_ATTENTION_BACKENDS = ("fa", "sage_attn") @@ -245,6 +246,7 @@ class ServerArgs(DisaggServerArgsMixin): hsdp_replicate_dim: int = 1 hsdp_shard_dim: Optional[int] = None dist_timeout: int | None = 3600 # 1 hour + scheduler_rpc_timeout: int | None = None pipeline_config: PipelineConfig = field(default_factory=PipelineConfig, repr=False) @@ -484,6 +486,7 @@ class ServerArgs(DisaggServerArgsMixin): def _validate_parameters(self): """check consistency and raise errors for invalid configs""" + self._validate_scheduler_rpc_timeout() self._validate_pipeline() self._validate_offload() if not current_platform.is_cpu(): @@ -493,6 +496,20 @@ class ServerArgs(DisaggServerArgsMixin): self._validate_breakable_cuda_graph() 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, ...]: """Sorted, de-duplicated, positive BCG text buckets. @@ -1553,6 +1570,16 @@ class ServerArgs(DisaggServerArgsMixin): help="Timeout for torch.distributed operations in seconds. " "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) diff --git a/python/sglang/multimodal_gen/test/unit/test_scheduler_client.py b/python/sglang/multimodal_gen/test/unit/test_scheduler_client.py new file mode 100644 index 000000000..b7ed61498 --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/test_scheduler_client.py @@ -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()) diff --git a/python/sglang/multimodal_gen/test/unit/test_server_args.py b/python/sglang/multimodal_gen/test/unit/test_server_args.py index 5d94d9b54..ba07357e7 100644 --- a/python/sglang/multimodal_gen/test/unit/test_server_args.py +++ b/python/sglang/multimodal_gen/test/unit/test_server_args.py @@ -54,7 +54,10 @@ from sglang.multimodal_gen.runtime.models.dits.qwen_image import ( from sglang.multimodal_gen.runtime.pipelines.minimax_h3_pipeline import ( 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 @@ -2196,6 +2199,40 @@ class TestDisaggTimeoutArgs(unittest.TestCase): 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): def test_transfer_backend_defaults_to_auto(self): args = _from_dict_without_model_resolution({"model_path": "/fake"}) diff --git a/python/sglang/multimodal_gen/test/unit/test_video_job_lifecycle.py b/python/sglang/multimodal_gen/test/unit/test_video_job_lifecycle.py new file mode 100644 index 000000000..82e07e599 --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/test_video_job_lifecycle.py @@ -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())