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:
Yassin Kortam 2026-09-15 15:03:42 -07:00 committed by GitHub
commit 579a30de83
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 119 additions and 3 deletions

View file

@ -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()

View file

@ -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}")

View file

@ -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

View file

@ -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():
"""