diff --git a/litellm/llms/chatgpt/authenticator.py b/litellm/llms/chatgpt/authenticator.py index 563826c2b93..a516fb2fe09 100644 --- a/litellm/llms/chatgpt/authenticator.py +++ b/litellm/llms/chatgpt/authenticator.py @@ -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() diff --git a/litellm/llms/chatgpt/images.py b/litellm/llms/chatgpt/images.py index ee6ba17e340..97a0da5f518 100644 --- a/litellm/llms/chatgpt/images.py +++ b/litellm/llms/chatgpt/images.py @@ -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 diff --git a/litellm/llms/chatgpt/realtime.py b/litellm/llms/chatgpt/realtime.py index 831f67d82c5..5be3e766b99 100644 --- a/litellm/llms/chatgpt/realtime.py +++ b/litellm/llms/chatgpt/realtime.py @@ -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, diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index 15339173a1d..9d0d2075870 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -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, diff --git a/tests/test_litellm/llms/chatgpt/test_images.py b/tests/test_litellm/llms/chatgpt/test_images.py index bfdd1e53e4d..911183ba50b 100644 --- a/tests/test_litellm/llms/chatgpt/test_images.py +++ b/tests/test_litellm/llms/chatgpt/test_images.py @@ -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" diff --git a/tests/test_litellm/llms/chatgpt/test_realtime.py b/tests/test_litellm/llms/chatgpt/test_realtime.py index 7d83d8c6856..f66a3a92ac4 100644 --- a/tests/test_litellm/llms/chatgpt/test_realtime.py +++ b/tests/test_litellm/llms/chatgpt/test_realtime.py @@ -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" + )