from __future__ import annotations import asyncio import concurrent.futures import logging import queue import threading import time from dataclasses import dataclass from types import SimpleNamespace from typing import Any import janus import pytest from vllm.outputs import CompletionOutput, RequestOutput from vllm.sampling_params import SamplingParams from vllm_omni.engine.messages import ( AbortRequestMessage, AddCompanionRequestMessage, CollectiveRPCRequestMessage, CollectiveRPCResultMessage, OutputMessage, ShutdownRequestMessage, StageSubmissionMessage, ) from vllm_omni.engine.orchestrator import ( Orchestrator, OrchestratorRequestState, _build_terminal_empty_output, _infer_stage_audio_sample_rate, ) from vllm_omni.engine.stage_pool import StagePool from vllm_omni.inputs.data import OmniDiffusionSamplingParams from vllm_omni.outputs import OmniRequestOutput pytestmark = [pytest.mark.core_model, pytest.mark.cpu] @dataclass class OrchestratorFixture: orchestrator: Orchestrator request_sync_q: Any output_sync_q: Any queues: tuple[janus.Queue, ...] thread: threading.Thread result_future: concurrent.futures.Future[None] class FakeStageClient: def __init__( self, *, stage_type: str = "llm", final_output: bool = False, final_output_type: str = "text", next_inputs: list[dict] | None = None, engine_input_source: list[int] | None = None, is_comprehension: bool = False, model_stage: str | None = None, kv_sender_info: dict[str, Any] | None = None, ) -> None: self.stage_id = 0 self.replica_id = 0 self.stage_type = stage_type self.final_output = final_output self.final_output_type = final_output_type self.default_sampling_params = SamplingParams(max_tokens=1) self.requires_multimodal_data = False self.engine_input_source = list(engine_input_source or [0]) self.is_comprehension = is_comprehension self.model_stage = model_stage self.next_inputs = list(next_inputs or []) self.custom_process_input_func = None self._kv_sender_info = dict(kv_sender_info) if kv_sender_info is not None else None self.add_request_calls: list[tuple] = [] self.abort_calls: list[list[str]] = [] self.collective_rpc_calls: list[tuple[str, float | None, tuple[Any, ...], dict[str, Any]]] = [] self.shutdown_calls = 0 self._engine_core_outputs = queue.Queue() self._diffusion_outputs = queue.Queue() # Orchestrator-facing interface. async def add_request_async(self, *args, **kwargs) -> None: self.add_request_calls.append(args) async def get_output_async(self): try: return self._engine_core_outputs.get_nowait() except queue.Empty: return SimpleNamespace(outputs=[]) def get_diffusion_output_nowait(self): try: return self._diffusion_outputs.get_nowait() except queue.Empty: return None def set_engine_outputs(self, outputs) -> None: return None def process_engine_inputs(self, source_outputs, prompt=None, streaming_context=None): return list(self.next_inputs) async def abort_requests_async(self, request_ids: list[str]) -> None: self.abort_calls.append(list(request_ids)) async def collective_rpc_async( self, *, method: str, timeout: float | None = None, args: tuple[Any, ...] = (), kwargs: dict[str, Any] | None = None, ) -> Any: normalized_kwargs = dict(kwargs or {}) self.collective_rpc_calls.append((method, timeout, args, normalized_kwargs)) return { "supported": False, "todo": True, "reason": f"{self.__class__.__name__}.collective_rpc_async is not implemented yet", } def get_kv_sender_info(self) -> dict[str, Any] | None: if self._kv_sender_info is None: return None return dict(self._kv_sender_info) def check_health(self) -> None: return None def shutdown(self) -> None: self.shutdown_calls += 1 # Test helpers for seeding fake stage outputs. def push_engine_core_outputs(self, outputs) -> None: self._engine_core_outputs.put_nowait(outputs) def push_diffusion_output(self, output) -> None: self._diffusion_outputs.put_nowait(output) def test_terminal_empty_audio_output_uses_stage_sample_rate() -> None: final_stage = FakeStageClient(final_output=True, final_output_type="audio") final_stage.sample_rate = 44100 final_pool = SimpleNamespace(stage_client=final_stage, _stage_vllm_config=None) terminal_output = _build_terminal_empty_output( "req-1", final_output_type="audio", audio_sample_rate=_infer_stage_audio_sample_rate(final_pool), ) assert terminal_output.outputs[0].multimodal_output["sr"] == 44100 class FakeCollectiveRpcStageClient(FakeStageClient): def __init__(self, *args, rpc_result: Any = None, **kwargs) -> None: super().__init__(*args, **kwargs) self.rpc_result = rpc_result async def collective_rpc_async( self, *, method: str, timeout: float | None = None, args: tuple[Any, ...] = (), kwargs: dict[str, Any] | None = None, ) -> Any: normalized_kwargs = dict(kwargs or {}) self.collective_rpc_calls.append((method, timeout, args, normalized_kwargs)) return self.rpc_result class FakeOutputProcessor: def __init__(self, *, request_outputs: list[object] | None = None) -> None: self.request_outputs = list(request_outputs or []) self.add_request_calls: list[tuple[tuple[Any, ...], dict[str, Any]]] = [] self.abort_calls: list[list[str]] = [] def add_request(self, *args, **kwargs) -> None: self.add_request_calls.append((args, kwargs)) return None def process_outputs(self, *_args, **_kwargs): return SimpleNamespace( request_outputs=list(self.request_outputs), reqs_to_abort=[], ) def abort_requests(self, request_ids, internal: bool = False): self.abort_calls.append(request_ids) return request_ids def update_scheduler_stats(self, _scheduler_stats) -> None: return None def _sampling_params(max_tokens: int = 4) -> SamplingParams: return SamplingParams(max_tokens=max_tokens) def _engine_core_outputs(tag: str, timestamp: float) -> SimpleNamespace: return SimpleNamespace(outputs=[tag], timestamp=timestamp, scheduler_stats=None) def _build_request_output( request_id: str, *, token_ids: list[int] | None = None, prompt_token_ids: list[int] | None = None, finished: bool = True, text: str = "test", ) -> RequestOutput: completion = CompletionOutput( index=0, text=text, token_ids=list(token_ids or [1, 2]), cumulative_logprob=0.0, logprobs=None, finish_reason="stop" if finished else None, stop_reason=None, ) return RequestOutput( request_id=request_id, prompt="prompt", prompt_token_ids=list(prompt_token_ids or [10, 11]), prompt_logprobs=None, outputs=[completion], finished=finished, metrics=None, lora_request=None, ) def _build_stage_pools( stage_clients: list[list[FakeStageClient]], *, output_processors: list[FakeOutputProcessor] | None = None, stage_vllm_configs: list[object] | None = None, ) -> list[StagePool]: """Build StagePool list from per-stage replica lists. ``stage_clients[i]`` is the list of FakeStageClient replicas for stage i. """ num_stages = len(stage_clients) if output_processors is None: output_processors = [FakeOutputProcessor() for _ in stage_clients] if stage_vllm_configs is None: stage_vllm_configs = [SimpleNamespace(model_config=SimpleNamespace(max_model_len=64)) for _ in stage_clients] pools: list[StagePool] = [] for stage_id in range(num_stages): clients = stage_clients[stage_id] if clients[0].stage_type == "diffusion": pools.append(StagePool(stage_id, clients[0])) else: pools.append( StagePool( stage_id, clients, output_processor=output_processors[stage_id], stage_vllm_config=stage_vllm_configs[stage_id], ) ) return pools def _build_harness( stage_clients: list[object], *, output_processors: list[object] | None = None, stage_vllm_configs: list[object] | None = None, async_chunk: bool = False, stage_pools: list[StagePool] | None = None, ) -> OrchestratorFixture: """Build an Orchestrator test harness. Accepts either pre-built ``stage_pools`` or flat lists of single-replica clients/processors. """ if stage_pools is None: # Wrap flat lists into per-stage single-replica lists. nested_clients = [[c] for c in stage_clients] stage_pools = _build_stage_pools( nested_clients, output_processors=output_processors, stage_vllm_configs=stage_vllm_configs, ) ready_future: concurrent.futures.Future[tuple[Orchestrator, janus.Queue, janus.Queue, janus.Queue]] = ( concurrent.futures.Future() ) result_future: concurrent.futures.Future[None] = concurrent.futures.Future() def _runner() -> None: loop = asyncio.new_event_loop() asyncio.set_event_loop(loop) async def _run() -> None: request_queue = janus.Queue() output_queue = janus.Queue() rpc_queue = janus.Queue() orchestrator = Orchestrator( request_async_queue=request_queue.async_q, output_async_queue=output_queue.async_q, rpc_async_queue=rpc_queue.async_q, stage_pools=stage_pools, async_chunk=async_chunk, ) ready_future.set_result((orchestrator, request_queue, output_queue, rpc_queue)) await orchestrator.run() try: loop.run_until_complete(_run()) result_future.set_result(None) except Exception as exc: result_future.set_exception(exc) finally: try: pending = [task for task in asyncio.all_tasks(loop) if not task.done()] for task in pending: task.cancel() if pending: loop.run_until_complete(asyncio.gather(*pending, return_exceptions=True)) loop.run_until_complete(loop.shutdown_asyncgens()) finally: asyncio.set_event_loop(None) loop.close() thread = threading.Thread(target=_runner, daemon=True, name="test-orchestrator") thread.start() orchestrator, request_queue, output_queue, rpc_queue = ready_future.result(timeout=5) return OrchestratorFixture( orchestrator=orchestrator, request_sync_q=request_queue.sync_q, output_sync_q=output_queue.sync_q, queues=(request_queue, output_queue, rpc_queue), thread=thread, result_future=result_future, ) async def _shutdown_orchestrator(orchestrator_fixture: OrchestratorFixture) -> None: orchestrator_fixture.request_sync_q.put_nowait(ShutdownRequestMessage()) await asyncio.to_thread(orchestrator_fixture.thread.join, 5) if orchestrator_fixture.thread.is_alive(): raise AssertionError("Timed out waiting for orchestrator thread shutdown") orchestrator_fixture.result_future.result(timeout=0) async def _wait_for(predicate, *, timeout: float = 2.0) -> None: deadline = time.monotonic() + timeout while not predicate(): if time.monotonic() >= deadline: raise AssertionError("Timed out waiting for predicate") await asyncio.sleep(0.01) async def _get_output_message(orchestrator_fixture: OrchestratorFixture, *, timeout: float = 2.0) -> OutputMessage: deadline = time.monotonic() + timeout while True: if time.monotonic() >= deadline: raise AssertionError("Timed out waiting for orchestrator output") try: msg = orchestrator_fixture.output_sync_q.get_nowait() except queue.Empty: await asyncio.sleep(0.01) continue if isinstance(msg, OutputMessage): return msg async def _get_rpc_message( orchestrator_fixture: OrchestratorFixture, *, timeout: float = 2.0, ) -> CollectiveRPCResultMessage: deadline = time.monotonic() + timeout rpc_sync_q = orchestrator_fixture.queues[2].sync_q while True: if time.monotonic() >= deadline: raise AssertionError("Timed out waiting for orchestrator rpc output") try: return rpc_sync_q.get_nowait() except queue.Empty: await asyncio.sleep(0.01) async def _enqueue_add_request( orchestrator_fixture: OrchestratorFixture, *, request_id: str, prompt, original_prompt, sampling_params_list, final_stage_id: int, ) -> None: orchestrator_fixture.request_sync_q.put_nowait( StageSubmissionMessage( type="add_request", request_id=request_id, prompt=prompt, original_prompt=original_prompt, output_prompt_text=None, sampling_params_list=sampling_params_list, final_stage_id=final_stage_id, preprocess_ms=0.0, request_timestamp=time.time(), enqueue_ts=time.perf_counter(), ) ) async def _enqueue_abort_request(orchestrator_fixture: OrchestratorFixture, request_ids: list[str]) -> None: orchestrator_fixture.request_sync_q.put_nowait(AbortRequestMessage(request_ids=request_ids)) @pytest.fixture def orchestrator_factory(): fixtures: list[OrchestratorFixture] = [] def _factory(*args, **kwargs) -> OrchestratorFixture: fixture = _build_harness(*args, **kwargs) fixtures.append(fixture) return fixture yield _factory for fixture in fixtures: if fixture.thread.is_alive(): fixture.request_sync_q.put_nowait(ShutdownRequestMessage()) fixture.thread.join(timeout=5) for q in fixture.queues: q.close() # --------------------------------------------------------------------------- # Existing single-replica tests (adapted to StagePool interface) # --------------------------------------------------------------------------- @pytest.mark.asyncio async def test_run_two_stage_llm(orchestrator_factory) -> None: stage0 = FakeStageClient(stage_type="llm", final_output=False) stage1 = FakeStageClient( stage_type="llm", final_output=True, next_inputs=[{"prompt_token_ids": [7, 8, 9]}], ) processors = [ FakeOutputProcessor(request_outputs=[_build_request_output("req-llm", token_ids=[3, 4], finished=True)]), FakeOutputProcessor(request_outputs=[_build_request_output("req-llm", token_ids=[10, 11], finished=True)]), ] orchestrator_fixture = orchestrator_factory([stage0, stage1], output_processors=processors) request = SimpleNamespace(request_id="req-llm", prompt_token_ids=[1, 2, 3]) try: await _enqueue_add_request( orchestrator_fixture, request_id="req-llm", prompt=request, original_prompt={"prompt": "hello"}, sampling_params_list=[_sampling_params(), _sampling_params()], final_stage_id=1, ) await _wait_for(lambda: len(stage0.add_request_calls) == 1) stage0.push_engine_core_outputs(_engine_core_outputs("stage0-raw", 1.0)) await _wait_for(lambda: len(stage1.add_request_calls) == 1) stage1_request = stage1.add_request_calls[0][0] assert stage1_request.request_id == "req-llm" assert stage1_request.prompt_token_ids == [7, 8, 9] stage1.push_engine_core_outputs(_engine_core_outputs("stage1-raw", 2.0)) output_msg = await _get_output_message(orchestrator_fixture) assert output_msg.request_id == "req-llm" assert output_msg.stage_id == 1 assert output_msg.finished is True assert output_msg.engine_outputs.request_id == "req-llm" assert "req-llm" not in orchestrator_fixture.orchestrator.request_states finally: await _shutdown_orchestrator(orchestrator_fixture) @pytest.mark.asyncio async def test_run_single_stage_diffusion(orchestrator_factory) -> None: stage0 = FakeStageClient(stage_type="diffusion", final_output=True, final_output_type="image") orchestrator_fixture = orchestrator_factory([stage0]) params = OmniDiffusionSamplingParams() try: await _enqueue_add_request( orchestrator_fixture, request_id="req-diff", prompt={"prompt": "draw a cat"}, original_prompt={"prompt": "draw a cat"}, sampling_params_list=[params], final_stage_id=0, ) await _wait_for(lambda: len(stage0.add_request_calls) == 1) stage0.push_diffusion_output( OmniRequestOutput.from_diffusion( request_id="req-diff", images=[], final_output_type="image", ) ) output_msg = await _get_output_message(orchestrator_fixture) assert output_msg.request_id == "req-diff" assert output_msg.stage_id == 0 assert output_msg.finished is True assert output_msg.engine_outputs.request_id == "req-diff" assert "req-diff" not in orchestrator_fixture.orchestrator.request_states finally: await _shutdown_orchestrator(orchestrator_fixture) @pytest.mark.asyncio async def test_run_single_stage_diffusion_streaming_forwards_intermediate_chunks(orchestrator_factory) -> None: """Intermediate diffusion chunks (finished=False) reach the frontend before the final chunk.""" stage0 = FakeStageClient(stage_type="diffusion", final_output=True, final_output_type="image") orchestrator_fixture = orchestrator_factory([stage0]) params = OmniDiffusionSamplingParams() try: await _enqueue_add_request( orchestrator_fixture, request_id="req-stream", prompt={"prompt": "draw a cat"}, original_prompt={"prompt": "draw a cat"}, sampling_params_list=[params], final_stage_id=0, ) await _wait_for(lambda: len(stage0.add_request_calls) == 1) stage0.push_diffusion_output( OmniRequestOutput.from_diffusion( request_id="req-stream", images=[], final_output_type="image", custom_output={"chunk": 0}, finished=False, ) ) stage0.push_diffusion_output( OmniRequestOutput.from_diffusion( request_id="req-stream", images=[], final_output_type="image", custom_output={"chunk": 1}, finished=True, ) ) output_msgs: list[OutputMessage] = [] deadline = time.monotonic() + 2.0 while not output_msgs or not output_msgs[-1].finished: if time.monotonic() >= deadline: raise AssertionError( f"Timed out waiting for finished orchestrator output, got {len(output_msgs)} message(s)" ) try: msg = orchestrator_fixture.output_sync_q.get_nowait() except queue.Empty: await asyncio.sleep(0.01) continue if isinstance(msg, OutputMessage): output_msgs.append(msg) assert [msg.request_id for msg in output_msgs] == ["req-stream", "req-stream"] assert [msg.finished for msg in output_msgs] == [False, True] assert [msg.engine_outputs.finished for msg in output_msgs] == [False, True] assert [msg.engine_outputs.custom_output["chunk"] for msg in output_msgs] == [0, 1] await _wait_for(lambda: "req-stream" not in orchestrator_fixture.orchestrator.request_states) finally: await _shutdown_orchestrator(orchestrator_fixture) @pytest.mark.asyncio async def test_run_llm_to_diffusion(orchestrator_factory) -> None: stage0 = FakeStageClient(stage_type="llm", final_output=False) stage1 = FakeStageClient(stage_type="diffusion", final_output=True, final_output_type="image") processors = [ FakeOutputProcessor(request_outputs=[_build_request_output("req-img", token_ids=[3, 4], finished=True)]), FakeOutputProcessor(), ] orchestrator_fixture = orchestrator_factory([stage0, stage1], output_processors=processors) request = SimpleNamespace(request_id="req-img", prompt_token_ids=[1, 2, 3]) params = OmniDiffusionSamplingParams() original_prompt = {"prompt": "draw a fox"} try: await _enqueue_add_request( orchestrator_fixture, request_id="req-img", prompt=request, original_prompt=original_prompt, sampling_params_list=[_sampling_params(), params], final_stage_id=1, ) await _wait_for(lambda: len(stage0.add_request_calls) == 1) stage0.push_engine_core_outputs(_engine_core_outputs("stage0-raw", 1.0)) await _wait_for(lambda: len(stage1.add_request_calls) == 1) assert stage1.add_request_calls[0] == ("req-img", original_prompt, params) stage1.push_diffusion_output( OmniRequestOutput.from_diffusion( request_id="req-img", images=[], final_output_type="image", ) ) output_msg = await _get_output_message(orchestrator_fixture) assert output_msg.request_id == "req-img" assert output_msg.stage_id == 1 assert output_msg.finished is True assert output_msg.engine_outputs.request_id == "req-img" assert "req-img" not in orchestrator_fixture.orchestrator.request_states finally: await _shutdown_orchestrator(orchestrator_fixture) @pytest.mark.asyncio async def test_run_async_chunk(orchestrator_factory) -> None: stage0 = FakeStageClient(stage_type="llm", final_output=False) stage1 = FakeStageClient(stage_type="llm", final_output=True) processors = [ FakeOutputProcessor(request_outputs=[_build_request_output("req-async", token_ids=[1], finished=True)]), FakeOutputProcessor(request_outputs=[_build_request_output("req-async", token_ids=[20, 21], finished=True)]), ] orchestrator_fixture = orchestrator_factory( [stage0, stage1], output_processors=processors, async_chunk=True, ) request = SimpleNamespace(request_id="req-async", prompt_token_ids=[1, 2, 3, 4]) try: await _enqueue_add_request( orchestrator_fixture, request_id="req-async", prompt=request, original_prompt={"prompt": "hello async"}, sampling_params_list=[_sampling_params(), _sampling_params()], final_stage_id=1, ) await _wait_for(lambda: len(stage1.add_request_calls) == 1) prewarmed_request = stage1.add_request_calls[0][0] assert prewarmed_request.request_id == "req-async" assert prewarmed_request.prompt_token_ids assert all(token_id == 0 for token_id in prewarmed_request.prompt_token_ids) stage1.push_engine_core_outputs(_engine_core_outputs("stage1-final", 3.0)) output_msg = await _get_output_message(orchestrator_fixture) assert output_msg.request_id == "req-async" assert output_msg.stage_id == 1 assert output_msg.finished is True assert "req-async" not in orchestrator_fixture.orchestrator.request_states finally: await _shutdown_orchestrator(orchestrator_fixture) @pytest.mark.asyncio async def test_run_shutdown(orchestrator_factory) -> None: stages = [ FakeStageClient(stage_type="llm", final_output=False), FakeStageClient(stage_type="diffusion", final_output=True, final_output_type="image"), ] orchestrator_fixture = orchestrator_factory(stages) await _shutdown_orchestrator(orchestrator_fixture) assert not orchestrator_fixture.thread.is_alive() for stage in stages: assert stage.shutdown_calls == 1 @pytest.mark.asyncio async def test_run_abort(orchestrator_factory) -> None: stages = [ FakeStageClient(stage_type="llm", final_output=False), FakeStageClient(stage_type="llm", final_output=True), ] processors = [ FakeOutputProcessor(request_outputs=[_build_request_output("req-abort", token_ids=[1], finished=True)]), FakeOutputProcessor(request_outputs=[_build_request_output("req-abort", token_ids=[2], finished=True)]), ] orchestrator_fixture = orchestrator_factory(stages, output_processors=processors) request = SimpleNamespace(request_id="req-abort", prompt_token_ids=[1, 2, 3]) try: await _enqueue_add_request( orchestrator_fixture, request_id="req-abort", prompt=request, original_prompt={"prompt": "cancel me"}, sampling_params_list=[_sampling_params(), _sampling_params()], final_stage_id=1, ) await _wait_for(lambda: len(stages[0].add_request_calls) == 1) await _enqueue_abort_request(orchestrator_fixture, ["req-abort"]) await _wait_for(lambda: bool(stages[0].abort_calls)) assert stages[0].abort_calls == [["req-abort"]] assert stages[1].abort_calls == [] assert "req-abort" not in orchestrator_fixture.orchestrator.request_states finally: await _shutdown_orchestrator(orchestrator_fixture) # --------------------------------------------------------------------------- # Multi-replica tests # --------------------------------------------------------------------------- @pytest.mark.asyncio async def test_multi_replica_round_robin_distribution(orchestrator_factory) -> None: """Two replicas at stage-0, single replica at stage-1. Send two requests — they should land on different stage-0 replicas (round-robin), then both forward to the single stage-1 replica. """ stage0_r0 = FakeStageClient(stage_type="llm", final_output=False) stage0_r1 = FakeStageClient(stage_type="llm", final_output=False) stage1 = FakeStageClient( stage_type="llm", final_output=True, next_inputs=[{"prompt_token_ids": [7, 8]}], ) proc0 = FakeOutputProcessor(request_outputs=[_build_request_output("req-0", token_ids=[3], finished=True)]) proc1 = FakeOutputProcessor(request_outputs=[_build_request_output("req-0", token_ids=[10], finished=True)]) default_vllm_cfg = SimpleNamespace(model_config=SimpleNamespace(max_model_len=64)) stage_pools = _build_stage_pools( [[stage0_r0, stage0_r1], [stage1]], output_processors=[proc0, proc1], stage_vllm_configs=[default_vllm_cfg, default_vllm_cfg], ) orchestrator_fixture = orchestrator_factory([], stage_pools=stage_pools) try: # Request 0 → should land on replica 0 (RR starts at 0) await _enqueue_add_request( orchestrator_fixture, request_id="req-0", prompt=SimpleNamespace(request_id="req-0", prompt_token_ids=[1, 2]), original_prompt={"prompt": "hello 0"}, sampling_params_list=[_sampling_params(), _sampling_params()], final_stage_id=1, ) await _wait_for(lambda: len(stage0_r0.add_request_calls) == 1) assert len(stage0_r1.add_request_calls) == 0 # Request 1 → should land on replica 1 (RR advances) await _enqueue_add_request( orchestrator_fixture, request_id="req-1", prompt=SimpleNamespace(request_id="req-1", prompt_token_ids=[5, 6]), original_prompt={"prompt": "hello 1"}, sampling_params_list=[_sampling_params(), _sampling_params()], final_stage_id=1, ) await _wait_for(lambda: len(stage0_r1.add_request_calls) == 1) assert len(stage0_r0.add_request_calls) == 1 # unchanged # Complete req-0 at stage-0 replica-0 → should forward to stage-1 stage0_r0.push_engine_core_outputs(_engine_core_outputs("s0r0-raw", 1.0)) await _wait_for(lambda: len(stage1.add_request_calls) == 1) assert stage1.add_request_calls[0][0].request_id == "req-0" # Complete req-0 at stage-1 → final output proc1.request_outputs = [_build_request_output("req-0", token_ids=[10], finished=True)] stage1.push_engine_core_outputs(_engine_core_outputs("s1-raw", 2.0)) output_msg = await _get_output_message(orchestrator_fixture) assert output_msg.request_id == "req-0" assert output_msg.stage_id == 1 assert output_msg.finished is True assert "req-0" not in orchestrator_fixture.orchestrator.request_states finally: await _shutdown_orchestrator(orchestrator_fixture) @pytest.mark.asyncio async def test_multi_replica_abort_broadcasts_to_all_replicas(orchestrator_factory) -> None: """Abort must be sent to every replica across all stages.""" stage0_r0 = FakeStageClient(stage_type="llm", final_output=False) stage0_r1 = FakeStageClient(stage_type="llm", final_output=False) stage1 = FakeStageClient(stage_type="llm", final_output=True) proc0 = FakeOutputProcessor() proc1 = FakeOutputProcessor() default_vllm_cfg = SimpleNamespace(model_config=SimpleNamespace(max_model_len=64)) stage_pools = _build_stage_pools( [[stage0_r0, stage0_r1], [stage1]], output_processors=[proc0, proc1], stage_vllm_configs=[default_vllm_cfg, default_vllm_cfg], ) orchestrator_fixture = orchestrator_factory([], stage_pools=stage_pools) try: await _enqueue_add_request( orchestrator_fixture, request_id="req-abort-mr", prompt=SimpleNamespace(request_id="req-abort-mr", prompt_token_ids=[1]), original_prompt={"prompt": "cancel"}, sampling_params_list=[_sampling_params(), _sampling_params()], final_stage_id=1, ) await _wait_for(lambda: len(stage0_r0.add_request_calls) == 1) await _enqueue_abort_request(orchestrator_fixture, ["req-abort-mr"]) await _wait_for(lambda: bool(stage0_r0.abort_calls)) assert stage0_r0.abort_calls == [["req-abort-mr"]] assert stage0_r1.abort_calls == [] assert stage1.abort_calls == [] assert "req-abort-mr" not in orchestrator_fixture.orchestrator.request_states finally: await _shutdown_orchestrator(orchestrator_fixture) @pytest.mark.asyncio async def test_multi_replica_shutdown_all_replicas(orchestrator_factory) -> None: """Shutdown must shut down every replica across all stages.""" stage0_r0 = FakeStageClient(stage_type="llm", final_output=False) stage0_r1 = FakeStageClient(stage_type="llm", final_output=False) stage1 = FakeStageClient(stage_type="llm", final_output=True) default_vllm_cfg = SimpleNamespace(model_config=SimpleNamespace(max_model_len=64)) stage_pools = _build_stage_pools( [[stage0_r0, stage0_r1], [stage1]], stage_vllm_configs=[default_vllm_cfg, default_vllm_cfg], ) orchestrator_fixture = orchestrator_factory([], stage_pools=stage_pools) await _shutdown_orchestrator(orchestrator_fixture) assert not orchestrator_fixture.thread.is_alive() for client in [stage0_r0, stage0_r1, stage1]: assert client.shutdown_calls == 1 @pytest.mark.asyncio async def test_stage_pool_submit_update_reuses_existing_binding() -> None: """A request admitted to one replica must keep using that replica on updates.""" stage0_r0 = FakeStageClient(stage_type="llm", final_output=False) stage0_r1 = FakeStageClient(stage_type="llm", final_output=False) pool = StagePool( 0, [stage0_r0, stage0_r1], output_processor=FakeOutputProcessor(), stage_vllm_config=SimpleNamespace(model_config=SimpleNamespace(max_model_len=64)), ) req0_state = OrchestratorRequestState( request_id="req-0", sampling_params_list=[_sampling_params()], final_stage_id=0, ) req1_state = OrchestratorRequestState( request_id="req-1", sampling_params_list=[_sampling_params()], final_stage_id=0, ) await pool.submit_initial("req-0", req0_state, SimpleNamespace(request_id="req-0", prompt_token_ids=[1, 2])) await pool.submit_update("req-0", req0_state, SimpleNamespace(request_id="req-0", prompt_token_ids=[3])) await pool.submit_initial("req-1", req1_state, SimpleNamespace(request_id="req-1", prompt_token_ids=[4, 5])) await pool.submit_update("req-1", req1_state, SimpleNamespace(request_id="req-1", prompt_token_ids=[6])) assert pool.get_bound_replica_id("req-0") == 0 assert pool.get_bound_replica_id("req-1") == 1 assert len(stage0_r0.add_request_calls) == 2 assert len(stage0_r1.add_request_calls) == 2 assert stage0_r0.add_request_calls[0][0].request_id == "req-0" assert stage0_r0.add_request_calls[1][0].request_id == "req-0" assert stage0_r1.add_request_calls[0][0].request_id == "req-1" assert stage0_r1.add_request_calls[1][0].request_id == "req-1" @pytest.mark.asyncio async def test_stage_pool_submit_update_refreshes_output_processor_state() -> None: output_processor = FakeOutputProcessor() class AssertingStageClient(FakeStageClient): async def add_request_async(self, *args, **kwargs) -> None: if len(self.add_request_calls) == 1: prompts = [call_kwargs["prompt"] for _, call_kwargs in output_processor.add_request_calls] assert prompts == ["seg-1", "seg-2"] await super().add_request_async(*args, **kwargs) stage0 = AssertingStageClient(stage_type="llm", final_output=False) pool = StagePool( 0, [stage0], output_processor=output_processor, stage_vllm_config=SimpleNamespace(model_config=SimpleNamespace(max_model_len=64)), ) req_state = OrchestratorRequestState( request_id="req-0", sampling_params_list=[_sampling_params()], final_stage_id=0, ) await pool.submit_initial( "req-0", req_state, SimpleNamespace(request_id="req-0", prompt_token_ids=[1, 2]), prompt_text="seg-1", ) await pool.submit_update( "req-0", req_state, SimpleNamespace(request_id="req-0", prompt_token_ids=[3], resumable=True), prompt_text="seg-2", ) assert len(output_processor.add_request_calls) == 2 assert output_processor.add_request_calls[1][1]["prompt"] == "seg-2" @pytest.mark.asyncio async def test_handle_streaming_update_passes_prompt_text_to_stage_pool() -> None: class RecordingPool: def __init__(self) -> None: self.calls: list[tuple[str, Any]] = [] async def submit_update(self, request_id, req_state, request, *, prompt_text=None) -> int: self.calls.append((request_id, prompt_text)) return 0 pool = RecordingPool() orchestrator = object.__new__(Orchestrator) orchestrator.async_chunk = False orchestrator.request_states = { "req-stream": OrchestratorRequestState( request_id="req-stream", sampling_params_list=[_sampling_params()], final_stage_id=0, ) } orchestrator.stage_pools = [pool] await orchestrator._handle_streaming_update( StageSubmissionMessage( type="streaming_update", request_id="req-stream", prompt=SimpleNamespace(request_id="req-stream", prompt_token_ids=[1], resumable=True), original_prompt={"prompt": "segment-2"}, output_prompt_text="segment-2", sampling_params_list=[_sampling_params()], final_stage_id=0, preprocess_ms=0.0, request_timestamp=time.time(), enqueue_ts=time.perf_counter(), ) ) assert pool.calls == [("req-stream", "segment-2")] assert orchestrator.request_states["req-stream"].streaming.enabled is True @pytest.mark.asyncio async def test_stage_pool_submit_initial_rolls_back_output_processor_when_client_submit_fails() -> None: class FailingStageClient(FakeStageClient): async def add_request_async(self, *args, **kwargs) -> None: raise RuntimeError("submit failed") class TrackingOutputProcessor(FakeOutputProcessor): def __init__(self) -> None: super().__init__() self.added_request_ids: list[str] = [] self.removed_request_ids: list[str] = [] def add_request(self, request, *_args, **_kwargs) -> None: self.added_request_ids.append(request.request_id) def remove_request(self, request_id: str) -> None: self.removed_request_ids.append(request_id) client = FailingStageClient(stage_type="llm", final_output=False) output_processor = TrackingOutputProcessor() pool = StagePool( 0, [client], output_processor=output_processor, stage_vllm_config=SimpleNamespace(model_config=SimpleNamespace(max_model_len=64)), ) req_state = OrchestratorRequestState( request_id="req-0", sampling_params_list=[_sampling_params()], final_stage_id=0, ) with pytest.raises(RuntimeError, match="submit failed"): await pool.submit_initial("req-0", req_state, SimpleNamespace(request_id="req-0", prompt_token_ids=[1, 2])) assert output_processor.added_request_ids == ["req-0"] assert output_processor.removed_request_ids == ["req-0"] assert pool.get_bound_replica_id("req-0") is None @pytest.mark.asyncio async def test_stage_pool_abort_requests_logs_when_binding_is_missing(caplog) -> None: stage0 = FakeStageClient(stage_type="llm", final_output=False) pool = StagePool( 0, [stage0], output_processor=FakeOutputProcessor(), stage_vllm_config=SimpleNamespace(model_config=SimpleNamespace(max_model_len=64)), ) target_logger = logging.getLogger("vllm_omni.engine.stage_pool") target_logger.addHandler(caplog.handler) prev_level = target_logger.level target_logger.setLevel(logging.DEBUG) try: await pool.abort_requests(["missing-req"]) finally: target_logger.removeHandler(caplog.handler) target_logger.setLevel(prev_level) assert not stage0.abort_calls assert "abort: no live binding for req=missing-req in stage-0" in caplog.text @pytest.mark.asyncio async def test_collective_rpc_ignores_invalid_stage_ids(orchestrator_factory, caplog) -> None: stage0 = FakeCollectiveRpcStageClient(stage_type="llm", final_output=True, rpc_result={"stage": 0}) stage1 = FakeCollectiveRpcStageClient(stage_type="llm", final_output=True, rpc_result={"stage": 1}) stage_pools = _build_stage_pools( [[stage0], [stage1]], output_processors=[FakeOutputProcessor(), FakeOutputProcessor()], stage_vllm_configs=[ SimpleNamespace(model_config=SimpleNamespace(max_model_len=64)), SimpleNamespace(model_config=SimpleNamespace(max_model_len=64)), ], ) orchestrator_fixture = orchestrator_factory([], stage_pools=stage_pools) try: target_logger = logging.getLogger("vllm_omni.engine.orchestrator") target_logger.addHandler(caplog.handler) prev_level = target_logger.level target_logger.setLevel(logging.WARNING) try: orchestrator_fixture.request_sync_q.put_nowait( CollectiveRPCRequestMessage( rpc_id="rpc-1", method="list_loras", timeout=None, args=(), kwargs={}, stage_ids=[99, 1], ) ) msg = await _get_rpc_message(orchestrator_fixture) finally: target_logger.removeHandler(caplog.handler) target_logger.setLevel(prev_level) assert msg.type == "collective_rpc_result" assert msg.rpc_id == "rpc-1" assert msg.stage_ids == [1] assert msg.results == [{"stage": 1}] assert not stage0.collective_rpc_calls assert len(stage1.collective_rpc_calls) == 1 assert "collective_rpc: ignoring invalid stage_id 99" in caplog.text finally: await _shutdown_orchestrator(orchestrator_fixture) @pytest.mark.asyncio async def test_multi_replica_cfg_companion_inherits_parent_affinity(orchestrator_factory) -> None: """CFG companions should be routed to the same stage-0 replica as their parent.""" stage0_r0 = FakeStageClient(stage_type="llm", final_output=False) stage0_r1 = FakeStageClient(stage_type="llm", final_output=False) default_vllm_cfg = SimpleNamespace(model_config=SimpleNamespace(max_model_len=64)) stage_pools = _build_stage_pools( [[stage0_r0, stage0_r1]], output_processors=[FakeOutputProcessor()], stage_vllm_configs=[default_vllm_cfg], ) orchestrator_fixture = orchestrator_factory([], stage_pools=stage_pools) try: # Consume replica-0 first so the parent request binds to replica-1. await _enqueue_add_request( orchestrator_fixture, request_id="warmup", prompt=SimpleNamespace(request_id="warmup", prompt_token_ids=[0]), original_prompt={"prompt": "warmup"}, sampling_params_list=[_sampling_params()], final_stage_id=0, ) await _wait_for(lambda: len(stage0_r0.add_request_calls) == 1) await _enqueue_add_request( orchestrator_fixture, request_id="parent", prompt=SimpleNamespace(request_id="parent", prompt_token_ids=[1, 2]), original_prompt={"prompt": "parent"}, sampling_params_list=[_sampling_params()], final_stage_id=0, ) await _wait_for(lambda: len(stage0_r1.add_request_calls) == 1) orchestrator_fixture.request_sync_q.put_nowait( AddCompanionRequestMessage( companion_id="parent-neg", parent_id="parent", role="negative", prompt=SimpleNamespace(request_id="parent-neg", prompt_token_ids=[9]), companion_prompt_text={"prompt": "negative"}, sampling_params_list=[_sampling_params()], ) ) await _wait_for(lambda: len(stage0_r1.add_request_calls) == 2) assert stage_pools[0].get_bound_replica_id("parent") == 1 assert stage_pools[0].get_bound_replica_id("parent-neg") == 1 assert len(stage0_r0.add_request_calls) == 1 assert stage0_r1.add_request_calls[0][0].request_id == "parent" assert stage0_r1.add_request_calls[1][0].request_id == "parent-neg" finally: await _shutdown_orchestrator(orchestrator_fixture) def test_orchestrator_does_not_re_introduce_global_stats_throttle() -> None: """Regression: each (stage, replica) must independently publish its wrapped vllm:* stats when its scheduler emits non-None scheduler_stats. A previous version of Orchestrator carried a global self._last_stats_ts / _stats_interval_s gate around _stat_logger.record(). Because OmniSchedulerMixin.make_stats() already throttles at 1 Hz per scheduler (one per (stage, replica)), the extra global gate starved every replica other than the first to emit within each second — their {stage, replica} gauges/counters went stale. The fix removed the global gate entirely; the only signal needed is 'this replica's scheduler emitted non-None scheduler_stats'. This test fails loudly if someone reintroduces the global throttle. """ import inspect from vllm_omni.engine.orchestrator import Orchestrator source = inspect.getsource(Orchestrator) assert "_last_stats_ts" not in source, ( "Orchestrator must not gate stat recording on a global timestamp. " "OmniSchedulerMixin.make_stats() already throttles per scheduler " "(per (stage, replica)); an outer global gate starves all but the " "first replica to emit within each 1s window." ) assert "_stats_interval_s" not in source assert "raw_outputs.scheduler_stats is not None" in source, ( "Orchestrator must gate stat recording solely on " "raw_outputs.scheduler_stats being non-None — the per-scheduler 1Hz " "throttle in OmniSchedulerMixin.make_stats() is the only gate needed." )