mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(chatgpt): honor configured gateways across Codex transports
This commit is contained in:
parent
852dee23d4
commit
96594e7b0e
6 changed files with 46 additions and 10 deletions
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue