fix(chatgpt): honor configured gateways across Codex transports

This commit is contained in:
jibanez-staticduo 2026-09-09 13:30:28 +02:00
parent 852dee23d4
commit 96594e7b0e
No known key found for this signature in database
6 changed files with 46 additions and 10 deletions

View file

@ -49,8 +49,8 @@ class Authenticator:
self.auth_file = os.path.join(self.token_dir, os.getenv("CHATGPT_AUTH_FILE", "auth.json"))
self._ensure_token_dir()
def get_api_base(self) -> str:
return os.getenv("CHATGPT_API_BASE") or os.getenv("OPENAI_CHATGPT_API_BASE") or CHATGPT_API_BASE
def get_api_base(self, default_base: str = CHATGPT_API_BASE) -> str:
return os.getenv("CHATGPT_API_BASE") or os.getenv("OPENAI_CHATGPT_API_BASE") or default_base
def get_access_token(self) -> str:
auth_data: Final = self._read_auth_file()

View file

@ -15,7 +15,7 @@ from litellm.llms.openai.image_generation.gpt_transformation import GPTImageGene
from litellm.types.llms.openai import AllMessageValues, FileTypes
from litellm.types.router import GenericLiteLLMParams
from .common_utils import CHATGPT_API_BASE
from .authenticator import Authenticator
from .responses.transformation import ChatGPTResponsesAPIConfig
@ -82,7 +82,7 @@ class ChatGPTImageGenerationConfig(GPTImageGenerationConfig):
litellm_params: Mapping[str, object],
stream: bool | None = None,
) -> str:
return f"{(api_base or CHATGPT_API_BASE).rstrip('/')}/images/generations"
return f"{(api_base or Authenticator().get_api_base()).rstrip('/')}/images/generations"
def transform_image_generation_request(
self,
@ -107,7 +107,7 @@ class ChatGPTImageEditConfig(OpenAIImageEditConfig):
return image_headers(headers, model, litellm_params or MappingProxyType({}))
def get_complete_url(self, model: str, api_base: str | None, litellm_params: Mapping[str, object]) -> str:
return f"{(api_base or CHATGPT_API_BASE).rstrip('/')}/images/edits"
return f"{(api_base or Authenticator().get_api_base()).rstrip('/')}/images/edits"
def use_multipart_form_data(self) -> bool:
return False

View file

@ -11,7 +11,7 @@ from litellm.types.realtime import RealtimeQueryParams
from litellm.types.router import GenericLiteLLMParams
from litellm.utils import get_model_info
from .common_utils import CHATGPT_API_BASE
from .authenticator import Authenticator
from .responses.transformation import ChatGPTResponsesAPIConfig
@ -44,6 +44,10 @@ def realtime_endpoint(model: str) -> str:
class ChatGPTRealtime(OpenAIRealtime):
@staticmethod
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")
def __init__(self, params: GenericLiteLLMParams, headers: Mapping[str, str]) -> None:
super().__init__()
self._profile_headers = realtime_headers(params, headers)
@ -90,7 +94,7 @@ class ChatGPTRealtimeHTTPConfig(OpenAIRealtimeHTTPConfig):
api_base: str | None,
**kwargs: object, # kwargs-ok: provider interface accepts optional credentials
) -> str:
return api_base or CHATGPT_API_BASE
return api_base or Authenticator().get_api_base()
def get_api_key(
self,

View file

@ -304,10 +304,12 @@ async def arealtime_calls(
api_version=litellm_params.api_version,
)
if custom_llm_provider == "chatgpt":
from litellm.llms.chatgpt.realtime import ChatGPTRealtime
response.extensions["chatgpt_realtime"] = MappingProxyType(
{
"model": model_name,
"api_base": litellm_params.api_base,
"api_base": ChatGPTRealtime.get_api_base(litellm_params.api_base),
}
)
return response
@ -462,7 +464,7 @@ async def _arealtime(
model=model,
websocket=websocket,
logging_obj=litellm_logging_obj,
api_base=api_base or "https://api.openai.com/v1",
api_base=ChatGPTRealtime.get_api_base(api_base),
api_key="chatgpt-oauth",
timeout=timeout,
query_params=query_params,

View file

@ -116,3 +116,16 @@ def test_edit_accepts_filesystem_path(tmp_path, as_tuple):
)
assert not files
assert data["images"] == ({"image_url": "data:image/png;base64," + base64.b64encode(image.read_bytes()).decode()},)
@pytest.mark.parametrize("env_name", ["CHATGPT_API_BASE", "OPENAI_CHATGPT_API_BASE"])
@pytest.mark.parametrize("api_base", [None, "https://deployment.example/codex"])
def test_image_routes_use_configured_gateway(monkeypatch, env_name, api_base):
monkeypatch.delenv("CHATGPT_API_BASE", raising=False)
monkeypatch.delenv("OPENAI_CHATGPT_API_BASE", raising=False)
monkeypatch.setenv(env_name, "https://gateway.example/codex/")
expected = api_base or "https://gateway.example/codex"
assert ChatGPTImageGenerationConfig().get_complete_url(api_base, None, "gpt-image-2", {}, {}) == (
expected + "/images/generations"
)
assert ChatGPTImageEditConfig().get_complete_url("gpt-image-2", api_base, {}) == expected + "/images/edits"

View file

@ -30,7 +30,7 @@ async def test_chatgpt_call_keeps_oauth_and_frameless_session(chatgpt_tokens, ap
extra_headers={"openai-alpha": "quicksilver=v2"},
client=client,
)
assert response.extensions["chatgpt_realtime"]["api_base"] == api_base
assert response.extensions["chatgpt_realtime"]["api_base"] == (api_base or "https://api.openai.com/v1")
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"
@ -82,3 +82,20 @@ def test_realtime_unknown_model_keeps_standard_endpoint(chatgpt_tokens, local_mo
assert handler._construct_url("https://api.openai.com/v1", {"model": "unknown-voice-model"}) == (
"wss://api.openai.com/v1/realtime?model=unknown-voice-model"
)
@pytest.mark.parametrize("env_name", ["CHATGPT_API_BASE", "OPENAI_CHATGPT_API_BASE"])
@pytest.mark.parametrize("api_base", [None, "https://deployment.example/codex"])
def test_realtime_routes_use_configured_gateway(monkeypatch, env_name, api_base, chatgpt_tokens):
from litellm.llms.chatgpt.realtime import ChatGPTRealtimeHTTPConfig
monkeypatch.delenv("CHATGPT_API_BASE", raising=False)
monkeypatch.delenv("OPENAI_CHATGPT_API_BASE", raising=False)
monkeypatch.setenv(env_name, "https://gateway.example/codex/")
expected = api_base or "https://gateway.example/codex"
config = ChatGPTRealtimeHTTPConfig(GenericLiteLLMParams())
assert config.get_realtime_calls_url(api_base, "gpt-live-1-codex") == expected + "/realtime/calls"
handler = ChatGPTRealtime(GenericLiteLLMParams(), {})
assert handler._construct_url(handler.get_api_base(api_base), {"model": "gpt-realtime-1.5"}) == (
expected.replace("https://", "wss://") + "/realtime?model=gpt-realtime-1.5"
)