mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
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:
parent
c6c3881d7f
commit
5545510316
5 changed files with 161 additions and 6 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
]
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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": {
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
]
|
||||
|
||||
|
|
|
|||
111
tests/integration/spend/test_responses_websocket_spend.py
Normal file
111
tests/integration/spend/test_responses_websocket_spend.py
Normal 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
|
||||
Loading…
Add table
Reference in a new issue