mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(chatgpt): retain gateway headers across realtime connections
This commit is contained in:
parent
bcac3b7234
commit
4b85d94de4
6 changed files with 67 additions and 6 deletions
|
|
@ -19,6 +19,7 @@ class CodexRealtimeCall(BaseModel):
|
||||||
model: str
|
model: str
|
||||||
alias: str
|
alias: str
|
||||||
api_base: str | None = None
|
api_base: str | None = None
|
||||||
|
extra_headers: Mapping[str, str] | None = None
|
||||||
owner: str
|
owner: str
|
||||||
expires_at: float
|
expires_at: float
|
||||||
|
|
||||||
|
|
@ -26,6 +27,7 @@ class CodexRealtimeCall(BaseModel):
|
||||||
class ChatGPTCallRouting(BaseModel):
|
class ChatGPTCallRouting(BaseModel):
|
||||||
model: str
|
model: str
|
||||||
api_base: str | None = None
|
api_base: str | None = None
|
||||||
|
extra_headers: Mapping[str, str] | None = None
|
||||||
|
|
||||||
|
|
||||||
class CodexSidebandRequest(TypedDict):
|
class CodexSidebandRequest(TypedDict):
|
||||||
|
|
@ -33,6 +35,7 @@ class CodexSidebandRequest(TypedDict):
|
||||||
model: ReadOnly[str]
|
model: ReadOnly[str]
|
||||||
chatgpt_realtime_call_id: ReadOnly[str]
|
chatgpt_realtime_call_id: ReadOnly[str]
|
||||||
query_params: ReadOnly[RealtimeQueryParams]
|
query_params: ReadOnly[RealtimeQueryParams]
|
||||||
|
extra_headers: ReadOnly[Mapping[str, str] | None]
|
||||||
|
|
||||||
|
|
||||||
def build_call_request(
|
def build_call_request(
|
||||||
|
|
@ -67,6 +70,7 @@ def parse_call_response(response: httpx.Response, alias: str, owner: str, expire
|
||||||
owner=owner,
|
owner=owner,
|
||||||
expires_at=expires_at,
|
expires_at=expires_at,
|
||||||
api_base=routing.api_base,
|
api_base=routing.api_base,
|
||||||
|
extra_headers=routing.extra_headers,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -76,4 +80,5 @@ def build_sideband_request(call: CodexRealtimeCall) -> CodexSidebandRequest:
|
||||||
model=f"chatgpt/{call.model}",
|
model=f"chatgpt/{call.model}",
|
||||||
chatgpt_realtime_call_id=call.call_id,
|
chatgpt_realtime_call_id=call.call_id,
|
||||||
query_params=RealtimeQueryParams(model=call.model),
|
query_params=RealtimeQueryParams(model=call.model),
|
||||||
|
extra_headers=call.extra_headers,
|
||||||
)
|
)
|
||||||
|
|
|
||||||
|
|
@ -12,11 +12,19 @@ from litellm.types.router import GenericLiteLLMParams
|
||||||
from litellm.utils import get_model_info
|
from litellm.utils import get_model_info
|
||||||
|
|
||||||
from .authenticator import Authenticator
|
from .authenticator import Authenticator
|
||||||
|
from .common_utils import without_oauth_identity_headers
|
||||||
from .responses.transformation import ChatGPTResponsesAPIConfig
|
from .responses.transformation import ChatGPTResponsesAPIConfig
|
||||||
|
|
||||||
|
|
||||||
|
def configured_realtime_headers(headers: Mapping[str, object] | None) -> Mapping[str, str]:
|
||||||
|
validated: Final = TypeAdapter(Mapping[str, str]).validate_python(
|
||||||
|
without_oauth_identity_headers(headers or MappingProxyType({}))
|
||||||
|
)
|
||||||
|
return MappingProxyType({key.lower(): value for key, value in validated.items()})
|
||||||
|
|
||||||
|
|
||||||
def realtime_headers(
|
def realtime_headers(
|
||||||
params: GenericLiteLLMParams, headers: Mapping[str, str]
|
params: GenericLiteLLMParams, headers: Mapping[str, str], extra_headers: Mapping[str, object] | None = None
|
||||||
) -> dict[str, str]: # mutable-ok: HTTP handler header contract
|
) -> dict[str, str]: # mutable-ok: HTTP handler header contract
|
||||||
forwarded: Final = MappingProxyType(
|
forwarded: Final = MappingProxyType(
|
||||||
{
|
{
|
||||||
|
|
@ -32,6 +40,7 @@ def realtime_headers(
|
||||||
litellm_params=params,
|
litellm_params=params,
|
||||||
),
|
),
|
||||||
**forwarded,
|
**forwarded,
|
||||||
|
**configured_realtime_headers(extra_headers),
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -48,9 +57,11 @@ class ChatGPTRealtime(OpenAIRealtime):
|
||||||
def get_api_base(api_base: str | None = None) -> str:
|
def get_api_base(api_base: str | None = None) -> str:
|
||||||
return api_base or Authenticator.get_api_base(default_base="https://api.openai.com/v1")
|
return api_base or Authenticator.get_api_base(default_base="https://api.openai.com/v1")
|
||||||
|
|
||||||
def __init__(self, params: GenericLiteLLMParams, headers: Mapping[str, str]) -> None:
|
def __init__(
|
||||||
|
self, params: GenericLiteLLMParams, headers: Mapping[str, str], extra_headers: Mapping[str, object] | None = None
|
||||||
|
) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self._profile_headers = realtime_headers(params, headers)
|
self._profile_headers = realtime_headers(params, headers, extra_headers)
|
||||||
self._call_id = TypeAdapter(str | None).validate_python(getattr(params, "chatgpt_realtime_call_id", None))
|
self._call_id = TypeAdapter(str | None).validate_python(getattr(params, "chatgpt_realtime_call_id", None))
|
||||||
|
|
||||||
def _get_additional_headers(
|
def _get_additional_headers(
|
||||||
|
|
|
||||||
|
|
@ -304,12 +304,13 @@ async def arealtime_calls(
|
||||||
api_version=litellm_params.api_version,
|
api_version=litellm_params.api_version,
|
||||||
)
|
)
|
||||||
if custom_llm_provider == "chatgpt":
|
if custom_llm_provider == "chatgpt":
|
||||||
from litellm.llms.chatgpt.realtime import ChatGPTRealtime
|
from litellm.llms.chatgpt.realtime import ChatGPTRealtime, configured_realtime_headers
|
||||||
|
|
||||||
response.extensions["chatgpt_realtime"] = MappingProxyType(
|
response.extensions["chatgpt_realtime"] = MappingProxyType(
|
||||||
{
|
{
|
||||||
"model": model_name,
|
"model": model_name,
|
||||||
"api_base": ChatGPTRealtime.get_api_base(litellm_params.api_base),
|
"api_base": ChatGPTRealtime.get_api_base(litellm_params.api_base),
|
||||||
|
"extra_headers": configured_realtime_headers(kwargs.get("extra_headers")),
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
return response
|
return response
|
||||||
|
|
@ -460,7 +461,7 @@ async def _arealtime(
|
||||||
elif _custom_llm_provider == "chatgpt":
|
elif _custom_llm_provider == "chatgpt":
|
||||||
from litellm.llms.chatgpt.realtime import ChatGPTRealtime
|
from litellm.llms.chatgpt.realtime import ChatGPTRealtime
|
||||||
|
|
||||||
await ChatGPTRealtime(litellm_params, websocket.headers).async_realtime(
|
await ChatGPTRealtime(litellm_params, websocket.headers, headers).async_realtime(
|
||||||
model=model,
|
model=model,
|
||||||
websocket=websocket,
|
websocket=websocket,
|
||||||
logging_obj=litellm_logging_obj,
|
logging_obj=litellm_logging_obj,
|
||||||
|
|
|
||||||
|
|
@ -14,10 +14,12 @@ def test_signaling_rejects_invalid_upstream_call_id(location):
|
||||||
|
|
||||||
def test_signaling_preserves_selected_model_for_sideband():
|
def test_signaling_preserves_selected_model_for_sideband():
|
||||||
response = httpx.Response(201, headers={"Location": "/v1/realtime/calls/rtc_provider"},
|
response = httpx.Response(201, headers={"Location": "/v1/realtime/calls/rtc_provider"},
|
||||||
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-route": "voice"}}})
|
||||||
call = parse_call_response(response, "voice", "owner", 1000)
|
call = parse_call_response(response, "voice", "owner", 1000)
|
||||||
request = build_sideband_request(call)
|
request = build_sideband_request(call)
|
||||||
assert request["api_base"] == "https://voice.example/codex"
|
assert request["api_base"] == "https://voice.example/codex"
|
||||||
assert request["model"] == "chatgpt/gpt-live-1-codex"
|
assert request["model"] == "chatgpt/gpt-live-1-codex"
|
||||||
assert request["chatgpt_realtime_call_id"] == "rtc_provider"
|
assert request["chatgpt_realtime_call_id"] == "rtc_provider"
|
||||||
assert request["query_params"] == {"model": "gpt-live-1-codex"}
|
assert request["query_params"] == {"model": "gpt-live-1-codex"}
|
||||||
|
assert request["extra_headers"] == {"x-gateway-route": "voice"}
|
||||||
|
|
|
||||||
|
|
@ -1,7 +1,11 @@
|
||||||
|
import asyncio
|
||||||
import json
|
import json
|
||||||
|
from types import SimpleNamespace
|
||||||
|
from unittest.mock import AsyncMock
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
import pytest
|
import pytest
|
||||||
|
from websockets.asyncio.server import serve
|
||||||
|
|
||||||
import litellm
|
import litellm
|
||||||
from litellm.llms.chatgpt.realtime import ChatGPTRealtime
|
from litellm.llms.chatgpt.realtime import ChatGPTRealtime
|
||||||
|
|
@ -9,6 +13,41 @@ from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||||
from litellm.types.router import GenericLiteLLMParams
|
from litellm.types.router import GenericLiteLLMParams
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.parametrize("model", ["gpt-realtime-1.5", "gpt-live-1-codex"])
|
||||||
|
@pytest.mark.parametrize("call_id", [None, "rtc_existing"])
|
||||||
|
async def test_websocket_forwards_configured_headers_without_client_identity(model, call_id, chatgpt_tokens):
|
||||||
|
captured = asyncio.get_running_loop().create_future()
|
||||||
|
|
||||||
|
async def receive_connection(connection):
|
||||||
|
captured.set_result(connection.request.headers)
|
||||||
|
await connection.wait_closed()
|
||||||
|
|
||||||
|
websocket = SimpleNamespace(
|
||||||
|
headers={"authorization": "Bearer client", "cookie": "private-cookie", "openai-alpha": "client-value"},
|
||||||
|
scope={},
|
||||||
|
receive_text=AsyncMock(side_effect=RuntimeError("client disconnected")),
|
||||||
|
send_text=AsyncMock(),
|
||||||
|
close=AsyncMock(),
|
||||||
|
)
|
||||||
|
async with serve(receive_connection, "127.0.0.1", 0) as gateway:
|
||||||
|
port = gateway.sockets[0].getsockname()[1]
|
||||||
|
await asyncio.wait_for(litellm._arealtime(
|
||||||
|
model=f"chatgpt/{model}", websocket=websocket, api_base=f"http://127.0.0.1:{port}",
|
||||||
|
chatgpt_realtime_call_id=call_id,
|
||||||
|
headers={"x-deployment-header": "configured"},
|
||||||
|
extra_headers={"X-Gateway-Route": "voice", "OpenAI-Alpha": "configured-value",
|
||||||
|
"aUtHoRiZaTiOn": "Bearer wrong", "CHATGPT-ACCOUNT-ID": "wrong"},
|
||||||
|
), timeout=10)
|
||||||
|
headers = await asyncio.wait_for(captured, timeout=5)
|
||||||
|
assert headers["x-deployment-header"] == "configured"
|
||||||
|
assert headers["x-gateway-route"] == "voice"
|
||||||
|
assert headers["openai-alpha"] == "configured-value"
|
||||||
|
assert headers["authorization"] == "Bearer test-token-default"
|
||||||
|
assert headers["chatgpt-account-id"] == "test-account-default"
|
||||||
|
assert "cookie" not in headers
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
@pytest.mark.parametrize("api_base", [None, "https://voice.example/backend-api/codex"])
|
@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):
|
async def test_chatgpt_call_keeps_oauth_and_frameless_session(chatgpt_tokens, api_base):
|
||||||
|
|
@ -36,6 +75,7 @@ async def test_chatgpt_call_keeps_oauth_and_frameless_session(chatgpt_tokens, ap
|
||||||
client=client,
|
client=client,
|
||||||
)
|
)
|
||||||
assert response.extensions["chatgpt_realtime"]["api_base"] == (api_base or "https://api.openai.com/v1")
|
assert response.extensions["chatgpt_realtime"]["api_base"] == (api_base or "https://api.openai.com/v1")
|
||||||
|
assert response.extensions["chatgpt_realtime"]["extra_headers"] == {"openai-alpha": "quicksilver=v2", "x-gateway-route": "voice"}
|
||||||
assert requests[0].url.host == ("voice.example" if api_base else "chatgpt.com")
|
assert requests[0].url.host == ("voice.example" if api_base else "chatgpt.com")
|
||||||
assert response.status_code == 201
|
assert response.status_code == 201
|
||||||
assert requests[0].url.path == "/backend-api/codex/realtime/calls"
|
assert requests[0].url.path == "/backend-api/codex/realtime/calls"
|
||||||
|
|
|
||||||
|
|
@ -52,11 +52,13 @@ def test_sideband_token_binds_owner_and_model(monkeypatch):
|
||||||
call_id="rtc_test",
|
call_id="rtc_test",
|
||||||
model="gpt-live-1-codex",
|
model="gpt-live-1-codex",
|
||||||
alias="gpt-live-1-codex",
|
alias="gpt-live-1-codex",
|
||||||
|
extra_headers={"x-gateway-secret": "configured-secret"},
|
||||||
owner=hashlib.sha256(b"Bearer test-owner").hexdigest(),
|
owner=hashlib.sha256(b"Bearer test-owner").hexdigest(),
|
||||||
expires_at=time.time() + 300,
|
expires_at=time.time() + 300,
|
||||||
)
|
)
|
||||||
token = encode_call(call)
|
token = encode_call(call)
|
||||||
assert "/" not in token
|
assert "/" not in token
|
||||||
|
assert "configured-secret" not in token
|
||||||
assert decode_call(token, "Bearer test-owner") == call
|
assert decode_call(token, "Bearer test-owner") == call
|
||||||
with pytest.raises(HTTPException) as error:
|
with pytest.raises(HTTPException) as error:
|
||||||
decode_call(token, "Bearer different-owner")
|
decode_call(token, "Bearer different-owner")
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue