Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/spend_tracking/spend_counter_batch.py: 47%

141 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-10 12:01 +0000

1"""One Redis MGET per phase (admission, reservation, post-call) for the spend counters it reads, not one GET each.""" 

2 

3import asyncio 

4from collections.abc import Iterator, Mapping, Sequence 

5from contextvars import ContextVar, Token 

6from dataclasses import dataclass 

7from types import MappingProxyType, TracebackType 

8from typing import Final 

9 

10from pydantic import TypeAdapter 

11 

12from litellm._logging import verbose_proxy_logger 

13from litellm.caching.redis_cache import RedisCache 

14from litellm.proxy._types import UserAPIKeyAuth 

15from litellm.proxy.common_utils.user_api_key_cache import ( 

16 model_access_group_spend_counter_key, 

17 project_spend_counter_key, 

18) 

19 

20_CounterValues: Final = TypeAdapter(dict[str, float | None]) 

21_NO_VALUES: Final[Mapping[str, float | None]] = MappingProxyType({}) 

22 

23 

24@dataclass(frozen=True, slots=True) 

25class PendingSpendIncrement: 

26 counter_key: str 

27 increment: float 

28 

29 

30class SpendCounterBatch: 

31 """Bound counters are read with one MGET on first use; counters bound later join the next MGET. 

32 ``async_batch_get_cache`` maps a clean miss to ``None`` and drops keys only when Redis failed, so an absent 

33 key means "read it yourself" and a present ``None`` is an authoritative miss.""" 

34 

35 __slots__ = ("_fetched", "_keys", "_loaded", "_lock", "_open", "_redis_cache") 

36 

37 def __init__(self, redis_cache: RedisCache) -> None: 

38 self._redis_cache: Final = redis_cache 

39 self._lock: Final = asyncio.Lock() 

40 self._open = True 

41 self._keys: frozenset[str] = frozenset() 

42 self._fetched: frozenset[str] = frozenset() 

43 self._loaded: Mapping[str, float | None] = _NO_VALUES 

44 

45 @property 

46 def counter_keys(self) -> frozenset[str]: 

47 return self._keys 

48 

49 @property 

50 def is_open(self) -> bool: 

51 return self._open 

52 

53 def bind(self, counter_keys: frozenset[str]) -> None: 

54 if self._open: 

55 self._keys = self._keys | counter_keys 

56 

57 def close(self) -> None: 

58 """Later reads go to Redis directly; call before any read-then-write on the counters.""" 

59 self._open = False 

60 

61 async def read(self, counter_key: str) -> tuple[float | None, bool] | None: 

62 """(value, authoritative) for a bound counter, None when the caller must read Redis itself.""" 

63 if not self._open or counter_key not in self._keys: 

64 return None 

65 loaded: Final = await self._load() 

66 if counter_key not in loaded: 

67 return None 

68 return loaded[counter_key], True 

69 

70 def record(self, counter_key: str, value: float) -> None: 

71 """A write returned the counter's new value; later reads in this scope see it instead of the MGET value.""" 

72 if not self._open: 

73 return 

74 key: Final = frozenset((counter_key,)) 

75 self._keys = self._keys | key 

76 self._fetched = self._fetched | key 

77 self._loaded = MappingProxyType({**self._loaded, counter_key: value}) 

78 

79 def forget(self, counter_key: str) -> None: 

80 """A write left the counter's value unknown; later reads in this scope go to Redis.""" 

81 key: Final = frozenset((counter_key,)) 

82 self._keys = self._keys - key 

83 self._fetched = self._fetched - key 

84 self._loaded = MappingProxyType({k: v for k, v in self._loaded.items() if k != counter_key}) 

85 

86 async def _load(self) -> Mapping[str, float | None]: 

87 async with self._lock: 

88 pending: Final = self._keys - self._fetched 

89 if pending: 

90 self._fetched = self._fetched | pending 

91 fetched: Final = await self._fetch(pending) 

92 self._loaded = MappingProxyType({**fetched, **self._loaded}) 

93 return self._loaded 

94 

95 async def _fetch(self, keys: frozenset[str]) -> Mapping[str, float | None]: 

96 try: 

97 return _CounterValues.validate_python( 

98 await self._redis_cache.async_batch_get_cache(key_list=sorted(keys)) # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # untyped cache API 

99 ) 

100 except Exception as e: # noqa: BLE001 # per-key reads take over and apply their own Redis fallback 

101 verbose_proxy_logger.debug("spend counter batch read failed, falling back to per-key reads: %s", e) 

102 return _NO_VALUES 

103 

104 

105_active_batch: Final[ContextVar[SpendCounterBatch | None]] = ContextVar("spend_counter_batch", default=None) 

106 

107 

108def active_spend_counter_batch() -> SpendCounterBatch | None: 

109 return _active_batch.get() 

110 

111 

112class spend_counter_batch_scope: 

113 """Reads inside the scope share one MGET for the keys bound here or by ``bind_*`` calls inside it. 

114 Opened inside a scope whose batch is still open, it binds into that batch so both phases share the MGET.""" 

115 

116 __slots__ = ("_counter_keys", "_redis_cache", "_token") 

117 

118 def __init__(self, redis_cache: RedisCache | None, counter_keys: frozenset[str] = frozenset()) -> None: 

119 self._redis_cache: Final = redis_cache 

120 self._counter_keys: Final = counter_keys 

121 self._token: Token[SpendCounterBatch | None] | None = None 

122 

123 def __enter__(self) -> None: 

124 if self._redis_cache is None: 124 ↛ 126line 124 didn't jump to line 126 because the condition on line 124 was always true

125 return 

126 outer: Final = _active_batch.get() 

127 if outer is not None and outer.is_open: 

128 outer.bind(self._counter_keys) 

129 return 

130 batch: Final = SpendCounterBatch(self._redis_cache) 

131 batch.bind(self._counter_keys) 

132 self._token = _active_batch.set(batch) 

133 

134 def __exit__( 

135 self, exc_type: type[BaseException] | None, exc: BaseException | None, tb: TracebackType | None 

136 ) -> None: 

137 if self._token is not None: 137 ↛ 138line 137 didn't jump to line 138 because the condition on line 137 was never true

138 _active_batch.reset(self._token) 

139 

140 

141def release_spend_counter_batch() -> None: 

142 batch: Final = _active_batch.get() 

143 if batch is not None: 143 ↛ 144line 143 didn't jump to line 144 because the condition on line 143 was never true

144 batch.close() 

145 

146 

147def _iter_admission_counter_keys(token: UserAPIKeyAuth, end_user_id: str | None) -> Iterator[str]: 

148 if token.token is not None: 148 ↛ 150line 148 didn't jump to line 150 because the condition on line 148 was always true

149 yield f"spend:key:{token.token}" 

150 if token.team_id is not None: 150 ↛ 151line 150 didn't jump to line 151 because the condition on line 150 was never true

151 yield f"spend:team:{token.team_id}" 

152 if token.user_id is not None: 

153 yield f"spend:team_member:{token.user_id}:{token.team_id}" 

154 if token.user_id is not None: 

155 yield f"spend:user:{token.user_id}" 

156 if end_user_id is not None: 

157 yield f"spend:end_user:{end_user_id}" 

158 if token.org_id is not None: 158 ↛ 159line 158 didn't jump to line 159 because the condition on line 158 was never true

159 yield f"spend:org:{token.org_id}" 

160 if token.project_id is not None: 160 ↛ 161line 160 didn't jump to line 161 because the condition on line 160 was never true

161 yield project_spend_counter_key(token.project_id) 

162 

163 

164def admission_counter_keys(token: UserAPIKeyAuth, end_user_id: str | None) -> frozenset[str]: 

165 return frozenset(_iter_admission_counter_keys(token, end_user_id)) 

166 

167 

168def post_call_counter_keys( 

169 token: str | None, 

170 team_id: str | None, 

171 user_id: str | None, 

172 org_id: str | None, 

173 end_user_id: str | None, 

174 tags: Sequence[object] | None, 

175 model_access_groups: Sequence[object] | None, 

176 project_id: str | None = None, 

177) -> frozenset[str]: 

178 """Every counter ``increment_spend_counters`` warm-checks, except budget windows which bind on read.""" 

179 entity_keys: Final = admission_counter_keys( 

180 UserAPIKeyAuth(token=token, team_id=team_id, user_id=user_id, org_id=org_id, project_id=project_id), 

181 end_user_id, 

182 ) 

183 tag_keys: Final = frozenset(f"spend:tag:{tag}" for tag in tags or () if tag and isinstance(tag, str)) 

184 group_keys: Final = frozenset( 

185 model_access_group_spend_counter_key(group) 

186 for group in model_access_groups or () 

187 if group and isinstance(group, str) 

188 ) 

189 return entity_keys | tag_keys | group_keys 

190 

191 

192def bind_admission_counter_keys(token: UserAPIKeyAuth, end_user_id: str | None) -> None: 

193 """Idempotent: call again after the token gains ids (end user, team org) so those counters join the MGET.""" 

194 bind_spend_counter_keys(admission_counter_keys(token, end_user_id)) 

195 

196 

197def bind_spend_counter_keys(counter_keys: frozenset[str]) -> None: 

198 batch: Final = _active_batch.get() 

199 if batch is None: 199 ↛ 201line 199 didn't jump to line 201 because the condition on line 199 was always true

200 return 

201 batch.bind(counter_keys) 

202 

203 

204def record_spend_counter_value(counter_key: str, value: float) -> None: 

205 batch: Final = _active_batch.get() 

206 if batch is None: 

207 return 

208 batch.record(counter_key, value) 

209 

210 

211def forget_spend_counter(counter_key: str) -> None: 

212 batch: Final = _active_batch.get() 

213 if batch is None: 

214 return 

215 batch.forget(counter_key) 

216 

217 

218async def read_batched_spend_counter(counter_key: str) -> tuple[float | None, bool] | None: 

219 """Bind-on-read for counters only known at read time (budget windows); the first reader pays the MGET.""" 

220 batch: Final = _active_batch.get() 

221 if batch is None: 

222 return None 

223 batch.bind(frozenset((counter_key,))) 

224 return await batch.read(counter_key)