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
« 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."""
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
10from pydantic import TypeAdapter
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)
20_CounterValues: Final = TypeAdapter(dict[str, float | None])
21_NO_VALUES: Final[Mapping[str, float | None]] = MappingProxyType({})
24@dataclass(frozen=True, slots=True)
25class PendingSpendIncrement:
26 counter_key: str
27 increment: float
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."""
35 __slots__ = ("_fetched", "_keys", "_loaded", "_lock", "_open", "_redis_cache")
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
45 @property
46 def counter_keys(self) -> frozenset[str]:
47 return self._keys
49 @property
50 def is_open(self) -> bool:
51 return self._open
53 def bind(self, counter_keys: frozenset[str]) -> None:
54 if self._open:
55 self._keys = self._keys | counter_keys
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
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
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})
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})
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
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
105_active_batch: Final[ContextVar[SpendCounterBatch | None]] = ContextVar("spend_counter_batch", default=None)
108def active_spend_counter_batch() -> SpendCounterBatch | None:
109 return _active_batch.get()
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."""
116 __slots__ = ("_counter_keys", "_redis_cache", "_token")
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
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)
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)
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()
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)
164def admission_counter_keys(token: UserAPIKeyAuth, end_user_id: str | None) -> frozenset[str]:
165 return frozenset(_iter_admission_counter_keys(token, end_user_id))
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
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))
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)
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)
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)
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)