mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge b717835645 into 79756cbb9b
This commit is contained in:
commit
923e4e5517
4 changed files with 170 additions and 6 deletions
|
|
@ -1027,8 +1027,6 @@ openai_compatible_providers: Final[list] = [
|
|||
"sail",
|
||||
]
|
||||
|
||||
OPENAI_AUDIO_TRANSCRIPTION_PROVIDERS: Final = frozenset({"openai"} | frozenset(openai_compatible_providers))
|
||||
|
||||
openai_text_completion_compatible_providers: Final[list] = [ # providers that support `/v1/completions`
|
||||
"together_ai",
|
||||
"fireworks_ai",
|
||||
|
|
|
|||
|
|
@ -3,10 +3,15 @@ JSON-based provider configuration loader for OpenAI-compatible providers.
|
|||
"""
|
||||
|
||||
import json
|
||||
from collections.abc import Mapping, Sequence
|
||||
from pathlib import Path
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.constants import openai_compatible_providers
|
||||
|
||||
OPENAI_AUDIO_ENDPOINTS: Final = frozenset({"/v1/audio/transcriptions", "/v1/audio/speech"})
|
||||
|
||||
|
||||
class SimpleProviderConfig:
|
||||
|
|
@ -21,7 +26,7 @@ class SimpleProviderConfig:
|
|||
self.param_mappings = data.get("param_mappings", {})
|
||||
self.constraints = data.get("constraints", {})
|
||||
self.special_handling = data.get("special_handling", {})
|
||||
self.supported_endpoints = data.get("supported_endpoints", [])
|
||||
self.supported_endpoints: Final[Sequence[str]] = data.get("supported_endpoints", [])
|
||||
|
||||
|
||||
class JSONProviderRegistry:
|
||||
|
|
@ -78,6 +83,11 @@ class JSONProviderRegistry:
|
|||
return False
|
||||
return "/v1/responses" in provider.supported_endpoints
|
||||
|
||||
@classmethod
|
||||
def declared_endpoints(cls) -> Mapping[str, Sequence[str]]:
|
||||
"""Endpoints each JSON provider declares for itself"""
|
||||
return MappingProxyType({slug: provider.supported_endpoints for slug, provider in cls._providers.items()})
|
||||
|
||||
@classmethod
|
||||
def list_providers(cls) -> list:
|
||||
"""List all registered provider slugs"""
|
||||
|
|
@ -86,3 +96,29 @@ class JSONProviderRegistry:
|
|||
|
||||
# Load on import
|
||||
JSONProviderRegistry.load()
|
||||
|
||||
|
||||
def derive_openai_audio_transcription_providers(
|
||||
compatible_providers: Sequence[str],
|
||||
declared_endpoints: Mapping[str, Sequence[str]],
|
||||
) -> frozenset[str]:
|
||||
"""Providers allowed to drive the OpenAI `/v1/audio/*` transport.
|
||||
|
||||
A JSON provider joins only when it declares an audio endpoint for itself, so a chat-only
|
||||
provider can never be sent the caller's OpenAI credentials against its third-party base url.
|
||||
"""
|
||||
non_json_providers: Final = frozenset(
|
||||
provider for provider in compatible_providers if provider not in declared_endpoints
|
||||
)
|
||||
declaring_providers: Final = frozenset(
|
||||
slug
|
||||
for slug, endpoints in declared_endpoints.items()
|
||||
if any(endpoint in OPENAI_AUDIO_ENDPOINTS for endpoint in endpoints)
|
||||
)
|
||||
return frozenset({"openai"}) | non_json_providers | declaring_providers
|
||||
|
||||
|
||||
OPENAI_AUDIO_TRANSCRIPTION_PROVIDERS: Final = derive_openai_audio_transcription_providers(
|
||||
compatible_providers=openai_compatible_providers,
|
||||
declared_endpoints=JSONProviderRegistry.declared_endpoints(),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -66,7 +66,6 @@ from litellm.constants import (
|
|||
DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT,
|
||||
DEFAULT_MOCK_RESPONSE_PROMPT_TOKEN_COUNT,
|
||||
NADIR_DEFAULT_API_BASE,
|
||||
OPENAI_AUDIO_TRANSCRIPTION_PROVIDERS,
|
||||
)
|
||||
from litellm.exceptions import LiteLLMUnknownProvider
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
|
@ -120,7 +119,10 @@ from litellm.llms.bedrock.common_utils import BedrockModelInfo
|
|||
from litellm.llms.cohere.common_utils import CohereModelInfo
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler, http2_enabled
|
||||
from litellm.llms.openai.chat.gpt_5_transformation import OpenAIGPT5Config
|
||||
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
|
||||
from litellm.llms.openai_like.json_loader import (
|
||||
OPENAI_AUDIO_TRANSCRIPTION_PROVIDERS,
|
||||
JSONProviderRegistry,
|
||||
)
|
||||
from litellm.llms.vertex_ai.common_utils import (
|
||||
VertexAIModelRoute,
|
||||
get_vertex_ai_model_route,
|
||||
|
|
@ -8315,7 +8317,7 @@ def speech(
|
|||
_is_async=aspeech or False,
|
||||
)
|
||||
elif custom_llm_provider == "openai" or (
|
||||
custom_llm_provider in litellm.openai_compatible_providers
|
||||
custom_llm_provider in OPENAI_AUDIO_TRANSCRIPTION_PROVIDERS
|
||||
and custom_llm_provider not in AZURE_OPENAI_AUDIO_PROVIDERS
|
||||
):
|
||||
if voice is None or not (isinstance(voice, str)):
|
||||
|
|
|
|||
128
tests/unit/llms/openai_like/test_json_loader.py
Normal file
128
tests/unit/llms/openai_like/test_json_loader.py
Normal file
|
|
@ -0,0 +1,128 @@
|
|||
"""
|
||||
A JSON-configured provider may use the OpenAI `/v1/audio/*` transport only when it declares
|
||||
that endpoint for itself, so `transcription()` and `speech()` cannot carry `OPENAI_API_KEY`
|
||||
to a provider whose entry in `providers.json` advertises chat only.
|
||||
"""
|
||||
|
||||
import io
|
||||
from collections.abc import Sequence
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
import respx
|
||||
|
||||
import litellm
|
||||
from litellm.constants import openai_compatible_providers
|
||||
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
|
||||
from litellm.llms.openai_like.json_loader import (
|
||||
OPENAI_AUDIO_ENDPOINTS,
|
||||
OPENAI_AUDIO_TRANSCRIPTION_PROVIDERS,
|
||||
JSONProviderRegistry,
|
||||
derive_openai_audio_transcription_providers,
|
||||
)
|
||||
|
||||
OPENAI_KEY_CANARY: Final = "sk-leaked-openai-key"
|
||||
|
||||
|
||||
def _declares_audio(endpoints: Sequence[str]) -> bool:
|
||||
return any(endpoint in OPENAI_AUDIO_ENDPOINTS for endpoint in endpoints)
|
||||
|
||||
|
||||
def _chat_only_json_provider() -> str:
|
||||
"""A JSON provider that is reachable as OpenAI-compatible yet declares no audio endpoint"""
|
||||
return next(
|
||||
slug
|
||||
for slug, endpoints in JSONProviderRegistry.declared_endpoints().items()
|
||||
if slug in openai_compatible_providers and not _declares_audio(endpoints)
|
||||
)
|
||||
|
||||
|
||||
def _audio_file() -> io.BytesIO:
|
||||
file: Final = io.BytesIO(b"riff-bytes")
|
||||
file.name = "audio.wav"
|
||||
return file
|
||||
|
||||
|
||||
def test_json_provider_is_audio_capable_only_when_it_declares_audio():
|
||||
for slug, endpoints in JSONProviderRegistry.declared_endpoints().items():
|
||||
assert (slug in OPENAI_AUDIO_TRANSCRIPTION_PROVIDERS) is _declares_audio(endpoints), slug
|
||||
|
||||
|
||||
def test_python_openai_compatible_providers_keep_the_audio_transport():
|
||||
for provider in openai_compatible_providers:
|
||||
if not JSONProviderRegistry.exists(provider):
|
||||
assert provider in OPENAI_AUDIO_TRANSCRIPTION_PROVIDERS, provider
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
["declared", "expected"],
|
||||
[
|
||||
(["/v1/audio/transcriptions"], True),
|
||||
(["/v1/audio/speech"], True),
|
||||
(["/v1/chat/completions", "/v1/audio/speech"], True),
|
||||
(["/v1/chat/completions", "/v1/responses"], False),
|
||||
([], False),
|
||||
],
|
||||
ids=["transcriptions", "speech", "audio_plus_chat", "chat_only", "nothing_declared"],
|
||||
)
|
||||
def test_derivation_grants_membership_on_declared_audio_endpoints_only(declared: list[str], expected: bool):
|
||||
derived: Final = derive_openai_audio_transcription_providers(
|
||||
compatible_providers=["groq", "chat-only"],
|
||||
declared_endpoints={"chat-only": ["/v1/chat/completions"], "voiced": declared},
|
||||
)
|
||||
|
||||
assert ("voiced" in derived) is expected
|
||||
assert "groq" in derived
|
||||
assert "chat-only" not in derived
|
||||
assert "openai" in derived
|
||||
|
||||
|
||||
def test_transcription_sends_no_openai_key_to_json_provider_without_audio_endpoint(monkeypatch):
|
||||
slug: Final = _chat_only_json_provider()
|
||||
monkeypatch.delenv(JSONProviderRegistry.get(slug).api_key_env, raising=False)
|
||||
monkeypatch.setenv("OPENAI_API_KEY", OPENAI_KEY_CANARY)
|
||||
|
||||
with respx.mock(assert_all_called=False) as upstream:
|
||||
route: Final = upstream.post(f"{JSONProviderRegistry.get(slug).base_url}/audio/transcriptions").respond(
|
||||
200, json={"text": "should never be reached"}
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="Unmapped provider"):
|
||||
litellm.transcription(model=f"{slug}/whisper-1", file=_audio_file())
|
||||
|
||||
assert route.called is False
|
||||
assert len(upstream.calls) == 0
|
||||
|
||||
|
||||
def test_speech_sends_no_openai_key_to_json_provider_without_audio_endpoint(monkeypatch):
|
||||
slug: Final = _chat_only_json_provider()
|
||||
monkeypatch.delenv(JSONProviderRegistry.get(slug).api_key_env, raising=False)
|
||||
monkeypatch.setenv("OPENAI_API_KEY", OPENAI_KEY_CANARY)
|
||||
|
||||
with respx.mock(assert_all_called=False) as upstream:
|
||||
route: Final = upstream.post(f"{JSONProviderRegistry.get(slug).base_url}/audio/speech").respond(
|
||||
200, content=b"should never be reached"
|
||||
)
|
||||
|
||||
with pytest.raises(Exception, match=f"Unable to map the custom llm provider={slug}"):
|
||||
litellm.speech(model=f"{slug}/tts-1", input="hello", voice="alloy")
|
||||
|
||||
assert route.called is False
|
||||
assert len(upstream.calls) == 0
|
||||
|
||||
|
||||
def test_transcription_still_routes_a_python_provider_with_its_own_credential(monkeypatch):
|
||||
monkeypatch.setenv("GROQ_API_KEY", "gsk-provider-key")
|
||||
_, _, _, api_base = get_llm_provider(
|
||||
model="groq/whisper-large-v3-turbo", custom_llm_provider=None, api_base=None, api_key=None
|
||||
)
|
||||
|
||||
with respx.mock() as upstream:
|
||||
route: Final = upstream.post(f"{api_base}/audio/transcriptions").respond(
|
||||
200, json={"text": "hello world", "language": "english"}
|
||||
)
|
||||
|
||||
response: Final = litellm.transcription(model="groq/whisper-large-v3-turbo", file=_audio_file())
|
||||
|
||||
assert response.text == "hello world"
|
||||
assert route.calls[0].request.headers["authorization"] == "Bearer gsk-provider-key"
|
||||
Loading…
Add table
Reference in a new issue