diff --git a/tests/engine/duplex/test_session_manager.py b/tests/engine/duplex/test_session_manager.py index 0323f16d8d4..ea4d0559f9c 100644 --- a/tests/engine/duplex/test_session_manager.py +++ b/tests/engine/duplex/test_session_manager.py @@ -192,7 +192,13 @@ def create_session_state(self) -> DuplexModelSessionState: def capabilities(self, *, max_sessions: int) -> DuplexCapabilities: del max_sessions - return DuplexCapabilities(supports_input_append=True) + # Match MiniCPM-o 4.5: a resident Stage0 request (``...r.stage0``). + # Leaving ``supports_core_resumable_request`` at its dataclass default + # (False) would mint turn-scoped ids, which this harness is not. + return DuplexCapabilities( + supports_input_append=True, + supports_core_resumable_request=True, + ) def validate_client_extra_body(self, extra_body: object) -> None: pass @@ -404,7 +410,10 @@ async def test_open_answers_with_capabilities_and_emits_session_created() -> Non assert result.control_id == "open-sid-open" assert result.session_id == "sid-open" assert result.lease_generation == 0 - assert result.capabilities == DuplexCapabilities(supports_input_append=True) + assert result.capabilities == DuplexCapabilities( + supports_input_append=True, + supports_core_resumable_request=True, + ) assert result.public_session is not None assert result.public_session["id"] == "sid-open" assert result.public_session["voice"] == "test" diff --git a/tests/engine/duplex/test_session_runner_ephemeral_bind.py b/tests/engine/duplex/test_session_runner_ephemeral_bind.py new file mode 100644 index 00000000000..dbb74b3a6bb --- /dev/null +++ b/tests/engine/duplex/test_session_runner_ephemeral_bind.py @@ -0,0 +1,856 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni project + +"""Ephemeral (non-resumable) Stage0 binding for turn-commit models. + +A model whose Stage0 cannot resume appends on a finished request gets one +ordinary, turn-scoped request per committed turn (``...r.stage0_t{N}``) +instead of the single resident ``...r.stage0`` id. These tests pin the id +shape, the capability branching, the two-turn submit flow through the real +session runner with a recording fake stage port, and the observe hook that +lets an intermediate stage reach the client without stopping the pipeline. +Silence continuation is refused on the non-resumable path, a finished +observed intermediate stage does not close the stream, a listen-only +append keeps the turn-scoped Stage0 id bound, and a later commit re-binds +a stale submitted id without aborting leftover downstream bindings. +""" + +from __future__ import annotations + +import asyncio +import base64 +import binascii +import struct +from collections.abc import Mapping, Sequence +from dataclasses import dataclass, field +from types import SimpleNamespace +from typing import Any + +import pytest +from vllm.sampling_params import SamplingParams + +from vllm_omni.config.stage_config import DuplexSessionRuntimeConfig +from vllm_omni.engine.duplex import commands +from vllm_omni.engine.duplex.commands import DuplexCommand +from vllm_omni.engine.duplex.config import DuplexCapabilities, DuplexSessionConfig +from vllm_omni.engine.duplex.contracts import ( + DuplexAppendPlan, + DuplexFence, + DuplexOutputContext, + DuplexOutputDecision, + DuplexRequestIdentity, + DuplexStagePort, + DuplexStageRequestContext, + DuplexStageSubmission, + DuplexStageSubmissionResult, + duplex_ephemeral_stage_request_id, + duplex_resource_request_id, + is_stable_stage0_placeholder, +) +from vllm_omni.engine.duplex.events import DuplexEvent +from vllm_omni.engine.duplex.messages import ( + DuplexControlResultMessage, + DuplexSessionCommandMessage, + DuplexSessionEventMessage, + OpenDuplexSessionMessage, +) +from vllm_omni.engine.duplex.plugin import ( + DuplexDataPlane, + DuplexModelPlugin, + DuplexModelSessionState, + PcmAppendBuffer, + PcmAppendReservation, +) +from vllm_omni.engine.duplex.session.helpers import stage0_request_id +from vllm_omni.engine.duplex.session.manager import DuplexSessionManager +from vllm_omni.engine.duplex.session.runner import DuplexSessionRunner + +pytestmark = [pytest.mark.core_model, pytest.mark.cpu] + +SESSION_ID = "duplex-ephemeral-test" + + +# --------------------------------------------------------------------------- # +# L1: id shape and capability branching # +# --------------------------------------------------------------------------- # + + +def test_stage_request_id_ephemeral_when_not_resumable() -> None: + fence = DuplexFence(SESSION_ID, epoch=0, turn_id=3) + request_id = DuplexSessionManager.stage_request_id(fence, stage_id=0, resumable=False) + assert request_id == duplex_ephemeral_stage_request_id(fence, stage_id=0) + assert request_id.endswith(".r.stage0_t3") + assert not is_stable_stage0_placeholder(request_id, session_id=SESSION_ID, epoch=0) + + +def test_stage_request_id_resumable_keeps_stable_role() -> None: + fence = DuplexFence(SESSION_ID, epoch=1, turn_id=3) + request_id = DuplexSessionManager.stage_request_id(fence, stage_id=0, resumable=True) + assert request_id == duplex_resource_request_id(fence, "stage0") + assert is_stable_stage0_placeholder(request_id, session_id=SESSION_ID, epoch=1) + + +def test_is_stable_stage0_placeholder_rejects_ephemeral_ids() -> None: + fence = DuplexFence(SESSION_ID, epoch=0, turn_id=3) + assert is_stable_stage0_placeholder(duplex_resource_request_id(fence, "stage0"), session_id=SESSION_ID, epoch=0) + assert not is_stable_stage0_placeholder( + duplex_ephemeral_stage_request_id(fence, stage_id=0), + session_id=SESSION_ID, + epoch=0, + ) + + +def test_stage_submission_resumable_defaults_to_true() -> None: + context = DuplexStageRequestContext( + request_id="req", + session_id=SESSION_ID, + fence=DuplexFence(SESSION_ID), + stage_id=0, + final_stage_id=1, + config_generation=0, + sampling_params=(SamplingParams(max_tokens=8),), + ) + submission = DuplexStageSubmission(context=context, prompt={"prompt_token_ids": [1]}, already_submitted=False) + assert submission.resumable is True + + +def test_helpers_stage0_request_id_branches_on_capability() -> None: + ephemeral_session = SimpleNamespace( + session_id=SESSION_ID, + turn_id=2, + capabilities=DuplexCapabilities(supports_core_resumable_request=False), + ) + ephemeral_id = stage0_request_id(ephemeral_session, 0) + assert ephemeral_id.endswith(".r.stage0_t2") + assert not is_stable_stage0_placeholder(ephemeral_id, session_id=SESSION_ID, epoch=0) + resident_session = SimpleNamespace( + session_id=SESSION_ID, + turn_id=2, + capabilities=DuplexCapabilities(supports_core_resumable_request=True), + ) + resident_id = stage0_request_id(resident_session, 0) + assert is_stable_stage0_placeholder(resident_id, session_id=SESSION_ID, epoch=0) + + +def test_observe_stage_output_defaults_off() -> None: + plugin = EphemeralFakePlugin() + assert plugin.observe_stage_output(stage_id=0, output=object(), context=object()) is False + + +# --------------------------------------------------------------------------- # +# Fakes: an ephemeral turn-commit plugin # +# --------------------------------------------------------------------------- # + + +class _Reservation(PcmAppendReservation): + def __init__(self, *, operation_id: str, payload: dict[str, object] | None, byte_count: int) -> None: + self.operation_id = operation_id + self.payload = payload + self._byte_count = byte_count + self._active = True + + @property + def active(self) -> bool: + return self._active + + @property + def byte_count(self) -> int: + return self._byte_count + + def commit(self) -> None: + self._active = False + + def rollback(self) -> None: + self._active = False + + +class _CommitOnlyPcmBuffer(PcmAppendBuffer): + """Accumulates PCM and emits the whole utterance only on commit.""" + + def __init__(self) -> None: + self._buffer = bytearray() + + @property + def pending_byte_count(self) -> int: + return len(self._buffer) + + def clear(self) -> None: + self._buffer.clear() + + def clear_force_listen(self) -> None: + return + + def has_pending(self) -> bool: + return bool(self._buffer) + + def has_reserved(self) -> bool: + return False + + def prepare_append( + self, + payload: dict[str, object], + *, + operation_id: str, + chunk_period_ms: int, + allow_emit: bool, + ) -> _Reservation | None: + del operation_id, chunk_period_ms, allow_emit + audio = payload.get("audio") + if isinstance(audio, str): + try: + self._buffer.extend(base64.b64decode(audio, validate=True)) + except (binascii.Error, ValueError): + pass + return None # commit-only: never emit per chunk + + def prepare_commit(self, *, operation_id: str, chunk_period_ms: int) -> _Reservation: + del chunk_period_ms + raw = bytes(self._buffer) + self._buffer.clear() + payload: dict[str, object] | None = None + if raw: + payload = { + "type": "audio", + "audio": base64.b64encode(raw).decode("ascii"), + "format": "pcm_f32le", + "sample_rate_hz": 16000, + "final": True, + "is_speech": True, + } + return _Reservation(operation_id=operation_id, payload=payload, byte_count=len(raw)) + + def flush(self, *, chunk_period_ms: int) -> dict[str, object] | None: + reservation = self.prepare_commit(operation_id="flush", chunk_period_ms=chunk_period_ms) + reservation.commit() + return reservation.payload + + +@dataclass(slots=True) +class _FakeSessionState(DuplexModelSessionState): + audio_buffer: _CommitOnlyPcmBuffer = field(default_factory=_CommitOnlyPcmBuffer) + input_since_commit: bool = False + speech_since_commit: bool = False + context_locked: bool = False + committed_audio_payload: dict[str, object] | None = None + committed_audio_operation_id: str | None = None + committed_audio_reserved_bytes: int = 0 + deferred_response_create: bool = False + deferred_precreate_response: bool = False + continuation_owner_id: str | None = None + continuation_units: int = 0 + pending_silence_task: asyncio.Task[bool] | None = None + pending_silence_owner_id: str | None = None + + def retain_committed_audio( + self, + payload: dict[str, object], + *, + operation_id: str | None, + reserved_bytes: int = 0, + ) -> None: + self.committed_audio_payload = payload + self.committed_audio_operation_id = operation_id + self.committed_audio_reserved_bytes += max(0, int(reserved_bytes)) + + def clear_committed_audio(self) -> int: + reserved_bytes = self.committed_audio_reserved_bytes + self.committed_audio_payload = None + self.committed_audio_operation_id = None + self.committed_audio_reserved_bytes = 0 + self.deferred_response_create = False + self.deferred_precreate_response = False + return reserved_bytes + + def clear_continuation(self) -> None: + self.continuation_owner_id = None + self.continuation_units = 0 + self.pending_silence_task = None + self.pending_silence_owner_id = None + + +class _FakeDataPlane(DuplexDataPlane): + """Projects every delivered output as one audio event for its request.""" + + def __init__(self) -> None: + self._terminal: set[str] = set() + self.closed_streams: list[str] = [] + + def begin_request(self, request_id: str) -> None: + self._terminal.discard(request_id) + + def is_terminal(self, request_id: str | None) -> bool: + return request_id in self._terminal if request_id is not None else False + + def mark_terminal(self, request_id: str) -> None: + self._terminal.add(request_id) + + def close_stream(self, request_id: str) -> None: + self.closed_streams.append(request_id) + + def close_session(self, session_id: str, *, active_request_id: str | None = None) -> None: + self._terminal.clear() + + def project(self, result: object, *, context: object | None = None) -> list[dict[str, object]]: + del context + if not isinstance(result, dict): + return [] + outputs = result.get("data_plane_outputs") + if not isinstance(outputs, list): + return [] + events: list[dict[str, object]] = [] + for output in outputs: + request_id = getattr(output, "request_id", None) + if not isinstance(request_id, str) or not request_id: + continue + mm = getattr(output, "multimodal_output", None) + model_turn_id = mm.get("model_turn_id") if isinstance(mm, Mapping) else None + if model_turn_id is None: + model_turn_id = getattr(output, "duplex_turn_id", None) + # Stage ``finished`` is not the duplex turn ending; only an explicit + # ``end_of_turn`` (final-stage / model turn_eos) closes the response. + end_of_turn = bool(mm.get("end_of_turn")) if isinstance(mm, Mapping) else False + events.append( + { + "stage_role": "tts", + "is_listen": False, + "data_plane_request_id": request_id, + "audio": "wav-fake", + "sample_rate_hz": 24000, + "model_turn_id": model_turn_id, + "end_of_turn": end_of_turn, + } + ) + return events + + +class EphemeralFakePlugin(DuplexModelPlugin): + """Turn-commit-only fake: one ordinary Stage0 request per committed turn.""" + + plugin_id = "fake-ephemeral" + + def __init__(self, *, observe_stage0: bool = False) -> None: + super().__init__(lambda audio, sample_rate_hz, fmt, speed: None) + self.data_plane = _FakeDataPlane() + self._observe_stage0 = observe_stage0 + self.planned_payloads: list[dict[str, object]] = [] + + def configure_sampling_params( + self, + *, + runtime_config: dict[str, object], + defaults: tuple[object, ...], + ) -> tuple[object, ...]: + del runtime_config + return tuple(defaults) + + def plan_append( + self, + *, + request_id: str, + fence: DuplexFence, + session_config: dict[str, object], + runtime_config: dict[str, object], + seq: int, + turn_seq: int, + payload: object, + final: bool, + sampling_params: object, + ) -> DuplexAppendPlan: + del session_config, runtime_config, seq, turn_seq, sampling_params + if not isinstance(payload, Mapping) or not final: + raise AssertionError("ephemeral fake only plans final commit payloads") + self.planned_payloads.append({"request_id": request_id, "turn_id": fence.turn_id, "payload": payload}) + return DuplexAppendPlan( + prompt={ + "prompt_token_ids": [1], + "additional_information": {"session_id": fence.session_id, "turn_id": fence.turn_id}, + } + ) + + def decide_output( + self, + *, + stage_id: int, + final_stage_id: int, + segment_finished: bool, + segment_token_ids: tuple[int, ...], + segment_output_metadata: dict[str, object], + output: object, + ) -> DuplexOutputDecision | None: + return None + + def observe_stage_output(self, *, stage_id: int, output: object, context: object) -> bool: + del output, context + return self._observe_stage0 and stage_id == 0 + + def create_session_state(self) -> _FakeSessionState: + return _FakeSessionState() + + def capabilities(self, *, max_sessions: int) -> DuplexCapabilities: + return DuplexCapabilities( + supports_model_native_turn_policy=False, + supports_barge_in=False, + supports_input_append=True, + supports_turn_commit_only=True, + supports_core_resumable_request=False, + supports_realtime_endpoint=True, + supports_chat_completions=True, + adapter_patterns=["turn_commit"], + chunk_period_ms=1000, + ) + + def validate_client_extra_body(self, extra_body: object) -> None: + return + + async def prepare_runtime_config( + self, config: DuplexSessionConfig, *, model_config: object | None + ) -> dict[str, object]: + del config, model_config + return {} + + def runtime_config_for_update( + self, config: DuplexSessionConfig, current: Mapping[str, object] + ) -> dict[str, object]: + del config + return dict(current) + + def data_plane_context( + self, + *, + epoch: int, + turn_id: int, + active_response_turn_id: int | None, + active_response_id: str | None, + auto_responds: bool, + response_format: str, + speed: float | None, + modalities: tuple[str, ...], + ) -> object: + return SimpleNamespace( + epoch=epoch, + turn_id=turn_id, + active_response_turn_id=active_response_turn_id, + active_response_id=active_response_id, + auto_responds=auto_responds, + response_format=response_format, + speed=speed, + modalities=modalities, + ) + + +# --------------------------------------------------------------------------- # +# L2 harness: real session runner + recording fake stage port # +# --------------------------------------------------------------------------- # + + +class RecordingStagePort(DuplexStagePort): + """Records what the runner asks of the orchestrator; never talks to a stage.""" + + def __init__(self, *, stage_count: int = 2) -> None: + self._stage_count = stage_count + self.ensured: list[DuplexStageRequestContext] = [] + self.submissions: list[DuplexStageSubmission] = [] + self.cleanups: list[tuple[list[str], bool]] = [] + self.aborts: list[list[str]] = [] + + @property + def stage_count(self) -> int: + return self._stage_count + + def sampling_defaults(self) -> tuple[object, ...]: + return tuple(SamplingParams(max_tokens=8) for _ in range(self._stage_count)) + + def ensure_request(self, context: DuplexStageRequestContext) -> None: + self.ensured.append(context) + + async def submit(self, submission: DuplexStageSubmission) -> DuplexStageSubmissionResult: + self.submissions.append(submission) + return DuplexStageSubmissionResult( + request_id=submission.context.request_id, + stage_id=submission.context.stage_id, + replica_id=0, + ) + + async def cleanup(self, request_ids: list[str], *, abort: bool = False) -> None: + self.cleanups.append((list(request_ids), abort)) + + async def abort_requests(self, request_ids: list[str]) -> None: + self.aborts.append(list(request_ids)) + + +@dataclass +class Harness: + manager: DuplexSessionManager + port: RecordingStagePort + plugin: EphemeralFakePlugin + output: asyncio.Queue[Any] + results: asyncio.Queue[Any] + runner: DuplexSessionRunner + events: list[DuplexEvent] = field(default_factory=list) + + @property + def session(self): # noqa: ANN202 + return self.runner.session + + def submit(self, command: DuplexCommand) -> None: + self.manager.dispatch(DuplexSessionCommandMessage(session_id=SESSION_ID, command=command)) + + async def settle(self, *, idle_s: float = 0.05, timeout_s: float = 3.0) -> list[DuplexEvent]: + """Run the loop until the runner mailbox and append tasks are quiet; return new events.""" + loop = asyncio.get_running_loop() + deadline = loop.time() + timeout_s + quiet_since: float | None = None + collected: list[DuplexEvent] = [] + while True: + drained = False + while not self.output.empty(): + message = self.output.get_nowait() + if isinstance(message, DuplexSessionEventMessage): + collected.append(message.event) + drained = True + busy = ( + drained + or not self.runner._mailbox.empty() + or any(not task.done() for task in self.runner.tasks.append_tasks) + or any(not task.done() for task in self.runner._background_tasks) + ) + now = loop.time() + if busy: + quiet_since = None + elif quiet_since is None: + quiet_since = now + elif now - quiet_since >= idle_s: + break + if now >= deadline: + break + await asyncio.sleep(0.005) + self.events.extend(collected) + return collected + + async def run(self, command: DuplexCommand) -> list[DuplexEvent]: + self.submit(command) + return await self.settle() + + def deliver( + self, + output: object, + *, + stage_id: int, + epoch: int | None = None, + ) -> bool: + session = self.session + fence = session.fence if epoch is None else DuplexFence(SESSION_ID, epoch=epoch) + context = DuplexOutputContext( + identity=DuplexRequestIdentity(session_id=SESSION_ID, fence=fence), + final_stage_id=self.port.stage_count - 1, + segment_finished=bool(getattr(output, "finished", False)), + ) + return self.runner.on_stage_output( + stage_id, output, None, request_id=getattr(output, "request_id"), context=context + ) + + async def deliver_and_settle(self, output: object, *, stage_id: int) -> list[DuplexEvent]: + self.deliver(output, stage_id=stage_id) + return await self.settle() + + +async def open_harness(*, observe_stage0: bool = False, stage_count: int = 2) -> Harness: + plugin = EphemeralFakePlugin(observe_stage0=observe_stage0) + port = RecordingStagePort(stage_count=stage_count) + output: asyncio.Queue[Any] = asyncio.Queue() + results: asyncio.Queue[Any] = asyncio.Queue() + manager = DuplexSessionManager( + plugin=plugin, + stage_port=port, + output_sink=output, + result_sink=results, + runtime_config=DuplexSessionRuntimeConfig(), + model_config=None, + ) + config = DuplexSessionConfig( + model="fake/ephemeral-turn-commit", + modalities=["text", "audio"], + instructions="You are a concise assistant.", + extra_body={}, + ) + await manager.handle(OpenDuplexSessionMessage(control_id="c-open", session_id=SESSION_ID, session_config=config)) + result = await asyncio.wait_for(results.get(), timeout=2.0) + assert isinstance(result, DuplexControlResultMessage) and result.ok, result + harness = Harness( + manager=manager, port=port, plugin=plugin, output=output, results=results, runner=manager.runners[SESSION_ID] + ) + await harness.settle() + return harness + + +async def close_harness(harness: Harness) -> None: + await harness.manager.shutdown() + + +def pcm_f32(samples: int, *, value: float = 0.05) -> bytes: + return struct.pack(f"<{samples}f", *([value] * samples)) + + +def append_audio(samples: int = 16000) -> commands.AppendAudio: + return commands.AppendAudio( + audio=pcm_f32(samples), + format="pcm_f32le", + sample_rate_hz=16000, + is_speech=True, + ) + + +def fake_output( + request_id: str, + *, + finished: bool, + turn_id: int | None = 0, + end_of_turn: bool = False, +) -> SimpleNamespace: + """A stage output the way the orchestrator hands it to the runner.""" + # ``model_turn_id`` must sit in multimodal_output to survive + # OmniRequestOutput.from_stage_output; a raw duplex_turn_id attribute does not. + multimodal_output: dict[str, object] = {"end_of_turn": end_of_turn} + if turn_id is not None: + multimodal_output["model_turn_id"] = turn_id + return SimpleNamespace( + request_id=request_id, + finished=finished, + duplex_turn_id=turn_id, + outputs=[SimpleNamespace(text="", token_ids=[], multimodal_output=multimodal_output)], + multimodal_output=multimodal_output, + ) + + +def types(events: Sequence[DuplexEvent]) -> list[str]: + return [event.type for event in events] + + +# --------------------------------------------------------------------------- # +# L2: two committed turns through the real runner # +# --------------------------------------------------------------------------- # + + +@pytest.mark.asyncio +async def test_ephemeral_two_committed_turns_get_fresh_stage0_ids() -> None: + h = await open_harness() + try: + # Admission reserves the turn-scoped Stage0 id, not the resident one. + assert [context.request_id for context in h.port.ensured] == [ + duplex_ephemeral_stage_request_id(DuplexFence(SESSION_ID, epoch=0, turn_id=0), stage_id=0) + ] + + # Turn 0: append + commit => one ephemeral ordinary request. + await h.run(append_audio()) + assert not h.port.submissions # commit-only buffer: nothing per chunk + events = await h.run(commands.Commit(create_response=True)) + assert "input_audio_buffer.committed" in types(events) + assert "response.created" in types(events) + assert len(h.port.submissions) == 1 + first = h.port.submissions[0] + assert first.context.request_id.endswith(".r.stage0_t0") + assert first.already_submitted is False + assert first.resumable is False + + # The model answers; the data plane reports model_turn_id so the turn completes. + events = await h.deliver_and_settle( + fake_output(first.context.request_id, finished=True, turn_id=0, end_of_turn=True), + stage_id=1, + ) + assert "response.done" in types(events) + assert h.session.turn_id == 1 + + # Turn 1: a fresh ephemeral id, again a first-time submit. + await h.run(append_audio()) + await h.run(commands.Commit(create_response=True)) + assert len(h.port.submissions) == 2 + second = h.port.submissions[1] + assert second.context.request_id.endswith(".r.stage0_t1") + assert second.context.request_id != first.context.request_id + assert second.already_submitted is False + assert second.resumable is False + # Happy path: no stale ephemeral had to be aborted. + assert h.port.cleanups == [] + finally: + await close_harness(h) + + +@pytest.mark.asyncio +async def test_ephemeral_rebind_after_turn_without_model_turn_id() -> None: + """A turn that never completed (no model_turn_id) re-binds on the next commit.""" + h = await open_harness() + try: + await h.run(append_audio()) + await h.run(commands.Commit(create_response=True)) + assert len(h.port.submissions) == 1 + first_id = h.port.submissions[0].context.request_id + assert first_id.endswith(".r.stage0_t0") + + # The response ends but the model turn is never completed (no id). + events = await h.deliver_and_settle( + fake_output(first_id, finished=True, turn_id=None, end_of_turn=True), + stage_id=1, + ) + assert "response.done" in types(events) + assert h.session.turn_id == 0 + + leftover_stage1 = duplex_resource_request_id(h.session.fence, "stage1") + h.session.bind_stage_request(1, leftover_stage1, fence=h.session.fence) + + # The next commit cannot reuse the finished t0 id: the turn is + # completed mechanically, a fresh t1 id is minted, and only the stale + # Stage0 request is aborted. A leftover downstream binding stays. + await h.run(append_audio()) + await h.run(commands.Commit(create_response=True)) + assert len(h.port.submissions) == 2 + second = h.port.submissions[1] + assert second.context.request_id.endswith(".r.stage0_t1") + assert second.already_submitted is False + assert h.session.turn_id == 1 + assert h.port.cleanups == [([first_id], True)] + assert leftover_stage1 not in {rid for ids, _abort in h.port.cleanups for rid in ids} + assert (1, leftover_stage1) in h.session.request_resources + finally: + await close_harness(h) + + +@pytest.mark.asyncio +async def test_listen_only_submit_rebinds_on_the_next_commit() -> None: + """A real listen-only Stage0 submit stays bound; the next commit re-binds it.""" + h = await open_harness() + try: + payload = { + "type": "audio", + "audio": base64.b64encode(pcm_f32(16000)).decode("ascii"), + "format": "pcm_f32le", + "sample_rate_hz": 16000, + "final": True, + "is_speech": True, + } + task = await h.runner._start_append(payload, final=True, precreate_response=False) + assert await task is True + await h.settle() + + assert len(h.port.submissions) == 1 + first = h.port.submissions[0] + first_id = first.context.request_id + assert first_id.endswith(".r.stage0_t0") + assert first.already_submitted is False + assert first.resumable is False + assert h.session.active_request_id == first_id + assert not is_stable_stage0_placeholder(first_id, session_id=SESSION_ID, epoch=0) + assert h.session.turn_id == 0 + assert h.session.active_response_id is None + assert "response.created" not in types(h.events) + + leftover_stage1 = duplex_resource_request_id(h.session.fence, "stage1") + h.session.bind_stage_request(1, leftover_stage1, fence=h.session.fence) + + await h.run(append_audio()) + events = await h.run(commands.Commit(create_response=True)) + assert "input_audio_buffer.committed" in types(events) + assert "response.created" in types(events) + assert len(h.port.submissions) == 2 + second = h.port.submissions[1] + assert second.context.request_id.endswith(".r.stage0_t1") + assert second.already_submitted is False + assert h.session.turn_id == 1 + assert h.port.cleanups == [([first_id], True)] + assert leftover_stage1 not in {rid for ids, _abort in h.port.cleanups for rid in ids} + assert (1, leftover_stage1) in h.session.request_resources + finally: + await close_harness(h) + + +@pytest.mark.asyncio +async def test_observe_stage0_projects_and_still_forwards() -> None: + h = await open_harness(observe_stage0=True) + try: + await h.run(append_audio()) + await h.run(commands.Commit(create_response=True)) + request_id = h.port.submissions[0].context.request_id + + # An intermediate-stage output is observed: projected to the client + # (an event is emitted) yet NOT consumed (on_stage_output returns + # False so the orchestrator still forwards it to the next stage). + forwarded = h.deliver(fake_output(request_id, finished=False, turn_id=0), stage_id=0) + assert forwarded is False + events = await h.settle() + assert "response.output_audio.delta" in types(events) + + # With the plugin's observe off, the same delivery emits nothing. + h.plugin._observe_stage0 = False + forwarded = h.deliver(fake_output(request_id, finished=False, turn_id=0), stage_id=0) + assert forwarded is False + events = await h.settle() + assert "response.output_audio.delta" not in types(events) + finally: + await close_harness(h) + + +@pytest.mark.asyncio +async def test_observe_finished_intermediate_does_not_close_the_stream() -> None: + """A finished observed stage is that stage ending, not the duplex turn.""" + h = await open_harness(observe_stage0=True) + try: + await h.run(append_audio()) + await h.run(commands.Commit(create_response=True)) + request_id = h.port.submissions[0].context.request_id + + forwarded = h.deliver(fake_output(request_id, finished=True, turn_id=0), stage_id=0) + events = await h.settle() + assert forwarded is False + assert "response.output_audio.delta" in types(events) + assert "response.done" not in types(events) + assert h.plugin.data_plane.closed_streams == [] + assert len(h.port.submissions) == 1 + assert h.runner.model_state.continuation_units == 0 + assert h.session.turn_id == 0 + finally: + await close_harness(h) + + +@pytest.mark.asyncio +async def test_finished_final_stage_does_not_schedule_silence_continuation() -> None: + """Non-resumable stage0 cannot submit_update a silence unit after TTS ends.""" + h = await open_harness() + try: + await h.run(append_audio()) + await h.run(commands.Commit(create_response=True)) + request_id = h.port.submissions[0].context.request_id + + events = await h.deliver_and_settle( + fake_output(request_id, finished=True, turn_id=0, end_of_turn=True), + stage_id=1, + ) + assert "response.done" in types(events) + assert request_id in h.plugin.data_plane.closed_streams + assert len(h.port.submissions) == 1 + assert h.runner.model_state.continuation_units == 0 + assert h.runner.model_state.continuation_owner_id is None + finally: + await close_harness(h) + + +@pytest.mark.asyncio +async def test_listen_only_append_keeps_ephemeral_stage0_bound( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Turn-scoped ids must survive a listen-only append; only ``...r.stage0`` is cleared.""" + h = await open_harness() + try: + + async def listen_only(*args: object, **kwargs: object) -> tuple[bool, bool]: + del args, kwargs + return True, False + + monkeypatch.setattr(h.runner.model, "append_runtime_input", listen_only) + await h.run(append_audio()) + await h.run(commands.Commit(create_response=True)) + bound = h.session.active_request_id + assert isinstance(bound, str) + assert bound.endswith(".r.stage0_t0") + assert not is_stable_stage0_placeholder(bound, session_id=SESSION_ID, epoch=0) + assert h.port.submissions == [] + finally: + await close_harness(h) diff --git a/tests/engine/test_duplex_orchestrator.py b/tests/engine/test_duplex_orchestrator.py index 5b8130c36a1..6153281f55c 100644 --- a/tests/engine/test_duplex_orchestrator.py +++ b/tests/engine/test_duplex_orchestrator.py @@ -23,7 +23,12 @@ from vllm_omni.config.stage_config import DuplexSessionRuntimeConfig from vllm_omni.engine.duplex import commands from vllm_omni.engine.duplex.config import DuplexSessionConfig, DuplexSessionState -from vllm_omni.engine.duplex.contracts import DuplexFence, duplex_resource_request_id +from vllm_omni.engine.duplex.contracts import ( + DuplexFence, + DuplexStageRequestContext, + DuplexStageSubmission, + duplex_resource_request_id, +) from vllm_omni.engine.duplex.messages import ( CloseDuplexSessionMessage, DuplexControlResultMessage, @@ -292,6 +297,53 @@ async def test_append_submits_the_resumable_stage0_request_and_counts_it_running assert counter.value == 0 +@pytest.mark.asyncio +async def test_submit_threads_resumable_false_into_the_engine_core_request() -> None: + """Ephemeral turn-commit submissions must not resume a finished ordinary request.""" + counter = FakeRunningCounter() + orchestrator, clients, rpc_q, _ = _build(running_counter=counter) + await _open(orchestrator, rpc_q) + request_id = _stage0_request_id() + request_state = orchestrator.request_states[request_id] + context = DuplexStageRequestContext( + request_id=request_id, + session_id=SESSION_ID, + fence=DuplexFence(SESSION_ID), + stage_id=0, + final_stage_id=0, + config_generation=request_state.config_generation, + sampling_params=tuple(request_state.sampling_params_list), + ) + result = await orchestrator.submit( + DuplexStageSubmission( + context=context, + prompt={"prompt_token_ids": [1]}, + already_submitted=False, + resumable=False, + ) + ) + assert result.request_id == request_id + assert len(clients[0].add_request_calls) == 1 + submitted = clients[0].add_request_calls[0][0] + assert submitted.request_id == request_id + assert submitted.resumable is False + assert counter.value == 1 + + await orchestrator.submit( + DuplexStageSubmission( + context=context, + prompt={"prompt_token_ids": [1, 2]}, + already_submitted=True, + resumable=False, + ) + ) + assert len(clients[0].add_request_calls) == 2 + assert clients[0].add_request_calls[1][0].resumable is False + assert counter.value == 1 + + await orchestrator.session_manager.shutdown() + + @pytest.mark.asyncio async def test_session_update_refreshes_the_next_append_sampling_params() -> None: orchestrator, clients, rpc_q, _ = _build() diff --git a/vllm_omni/engine/duplex/contracts.py b/vllm_omni/engine/duplex/contracts.py index ed14f5fb818..5cc4cd0fb60 100644 --- a/vllm_omni/engine/duplex/contracts.py +++ b/vllm_omni/engine/duplex/contracts.py @@ -78,6 +78,12 @@ class DuplexStageSubmission: context: DuplexStageRequestContext prompt: Mapping[str, object] already_submitted: bool + #: True: resume/update an existing stage0 id. False: open a new ephemeral + #: (turn-scoped) request. Distinct from + #: ``DuplexCapabilities.supports_core_resumable_request``, which selects + #: the request-id shape; this flag tells the stage port which submit + #: semantics the id carries. + resumable: bool = True def __post_init__(self) -> None: object.__setattr__(self, "prompt", MappingProxyType(dict(self.prompt))) @@ -157,6 +163,25 @@ def duplex_resource_request_id(fence: DuplexFence, role: str) -> str: return f"duplex-s.{encoded_session_id}.e.{fence.epoch}.r.{role}" +def duplex_ephemeral_stage_request_id(fence: DuplexFence, *, stage_id: int) -> str: + """Turn-scoped stage request id for non-resumable (ephemeral) duplex models. + + A model whose Stage0 cannot resume appends on a finished request gets one + ordinary request per committed turn instead of the single resident + ``...r.stage0`` id. + """ + return duplex_resource_request_id(fence, f"stage{stage_id}_t{fence.turn_id}") + + +def is_stable_stage0_placeholder(request_id: str, *, session_id: str, epoch: int) -> bool: + """Return whether ``request_id`` is the resident ``...r.stage0`` placeholder. + + Turn-scoped ephemeral ids (``...r.stage0_t{N}``) use a different role and + do not match. + """ + return request_id == duplex_resource_request_id(DuplexFence(session_id, epoch=epoch), "stage0") + + def duplex_resource_request_belongs_to_session(request_id: str, session_id: str) -> bool: """Return whether a current-format resource request belongs to a session.""" parts = request_id.split(".") @@ -187,6 +212,8 @@ def duplex_resource_request_belongs_to_session(request_id: str, session_id: str) "DuplexStageSubmission", "DuplexStageSubmissionResult", "duplex_data_plane_request_info", + "duplex_ephemeral_stage_request_id", "duplex_resource_request_belongs_to_session", "duplex_resource_request_id", + "is_stable_stage0_placeholder", ] diff --git a/vllm_omni/engine/duplex/plugin.py b/vllm_omni/engine/duplex/plugin.py index 130d5204d6a..5d95981b311 100644 --- a/vllm_omni/engine/duplex/plugin.py +++ b/vllm_omni/engine/duplex/plugin.py @@ -221,6 +221,23 @@ def decide_output( output: object, ) -> DuplexOutputDecision | None: ... + def observe_stage_output( + self, + *, + stage_id: int, + output: object, + context: object, + ) -> bool: + """Return True to project this intermediate stage to the client. + + Unlike ``decide_output``, observing does **not** short-circuit the + pipeline: the stage output is still forwarded to the next stage. A + multi-stage turn-commit model uses this to surface e.g. its text + stage's tokens while the audio stages keep running. Default is off. + """ + del stage_id, output, context + return False + # ---- session policy (was ServingRuntimeAdapter) ---- @abstractmethod diff --git a/vllm_omni/engine/duplex/session/append_task.py b/vllm_omni/engine/duplex/session/append_task.py index 724dc7f8224..be6a958e953 100644 --- a/vllm_omni/engine/duplex/session/append_task.py +++ b/vllm_omni/engine/duplex/session/append_task.py @@ -24,8 +24,8 @@ from vllm.logger import init_logger from vllm_omni.engine.duplex.config import DuplexSessionState, DuplexTurnEventType +from vllm_omni.engine.duplex.contracts import is_stable_stage0_placeholder from vllm_omni.engine.duplex.plugin import PcmAppendReservation -from vllm_omni.engine.duplex.session import helpers from vllm_omni.engine.duplex.session.context import DuplexSessionContext from vllm_omni.engine.duplex.session.emitter import SessionEmitter from vllm_omni.engine.duplex.session.model_channel import ModelChannel @@ -155,8 +155,14 @@ async def _submit(self) -> bool: self.ctx.run.runtime_closed = True return False if not emitted_response and session.epoch == self.epoch: - if session.active_request_id == helpers.stage0_request_id(session, self.epoch): - session.clear_request(self.request_id) + # Only clear the resident ...r.stage0 placeholder. Turn-scoped + # ephemeral ids (...r.stage0_tN) stay bound after a listen-only + # append; the next commit re-binds them. + active = session.active_request_id + if isinstance(active, str) and is_stable_stage0_placeholder( + active, session_id=session.session_id, epoch=self.epoch + ): + session.clear_request(active) if self.final: self.out.emit_events([session.signal_turn(DuplexTurnEventType.USER_STARTED.value)]) return append_ok diff --git a/vllm_omni/engine/duplex/session/helpers.py b/vllm_omni/engine/duplex/session/helpers.py index f8b705e6d2e..9dd8fdf9769 100644 --- a/vllm_omni/engine/duplex/session/helpers.py +++ b/vllm_omni/engine/duplex/session/helpers.py @@ -19,7 +19,11 @@ import pybase64 as base64 from vllm_omni.engine.duplex.config import DuplexPlaybackCommitPolicy -from vllm_omni.engine.duplex.contracts import DuplexFence, duplex_resource_request_id +from vllm_omni.engine.duplex.contracts import ( + DuplexFence, + duplex_ephemeral_stage_request_id, + duplex_resource_request_id, +) from vllm_omni.engine.duplex.events import ErrorEvent, OverlapDecision, error_event if TYPE_CHECKING: @@ -35,7 +39,11 @@ def stage0_request_id(session: DuplexEngineSession, epoch: int) -> str: - return duplex_resource_request_id(DuplexFence(session.session_id, epoch=epoch), "stage0") + fence = DuplexFence(session.session_id, epoch=epoch, turn_id=session.turn_id) + if session.capabilities.supports_core_resumable_request: + return duplex_resource_request_id(fence, "stage0") + # A non-resumable Stage0 gets one ordinary request per committed turn. + return duplex_ephemeral_stage_request_id(fence, stage_id=0) def response_in_progress(session: DuplexEngineSession, tasks: DuplexSessionTasks) -> bool: diff --git a/vllm_omni/engine/duplex/session/manager.py b/vllm_omni/engine/duplex/session/manager.py index 0b7a7a07ea0..ae4965d92a2 100644 --- a/vllm_omni/engine/duplex/session/manager.py +++ b/vllm_omni/engine/duplex/session/manager.py @@ -31,6 +31,7 @@ DuplexFence, DuplexStagePort, DuplexStageRequestContext, + duplex_ephemeral_stage_request_id, duplex_resource_request_belongs_to_session, duplex_resource_request_id, ) @@ -408,8 +409,10 @@ def sampling_params_for(self, session: DuplexEngineSession) -> tuple[object, ... return configured @staticmethod - def stage_request_id(fence: DuplexFence, *, stage_id: int) -> str: - return duplex_resource_request_id(fence, f"stage{stage_id}") + def stage_request_id(fence: DuplexFence, *, stage_id: int, resumable: bool = True) -> str: + if resumable: + return duplex_resource_request_id(fence, f"stage{stage_id}") + return duplex_ephemeral_stage_request_id(fence, stage_id=stage_id) def ensure_stage_request( self, @@ -422,7 +425,8 @@ def ensure_stage_request( if stage_id >= self.stage_port.stage_count: return None effective_fence = fence or session.fence - request_id = self.stage_request_id(effective_fence, stage_id=stage_id) + resumable = bool(session.capabilities.supports_core_resumable_request) + request_id = self.stage_request_id(effective_fence, stage_id=stage_id, resumable=resumable) session.reserve_stage_request(stage_id, request_id, fence=effective_fence) context = DuplexStageRequestContext( request_id=request_id, diff --git a/vllm_omni/engine/duplex/session/model_channel.py b/vllm_omni/engine/duplex/session/model_channel.py index 322db7edeaf..825cf1ff7d4 100644 --- a/vllm_omni/engine/duplex/session/model_channel.py +++ b/vllm_omni/engine/duplex/session/model_channel.py @@ -208,15 +208,39 @@ async def _append_via_data_plane( lease_operation_id = f"append:{operation_id or uuid.uuid4().hex}" operation_started = False stage_id = 0 - request_id = self._ctx.manager.stage_request_id(fence, stage_id=stage_id) + resumable = bool(session.capabilities.supports_core_resumable_request) + request_id = self._ctx.manager.stage_request_id(fence, stage_id=stage_id, resumable=resumable) + if not resumable and session.stage_request_submitted(stage_id, request_id): + # Ephemeral turn-commit cannot submit_update on a finished stage0 + # id: this turn never completed (e.g. a listen-only turn), so its + # id is still bound. Complete the turn to advance turn_id, mint a + # fresh ephemeral id, and abort only that stale Stage0 request. + # Downstream stage bindings are left alone — they are created by + # orchestrator forward, not by this Stage0 append path. + stale_ephemeral_id = request_id + session.complete_model_turn(fence.turn_id) + fence = DuplexFence(session.session_id, epoch=session.epoch, turn_id=session.turn_id) + request_id = self._ctx.manager.stage_request_id(fence, stage_id=stage_id, resumable=False) + session.request_resources.pop((stage_id, stale_ephemeral_id), None) + try: + await self._ctx.stage_port.cleanup([stale_ephemeral_id], abort=True) + except Exception: + logger.warning( + "duplex abort of stale ephemeral request failed session=%s id=%s", + session.session_id, + stale_ephemeral_id, + exc_info=True, + ) try: session.begin_lease_operation(fence, lease_operation_id) operation_started = True reservation = session.prepare_append(fence) - already_submitted = session.stage_request_submitted(stage_id, request_id) + already_submitted = False if not resumable else session.stage_request_submitted(stage_id, request_id) request_context = self._ctx.manager.ensure_stage_request(session, stage_id=stage_id, fence=fence) if request_context is None: raise RuntimeError("duplex_data_plane_has_no_stage") + if request_context.request_id != request_id: + request_id = request_context.request_id append_plan = self._ctx.plugin.plan_append( request_id=request_id, fence=fence, @@ -234,6 +258,7 @@ async def _append_via_data_plane( context=request_context, prompt=append_plan.prompt, already_submitted=already_submitted, + resumable=resumable, ) submission_result = await self._ctx.stage_port.submit(submission) try: @@ -254,6 +279,7 @@ async def _append_via_data_plane( ) raise session.touch_lease(DuplexLeaseActivity.APPEND) + session.bind_request(request_id) return { "ok": True, "operation": "append", @@ -270,7 +296,7 @@ async def _append_via_data_plane( "seq": update.seq, "turn_id": update.turn_id, "turn_seq": update.turn_seq, - "resumable": True, + "resumable": resumable, }, } ], @@ -367,6 +393,16 @@ def decide_output( raise TypeError("duplex plugin decide_output() must return DuplexOutputDecision or None") return decision + def observe_stage_output(self, stage_id: int, output: RequestOutput, context: DuplexOutputContext) -> bool: + """Project an intermediate stage to the client without short-circuiting the pipeline.""" + return bool( + self._ctx.plugin.observe_stage_output( + stage_id=stage_id, + output=output, + context=context, + ) + ) + @staticmethod def stage_metrics_snapshot(stage_id: int, metrics: object, output: object) -> dict[str, dict[str, object]] | None: if not isinstance(metrics, StageRequestStats): @@ -452,7 +488,12 @@ async def on_stage_output_item(self, item: StageOutput) -> None: await self._close_from_runtime(close_reason) return finished = self._data_plane_outputs_finished(drain_result) - if finished and emitted_response and not self._out.auto_responds(): + # An observed intermediate stage finishing means that stage is done, + # not the duplex turn: only the final stage — or a stage whose direct + # decision short-circuited the pipeline — closes the stream (and + # offers the model another silence unit). + response_completes_here = item.stage_id >= item.context.final_stage_id or item.decision is not None + if finished and emitted_response and not self._out.auto_responds() and response_completes_here: # A finished, emitted response releases the per-request projector # cursor on its way out and offers the model another # silence unit. @@ -1046,6 +1087,11 @@ async def maybe_continue_response( if session.state == DuplexSessionState.CLOSED or self._ctx.run.closing: model_state.clear_continuation() return + if not session.capabilities.supports_core_resumable_request: + # Non-resumable stage0 cannot submit_update after the request + # finishes; a turn-commit model has no silence continuation. + model_state.clear_continuation() + return request_id = session.active_request_id if request_id is None: model_state.clear_continuation() diff --git a/vllm_omni/engine/duplex/session/runner.py b/vllm_omni/engine/duplex/session/runner.py index 7915500b739..6782b592151 100644 --- a/vllm_omni/engine/duplex/session/runner.py +++ b/vllm_omni/engine/duplex/session/runner.py @@ -215,10 +215,15 @@ def on_stage_output( ) -> bool: """Accept one stage output (orchestrator loop); return True when it must not be forwarded.""" decision: DuplexOutputDecision | None = None + observe = False if stage_id < context.final_stage_id: decision = self.model.decide_output(stage_id, output, context) + # Optional mid-pipeline observe: project to the client without + # short-circuiting the pipeline (still forwarded downstream). + if decision is None: + observe = self.model.observe_stage_output(stage_id, output, context) consume = decision is not None or stage_id >= context.final_stage_id - if not consume: + if not consume and not observe: # Stage0 text without a direct decision feeds the TTS stage as before. # Its metrics still have to reach the client: before sessions moved # into the engine the orchestrator published them as a standalone @@ -241,7 +246,8 @@ def on_stage_output( decision=decision, ) ) - return True + # An observe-only projection must still forward to the next stage. + return consume def on_stage_failure(self, stage_id: int, exc: BaseException) -> None: """A stage rejected this session's request: fail the active response now. diff --git a/vllm_omni/engine/duplex_orchestrator.py b/vllm_omni/engine/duplex_orchestrator.py index 09a767fb135..c2ede5f320e 100644 --- a/vllm_omni/engine/duplex_orchestrator.py +++ b/vllm_omni/engine/duplex_orchestrator.py @@ -311,7 +311,7 @@ async def submit(self, submission: DuplexStageSubmission) -> DuplexStageSubmissi prompt=dict(submission.prompt), params=context.stage_sampling_params, model_config=self.stage_pools[context.stage_id].stage_vllm_config.model_config, - resumable=True, + resumable=submission.resumable, ) request.external_req_id = request.request_id pool = self.stage_pools[context.stage_id]