diff --git a/app/routes/chat.py b/app/routes/chat.py index ab9200d6..15447c8c 100644 --- a/app/routes/chat.py +++ b/app/routes/chat.py @@ -11,7 +11,7 @@ import json import time import uuid -from collections.abc import AsyncGenerator, AsyncIterable, Callable +from collections.abc import AsyncGenerator, AsyncIterable, Awaitable, Callable import anyio import structlog @@ -31,6 +31,12 @@ from app.protocols.sse import AdapterError from app.quality_scores import resolve_model_metrics from app.schemas import ChatCompletionRequest +from packages.auth.spend import ( + MICROCENTS_PER_CENT, + budget_precheck, + charge_budget, + record_unsettled_spend, +) from packages.auth.types import KeyContext from packages.db.models.request_log import RequestLog from packages.litellm_adapter.catalog import CATALOG, CATALOG_BY_ID @@ -58,6 +64,303 @@ # retries hold the (already [DONE]) stream open. Tests shrink this. _LOG_COMMIT_BACKOFF_S: tuple[float, ...] = (0.1, 0.4) +# Crude character→token divisor used only to price a delivery the provider +# never measured (see `_estimate_usage`). +_CHARS_PER_TOKEN = 4 + + +def _text_chars(content) -> int: + """Character count of a message's text, across str and content-part lists.""" + if isinstance(content, str): + return len(content) + if isinstance(content, list): + return sum( + len(part["text"]) + for part in content + if isinstance(part, dict) and isinstance(part.get("text"), str) + ) + return 0 + + +_COUNTABLE_USAGE_KEYS = ( + "prompt_tokens", "completion_tokens", "input_tokens", "output_tokens", +) + + +def _countable_usage(usage) -> bool: + """Whether a usage dict carries at least one countable token total. + + A truthy `usage` is not proof the provider measured anything: some + upstreams report `{"total_tokens": 123}` and nothing else. Settling on that + normalizes every token read to zero, the row records a zero cost, and a + delivered completion bills the key nothing — the cap becomes steerable by + the upstream's field choice. Anything short of a countable key is treated + as no measurement at all. + """ + if not isinstance(usage, dict) or not usage: + return False + return any(usage.get(k) for k in _COUNTABLE_USAGE_KEYS) + + +def _estimate_usage(prompt_chars: int, completion_chars: int) -> dict: + """Token counts priced from characters, for a delivery the provider never measured.""" + return { + "prompt_tokens": max(1, prompt_chars // _CHARS_PER_TOKEN), + "completion_tokens": max(1, completion_chars // _CHARS_PER_TOKEN), + } + + +def _blocking_delivery_chars(response: dict) -> tuple[bool, int]: + """Delivered-content flag and completion characters of a blocking response. + + Returns `(has_content, completion_chars)` so the fail-closed gate and the + delivery estimate read the same evidence — a gate that counts content as + "delivered" while the estimator cannot see it would price the same delivery + two different ways. + """ + choices = response.get("choices") + if not isinstance(choices, list): + return False, 0 + has_content = False + chars = 0 + for choice in choices: + if not isinstance(choice, dict): + continue + candidates: list = [] + text = choice.get("text") + if isinstance(text, str): + candidates.append(text) + message = choice.get("message") + if isinstance(message, dict): + content = message.get("content") + if isinstance(content, str): + candidates.append(content) + elif isinstance(content, list): + for part in content: + if isinstance(part, str): + candidates.append(part) + elif isinstance(part, dict) and isinstance(part.get("text"), str): + candidates.append(part["text"]) + for value in candidates: + n = len(value) + chars += n + if value.strip(): + has_content = True + return has_content, chars + + +def _blocking_delivery_has_content(response: dict) -> bool: + """Whether a blocking ChatCompletion response carried delivered content. + + Unused on this rung — the gate here needs the character count too, so it + reads `_blocking_delivery_chars` directly. This exists for the + durable-recovery rung, whose blocking gate checks content alone. + """ + return _blocking_delivery_chars(response)[0] + + +# Endings a settlement can have. Handlers record which one happened; the +# charge is decided from it in exactly one place (see `_unmeasured_charge`), +# so no handler can price its own ending differently from its siblings. +_STREAM_IN_FLIGHT = "in_flight" +_STREAM_COMPLETED = "completed" +_STREAM_CLIENT_DISCONNECT = "client_disconnect" +_STREAM_UPSTREAM_ERROR = "upstream_error" +_STREAM_ADAPTER_ERROR = "adapter_error" + +# How an unmeasured-but-delivered request is charged when it ended normally. +_REMAINING = "remaining" +_ESTIMATE = "estimate" + + +def _cost_of_usage( + usage: dict, + *, + model_id: str | None, + fallback_model: str | None = None, +) -> int: + """Catalog cost of a token count, for a delivery the provider never priced.""" + return _compute_cost_microcents( + litellm_cost_usd=None, + model_id=model_id, + fallback_model=fallback_model, + input_tokens=usage.get("prompt_tokens", 0) or 0, + output_tokens=usage.get("completion_tokens", 0) or 0, + ) + + +def _unmeasured_charge( + *, + delivered: bool, + ending: str, + prompt_chars: int, + completion_chars: int, + policy: str, + cap: int, + spent: int, + model_id: str | None, + fallback_model: str | None = None, +) -> tuple[int, dict | None]: + """Charge for a request the provider never measured, from stream facts. + + The single decision point for every unmeasured ending, so the streaming and + blocking paths cannot drift apart. The rules, in order: + + - Nothing delivered: a provider or adapter failure that returned nothing + cost nothing measurable and settles at zero; a client that hung up + before the first byte still caused the prompt to go upstream, so it pays + the prompt estimate. + - Content delivered and the stream ended normally: charged the remaining + allowance (`policy=_REMAINING`). A stream that completed without ever + reporting usage is anomalous, and a client must not be able to opt out of + measurement and keep streaming for free. + - Content delivered and the ending was a hangup or a fault on our side or + the provider's: priced from the delivery estimate, which is the most we + can honestly know about a delivery nobody measured. + + A disconnect *after* a completed stream is not a separate ending: `ending` + is already `_STREAM_COMPLETED` by then, so where the client chose to hang up + cannot change the bill. + + Returns the charge together with the token counts it was derived from, so + the request-log row explains its own `cost_microcents` instead of recording + a cost with zero tokens behind it. The counts are `None` for the branches + that charge nothing and for the fail-closed raise, where no honest token + estimate exists. + """ + if ending == _STREAM_IN_FLIGHT: + # Settlement ran without an ending being recorded, which can only mean + # the consumer left before the provider finished. Priced as the hangup + # it is — never as a normal completion, which would let an unrecorded + # ending charge the whole remaining allowance. + ending = _STREAM_CLIENT_DISCONNECT + if not delivered: + if ending == _STREAM_CLIENT_DISCONNECT: + estimate = _estimate_usage(prompt_chars, 0) + return ( + max(1, _cost_of_usage( + estimate, + model_id=model_id, + fallback_model=fallback_model, + )), + estimate, + ) + return 0, None + if ending == _STREAM_COMPLETED and policy == _REMAINING: + return max(0, cap - spent), None + estimate = _estimate_usage(prompt_chars, completion_chars) + # Floored at one microcent: a short answer estimates to a fraction of a + # microcent, which truncates to zero, and a delivery that costs nothing is + # exactly what the cap must not allow. One microcent is 0.0001 cent — the + # smallest amount the counter can honestly represent. + return ( + max(1, _cost_of_usage( + estimate, model_id=model_id, fallback_model=fallback_model, + )), + estimate, + ) + + +def _blocking_delivery_has_content(response: dict) -> bool: + """Whether a blocking ChatCompletion response carried delivered content.""" + choices = response.get("choices") + if not isinstance(choices, list): + return False + for choice in choices: + if not isinstance(choice, dict): + continue + text = choice.get("text") + if isinstance(text, str) and text.strip(): + return True + message = choice.get("message") + if isinstance(message, dict): + content = message.get("content") + if isinstance(content, str) and content.strip(): + return True + if isinstance(content, list): + for part in content: + if isinstance(part, str) and part.strip(): + return True + if isinstance(part, dict): + t = part.get("text") + if isinstance(t, str) and t.strip(): + return True + return False +async def _give_up_settlement( + kc: KeyContext, + trace_id: str, + amount: int, + attempts: int, + error: BaseException, + persisted: Callable[[], Awaitable[bool]], +) -> None: + """Last resort for a settlement that is not durable and will not be retried. + + The log row dies with the charge (one transaction), so nothing anywhere + remembers this cost. Park it — one idempotent row keyed by this + settlement's `trace_id` — against the key's cap, rather than leaving the + cap open for whoever reads the warning (see `packages.auth.spend`). + + `persisted()` is asked first, because the last attempt has no retry left to + run the trace_id check: a commit that applied but whose ack was lost (or a + cancellation that landed after it) leaves the charge already in + `spent_microcents`, and parking on top of it would bill one delivery twice. + """ + logger.warning("request_log_commit_failed", error=str(error), attempts=attempts) + if getattr(kc, "_budget_cap", None) is None: + return + try: + durable = await persisted() + except asyncio.CancelledError: + # Torn down mid-probe, so the outcome is still unknown: park before + # propagating rather than letting this cost vanish with the coroutine. + await record_unsettled_spend( + trace_id=trace_id, api_key_id=str(kc.key_id), microcents=amount + ) + raise + if not durable: + await record_unsettled_spend( + trace_id=trace_id, api_key_id=str(kc.key_id), microcents=amount + ) + + +async def _trace_is_persisted(session: AsyncSession, trace_id: str) -> bool: + from sqlalchemy import select + + return ( + await session.scalar( + select(RequestLog.id).where(RequestLog.trace_id == trace_id) + ) + ) is not None + + +async def _settlement_is_durable(db: AsyncSession, trace_id: str) -> bool: + """Whether the request-log row for this trace_id is committed. + + The row and the budget charge are one transaction, so the row is proof the + spend is already counted. It reads on its own session because the failed + attempt's session is closed or rolled back by the time the give-up runs, and + answers False when it cannot read at all: during a real outage parking is the + only thing keeping the cap honest, so an unreadable DB must look + not-durable rather than durable. + """ + from packages.db import session as session_mod + + try: + if session_mod._session_factory is None: + # Test-only fallback (the app always installs a factory). + return await _trace_is_persisted(db, trace_id) + s = session_mod._session_factory() + try: + return await _trace_is_persisted(s, trace_id) + finally: + try: + await s.close() + except Exception as close_err: + logger.debug("request_log_session_close_failed", error=str(close_err)) + except Exception: + return False + def _chunk_to_dict(chunk) -> dict: """Normalize a litellm chunk (Pydantic model or dict) into a plain dict. @@ -209,8 +512,13 @@ async def _build_log_row( meta = response.get("_orca_meta", {}) or {} usage = response.get("usage", {}) or {} resolved = actual_resolved or response.get("model") or requested_model - input_t = usage.get("prompt_tokens", 0) or 0 - output_t = usage.get("completion_tokens", 0) or 0 + # Providers report usage under different keys: OpenAI-style + # prompt_tokens/completion_tokens, Anthropic-style input_tokens/ + # output_tokens (what /v1/messages forwards). Read both so a + # differently-keyed frame does not normalize to zero tokens — the + # token count is what the fail-closed budget gates key on. + input_t = usage.get("prompt_tokens", 0) or usage.get("input_tokens", 0) or 0 + output_t = usage.get("completion_tokens", 0) or usage.get("output_tokens", 0) or 0 return RequestLog( workspace_id=str(kc.workspace_id), api_key_id=str(kc.key_id), @@ -315,6 +623,28 @@ def _lookup_priced_model(model_id: str | None): return None +def _has_known_price( + *, + litellm_cost_usd: float | None, + model_id: str | None, + fallback_model: str | None = None, +) -> bool: + """True when a delivered completion's cost is measurable, even if it is 0. + + Mirrors `_compute_cost_microcents`' two tiers without computing: an + authoritative LiteLLM cost, or any catalog entry. A 0.0/0.0 entry is a + known-free model, not an unknown cost. Tokens with neither are + unpriceable — a custom upstream LiteLLM can't cost, or a model absent + from our catalog — and must not settle at 0 for a budgeted key, or the + lifetime cap would stand still while the upstream still bills us. + """ + if litellm_cost_usd is not None and litellm_cost_usd > 0: + return True + return ( + _lookup_priced_model(model_id) or _lookup_priced_model(fallback_model) + ) is not None + + @router.post("/chat/completions") async def chat_completions( body: ChatCompletionRequest, @@ -368,6 +698,21 @@ async def execute_chat( detail=f"Model '{body.model}' is not allowed for this API key", ) + async def _settle_budget(session, actual_microcents: int, *, commit: bool = True) -> None: + """Record `actual_microcents` of spend against the cap, if any. + + No-op when the key has no budget cap. When `commit` is False the UPDATE is + executed but not committed, so the caller commits it in the same + transaction as the request-log write — making the row and the charge one + atomic unit. Idempotency across retries comes from the row's trace_id + (a persisted trace_id proves the charge also landed), not from a + process-local flag. + """ + cap = getattr(kc, "_budget_cap", None) + if cap is None: + return + await charge_budget(session, str(kc.key_id), cap, actual_microcents, commit=commit) + client = await router_cache.get_router(db) raw_strategy = getattr(client, "strategy", None) strategy = raw_strategy if isinstance(raw_strategy, str) and raw_strategy else "balanced" @@ -482,6 +827,28 @@ async def execute_chat( resolved_model = candidates[0] body.model = candidates[0] # mutate for downstream completion call + # Budget enforcement: `budget_limit_cents` is a hard lifetime cap. The check + # runs only after the request has passed every pre-dispatch validation (model + # allowlist, provider deployability), so a request we reject before touching + # an upstream never consumes budget. The real cost is only known once the + # upstream response/stream completes, so we record it atomically in + # `_settle_budget` — the `UPDATE spent = spent + actual WHERE spent + actual + # <= cap` guard makes this safe under concurrency and never lets the counter + # exceed the cap (fail-closed, never over-recorded). + if kc.budget_limit_cents is not None: + cap = kc.budget_limit_cents * MICROCENTS_PER_CENT + # Fold-aware on this rung: the pre-check has to re-file and settle what + # an outage left parked, or a key whose settlement failed would keep + # dispatching as if it had never spent anything. + spent = await budget_precheck(db, str(kc.key_id), cap) + if spent >= cap: + raise HTTPException( + status_code=429, + detail="API key budget exhausted (lifetime cap reached).", + ) + kc._budget_cap = cap + kc._budget_spent = spent + started_perf = time.perf_counter() completion_kwargs = body.model_dump(exclude_none=True) @@ -553,6 +920,7 @@ async def execute_chat( log.cost_microcents = 0 db.add(log) try: + await _settle_budget(db, 0, commit=False) await db.commit() except Exception as commit_err: logger.warning("request_log_commit_failed", error=str(commit_err)) @@ -573,14 +941,14 @@ async def execute_chat( # mid-flight cascade is impossible — we have to surface the error and let # the client decide what to do. if body.stream: - # Auto-inject `stream_options.include_usage=True` if the client - # didn't set it. Without this, OpenAI/LiteLLM streaming responses - # omit the `usage` field entirely — chunks have no token counts, - # so our log row gets input=0, output=0 and the cost calculation - # rounds to zero. Almost no client knows to opt-in to this flag, - # which would silently zero out streaming spend in the dashboard. - # Honor an explicit `include_usage=False` from the client if they - # really want to disable it (e.g. wire-format compatibility tests). + # Auto-inject `stream_options.include_usage=True` if the client didn't + # set it, so streaming responses carry token counts and we bill the + # measured cost. A budgeted key needs it: without a usage frame the + # settlement falls back to pricing the delivery, which under-counts. + # A client that explicitly asked for `include_usage: False` still gets + # it — overriding that would hand every budgeted client an extra + # usage-only frame they asked not to receive. The cost of honouring it + # is bounded and visible: the stream settles from what it delivered. existing_so = completion_kwargs.get("stream_options") or {} if "include_usage" not in existing_so: completion_kwargs["stream_options"] = {**existing_so, "include_usage": True} @@ -606,6 +974,7 @@ async def _log_pre_stream_failure(status: int, err_type: str | None) -> None: ) db.add(log) try: + await _settle_budget(db, 0, commit=False) await db.commit() except Exception as commit_err: # Roll back so the request-scoped session is not left in a @@ -649,12 +1018,33 @@ async def sse() -> AsyncGenerator[str, None]: agg_provider = "unknown" agg_fallback = False agg_latency = 0 + # Characters of assistant text handed to the client — the only + # measure of what a stream delivered when the provider never + # reported usage (see `_estimate_usage`). + agg_output_chars = 0 # The first chunk's `model` field tells us what LiteLLM actually # served (could be a cascaded fallback, not the resolved primary). agg_model: str | None = None status_code = 200 error_type: str | None = None log_written = False + # How the stream ended. Handlers record the ending and nothing + # else — the charge is decided from these facts in one place + # (`_unmeasured_charge`), so a handler cannot price its own ending + # differently from a sibling. `usage_seen` is deliberately gone: a + # truthy usage dict is not a measurement (see `_countable_usage`), + # and where the client chose to hang up must not be able to change + # the bill. + stream_ending = _STREAM_IN_FLIGHT + # Whether this request asked the provider for a usage frame. We ask + # unless the client explicitly opted out. The distinction decides + # what an unmeasured completion costs: a provider that ignored our + # request is failing to measure something we needed (fail closed), + # while a client that declined usage told us up front, and is then + # charged for what it actually received. + usage_requested = bool( + (completion_kwargs.get("stream_options") or {}).get("include_usage") + ) async def _finalize() -> None: """Write the request log row exactly once. @@ -752,27 +1142,118 @@ async def _already_persisted(s) -> bool: select(RequestLog.id).where(RequestLog.trace_id == row_values["trace_id"]) )) is not None + async def _durable() -> bool: + return await _settlement_is_durable(db, row_values["trace_id"]) + + def _settlement_amount() -> int: + """Budget charge for this request, in microcents. + + A usage frame carrying countable token totals is the + billing signal: the cost is known, and charging more would + over-bill a quantity the row already accounts for. A frame + without those keys is not a measurement at all. + + A frame that carries tokens but no price — a custom + upstream LiteLLM cannot cost, or a model absent from our + catalog — is unknown too: charging the recorded 0 would + let the cap stand still while the upstream still bills us. + That arm is gated on the stream having completed normally, + because an error ending (disconnect, upstream or adapter + fault) is already priced from its delivery and keeps that + estimate. A catalog-listed free model, or an empty + delivery, is known-zero and settles at the 0 the row + records. + + When nothing was measured, the charge comes from + `_unmeasured_charge` and the stream's recorded ending — + never from which exception handler happened to run. The + rules that matters here: + + - we asked for usage and the provider sent none: the + completed delivery is charged the remaining allowance + (fail-closed), because that is an upstream failing to + report what the cap needs; + - the client declined usage frames up front: priced from + the delivery, since we were told not to ask; + - a hangup or a fault on either side: priced from what + reached the client; + - nothing delivered: charged nothing, unless a client hung + up, in which case the prompt it caused was still billed. + + The amount is written back into the row, not just into the + counter: a charge only the counter saw would leave a key + exhausted by an amount nothing in its own request history + accounts for. + """ + actual = row_values.get("cost_microcents") or 0 + if not _countable_usage(agg_usage): + charge, estimate = _unmeasured_charge( + delivered=agg_output_chars > 0, + ending=stream_ending, + prompt_chars=sum( + _text_chars(m.content) for m in body.messages + ), + completion_chars=agg_output_chars, + policy=_REMAINING if usage_requested else _ESTIMATE, + cap=getattr(kc, "_budget_cap", 0) or 0, + spent=getattr(kc, "_budget_spent", 0) or 0, + model_id=row_values.get("model_resolved"), + fallback_model=row_values.get("model_requested"), + ) + actual = max(actual, charge) + row_values["cost_microcents"] = actual + if estimate is not None: + # The row explains its own charge: an estimated + # delivery records the estimated tokens, so a + # non-zero cost never sits on a zero-token row. + row_values["input_tokens"] = estimate["prompt_tokens"] + row_values["output_tokens"] = estimate["completion_tokens"] + elif ( + stream_ending == _STREAM_COMPLETED + and not actual + and ( + row_values.get("input_tokens") + or row_values.get("output_tokens") + ) + and not _has_known_price( + litellm_cost_usd=(agg_usage or {}).get("cost_usd"), + model_id=row_values.get("model_resolved"), + fallback_model=row_values.get("model_requested"), + ) + ): + actual = max( + actual, + (getattr(kc, "_budget_cap", 0) or 0) + - (getattr(kc, "_budget_spent", 0) or 0), + ) + row_values["cost_microcents"] = actual + return actual + async def _commit_row(*, retry: bool) -> None: - """INSERT + COMMIT the row on a session of its own. - - Only a failing `commit()` propagates; a failure while - closing the session AFTER the commit returned is - swallowed — the row is already in. A retry is - idempotent: it first looks the trace_id up, so a COMMIT - that landed but whose ack was lost on the wire - (PostgreSQL, connection dropped mid-ack) is not - inserted a second time — and the shared primary key - would reject a duplicate anyway. + """Persist the request-log row and charge the budget in ONE commit. + + The INSERT and the budget charge share a single transaction. If it + commits, both are durable; if it fails, both roll back and the + retry re-runs both. Because the charge lands in the same commit as + the row, a persisted trace_id proves the charge also landed — so a + retry returns without re-charging. The charge is therefore applied + exactly once per request: never doubled (on a commit-ack-loss + retry) and never dropped. """ + # Resolved before the row is built: a fail-closed charge + # raises `row_values["cost_microcents"]`, and the object + # inserted must carry the amount the counter will move by. + settlement = _settlement_amount() log = RequestLog(**row_values) if session_mod._session_factory is None: # Test-only fallback (the app always installs a # factory): the request-scoped session has to be # rolled back before a retry can reuse it. - if retry and await _already_persisted(db): + if retry and (await _already_persisted(db)): return db.add(log) try: + await _settle_budget(db, settlement, commit=False) await db.commit() except Exception: try: @@ -783,9 +1264,10 @@ async def _commit_row(*, retry: bool) -> None: return s = session_mod._session_factory() try: - if retry and await _already_persisted(s): + if retry and (await _already_persisted(s)): return s.add(log) + await _settle_budget(s, settlement, commit=False) await s.commit() finally: try: @@ -825,31 +1307,42 @@ async def _commit_row(*, retry: bool) -> None: # Cancelled during the backoff: nothing is # in flight, the row is given up on — say # so, then propagate like the arm below. - logger.warning( - "request_log_commit_failed", - error=str(commit_err), attempts=attempt, + await _give_up_settlement( + kc, row_values["trace_id"], _settlement_amount(), + attempt, commit_err, _durable, ) raise continue - logger.warning( - "request_log_commit_failed", - error=str(commit_err), attempts=attempt, + await _give_up_settlement( + kc, row_values["trace_id"], _settlement_amount(), + attempt, commit_err, _durable, ) except BaseException: # CancelledError aimed at us, not at the commit — # wait the in-flight attempt out so a row about to # land isn't dropped, then let the cancellation - # propagate (no further attempts: we are being torn - # down). + # propagate. try: await commit_task + return except Exception as commit_err: - logger.warning( - "request_log_commit_failed", - error=str(commit_err), attempts=attempt, + await _give_up_settlement( + kc, row_values["trace_id"], _settlement_amount(), + attempt, commit_err, _durable, + ) + except BaseException as commit_cancel: + # The in-flight commit task is what was cancelled + # (or the loop is tearing it down): it can no longer + # write, so probe durability and park the + # obligation like every other arm here. This rung + # does not also re-run the write inside a shield — + # #161 does that, and the park is the stronger + # guarantee: the cost keeps counting even when the + # write can no longer land at all. + await _give_up_settlement( + kc, row_values["trace_id"], _settlement_amount(), + attempt, commit_cancel, _durable, ) - except BaseException: - pass raise last_d: dict = {} @@ -869,7 +1362,17 @@ async def _commit_row(*, retry: bool) -> None: if d.get("model"): agg_model = d["model"] last_d = d + for choice in d.get("choices") or []: + if isinstance(choice, dict): + delta = choice.get("delta") or {} + agg_output_chars += _text_chars(delta.get("content")) yield f"data: {json.dumps(d, separators=(',', ':'))}\n\n" + # The provider's stream is done: everything after this point is + # our own framing. Recording it HERE is what makes a client that + # hangs up at the trailing [DONE] frame pay the same charge as + # one that stays connected — where it disconnects is its choice, + # not a fact about the delivery. + stream_ending = _STREAM_COMPLETED # A trailing frame that already carries `usage` is the usage # frame, whether or not `choices` is empty. LiteLLM's # include_usage chunk uses @@ -929,6 +1432,17 @@ async def _commit_row(*, retry: bool) -> None: # awaits inside the shielded scope ignore outer cancellation # and run to completion. The scope exits normally and we # re-raise the original CancelledError below. + # + # Record the ending; the charge is decided from it in + # `_settlement_amount`, together with every other ending. A + # stream that already completed keeps `completed`: hanging up + # at the trailing [DONE] frame changes nothing about what was + # delivered, and pricing it as a disconnect would let a client + # buy the estimate for a full answer by leaving at the last + # byte. A mid-stream bail keeps its own pricing (including the + # prompt-only charge for a hangup before the first byte). + if stream_ending != _STREAM_COMPLETED: + stream_ending = _STREAM_CLIENT_DISCONNECT aclose = getattr(stream_obj, "aclose", None) with anyio.CancelScope(shield=True): if aclose is not None: @@ -943,16 +1457,16 @@ async def _commit_row(*, retry: bool) -> None: # Re-raise so asyncio/Starlette see proper cancel propagation. raise except AdapterError: - # The protocol adapter downstream of us failed and threw - # this in rather than closing us: our own fault, not the - # caller's. Without this branch the close would be - # indistinguishable from a disconnect and every adapter bug - # would be filed as 499/client_disconnect. error_type = "adapter_error" status_code = 500 logger.warning( "chat_completion_stream_adapter_error", served_model=agg_model, ) + # An adapter fault is our bug, not a choice the caller made, so + # it prices like the other faults: from what reached the client. + # Nothing delivered settles at the 0 the row already records. + if stream_ending != _STREAM_COMPLETED: + stream_ending = _STREAM_ADAPTER_ERROR aclose = getattr(stream_obj, "aclose", None) with anyio.CancelScope(shield=True): if aclose is not None: @@ -987,6 +1501,13 @@ async def _commit_row(*, retry: bool) -> None: "chat_completion_stream_error", error=str(exc), error_type=error_type, ) + # Record the ending BEFORE yielding the error frame and the + # sentinel. If the client disconnects during that yield, + # GeneratorExit unwinds straight through finally, so anything + # decided after the yield would be skipped and this fault would + # be priced as a normal completion. + if stream_ending != _STREAM_COMPLETED: + stream_ending = _STREAM_UPSTREAM_ERROR err_body = { "error": { "message": f"Upstream provider error: {exc}", @@ -1088,11 +1609,193 @@ async def _commit_row(*, retry: bool) -> None: # _build_log_row would otherwise default to via requested_model). actual_resolved=actual_resolved or resolved_model, ) - db.add(log) - try: - await db.commit() - except Exception as commit_err: - logger.warning("request_log_commit_failed", error=str(commit_err)) + # Persist the log row and the budget charge atomically (same transaction), + # retrying transient commit failures so a budgeted key is never under- + # charged when the DB is stressed — mirroring the streaming path. A + # persisted trace_id proves both landed, so a retry skips rather than + # double-charging. + from sqlalchemy import select + + settle_amount = log.cost_microcents + # A budgeted key whose successful response carries no countable usage + # (a provider that answered without token counts, or with a frame that + # has none) has an unknown cost. Unlike the streaming path — where a + # stream that completed without ever reporting usage is anomalous and + # priced fail-closed — a blocking response is here in full, so what it + # actually returned can be priced: the same delivery-estimate the fault + # endings use. Charging the whole remaining allowance here would bill a + # customer their lifetime budget for one ordinary request whose upstream + # omits a field, which is a far worse failure than the one the rule + # exists to prevent. + # + # Gated on a delivered completion dict AND delivered content: a request + # that failed before the upstream answered (response == {}, e.g. the + # re-raised HTTPException above, whose status_code never left 200) and + # a successful but empty completion both charge their recorded ~0 cost, + # mirroring the cache-hit and pre-stream-failure paths. A completion the + # provider priced through _orca_meta keeps that authoritative cost. + # Fail-closed mirror of the streaming path's cost-unknown rule: a + # budgeted key whose successful response carries no usage (a provider + # that answered without the field at all) has an unknown cost — charge + # the full remaining allowance so a delivered completion can never cost + # nothing, and record that amount on the row: a charge only the counter + # saw would leave a key exhausted by an amount nothing in its own + # request history accounts for. Usage with tokens but no price is the + # same unknown (a custom upstream LiteLLM can't cost, or a model + # absent from our catalog); a catalog-listed free model is known-zero + # and keeps its 0. As in the streaming arm, the unpriceable branch is + # gated on the usage having carried tokens: a usage of {0, 0} is an + # empty delivery, whose cost is known to be zero, and charging the whole + # remaining allowance for it would let an upstream that reports no tokens + # on a 200 permanently exhaust a capped key. Gated on having actually + # received a completion dict: a request that failed before the upstream + # answered (response == {}, e.g. the re-raised HTTPException above, whose + # status_code never left 200) charges its recorded ~0 cost instead — + # mirroring the cache-hit and pre-stream-failure paths. The missing- + # usage disjunct is gated on delivered content too: an empty 200 + # where usage is absent (or "{}") is a known-zero delivery, not the + # unknown-cost case fail-closed exists for. + if ( + getattr(kc, "_budget_cap", None) is not None + and status_code < 400 + and isinstance(response, dict) + and response + and ( + # Unmeasured: no countable usage at all. Priced from what the + # response returned, which is in full here. + ( + not _countable_usage(response.get("usage")) + and _blocking_delivery_has_content(response) + ) + or ( + # Measured but unpriceable: tokens with no price we can + # honour (a custom upstream cannot cost, or the model is + # absent from the catalog). Fail closed — the counter must + # not stand still while the upstream still bills us. + (log.input_tokens or log.output_tokens) + and not (log.cost_microcents or 0) + and not _has_known_price( + litellm_cost_usd=(response.get("usage") or {}).get("cost_usd") + or (response.get("_orca_meta") or {}).get("cost_usd"), + model_id=log.model_resolved, + fallback_model=log.model_requested, + ) + ) + ) + ): + if _countable_usage(response.get("usage")): + settle_amount = max( + log.cost_microcents or 0, + kc._budget_cap - (getattr(kc, "_budget_spent", 0) or 0), + ) + else: + delivered, completion_chars = _blocking_delivery_chars(response) + charge, estimate = _unmeasured_charge( + delivered=delivered, + ending=( + _STREAM_COMPLETED if delivered else _STREAM_UPSTREAM_ERROR + ), + prompt_chars=sum(_text_chars(m.content) for m in body.messages), + completion_chars=completion_chars, + # A delivered blocking completion is priced from what it + # returned; `_REMAINING` is the streaming path's rule for a + # stream that completed without ever reporting usage. + policy=_ESTIMATE, + cap=kc._budget_cap, + spent=getattr(kc, "_budget_spent", 0) or 0, + model_id=log.model_resolved, + fallback_model=log.model_requested, + ) + settle_amount = max(log.cost_microcents or 0, charge) + if estimate is not None: + # The row accounts for its own cost, as on the streaming path. + log.input_tokens = estimate["prompt_tokens"] + log.output_tokens = estimate["completion_tokens"] + log.cost_microcents = settle_amount + + # Values are snapshotted once (latency is measured in _build_log_row, + # before any commit attempt, so retry backoff never inflates it) and + # each attempt inserts a fresh ORM object carrying the same id/trace_id + # — mirroring the streaming path, so a retry works regardless of what + # the rollback left the old object as. + log_values = { + c.key: getattr(log, c.key) + for c in RequestLog.__table__.columns + if getattr(log, c.key) is not None + } + max_attempts = len(_LOG_COMMIT_BACKOFF_S) + 1 + + async def _durable() -> bool: + # The snapshot, not `log`: the give-up can run after a rollback has + # expired the session's state. + return await _settlement_is_durable(db, log_values["trace_id"]) + + for attempt in range(1, max_attempts + 1): + try: + if attempt > 1 and ( + await db.scalar( + select(RequestLog.id).where(RequestLog.trace_id == log.trace_id) + ) + ) is not None: + break # already durable (log + charge committed) + db.add(RequestLog(**log_values)) + await _settle_budget(db, settle_amount, commit=False) + await db.commit() + break + except Exception as commit_err: + try: + await db.rollback() + except BaseException: + # A cancellation here must not escape the handler: the + # give-ups below are the only thing that keeps this cost + # counting, and an exception raised from inside a handler is + # not caught by this `try`'s other arms. + pass + if attempt == max_attempts: + # Last attempt: probe whether the write landed and, if it + # did not, park the obligation. Nothing anywhere else + # remembers this cost — the row dies with the charge in one + # transaction — so a silent drop here would reopen the cap + # for exactly the key whose settlement failed. + await _give_up_settlement( + kc, log_values["trace_id"], settle_amount, + attempt, commit_err, _durable, + ) + break + logger.info( + "request_log_commit_retry", error=str(commit_err), attempt=attempt, + ) + try: + await asyncio.sleep(_LOG_COMMIT_BACKOFF_S[attempt - 1]) + except BaseException: + # Cancelled during the backoff: nothing is in flight and the + # row is given up on — say so, then propagate. An exception + # raised from a handler is not caught by this `try`'s other + # arms, so the give-up below does not run a second time for + # this settlement. + await _give_up_settlement( + kc, log_values["trace_id"], settle_amount, + attempt, commit_err, _durable, + ) + raise + except BaseException as cancel_err: + # Cancelled while the write was in flight. Unlike the streaming + # path there is no detached task left to land it: the transaction + # dies with this coroutine. Release it first — it holds the + # pending write's locks, which the durability read below needs + # past — then park unless it did commit, or this request's + # delivered cost would vanish with the coroutine. + try: + await db.rollback() + except BaseException: + # A second cancellation here must not skip the park: the + # give-up below propagates the original either way. + pass + await _give_up_settlement( + kc, log_values["trace_id"], settle_amount, + attempt, cancel_err, _durable, + ) + raise hosted_fallback = _meta_hosted_fallback(response) if isinstance(response, dict) and "_orca_meta" in response: diff --git a/packages/auth/spend.py b/packages/auth/spend.py index 27649968..9d93cf29 100644 --- a/packages/auth/spend.py +++ b/packages/auth/spend.py @@ -31,16 +31,310 @@ from __future__ import annotations -from sqlalchemy import select, update +import asyncio + +from sqlalchemy import delete, func, select, update +from sqlalchemy.exc import IntegrityError from sqlalchemy.ext.asyncio import AsyncSession from packages.db.models.api_key import ApiKey +from packages.db.models.budget_park import BudgetPark +from packages.db.models.request_log import RequestLog # Defined in `packages.db.units` so the boot repair that clamps a counter to the # same cap scales by the identical number; re-exported here because this module # is the documented home of the budget accounting primitives. from packages.db.units import MICROCENTS_PER_CENT as MICROCENTS_PER_CENT +# A settlement that gives up after every retry leaves a delivered cost with no +# record anywhere: the log row and the charge are one transaction, so both roll +# back and the counter never moves. The obligation is parked — one row per +# settlement, keyed by its `trace_id` — and keeps counting against the cap until +# a budget pre-check folds it into `spent_microcents`. +# +# The park is a database table, not process memory, because the deployment +# stops its machine whenever it goes idle: an in-memory obligation is lost on +# the next cold start, which reopens the cap for exactly the key the failure +# was about to protect. `_unsettled` below is the hold for when even that write +# cannot be made — it maps `(api_key_id, trace_id)` to the amount, and a later +# pre-check re-files it once a write goes through again. +_unsettled: dict[tuple[str, str], int] = {} + + +class _FoldConflict(Exception): + """A concurrent worker folded the same park first; the loser retries later.""" + + +async def _park_is_durable(trace_id: str) -> bool: + """Whether a park row for this `trace_id` is committed. + + A commit that applied but whose ack never came back raises exactly like a + failure, and holding a memory copy beside the durable row puts one + obligation in both ledgers — the double-bill the `trace_id` key exists to + absorb. It reads on a fresh session because the failed one is closed by the + time this runs, and answers False when it cannot read at all: with the + database truly unreachable the memory hold is all that keeps the cap honest. + """ + from packages.db import session as session_mod + + factory = session_mod._session_factory + if factory is None: + return False + try: + async with factory() as s: + return ( + await s.scalar( + select(BudgetPark.trace_id).where(BudgetPark.trace_id == trace_id) + ) + ) is not None + except Exception: + return False + + +async def _insert_park(*, trace_id: str, api_key_id: str, microcents: int) -> bool: + """Persist one parked obligation. Returns True when it is durable. + + Idempotent on `trace_id`: a commit that applied but whose ack was lost + retries into the same primary key instead of recording the obligation a + second time — and when the retry does not reach the server to be told that, + the probe below asks the database directly rather than reporting a write that + landed as one that did not. Returns False only when the obligation is + genuinely not durable, leaving the caller to hold the amount in memory. A + cancellation propagates with the memory copy still held. + """ + from packages.db import session as session_mod + + factory = session_mod._session_factory + if factory is None: + return False + try: + async with factory() as s: + s.add( + BudgetPark( + trace_id=trace_id, api_key_id=api_key_id, microcents=microcents, + ) + ) + await s.commit() + return True + except IntegrityError: + # The obligation is already parked — either by our own retried write + # after an ack loss, or by a concurrent give-up for the same trace. + # Either way there is exactly one durable copy, so report success. + return True + except asyncio.CancelledError: + raise + except Exception: + return await _park_is_durable(trace_id) + + +async def record_unsettled_spend( + *, trace_id: str, api_key_id: str, microcents: int +) -> None: + """Keep a settlement that gave up counting against the key's cap. + + The caller is unwinding a failure, so this never loses the amount: a park + that cannot be written durably is held in memory under its `trace_id` for + the next pre-check to re-file, and a cancellation holds it before + propagating rather than taking the obligation with it. + """ + if microcents <= 0 or not trace_id or not api_key_id: + return + key = (str(api_key_id), str(trace_id)) + try: + if await _insert_park( + trace_id=key[1], api_key_id=key[0], microcents=microcents + ): + return + except asyncio.CancelledError: + # `_insert_park` only raises the cancellation unwinding the caller — + # and the amount still has to be held before it propagates. + _unsettled[key] = microcents + raise + _unsettled[key] = microcents + + +async def pending_parked_spend(api_key_id: str) -> int | None: + """The outstanding park for a key: durable rows plus whatever is memory-only. + + ``None`` means the durable ledger could not be read, which is a different + answer from ``0``: the table is the only record of a park written by another + worker, or by this one before it stopped, so folding a failed read into the + total reopens the cap for exactly the key the park exists to hold shut. The + memory half still counts on the way to a real durable read failing, because + that half is what an outage is expected to lose and what recovers after it. + """ + from packages.db import session as session_mod + + key = str(api_key_id) + total = sum( + amount for (held_key, _trace), amount in _unsettled.items() if held_key == key + ) + factory = session_mod._session_factory + if factory is None: + return total + try: + async with factory() as s: + stored = ( + await s.execute( + select(func.sum(BudgetPark.microcents)).where( + BudgetPark.api_key_id == key + ) + ) + ).scalar() + except Exception: + return None + return total + int(stored or 0) + + +async def settle_parked_spend(api_key_id: str, cap_microcents: int) -> int: + """Fold a key's parked obligations into its recorded spend. Returns what moved. + + The park exists because a charge could not be recorded; leaving it parked + forever would mean a key at its cap is rejected by an amount that never + settles and never clears, so every pre-check tries to move it. Park rows are + read oldest-`created_at` first and the remaining allowance is applied to them + in that order; a park larger than the allowance bills what fits and is + rewritten to its remainder, rather than staying parked whole. That keeps the + invariant the fold exists to hold: either the queue is empty, or the counter + sits exactly on the cap. Without it a key can be refused at a lifetime spend + below its limit with a row that nothing will ever shrink, which is the state + this function is supposed to drain. The remainder is still a real debt — the + over-claim is the fail-closed policy — so it stays visible and keeps + `is_exhausted` blocking; + it is never written off, and it folds for free the moment the cap is raised. + + Two parks carrying the same `created_at` fall through to the `trace_id` + tiebreak, and `trace_id` is a uuid4 — so for obligations stamped within the + same clock tick the order is arbitrary rather than oldest-first. The + accounting does not depend on which row wins: the parked total is conserved + either way, the counter still reaches `min(cap, spent + debt)`, and only + which row is trimmed differs. + + The charge and the row writes share one transaction with compare-and-swap + guards on each: two workers folding the same park cannot double-bill it, + because the loser's UPDATE or DELETE matches nothing and its next request + folds what the winner left. + """ + from packages.db import session as session_mod + + key = str(api_key_id) + factory = session_mod._session_factory + if factory is None: + return 0 + for (held_key, trace_id), amount in list(_unsettled.items()): + if held_key != key: + continue + # A cancellation here propagates with the entry still held; a later + # pre-check re-files it, and the `trace_id` key keeps the retry from + # duplicating it. + if await _insert_park( + trace_id=trace_id, api_key_id=key, microcents=amount + ): + _unsettled.pop((held_key, trace_id), None) + move = 0 + settling: list[str] = [] + trim: tuple[str, int, int] | None = None + try: + async with factory() as s: + async with s.begin(): + spent = ( + await s.execute( + select(ApiKey.spent_microcents).where(ApiKey.id == key) + ) + ).scalar_one_or_none() + if spent is None: + return 0 + spent = int(spent) + rows = ( + await s.execute( + select(BudgetPark.trace_id, BudgetPark.microcents) + .where(BudgetPark.api_key_id == key) + # Oldest debt first. `created_at` is stamped + # Python-side with sub-second resolution, but two + # workers computing the same fold still have to agree on + # which row is the partial one, so `trace_id` breaks any + # tie rather than letting the order depend on a race. + .order_by(BudgetPark.created_at, BudgetPark.trace_id) + ) + ).all() + # A log row and its budget charge are committed atomically. + # If a commit acknowledgement was lost and the durability + # probe also failed, the matching park is only a fallback + # record; the request-log row proves the charge already landed. + logged = set( + ( + await s.scalars( + select(RequestLog.trace_id).where( + RequestLog.trace_id.in_( + [trace_id for trace_id, _amount in rows] + ) + ) + ) + ).all() + ) if rows else set() + already_charged = [trace_id for trace_id, _amount in rows if trace_id in logged] + if already_charged: + cleared = await s.execute( + delete(BudgetPark).where( + BudgetPark.trace_id.in_(already_charged) + ) + ) + if cleared.rowcount != len(already_charged): + raise _FoldConflict + rows = [row for row in rows if row[0] not in logged] + room = cap_microcents - spent + for trace_id, microcents in rows: + microcents = int(microcents) + if microcents <= room: + room -= microcents + move += microcents + settling.append(trace_id) + continue + if room > 0: + trim = (trace_id, microcents, microcents - room) + move += room + break + if move <= 0: + return 0 + charged = await s.execute( + update(ApiKey) + .where(ApiKey.id == key, ApiKey.spent_microcents == spent) + .values(spent_microcents=spent + move) + ) + if charged.rowcount != 1: + raise _FoldConflict + if settling: + cleared = await s.execute( + delete(BudgetPark).where(BudgetPark.trace_id.in_(settling)) + ) + if cleared.rowcount != len(settling): + raise _FoldConflict + if trim is not None: + trace_id, whole, remainder = trim + trimmed = await s.execute( + update(BudgetPark) + .where( + BudgetPark.trace_id == trace_id, + BudgetPark.microcents == whole, + ) + .values(microcents=remainder) + ) + if trimmed.rowcount != 1: + raise _FoldConflict + except _FoldConflict: + return 0 + # Past the commit, so everything in `settling` is billed and gone from the + # table. A memory hold for one of those rows is now worse than stale: the + # next pre-check re-files it as a brand-new park — the `trace_id` no longer + # collides, this transaction deleted the row — and the same delivery is + # charged twice against a cap that has no idea it moved. Only `settling`: + # a trimmed row is still owed, and the hold on it has to stay. + for _settled in settling: + _unsettled.pop((key, _settled), None) + for _already_charged in already_charged: + _unsettled.pop((key, _already_charged), None) + return move + async def read_spent(db: AsyncSession, api_key_id: str) -> int: """Return the key's currently-recorded lifetime spend in microcents.""" @@ -50,15 +344,68 @@ async def read_spent(db: AsyncSession, api_key_id: str) -> int: return int(spent or 0) -async def is_exhausted(db: AsyncSession, api_key_id: str, cap_microcents: int) -> bool: - """Fast pre-check: has the key already reached its lifetime cap? +async def budget_precheck(db: AsyncSession, api_key_id: str, cap_microcents: int) -> int: + """The key's spend as of this pre-check, in microcents. + + One read behind both halves of the caller's decision: whether to reject the + request, and what allowance is left for a cost that is not known yet. It + folds a parked obligation into the counter before answering, so a write + outage is neither a window of free requests nor a park that can never clear, + and the result adds whatever is still pending rather than reading only the + counter: when the fold could not commit, the obligation still has to block + dispatch. A park ledger that cannot be read at all is answered as the cap + rather than as no debt. ``cap_microcents`` is ``ApiKey.budget_limit_cents`` scaled by ``MICROCENTS_PER_CENT``, not the column itself — passing the raw cents value asks whether the key has spent a ten-thousandth of its budget. """ - spent = await read_spent(db, api_key_id) - return spent >= cap_microcents + key = str(api_key_id) + pending = await pending_parked_spend(key) + if pending is None: + # Unknown debt is answered as full debt, the way the durability probe + # that wrote the park reads an unreadable database as "not settled". + # The alternative is a key whose cap is held shut only by parked rows + # dispatching freely while the ledger is down. + return cap_microcents + # The park ledger is read before the counter, and the counter on a fresh + # snapshot: a fold both moves `spent_microcents` and empties the park + # queue, so a counter read first can pair a pre-fold spend with a zero + # pending — understating the spend by exactly the folded amount with no + # evidence left to trigger the re-read. Parks-first keeps the evidence: + # a fold racing these reads leaves `pending > 0`, which takes the + # re-read below. The rollback ends whatever pending work the request + # session carries; a cancellation landing on it must not escape — the + # pre-check has no retry of its own and the caller's loop owns the + # session, so swallow it and keep reading on the (stale) snapshot. + try: + await db.rollback() + except BaseException: + pass + spent = await read_spent(db, key) + if pending: + try: + await settle_parked_spend(key, cap_microcents) + except Exception: + # The fold runs on sessions of its own, so there is nothing to + # roll back here — and the re-read below still counts the park. + pass + # End db's read transaction so its SQLite/Postgres snapshot doesn't stay + # pinned to the pre-fold spend counter, then re-read on a fresh snapshot. + try: + await db.rollback() + except BaseException: + pass + spent = await read_spent(db, key) + pending = await pending_parked_spend(key) + if pending is None: + return cap_microcents + return spent + pending + + +async def is_exhausted(db: AsyncSession, api_key_id: str, cap_microcents: int) -> bool: + """Fast pre-check: has the key already reached its lifetime cap?""" + return await budget_precheck(db, api_key_id, cap_microcents) >= cap_microcents async def charge_budget( diff --git a/packages/db/migrate.py b/packages/db/migrate.py index 34c5e041..b5831b55 100644 --- a/packages/db/migrate.py +++ b/packages/db/migrate.py @@ -14,26 +14,29 @@ from __future__ import annotations +from collections.abc import Callable +from typing import Any + from sqlalchemy import BigInteger, inspect, text from sqlalchemy.exc import DBAPIError +from packages.db.models.budget_park import BudgetPark from packages.db.units import MICROCENTS_PER_CENT def _already_applied(err: DBAPIError) -> bool: """Whether a DDL failure means someone else applied the change first.""" msg = str(err).lower() - return ( - "already exists" in msg - or "duplicate column" in msg - # Concurrent CREATE INDEX on Postgres can lose the race at the catalog - # insert rather than the IF NOT EXISTS probe, surfacing as a verror on - # pg_class's unique index instead of the usual "already exists". - or "pg_class_relname_nsp_index" in msg - ) + if "already exists" in msg or "duplicate column" in msg: + return True + # Two concurrent CREATE TABLE are serialised by the catalog rather than by + # the wording of a complaint, so the loser gets a unique violation on + # pg_class/pg_type (`*_relname_nsp_index`, `*_typname_nsp_index`) instead of + # "already exists". Same meaning: the object is there now. + return "duplicate key value" in msg and "nsp_index" in msg -async def _apply_ddl(conn, statement: str) -> None: +async def _apply_ddl(conn, statement: str | Callable[..., Any]) -> None: """Run one startup DDL statement, tolerating a boot that raced us to it. Every worker runs this in its lifespan, so the first boot after an upgrade @@ -43,10 +46,16 @@ async def _apply_ddl(conn, statement: str) -> None: from here, not a reason to keep the worker from booting. The failure is caught inside a SAVEPOINT because on Postgres an error would otherwise abort the whole transaction and take the rest of the startup with it. + + `statement` is raw SQL, or a callable handed to `run_sync` for DDL that only + the dialect's own generator can emit (a `Table.create`). """ try: async with conn.begin_nested(): - await conn.execute(text(statement)) + if callable(statement): + await conn.run_sync(statement) + else: + await conn.execute(text(statement)) except DBAPIError as err: if not _already_applied(err): raise @@ -59,18 +68,37 @@ async def ensure_budget_columns(engine) -> None: against historical request logs on every boot so the `ALTER` is never mistaken for proof that the seed ran. Also widens `budget_limit_cents` to BIGINT on Postgres (the column is scaled into microcents for every comparison against - spend, and an int4 ceiling is about 214,748 dollars of lifetime budget) and + spend, and an int4 ceiling is about 214,748 dollars of lifetime budget), creates the `ix_requests_log_api_key_spend` index that create_all only builds - on fresh databases. Each step costs nothing on a database that needs none of - it — there the repair is a single indexed UPDATE that matches no row. + on fresh databases, and creates `budget_parks` for deployments that predate + the durable-recovery release — a lost settlement needs somewhere every + worker, and every reboot, can see it. Each step costs nothing on a database + that needs none of it — there the repair is a single indexed UPDATE that + matches no row. """ async with engine.begin() as conn: + tables = set( + await conn.run_sync(lambda sync: inspect(sync).get_table_names()) + ) cols = { c["name"]: c["type"] for c in await conn.run_sync(lambda sync: inspect(sync).get_columns("api_keys")) } is_postgres = engine.dialect.name == "postgresql" + if BudgetPark.__tablename__ not in tables: + # `create_all` covers fresh databases; this covers upgrades whose + # schema predates the table. `checkfirst` re-reads the catalog, and + # that read cannot see another boot's uncommitted CREATE — so two + # workers both get here and one loses anyway. It goes through + # `_apply_ddl` for the same reason the ALTER does: losing that race + # has to count as having won, and the savepoint is what keeps the + # error from aborting the transaction the rest of startup runs in. + await _apply_ddl( + conn, + lambda sync: BudgetPark.__table__.create(sync, checkfirst=True), + ) + # The model declares ix_requests_log_api_key_spend (api_key_id, # is_deleted); create_all only builds it on fresh databases, so an # upgraded deployment would drift. Built before the seed below, which is diff --git a/packages/db/models/__init__.py b/packages/db/models/__init__.py index ca728ffa..f8fc00a3 100644 --- a/packages/db/models/__init__.py +++ b/packages/db/models/__init__.py @@ -2,6 +2,7 @@ from packages.db.models.api_key import ApiKey from packages.db.models.base import Base, SoftDeleteMixin, TimestampMixin, UUIDMixin +from packages.db.models.budget_park import BudgetPark from packages.db.models.provider_key import ProviderKey from packages.db.models.quality_score_override import QualityScoreOverride from packages.db.models.quality_score_snapshot import QualityScoreSnapshot @@ -15,6 +16,7 @@ "TimestampMixin", "UUIDMixin", "ApiKey", + "BudgetPark", "ProviderKey", "QualityScoreOverride", "QualityScoreSnapshot", diff --git a/packages/db/models/budget_park.py b/packages/db/models/budget_park.py new file mode 100644 index 00000000..4a55f26f --- /dev/null +++ b/packages/db/models/budget_park.py @@ -0,0 +1,51 @@ +"""Unsettled budget obligations — delivered spend the ledger never recorded. + +A settlement that gives up after every retry leaves a delivered cost with no +record anywhere: the log row and the charge are one transaction, so both roll +back and `spent_microcents` never moves. The obligation is parked here, one row +per settlement, so it outlives the process that lost it and is visible to every +worker behind the same database. + +Rows are keyed by the settlement's `trace_id`, which makes every park write +idempotent: a commit that applied but whose ack was lost retries into the same +primary key instead of recording the obligation twice, and a write that fails +outright is checked for having landed anyway before the process falls back to +holding it in memory. A fold that bills a row either deletes it or shrinks it to +what the cap could not absorb, in the same transaction that moves +`spent_microcents`, and it drops the writer's memory hold for any row it fully +billed. + +What that leaves is a double charge needing three faults at once — the ack lost, +the durability probe failing alongside it, and another worker folding the row +before this one retries — at which point the extra charge lands on a key that had +already breached its cap. Closing that last window means a fold leaving a tombstone +behind instead of deleting, so a re-file always collides with something; that is a +second state the queue has to drain, and it is not worth its weight here. +""" + +from datetime import datetime, timezone + +from sqlalchemy import BigInteger, DateTime, String +from sqlalchemy.orm import Mapped, mapped_column + +from packages.db.models.base import Base, TimestampMixin, UUIDMixin + + +class BudgetPark(Base, UUIDMixin, TimestampMixin): + __tablename__ = "budget_parks" + + trace_id: Mapped[str] = mapped_column(String(36), nullable=False, unique=True) + api_key_id: Mapped[str] = mapped_column(String(36), nullable=False, index=True) + microcents: Mapped[int] = mapped_column(BigInteger, nullable=False) + # Overrides TimestampMixin's column, whose `server_default=func.now()` is + # CURRENT_TIMESTAMP: one second wide on SQLite, where a recovered outage + # re-files a whole batch of obligations in a single pre-check and every row + # in it ties. The fold bills oldest debt first, and two workers have to + # agree on which row the partial one is, so the stamp needs enough + # resolution to settle that on its own instead of falling through to the + # `trace_id` tiebreak — which is a uuid4, and so picks at random. + created_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), + default=lambda: datetime.now(timezone.utc), + nullable=False, + ) diff --git a/tests/integration/test_adapter_failure_attribution.py b/tests/integration/test_adapter_failure_attribution.py index 40f2966a..a4738f39 100644 --- a/tests/integration/test_adapter_failure_attribution.py +++ b/tests/integration/test_adapter_failure_attribution.py @@ -27,6 +27,10 @@ _messages_payload, native_client, ) +from tests.integration.test_budget_enforcement import ( # noqa: F401 + _get_spent, + _make_budgeted_key, +) _GEMINI_PAYLOAD = {"contents": [{"role": "user", "parts": [{"text": "hi"}]}]} @@ -129,3 +133,90 @@ async def _stream_router(**kwargs): assert slow.closed is True assert await _log_rows() == [(499, "client_disconnect", True)] + + +# ── Budget settlement on an adapter fault ────────────────────────────── +# The adapter IS the response body, so its failure unwinds through the +# engine's SSE generator. The engine must settle that request on what it +# delivered — leaving the settlement "unknown" would charge a budgeted key +# its entire remaining lifetime budget for a bug of ours. + +_BAD_CHUNK = { + "id": "chatcmpl-2", "object": "chat.completion.chunk", "model": "gpt-4o-mini", + "created": int(time.time()), "choices": "boom", +} + + +def _content_chunk(text: str) -> dict: + return { + "id": "chatcmpl-1", "object": "chat.completion.chunk", "model": "gpt-4o-mini", + "created": int(time.time()), + "choices": [{"index": 0, "delta": {"content": text}, "finish_reason": None}], + } + + +def _budget_stream_router(fake, chunks) -> None: + async def _acompletion(**kwargs): + assert kwargs.get("stream") + + async def _gen(): + for c in chunks: + yield c + + return _gen() + + fake.acompletion = AsyncMock(side_effect=_acompletion) + + +async def _log_row_for(api_key_id: str): + from packages.db import session as session_mod + from packages.db.models.request_log import RequestLog + + async with session_mod._session_factory() as s: + return ( + await s.execute(select(RequestLog).where(RequestLog.api_key_id == api_key_id)) + ).scalars().one() + + +async def test_budgeted_adapter_fault_after_content_charges_delivery_not_remaining( + native_client, +): + client, fake, _root = native_client + from packages.db import session as session_mod + + factory = session_mod._session_factory + # 100 cents = 1_000_000 microcents of lifetime budget. + key, key_id = await _make_budgeted_key(factory, budget_limit_cents=100) + delivered = "the quick brown fox jumps over the lazy dog. " * 200 + _budget_stream_router(fake, [_content_chunk(delivered), _BAD_CHUNK]) + + r = await client.post( + "/v1/messages", json=_messages_payload(stream=True), headers={"x-api-key": key}, + ) + assert r.status_code == 200 + await asyncio.sleep(0.1) + + assert await _log_rows() == [(500, "adapter_error", True)] + row = await _log_row_for(key_id) + spent = await _get_spent(factory, key_id) + assert spent == row.cost_microcents # charged == accounted + assert 0 < spent < 1_000_000 # priced, not the whole remaining budget + + +async def test_budgeted_adapter_fault_before_content_charges_nothing(native_client): + """A fault with nothing delivered cost nothing, so it must bill nothing.""" + client, fake, _root = native_client + from packages.db import session as session_mod + + factory = session_mod._session_factory + key, key_id = await _make_budgeted_key(factory, budget_limit_cents=100) + _budget_stream_router(fake, [_BAD_CHUNK]) + + r = await client.post( + "/v1/messages", json=_messages_payload(stream=True), headers={"x-api-key": key}, + ) + assert r.status_code == 200 + await asyncio.sleep(0.1) + + assert await _log_rows() == [(500, "adapter_error", True)] + assert await _get_spent(factory, key_id) == 0 diff --git a/tests/integration/test_budget_enforcement.py b/tests/integration/test_budget_enforcement.py new file mode 100644 index 00000000..178ea24f --- /dev/null +++ b/tests/integration/test_budget_enforcement.py @@ -0,0 +1,1869 @@ +"""Budget enforcement on /v1/chat/completions. + +`budget_limit_cents` was loaded into KeyContext but never enforced anywhere — +a leaked key meant unbounded spend. These tests pin the new behavior: an +exhausted key gets 429 before any routing / cache / upstream work and +unbudgeted keys are unaffected. Provisioning of budgeted/allowlisted keys +is covered in the keys-authz PR. +""" + +from __future__ import annotations + +import asyncio +import time +from unittest.mock import AsyncMock + +import pytest +from sqlalchemy.ext.asyncio import AsyncSession + + +@pytest.fixture +async def budget_env(tmp_sqlite_url, monkeypatch): + """Full app + seeded root key, with the router client mocked out. + + Yields (make_client, fake_client, session_factory, root_key). + """ + monkeypatch.setenv("DATABASE_URL", tmp_sqlite_url) + monkeypatch.setenv("OPENAI_API_KEY", "sk-test-openai") + + from app import config as cfg + cfg.get_settings.cache_clear() + + from packages.db.engine import build_engine + from packages.db.models.base import Base + + engine = build_engine(tmp_sqlite_url) + async with engine.begin() as conn: + await conn.run_sync(Base.metadata.create_all) + + from sqlalchemy.ext.asyncio import async_sessionmaker + + from packages.db import session as session_mod + factory = async_sessionmaker(engine, expire_on_commit=False) + session_mod._session_factory = factory + + from app.seed import seed_initial_state + async with factory() as s: + seed = await seed_initial_state(s) + + fake_client = AsyncMock() + fake_client.acompletion = AsyncMock( + return_value={ + "id": "chatcmpl-budget-test", + "model": "gpt-4o-mini", + "object": "chat.completion", + "created": int(time.time()), + "choices": [{ + "index": 0, + "message": {"role": "assistant", "content": "Hello!"}, + "finish_reason": "stop", + }], + "usage": {"prompt_tokens": 5, "completion_tokens": 2, "total_tokens": 7}, + "_orca_meta": { + "provider": "openai", + "litellm_model": "openai/gpt-4o-mini", + "latency_ms": 42, + }, + } + ) + + from app import router_cache + router_cache.invalidate_router() + + async def _fake_get_router(_session): + return fake_client + + monkeypatch.setattr(router_cache, "get_router", _fake_get_router) + + from httpx import ASGITransport, AsyncClient + + from app.main import create_app + app = create_app() + + async def make_client(api_key: str): + return AsyncClient( + transport=ASGITransport(app=app), + base_url="http://t", + headers={"Authorization": f"Bearer {api_key}"}, + ) + + yield make_client, fake_client, factory, seed.api_key + + await engine.dispose() + session_mod._session_factory = None + + +async def _make_budgeted_key( + factory, *, budget_limit_cents: int | None +) -> tuple[str, str]: + """Insert a budgeted child key; return (plaintext_key, key_id).""" + from packages.auth.hashing import generate_api_key + from packages.db.models.api_key import ApiKey + + full_key, key_hash, key_prefix = generate_api_key() + async with factory() as s: + row = ApiKey( + workspace_id="default", + name="budgeted", + key_hash=key_hash, + key_prefix=key_prefix, + budget_limit_cents=budget_limit_cents, + ) + s.add(row) + await s.commit() + await s.refresh(row) + return full_key, row.id + + +async def _add_billable_spend(factory, key_id: str, microcents: int) -> None: + from packages.db.models.api_key import ApiKey + from packages.db.models.request_log import RequestLog + + async with factory() as s: + s.add(RequestLog( + workspace_id="default", + api_key_id=key_id, + trace_id="budget-test-trace", + model_requested="gpt-4o-mini", + model_resolved="gpt-4o-mini", + provider="openai", + routing_strategy="balanced", + input_tokens=5, + output_tokens=2, + cost_microcents=microcents, + latency_ms=10, + status_code=200, + )) + # The budget counter lives on the key, not the request-log rows, so + # pre-load it directly to simulate prior spend. + await s.execute( + ApiKey.__table__.update() + .where(ApiKey.id == key_id) + .values(spent_microcents=ApiKey.spent_microcents + microcents) + ) + await s.commit() + + +async def test_exhausted_budget_returns_429_without_upstream_call(budget_env): + make_client, fake, factory, _root = budget_env + key, key_id = await _make_budgeted_key(factory, budget_limit_cents=1) + # Pre-load spend past the 1-cent cap (10_000 microcents). + await _add_billable_spend(factory, key_id, microcents=20_000) + + async with await make_client(key) as c: + r = await c.post( + "/v1/chat/completions", + json={"model": "gpt-4o-mini", + "messages": [{"role": "user", "content": "hi"}]}, + ) + + assert r.status_code == 429, r.text + assert r.json()["error"]["type"] == "rate_limit_error" + fake.acompletion.assert_not_awaited() + + +async def test_blocked_request_writes_no_log_row(budget_env): + make_client, _fake, factory, _root = budget_env + key, key_id = await _make_budgeted_key(factory, budget_limit_cents=1) + await _add_billable_spend(factory, key_id, microcents=99_999) + + async with await make_client(key) as c: + await c.post( + "/v1/chat/completions", + json={"model": "gpt-4o-mini", + "messages": [{"role": "user", "content": "hi"}]}, + ) + + from sqlalchemy import func, select + + from packages.db.models.request_log import RequestLog + + async with factory() as s: + count = ( + await s.execute( + select(func.count()).select_from(RequestLog).where( + RequestLog.api_key_id == key_id + ) + ) + ).scalar_one() + assert count == 1 # only the pre-loaded history row + + +async def test_under_budget_key_serves_normally(budget_env): + make_client, fake, factory, _root = budget_env + key, _key_id = await _make_budgeted_key(factory, budget_limit_cents=100) + + async with await make_client(key) as c: + r = await c.post( + "/v1/chat/completions", + json={"model": "gpt-4o-mini", + "messages": [{"role": "user", "content": "hi"}]}, + ) + + assert r.status_code == 200, r.text + fake.acompletion.assert_awaited_once() + + +async def test_unbudgeted_root_key_unaffected(budget_env): + make_client, fake, _factory, root = budget_env + + async with await make_client(root) as c: + r = await c.post( + "/v1/chat/completions", + json={"model": "gpt-4o-mini", + "messages": [{"role": "user", "content": "hi"}]}, + ) + + assert r.status_code == 200, r.text + fake.acompletion.assert_awaited_once() + + +async def _budgeted_stream( + budget_env, *, chunks, budget_limit_cents=10, stream_options=None, +): + """Drive a streaming request for a budgeted key and return its final spend.""" + make_client, fake, factory, _root = budget_env + key, key_id = await _make_budgeted_key(factory, budget_limit_cents=budget_limit_cents) + + async def _stream(): + for ch in chunks: + yield ch + + fake.acompletion = AsyncMock(return_value=_stream()) + + async with await make_client(key) as c: + async with c.stream( + "POST", + "/v1/chat/completions", + json={ + "model": "gpt-4o-mini", + "stream": True, + "messages": [{"role": "user", "content": "hi"}], + "stream_options": ( + {"include_usage": False} if stream_options is None + else stream_options + ), + }, + ) as r: + async for _ in r.aiter_lines(): + pass + + from sqlalchemy import select + + from packages.db.models.api_key import ApiKey + + async with factory() as s: + return ( + await s.execute(select(ApiKey.spent_microcents).where(ApiKey.id == key_id)) + ).scalar_one(), fake.acompletion.call_args + + +async def test_budgeted_stream_usage_declined_by_client_prices_the_delivery( + budget_env, +): + # The client explicitly declined usage frames. The engine does not override + # that (it would hand the client a usage-only frame it asked not to + # receive), and the completed stream is then priced from what it delivered + # rather than fail-closed: the client told us up front that it does not + # want usage frames, so charging it a lifetime budget for a normal answer + # would punish an explicit, documented preference. + spent, call_args = await _budgeted_stream( + budget_env, + chunks=[ + {"choices": [{"delta": {"content": "hi"}, "finish_reason": None}]}, + {"choices": [{"delta": {}, "finish_reason": "stop"}]}, + ], + stream_options={"include_usage": False}, + ) + assert call_args.kwargs["stream_options"]["include_usage"] is False + assert 0 < spent < 100_000 + + +async def test_budgeted_stream_usage_requested_but_missing_charges_remaining(budget_env): + # Nothing suppressed measurement: the engine asked for a usage frame and + # the provider never sent one. A completed stream with content and no + # measurement is the one unmeasured case still charged fail-closed — that + # is what keeps a capped key from being streamed free behind an upstream + # that silently omits its usage. + spent, call_args = await _budgeted_stream( + budget_env, + chunks=[ + {"choices": [{"delta": {"content": "hi"}, "finish_reason": None}]}, + {"choices": [{"delta": {}, "finish_reason": "stop"}]}, + ], + stream_options={"include_usage": True}, + ) + assert call_args.kwargs["stream_options"]["include_usage"] is True + assert spent == 100_000 + + +async def test_budgeted_stream_with_usage_frame_charges_actual(budget_env): + # A usage frame was observed, so only the real (tiny) cost is charged, not the + # full remaining allowance. + spent, _call_args = await _budgeted_stream( + budget_env, + budget_limit_cents=100, + chunks=[ + {"choices": [{"delta": {"content": "hi"}, "finish_reason": None}]}, + { + "usage": {"prompt_tokens": 5000, "completion_tokens": 2000, "total_tokens": 7000}, + "choices": [{"delta": {}, "finish_reason": "stop"}], + }, + ], + ) + assert 0 <= spent < 100_000 + + +async def test_budgeted_blocking_request_sends_no_stream_options(budget_env): + """A cap must not put a streaming-only parameter on a blocking request. + + `include_usage` only decides whether the last frame of a *stream* reports + usage — a non-streaming completion always carries it. LiteLLM forwards the + parameter without looking at `stream`, and OpenAI rejects it on a request + where stream is false, so forcing it here turned every request for a + budgeted key into an upstream 400: the cap made the endpoint unusable + instead of enforced. + """ + make_client, fake, factory, _root = budget_env + key, _key_id = await _make_budgeted_key(factory, budget_limit_cents=100) + + async with await make_client(key) as c: + r = await c.post( + "/v1/chat/completions", + json={ + "model": "gpt-4o-mini", + "messages": [{"role": "user", "content": "hi"}], + }, + ) + + assert r.status_code == 200, r.text + assert "stream_options" not in fake.acompletion.call_args.kwargs + + +async def test_fail_closed_charge_is_recorded_on_the_row_it_charges(budget_env): + """The counter and the request history are one quantity and may not diverge. + + `spent_microcents` is seeded from, and reconciled against, the sum of + `cost_microcents`, so a fail-closed charge that only the counter saw leaves + a key exhausted by an amount no query over its requests can reproduce. + """ + spent, _call_args = await _budgeted_stream( + budget_env, + chunks=[ + {"choices": [{"delta": {"content": "hi"}, "finish_reason": None}]}, + {"choices": [{"delta": {}, "finish_reason": "stop"}]}, + ], + stream_options={"include_usage": True}, + ) + assert spent == 100_000 # the full remaining allowance + + _make_client, _fake, factory, _root = budget_env + from sqlalchemy import select + + from packages.db.models.request_log import RequestLog + + async with factory() as s: + rows = (await s.execute(select(RequestLog.cost_microcents))).scalars().all() + assert rows == [spent] + + +async def _get_spent(factory, key_id: str) -> int: + from sqlalchemy import select + + from packages.db.models.api_key import ApiKey + + async with factory() as s: + return ( + await s.execute(select(ApiKey.spent_microcents).where(ApiKey.id == key_id)) + ).scalar_one() + + +async def test_budgeted_stream_midstream_error_charges_actual_only(budget_env): + # A mid-stream provider error is delivered as a complete error response + # (SSE error frame + terminal [DONE]); the log row records its ~0 cost, so + # settlement is KNOWN and must charge the actual cost only. Before the fix, + # usage_seen stayed False in that branch and every transient provider + # failure permanently exhausted the key (charged cap - spent). + make_client, fake, factory, _root = budget_env + key, key_id = await _make_budgeted_key(factory, budget_limit_cents=10) + + def _failing_stream(): + async def _gen(): + yield {"choices": [{"delta": {"content": "partial"}, "finish_reason": None}]} + raise RuntimeError("upstream exploded") + return _gen() + + fake.acompletion = AsyncMock(return_value=_failing_stream()) + + async with await make_client(key) as c: + async with c.stream( + "POST", + "/v1/chat/completions", + json={ + "model": "gpt-4o-mini", + "stream": True, + "messages": [{"role": "user", "content": "hi"}], + }, + ) as r: + text = "\n".join([line async for line in r.aiter_lines()]) + + # The error response was delivered in full. + assert "Upstream provider error" in text + assert "[DONE]" in text + # The delivery is priced from its characters, floored at one microcent so a + # short answer is never free — and nowhere near the 100_000-microcent cap. + spent = await _get_spent(factory, key_id) + assert 0 < spent < 100_000 + + # The key is NOT exhausted: a follow-up streaming request is still served. + fake.acompletion = AsyncMock(return_value=_ok_stream()) + async with await make_client(key) as c: + async with c.stream( + "POST", + "/v1/chat/completions", + json={ + "model": "gpt-4o-mini", + "stream": True, + "messages": [{"role": "user", "content": "hi again"}], + }, + ) as r2: + assert r2.status_code == 200 + async for _ in r2.aiter_lines(): + pass + + +def _ok_stream(): + async def _gen(): + yield {"choices": [{"delta": {"content": "hi"}, "finish_reason": None}]} + yield { + "usage": {"prompt_tokens": 5, "completion_tokens": 2, "total_tokens": 7}, + "choices": [{"delta": {}, "finish_reason": "stop"}], + } + return _gen() + + +async def test_budgeted_stream_error_after_unmeasured_content_charges_estimate(budget_env): + """Partial content the provider never measured must still cost something. + + The upstream dies mid-generation after a long delivery and no usage frame + ever arrives, so nothing measures it. Charging zero — what the recorded cost + says — would let a capped key stream unbounded tokens free of charge behind + a flaky provider; charging the whole remaining allowance would exhaust the + key for a failure it cannot steer. Settlement is therefore priced from the + delivered characters, and the same number lands on the row and on the key. + """ + make_client, fake, factory, _root = budget_env + key, key_id = await _make_budgeted_key(factory, budget_limit_cents=100) + + delivered = "the quick brown fox " * 2_000 # ~40k chars ≈ 10k tokens + + def _failing_stream(): + async def _gen(): + yield {"choices": [{"delta": {"content": delivered}, "finish_reason": None}]} + raise RuntimeError("upstream exploded mid-generation") + return _gen() + + fake.acompletion = AsyncMock(return_value=_failing_stream()) + + async with await make_client(key) as c: + async with c.stream( + "POST", + "/v1/chat/completions", + json={ + "model": "gpt-4o-mini", + "stream": True, + "messages": [{"role": "user", "content": "say it again " * 400}], + }, + ) as r: + text = "\n".join([line async for line in r.aiter_lines()]) + + assert "[DONE]" in text + + from sqlalchemy import select + + from packages.db.models.request_log import RequestLog + + async with factory() as s: + row = ( + await s.execute( + select(RequestLog).where(RequestLog.api_key_id == key_id) + ) + ).scalars().one() + assert row.output_tokens > 0 # the delivery is recorded, not erased + spent = await _get_spent(factory, key_id) + assert spent == row.cost_microcents # charged == accounted + assert 0 < spent < 1_000_000 # not free, and not the 100-cent cap + + +async def test_budgeted_blocking_without_usage_prices_the_delivery(budget_env): + # A budgeted key whose provider answers without token counts has an unknown + # cost — but a blocking response is here in full, so it is priced from what + # it actually returned, exactly like a mid-stream fault. Charging the whole + # remaining lifetime allowance instead (the old rule) billed a customer + # their entire budget for one ordinary request whose upstream omitted a + # field. + make_client, fake, factory, _root = budget_env + key, key_id = await _make_budgeted_key(factory, budget_limit_cents=10) + + fake.acompletion = AsyncMock(return_value={ + "id": "chatcmpl-no-usage", + "model": "gpt-4o-mini", + "object": "chat.completion", + "created": int(time.time()), + "choices": [{ + "index": 0, + "message": {"role": "assistant", "content": "Hello!"}, + "finish_reason": "stop", + }], + }) + + async with await make_client(key) as c: + r = await c.post( + "/v1/chat/completions", + json={"model": "gpt-4o-mini", + "messages": [{"role": "user", "content": "hi"}]}, + ) + + assert r.status_code == 200, r.text + spent = await _get_spent(factory, key_id) + # Never free (floored at one microcent), never the cap. + assert 0 < spent < 100_000 + + from sqlalchemy import select + + from packages.db.models.request_log import RequestLog + + async with factory() as s: + row = ( + await s.execute( + select(RequestLog).where(RequestLog.api_key_id == key_id) + ) + ).scalar_one() + assert spent == row.cost_microcents # charged == accounted + assert row.output_tokens > 0 # the row explains its own cost + + +async def test_budgeted_blocking_httpexception_charges_recorded_cost(budget_env): + # A budgeted blocking request whose upstream call raised HTTPException never + # received a completion (response == {}, status_code never left 200). The + # fail-closed remaining-charge rule applies only to *delivered* usage-less + # completions — charging the cap here would repeat the mid-stream-error bug + # class on the blocking path. The key must be charged its recorded ~0 cost. + from fastapi import HTTPException + + make_client, fake, factory, _root = budget_env + key, key_id = await _make_budgeted_key(factory, budget_limit_cents=10) + + fake.acompletion = AsyncMock( + side_effect=HTTPException(status_code=429, detail="upstream rate limit") + ) + + async with await make_client(key) as c: + r = await c.post( + "/v1/chat/completions", + json={"model": "gpt-4o-mini", + "messages": [{"role": "user", "content": "hi"}]}, + ) + + assert r.status_code == 429, r.text + assert await _get_spent(factory, key_id) == 0 + + +async def test_budgeted_blocking_with_usage_charges_actual(budget_env): + # Control for the test above: a blocking response WITH usage must charge only + # the recorded cost (never the remaining allowance) — no over-charging. + make_client, fake, factory, _root = budget_env + key, key_id = await _make_budgeted_key(factory, budget_limit_cents=100) + + async with await make_client(key) as c: + r = await c.post( + "/v1/chat/completions", + json={"model": "gpt-4o-mini", + "messages": [{"role": "user", "content": "hi"}]}, + ) + + assert r.status_code == 200, r.text # fixture response carries usage + + from sqlalchemy import select + + from packages.db.models.request_log import RequestLog + + async with factory() as s: + row_cost = ( + await s.execute( + select(RequestLog.cost_microcents).where( + RequestLog.api_key_id == key_id + ) + ) + ).scalar_one() + assert await _get_spent(factory, key_id) == row_cost + + +async def test_budgeted_stream_disconnect_after_usage_charges_actual(budget_env): + """Measured spend must not be re-opened by a later hang-up. + + The usage frame arrives, then the client disconnects. Cost is therefore + KNOWN (the row records it), so settlement charges that cost. Keying the + fail-closed rule on stream completion instead charged the whole remaining + allowance for a request whose tokens were already accounted for. + """ + make_client, fake, factory, _root = budget_env + key, key_id = await _make_budgeted_key(factory, budget_limit_cents=100) + + class _CancelAfterUsage: + def __init__(self): + self._n = 0 + + def __aiter__(self): + return self + + async def __anext__(self): + self._n += 1 + if self._n == 1: + return {"choices": [{"delta": {"content": "hi"}, + "finish_reason": None}]} + if self._n == 2: + return { + "usage": { + "prompt_tokens": 100_000, + "completion_tokens": 50_000, + "total_tokens": 150_000, + }, + "choices": [{"delta": {}, "finish_reason": "stop"}], + } + # Mirrors Starlette cancelling the response task on http.disconnect. + raise asyncio.CancelledError() + + async def aclose(self): + pass + + fake.acompletion = AsyncMock(return_value=_CancelAfterUsage()) + + async with await make_client(key) as c: + try: + async with c.stream( + "POST", + "/v1/chat/completions", + json={ + "model": "gpt-4o-mini", + "stream": True, + "messages": [{"role": "user", "content": "hi"}], + }, + ) as r: + async for _ in r.aiter_lines(): + pass + except Exception: + pass # the injected cancel may surface to the test transport + + from sqlalchemy import select + + from packages.db.models.request_log import RequestLog + + async with factory() as s: + row = ( + await s.execute( + select(RequestLog).where(RequestLog.api_key_id == key_id) + ) + ).scalars().one() + assert row.status_code == 499 + assert row.error_type == "client_disconnect" + assert row.cost_microcents > 0 + spent = await _get_spent(factory, key_id) + assert spent == row.cost_microcents + # The disconnect is not a cost-unknown bail: it must not exhaust the key. + assert spent < 1_000_000 # cap is 100 cents = 1_000_000 microcents + + +async def test_budgeted_stream_disconnect_before_usage_charges_delivery(budget_env): + """A hangup before the usage frame must cost the delivery, not the cap. + + The usage frame is the last chunk, so a user pressing stop mid-answer means + it never arrives and nothing measured the cost. Treating that as + cost-unknown-and-therefore-max charged the entire remaining allowance for + reading a few sentences, which is a normal action and bricks the key + permanently. The delivery is priced from the characters that reached the + client, exactly as the provider-error branch does it. + """ + make_client, fake, factory, _root = budget_env + key, key_id = await _make_budgeted_key(factory, budget_limit_cents=100) + + delivered = "the quick brown fox " * 200 # 4000 chars, no usage frame ever + + class _CancelBeforeUsage: + def __init__(self): + self._n = 0 + + def __aiter__(self): + return self + + async def __anext__(self): + self._n += 1 + if self._n == 1: + return {"choices": [{"delta": {"content": delivered}, + "finish_reason": None}]} + # The client hangs up mid-answer: no usage frame was ever produced. + raise asyncio.CancelledError() + + async def aclose(self): + pass + + fake.acompletion = AsyncMock(return_value=_CancelBeforeUsage()) + + async with await make_client(key) as c: + try: + async with c.stream( + "POST", + "/v1/chat/completions", + json={ + "model": "gpt-4o-mini", + "stream": True, + "messages": [{"role": "user", "content": "hi"}], + }, + ) as r: + async for _ in r.aiter_lines(): + pass + except Exception: + pass # the injected cancel may surface to the test transport + + from sqlalchemy import select + + from packages.db.models.request_log import RequestLog + + async with factory() as s: + row = ( + await s.execute( + select(RequestLog).where(RequestLog.api_key_id == key_id) + ) + ).scalars().one() + assert row.status_code == 499 + assert row.error_type == "client_disconnect" + + spent = await _get_spent(factory, key_id) + # Something real was delivered, so it is not free... + assert spent > 0 + # ...and it is the delivery, not the 1_000_000-microcent remainder. + assert spent == row.cost_microcents + assert spent < 1_000_000 + + # The key still works: a follow-up streaming request is served. + fake.acompletion = AsyncMock(return_value=_ok_stream()) + async with await make_client(key) as c: + async with c.stream( + "POST", + "/v1/chat/completions", + json={ + "model": "gpt-4o-mini", + "stream": True, + "messages": [{"role": "user", "content": "hi again"}], + }, + ) as r2: + assert r2.status_code == 200 + async for _ in r2.aiter_lines(): + pass + + +async def test_budgeted_stream_hangup_before_first_chunk_still_costs_the_prompt(budget_env): + """Bailing at the first byte must not be a way to read a capped key for free. + + Pricing an unmeasured stream from what was delivered is right for a failure + the caller cannot steer, but a disconnect is the caller's own choice, and it + happens after `acompletion` has already sent the prompt upstream. Settling an + empty delivery at zero made "hang up immediately" the cheapest request of + all: the counter never moved, so the lifetime cap stopped applying entirely + and every retry was dispatched upstream and billed there. + """ + make_client, fake, factory, _root = budget_env + key, key_id = await _make_budgeted_key(factory, budget_limit_cents=100) + + class _CancelBeforeAnyChunk: + def __aiter__(self): + return self + + async def __anext__(self): + # The client is gone before the first delta is forwarded. + raise asyncio.CancelledError() + + async def aclose(self): + pass + + fake.acompletion = AsyncMock(return_value=_CancelBeforeAnyChunk()) + prompt = "a very long prompt " * 400 # 7600 chars, priced upstream + + async with await make_client(key) as c: + try: + async with c.stream( + "POST", + "/v1/chat/completions", + json={ + "model": "gpt-4o-mini", + "stream": True, + "messages": [{"role": "user", "content": prompt}], + }, + ) as r: + async for _ in r.aiter_lines(): + pass + except Exception: + pass + + spent = await _get_spent(factory, key_id) + # Nothing was delivered, yet the request was not free. + assert spent > 0 + # And it is the prompt, not the whole remaining allowance. + assert spent < 1_000_000 + + # Repeating the trick keeps charging, so the cap still closes. + for _ in range(3): + fake.acompletion = AsyncMock(return_value=_CancelBeforeAnyChunk()) + async with await make_client(key) as c: + try: + async with c.stream( + "POST", + "/v1/chat/completions", + json={ + "model": "gpt-4o-mini", + "stream": True, + "messages": [{"role": "user", "content": prompt}], + }, + ) as r: + async for _ in r.aiter_lines(): + pass + except Exception: + pass + assert await _get_spent(factory, key_id) > spent + spent = await _get_spent(factory, key_id) + + +async def test_budgeted_blocking_commit_failure_persists_row_and_charge(budget_env): + """A transient write failure must drop neither the row nor the charge. + + The blocking path retries with a FRESH ORM object (the failed attempt's + INSERT was rolled back) and skips the retry when the trace_id is already + durable, so the atomic row+charge unit lands exactly once. + """ + from sqlalchemy import event, select + + from packages.db.models.request_log import RequestLog + + make_client, fake, factory, _root = budget_env + key, key_id = await _make_budgeted_key(factory, budget_limit_cents=100) + + # Fail the log INSERT once, at the cursor: by the time commit runs, the + # row is already flushed (the budget UPDATE autoflushes it), so this is the + # only seam that reproduces a real "database is locked" mid-write. + sync_engine = factory.kw["bind"].sync_engine + failures = {"n": 0} + + def _fail_first_log_insert(conn, cursor, statement, parameters, context, executemany): + if "INSERT INTO requests_log" in statement and failures["n"] == 0: + failures["n"] += 1 + raise RuntimeError("database is locked") + + event.listen(sync_engine, "before_cursor_execute", _fail_first_log_insert) + try: + async with await make_client(key) as c: + r = await c.post( + "/v1/chat/completions", + json={"model": "gpt-4o-mini", + "messages": [{"role": "user", "content": "hi"}]}, + ) + finally: + event.remove(sync_engine, "before_cursor_execute", _fail_first_log_insert) + + assert r.status_code == 200, r.text + assert failures["n"] == 1 # the retry is what saved the write + async with factory() as s: + rows = ( + await s.execute(select(RequestLog).where(RequestLog.api_key_id == key_id)) + ).scalars().all() + assert len(rows) == 1 # never doubled + assert await _get_spent(factory, key_id) == rows[0].cost_microcents + +async def test_budgeted_blocking_missing_usage_empty_content_settles_zero(budget_env): + # An empty delivery with the usage dict missing entirely must NOT fail + # closed: gating on content (not just `response` truthiness) means a 200 + # that carried nothing bills nothing. The alternative (charging the whole + # remaining allowance for an empty answer) permanently exhausts the key + # for a no-op delivery. + make_client, fake, factory, _root = budget_env + key, key_id = await _make_budgeted_key(factory, budget_limit_cents=10) + + fake.acompletion = AsyncMock(return_value={ + "id": "chatcmpl-empty", + "model": "gpt-4o-mini", + "object": "chat.completion", + "created": int(time.time()), + "choices": [{ + "index": 0, + "message": {"role": "assistant", "content": ""}, + "finish_reason": "stop", + }], + }) + + async with await make_client(key) as c: + r = await c.post( + "/v1/chat/completions", + json={"model": "gpt-4o-mini", + "messages": [{"role": "user", "content": "hi"}]}, + ) + + assert r.status_code == 200, r.text + assert await _get_spent(factory, key_id) == 0 + + +async def test_budgeted_blocking_empty_usage_dict_settles_zero(budget_env): + # `usage: {}` on an empty completion is the same empty delivery as a + # usage of explicit zeros — it must not hit the fail-closed arm. + make_client, fake, factory, _root = budget_env + key, key_id = await _make_budgeted_key(factory, budget_limit_cents=10) + + fake.acompletion = AsyncMock(return_value={ + "id": "chatcmpl-empty-usage", + "model": "gpt-4o-mini", + "object": "chat.completion", + "created": int(time.time()), + "choices": [{ + "index": 0, + "message": {"role": "assistant", "content": ""}, + "finish_reason": "stop", + }], + "usage": {}, + }) + + async with await make_client(key) as c: + r = await c.post( + "/v1/chat/completions", + json={"model": "gpt-4o-mini", + "messages": [{"role": "user", "content": "hi"}]}, + ) + + assert r.status_code == 200, r.text + assert await _get_spent(factory, key_id) == 0 + + +async def test_build_log_row_normalizes_anthropic_style_usage(): + # A usage frame keyed input_tokens/output_tokens (the keys /v1/messages + # forwards) must normalize to the request row's token fields, or every + # token-keyed settlement gate treats the frame as zero tokens. + from app.routes.chat import _build_log_row + + class _Body: + messages = [] + model = "claude-3-5-haiku" + stream = False + + class _Kc: + workspace_id = "default" + key_id = "k" + + row = await _build_log_row( + body=_Body(), + kc=_Kc(), + response={ + "model": "claude-3-5-haiku", + "usage": {"input_tokens": 12, "output_tokens": 7}, + }, + status_code=200, + error_type=None, + started_perf=0.0, + strategy="balanced", + requested_model="claude-3-5-haiku", + ) + assert row.input_tokens == 12 + assert row.output_tokens == 7 + + +async def test_budgeted_empty_stream_without_usage_settles_zero(budget_env): + """Nothing delivered and nothing measured must bill nothing. + + The fail-closed charge exists for a delivery the provider never priced. + A stream that yields no content and no usage frame delivered nothing at + all, so charging the key its whole remaining allowance would exhaust it for + an empty response — and every later request with it would 429. + """ + make_client, fake, factory, _root = budget_env + key, key_id = await _make_budgeted_key(factory, budget_limit_cents=10) + + spent, _args = await _budgeted_stream( + budget_env, + chunks=[{"choices": [{"delta": {}, "finish_reason": "stop"}]}], + ) + assert spent == 0 + + # And the key is not exhausted: a follow-up request is still served. + fake.acompletion = AsyncMock(return_value={ + "id": "chatcmpl-after-empty", + "model": "gpt-4o-mini", + "object": "chat.completion", + "created": int(time.time()), + "choices": [{ + "index": 0, + "message": {"role": "assistant", "content": "Hello!"}, + "finish_reason": "stop", + }], + "usage": {"prompt_tokens": 5, "completion_tokens": 2, "total_tokens": 7}, + }) + async with await make_client(key) as c: + r = await c.post( + "/v1/chat/completions", + json={"model": "gpt-4o-mini", + "messages": [{"role": "user", "content": "hi"}]}, + ) + assert r.status_code == 200, r.text + + +async def test_budgeted_disconnect_at_done_frame_still_charges_remaining(budget_env): + """Where the client hangs up must not change what it is charged. + + A stream that completed with content but no usage frame is settled + fail-closed. If the client happened to disconnect at the trailing [DONE] + yield instead of waiting for it, the identical delivery must not be + re-priced as a mid-stream hangup — otherwise reading the whole answer and + leaving is the cheapest way to use a capped key. + """ + make_client, fake, factory, _root = budget_env + key, key_id = await _make_budgeted_key(factory, budget_limit_cents=100) + + delivered = "the quick brown fox " * 200 + + class _HangupAtDone: + def __init__(self): + self._n = 0 + + def __aiter__(self): + return self + + async def __anext__(self): + self._n += 1 + if self._n == 1: + return {"choices": [{"delta": {"content": delivered}, + "finish_reason": None}]} + # The provider is done (the next pull would raise StopAsyncIteration), + # so the engine finishes its own framing; the consumer leaves during + # the [DONE] yield, i.e. after completion. + raise StopAsyncIteration + + async def aclose(self): + pass + + async def _stream(): + yield {"choices": [{"delta": {"content": delivered}, "finish_reason": None}]} + # The consumer goes away the moment the engine starts its trailing + # framing, which is after the provider's stream ended. + await asyncio.sleep(0) + + fake.acompletion = AsyncMock(return_value=_HangupAtDone()) + + async with await make_client(key) as c: + try: + async with c.stream( + "POST", + "/v1/chat/completions", + json={"model": "gpt-4o-mini", "stream": True, + "messages": [{"role": "user", "content": "hi"}]}, + ) as r: + async for _ in r.aiter_lines(): + pass + except Exception: + pass + + await asyncio.sleep(0.05) + spent = await _get_spent(factory, key_id) + assert spent == 1_000_000 # the full cap, i.e. the fail-closed charge + + +async def test_budgeted_stream_usage_without_token_keys_is_unmeasured(budget_env): + """A usage frame with no countable token total is not a measurement. + + `{"total_tokens": 123}` carries no per-direction counts, so every token + read normalizes to zero: charging the recorded cost would bill a delivered + completion nothing, leaving the cap steerable by the upstream's choice of + fields. It must be treated as no usage at all. + """ + make_client, fake, factory, _root = budget_env + key, key_id = await _make_budgeted_key(factory, budget_limit_cents=10) + + spent, _args = await _budgeted_stream( + budget_env, + chunks=[ + {"choices": [{"delta": {"content": "hi"}, "finish_reason": None}]}, + { + "usage": {"total_tokens": 123}, + "choices": [{"delta": {}, "finish_reason": "stop"}], + }, + ], + stream_options={"include_usage": True}, + ) + # A stream that completed with delivered content but no countable usage is + # the fail-closed case, so the key pays its remaining allowance. What must + # NOT happen is the old behavior — the tokenless frame counting as measured + # and settling the delivery at zero. + assert spent == 100_000 + + +async def test_budgeted_blocking_usage_without_token_keys_is_unmeasured(budget_env): + """Same rule on the blocking path: tokenless usage is unmeasured.""" + make_client, fake, factory, _root = budget_env + key, key_id = await _make_budgeted_key(factory, budget_limit_cents=10) + + fake.acompletion = AsyncMock(return_value={ + "id": "chatcmpl-total-only", + "model": "gpt-4o-mini", + "object": "chat.completion", + "created": int(time.time()), + "choices": [{ + "index": 0, + "message": {"role": "assistant", "content": "a delivered answer"}, + "finish_reason": "stop", + }], + "usage": {"total_tokens": 999}, + }) + + async with await make_client(key) as c: + r = await c.post( + "/v1/chat/completions", + json={"model": "gpt-4o-mini", + "messages": [{"role": "user", "content": "hi"}]}, + ) + + assert r.status_code == 200, r.text + spent = await _get_spent(factory, key_id) + assert 0 < spent < 100_000 # priced from the delivery, never free, never the cap +async def test_cancelled_during_backoff_gives_up_exactly_once( + budget_env, monkeypatch +): + """A request torn down between retries parks its cost once. + + The write failed, the backoff is cancelled, and the transaction is abandoned + with the completion already delivered: the obligation has to be parked, and + parked exactly once. The `raise` out of the give-up is what keeps the + write-in-flight arm from running a second give-up for the same settlement — + two warnings for one abandoned row, and a second durability probe while the + process is going away. + + The row is a usage-less completion for a 10-cent key, so the obligation is + the fail-closed 100_000 microcents. + """ + from sqlalchemy import event, select + + import app.routes.chat as chat + from packages.auth.spend import pending_parked_spend + from packages.db.models.request_log import RequestLog + + make_client, fake, factory, _root = budget_env + key, key_id = await _make_budgeted_key(factory, budget_limit_cents=10) + fake.acompletion = AsyncMock(return_value=_completion("Hello!")) + + give_ups: list[int] = [] + real_give_up = chat._give_up_settlement + + async def _count_give_ups(*args, **kwargs): + give_ups.append(1) + return await real_give_up(*args, **kwargs) + + monkeypatch.setattr(chat, "_give_up_settlement", _count_give_ups) + + # Aim the cancellation at the backoff and nowhere earlier: a task is + # cancelled at its next suspension, and the handler's rollback is one, so + # make that call a coroutine that never yields. The sleep is then the only + # place the cancellation can land. + from sqlalchemy.ext.asyncio import AsyncSession + + async def _suspendless_rollback(self, *args, **kwargs): + return None + + monkeypatch.setattr(AsyncSession, "rollback", _suspendless_rollback) + + in_retry = asyncio.Event() + failures = {"n": 0} + + def _fail_the_write(conn, cursor, statement, parameters, context, executemany): + if "INSERT INTO requests_log" in statement and failures["n"] == 0: + failures["n"] += 1 + in_retry.set() + raise RuntimeError("database is locked") + + sync_engine = factory.kw["bind"].sync_engine + event.listen(sync_engine, "before_cursor_execute", _fail_the_write) + + async def _request(): + async with await make_client(key) as c: + return await c.post( + "/v1/chat/completions", + json={"model": "gpt-4o-mini", + "messages": [{"role": "user", "content": "hi"}]}, + ) + + try: + task = asyncio.ensure_future(_request()) + await in_retry.wait() + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + finally: + event.remove(sync_engine, "before_cursor_execute", _fail_the_write) + + assert failures["n"] == 1 # the write really did fail and start backing off + assert give_ups == [1] # abandoned once, not twice and not never + async with factory() as s: + assert not ( + await s.execute(select(RequestLog.id).where(RequestLog.api_key_id == key_id)) + ).all() + # The parked obligation is the charge this settlement decided on: the + # delivery-priced amount for a usage-less blocking completion, never the + # whole remaining allowance. + parked = await pending_parked_spend(key_id) + assert 0 < parked < 100_000 + + +async def test_cancelled_rollback_does_not_skip_the_give_up( + budget_env, monkeypatch +): + """A cancellation inside the retry handler still accounts for the cost. + + Same teardown, other delivery point: the rollback that opens the handler is + itself an await, so it can be where the cancellation lands. An exception + raised from inside a handler is not caught by this `try`'s other arms, so it + used to escape straight past the give-up below — the write lost, the + obligation unparked, and the cap reopened for exactly the key whose write + had just failed. + """ + from sqlalchemy.ext.asyncio import AsyncSession + + import app.routes.chat as chat + from packages.auth.spend import pending_parked_spend + + make_client, fake, factory, _root = budget_env + key, key_id = await _make_budgeted_key(factory, budget_limit_cents=10) + fake.acompletion = AsyncMock(return_value=_completion("Hello!")) + + give_ups: list[int] = [] + real_give_up = chat._give_up_settlement + + async def _count_give_ups(*args, **kwargs): + give_ups.append(1) + return await real_give_up(*args, **kwargs) + + rollbacks = {"n": 0} + real_rollback = AsyncSession.rollback + + async def _cancel_the_first_rollback(self, *args, **kwargs): + rollbacks["n"] += 1 + if rollbacks["n"] == 1: + raise asyncio.CancelledError + return await real_rollback(self, *args, **kwargs) + + failures = {"n": 0} + + def _fail_the_write(conn, cursor, statement, parameters, context, executemany): + if "INSERT INTO requests_log" in statement and failures["n"] == 0: + failures["n"] += 1 + raise RuntimeError("database is locked") + + monkeypatch.setattr(chat, "_give_up_settlement", _count_give_ups) + monkeypatch.setattr(AsyncSession, "rollback", _cancel_the_first_rollback) + sync_engine = factory.kw["bind"].sync_engine + from sqlalchemy import event + + event.listen(sync_engine, "before_cursor_execute", _fail_the_write) + try: + async with await make_client(key) as c: + await c.post( + "/v1/chat/completions", + json={"model": "gpt-4o-mini", + "messages": [{"role": "user", "content": "hi"}]}, + ) + finally: + event.remove(sync_engine, "before_cursor_execute", _fail_the_write) + + assert rollbacks["n"] >= 1 # the cancellation really did land there + assert failures["n"] == 1 + # The cancellation is swallowed rather than escaping, so the loop keeps + # going: the next attempt still finds the session poisoned (the failed flush + # was never rolled back), and the real rollback at the head of that handler + # clears it, so a later attempt lands the row and its charge. Pinned: the + # cost is accounted for once, and no give-up was needed to do it — before + # this the request died here with nothing billed and nothing parked. + assert give_ups == [] + landed = await _get_spent(factory, key_id) + assert 0 < landed < 100_000 # the delivery-priced charge, once + assert await pending_parked_spend(key_id) == 0 + from sqlalchemy import select + + from packages.db.models.request_log import RequestLog + + async with factory() as s: + rows = ( + await s.execute(select(RequestLog.id).where(RequestLog.api_key_id == key_id)) + ).all() + assert len(rows) == 1 + + +# ── Durable recovery: the park outlives the process that lost it ────── + +class _AckLossSession(AsyncSession): + """Commit for real, then report failure as if the ack never came back. + + Armed for settlement commits only — a commit carrying a `RequestLog` row. + The row is durable while its caller still sees an exception: the case the + retry loops' trace_id check exists for. + """ + + drop_ack = False + drops = 0 + + async def commit(self): + settles = self.info.pop("settles_request_log", False) + await super().commit() + if settles and _AckLossSession.drop_ack: + _AckLossSession.drop_ack = False + _AckLossSession.drops += 1 + raise ConnectionError("connection dropped mid-ack") + + async def rollback(self): + self.info.pop("settles_request_log", None) + await super().rollback() + + +def _note_settlement_flush(session, flush_context, instances): + # `commit()` runs after autoflush has already emptied `session.new`, so + # the settlement commit is recognised here, while the row is still new. + from packages.db.models.request_log import RequestLog + + if any(isinstance(o, RequestLog) for o in session.new): + session.info["settles_request_log"] = True + + +from sqlalchemy import event as _sa_event +from sqlalchemy.orm import Session as _SyncSession + +_sa_event.listen(_SyncSession, "before_flush", _note_settlement_flush) + + +class _WriteBlackout: + """Fail every settlement write at the cursor — a sustained write outage. + + Reads still work, which is what makes this the dangerous shape: the key + keeps being served on its pre-check while nothing it spends can be + recorded — not the charge, and not even the park. + """ + + _MATCHES = ( + "INSERT INTO requests_log", + "UPDATE api_keys SET spent_microcents", + "INSERT INTO budget_parks", + ) + + def __init__(self, factory): + self.active = False + self._engine = factory.kw["bind"].sync_engine + from sqlalchemy import event + + event.listen(self._engine, "before_cursor_execute", self._handle) + + def _handle(self, conn, cursor, statement, parameters, context, executemany): + if self.active and any(m in statement for m in self._MATCHES): + raise RuntimeError("database is locked") + + def close(self): + from sqlalchemy import event + + event.remove(self._engine, "before_cursor_execute", self._handle) + + +def _completion(text: str, *, usage: dict | None = None) -> dict: + response = { + "id": "chatcmpl-blackout", + "model": "gpt-4o-mini", + "object": "chat.completion", + "created": int(time.time()), + "choices": [{ + "index": 0, + "message": {"role": "assistant", "content": text}, + "finish_reason": "stop", + }], + "_orca_meta": {"provider": "openai", "litellm_model": "openai/gpt-4o-mini", "latency_ms": 42}, + } + if usage: + response["usage"] = usage + return response + + +async def _lossy_factory(factory): + """A session factory on the same engine whose settlement commits drop acks.""" + from sqlalchemy.ext.asyncio import async_sessionmaker + + return async_sessionmaker( + factory.kw["bind"], expire_on_commit=False, class_=_AckLossSession, + ) + + +async def test_parked_spend_survives_a_real_restart(tmp_sqlite_url): + """A lost settlement must outlive the process that lost it. + + A new engine on the same file, with no process memory carried over, still + folds the park exactly once: the obligation lives in the database, not in + the worker that recorded it. + """ + from sqlalchemy.ext.asyncio import async_sessionmaker + + from packages.auth import spend as spend_mod + from packages.auth.hashing import generate_api_key + from packages.auth.spend import ( + charge_budget, + is_exhausted, + pending_parked_spend, + read_spent, + record_unsettled_spend, + ) + from packages.db import session as session_mod + from packages.db.engine import build_engine + from packages.db.models.api_key import ApiKey + from packages.db.models.base import Base + + cap = 10_000 + engine1 = build_engine(tmp_sqlite_url) + try: + async with engine1.begin() as conn: + await conn.run_sync(Base.metadata.create_all) + factory1 = async_sessionmaker(engine1, expire_on_commit=False) + session_mod._session_factory = factory1 + + full_key, key_hash, key_prefix = generate_api_key() + async with factory1() as s: + row = ApiKey( + workspace_id="default", name="restart", + key_hash=key_hash, key_prefix=key_prefix, + budget_limit_cents=1, + ) + s.add(row) + await s.commit() + await s.refresh(row) + key_id = row.id + await charge_budget(s, key_id, cap, 4_000) + + await record_unsettled_spend( + trace_id="t-restart", api_key_id=key_id, microcents=3_000 + ) + assert await pending_parked_spend(key_id) == 3_000 + finally: + await engine1.dispose() + + spend_mod._unsettled.clear() # the machine stopped: no memory survives + session_mod._session_factory = None + + engine2 = build_engine(tmp_sqlite_url) + try: + factory2 = async_sessionmaker(engine2, expire_on_commit=False) + session_mod._session_factory = factory2 + async with factory2() as s: + assert await is_exhausted(s, key_id, cap) is False # 7_000 of 10_000 + async with factory2() as s: + assert await read_spent(s, key_id) == 7_000 + assert await pending_parked_spend(key_id) == 0 + # A second pre-check must not move the same microcents a second time. + async with factory2() as s: + assert await is_exhausted(s, key_id, cap) is False + async with factory2() as s: + assert await read_spent(s, key_id) == 7_000 + finally: + session_mod._session_factory = None + await engine2.dispose() + + +async def test_park_is_visible_to_another_worker(budget_env, monkeypatch): + """A park recorded on one worker blocks and folds on another.""" + from sqlalchemy.ext.asyncio import async_sessionmaker + + from packages.auth.spend import ( + charge_budget, + is_exhausted, + pending_parked_spend, + read_spent, + record_unsettled_spend, + ) + from packages.db import session as session_mod + + _make_client, _fake, factory, _root = budget_env + _key, key_id = await _make_budgeted_key(factory, budget_limit_cents=1) + cap = 10_000 + async with factory() as s: + await charge_budget(s, key_id, cap, 4_000) + await record_unsettled_spend( + trace_id="t-worker", api_key_id=key_id, microcents=3_000 + ) + + worker_b = async_sessionmaker( + factory.kw["bind"], expire_on_commit=False, + ) + monkeypatch.setattr(session_mod, "_session_factory", worker_b) + async with worker_b() as s: + assert await is_exhausted(s, key_id, cap) is False + assert await pending_parked_spend(key_id) == 0 + async with worker_b() as s: + assert await read_spent(s, key_id) == 7_000 + + +async def test_budgeted_blocking_write_outage_still_bills_the_delivery(budget_env): + """Spend a write outage could not record must not simply disappear. + + Three requests settle while every write — charges and park inserts alike — + fails, so no row and no durable park survives anywhere; the obligations sit + in memory. When the DB recovers, the next settlement re-files and pays for + what was already delivered as well. + """ + from sqlalchemy import select + + from packages.auth.spend import pending_parked_spend + from packages.db.models.request_log import RequestLog + + make_client, fake, factory, _root = budget_env + key, key_id = await _make_budgeted_key(factory, budget_limit_cents=100) + usage = {"prompt_tokens": 10_000, "completion_tokens": 5_000, "total_tokens": 15_000} + fake.acompletion = AsyncMock(side_effect=lambda **kw: _completion("hello", usage=usage)) + + async def _ask(i: int): + async with await make_client(key) as c: + r = await c.post( + "/v1/chat/completions", + json={"model": "gpt-4o-mini", + "messages": [{"role": "user", "content": f"hi {i}"}]}, + ) + assert r.status_code == 200, r.text + + blackout = _WriteBlackout(factory) + blackout.active = True + try: + for i in range(3): + await _ask(i) + assert await _get_spent(factory, key_id) == 0 # nothing was recordable + parked = await pending_parked_spend(key_id) + assert parked > 0 # held in memory: even the park writes failed + blackout.active = False + await _ask(3) + finally: + blackout.close() + + async with factory() as s: + rows = ( + await s.execute(select(RequestLog).where(RequestLog.api_key_id == key_id)) + ).scalars().all() + assert len(rows) == 1 # the three lost settlements left no rows, no doubles + cost = rows[0].cost_microcents + assert cost > 0 + assert await pending_parked_spend(key_id) == 0 + assert await _get_spent(factory, key_id) == 4 * cost + + +async def test_budgeted_stream_write_outage_still_blocks_the_next_request(budget_env): + """The streaming loop's give-up must clamp the next request too. + + A budgeted stream with no usage frame settles fail-closed at the whole + remaining allowance; if that commit is impossible, the amount has to keep + counting, or the outage leaves the key uncapped and the very next request + is served for free. + """ + from packages.auth.spend import pending_parked_spend + + make_client, fake, factory, _root = budget_env + key, key_id = await _make_budgeted_key(factory, budget_limit_cents=10) + + async def _no_usage(): + yield {"choices": [{"delta": {"content": "hi"}, "finish_reason": None}]} + yield {"choices": [{"delta": {}, "finish_reason": "stop"}]} + + fake.acompletion = AsyncMock(return_value=_no_usage()) + + payload = { + "model": "gpt-4o-mini", + "stream": True, + "messages": [{"role": "user", "content": "hi"}], + } + blackout = _WriteBlackout(factory) + blackout.active = True + try: + async with await make_client(key) as c: + async with c.stream("POST", "/v1/chat/completions", json=payload) as r: + assert r.status_code == 200 + async for _ in r.aiter_lines(): + pass + await asyncio.sleep(1.0) # the bounded retries run out after the response + finally: + blackout.close() + + assert await _get_spent(factory, key_id) == 0 + assert await pending_parked_spend(key_id) == 100_000 # the full remainder, held + fake.acompletion = AsyncMock(return_value=_completion( + "hello", usage={"prompt_tokens": 10_000, "completion_tokens": 5_000, "total_tokens": 15_000}, + )) + async with await make_client(key) as c: + r = await c.post( + "/v1/chat/completions", + json={"model": "gpt-4o-mini", + "messages": [{"role": "user", "content": "hi again"}]}, + ) + assert r.status_code == 429, r.text + assert r.json()["error"]["type"] == "rate_limit_error" + + +async def test_budgeted_blocking_commit_ack_loss_bills_the_delivery_once( + budget_env, monkeypatch, +): + """A commit that lands but loses its ack must bill once and park nothing.""" + from sqlalchemy import select + + from packages.auth.spend import pending_parked_spend + from packages.db import session as session_mod + from packages.db.models.request_log import RequestLog + + make_client, fake, factory, _root = budget_env + key, key_id = await _make_budgeted_key(factory, budget_limit_cents=100) + monkeypatch.setattr(session_mod, "_session_factory", await _lossy_factory(factory)) + + _AckLossSession.drop_ack = True + _AckLossSession.drops = 0 + async with await make_client(key) as c: + r = await c.post( + "/v1/chat/completions", + json={"model": "gpt-4o-mini", + "messages": [{"role": "user", "content": "hi"}]}, + ) + assert r.status_code == 200, r.text + assert _AckLossSession.drops == 1 # the scenario actually dropped the ack + + async with factory() as s: + rows = ( + await s.execute(select(RequestLog).where(RequestLog.api_key_id == key_id)) + ).scalars().all() + cost = rows[0].cost_microcents + assert cost > 0 + assert await _get_spent(factory, key_id) == cost + assert await pending_parked_spend(key_id) == 0 # durable, so nothing to park + + await _ask_blocking_once(make_client, key) + assert await _get_spent(factory, key_id) == 2 * cost # not three + + +async def _ask_blocking_once(make_client, key: str) -> None: + async with await make_client(key) as c: + r = await c.post( + "/v1/chat/completions", + json={"model": "gpt-4o-mini", + "messages": [{"role": "user", "content": "hi again"}]}, + ) + assert r.status_code == 200, r.text + + +async def test_budgeted_stream_commit_ack_loss_bills_the_delivery_once( + budget_env, monkeypatch, +): + """A streaming commit that lands but loses its ack bills once, parks nothing.""" + from sqlalchemy import select + + from packages.auth.spend import pending_parked_spend + from packages.db import session as session_mod + from packages.db.models.request_log import RequestLog + + make_client, fake, factory, _root = budget_env + key, key_id = await _make_budgeted_key(factory, budget_limit_cents=100) + fake.acompletion = AsyncMock(return_value=_ok_stream()) + monkeypatch.setattr(session_mod, "_session_factory", await _lossy_factory(factory)) + + _AckLossSession.drop_ack = True + _AckLossSession.drops = 0 + async with await make_client(key) as c: + async with c.stream( + "POST", + "/v1/chat/completions", + json={ + "model": "gpt-4o-mini", + "stream": True, + "messages": [{"role": "user", "content": "hi"}], + }, + ) as r: + assert r.status_code == 200 + async for _ in r.aiter_lines(): + pass + await asyncio.sleep(0.5) # the shielded retry runs out after the response + assert _AckLossSession.drops == 1 + + async with factory() as s: + rows = ( + await s.execute(select(RequestLog).where(RequestLog.api_key_id == key_id)) + ).scalars().all() + cost = rows[0].cost_microcents + assert cost > 0 + assert await _get_spent(factory, key_id) == cost + assert await pending_parked_spend(key_id) == 0 + + +async def _blocking_with_usage(budget_env, monkeypatch, *, usage: dict) -> tuple[int, int]: + """Drive a blocking request whose model has no price; return (key spend, row cost). + + `_lookup_priced_model` is the single oracle both the cost tier and + `_has_known_price` consult, so forcing it to miss makes the response + genuinely "usage with no price" — the shape a custom upstream produces — + without depending on which models the catalog happens to list today. + """ + from sqlalchemy import select + + from app.routes import chat + from packages.db.models.api_key import ApiKey + from packages.db.models.request_log import RequestLog + + monkeypatch.setattr(chat, "_lookup_priced_model", lambda model_id: None) + + make_client, fake, factory, _root = budget_env + key, key_id = await _make_budgeted_key(factory, budget_limit_cents=10) + + fake.acompletion = AsyncMock(return_value={ + "id": "chatcmpl-usage-shape", + "model": "gpt-4o-mini", + "object": "chat.completion", + "created": int(time.time()), + "choices": [{ + "index": 0, + "message": {"role": "assistant", "content": "ok"}, + "finish_reason": "stop", + }], + "usage": usage, + "_orca_meta": { + "provider": "openai", + "litellm_model": "openai/gpt-4o-mini", + "latency_ms": 10, + }, + }) + + async with await make_client(key) as c: + r = await c.post( + "/v1/chat/completions", + json={"model": "gpt-4o-mini", + "messages": [{"role": "user", "content": "hi"}]}, + ) + assert r.status_code == 200, r.text + + async with factory() as s: + spent = ( + await s.execute(select(ApiKey.spent_microcents).where(ApiKey.id == key_id)) + ).scalar_one() + cost = ( + await s.execute( + select(RequestLog.cost_microcents).where(RequestLog.api_key_id == key_id) + ) + ).scalar_one() + return spent, cost + + +async def test_budgeted_blocking_empty_usage_is_not_billed_the_cap(budget_env, monkeypatch): + """An empty delivery costs nothing, blocking or streaming. + + The blocking fail-closed arm is the streaming path's mirror but lost its token + guard, so any 200 with no price charged the key's ENTIRE remaining allowance + regardless of usage — a provider that reports zero tokens on an empty prompt, + or a custom upstream that leaves the field unset, exhausted a capped key with + one request and 429-blocked it for good. A usage frame that carried no tokens + is a measured zero; only tokens without a price are an unknown cost. + """ + spent, cost = await _blocking_with_usage( + budget_env, + monkeypatch, + usage={"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0}, + ) + # Not the cap — that is the point of the test. And not zero either: the + # completion carried content, so it is priced from that delivery (floored at + # one microcent) rather than being treated as a measured zero, which would + # let a capped key be served free behind an unpriceable upstream. + assert 0 < spent < 100_000 + assert cost == spent + + +async def test_budgeted_blocking_unpriced_usage_with_tokens_bills_the_cap( + budget_env, monkeypatch, +): + """The token guard must not disarm the fail-closed arm standing beside it. + + Same unpriceable model, but the usage carried tokens: the cost is unknown and + a budgeted key must not get a delivered completion for free. This is the case + the arm exists for, so it is pinned next to the empty-delivery case — together + they fix the arm's meaning to exactly the streaming path's. + """ + spent, cost = await _blocking_with_usage( + budget_env, + monkeypatch, + usage={"prompt_tokens": 5, "completion_tokens": 2, "total_tokens": 7}, + ) + assert spent == 100_000 # the full remaining allowance (10 cents) + assert cost == 100_000 # recorded on the row it charges, not just the counter + + + +async def test_streaming_commit_task_cancellation_gives_up(budget_env, monkeypatch): + """A cancelled detached commit task still fires the give-up path. + + The detached commit runs so the in-flight stream cannot abort it; but + when the task is cancelled itself (loop teardown, direct cancel), the old + arm swallowed the CancelledError with `pass` and the delivered stream's + cost vanished with no row, no charge, and no park -- the cap reopened + for exactly the key whose settlement failed. + """ + from sqlalchemy.ext.asyncio import AsyncSession + + import app.routes.chat as chat + + make_client, fake, factory, _root = budget_env + key, key_id = await _make_budgeted_key(factory, budget_limit_cents=10) + + give_ups: list[int] = [] + real_give_up = chat._give_up_settlement + + async def _count_give_ups(*args, **kwargs): + give_ups.append(1) + return await real_give_up(*args, **kwargs) + + monkeypatch.setattr(chat, "_give_up_settlement", _count_give_ups) + + in_commit = asyncio.Event() + real_commit = AsyncSession.commit + commits = {"n": 0} + + async def _armed_stolen_commit(self, *args, **kwargs): + # Let the auth-time commit through, kill the request's detached + # settlement commit the way a loop teardown cancels it. + commits["n"] += 1 + if commits["n"] == 1: + return await real_commit(self, *args, **kwargs) + in_commit.set() + raise asyncio.CancelledError() + + monkeypatch.setattr(AsyncSession, "commit", _armed_stolen_commit) + + async def _stream(): + yield {"choices": [{"delta": {"content": "hi"}, "finish_reason": None}]} + yield {"choices": [{"delta": {}, "finish_reason": "stop"}]} + + fake.acompletion = AsyncMock(return_value=_stream()) + + async def _request(): + async with await make_client(key) as c: + try: + async with c.stream( + "POST", "/v1/chat/completions", + json={"model": "gpt-4o-mini", "stream": True, + "messages": [{"role": "user", "content": "hi"}]}, + ) as r: + async for _ in r.aiter_lines(): + pass + except Exception: + pass + + task = asyncio.ensure_future(_request()) + await asyncio.wait_for(in_commit.wait(), timeout=10) + await asyncio.sleep(0.05) # let the detached commit_task finish cancelled + + # The cancellation surfaced inside the retry loop's shield await, so the + # new arm must have counted the give-up exactly once. + assert give_ups == [1] + task.cancel() + try: + await task + except BaseException: + pass diff --git a/tests/unit/test_budget_migration.py b/tests/unit/test_budget_migration.py index e1ec8ffd..8c35fceb 100644 --- a/tests/unit/test_budget_migration.py +++ b/tests/unit/test_budget_migration.py @@ -14,12 +14,14 @@ from sqlalchemy import inspect as sa_inspect from sqlalchemy import select, text +from sqlalchemy.exc import OperationalError from sqlalchemy.ext.asyncio import async_sessionmaker from packages.db.engine import build_engine from packages.db.migrate import ensure_budget_columns from packages.db.models.api_key import ApiKey from packages.db.models.base import Base +from packages.db.models.budget_park import BudgetPark from packages.db.models.request_log import RequestLog @@ -357,3 +359,97 @@ def _record(conn, cursor, statement, parameters, context, executemany): finally: await engine.dispose() + +async def test_ensure_budget_columns_creates_budget_parks_table(tmp_sqlite_url): + """A deployment that predates durable recovery gets the park table. + + `create_all` covers fresh databases, but an upgraded SQLite volume keeps + its old schema — without this step the first give-up would have nowhere + durable to park, and the cap would silently reopen after every restart. + """ + engine = await _legacy_deploy_engine(tmp_sqlite_url) + try: + async with engine.begin() as conn: + await conn.execute(text("DROP TABLE IF EXISTS budget_parks")) + await ensure_budget_columns(engine) + + async with engine.connect() as conn: + tables = await conn.run_sync( + lambda sync: sa_inspect(sync).get_table_names() + ) + assert "budget_parks" in tables + + # The table the migration created actually holds a parked obligation. + factory = async_sessionmaker(engine, expire_on_commit=False) + async with factory() as s: + s.add(BudgetPark(trace_id="park-t1", api_key_id="k1", microcents=900)) + await s.commit() + async with factory() as s: + row = ( + await s.execute( + select(BudgetPark).where(BudgetPark.trace_id == "park-t1") + ) + ).scalar_one() + assert row.microcents == 900 + + # Idempotent across restarts: a second boot changes nothing. + await ensure_budget_columns(engine) + finally: + await engine.dispose() + + +async def test_losing_the_park_table_race_still_boots(tmp_sqlite_url, monkeypatch): + """The one startup DDL a racing boot used to die on. + + `checkfirst` asks the catalog and then creates, so two workers inspecting an + un-migrated schema both hear "no" and one is rejected anyway — and Postgres + reports that collision as a catalog unique violation rather than an + "already exists" message. The create has to land in the same + tolerate-and-continue path as the ALTER, because a worker that dies here + never gets far enough to fold the obligation this table holds. + """ + engine = await _legacy_deploy_engine(tmp_sqlite_url) + async with engine.begin() as conn: + await conn.execute(text("DROP TABLE IF EXISTS budget_parks")) + + attempts: list[str] = [] + + def _lose_the_race(*args, **kwargs): + attempts.append("create") + raise OperationalError( + "CREATE TABLE budget_parks (...)", + {}, + Exception( + 'duplicate key value violates unique constraint ' + '"pg_class_relname_nsp_index"' + ), + ) + + monkeypatch.setattr(BudgetPark.__table__, "create", _lose_the_race) + try: + await ensure_budget_columns(engine) + + assert attempts == ["create"] # it really did take the race + async with engine.connect() as conn: + # and it went on to do the rest of startup. + assert await conn.scalar( + text("SELECT spent_microcents FROM api_keys WHERE workspace_id = 'w1'") + ) == 2500 + finally: + await engine.dispose() + + +def test_catalog_collision_reads_as_already_applied(): + """Only a collision on the object's own name counts as someone winning.""" + from packages.db.migrate import _already_applied + + def _err(msg: str) -> OperationalError: + return OperationalError("CREATE TABLE budget_parks (...)", {}, Exception(msg)) + + assert _already_applied( + _err('duplicate key value violates unique constraint "pg_class_relname_nsp_index"') + ) + assert _already_applied(_err("table budget_parks already exists")) + assert _already_applied(_err("duplicate column name: spent_microcents")) + assert not _already_applied(_err("permission denied to create relation")) + assert not _already_applied(_err('near "TABL": syntax error')) diff --git a/tests/unit/test_budget_spend.py b/tests/unit/test_budget_spend.py index ffb1a0af..7cd10351 100644 --- a/tests/unit/test_budget_spend.py +++ b/tests/unit/test_budget_spend.py @@ -1,15 +1,56 @@ """Unit tests for packages.auth.spend — atomic budget charge under a hard cap.""" import asyncio +import contextlib +from datetime import datetime, timedelta, timezone import pytest +from sqlalchemy import select, update +from sqlalchemy.exc import OperationalError +from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker from packages.auth.spend import ( MICROCENTS_PER_CENT, + budget_precheck, charge_budget, is_exhausted, + pending_parked_spend, read_spent, + record_unsettled_spend, + settle_parked_spend, ) +from packages.db.models.budget_park import BudgetPark + +_PARK_EPOCH = datetime(2026, 1, 1, tzinfo=timezone.utc) + + +async def _stamp_parks(env, offsets: dict[str, int]) -> None: + """Give each park a distinct, host-independent `created_at`. + + The column default is wall-clock and its resolution is platform-dependent + (two back-to-back `datetime.now()` calls return the same value on Windows, + whose clock ticks every ~15.6 ms), so parks written moments apart can share + a stamp and fall through to the `trace_id` tiebreak. Pinning the stamps + keeps these tests about fold ordering rather than about the host clock. + """ + async with env() as s: + for trace_id, offset in offsets.items(): + await s.execute( + update(BudgetPark) + .where(BudgetPark.trace_id == trace_id) + .values(created_at=_PARK_EPOCH + timedelta(seconds=offset)) + ) + await s.commit() + + +@pytest.fixture(autouse=True) +def _isolated_memory_holds(): + """`_unsettled` is process state, so one test's leftover hold is another's bug.""" + from packages.auth import spend as spend_mod + + spend_mod._unsettled.clear() + yield + spend_mod._unsettled.clear() @pytest.fixture @@ -56,8 +97,6 @@ async def test_concurrent_charges_never_exceed_cap(tmp_sqlite_url): still makes the outcome deterministic — one charge fits, the other's guard matches no row and its clamp fills the counter to exactly `cap`. """ - from sqlalchemy.ext.asyncio import async_sessionmaker - from packages.db.engine import build_engine from packages.db.models.api_key import ApiKey from packages.db.models.base import Base @@ -138,3 +177,371 @@ async def test_stale_identity_map_cannot_clobber_a_concurrent_charge(tmp_sqlite_ def test_microcent_conversion_constant(): assert MICROCENTS_PER_CENT == 10_000 + + +@pytest.fixture +async def parked_env(tmp_sqlite_url): + """Engine + global session factory, so the park ledger is durable here.""" + async with _park_ledger(tmp_sqlite_url) as factory: + yield factory + + +class _LostAckParkSession(AsyncSession): + """``AsyncSession`` whose park COMMIT applies and then reports failure. + + Losing the ack is the state that matters: a commit that landed and a commit + that failed look identical to the caller, and treating the first as the + second is what puts one obligation in the table and in memory at once. Only + a session inserting a park is rigged, so the fixture's own bookkeeping + commits run untouched. + """ + + lost_acks_left = 0 + + async def commit(self): + if type(self).lost_acks_left > 0 and any( + isinstance(o, BudgetPark) for o in self.sync_session.new + ): + type(self).lost_acks_left -= 1 + await super().commit() + raise OperationalError("COMMIT", {}, Exception("connection reset before ack")) + return await super().commit() + + +@contextlib.asynccontextmanager +async def _park_ledger(tmp_sqlite_url, session_class=AsyncSession): + from packages.db import session as session_mod + from packages.db.engine import build_engine + from packages.db.models.base import Base + + engine = build_engine(tmp_sqlite_url) + async with engine.begin() as conn: + await conn.run_sync(Base.metadata.create_all) + factory = async_sessionmaker(engine, expire_on_commit=False, class_=session_class) + old, session_mod._session_factory = session_mod._session_factory, factory + try: + yield factory + finally: + session_mod._session_factory = old + await engine.dispose() + + +async def _parks(factory, api_key_id: str) -> dict[str, int]: + """What is still parked for a key, by `trace_id`.""" + async with factory() as s: + rows = ( + await s.execute( + select(BudgetPark.trace_id, BudgetPark.microcents).where( + BudgetPark.api_key_id == api_key_id + ) + ) + ).all() + return {trace_id: int(microcents) for trace_id, microcents in rows} + + +async def _parked_key(factory, *, spent: int = 0): + from packages.db.models.api_key import ApiKey + + async with factory() as s: + k = ApiKey( + workspace_id="default", name="p", key_hash="h-p", key_prefix="p-p", + ) + s.add(k) + await s.commit() + await s.refresh(k) + if spent: + await charge_budget(s, k.id, 10_000_000, spent) + return k.id + + +async def test_recorded_park_folds_exactly_once(parked_env): + """A lost settlement is billed once, by the next pre-check — never twice.""" + from packages.auth import spend as spend_mod + + key_id = await _parked_key(parked_env, spent=4_000) + cap = 10_000 + await record_unsettled_spend(trace_id="t-fold", api_key_id=key_id, microcents=3_000) + assert await pending_parked_spend(key_id) == 3_000 # durable, not just memory + + spend_mod._unsettled.clear() # the machine stopped and cold-started + + async with parked_env() as s: + assert await is_exhausted(s, key_id, cap) is False # 7_000 of 10_000 + async with parked_env() as s: + assert await read_spent(s, key_id) == 7_000 + assert await pending_parked_spend(key_id) == 0 + + # A second pre-check must not move the same microcents a second time. + async with parked_env() as s: + assert await is_exhausted(s, key_id, cap) is False + async with parked_env() as s: + assert await read_spent(s, key_id) == 7_000 + + +async def test_a_park_that_lost_its_ack_is_not_also_held_in_memory( + tmp_sqlite_url, monkeypatch +): + """A commit that applied must not report itself as one that did not. + + The old test of this name called `record_unsettled_spend` twice against a + working database, which only exercises the retried insert — the durability + probe never ran. Here the COMMIT lands and the ack does not, so `_insert_park` + has to ask the table instead of trusting its own failure. Reporting it as + not durable is what left the obligation parked *and* held in memory: another + worker folds and deletes the row, this process then re-files its stale copy + under a `trace_id` that no longer collides, and the same delivery is billed + twice against a key that is now pinned on its cap with no debt left to fold. + """ + from packages.auth import spend as spend_mod + + async with _park_ledger(tmp_sqlite_url, _LostAckParkSession) as factory: + key_id = await _parked_key(factory) + monkeypatch.setattr(_LostAckParkSession, "lost_acks_left", 1) + await record_unsettled_spend( + trace_id="t-ack", api_key_id=key_id, microcents=900 + ) + + assert spend_mod._unsettled == {} # the probe found the durable row + assert await _parks(factory, key_id) == {"t-ack": 900} + assert await pending_parked_spend(key_id) == 900 + + assert await settle_parked_spend(key_id, 10_000) == 900 + assert await pending_parked_spend(key_id) == 0 + async with factory() as s: + assert await read_spent(s, key_id) == 900 + + +async def test_folding_a_row_this_process_also_holds_clears_the_hold( + parked_env, monkeypatch +): + """The fold that bills a durable row drops the memory hold for it too. + + A path the durability probe does not close: the probe reads on a session of + its own and can fail while the row is real, and this same call then picks + that row up and bills it. Leaving the hold behind means the next pre-check + re-files it as a new park — the row it mirrored is gone, so nothing collides. + Only fully-billed rows are dropped; a trimmed one is still owed. + """ + from packages.auth import spend as spend_mod + + key_id = await _parked_key(parked_env) + await record_unsettled_spend(trace_id="t-both", api_key_id=key_id, microcents=1_200) + spend_mod._unsettled[(key_id, "t-both")] = 1_200 + + async def _unreachable(**kwargs): + return False + + monkeypatch.setattr(spend_mod, "_insert_park", _unreachable) + assert await settle_parked_spend(key_id, 10_000) == 1_200 + assert (key_id, "t-both") not in spend_mod._unsettled + assert await pending_parked_spend(key_id) == 0 + + monkeypatch.undo() + assert await settle_parked_spend(key_id, 10_000) == 0 # nothing re-files + async with parked_env() as s: + assert await read_spent(s, key_id) == 1_200 + + +async def test_a_fold_that_overshoots_the_cap_keeps_its_memory_hold( + parked_env, monkeypatch +): + """A trimmed row is still owed, so the hold beside it must stay. + + The reconcile in the test above is deliberately narrow: running it over the + trimmed row as well would write off the part of a delivery the cap could not + absorb. + """ + from packages.auth import spend as spend_mod + + key_id = await _parked_key(parked_env, spent=9_000) + await record_unsettled_spend(trace_id="t-trim", api_key_id=key_id, microcents=2_000) + spend_mod._unsettled[(key_id, "t-trim")] = 2_000 + + async def _unreachable(**kwargs): + return False + + monkeypatch.setattr(spend_mod, "_insert_park", _unreachable) + assert await settle_parked_spend(key_id, 10_000) == 1_000 + assert spend_mod._unsettled[(key_id, "t-trim")] == 2_000 + assert await _parks(parked_env, key_id) == {"t-trim": 1_000} + + +async def test_park_beyond_the_remainder_bills_what_fits(parked_env): + """An oversized park converges the counter on the cap instead of freezing. + + Moving the whole 2_000 row would push the counter past the cap, and leaving + it parked whole would refuse the key at a lifetime spend below its limit on + the strength of a row nothing ever shrinks. So the 1_000 the cap can absorb + bills and the row keeps the rest: the over-claim stays visible and still + blocks, and with a park outstanding the counter is now exactly on the cap. + """ + key_id = await _parked_key(parked_env, spent=9_000) + cap = 10_000 + await record_unsettled_spend(trace_id="t-big", api_key_id=key_id, microcents=2_000) + + async with parked_env() as s: + assert await is_exhausted(s, key_id, cap) is True + async with parked_env() as s: + assert await read_spent(s, key_id) == 10_000 # the cap, never past it + assert await pending_parked_spend(key_id) == 1_000 # the debt stays visible + + # A second pre-check must not bill the microcents the first one moved, and + # the remainder must not clear on its own. + async with parked_env() as s: + assert await is_exhausted(s, key_id, cap) is True + async with parked_env() as s: + assert await read_spent(s, key_id) == 10_000 + assert await pending_parked_spend(key_id) == 1_000 + + +async def test_parked_queue_drains_oldest_first_as_the_cap_opens(parked_env): + """What the cap could not absorb waits, and folds the moment it can. + + The oversized head takes the whole allowance, so the younger park behind it + waits — not lost, just queued. Raising the cap reopens room, and the fold + keeps its promise that the counter reaches `min(cap, spent + debt)`. + + Which row shrinks is the assertion that matters: the totals come out the + same either way, so only the ledger shows whether the head of the queue was + the older obligation or whichever `trace_id` sorts first. + """ + key_id = await _parked_key(parked_env, spent=9_000) + await record_unsettled_spend(trace_id="t-old", api_key_id=key_id, microcents=2_000) + await record_unsettled_spend(trace_id="t-new", api_key_id=key_id, microcents=500) + # Pinned so the assertion is about fold order and not about the host clock + # (see `_stamp_parks`); `t-old` is the older obligation. + await _stamp_parks(parked_env, {"t-old": 0, "t-new": 1}) + + assert await settle_parked_spend(key_id, 10_000) == 1_000 + # `t-old` absorbed the whole allowance and kept its remainder; the younger, + # smaller park behind it is untouched. + assert await _parks(parked_env, key_id) == {"t-old": 1_000, "t-new": 500} + assert await settle_parked_spend(key_id, 10_000) == 0 # no room left + assert await pending_parked_spend(key_id) == 1_500 + + assert await settle_parked_spend(key_id, 11_000) == 1_000 + assert await _parks(parked_env, key_id) == {"t-new": 500} + assert await pending_parked_spend(key_id) == 500 + assert await settle_parked_spend(key_id, 12_000) == 500 + assert await pending_parked_spend(key_id) == 0 + async with parked_env() as s: + assert await read_spent(s, key_id) == 11_500 + + +async def test_tied_park_stamps_still_bill_to_the_cap(parked_env): + """Parks sharing a `created_at` fall to `trace_id`, and that order is arbitrary. + + The column default's resolution is platform-dependent, so two obligations + filed in the same clock tick tie — and `trace_id` is a uuid4 in production, + so which row is trimmed first is arbitrary by construction. What must hold + regardless of which row wins is the part every caller depends on: the whole + parked total is accounted, the counter lands on the cap, and the remainder + stays visible as debt rather than being written off. + """ + key_id = await _parked_key(parked_env, spent=9_000) + await record_unsettled_spend(trace_id="t-old", api_key_id=key_id, microcents=2_000) + await record_unsettled_spend(trace_id="t-new", api_key_id=key_id, microcents=500) + await _stamp_parks(parked_env, {"t-old": 0, "t-new": 0}) # the tie + + assert await settle_parked_spend(key_id, 10_000) == 1_000 + async with parked_env() as s: + assert await read_spent(s, key_id) == 10_000 + # Whichever row absorbed the allowance, the untouched other row plus the + # trimmed remainder still adds up to the 1_500 that could not be absorbed. + assert sum((await _parks(parked_env, key_id)).values()) == 1_500 + assert await pending_parked_spend(key_id) == 1_500 + + assert await settle_parked_spend(key_id, 20_000) == 1_500 + assert await pending_parked_spend(key_id) == 0 + + +async def test_an_unreadable_park_ledger_does_not_read_as_no_debt(parked_env): + """A failed durable SUM has to block dispatch, not clear the key's cap. + + The durable table is the only record of what another worker, or this process + before it stopped, still owes — and it is a *different* read from the rest of + the pre-check: the counter comes in on the request's own connection, while + the ledger opens a fresh one from the factory. So a pool checkout that times + out can hide every park while the request itself would otherwise work, and + folding that failure into the total as zero dispatches the exact key the park + exists to hold shut. + + The factory is swapped by hand rather than with `monkeypatch`: that teardown + runs after the fixture's and would restore the rigged one, leaving every + later test in the session reading a disposed engine. + """ + from packages.db import session as session_mod + + key_id = await _parked_key(parked_env, spent=1_000) + cap = 10_000 + async with parked_env() as s: + assert await is_exhausted(s, key_id, cap) is False + + def _blinded(): + raise TimeoutError("database connection checkout timed out") + + real = session_mod._session_factory + try: + session_mod._session_factory = _blinded + assert await pending_parked_spend(key_id) is None + async with parked_env() as s: + assert await budget_precheck(s, key_id, cap) == cap + assert await is_exhausted(s, key_id, cap) is True + finally: + session_mod._session_factory = real + + +async def test_concurrent_folds_bill_the_park_once(parked_env): + """Two workers folding the same park move it exactly once.""" + key_id = await _parked_key(parked_env) + await record_unsettled_spend(trace_id="t-race", api_key_id=key_id, microcents=3_000) + + moved = await asyncio.gather( + settle_parked_spend(key_id, 10_000), settle_parked_spend(key_id, 10_000), + ) + assert sorted(moved) == [0, 3_000] + async with parked_env() as s: + assert await read_spent(s, key_id) == 3_000 + assert await pending_parked_spend(key_id) == 0 + + +async def test_memory_fallback_refiles_once_the_database_recovers(parked_env): + """An amount held in memory during an outage becomes exactly one park.""" + from packages.db import session as session_mod + + key_id = await _parked_key(parked_env) + old = session_mod._session_factory + session_mod._session_factory = None + try: + # No session factory: the database might as well be down. + await record_unsettled_spend(trace_id="t-mem", api_key_id=key_id, microcents=700) + assert await pending_parked_spend(key_id) == 700 + finally: + session_mod._session_factory = old + + async with parked_env() as s: + assert await is_exhausted(s, key_id, 10_000) is False + assert await pending_parked_spend(key_id) == 0 + async with parked_env() as s: + assert await read_spent(s, key_id) == 700 + + +async def test_cancelled_refile_keeps_the_memory_copy(parked_env, monkeypatch): + """Cancelling a memory-to-database refile must not drop the obligation.""" + from packages.auth import spend as spend_mod + + key_id = await _parked_key(parked_env) + spend_mod._unsettled[(key_id, "t-cancel")] = 500 + + async def _cancelled(**kwargs): + raise asyncio.CancelledError + + monkeypatch.setattr(spend_mod, "_insert_park", _cancelled) + with pytest.raises(asyncio.CancelledError): + await settle_parked_spend(key_id, 10_000) + assert spend_mod._unsettled[(key_id, "t-cancel")] == 500 + + monkeypatch.undo() + assert await settle_parked_spend(key_id, 10_000) == 500 + async with parked_env() as s: + assert await read_spent(s, key_id) == 500 diff --git a/tests/unit/test_give_up_settlement.py b/tests/unit/test_give_up_settlement.py new file mode 100644 index 00000000..cd484d4c --- /dev/null +++ b/tests/unit/test_give_up_settlement.py @@ -0,0 +1,74 @@ +"""The give-up's durability gate (app/routes/chat.py::_give_up_settlement). + +Parking is the last resort for a settlement the database never accepted: it +keeps a delivered response counting against the key's cap until a budget +pre-check can fold it back into the counter. But the final attempt has no retry +left to run the trace_id check, so it gives up on the exception alone — and an +exception can follow a commit that applied (ack lost, or a cancellation that +landed after it). Parking that cost again would bill one delivery twice. + +These run without a session factory, so the park lands in the process-memory +overflow `record_unsettled_spend` falls back to; the durable ledger itself is +covered in tests/unit/test_budget_spend.py. +""" + +from __future__ import annotations + +import asyncio + +import pytest + +from app.routes.chat import _give_up_settlement +from packages.auth.spend import pending_parked_spend + + +class _Kc: + """The two attributes the give-up reads off a KeyContext.""" + + def __init__(self, key_id: str, *, cap: int = 10_000): + self.key_id = key_id + self._budget_cap = cap + + +def _failed(error: str = "connection dropped mid-ack") -> RuntimeError: + return RuntimeError(error) + + +async def test_durable_settlement_is_not_parked_again(): + """The row is committed, so its charge already counts: park nothing.""" + kc = _Kc("durable-key") + + async def persisted() -> bool: + return True + + await _give_up_settlement(kc, "trace-durable", 900, 3, _failed(), persisted) + assert await pending_parked_spend(kc.key_id) == 0 + + +async def test_lost_settlement_parks_its_own_cost(): + """Nothing is durable: the delivery stays on the cap until it is folded.""" + kc = _Kc("lost-key") + + async def persisted() -> bool: + return False + + await _give_up_settlement( + kc, "trace-lost", 900, 3, _failed("database is locked"), persisted + ) + assert await pending_parked_spend(kc.key_id) == 900 + + +async def test_probe_cancelled_parks_before_propagating(): + """Torn down mid-read: the outcome is unknown, so park and re-raise. + + Dropping it here would be fail-open — the cancellation would take this + request's cost out with the coroutine. + """ + kc = _Kc("cancelled-probe-key") + + async def persisted() -> bool: + raise asyncio.CancelledError + + with pytest.raises(asyncio.CancelledError): + await _give_up_settlement(kc, "trace-cancelled", 900, 1, _failed(), persisted) + assert await pending_parked_spend(kc.key_id) == 900 diff --git a/tests/unit/test_unmeasured_stream_settlement.py b/tests/unit/test_unmeasured_stream_settlement.py new file mode 100644 index 00000000..110eb36d --- /dev/null +++ b/tests/unit/test_unmeasured_stream_settlement.py @@ -0,0 +1,41 @@ +"""Delivery estimates only price what nothing else measured.""" + +from __future__ import annotations + +from types import SimpleNamespace + +from app.routes.chat import _countable_usage, _estimate_usage, _text_chars + + +def _body(prompt: str): + return SimpleNamespace(messages=[SimpleNamespace(content=prompt)]) + + +def test_measured_usage_is_never_replaced_by_an_estimate(): + assert _countable_usage({"prompt_tokens": 11, "completion_tokens": 22}) + assert _countable_usage({"input_tokens": 3, "output_tokens": 4}) + + +def test_nothing_delivered_stays_unbilled(): + # A failure before the first content chunk delivered no tokens; this is the + # input the table treats as known-zero rather than fail-closed. + assert not _countable_usage({}) + assert not _countable_usage({"total_tokens": 123}) + + +def test_empty_client_bail_still_costs_the_prompt(): + # Character math behind the prompt-only estimate for a disconnect before + # the first byte. + assert _estimate_usage(400, 0) == {"prompt_tokens": 100, "completion_tokens": 1} + + +def test_estimate_prices_prompt_and_delivery_at_char_quarter(): + got = _estimate_usage(400, 4_000) + assert got == {"prompt_tokens": 100, "completion_tokens": 1000} + + +def test_content_part_lists_count_their_text(): + content = [{"type": "text", "text": "z" * 800}, {"type": "image_url"}] + assert _text_chars(content) == 800 + assert _text_chars("plain") == 5 + assert _text_chars(None) == 0