mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
perf: add LRU cache for get_llm_provider when called with (model, custom_llm_provider) only
This commit is contained in:
parent
af89c64669
commit
941d89cd9f
2 changed files with 61 additions and 19 deletions
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue