litellm/tests/integration/_support/database_relay.py
devin-ai-integration[bot] a1a42768c1
fix(spend-tracking): stop caching failed spend-log metadata lookups as confirmed misses (#43560)
* fix(spend-tracking): stop caching failed spend-log metadata lookups as confirmed misses

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* refactor(spend-tracking): share the short-lived miss cache write

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(spend): cover key alias recovery after spend log lookup failures across usage routes

* test(spend): bound outage alias lookups per miss window instead of a fixed count

* fix(spend-tracking): treat any spend-log lookup failure as a short-lived miss

The Prisma client raises a plain AttributeError when the database drops
the connection mid-query, so the PrismaError catch let it through and
the whole usage call answered 500. Any failure now keeps the 30 second
backoff only, and the integration proxy patches its test entitlement at
import so uvicorn's spawned workers inherit it

* test(integration): audit spend-log metadata recovery under timeouts and dropped connections

Cover the daily activity routes, the usage AI chat, the Vantage and
CloudZero dry runs and exports under a locked spend-log table and under
a database connection dropped mid-lookup, on a two-worker proxy, with
the recovery after the outage asserted through the proxy's own miss TTL.

Add a dropped_connection_relay that closes only the connection whose
bytes carry a trigger, so a cell can drop the one connection the
recovery query runs on while the rest of the pool keeps serving. Rewrite
the sweep and JWT cells for the merged main: the export route reads
metadata by SQL join and never calls the recovery, the search routes
answer key rows and find deleted keys by alias, and the daily-spend
owner recovery names the user while the alias stays blank. The sweep
cell now times out a second lookup under the same lock, which pins the
keys blank on the merge base and recovers on this branch.

* test(integration): match a dropped-connection trigger split across two reads

The dropped-connection relay checked each TCP read on its own, so a SQL
marker that straddled two reads never tripped it and the outage cells
would run without the outage they meant to exercise. Carry the tail of
the previous read into the next check, as the held-statement relay
already does, and pin that with a unit test that splits the trigger
across two writes.

* test(integration): scan relay triggers through an in-process helper

The dropped-connection relay now matches its SQL trigger through a TriggerScanner that carries the previous read's tail, and the unit test exercises that scanner directly instead of opening loopback sockets, which tests/unit forbids. The relay's end to end behavior stays covered by the integration cells

---------

Co-authored-by: gabriele <gabriele@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
2026-10-05 22:36:25 +00:00

256 lines
9.6 KiB
Python

import asyncio
import socket
import threading
import time
from collections.abc import Generator
from contextlib import contextmanager
from typing import Final
from urllib.parse import urlsplit, urlunsplit
from pydantic import TypeAdapter
PORT: Final = TypeAdapter(int)
OUTAGE_SECONDS: Final = 10.0
def _free_port() -> int:
with socket.socket() as reserve:
reserve.bind(("127.0.0.1", 0))
return PORT.validate_python(reserve.getsockname()[1])
class DatabaseRelay:
def __init__(self, upstream_host: str, upstream_port: int, trigger: bytes) -> None:
self.port: Final = _free_port()
self._upstream_host: Final = upstream_host
self._upstream_port: Final = upstream_port
self._trigger: Final = trigger
self._loop: Final = asyncio.new_event_loop()
self._armed: Final = threading.Event()
self.tripped: Final = threading.Event()
self.refused = 0
self.reconnected: Final = threading.Event()
self._tripped_at = 0.0
self._writers: tuple[asyncio.StreamWriter, ...] = ()
self._ready: Final = threading.Event()
self._thread: Final = threading.Thread(target=self._run, daemon=True)
def arm(self) -> None:
self._armed.set()
def start(self) -> None:
self._thread.start()
assert self._ready.wait(10), "Database relay did not start"
def stop(self) -> None:
self._loop.call_soon_threadsafe(self._loop.stop)
self._thread.join(10)
def _run(self) -> None:
asyncio.set_event_loop(self._loop)
self._loop.run_until_complete(asyncio.start_server(self._serve, "127.0.0.1", self.port))
self._ready.set()
self._loop.run_forever()
def _drop_all(self) -> None:
for writer in self._writers:
writer.close()
self._writers = ()
async def _serve(self, client_reader: asyncio.StreamReader, client_writer: asyncio.StreamWriter) -> None:
if self.tripped.is_set() and time.monotonic() - self._tripped_at < OUTAGE_SECONDS:
self.refused += 1
client_writer.close()
return
if self.tripped.is_set():
self.reconnected.set()
server_reader, server_writer = await asyncio.open_connection(self._upstream_host, self._upstream_port)
self._writers = (*self._writers, client_writer, server_writer)
async def forward(reader: asyncio.StreamReader, writer: asyncio.StreamWriter, inspect: bool) -> None:
try:
while chunk := await reader.read(65536):
if inspect and self._armed.is_set() and not self.tripped.is_set() and self._trigger in chunk:
self._tripped_at = time.monotonic()
self.tripped.set()
self._drop_all()
return
writer.write(chunk)
await writer.drain()
except (ConnectionError, asyncio.IncompleteReadError):
return
finally:
writer.close()
await asyncio.gather(
forward(client_reader, server_writer, True),
forward(server_reader, client_writer, False),
)
class HeldStatementRelay:
def __init__(self, upstream_host: str, upstream_port: int, trigger: bytes) -> None:
self.port: Final = _free_port()
self._upstream_host: Final = upstream_host
self._upstream_port: Final = upstream_port
self._trigger: Final = trigger
self._loop: Final = asyncio.new_event_loop()
self._released: Final = asyncio.Event()
self.held: Final = threading.Event()
self._ready: Final = threading.Event()
self._thread: Final = threading.Thread(target=self._run, daemon=True)
def release(self) -> None:
self._loop.call_soon_threadsafe(self._released.set)
def start(self) -> None:
self._thread.start()
assert self._ready.wait(10), "Database relay did not start"
def stop(self) -> None:
self.release()
self._loop.call_soon_threadsafe(self._loop.stop)
self._thread.join(10)
def _run(self) -> None:
asyncio.set_event_loop(self._loop)
self._loop.run_until_complete(asyncio.start_server(self._serve, "127.0.0.1", self.port))
self._ready.set()
self._loop.run_forever()
def _holds(self, window: bytes) -> bool:
return not self.held.is_set() and self._trigger in window
async def _serve(self, client_reader: asyncio.StreamReader, client_writer: asyncio.StreamWriter) -> None:
server_reader, server_writer = await asyncio.open_connection(self._upstream_host, self._upstream_port)
async def forward(reader: asyncio.StreamReader, writer: asyncio.StreamWriter, inspect: bool) -> None:
tail = b"" # rebind-ok: carries the previous read's end so a trigger split across reads still matches
try:
while chunk := await reader.read(65536):
window: Final = tail + chunk
if inspect and self._holds(window):
self.held.set()
await self._released.wait()
tail = window[-(len(self._trigger) - 1) :]
writer.write(chunk)
await writer.drain()
except (ConnectionError, asyncio.IncompleteReadError):
return
finally:
writer.close()
await asyncio.gather(
forward(client_reader, server_writer, True),
forward(server_reader, client_writer, False),
)
class TriggerScanner:
def __init__(self, trigger: bytes) -> None:
self._trigger: Final = trigger
self._tail: bytes = b""
def feed(self, chunk: bytes) -> bool:
window: Final = self._tail + chunk
self._tail = window[-(len(self._trigger) - 1) :]
return self._trigger in window
class DroppedConnectionRelay:
def __init__(self, upstream_host: str, upstream_port: int, trigger: bytes) -> None:
self.port: Final = _free_port()
self._upstream_host: Final = upstream_host
self._upstream_port: Final = upstream_port
self._trigger: Final = trigger
self._loop: Final = asyncio.new_event_loop()
self._armed: Final = threading.Event()
self.dropped: Final = threading.Event()
self._ready: Final = threading.Event()
self._thread: Final = threading.Thread(target=self._run, daemon=True)
def arm(self) -> None:
self._armed.set()
def disarm(self) -> None:
self._armed.clear()
def start(self) -> None:
self._thread.start()
assert self._ready.wait(10), "Database relay did not start"
def stop(self) -> None:
self._loop.call_soon_threadsafe(self._loop.stop)
self._thread.join(10)
def _run(self) -> None:
asyncio.set_event_loop(self._loop)
self._loop.run_until_complete(asyncio.start_server(self._serve, "127.0.0.1", self.port))
self._ready.set()
self._loop.run_forever()
async def _serve(self, client_reader: asyncio.StreamReader, client_writer: asyncio.StreamWriter) -> None:
server_reader, server_writer = await asyncio.open_connection(self._upstream_host, self._upstream_port)
async def forward(reader: asyncio.StreamReader, writer: asyncio.StreamWriter, inspect: bool) -> None:
scanner: Final = TriggerScanner(self._trigger)
try:
while chunk := await reader.read(65536):
matched: Final = scanner.feed(chunk)
if inspect and self._armed.is_set() and matched:
self.dropped.set()
client_writer.close()
return
writer.write(chunk)
await writer.drain()
except (ConnectionError, asyncio.IncompleteReadError):
return
finally:
writer.close()
await asyncio.gather(
forward(client_reader, server_writer, True),
forward(server_reader, client_writer, False),
)
def _relayed_url(database_url: str, port: int) -> str:
parts: Final = urlsplit(database_url)
credentials: Final = f"{parts.username}:{parts.password}@" if parts.username else ""
return urlunsplit(parts._replace(netloc=f"{credentials}127.0.0.1:{port}"))
@contextmanager
def database_relay(database_url: str, trigger: bytes) -> Generator[tuple[DatabaseRelay, str]]:
parts: Final = urlsplit(database_url)
assert parts.hostname is not None and parts.port is not None, database_url
relay: Final = DatabaseRelay(parts.hostname, parts.port, trigger)
relay.start()
try:
yield relay, _relayed_url(database_url, relay.port)
finally:
relay.stop()
@contextmanager
def held_statement_relay(database_url: str, trigger: bytes) -> Generator[tuple[HeldStatementRelay, str]]:
parts: Final = urlsplit(database_url)
assert parts.hostname is not None and parts.port is not None, database_url
relay: Final = HeldStatementRelay(parts.hostname, parts.port, trigger)
relay.start()
try:
yield relay, _relayed_url(database_url, relay.port)
finally:
relay.stop()
@contextmanager
def dropped_connection_relay(database_url: str, trigger: bytes) -> Generator[tuple[DroppedConnectionRelay, str]]:
parts: Final = urlsplit(database_url)
assert parts.hostname is not None and parts.port is not None, database_url
relay: Final = DroppedConnectionRelay(parts.hostname, parts.port, trigger)
relay.start()
try:
yield relay, _relayed_url(database_url, relay.port)
finally:
relay.stop()