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
« 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.
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.
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"""
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
31from pydantic import AliasChoices, Field
32from pydantic_settings import BaseSettings, SettingsConfigDict
34from litellm._logging import verbose_proxy_logger
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
42UnavailablePolicy: TypeAlias = Literal["fallback", "drop"]
43PublishOutcome: TypeAlias = Literal["queued", "fallback", "dropped"]
46class CollectorSettings(BaseSettings):
47 """``LITELLM_COLLECTOR_*`` env vars, shared by the gateway producer and the sidecar consumer."""
49 model_config = SettingsConfigDict(
50 env_prefix=COLLECTOR_ENV_PREFIX, case_sensitive=False, extra="ignore", frozen=True, populate_by_name=True
51 )
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"))
61 @property
62 def produces(self) -> bool:
63 return self.enabled and self.job_role != COLLECTOR_JOB_ROLE
66@dataclass(frozen=True, slots=True)
67class UnixAddress:
68 path: str
71@dataclass(frozen=True, slots=True)
72class TcpAddress:
73 host: str
74 port: int
77@dataclass(frozen=True, slots=True)
78class AddressError:
79 reason: str
82CollectorAddress: TypeAlias = UnixAddress | TcpAddress
85def _is_loopback(host: str) -> bool:
86 try:
87 return ipaddress.ip_address(host).is_loopback
88 except ValueError:
89 return host == "localhost"
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}")
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)
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 )
139@dataclass(frozen=True, slots=True)
140class _Connection:
141 reader: asyncio.StreamReader
142 writer: asyncio.StreamWriter
144 @property
145 def alive(self) -> bool:
146 return not self.writer.is_closing() and not self.reader.at_eof()
149@dataclass(frozen=True, slots=True)
150class SpendEventProducerStats:
151 queued: int
152 sent: int
153 fallback: int
154 dropped: int
155 connected: bool
158class SpendEventProducer:
159 """Bounded buffer plus one writer task per process; see the module docstring for the contract."""
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
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 )
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"
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")
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
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
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()
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
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
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
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"
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)