diff --git a/litellm/llms/chatgpt/codex.py b/litellm/llms/chatgpt/codex.py index 85388d8c386..49e4227e5d2 100644 --- a/litellm/llms/chatgpt/codex.py +++ b/litellm/llms/chatgpt/codex.py @@ -18,15 +18,18 @@ class CodexRealtimeCall(BaseModel): call_id: str = Field(pattern=r"^rtc_[A-Za-z0-9_-]+$") model: str alias: str + api_base: str | None = None owner: str expires_at: float class ChatGPTCallRouting(BaseModel): model: str + api_base: str | None = None class CodexSidebandRequest(TypedDict): + api_base: ReadOnly[str | None] model: ReadOnly[str] chatgpt_realtime_call_id: ReadOnly[str] query_params: ReadOnly[RealtimeQueryParams] @@ -63,11 +66,13 @@ def parse_call_response(response: httpx.Response, alias: str, owner: str, expire alias=alias, owner=owner, expires_at=expires_at, + api_base=routing.api_base, ) def build_sideband_request(call: CodexRealtimeCall) -> CodexSidebandRequest: return CodexSidebandRequest( + api_base=call.api_base, model=f"chatgpt/{call.model}", chatgpt_realtime_call_id=call.call_id, query_params=RealtimeQueryParams(model=call.model), diff --git a/litellm/proxy/realtime_endpoints/call_sessions.py b/litellm/proxy/realtime_endpoints/call_sessions.py index f0105ea6cfd..b1d1112af21 100644 --- a/litellm/proxy/realtime_endpoints/call_sessions.py +++ b/litellm/proxy/realtime_endpoints/call_sessions.py @@ -10,6 +10,8 @@ from fastapi import HTTPException, Request, Response, WebSocket from starlette.types import Message from litellm._logging import verbose_proxy_logger +from litellm.litellm_core_utils.litellm_logging import Logging +from litellm.litellm_core_utils.realtime_streaming import REALTIME_SESSION_SUCCESS_LOGGED_KEY from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.llms.chatgpt.codex import ( CodexRealtimeCall, @@ -60,12 +62,12 @@ async def process_codex_request( auth: UserAPIKeyAuth, model: str, route_type: Literal["arealtime_calls", "_arealtime"], -) -> dict[str, object]: # mutable-ok: common request processor returns enriched routing arguments +) -> tuple[dict[str, object], Logging]: # mutable-ok: common request processor returns enriched routing arguments from litellm.proxy import proxy_server as server from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing processor: Final = ProxyBaseLLMRequestProcessing(data=data) - processed, _ = await processor.common_processing_pre_call_logic( + processed, logging_obj = await processor.common_processing_pre_call_logic( request=request, general_settings=server.general_settings, user_api_key_dict=auth, @@ -80,7 +82,7 @@ async def process_codex_request( model=model, route_type=route_type, ) - return processed + return processed, logging_obj async def create_codex_realtime_call(request: Request) -> Response: @@ -110,7 +112,7 @@ async def create_codex_realtime_call(request: Request) -> Response: llm_router=server.llm_router, ) data: Final = build_call_request(offer, request.query_params, request.headers) - processed: Final = await process_codex_request(request, data, auth, model, "arealtime_calls") + processed, _ = await process_codex_request(request, data, auth, model, "arealtime_calls") result: Final = await server.route_request( data=processed, route_type="arealtime_calls", @@ -156,6 +158,7 @@ async def codex_realtime_sideband(websocket: WebSocket, token: str, auth: UserAP (p.removeprefix("openai-insecure-api-key.") for p in protocols if p.startswith("openai-insecure-api-key.")), "" ) authorization: Final = websocket.headers.get("authorization") or f"Bearer {alternate_key}" + logging_obj: Logging | None = None # rebind-ok: cleanup needs the logger only after pre-call succeeds try: try: call: Final = decode_call(token, authorization) @@ -194,7 +197,7 @@ async def codex_realtime_sideband(websocket: WebSocket, token: str, auth: UserAP ], } try: - processed: Final = await process_codex_request(request, data, auth, call.alias, "_arealtime") + processed, logging_obj = await process_codex_request(request, data, auth, call.alias, "_arealtime") except Exception: # noqa: BLE001 # custom hook exceptions must reject the connection verbose_proxy_logger.exception("Realtime sideband pre-call rejected") await websocket.close(code=1008, reason="Realtime pre-call rejected") @@ -211,4 +214,5 @@ async def codex_realtime_sideband(websocket: WebSocket, token: str, auth: UserAP } ) finally: - await release_or_invalidate_budget_reservation(budget_reservation=auth.budget_reservation) + if logging_obj is None or not logging_obj.model_call_details.get(REALTIME_SESSION_SUCCESS_LOGGED_KEY): + await release_or_invalidate_budget_reservation(budget_reservation=auth.budget_reservation) diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index 4a5859d718d..15339173a1d 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -99,7 +99,11 @@ def _get_realtime_http_provider_config( provider=LlmProviders(custom_llm_provider), ) - raw_api_base: Final = dynamic_api_base or litellm_params.api_base + raw_api_base: Final = ( + litellm_params.api_base or dynamic_api_base + if custom_llm_provider == "chatgpt" + else dynamic_api_base or litellm_params.api_base + ) raw_api_key: Final = dynamic_api_key or litellm_params.api_key if provider_config is not None: @@ -303,6 +307,7 @@ async def arealtime_calls( response.extensions["chatgpt_realtime"] = MappingProxyType( { "model": model_name, + "api_base": litellm_params.api_base, } ) return response diff --git a/tests/test_litellm/llms/chatgpt/test_codex.py b/tests/test_litellm/llms/chatgpt/test_codex.py index a893ddbfb12..3ea5b1fffd2 100644 --- a/tests/test_litellm/llms/chatgpt/test_codex.py +++ b/tests/test_litellm/llms/chatgpt/test_codex.py @@ -14,9 +14,10 @@ def test_signaling_rejects_invalid_upstream_call_id(location): def test_signaling_preserves_selected_model_for_sideband(): response = httpx.Response(201, headers={"Location": "/v1/realtime/calls/rtc_provider"}, - extensions={"chatgpt_realtime": {"model": "gpt-live-1-codex"}}) + extensions={"chatgpt_realtime": {"model": "gpt-live-1-codex", "api_base": "https://voice.example/codex"}}) call = parse_call_response(response, "voice", "owner", 1000) request = build_sideband_request(call) + assert request["api_base"] == "https://voice.example/codex" assert request["model"] == "chatgpt/gpt-live-1-codex" assert request["chatgpt_realtime_call_id"] == "rtc_provider" assert request["query_params"] == {"model": "gpt-live-1-codex"} diff --git a/tests/test_litellm/llms/chatgpt/test_realtime.py b/tests/test_litellm/llms/chatgpt/test_realtime.py index 239a5a12280..7d83d8c6856 100644 --- a/tests/test_litellm/llms/chatgpt/test_realtime.py +++ b/tests/test_litellm/llms/chatgpt/test_realtime.py @@ -10,7 +10,8 @@ from litellm.types.router import GenericLiteLLMParams @pytest.mark.asyncio -async def test_chatgpt_call_keeps_oauth_and_frameless_session(chatgpt_tokens): +@pytest.mark.parametrize("api_base", [None, "https://voice.example/backend-api/codex"]) +async def test_chatgpt_call_keeps_oauth_and_frameless_session(chatgpt_tokens, api_base): requests = [] def respond(request): @@ -21,6 +22,7 @@ async def test_chatgpt_call_keeps_oauth_and_frameless_session(chatgpt_tokens): client.client = httpx.AsyncClient(transport=httpx.MockTransport(respond)) response = await litellm.arealtime_calls( model="chatgpt/gpt-live-1-codex", + api_base=api_base, openai_ephemeral_key="", sdp_body=b"v=0\r\n", session={"model": "chatgpt/gpt-live-1-codex", "audio": {"output": {"voice": "sol"}}}, @@ -28,6 +30,8 @@ async def test_chatgpt_call_keeps_oauth_and_frameless_session(chatgpt_tokens): extra_headers={"openai-alpha": "quicksilver=v2"}, client=client, ) + assert response.extensions["chatgpt_realtime"]["api_base"] == api_base + assert requests[0].url.host == ("voice.example" if api_base else "chatgpt.com") assert response.status_code == 201 assert requests[0].url.path == "/backend-api/codex/realtime/calls" assert requests[0].url.params["architecture"] == "avas" 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 617cd644702..59c08c09b00 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py @@ -1,5 +1,6 @@ import hashlib import time +from types import SimpleNamespace import pytest from fastapi import HTTPException, WebSocket @@ -10,6 +11,41 @@ from litellm.llms.chatgpt.codex import CodexRealtimeCall from litellm.proxy.realtime_endpoints.call_sessions import decode_call, encode_call +@pytest.mark.asyncio +@pytest.mark.parametrize("logged_success", [False, True]) +@pytest.mark.parametrize("disconnect_error", [False, True]) +async def test_sideband_preserves_pending_cost_reconciliation(monkeypatch, logged_success, disconnect_error): + import litellm + from unittest.mock import AsyncMock + from litellm.litellm_core_utils.realtime_streaming import REALTIME_SESSION_SUCCESS_LOGGED_KEY + + monkeypatch.setenv("LITELLM_SALT_KEY", "test-only-salt-for-codex-realtime") + call = CodexRealtimeCall(call_id="rtc_test", model="gpt-live-1-codex", alias="voice", + owner=hashlib.sha256(b"Bearer owner").hexdigest(), expires_at=time.time()+300) + auth = UserAPIKeyAuth() + auth.budget_reservation = {"reserved_cost": 0.55, "input_cost": 0.0, "finalized": False, "entries": []} + logger = SimpleNamespace(model_call_details={}) + monkeypatch.setattr(codex, "can_key_call_resolved_model", AsyncMock()) + monkeypatch.setattr(codex, "process_codex_request", AsyncMock(return_value=({}, logger))) + + async def forward(**kwargs): + if logged_success: + logger.model_call_details[REALTIME_SESSION_SUCCESS_LOGGED_KEY] = True + if disconnect_error: + raise RuntimeError("Backend disconnected") + + monkeypatch.setattr(litellm, "_arealtime", forward) + websocket = WebSocket({"type": "websocket", "path": "/v1/live/opaque", "query_string": b"", + "headers": [(b"authorization", b"Bearer owner")]}, + AsyncMock(return_value={"type": "websocket.connect"}), AsyncMock()) + if disconnect_error: + with pytest.raises(RuntimeError, match="Backend disconnected"): + await codex.codex_realtime_sideband(websocket, encode_call(call), auth) + else: + await codex.codex_realtime_sideband(websocket, encode_call(call), auth) + assert auth.budget_reservation["finalized"] is not logged_success + + def test_sideband_token_binds_owner_and_model(monkeypatch): monkeypatch.setenv("LITELLM_SALT_KEY", "test-only-salt-for-codex-realtime") call = CodexRealtimeCall( @@ -166,7 +202,7 @@ async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch, async def respond(): return httpx.Response(201, content=b"v=0\r\nanswer", headers={"Location": "/v1/realtime/calls/rtc_private"}, - extensions={"chatgpt_realtime": {"model": "gpt-live-1-codex"}}) + extensions={"chatgpt_realtime": {"model": "gpt-live-1-codex", "api_base": "https://voice.example/codex"}}) return respond() monkeypatch.setattr(proxy_server, "route_request", route) @@ -206,6 +242,7 @@ async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch, assert forward.await_args.kwargs["metadata"] == {"guardrails": ["policy-guardrail"], "user_api_key_team_id": "team"} assert forward.await_args.kwargs["chatgpt_realtime_call_id"] == "rtc_private" assert forward.await_args.kwargs["model"] == "chatgpt/gpt-live-1-codex" + assert forward.await_args.kwargs["api_base"] == "https://voice.example/codex" assert authorize.await_count == 2