test(integration): native Responses WebSocket sessions record summed usage and spend (Pylon #7872)

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
kerry 2026-09-22 20:50:13 +00:00
parent c6c3881d7f
commit 5545510316
5 changed files with 161 additions and 6 deletions

View file

@ -6714,11 +6714,13 @@ class BaseLLMHTTPHandler:
ws_url = urlunparse(_parsed._replace(query=urlencode({k: v[0] for k, v in _qs.items()})))
try:
ssl_context = get_shared_realtime_ssl_context()
if ws_url.startswith("wss://") and ssl_context is False:
ssl_context = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT)
ssl_context.check_hostname = False
ssl_context.verify_mode = ssl.CERT_NONE
ssl_context = None
if ws_url.startswith("wss://"):
ssl_context = get_shared_realtime_ssl_context()
if ssl_context is False:
ssl_context = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT)
ssl_context.check_hostname = False
ssl_context.verify_mode = ssl.CERT_NONE
logging_obj.pre_call(
input=None,

View file

@ -24,6 +24,7 @@ from integration.cost_calculation.cost_tracking_case import (
EventStreamResponse,
JsonResponse,
RealtimeResponse,
ResponsesWebSocketResponse,
RoutedResponse,
SseResponse,
StoredResponse,
@ -261,6 +262,28 @@ class Provider:
)
await websocket.send_json(rendered)
async def responses_websocket(self, websocket: WebSocket) -> None:
scenario_id: Final = websocket.headers.get("authorization", "").removeprefix("Bearer ")
response: Final = self.scenario_store.get(scenario_id)
if not isinstance(response, ResponsesWebSocketResponse):
await websocket.close(code=4404)
return
await websocket.accept()
event_index: Final = iter(response.events)
async for message in websocket.iter_json():
payload: Final = JSON_OBJECT.validate_python(message)
if payload.get("type") != "response.create":
continue
event: Final = next(event_index, None)
if event is None:
continue
rendered: Final = JSON_OBJECT.validate_json(
json.dumps(event, separators=(",", ":"))
.replace("$REQUEST_ID", scenario_id)
.replace("$UNIQUE_ID", f"{scenario_id}-{uuid.uuid4().hex[:8]}")
)
await websocket.send_json(rendered)
@staticmethod
def _response(response: StoredResponse, scenario_id: str) -> Response:
unique_id: Final = f"{scenario_id}-{uuid.uuid4().hex[:8]}"
@ -323,6 +346,8 @@ class Provider:
_aws_event_frame(event.event_type, event.payload, scenario_id, unique_id) for event in events
)
return Response(content=event_body, media_type=response.content_type)
case RealtimeResponse() | ResponsesWebSocketResponse():
return JSONResponse({"error": "not an HTTP scenario"}, status_code=400)
def app(self) -> Starlette:
return Starlette(
@ -341,6 +366,7 @@ class Provider:
Route("/{path:path}", self.scripted, methods=["POST"]),
Route("/{path:path}", self.scripted, methods=["GET"]),
WebSocketRoute("/v1/realtime", self.realtime),
WebSocketRoute("/v1/responses", self.responses_websocket),
]
)

View file

@ -1764,6 +1764,9 @@
],
"tests/integration/mcp/test_mcp_lifecycle.py::test_same_url_server_grants_scope_discovery_and_direct_or_virtual_execution[bearer]": [
"other.mcp.permissions.same_url_servers_enforce_discovery_and_execution"
],
"tests/integration/spend/test_responses_websocket_spend.py::test_native_responses_websocket_session_records_summed_usage_and_spend": [
"spend.responses_websocket.native_session_usage_is_billed"
]
},
"browser": {

View file

@ -172,8 +172,21 @@ class RealtimeResponse(BaseModel):
session_model: str | None = None
class ResponsesWebSocketResponse(BaseModel):
model_config = ConfigDict(frozen=True, extra="forbid")
content_type: Literal["application/x-responses-websocket"]
events: tuple[dict[str, JsonValue], ...]
StoredResponse: TypeAlias = Annotated[
JsonResponse | SseResponse | EventStreamResponse | BinaryResponse | RoutedResponse | RealtimeResponse,
JsonResponse
| SseResponse
| EventStreamResponse
| BinaryResponse
| RoutedResponse
| RealtimeResponse
| ResponsesWebSocketResponse,
Field(discriminator="content_type"),
]

View file

@ -0,0 +1,111 @@
from __future__ import annotations
import asyncio
import json
import os
import uuid
from hashlib import sha256
from typing import Final
import pytest
import websockets
from integration._support.client import JSON_OBJECT, Gateway, eventually
from integration._support.database import read_rows
from integration._support.upstream import delete_scenario, register_scenario
from integration.cost_calculation.cost_tracking_case import ResponsesWebSocketResponse
from pydantic import JsonValue
_TERMINAL_TYPES: Final = frozenset({"response.completed", "response.incomplete"})
def _response_event(
event_type: str, status: str, text: str, input_tokens: int, output_tokens: int
) -> dict[str, JsonValue]:
return {
"type": event_type,
"response": {
"id": "resp_$UNIQUE_ID",
"object": "response",
"status": status,
**({"incomplete_details": {"reason": "max_output_tokens"}} if status == "incomplete" else {}),
"model": "gpt-4o-mini",
"output": [
{
"type": "message",
"id": "msg_$UNIQUE_ID",
"role": "assistant",
"status": "completed",
"content": [{"type": "output_text", "text": text, "annotations": []}],
}
],
"usage": {
"input_tokens": input_tokens,
"output_tokens": output_tokens,
"total_tokens": input_tokens + output_tokens,
"input_tokens_details": {"cached_tokens": 0},
"output_tokens_details": {"reasoning_tokens": 0},
},
},
}
async def _run_responses_websocket(url: str, key: str, model_name: str) -> tuple[dict[str, JsonValue], ...]:
ws_url: Final = f"{url.replace('http://', 'ws://').replace('https://', 'wss://')}/v1/responses?model={model_name}"
async with websockets.connect(ws_url, additional_headers={"Authorization": f"Bearer {key}"}) as websocket:
terminal_events: list[dict[str, JsonValue]] = []
for index in range(2):
await websocket.send(json.dumps({"type": "response.create", "input": f"turn {index}"}))
while True:
event: Final = JSON_OBJECT.validate_json(await websocket.recv())
if event.get("type") in _TERMINAL_TYPES:
terminal_events.append(event)
break
return tuple(terminal_events)
@pytest.mark.covers("spend.responses_websocket.native_session_usage_is_billed")
def test_native_responses_websocket_session_records_summed_usage_and_spend(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
scenario_id: Final = f"responses-ws-{uuid.uuid4().hex[:12]}"
completed_event: Final = _response_event("response.completed", "completed", "first turn", 10, 5)
incomplete_event: Final = _response_event("response.incomplete", "incomplete", "second turn", 7, 4)
handle: Final = register_scenario(
scenario_id,
ResponsesWebSocketResponse(
content_type="application/x-responses-websocket",
events=(completed_event, incomplete_event),
),
)
scenario.cleanups.callback(delete_scenario, handle)
model: Final = scenario.model(
input_cost_per_token=0.001,
output_cost_per_token=0.002,
api_key=scenario_id,
api_base=f"{gateway.upstream_url}/v1",
)
key: Final = scenario.key(models=[model])
terminal_events: Final = asyncio.run(
_run_responses_websocket(os.environ["INTEGRATION_PROXY_URL"].rstrip("/"), key, model)
)
totals: Final = tuple(
JSON_OBJECT.validate_python(JSON_OBJECT.validate_python(event["response"])["usage"])["total_tokens"]
for event in terminal_events
)
assert totals == (15, 11), terminal_events
rows: Final = eventually(
lambda: read_rows(
'SELECT prompt_tokens, completion_tokens, spend, call_type, status FROM "LiteLLM_SpendLogs" WHERE api_key=%s',
(sha256(key.encode()).hexdigest(),),
),
lambda values: len(values) == 1,
seconds=70,
)
assert rows == [
{
"prompt_tokens": 17,
"completion_tokens": 9,
"spend": pytest.approx(17 * 0.001 + 9 * 0.002),
"call_type": "_aresponses_websocket",
"status": "success",
}
], rows