mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
[Refactor] litellm/init.py: lazy load LLMClientCache (#18008)
This commit is contained in:
parent
e31d8a1cf6
commit
6ca812130b
6 changed files with 112 additions and 40 deletions
|
|
@ -26,7 +26,6 @@ from typing import (
|
|||
)
|
||||
from litellm.types.integrations.datadog_llm_obs import DatadogLLMObsInitParams
|
||||
from litellm.types.integrations.datadog import DatadogInitParams
|
||||
from litellm.caching.llm_caching_handler import LLMClientCache
|
||||
from litellm.types.llms.bedrock import COHERE_EMBEDDING_INPUT_TYPES
|
||||
from litellm.types.utils import (
|
||||
ImageObject,
|
||||
|
|
@ -285,7 +284,7 @@ disable_token_counter: bool = False
|
|||
disable_add_transform_inline_image_block: bool = False
|
||||
disable_add_user_agent_to_request_tags: bool = False
|
||||
extra_spend_tag_headers: Optional[List[str]] = None
|
||||
in_memory_llm_clients_cache: LLMClientCache = LLMClientCache()
|
||||
in_memory_llm_clients_cache: "LLMClientCache"
|
||||
safe_memory_mode: bool = False
|
||||
enable_azure_ad_token_refresh: Optional[bool] = False
|
||||
### DEFAULT AZURE API VERSION ###
|
||||
|
|
@ -1515,6 +1514,7 @@ if TYPE_CHECKING:
|
|||
from litellm.types.utils import ModelInfo as _ModelInfoType
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
|
||||
from litellm.caching.caching import Cache
|
||||
from litellm.caching.llm_caching_handler import LLMClientCache
|
||||
|
||||
# Cost calculator functions
|
||||
cost_per_token: Callable[..., Tuple[float, float]]
|
||||
|
|
@ -1567,6 +1567,7 @@ def __getattr__(name: str) -> Any:
|
|||
LITELLM_LOGGING_NAMES,
|
||||
UTILS_NAMES,
|
||||
TOKEN_COUNTER_NAMES,
|
||||
LLM_CLIENT_CACHE_NAMES,
|
||||
CACHING_NAMES,
|
||||
HTTP_HANDLER_NAMES,
|
||||
)
|
||||
|
|
@ -1591,6 +1592,11 @@ def __getattr__(name: str) -> Any:
|
|||
from ._lazy_imports import _lazy_import_token_counter
|
||||
return _lazy_import_token_counter(name)
|
||||
|
||||
# Lazy load LLM client cache and its singleton
|
||||
if name in LLM_CLIENT_CACHE_NAMES:
|
||||
from ._lazy_imports import _lazy_import_llm_client_cache
|
||||
return _lazy_import_llm_client_cache(name)
|
||||
|
||||
# Lazy load caching classes
|
||||
if name in CACHING_NAMES:
|
||||
from ._lazy_imports import _lazy_import_caching
|
||||
|
|
|
|||
|
|
@ -39,6 +39,12 @@ TOKEN_COUNTER_NAMES = (
|
|||
"get_modified_max_tokens",
|
||||
)
|
||||
|
||||
# LLM client cache names that support lazy loading via _lazy_import_llm_client_cache
|
||||
LLM_CLIENT_CACHE_NAMES = (
|
||||
"LLMClientCache",
|
||||
"in_memory_llm_clients_cache",
|
||||
)
|
||||
|
||||
# Caching / cache classes that support lazy loading via _lazy_import_caching
|
||||
CACHING_NAMES = (
|
||||
"Cache",
|
||||
|
|
@ -330,6 +336,28 @@ def _lazy_import_caching(name: str) -> Any:
|
|||
raise AttributeError(f"Caching lazy import: unknown attribute {name!r}")
|
||||
|
||||
|
||||
def _lazy_import_llm_client_cache(name: str) -> Any:
|
||||
"""Lazy import for LLM client cache class and singleton."""
|
||||
_globals = _get_litellm_globals()
|
||||
|
||||
if name == "LLMClientCache":
|
||||
from litellm.caching.llm_caching_handler import LLMClientCache as _LLMClientCache
|
||||
|
||||
_globals["LLMClientCache"] = _LLMClientCache
|
||||
return _LLMClientCache
|
||||
|
||||
if name == "in_memory_llm_clients_cache":
|
||||
from litellm.caching.llm_caching_handler import LLMClientCache as _LLMClientCache
|
||||
|
||||
instance = _LLMClientCache()
|
||||
# Only populate the requested singleton name to keep lazy-import
|
||||
# semantics consistent with other helpers (no extra symbols).
|
||||
_globals["in_memory_llm_clients_cache"] = instance
|
||||
return instance
|
||||
|
||||
raise AttributeError(f"LLM client cache lazy import: unknown attribute {name!r}")
|
||||
|
||||
|
||||
def _lazy_import_litellm_logging(name: str) -> Any:
|
||||
"""Lazy import for litellm_logging module."""
|
||||
_globals = _get_litellm_globals()
|
||||
|
|
|
|||
|
|
@ -6,12 +6,11 @@ This module provides fake streaming by converting non-streaming responses into s
|
|||
"""
|
||||
|
||||
import asyncio
|
||||
from typing import Any, AsyncIterator, Dict
|
||||
from typing import Any, AsyncIterator, Dict, cast
|
||||
from uuid import uuid4
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, get_async_httpx_client
|
||||
|
||||
|
||||
class PydanticAITransformation:
|
||||
|
|
@ -78,7 +77,7 @@ class PydanticAITransformation:
|
|||
|
||||
@staticmethod
|
||||
async def _poll_for_completion(
|
||||
client: httpx.AsyncClient,
|
||||
client: AsyncHTTPHandler,
|
||||
endpoint: str,
|
||||
task_id: str,
|
||||
request_id: str,
|
||||
|
|
@ -179,34 +178,37 @@ class PydanticAITransformation:
|
|||
f"Pydantic AI: Sending non-streaming request to {endpoint}"
|
||||
)
|
||||
|
||||
# Send request to Pydantic AI agent
|
||||
async with httpx.AsyncClient(timeout=timeout) as client:
|
||||
response = await client.post(
|
||||
endpoint,
|
||||
json=a2a_request,
|
||||
headers={"Content-Type": "application/json"},
|
||||
)
|
||||
response.raise_for_status()
|
||||
response_data = response.json()
|
||||
|
||||
# Check if task is already completed
|
||||
result = response_data.get("result", {})
|
||||
status = result.get("status", {})
|
||||
state = status.get("state", "")
|
||||
|
||||
if state != "completed":
|
||||
# Need to poll for completion
|
||||
task_id = result.get("id")
|
||||
if task_id:
|
||||
verbose_logger.info(
|
||||
f"Pydantic AI: Task {task_id} submitted, polling for completion..."
|
||||
)
|
||||
response_data = await PydanticAITransformation._poll_for_completion(
|
||||
client=client,
|
||||
endpoint=endpoint,
|
||||
task_id=task_id,
|
||||
request_id=request_id,
|
||||
)
|
||||
# Send request to Pydantic AI agent using shared async HTTP client
|
||||
client = get_async_httpx_client(
|
||||
llm_provider=cast(Any, "pydantic_ai_agent"),
|
||||
params={"timeout": timeout},
|
||||
)
|
||||
response = await client.post(
|
||||
endpoint,
|
||||
json=a2a_request,
|
||||
headers={"Content-Type": "application/json"},
|
||||
)
|
||||
response.raise_for_status()
|
||||
response_data = response.json()
|
||||
|
||||
# Check if task is already completed
|
||||
result = response_data.get("result", {})
|
||||
status = result.get("status", {})
|
||||
state = status.get("state", "")
|
||||
|
||||
if state != "completed":
|
||||
# Need to poll for completion
|
||||
task_id = result.get("id")
|
||||
if task_id:
|
||||
verbose_logger.info(
|
||||
f"Pydantic AI: Task {task_id} submitted, polling for completion..."
|
||||
)
|
||||
response_data = await PydanticAITransformation._poll_for_completion(
|
||||
client=client,
|
||||
endpoint=endpoint,
|
||||
task_id=task_id,
|
||||
request_id=request_id,
|
||||
)
|
||||
|
||||
verbose_logger.info(f"Pydantic AI: Received completed response for request_id={request_id}")
|
||||
|
||||
|
|
|
|||
|
|
@ -16,7 +16,6 @@ from typing import (
|
|||
from pydantic import BaseModel
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH_SENSITIVE_DATA_MASKER
|
||||
from litellm.types.integrations.argilla import ArgillaItem
|
||||
from litellm.types.llms.openai import AllMessageValues, ChatCompletionRequest
|
||||
|
|
@ -33,6 +32,7 @@ from litellm.types.utils import (
|
|||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.caching.caching import DualCache
|
||||
from opentelemetry.trace import Span as _Span
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
|
|
@ -334,7 +334,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
|
|||
async def async_pre_call_hook(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
cache: DualCache,
|
||||
cache: "DualCache",
|
||||
data: dict,
|
||||
call_type: CallTypesLiteral,
|
||||
) -> Optional[
|
||||
|
|
|
|||
|
|
@ -1153,7 +1153,17 @@ def get_async_httpx_client(
|
|||
pass
|
||||
|
||||
_cache_key_name = "async_httpx_client" + _params_key_name + llm_provider
|
||||
_cached_client = litellm.in_memory_llm_clients_cache.get_cache(_cache_key_name)
|
||||
|
||||
# Lazily initialize the global in-memory client cache to avoid relying on
|
||||
# litellm globals being fully populated during import time.
|
||||
cache = getattr(litellm, "in_memory_llm_clients_cache", None)
|
||||
if cache is None:
|
||||
from litellm.caching.llm_caching_handler import LLMClientCache
|
||||
|
||||
cache = LLMClientCache()
|
||||
setattr(litellm, "in_memory_llm_clients_cache", cache)
|
||||
|
||||
_cached_client = cache.get_cache(_cache_key_name)
|
||||
if _cached_client:
|
||||
return _cached_client
|
||||
|
||||
|
|
@ -1166,7 +1176,7 @@ def get_async_httpx_client(
|
|||
shared_session=shared_session,
|
||||
)
|
||||
|
||||
litellm.in_memory_llm_clients_cache.set_cache(
|
||||
cache.set_cache(
|
||||
key=_cache_key_name,
|
||||
value=_new_client,
|
||||
ttl=_DEFAULT_TTL_FOR_HTTPX_CLIENTS,
|
||||
|
|
@ -1191,7 +1201,16 @@ def _get_httpx_client(params: Optional[dict] = None) -> HTTPHandler:
|
|||
|
||||
_cache_key_name = "httpx_client" + _params_key_name
|
||||
|
||||
_cached_client = litellm.in_memory_llm_clients_cache.get_cache(_cache_key_name)
|
||||
# Lazily initialize the global in-memory client cache to avoid relying on
|
||||
# litellm globals being fully populated during import time.
|
||||
cache = getattr(litellm, "in_memory_llm_clients_cache", None)
|
||||
if cache is None:
|
||||
from litellm.caching.llm_caching_handler import LLMClientCache
|
||||
|
||||
cache = LLMClientCache()
|
||||
setattr(litellm, "in_memory_llm_clients_cache", cache)
|
||||
|
||||
_cached_client = cache.get_cache(_cache_key_name)
|
||||
if _cached_client:
|
||||
return _cached_client
|
||||
|
||||
|
|
@ -1200,7 +1219,7 @@ def _get_httpx_client(params: Optional[dict] = None) -> HTTPHandler:
|
|||
else:
|
||||
_new_client = HTTPHandler(timeout=httpx.Timeout(timeout=600.0, connect=5.0))
|
||||
|
||||
litellm.in_memory_llm_clients_cache.set_cache(
|
||||
cache.set_cache(
|
||||
key=_cache_key_name,
|
||||
value=_new_client,
|
||||
ttl=_DEFAULT_TTL_FOR_HTTPX_CLIENTS,
|
||||
|
|
|
|||
|
|
@ -14,12 +14,14 @@ from litellm._lazy_imports import (
|
|||
UTILS_NAMES,
|
||||
TOKEN_COUNTER_NAMES,
|
||||
CACHING_NAMES,
|
||||
LLM_CLIENT_CACHE_NAMES,
|
||||
HTTP_HANDLER_NAMES,
|
||||
_lazy_import_cost_calculator,
|
||||
_lazy_import_litellm_logging,
|
||||
_lazy_import_utils,
|
||||
_lazy_import_token_counter,
|
||||
_lazy_import_caching,
|
||||
_lazy_import_llm_client_cache,
|
||||
_lazy_import_http_handlers,
|
||||
)
|
||||
|
||||
|
|
@ -111,6 +113,18 @@ def test_token_counter_lazy_imports():
|
|||
_verify_only_requested_name_imported(name, TOKEN_COUNTER_NAMES)
|
||||
|
||||
|
||||
def test_llm_client_cache_lazy_imports():
|
||||
"""Test that LLM client cache class and singleton can be lazy imported."""
|
||||
for name in LLM_CLIENT_CACHE_NAMES:
|
||||
_clear_names_from_globals(LLM_CLIENT_CACHE_NAMES)
|
||||
|
||||
obj = _lazy_import_llm_client_cache(name)
|
||||
assert obj is not None
|
||||
assert name in litellm.__dict__
|
||||
|
||||
_verify_only_requested_name_imported(name, LLM_CLIENT_CACHE_NAMES)
|
||||
|
||||
|
||||
def test_http_handler_lazy_imports():
|
||||
"""Test that HTTP handler singletons can be lazy imported."""
|
||||
for name in HTTP_HANDLER_NAMES:
|
||||
|
|
@ -140,3 +154,6 @@ def test_unknown_attribute_raises_error():
|
|||
with pytest.raises(AttributeError):
|
||||
_lazy_import_token_counter("unknown")
|
||||
|
||||
with pytest.raises(AttributeError):
|
||||
_lazy_import_llm_client_cache("unknown")
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue