test: cover ChatGPT realtime and image routes

This commit is contained in:
jibanez-staticduo 2026-09-18 05:47:34 +02:00
parent ccefa07cea
commit fbace1f2c6
No known key found for this signature in database
15 changed files with 1497 additions and 2 deletions

View file

@ -4410,6 +4410,41 @@ def test_image_response_input_image_tokens_priced_at_image_rate(details_as_dict)
expected = 19 * 5e-6 + 512 * 8e-6 + 158 * 3e-5
assert cost is not None
assert round(cost, 12) == round(expected, 12)
@pytest.mark.parametrize("excess", ["text", "image"])
def test_image_response_cached_modality_counts_cannot_exceed_inputs(excess):
"""
A cached_tokens_details entry larger than the matching input modality count
would turn cache reads into negative savings; the calculation must reject
the inconsistent usage instead of pricing it.
"""
from unittest.mock import patch
from litellm.litellm_core_utils.llm_cost_calc.utils import (
calculate_image_response_cost_from_usage,
)
from litellm.types.utils import Usage
cached: dict = (
{"text_tokens": 11, "image_tokens": 0} if excess == "text" else {"text_tokens": 0, "image_tokens": 101}
)
image_response = ImageResponse(data=[ImageObject(b64_json="x")])
image_response.usage = Usage(
prompt_tokens=0,
completion_tokens=0,
total_tokens=212,
input_tokens=110,
input_tokens_details={"text_tokens": 10, "image_tokens": 100, "cached_tokens_details": cached},
output_tokens=102,
output_tokens_details={"image_tokens": 102, "text_tokens": 0},
)
with pytest.raises(ValueError, match="Image cached token counts exceed their input modality counts"):
calculate_image_response_cost_from_usage(
model="gpt-image-2",
image_response=image_response,
custom_llm_provider="openai",
)
GEMINI_DAY0_LAUNCH_PRICING = [
("gemini-3.6-flash", 7.5e-07, 3.75e-06, 7.5e-08),
("gemini/gemini-3.6-flash", 7.5e-07, 3.75e-06, 7.5e-08),

View file

@ -3425,6 +3425,22 @@ async def test_live_attachment_does_not_dispatch_duplicate_usage():
logger.dispatch_success_handlers.assert_not_called()
@pytest.mark.asyncio
async def test_log_messages_flush_awaits_dispatch_instead_of_enqueueing():
worker = MagicMock()
logger = MagicMock()
logger.model_call_details = {}
logger.dispatch_success_handlers = AsyncMock()
stream = RealTimeStreaming(MagicMock(), MagicMock(), logger, logging_worker=worker)
stream.store_message({"type": "session.created"})
await stream.log_messages(wait_for_dispatch=True)
logger.dispatch_success_handlers.assert_awaited_once_with(stream.messages, prefer_async_handlers=True)
worker.ensure_initialized_and_enqueue.assert_not_called()
assert logger.model_call_details[REALTIME_SESSION_SUCCESS_LOGGED_KEY] is True
@pytest.mark.asyncio
@pytest.mark.parametrize("account_usage", [False, True])
async def test_attachment_cleanup_runs_in_owning_context_only(account_usage):

View file

@ -55,3 +55,9 @@ def test_signaling_preserves_selected_model_for_sideband(extra_query):
assert request["query_params"] == {"model": "gpt-live-1-codex"}
assert request["extra_headers"] == {"x-gateway-route": "voice"}
assert request["extra_query"] == extra_query
def test_signaling_requires_chatgpt_routing_extension():
response = httpx.Response(201, headers={"Location": "/v1/realtime/calls/rtc_unrouted"})
with pytest.raises(ValueError, match="Direct call signaling requires a ChatGPT deployment"):
parse_call_response(response, "voice", "owner", 1000)

View file

@ -227,3 +227,57 @@ def test_image_routes_use_configured_gateway(monkeypatch, env_name, api_base, tm
expected + "/images/generations"
)
assert ChatGPTImageEditConfig().get_complete_url("gpt-image-2", api_base, {}) == expected + "/images/edits"
@pytest.mark.parametrize(
"reference",
[
b"GIF89a" + b"\x00" * 32,
("reference.gif", b"hello", "image/gif"),
],
ids=["detected-gif", "declared-gif"],
)
def test_edit_rejects_non_bitmap_reference_content_type(reference):
with pytest.raises(ValueError, match="Reference images must be PNG, JPEG, or WEBP"):
ChatGPTImageEditConfig().transform_image_edit_request(
"gpt-image-2", "edit", reference, {}, GenericLiteLLMParams(), {}
)
def test_edit_rejects_mask_before_any_provider_call():
with pytest.raises(ValueError, match="ChatGPT image editing does not support masks"):
ChatGPTImageEditConfig().transform_image_edit_request(
"gpt-image-2",
"edit",
"data:image/png;base64,aGVsbG8=",
{"mask": "data:image/png;base64,aGVsbG8="},
GenericLiteLLMParams(),
{},
)
def test_edit_rejects_image_and_images_together():
with pytest.raises(ValueError, match="Specify only one of image or images"):
ChatGPTImageEditConfig().transform_image_edit_request(
"gpt-image-2",
"edit",
"data:image/png;base64,aGVsbG8=",
{},
GenericLiteLLMParams(images=[{"image_url": "data:image/png;base64,aGVsbG8="}]),
{},
)
@pytest.mark.parametrize(
"images",
[
[],
["data:image/png;base64,aGVsbG8="] * 6,
],
ids=["zero", "six"],
)
def test_edit_enforces_one_to_five_reference_images(images):
with pytest.raises(ValueError, match="images must contain between 1 and 5 reference images"):
ChatGPTImageEditConfig().transform_image_edit_request(
"gpt-image-2", "edit", images, {}, GenericLiteLLMParams(), {}
)

View file

@ -192,3 +192,31 @@ async def test_live_websocket_does_not_redirect_credentials():
with pytest.raises(InvalidStatus) as failure:
await transport.connect("live/sessions")
assert failure.value.response.status_code == 307
@pytest.mark.asyncio
@pytest.mark.parametrize(
"api_base",
[
"ftp://gateway.example/v1",
"https://user:secret@gateway.example/v1",
"https://gateway.example/v1#fragment",
"not a url",
],
)
async def test_live_rejects_invalid_api_base_before_network(api_base):
requests: list = []
def respond(request):
requests.append(request)
return httpx.Response(200, json={})
async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as client:
transport = LiveTransport(
LiveDeployment("deployment-model", provider="openai", api_key="deployment-key", api_base=api_base),
{},
http_client=client,
)
with pytest.raises(ValueError, match="Invalid Live API base"):
await transport.request("POST", "live/sessions", {})
assert requests == []

View file

@ -475,3 +475,22 @@ async def test_supervisor_connection_preserves_call_routing(model, chatgpt_token
assert url.params["gateway_token"] == "a+b&c"
assert connect.call_args.kwargs["additional_headers"]["x-gateway-token"] == "configured"
assert url.path.endswith("/rtc_owner") if model == "gpt-live-1-codex" else url.params["call_id"] == "rtc_owner"
@pytest.mark.asyncio
@pytest.mark.parametrize("model", ["gpt-live-1-codex", "gpt-realtime-1.5"])
async def test_close_call_prefers_session_close_only_for_live_models(model, chatgpt_tokens):
handler = ChatGPTRealtime(
GenericLiteLLMParams(chatgpt_realtime_call_id="rtc_close", chatgpt_token_dir=chatgpt_tokens),
{},
{},
)
connection = SimpleNamespace(send=AsyncMock())
handler.hangup_call = AsyncMock()
await handler.close_call(connection, model, "https://gateway.example/v1")
if model == "gpt-live-1-codex":
connection.send.assert_awaited_once_with('{"type":"session.close"}')
handler.hangup_call.assert_not_awaited()
else:
connection.send.assert_not_awaited()
handler.hangup_call.assert_awaited_once_with("https://gateway.example/v1")

View file

@ -25,6 +25,7 @@ from litellm.llms.base_llm.search.transformation import BaseSearchConfig, Search
from litellm.llms.bedrock.base_aws_llm import SignsRequestsWithAWS
from litellm.llms.brave.search.transformation import BraveSearchConfig
from litellm.llms.base_llm.image_edit.transformation import BaseImageEditConfig
from litellm.llms.base_llm.image_generation.transformation import BaseImageGenerationConfig
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.llms.custom_httpx.llm_http_handler import (
BaseLLMHTTPHandler,
@ -3772,3 +3773,63 @@ async def test_realtime_http_sessions_preserve_provider_identity(
else:
assert requests[0].headers.get_list("authorization")[-1] == "Bearer test-override"
assert requests[0].headers["chatgpt-account-id"] == "test-other-account"
class _ImageGenerationRecordingConfig(BaseImageGenerationConfig):
def get_supported_openai_params(self, model):
return ["size"]
def map_openai_params(self, non_default_params, optional_params, model, drop_params):
optional_params.update(non_default_params)
return optional_params
def validate_environment(self, headers, model, messages, optional_params, litellm_params, api_key=None, api_base=None):
return {"authorization": f"Bearer {api_key}"}
def get_complete_url(self, api_base, api_key, model, optional_params, litellm_params, stream=None):
return "https://images.example/v1/generations"
def transform_image_generation_request(self, model, prompt, optional_params, litellm_params, headers):
return {"model": model, "prompt": prompt}
def transform_image_generation_response(self, model, raw_response, model_response, logging_obj, request_data, optional_params, litellm_params, encoding=None, api_key=None, json_mode=None):
return ImageResponse(data=[ImageObject(b64_json=raw_response.json()["created"])])
def test_image_extra_headers_strips_oauth_identity_only_for_chatgpt():
headers: Final = {"authorization": "Bearer oauth", "chatgpt-account-id": "acct-1", "x-router": "keep"}
assert BaseLLMHTTPHandler._image_extra_headers("openai", headers) is headers
stripped: Final = BaseLLMHTTPHandler._image_extra_headers("chatgpt", headers)
assert dict(stripped) == {"x-router": "keep"}
@pytest.mark.asyncio
async def test_async_image_generation_handler_merges_extra_headers_for_non_chatgpt():
requests: Final = []
def respond(request):
requests.append(request)
return httpx.Response(200, json={"created": "ok"})
client: Final = AsyncHTTPHandler()
client.client = httpx.AsyncClient(transport=httpx.MockTransport(respond))
try:
response: Final = await BaseLLMHTTPHandler().async_image_generation_handler(
model="image-model",
prompt="a red circle",
image_generation_provider_config=_ImageGenerationRecordingConfig(),
image_generation_optional_request_params={},
custom_llm_provider="openai",
litellm_params={"api_key": "sk-image"},
logging_obj=Mock(),
timeout=10,
extra_headers={"x-router-header": "routed"},
api_key="sk-image",
client=client,
)
finally:
await client.client.aclose()
assert requests[0].headers["x-router-header"] == "routed"
assert requests[0].headers["authorization"] == "Bearer sk-image"
assert requests[0].url == "https://images.example/v1/generations"
assert response.data[0].b64_json == "ok"

View file

@ -26,6 +26,15 @@ def test_openai_realtime_handler_url_construction(api_base):
assert "model=gpt-4o-realtime-preview-2024-10-01" in url
def test_openai_realtime_handler_requires_api_key():
from litellm.llms.openai.realtime.handler import OpenAIRealtime
handler = OpenAIRealtime()
with pytest.raises(ValueError, match="api_key is required for OpenAI realtime calls"):
handler._resolve_api_key(None)
assert handler._resolve_api_key("sk-realtime-key") == "sk-realtime-key"
def test_openai_realtime_handler_url_with_extra_params():
from litellm.llms.openai.realtime.handler import OpenAIRealtime
from litellm.types.realtime import RealtimeQueryParams

View file

@ -364,7 +364,7 @@ async def test_custom_auth_does_not_enforce_key_model_access_by_default():
async def test_post_custom_auth_expired_key_returns_unauthorized():
expired_token = UserAPIKeyAuth(
token="test_token",
expires=datetime.now() - timedelta(minutes=1),
expires=datetime.now(timezone.utc) - timedelta(minutes=1),
)
with pytest.raises(ProxyException) as exc_info:
@ -7755,3 +7755,65 @@ async def test_claude_view_never_reinterprets_explicit_names(monkeypatch, layer)
assert data["model"] == ("foo" if layer == "unclaimed" else encoded)
await _normalize_claude_model(data, token, request, "/v1/messages")
assert data["model"] == ("foo" if layer == "unclaimed" else encoded)
def _malformed_authorization_websocket(send):
from unittest.mock import AsyncMock
from fastapi import WebSocket
return WebSocket(
{
"type": "websocket", "scheme": "ws", "server": ("localhost", 4000),
"path": "/v1/realtime", "query_string": b"",
"headers": [(b"authorization", b"Token malformed")],
},
AsyncMock(),
send,
)
@pytest.mark.parametrize("authorization_value", ["Token malformed", "bearer lowercase"])
def test_get_websocket_api_key_rejects_malformed_authorization(monkeypatch, authorization_value):
import importlib
from unittest.mock import AsyncMock
from fastapi import HTTPException, WebSocket
from litellm.proxy import proxy_server
auth_module = importlib.import_module("litellm.proxy.auth.user_api_key_auth")
monkeypatch.setattr(proxy_server, "general_settings", {})
websocket = WebSocket(
{
"type": "websocket", "scheme": "ws", "server": ("localhost", 4000),
"path": "/v1/realtime", "query_string": b"",
"headers": [(b"authorization", authorization_value.encode())],
},
AsyncMock(),
AsyncMock(),
)
with pytest.raises(HTTPException) as error:
auth_module.get_websocket_api_key(websocket)
assert error.value.status_code == 403
assert error.value.detail == "Invalid Authorization header format"
@pytest.mark.asyncio
async def test_websocket_auth_closes_policy_violation_on_malformed_authorization(monkeypatch):
import importlib
from unittest.mock import AsyncMock
from fastapi import HTTPException
from litellm.proxy import proxy_server
auth_module = importlib.import_module("litellm.proxy.auth.user_api_key_auth")
monkeypatch.setattr(proxy_server, "general_settings", {})
send = AsyncMock()
websocket = _malformed_authorization_websocket(send)
with pytest.raises(HTTPException) as error:
await auth_module.user_api_key_auth_websocket(websocket)
assert error.value.status_code == 403
assert error.value.detail == "Invalid Authorization header format"
send.assert_awaited_once_with({"type": "websocket.close", "code": 1008, "reason": ""})

View file

@ -178,3 +178,51 @@ async def test_async_log_success_event_counts_non_chat_response_tokens(response_
litellm_parent_otel_span=None,
)
assert current["current_tpm"] == 50, f"expected 50 tokens counted for {scope_id}, got {current['current_tpm']}"
@pytest.mark.asyncio
async def test_realtime_attachment_release_without_receipt_never_touches_counters():
dual_cache = MagicMock()
handler = _PROXY_MaxParallelRequestsHandler(InternalUsageCache(dual_cache))
auth = UserAPIKeyAuth(api_key="no-receipt")
await handler.async_release_realtime_attachment({}, auth)
await handler.async_release_realtime_attachment(
{"_legacy_realtime_attachment_reservations": {"cache_keys": [], "global_acquired": True}}, auth
)
# A release without a matching begin (or with a foreign receipt shape) must not decrement anything.
assert dual_cache.mock_calls == []
@pytest.mark.asyncio
async def test_failure_event_skips_realtime_observer_without_decrementing_slots():
from datetime import datetime
from litellm.proxy._types import InternalRequestOrigin
def failure_kwargs() -> dict:
return {
"litellm_params": {"metadata": {"user_api_key": "observer-hash", "global_max_parallel_requests": 5}},
"exception": RuntimeError("backend disconnected"),
}
dual_cache = MagicMock()
dual_cache.async_get_cache = AsyncMock(return_value=None)
dual_cache.async_increment_cache = AsyncMock()
dual_cache.async_batch_set_cache = AsyncMock()
handler = _PROXY_MaxParallelRequestsHandler(InternalUsageCache(dual_cache))
start = datetime.now()
end = datetime.now()
kwargs = failure_kwargs()
kwargs["internal_request_origin"] = InternalRequestOrigin.REALTIME_OBSERVER
await handler.async_log_failure_event(kwargs, None, start, end)
# The observer-internal failure mirror must leave the client-facing slot untouched.
assert dual_cache.mock_calls == []
dual_cache.mock_calls.clear()
await handler.async_log_failure_event(failure_kwargs(), None, start, end)
assert dual_cache.async_increment_cache.await_count >= 1
assert any(
call.kwargs.get("key") == "global_max_parallel_requests" and call.kwargs.get("value") == -1
for call in dual_cache.async_increment_cache.await_args_list
)

View file

@ -6105,3 +6105,49 @@ async def test_an_open_circuit_breaker_reads_the_sliding_window_locally_without_
assert isinstance(values, list)
assert [record.getMessage() for record in caplog.records if record.levelno >= logging.WARNING] == []
assert any("circuit breaker is open" in record.getMessage() for record in caplog.records)
@pytest.mark.asyncio
async def test_cluster_gauge_guards_fail_fast_when_scripts_are_unavailable(monkeypatch):
# _check_parallel_request_gauges only enters the cluster path with an acquire script in hand, so the
# guards below run only for direct cluster callers or script resets racing an in-flight batch.
handler, transport, gauges = _cluster_parallel_fixture(monkeypatch)
monkeypatch.setattr(handler, "parallel_count_script", None)
with pytest.raises(RuntimeError, match="Redis cluster parallel count script is unavailable"):
await handler._check_cluster_parallel_gauges(gauges, "reader", None, read_only=True)
monkeypatch.setattr(handler, "parallel_count_script", transport.script("count"))
monkeypatch.setattr(handler, "parallel_acquire_script", None)
with pytest.raises(RuntimeError, match="Redis cluster parallel acquire script is unavailable"):
await handler._check_cluster_parallel_gauges(gauges, "owner", None, read_only=False)
assert transport.calls == []
@pytest.mark.asyncio
async def test_cluster_release_guards_when_release_script_unavailable(monkeypatch):
handler, transport, gauges = _cluster_parallel_fixture(monkeypatch)
keys = tuple(gauge["counter_key"] for gauge in gauges)
monkeypatch.setattr(handler, "parallel_release_script", None)
with pytest.raises(RuntimeError, match="Redis cluster parallel release script is unavailable"):
await handler._release_cluster_parallel_slots(keys, "owner", None)
# Every shard group is still attempted; the first shard's error is the one re-raised.
assert [operation for operation, _ in transport.calls] == []
@pytest.mark.asyncio
async def test_cluster_rollback_swallows_release_failure_and_logs(monkeypatch, caplog):
handler, transport, gauges = _cluster_parallel_fixture(monkeypatch)
keys = tuple(gauge["counter_key"] for gauge in gauges)
assert (await handler._check_parallel_request_gauges(gauges, "owner"))["overall_code"] == "OK"
transport.fail = ("release", keys[0])
await handler._rollback_cluster_parallel_slots(keys, "owner", None)
# The admission error must not be replaced by an unreachable compensation shard: the rollback
# exception is retrieved, reported once, and swallowed so the caller keeps its original failure.
assert "Could not roll back all Redis cluster parallel request slots" in caplog.text
released = {key for operation, group in transport.calls if operation == "release" for key in group}
assert released == set(keys)
transport.fail = None
assert "owner" not in transport.members[keys[1]]
# Note: _renew_realtime_call_slot's in-memory "return False" isinstance guard after the any() scan is
# unreachable for any real cache state (the scan rejects non-dict values first), so no test drives it.

View file

@ -894,7 +894,11 @@ async def test_supervisor_policy_failure_hangs_up_before_releasing(monkeypatch):
@pytest.mark.asyncio
@pytest.mark.parametrize("hangup_fails", [False, True])
async def test_supervisor_constructor_failure_closes_effective_connection(monkeypatch, hangup_fails, caplog):
@pytest.mark.parametrize("invalidate_fails", [False, True])
@pytest.mark.parametrize("close_fails", [False, True])
async def test_supervisor_constructor_failure_closes_effective_connection(
monkeypatch, hangup_fails, invalidate_fails, close_fails, caplog
):
from unittest.mock import AsyncMock, MagicMock
from fastapi import Request
@ -906,8 +910,12 @@ async def test_supervisor_constructor_failure_closes_effective_connection(monkey
logger = MagicMock()
logger.litellm_params = {}
connection = AsyncMock()
if close_fails:
connection.close = AsyncMock(side_effect=RuntimeError("socket cleanup secret"))
handlers = []
invalidate = AsyncMock()
if invalidate_fails:
invalidate.side_effect = RuntimeError("counter cleanup secret")
release = AsyncMock()
monkeypatch.setattr(codex, "invalidate_budget_reservation_counters", invalidate, raising=False)
monkeypatch.setattr(codex, "release_or_invalidate_budget_reservation", release)
@ -946,6 +954,14 @@ async def test_supervisor_constructor_failure_closes_effective_connection(monkey
release.assert_awaited_once_with(budget_reservation=auth.budget_reservation)
invalidate.assert_not_awaited()
assert "private-cleanup-credential" not in caplog.text
if hangup_fails and invalidate_fails:
assert "Realtime startup cleanup could not invalidate budget counters" in caplog.text
elif invalidate_fails:
assert "Realtime startup cleanup could not invalidate budget counters" not in caplog.text
if close_fails:
assert "Realtime startup cleanup could not close observer socket" in caplog.text
assert "socket cleanup secret" not in caplog.text
assert "counter cleanup secret" not in caplog.text
@pytest.mark.asyncio
@ -1312,3 +1328,325 @@ async def test_signaling_settles_tokens_once_with_isolated_sdk_callbacks(monkeyp
await asyncio.gather(*callbacks)
assert await counter("tokens") == 0
assert limiter._gauge_in_flight_from_cache_value(await counter("max_parallel_requests")) == 0
@pytest.mark.asyncio
async def test_observer_startup_owns_call_lifecycle_with_synthetic_sockets(monkeypatch):
import asyncio
import json
from unittest.mock import AsyncMock, MagicMock
from fastapi import Request
from starlette.websockets import WebSocketState
from litellm.proxy.realtime_endpoints import call_supervision
call = CodexRealtimeCall(
call_id="rtc_open",
model="gpt-live-1-codex",
alias="voice",
owner="owner",
api_base="https://voice.example/codex",
expires_at=time.time() + 60,
)
auth = UserAPIKeyAuth()
logger = MagicMock()
logger.litellm_params = {}
observer: dict = {}
async def process(request, data, _auth, _model, route_type, *, internal_realtime_observer=False):
observer["request"] = request
observer["data"] = data
observer["route_type"] = route_type
observer["internal_realtime_observer"] = internal_realtime_observer
return {"extra_headers": {}}, logger
monkeypatch.setattr(codex, "process_codex_request", process)
class Connection:
def __init__(self):
self.messages = asyncio.Queue()
self.close = AsyncMock()
def __aiter__(self):
return self
async def __anext__(self):
message = await self.messages.get()
if message is None:
raise StopAsyncIteration
return message
connection = Connection()
await connection.messages.put(json.dumps({"type": "session.started"}))
stream_instance = MagicMock()
stream_instance.log_messages = AsyncMock()
terminations: list = []
class Handler:
def __init__(self, params, headers, extra_headers):
pass
@staticmethod
def get_api_base(base):
return "https://gateway.test/v1"
async def open_call_connection(self, model, base):
return connection
async def close_call(self, opened, model, base):
terminations.append(("close", opened, model, base))
async def hangup_call(self, base):
terminations.append(("hangup", base))
monkeypatch.setattr(codex, "ChatGPTRealtime", Handler)
stream = MagicMock(return_value=stream_instance)
monkeypatch.setattr(codex, "RealTimeStreaming", stream)
started: list = []
original_start = call_supervision.CALL_SUPERVISORS.start
async def capture_start(supervisor):
started.append(supervisor)
await original_start(supervisor)
monkeypatch.setattr(call_supervision.CALL_SUPERVISORS, "start", capture_start)
request = Request({"type": "http", "headers": [], "method": "POST", "path": "/v1/realtime/calls"})
await codex.supervise_codex_call(request, call, auth)
# The observer request serves the synthetic aliased-model body through the ASGI receive closure.
assert await observer["request"].json() == {"model": "voice"}
assert observer["data"]["model"] == "voice"
assert observer["route_type"] == "_arealtime"
assert observer["internal_realtime_observer"] is True
# The synthetic frontend completes the raw ASGI handshake through the swallow-and-return send closure.
frontend = stream.call_args.args[0]
await frontend.send({"type": "websocket.accept"})
assert frontend.application_state is WebSocketState.CONNECTED
# The supervisor owns the opened connection: the startup socket stack released it without closing it.
assert len(started) == 1
supervisor = started[0]
assert isinstance(supervisor, call_supervision.CallSupervisor)
assert supervisor._lease is None
assert supervisor._terminal_usage_required is (codex.realtime_endpoint(call.model) == "live")
connection.close.assert_not_awaited()
await supervisor._close_call()
assert terminations == [("close", connection, "gpt-live-1-codex", "https://gateway.test/v1")]
await supervisor._force_close_call()
assert terminations[-1] == ("hangup", "https://gateway.test/v1")
await connection.messages.put(
json.dumps({"type": "session.closed", "usage": {"total_tokens": 1}})
)
await supervisor.wait()
await call_supervision.CALL_SUPERVISORS.shutdown()
connection.close.assert_awaited_once()
@pytest.mark.asyncio
async def test_mixed_case_multipart_without_boundary_returns_400():
from fastapi import Request
from starlette.formparsers import MultiPartException
async def receive():
return {"type": "http.request", "body": b"sdp payload", "more_body": False}
request = Request(
{"type": "http", "headers": [(b"content-type", b"Multipart/Form-Data; charset=utf-8")]},
receive,
)
assert not await request.form()
with pytest.raises(HTTPException) as error:
await codex.read_codex_offer(request)
assert error.value.status_code == 400
assert error.value.detail == "Invalid realtime multipart offer"
assert isinstance(error.value.__cause__, MultiPartException)
@pytest.mark.asyncio
@pytest.mark.parametrize("observer", [False, True])
async def test_observer_processing_stamps_internal_request_origin(monkeypatch, observer):
from types import SimpleNamespace
from fastapi import Request
from litellm.proxy import common_request_processing
from litellm.proxy._types import InternalRequestOrigin
from litellm.proxy.realtime_endpoints.call_sessions import process_codex_request
recorded: dict = {}
class PassthroughProcessor:
def __init__(self, data):
self.data = data
async def common_processing_pre_call_logic(self, **kwargs):
recorded["observer"] = kwargs.get("internal_realtime_observer", False)
return {**self.data, "extra_headers": {}}, SimpleNamespace(model_call_details={})
monkeypatch.setattr(common_request_processing, "ProxyBaseLLMRequestProcessing", PassthroughProcessor)
request = Request({"type": "http", "method": "POST", "path": "/v1/realtime", "headers": []})
processed, logging_obj = await process_codex_request(
request,
{"model": "voice"},
UserAPIKeyAuth(),
"voice",
"_arealtime",
internal_realtime_observer=observer,
)
assert recorded["observer"] is observer
assert processed["model"] == "voice"
if observer:
assert logging_obj.model_call_details["internal_request_origin"] is InternalRequestOrigin.REALTIME_OBSERVER
else:
assert "internal_request_origin" not in logging_obj.model_call_details
@pytest.mark.asyncio
@pytest.mark.parametrize("mode", ["upstream_exception", "invalid_response", "upstream_error_status", "unroutable"])
async def test_signaling_response_paths_map_upstream_outcomes_to_http(monkeypatch, mode):
import json
from unittest.mock import AsyncMock
import httpx
from fastapi import Request
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.proxy import common_request_processing, proxy_server
monkeypatch.setenv("LITELLM_SALT_KEY", "test-only-salt-for-codex-realtime")
body = json.dumps({"sdp": "v=0\r\n", "session": {"model": "voice-alias"}}).encode()
async def receive():
return {"type": "http.request", "body": body, "more_body": False}
request = Request(
{
"type": "http",
"method": "POST",
"path": "/v1/realtime/calls",
"scheme": "http",
"server": ("localhost", 80),
"query_string": b"",
"headers": [(b"content-type", b"application/json"), (b"authorization", b"Bearer owner")],
},
receive,
)
monkeypatch.setattr(proxy_server, "master_key", "owner")
monkeypatch.setattr(proxy_server, "general_settings", {})
monkeypatch.setattr(codex, "can_key_call_resolved_model", AsyncMock())
class PassthroughProcessor:
def __init__(self, data):
self.data = data
async def common_processing_pre_call_logic(self, **kwargs):
return self.data, None
monkeypatch.setattr(common_request_processing, "ProxyBaseLLMRequestProcessing", PassthroughProcessor)
supervise = AsyncMock()
monkeypatch.setattr(codex, "supervise_codex_call", supervise)
def route_returning(value):
async def route(**kwargs):
async def respond():
return value
return respond()
return route
def route_raising(exc):
async def route(**kwargs):
async def boom():
raise exc
return boom()
return route
if mode == "upstream_exception":
monkeypatch.setattr(proxy_server, "route_request", route_raising(BaseLLMException(429, "provider saturated")))
with pytest.raises(HTTPException) as error:
await codex.create_codex_realtime_call(request)
assert error.value.status_code == 429
assert "provider saturated" in str(error.value.detail)
elif mode == "invalid_response":
monkeypatch.setattr(proxy_server, "route_request", route_returning("not-an-http-response"))
with pytest.raises(HTTPException) as error:
await codex.create_codex_realtime_call(request)
assert error.value.status_code == 502
assert error.value.detail == "Invalid realtime signaling response"
elif mode == "upstream_error_status":
monkeypatch.setattr(
proxy_server, "route_request", route_returning(httpx.Response(422, content=b'{"error":"invalid sdp"}'))
)
response = await codex.create_codex_realtime_call(request)
assert response.status_code == 422
assert response.body == b'{"error":"invalid sdp"}'
assert response.media_type == "application/json"
else:
monkeypatch.setattr(proxy_server, "route_request", route_returning(httpx.Response(201, content=b"v=0\r\n")))
with pytest.raises(HTTPException) as error:
await codex.create_codex_realtime_call(request)
assert error.value.status_code == 400
assert "ChatGPT deployment" in error.value.detail
supervise.assert_not_awaited()
@pytest.mark.asyncio
async def test_sideband_begins_realtime_attachment_on_legacy_limiter(monkeypatch):
from unittest.mock import AsyncMock, MagicMock
import litellm
from litellm.caching.caching import DualCache
from litellm.proxy import proxy_server as server
from litellm.proxy.hooks.parallel_request_limiter import (
_PROXY_MaxParallelRequestsHandler,
_RealtimeAttachmentReservations,
)
from litellm.proxy.utils import InternalUsageCache
monkeypatch.setenv("LITELLM_SALT_KEY", "test-only-salt-for-codex-realtime")
call = CodexRealtimeCall(
call_id="rtc_test",
model="gpt-live-1-codex",
alias="voice",
usage_supervised=True,
owner=hashlib.sha256(b"Bearer owner").hexdigest(),
expires_at=time.time() + 300,
)
auth = UserAPIKeyAuth()
logger = SimpleNamespace(model_call_details={})
limiter = _PROXY_MaxParallelRequestsHandler(InternalUsageCache(DualCache()))
proxy_logging = MagicMock()
proxy_logging.get_proxy_hook.return_value = limiter
monkeypatch.setattr(server, "proxy_logging_obj", proxy_logging)
monkeypatch.setattr(codex, "can_key_call_resolved_model", AsyncMock())
captured: dict = {}
async def process(request, data, *_args, **_kwargs):
receipt = data.get("_legacy_realtime_attachment_reservations")
assert isinstance(receipt, _RealtimeAttachmentReservations)
assert receipt.cache_keys == () and receipt.global_acquired is False
captured["data"] = data
return {}, logger
monkeypatch.setattr(codex, "process_codex_request", process)
monkeypatch.setattr(litellm, "_arealtime", AsyncMock())
websocket = WebSocket(
{
"type": "websocket",
"path": "/v1/live/opaque",
"query_string": b"",
"headers": [(b"authorization", b"Bearer owner")],
},
AsyncMock(return_value={"type": "websocket.connect"}),
AsyncMock(),
)
await codex.codex_realtime_sideband(websocket, encode_call(call), auth)
# The release consumed the receipt opened by begin_realtime_attachment before pre-call processing.
assert captured["data"]["_legacy_realtime_attachment_reservations"].take() == ((), False)

View file

@ -751,3 +751,108 @@ async def test_live_missing_backend_accounting_invalidates_budget_after_dispatch
assert logger.model_call_details["realtime_accounting_incomplete"] is True
assert "realtime_usage_incomplete" not in logger.model_call_details
invalidate.assert_awaited_once()
@pytest.mark.asyncio
async def test_repeated_start_is_rejected_while_observer_is_running():
socket, sink, close_call, supervisor = fixture()
await socket.messages.put({"type": "session.started"})
await supervisor.start()
with pytest.raises(RuntimeError, match="Call observer already started"):
await supervisor.start()
# The rejected second start leaves the running observer untouched.
await socket.messages.put({"type": "session.closed", "usage": {"total_tokens": 7}})
await supervisor.wait()
await supervisor.close()
close_call.assert_not_awaited()
assert sink.logs == 1
assert socket.closed
@pytest.mark.asyncio
async def test_start_rejects_quota_reservation_lost_during_startup():
from litellm.proxy.hooks.realtime_call_lease import RealtimeCallLease
socket = Socket()
logger = MagicMock(spec=Logging)
logger.model_call_details = {}
sink = Sink(logger)
async def hangup():
await socket.messages.put({"type": "session.closed", "usage": {"total_tokens": 42}})
close_call = AsyncMock(side_effect=hangup)
renew = AsyncMock(return_value=False)
release = AsyncMock()
lease = RealtimeCallLease(renew=renew, release=release, interval=3600)
lease.start()
supervisor = CallSupervisor(
socket,
sink,
logger,
UserAPIKeyAuth(),
close_call,
ready_timeout=10,
lifetime=10,
drain_timeout=0.05,
lease=lease,
)
await socket.messages.put({"type": "session.started"})
with pytest.raises(RuntimeError, match="lost its quota reservation during startup"):
await supervisor.start()
assert socket.closed
close_call.assert_awaited_once()
assert renew.await_count >= 1
release.assert_awaited_once()
assert sink.logs == 1
@pytest.mark.asyncio
async def test_registry_watch_logs_observer_accounting_failure_without_payload(caplog, monkeypatch):
import logging
from litellm.proxy.realtime_endpoints import call_supervision
caplog.set_level(logging.ERROR, logger="LiteLLM Proxy")
class FailingSink(Sink):
async def log_messages(self, *, wait_for_dispatch=False):
raise RuntimeError("observer accounting secret-token")
socket = Socket()
logger = MagicMock(spec=Logging)
logger.model_call_details = {}
sink = FailingSink(logger)
async def hangup():
await socket.messages.put({"type": "session.closed", "usage": {"total_tokens": 42}})
close_call = AsyncMock(side_effect=hangup)
invalidate = AsyncMock()
monkeypatch.setattr(call_supervision, "invalidate_budget_reservation_counters", invalidate)
supervisor = CallSupervisor(
socket,
sink,
logger,
UserAPIKeyAuth(),
close_call,
ready_timeout=5,
lifetime=5,
drain_timeout=0.05,
)
registry = CallSupervisors()
await socket.messages.put({"type": "session.created"})
await registry.start(supervisor)
await socket.messages.put({"type": "session.closed", "usage": {"total_tokens": 42}})
with pytest.raises(RuntimeError, match="observer accounting secret-token"):
await supervisor.wait()
await registry.shutdown()
assert registry._calls == ()
assert registry._tasks == ()
assert socket.closed
close_call.assert_not_awaited()
invalidate.assert_awaited_once_with(budget_reservation=None)
assert logger.model_call_details["realtime_accounting_incomplete"] is True
proxy_logs = [record.getMessage() for record in caplog.records if record.name == "LiteLLM Proxy"]
assert any("Realtime observer accounting failed" in message for message in proxy_logs)
assert not any("secret-token" in message for message in proxy_logs)

View file

@ -1412,3 +1412,647 @@ async def test_managed_budget_fails_closed_when_member_scope_lookup_fails_after_
)
assert rejected.value.status_code == 503
def test_json_conversion_rejects_unsupported_parent_container(monkeypatch):
class _ForeignMapping:
def validate_python(self, value):
return {1: "coerced"}
monkeypatch.setattr(live, "_MAPPING", _ForeignMapping())
with pytest.raises(TypeError, match="Invalid Live JSON conversion target"):
live._json_value(MappingProxyType({"nested": "value"}))
def test_rewrite_session_ids_serializes_non_object_events_without_touching_ids():
assert live.rewrite_session_ids(["live", {"id": "raw"}], "raw", "public") == ["live", {"id": "raw"}]
assert live.rewrite_session_ids("public", "raw", "public") == "public"
def test_owner_requires_authenticated_api_key():
with pytest.raises(HTTPException) as rejected:
live._owner(UserAPIKeyAuth())
assert rejected.value.status_code == 403
assert rejected.value.detail == "Live sessions require an authenticated API key"
def _streamed_request(chunks: list[bytes]) -> Request:
pending = list(chunks)
async def receive():
return {"type": "http.request", "body": pending.pop(0), "more_body": bool(pending)}
return Request({"type": "http", "method": "POST", "headers": [], "query_string": b""}, receive=receive)
@pytest.mark.asyncio
async def test_body_rejects_streams_larger_than_the_offer_limit():
with pytest.raises(HTTPException) as rejected:
await live._body(_streamed_request([b"a" * (8 * 1024 * 1024), b"b"]))
assert rejected.value.status_code == 413
assert rejected.value.detail == "Live request exceeds the 8 MiB limit"
@pytest.mark.asyncio
async def test_body_rejects_json_that_is_not_an_object():
with pytest.raises(HTTPException) as rejected:
await live._body(_streamed_request([b"[1,2]"]))
assert rejected.value.status_code == 400
assert rejected.value.detail == "Expected a JSON object"
def test_session_model_requires_object_and_model():
with pytest.raises(HTTPException) as not_object:
live._session_model({"session": "voice"})
assert not_object.value.status_code == 400 and not_object.value.detail == "session must be a JSON object"
with pytest.raises(HTTPException) as no_model:
live._session_model({"session": {}})
assert no_model.value.status_code == 400 and no_model.value.detail == "session.model is required"
@pytest.mark.asyncio
async def test_team_organization_lookup_maps_failures_to_service_unavailable(monkeypatch):
from litellm.proxy import proxy_server
monkeypatch.setattr(proxy_server, "prisma_client", object())
monkeypatch.setattr(proxy_server, "user_api_key_cache", object())
monkeypatch.setattr(proxy_server, "proxy_logging_obj", object())
monkeypatch.setattr(live, "get_team_object", AsyncMock(side_effect=RuntimeError("database unavailable")))
with pytest.raises(HTTPException) as rejected:
await live._live_organization_id(UserAPIKeyAuth(api_key="owner", team_id="team"))
assert rejected.value.status_code == 503
assert rejected.value.detail == "Could not verify Live team organization model access"
@pytest.mark.asyncio
async def test_direct_user_authorization_fails_closed_when_user_lookup_fails(monkeypatch):
from litellm.proxy import proxy_server
monkeypatch.setattr(proxy_server, "llm_model_list", [])
monkeypatch.setattr(proxy_server, "llm_router", None)
monkeypatch.setattr(proxy_server, "prisma_client", object())
monkeypatch.setattr(proxy_server, "user_api_key_cache", object())
monkeypatch.setattr(proxy_server, "proxy_logging_obj", object())
monkeypatch.setattr(live, "can_key_call_resolved_model", AsyncMock())
monkeypatch.setattr(live, "get_user_object", AsyncMock(side_effect=RuntimeError("database unavailable")))
with pytest.raises(HTTPException) as rejected:
await live._authorize("voice", UserAPIKeyAuth(api_key="owner", user_id="user"))
assert rejected.value.status_code == 503
assert rejected.value.detail == "Could not verify Live user model access"
@pytest.mark.asyncio
async def test_direct_org_authorization_fails_closed_when_org_lookup_fails(monkeypatch):
from litellm.proxy import proxy_server
monkeypatch.setattr(proxy_server, "llm_model_list", [])
monkeypatch.setattr(proxy_server, "llm_router", None)
monkeypatch.setattr(proxy_server, "prisma_client", object())
monkeypatch.setattr(proxy_server, "user_api_key_cache", object())
monkeypatch.setattr(proxy_server, "proxy_logging_obj", object())
monkeypatch.setattr(live, "can_key_call_resolved_model", AsyncMock())
monkeypatch.setattr(live, "get_org_object", AsyncMock(side_effect=RuntimeError("database unavailable")))
with pytest.raises(HTTPException) as rejected:
await live._authorize("voice", UserAPIKeyAuth(api_key="owner", org_id="org-1"))
assert rejected.value.status_code == 503
assert rejected.value.detail == "Could not verify Live organization model access"
@pytest.mark.asyncio
@pytest.mark.parametrize(
"auth,lookup,expected",
[
(UserAPIKeyAuth(api_key="owner", user_id="user"), "get_user_object", "Could not verify Live user model access"),
(
UserAPIKeyAuth(api_key="owner", org_id="org-1"),
"get_org_object",
"Could not verify Live organization model access",
),
],
ids=["user", "organization"],
)
async def test_authorization_fails_closed_when_the_principal_row_is_missing(monkeypatch, auth, lookup, expected):
from litellm.proxy import proxy_server
monkeypatch.setattr(proxy_server, "llm_model_list", [])
monkeypatch.setattr(proxy_server, "llm_router", None)
monkeypatch.setattr(proxy_server, "prisma_client", object())
monkeypatch.setattr(proxy_server, "user_api_key_cache", object())
monkeypatch.setattr(proxy_server, "proxy_logging_obj", object())
monkeypatch.setattr(live, "can_key_call_resolved_model", AsyncMock())
monkeypatch.setattr(live, "can_user_call_model", AsyncMock())
monkeypatch.setattr(live, "can_org_access_model", Mock())
monkeypatch.setattr(live, lookup, AsyncMock(return_value=None))
with pytest.raises(HTTPException) as rejected:
await live._authorize("voice", auth)
assert rejected.value.status_code == 503
assert rejected.value.detail == expected
@pytest.mark.asyncio
async def test_deployment_requires_router_and_supported_provider(monkeypatch):
from litellm.proxy import proxy_server
monkeypatch.setattr(proxy_server, "llm_router", None)
with pytest.raises(HTTPException) as without_router:
await live._deployment("voice", {})
assert without_router.value.status_code == 503
assert without_router.value.detail == "Live requires a configured model deployment"
monkeypatch.setattr(
proxy_server,
"llm_router",
SimpleNamespace(
async_get_available_deployment=AsyncMock(
return_value={
"litellm_params": {"model": "bedrock/voice"},
"model_info": {"id": "deployment-a"},
}
),
async_routing_strategy_pre_call_checks=AsyncMock(),
),
)
with pytest.raises(HTTPException) as wrong_provider:
await live._deployment("voice", {})
assert wrong_provider.value.status_code == 400
assert "OpenAI or ChatGPT" in wrong_provider.value.detail
@pytest.mark.parametrize("payload", [{}, {"session": {"id": 5}}, {"session": None}])
def test_session_id_requires_upstream_string_id(payload):
with pytest.raises(HTTPException) as rejected:
live._session_id(payload)
assert rejected.value.status_code == 502
assert rejected.value.detail == "Upstream did not return a Live session ID"
@pytest.mark.asyncio
async def test_live_team_membership_prefers_reservation_cache_and_sentinel(monkeypatch):
from litellm.proxy import proxy_server
from litellm.proxy.common_utils.cache_pydantic_utils import CacheCodec
from litellm.proxy.common_utils.user_api_key_cache import NO_TEAM_MEMBERSHIP_SENTINEL
membership = LiteLLM_TeamMembership(user_id="user", team_id="team")
auth = UserAPIKeyAuth(api_key="owner", user_id="user", team_id="team")
monkeypatch.setattr(
proxy_server,
"user_api_key_cache",
SimpleNamespace(
async_get_cache=AsyncMock(return_value=CacheCodec.serialize(membership, model_type=LiteLLM_TeamMembership))
),
)
restored = await live._live_team_membership(auth)
assert restored is not None and restored.user_id == "user" and restored.team_id == "team"
monkeypatch.setattr(
proxy_server,
"user_api_key_cache",
SimpleNamespace(async_get_cache=AsyncMock(return_value=NO_TEAM_MEMBERSHIP_SENTINEL)),
)
assert await live._live_team_membership(auth) is None
@pytest.mark.asyncio
async def test_live_team_uses_team_cache_before_database(monkeypatch):
from litellm.proxy import proxy_server
team = SimpleNamespace(team_id="team", models=["*"])
monkeypatch.setattr(
proxy_server, "user_api_key_cache", SimpleNamespace(async_get_cache=AsyncMock(return_value=team))
)
assert await live._live_team(UserAPIKeyAuth(api_key="owner", team_id="team")) is team
@pytest.mark.parametrize(
"team",
[
SimpleNamespace(budget_limits=3, rpm_limit=None, tpm_limit=None, max_budget=None, model_max_budget=None),
SimpleNamespace(budget_limits=None, rpm_limit=5, tpm_limit=None, max_budget=None, model_max_budget=None),
SimpleNamespace(
budget_limits=None, rpm_limit=None, tpm_limit=None, max_budget=None, model_max_budget={"voice": 1}
),
],
ids=["scalar-windows", "scalar-rpm", "model-max-budget"],
)
def test_team_budget_fields_short_circuit_before_metadata_scan(team):
assert live._live_team_budget_configured(UserAPIKeyAuth(api_key="owner"), team) is True
@pytest.mark.parametrize(
"value,zero_is_limit,expected",
[
(None, False, False),
({"max_budget": 0}, True, True),
({"max_budget": 0}, False, False),
({"max_budget": "unlimited"}, False, True),
({"rpm_limit": 2}, False, True),
],
ids=["missing", "zero-as-limit", "zero-unlimited", "non-numeric-limit", "other-limit"],
)
def test_live_budget_configured_separates_zero_from_non_numeric_limits(value, zero_is_limit, expected):
assert live._live_budget_configured(value, zero_is_limit=zero_is_limit) is expected
@pytest.mark.asyncio
async def test_live_default_budget_uses_cached_team_member_budget(monkeypatch):
from litellm.proxy import proxy_server
budget = LiteLLM_BudgetTable(max_budget=1)
monkeypatch.setattr(
proxy_server, "user_api_key_cache", SimpleNamespace(async_get_cache=AsyncMock(return_value=budget))
)
team = SimpleNamespace(metadata={"team_member_budget_id": "budget-1"})
auth = UserAPIKeyAuth(api_key="owner", user_id="user", team_id="team")
assert await live._live_default_budget(auth, team) is budget
@pytest.mark.asyncio
async def test_live_project_uses_cache_or_reports_missing_row(monkeypatch):
from litellm.proxy import proxy_server
project = SimpleNamespace(project_id="project-1")
monkeypatch.setattr(
proxy_server, "user_api_key_cache", SimpleNamespace(async_get_cache=AsyncMock(return_value=project))
)
auth = UserAPIKeyAuth(api_key="owner", project_id="project-1")
assert await live._live_project(auth) is project
monkeypatch.setattr(
proxy_server, "user_api_key_cache", SimpleNamespace(async_get_cache=AsyncMock(return_value=None))
)
monkeypatch.setattr(
live,
"ProjectRepository",
lambda client: SimpleNamespace(table=SimpleNamespace(find_unique=AsyncMock(return_value=None))),
)
assert await live._live_project(auth) is None
@pytest.mark.asyncio
@pytest.mark.parametrize(
"project, managed",
[
(
SimpleNamespace(
litellm_budget_table=None,
budget_id=None,
model_rpm_limit={"voice": 5},
model_tpm_limit=None,
metadata=None,
),
True,
),
(
SimpleNamespace(
litellm_budget_table=None,
budget_id=None,
model_rpm_limit=None,
model_tpm_limit=None,
metadata={"rpm_limit": 5},
),
True,
),
(
SimpleNamespace(
litellm_budget_table=None, budget_id=None, model_rpm_limit=None, model_tpm_limit=None, metadata=None
),
False,
),
],
ids=["model-rate-limit", "metadata-limit", "nothing"],
)
async def test_project_budget_falls_back_to_rate_limits_and_metadata(monkeypatch, project, managed):
from litellm.proxy import proxy_server
monkeypatch.setattr(proxy_server, "prisma_client", object())
assert await live._live_project_budget_configured(UserAPIKeyAuth(api_key="owner"), project) is managed
@pytest.mark.asyncio
async def test_model_group_budget_requires_model_name():
assert await live._live_model_group_budget_configured(UserAPIKeyAuth(api_key="owner"), None, None, None) is False
@pytest.mark.asyncio
async def test_managed_member_budget_fails_closed_without_database(monkeypatch):
from litellm.proxy import proxy_server
monkeypatch.setattr(proxy_server, "prisma_client", None)
monkeypatch.setattr(proxy_server, "llm_router", object())
with pytest.raises(HTTPException) as rejected:
await live._managed_member_budget(UserAPIKeyAuth(api_key="owner", team_id="team"), "voice")
assert rejected.value.status_code == 503
assert rejected.value.detail == "Could not verify Live managed budgets"
@pytest.mark.asyncio
async def test_backend_delegation_without_responses_contract_requires_named_model(monkeypatch):
monkeypatch.setattr(live, "_managed_member_budget", AsyncMock(return_value=False))
await live._authorize_delegation(
{"session": {"delegation": {"type": "backend"}}},
UserAPIKeyAuth(api_key="owner"),
)
with pytest.raises(HTTPException) as rejected:
await live._authorize_delegation(
{"type": "session.start", "session": {"delegation": {"type": "responses"}}},
UserAPIKeyAuth(api_key="owner", models=["voice"]),
)
assert rejected.value.status_code == 400
assert "explicit authorized delegation.responses.model" in rejected.value.detail
@pytest.mark.asyncio
async def test_precall_aborts_when_transferred_quota_lease_cannot_renew(monkeypatch):
from litellm.proxy import proxy_server
from litellm.proxy.hooks.parallel_request_limiter_v3 import _PROXY_MaxParallelRequestsHandler_v3
lease = SimpleNamespace(start=Mock(), renew=AsyncMock(return_value=False), close=AsyncMock())
limiter = Mock(spec=_PROXY_MaxParallelRequestsHandler_v3)
limiter.transfer_realtime_call_slot = Mock(return_value=lease)
limiter.async_post_call_failure_hook = AsyncMock()
monkeypatch.setattr(proxy_server, "proxy_logging_obj", SimpleNamespace(get_proxy_hook=lambda _: limiter))
monkeypatch.setattr(proxy_server, "general_settings", {})
monkeypatch.setattr(live, "_authorize", AsyncMock())
monkeypatch.setattr(live, "process_codex_request", AsyncMock(return_value=({"model": "voice"}, Mock())))
request = live._request(Request({"type": "http", "headers": []}), {"session": {"model": "voice"}})
with pytest.raises(HTTPException) as lost:
async with live._precall(request, UserAPIKeyAuth(api_key="owner"), "voice"):
pytest.fail("session must not start when the quota reservation is lost")
assert lost.value.status_code == 503 and "quota reservation was lost" in lost.value.detail
lease.start.assert_called_once()
lease.close.assert_awaited_once()
limiter.async_post_call_failure_hook.assert_awaited_once()
@pytest.mark.asyncio
async def test_supervise_starts_observer_under_isolated_request_stash(monkeypatch):
start = AsyncMock(return_value="stream")
monkeypatch.setattr(live, "_start_supervisor", start)
request = Request({"type": "http", "headers": []})
auth = UserAPIKeyAuth(api_key="owner")
source = handle()
assert await live._supervise(request, source, auth, None, None) == "stream"
assert (
start.await_args.args[0] is request and start.await_args.args[1] is source and start.await_args.args[2] is auth
)
@pytest.mark.asyncio
async def test_observer_frontend_swallows_traffic_and_hangup_checks_upstream_status(monkeypatch):
from starlette.websockets import WebSocketState
connection = SimpleNamespace(close=AsyncMock())
transport = SimpleNamespace(
connect=AsyncMock(return_value=connection),
request=AsyncMock(
return_value=httpx.Response(
502, request=httpx.Request("POST", "http://upstream.test/live/sessions/sess_upstream/hangup")
)
),
)
monkeypatch.setattr(live, "LiveTransport", Mock(return_value=transport))
monkeypatch.setattr(
live, "process_codex_request", AsyncMock(return_value=({"model": "voice"}, Mock(litellm_params={})))
)
monkeypatch.setattr(live, "CALL_SUPERVISORS", SimpleNamespace(start=AsyncMock()))
supervisor_init = Mock()
monkeypatch.setattr(live, "CallSupervisor", supervisor_init)
stream_cls = Mock()
monkeypatch.setattr(live, "RealTimeStreaming", stream_cls)
request = Request({"type": "http", "headers": []})
await live._start_supervisor(request, handle(), UserAPIKeyAuth(api_key="owner"), None)
frontend = stream_cls.call_args.args[0]
frontend.client_state = WebSocketState.CONNECTED
frontend.application_state = WebSocketState.CONNECTED
assert await frontend.receive() == {"type": "websocket.disconnect", "code": 1000}
assert await frontend.send({"type": "websocket.send", "text": "tick"}) is None
hangup = supervisor_init.call_args.args[4]
with pytest.raises(httpx.HTTPStatusError):
await hangup()
transport.request.assert_awaited_once_with("POST", "live/sessions/sess_upstream/hangup")
@pytest.mark.asyncio
async def test_observer_startup_failure_still_hangs_up_and_keeps_the_original_error(monkeypatch):
from litellm.proxy.spend_tracking import budget_reservation
connection = SimpleNamespace(close=AsyncMock())
transport = SimpleNamespace(
connect=AsyncMock(return_value=connection),
request=AsyncMock(
return_value=httpx.Response(
200, request=httpx.Request("POST", "http://upstream.test/live/sessions/sess_upstream/hangup")
)
),
)
monkeypatch.setattr(live, "LiveTransport", Mock(return_value=transport))
monkeypatch.setattr(
live, "process_codex_request", AsyncMock(return_value=({"model": "voice"}, Mock(litellm_params={})))
)
monkeypatch.setattr(live, "CALL_SUPERVISORS", SimpleNamespace(start=AsyncMock()))
monkeypatch.setattr(live, "RealTimeStreaming", Mock())
monkeypatch.setattr(live, "CallSupervisor", Mock(side_effect=RuntimeError("supervisor refused the call")))
invalidate = AsyncMock()
monkeypatch.setattr(budget_reservation, "invalidate_budget_reservation_counters", invalidate)
request = Request({"type": "http", "headers": []})
with pytest.raises(RuntimeError, match="supervisor refused the call"):
await live._start_supervisor(request, handle(), UserAPIKeyAuth(api_key="owner"), None)
transport.request.assert_awaited_once_with("POST", "live/sessions/sess_upstream/hangup")
invalidate.assert_not_awaited()
connection.close.assert_awaited_once()
def test_admin_sip_accept_rejects_model_mismatch_and_passes_upstream_errors_through(route_client, monkeypatch):
from litellm.proxy import proxy_server
from litellm.proxy._types import LitellmUserRoles
route_client.auth.user_role = LitellmUserRoles.PROXY_ADMIN
monkeypatch.setattr(proxy_server, "llm_model_list", [{"model_name": "voice"}])
mismatch = route_client.client.post(
"/v1/live/sessions/sess_incoming/accept",
json={"session": {"model": "other", "type": "live"}},
headers={"x-litellm-live-model": "voice"},
)
assert mismatch.status_code == 400 and "must match" in mismatch.json()["detail"]
route_client.transport.request.assert_not_awaited()
route_client.transport.request.return_value = httpx.Response(503, json={"error": "gateway down"})
failed = route_client.client.post(
"/v1/live/sessions/sess_incoming/accept",
json={"session": {"model": "voice", "type": "live"}},
headers={"x-litellm-live-model": "voice"},
)
assert failed.status_code == 503 and "x-litellm-live-session-id" not in failed.headers
route_client.supervised.assert_not_awaited()
@pytest.mark.asyncio
async def test_public_sideband_rejects_restart_and_model_change_then_rewrites_ids():
websocket = SimpleNamespace(
receive_text=AsyncMock(return_value=json.dumps({"type": "session.start"})),
send_text=AsyncMock(),
close=AsyncMock(),
scope={},
headers=Mock(),
)
public = live._PublicSocket(websocket, handle(), "public", UserAPIKeyAuth(api_key="owner"))
with pytest.raises(HTTPException) as restart:
await public.receive_text()
assert restart.value.status_code == 400 and restart.value.detail == "Session has already started"
websocket.receive_text.return_value = json.dumps({"type": "session.update", "session": {"model": "other"}})
with pytest.raises(HTTPException) as model:
await public.receive_text()
assert model.value.status_code == 400 and model.value.detail == "Session model cannot change"
websocket.receive_text.return_value = json.dumps({"type": "custom", "session_id": "public"})
assert json.loads(await public.receive_text()) == {"type": "custom", "session_id": "sess_upstream"}
def test_startup_events_overflow_fails_closed():
events = live._StartupEvents()
for _ in range(128):
events.store({"type": "info"})
with pytest.raises(HTTPException) as overflowed:
events.store({"type": "info"})
assert overflowed.value.status_code == 502
def test_websocket_requires_api_key_then_session_start(route_client):
from starlette.websockets import WebSocketDisconnect
with pytest.raises(WebSocketDisconnect) as anonymous:
with route_client.client.websocket_connect("/v1/live/sessions") as ws:
ws.receive_json()
assert anonymous.value.code == 1008
with route_client.client.websocket_connect(
"/v1/live/sessions", headers={"Authorization": "Bearer owner"}
) as ws:
ws.send_json({"type": "ping"})
with pytest.raises(WebSocketDisconnect) as wrong_first:
ws.receive_json()
assert wrong_first.value.code == 1008
route_client.transport.request.assert_not_awaited()
@pytest.mark.asyncio
async def test_attached_socket_authorizes_policy_and_reuses_source_without_start(route_client, monkeypatch):
from starlette.websockets import WebSocket
streams: list = []
class AttachedStream:
def __init__(self, *args, **kwargs):
self.args = args
self.bidirectional_forward = AsyncMock()
streams.append(self)
backend = SimpleNamespace(send=AsyncMock(), close=AsyncMock(), recv=AsyncMock())
route_client.transport.connect = AsyncMock(return_value=backend)
monkeypatch.setattr(live, "RealTimeStreaming", AttachedStream)
authorize = AsyncMock()
monkeypatch.setattr(live, "_authorize_delegation", authorize)
token = live.encode_session(handle())
inbound = iter([{"type": "websocket.connect"}, {"type": "websocket.disconnect", "code": 1000}])
sent: list = []
async def receive():
return next(inbound)
async def send(message):
sent.append(message)
websocket = WebSocket(
{
"type": "websocket",
"path": f"/v1/live/sessions/{token}/attach",
"query_string": b"",
"headers": [(b"authorization", b"Bearer owner")],
"scheme": "ws",
"server": ("testserver", 80),
"client": ("testclient", 50000),
"subprotocols": [],
},
receive,
send,
)
await live.websocket_live_session(websocket, token)
assert sent == [{"type": "websocket.accept", "subprotocol": None, "headers": []}]
authorize.assert_awaited_once()
assert authorize.await_args.args[1] is route_client.auth
route_client.transport.connect.assert_awaited_once_with("live/sessions/sess_upstream/attach")
backend.send.assert_not_awaited()
backend.close.assert_awaited_once()
frontend = streams[0].args[0]
assert frontend.public_id == token and frontend.handle.session_id == "sess_upstream" and frontend.observer is None
route_client.supervised.assert_not_awaited()
streams[0].bidirectional_forward.assert_awaited_once()
def test_websocket_connection_failure_closes_with_internal_error(route_client):
from starlette.websockets import WebSocketDisconnect
route_client.transport.connect = AsyncMock(return_value=None)
with route_client.client.websocket_connect(
"/v1/live/sessions", headers={"Authorization": "Bearer owner"}
) as ws:
ws.send_json({"type": "session.start", "session": {"model": "voice"}})
with pytest.raises(WebSocketDisconnect) as internal:
ws.receive_json()
assert internal.value.code == 1011
route_client.transport.request.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize("stage", ["rejected", "crashed"])
async def test_websocket_close_after_asgi_completion_is_swallowed(monkeypatch, stage):
from fastapi import WebSocket
class CompletedWebSocket(WebSocket):
async def close(self, code=1000, reason=None):
raise RuntimeError("ASGI send channel already completed")
headers = [] if stage == "rejected" else [(b"authorization", b"Bearer owner")]
sent = []
async def receive():
return {"type": "websocket.disconnect"}
async def send(message):
sent.append(message)
websocket = CompletedWebSocket(
{
"type": "websocket",
"path": "/v1/live/sessions",
"query_string": b"",
"headers": headers,
"scheme": "ws",
"server": ("localhost", 4000),
},
receive,
send,
)
if stage == "crashed":
monkeypatch.setattr(live, "_auth", AsyncMock(side_effect=ConnectionError("redis down")))
await live.websocket_live_session(websocket)
assert sent == []

View file

@ -5387,3 +5387,27 @@ def test_live_missing_backend_price_preserves_duration_and_marks_accounting_inco
litellm_logging_obj=logger,
) == pytest.approx(0.75)
assert logger.model_call_details["realtime_backend_accounting_incomplete"] is True
@pytest.mark.parametrize(
"envelope",
[
{"type": "response.event"},
{"type": "response.event", "event": {"response": {"id": "resp"}}},
{"type": "response.event", "event": "not-an-object"},
],
)
def test_live_backend_malformed_envelope_is_dropped_without_accounting_flag(envelope):
"""
Malformed event envelopes (missing event, missing event.type, wrong shape)
are skipped silently: unlike a terminal response.completed that fails
response validation, they must not mark the call's accounting incomplete.
"""
from unittest.mock import MagicMock
from litellm.cost_calculator import _live_backend_response
logger = MagicMock(spec=Logging)
logger.model_call_details = {}
assert _live_backend_response(envelope, logger) is None
assert "realtime_backend_accounting_incomplete" not in logger.model_call_details