mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(chatgpt): retain sideband routing and pending usage reservations
This commit is contained in:
parent
b0022e5d2e
commit
852dee23d4
6 changed files with 66 additions and 10 deletions
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue