diff --git a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py index cbe6fe198c9..465c9dba64a 100644 --- a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py +++ b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py @@ -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), diff --git a/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py b/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py index b3a41d6cf81..045af34ec8a 100644 --- a/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py +++ b/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py @@ -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): diff --git a/tests/test_litellm/llms/chatgpt/test_codex.py b/tests/test_litellm/llms/chatgpt/test_codex.py index 35606931e24..59107d9acff 100644 --- a/tests/test_litellm/llms/chatgpt/test_codex.py +++ b/tests/test_litellm/llms/chatgpt/test_codex.py @@ -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) diff --git a/tests/test_litellm/llms/chatgpt/test_images.py b/tests/test_litellm/llms/chatgpt/test_images.py index 7cd8d727a94..3d966688395 100644 --- a/tests/test_litellm/llms/chatgpt/test_images.py +++ b/tests/test_litellm/llms/chatgpt/test_images.py @@ -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(), {} + ) diff --git a/tests/test_litellm/llms/chatgpt/test_live.py b/tests/test_litellm/llms/chatgpt/test_live.py index 2e1a35473ef..765b1723a5f 100644 --- a/tests/test_litellm/llms/chatgpt/test_live.py +++ b/tests/test_litellm/llms/chatgpt/test_live.py @@ -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 == [] diff --git a/tests/test_litellm/llms/chatgpt/test_realtime.py b/tests/test_litellm/llms/chatgpt/test_realtime.py index 2e3d091c088..4f367dd8229 100644 --- a/tests/test_litellm/llms/chatgpt/test_realtime.py +++ b/tests/test_litellm/llms/chatgpt/test_realtime.py @@ -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") diff --git a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py index bed15399fbb..2f58155eb72 100644 --- a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py +++ b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py @@ -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" diff --git a/tests/test_litellm/llms/openai/realtime/test_openai_realtime_handler.py b/tests/test_litellm/llms/openai/realtime/test_openai_realtime_handler.py index 4221954d787..2cde9acffc3 100644 --- a/tests/test_litellm/llms/openai/realtime/test_openai_realtime_handler.py +++ b/tests/test_litellm/llms/openai/realtime/test_openai_realtime_handler.py @@ -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 diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index b19d7b1d192..18e4b4f98c0 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -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": ""}) diff --git a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter.py index 589c8d1bd06..3b0ce1d4c7b 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter.py @@ -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 + ) diff --git a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py index 14d078820af..5c6a61cbc3f 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py @@ -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. diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py index cf3be1e40b8..cd139563dc2 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py @@ -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) diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py b/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py index 4455ef25cd7..e7581faab0d 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py @@ -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) diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_live.py b/tests/test_litellm/proxy/realtime_endpoints/test_live.py index 39592e8ef3d..4cc8f9dd16a 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_live.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_live.py @@ -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 == [] diff --git a/tests/test_litellm/test_cost_calculator.py b/tests/test_litellm/test_cost_calculator.py index 0364e609a30..a756224aa8c 100644 --- a/tests/test_litellm/test_cost_calculator.py +++ b/tests/test_litellm/test_cost_calculator.py @@ -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