Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/db/gateway_request_tracking.py: 52%

113 statements  

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

1""" 

2Accumulates gateway request counts (SGR) recorded at the ASGI edge and commits 

3them to ``LiteLLM_DailyGatewayRequests``. 

4 

5Unlike the spend queues this keeps no per-request item. A count is a pure 

6aggregate, so requests fold into an in-memory map as they finish. Every 

7dimension of the key is server-chosen and drawn from a fixed set: the date, the 

8category, and a route that the classifier maps to one of a closed list of 

9strings rather than passing the raw path through. Nothing a caller sends can 

10add a key, so the fold and the table it commits to are bounded by (days x 

11routes) however much traffic arrives, and the response path carries no 

12unbounded queue that would block once full. 

13 

14A flush commits its whole snapshot as one multi-row ``INSERT ... ON CONFLICT DO 

15UPDATE`` rather than one upsert per key, so a worker costs the primary one 

16statement per interval however many routes it served. With 

17``use_redis_transaction_buffer`` on, workers instead push their snapshot to a 

18Redis list and one lock-holding pod folds every entry and writes the table, so 

19the deployment as a whole costs the primary one statement per interval. 

20""" 

21 

22import json 

23from collections.abc import AsyncIterator, Iterable 

24from datetime import datetime, timezone 

25from itertools import chain 

26from types import MappingProxyType 

27from typing import TYPE_CHECKING, Final, TypeAlias 

28 

29from pydantic import TypeAdapter 

30 

31from litellm._logging import verbose_proxy_logger 

32from litellm.caching import RedisCache 

33from litellm.constants import MAX_REDIS_BUFFER_DEQUEUE_COUNT, REDIS_GATEWAY_REQUESTS_BUFFER_KEY 

34from litellm.proxy.db.db_transaction_queue.pod_lock_manager import PodLockManager 

35from litellm.proxy.middleware.billable_request_metrics_middleware import BillableCategory 

36from litellm.types.proxy.gateway_requests import ( 

37 GatewayRequestCounts, 

38 GatewayRequestKey, 

39 GatewayRequestSnapshot, 

40) 

41 

42if TYPE_CHECKING: 42 ↛ 43line 42 didn't jump to line 43 because the condition on line 42 was never true

43 from litellm.proxy.utils import PrismaClient 

44 

45_EMPTY: Final = GatewayRequestCounts(successful_requests=0, failed_requests=0) 

46_TABLE: Final = '"LiteLLM_DailyGatewayRequests"' 

47_COLUMNS_PER_ROW: Final = 5 

48_UTC_NOW: Final = "(NOW() AT TIME ZONE 'UTC')" 

49GATEWAY_REQUESTS_JOB_NAME: Final = "update_gateway_requests_job" 

50 

51_BufferedRows: TypeAlias = tuple[tuple[str, str, str, int, int], ...] 

52_BUFFERED_ROWS: Final = TypeAdapter(_BufferedRows) 

53_BUFFERED_ENTRIES: Final = TypeAdapter(tuple[str | bytes, ...]) 

54_NO_COUNTS: Final[GatewayRequestSnapshot] = MappingProxyType({}) 

55 

56 

57def _utc_date() -> str: 

58 return datetime.now(timezone.utc).strftime("%Y-%m-%d") 

59 

60 

61class GatewayRequestAccumulator: 

62 """Sink for the request-metrics middleware. ``record`` is sync and never awaits.""" 

63 

64 def __init__(self) -> None: 

65 self._counts: dict[GatewayRequestKey, GatewayRequestCounts] = {} # mutable-ok: bounded fold, drained per flush 

66 

67 def record(self, *, category: BillableCategory, route: str, status_code: int) -> None: 

68 key: Final = GatewayRequestKey(date=_utc_date(), category=category.value, route=route) 

69 self._counts[key] = self._counts.get(key, _EMPTY).plus(succeeded=200 <= status_code < 300) 

70 

71 def drain(self) -> GatewayRequestSnapshot: 

72 drained: Final = self._counts 

73 self._counts = {} # mutable-ok: the fold restarts empty; the drained map is handed off whole 

74 return drained 

75 

76 def restore(self, snapshot: GatewayRequestSnapshot) -> None: 

77 """ 

78 Merge un-committed counts back so the next flush retries them. 

79 

80 A dropped flush would silently undercount the metric the dashboard now 

81 treats as the source of truth. Merging cannot grow without bound: keys 

82 collapse on collision, so the fold stays bounded by (date x category x 

83 route) however long the database is unreachable. 

84 

85 This buys at-least-once, not exactly-once, and the cost is worth stating. 

86 The statement commits on the server before its acknowledgement is read, so 

87 a failure raised after the commit (a connection dropped while reading the 

88 acknowledgement) restores counts that are already persisted, and the next 

89 flush increments them a second time. Exactly-once would need a dedup key 

90 the upsert could ignore on replay. For a traffic-volume metric a rare 

91 overcount on a dropped acknowledgement beats losing a whole interval to 

92 every database blip, so the trade is deliberate. 

93 """ 

94 self._counts = dict(fold_counts(chain(self._counts.items(), snapshot.items()))) # mutable-ok: fold replaced 

95 

96 

97def fold_counts(items: Iterable[tuple[GatewayRequestKey, GatewayRequestCounts]]) -> GatewayRequestSnapshot: 

98 """Sum counts key-wise; the result stays bounded by (date x category x route).""" 

99 folded: Final[dict[GatewayRequestKey, GatewayRequestCounts]] = {} # mutable-ok: local fold returned once 

100 for key, counts in items: 

101 existing = folded.get(key, _EMPTY) 

102 folded[key] = GatewayRequestCounts( 

103 successful_requests=existing.successful_requests + counts.successful_requests, 

104 failed_requests=existing.failed_requests + counts.failed_requests, 

105 ) 

106 return folded 

107 

108 

109def build_gateway_requests_upsert(snapshot: GatewayRequestSnapshot) -> tuple[str, tuple[str | int, ...]]: 

110 """ 

111 One ``INSERT ... ON CONFLICT DO UPDATE`` that increments every (date, category, 

112 route) in the snapshot. Rows are ordered by the conflict key so concurrent 

113 writers lock rows in the same order and cannot deadlock. 

114 """ 

115 ordered: Final = sorted(snapshot.items(), key=lambda item: (item[0].date, item[0].category, item[0].route)) 

116 rows: Final = ", ".join( 

117 f"(${base + 1}::text, ${base + 2}::text, ${base + 3}::text, ${base + 4}::bigint, ${base + 5}::bigint, {_UTC_NOW})" 

118 for base in range(0, len(ordered) * _COLUMNS_PER_ROW, _COLUMNS_PER_ROW) 

119 ) 

120 sql: Final = ( 

121 f'INSERT INTO {_TABLE} ("date", "category", "route", "successful_requests", "failed_requests", "updated_at")\n' 

122 f"VALUES {rows}\n" 

123 'ON CONFLICT ("date", "category", "route") DO UPDATE SET\n' 

124 f' "successful_requests" = {_TABLE}."successful_requests" + EXCLUDED."successful_requests",\n' 

125 f' "failed_requests" = {_TABLE}."failed_requests" + EXCLUDED."failed_requests",\n' 

126 f' "updated_at" = {_UTC_NOW}' 

127 ) 

128 params: Final[tuple[str | int, ...]] = tuple( 

129 value 

130 for key, counts in ordered 

131 for value in (key.date, key.category, key.route, counts.successful_requests, counts.failed_requests) 

132 ) 

133 return sql, params 

134 

135 

136async def commit_gateway_requests_to_db( 

137 *, 

138 prisma_client: "PrismaClient", 

139 snapshot: GatewayRequestSnapshot, 

140) -> None: 

141 """Increment every (date, category, route) in the snapshot with a single statement.""" 

142 if not snapshot: 

143 return 

144 

145 sql, params = build_gateway_requests_upsert(snapshot) 

146 await prisma_client.db.execute_raw(sql, *params) # pyright: ignore[reportAny] # untyped prisma client 

147 

148 verbose_proxy_logger.debug( 

149 "Gateway request tracking - committed %d aggregated rows in one statement", len(snapshot) 

150 ) 

151 

152 

153class GatewayRequestRedisBuffer: 

154 """ 

155 Folds every worker's snapshot through one Redis list so a single pod per 

156 interval writes the table, mirroring the spend writer's transaction buffer. 

157 

158 Each entry is one worker's snapshot as JSON rows; the lock holder pops them, 

159 sums them, and commits one statement. A commit failure pushes the summed 

160 rows back so the next holder retries, keeping the at-least-once guarantee. 

161 If that push fails too, the rows go back to the holder's own accumulator so 

162 they ride along with its next flush instead of vanishing with the pop. 

163 """ 

164 

165 def __init__(self, *, redis_cache: RedisCache, pod_lock_manager: PodLockManager) -> None: 

166 self._redis_cache: Final = redis_cache 

167 self._pod_lock_manager: Final = pod_lock_manager 

168 

169 async def push(self, snapshot: GatewayRequestSnapshot) -> None: 

170 if not snapshot: 

171 return 

172 rows: Final[_BufferedRows] = tuple( 

173 (key.date, key.category, key.route, counts.successful_requests, counts.failed_requests) 

174 for key, counts in snapshot.items() 

175 ) 

176 await self._redis_cache.async_rpush(key=REDIS_GATEWAY_REQUESTS_BUFFER_KEY, values=(json.dumps(rows),)) 

177 

178 async def _pop_batch(self) -> tuple[str | bytes, ...]: 

179 popped: Final[object] = await self._redis_cache.async_lpop( # pyright: ignore[reportAny] # redis returns Any 

180 key=REDIS_GATEWAY_REQUESTS_BUFFER_KEY, count=MAX_REDIS_BUFFER_DEQUEUE_COUNT 

181 ) 

182 if not popped: 

183 return () 

184 return _BUFFERED_ENTRIES.validate_python(popped if isinstance(popped, list) else (popped,)) 

185 

186 async def _pop_all(self) -> AsyncIterator[str | bytes]: 

187 while True: 

188 batch = await self._pop_batch() 

189 for entry in batch: 

190 yield entry 

191 if len(batch) < MAX_REDIS_BUFFER_DEQUEUE_COUNT: 

192 return 

193 

194 async def pop(self) -> GatewayRequestSnapshot: 

195 entries: Final = tuple([entry async for entry in self._pop_all()]) 

196 return fold_counts( 

197 ( 

198 GatewayRequestKey(date=date, category=category, route=route), 

199 GatewayRequestCounts(successful_requests=succeeded, failed_requests=failed), 

200 ) 

201 for entry in entries 

202 for date, category, route, succeeded, failed in _BUFFERED_ROWS.validate_json(entry) 

203 ) 

204 

205 async def commit_if_leader(self, prisma_client: "PrismaClient") -> GatewayRequestSnapshot: 

206 """ 

207 Drain the list and write it as one statement, but only on the pod holding the job lock. 

208 

209 The lock is a lease, never released: the holder re-enters it on every flush and 

210 keeps committing alone until the TTL lapses, so the primary sees one statement 

211 per flush interval deployment-wide instead of one per worker. 

212 

213 Returns the popped rows that could be neither committed nor re-queued, for the 

214 caller to keep in memory. Empty on success. 

215 """ 

216 if not await self._pod_lock_manager.acquire_lock(cronjob_id=GATEWAY_REQUESTS_JOB_NAME): 

217 return _NO_COUNTS 

218 buffered: Final = await self.pop() 

219 try: 

220 await commit_gateway_requests_to_db(prisma_client=prisma_client, snapshot=buffered) 

221 except Exception: # noqa: BLE001 -- a failed commit must not stop the scheduler 

222 verbose_proxy_logger.warning( 

223 "Gateway request tracking - failed to commit %d buffered rows, re-queuing to Redis for the next flush", 

224 len(buffered), 

225 exc_info=True, 

226 ) 

227 return await self._requeue(buffered) 

228 return _NO_COUNTS 

229 

230 async def _requeue(self, snapshot: GatewayRequestSnapshot) -> GatewayRequestSnapshot: 

231 try: 

232 await self.push(snapshot) 

233 except Exception: # noqa: BLE001 -- the rows go back to the caller's accumulator instead 

234 verbose_proxy_logger.warning( 

235 "Gateway request tracking - Redis re-queue failed, keeping %d rows in memory for the next flush", 

236 len(snapshot), 

237 exc_info=True, 

238 ) 

239 return snapshot 

240 return _NO_COUNTS 

241 

242 

243async def flush_gateway_requests( 

244 prisma_client: "PrismaClient", 

245 accumulator: GatewayRequestAccumulator, 

246 redis_buffer: GatewayRequestRedisBuffer | None = None, 

247) -> None: 

248 """ 

249 Scheduler entrypoint. Never raises: a metering failure must not kill the job. 

250 

251 With ``redis_buffer`` the snapshot goes to Redis and only the lease holder 

252 writes to Postgres. Shutdown passes no buffer so a departing worker writes its 

253 own counts directly instead of parking them behind a lease it may not hold. 

254 

255 ``CancelledError`` is deliberately not caught, so a flush cancelled during 

256 shutdown drops its snapshot rather than restoring counts onto an accumulator 

257 the process is about to discard. 

258 """ 

259 snapshot: Final = accumulator.drain() 

260 try: 

261 if redis_buffer is None: 261 ↛ 264line 261 didn't jump to line 264 because the condition on line 261 was always true

262 await commit_gateway_requests_to_db(prisma_client=prisma_client, snapshot=snapshot) 

263 else: 

264 await redis_buffer.push(snapshot) 

265 except Exception: # noqa: BLE001 -- a failed flush must not stop the scheduler 

266 accumulator.restore(snapshot) 

267 verbose_proxy_logger.warning( 

268 "Gateway request tracking - failed to commit %d rows, retrying on the next flush", 

269 len(snapshot), 

270 exc_info=True, 

271 ) 

272 return 

273 if redis_buffer is None: 273 ↛ 275line 273 didn't jump to line 275 because the condition on line 273 was always true

274 return 

275 try: 

276 accumulator.restore(await redis_buffer.commit_if_leader(prisma_client)) 

277 except Exception: # noqa: BLE001 -- entries still in Redis are drained by the next flush 

278 verbose_proxy_logger.warning( 

279 "Gateway request tracking - leader drain failed, buffered rows stay in Redis for the next flush", 

280 exc_info=True, 

281 )