Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/hooks/prompt_cache_prediction.py: 46%
94 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1from __future__ import annotations
3import asyncio
4import time
5from collections.abc import Callable, Mapping
6from datetime import datetime
7from typing import TYPE_CHECKING, Final, Literal
9import httpx
10from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, ValidationError
12from litellm.caching.dual_cache import DualCache
13from litellm.integrations.custom_logger import CustomLogger
14from litellm.llms.anthropic.prompt_cache_prediction import PromptPrefix, parse_observed_cache
15from litellm.types.utils import ModelResponse
17if TYPE_CHECKING: 17 ↛ 18line 17 didn't jump to line 18 because the condition on line 17 was never true
18 from litellm.proxy.utils import InternalUsageCache
20_RETENTION_SECONDS: Final = 86_400
23class CacheObservation(BaseModel):
24 model_config = ConfigDict(extra="forbid", frozen=True, strict=True)
26 fingerprint: str = Field(pattern=r"^[0-9a-f]{64}$")
27 cached_tokens: int = Field(gt=0)
28 observed_at: float = Field(ge=0, allow_inf_nan=False)
29 expires_at: float = Field(ge=0, allow_inf_nan=False)
32_CACHE_ENTRY: Final[TypeAdapter[CacheObservation | str | None]] = TypeAdapter(CacheObservation | str | None)
35def _cache_key(scope: str, fingerprint: str) -> str:
36 return f"prompt-cache-observation:{scope}:{fingerprint}"
39async def lookup(
40 cache: DualCache, scope: str, prefix: PromptPrefix, now: float | None = None
41) -> CacheObservation | None:
42 checked_at: Final = time.time() if now is None else now
43 exact: Final = await _read_exact(cache, scope, prefix.fingerprint)
44 if exact is not None and exact.expires_at > checked_at:
45 return exact
46 older: Final = await asyncio.gather(
47 *(_read_exact(cache, scope, fingerprint) for fingerprint in prefix.fingerprints[1:])
48 )
49 observations: Final = tuple(observation for observation in (exact, *older) if observation is not None)
50 return next(
51 (observation for observation in observations if observation.expires_at > checked_at),
52 next(iter(observations), None),
53 )
56async def _read_exact(cache: DualCache, scope: str, fingerprint: str) -> CacheObservation | None:
57 try:
58 value: Final = _CACHE_ENTRY.validate_python(await cache.async_get_cache(_cache_key(scope, fingerprint), ttl=1)) # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # validate the legacy cache's untyped result at the I/O boundary
59 if value is None:
60 return None
61 observation: Final = CacheObservation.model_validate_json(value) if isinstance(value, str) else value
62 except ValidationError:
63 return None
64 return observation if observation.fingerprint == fingerprint else None
67class _Metadata(BaseModel):
68 model_config = ConfigDict(strict=True)
69 user_api_key_hash: str = Field(min_length=1)
72class _Logged(BaseModel):
73 model_config = ConfigDict(strict=True)
74 status: Literal["success"]
75 model_id: str = Field(min_length=1)
76 metadata: _Metadata
79class _Event(BaseModel):
80 model_config = ConfigDict(strict=True, arbitrary_types_allowed=True)
81 call_type: Literal["anthropic_messages"]
82 custom_llm_provider: Literal["anthropic"]
83 cache_hit: bool | None = None
84 httpx_response: httpx.Response
85 first_api_call_start_time: datetime
86 standard_logging_object: _Logged
87 stream: bool = False
88 prompt_cache_response_complete: bool = False
91class PromptCacheObserver(CustomLogger):
92 def __init__(self, internal_usage_cache: InternalUsageCache, clock: Callable[[], float] = time.time) -> None:
93 super().__init__() # pyright: ignore[reportUnknownMemberType] # base callback constructor accepts untyped kwargs
94 self.cache = internal_usage_cache.dual_cache
95 self.clock = clock
97 async def async_log_success_event(
98 self, kwargs: Mapping[str, object], response_obj: object, start_time: datetime, end_time: datetime
99 ) -> None:
100 if not isinstance(response_obj, ModelResponse): 100 ↛ 102line 100 didn't jump to line 102 because the condition on line 100 was always true
101 return
102 try:
103 event: Final = _Event.model_validate(kwargs)
104 wire: Final = event.httpx_response.request
105 except (ValidationError, RuntimeError, httpx.RequestNotRead):
106 return
107 if (
108 event.cache_hit
109 or event.httpx_response.status_code != 200
110 or (event.stream and not event.prompt_cache_response_complete)
111 ):
112 return
113 observed: Final = parse_observed_cache(
114 wire,
115 response_obj,
116 event.standard_logging_object.metadata.user_api_key_hash,
117 event.standard_logging_object.model_id,
118 )
119 if observed is None:
120 return
121 prefix: Final = observed.prefix
122 scope: Final = observed.scope
123 cache_tokens: Final = observed.cached_tokens
124 now: Final = self.clock()
125 started: Final = event.first_api_call_start_time.timestamp()
126 if started > now:
127 return
128 if observed.cache_creation_tokens == 0:
129 previous: Final = await _read_exact(self.cache, scope, prefix.fingerprint)
130 if previous is None or previous.fingerprint != prefix.fingerprint or previous.cached_tokens != cache_tokens:
131 return
132 observation: Final = CacheObservation(
133 fingerprint=prefix.fingerprint,
134 cached_tokens=cache_tokens,
135 observed_at=now,
136 expires_at=started + prefix.ttl_seconds,
137 )
138 key: Final = _cache_key(scope, prefix.fingerprint)
139 payload: Final = observation.model_dump_json()
140 await self.cache.async_set_cache(key, payload, ttl=_RETENTION_SECONDS) # pyright: ignore[reportUnknownMemberType] # legacy cache accepts a serialized validated observation
141 if self.cache.redis_cache is not None:
142 await self.cache.async_set_cache(key, payload, local_only=True, ttl=1) # pyright: ignore[reportUnknownMemberType] # keep the local copy short-lived while Redis retains stale evidence