diff --git a/docs/design/metrics.md b/docs/design/metrics.md index 85eb7ac3876..874cbd310c4 100644 --- a/docs/design/metrics.md +++ b/docs/design/metrics.md @@ -60,7 +60,7 @@ Upstream per-engine metrics retain the `vllm:` prefix but are now registered by There are five independent paths for metric collection. -**Path 1: Pipeline-level metrics (`vllm_omni:*`)** +#### Path 1: Pipeline-level metrics (`vllm_omni:*`) `OmniPrometheusMetrics` registers the Gauge / Counter / Histogram collectors at import time. It is instantiated once per entrypoint, labeled with the model name. The entrypoint calls its methods as requests progress: @@ -69,7 +69,7 @@ There are five independent paths for metric collection. - `request_failed()` — recorded by the cleanup path when a request exits without natural completion. Internally maps to `finished_reason="abort"` on the existing completion Counter. - `inc_requests_failed(reason)` — increments the dedicated failure-attribution Counter once per request. Reasons are normalized to the bounded set `client_abort`, `client_disconnect`, `stage_error`, and `unknown`. -**Path 2: Image and diffusion metrics** +#### Path 2: Image and diffusion metrics Image and diffusion metrics use two cooperating collectors. `OmniPrometheusMetrics` owns pipeline/stage workload families, while `OmniModalityMetrics` owns per-replica diffusion breakdowns. Finished stage messages are consumed once in `OmniBase._process_single_result`, which prevents replayed terminal messages from incrementing Counters or Histograms twice. @@ -81,15 +81,15 @@ Image and diffusion metrics use two cooperating collectors. `OmniPrometheusMetri Zero and missing values have different meanings. Valid zero-duration queue observations are recorded, while unavailable optional measurements are omitted instead of being exported as synthetic zeroes. `peak_memory_mb`, `diffusion_forward_s`, `diffusion_kv_load_s`, `vae_decode_s`, and `kv_wait_s` therefore appear only when their data source is active. -**Path 3: Audio modality metrics (`vllm_omni:audio_*`)** +#### Path 3: Audio modality metrics (`vllm_omni:audio_*`) `OmniModalityMetrics` registers seven audio families with `{model_name, stage, replica}` (plus an extra `threshold_ms` / `reason` label on the two extra-cardinality Counters). Three observation sites: - `observe_modality_at_finalize(...)` — called from `omni_base._process_single_result` inside the existing `e2e_done` finalize guard. For `output_type == "audio"` it emits `audio_frames_total`, `audio_duration_s`, `audio_rtf` (or `audio_skipped_requests_total{reason="no_audio_data"}` when no audio was produced). Sample rate is resolved from `engine_outputs.multimodal_output` via `definitions.resolve_audio_sample_rate(...)` (fallback chain mirrors `serving_chat.py`'s audio response path). -- `observe_audio_first_packet(...)` — called from the OpenAI SSE audio branch in `serving_chat.py` on the first audio packet for a request. The once-per-request guard is held by `ClientRequestState.first_audio_ts`. The `request_arrival_ts` anchor is stored in `ClientRequestState` by `async_omni.generate()`, computed at request entry. -- `observe_audio_streaming_finalize(...)` — called from `serving_chat.py` after the streaming chunk loop exhausts. It runs the per-chunk player simulation from `vllm_omni/benchmarks/audio_continuity.py` to compute the worst-case underrun and emits `audio_underrun_s` plus (when the request stayed below the threshold) `audio_continuity_ok_total{threshold_ms}`. Per-chunk PCM byte counts and arrival timestamps are recorded by the same audio branch that updates `first_audio_ts`. +- `observe_audio_first_packet(...)` — called from the OpenAI SSE audio branch in `serving_chat.py` and the streaming TTS output paths in `serving_speech.py` on the first non-empty audio payload for a request. Streaming Speech observes the first PCM payload rather than a header-only WAV chunk. HTTP carries the middleware request timestamp through response generation; WebSocket uses the start of each sentence synthesis request. Non-streaming Speech does not emit TTFP and uses `e2e_request_latency_s` for overall latency instead. +- `observe_audio_streaming_finalize(...)` — called from `serving_chat.py` and `serving_speech.py` after a streaming chunk loop exhausts normally. It runs the per-chunk player simulation from `vllm_omni/benchmarks/audio_continuity.py` to compute the worst-case underrun and emits `audio_underrun_s` plus (when the request stayed below the threshold) `audio_continuity_ok_total{threshold_ms}`. Per-chunk PCM byte counts and arrival timestamps are recorded by the same output paths that update TTFP. Speech records raw PCM payloads before raw/SSE/WebSocket encoding, excludes WAV headers and empty chunks, and does not emit continuity samples for failed, cancelled, or non-streaming requests. -**Path 4: Cross-stage transfer metrics (`vllm_omni:transfer_*`)** +#### Path 4: Cross-stage transfer metrics (`vllm_omni:transfer_*`) `OmniTransferMetrics` registers four Histogram families with `{model_name, from_stage, from_replica, to_stage, to_replica}` labels. Each observation corresponds to one physical transfer hop (one chunk between adjacent stages), not the per-request accumulated total — so the histograms track per-transfer distribution. @@ -100,7 +100,7 @@ The hook lives in `OrchestratorAggregator.record_transfer_tx` and `record_transf Defensive fail-safe: if `transfer_emitter` or `replica_resolver` is missing, or the resolver returns `None` for either side, the emit is skipped silently (the underlying `TransferEdgeStats` accumulation is unaffected). -**Path 5: Per-engine metrics (`vllm:*`, stage/replica wrap)** +#### Path 5: Per-engine metrics (`vllm:*`, stage/replica wrap) The Orchestrator instantiates `OmniPrometheusStatLogger` (a thin subclass of upstream `vllm.v1.metrics.loggers.PrometheusStatLogger`) and feeds it scheduler stats and iteration stats after processing each batch of engine outputs. This populates the standard ~37 vLLM metric families (TTFT, ITL, TPOT, KV cache usage, etc.) using the same upstream code path — but with the `engine` label reshaped into `stage` + `replica` so multi-replica deployments produce distinct series per replica. See the next section for the wrap mechanics. @@ -174,7 +174,7 @@ second. ### Pipeline (4) | Metric | Type | Labels | Description | -|--------|------|--------|-------------| +| -------- | ------ | -------- | ------------- | | `vllm_omni:num_requests_running` | Gauge | `model_name` | Requests currently executing across all stages | | `vllm_omni:num_requests_waiting` | Gauge | `model_name` | Requests queued but not yet scheduled | | `vllm_omni:requests_success_total` | Counter | `model_name`, `finished_reason` | Total requests by completion reason ({stop, length, abort, ...}); aborts cover client-disconnect / cancellation paths in addition to upstream `FinishReason.ABORT` | @@ -183,7 +183,7 @@ second. ### Image and diffusion service-level metrics | Metric | Type | Labels | Description | -|--------|------|--------|-------------| +| -------- | ------ | -------- | ------------- | | `vllm_omni:stage_gen_time_s` | Histogram | `model_name`, `stage`, `stage_type` | Stage submit to finished output; includes in-stage queueing | | `vllm_omni:request_queue_wait_s` | Histogram | `model_name` | Orchestration-layer queue wait; a present zero is recorded | | `vllm_omni:stage_waiting_requests` | Gauge | `model_name`, `stage` | Sum of the latest waiting snapshots across live replicas in the stage | @@ -200,7 +200,7 @@ second. Labels: `{model_name, stage, replica}`. | Metric | Type | Description | -|--------|------|-------------| +| -------- | ------ | ------------- | | `vllm_omni:diffusion_exec_s` | Histogram | Core diffusion-step execution time per request | | `vllm_omni:diffusion_exec_per_step_s` | Histogram | Core diffusion execution divided by inference-step count | | `vllm_omni:diffusion_preprocess_s` | Histogram | Diffusion input preprocessing time | @@ -218,7 +218,7 @@ The optional breakdown and memory families are sparse by design: when an engine Labels: `{model_name, stage, replica}` plus the listed extra label. | Metric | Type | Extra label | Description | -|--------|------|-------------|-------------| +| -------- | ------ | ------------- | ------------- | | `vllm_omni:audio_ttfp_s` | Histogram | — | Time from request arrival to first audio packet/frame | | `vllm_omni:audio_duration_s` | Histogram | — | Audio content duration (`audio_frames / sample_rate`) | | `vllm_omni:audio_rtf` | Histogram | — | Real-time factor `stage_gen_time_s / audio_duration_s` (SLO `< 1`); uses `RTF_BUCKETS` | @@ -227,12 +227,29 @@ Labels: `{model_name, stage, replica}` plus the listed extra label. | `vllm_omni:audio_continuity_ok_total` | Counter | `threshold_ms` | Incremented when the request's worst underrun stayed below `threshold_ms` | | `vllm_omni:audio_skipped_requests_total` | Counter | `reason` | Silent-loss counter — code2wav rejected malformed codec input and returned `200 OK` with empty audio | +### Speech streaming (2) + +Labels: `{model_name}` plus the listed extra label. + +| Metric | Type | Extra label | Description | +| -------- | ------ | ------------- | ------------- | +| `vllm_omni:speech_stream_aborted_total` | Counter | `reason` | Interrupted audio generators, including before first PCM; reasons: `cancelled`, `closed`, `engine_dead`, `error` | +| `vllm_omni:speech_stream_completed_total` | Counter | — | Normally completed audio generators; does not confirm client receipt | + +Scope: started Speech audio generators (raw/SSE/WebSocket), counted once per +generation; WebSocket counts each sentence. Chat, non-streaming Speech, and +failures before generator execution are excluded. Interrupted streams do not +emit continuity or underrun samples. + +Interruption ratio: `aborted / (aborted + completed)`, using increments over the +same time window and summing all reasons per model. + ### Cross-stage transfer (4) Labels: `{model_name, from_stage, from_replica, to_stage, to_replica}`. | Metric | Type | Description | -|--------|------|-------------| +| -------- | ------ | ------------- | | `vllm_omni:transfer_size_bytes` | Histogram | Per-transfer payload size in bytes | | `vllm_omni:transfer_tx_s` | Histogram | Sender-side time (serialize + submit to connector) | | `vllm_omni:transfer_rx_s` | Histogram | Receiver-side time (recv + deserialize) | @@ -245,8 +262,8 @@ After the wrap, every upstream `vllm:*` family — TTFT, ITL, TPOT, e2e latency, ## Naming Convention - All time-bearing metrics use the `_s` suffix (values in seconds). Two bucket families are used: - - `SECONDS_BUCKETS` (0.05 s – 300 s) for e2e / generation / TTFP style values. - - `SECONDS_FAST_BUCKETS` (0.001 s – 60 s) for fine-grained cross-stage transfer and audio-underrun values that need millisecond-level resolution. + - `SECONDS_BUCKETS` (0.05 s – 300 s) for e2e / generation / TTFP style values. + - `SECONDS_FAST_BUCKETS` (0.001 s – 60 s) for fine-grained cross-stage transfer and audio-underrun values that need millisecond-level resolution. - Counters use the `_total` suffix (auto-appended by `prometheus_client`). - Sizes use the `_bytes` suffix. - All omni-specific families are prefixed `vllm_omni:`. The upstream `unregister_vllm_metrics()` function is monkey-patched to a scoped version that still strips upstream `vllm:*` collectors (so multi-engine init within one process does not crash on duplicate registration) but preserves anything prefixed `vllm_omni:` / `vllm_omni`. diff --git a/docs/usage/metrics.md b/docs/usage/metrics.md index 184a76eefad..1e0c1d06876 100644 --- a/docs/usage/metrics.md +++ b/docs/usage/metrics.md @@ -12,7 +12,7 @@ curl http://localhost:8000/metrics ## Metric Namespaces | Prefix | Source | Present when | -|--------|--------|--------------| +| -------- | -------- | -------------- | | `vllm_omni:` | vLLM-Omni orchestrator / audio modality / cross-stage transfer | Pipeline-dependent | | `vllm:` | Upstream vLLM engine, wrapped by `OmniPrometheusStatLogger` to expose `{stage, replica}` | Pipeline includes an LLM (AR) stage | | `http_` / `process_` | Uvicorn / Python runtime | Always | @@ -24,7 +24,7 @@ Defined in `vllm_omni/metrics/prometheus.py`. Track request lifecycle across the ### Request counts | Metric | Type | Labels | Description | -|--------|------|--------|-------------| +| -------- | ------ | -------- | ------------- | | `vllm_omni:num_requests_running` | Gauge | `model_name` | Pipeline-global in-flight requests (dispatched to engine, not yet finalized) | | `vllm_omni:num_requests_waiting` | Gauge | `model_name` | Requests waiting in the Orchestrator queue | | `vllm_omni:requests_success_total` | Counter | `model_name`, `finished_reason` | Total requests by completion reason. `finished_reason` ∈ {`stop`, `length`, `abort`, ...} mirroring upstream `vllm:request_success_total`; aborts cover client disconnect / cancellation paths in addition to upstream `FinishReason.ABORT` | @@ -32,15 +32,15 @@ Defined in `vllm_omni/metrics/prometheus.py`. Track request lifecycle across the ### Latency | Metric | Type | Labels | Description | -|--------|------|--------|-------------| +| -------- | ------ | -------- | ------------- | | `vllm_omni:e2e_request_latency_s` | Histogram | `model_name` | Pipeline-global end-to-end request latency in seconds | ## Audio Modality Metrics (`vllm_omni:`) -Emitted at request finalize, except for `audio_ttfp_s` (streaming-hook at the first audio packet) and `audio_underrun_s` / `audio_continuity_ok_total` (streaming finalize, after the chunk stream is exhausted). All carry `{model_name, stage, replica}` plus the listed extra label. +Emitted at request finalize, except for `audio_ttfp_s` (the first audio packet/frame from streaming Chat or Speech APIs) and `audio_underrun_s` / `audio_continuity_ok_total` (Chat or Speech streaming finalize, after the chunk stream is exhausted). Non-streaming Speech latency is represented by `e2e_request_latency_s`, not TTFP. All audio metrics carry `{model_name, stage, replica}` plus the listed extra label. | Metric | Type | Extra label | Description | -|--------|------|-------------|-------------| +| -------- | ------ | ------------- | ------------- | | `vllm_omni:audio_ttfp_s` | Histogram | — | Time from request arrival to first audio packet/frame | | `vllm_omni:audio_duration_s` | Histogram | — | Audio content duration (`audio_frames / sample_rate`) | | `vllm_omni:audio_rtf` | Histogram | — | Real-time factor (`stage_gen_time_s / audio_duration_s`); streaming TTS SLO red line `< 1`; uses `RTF_BUCKETS` | @@ -50,13 +50,22 @@ Emitted at request finalize, except for `audio_ttfp_s` (streaming-hook at the fi | `vllm_omni:audio_skipped_requests_total` | Counter | `reason` | Silent-loss counter — code2wav rejected malformed codec input and returned `200 OK` with empty audio | The continuity math comes from `vllm_omni/benchmarks/audio_continuity.py::compute_continuity_stats` so the server-side observation aligns with the bench-side definition. +For Speech, continuity applies to raw audio, SSE, and WebSocket streaming. WAV headers and empty chunks are excluded; non-streaming Speech requests do not emit continuity samples. + +Continuity rate for the default 100 ms threshold can be derived from the successful-continuity counter and the underrun Histogram's request count: + +```promql +sum by (model_name, stage, replica) (rate(vllm_omni:audio_continuity_ok_total{threshold_ms="100"}[5m])) +/ +sum by (model_name, stage, replica) (rate(vllm_omni:audio_underrun_s_count[5m])) +``` ## Diffusion Metrics (`vllm_omni:`) Per-request timing breakdowns for diffusion (image/video) stages. Emitted at request finalize when `stage_metrics.diffusion_metrics` is present. All carry `{model_name, stage, replica}`. | Metric | Type | Description | -|--------|------|-------------| +| -------- | ------ | ------------- | | `vllm_omni:diffusion_exec_s` | Histogram | DiT forward pass execution time per request in seconds | | `vllm_omni:diffusion_exec_per_step_s` | Histogram | DiT forward pass execution time per denoising step in seconds (`exec_s / num_inference_steps`) | | `vllm_omni:diffusion_preprocess_s` | Histogram | Diffusion input preprocessing time per request in seconds | @@ -80,7 +89,7 @@ rate(vllm_omni:diffusion_preprocess_s_sum[5m]) Per-physical-transfer histograms tracking the data hop between adjacent stages. Labels `{model_name, from_stage, from_replica, to_stage, to_replica}` let dashboards attribute latency to specific replica edges. `from_replica` / `to_replica` are resolved from the orchestrator's sticky-routing binding (`stage_pool.get_bound_replica_id(request_id)`), so no extra plumbing through `TransferEdgeStats` is needed. | Metric | Type | Description | -|--------|------|-------------| +| -------- | ------ | ------------- | | `vllm_omni:transfer_size_bytes` | Histogram | Per-transfer payload size in bytes | | `vllm_omni:transfer_tx_s` | Histogram | Sender-side time (serialize + submit to connector) | | `vllm_omni:transfer_rx_s` | Histogram | Receiver-side time (recv + deserialize) | @@ -106,7 +115,7 @@ For the full list of upstream metrics, see [the vLLM docs](https://github.com/vl ## Metric Availability by Pipeline Type | Metric group | Multi-stage LLM (Qwen3-Omni) | -|---|---| +| --- | --- | | `vllm_omni:` request tracking + latency | With `--log-stats` | | `vllm_omni:` audio modality | With `--log-stats`, if pipeline has a talker stage | | `vllm_omni:` transfer | With `--log-stats`, if pipeline has ≥ 2 stages | diff --git a/tests/entrypoints/openai_api/test_serving_speech.py b/tests/entrypoints/openai_api/test_serving_speech.py index 6040bede35e..732f86f2f53 100644 --- a/tests/entrypoints/openai_api/test_serving_speech.py +++ b/tests/entrypoints/openai_api/test_serving_speech.py @@ -3287,7 +3287,7 @@ def test_raw_stream_passes_tts_params_to_common_guard(self, streaming_app, monke finalized_tts_params = {"_qwen3_tts_effective_max_tokens": [192]} captured: dict = {} - async def prepare(_request, request_id=None): + async def prepare(_request, request_id=None, arrival_time=None): return request_id, object(), finalized_tts_params async def generate_chunks( @@ -3296,6 +3296,7 @@ async def generate_chunks( _response_format="pcm", raw_request=None, request_start_s=None, + request_arrival_ts=None, include_sample_rate=False, usage_acc=None, tts_params=None, diff --git a/tests/entrypoints/openai_api/test_serving_speech_metrics.py b/tests/entrypoints/openai_api/test_serving_speech_metrics.py new file mode 100644 index 00000000000..a5e284a0177 --- /dev/null +++ b/tests/entrypoints/openai_api/test_serving_speech_metrics.py @@ -0,0 +1,465 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni project + +"""Prometheus metric coverage for the Speech API audio path.""" + +import asyncio +from types import SimpleNamespace + +import numpy as np +import pytest +import torch + +from vllm_omni.entrypoints.openai import serving_speech as speech_module +from vllm_omni.entrypoints.openai.protocol.audio import OpenAICreateSpeechRequest +from vllm_omni.entrypoints.openai.serving_speech import OmniOpenAIServingSpeech + +pytestmark = [pytest.mark.core_model, pytest.mark.cpu] + + +class _MetricsStub: + def __init__(self) -> None: + self.ttfp_calls: list[tuple[str, str, float]] = [] + self.underrun_calls: list[tuple[str, str, float]] = [] + self.continuity_calls: list[tuple[str, str, int]] = [] + self.abort_calls: list[str] = [] + self.completed = 0 + + def inc_speech_stream_completed(self) -> None: + self.completed += 1 + + def inc_speech_stream_aborted(self, reason: str) -> None: + self.abort_calls.append(reason) + + def observe_audio_ttfp(self, stage: str, replica: str, seconds: float) -> None: + self.ttfp_calls.append((stage, replica, seconds)) + + def observe_audio_underrun(self, stage: str, replica: str, seconds: float) -> None: + self.underrun_calls.append((stage, replica, seconds)) + + def inc_audio_continuity_ok(self, stage: str, replica: str, threshold_ms: int) -> None: + self.continuity_calls.append((stage, replica, threshold_ms)) + + +def _serving(metrics: _MetricsStub, *, adapter=None) -> OmniOpenAIServingSpeech: + serving = OmniOpenAIServingSpeech.__new__(OmniOpenAIServingSpeech) + serving._tts_model_type = "qwen3_tts" + serving.engine_client = SimpleNamespace(mod_metrics=metrics, request_states={}) + serving._get_tts_adapter = lambda: adapter + serving.create_audio = lambda audio_obj: SimpleNamespace( + audio_data=b"\0\0" * int(audio_obj.audio_tensor.size), + media_type="audio/pcm", + ) + serving._mark_ref_audio_artifact_ready_for_request = lambda request_id: None + serving._discard_ref_audio_artifact_warmup = lambda request_id: None + return serving + + +def _result(samples: int = 320, *, stage_id: int = 1, replica_id: int | None = 2) -> SimpleNamespace: + return SimpleNamespace( + multimodal_output={"audio": torch.zeros(samples), "sr": 16000}, + stage_id=stage_id, + replica_id=replica_id, + ) + + +async def _generate(*results): + for result in results: + yield result + + +@pytest.mark.asyncio +async def test_streaming_speech_observes_ttfp_once_on_first_pcm_payload(monkeypatch): + metrics = _MetricsStub() + serving = _serving(metrics) + monkeypatch.setattr("vllm_omni.entrypoints.openai.serving_speech.time.time", lambda: 100.25) + + chunks = serving._generate_audio_chunks( + _generate(_result(), _result()), + request_id="speech-test", + request_arrival_ts=100.0, + ) + assert len([chunk async for chunk in chunks]) == 2 + assert metrics.ttfp_calls == [("1", "2", pytest.approx(0.25))] + assert metrics.underrun_calls == [("1", "2", pytest.approx(0.0))] + assert metrics.continuity_calls == [("1", "2", 100)] + assert metrics.abort_calls == [] + assert metrics.completed == 1 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("before_pcm", [True, False]) +@pytest.mark.parametrize( + ("error_type", "reason"), + [(asyncio.CancelledError, "cancelled"), (speech_module.EngineDeadError, "engine_dead"), (RuntimeError, "error")], +) +async def test_speech_stream_abort_counted_once_without_continuity(before_pcm, error_type, reason): + metrics = _MetricsStub() + serving = _serving(metrics) + + async def results(): + if not before_pcm: + yield _result() + raise error_type() + + chunks = serving._generate_audio_chunks(results(), request_id="speech-test", request_arrival_ts=100.0) + with pytest.raises(error_type): + _ = [chunk async for chunk in chunks] + + assert metrics.abort_calls == [reason] + assert metrics.completed == 0 + assert len(metrics.ttfp_calls) == (0 if before_pcm else 1) + assert metrics.underrun_calls == [] + assert metrics.continuity_calls == [] + + +@pytest.mark.asyncio +async def test_speech_stream_close_counted_once_without_continuity(): + metrics = _MetricsStub() + serving = _serving(metrics) + chunks = serving._generate_audio_chunks( + _generate(_result(), _result()), request_id="speech-test", request_arrival_ts=100.0 + ) + assert await anext(chunks) + await chunks.aclose() + await chunks.aclose() + assert metrics.abort_calls == ["closed"] + assert metrics.completed == 0 + assert metrics.underrun_calls == [] + assert metrics.continuity_calls == [] + + +@pytest.mark.asyncio +async def test_streaming_speech_retries_ttfp_labels_without_moving_first_packet_time(monkeypatch): + metrics = _MetricsStub() + serving = _serving(metrics) + clock = {"now": 100.25, "perf": 0.25} + monkeypatch.setattr(speech_module.time, "time", lambda: clock["now"]) + monkeypatch.setattr(speech_module.time, "perf_counter", lambda: clock["perf"]) + + async def results(): + yield _result(replica_id=None) + clock["now"] = 101.0 + clock["perf"] = 0.26 + yield _result(replica_id=2) + clock["now"] = 101.25 + clock["perf"] = 0.27 + yield _result(replica_id=2) + + chunks = serving._generate_audio_chunks( + results(), + request_id="speech-test", + request_arrival_ts=100.0, + ) + assert len([chunk async for chunk in chunks]) == 3 + assert metrics.ttfp_calls == [("1", "2", pytest.approx(0.25))] + assert metrics.underrun_calls == [("1", "2", pytest.approx(0.0))] + assert metrics.continuity_calls == [("1", "2", 100)] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("mode", ["raw", "sse", "pcm"]) +async def test_speech_disconnect_during_send_closes_entire_generator_chain(mode): + metrics = _MetricsStub() + serving = _serving(metrics) + closed = [] + + async def source(): + try: + yield _result() + yield _result() + finally: + closed.append(True) + + # Retain the source to prove cleanup does not depend on garbage collection. + generator = source() + method = { + "raw": serving._generate_audio_chunks, + "sse": serving._generate_audio_sse_events, + "pcm": serving._generate_pcm_chunks, + }[mode] + stream = method(generator, request_id="speech-test", request_arrival_ts=100.0) + + async def send(message): + if message["type"] == "http.response.body": + raise OSError("client disconnected") + + if mode == "pcm": + # WebSocket handlers already explicitly close their PCM iterator. + assert await anext(stream) + await stream.aclose() + else: + response = speech_module._SpeechStreamingResponse(stream) + with pytest.raises(OSError, match="client disconnected"): + await response.stream_response(send) + assert closed == [True] + assert metrics.abort_calls == ["closed"] + assert metrics.completed == 0 + assert metrics.underrun_calls == [] + assert metrics.continuity_calls == [] + + +@pytest.mark.parametrize("arrival_ts", [0.0, -1.0]) +def test_speech_ttfp_invalid_arrival_does_not_consume_guard(arrival_ts): + metrics = _MetricsStub() + serving = _serving(metrics) + state = SimpleNamespace( + external_request_id="speech-test", + request_arrival_ts=arrival_ts, + first_audio_ts=None, + audio_emit_stage_id=None, + audio_emit_replica_id=None, + ) + serving.engine_client.request_states = {"internal-test": state} + + assert serving._observe_speech_audio_ttfp(request_id="speech-test", result=_result(), first_packet_ts=100.25) == ( + 1, + 2, + False, + ) + assert metrics.ttfp_calls == [] + assert state.first_audio_ts is None + assert state.audio_emit_stage_id is None + assert state.audio_emit_replica_id is None + + state.request_arrival_ts = 100.0 + assert serving._observe_speech_audio_ttfp(request_id="speech-test", result=_result(), first_packet_ts=100.25) == ( + 1, + 2, + True, + ) + assert metrics.ttfp_calls == [("1", "2", pytest.approx(0.25))] + assert state.first_audio_ts == 100.25 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("during_send", [False, True]) +async def test_speech_response_task_cancellation_closes_source(during_send): + metrics = _MetricsStub() + serving = _serving(metrics) + waiting = asyncio.Event() + blocker = asyncio.Event() + closed = [] + + async def source(): + try: + if not during_send: + waiting.set() + await blocker.wait() + yield _result() + finally: + closed.append(True) + + async def send(message): + if during_send and message["type"] == "http.response.body": + waiting.set() + await blocker.wait() + + generator = source() + stream = serving._generate_audio_chunks(generator, request_id="speech-test", request_arrival_ts=100.0) + response = speech_module._SpeechStreamingResponse(stream) + task = asyncio.create_task(response.stream_response(send)) + try: + await asyncio.wait_for(waiting.wait(), timeout=5) + finally: + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + assert closed == [True] + assert metrics.abort_calls == ["closed" if during_send else "cancelled"] + assert metrics.completed == 0 + assert metrics.underrun_calls == [] + + +@pytest.mark.asyncio +async def test_streaming_speech_does_not_count_empty_payload_as_first_packet(monkeypatch): + metrics = _MetricsStub() + serving = _serving(metrics) + monkeypatch.setattr("vllm_omni.entrypoints.openai.serving_speech.time.time", lambda: 100.25) + + chunks = serving._generate_audio_chunks( + _generate(_result(samples=0), _result()), + request_id="speech-test", + request_arrival_ts=100.0, + ) + assert len([chunk async for chunk in chunks]) == 2 + assert metrics.ttfp_calls == [("1", "2", pytest.approx(0.25))] + assert metrics.underrun_calls == [("1", "2", pytest.approx(0.0))] + assert metrics.continuity_calls == [("1", "2", 100)] + + +@pytest.mark.asyncio +async def test_streaming_speech_does_not_finalize_continuity_on_error(monkeypatch): + metrics = _MetricsStub() + serving = _serving(metrics) + monkeypatch.setattr("vllm_omni.entrypoints.openai.serving_speech.time.time", lambda: 100.25) + + async def failing_stream(): + yield _result() + raise RuntimeError("stream failed") + + chunks = serving._generate_audio_chunks( + failing_stream(), + request_id="speech-test", + request_arrival_ts=100.0, + ) + with pytest.raises(RuntimeError, match="stream failed"): + _ = [chunk async for chunk in chunks] + + assert len(metrics.ttfp_calls) == 1 + assert metrics.underrun_calls == [] + assert metrics.continuity_calls == [] + + +@pytest.mark.asyncio +async def test_streaming_speech_does_not_finalize_continuity_on_validation_error(monkeypatch): + metrics = _MetricsStub() + + def reject_generation(_tts_params, **_kwargs): + raise RuntimeError("generation validation failed") + + adapter = SimpleNamespace(validates_generation=True, validate_generation=reject_generation) + serving = _serving(metrics, adapter=adapter) + monkeypatch.setattr("vllm_omni.entrypoints.openai.serving_speech.time.time", lambda: 100.25) + + chunks = serving._generate_audio_chunks( + _generate(_result()), + request_id="speech-test", + request_arrival_ts=100.0, + tts_params={"task_type": ["Base"]}, + ) + with pytest.raises(RuntimeError, match="generation validation failed"): + _ = [chunk async for chunk in chunks] + + assert len(metrics.ttfp_calls) == 1 + assert metrics.underrun_calls == [] + assert metrics.continuity_calls == [] + + +@pytest.mark.asyncio +async def test_streaming_speech_reports_late_chunk_as_underrun(monkeypatch): + metrics = _MetricsStub() + serving = _serving(metrics) + monkeypatch.setattr("vllm_omni.entrypoints.openai.serving_speech.time.time", lambda: 100.25) + perf_times = iter((0.0, 0.1, 0.1, 2.0, 2.0)) + monkeypatch.setattr("vllm_omni.entrypoints.openai.serving_speech.time.perf_counter", lambda: next(perf_times)) + + chunks = serving._generate_audio_chunks( + _generate(_result(), _result()), + request_id="speech-test", + request_arrival_ts=100.0, + ) + assert len([chunk async for chunk in chunks]) == 2 + assert metrics.underrun_calls[0][:2] == ("1", "2") + assert metrics.underrun_calls[0][2] > 0.1 + assert metrics.continuity_calls == [] + + +@pytest.mark.asyncio +async def test_non_streaming_speech_does_not_observe_ttfp(): + metrics = _MetricsStub() + serving = _serving(metrics) + serving._audio_encode_speed = lambda _request: 1.0 + + async def prepare(_request, **_kwargs): + result = _result() + result.metrics = {} + return "speech-test", _generate(result), {} + + serving._prepare_speech_generation = prepare + request = OpenAICreateSpeechRequest(input="hello", response_format="pcm") + + audio_data, media_type = await serving._generate_audio_bytes(request, request_arrival_ts=100.0) + + assert audio_data + assert media_type == "audio/pcm" + assert metrics.ttfp_calls == [] + assert metrics.completed == 0 + assert metrics.abort_calls == [] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("flush_only", [False, True]) +@pytest.mark.parametrize("response_format", ["pcm", "wav"]) +async def test_resampled_speech_records_flush_with_audio_producer(monkeypatch, flush_only, response_format): + metrics = _MetricsStub() + serving = _serving(metrics) + clock = {"now": 0.0} + finalized = [] + original_finalize = speech_module.observe_audio_streaming_finalize + + def capture_finalize(*args, **kwargs): + finalized.append(kwargs) + original_finalize(*args, **kwargs) + + class BufferedResampler: + def __init__(self, source_rate, target_rate): + assert (source_rate, target_rate) == (16000, 24000) + + def process(self, chunk, *, final=False): + clock["now"] = 0.5 if final else 0.25 + if not final and flush_only: + return np.empty(0, dtype=np.float32) + return np.zeros(2400, dtype=np.float32) + + monkeypatch.setattr(speech_module, "StreamingAudioResampler", BufferedResampler) + monkeypatch.setattr(speech_module, "observe_audio_streaming_finalize", capture_finalize) + monkeypatch.setattr(speech_module.time, "time", lambda: 100.0 + clock["now"]) + monkeypatch.setattr(speech_module.time, "perf_counter", lambda: clock["now"]) + # The last result is not audio and must not supply the flush metric labels. + non_audio = SimpleNamespace(multimodal_output={"timestamps": []}, stage_id=9, replica_id=8) + chunks = [ + chunk + async for chunk in serving._generate_audio_chunks( + _generate(_result(), non_audio), + request_id="speech-test", + response_format=response_format, + request_start_s=0.0, + request_arrival_ts=100.0, + target_sample_rate=24000, + ) + ] + + if response_format == "wav": + assert chunks.pop(0).startswith(b"RIFF") + expected_arrivals = [0.5] if flush_only else [0.25, 0.5] + assert [len(chunk) for chunk in chunks] == [4800] * len(expected_arrivals) + assert metrics.ttfp_calls == [("1", "2", pytest.approx(expected_arrivals[0]))] + assert len(finalized) == 1 + assert finalized[0]["sample_rate"] == 24000 + assert finalized[0]["channels"] == 1 + assert finalized[0]["chunk_bytes"] == [4800] * len(expected_arrivals) + assert finalized[0]["chunk_arrival_times_s"] == expected_arrivals + assert metrics.underrun_calls == [("1", "2", pytest.approx(0.0 if flush_only else 0.15))] + assert metrics.continuity_calls == ([("1", "2", 100)] if flush_only else []) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("empty_pcm", [False, True]) +async def test_resampled_speech_without_pcm_does_not_emit_metrics(monkeypatch, empty_pcm): + metrics = _MetricsStub() + serving = _serving(metrics) + + class EmptyResampler: + def __init__(self, source_rate, target_rate): + pass + + def process(self, chunk, *, final=False): + # Also cover a nonempty flush waveform whose encoder emits no PCM. + return np.zeros(10 if final and empty_pcm else 0, dtype=np.float32) + + monkeypatch.setattr(speech_module, "StreamingAudioResampler", EmptyResampler) + serving.create_audio = lambda audio_obj: SimpleNamespace(audio_data=b"", media_type="audio/pcm") + chunks = [ + chunk + async for chunk in serving._generate_audio_chunks( + _generate(_result()), + request_id="speech-test", + request_arrival_ts=100.0, + target_sample_rate=24000, + ) + ] + assert not any(chunks) + assert metrics.ttfp_calls == [] + assert metrics.underrun_calls == [] + assert metrics.continuity_calls == [] diff --git a/tests/entrypoints/openai_api/test_serving_speech_stream.py b/tests/entrypoints/openai_api/test_serving_speech_stream.py index 6414ab4ff8b..39c9e66d950 100644 --- a/tests/entrypoints/openai_api/test_serving_speech_stream.py +++ b/tests/entrypoints/openai_api/test_serving_speech_stream.py @@ -48,7 +48,14 @@ def _build_test_app( speech_service.forced_aligner_enabled = False async def mock_generate_pcm_chunks( - _generator, _request_id, *, include_sample_rate=False, tts_params=None, collect=None + _generator, + _request_id, + *, + request_start_s=None, + request_arrival_ts=None, + include_sample_rate=False, + tts_params=None, + collect=None, ): for chunk in (b"\x01\x02", b"\x03\x04\x05"): yield (chunk, 24000) if include_sample_rate else chunk @@ -249,6 +256,7 @@ def test_idle_timeout_closes_reused_connection(self, mocker: MockerFixture): def test_streaming_multiple_binary_frames(self, mocker: MockerFixture): captured_requests = [] + captured_timing = {} captured_tts_params = [] speech_service = mocker.MagicMock(spec=OmniOpenAIServingSpeech) @@ -257,13 +265,27 @@ def test_streaming_multiple_binary_frames(self, mocker: MockerFixture): speech_service.engine_client.abort = mocker.AsyncMock() speech_service.forced_aligner_enabled = False - async def mock_prepare_speech_generation(request): + async def mock_prepare_speech_generation(request, *, arrival_time=None): captured_requests.append(request) + assert arrival_time is not None + captured_timing["prepare_arrival"] = arrival_time return "req-stream", object(), {"_qwen3_tts_effective_max_tokens": [192]} speech_service._prepare_speech_generation = mock_prepare_speech_generation - async def mock_generate_pcm_chunks(_generator, _request_id, *, include_sample_rate=False, tts_params=None): + async def mock_generate_pcm_chunks( + _generator, + _request_id, + *, + request_start_s=None, + request_arrival_ts=None, + include_sample_rate=False, + tts_params=None, + ): + assert request_start_s is not None + assert request_arrival_ts is not None + captured_timing["chunk_start"] = request_start_s + captured_timing["chunk_arrival"] = request_arrival_ts captured_tts_params.append(tts_params) for chunk in (b"\x01\x02", b"\x03\x04\x05", b"\x06"): yield (chunk, 24000) if include_sample_rate else chunk @@ -309,6 +331,8 @@ async def mock_generate_pcm_chunks(_generator, _request_id, *, include_sample_ra assert captured_requests[0].stream is True assert captured_requests[0].response_format == "pcm" assert captured_requests[0].initial_codec_chunk_frames == 12 + assert captured_timing["prepare_arrival"] == captured_timing["chunk_arrival"] + assert captured_timing["chunk_start"] > 0 assert captured_tts_params == [{"_qwen3_tts_effective_max_tokens": [192]}] assert speech_service._generate_audio_bytes.await_count == 0 @@ -341,8 +365,9 @@ def test_word_timestamps_emit_pipeline_json_frame(self, mocker: MockerFixture): speech_service.engine_client.abort = mocker.AsyncMock() speech_service.forced_aligner_enabled = True - async def mock_prepare_speech_generation(request): + async def mock_prepare_speech_generation(request, *, arrival_time=None): captured_requests.append(request) + assert arrival_time is not None return "req-stream", object(), {} speech_service._prepare_speech_generation = mock_prepare_speech_generation @@ -353,8 +378,17 @@ async def mock_prepare_speech_generation(request): # The forced-aligner stage rides the same generator: its pooling output # is surfaced via the ``collect`` channel once the audio has streamed. async def mock_generate_pcm_chunks( - _generator, _request_id, *, include_sample_rate=False, tts_params=None, collect=None + _generator, + _request_id, + *, + request_start_s=None, + request_arrival_ts=None, + include_sample_rate=False, + tts_params=None, + collect=None, ): + assert request_start_s is not None + assert request_arrival_ts is not None for chunk in (first_chunk, second_chunk): yield (chunk, 1000) if include_sample_rate else chunk if collect is not None: @@ -438,7 +472,14 @@ def test_word_timestamps_emit_word_dicts(self, mocker: MockerFixture): speech_service._prepare_speech_generation = mocker.AsyncMock(return_value=("req", object(), {})) async def mock_generate_pcm_chunks( - _generator, _request_id, *, include_sample_rate=False, tts_params=None, collect=None + _generator, + _request_id, + *, + request_start_s=None, + request_arrival_ts=None, + include_sample_rate=False, + tts_params=None, + collect=None, ): chunk = b"\x01" * 1000 yield (chunk, 1000) if include_sample_rate else chunk @@ -620,7 +661,15 @@ def test_streaming_generation_error_marks_audio_done(self, mocker: MockerFixture speech_service.engine_client.abort = mocker.AsyncMock() speech_service.forced_aligner_enabled = False - async def mock_generate_pcm_chunks(_generator, _request_id, *, include_sample_rate=False, tts_params=None): + async def mock_generate_pcm_chunks( + _generator, + _request_id, + *, + request_start_s=None, + request_arrival_ts=None, + include_sample_rate=False, + tts_params=None, + ): yield b"\x01\x02" raise RuntimeError("stream boom") @@ -716,7 +765,15 @@ def test_disconnect_aborts_streaming_request(self, mocker: MockerFixture): speech_service.engine_client.abort = mocker.AsyncMock() speech_service.forced_aligner_enabled = False - async def mock_generate_pcm_chunks(_generator, _request_id, *, include_sample_rate=False, tts_params=None): + async def mock_generate_pcm_chunks( + _generator, + _request_id, + *, + request_start_s=None, + request_arrival_ts=None, + include_sample_rate=False, + tts_params=None, + ): yield b"\x01\x02" speech_service._generate_pcm_chunks = mock_generate_pcm_chunks diff --git a/tests/metrics/test_modality.py b/tests/metrics/test_modality.py index 14bb4f309ca..a918dce53c7 100644 --- a/tests/metrics/test_modality.py +++ b/tests/metrics/test_modality.py @@ -48,6 +48,8 @@ def _sample_value(output: str, line_prefix: str) -> float | None: defs.AUDIO_UNDERRUN_S, defs.AUDIO_CONTINUITY_OK_METRIC, defs.AUDIO_SKIPPED_REQUESTS_METRIC, + defs.SPEECH_STREAM_ABORTED_METRIC, + defs.SPEECH_STREAM_COMPLETED_METRIC, defs.DIFFUSION_EXEC_S, defs.DIFFUSION_EXEC_PER_STEP_S, defs.DIFFUSION_PREPROCESS_S, @@ -70,6 +72,8 @@ def test_all_locked_families_present(self, mod: OmniModalityMetrics) -> None: mod.observe_audio_underrun("s", "r", 0.01) mod.inc_audio_continuity_ok("s", "r", 100) mod.inc_audio_skipped("s", "r", "malformed_codec") + mod.inc_speech_stream_aborted("error") + mod.inc_speech_stream_completed() mod.observe_diffusion_exec("s", "r", 0.5) mod.observe_diffusion_exec_per_step("s", "r", 0.01) mod.observe_diffusion_preprocess("s", "r", 0.01) @@ -91,6 +95,23 @@ def test_all_locked_families_present(self, mod: OmniModalityMetrics) -> None: class TestAudio: + def test_speech_stream_abort_counter_and_disabled_logging(self): + model = "test-speech-stream-abort-counter" + metrics = OmniModalityMetrics(model_name=model) + disabled = OmniModalityMetrics(model_name=model, log_stats=False) + metrics.inc_speech_stream_completed() + disabled.inc_speech_stream_completed() + assert REGISTRY.get_sample_value(defs.SPEECH_STREAM_COMPLETED_METRIC + "_total", {"model_name": model}) == 1.0 + for reason in ("cancelled", "closed", "engine_dead", "error"): + metrics.inc_speech_stream_aborted(reason) + disabled.inc_speech_stream_aborted(reason) + assert ( + REGISTRY.get_sample_value( + defs.SPEECH_STREAM_ABORTED_METRIC + "_total", {"model_name": model, "reason": reason} + ) + == 1.0 + ) + def test_audio_ttfp_observed(self, mod: OmniModalityMetrics) -> None: stage, replica = "talker_ttfp", "0" mod.observe_audio_ttfp(stage, replica, 0.42) @@ -242,6 +263,9 @@ def observe_audio_duration(self, s, r, d): def observe_audio_rtf(self, s, r, rtf): self.calls.append(("observe_audio_rtf", s, r, rtf)) + def observe_audio_ttfp(self, s, r, t): + self.calls.append(("observe_audio_ttfp", s, r, t)) + def observe_audio_underrun(self, s, r, u): self.calls.append(("observe_audio_underrun", s, r, u)) @@ -599,7 +623,6 @@ def test_audio_diffusion_stage_fires_both_paths(self): class TestObserveAudioFirstPacket: def test_observes_with_valid_inputs(self): stub = _StubModMetrics() - stub.observe_audio_ttfp = lambda s, r, t: stub.calls.append(("observe_audio_ttfp", s, r, t)) observe_audio_first_packet( stub, @@ -612,19 +635,16 @@ def test_observes_with_valid_inputs(self): def test_replica_none_skipped(self): stub = _StubModMetrics() - stub.observe_audio_ttfp = lambda s, r, t: stub.calls.append(("observe_audio_ttfp", s, r, t)) observe_audio_first_packet(stub, stage_id=1, replica_id=None, arrival_ts=100.0, now_ts=100.5) assert stub.calls == [] def test_arrival_ts_zero_skipped(self): stub = _StubModMetrics() - stub.observe_audio_ttfp = lambda s, r, t: stub.calls.append(("observe_audio_ttfp", s, r, t)) observe_audio_first_packet(stub, stage_id=1, replica_id=0, arrival_ts=0.0, now_ts=100.5) assert stub.calls == [] def test_clock_skew_clamped_to_zero(self): stub = _StubModMetrics() - stub.observe_audio_ttfp = lambda s, r, t: stub.calls.append(("observe_audio_ttfp", s, r, t)) observe_audio_first_packet(stub, stage_id=1, replica_id=0, arrival_ts=100.5, now_ts=100.0) assert stub.calls == [("observe_audio_ttfp", "1", "0", 0.0)] @@ -668,6 +688,24 @@ def test_late_chunk_emits_nonzero_underrun_and_no_continuity_inc(self): assert underrun_calls[0][-1] > 0.1 assert not any(c[0] == "inc_audio_continuity_ok" for c in stub.calls) + def test_stereo_channel_count_is_used_for_playback_rate(self): + stub = _StubModMetrics() + # At 100 Hz, s16le stereo consumes 400 bytes/s. The first 200-byte + # chunk buffers 0.5s, so a second chunk arriving at 1.0s underruns by 0.5s. + observe_audio_streaming_finalize( + stub, + stage_id=1, + replica_id=0, + chunk_arrival_times_s=[0.0, 1.0], + chunk_bytes=[200, 200], + sample_rate=100, + channels=2, + threshold_s=0.1, + ) + underrun_calls = [c for c in stub.calls if c[0] == "observe_audio_underrun"] + assert underrun_calls == [("observe_audio_underrun", "1", "0", pytest.approx(0.5))] + assert not any(c[0] == "inc_audio_continuity_ok" for c in stub.calls) + def test_empty_arrivals_skipped(self): stub = _StubModMetrics() observe_audio_streaming_finalize( diff --git a/vllm_omni/entrypoints/openai/serving_speech.py b/vllm_omni/entrypoints/openai/serving_speech.py index 2f4e472f824..d4aa72818ae 100644 --- a/vllm_omni/entrypoints/openai/serving_speech.py +++ b/vllm_omni/entrypoints/openai/serving_speech.py @@ -13,12 +13,14 @@ import time from collections import OrderedDict from concurrent.futures import ThreadPoolExecutor +from contextlib import aclosing from http import HTTPStatus from pathlib import Path -from typing import Any +from typing import Any, cast from urllib.parse import urlparse from urllib.request import url2pathname +import anyio import numpy as np import soundfile as sf import torch @@ -59,6 +61,7 @@ tts_entry_stage_archs, ) from vllm_omni.entrypoints.utils import coerce_param_message_types +from vllm_omni.metrics.modality import observe_audio_first_packet, observe_audio_streaming_finalize from vllm_omni.outputs import OmniRequestOutput from vllm_omni.utils.speaker_cache import get_speaker_cache @@ -194,6 +197,17 @@ def _validate_path_within_directory(file_path: Path, directory: Path) -> bool: return False +class _SpeechStreamingResponse(StreamingResponse): + async def stream_response(self, send) -> None: + try: + await super().stream_response(send) + finally: + # A disconnect during send leaves the iterator suspended at yield. + # Close it explicitly, even inside Starlette's cancelled scope. + with anyio.CancelScope(shield=True): + await cast(Any, self.body_iterator).aclose() + + class OmniOpenAIServingSpeech(OpenAIServing, AudioMixin): _diffusion_mode: bool = False _media_connector: MediaConnector | None = None @@ -451,9 +465,7 @@ def _find_tts_stage(self): """Find and return the TTS stage config, or None if not found.""" tts_stage_keys = all_tts_stage_keys() entry_stage_archs = tts_entry_stage_archs() - all_stages = frozenset( - getattr(stage.engine_args, "model_stage", None) for stage in self.engine_client.stage_configs - ) + all_stages = frozenset(stage.engine_args.model_stage for stage in self.engine_client.stage_configs) for stage in self.engine_client.stage_configs: engine_args = stage.engine_args model_stage = engine_args.model_stage @@ -1083,7 +1095,7 @@ def _validate_speech_sample_rate(self, request: OpenAICreateSpeechRequest) -> st return "sample_rate is not supported by the current TTS model" return None - def _validate_ref_audio_format(self, ref_audio: str) -> str | None: + def _validate_ref_audio_format(self, ref_audio: str | list[str] | None) -> str | None: """Validate ref_audio is a supported URI format. Returns error or None.""" if not isinstance(ref_audio, str): return "ref_audio must be a URL (http/https), base64 data URL (data:...), or file URI (file://...)" @@ -1362,6 +1374,57 @@ async def _resolve_ref_audio_many(self, ref_audio_list: list[str]) -> list[tuple resolved.append((wav_list, sr)) return resolved + def _observe_speech_audio_ttfp( + self, + *, + request_id: str, + result: Any, + request_arrival_ts: float | None = None, + first_packet_ts: float | None = None, + ) -> tuple[int | None, int | None, bool]: + """Emit the Speech API TTFP sample once the first PCM payload exists.""" + engine_client = getattr(self, "engine_client", None) + mod_metrics = getattr(engine_client, "mod_metrics", None) + if mod_metrics is None: + return None, None, False + + req_state = next( + ( + state + for state in getattr(engine_client, "request_states", {}).values() + if state.external_request_id == request_id + ), + None, + ) + if req_state is not None and req_state.first_audio_ts is not None: + return req_state.audio_emit_stage_id, req_state.audio_emit_replica_id, True + + arrival_ts = request_arrival_ts + if arrival_ts is None and req_state is not None: + arrival_ts = req_state.request_arrival_ts + stage_id = getattr(result, "stage_id", None) + replica_id = getattr(result, "replica_id", None) + if replica_id is None and req_state is not None and stage_id is not None: + stage_pools = getattr(getattr(engine_client, "engine", None), "stage_pools", None) + if stage_pools is not None and 0 <= stage_id < len(stage_pools): + replica_id = stage_pools[stage_id].get_bound_replica_id(req_state.request_id) + if arrival_ts is None or arrival_ts <= 0 or stage_id is None or replica_id is None: + return stage_id, replica_id, False + + now_ts = first_packet_ts if first_packet_ts is not None else time.time() + observe_audio_first_packet( + mod_metrics, + stage_id=stage_id, + replica_id=replica_id, + arrival_ts=arrival_ts, + now_ts=now_ts, + ) + if req_state is not None: + req_state.first_audio_ts = now_ts + req_state.audio_emit_stage_id = stage_id + req_state.audio_emit_replica_id = replica_id + return stage_id, replica_id, True + async def _generate_audio_chunks( self, generator, @@ -1369,6 +1432,7 @@ async def _generate_audio_chunks( response_format: str = "pcm", raw_request: Request | None = None, request_start_s: float | None = None, + request_arrival_ts: float | None = None, include_sample_rate: bool = False, usage_acc: SpeechOutputTokenCounter | None = None, tts_params: dict[str, Any] | None = None, @@ -1395,10 +1459,42 @@ async def _generate_audio_chunks( sample_rate_val = 24000 first_chunk = True first_audio_chunk_s: float | None = None + first_audio_packet_ts: float | None = None + ttfp_observed = False stream_start_s = request_start_s if request_start_s is not None else time.perf_counter() artifact_ready = False + audio_chunk_arrivals_s: list[float] = [] + audio_chunk_bytes: list[int] = [] + audio_stage_id: int | None = None + audio_replica_id: int | None = None + audio_channels = 1 source_sample_rate: int | None = None resampler: StreamingAudioResampler | None = None + last_audio_result: Any = None + + def record_stream_abort(reason: str) -> None: + mod_metrics = getattr(getattr(self, "engine_client", None), "mod_metrics", None) + if mod_metrics is not None: + mod_metrics.inc_speech_stream_aborted(reason) + + def record_audio_chunk(audio_bytes: bytes, chunk_np: np.ndarray, result: Any) -> None: + nonlocal first_audio_chunk_s, first_audio_packet_ts, ttfp_observed + nonlocal audio_stage_id, audio_replica_id, audio_channels + if not audio_bytes: + return + if first_audio_chunk_s is None: + first_audio_chunk_s = time.perf_counter() + first_audio_packet_ts = time.time() + audio_channels = _infer_audio_num_channels(np.asarray(chunk_np)) + if not ttfp_observed: + audio_stage_id, audio_replica_id, ttfp_observed = self._observe_speech_audio_ttfp( + request_id=request_id, + result=result, + request_arrival_ts=request_arrival_ts, + first_packet_ts=first_audio_packet_ts, + ) + audio_chunk_arrivals_s.append(max(time.perf_counter() - stream_start_s, 0.0)) + audio_chunk_bytes.append(len(audio_bytes)) # SSE supplies an accumulator for usage output. Raw-audio and WebSocket # streams retain terminal metrics only when their model adapter needs @@ -1419,6 +1515,7 @@ async def _generate_audio_chunks( if collect is not None and self._is_timestamps_output(res): collect["aligner_res"] = res continue + audio_output = cast(dict, audio_output) sr_raw = audio_output.get("sr") if sr_raw is not None: @@ -1452,6 +1549,9 @@ async def _generate_audio_chunks( output_sample_rate = target_sample_rate or sample_rate_val for chunk_tensor in new_chunks: + # Flush may run after a non-audio (e.g. aligner) result. + # Retain the producer of the audio buffered by the resampler. + last_audio_result = res chunk_np = ( chunk_tensor.float().detach().cpu().numpy() if hasattr(chunk_tensor, "float") else chunk_tensor ) @@ -1466,6 +1566,7 @@ async def _generate_audio_chunks( # as first audio; the post-loop guard below needs to # see an audio-less stream to fail the request. continue + wav_header: bytes | None = None # For WAV format, emit header before first audio chunk if response_format == "wav" and first_chunk: num_channels = _infer_audio_num_channels(np.asarray(chunk_np)) @@ -1474,7 +1575,6 @@ async def _generate_audio_chunks( num_channels=num_channels, bits_per_sample=16, ) - yield wav_header first_chunk = False # Convert audio to PCM bytes @@ -1485,9 +1585,10 @@ async def _generate_audio_chunks( speed=1.0, base64_encode=False, ) - if first_audio_chunk_s is None: - first_audio_chunk_s = time.perf_counter() - audio_bytes = self.create_audio(audio_obj).audio_data + audio_bytes = cast(bytes, self.create_audio(audio_obj).audio_data) + record_audio_chunk(audio_bytes, chunk_np, res) + if wav_header is not None: + yield wav_header if include_sample_rate: yield audio_bytes, output_sample_rate else: @@ -1497,8 +1598,9 @@ async def _generate_audio_chunks( final_chunk = resampler.process(np.empty((0,), dtype=np.float32), final=True) if final_chunk.size: output_sample_rate = target_sample_rate or sample_rate_val + wav_header = None if response_format == "wav" and first_chunk: - yield _create_wav_header( + wav_header = _create_wav_header( sample_rate=output_sample_rate, num_channels=1, bits_per_sample=16, @@ -1511,9 +1613,10 @@ async def _generate_audio_chunks( speed=1.0, base64_encode=False, ) - if first_audio_chunk_s is None: - first_audio_chunk_s = time.perf_counter() - audio_bytes = self.create_audio(audio_obj).audio_data + audio_bytes = cast(bytes, self.create_audio(audio_obj).audio_data) + record_audio_chunk(audio_bytes, final_chunk, last_audio_result) + if wav_header is not None: + yield wav_header if include_sample_rate: yield audio_bytes, output_sample_rate else: @@ -1527,8 +1630,21 @@ async def _generate_audio_chunks( # bytes, but they must terminate as an error rather than cleanly. if tts_params is not None and usage_acc is not None: self._validate_tts_generation(tts_params, usage_acc) + mod_metrics = getattr(getattr(self, "engine_client", None), "mod_metrics", None) + if mod_metrics is not None and audio_stage_id is not None and audio_replica_id is not None: + observe_audio_streaming_finalize( + mod_metrics, + stage_id=audio_stage_id, + replica_id=audio_replica_id, + chunk_arrival_times_s=audio_chunk_arrivals_s, + chunk_bytes=audio_chunk_bytes, + sample_rate=target_sample_rate or sample_rate_val, + channels=audio_channels, + ) self._mark_ref_audio_artifact_ready_for_request(request_id) artifact_ready = True + if mod_metrics is not None: + mod_metrics.inc_speech_stream_completed() total_ms = (time.perf_counter() - stream_start_s) * 1000.0 if first_audio_chunk_s is not None: first_chunk_ms = (first_audio_chunk_s - stream_start_s) * 1000.0 @@ -1544,7 +1660,11 @@ async def _generate_audio_chunks( request_id, total_ms, ) + except GeneratorExit: + record_stream_abort("closed") + raise except asyncio.CancelledError: + record_stream_abort("cancelled") total_ms = (time.perf_counter() - stream_start_s) * 1000.0 logger.info( "[SpeechE2E] request_id=%s stream=true status=cancelled total_ms=%.2f", @@ -1554,6 +1674,7 @@ async def _generate_audio_chunks( logger.info("Streaming request %s cancelled by client", request_id) raise except EngineDeadError as e: + record_stream_abort("engine_dead") total_ms = (time.perf_counter() - stream_start_s) * 1000.0 logger.error( "[SpeechE2E] request_id=%s stream=true status=engine_dead total_ms=%.2f", @@ -1573,6 +1694,7 @@ async def _generate_audio_chunks( ) raise except Exception as e: + record_stream_abort("error") total_ms = (time.perf_counter() - stream_start_s) * 1000.0 logger.exception( "[SpeechE2E] request_id=%s stream=true status=error total_ms=%.2f error=%s", @@ -1585,6 +1707,10 @@ async def _generate_audio_chunks( finally: if not artifact_ready: self._discard_ref_audio_artifact_warmup(request_id) + close = getattr(generator, "aclose", None) + if close is not None: + with anyio.CancelScope(shield=True): + await close() async def _generate_audio_sse_events( self, @@ -1593,6 +1719,7 @@ async def _generate_audio_sse_events( response_format: str = "pcm", raw_request: Request | None = None, request_start_s: float | None = None, + request_arrival_ts: float | None = None, request: OpenAICreateSpeechRequest | None = None, tts_params: dict[str, Any] | None = None, ): @@ -1613,24 +1740,28 @@ async def _generate_audio_sse_events( usage_acc = SpeechOutputTokenCounter() emitted_audio = False try: - async for chunk in self._generate_audio_chunks( - generator, - request_id, - response_format, - raw_request=raw_request, - request_start_s=request_start_s, - usage_acc=usage_acc, - tts_params=tts_params, - target_sample_rate=request.sample_rate if request is not None else None, - ): - payload = { - "type": "speech.audio.delta", - "audio": base64.b64encode(chunk).decode("ascii"), - "response_format": response_format, - } - data = json.dumps(payload, separators=(",", ":")) - emitted_audio = True - yield f"event: speech.audio.delta\ndata: {data}\n\n" + async with aclosing( + self._generate_audio_chunks( + generator, + request_id, + response_format, + raw_request=raw_request, + request_start_s=request_start_s, + request_arrival_ts=request_arrival_ts, + usage_acc=usage_acc, + tts_params=tts_params, + target_sample_rate=request.sample_rate if request is not None else None, + ) + ) as chunks: + async for chunk in chunks: + payload = { + "type": "speech.audio.delta", + "audio": base64.b64encode(chunk).decode("ascii"), + "response_format": response_format, + } + data = json.dumps(payload, separators=(",", ":")) + emitted_audio = True + yield f"event: speech.audio.delta\ndata: {data}\n\n" done_payload: dict[str, Any] = {"type": "speech.audio.done"} if request is not None: # Streaming path: output_tokens = sum of stage-0 deltas. @@ -1695,6 +1826,7 @@ async def _prepare_speech_generation( request: OpenAICreateSpeechRequest, request_id: str | None = None, has_inline_ref_audio: bool | None = None, + arrival_time: float | None = None, ) -> tuple[str, Any, dict[str, Any]]: if self.engine_client.errored: raise self.engine_client.dead_error @@ -1823,6 +1955,7 @@ async def _prepare_speech_generation( request_id=request_id, sampling_params_list=sampling_params_list, output_modalities=output_modalities, + arrival_time=arrival_time, ) self._track_ref_audio_artifact_warmup( request_id, @@ -1836,6 +1969,8 @@ async def _generate_pcm_chunks( generator, request_id: str, *, + request_start_s: float | None = None, + request_arrival_ts: float | None = None, include_sample_rate: bool = False, tts_params: dict[str, Any] | None = None, collect: dict | None = None, @@ -1848,28 +1983,36 @@ async def _generate_pcm_chunks( ``collect`` (when given) receives the forced-aligner stage's pooling output under ``"aligner_res"`` for downstream word-timestamp extraction. """ - async for chunk in self._generate_audio_chunks( - generator, - request_id, - response_format="pcm", - include_sample_rate=include_sample_rate, - tts_params=tts_params, - collect=collect, - target_sample_rate=target_sample_rate, - ): - yield chunk + async with aclosing( + self._generate_audio_chunks( + generator, + request_id, + response_format="pcm", + request_start_s=request_start_s, + request_arrival_ts=request_arrival_ts, + include_sample_rate=include_sample_rate, + tts_params=tts_params, + collect=collect, + target_sample_rate=target_sample_rate, + ) + ) as chunks: + async for chunk in chunks: + yield chunk async def _iter_pcm_audio_bytes(self, request: OpenAICreateSpeechRequest): """Yield raw PCM bytes for a speech request as soon as chunks are decoded.""" request_id, generator, tts_params = await self._prepare_speech_generation(request) try: - async for chunk in self._generate_pcm_chunks( - generator, - request_id, - tts_params=tts_params, - target_sample_rate=request.sample_rate, - ): - yield chunk + async with aclosing( + self._generate_pcm_chunks( + generator, + request_id, + tts_params=tts_params, + target_sample_rate=request.sample_rate, + ) + ) as chunks: + async for chunk in chunks: + yield chunk finally: self._discard_ref_audio_artifact_warmup(request_id) @@ -1881,14 +2024,19 @@ async def _generate_audio_bytes( usage_out: list[SpeechTokenUsage] | None = None, has_inline_ref_audio: bool | None = None, collect: dict | None = None, + request_arrival_ts: float | None = None, ) -> tuple[bytes | str, str]: # ``usage_out`` is an opt-in output channel: when a list is passed, the # computed SpeechTokenUsage is appended to it. The return stays a # 2-tuple so existing callers (and their test mocks) are unaffected; # batch and non-streaming response-header paths opt in when surfacing # usage outside the raw audio body. + request_arrival_ts = request_arrival_ts if request_arrival_ts is not None else time.time() request_id, generator, bytes_tts_params = await self._prepare_speech_generation( - request, request_id=request_id, has_inline_ref_audio=has_inline_ref_audio + request, + request_id=request_id, + has_inline_ref_audio=has_inline_ref_audio, + arrival_time=request_arrival_ts, ) artifact_ready = False @@ -1927,6 +2075,7 @@ async def _generate_audio_bytes( continue if step_key is None: continue + step_audio = cast(dict, step_audio) chunk = step_audio[step_key] candidates = chunk if isinstance(chunk, list) else [chunk] for cand in candidates: @@ -1948,6 +2097,7 @@ async def _generate_audio_bytes( audio_output, audio_key = self._extract_audio_output(audio_source) if audio_key is None: raise ValueError("TTS model did not produce audio output.") + audio_output = cast(dict, audio_output) # Surface forced-aligner word timestamps to the caller (set as a # response header) when requested and an aligner stage produced them. @@ -2074,7 +2224,7 @@ async def _create_diffusion_speech( request_id = f"speech-{random_uuid()}" prompt: dict[str, Any] = {"input": request.input} if request.ref_audio: - wav, sr, _ = await self._resolve_ref_audio(request.ref_audio) + wav, sr, _ = await self._resolve_ref_audio(cast(str, request.ref_audio)) prompt["ref_audio"] = (np.asarray(wav, dtype=np.float32), sr) if request.ref_text: prompt["ref_text"] = request.ref_text @@ -2151,6 +2301,7 @@ async def _create_diffusion_speech( audio_output, audio_key = self._extract_audio_output(final_output) if audio_key is None: raise ValueError("TTS model did not produce audio output.") + audio_output = cast(dict, audio_output) audio_tensor = audio_output[audio_key] sr_raw = audio_output.get("sr", 24000) @@ -2265,6 +2416,11 @@ async def create_speech( request_id = f"speech-{random_uuid()}" request_start_s = time.perf_counter() + request_arrival_ts = ( + float(getattr(raw_request.state, "request_timestamp", time.time())) + if raw_request is not None + else time.time() + ) if raw_request: raw_request.state.request_metadata = RequestResponseMetadata( request_id=request_id, @@ -2292,14 +2448,19 @@ async def create_speech( return error media_type = "audio/wav" if response_format == "wav" else "audio/pcm" - _, generator, raw_tts_params = await self._prepare_speech_generation(request, request_id=request_id) - return StreamingResponse( + _, generator, raw_tts_params = await self._prepare_speech_generation( + request, + request_id=request_id, + arrival_time=request_arrival_ts, + ) + return _SpeechStreamingResponse( self._generate_audio_chunks( generator, request_id, response_format, raw_request=raw_request, request_start_s=request_start_s, + request_arrival_ts=request_arrival_ts, tts_params=raw_tts_params, target_sample_rate=request.sample_rate, ), @@ -2314,14 +2475,19 @@ async def create_speech( if error is not None: return error - _, generator, sse_tts_params = await self._prepare_speech_generation(request, request_id=request_id) - return StreamingResponse( + _, generator, sse_tts_params = await self._prepare_speech_generation( + request, + request_id=request_id, + arrival_time=request_arrival_ts, + ) + return _SpeechStreamingResponse( self._generate_audio_sse_events( generator, request_id, response_format, raw_request=raw_request, request_start_s=request_start_s, + request_arrival_ts=request_arrival_ts, request=request, tts_params=sse_tts_params, ), @@ -2332,7 +2498,11 @@ async def create_speech( usage_box: list[SpeechTokenUsage] = [] try: audio_bytes, media_type = await self._generate_audio_bytes( - request, request_id=request_id, usage_out=usage_box, collect=collect + request, + request_id=request_id, + usage_out=usage_box, + collect=collect, + request_arrival_ts=request_arrival_ts, ) except TTSGenerationError as error: # An adapter can reject otherwise completed audio. Retry only @@ -2358,6 +2528,7 @@ async def create_speech( request_id=retry_request_id, usage_out=usage_box, collect=collect, + request_arrival_ts=request_arrival_ts, ) total_ms = (time.perf_counter() - request_start_s) * 1000.0 logger.info( diff --git a/vllm_omni/entrypoints/openai/serving_speech_stream.py b/vllm_omni/entrypoints/openai/serving_speech_stream.py index 32e9fa40d01..ceeb208c92d 100644 --- a/vllm_omni/entrypoints/openai/serving_speech_stream.py +++ b/vllm_omni/entrypoints/openai/serving_speech_stream.py @@ -55,6 +55,7 @@ import asyncio import base64 import json +import time from contextlib import aclosing from fastapi import WebSocket, WebSocketDisconnect @@ -284,6 +285,8 @@ async def _generate_and_send( ``utterance_index`` identifies the flush this sentence belongs to and ``sentence_index`` its position inside that flush. """ + request_arrival_ts = time.time() + request_start_s = time.perf_counter() response_format = config.response_format or "wav" # Reject unmet word-timestamps preconditions early with a clear reason. @@ -346,7 +349,10 @@ async def _generate_and_send( request_id = None try: if config.stream_audio: - request_id, generator, tts_params = await self._speech_service._prepare_speech_generation(request) + request_id, generator, tts_params = await self._speech_service._prepare_speech_generation( + request, + arrival_time=request_arrival_ts, + ) if config.word_timestamps: total_bytes = await self._stream_audio_with_alignments( websocket=websocket, @@ -356,6 +362,8 @@ async def _generate_and_send( utterance_index=utterance_index, sentence_index=sentence_index, language=config.language, + request_start_s=request_start_s, + request_arrival_ts=request_arrival_ts, tts_params=tts_params, ) else: @@ -363,6 +371,8 @@ async def _generate_and_send( self._speech_service._generate_pcm_chunks( generator, request_id, + request_start_s=request_start_s, + request_arrival_ts=request_arrival_ts, tts_params=tts_params, ) ) as stream: @@ -370,7 +380,10 @@ async def _generate_and_send( total_bytes += len(chunk) await websocket.send_bytes(chunk) else: - audio_bytes, _ = await self._speech_service._generate_audio_bytes(request) + audio_bytes, _ = await self._speech_service._generate_audio_bytes( + request, + request_arrival_ts=request_arrival_ts, + ) total_bytes = len(audio_bytes) await websocket.send_bytes(audio_bytes) except WebSocketDisconnect: @@ -416,6 +429,8 @@ async def _stream_audio_with_alignments( sentence_text: str, utterance_index: int, sentence_index: int, + request_start_s: float, + request_arrival_ts: float, language: str | None = None, tts_params: dict | None = None, ) -> int: @@ -463,6 +478,8 @@ async def send_chunk( self._speech_service._generate_pcm_chunks( generator, request_id, + request_start_s=request_start_s, + request_arrival_ts=request_arrival_ts, include_sample_rate=True, tts_params=tts_params, collect=collect, diff --git a/vllm_omni/metrics/definitions.py b/vllm_omni/metrics/definitions.py index 6cc2825c1ea..a09fbcc2ff2 100644 --- a/vllm_omni/metrics/definitions.py +++ b/vllm_omni/metrics/definitions.py @@ -160,6 +160,8 @@ AUDIO_UNDERRUN_S = METRIC_PREFIX + AUDIO_UNDERRUN + "_s" AUDIO_CONTINUITY_OK_METRIC = METRIC_PREFIX + AUDIO_CONTINUITY_OK AUDIO_SKIPPED_REQUESTS_METRIC = METRIC_PREFIX + AUDIO_SKIPPED_REQUESTS +SPEECH_STREAM_ABORTED_METRIC = METRIC_PREFIX + "speech_stream_aborted" +SPEECH_STREAM_COMPLETED_METRIC = METRIC_PREFIX + "speech_stream_completed" # Realtime Server VAD serving metrics. REALTIME_VAD_ACTIVE_SESSIONS = METRIC_PREFIX + "realtime_vad_active_sessions" diff --git a/vllm_omni/metrics/modality.py b/vllm_omni/metrics/modality.py index e0d51e2b0b4..bb7842e36ab 100644 --- a/vllm_omni/metrics/modality.py +++ b/vllm_omni/metrics/modality.py @@ -1,3 +1,6 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni project + """OmniModalityMetrics — per-modality Prometheus families (audio path only). 7 audio business-semantic metric families. Text-path metrics (TTFT / ITL / @@ -76,6 +79,16 @@ "Silent-loss counter — code2wav rejected malformed codec input and returned 200 OK with empty audio.", labelnames=list(defs.AUDIO_SKIPPED_LABELS), ) +_speech_stream_aborted_family = Counter( + defs.SPEECH_STREAM_ABORTED_METRIC, + "Speech audio generators terminated before normal completion, including before the first PCM payload.", + labelnames=["model_name", "reason"], +) +_speech_stream_completed_family = Counter( + defs.SPEECH_STREAM_COMPLETED_METRIC, + "Speech audio generators that completed normally; does not confirm client receipt.", + labelnames=["model_name"], +) # ---------------------------------------------------------------------------- @@ -153,6 +166,17 @@ def __init__(self, model_name: str, log_stats: bool = True) -> None: # ---- Audio ------------------------------------------------------------ + def inc_speech_stream_aborted(self, reason: str) -> None: + if not self._log_stats: + return + if reason not in {"cancelled", "closed", "engine_dead", "error"}: + reason = "error" + _speech_stream_aborted_family.labels(model_name=self._model_name, reason=reason).inc() + + def inc_speech_stream_completed(self) -> None: + if self._log_stats: + _speech_stream_completed_family.labels(model_name=self._model_name).inc() + def observe_audio_ttfp(self, stage: str, replica: str, ttfp_seconds: float) -> None: if not self._log_stats: return @@ -335,6 +359,7 @@ def observe_audio_streaming_finalize( chunk_arrival_times_s: list[float], chunk_bytes: list[int], sample_rate: int, + channels: int = defs.DEFAULT_AUDIO_CHANNELS, threshold_s: float = defs.AUDIO_CONTINUITY_DEFAULT_THRESHOLD_S, ) -> None: """Emit audio_underrun_s + audio_continuity_ok_total at request end. @@ -353,6 +378,7 @@ def observe_audio_streaming_finalize( chunk_arrival_times_s=chunk_arrival_times_s, chunk_bytes=chunk_bytes, sample_rate=sample_rate, + channels=channels, threshold_s=threshold_s, ) stage_label = str(stage_id)