mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(chatgpt): retain alternate signaling keys and hook headers
This commit is contained in:
parent
297c9fdfe8
commit
bba4fd50e8
2 changed files with 41 additions and 11 deletions
|
|
@ -20,9 +20,10 @@ from litellm.llms.chatgpt.codex import (
|
|||
build_sideband_request,
|
||||
parse_call_response,
|
||||
)
|
||||
from litellm.llms.chatgpt.realtime import configured_realtime_headers
|
||||
from litellm.proxy._types import ProxyException, UserAPIKeyAuth
|
||||
from litellm.proxy.auth.auth_checks import can_key_call_resolved_model
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.auth.user_api_key_auth import get_api_key, get_api_key_from_custom_header, user_api_key_auth
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper, encrypt_value_helper
|
||||
from litellm.proxy.spend_tracking.budget_reservation import release_or_invalidate_budget_reservation
|
||||
|
||||
|
|
@ -99,11 +100,26 @@ async def create_codex_realtime_call(request: Request) -> Response:
|
|||
auth: Final = await user_api_key_auth(
|
||||
request=request,
|
||||
api_key=request.headers.get("authorization", ""),
|
||||
azure_api_key_header="",
|
||||
azure_api_key_header=request.headers.get("api-key", ""),
|
||||
anthropic_api_key_header=None,
|
||||
google_ai_studio_api_key_header=None,
|
||||
azure_apim_header=None,
|
||||
custom_litellm_key_header=None,
|
||||
custom_litellm_key_header=request.headers.get("x-litellm-api-key"),
|
||||
)
|
||||
selected_key, _ = get_api_key(
|
||||
request=request,
|
||||
api_key=request.headers.get("authorization", ""),
|
||||
azure_api_key_header=request.headers.get("api-key", ""),
|
||||
custom_litellm_key_header=request.headers.get("x-litellm-api-key"),
|
||||
anthropic_api_key_header=None,
|
||||
google_ai_studio_api_key_header=None,
|
||||
azure_apim_header=None,
|
||||
pass_through_endpoints=None,
|
||||
route="/v1/realtime/calls",
|
||||
)
|
||||
custom_header: Final = server.general_settings.get("litellm_key_header_name")
|
||||
owner_key: Final = (
|
||||
get_api_key_from_custom_header(request, custom_header) if isinstance(custom_header, str) else selected_key
|
||||
)
|
||||
try:
|
||||
await can_key_call_resolved_model(
|
||||
|
|
@ -132,7 +148,7 @@ async def create_codex_realtime_call(request: Request) -> Response:
|
|||
call: Final = parse_call_response(
|
||||
response,
|
||||
alias=model,
|
||||
owner=hashlib.sha256(request.headers.get("authorization", "").encode()).hexdigest(),
|
||||
owner=hashlib.sha256(f"Bearer {owner_key}".encode()).hexdigest(),
|
||||
expires_at=time.time() + 3600,
|
||||
)
|
||||
except ValueError as exc:
|
||||
|
|
@ -210,6 +226,12 @@ async def codex_realtime_sideband(websocket: WebSocket, token: str, auth: UserAP
|
|||
**{ # mutable-ok: retain processed policy metadata while pinning the existing call's routing
|
||||
**processed,
|
||||
**build_sideband_request(call),
|
||||
"extra_headers": MappingProxyType(
|
||||
{
|
||||
**configured_realtime_headers(processed.get("extra_headers")),
|
||||
**configured_realtime_headers(call.extra_headers),
|
||||
}
|
||||
),
|
||||
"websocket": websocket,
|
||||
"user_api_key_dict": auth,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -173,7 +173,8 @@ async def test_realtime_endpoint_rejects_untrusted_call_ids(monkeypatch, call_id
|
|||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("multipart", [False, True])
|
||||
@pytest.mark.parametrize("credential", ["authorization", "api-key", "subprotocol"])
|
||||
async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch, multipart, credential):
|
||||
@pytest.mark.parametrize("signaling_credential", ["authorization", "api-key", "x-litellm-api-key", "mixed"])
|
||||
async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch, multipart, credential, signaling_credential):
|
||||
import json
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
|
|
@ -197,15 +198,21 @@ async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch,
|
|||
async def receive():
|
||||
return {"type": "http.request", "body": body, "more_body": False}
|
||||
|
||||
signaling_headers = (
|
||||
[(b"authorization", b"Bearer other-owner"), (b"x-litellm-api-key", b"owner")]
|
||||
if signaling_credential == "mixed"
|
||||
else [(signaling_credential.encode(), b"Bearer owner" if signaling_credential == "authorization" else b"owner")]
|
||||
)
|
||||
request = Request({"type": "http", "method": "POST", "path": "/v1/realtime/calls",
|
||||
"scheme": "http", "server": ("localhost", 80),
|
||||
"query_string": b"intent=quicksilver&architecture=avas&untrusted=bad",
|
||||
"headers": [(b"content-type", body_request.headers["content-type"].encode()),
|
||||
(b"authorization", b"Bearer owner"), (b"openai-alpha", b"quicksilver=v2"),
|
||||
*signaling_headers, (b"openai-alpha", b"quicksilver=v2"),
|
||||
(b"x-untrusted", b"bad")]}, receive)
|
||||
auth = UserAPIKeyAuth()
|
||||
authenticate = AsyncMock(return_value=auth)
|
||||
authorize = AsyncMock()
|
||||
monkeypatch.setattr(codex, "user_api_key_auth", authenticate)
|
||||
monkeypatch.setattr(proxy_server, "master_key", "owner")
|
||||
monkeypatch.setattr(proxy_server, "general_settings", {})
|
||||
monkeypatch.setattr(codex, "can_key_call_resolved_model", authorize)
|
||||
|
||||
class Processor:
|
||||
|
|
@ -213,12 +220,12 @@ async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch,
|
|||
self.data = data
|
||||
|
||||
async def common_processing_pre_call_logic(self, **kwargs):
|
||||
assert kwargs["user_api_key_dict"] is auth
|
||||
assert isinstance(kwargs["user_api_key_dict"], UserAPIKeyAuth)
|
||||
if kwargs["route_type"] == "_arealtime":
|
||||
assert self.data["model"] == "voice-alias"
|
||||
assert self.data["guardrails"] == ["query-guardrail"]
|
||||
assert await kwargs["request"].json() == {"model": "voice-alias"}
|
||||
return {**self.data, "metadata": {"guardrails": ["policy-guardrail"], "user_api_key_team_id": "team"}}, None
|
||||
return {**self.data, "extra_headers": {"X-Hook-Required": "policy-value", "x-gateway-token": "untrusted-override", "Authorization": "Bearer untrusted"}, "metadata": {"guardrails": ["policy-guardrail"], "user_api_key_team_id": "team"}}, None
|
||||
return self.data, None
|
||||
|
||||
monkeypatch.setattr(common_request_processing, "ProxyBaseLLMRequestProcessing", Processor)
|
||||
|
|
@ -233,7 +240,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", "api_base": "https://voice.example/codex"}})
|
||||
extensions={"chatgpt_realtime": {"model": "gpt-live-1-codex", "api_base": "https://voice.example/codex", "extra_headers": {"X-Gateway-Token": "pinned-value"}}})
|
||||
return respond()
|
||||
|
||||
monkeypatch.setattr(proxy_server, "route_request", route)
|
||||
|
|
@ -270,6 +277,7 @@ async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch,
|
|||
assert sent[0]["type"] == "websocket.accept"
|
||||
if credential == "subprotocol":
|
||||
assert sent[0]["subprotocol"] == "realtime"
|
||||
assert forward.await_args.kwargs["extra_headers"] == {"x-hook-required": "policy-value", "x-gateway-token": "pinned-value"}
|
||||
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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue