Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
89 changes: 89 additions & 0 deletions tests/entrypoints/test_async_omni.py
Original file line number Diff line number Diff line change
Expand Up @@ -329,6 +329,95 @@ async def empty_abort_async(request_ids):
asyncio.run(run())


@pytest.mark.cpu
def test_abort_drops_consumed_metric_message_ids_with_state():
"""Exercise the production leak path from #6462: a metric-bearing
``OutputMessage`` goes through ``_handle_output_message`` (which populates
the per-request de-dup set), the request is aborted (which, since #6367,
keeps ``request_states`` registered and enqueues a terminal abort output),
and the set must be released with the state by ``generate()``'s normal
cleanup — with no independent request-keyed metric state surviving on the
instance at any point. A reintroduced
class-level ``_consumed_metric_messages``-style map populated by the
production handler fails the sweep below.
"""
import time as _time
from collections.abc import Mapping

from vllm_omni.engine.messages import OutputMessage
from vllm_omni.entrypoints.client_request_state import ClientRequestState
from vllm_omni.metrics.stats import OrchestratorAggregator

async def run():
omni = get_async_omni_instance()

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

[P2] Exercise the production leak path in this regression

This helper uses object.__new__, manually seeds the new field, and never runs OmniBase.__init__ or a metric-bearing output through _handle_output_message. Reintroducing and populating the old request-keyed map in production could therefore still pass this test. Drive a metric message through the production handler, abort the request, and assert that no independent request-keyed metric state survives.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Done in d9d2b69 — the regression now drives a metric-bearing OutputMessage through _handle_output_message (asserting the handler records the de-dup entry on the state and de-duplicates a replay of the same message object), aborts via the public abort(), and then sweeps vars(omni) so no mapping on the instance still keys anything by the request id. A reintroduced request-keyed map populated by the production handler fails either the sweep or the hasattr guard.

omni.engine.get_stage_metadata = lambda stage_id: SimpleNamespace(
stage_type="llm",
final_output=True,
final_output_type="text",
)

rid = "req-1-cccc"
state = ClientRequestState(request_id=rid, external_request_id="req-1")
state.metrics = OrchestratorAggregator(
num_stages=1,
log_stats=False,
wall_start_ts=_time.time(),
final_stage_id_for_e2e=0,
)
state.input_stream_task = None
omni.request_states[rid] = state

msg = OutputMessage(
request_id=rid,
stage_id=0,
engine_outputs=SimpleNamespace(final_output_type="text"),
metrics=SimpleNamespace(
num_tokens_in=3,
num_tokens_out=5,
stage_gen_time_ms=1.0,
rx_transfer_bytes=0,
rx_decode_time_ms=0.0,
rx_in_flight_time_ms=0.0,
),
finished=False,
)
handled, out_rid, out_stage, out_state = omni._handle_output_message(msg)
assert handled is False and out_state is state
# The production handler recorded the de-dup entry on the state...
assert id(msg) in state.consumed_metric_message_ids
# ...and de-duplicates a replay of the same message object.
omni._handle_output_message(msg)
assert len(state.consumed_metric_message_ids) == 1

assert not hasattr(omni, "_consumed_metric_messages")

# #6367 contract: abort() keeps the state registered and enqueues a
# terminal abort output; generate()'s ``finally`` owns the cleanup.
await omni.abort("req-1")
assert rid in omni.request_states
terminal = state.queue.get_nowait()
assert terminal.finished is True
assert terminal.engine_outputs.outputs[0].finish_reason == "abort"
# The de-dup set still lives only on the (still-registered) state.
for name, value in vars(omni).items():
if name == "request_states":
continue
assert not (isinstance(value, Mapping) and rid in value), (
f"request-keyed metric state kept outside request_states in {name!r}"
)

# generate()'s normal cleanup releases the state and the set with it.
omni._log_summary_and_cleanup(rid)
assert rid not in omni.request_states
for name, value in vars(omni).items():
assert not (isinstance(value, Mapping) and rid in value), (
f"request-keyed state survived cleanup in {name!r}"
)
assert not hasattr(omni, "_consumed_metric_messages")

asyncio.run(run())


@pytest.mark.cpu
def test_generate_accepts_request_after_repeated_cancellations():
async def run_test():
Expand Down
3 changes: 1 addition & 2 deletions tests/metrics/test_emit_calls.py
Original file line number Diff line number Diff line change
Expand Up @@ -41,7 +41,6 @@ def _make_omni_base_with_mock_prom(mocker):
obj = object.__new__(OmniBase)
obj.prom_metrics = mocker.Mock(spec=OmniPrometheusMetrics)
obj.request_states = {}
obj._consumed_metric_messages = {}
obj.log_stats = True
return obj, obj.prom_metrics

Expand Down Expand Up @@ -270,7 +269,7 @@ def test_same_finished_image_message_is_observed_exactly_once(self, mocker, peak

obj = object.__new__(OmniBase)
obj._enable_ar_profiler = False
obj._consumed_metric_messages = {}
obj.request_states = {"req-replay": SimpleNamespace(consumed_metric_message_ids=set())}
obj.prom_metrics = mocker.Mock(spec=OmniPrometheusMetrics)
obj.mod_metrics = mocker.Mock()
obj.engine = SimpleNamespace(
Expand Down
2 changes: 0 additions & 2 deletions tests/metrics/test_prometheus.py
Original file line number Diff line number Diff line change
Expand Up @@ -154,7 +154,6 @@ def test_running_and_waiting_zero_after_request_completes(self, registry: Collec
obj.engine = SimpleNamespace(_running_counter=OmniRequestCounter())
obj.prom_metrics = OmniPrometheusMetrics(model_name="lifecycle-test")
obj.request_states = {}
obj._consumed_metric_messages = {}
obj.log_stats = False

# Simulate request lifecycle: start (counter 0→1, dict {} → {req}),
Expand Down Expand Up @@ -183,7 +182,6 @@ def test_gauges_reflect_remaining_requests_after_one_completes(self, registry: C
obj.engine = SimpleNamespace(_running_counter=OmniRequestCounter())
obj.prom_metrics = OmniPrometheusMetrics(model_name="lifecycle-test-2")
obj.request_states = {}
obj._consumed_metric_messages = {}
obj.log_stats = False

# Two in flight, one finalizes — running should report 1, waiting 0.
Expand Down
6 changes: 6 additions & 0 deletions vllm_omni/entrypoints/client_request_state.py
Original file line number Diff line number Diff line change
Expand Up @@ -39,3 +39,9 @@ def __init__(
# without re-querying stage_pools.
self.audio_emit_stage_id: int | None = None
self.audio_emit_replica_id: int | None = None
# De-dup set for metric messages: OmniBase populates this in
# ``_handle_output_message`` / ``_process_single_result`` so the same
# ``id(msg)`` isn't counted twice into per-request metrics. Kept on
# the request state (not a class-level dict) so it is released with
# the state — see #6462 / #6561.
self.consumed_metric_message_ids: set[int] = set()
18 changes: 4 additions & 14 deletions vllm_omni/entrypoints/omni_base.py
Original file line number Diff line number Diff line change
Expand Up @@ -231,7 +231,6 @@ def __init__(
self.async_chunk = bool(getattr(self.engine, "async_chunk", False))

self.request_states: dict[str, ClientRequestState] = {}
self._consumed_metric_messages: dict[str, set[int]] = {}
self.mod_metrics = OmniModalityMetrics(model_name=model, log_stats=log_stats)

self.default_sampling_params_list = self.engine.default_sampling_params_list
Expand Down Expand Up @@ -298,13 +297,6 @@ def _stage_has_no_live_replica(self, pool: StagePool) -> bool:
"""True when a non-empty stage pool has lost all of its replicas."""
return len(pool.clients) > 0 and self._live_replica_count(pool) == 0

def _consumed_metric_message_ids(self, request_id: str) -> set[int]:
consumed_by_request = getattr(self, "_consumed_metric_messages", None)
if consumed_by_request is None:
consumed_by_request = {}
self._consumed_metric_messages = consumed_by_request
return consumed_by_request.setdefault(request_id, set())

@property
def is_running(self) -> bool:
return self.engine.is_alive()
Expand Down Expand Up @@ -449,9 +441,6 @@ def _log_summary_and_cleanup(self, request_id: str, reason: str = "stage_error")
)
finally:
self.request_states.pop(request_id, None)
consumed_by_request = getattr(self, "_consumed_metric_messages", None)
if consumed_by_request is not None:
consumed_by_request.pop(request_id, None)
# Republish gauges so any stale value left by the per-stage
# publish in _process_single_result (which runs while the request
# is still in self.request_states) is corrected after the pop.
Expand Down Expand Up @@ -532,7 +521,7 @@ def _handle_output_message(
stage_meta = self.engine.get_stage_metadata(stage_id)
output_type = getattr(msg.engine_outputs, "final_output_type", stage_meta.final_output_type)
msg_id = id(msg)
consumed = self._consumed_metric_message_ids(req_id)
consumed = req_state.consumed_metric_message_ids
if msg_id not in consumed:
req_state.metrics.on_stage_metrics(stage_id, req_id, msg.metrics, output_type)
submit_ts = msg.stage_submit_ts
Expand Down Expand Up @@ -642,8 +631,9 @@ def _process_single_result(
output_type = getattr(engine_outputs, "final_output_type", stage_meta.final_output_type)
if finished and _m is not None:
msg_id = id(result)
consumed = self._consumed_metric_message_ids(req_id)
if msg_id not in consumed:
req_state = self.request_states.get(req_id)
consumed = req_state.consumed_metric_message_ids if req_state is not None else None
if consumed is not None and msg_id not in consumed:
metrics.accumulate_diffusion_metrics(stage_meta.stage_type, req_id, engine_outputs)
metrics.on_stage_metrics(stage_id, req_id, _m, output_type)
consumed.add(msg_id)
Expand Down
Loading