Track the response model for WebSocket session metrics

This commit is contained in:
Gil Korzen 2026-08-14 16:26:35 +03:00
parent 2d88e31a40
commit d5d8d7cab8
2 changed files with 63 additions and 7 deletions

View file

@ -7160,7 +7160,6 @@ class OpenAIHandlerMixin:
if isinstance(first_response_body, dict)
else None
)
# Hot-fix follow-up to PR #406 — inline Rust compression on the
# WS first frame before forwarding upstream. PR #406 enabled
# the same call for HTTP /v1/responses; PR-C5's "WS-side
@ -8000,6 +7999,7 @@ class OpenAIHandlerMixin:
response_output_items.clear()
response_started_ms: float | None = None
completed_response_model = "unknown"
async def _record_ws_response_metrics() -> None:
"""Record one completed Responses turn on long-lived WS sessions."""
@ -8056,7 +8056,7 @@ class OpenAIHandlerMixin:
):
return
model_for_metrics = str(body.get("model") or "unknown")
model_for_metrics = completed_response_model
latency_ms = (
(time.perf_counter() * 1000.0 - response_started_ms)
if response_started_ms is not None
@ -8209,6 +8209,13 @@ class OpenAIHandlerMixin:
upstream_frame_index,
ws_last_upstream_frame_type,
)
response = event.get("response")
completed_response_model = (
str(response.get("model") or "unknown")
if isinstance(response, dict)
else "unknown"
)
if event_type == "response.created":
response_started_ms = time.perf_counter() * 1000.0
(
@ -8581,11 +8588,7 @@ class OpenAIHandlerMixin:
)
if not isinstance(ws_inner_for_telemetry, dict):
ws_inner_for_telemetry = {}
model_name = (
ws_inner_for_telemetry.get("model")
or (body.get("model") if isinstance(body, dict) else None)
or "unknown"
)
model_name = str(current_response_template.get("model") or "unknown")
_final_auth_mode = classify_auth_mode(ws_headers)
residual_input_tokens = max(0, ws_input_tokens_total - ws_recorded_input_tokens_total)
residual_output_tokens = max(

View file

@ -2243,3 +2243,56 @@ async def test_ws_memory_continuation_continues_pre_stream_and_passes_late_call(
assert second_response[6]["item"] == function_call_two
assert second_response[7]["response"]["id"] == "r-2"
assert executed == [("memory_search", {}, "user-1", "openai")]
@pytest.mark.asyncio
async def test_ws_session_metrics_track_model_per_response_create():
"""A model switch on one WS session must affect the next request outcome."""
upstream_events = [
json.dumps({"type": "response.created", "response": {"id": "r_1"}}),
json.dumps(
{
"type": "response.completed",
"response": {
"id": "r_1",
"model": "model-a",
"usage": {"input_tokens": 10, "output_tokens": 1},
},
}
),
json.dumps({"type": "response.created", "response": {"id": "r_2"}}),
json.dumps(
{
"type": "response.completed",
"response": {
"id": "r_2",
"model": "model-b",
"usage": {"input_tokens": 10, "output_tokens": 1},
},
}
),
]
first_frame = json.dumps(
{
"type": "response.create",
"response": {"model": "model-a", "input": "first turn"},
}
)
second_frame = json.dumps(
{
"type": "response.create",
"response": {"model": "model-b", "input": "second turn"},
}
)
upstream = _FakeUpstream(upstream_events)
fake_ws_mod = _make_fake_websockets_module(upstream)
client_ws = _FakeWebSocket(frames=[first_frame, second_frame])
handler = _DummyOpenAIHandler()
with patch.dict(sys.modules, {"websockets": fake_ws_mod}):
await handler.handle_openai_responses_ws(client_ws)
assert [request["model"] for request in handler.metrics.recorded_requests] == [
"model-a",
"model-b",
]