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

1from __future__ import annotations 

2 

3import asyncio 

4import time 

5from collections.abc import Callable, Mapping 

6from datetime import datetime 

7from typing import TYPE_CHECKING, Final, Literal 

8 

9import httpx 

10from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, ValidationError 

11 

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 

16 

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 

19 

20_RETENTION_SECONDS: Final = 86_400 

21 

22 

23class CacheObservation(BaseModel): 

24 model_config = ConfigDict(extra="forbid", frozen=True, strict=True) 

25 

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) 

30 

31 

32_CACHE_ENTRY: Final[TypeAdapter[CacheObservation | str | None]] = TypeAdapter(CacheObservation | str | None) 

33 

34 

35def _cache_key(scope: str, fingerprint: str) -> str: 

36 return f"prompt-cache-observation:{scope}:{fingerprint}" 

37 

38 

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 ) 

54 

55 

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 

65 

66 

67class _Metadata(BaseModel): 

68 model_config = ConfigDict(strict=True) 

69 user_api_key_hash: str = Field(min_length=1) 

70 

71 

72class _Logged(BaseModel): 

73 model_config = ConfigDict(strict=True) 

74 status: Literal["success"] 

75 model_id: str = Field(min_length=1) 

76 metadata: _Metadata 

77 

78 

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 

89 

90 

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 

96 

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