diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 8cbc28362a8..e0068188413 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -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, diff --git a/tests/integration/_support/upstream.py b/tests/integration/_support/upstream.py index e9c50ea7966..1b300ca8245 100644 --- a/tests/integration/_support/upstream.py +++ b/tests/integration/_support/upstream.py @@ -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), ] ) diff --git a/tests/integration/contracts.json b/tests/integration/contracts.json index d12e3ae4620..74b0cef6e31 100644 --- a/tests/integration/contracts.json +++ b/tests/integration/contracts.json @@ -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": { diff --git a/tests/integration/cost_calculation/cost_tracking_case.py b/tests/integration/cost_calculation/cost_tracking_case.py index ac3fcd33d2e..aa5529aec27 100644 --- a/tests/integration/cost_calculation/cost_tracking_case.py +++ b/tests/integration/cost_calculation/cost_tracking_case.py @@ -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"), ] diff --git a/tests/integration/spend/test_responses_websocket_spend.py b/tests/integration/spend/test_responses_websocket_spend.py new file mode 100644 index 00000000000..b0a6c7689c1 --- /dev/null +++ b/tests/integration/spend/test_responses_websocket_spend.py @@ -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