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

221 statements  

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

1"""Fire-and-forget push of serialized spend events from an inference worker to the pod-local sidecar. 

2 

3``LITELLM_COLLECTOR_ENABLED=true`` turns the push on in the gateway; the sidecar process sets 

4``LITELLM_JOB_ROLE=collector`` and always runs the pipeline in-process. Events queue in a bounded 

5in-memory buffer that a single writer task flushes over a unix socket or loopback TCP connection. 

6When the sidecar is unreachable, the buffer is full, or the connection breaks mid-write, each affected 

7event follows ``LITELLM_COLLECTOR_ON_UNAVAILABLE``: ``fallback`` runs the existing cost pipeline in 

8the worker, ``drop`` counts it and moves on. Transitions are logged with the counters, so a sidecar 

9outage is visible without scraping anything. 

10 

11Delivery is at-most-once: a sidecar crash loses the events the kernel already took from its socket. 

12A sidecar that stops gracefully half-closes each connection first (EOF towards the producer) and 

13keeps reading until the producer hangs up, so the producer switches to the unavailable policy without 

14losing the events in flight. A write that fails part-way follows the unavailable policy without double 

15counting: ``drain()`` only fails while part of the line is still buffered in this process, so the 

16sidecar can at most have read a truncated line, which it discards. When the gateway itself stops with 

17the writer stuck mid-send, only an event whose bytes are still in the producer's write buffer follows 

18the unavailable policy; the connection is aborted first so the sidecar discards the truncated line 

19instead of also counting it. Events from one uvicorn worker are handled in the order it produced them; 

20events from different workers interleave, exactly like the in-process callbacks do today. 

21""" 

22 

23import asyncio 

24import ipaddress 

25import time 

26from collections.abc import Awaitable, Callable 

27from dataclasses import dataclass 

28from typing import Final, Literal, TypeAlias 

29from urllib.parse import urlsplit 

30 

31from pydantic import AliasChoices, Field 

32from pydantic_settings import BaseSettings, SettingsConfigDict 

33 

34from litellm._logging import verbose_proxy_logger 

35 

36COLLECTOR_ENV_PREFIX: Final = "LITELLM_COLLECTOR_" 

37COLLECTOR_JOB_ROLE: Final = "collector" 

38DEFAULT_COLLECTOR_ADDRESS: Final = "unix:///var/run/litellm/collector.sock" 

39RECONNECT_BACKOFF_SECONDS: Final = 1.0 

40DROP_LOG_EVERY: Final = 1000 

41 

42UnavailablePolicy: TypeAlias = Literal["fallback", "drop"] 

43PublishOutcome: TypeAlias = Literal["queued", "fallback", "dropped"] 

44 

45 

46class CollectorSettings(BaseSettings): 

47 """``LITELLM_COLLECTOR_*`` env vars, shared by the gateway producer and the sidecar consumer.""" 

48 

49 model_config = SettingsConfigDict( 

50 env_prefix=COLLECTOR_ENV_PREFIX, case_sensitive=False, extra="ignore", frozen=True, populate_by_name=True 

51 ) 

52 

53 enabled: bool = False 

54 address: str = DEFAULT_COLLECTOR_ADDRESS 

55 buffer_size: int = Field(default=1000, ge=1) 

56 on_unavailable: UnavailablePolicy = "fallback" 

57 drain_timeout_seconds: float = Field(default=10.0, gt=0) 

58 connect_timeout_seconds: float = Field(default=1.0, gt=0) 

59 job_role: str | None = Field(default=None, validation_alias=AliasChoices("LITELLM_JOB_ROLE")) 

60 

61 @property 

62 def produces(self) -> bool: 

63 return self.enabled and self.job_role != COLLECTOR_JOB_ROLE 

64 

65 

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

67class UnixAddress: 

68 path: str 

69 

70 

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

72class TcpAddress: 

73 host: str 

74 port: int 

75 

76 

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

78class AddressError: 

79 reason: str 

80 

81 

82CollectorAddress: TypeAlias = UnixAddress | TcpAddress 

83 

84 

85def _is_loopback(host: str) -> bool: 

86 try: 

87 return ipaddress.ip_address(host).is_loopback 

88 except ValueError: 

89 return host == "localhost" 

90 

91 

92def parse_collector_address(address: str) -> CollectorAddress | AddressError: 

93 """``unix:///path/to.sock`` or ``tcp://127.0.0.1:port``; the socket carries unauthenticated spend events.""" 

94 parsed: Final = urlsplit(address) 

95 if parsed.scheme == "unix" and parsed.path: 

96 return UnixAddress(path=parsed.path) 

97 if parsed.scheme == "tcp" and parsed.hostname and parsed.port is not None: 

98 if not _is_loopback(parsed.hostname): 

99 return AddressError(reason=f"tcp collector address must be a loopback host, got {address!r}") 

100 return TcpAddress(host=parsed.hostname, port=parsed.port) 

101 return AddressError(reason=f"expected unix:///path or tcp://127.0.0.1:port, got {address!r}") 

102 

103 

104async def open_collector_connection( 

105 address: CollectorAddress, timeout: float 

106) -> tuple[asyncio.StreamReader, asyncio.StreamWriter]: 

107 match address: 

108 case UnixAddress(path=path): 

109 return await asyncio.wait_for(asyncio.open_unix_connection(path), timeout) 

110 case TcpAddress(host=host, port=port): 

111 return await asyncio.wait_for(asyncio.open_connection(host, port), timeout) 

112 

113 

114def build_spend_event_producer( 

115 settings: CollectorSettings, fallback: Callable[[bytes], Awaitable[None]] 

116) -> "SpendEventProducer | None": 

117 """The gateway producer for these settings, or ``None`` when the pipeline stays in-process.""" 

118 if not settings.produces: 118 ↛ 120line 118 didn't jump to line 120 because the condition on line 118 was always true

119 return None 

120 address: Final = parse_collector_address(settings.address) 

121 if isinstance(address, AddressError): 

122 verbose_proxy_logger.error("collector: %s; running the spend pipeline in-process", address.reason) 

123 return None 

124 verbose_proxy_logger.info( 

125 "collector: offloading spend tracking to %s (buffer=%d, on_unavailable=%s)", 

126 settings.address, 

127 settings.buffer_size, 

128 settings.on_unavailable, 

129 ) 

130 return SpendEventProducer( 

131 address=address, 

132 on_unavailable=settings.on_unavailable, 

133 buffer_size=settings.buffer_size, 

134 connect_timeout=settings.connect_timeout_seconds, 

135 fallback=fallback, 

136 ) 

137 

138 

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

140class _Connection: 

141 reader: asyncio.StreamReader 

142 writer: asyncio.StreamWriter 

143 

144 @property 

145 def alive(self) -> bool: 

146 return not self.writer.is_closing() and not self.reader.at_eof() 

147 

148 

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

150class SpendEventProducerStats: 

151 queued: int 

152 sent: int 

153 fallback: int 

154 dropped: int 

155 connected: bool 

156 

157 

158class SpendEventProducer: 

159 """Bounded buffer plus one writer task per process; see the module docstring for the contract.""" 

160 

161 def __init__( 

162 self, 

163 address: CollectorAddress, 

164 on_unavailable: UnavailablePolicy, 

165 buffer_size: int, 

166 connect_timeout: float, 

167 fallback: Callable[[bytes], Awaitable[None]], 

168 clock: Callable[[], float] = time.monotonic, 

169 open_connection: Callable[ 

170 [CollectorAddress, float], Awaitable[tuple[asyncio.StreamReader, asyncio.StreamWriter]] 

171 ] = open_collector_connection, 

172 ) -> None: 

173 self._address = address 

174 self._on_unavailable = on_unavailable 

175 self._buffer_size = buffer_size 

176 self._connect_timeout = connect_timeout 

177 self._fallback = fallback 

178 self._clock = clock 

179 self._open_connection = open_connection 

180 self._queue: asyncio.Queue[bytes] | None = None 

181 self._writer_task: asyncio.Task[None] | None = None 

182 self._connection: _Connection | None = None 

183 self._in_flight: bytes | None = None 

184 self._closing = False 

185 self._next_connect_at = 0.0 

186 self._queued = 0 

187 self._sent = 0 

188 self._fallback_count = 0 

189 self._dropped = 0 

190 

191 def stats(self) -> SpendEventProducerStats: 

192 return SpendEventProducerStats( 

193 queued=self._queued, 

194 sent=self._sent, 

195 fallback=self._fallback_count, 

196 dropped=self._dropped, 

197 connected=self._connection is not None, 

198 ) 

199 

200 async def publish(self, line: bytes) -> PublishOutcome: 

201 """Hand one serialized event to the writer task, or apply the unavailable policy right away.""" 

202 if self._closing or self._clock() < self._next_connect_at: 

203 return await self._unavailable(line, "sidecar unreachable") 

204 queue: Final = self._ensure_writer() 

205 try: 

206 queue.put_nowait(line) 

207 except asyncio.QueueFull: 

208 return await self._unavailable(line, "buffer full") 

209 self._queued += 1 

210 return "queued" 

211 

212 async def close(self, drain_timeout: float) -> None: 

213 """Flush the buffer for up to ``drain_timeout`` seconds, then apply the unavailable policy to the rest.""" 

214 self._closing = True 

215 queue: Final = self._queue 

216 task: Final = self._writer_task 

217 if queue is None or task is None: 

218 return 

219 try: 

220 await asyncio.wait_for(queue.join(), drain_timeout) 

221 except asyncio.TimeoutError: 

222 verbose_proxy_logger.warning( 

223 "collector: %s events still buffered after %.1fs drain timeout", queue.qsize(), drain_timeout 

224 ) 

225 task.cancel() 

226 try: 

227 await task 

228 except asyncio.CancelledError: 

229 pass 

230 unsent: Final = self._take_unsent() 

231 await self._disconnect() 

232 if unsent is not None: 

233 await self._unavailable(unsent, "shutdown") 

234 while not queue.empty(): 

235 await self._unavailable(queue.get_nowait(), "shutdown") 

236 

237 def _take_unsent(self) -> bytes | None: 

238 """The in-flight event if any of its bytes never left this process, aborting the half-written connection.""" 

239 in_flight: Final = self._in_flight 

240 self._in_flight = None 

241 connection: Final = self._connection 

242 if in_flight is None: 

243 return None 

244 if connection is None: 

245 return in_flight 

246 if connection.writer.transport.get_write_buffer_size() == 0: 

247 return None 

248 connection.writer.transport.abort() 

249 return in_flight 

250 

251 def _ensure_writer(self) -> asyncio.Queue[bytes]: 

252 if self._queue is None: 

253 self._queue = asyncio.Queue(maxsize=self._buffer_size) 

254 if self._writer_task is None or self._writer_task.done(): 

255 self._writer_task = asyncio.get_running_loop().create_task(self._run_writer(self._queue)) 

256 return self._queue 

257 

258 async def _run_writer(self, queue: asyncio.Queue[bytes]) -> None: 

259 while True: 

260 line = await queue.get() 

261 try: 

262 await self._send(line) 

263 finally: 

264 queue.task_done() 

265 

266 async def _send(self, line: bytes) -> None: 

267 self._in_flight = line 

268 connection: Final = await self._connect() 

269 if connection is None: 

270 self._in_flight = None 

271 await self._unavailable(line, "sidecar unreachable") 

272 return 

273 try: 

274 connection.writer.write(line) 

275 await connection.writer.drain() 

276 except (ConnectionError, OSError, RuntimeError) as error: # uvloop: RuntimeError on a closed transport 

277 self._in_flight = None 

278 await self._disconnect() 

279 self._next_connect_at = self._clock() + RECONNECT_BACKOFF_SECONDS 

280 await self._unavailable(line, f"write failed: {error}") 

281 return 

282 self._in_flight = None 

283 self._sent += 1 

284 

285 async def _connect(self) -> _Connection | None: 

286 if self._connection is not None and self._connection.alive: 

287 return self._connection 

288 await self._disconnect() 

289 if self._clock() < self._next_connect_at: 

290 return None 

291 try: 

292 reader, writer = await self._open_connection(self._address, self._connect_timeout) 

293 except (ConnectionError, OSError, asyncio.TimeoutError) as error: 

294 self._next_connect_at = self._clock() + RECONNECT_BACKOFF_SECONDS 

295 verbose_proxy_logger.warning( 

296 "collector: cannot reach %s (%s); applying %s policy for %.0fs. stats=%s", 

297 self._address, 

298 error, 

299 self._on_unavailable, 

300 RECONNECT_BACKOFF_SECONDS, 

301 self.stats(), 

302 ) 

303 return None 

304 self._connection = _Connection(reader=reader, writer=writer) 

305 verbose_proxy_logger.info("collector: connected to %s. stats=%s", self._address, self.stats()) 

306 return self._connection 

307 

308 async def _disconnect(self) -> None: 

309 connection: Final = self._connection 

310 self._connection = None 

311 if connection is None: 

312 return 

313 connection.writer.close() 

314 try: 

315 await connection.writer.wait_closed() 

316 except (ConnectionError, OSError): 

317 pass 

318 

319 async def _unavailable(self, line: bytes, reason: str) -> PublishOutcome: 

320 if self._on_unavailable == "fallback": 

321 self._fallback_count += 1 

322 fallback: Final = asyncio.ensure_future(self._run_fallback(line, reason)) 

323 try: 

324 await asyncio.shield(fallback) 

325 except asyncio.CancelledError: 

326 await fallback 

327 raise 

328 return "fallback" 

329 self._dropped += 1 

330 if self._dropped % DROP_LOG_EVERY == 1: 

331 verbose_proxy_logger.warning("collector: dropping spend event (%s). stats=%s", reason, self.stats()) 

332 return "dropped" 

333 

334 async def _run_fallback(self, line: bytes, reason: str) -> None: 

335 try: 

336 await self._fallback(line) 

337 except Exception: # noqa: BLE001 # one failing event must not kill the writer task 

338 verbose_proxy_logger.exception("collector: in-process fallback failed (%s)", reason)