mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-15 23:31:29 +00:00
Merge pull request #41300 from BerriAI/litellm_lit_1742_custom_provider_map_router
fix(router): accept custom_provider_map providers before the first completion call
This commit is contained in:
commit
579a30de83
4 changed files with 119 additions and 3 deletions
|
|
@ -238,6 +238,8 @@ def get_llm_provider(
|
|||
if dynamic_api_key is not None and not isinstance(dynamic_api_key, str):
|
||||
raise Exception(f"dynamic_api_key needs to be a string. Got type={type(dynamic_api_key).__name__}")
|
||||
return model, custom_llm_provider, dynamic_api_key, api_base
|
||||
if "/" in model and is_registered_custom_provider(provider_prefix):
|
||||
return model.split("/", 1)[1], provider_prefix, dynamic_api_key, api_base
|
||||
# check if api base is a known openai compatible endpoint
|
||||
if api_base:
|
||||
for endpoint in litellm.openai_compatible_endpoints:
|
||||
|
|
@ -536,6 +538,10 @@ def get_llm_provider(
|
|||
)
|
||||
|
||||
|
||||
def is_registered_custom_provider(custom_llm_provider: str | None) -> bool:
|
||||
return any(item["provider"] == custom_llm_provider for item in litellm.custom_provider_map)
|
||||
|
||||
|
||||
def _dashscope_family_chat_config(custom_llm_provider: str) -> "litellm.DashScopeChatConfig":
|
||||
if custom_llm_provider == "qwencloud":
|
||||
return litellm.QwenCloudChatConfig()
|
||||
|
|
|
|||
|
|
@ -79,7 +79,10 @@ from litellm.litellm_core_utils.core_helpers import (
|
|||
from litellm.litellm_core_utils.coroutine_checker import coroutine_checker
|
||||
from litellm.litellm_core_utils.credential_accessor import CredentialAccessor
|
||||
from litellm.litellm_core_utils.dd_tracing import tracer
|
||||
from litellm.litellm_core_utils.get_llm_provider_logic import declared_authenticating_provider
|
||||
from litellm.litellm_core_utils.get_llm_provider_logic import (
|
||||
declared_authenticating_provider,
|
||||
is_registered_custom_provider,
|
||||
)
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
|
||||
from litellm.litellm_core_utils.ptu_pricing import (
|
||||
PTU_COST_ATTRIBUTION_ENV_VAR,
|
||||
|
|
@ -9563,8 +9566,10 @@ class Router:
|
|||
)
|
||||
# done reading model["litellm_params"]
|
||||
# Check if provider is supported: either in enum or JSON-configured
|
||||
if custom_llm_provider not in litellm.provider_list and not JSONProviderRegistry.exists(
|
||||
custom_llm_provider
|
||||
if (
|
||||
custom_llm_provider not in litellm.provider_list
|
||||
and not JSONProviderRegistry.exists(custom_llm_provider)
|
||||
and not is_registered_custom_provider(custom_llm_provider)
|
||||
):
|
||||
raise Exception(f"Unsupported provider - {custom_llm_provider}")
|
||||
|
||||
|
|
|
|||
|
|
@ -0,0 +1,55 @@
|
|||
from typing import Final
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm import CustomLLM
|
||||
from litellm.litellm_core_utils.get_llm_provider_logic import (
|
||||
get_llm_provider,
|
||||
is_registered_custom_provider,
|
||||
)
|
||||
|
||||
CUSTOM_PROVIDER: Final = "test-onprem-llm"
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def registered_custom_provider(monkeypatch: pytest.MonkeyPatch) -> str:
|
||||
monkeypatch.setattr(litellm, "custom_provider_map", [{"provider": CUSTOM_PROVIDER, "custom_handler": CustomLLM()}])
|
||||
monkeypatch.setattr(litellm, "provider_list", list(litellm.provider_list))
|
||||
monkeypatch.setattr(litellm, "_custom_providers", list(litellm._custom_providers))
|
||||
return CUSTOM_PROVIDER
|
||||
|
||||
|
||||
def test_get_llm_provider_resolves_custom_provider_map_prefix_before_first_completion(
|
||||
registered_custom_provider: str,
|
||||
) -> None:
|
||||
assert registered_custom_provider not in litellm.provider_list
|
||||
|
||||
model, provider, dynamic_api_key, api_base = get_llm_provider(model=f"{registered_custom_provider}/my-model")
|
||||
|
||||
assert (model, provider, dynamic_api_key, api_base) == ("my-model", registered_custom_provider, None, None)
|
||||
|
||||
|
||||
def test_get_llm_provider_strips_prefix_when_custom_provider_passed_explicitly(
|
||||
registered_custom_provider: str,
|
||||
) -> None:
|
||||
model, provider, _, api_base = get_llm_provider(
|
||||
model="my-model",
|
||||
custom_llm_provider=registered_custom_provider,
|
||||
api_base="http://onprem.internal:8080",
|
||||
)
|
||||
|
||||
assert (model, provider, api_base) == ("my-model", registered_custom_provider, "http://onprem.internal:8080")
|
||||
|
||||
|
||||
def test_get_llm_provider_still_rejects_unregistered_prefix(registered_custom_provider: str) -> None:
|
||||
with pytest.raises(litellm.BadRequestError, match="LLM Provider NOT provided"):
|
||||
get_llm_provider(model="not-registered-llm/my-model")
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("candidate", "expected"),
|
||||
[(CUSTOM_PROVIDER, True), ("not-registered-llm", False), (None, False), ("", False)],
|
||||
)
|
||||
def test_is_registered_custom_provider(registered_custom_provider: str, candidate: str | None, expected: bool) -> None:
|
||||
assert is_registered_custom_provider(candidate) is expected
|
||||
|
|
@ -1360,6 +1360,56 @@ def test_add_invalid_provider_to_router():
|
|||
assert router.pattern_router.patterns == {}
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def registered_custom_provider(monkeypatch: pytest.MonkeyPatch) -> str:
|
||||
from litellm import CustomLLM
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
class OnPremLLM(CustomLLM):
|
||||
def completion(self, *args, **kwargs) -> ModelResponse:
|
||||
return litellm.completion(
|
||||
model="gpt-5.6", messages=[{"role": "user", "content": "hi"}], mock_response="served by onprem handler"
|
||||
)
|
||||
|
||||
monkeypatch.setattr(litellm, "custom_provider_map", [{"provider": "test-onprem-llm", "custom_handler": OnPremLLM()}])
|
||||
monkeypatch.setattr(litellm, "provider_list", list(litellm.provider_list))
|
||||
monkeypatch.setattr(litellm, "_custom_providers", list(litellm._custom_providers))
|
||||
return "test-onprem-llm"
|
||||
|
||||
|
||||
def test_router_init_accepts_custom_provider_map_prefix_before_first_completion(registered_custom_provider: str):
|
||||
assert registered_custom_provider not in litellm.provider_list
|
||||
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{"model_name": "onprem", "litellm_params": {"model": f"{registered_custom_provider}/my-model"}},
|
||||
],
|
||||
)
|
||||
|
||||
assert router.get_model_list(model_name="onprem")[0]["litellm_params"]["model"] == (
|
||||
f"{registered_custom_provider}/my-model"
|
||||
)
|
||||
response = router.completion(model="onprem", messages=[{"role": "user", "content": "hi"}])
|
||||
assert response.choices[0].message.content == "served by onprem handler"
|
||||
|
||||
|
||||
def test_router_add_deployment_accepts_explicit_custom_provider_from_custom_provider_map(
|
||||
registered_custom_provider: str,
|
||||
):
|
||||
from litellm.types.router import Deployment
|
||||
|
||||
router = litellm.Router(model_list=[])
|
||||
|
||||
router.add_deployment(
|
||||
Deployment(
|
||||
model_name="onprem",
|
||||
litellm_params={"model": "my-model", "custom_llm_provider": registered_custom_provider},
|
||||
)
|
||||
)
|
||||
|
||||
assert router.get_model_list(model_name="onprem")[0]["litellm_params"]["model"] == "my-model"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_router_ageneric_api_call_with_fallbacks_helper():
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue