mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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>
This commit is contained in:
parent
b9251dafad
commit
a1a42768c1
6 changed files with 1943 additions and 7 deletions
|
|
@ -172,11 +172,9 @@ async def _db_or_empty(
|
|||
warning: str,
|
||||
count: int,
|
||||
) -> _T | None:
|
||||
from prisma.errors import PrismaError
|
||||
|
||||
try:
|
||||
return await load()
|
||||
except PrismaError as e:
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.warning(warning, count, e)
|
||||
return None
|
||||
|
||||
|
|
@ -441,6 +439,10 @@ async def _query_spend_log_metadata(
|
|||
)
|
||||
|
||||
|
||||
def _remember_short_lived_miss(cache: InMemoryCache, key: str) -> None:
|
||||
cache.set_cache(key, KeyMetadataDict(), ttl=SPEND_LOG_KEY_METADATA_MISS_CACHE_TTL)
|
||||
|
||||
|
||||
def _remember_spend_log_metadata(
|
||||
cache: InMemoryCache, digest: str, window: tuple[datetime, datetime], meta: KeyMetadataDict | None
|
||||
) -> None:
|
||||
|
|
@ -452,7 +454,7 @@ def _remember_spend_log_metadata(
|
|||
if cache.get_cache(missed_before) is not None:
|
||||
cache.set_cache(key, KeyMetadataDict())
|
||||
return
|
||||
cache.set_cache(key, KeyMetadataDict(), ttl=SPEND_LOG_KEY_METADATA_MISS_CACHE_TTL)
|
||||
_remember_short_lived_miss(cache, key)
|
||||
cache.set_cache(missed_before, True)
|
||||
|
||||
|
||||
|
|
@ -469,10 +471,13 @@ async def _spend_log_metadata_one_query_at_a_time(
|
|||
fresh: Final = (
|
||||
await _query_spend_log_metadata(prisma_client, pending, window) if pending else _EMPTY_KEY_METADATA
|
||||
)
|
||||
found: Final = fresh if fresh is not None else _EMPTY_KEY_METADATA
|
||||
if fresh is None:
|
||||
for digest in pending:
|
||||
_remember_short_lived_miss(cache, _spend_log_cache_key(digest, window))
|
||||
return settled
|
||||
for digest in pending:
|
||||
_remember_spend_log_metadata(cache, digest, window, found.get(digest))
|
||||
return MappingProxyType({**settled, **found})
|
||||
_remember_spend_log_metadata(cache, digest, window, fresh.get(digest))
|
||||
return MappingProxyType({**settled, **fresh})
|
||||
|
||||
|
||||
async def recover_key_metadata_from_spend_logs(
|
||||
|
|
|
|||
|
|
@ -146,6 +146,74 @@ class HeldStatementRelay:
|
|||
)
|
||||
|
||||
|
||||
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 ""
|
||||
|
|
@ -174,3 +242,15 @@ def held_statement_relay(database_url: str, trigger: bytes) -> Generator[tuple[H
|
|||
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()
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -0,0 +1,626 @@
|
|||
import csv
|
||||
import io
|
||||
import json
|
||||
import re
|
||||
import socket
|
||||
import socketserver
|
||||
import threading
|
||||
import uuid
|
||||
from collections.abc import Generator, Mapping
|
||||
from contextlib import contextmanager, suppress
|
||||
from dataclasses import dataclass
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from email.message import Message
|
||||
from email.parser import BytesParser
|
||||
from email.policy import HTTP
|
||||
from hashlib import sha256
|
||||
from pathlib import Path
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from integration._support.client import Gateway, object_value, string_value
|
||||
from integration._support.daily_activity import (
|
||||
DAY,
|
||||
ROUTES,
|
||||
TEAM_SPEND,
|
||||
USER_SPEND,
|
||||
Route,
|
||||
activity_of_key,
|
||||
assert_key_reported,
|
||||
daily_rows,
|
||||
key_metadata,
|
||||
key_no_key_table_holds,
|
||||
seeded_metrics,
|
||||
seeded_row,
|
||||
user_row,
|
||||
)
|
||||
from integration._support.database import scratch_database
|
||||
from integration._support.database_relay import dropped_connection_relay
|
||||
from integration._support.process import OwnedProxy, owned_proxy_process
|
||||
from integration._support.tls import server_context, write_self_signed_cert
|
||||
from integration._support.wire import Reply, Request, Wire, wire_server
|
||||
from pydantic import JsonValue
|
||||
|
||||
WORKERS: Final = 2
|
||||
PROXY_CONFIG: Final = Path(__file__).resolve().parents[1] / "proxy_config.yaml"
|
||||
REVERSE_HASH_TRIGGER: Final = b"encode(sha256(convert_to(token, 'UTF8')), 'hex') = ANY("
|
||||
OWNER_RECOVERY_TRIGGER: Final = b"MIN(user_id) AS first_owner"
|
||||
CLOUDZERO_HOST: Final = "api.cloudzero.com"
|
||||
CLOUDZERO_AUTHORITY: Final = f"{CLOUDZERO_HOST}:443"
|
||||
EXPORT_FILENAME: Final = re.compile(r"usage_\d{8}T\d{6}Z_\d{8}T\d{6}Z\.csv")
|
||||
FILENAME_STAMP: Final = "%Y%m%dT%H%M%SZ"
|
||||
VANTAGE_DONE: Final = "Vantage export completed successfully"
|
||||
CLOUDZERO_DONE: Final = "CloudZero export completed successfully"
|
||||
NO_ENVIRONMENT: Final[Mapping[str, str]] = MappingProxyType({})
|
||||
ENTITY_TABLES: Final = tuple(
|
||||
dict.fromkeys((route.table, route.entity_column) for route in ROUTES if route.table != USER_SPEND)
|
||||
)
|
||||
FOCUS_ALIAS_FIELDS: Final = frozenset({"BillingAccountName", "Tags"})
|
||||
CBF_ALIAS_FIELDS: Final = frozenset({"resource/account", "resource/tag:api_key_alias"})
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Owned:
|
||||
owner: str
|
||||
email: str
|
||||
alias: str
|
||||
key: str
|
||||
token: str
|
||||
double: str
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Sink:
|
||||
forbidden: threading.Event
|
||||
missing: threading.Event
|
||||
|
||||
def respond(self, request: Request) -> Reply:
|
||||
if self.forbidden.is_set():
|
||||
return Reply(status=403, body=b'{"error":"forbidden"}')
|
||||
if self.missing.is_set():
|
||||
return Reply(status=404, body=b'{"error":"missing"}')
|
||||
return Reply()
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Upload:
|
||||
target: str
|
||||
authorization: str
|
||||
filename: str
|
||||
row: Mapping[str, str]
|
||||
|
||||
|
||||
def _identity(label: str) -> Owned:
|
||||
stamp: Final = uuid.uuid4().hex[:8]
|
||||
key: Final = f"sk-{uuid.uuid4().hex}"
|
||||
token: Final = sha256(key.encode()).hexdigest()
|
||||
return Owned(
|
||||
owner=f"{label}-owner-{stamp}",
|
||||
email=f"{label}-{uuid.uuid4().hex[:8]}@example.com",
|
||||
alias=f"{label}-key-{stamp}",
|
||||
key=key,
|
||||
token=token,
|
||||
double=sha256(token.encode()).hexdigest(),
|
||||
)
|
||||
|
||||
|
||||
def _register(proxy: Gateway, owned: Owned) -> None:
|
||||
proxy.post("/user/new", {"user_id": owned.owner, "user_email": owned.email, "auto_create_key": False})
|
||||
proxy.post("/key/generate", {"key": owned.key, "user_id": owned.owner, "key_alias": owned.alias})
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _relayed_proxy(
|
||||
gateway: Gateway,
|
||||
directory: Path,
|
||||
relayed_url: str,
|
||||
environment: Mapping[str, str] = NO_ENVIRONMENT,
|
||||
config: Path | None = None,
|
||||
) -> Generator[OwnedProxy]:
|
||||
with owned_proxy_process(
|
||||
gateway,
|
||||
directory,
|
||||
{
|
||||
"DATABASE_URL": relayed_url,
|
||||
"PRISMA_HEALTH_WATCHDOG_ENABLED": "false",
|
||||
"LITELLM_DISABLE_NO_REDIS_WARNING": "true",
|
||||
**environment,
|
||||
},
|
||||
config=config,
|
||||
remove_environment=("DATABASE_URL_READ_REPLICA",),
|
||||
workers=WORKERS,
|
||||
) as owned:
|
||||
yield owned
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _reader(owned: OwnedProxy) -> Generator[Gateway]:
|
||||
with httpx.Client(base_url=str(owned.gateway.client.base_url), timeout=90, trust_env=False) as client:
|
||||
yield Gateway(client, owned.gateway.key, owned.gateway.upstream_url)
|
||||
|
||||
|
||||
def _filters(route: Route, entity: str) -> dict[str, str]:
|
||||
return {} if route.entity_filter is None else {route.entity_filter: entity}
|
||||
|
||||
|
||||
def _config_allowing_a_base_url_in_the_body(directory: Path) -> Path:
|
||||
config: Final = directory / "client_side_credentials.yaml"
|
||||
config.write_text(
|
||||
PROXY_CONFIG.read_text().replace(
|
||||
"general_settings:\n", "general_settings:\n allow_client_side_credentials: true\n", 1
|
||||
)
|
||||
)
|
||||
return config
|
||||
|
||||
|
||||
def _records(value: JsonValue) -> tuple[dict[str, JsonValue], ...]:
|
||||
assert isinstance(value, list), value
|
||||
return tuple(object_value(item) for item in value)
|
||||
|
||||
|
||||
def _first(body: Mapping[str, JsonValue], name: str) -> dict[str, JsonValue]:
|
||||
records: Final = _records(body[name])
|
||||
assert len(records) == 1, body
|
||||
return records[0]
|
||||
|
||||
|
||||
def _without(record: Mapping[str, JsonValue], names: frozenset[str]) -> dict[str, JsonValue]:
|
||||
return {name: value for name, value in record.items() if name not in names}
|
||||
|
||||
|
||||
def _tags(record: Mapping[str, JsonValue]) -> dict[str, JsonValue]:
|
||||
return object_value(json.loads(string_value(record["Tags"])))
|
||||
|
||||
|
||||
def _dry_run(reader: Gateway, path: str) -> dict[str, JsonValue]:
|
||||
return object_value(reader.post(path, {})["dry_run_data"])
|
||||
|
||||
|
||||
def _focus_tags(owned: Owned, *, alias: bool) -> dict[str, str]:
|
||||
return {
|
||||
**({"api_key_alias": owned.alias} if alias else {}),
|
||||
"user_id": owned.owner,
|
||||
"user_email": owned.email,
|
||||
"model": "gpt-4o-mini",
|
||||
"model_group": "gpt-4o-mini",
|
||||
"custom_llm_provider": "openai",
|
||||
}
|
||||
|
||||
|
||||
def _assert_vantage_dry_run(body: Mapping[str, JsonValue], owned: Owned, *, alias: str | None) -> None:
|
||||
usage: Final = _first(body, "usage_data")
|
||||
assert {name: usage.get(name) for name in ("api_key", "api_key_alias", "user_id", "user_email", "spend")} == {
|
||||
"api_key": owned.double,
|
||||
"api_key_alias": alias,
|
||||
"user_id": owned.owner,
|
||||
"user_email": owned.email,
|
||||
"spend": 0.25,
|
||||
}, body
|
||||
focus: Final = _first(body, "normalized_data")
|
||||
assert {
|
||||
name: focus.get(name)
|
||||
for name in ("BillingAccountName", "BillingAccountId", "BilledCost", "ChargePeriodStart", "ChargePeriodEnd")
|
||||
} == {
|
||||
"BillingAccountName": alias,
|
||||
"BillingAccountId": owned.double,
|
||||
"BilledCost": 0.25,
|
||||
"ChargePeriodStart": "2026-02-03T00:00:00Z",
|
||||
"ChargePeriodEnd": "2026-02-04T00:00:00Z",
|
||||
}, body
|
||||
assert _tags(focus) == _focus_tags(owned, alias=alias is not None), body
|
||||
|
||||
|
||||
def _assert_cloudzero_dry_run(body: Mapping[str, JsonValue], owned: Owned, *, alias: str | None) -> None:
|
||||
usage: Final = _first(body, "usage_data")
|
||||
assert {name: usage.get(name) for name in ("api_key", "api_key_alias", "user_id", "user_email")} == {
|
||||
"api_key": owned.double,
|
||||
"api_key_alias": alias,
|
||||
"user_id": owned.owner,
|
||||
"user_email": owned.email,
|
||||
}, body
|
||||
cbf: Final = _first(body, "cbf_data")
|
||||
assert {name: cbf.get(name) for name in ("resource/account", "resource/tag:api_key_alias", "cost/cost")} == {
|
||||
"resource/account": f"{alias}|{owned.double[:8]}" if alias else owned.double[:8],
|
||||
"resource/tag:api_key_alias": str(alias),
|
||||
"cost/cost": 0.25,
|
||||
}, body
|
||||
|
||||
|
||||
def _cbf_record(owned: Owned, *, alias: str | None) -> dict[str, str]:
|
||||
prefix: Final = owned.double[:8]
|
||||
return {
|
||||
"time/usage_start": "2026-02-03T00:00:00Z",
|
||||
"cost/cost": "0.25",
|
||||
"resource/id": "czrn:litellm:openai:cross-region:unknown:llm-usage:gpt-4o-mini",
|
||||
"usage/amount": "15",
|
||||
"usage/units": "tokens",
|
||||
"resource/service": "gpt-4o-mini",
|
||||
"resource/account": f"{alias}|{prefix}" if alias else prefix,
|
||||
"resource/region": "cross-region",
|
||||
"resource/usage_family": "openai",
|
||||
"action/operation": "",
|
||||
"lineitem/type": "Usage",
|
||||
"resource/tag:provider": "openai",
|
||||
"resource/tag:model": "gpt-4o-mini",
|
||||
"resource/tag:entity_type": "team",
|
||||
"resource/tag:model_group": "gpt-4o-mini",
|
||||
"resource/tag:api_key_prefix": prefix,
|
||||
"resource/tag:api_key_alias": str(alias),
|
||||
"resource/tag:user_email": owned.email,
|
||||
"resource/tag:api_requests": "1",
|
||||
"resource/tag:successful_requests": "1",
|
||||
"resource/tag:failed_requests": "0",
|
||||
"resource/tag:cache_creation_tokens": "0",
|
||||
"resource/tag:cache_read_tokens": "0",
|
||||
"resource/tag:prompt_tokens": "10",
|
||||
"resource/tag:completion_tokens": "5",
|
||||
}
|
||||
|
||||
|
||||
def _window() -> tuple[datetime, datetime]:
|
||||
now: Final = datetime.now(UTC).replace(microsecond=0)
|
||||
return now - timedelta(hours=1), now + timedelta(hours=1)
|
||||
|
||||
|
||||
def _window_body(window: tuple[datetime, datetime]) -> dict[str, JsonValue]:
|
||||
return {"start_time_utc": window[0].isoformat(), "end_time_utc": window[1].isoformat()}
|
||||
|
||||
|
||||
def _window_filename(window: tuple[datetime, datetime]) -> str:
|
||||
return f"usage_{window[0].strftime(FILENAME_STAMP)}_{window[1].strftime(FILENAME_STAMP)}.csv"
|
||||
|
||||
|
||||
def _last_received(wire: Wire) -> Request:
|
||||
received: Final = wire.drain()
|
||||
assert received, "The sink received nothing"
|
||||
return received[-1]
|
||||
|
||||
|
||||
def _csv_part(request: Request) -> Message:
|
||||
message: Final = BytesParser(policy=HTTP).parsebytes(
|
||||
b"content-type: " + request.headers["content-type"].encode() + b"\r\n\r\n" + request.body
|
||||
)
|
||||
parts: Final = message.get_payload()
|
||||
assert isinstance(parts, list) and len(parts) == 1, request.body
|
||||
part: Final = parts[0]
|
||||
assert isinstance(part, Message), request.body
|
||||
assert part.get_param("name", header="content-disposition") == "csv", request.body
|
||||
return part
|
||||
|
||||
|
||||
def _upload(wire: Wire) -> Upload:
|
||||
request: Final = _last_received(wire)
|
||||
part: Final = _csv_part(request)
|
||||
payload: Final = part.get_payload(decode=True)
|
||||
assert isinstance(payload, bytes), request.body
|
||||
rows: Final = tuple(csv.DictReader(io.StringIO(payload.decode())))
|
||||
assert len(rows) == 1, payload
|
||||
filename: Final = part.get_filename()
|
||||
assert filename is not None, request.body
|
||||
return Upload(request.target, request.headers["authorization"], filename, rows[0])
|
||||
|
||||
|
||||
def _assert_csv_row(row: Mapping[str, str], owned: Owned, *, alias: str | None) -> None:
|
||||
assert {
|
||||
name: row.get(name)
|
||||
for name in (
|
||||
"BilledCost",
|
||||
"BillingAccountId",
|
||||
"BillingAccountName",
|
||||
"ChargeDescription",
|
||||
"ChargePeriodStart",
|
||||
"ChargePeriodEnd",
|
||||
)
|
||||
} == {
|
||||
"BilledCost": "0.25",
|
||||
"BillingAccountId": owned.double,
|
||||
"BillingAccountName": alias or "",
|
||||
"ChargeDescription": "gpt-4o-mini",
|
||||
"ChargePeriodStart": "2026-02-03T00:00:00Z",
|
||||
"ChargePeriodEnd": "2026-02-04T00:00:00Z",
|
||||
}, row
|
||||
assert json.loads(row["Tags"]) == _focus_tags(owned, alias=alias is not None), row
|
||||
|
||||
|
||||
def _pipe(source: socket.socket, sink: socket.socket) -> None:
|
||||
with suppress(OSError):
|
||||
for chunk in iter(lambda: source.recv(65536), b""):
|
||||
sink.sendall(chunk)
|
||||
with suppress(OSError):
|
||||
sink.shutdown(socket.SHUT_WR)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _connect_tunnel(destination: Wire, authority: str) -> Generator[str]:
|
||||
destination_port: Final = int(destination.url.rsplit(":", 1)[1])
|
||||
|
||||
class Tunnel(socketserver.StreamRequestHandler):
|
||||
rbufsize = 0
|
||||
request: socket.socket
|
||||
|
||||
def handle(self) -> None:
|
||||
requested: Final = self.rfile.readline().decode().split()[1]
|
||||
while self.rfile.readline() not in (b"\r\n", b""):
|
||||
pass
|
||||
if requested != authority:
|
||||
self.wfile.write(b"HTTP/1.1 403 Forbidden\r\ncontent-length: 0\r\n\r\n")
|
||||
return
|
||||
self.wfile.write(b"HTTP/1.1 200 Connection established\r\n\r\n")
|
||||
self.request.settimeout(10)
|
||||
with socket.create_connection(("127.0.0.1", destination_port), timeout=10) as upstream:
|
||||
outbound: Final = threading.Thread(target=_pipe, args=(self.request, upstream))
|
||||
outbound.start()
|
||||
_pipe(upstream, self.request)
|
||||
outbound.join(timeout=12)
|
||||
|
||||
with socketserver.ThreadingTCPServer(("127.0.0.1", 0), Tunnel) as server:
|
||||
thread: Final = threading.Thread(target=server.serve_forever, kwargs={"poll_interval": 0.05})
|
||||
thread.start()
|
||||
try:
|
||||
yield f"http://127.0.0.1:{server.server_address[1]}"
|
||||
finally:
|
||||
server.shutdown()
|
||||
thread.join(timeout=6)
|
||||
|
||||
|
||||
@pytest.mark.timeout(600)
|
||||
def test_every_usage_route_survives_a_dropped_database_connection_during_reverse_hash_recovery(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
owned_key: Final = _identity("dropped-reverse-hash")
|
||||
entity: Final = f"dropped-reverse-hash-{uuid.uuid4().hex[:8]}"
|
||||
rows: Final = (
|
||||
user_row(None, owned_key.double, DAY),
|
||||
*(seeded_row(table, column, entity, owned_key.double, DAY) for table, column in ENTITY_TABLES),
|
||||
)
|
||||
with (
|
||||
scratch_database() as database_url,
|
||||
dropped_connection_relay(database_url, REVERSE_HASH_TRIGGER) as (relay, relayed_url),
|
||||
_relayed_proxy(gateway, tmp_path, relayed_url) as owned,
|
||||
_reader(owned) as reader,
|
||||
daily_rows(rows, database_url=database_url),
|
||||
):
|
||||
_register(owned.gateway, owned_key)
|
||||
relay.arm()
|
||||
for route in ROUTES:
|
||||
relay.dropped.clear()
|
||||
dropped: Final = activity_of_key(reader, route.path, owned_key.double, **_filters(route, entity))
|
||||
assert relay.dropped.is_set(), f"{route.path}: {dropped.text}"
|
||||
assert_key_reported(dropped, owned_key.double, DAY, key_metadata(), seeded_metrics(1))
|
||||
relay.disarm()
|
||||
named: Final = key_metadata(alias=owned_key.alias, user=owned_key.owner, email=owned_key.email)
|
||||
for route in ROUTES:
|
||||
recovered: Final = activity_of_key(reader, route.path, owned_key.double, **_filters(route, entity))
|
||||
assert_key_reported(recovered, owned_key.double, DAY, named, seeded_metrics(1))
|
||||
|
||||
|
||||
@pytest.mark.timeout(300)
|
||||
def test_usage_page_survives_a_dropped_database_connection_during_user_detail_recovery(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
owned_key: Final = _identity("dropped-user-detail")
|
||||
with (
|
||||
scratch_database() as database_url,
|
||||
dropped_connection_relay(database_url, owned_key.owner.encode()) as (relay, relayed_url),
|
||||
_relayed_proxy(gateway, tmp_path, relayed_url) as owned,
|
||||
_reader(owned) as reader,
|
||||
daily_rows((user_row(None, owned_key.token, DAY),), database_url=database_url),
|
||||
):
|
||||
_register(owned.gateway, owned_key)
|
||||
relay.arm()
|
||||
dropped: Final = activity_of_key(reader, "/user/daily/activity", owned_key.token)
|
||||
assert relay.dropped.is_set(), dropped.text
|
||||
assert_key_reported(
|
||||
dropped,
|
||||
owned_key.token,
|
||||
DAY,
|
||||
key_metadata(alias=owned_key.alias, user=owned_key.owner, exists=True),
|
||||
seeded_metrics(1),
|
||||
)
|
||||
relay.disarm()
|
||||
recovered: Final = activity_of_key(reader, "/user/daily/activity", owned_key.token)
|
||||
assert_key_reported(
|
||||
recovered,
|
||||
owned_key.token,
|
||||
DAY,
|
||||
key_metadata(alias=owned_key.alias, user=owned_key.owner, email=owned_key.email, exists=True),
|
||||
seeded_metrics(1),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.timeout(300)
|
||||
def test_team_usage_page_survives_a_dropped_database_connection_during_owner_recovery(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
owner: Final = f"dropped-owner-{uuid.uuid4().hex[:8]}"
|
||||
email: Final = f"owner-{uuid.uuid4().hex[:8]}@example.com"
|
||||
team: Final = f"dropped-owner-team-{uuid.uuid4().hex[:8]}"
|
||||
api_key: Final = key_no_key_table_holds()
|
||||
rows: Final = (user_row(owner, api_key, DAY), seeded_row(TEAM_SPEND, "team_id", team, api_key, DAY))
|
||||
with (
|
||||
scratch_database() as database_url,
|
||||
dropped_connection_relay(database_url, OWNER_RECOVERY_TRIGGER) as (relay, relayed_url),
|
||||
_relayed_proxy(gateway, tmp_path, relayed_url) as owned,
|
||||
_reader(owned) as reader,
|
||||
daily_rows(rows, database_url=database_url),
|
||||
):
|
||||
owned.gateway.post("/user/new", {"user_id": owner, "user_email": email, "auto_create_key": False})
|
||||
relay.arm()
|
||||
dropped: Final = activity_of_key(reader, "/team/daily/activity", api_key, team_ids=team)
|
||||
assert relay.dropped.is_set(), dropped.text
|
||||
assert_key_reported(dropped, api_key, DAY, key_metadata(), seeded_metrics(1))
|
||||
relay.disarm()
|
||||
recovered: Final = activity_of_key(reader, "/team/daily/activity", api_key, team_ids=team)
|
||||
assert_key_reported(recovered, api_key, DAY, key_metadata(user=owner, email=email), seeded_metrics(1))
|
||||
|
||||
|
||||
@pytest.mark.timeout(300)
|
||||
def test_vantage_and_cloudzero_dry_runs_survive_a_dropped_database_connection_during_reverse_hash_recovery(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
owned_key: Final = _identity("dropped-dry-run")
|
||||
with (
|
||||
scratch_database() as database_url,
|
||||
dropped_connection_relay(database_url, owned_key.double.encode()) as (relay, relayed_url),
|
||||
_relayed_proxy(gateway, tmp_path, relayed_url) as owned,
|
||||
_reader(owned) as reader,
|
||||
daily_rows((user_row(owned_key.owner, owned_key.double, DAY),), database_url=database_url),
|
||||
):
|
||||
_register(owned.gateway, owned_key)
|
||||
relay.arm()
|
||||
vantage_dropped: Final = _dry_run(reader, "/vantage/dry-run")
|
||||
assert relay.dropped.is_set(), vantage_dropped
|
||||
_assert_vantage_dry_run(vantage_dropped, owned_key, alias=None)
|
||||
relay.dropped.clear()
|
||||
cloudzero_dropped: Final = _dry_run(reader, "/cloudzero/dry-run")
|
||||
assert relay.dropped.is_set(), cloudzero_dropped
|
||||
_assert_cloudzero_dry_run(cloudzero_dropped, owned_key, alias=None)
|
||||
relay.disarm()
|
||||
vantage: Final = _dry_run(reader, "/vantage/dry-run")
|
||||
_assert_vantage_dry_run(vantage, owned_key, alias=owned_key.alias)
|
||||
cloudzero: Final = _dry_run(reader, "/cloudzero/dry-run")
|
||||
_assert_cloudzero_dry_run(cloudzero, owned_key, alias=owned_key.alias)
|
||||
assert _without(_first(vantage_dropped, "normalized_data"), FOCUS_ALIAS_FIELDS) == _without(
|
||||
_first(vantage, "normalized_data"), FOCUS_ALIAS_FIELDS
|
||||
)
|
||||
assert _without(_first(cloudzero_dropped, "cbf_data"), CBF_ALIAS_FIELDS) == _without(
|
||||
_first(cloudzero, "cbf_data"), CBF_ALIAS_FIELDS
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.timeout(300)
|
||||
def test_vantage_and_cloudzero_dry_runs_survive_a_dropped_database_connection_during_user_detail_recovery(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
owned_key: Final = _identity("dropped-dry-run-detail")
|
||||
with (
|
||||
scratch_database() as database_url,
|
||||
dropped_connection_relay(database_url, owned_key.owner.encode()) as (relay, relayed_url),
|
||||
_relayed_proxy(gateway, tmp_path, relayed_url) as owned,
|
||||
_reader(owned) as reader,
|
||||
daily_rows((user_row(owned_key.owner, owned_key.double, DAY),), database_url=database_url),
|
||||
):
|
||||
_register(owned.gateway, owned_key)
|
||||
relay.arm()
|
||||
vantage_dropped: Final = _dry_run(reader, "/vantage/dry-run")
|
||||
assert relay.dropped.is_set(), vantage_dropped
|
||||
_assert_vantage_dry_run(vantage_dropped, owned_key, alias=owned_key.alias)
|
||||
relay.dropped.clear()
|
||||
cloudzero_dropped: Final = _dry_run(reader, "/cloudzero/dry-run")
|
||||
assert relay.dropped.is_set(), cloudzero_dropped
|
||||
_assert_cloudzero_dry_run(cloudzero_dropped, owned_key, alias=owned_key.alias)
|
||||
relay.disarm()
|
||||
vantage: Final = _dry_run(reader, "/vantage/dry-run")
|
||||
cloudzero: Final = _dry_run(reader, "/cloudzero/dry-run")
|
||||
assert _first(vantage_dropped, "normalized_data") == _first(vantage, "normalized_data")
|
||||
assert _first(cloudzero_dropped, "cbf_data") == _first(cloudzero, "cbf_data")
|
||||
|
||||
|
||||
@pytest.mark.timeout(420)
|
||||
def test_vantage_export_delivers_a_blank_alias_under_a_dropped_database_connection_and_reports_sink_errors(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
owned_key: Final = _identity("dropped-vantage-export")
|
||||
sink: Final = Sink(threading.Event(), threading.Event())
|
||||
api_key: Final = f"vantage-api-key-{uuid.uuid4().hex}"
|
||||
integration_token: Final = f"vantage-token-{uuid.uuid4().hex}"
|
||||
costs_target: Final = f"/v2/integrations/{integration_token}/costs.csv"
|
||||
with (
|
||||
scratch_database() as database_url,
|
||||
dropped_connection_relay(database_url, owned_key.double.encode()) as (relay, relayed_url),
|
||||
wire_server(sink.respond) as wire,
|
||||
_relayed_proxy(
|
||||
gateway, tmp_path, relayed_url, config=_config_allowing_a_base_url_in_the_body(tmp_path)
|
||||
) as owned,
|
||||
_reader(owned) as reader,
|
||||
daily_rows((user_row(owned_key.owner, owned_key.double, DAY),), database_url=database_url),
|
||||
):
|
||||
_register(owned.gateway, owned_key)
|
||||
owned.gateway.post(
|
||||
"/vantage/init", {"api_key": api_key, "integration_token": integration_token, "base_url": wire.url}
|
||||
)
|
||||
relay.arm()
|
||||
dropped: Final = reader.post("/vantage/export", {})
|
||||
assert relay.dropped.is_set(), dropped
|
||||
assert dropped["message"] == VANTAGE_DONE, dropped
|
||||
blank: Final = _upload(wire)
|
||||
assert (blank.target, blank.authorization) == (costs_target, f"Bearer {api_key}"), blank
|
||||
assert EXPORT_FILENAME.fullmatch(blank.filename), blank
|
||||
_assert_csv_row(blank.row, owned_key, alias=None)
|
||||
relay.dropped.clear()
|
||||
window: Final = _window()
|
||||
windowed: Final = reader.post("/vantage/export", _window_body(window))
|
||||
assert relay.dropped.is_set(), windowed
|
||||
assert windowed["message"] == VANTAGE_DONE, windowed
|
||||
bounded: Final = _upload(wire)
|
||||
assert bounded.filename == _window_filename(window), bounded
|
||||
_assert_csv_row(bounded.row, owned_key, alias=None)
|
||||
relay.disarm()
|
||||
healthy: Final = reader.post("/vantage/export", {})
|
||||
assert healthy["message"] == VANTAGE_DONE, healthy
|
||||
named: Final = _upload(wire)
|
||||
_assert_csv_row(named.row, owned_key, alias=owned_key.alias)
|
||||
assert _without(dict(blank.row), FOCUS_ALIAS_FIELDS) == _without(dict(named.row), FOCUS_ALIAS_FIELDS)
|
||||
sink.forbidden.set()
|
||||
forbidden: Final = reader.request("POST", "/vantage/export", {})
|
||||
assert forbidden.status_code == 500 and "Failed to perform Vantage export" in forbidden.text, forbidden.text
|
||||
assert _last_received(wire).target == costs_target
|
||||
sink.forbidden.clear()
|
||||
sink.missing.set()
|
||||
missing: Final = reader.request("POST", "/vantage/export", {})
|
||||
assert missing.status_code == 500 and "Failed to perform Vantage export" in missing.text, missing.text
|
||||
assert _last_received(wire).target == costs_target
|
||||
alive: Final = reader.request("GET", "/health/liveliness")
|
||||
assert alive.status_code == 200, alive.text
|
||||
|
||||
|
||||
@pytest.mark.timeout(420)
|
||||
def test_cloudzero_export_delivers_a_blank_alias_under_a_dropped_database_connection(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
owned_key: Final = _identity("dropped-cloudzero-export")
|
||||
api_key: Final = f"cloudzero-api-key-{uuid.uuid4().hex}"
|
||||
connection_id: Final = f"cloudzero-connection-{uuid.uuid4().hex[:8]}"
|
||||
drops_target: Final = f"/v2/connections/billing/anycost/{connection_id}/billing_drops"
|
||||
cert, key = write_self_signed_cert(tmp_path, (CLOUDZERO_HOST,))
|
||||
with (
|
||||
scratch_database() as database_url,
|
||||
dropped_connection_relay(database_url, owned_key.double.encode()) as (relay, relayed_url),
|
||||
wire_server(lambda request: Reply(), tls=server_context(cert, key)) as wire,
|
||||
_connect_tunnel(wire, CLOUDZERO_AUTHORITY) as tunnel_url,
|
||||
_relayed_proxy(
|
||||
gateway, tmp_path, relayed_url, {"HTTPS_PROXY": tunnel_url, "SSL_CERT_FILE": str(cert)}
|
||||
) as owned,
|
||||
_reader(owned) as reader,
|
||||
daily_rows((user_row(owned_key.owner, owned_key.double, DAY),), database_url=database_url),
|
||||
):
|
||||
_register(owned.gateway, owned_key)
|
||||
owned.gateway.post("/cloudzero/init", {"api_key": api_key, "connection_id": connection_id, "timezone": "UTC"})
|
||||
relay.arm()
|
||||
dropped: Final = reader.post("/cloudzero/export", {})
|
||||
assert relay.dropped.is_set(), dropped
|
||||
assert dropped["message"] == CLOUDZERO_DONE, dropped
|
||||
blank: Final = _last_received(wire)
|
||||
assert (blank.method, blank.target) == ("POST", drops_target), blank
|
||||
assert blank.headers["authorization"] == f"Bearer {api_key}", blank.headers
|
||||
assert json.loads(blank.body) == {
|
||||
"month": "2026-02",
|
||||
"operation": "replace_hourly",
|
||||
"data": [_cbf_record(owned_key, alias=None)],
|
||||
}, blank.body
|
||||
relay.dropped.clear()
|
||||
windowed: Final = reader.post("/cloudzero/export", _window_body(_window()))
|
||||
assert relay.dropped.is_set(), windowed
|
||||
assert windowed["message"] == CLOUDZERO_DONE, windowed
|
||||
assert json.loads(_last_received(wire).body) == json.loads(blank.body)
|
||||
relay.disarm()
|
||||
healthy: Final = reader.post("/cloudzero/export", {})
|
||||
assert healthy["message"] == CLOUDZERO_DONE, healthy
|
||||
named: Final = _last_received(wire)
|
||||
assert json.loads(named.body) == {
|
||||
"month": "2026-02",
|
||||
"operation": "replace_hourly",
|
||||
"data": [_cbf_record(owned_key, alias=owned_key.alias)],
|
||||
}, named.body
|
||||
28
tests/unit/integration_support/test_database_relay.py
Normal file
28
tests/unit/integration_support/test_database_relay.py
Normal file
|
|
@ -0,0 +1,28 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
|
||||
from tests.integration._support.database_relay import TriggerScanner
|
||||
|
||||
TRIGGER: Final = b'SELECT "startTime" FROM "LiteLLM_SpendLogs"'
|
||||
|
||||
|
||||
@pytest.mark.parametrize("split_at", range(1, len(TRIGGER)))
|
||||
def test_trigger_scanner_matches_a_trigger_split_across_two_reads(split_at: int) -> None:
|
||||
scanner: Final = TriggerScanner(TRIGGER)
|
||||
assert not scanner.feed(TRIGGER[:split_at])
|
||||
assert scanner.feed(TRIGGER[split_at:])
|
||||
|
||||
|
||||
def test_trigger_scanner_matches_a_trigger_arriving_one_byte_at_a_time() -> None:
|
||||
scanner: Final = TriggerScanner(TRIGGER)
|
||||
hits: Final = tuple(scanner.feed(TRIGGER[i : i + 1]) for i in range(len(TRIGGER)))
|
||||
assert hits == (False,) * (len(TRIGGER) - 1) + (True,)
|
||||
|
||||
|
||||
def test_trigger_scanner_reports_a_match_once() -> None:
|
||||
scanner: Final = TriggerScanner(TRIGGER)
|
||||
assert scanner.feed(b"x" + TRIGGER + b"y")
|
||||
assert not scanner.feed(b"z")
|
||||
|
|
@ -443,6 +443,65 @@ async def test_recover_key_metadata_from_spend_logs_retries_a_failed_query_only_
|
|||
assert result[digest]["key_alias"] == "back-online"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_recover_key_metadata_from_spend_logs_never_caches_repeated_query_failures_as_long_as_a_hit():
|
||||
digest: Final = hash_token("cli-session-repeated-failure")
|
||||
window: Final = (datetime(2026, 9, 7), datetime(2026, 9, 10))
|
||||
cache: Final = InMemoryCache(default_ttl=SPEND_LOG_KEY_METADATA_CACHE_TTL)
|
||||
mock_prisma: Final = MagicMock()
|
||||
_spend_log_transaction(
|
||||
mock_prisma,
|
||||
AsyncMock(
|
||||
side_effect=[
|
||||
PrismaError("statement timeout"),
|
||||
PrismaError("statement timeout"),
|
||||
[_spend_log_row(digest, "back-online", None, None)],
|
||||
]
|
||||
),
|
||||
)
|
||||
await recover_key_metadata_from_spend_logs(mock_prisma, {digest}, window, cache=cache)
|
||||
miss_key: Final = next(key for key in cache.ttl_dict if digest in key and not key.endswith(":missed-before"))
|
||||
cache.ttl_dict[miss_key] = time.time() - 1
|
||||
second_query_started: Final = time.time()
|
||||
|
||||
await recover_key_metadata_from_spend_logs(mock_prisma, {digest}, window, cache=cache)
|
||||
|
||||
assert cache.ttl_dict[miss_key] - second_query_started <= SPEND_LOG_KEY_METADATA_MISS_CACHE_TTL + 1
|
||||
cache.ttl_dict[miss_key] = time.time() - 1
|
||||
|
||||
result: Final = await recover_key_metadata_from_spend_logs(mock_prisma, {digest}, window, cache=cache)
|
||||
|
||||
assert result[digest]["key_alias"] == "back-online"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_recover_key_metadata_from_spend_logs_treats_a_dropped_connection_error_as_a_short_lived_miss():
|
||||
digest: Final = hash_token("cli-session-dropped-connection")
|
||||
window: Final = (datetime(2026, 9, 7), datetime(2026, 9, 10))
|
||||
cache: Final = InMemoryCache(default_ttl=SPEND_LOG_KEY_METADATA_CACHE_TTL)
|
||||
mock_prisma: Final = MagicMock()
|
||||
_spend_log_transaction(
|
||||
mock_prisma,
|
||||
AsyncMock(
|
||||
side_effect=[
|
||||
AttributeError("'NoneType' object has no attribute 'get'"),
|
||||
[_spend_log_row(digest, "back-online", None, None)],
|
||||
]
|
||||
),
|
||||
)
|
||||
started: Final = time.time()
|
||||
|
||||
assert await recover_key_metadata_from_spend_logs(mock_prisma, {digest}, window, cache=cache) == {}
|
||||
|
||||
miss_key: Final = next(key for key in cache.ttl_dict if digest in key and not key.endswith(":missed-before"))
|
||||
assert cache.ttl_dict[miss_key] - started <= SPEND_LOG_KEY_METADATA_MISS_CACHE_TTL + 1
|
||||
assert f"{miss_key}:missed-before" not in cache.ttl_dict
|
||||
cache.ttl_dict[miss_key] = time.time() - 1
|
||||
|
||||
result: Final = await recover_key_metadata_from_spend_logs(mock_prisma, {digest}, window, cache=cache)
|
||||
|
||||
assert result[digest]["key_alias"] == "back-online"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_recover_key_metadata_from_spend_logs_drops_the_owner_of_a_digest_shared_by_several_users():
|
||||
shared_ui_digest = hash_token("ui-token")
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue