From 063c3dec77b9d7d5ec86338d777f840cbfa2040b Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 2 Oct 2026 18:06:30 -0700 Subject: [PATCH] test(integration): audit durable background interaction settlement across replicas Twenty-six deterministic cells drive a one-worker creator and a two-worker settler against an owned scripted Gemini upstream: cross-replica deletes bill once, failed and cancelled interactions release, a later replica resumes unclaimed rows, custom deployment pricing bills at the deployment rate, a fetch the settler cannot make fails the delete closed, odd ids are refused, a missing settlement table keeps in-process billing, polling disabled registers nothing, the budget reservation is released by the settler, an upstream outage mid-burst fails closed and recovers, killed workers hand their polls to the respawned ones, and concurrent deletes on a slow upstream settle exactly once. The support upstream gains a scripted interaction store with per-id GET status and delay, and the process helper gains an owned upstream a test can stop and restart --- tests/integration/_support/process.py | 64 +- tests/integration/_support/upstream.py | 170 +++- .../test_background_interaction_settlement.py | 775 ++++++++++++++++++ 3 files changed, 989 insertions(+), 20 deletions(-) create mode 100644 tests/integration/spend/test_background_interaction_settlement.py diff --git a/tests/integration/_support/process.py b/tests/integration/_support/process.py index 8cfdf0db2c3..9787c1fba80 100644 --- a/tests/integration/_support/process.py +++ b/tests/integration/_support/process.py @@ -5,7 +5,7 @@ import subprocess import sys import time import uuid -from collections.abc import Iterator, Mapping +from collections.abc import Generator, Iterator, Mapping from contextlib import contextmanager from dataclasses import dataclass from pathlib import Path @@ -220,3 +220,65 @@ def owned_proxy_process( yield OwnedProxy(Gateway(client, gateway.key, gateway.upstream_url), process, launch.log) finally: _stop(process) + + +_UPSTREAM_READY_SECONDS: Final = 60 + + +class UpstreamSlot: + """A scripted upstream a test module owns on a fixed port, so a cell can take it down and bring it back.""" + + __slots__ = ("directory", "port", "process", "root") + + def __init__(self, directory: Path, port: int, root: Path) -> None: + self.directory = directory + self.port = port + self.root = root + self.process: subprocess.Popen[bytes] | None = None + + @property + def url(self) -> str: + return f"http://127.0.0.1:{self.port}" + + def start(self) -> None: + assert self.process is None, "Owned upstream is already running" + output: Final = Path(os.environ.get("INTEGRATION_RESULTS_DIR") or self.directory) + log_path: Final = output / f"owned-upstream-{self.port}-{uuid.uuid4().hex}.log" + with log_path.open("w") as log: + process: Final = subprocess.Popen( + [sys.executable, "-m", "integration._support.upstream", "--port", str(self.port)], + cwd=self.root, + env=dict(os.environ), + stdout=log, + stderr=subprocess.STDOUT, + start_new_session=True, + ) + self.process = process + deadline: Final = time.monotonic() + _UPSTREAM_READY_SECONDS + while process.poll() is None: + try: + if httpx.get(f"{self.url}/health", timeout=2, trust_env=False).status_code == 200: + return + except httpx.TransportError: + pass + assert time.monotonic() < deadline, f"Owned upstream readiness deadline exceeded: {log_path}" + time.sleep(0.1) + raise AssertionError(f"Owned upstream exited before readiness: {log_path}") + + def stop(self) -> None: + process: Final = self.process + assert process is not None, "Owned upstream is not running" + self.process = None + _stop(process) + + +@contextmanager +def owned_upstream(directory: Path) -> Generator[UpstreamSlot]: + root: Final = Path(os.environ.get("INTEGRATION_PROXY_ROOT") or Path(__file__).resolve().parents[3]) + slot: Final = UpstreamSlot(directory, _free_port(), root) + slot.start() + try: + yield slot + finally: + if slot.process is not None: + slot.stop() diff --git a/tests/integration/_support/upstream.py b/tests/integration/_support/upstream.py index e645126032b..b21ec10254f 100644 --- a/tests/integration/_support/upstream.py +++ b/tests/integration/_support/upstream.py @@ -29,7 +29,7 @@ from integration.cost_calculation.cost_tracking_case import ( StoredResponse, TextResponse, ) -from pydantic import BaseModel, JsonValue, TypeAdapter, ValidationError +from pydantic import BaseModel, ConfigDict, JsonValue, TypeAdapter, ValidationError from starlette.applications import Starlette from starlette.requests import Request from starlette.responses import JSONResponse, Response, StreamingResponse @@ -64,6 +64,19 @@ class Observation: path: str authorization: str body: dict[str, JsonValue] + method: str = "POST" + api_key: str = "" + + +class InteractionState(BaseModel): + """What the scripted Interactions API answers for one interaction id until a DELETE drops it.""" + + model_config = ConfigDict(extra="forbid") + + status: str + usage: dict[str, JsonValue] | None = None + get_status: int = 200 + delay_seconds: float = 0 class _ScenarioRegistration(BaseModel): @@ -89,9 +102,12 @@ def _aws_event_frame( scenario_id: str, unique_id: str, ) -> bytes: - payload_bytes: Final = json.dumps(payload, separators=(",", ":")).replace( - "$REQUEST_ID", scenario_id - ).replace("$UNIQUE_ID", unique_id).encode() + payload_bytes: Final = ( + json.dumps(payload, separators=(",", ":")) + .replace("$REQUEST_ID", scenario_id) + .replace("$UNIQUE_ID", unique_id) + .encode() + ) headers_bytes: Final = ( _aws_str_header(":event-type", event_type) + _aws_str_header(":content-type", "application/json") @@ -123,6 +139,7 @@ class Provider: observations: SimpleQueue[Observation] = field(default_factory=SimpleQueue) scripts: dict[str, deque[int]] = field(default_factory=dict) scenario_store: ScenarioStore = field(default_factory=ScenarioStore) + interactions: dict[str, InteractionState] = field(default_factory=dict) async def chat(self, request: Request) -> Response: body: Final = JSON_OBJECT.validate_json(await request.body()) @@ -147,7 +164,13 @@ class Provider: status: Final = script.popleft() if status != 200: return JSONResponse( - {"error": {"message": "Controlled provider failure", "type": error_type(status), "code": str(status)}}, + { + "error": { + "message": "Controlled provider failure", + "type": error_type(status), + "code": str(status), + } + }, status_code=status, ) return await chat_completions(request) @@ -198,7 +221,14 @@ class Provider: return JSONResponse( { "requests": [ - {"path": value.path, "authorization": value.authorization, "body": value.body} for value in values + { + "path": value.path, + "authorization": value.authorization, + "body": value.body, + "method": value.method, + "api_key": value.api_key, + } + for value in values ] } ) @@ -245,7 +275,12 @@ class Provider: body: Final = JSON_OBJECT.validate_json(raw_body) if isinstance(body, dict): self.observations.put( - Observation(request.url.path, request.headers.get("authorization", ""), body) + Observation( + request.url.path, + request.headers.get("authorization", ""), + body, + api_key=request.headers.get("x-goog-api-key", ""), + ) ) if isinstance(response, RoutedResponse): route_key: Final = f"{request.method} /{'/'.join(segments[1:])}" @@ -262,6 +297,57 @@ class Provider: return self._response(route, scenario_id) return self._response(response, scenario_id) + async def interaction_state(self, request: Request) -> Response: + interaction_id: Final = cast(str, request.path_params["interaction_id"]) + if request.method == "DELETE": + self.interactions.pop(interaction_id, None) + return JSONResponse({"interaction_id": interaction_id, "registered": False}) + try: + state: Final = InteractionState.model_validate_json(await request.body()) + except ValidationError as exc: + return JSONResponse({"error": str(exc)}, status_code=400) + self.interactions[interaction_id] = state + return JSONResponse({"interaction_id": interaction_id, "registered": True}) + + def _observe_interaction(self, request: Request) -> None: + self.observations.put( + Observation( + request.url.path, + request.headers.get("authorization", ""), + {}, + method=request.method, + api_key=request.headers.get("x-goog-api-key", ""), + ) + ) + + async def interaction(self, request: Request) -> Response: + self._observe_interaction(request) + interaction_id: Final = cast(str, request.path_params["interaction_id"]) + state: Final = self.interactions.get(interaction_id) + if state is None: + return JSONResponse(_interaction_not_found(interaction_id), status_code=404) + if state.delay_seconds: + await asyncio.sleep(state.delay_seconds) + if request.method == "DELETE": + del self.interactions[interaction_id] + return JSONResponse({}) + if state.get_status != 200: + return JSONResponse( + {"error": {"code": state.get_status, "message": "Scripted interaction fetch failure"}}, + status_code=state.get_status, + ) + return JSONResponse(_interaction_body(interaction_id, state)) + + async def cancel_interaction(self, request: Request) -> Response: + self._observe_interaction(request) + interaction_id: Final = cast(str, request.path_params["interaction_id"]) + state: Final = self.interactions.get(interaction_id) + if state is None: + return JSONResponse(_interaction_not_found(interaction_id), status_code=404) + cancelled: Final = InteractionState(status="cancelled", usage=state.usage, get_status=state.get_status) + self.interactions[interaction_id] = cancelled + return JSONResponse(_interaction_body(interaction_id, cancelled)) + async def realtime(self, websocket: WebSocket) -> None: scenario_id: Final = websocket.headers.get("authorization", "").removeprefix("Bearer ") response: Final = self.scenario_store.get(scenario_id) @@ -300,11 +386,10 @@ class Provider: match response: case JsonResponse(): return Response( - content=json.dumps(response.body, separators=(",", ":")).replace( - "$REQUEST_ID", scenario_id - ).replace( - "$UNIQUE_ID", unique_id - ).encode(), + content=json.dumps(response.body, separators=(",", ":")) + .replace("$REQUEST_ID", scenario_id) + .replace("$UNIQUE_ID", unique_id) + .encode(), media_type=response.content_type, status_code=response.status, ) @@ -321,6 +406,7 @@ class Provider: ) case SseResponse(): if response.frame_delay_ms > 0: + async def stream() -> AsyncIterator[bytes]: for frame in response.frames: yield ( @@ -329,9 +415,11 @@ class Provider: await asyncio.sleep(response.frame_delay_ms / 1000) return StreamingResponse(stream(), media_type=response.content_type) - stream_body: Final = ("\n\n".join(response.frames) + "\n\n").replace( - "$REQUEST_ID", scenario_id - ).replace("$UNIQUE_ID", unique_id) + stream_body: Final = ( + ("\n\n".join(response.frames) + "\n\n") + .replace("$REQUEST_ID", scenario_id) + .replace("$UNIQUE_ID", unique_id) + ) return Response(content=stream_body.encode(), media_type=response.content_type) case EventStreamResponse(): events: Final = ( @@ -372,6 +460,19 @@ class Provider: Route("/v1/embeddings", embeddings, methods=["POST"]), Route("/v1/moderations", moderations, methods=["POST"]), Route("/vector_stores/{vector_store_id}/search", self.vector_store_search, methods=["POST"]), + Route("/__interactions/{interaction_id}", self.interaction_state, methods=["PUT", "DELETE"]), + Route("/v1beta/interactions/{interaction_id}:cancel", self.cancel_interaction, methods=["POST"]), + Route("/v1beta/interactions/{interaction_id}", self.interaction, methods=["GET", "DELETE"]), + Route( + "/{prefix:path}/v1beta/interactions/{interaction_id}:cancel", + self.cancel_interaction, + methods=["POST"], + ), + Route( + "/{prefix:path}/v1beta/interactions/{interaction_id}", + self.interaction, + methods=["GET", "DELETE"], + ), Route("/{path:path}", self.scripted, methods=["POST"]), Route("/{path:path}", self.scripted, methods=["GET"]), WebSocketRoute("/v1/realtime", self.realtime), @@ -382,6 +483,21 @@ class Provider: CONTROL_URL: Final = os.environ.get("INTEGRATION_UPSTREAM_URL", "http://127.0.0.1:8190").rstrip("/") +def _interaction_not_found(interaction_id: str) -> dict[str, JsonValue]: + return {"error": {"code": 404, "message": f"Interaction {interaction_id} not found", "status": "NOT_FOUND"}} + + +def _interaction_body(interaction_id: str, state: InteractionState) -> dict[str, JsonValue]: + return { + "id": interaction_id, + "object": "interaction", + "model": "gemini-3.8-flash", + "status": state.status, + "steps": [], + "usage": state.usage, + } + + @dataclass(frozen=True, slots=True) class ScenarioHandle: scenario_id: str @@ -391,9 +507,9 @@ class ScenarioHandle: return f"{self.control_url}/{self.scenario_id}" -def register_scenario(scenario_id: str, response: StoredResponse) -> ScenarioHandle: +def register_scenario(scenario_id: str, response: StoredResponse, *, control_url: str = CONTROL_URL) -> ScenarioHandle: http_response: Final = httpx.post( - f"{CONTROL_URL}/__scenarios", + f"{control_url}/__scenarios", json={"scenario_id": scenario_id, "response": response.model_dump(mode="json")}, trust_env=False, timeout=15, @@ -401,19 +517,35 @@ def register_scenario(scenario_id: str, response: StoredResponse) -> ScenarioHan http_response.raise_for_status() return ScenarioHandle( scenario_id=scenario_id, - control_url=CONTROL_URL, + control_url=control_url, ) def delete_scenario(handle: ScenarioHandle) -> None: response: Final = httpx.delete( - f"{CONTROL_URL}/__scenarios/{handle.scenario_id}", + f"{handle.control_url}/__scenarios/{handle.scenario_id}", trust_env=False, timeout=15, ) response.raise_for_status() +def set_interaction_state(control_url: str, interaction_id: str, state: InteractionState) -> None: + response: Final = httpx.put( + f"{control_url}/__interactions/{interaction_id}", + content=state.model_dump_json(), + headers={"content-type": "application/json"}, + trust_env=False, + timeout=15, + ) + response.raise_for_status() + + +def clear_interaction_state(control_url: str, interaction_id: str) -> None: + response: Final = httpx.delete(f"{control_url}/__interactions/{interaction_id}", trust_env=False, timeout=15) + response.raise_for_status() + + def main() -> None: parser: Final = argparse.ArgumentParser() parser.add_argument("--port", type=int, default=8190) diff --git a/tests/integration/spend/test_background_interaction_settlement.py b/tests/integration/spend/test_background_interaction_settlement.py new file mode 100644 index 00000000000..2b11a27e4f0 --- /dev/null +++ b/tests/integration/spend/test_background_interaction_settlement.py @@ -0,0 +1,775 @@ +import math +import socket +import uuid +from collections.abc import Iterator, Mapping, Sequence +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass +from pathlib import Path +from typing import Final + +import httpx +import psutil +import pytest +import yaml +from integration._support.client import ( + JSON_OBJECT, + Gateway, + eventually, + gateway_from_environment, + object_value, + string_value, +) +from integration._support.database import read_rows, write_rows +from integration._support.process import ( + UpstreamSlot, + group_members, + owned_proxy, + owned_proxy_process, + owned_upstream, +) +from integration._support.upstream import ( + InteractionState, + clear_interaction_state, + register_scenario, + set_interaction_state, +) +from integration.cost_calculation.cost_tracking_case import JsonResponse, RoutedResponse +from pydantic import JsonValue + +from litellm.proxy.spend_tracking.budget_reservation import DEFAULT_MAX_OUTPUT_TOKENS_FALLBACK + +pytestmark: Final = pytest.mark.timeout(900) + +_MODEL: Final = "gemini/gemini-3.8-flash" +_INPUT_TOKENS: Final = 300 +_OUTPUT_TOKENS: Final = 41 +_USAGE: Final[dict[str, JsonValue]] = { + "total_input_tokens": _INPUT_TOKENS, + "total_output_tokens": _OUTPUT_TOKENS, + "total_tool_use_tokens": 0, + "total_reasoning_tokens": 0, +} +_CUSTOM_INPUT_RATE: Final = 2e-06 +_CUSTOM_OUTPUT_RATE: Final = 4e-05 +_ENV_KEY: Final = "integration-gemini-env-key" +_DEPLOYMENT_KEY: Final = "integration-gemini-deployment-key" +_CREATOR_POLL: Final = {"BACKGROUND_INTERACTION_COST_POLL_INITIAL_INTERVAL_SECONDS": "300"} +_SETTLER_POLL: Final = { + "BACKGROUND_INTERACTION_COST_POLL_INITIAL_INTERVAL_SECONDS": "1", + "BACKGROUND_INTERACTION_COST_POLL_MAX_INTERVAL_SECONDS": "1", + "BACKGROUND_INTERACTION_COST_POLL_TIMEOUT_SECONDS": "8", +} +_RESUMER_POLL: Final = { + "BACKGROUND_INTERACTION_COST_POLL_INITIAL_INTERVAL_SECONDS": "1", + "BACKGROUND_INTERACTION_COST_POLL_MAX_INTERVAL_SECONDS": "1", + "BACKGROUND_INTERACTION_COST_POLL_TIMEOUT_SECONDS": "120", +} +_SPEND_QUERY: Final = ( + "SELECT request_id, spend, call_type, status, model, prompt_tokens, completion_tokens " + 'FROM "LiteLLM_SpendLogs" WHERE request_id = %s' +) +_SETTLEMENT_QUERY: Final = ( + "SELECT interaction_id, claimed_by, outcome, claimed_at IS NOT NULL AS claimed, " + 'settled_at IS NOT NULL AS settled, create_context FROM "LiteLLM_BackgroundInteractionSettlement" ' + "WHERE interaction_id = %s" +) +_SETTLEMENT_TABLE_PRESENT_QUERY: Final = "SELECT to_regclass(%s) IS NOT NULL AS present" +_SETTLEMENT_TABLE: Final = '"LiteLLM_BackgroundInteractionSettlement"' +_SETTLEMENT_BY_CALL_QUERY: Final = ( + 'SELECT interaction_id FROM "LiteLLM_BackgroundInteractionSettlement" WHERE create_context->>%s = %s' +) +_OUTAGE_RENAME: Final = ( + 'ALTER TABLE IF EXISTS "LiteLLM_BackgroundInteractionSettlement" ' + 'RENAME TO "LiteLLM_BackgroundInteractionSettlement_outage"' +) +_OUTAGE_RESTORE: Final = ( + 'ALTER TABLE IF EXISTS "LiteLLM_BackgroundInteractionSettlement_outage" ' + 'RENAME TO "LiteLLM_BackgroundInteractionSettlement"' +) + + +@dataclass(frozen=True, slots=True) +class Deployments: + """Config deployments every replica boots with, so no worker ever misses a model added at run time.""" + + in_progress: str + completed_at_once: str + failing_create: str + custom_priced: str + + +@dataclass(frozen=True, slots=True) +class Rig: + gateway: Gateway + upstream: UpstreamSlot + config: Path + models: Deployments + creator: Gateway + settler: Gateway + settler_pid: int + directory: Path + + def environment(self, **poll: str) -> dict[str, str]: + return {"GEMINI_API_BASE": self.upstream.url, "GEMINI_API_KEY": _ENV_KEY, **poll} + + +@pytest.fixture(scope="module") +def rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Rig]: + directory: Final = tmp_path_factory.mktemp("settlement") + with gateway_from_environment() as gateway, owned_upstream(directory) as upstream: + models: Final = _register_deployments(upstream.url) + config: Final = _write_config(directory, upstream.url, models) + environment: Final = {"GEMINI_API_BASE": upstream.url, "GEMINI_API_KEY": _ENV_KEY} + with ( + owned_proxy(gateway, directory, {**environment, **_CREATOR_POLL}, config=config, workers=1) as creator, + owned_proxy_process( + gateway, directory, {**environment, **_SETTLER_POLL}, config=config, workers=2 + ) as settler, + ): + yield Rig(gateway, upstream, config, models, creator, settler.gateway, settler.process.pid, directory) + + +def _register_deployments(upstream_url: str) -> Deployments: + suffix: Final = uuid.uuid4().hex[:8] + models: Final = Deployments( + in_progress=f"settle-in-progress-{suffix}", + completed_at_once=f"settle-completed-at-once-{suffix}", + failing_create=f"settle-failing-create-{suffix}", + custom_priced=f"settle-custom-priced-{suffix}", + ) + _register_scenarios(upstream_url, models) + return models + + +def _register_scenarios(upstream_url: str, models: Deployments) -> None: + scripted: Final = { + models.in_progress: _interaction("in_progress", None), + models.completed_at_once: _interaction("completed", _USAGE), + models.failing_create: JsonResponse( + content_type="application/json", body={"error": {"message": "boom"}}, status=500 + ), + models.custom_priced: _interaction("in_progress", None), + } + for name, response in scripted.items(): + register_scenario( + name, + RoutedResponse(content_type="application/x-routed", routes={"POST /v1beta/interactions": response}), + control_url=upstream_url, + ) + + +def _write_config(directory: Path, upstream_url: str, models: Deployments) -> Path: + base: Final = JSON_OBJECT.validate_python(yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())) + custom_pricing: Final = {"input_cost_per_token": _CUSTOM_INPUT_RATE, "output_cost_per_token": _CUSTOM_OUTPUT_RATE} + model_list: Final = [ + { + "model_name": name, + "litellm_params": { + "model": _MODEL, + "api_base": f"{upstream_url}/{name}", + "api_key": _DEPLOYMENT_KEY, + **(custom_pricing if name == models.custom_priced else {}), + }, + } + for name in (models.in_progress, models.completed_at_once, models.failing_create, models.custom_priced) + ] + path: Final = directory / "settlement_config.yaml" + path.write_text(yaml.safe_dump({**base, "model_list": model_list})) + return path + + +def _interaction(status: str, usage: dict[str, JsonValue] | None, http_status: int = 200) -> JsonResponse: + return JsonResponse( + content_type="application/json", + body={ + "id": "$UNIQUE_ID", + "object": "interaction", + "model": "gemini-3.8-flash", + "status": status, + "steps": [], + "usage": usage, + }, + status=http_status, + ) + + +def _completed() -> InteractionState: + return InteractionState(status="completed", usage=_USAGE) + + +def _create( + replica: Gateway, + model: str, + key: str, + *, + path: str = "/v1beta/interactions", + background: bool = True, + text: str | None = None, +) -> str: + response: Final = replica.request( + "POST", + path, + {"model": model, "input": text or f"settle {uuid.uuid4().hex}", "background": background}, + key=key, + ) + assert response.status_code == 200, response.text + return string_value(JSON_OBJECT.validate_json(response.content)["id"]) + + +def _state(rig: Rig, interaction_id: str, state: InteractionState) -> None: + set_interaction_state(rig.upstream.url, interaction_id, state) + + +def _delete(replica: Gateway, interaction_id: str, key: str, *, path: str = "/v1beta/interactions") -> httpx.Response: + return replica.request("DELETE", f"{path}/{interaction_id}", key=key) + + +def _delete_ok(replica: Gateway, interaction_id: str, key: str) -> None: + deleted: Final = _delete(replica, interaction_id, key) + assert deleted.status_code == 200, deleted.text + + +def _delete_concurrently(replica: Gateway, interaction_ids: Sequence[str], key: str) -> tuple[int, ...]: + def status(interaction_id: str) -> int: + return _delete(replica, interaction_id, key).status_code + + with ThreadPoolExecutor(max_workers=8) as pool: + return tuple(pool.map(status, interaction_ids)) + + +def _assert_unclaimed(interaction_id: str) -> None: + row: Final = _settlement(interaction_id) + assert row is not None and row["claimed"] is False and row["outcome"] is None, row + + +def _spend_rows(request_id: str) -> list[dict[str, JsonValue]]: + return read_rows(_SPEND_QUERY, (request_id,)) + + +def _settlement(interaction_id: str) -> dict[str, JsonValue] | None: + rows: Final = read_rows(_SETTLEMENT_QUERY, (interaction_id,)) + return rows[0] if rows else None + + +def _settlement_table_present() -> bool: + return read_rows(_SETTLEMENT_TABLE_PRESENT_QUERY, (_SETTLEMENT_TABLE,))[0]["present"] is True + + +def _settlement_if_stored(interaction_id: str) -> dict[str, JsonValue] | None: + return _settlement(interaction_id) if _settlement_table_present() else None + + +def _settlements_by_call_if_stored(call_id: str) -> list[dict[str, JsonValue]]: + return read_rows(_SETTLEMENT_BY_CALL_QUERY, ("litellm_call_id", call_id)) if _settlement_table_present() else [] + + +def _await_spend_row(interaction_id: str, seconds: float = 30) -> dict[str, JsonValue]: + return eventually(lambda: _spend_rows(interaction_id), lambda rows: len(rows) == 1, seconds=seconds)[0] + + +def _await_outcome(interaction_id: str, outcome: str, seconds: float = 30) -> dict[str, JsonValue]: + row: Final = eventually( + lambda: _settlement(interaction_id), + lambda value: value is not None and value["outcome"] == outcome, + seconds=seconds, + ) + assert row is not None + return row + + +def _model_info(replica: Gateway, model: str) -> Mapping[str, JsonValue]: + entries: Final = replica.get("/model/info")["data"] + assert isinstance(entries, list), entries + return object_value( + next(object_value(entry)["model_info"] for entry in entries if object_value(entry)["model_name"] == model) + ) + + +def _rates(replica: Gateway, model: str) -> tuple[float, float]: + info: Final = _model_info(replica, model) + input_rate: Final = info["input_cost_per_token"] + output_rate: Final = info["output_cost_per_token"] + assert isinstance(input_rate, float) and isinstance(output_rate, float), info + return input_rate, output_rate + + +def _reservation_pin(replica: Gateway, model: str) -> float: + """What one background create reserves on the key before its usage is known: the output tokens the + estimator assumes, at the deployment's output rate, with the prompt's few input tokens left as slack.""" + info: Final = _model_info(replica, model) + max_output: Final = info["max_output_tokens"] + output_rate: Final = info["output_cost_per_token"] + assert isinstance(max_output, int) and isinstance(output_rate, float), info + return min(max_output, DEFAULT_MAX_OUTPUT_TOKENS_FALLBACK) * output_rate + + +def _assert_billed(row: Mapping[str, JsonValue], rates: tuple[float, float]) -> float: + expected: Final = _INPUT_TOKENS * rates[0] + _OUTPUT_TOKENS * rates[1] + spend: Final = row["spend"] + assert isinstance(spend, float) and math.isclose(spend, expected, rel_tol=1e-9), (row, expected) + assert row["call_type"] == "acreate_interaction", row + assert row["status"] == "success", row + assert row["prompt_tokens"] == _INPUT_TOKENS and row["completion_tokens"] == _OUTPUT_TOKENS, row + return spend + + +def _key_spend(replica: Gateway, key: str) -> float: + spend: Final = object_value(replica.get("/key/info", {"key": key})["info"])["spend"] + assert isinstance(spend, float | int), spend + return float(spend) + + +def _await_key_spend(replica: Gateway, key: str, expected: float) -> None: + eventually(lambda: _key_spend(replica, key), lambda spend: math.isclose(spend, expected, rel_tol=1e-9), seconds=30) + + +def _drain(rig: Rig) -> list[JsonValue]: + observed: Final = httpx.get(f"{rig.upstream.url}/__observations", trust_env=False, timeout=15) + observed.raise_for_status() + requests: Final = JSON_OBJECT.validate_json(observed.content)["requests"] + assert isinstance(requests, list), requests + return requests + + +def _calls(rig: Rig, interaction_id: str) -> tuple[tuple[str, str], ...]: + suffix: Final = f"/v1beta/interactions/{interaction_id}" + return tuple( + (string_value(object_value(entry)["method"]), string_value(object_value(entry)["api_key"])) + for entry in _drain(rig) + if string_value(object_value(entry)["path"]).endswith(suffix) + ) + + +def _claimer_pid(row: Mapping[str, JsonValue]) -> int: + claimed_by: Final = string_value(row["claimed_by"]) + host, _, pid = claimed_by.rpartition(":") + assert host == socket.gethostname(), claimed_by + return int(pid) + + +def _worker_pids(root_pid: int) -> frozenset[int]: + return frozenset( + process.pid for process in group_members(root_pid) if process.pid != root_pid and _is_spawned_worker(process) + ) + + +def _is_spawned_worker(process: psutil.Process) -> bool: + try: + return process.name().lower().startswith("python") and "resource_tracker" not in " ".join(process.cmdline()) + except psutil.Error: + return False + + +def _readiness(replica: Gateway) -> int: + try: + return replica.request("GET", "/health/readiness").status_code + except httpx.TransportError: + return 0 + + +def test_creator_poll_bills_a_completed_background_interaction_once(rig: Rig) -> None: + with rig.settler.scenario() as scenario: + model: Final = rig.models.in_progress + key: Final = scenario.key() + created: Final = _create(rig.settler, model, key) + _state(rig, created, _completed()) + spend: Final = _assert_billed(_await_spend_row(created), _rates(rig.settler, model)) + _await_key_spend(rig.settler, key, spend) + assert len(_spend_rows(created)) == 1 + + +def test_creator_poll_records_its_settlement_durably(rig: Rig) -> None: + with rig.settler.scenario() as scenario: + model: Final = rig.models.in_progress + created: Final = _create(rig.settler, model, scenario.key()) + _state(rig, created, _completed()) + _await_spend_row(created) + row: Final = _await_outcome(created, "billed") + assert row["claimed"] is True and row["settled"] is True, row + assert row["create_context"] == {}, row + _claimer_pid(row) + + +@pytest.mark.parametrize("path", ["/v1beta/interactions", "/interactions"]) +def test_delete_on_another_replica_bills_the_creators_interaction_once(rig: Rig, path: str) -> None: + with rig.creator.scenario() as scenario: + model: Final = rig.models.in_progress + key: Final = scenario.key() + created: Final = _create(rig.creator, model, key, path=path) + _state(rig, created, _completed()) + _drain(rig) + deleted: Final = _delete(rig.settler, created, key, path=path) + assert deleted.status_code == 200, deleted.text + spend: Final = _assert_billed(_await_spend_row(created), _rates(rig.creator, model)) + row: Final = _await_outcome(created, "billed") + assert _claimer_pid(row) in _worker_pids(rig.settler_pid), row + assert _calls(rig, created) == (("GET", _ENV_KEY), ("DELETE", _ENV_KEY)) + _await_key_spend(rig.creator, key, spend) + assert len(_spend_rows(created)) == 1 + + +def test_delete_of_a_failed_interaction_releases_without_a_spend_row(rig: Rig) -> None: + with rig.creator.scenario() as scenario: + model: Final = rig.models.in_progress + key: Final = scenario.key() + created: Final = _create(rig.creator, model, key) + _state(rig, created, InteractionState(status="failed", usage=None)) + deleted: Final = _delete(rig.settler, created, key) + assert deleted.status_code == 200, deleted.text + _await_outcome(created, "released") + assert _spend_rows(created) == [] + assert _key_spend(rig.creator, key) == 0 + + +def test_delete_of_a_requires_action_interaction_bills_its_usage(rig: Rig) -> None: + with rig.creator.scenario() as scenario: + model: Final = rig.models.in_progress + key: Final = scenario.key() + created: Final = _create(rig.creator, model, key) + _state(rig, created, InteractionState(status="requires_action", usage=_USAGE)) + deleted: Final = _delete(rig.settler, created, key) + assert deleted.status_code == 200, deleted.text + _assert_billed(_await_spend_row(created), _rates(rig.creator, model)) + _await_outcome(created, "billed") + + +def test_a_replica_booting_later_resumes_and_bills_unclaimed_interactions(rig: Rig) -> None: + with rig.creator.scenario() as scenario: + model: Final = rig.models.in_progress + key: Final = scenario.key() + created: Final = tuple(_create(rig.creator, model, key) for _ in range(3)) + for item in created: + _state(rig, item, _completed()) + rates: Final = _rates(rig.creator, model) + with owned_proxy_process( + rig.gateway, rig.directory, rig.environment(**_RESUMER_POLL), config=rig.config, workers=2 + ) as resumer: + pids: Final = _worker_pids(resumer.process.pid) + assert len(pids) == 2, pids + for item in created: + _assert_billed(_await_spend_row(item, seconds=90), rates) + assert _claimer_pid(_await_outcome(item, "billed")) in pids + for item in created: + assert len(_spend_rows(item)) == 1 + + +def test_deletes_on_the_creating_proxy_bill_each_interaction_once(rig: Rig) -> None: + with rig.settler.scenario() as scenario: + model: Final = rig.models.in_progress + key: Final = scenario.key() + created: Final = tuple(_create(rig.settler, model, key) for _ in range(8)) + for item in created: + _state(rig, item, _completed()) + assert _delete_concurrently(rig.settler, created, key) == (200,) * 8 + rates: Final = _rates(rig.settler, model) + for item in created: + _assert_billed(_await_spend_row(item), rates) + _await_outcome(item, "billed") + _await_key_spend(rig.settler, key, 8 * (_INPUT_TOKENS * rates[0] + _OUTPUT_TOKENS * rates[1])) + for item in created: + assert len(_spend_rows(item)) == 1 + + +def test_custom_deployment_pricing_bills_at_the_deployment_rate_on_another_replica(rig: Rig) -> None: + with rig.creator.scenario() as scenario: + model: Final = rig.models.custom_priced + key: Final = scenario.key() + created: Final = _create(rig.creator, model, key) + _state(rig, created, _completed()) + deleted: Final = _delete(rig.settler, created, key) + assert deleted.status_code == 200, deleted.text + _assert_billed(_await_spend_row(created), (_CUSTOM_INPUT_RATE, _CUSTOM_OUTPUT_RATE)) + _await_outcome(created, "billed") + + +def test_cancel_then_delete_on_another_replica_releases_without_a_spend_row(rig: Rig) -> None: + with rig.creator.scenario() as scenario: + model: Final = rig.models.in_progress + key: Final = scenario.key() + created: Final = _create(rig.creator, model, key) + _state(rig, created, InteractionState(status="in_progress")) + cancelled: Final = rig.settler.request("POST", f"/v1beta/interactions/{created}/cancel", {}, key=key) + assert cancelled.status_code == 200, cancelled.text + before_delete: Final = _settlement(created) + assert before_delete is not None and before_delete["claimed"] is False, before_delete + deleted: Final = _delete(rig.settler, created, key) + assert deleted.status_code == 200, deleted.text + _await_outcome(created, "released") + assert _spend_rows(created) == [] + + +def test_delete_fails_closed_when_the_settling_replica_cannot_fetch(rig: Rig) -> None: + with rig.creator.scenario() as scenario: + model: Final = rig.models.in_progress + key: Final = scenario.key() + created: Final = _create(rig.creator, model, key) + _state(rig, created, InteractionState(status="completed", usage=_USAGE, get_status=500)) + _drain(rig) + refused: Final = _delete(rig.settler, created, key) + assert refused.status_code >= 500, refused.text + assert "Scripted interaction fetch failure" in refused.text, refused.text + assert _calls(rig, created) == (("GET", _ENV_KEY),) + _assert_unclaimed(created) + assert _spend_rows(created) == [] + _state(rig, created, _completed()) + deleted: Final = _delete(rig.settler, created, key) + assert deleted.status_code == 200, deleted.text + _assert_billed(_await_spend_row(created), _rates(rig.creator, model)) + _await_outcome(created, "billed") + + +def test_delete_of_an_interaction_the_vendor_purged_sends_no_delete_and_keeps_the_row(rig: Rig) -> None: + with rig.creator.scenario() as scenario: + model: Final = rig.models.in_progress + key: Final = scenario.key() + created: Final = _create(rig.creator, model, key) + clear_interaction_state(rig.upstream.url, created) + _drain(rig) + deleted: Final = _delete(rig.settler, created, key) + assert deleted.status_code == 404, deleted.text + assert _calls(rig, created) == (("GET", _ENV_KEY),) + row: Final = _settlement(created) + assert row is not None and row["claimed"] is False, row + assert _spend_rows(created) == [] + + +def test_reading_an_interaction_never_bills_it(rig: Rig) -> None: + with rig.creator.scenario() as scenario: + model: Final = rig.models.in_progress + key: Final = scenario.key() + created: Final = _create(rig.creator, model, key) + _state(rig, created, InteractionState(status="in_progress")) + read_ids: Final = tuple(str(uuid.uuid4()) for _ in range(2)) + first: Final = rig.settler.request( + "GET", f"/v1beta/interactions/{created}", key=key, headers={"x-litellm-call-id": read_ids[0]} + ) + assert first.status_code == 200 and JSON_OBJECT.validate_json(first.content)["status"] == "in_progress", ( + first.text + ) + _state(rig, created, _completed()) + second: Final = rig.settler.request( + "GET", f"/v1beta/interactions/{created}", key=key, headers={"x-litellm-call-id": read_ids[1]} + ) + assert second.status_code == 200 and JSON_OBJECT.validate_json(second.content)["usage"] == _USAGE, second.text + assert _key_spend(rig.creator, key) == 0 + assert _spend_rows(created) == [] + deleted: Final = _delete(rig.settler, created, key) + assert deleted.status_code == 200, deleted.text + spend: Final = _assert_billed(_await_spend_row(created), _rates(rig.creator, model)) + _await_key_spend(rig.creator, key, spend) + for read_id in read_ids: + assert all(row["spend"] == 0 for row in _spend_rows(read_id)), _spend_rows(read_id) + + +@pytest.mark.parametrize( + "interaction_id", + [f"missing-{uuid.uuid4().hex}", "x" * 5000, "a.b:c", "%2F..%2Fup"], + ids=["unknown", "five-kilobytes", "punctuation", "encoded-traversal"], +) +def test_delete_of_an_odd_or_unknown_id_is_refused_and_the_proxy_keeps_serving(rig: Rig, interaction_id: str) -> None: + with rig.settler.scenario() as scenario: + key: Final = scenario.key() + deleted: Final = rig.settler.request("DELETE", f"/v1beta/interactions/{interaction_id}", key=key) + assert 400 <= deleted.status_code < 500, deleted.text + assert _readiness(rig.settler) == 200 + assert _key_spend(rig.settler, key) == 0 + + +def test_a_missing_settlement_table_leaves_in_process_billing_intact(rig: Rig) -> None: + write_rows(_OUTAGE_RENAME, ()) + try: + with rig.settler.scenario() as scenario: + model: Final = rig.models.in_progress + key: Final = scenario.key() + created: Final = _create(rig.settler, model, key) + _state(rig, created, _completed()) + spend: Final = _assert_billed(_await_spend_row(created), _rates(rig.settler, model)) + _await_key_spend(rig.settler, key, spend) + deleted: Final = _delete(rig.creator, created, key) + assert deleted.status_code == 200, deleted.text + assert len(_spend_rows(created)) == 1 + finally: + write_rows(_OUTAGE_RESTORE, ()) + + +def test_a_failed_create_registers_nothing(rig: Rig) -> None: + with rig.creator.scenario() as scenario: + model: Final = rig.models.failing_create + key: Final = scenario.key() + call_id: Final = str(uuid.uuid4()) + response: Final = rig.creator.request( + "POST", + "/v1beta/interactions", + {"model": model, "input": f"settle {uuid.uuid4().hex}", "background": True}, + key=key, + headers={"x-litellm-call-id": call_id}, + ) + assert response.status_code >= 500, response.text + assert _key_spend(rig.creator, key) == 0 + assert all(row["spend"] == 0 for row in _spend_rows(call_id)), _spend_rows(call_id) + assert _settlements_by_call_if_stored(call_id) == [] + + +def test_polling_disabled_replica_registers_nothing_and_never_bills(rig: Rig) -> None: + disabled: Final = rig.environment(BACKGROUND_INTERACTION_COST_POLLING_ENABLED="false") + with ( + owned_proxy(rig.gateway, rig.directory, disabled, config=rig.config, workers=1) as quiet, + quiet.scenario() as scenario, + ): + model: Final = rig.models.in_progress + key: Final = scenario.key() + created: Final = _create(quiet, model, key) + _state(rig, created, _completed()) + deleted: Final = _delete(rig.settler, created, key) + assert deleted.status_code == 200, deleted.text + assert _settlement_if_stored(created) is None + assert _key_spend(quiet, key) == 0 + assert _spend_rows(created) == [] + + +@pytest.mark.parametrize("background", [False, True], ids=["synchronous", "background"]) +def test_a_create_that_completes_at_once_is_billed_by_the_create_alone(rig: Rig, background: bool) -> None: + with rig.creator.scenario() as scenario: + model: Final = rig.models.completed_at_once + key: Final = scenario.key() + created: Final = _create(rig.creator, model, key, background=background) + spend: Final = _assert_billed(_await_spend_row(created), _rates(rig.creator, model)) + _await_key_spend(rig.creator, key, spend) + assert _settlement_if_stored(created) is None + _state(rig, created, _completed()) + deleted: Final = _delete(rig.settler, created, key) + assert deleted.status_code == 200, deleted.text + _await_key_spend(rig.creator, key, spend) + assert len(_spend_rows(created)) == 1 + + +def test_identical_creates_settle_as_separate_interactions(rig: Rig) -> None: + with rig.creator.scenario() as scenario: + model: Final = rig.models.in_progress + key: Final = scenario.key() + text: Final = f"settle {uuid.uuid4().hex}" + created: Final = tuple(_create(rig.creator, model, key, text=text) for _ in range(3)) + assert len({item for item in created}) == 3, created + for item in created: + _state(rig, item, _completed()) + _delete_ok(rig.settler, item, key) + rates: Final = _rates(rig.creator, model) + for item in created: + _assert_billed(_await_spend_row(item), rates) + _await_outcome(item, "billed") + _await_key_spend(rig.creator, key, 3 * (_INPUT_TOKENS * rates[0] + _OUTPUT_TOKENS * rates[1])) + + +def test_settlement_on_another_replica_releases_the_creators_budget_reservation(rig: Rig) -> None: + with rig.creator.scenario() as scenario: + model: Final = rig.models.in_progress + key: Final = scenario.key(max_budget=1.5 * _reservation_pin(rig.creator, model)) + first: Final = _create(rig.creator, model, key) + pinned: Final = rig.creator.request( + "POST", "/v1beta/interactions", {"model": model, "input": "settle pinned", "background": True}, key=key + ) + assert pinned.status_code == 400, pinned.text + _state(rig, first, _completed()) + deleted: Final = _delete(rig.settler, first, key) + assert deleted.status_code == 200, deleted.text + spend: Final = _assert_billed(_await_spend_row(first), _rates(rig.creator, model)) + _await_key_spend(rig.creator, key, spend) + released: Final = eventually( + lambda: ( + rig.creator.request( + "POST", + "/v1beta/interactions", + {"model": model, "input": "settle released", "background": True}, + key=key, + ).status_code + ), + lambda status: status == 200, + seconds=20, + return_last_on_timeout=True, + ) + assert released == 200 + + +def test_a_poll_that_never_sees_a_terminal_status_records_unsettled_and_releases(rig: Rig) -> None: + with rig.settler.scenario() as scenario: + model: Final = rig.models.in_progress + key: Final = scenario.key() + created: Final = _create(rig.settler, model, key) + _state(rig, created, InteractionState(status="in_progress")) + row: Final = _await_outcome(created, "unsettled", seconds=40) + assert row["create_context"] == {}, row + assert _spend_rows(created) == [] + assert _key_spend(rig.settler, key) == 0 + deleted: Final = _delete(rig.creator, created, key) + assert deleted.status_code == 200, deleted.text + assert _spend_rows(created) == [] + + +def test_an_upstream_outage_fails_deletes_closed_and_every_interaction_bills_once_after_recovery(rig: Rig) -> None: + with rig.creator.scenario() as scenario: + model: Final = rig.models.in_progress + key: Final = scenario.key() + created: Final = tuple(_create(rig.creator, model, key) for _ in range(16)) + rates: Final = _rates(rig.creator, model) + rig.upstream.stop() + try: + refused: Final = _delete_concurrently(rig.settler, created, key) + assert all(status >= 500 for status in refused), refused + for item in created: + _assert_unclaimed(item) + assert _readiness(rig.creator) == 200 and _readiness(rig.settler) == 200 + finally: + rig.upstream.start() + _register_scenarios(rig.upstream.url, rig.models) + for item in created: + _state(rig, item, _completed()) + assert _delete_concurrently(rig.settler, created, key) == (200,) * 16 + for item in created: + _assert_billed(_await_spend_row(item), rates) + _await_outcome(item, "billed") + _await_key_spend(rig.creator, key, 16 * (_INPUT_TOKENS * rates[0] + _OUTPUT_TOKENS * rates[1])) + for item in created: + assert len(_spend_rows(item)) == 1 + + +def test_killed_workers_leave_their_polls_to_the_respawned_workers(rig: Rig) -> None: + with ( + owned_proxy_process( + rig.gateway, rig.directory, rig.environment(**_RESUMER_POLL), config=rig.config, workers=2 + ) as resumer, + resumer.gateway.scenario() as scenario, + ): + model: Final = rig.models.in_progress + key: Final = scenario.key() + created: Final = tuple(_create(resumer.gateway, model, key) for _ in range(16)) + for item in created: + _state(rig, item, InteractionState(status="in_progress")) + rates: Final = _rates(resumer.gateway, model) + killed: Final = _worker_pids(resumer.process.pid) + assert len(killed) == 2, killed + victims: Final = tuple(psutil.Process(pid) for pid in killed) + for victim in victims: + victim.kill() + psutil.wait_procs(victims, timeout=15) + for item in created: + _state(rig, item, _completed()) + for item in created: + _assert_billed(_await_spend_row(item, seconds=150), rates) + assert _claimer_pid(_await_outcome(item, "billed")) not in killed + assert eventually(lambda: _readiness(resumer.gateway), lambda status: status == 200, seconds=60) == 200 + for item in created: + assert len(_spend_rows(item)) == 1 + + +def test_concurrent_deletes_on_a_slow_upstream_settle_exactly_once(rig: Rig) -> None: + with rig.creator.scenario() as scenario: + model: Final = rig.models.in_progress + key: Final = scenario.key() + created: Final = _create(rig.creator, model, key) + _state(rig, created, InteractionState(status="completed", usage=_USAGE, delay_seconds=1.5)) + statuses: Final = _delete_concurrently(rig.settler, (created, created), key) + assert sorted(statuses) == [200, 404], statuses + spend: Final = _assert_billed(_await_spend_row(created), _rates(rig.creator, model)) + _await_outcome(created, "billed") + _await_key_spend(rig.creator, key, spend) + assert len(_spend_rows(created)) == 1