[Refactor] litellm/init.py: lazy load LLMClientCache (#18008)

This commit is contained in:
Alexsander Hamir 2025-12-16 05:44:06 -08:00 • committed by GitHub
parent e31d8a1cf6
commit 6ca812130b
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 112 additions and 40 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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