mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
test: cover ChatGPT realtime and image routes
This commit is contained in:
parent
ccefa07cea
commit
fbace1f2c6
15 changed files with 1497 additions and 2 deletions
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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(), {}
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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 == []
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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": ""})
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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 == []
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue