mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(chatgpt): apply model guardrails and realtime gateway URLs
This commit is contained in:
parent
a47253abcb
commit
297c9fdfe8
5 changed files with 72 additions and 8 deletions
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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])
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue