perf: add LRU cache for get_llm_provider when called with (model, custom_llm_provider) only

This commit is contained in:
Alexsander Hamir 2026-01-31 13:50:30 -08:00
parent af89c64669
commit 941d89cd9f
2 changed files with 61 additions and 19 deletions

View file

@ -1,3 +1,4 @@
from functools import lru_cache
from typing import Optional, Tuple
import httpx
@ -98,7 +99,7 @@ def handle_anthropic_text_model_custom_llm_provider(
return model, custom_llm_provider
def get_llm_provider( # noqa: PLR0915
def _get_llm_provider_impl( # noqa: PLR0915
model: str,
custom_llm_provider: Optional[str] = None,
api_base: Optional[str] = None,
@ -106,13 +107,8 @@ def get_llm_provider( # noqa: PLR0915
litellm_params: Optional[LiteLLM_Params] = None,
) -> Tuple[str, str, Optional[str], Optional[str]]:
"""
Returns the provider for a given model name - e.g. 'azure/chatgpt-v-2' -> 'azure'
For router -> Can also give the whole litellm param dict -> this function will extract the relevant details
Raises Error - if unable to map model to a provider
Return model, custom_llm_provider, dynamic_api_key, api_base
Implementation of get_llm_provider. Use get_llm_provider() which adds LRU caching
when called with only (model, custom_llm_provider).
"""
try:
if litellm.LiteLLMProxyChatConfig._should_use_litellm_proxy_by_default(
@ -490,6 +486,47 @@ def get_llm_provider( # noqa: PLR0915
)
@lru_cache(maxsize=1024)
def _get_llm_provider_cached(
model: str, custom_llm_provider: Optional[str]
) -> Tuple[str, str, Optional[str], Optional[str]]:
"""Cached path for get_llm_provider when api_base, api_key, litellm_params are None."""
return _get_llm_provider_impl(
model=model,
custom_llm_provider=custom_llm_provider,
api_base=None,
api_key=None,
litellm_params=None,
)
def get_llm_provider( # noqa: PLR0915
model: str,
custom_llm_provider: Optional[str] = None,
api_base: Optional[str] = None,
api_key: Optional[str] = None,
litellm_params: Optional[LiteLLM_Params] = None,
) -> Tuple[str, str, Optional[str], Optional[str]]:
"""
Returns the provider for a given model name - e.g. 'azure/chatgpt-v-2' -> 'azure'
For router -> Can also give the whole litellm param dict -> this function will extract the relevant details
Raises Error - if unable to map model to a provider
Return model, custom_llm_provider, dynamic_api_key, api_base
"""
if api_base is None and api_key is None and litellm_params is None:
return _get_llm_provider_cached(model, custom_llm_provider)
return _get_llm_provider_impl(
model=model,
custom_llm_provider=custom_llm_provider,
api_base=api_base,
api_key=api_key,
litellm_params=litellm_params,
)
def _get_openai_compatible_provider_info( # noqa: PLR0915
model: str,
api_base: Optional[str],

View file

@ -968,11 +968,22 @@ def function_setup( # noqa: PLR0915
# signatures to ensure compatibility.
if isinstance(messages, list) and len(messages) > 0:
try:
from litellm.litellm_core_utils.get_llm_provider_logic import (
get_llm_provider,
)
from litellm.litellm_core_utils.prompt_templates.factory import (
THOUGHT_SIGNATURE_SEPARATOR,
)
custom_llm_provider = kwargs.get("custom_llm_provider")
if not custom_llm_provider and model:
try:
_, custom_llm_provider, _, _ = get_llm_provider(
model=model,
custom_llm_provider=custom_llm_provider,
)
except Exception:
pass
if not _is_gemini_model(model, custom_llm_provider):
verbose_logger.debug(
@ -1858,7 +1869,10 @@ def client(original_function): # noqa: PLR0915
end_time=end_time,
)
_update_response_metadata(
update_response_metadata = getattr(
sys.modules[__name__], "update_response_metadata"
)
update_response_metadata(
result=result,
logging_obj=logging_obj,
model=model,
@ -5682,16 +5696,7 @@ def get_model_info(model: str, custom_llm_provider: Optional[str] = None) -> Mod
custom_llm_provider=custom_llm_provider,
)
provider_info = get_provider_info(
model=model, custom_llm_provider=custom_llm_provider
)
if provider_info:
for key, value in provider_info.items():
if value is not None:
_model_info[key] = value # type: ignore
if verbose_logger.isEnabledFor(logging.DEBUG):
verbose_logger.debug(f"model_info: {_model_info}")
verbose_logger.debug(f"model_info: {_model_info}")
returned_model_info = ModelInfo(
**_model_info, supported_openai_params=supported_openai_params