fix(chatgpt): apply model guardrails and realtime gateway URLs

This commit is contained in:
jibanez-staticduo 2026-09-10 04:47:54 +02:00
parent a47253abcb
commit 297c9fdfe8
No known key found for this signature in database
5 changed files with 72 additions and 8 deletions

View file

@ -119,15 +119,16 @@ class ChatGPTRealtime(OpenAIRealtime):
class ChatGPTRealtimeHTTPConfig(OpenAIRealtimeHTTPConfig):
realtime_calls_json: Final = True
def __init__(self, params: GenericLiteLLMParams) -> None:
def __init__(self, params: GenericLiteLLMParams, use_codex_backend: bool = True) -> None:
self._params = params
self._use_codex_backend = use_codex_backend
def get_api_base(
self,
api_base: str | None,
**kwargs: object, # kwargs-ok: provider interface accepts optional credentials
) -> str:
return api_base or Authenticator.get_api_base()
return api_base or (Authenticator.get_api_base() if self._use_codex_backend else ChatGPTRealtime.get_api_base())
def get_api_key(
self,
@ -159,7 +160,7 @@ class ChatGPTRealtimeHTTPConfig(OpenAIRealtimeHTTPConfig):
}
def get_complete_url(self, api_base: str | None, model: str, api_version: str | None = None) -> str:
return "https://api.openai.com/v1/realtime/client_secrets"
return f"{self.get_api_base(api_base).rstrip('/')}/realtime/client_secrets"
def get_transcription_session_url(
self,
@ -167,4 +168,4 @@ class ChatGPTRealtimeHTTPConfig(OpenAIRealtimeHTTPConfig):
model: str,
api_version: str | None = None,
) -> str:
return "https://api.openai.com/v1/realtime/transcription_sessions"
return f"{self.get_api_base(api_base).rstrip('/')}/realtime/transcription_sessions"

View file

@ -74,6 +74,7 @@ async def process_codex_request(
version=server.version,
proxy_logging_obj=server.proxy_logging_obj,
proxy_config=server.proxy_config,
llm_router=server.llm_router,
user_model=server.user_model,
user_temperature=server.user_temperature,
user_request_timeout=server.user_request_timeout,

View file

@ -76,6 +76,7 @@ def _get_realtime_http_provider_config(
dynamic_api_base: str | None,
dynamic_api_key: str | None,
litellm_params: GenericLiteLLMParams,
use_codex_backend: bool = False,
) -> tuple["BaseRealtimeHTTPConfig | None", str, str]:
"""
Return (provider_config, resolved_api_base, resolved_api_key) for the
@ -92,7 +93,7 @@ def _get_realtime_http_provider_config(
if custom_llm_provider == "chatgpt":
from litellm.llms.chatgpt.realtime import ChatGPTRealtimeHTTPConfig
provider_config = ChatGPTRealtimeHTTPConfig(litellm_params)
provider_config = ChatGPTRealtimeHTTPConfig(litellm_params, use_codex_backend=use_codex_backend)
elif custom_llm_provider in LlmProviders._member_map_.values():
provider_config = ProviderConfigManager.get_provider_realtime_http_config(
model="",
@ -100,9 +101,7 @@ def _get_realtime_http_provider_config(
)
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
litellm_params.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
@ -282,6 +281,7 @@ async def arealtime_calls(
dynamic_api_base=dynamic_api_base,
dynamic_api_key=dynamic_api_key,
litellm_params=litellm_params,
use_codex_backend=True,
)
if session is not None:
session = _with_resolved_session_model(session, model_name)

View file

@ -11,6 +11,40 @@ from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.types.router import GenericLiteLLMParams
@pytest.mark.asyncio
@pytest.mark.parametrize("endpoint", ["client_secrets", "transcription_sessions"])
@pytest.mark.parametrize("source", ["default", "explicit", "CHATGPT_API_BASE", "OPENAI_CHATGPT_API_BASE"])
async def test_realtime_session_urls_honor_gateway(endpoint, source, chatgpt_tokens, monkeypatch):
monkeypatch.setenv("CHATGPT_TOKEN_DIR", chatgpt_tokens)
monkeypatch.delenv("CHATGPT_API_BASE", raising=False)
monkeypatch.delenv("OPENAI_CHATGPT_API_BASE", raising=False)
gateway = "https://voice.example/custom/v1/"
if source in ("CHATGPT_API_BASE", "OPENAI_CHATGPT_API_BASE"):
monkeypatch.setenv(source, gateway)
requests = []
def respond(request):
requests.append(request)
return httpx.Response(200, json={"client_secret": {"value": "test-secret"}})
client = AsyncHTTPHandler()
client.client = httpx.AsyncClient(transport=httpx.MockTransport(respond))
kwargs = {"model": "chatgpt/gpt-realtime-1.5", "client": client}
if source == "explicit":
kwargs["api_base"] = gateway
try:
if endpoint == "client_secrets":
await litellm.acreate_realtime_client_secret(**kwargs)
else:
await litellm.acreate_realtime_transcription_session(**kwargs)
finally:
await client.client.aclose()
base = "https://api.openai.com/v1" if source == "default" else gateway.rstrip("/")
assert len(requests) == 1
assert str(requests[0].url) == f"{base}/realtime/{endpoint}"
assert requests[0].headers["authorization"] == "Bearer test-token-default"
@pytest.mark.asyncio
@pytest.mark.parametrize("inbound_headers", [{}, {"openai-alpha": "quicksilver=v2"}])
async def test_routed_call_preserves_deployment_gateway_headers(inbound_headers, chatgpt_tokens, monkeypatch):

View file

@ -11,6 +11,34 @@ 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("route_type", ["arealtime_calls", "_arealtime"])
async def test_codex_processing_merges_model_guardrails(monkeypatch, route_type):
from fastapi import Request
from litellm import Router
from litellm.proxy import proxy_server as server
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.realtime_endpoints.call_sessions import process_codex_request
class PolicyHook:
async def pre_call_hook(self, user_api_key_dict, data, call_type):
if "model-policy" in data.get("metadata", {}).get("guardrails", []):
raise HTTPException(403, "Model policy rejected request")
return data
router = Router(model_list=[{
"model_name": "voice-policy",
"litellm_params": {"model": "openai/gpt-realtime-1.5", "api_key": "test", "guardrails": ["model-policy"]},
}])
monkeypatch.setattr(server, "llm_router", router)
monkeypatch.setattr(server, "proxy_logging_obj", PolicyHook())
request = Request({"type": "http", "method": "POST", "path": "/v1/realtime/calls", "headers": [], "query_string": b"", "scheme": "http", "server": ("localhost", 80)})
with pytest.raises(HTTPException) as error:
await process_codex_request(request, {"model": "voice-policy"}, UserAPIKeyAuth(), "voice-policy", route_type)
assert error.value.status_code == 403
assert error.value.detail == "Model policy rejected request"
@pytest.mark.asyncio
@pytest.mark.parametrize("logged_success", [False, True])
@pytest.mark.parametrize("disconnect_error", [False, True])