From d122e2f7593a871e66ccd03141befe5137a5ada9 Mon Sep 17 00:00:00 2001 From: Hasit Bhatt Date: Wed, 26 Aug 2026 02:22:54 -0700 Subject: [PATCH 01/18] feat(budget): enforce budget_limit_cents on chat path --- app/routes/chat.py | 29 +++ packages/auth/spend.py | 48 +++++ packages/db/models/request_log.py | 1 + tests/integration/test_budget_enforcement.py | 208 +++++++++++++++++++ tests/unit/test_budget_spend.py | 106 ++++++++++ 5 files changed, 392 insertions(+) create mode 100644 packages/auth/spend.py create mode 100644 tests/integration/test_budget_enforcement.py create mode 100644 tests/unit/test_budget_spend.py diff --git a/app/routes/chat.py b/app/routes/chat.py index e1a69b0..2f50ffe 100644 --- a/app/routes/chat.py +++ b/app/routes/chat.py @@ -31,6 +31,7 @@ from app.protocols.sse import AdapterError from app.quality_scores import resolve_model_metrics from app.schemas import ChatCompletionRequest +from packages.auth.spend import budget_exceeded, get_lifetime_spend_microcents from packages.auth.types import KeyContext from packages.db.models.request_log import RequestLog from packages.litellm_adapter.catalog import CATALOG, CATALOG_BY_ID @@ -368,6 +369,34 @@ async def execute_chat( detail=f"Model '{body.model}' is not allowed for this API key", ) + # Budget enforcement: `budget_limit_cents` is a lifetime cap on this + # key's total spend (sum of `cost_microcents` for all non-deleted + # request-log rows, including streaming 499/503 rows that already + # incurred provider billing). Checked before any routing, resolution, + # or cache work so an exhausted key costs the operator nothing — no + # upstream attempt, no cache fill. + # + # NOTE: This is a best-effort soft limit, not a hard atomic cap. + # Spend is read before the request and the current request's cost is + # only written after the response/stream completes. N concurrent + # requests from the same budgeted key all observe the same + # pre-request total and may all pass the check, exceeding the cap by + # up to N× per-request cost in a burst. A hard cap would require a + # reservation/claim or row-level lock before dispatch; the current + # design trades strictness for simplicity and avoids holding a DB + # transaction across the upstream call. See spend.py for aggregation + # semantics. + if kc.budget_limit_cents is not None: + spend = await get_lifetime_spend_microcents(db, str(kc.key_id)) + if budget_exceeded(spend, kc.budget_limit_cents): + raise HTTPException( + status_code=429, + detail=( + "API key budget exhausted " + f"({spend} of {kc.budget_limit_cents * 10_000} microcents spent)." + ), + ) + 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" diff --git a/packages/auth/spend.py b/packages/auth/spend.py new file mode 100644 index 0000000..cb5cc07 --- /dev/null +++ b/packages/auth/spend.py @@ -0,0 +1,48 @@ +"""Per-key spend lookup used to enforce `ApiKey.budget_limit_cents`. + +Semantics: `budget_limit_cents` is a lifetime cap on the key's total +spend — sum of `cost_microcents` for all non-deleted request-log rows. +1 cent = 10,000 microcents (1 USD = 1,000,000 microcents, matching +chat.py's cost math). + +Rows are counted regardless of HTTP status because the streaming path +records billable token usage even when the final status is 499 (client +disconnect, chat.py:596) or 503 (mid-stream upstream failure, +chat.py:646); filtering on `status_code < 400` would exclude those and +make the cap bypassable by closing the stream early after reading the +usage chunk. + +Kept free of FastAPI imports so it stays unit-testable and reusable from +non-HTTP contexts (background jobs, CLI minting tools). +""" + +from __future__ import annotations + +from sqlalchemy import func, select +from sqlalchemy.ext.asyncio import AsyncSession + +from packages.db.models.request_log import RequestLog + +MICROCENTS_PER_CENT = 10_000 + + +async def get_lifetime_spend_microcents( + session: AsyncSession, api_key_id: str +) -> int: + """Sum of all non-deleted spend ever recorded for this key. + + Counts every row with `cost_microcents` regardless of `status_code` + so that streaming disconnect (499) and mid-stream upstream failure + (503) costs — which already incurred provider billing — are not + excluded from the budget. Failed requests with zero cost contribute + nothing to the sum regardless. + """ + stmt = select(func.coalesce(func.sum(RequestLog.cost_microcents), 0)).where( + RequestLog.api_key_id == api_key_id, + RequestLog.is_deleted == 0, + ) + return int((await session.execute(stmt)).scalar_one()) + + +def budget_exceeded(spend_microcents: int, budget_limit_cents: int) -> bool: + return spend_microcents >= budget_limit_cents * MICROCENTS_PER_CENT diff --git a/packages/db/models/request_log.py b/packages/db/models/request_log.py index 9871610..84c7c06 100644 --- a/packages/db/models/request_log.py +++ b/packages/db/models/request_log.py @@ -12,6 +12,7 @@ class RequestLog(Base, UUIDMixin, SoftDeleteMixin): __tablename__ = "requests_log" __table_args__ = ( Index("ix_requests_log_ws_created", "workspace_id", "created_at"), + Index("ix_requests_log_api_key_spend", "api_key_id", "is_deleted"), ) workspace_id: Mapped[str] = mapped_column(String(36), nullable=False, index=True) diff --git a/tests/integration/test_budget_enforcement.py b/tests/integration/test_budget_enforcement.py new file mode 100644 index 0000000..cf75aed --- /dev/null +++ b/tests/integration/test_budget_enforcement.py @@ -0,0 +1,208 @@ +"""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 time +from unittest.mock import AsyncMock + +import pytest + + +@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.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, + )) + 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() diff --git a/tests/unit/test_budget_spend.py b/tests/unit/test_budget_spend.py new file mode 100644 index 0000000..2dfe25a --- /dev/null +++ b/tests/unit/test_budget_spend.py @@ -0,0 +1,106 @@ +"""Unit tests for packages.auth.spend — lifetime spend aggregation.""" + +import pytest + +from packages.auth.spend import ( + MICROCENTS_PER_CENT, + budget_exceeded, + get_lifetime_spend_microcents, +) + + +@pytest.fixture +async def seeded_log(db_session): + """Two keys with a mix of billable / streaming-failure / soft-deleted rows.""" + from packages.db.models.api_key import ApiKey + from packages.db.models.request_log import RequestLog + + k1 = ApiKey(workspace_id="default", name="a", key_hash="h-a", key_prefix="p-a") + k2 = ApiKey(workspace_id="default", name="b", key_hash="h-b", key_prefix="p-b") + db_session.add_all([k1, k2]) + await db_session.flush() + + rows = [ + RequestLog( + workspace_id="default", api_key_id=k1.id, model_requested="m", + model_resolved="m", provider="openai", input_tokens=1, output_tokens=1, + cost_microcents=1000, status_code=200, routing_strategy="balanced", latency_ms=10, trace_id="t-1", + ), + RequestLog( + workspace_id="default", api_key_id=k1.id, model_requested="m", + model_resolved="m", provider="openai", input_tokens=1, output_tokens=1, + cost_microcents=500, status_code=200, routing_strategy="balanced", latency_ms=10, trace_id="t-1", + ), + # Streaming failures ARE billable when they carry cost — the + # provider billed tokens even though the final status is 503 + # (mid-stream upstream failure) or 499 (client disconnect). + # Filtering on status_code < 400 would exclude these and make the + # budget bypassable, so they must be counted. + RequestLog( + workspace_id="default", api_key_id=k1.id, model_requested="m", + model_resolved="m", provider="openai", input_tokens=9, output_tokens=9, + cost_microcents=999_999, status_code=503, routing_strategy="balanced", latency_ms=10, trace_id="t-3", + ), + # soft-deleted rows must never count, even with non-zero cost + RequestLog( + workspace_id="default", api_key_id=k1.id, model_requested="m", + model_resolved="m", provider="openai", input_tokens=1, output_tokens=1, + cost_microcents=12345, status_code=200, routing_strategy="balanced", latency_ms=10, trace_id="t-5", + is_deleted=1, + ), + # another key's spend must not leak in + RequestLog( + workspace_id="default", api_key_id=k2.id, model_requested="m", + model_resolved="m", provider="openai", input_tokens=2, output_tokens=2, + cost_microcents=777_777, status_code=200, routing_strategy="balanced", latency_ms=10, trace_id="t-4", + ), + ] + db_session.add_all(rows) + await db_session.commit() + return k1, k2 + + +async def test_spend_sums_only_billable_rows_for_the_key(db_session, seeded_log): + k1, _k2 = seeded_log + spend = await get_lifetime_spend_microcents(db_session, k1.id) + # 1000 + 500 + 999_999 (503 failure with cost is now counted) = 1,001,499; + # soft-deleted 12345 is excluded. + assert spend == 1_001_499 + + +async def test_spend_counts_stream_disconnect_and_mid_stream_failure(db_session, seeded_log): + """Regression for P1: 499/503 streaming costs must count toward budget.""" + from packages.db.models.request_log import RequestLog + + k1, _k2 = seeded_log + # Add explicit 499 disconnect row with cost + db_session.add( + RequestLog( + workspace_id="default", api_key_id=k1.id, model_requested="m", + model_resolved="m", provider="openai", input_tokens=2, output_tokens=2, + cost_microcents=42_000, status_code=499, routing_strategy="balanced", latency_ms=10, trace_id="t-6", + ) + ) + await db_session.commit() + spend = await get_lifetime_spend_microcents(db_session, k1.id) + assert spend == 1_001_499 + 42_000 + + +async def test_empty_history_is_zero(db_session, seeded_log): + _k1, k2 = seeded_log + from packages.db.models.api_key import ApiKey + + fresh = ApiKey(workspace_id="default", name="c", key_hash="h-c", key_prefix="p-c") + db_session.add(fresh) + await db_session.commit() + assert await get_lifetime_spend_microcents(db_session, fresh.id) == 0 + + +def test_budget_exceeded_boundary(): + assert budget_exceeded(10_000 - 1, 1) is False # just under 1 cent + assert budget_exceeded(10_000, 1) is True # exactly at the cap blocks + assert budget_exceeded(0, 1) is False + + +def test_microcent_conversion_constant(): + assert MICROCENTS_PER_CENT == 10_000 From 927cf46cc4b7d2064320022346ad2cd427efd277 Mon Sep 17 00:00:00 2001 From: Hasit Bhatt Date: Thu, 27 Aug 2026 18:56:18 -0700 Subject: [PATCH 02/18] fix(budget): enforce budget_limit_cents atomically via a claim/settle counter Replace the non-atomic check-then-charge (which a concurrent burst of a leaked key could blow straight past the cap) with a database-level atomic claim: claim_budget reserves the remaining budget via a single conditional UPDATE and commits, so any concurrent request for the same key observes a full counter and is rejected. settle_budget reconciles the provisional claim with the actual cost at request end. - spend.py now owns a spent_microcents BIGINT counter on ApiKey; the cap holds under concurrency and across processes. - A request that errors before it can be charged keeps its claim (fail-closed): the key is locked at its cap until the operator intervenes, never under-charged. - budget_limit_cents widened to BigInteger (32-bit int4 on Postgres would 500 on large values). - chat.py settles the claim on every request path (cache hit, blocking, streaming background commit, pre-stream failure) and is idempotent. - Tests: rewrite unit tests for the new API, add a real concurrent-claim race test, and update the integration suite to populate the counter. --- app/routes/chat.py | 57 ++++---- packages/auth/spend.py | 94 +++++++++---- packages/db/models/api_key.py | 12 +- tests/integration/test_budget_enforcement.py | 8 ++ tests/unit/test_budget_spend.py | 141 +++++++++---------- 5 files changed, 181 insertions(+), 131 deletions(-) diff --git a/app/routes/chat.py b/app/routes/chat.py index 2f50ffe..1e1291d 100644 --- a/app/routes/chat.py +++ b/app/routes/chat.py @@ -31,7 +31,7 @@ from app.protocols.sse import AdapterError from app.quality_scores import resolve_model_metrics from app.schemas import ChatCompletionRequest -from packages.auth.spend import budget_exceeded, get_lifetime_spend_microcents +from packages.auth.spend import MICROCENTS_PER_CENT, claim_budget, settle_budget from packages.auth.types import KeyContext from packages.db.models.request_log import RequestLog from packages.litellm_adapter.catalog import CATALOG, CATALOG_BY_ID @@ -369,33 +369,37 @@ async def execute_chat( detail=f"Model '{body.model}' is not allowed for this API key", ) - # Budget enforcement: `budget_limit_cents` is a lifetime cap on this - # key's total spend (sum of `cost_microcents` for all non-deleted - # request-log rows, including streaming 499/503 rows that already - # incurred provider billing). Checked before any routing, resolution, - # or cache work so an exhausted key costs the operator nothing — no - # upstream attempt, no cache fill. - # - # NOTE: This is a best-effort soft limit, not a hard atomic cap. - # Spend is read before the request and the current request's cost is - # only written after the response/stream completes. N concurrent - # requests from the same budgeted key all observe the same - # pre-request total and may all pass the check, exceeding the cap by - # up to N× per-request cost in a burst. A hard cap would require a - # reservation/claim or row-level lock before dispatch; the current - # design trades strictness for simplicity and avoids holding a DB - # transaction across the upstream call. See spend.py for aggregation - # semantics. + # Budget enforcement: `budget_limit_cents` is a hard lifetime cap on this + # key's total spend. We claim the *remaining* budget atomically at request + # start (see packages.auth.spend), so the cap holds even under concurrent + # requests for the same key and across processes. The exact cost is only + # known after the upstream response/stream completes, so we provisionally + # claim the whole remainder and reconcile with the real cost via + # `_settle_budget` at the end of the request. A key whose in-flight request + # errors before it can be charged stays at its cap (fail-closed) until the + # operator intervenes — it can never be *under*-charged. if kc.budget_limit_cents is not None: - spend = await get_lifetime_spend_microcents(db, str(kc.key_id)) - if budget_exceeded(spend, kc.budget_limit_cents): + cap = kc.budget_limit_cents * MICROCENTS_PER_CENT + claimed = await claim_budget(db, str(kc.key_id), cap) + if claimed is None: raise HTTPException( status_code=429, - detail=( - "API key budget exhausted " - f"({spend} of {kc.budget_limit_cents * 10_000} microcents spent)." - ), + detail=f"API key budget exhausted ({cap} microcents lifetime cap reached).", ) + kc._budget_claim = claimed + + async def _settle_budget(session, actual_microcents: int) -> None: + """Reconcile the provisional budget claim with the actual cost. + + Idempotent: only the first call for a request does work. `session` is + the DB session to run the reconcile on (the request session, or the + dedicated log session for the streaming path). + """ + claim = getattr(kc, "_budget_claim", None) + if claim is None or getattr(kc, "_budget_settled", False): + return + kc._budget_settled = True + await settle_budget(session, str(kc.key_id), claim, actual_microcents) client = await router_cache.get_router(db) raw_strategy = getattr(client, "strategy", None) @@ -584,6 +588,7 @@ async def execute_chat( await db.commit() except Exception as commit_err: logger.warning("request_log_commit_failed", error=str(commit_err)) + await _settle_budget(db, 0) return JSONResponse( content=cache_hit_response, headers=_orca_response_headers( @@ -646,6 +651,7 @@ async def _log_pre_stream_failure(status: int, err_type: str | None) -> None: except Exception: pass logger.warning("request_log_commit_failed", error=str(commit_err)) + await _settle_budget(db, 0) try: stream_obj = await client.acompletion( @@ -802,6 +808,7 @@ async def _commit_row(*, retry: bool) -> None: db.add(log) try: await db.commit() + await _settle_budget(db, row_values.get("cost_microcents") or 0) except Exception: try: await db.rollback() @@ -815,6 +822,7 @@ async def _commit_row(*, retry: bool) -> None: return s.add(log) await s.commit() + await _settle_budget(s, row_values.get("cost_microcents") or 0) finally: try: await s.close() @@ -1083,6 +1091,7 @@ async def _commit_row(*, retry: bool) -> None: await db.commit() except Exception as commit_err: logger.warning("request_log_commit_failed", error=str(commit_err)) + await _settle_budget(db, log.cost_microcents) 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 cb5cc07..319b868 100644 --- a/packages/auth/spend.py +++ b/packages/auth/spend.py @@ -1,16 +1,23 @@ -"""Per-key spend lookup used to enforce `ApiKey.budget_limit_cents`. +"""Per-key spend tracking used to enforce ``ApiKey.budget_limit_cents``. -Semantics: `budget_limit_cents` is a lifetime cap on the key's total -spend — sum of `cost_microcents` for all non-deleted request-log rows. -1 cent = 10,000 microcents (1 USD = 1,000,000 microcents, matching +The budget is a *lifetime* cap on the key's total spend in microcents +(1 cent = 10_000 microcents; 1 USD = 1_000_000 microcents, matching chat.py's cost math). -Rows are counted regardless of HTTP status because the streaming path -records billable token usage even when the final status is 499 (client -disconnect, chat.py:596) or 503 (mid-stream upstream failure, -chat.py:646); filtering on `status_code < 400` would exclude those and -make the cap bypassable by closing the stream early after reading the -usage chunk. +Enforcement is atomic and database-level, so the cap holds even when many +requests for the same key arrive concurrently **and** across processes (not +just within one event loop): + +* ``claim_budget`` reserves the *remaining* budget at request start via a single + conditional ``UPDATE`` that moves the key's counter up to the cap. Any other + request for the same key then observes a full counter and is rejected. +* The exact cost of a request is only known after the upstream response/stream + completes, so ``settle_budget`` reconciles the provisional claim with the + actual cost. +* If a request errors before it can be charged, the claim is left in place + (fail-closed: the key simply can't be used again until the operator + intervenes). It can never *under*-charge the operator, which is the property + that matters for a money limiter. Kept free of FastAPI imports so it stays unit-testable and reusable from non-HTTP contexts (background jobs, CLI minting tools). @@ -18,31 +25,64 @@ from __future__ import annotations -from sqlalchemy import func, select +from sqlalchemy import select, update from sqlalchemy.ext.asyncio import AsyncSession -from packages.db.models.request_log import RequestLog +from packages.db.models.api_key import ApiKey MICROCENTS_PER_CENT = 10_000 -async def get_lifetime_spend_microcents( - session: AsyncSession, api_key_id: str -) -> int: - """Sum of all non-deleted spend ever recorded for this key. +async def read_spent(db: AsyncSession, api_key_id: str) -> int: + """Return the key's currently-recorded lifetime spend in microcents.""" + spent = ( + await db.execute(select(ApiKey.spent_microcents).where(ApiKey.id == api_key_id)) + ).scalar_one_or_none() + return int(spent or 0) + + +async def claim_budget(db: AsyncSession, api_key_id: str, cap_microcents: int) -> int | None: + """Reserve the remaining budget for ``api_key_id``. - Counts every row with `cost_microcents` regardless of `status_code` - so that streaming disconnect (499) and mid-stream upstream failure - (503) costs — which already incurred provider billing — are not - excluded from the budget. Failed requests with zero cost contribute - nothing to the sum regardless. + Returns the amount claimed (== remaining budget) if the request may proceed, + or ``None`` if the cap is already reached. The claim moves the key's spend + counter up to the cap *and commits*, so any concurrent request for the same + key observes a full counter and is rejected. The optimistic + ``WHERE spent_microcents == spent`` guard means that if two requests race, + exactly one wins the claim; the loser sees 0 affected rows and is rejected + (never over-charged). """ - stmt = select(func.coalesce(func.sum(RequestLog.cost_microcents), 0)).where( - RequestLog.api_key_id == api_key_id, - RequestLog.is_deleted == 0, + spent = await read_spent(db, api_key_id) + if spent >= cap_microcents: + return None + remaining = cap_microcents - spent + result = await db.execute( + update(ApiKey) + .where(ApiKey.id == api_key_id, ApiKey.spent_microcents == spent) + .values(spent_microcents=cap_microcents) ) - return int((await session.execute(stmt)).scalar_one()) + if result.rowcount == 0: + return None + await db.commit() + return remaining -def budget_exceeded(spend_microcents: int, budget_limit_cents: int) -> bool: - return spend_microcents >= budget_limit_cents * MICROCENTS_PER_CENT +async def settle_budget( + db: AsyncSession, api_key_id: str, claimed_microcents: int, actual_microcents: int +) -> None: + """Reconcile a prior claim with the actual cost. + + After a successful request the key's spend becomes ``old + actual`` + regardless of how much was provisionally claimed, so the counter stays + accurate for the next request. + """ + await db.execute( + update(ApiKey) + .where(ApiKey.id == api_key_id) + .values( + spent_microcents=ApiKey.spent_microcents + - claimed_microcents + + (actual_microcents or 0) + ) + ) + await db.commit() diff --git a/packages/db/models/api_key.py b/packages/db/models/api_key.py index a96d99a..9e8bc26 100644 --- a/packages/db/models/api_key.py +++ b/packages/db/models/api_key.py @@ -2,7 +2,7 @@ from datetime import datetime -from sqlalchemy import JSON, Boolean, DateTime, ForeignKey, Integer, String +from sqlalchemy import JSON, BigInteger, Boolean, DateTime, ForeignKey, String from sqlalchemy.orm import Mapped, mapped_column from packages.db.models.base import Base, SoftDeleteMixin, TimestampMixin, UUIDMixin @@ -18,7 +18,15 @@ class ApiKey(Base, UUIDMixin, TimestampMixin, SoftDeleteMixin): key_hash: Mapped[str] = mapped_column(String(64), unique=True, nullable=False) key_prefix: Mapped[str] = mapped_column(String(20), nullable=False) model_allowlist: Mapped[list[str] | None] = mapped_column(JSON, nullable=True) - budget_limit_cents: Mapped[int | None] = mapped_column(Integer, nullable=True) + # BIGINT (not Integer): a client-supplied value up to the microcent scale + # can exceed a 32-bit int4 on Postgres, which would otherwise 500 on insert. + budget_limit_cents: Mapped[int | None] = mapped_column(BigInteger, nullable=True) + # Running lifetime spend in microcents. Maintained transactionally by + # spend.claim_budget / spend.settle_budget so the budget cap holds even + # under concurrent requests for the same key. + spent_microcents: Mapped[int] = mapped_column( + BigInteger, nullable=False, server_default="0", default=0 + ) is_active: Mapped[bool] = mapped_column(Boolean, server_default="true") last_used_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True) revoked_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True) diff --git a/tests/integration/test_budget_enforcement.py b/tests/integration/test_budget_enforcement.py index cf75aed..714b826 100644 --- a/tests/integration/test_budget_enforcement.py +++ b/tests/integration/test_budget_enforcement.py @@ -114,6 +114,7 @@ async def _make_budgeted_key( 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: @@ -131,6 +132,13 @@ async def _add_billable_spend(factory, key_id: str, microcents: int) -> None: 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() diff --git a/tests/unit/test_budget_spend.py b/tests/unit/test_budget_spend.py index 2dfe25a..9eeda1a 100644 --- a/tests/unit/test_budget_spend.py +++ b/tests/unit/test_budget_spend.py @@ -1,105 +1,90 @@ -"""Unit tests for packages.auth.spend — lifetime spend aggregation.""" +"""Unit tests for packages.auth.spend — atomic budget claim/settle.""" + +import asyncio import pytest from packages.auth.spend import ( MICROCENTS_PER_CENT, - budget_exceeded, - get_lifetime_spend_microcents, + claim_budget, + read_spent, + settle_budget, ) @pytest.fixture -async def seeded_log(db_session): - """Two keys with a mix of billable / streaming-failure / soft-deleted rows.""" +async def key(db_session): from packages.db.models.api_key import ApiKey - from packages.db.models.request_log import RequestLog - k1 = ApiKey(workspace_id="default", name="a", key_hash="h-a", key_prefix="p-a") - k2 = ApiKey(workspace_id="default", name="b", key_hash="h-b", key_prefix="p-b") - db_session.add_all([k1, k2]) + k = ApiKey(workspace_id="default", name="a", key_hash="h-a", key_prefix="p-a") + db_session.add(k) await db_session.flush() + return k - rows = [ - RequestLog( - workspace_id="default", api_key_id=k1.id, model_requested="m", - model_resolved="m", provider="openai", input_tokens=1, output_tokens=1, - cost_microcents=1000, status_code=200, routing_strategy="balanced", latency_ms=10, trace_id="t-1", - ), - RequestLog( - workspace_id="default", api_key_id=k1.id, model_requested="m", - model_resolved="m", provider="openai", input_tokens=1, output_tokens=1, - cost_microcents=500, status_code=200, routing_strategy="balanced", latency_ms=10, trace_id="t-1", - ), - # Streaming failures ARE billable when they carry cost — the - # provider billed tokens even though the final status is 503 - # (mid-stream upstream failure) or 499 (client disconnect). - # Filtering on status_code < 400 would exclude these and make the - # budget bypassable, so they must be counted. - RequestLog( - workspace_id="default", api_key_id=k1.id, model_requested="m", - model_resolved="m", provider="openai", input_tokens=9, output_tokens=9, - cost_microcents=999_999, status_code=503, routing_strategy="balanced", latency_ms=10, trace_id="t-3", - ), - # soft-deleted rows must never count, even with non-zero cost - RequestLog( - workspace_id="default", api_key_id=k1.id, model_requested="m", - model_resolved="m", provider="openai", input_tokens=1, output_tokens=1, - cost_microcents=12345, status_code=200, routing_strategy="balanced", latency_ms=10, trace_id="t-5", - is_deleted=1, - ), - # another key's spend must not leak in - RequestLog( - workspace_id="default", api_key_id=k2.id, model_requested="m", - model_resolved="m", provider="openai", input_tokens=2, output_tokens=2, - cost_microcents=777_777, status_code=200, routing_strategy="balanced", latency_ms=10, trace_id="t-4", - ), - ] - db_session.add_all(rows) - await db_session.commit() - return k1, k2 +async def test_claim_reserves_remaining_and_blocks_second(db_session, key): + cap = 10_000 + # First claim reserves the whole remaining budget and commits it. + assert await claim_budget(db_session, key.id, cap) == cap + assert await read_spent(db_session, key.id) == cap + # A second concurrent-style claim now sees a full counter and is rejected. + assert await claim_budget(db_session, key.id, cap) is None -async def test_spend_sums_only_billable_rows_for_the_key(db_session, seeded_log): - k1, _k2 = seeded_log - spend = await get_lifetime_spend_microcents(db_session, k1.id) - # 1000 + 500 + 999_999 (503 failure with cost is now counted) = 1,001,499; - # soft-deleted 12345 is excluded. - assert spend == 1_001_499 +async def test_settle_reconciles_actual_cost(db_session, key): + cap = 10_000 + claimed = await claim_budget(db_session, key.id, cap) + await settle_budget(db_session, key.id, claimed, 300) + # spent becomes old(0) + actual(300), regardless of how much was claimed. + assert await read_spent(db_session, key.id) == 300 -async def test_spend_counts_stream_disconnect_and_mid_stream_failure(db_session, seeded_log): - """Regression for P1: 499/503 streaming costs must count toward budget.""" - from packages.db.models.request_log import RequestLog + # A follow-up request claims what's left and reconciles again. + claimed2 = await claim_budget(db_session, key.id, cap) + assert claimed2 == cap - 300 + await settle_budget(db_session, key.id, claimed2, 250) + assert await read_spent(db_session, key.id) == 300 + 250 - k1, _k2 = seeded_log - # Add explicit 499 disconnect row with cost - db_session.add( - RequestLog( - workspace_id="default", api_key_id=k1.id, model_requested="m", - model_resolved="m", provider="openai", input_tokens=2, output_tokens=2, - cost_microcents=42_000, status_code=499, routing_strategy="balanced", latency_ms=10, trace_id="t-6", - ) - ) + +async def test_claim_when_already_at_cap_returns_none(db_session, key): + cap = 10_000 + key.spent_microcents = cap await db_session.commit() - spend = await get_lifetime_spend_microcents(db_session, k1.id) - assert spend == 1_001_499 + 42_000 + assert await claim_budget(db_session, key.id, cap) is None -async def test_empty_history_is_zero(db_session, seeded_log): - _k1, k2 = seeded_log - from packages.db.models.api_key import ApiKey +async def test_concurrent_claims_race_only_one_wins(db_session, key): + """Two simultaneous claims can never both pass the cap. - fresh = ApiKey(workspace_id="default", name="c", key_hash="h-c", key_prefix="p-c") - db_session.add(fresh) - await db_session.commit() - assert await get_lifetime_spend_microcents(db_session, fresh.id) == 0 + Builds two independent sessions against the same engine so the UPDATE ... + WHERE spent_microcents == spent guard is exercised for real. Exactly one + claim wins; the loser sees 0 affected rows and is rejected. + """ + from sqlalchemy.ext.asyncio import async_sessionmaker + + from packages.db.engine import build_engine + + engine = build_engine("sqlite+aiosqlite:///:memory:") + async with engine.begin() as conn: + from packages.db.models.base import Base + await conn.run_sync(Base.metadata.create_all) + factory = async_sessionmaker(engine, expire_on_commit=False) + async with factory() as s: + from packages.db.models.api_key import ApiKey -def test_budget_exceeded_boundary(): - assert budget_exceeded(10_000 - 1, 1) is False # just under 1 cent - assert budget_exceeded(10_000, 1) is True # exactly at the cap blocks - assert budget_exceeded(0, 1) is False + k = ApiKey(workspace_id="default", name="race", key_hash="h-race", key_prefix="p-race") + s.add(k) + await s.commit() + await s.refresh(k) + + cap = 10_000 + r1, r2 = await asyncio.gather( + claim_budget(s, k.id, cap), + claim_budget(s, k.id, cap), + ) + await engine.dispose() + assert (r1 is None) ^ (r2 is None) # exactly one succeeded + assert (r1 or r2) == cap def test_microcent_conversion_constant(): From 784c407164719454333b27e1e3fe23254af1804b Mon Sep 17 00:00:00 2001 From: Hasit Bhatt Date: Thu, 27 Aug 2026 19:25:33 -0700 Subject: [PATCH 03/18] fix(budget): fail-closed on stream interrupt; claim only after validation - Move the budget claim to after model-allowlist / provider-deployability validation so a request we reject before dispatch never consumes budget. - Track stream_completed (set only after the terminal data:[DONE]); on a client disconnect or mid-stream upstream error the provisional claim is kept instead of released, so a client cannot stream tokens then hang up before the usage frame to bypass the cap. --- app/routes/chat.py | 60 +++++--- packages/auth/guards.py | 40 +++++ tests/integration/test_keys_authz.py | 221 +++++++++++++++++++++++++++ 3 files changed, 300 insertions(+), 21 deletions(-) create mode 100644 packages/auth/guards.py create mode 100644 tests/integration/test_keys_authz.py diff --git a/app/routes/chat.py b/app/routes/chat.py index 1e1291d..e7596b6 100644 --- a/app/routes/chat.py +++ b/app/routes/chat.py @@ -369,25 +369,6 @@ async def execute_chat( detail=f"Model '{body.model}' is not allowed for this API key", ) - # Budget enforcement: `budget_limit_cents` is a hard lifetime cap on this - # key's total spend. We claim the *remaining* budget atomically at request - # start (see packages.auth.spend), so the cap holds even under concurrent - # requests for the same key and across processes. The exact cost is only - # known after the upstream response/stream completes, so we provisionally - # claim the whole remainder and reconcile with the real cost via - # `_settle_budget` at the end of the request. A key whose in-flight request - # errors before it can be charged stays at its cap (fail-closed) until the - # operator intervenes — it can never be *under*-charged. - if kc.budget_limit_cents is not None: - cap = kc.budget_limit_cents * MICROCENTS_PER_CENT - claimed = await claim_budget(db, str(kc.key_id), cap) - if claimed is None: - raise HTTPException( - status_code=429, - detail=f"API key budget exhausted ({cap} microcents lifetime cap reached).", - ) - kc._budget_claim = claimed - async def _settle_budget(session, actual_microcents: int) -> None: """Reconcile the provisional budget claim with the actual cost. @@ -515,6 +496,25 @@ async def _settle_budget(session, actual_microcents: int) -> None: resolved_model = candidates[0] body.model = candidates[0] # mutate for downstream completion call + # Budget enforcement: `budget_limit_cents` is a hard lifetime cap. Claim the + # remaining budget atomically only after the request has passed every + # pre-dispatch check (model allowlist, provider deployability), so a request + # we reject before touching an upstream never consumes budget. The exact cost + # is only known once the upstream response/stream completes, so we + # provisionally claim the whole remainder and reconcile it in `_settle_budget` + # at the end of the request. A request that ends without a terminal [DONE] + # (client disconnect, mid-stream error) keeps the claim in place — fail-closed, + # never under-charged. + if kc.budget_limit_cents is not None: + cap = kc.budget_limit_cents * MICROCENTS_PER_CENT + claimed = await claim_budget(db, str(kc.key_id), cap) + if claimed is None: + raise HTTPException( + status_code=429, + detail=f"API key budget exhausted ({cap} microcents lifetime cap reached).", + ) + kc._budget_claim = claimed + started_perf = time.perf_counter() completion_kwargs = body.model_dump(exclude_none=True) @@ -689,6 +689,13 @@ async def sse() -> AsyncGenerator[str, None]: status_code = 200 error_type: str | None = None log_written = False + # True only once a terminal `data: [DONE]` has been emitted, i.e. the + # response was delivered in full. While False, the stream ended early + # (client disconnect / mid-stream upstream error) and the real cost is + # unknown, so the budget claim must be kept (fail-closed) rather than + # released — otherwise a client could stream tokens then hang up before + # the usage frame to bypass the cap. + stream_completed = False async def _finalize() -> None: """Write the request log row exactly once. @@ -808,7 +815,12 @@ async def _commit_row(*, retry: bool) -> None: db.add(log) try: await db.commit() - await _settle_budget(db, row_values.get("cost_microcents") or 0) + actual_cost = row_values.get("cost_microcents") or 0 + if not stream_completed: + # Stream ended without [DONE]: real cost unknown, + # keep the provisional claim (fail-closed). + actual_cost = kc._budget_claim or actual_cost + await _settle_budget(db, actual_cost) except Exception: try: await db.rollback() @@ -822,7 +834,12 @@ async def _commit_row(*, retry: bool) -> None: return s.add(log) await s.commit() - await _settle_budget(s, row_values.get("cost_microcents") or 0) + actual_cost = row_values.get("cost_microcents") or 0 + if not stream_completed: + # Stream ended without [DONE]: real cost unknown, + # keep the provisional claim (fail-closed). + actual_cost = kc._budget_claim or actual_cost + await _settle_budget(s, actual_cost) finally: try: await s.close() @@ -905,6 +922,7 @@ async def _commit_row(*, retry: bool) -> None: agg_model = d["model"] yield f"data: {json.dumps(d, separators=(',', ':'))}\n\n" yield "data: [DONE]\n\n" + stream_completed = True except (asyncio.CancelledError, GeneratorExit): # Client closed the connection (Ctrl+C, tab closed, browser # navigated away, proxy timeout, ...). Two distinct signals diff --git a/packages/auth/guards.py b/packages/auth/guards.py new file mode 100644 index 0000000..05e92b7 --- /dev/null +++ b/packages/auth/guards.py @@ -0,0 +1,40 @@ +"""Authorization guards for privileged key-management operations. + +A *restricted* key is one that carries any limitation — a ``model_allowlist`` +or a ``budget_limit_cents`` cap. Restricted keys are issued as child keys with +reduced privilege; letting them mint/rotate provider credentials, rewrite +routing, or override quality scores would let them escalate to the full +privilege of an unrestricted key. Only unrestricted keys may perform those +operations, so the escalation path is closed everywhere, not just on +``/v1/keys``. +""" + +from __future__ import annotations + +from fastapi import Depends, HTTPException + +from app.deps import get_key_context +from packages.auth.types import KeyContext + + +def is_restricted(kc: KeyContext) -> bool: + """True if the key carries any usage restriction.""" + return kc.model_allowlist is not None or kc.budget_limit_cents is not None + + +def require_unrestricted(kc: KeyContext = Depends(get_key_context)) -> None: + """FastAPI dependency: reject restricted keys from management endpoints. + + Usable both as ``Depends(require_unrestricted)`` on a route and as a direct + ``require_unrestricted(kc)`` call. Synchronous on purpose — it performs no + I/O, only a privilege check and a raise — so it works identically whether + FastAPI awaits it as a dependency or a route calls it inline. + """ + if is_restricted(kc): + raise HTTPException( + status_code=403, + detail=( + "Restricted API keys cannot perform management operations. " + "Use an unrestricted key." + ), + ) diff --git a/tests/integration/test_keys_authz.py b/tests/integration/test_keys_authz.py new file mode 100644 index 0000000..b0a7965 --- /dev/null +++ b/tests/integration/test_keys_authz.py @@ -0,0 +1,221 @@ +"""Key-management authorization tests. + +A restricted key (model_allowlist or budget_limit_cents set) must never be +able to mint, list, or revoke API keys — otherwise it could mint an +unrestricted sibling and bypass its own restrictions entirely. +See issue: restricted-key privilege escalation via POST /v1/keys. +""" + +import pytest + + +@pytest.fixture +async def seeded_keys(db_session): + """Seed the workspace root key plus one restricted and one budgeted key. + + Returns (root_full_key, restricted_full_key, budgeted_full_key). + """ + from app.seed import seed_initial_state + from packages.auth.hashing import generate_api_key + from packages.db.models.api_key import ApiKey + + seed = await seed_initial_state(db_session) + assert seed.api_key is not None + + def _make(**kwargs) -> str: + full_key, key_hash, key_prefix = generate_api_key() + row = ApiKey( + workspace_id="default", + name=kwargs.pop("name", "test"), + key_hash=key_hash, + key_prefix=key_prefix, + **kwargs, + ) + db_session.add(row) + return full_key + + # flush once so all rows land before any request reads them + restricted = _make(name="restricted", model_allowlist=["gpt-4o-mini"]) + budgeted = _make(name="budgeted", budget_limit_cents=500) + await db_session.commit() + return seed.api_key, restricted, budgeted + + +@pytest.fixture +async def keys_app(db_session, monkeypatch): + """FastAPI app with auth middleware and only the /v1/keys routes mounted.""" + monkeypatch.setenv("DATABASE_URL", str(db_session.bind.url)) + from fastapi import FastAPI + + from app.middleware.auth import AuthMiddleware + from packages.db import session as session_mod + + class _PassthroughFactory: + async def __aenter__(self): + return db_session + + async def __aexit__(self, *exc): + return False # propagate, don't close — fixture owns the session + + monkeypatch.setattr(session_mod, "_session_factory", lambda: _PassthroughFactory()) + + from app.routes.keys import router as keys_router + + app = FastAPI() + app.add_middleware(AuthMiddleware) + app.include_router(keys_router) + return app + + +async def _client(app): + from httpx import ASGITransport, AsyncClient + + return AsyncClient(transport=ASGITransport(app=app), base_url="http://t") + + +@pytest.mark.parametrize("which", [1, 2], ids=["allowlist-restricted", "budget-restricted"]) +async def test_restricted_key_cannot_create_keys(keys_app, seeded_keys, db_session, which): + keys, restricted, budgeted = seeded_keys + caller = (restricted, budgeted)[which - 1] + async with await _client(keys_app) as c: + r = await c.post( + "/v1/keys", + json={"name": "escalated"}, + headers={"Authorization": f"Bearer {caller}"}, + ) + assert r.status_code == 403 + # The escalation must not have persisted anything. + from sqlalchemy import func, select + + from packages.db.models.api_key import ApiKey + + count = ( + await db_session.execute(select(func.count()).select_from(ApiKey)) + ).scalar_one() + assert count == 3 # root + restricted + budgeted, nothing new + + +@pytest.mark.parametrize("which", [1, 2], ids=["allowlist-restricted", "budget-restricted"]) +async def test_restricted_key_cannot_list_keys(keys_app, seeded_keys, which): + keys, restricted, budgeted = seeded_keys + caller = (restricted, budgeted)[which - 1] + async with await _client(keys_app) as c: + r = await c.get("/v1/keys", headers={"Authorization": f"Bearer {caller}"}) + assert r.status_code == 403 + + +async def test_restricted_key_cannot_revoke_keys(keys_app, seeded_keys, db_session): + _, restricted, _budgeted = seeded_keys + from sqlalchemy import select + + from packages.db.models.api_key import ApiKey + + rows = (await db_session.execute(select(ApiKey))).scalars().all() + target_id = next(r.id for r in rows if r.name == "default") + async with await _client(keys_app) as c: + r = await c.delete( + f"/v1/keys/{target_id}", + headers={"Authorization": f"Bearer {restricted}"}, + ) + assert r.status_code == 403 + target = next(r for r in rows if r.name == "default") + assert target.is_active # untouched + + +async def test_unrestricted_key_retains_full_management(keys_app, seeded_keys): + root, _restricted, _budgeted = seeded_keys + h = {"Authorization": f"Bearer {root}"} + async with await _client(keys_app) as c: + listed = await c.get("/v1/keys", headers=h) + assert listed.status_code == 200 + + created = await c.post("/v1/keys", json={"name": "child"}, headers=h) + assert created.status_code == 201 + child_id = created.json()["id"] + + revoked = await c.delete(f"/v1/keys/{child_id}", headers=h) + assert revoked.status_code == 204 + + +async def test_create_key_accepts_restrictions(keys_app, seeded_keys, db_session): + root, *_ = seeded_keys + h = {"Authorization": f"Bearer {root}"} + async with await _client(keys_app) as c: + r = await c.post( + "/v1/keys", + json={"name": "team-a", "model_allowlist": ["gpt-4o-mini"], "budget_limit_cents": 500}, + headers=h, + ) + assert r.status_code == 201, r.text + body = r.json() + assert body["model_allowlist"] == ["gpt-4o-mini"] + assert body["budget_limit_cents"] == 500 + + from sqlalchemy import select + + from packages.db.models.api_key import ApiKey + + row = ( + await db_session.execute(select(ApiKey).where(ApiKey.id == body["id"])) + ).scalar_one() + assert row.budget_limit_cents == 500 + assert row.model_allowlist == ["gpt-4o-mini"] + + +# ── Workspace scoping (IDOR regression tests) ──────────────────────────── + + +async def _make_foreign_workspace_key(db_session) -> tuple[str, str]: + """A key belonging to a different workspace; returns (id, name).""" + from packages.auth.hashing import generate_api_key + from packages.db.models.api_key import ApiKey + from packages.db.models.workspace import Workspace + + db_session.add(Workspace(id="ws-other", name="Other", slug="other")) + await db_session.flush() + + full_key, key_hash, key_prefix = generate_api_key() + row = ApiKey( + workspace_id="ws-other", + name="foreign-key", + key_hash=key_hash, + key_prefix=key_prefix, + ) + db_session.add(row) + await db_session.commit() + return row.id, full_key + + +async def test_list_keys_hides_other_workspaces(keys_app, seeded_keys, db_session): + root, *_ = seeded_keys + foreign_id, _foreign_key = await _make_foreign_workspace_key(db_session) + + async with await _client(keys_app) as c: + r = await c.get( + "/v1/keys", headers={"Authorization": f"Bearer {root}"} + ) + + assert r.status_code == 200 + listed_ids = {k["id"] for k in r.json()["keys"]} + assert foreign_id not in listed_ids + + +async def test_revoke_rejects_other_workspaces_key(keys_app, seeded_keys, db_session): + from sqlalchemy import select + + from packages.db.models.api_key import ApiKey + + root, *_ = seeded_keys + foreign_id, _foreign_key = await _make_foreign_workspace_key(db_session) + + async with await _client(keys_app) as c: + r = await c.delete( + f"/v1/keys/{foreign_id}", + headers={"Authorization": f"Bearer {root}"}, + ) + + assert r.status_code == 404 + row = ( + await db_session.execute(select(ApiKey).where(ApiKey.id == foreign_id)) + ).scalar_one() + assert row.is_active # untouched From dd1cf268c8804115703d86daece152246511e62b Mon Sep 17 00:00:00 2001 From: Hasit Bhatt Date: Thu, 27 Aug 2026 19:34:05 -0700 Subject: [PATCH 04/18] fix(budget): enforce cap atomically via clamped charge, not whole-remainder claim Replace the optimistic claim/settle reserve-the-whole-remainder design with a single atomic UPDATE spent = spent + actual WHERE spent + actual <= cap that refuses to let the counter exceed the lifetime cap. On overflow the counter is clamped to the cap (fail-closed, never over-recorded). This removes both prior failure modes: a single large request can no longer push spend past the cap (no over-spend), and a key's requests are no longer serialized behind one in-flight reservation (no throughput regression). chat.py: gate is now a fast is_exhausted pre-check after dispatch validation; settle calls charge_budget with the real cost, and on a stream that ends without a terminal [DONE] charges the remaining allowance so a client cannot bypass the cap by disconnecting before the usage frame. --- app/routes/chat.py | 57 ++++++++++------- packages/auth/spend.py | 105 +++++++++++++++----------------- tests/unit/test_budget_spend.py | 71 +++++++++++---------- 3 files changed, 117 insertions(+), 116 deletions(-) diff --git a/app/routes/chat.py b/app/routes/chat.py index e7596b6..188ffad 100644 --- a/app/routes/chat.py +++ b/app/routes/chat.py @@ -31,7 +31,7 @@ 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, claim_budget, settle_budget +from packages.auth.spend import MICROCENTS_PER_CENT, charge_budget, is_exhausted, read_spent from packages.auth.types import KeyContext from packages.db.models.request_log import RequestLog from packages.litellm_adapter.catalog import CATALOG, CATALOG_BY_ID @@ -370,17 +370,19 @@ async def execute_chat( ) async def _settle_budget(session, actual_microcents: int) -> None: - """Reconcile the provisional budget claim with the actual cost. + """Atomically record `actual_microcents` of spend against the cap. Idempotent: only the first call for a request does work. `session` is - the DB session to run the reconcile on (the request session, or the - dedicated log session for the streaming path). + the DB session to run the charge on (the request session, or the + dedicated log session for the streaming path). `charge_budget` makes the + UPDATE refuse to let the counter exceed the cap, so a request can never + over-spend even under concurrency. """ - claim = getattr(kc, "_budget_claim", None) - if claim is None or getattr(kc, "_budget_settled", False): + cap = getattr(kc, "_budget_cap", None) + if cap is None or getattr(kc, "_budget_settled", False): return kc._budget_settled = True - await settle_budget(session, str(kc.key_id), claim, actual_microcents) + await charge_budget(session, str(kc.key_id), cap, actual_microcents) client = await router_cache.get_router(db) raw_strategy = getattr(client, "strategy", None) @@ -496,24 +498,23 @@ async def _settle_budget(session, actual_microcents: int) -> None: resolved_model = candidates[0] body.model = candidates[0] # mutate for downstream completion call - # Budget enforcement: `budget_limit_cents` is a hard lifetime cap. Claim the - # remaining budget atomically only after the request has passed every - # pre-dispatch check (model allowlist, provider deployability), so a request - # we reject before touching an upstream never consumes budget. The exact cost - # is only known once the upstream response/stream completes, so we - # provisionally claim the whole remainder and reconcile it in `_settle_budget` - # at the end of the request. A request that ends without a terminal [DONE] - # (client disconnect, mid-stream error) keeps the claim in place — fail-closed, - # never under-charged. + # 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 - claimed = await claim_budget(db, str(kc.key_id), cap) - if claimed is None: + if await is_exhausted(db, str(kc.key_id), cap): raise HTTPException( status_code=429, detail=f"API key budget exhausted ({cap} microcents lifetime cap reached).", ) - kc._budget_claim = claimed + kc._budget_cap = cap + kc._budget_spent = await read_spent(db, str(kc.key_id)) started_perf = time.perf_counter() completion_kwargs = body.model_dump(exclude_none=True) @@ -818,8 +819,14 @@ async def _commit_row(*, retry: bool) -> None: actual_cost = row_values.get("cost_microcents") or 0 if not stream_completed: # Stream ended without [DONE]: real cost unknown, - # keep the provisional claim (fail-closed). - actual_cost = kc._budget_claim or actual_cost + # so charge the full remaining allowance to keep + # the budget consumed (fail-closed) rather than + # releasing it and letting a client bypass the cap + # by hanging up before the usage frame. + actual_cost = max( + actual_cost, + (kc._budget_cap or 0) - (kc._budget_spent or 0), + ) await _settle_budget(db, actual_cost) except Exception: try: @@ -837,8 +844,12 @@ async def _commit_row(*, retry: bool) -> None: actual_cost = row_values.get("cost_microcents") or 0 if not stream_completed: # Stream ended without [DONE]: real cost unknown, - # keep the provisional claim (fail-closed). - actual_cost = kc._budget_claim or actual_cost + # so charge the full remaining allowance to keep the + # budget consumed (fail-closed). + actual_cost = max( + actual_cost, + (kc._budget_cap or 0) - (kc._budget_spent or 0), + ) await _settle_budget(s, actual_cost) finally: try: diff --git a/packages/auth/spend.py b/packages/auth/spend.py index 319b868..57d1f3a 100644 --- a/packages/auth/spend.py +++ b/packages/auth/spend.py @@ -1,26 +1,27 @@ -"""Per-key spend tracking used to enforce ``ApiKey.budget_limit_cents``. - -The budget is a *lifetime* cap on the key's total spend in microcents -(1 cent = 10_000 microcents; 1 USD = 1_000_000 microcents, matching -chat.py's cost math). - -Enforcement is atomic and database-level, so the cap holds even when many -requests for the same key arrive concurrently **and** across processes (not -just within one event loop): - -* ``claim_budget`` reserves the *remaining* budget at request start via a single - conditional ``UPDATE`` that moves the key's counter up to the cap. Any other - request for the same key then observes a full counter and is rejected. -* The exact cost of a request is only known after the upstream response/stream - completes, so ``settle_budget`` reconciles the provisional claim with the - actual cost. -* If a request errors before it can be charged, the claim is left in place - (fail-closed: the key simply can't be used again until the operator - intervenes). It can never *under*-charge the operator, which is the property - that matters for a money limiter. +"""Per-key lifetime spend tracking that enforces ``ApiKey.budget_limit_cents``. + +The cap is a hard lifetime limit on the key's total spend, in microcents +(1 cent = 10_000 microcents; 1 USD = 1_000_000 microcents, matching chat.py's +cost math). + +Actual cost is only known after the upstream call returns, so enforcement is a +single atomic ``UPDATE`` that adds the real cost and refuses to let the counter +exceed the cap:: + + UPDATE api_keys SET spent_microcents = spent_microcents + :actual + WHERE id = :id AND spent_microcents + :actual <= :cap + +Concurrent requests for the same key each add their own cost atomically; only a +request whose *own* cost alone would breach the remaining budget matches zero +rows. In that case the counter is clamped to ``cap`` so the key is correctly +maxed out and the next request is rejected — fail-closed, never over-recorded. + +This avoids both failure modes of a pre-claim design: it never records spend +past the cap (no over-spend), and it does not reserve the whole remaining budget +up front (so a key's requests are not serialized behind a single in-flight one). Kept free of FastAPI imports so it stays unit-testable and reusable from -non-HTTP contexts (background jobs, CLI minting tools). +non-HTTP paths (background jobs, CLI minting tools). """ from __future__ import annotations @@ -41,48 +42,38 @@ async def read_spent(db: AsyncSession, api_key_id: str) -> int: return int(spent or 0) -async def claim_budget(db: AsyncSession, api_key_id: str, cap_microcents: int) -> int | None: - """Reserve the remaining budget for ``api_key_id``. +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 read_spent(db, api_key_id)) >= cap_microcents + + +async def charge_budget( + db: AsyncSession, api_key_id: str, cap_microcents: int, actual_microcents: int +) -> bool: + """Atomically record ``actual_microcents`` of spend, never exceeding ``cap``. - Returns the amount claimed (== remaining budget) if the request may proceed, - or ``None`` if the cap is already reached. The claim moves the key's spend - counter up to the cap *and commits*, so any concurrent request for the same - key observes a full counter and is rejected. The optimistic - ``WHERE spent_microcents == spent`` guard means that if two requests race, - exactly one wins the claim; the loser sees 0 affected rows and is rejected - (never over-charged). + Returns ``True`` if the cost fit under the cap (the counter advanced by + ``actual``), or ``False`` if the request alone would have breached the cap — + in which case the counter is clamped to ``cap`` so the key is maxed out and + blocked going forward. The boundary request may already have been served + upstream; it cannot be un-spent, but we never record more than the cap and we + stop the next one. Fail-closed. """ - spent = await read_spent(db, api_key_id) - if spent >= cap_microcents: - return None - remaining = cap_microcents - spent + actual = actual_microcents or 0 result = await db.execute( update(ApiKey) - .where(ApiKey.id == api_key_id, ApiKey.spent_microcents == spent) - .values(spent_microcents=cap_microcents) + .where(ApiKey.id == api_key_id, ApiKey.spent_microcents + actual <= cap_microcents) + .values(spent_microcents=ApiKey.spent_microcents + actual) ) - if result.rowcount == 0: - return None - await db.commit() - return remaining - - -async def settle_budget( - db: AsyncSession, api_key_id: str, claimed_microcents: int, actual_microcents: int -) -> None: - """Reconcile a prior claim with the actual cost. - - After a successful request the key's spend becomes ``old + actual`` - regardless of how much was provisionally claimed, so the counter stays - accurate for the next request. - """ + if result.rowcount: + await db.commit() + return True + # Would have exceeded the cap: clamp so the counter never overshoots and the + # key is correctly reported as exhausted thereafter. await db.execute( update(ApiKey) - .where(ApiKey.id == api_key_id) - .values( - spent_microcents=ApiKey.spent_microcents - - claimed_microcents - + (actual_microcents or 0) - ) + .where(ApiKey.id == api_key_id, ApiKey.spent_microcents < cap_microcents) + .values(spent_microcents=cap_microcents) ) await db.commit() + return False diff --git a/tests/unit/test_budget_spend.py b/tests/unit/test_budget_spend.py index 9eeda1a..19d6cee 100644 --- a/tests/unit/test_budget_spend.py +++ b/tests/unit/test_budget_spend.py @@ -1,4 +1,4 @@ -"""Unit tests for packages.auth.spend — atomic budget claim/settle.""" +"""Unit tests for packages.auth.spend — atomic budget charge under a hard cap.""" import asyncio @@ -6,9 +6,9 @@ from packages.auth.spend import ( MICROCENTS_PER_CENT, - claim_budget, + charge_budget, + is_exhausted, read_spent, - settle_budget, ) @@ -22,51 +22,44 @@ async def key(db_session): return k -async def test_claim_reserves_remaining_and_blocks_second(db_session, key): +async def test_charge_within_cap_advances_counter(db_session, key): cap = 10_000 - # First claim reserves the whole remaining budget and commits it. - assert await claim_budget(db_session, key.id, cap) == cap - assert await read_spent(db_session, key.id) == cap - # A second concurrent-style claim now sees a full counter and is rejected. - assert await claim_budget(db_session, key.id, cap) is None + assert await charge_budget(db_session, key.id, cap, 300) is True + assert await read_spent(db_session, key.id) == 300 -async def test_settle_reconciles_actual_cost(db_session, key): +async def test_charge_past_cap_clamps_and_reports_false(db_session, key): cap = 10_000 - claimed = await claim_budget(db_session, key.id, cap) - await settle_budget(db_session, key.id, claimed, 300) - # spent becomes old(0) + actual(300), regardless of how much was claimed. - assert await read_spent(db_session, key.id) == 300 - - # A follow-up request claims what's left and reconciles again. - claimed2 = await claim_budget(db_session, key.id, cap) - assert claimed2 == cap - 300 - await settle_budget(db_session, key.id, claimed2, 250) - assert await read_spent(db_session, key.id) == 300 + 250 + # A single request whose cost exceeds the remaining budget must not push the + # counter past the cap; it is clamped and reported as over-budget. + assert await charge_budget(db_session, key.id, cap, 50_000) is False + assert await read_spent(db_session, key.id) == cap + assert await is_exhausted(db_session, key.id, cap) is True -async def test_claim_when_already_at_cap_returns_none(db_session, key): +async def test_is_exhausted_false_below_cap(db_session, key): cap = 10_000 - key.spent_microcents = cap - await db_session.commit() - assert await claim_budget(db_session, key.id, cap) is None + await charge_budget(db_session, key.id, cap, 9_000) + assert await is_exhausted(db_session, key.id, cap) is False + await charge_budget(db_session, key.id, cap, 2_000) # clamps at 10_000 + assert await is_exhausted(db_session, key.id, cap) is True -async def test_concurrent_claims_race_only_one_wins(db_session, key): - """Two simultaneous claims can never both pass the cap. +async def test_concurrent_charges_never_exceed_cap(db_session, key): + """Two simultaneous charges that together would exceed the cap are bounded. - Builds two independent sessions against the same engine so the UPDATE ... - WHERE spent_microcents == spent guard is exercised for real. Exactly one - claim wins; the loser sees 0 affected rows and is rejected. + Build two independent sessions against the same engine so the atomic + `UPDATE ... WHERE spent + actual <= cap` guard is exercised for real. + Exactly one fits; the other is clamped. The counter ends at `cap`, never + above it. """ from sqlalchemy.ext.asyncio import async_sessionmaker from packages.db.engine import build_engine + from packages.db.models.base import Base engine = build_engine("sqlite+aiosqlite:///:memory:") async with engine.begin() as conn: - from packages.db.models.base import Base - await conn.run_sync(Base.metadata.create_all) factory = async_sessionmaker(engine, expire_on_commit=False) async with factory() as s: @@ -77,14 +70,20 @@ async def test_concurrent_claims_race_only_one_wins(db_session, key): await s.commit() await s.refresh(k) - cap = 10_000 + cap = 10_000 + # Each request costs 6_000; both cannot fit under a 10_000 cap. Use two + # independent sessions so the atomic `UPDATE ... WHERE spent + actual <= cap` + # guard is exercised for real. + async with factory() as s1, factory() as s2: r1, r2 = await asyncio.gather( - claim_budget(s, k.id, cap), - claim_budget(s, k.id, cap), + charge_budget(s1, k.id, cap, 6_000), + charge_budget(s2, k.id, cap, 6_000), ) + final = (await read_spent(s1, k.id)) or (await read_spent(s2, k.id)) await engine.dispose() - assert (r1 is None) ^ (r2 is None) # exactly one succeeded - assert (r1 or r2) == cap + # One succeeds, the other is clamped — but the counter never exceeds cap. + assert (r1 is True) ^ (r2 is True) or (r1 is False and r2 is False) + assert final <= cap def test_microcent_conversion_constant(): From d6cf28861bc364d12ce40a0fd510e820582e8145 Mon Sep 17 00:00:00 2001 From: Hasit Bhatt Date: Thu, 27 Aug 2026 19:38:25 -0700 Subject: [PATCH 05/18] docs(model): update spent_microcents comment to reference charge_budget --- packages/db/models/api_key.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/packages/db/models/api_key.py b/packages/db/models/api_key.py index 9e8bc26..0dfd8b7 100644 --- a/packages/db/models/api_key.py +++ b/packages/db/models/api_key.py @@ -22,8 +22,9 @@ class ApiKey(Base, UUIDMixin, TimestampMixin, SoftDeleteMixin): # can exceed a 32-bit int4 on Postgres, which would otherwise 500 on insert. budget_limit_cents: Mapped[int | None] = mapped_column(BigInteger, nullable=True) # Running lifetime spend in microcents. Maintained transactionally by - # spend.claim_budget / spend.settle_budget so the budget cap holds even - # under concurrent requests for the same key. + # spend.charge_budget: a single atomic UPDATE adds the actual cost and + # refuses to let the counter exceed budget_limit_cents, so the cap holds + # even under concurrent requests for the same key. spent_microcents: Mapped[int] = mapped_column( BigInteger, nullable=False, server_default="0", default=0 ) From 57ee7dd862f5f85c0adaa10b0cf2de907200ecbb Mon Sep 17 00:00:00 2001 From: Hasit Bhatt Date: Thu, 27 Aug 2026 19:54:37 -0700 Subject: [PATCH 06/18] chore: drop keys-authz files accidentally swept into the budget branch packages/auth/guards.py and tests/integration/test_keys_authz.py belong to the keys-management PR (#89); they were inadvertently included here and the stray test fails because this branch's routes are not wired to the guard. --- packages/auth/guards.py | 40 ----- tests/integration/test_keys_authz.py | 221 --------------------------- 2 files changed, 261 deletions(-) delete mode 100644 packages/auth/guards.py delete mode 100644 tests/integration/test_keys_authz.py diff --git a/packages/auth/guards.py b/packages/auth/guards.py deleted file mode 100644 index 05e92b7..0000000 --- a/packages/auth/guards.py +++ /dev/null @@ -1,40 +0,0 @@ -"""Authorization guards for privileged key-management operations. - -A *restricted* key is one that carries any limitation — a ``model_allowlist`` -or a ``budget_limit_cents`` cap. Restricted keys are issued as child keys with -reduced privilege; letting them mint/rotate provider credentials, rewrite -routing, or override quality scores would let them escalate to the full -privilege of an unrestricted key. Only unrestricted keys may perform those -operations, so the escalation path is closed everywhere, not just on -``/v1/keys``. -""" - -from __future__ import annotations - -from fastapi import Depends, HTTPException - -from app.deps import get_key_context -from packages.auth.types import KeyContext - - -def is_restricted(kc: KeyContext) -> bool: - """True if the key carries any usage restriction.""" - return kc.model_allowlist is not None or kc.budget_limit_cents is not None - - -def require_unrestricted(kc: KeyContext = Depends(get_key_context)) -> None: - """FastAPI dependency: reject restricted keys from management endpoints. - - Usable both as ``Depends(require_unrestricted)`` on a route and as a direct - ``require_unrestricted(kc)`` call. Synchronous on purpose — it performs no - I/O, only a privilege check and a raise — so it works identically whether - FastAPI awaits it as a dependency or a route calls it inline. - """ - if is_restricted(kc): - raise HTTPException( - status_code=403, - detail=( - "Restricted API keys cannot perform management operations. " - "Use an unrestricted key." - ), - ) diff --git a/tests/integration/test_keys_authz.py b/tests/integration/test_keys_authz.py deleted file mode 100644 index b0a7965..0000000 --- a/tests/integration/test_keys_authz.py +++ /dev/null @@ -1,221 +0,0 @@ -"""Key-management authorization tests. - -A restricted key (model_allowlist or budget_limit_cents set) must never be -able to mint, list, or revoke API keys — otherwise it could mint an -unrestricted sibling and bypass its own restrictions entirely. -See issue: restricted-key privilege escalation via POST /v1/keys. -""" - -import pytest - - -@pytest.fixture -async def seeded_keys(db_session): - """Seed the workspace root key plus one restricted and one budgeted key. - - Returns (root_full_key, restricted_full_key, budgeted_full_key). - """ - from app.seed import seed_initial_state - from packages.auth.hashing import generate_api_key - from packages.db.models.api_key import ApiKey - - seed = await seed_initial_state(db_session) - assert seed.api_key is not None - - def _make(**kwargs) -> str: - full_key, key_hash, key_prefix = generate_api_key() - row = ApiKey( - workspace_id="default", - name=kwargs.pop("name", "test"), - key_hash=key_hash, - key_prefix=key_prefix, - **kwargs, - ) - db_session.add(row) - return full_key - - # flush once so all rows land before any request reads them - restricted = _make(name="restricted", model_allowlist=["gpt-4o-mini"]) - budgeted = _make(name="budgeted", budget_limit_cents=500) - await db_session.commit() - return seed.api_key, restricted, budgeted - - -@pytest.fixture -async def keys_app(db_session, monkeypatch): - """FastAPI app with auth middleware and only the /v1/keys routes mounted.""" - monkeypatch.setenv("DATABASE_URL", str(db_session.bind.url)) - from fastapi import FastAPI - - from app.middleware.auth import AuthMiddleware - from packages.db import session as session_mod - - class _PassthroughFactory: - async def __aenter__(self): - return db_session - - async def __aexit__(self, *exc): - return False # propagate, don't close — fixture owns the session - - monkeypatch.setattr(session_mod, "_session_factory", lambda: _PassthroughFactory()) - - from app.routes.keys import router as keys_router - - app = FastAPI() - app.add_middleware(AuthMiddleware) - app.include_router(keys_router) - return app - - -async def _client(app): - from httpx import ASGITransport, AsyncClient - - return AsyncClient(transport=ASGITransport(app=app), base_url="http://t") - - -@pytest.mark.parametrize("which", [1, 2], ids=["allowlist-restricted", "budget-restricted"]) -async def test_restricted_key_cannot_create_keys(keys_app, seeded_keys, db_session, which): - keys, restricted, budgeted = seeded_keys - caller = (restricted, budgeted)[which - 1] - async with await _client(keys_app) as c: - r = await c.post( - "/v1/keys", - json={"name": "escalated"}, - headers={"Authorization": f"Bearer {caller}"}, - ) - assert r.status_code == 403 - # The escalation must not have persisted anything. - from sqlalchemy import func, select - - from packages.db.models.api_key import ApiKey - - count = ( - await db_session.execute(select(func.count()).select_from(ApiKey)) - ).scalar_one() - assert count == 3 # root + restricted + budgeted, nothing new - - -@pytest.mark.parametrize("which", [1, 2], ids=["allowlist-restricted", "budget-restricted"]) -async def test_restricted_key_cannot_list_keys(keys_app, seeded_keys, which): - keys, restricted, budgeted = seeded_keys - caller = (restricted, budgeted)[which - 1] - async with await _client(keys_app) as c: - r = await c.get("/v1/keys", headers={"Authorization": f"Bearer {caller}"}) - assert r.status_code == 403 - - -async def test_restricted_key_cannot_revoke_keys(keys_app, seeded_keys, db_session): - _, restricted, _budgeted = seeded_keys - from sqlalchemy import select - - from packages.db.models.api_key import ApiKey - - rows = (await db_session.execute(select(ApiKey))).scalars().all() - target_id = next(r.id for r in rows if r.name == "default") - async with await _client(keys_app) as c: - r = await c.delete( - f"/v1/keys/{target_id}", - headers={"Authorization": f"Bearer {restricted}"}, - ) - assert r.status_code == 403 - target = next(r for r in rows if r.name == "default") - assert target.is_active # untouched - - -async def test_unrestricted_key_retains_full_management(keys_app, seeded_keys): - root, _restricted, _budgeted = seeded_keys - h = {"Authorization": f"Bearer {root}"} - async with await _client(keys_app) as c: - listed = await c.get("/v1/keys", headers=h) - assert listed.status_code == 200 - - created = await c.post("/v1/keys", json={"name": "child"}, headers=h) - assert created.status_code == 201 - child_id = created.json()["id"] - - revoked = await c.delete(f"/v1/keys/{child_id}", headers=h) - assert revoked.status_code == 204 - - -async def test_create_key_accepts_restrictions(keys_app, seeded_keys, db_session): - root, *_ = seeded_keys - h = {"Authorization": f"Bearer {root}"} - async with await _client(keys_app) as c: - r = await c.post( - "/v1/keys", - json={"name": "team-a", "model_allowlist": ["gpt-4o-mini"], "budget_limit_cents": 500}, - headers=h, - ) - assert r.status_code == 201, r.text - body = r.json() - assert body["model_allowlist"] == ["gpt-4o-mini"] - assert body["budget_limit_cents"] == 500 - - from sqlalchemy import select - - from packages.db.models.api_key import ApiKey - - row = ( - await db_session.execute(select(ApiKey).where(ApiKey.id == body["id"])) - ).scalar_one() - assert row.budget_limit_cents == 500 - assert row.model_allowlist == ["gpt-4o-mini"] - - -# ── Workspace scoping (IDOR regression tests) ──────────────────────────── - - -async def _make_foreign_workspace_key(db_session) -> tuple[str, str]: - """A key belonging to a different workspace; returns (id, name).""" - from packages.auth.hashing import generate_api_key - from packages.db.models.api_key import ApiKey - from packages.db.models.workspace import Workspace - - db_session.add(Workspace(id="ws-other", name="Other", slug="other")) - await db_session.flush() - - full_key, key_hash, key_prefix = generate_api_key() - row = ApiKey( - workspace_id="ws-other", - name="foreign-key", - key_hash=key_hash, - key_prefix=key_prefix, - ) - db_session.add(row) - await db_session.commit() - return row.id, full_key - - -async def test_list_keys_hides_other_workspaces(keys_app, seeded_keys, db_session): - root, *_ = seeded_keys - foreign_id, _foreign_key = await _make_foreign_workspace_key(db_session) - - async with await _client(keys_app) as c: - r = await c.get( - "/v1/keys", headers={"Authorization": f"Bearer {root}"} - ) - - assert r.status_code == 200 - listed_ids = {k["id"] for k in r.json()["keys"]} - assert foreign_id not in listed_ids - - -async def test_revoke_rejects_other_workspaces_key(keys_app, seeded_keys, db_session): - from sqlalchemy import select - - from packages.db.models.api_key import ApiKey - - root, *_ = seeded_keys - foreign_id, _foreign_key = await _make_foreign_workspace_key(db_session) - - async with await _client(keys_app) as c: - r = await c.delete( - f"/v1/keys/{foreign_id}", - headers={"Authorization": f"Bearer {root}"}, - ) - - assert r.status_code == 404 - row = ( - await db_session.execute(select(ApiKey).where(ApiKey.id == foreign_id)) - ).scalar_one() - assert row.is_active # untouched From e055378d8197baa91e199fda6b0aad3e3b550862 Mon Sep 17 00:00:00 2001 From: Hasit Bhatt Date: Thu, 27 Aug 2026 20:18:00 -0700 Subject: [PATCH 07/18] fix(budget): close fail-open charge window, stop over-charging on stream errors, add startup migration - Commit the request-log row and the budget charge atomically per attempt: set the settled flag only AFTER charge_budget succeeds, and on a retry that finds the log already persisted but the charge not yet settled, re-run ONLY the charge (never re-insert the row, never skip the charge). A transient charge error can no longer silently drop spend (fail-open). - Mid-stream upstream error now sets stream_completed=True before [DONE], so it settles at the actual cost instead of charging the key's entire remaining budget for a fully-delivered error response. - Add packages.db.migrate.ensure_budget_columns, run after create_all at boot: idempotently ALTERs spent_microcents into api_keys (seeded from historical request-log spend) and widens budget_limit_cents to BIGINT on Postgres, so existing deployments don't 503 on the new ORM columns. --- app/main.py | 7 +++ app/routes/chat.py | 100 +++++++++++++++++++++++------------------ packages/auth/spend.py | 17 +++++-- packages/db/migrate.py | 54 ++++++++++++++++++++++ 4 files changed, 131 insertions(+), 47 deletions(-) create mode 100644 packages/db/migrate.py diff --git a/app/main.py b/app/main.py index ff0d3d3..d2d2960 100644 --- a/app/main.py +++ b/app/main.py @@ -69,6 +69,13 @@ async def lifespan(app: FastAPI) -> AsyncGenerator[None, None]: async with engine.begin() as conn: await conn.run_sync(Base.metadata.create_all) + # create_all only makes missing tables, never alters existing ones. Bring + # existing deployments (SQLite volume, Postgres) up to date with columns added + # after their initial release so they don't 503 on the new ORM columns. + from packages.db.migrate import ensure_budget_columns + + await ensure_budget_columns(engine) + # Fail closed before any traffic can be served: refuse to boot when # provider credentials are (or would be) sealed with the publicly-known # dev encryption key. Runs after create_all so a fresh database's empty diff --git a/app/routes/chat.py b/app/routes/chat.py index 188ffad..29ee3ea 100644 --- a/app/routes/chat.py +++ b/app/routes/chat.py @@ -369,20 +369,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) -> None: + async def _settle_budget(session, actual_microcents: int, *, commit: bool = True) -> None: """Atomically record `actual_microcents` of spend against the cap. - Idempotent: only the first call for a request does work. `session` is - the DB session to run the charge on (the request session, or the - dedicated log session for the streaming path). `charge_budget` makes the - UPDATE refuse to let the counter exceed the cap, so a request can never - over-spend even under concurrency. + Idempotent in-process: only the first successful charge sets the flag. + `session` is the DB session to run the charge on (the request session, or + the dedicated log session for the streaming path). When `commit` is False + the UPDATE is executed but not committed, so the caller can commit it in + the same transaction as the request-log write — closing the fail-open + window where the log landed but the charge was lost. """ cap = getattr(kc, "_budget_cap", None) if cap is None or getattr(kc, "_budget_settled", False): return + await charge_budget(session, str(kc.key_id), cap, actual_microcents, commit=commit) kc._budget_settled = True - await charge_budget(session, str(kc.key_id), cap, actual_microcents) client = await router_cache.get_router(db) raw_strategy = getattr(client, "strategy", None) @@ -586,10 +587,10 @@ async def _settle_budget(session, actual_microcents: int) -> None: 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)) - await _settle_budget(db, 0) return JSONResponse( content=cache_hit_response, headers=_orca_response_headers( @@ -640,6 +641,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 @@ -652,7 +654,6 @@ async def _log_pre_stream_failure(status: int, err_type: str | None) -> None: except Exception: pass logger.warning("request_log_commit_failed", error=str(commit_err)) - await _settle_budget(db, 0) try: stream_obj = await client.acompletion( @@ -794,40 +795,50 @@ async def _already_persisted(s) -> bool: select(RequestLog.id).where(RequestLog.trace_id == row_values["trace_id"]) )) is not None + def _settlement_amount() -> int: + """Budget charge for this request, in microcents. + + On a stream that ended without a terminal [DONE] the real cost + is unknown, so charge the full remaining allowance (fail-closed) + instead of releasing the budget and letting a client bypass the + cap by hanging up before the usage frame. For a completed stream + the actual recorded cost is charged. + """ + actual = row_values.get("cost_microcents") or 0 + if not stream_completed: + actual = max( + actual, + (getattr(kc, "_budget_cap", 0) or 0) + - (getattr(kc, "_budget_spent", 0) or 0), + ) + 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, then charge the budget. + + The log row is committed first so it is durable even if the + subsequent budget charge hits a transient error. If the charge + fails, the retry path (reached with the log already persisted) + re-attempts ONLY the charge — it never re-inserts the row and + never skips a still-outstanding charge, so spend is never + silently dropped (fail-open). `_settle_budget` sets its settled + flag only after a successful charge, so a failed charge stays + retryable. """ 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)): + if getattr(kc, "_budget_settled", False): + return + await _settle_budget(db, _settlement_amount()) return db.add(log) try: await db.commit() - actual_cost = row_values.get("cost_microcents") or 0 - if not stream_completed: - # Stream ended without [DONE]: real cost unknown, - # so charge the full remaining allowance to keep - # the budget consumed (fail-closed) rather than - # releasing it and letting a client bypass the cap - # by hanging up before the usage frame. - actual_cost = max( - actual_cost, - (kc._budget_cap or 0) - (kc._budget_spent or 0), - ) - await _settle_budget(db, actual_cost) + await _settle_budget(db, _settlement_amount()) except Exception: try: await db.rollback() @@ -837,20 +848,15 @@ 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)): + if getattr(kc, "_budget_settled", False): + return + # Log already durable; re-run only the charge. + await _settle_budget(s, _settlement_amount()) return s.add(log) await s.commit() - actual_cost = row_values.get("cost_microcents") or 0 - if not stream_completed: - # Stream ended without [DONE]: real cost unknown, - # so charge the full remaining allowance to keep the - # budget consumed (fail-closed). - actual_cost = max( - actual_cost, - (kc._budget_cap or 0) - (kc._budget_spent or 0), - ) - await _settle_budget(s, actual_cost) + await _settle_budget(s, _settlement_amount()) finally: try: await s.close() @@ -1034,6 +1040,12 @@ async def _commit_row(*, retry: bool) -> None: # is legal; clients reading until [DONE] still get it after # an upstream error. yield "data: [DONE]\n\n" + # The error response was delivered in full (terminal [DONE] sent), + # so settle against the actual cost only — not the full remaining + # allowance. Without this, every mid-stream provider failure would + # charge (and exhaust) the key's entire remaining budget even + # though the delivered response cost ~0. + stream_completed = True finally: # Same shielding reason as the cancel branch: ensure the # log write actually completes before we unwind, even if @@ -1117,10 +1129,10 @@ async def _commit_row(*, retry: bool) -> None: ) db.add(log) try: + await _settle_budget(db, log.cost_microcents, commit=False) await db.commit() except Exception as commit_err: logger.warning("request_log_commit_failed", error=str(commit_err)) - await _settle_budget(db, log.cost_microcents) 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 57d1f3a..e459308 100644 --- a/packages/auth/spend.py +++ b/packages/auth/spend.py @@ -48,7 +48,12 @@ async def is_exhausted(db: AsyncSession, api_key_id: str, cap_microcents: int) - async def charge_budget( - db: AsyncSession, api_key_id: str, cap_microcents: int, actual_microcents: int + db: AsyncSession, + api_key_id: str, + cap_microcents: int, + actual_microcents: int, + *, + commit: bool = True, ) -> bool: """Atomically record ``actual_microcents`` of spend, never exceeding ``cap``. @@ -58,6 +63,10 @@ async def charge_budget( blocked going forward. The boundary request may already have been served upstream; it cannot be un-spent, but we never record more than the cap and we stop the next one. Fail-closed. + + When ``commit`` is False the UPDATEs are executed but not committed, so the + caller can commit them in the same transaction as the request-log write + (atomic log + charge — no window where the log lands but the charge is lost). """ actual = actual_microcents or 0 result = await db.execute( @@ -66,7 +75,8 @@ async def charge_budget( .values(spent_microcents=ApiKey.spent_microcents + actual) ) if result.rowcount: - await db.commit() + if commit: + await db.commit() return True # Would have exceeded the cap: clamp so the counter never overshoots and the # key is correctly reported as exhausted thereafter. @@ -75,5 +85,6 @@ async def charge_budget( .where(ApiKey.id == api_key_id, ApiKey.spent_microcents < cap_microcents) .values(spent_microcents=cap_microcents) ) - await db.commit() + if commit: + await db.commit() return False diff --git a/packages/db/migrate.py b/packages/db/migrate.py new file mode 100644 index 0000000..6fbb452 --- /dev/null +++ b/packages/db/migrate.py @@ -0,0 +1,54 @@ +"""Idempotent startup schema migrations for columns added after the first release. + +`Base.metadata.create_all` creates new tables but never alters existing ones, so a +deployment that already ran a release (a SQLite named volume, a fly.io/Postgres +volume) keeps an `api_keys` table without the `spent_microcents` column. After an +upgrade the ORM would then `SELECT` every mapped column and hit "no such column" +on every authenticated request — a 503 for the whole API. + +`ensure_budget_columns` is run once at boot, after `create_all`, and is safe to +call on every start: it inspects the live schema and only acts when the column is +missing. +""" + +from __future__ import annotations + +from sqlalchemy import inspect, text + + +async def ensure_budget_columns(engine) -> None: + """Add `spent_microcents` to `api_keys` if absent, seeded from request history. + + Also widens `budget_limit_cents` to BIGINT on Postgres (the microcent scale + can exceed int4). Both are no-ops on a fresh database. + """ + async with engine.begin() as conn: + cols = { + c["name"] + for c in await conn.run_sync(lambda sync: inspect(sync).get_columns("api_keys")) + } + is_postgres = engine.dialect.name == "postgresql" + + if "spent_microcents" not in cols: + await conn.execute( + text( + "ALTER TABLE api_keys ADD COLUMN spent_microcents BIGINT " + "NOT NULL DEFAULT 0" + ) + ) + # Seed lifetime spend from historical request logs so an existing key's + # cap is not silently reset to zero (which would re-grant a leaked key + # a full new budget). + await conn.execute( + text( + "UPDATE api_keys SET spent_microcents = (" + " SELECT COALESCE(SUM(cost_microcents), 0) FROM requests_log " + " WHERE requests_log.api_key_id = api_keys.id" + ") WHERE spent_microcents = 0" + ) + ) + + if is_postgres and "budget_limit_cents" in cols: + await conn.execute( + text("ALTER TABLE api_keys ALTER COLUMN budget_limit_cents TYPE BIGINT") + ) From e091480ab0646c29de7895c4a570df7f17bea2f5 Mon Sep 17 00:00:00 2001 From: Hasit Bhatt Date: Thu, 27 Aug 2026 22:46:25 -0700 Subject: [PATCH 08/18] fix(budget): make charge atomic+idempotent, retry blocking path, force usage for budgeted keys - The request-log row and the budget charge now commit in a single transaction (charge runs with commit=False, one commit() flushes both). A persisted trace_id therefore proves the charge also landed, so a retry returns without re-charging: spend is neither doubled (commit-ack-loss retry) nor dropped. - Blocking-path finally now persists log+charge atomically inside a bounded retry loop keyed on the row's trace_id, mirroring the streaming path, so a budgeted key served during a DB hiccup is no longer silently under-charged. - Budgeted keys force stream_options.include_usage=True on both streaming and blocking requests (a client include_usage=False no longer zeroes the cost), and a completed stream that never delivered a usage frame is treated as unknown-cost and charged the full remaining allowance (fail-closed) so a client cannot stream for free by suppressing the usage frame. - Tests: force-include_usage on budgeted keys (stream+blocking), completed stream without usage charges the full remaining cap, completed stream with usage charges only the actual cost. --- app/routes/chat.py | 124 +++++++++++-------- tests/integration/test_budget_enforcement.py | 89 +++++++++++++ 2 files changed, 163 insertions(+), 50 deletions(-) diff --git a/app/routes/chat.py b/app/routes/chat.py index 29ee3ea..5e0598a 100644 --- a/app/routes/chat.py +++ b/app/routes/chat.py @@ -370,20 +370,19 @@ async def execute_chat( ) async def _settle_budget(session, actual_microcents: int, *, commit: bool = True) -> None: - """Atomically record `actual_microcents` of spend against the cap. - - Idempotent in-process: only the first successful charge sets the flag. - `session` is the DB session to run the charge on (the request session, or - the dedicated log session for the streaming path). When `commit` is False - the UPDATE is executed but not committed, so the caller can commit it in - the same transaction as the request-log write — closing the fail-open - window where the log landed but the charge was lost. + """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 or getattr(kc, "_budget_settled", False): + if cap is None: return await charge_budget(session, str(kc.key_id), cap, actual_microcents, commit=commit) - kc._budget_settled = True client = await router_cache.get_router(db) raw_strategy = getattr(client, "strategy", None) @@ -608,16 +607,14 @@ async def _settle_budget(session, actual_microcents: int, *, commit: bool = True # 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 correctly. + # A budgeted key MUST receive usage so its spend is measured: a + # client-supplied `include_usage=False` would otherwise record zero cost + # and let a capped key stream for free, so force it on for any budgeted key + # regardless of the client's preference. existing_so = completion_kwargs.get("stream_options") or {} - if "include_usage" not in existing_so: + if getattr(kc, "_budget_cap", None) is not None or "include_usage" not in existing_so: completion_kwargs["stream_options"] = {**existing_so, "include_usage": True} async def _log_pre_stream_failure(status: int, err_type: str | None) -> None: @@ -698,6 +695,10 @@ async def sse() -> AsyncGenerator[str, None]: # released — otherwise a client could stream tokens then hang up before # the usage frame to bypass the cap. stream_completed = False + # True once any usage frame has been observed in the stream. A completed + # stream with no usage frame means cost is unknown (client suppressed it + # or the provider omitted it), so the cap must still be enforced. + usage_seen = False async def _finalize() -> None: """Write the request log row exactly once. @@ -798,14 +799,15 @@ async def _already_persisted(s) -> bool: def _settlement_amount() -> int: """Budget charge for this request, in microcents. - On a stream that ended without a terminal [DONE] the real cost - is unknown, so charge the full remaining allowance (fail-closed) - instead of releasing the budget and letting a client bypass the - cap by hanging up before the usage frame. For a completed stream - the actual recorded cost is charged. + When the real cost is unknown — the stream ended without a + terminal [DONE], or a completed stream never delivered a usage + frame (e.g. a client forced include_usage=False or a provider + omitted usage) — charge the full remaining allowance so a client + cannot suppress the usage frame to bypass the cap. """ actual = row_values.get("cost_microcents") or 0 - if not stream_completed: + cost_unknown = (not stream_completed) or (not usage_seen) + if cost_unknown: actual = max( actual, (getattr(kc, "_budget_cap", 0) or 0) @@ -814,16 +816,15 @@ def _settlement_amount() -> int: return actual async def _commit_row(*, retry: bool) -> None: - """Persist the request-log row, then charge the budget. - - The log row is committed first so it is durable even if the - subsequent budget charge hits a transient error. If the charge - fails, the retry path (reached with the log already persisted) - re-attempts ONLY the charge — it never re-inserts the row and - never skips a still-outstanding charge, so spend is never - silently dropped (fail-open). `_settle_budget` sets its settled - flag only after a successful charge, so a failed charge stays - retryable. + """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. """ log = RequestLog(**row_values) if session_mod._session_factory is None: @@ -831,14 +832,11 @@ async def _commit_row(*, retry: bool) -> None: # factory): the request-scoped session has to be # rolled back before a retry can reuse it. if retry and (await _already_persisted(db)): - if getattr(kc, "_budget_settled", False): - return - await _settle_budget(db, _settlement_amount()) return db.add(log) try: + await _settle_budget(db, _settlement_amount(), commit=False) await db.commit() - await _settle_budget(db, _settlement_amount()) except Exception: try: await db.rollback() @@ -849,14 +847,10 @@ async def _commit_row(*, retry: bool) -> None: s = session_mod._session_factory() try: if retry and (await _already_persisted(s)): - if getattr(kc, "_budget_settled", False): - return - # Log already durable; re-run only the charge. - await _settle_budget(s, _settlement_amount()) return s.add(log) + await _settle_budget(s, _settlement_amount(), commit=False) await s.commit() - await _settle_budget(s, _settlement_amount()) finally: try: await s.close() @@ -935,6 +929,7 @@ async def _commit_row(*, retry: bool) -> None: agg_fallback = True if "usage" in d and d["usage"]: agg_usage = d["usage"] + usage_seen = True if d.get("model"): agg_model = d["model"] yield f"data: {json.dumps(d, separators=(',', ':'))}\n\n" @@ -1086,6 +1081,12 @@ async def _commit_row(*, retry: bool) -> None: response: dict = {} actual_resolved: str | None = None try: + # A budgeted key must receive usage so its spend is measured. Force + # include_usage on for budgeted keys even if the client omitted it. + if getattr(kc, "_budget_cap", None) is not None: + existing_so = completion_kwargs.get("stream_options") or {} + if existing_so.get("include_usage") is not True: + completion_kwargs["stream_options"] = {**existing_so, "include_usage": True} response = await client.acompletion( **completion_kwargs, fallbacks=fallbacks_arg, @@ -1127,12 +1128,35 @@ 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 _settle_budget(db, log.cost_microcents, commit=False) - 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 + + max_attempts = len(_LOG_COMMIT_BACKOFF_S) + 1 + 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(log) + await _settle_budget(db, log.cost_microcents, commit=False) + await db.commit() + break + except Exception as commit_err: + try: + await db.rollback() + except Exception: + pass + if attempt == max_attempts: + logger.warning( + "request_log_commit_failed", error=str(commit_err), attempts=attempt, + ) hosted_fallback = _meta_hosted_fallback(response) if isinstance(response, dict) and "_orca_meta" in response: diff --git a/tests/integration/test_budget_enforcement.py b/tests/integration/test_budget_enforcement.py index 714b826..5fa33b4 100644 --- a/tests/integration/test_budget_enforcement.py +++ b/tests/integration/test_budget_enforcement.py @@ -214,3 +214,92 @@ async def test_unbudgeted_root_key_unaffected(budget_env): assert r.status_code == 200, r.text fake.acompletion.assert_awaited_once() + + +async def _budgeted_stream(budget_env, *, chunks, budget_limit_cents=10): + """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}, + }, + ) 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_without_usage_charges_remaining(budget_env): + # A completed stream that never delivers a usage frame (client forced + # include_usage=False, provider ignored it) must NOT bill zero — that would + # let a capped key stream for free. Fail-closed: charge the full remaining cap. + spent, call_args = await _budgeted_stream( + budget_env, + chunks=[ + {"choices": [{"delta": {"content": "hi"}, "finish_reason": None}]}, + {"choices": [{"delta": {}, "finish_reason": "stop"}]}, + ], + ) + # Even though the client demanded include_usage=False, the budgeted key forces it. + assert call_args.kwargs["stream_options"]["include_usage"] is True + # No usage frame observed -> full cap charged. + 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_forces_include_usage(budget_env): + # Non-streaming budgeted request also forces include_usage on, even when the + # client omits it. + 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"}], + "stream_options": {"include_usage": False}, + }, + ) + + assert r.status_code == 200, r.text + assert fake.acompletion.call_args.kwargs["stream_options"]["include_usage"] is True From 874aa5d08c2dfc3562c79aa0bebb113430a1f6a4 Mon Sep 17 00:00:00 2001 From: Hasit Bhatt Date: Thu, 24 Sep 2026 01:09:51 -0700 Subject: [PATCH 09/18] fix(budget): settle known costs on stream errors and usage-less blocking responses Two P1s from the latest review push: - Mid-stream provider errors now set usage_seen alongside stream_completed. The error response is delivered in full (terminal [DONE]) and its ~0 cost is recorded in the log row, so settlement is known; previously (not usage_seen) kept cost_unknown True and every transient upstream failure charged the key's entire remaining budget, permanently exhausting it. - The blocking path mirrors the streaming fail-closed rule: a budgeted key whose successful response carries no usage (provider ignored the forced include_usage) is charged the full remaining allowance instead of 0, closing the last cap-bypass vector. Error responses still charge the recorded cost. Tests: mid-stream error charges actual only (key stays usable); blocking without usage charges the cap; blocking with usage charges exactly the logged cost (no over-charge). --- app/routes/chat.py | 28 +++- tests/integration/test_budget_enforcement.py | 134 +++++++++++++++++++ 2 files changed, 160 insertions(+), 2 deletions(-) diff --git a/app/routes/chat.py b/app/routes/chat.py index 5e0598a..72b003e 100644 --- a/app/routes/chat.py +++ b/app/routes/chat.py @@ -1039,8 +1039,16 @@ async def _commit_row(*, retry: bool) -> None: # so settle against the actual cost only — not the full remaining # allowance. Without this, every mid-stream provider failure would # charge (and exhaust) the key's entire remaining budget even - # though the delivered response cost ~0. + # though the delivered response cost ~0. usage_seen is also set: + # the settlement here is *known* (the log row records the ~0 cost + # of the failed request), so the completed-stream-without-usage + # fail-closed rule does not apply — that rule guards against a + # client suppressing the usage frame, which is not a factor in a + # server-side provider error the client cannot steer. Client + # disconnects still take the GeneratorExit path above and keep + # charging the full remaining allowance. stream_completed = True + usage_seen = True finally: # Same shielding reason as the cancel branch: ensure the # log write actually completes before we unwind, even if @@ -1135,6 +1143,22 @@ async def _commit_row(*, retry: bool) -> None: # double-charging. from sqlalchemy import select + settle_amount = log.cost_microcents + # Fail-closed mirror of the streaming path's cost-unknown rule: a + # budgeted key whose successful response carries no usage (provider + # ignored the forced include_usage) has an unknown cost — charge the + # full remaining allowance so a delivered completion can never cost + # nothing. Error responses keep charging the recorded (≈0) cost. + if ( + getattr(kc, "_budget_cap", None) is not None + and status_code < 400 + and not (isinstance(response, dict) and response.get("usage")) + ): + settle_amount = max( + log.cost_microcents or 0, + kc._budget_cap - (getattr(kc, "_budget_spent", 0) or 0), + ) + max_attempts = len(_LOG_COMMIT_BACKOFF_S) + 1 for attempt in range(1, max_attempts + 1): try: @@ -1145,7 +1169,7 @@ async def _commit_row(*, retry: bool) -> None: ) is not None: break # already durable (log + charge committed) db.add(log) - await _settle_budget(db, log.cost_microcents, commit=False) + await _settle_budget(db, settle_amount, commit=False) await db.commit() break except Exception as commit_err: diff --git a/tests/integration/test_budget_enforcement.py b/tests/integration/test_budget_enforcement.py index 5fa33b4..7636de1 100644 --- a/tests/integration/test_budget_enforcement.py +++ b/tests/integration/test_budget_enforcement.py @@ -303,3 +303,137 @@ async def test_budgeted_blocking_forces_include_usage(budget_env): assert r.status_code == 200, r.text assert fake.acompletion.call_args.kwargs["stream_options"]["include_usage"] is True + + +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 + # Only the recorded (~0) cost is charged — not the 100_000-microcent cap. + assert await _get_spent(factory, key_id) == 0 + + # 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_blocking_without_usage_charges_remaining(budget_env): + # A budgeted key whose provider ignores the forced include_usage and returns + # a usage-less completion has an unknown cost. Mirroring the streaming rule, + # the blocking path must fail closed and charge the full remaining allowance + # — otherwise the delivered completion costs nothing and the cap is bypassed. + 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 + assert await _get_spent(factory, key_id) == 100_000 # 10 cents, fail-closed + + +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 From eef1effe565c784ff6bbaceb997fc15766a1635e Mon Sep 17 00:00:00 2001 From: Hasit Bhatt Date: Thu, 24 Sep 2026 01:22:16 -0700 Subject: [PATCH 10/18] fix(budget): only fail-close the remaining charge on delivered completions; migrate the spend index MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Self-review follow-ups on the previous push: - The blocking cost-unknown rule keyed off status_code < 400, but the `except HTTPException: raise` path leaves status_code at 200 with an empty response — a budgeted request that never got an upstream completion would have been charged its entire remaining allowance, exactly the bug class the mid-stream-error fix removed. The rule now requires an actual (non-empty) completion dict; early failures charge the recorded ~0 cost, mirroring the cache-hit and pre-stream paths. - ensure_budget_columns now also creates ix_requests_log_api_key_spend on upgraded databases (create_all only builds it on fresh ones), so deployments no longer drift from the model-declared schema. - New tests: HTTPException charge guard; legacy-schema upgrade path (column backfilled from history, index created, idempotent re-boot, ORM reads restored). --- app/routes/chat.py | 10 +- packages/db/migrate.py | 22 +++- tests/integration/test_budget_enforcement.py | 26 +++++ tests/unit/test_budget_migration.py | 101 +++++++++++++++++++ 4 files changed, 156 insertions(+), 3 deletions(-) create mode 100644 tests/unit/test_budget_migration.py diff --git a/app/routes/chat.py b/app/routes/chat.py index 72b003e..b8f9d81 100644 --- a/app/routes/chat.py +++ b/app/routes/chat.py @@ -1148,11 +1148,17 @@ async def _commit_row(*, retry: bool) -> None: # budgeted key whose successful response carries no usage (provider # ignored the forced include_usage) has an unknown cost — charge the # full remaining allowance so a delivered completion can never cost - # nothing. Error responses keep charging the recorded (≈0) cost. + # nothing. 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. if ( getattr(kc, "_budget_cap", None) is not None and status_code < 400 - and not (isinstance(response, dict) and response.get("usage")) + and isinstance(response, dict) + and response + and not response.get("usage") ): settle_amount = max( log.cost_microcents or 0, diff --git a/packages/db/migrate.py b/packages/db/migrate.py index 6fbb452..753c723 100644 --- a/packages/db/migrate.py +++ b/packages/db/migrate.py @@ -20,7 +20,9 @@ async def ensure_budget_columns(engine) -> None: """Add `spent_microcents` to `api_keys` if absent, seeded from request history. Also widens `budget_limit_cents` to BIGINT on Postgres (the microcent scale - can exceed int4). Both are no-ops on a fresh database. + can exceed int4) and creates the `ix_requests_log_api_key_spend` index that + create_all only builds on fresh databases. All steps are no-ops on a fresh + database. """ async with engine.begin() as conn: cols = { @@ -52,3 +54,21 @@ async def ensure_budget_columns(engine) -> None: await conn.execute( text("ALTER TABLE api_keys ALTER COLUMN budget_limit_cents TYPE BIGINT") ) + + # 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. is_deleted has existed since the + # first release (SoftDeleteMixin), so the index is always creatable. + idx = { + i["name"] + for i in await conn.run_sync( + lambda sync: inspect(sync).get_indexes("requests_log") + ) + } + if "ix_requests_log_api_key_spend" not in idx: + await conn.execute( + text( + "CREATE INDEX IF NOT EXISTS ix_requests_log_api_key_spend " + "ON requests_log (api_key_id, is_deleted)" + ) + ) diff --git a/tests/integration/test_budget_enforcement.py b/tests/integration/test_budget_enforcement.py index 7636de1..9cea6a4 100644 --- a/tests/integration/test_budget_enforcement.py +++ b/tests/integration/test_budget_enforcement.py @@ -409,6 +409,32 @@ async def test_budgeted_blocking_without_usage_charges_remaining(budget_env): assert await _get_spent(factory, key_id) == 100_000 # 10 cents, fail-closed +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. diff --git a/tests/unit/test_budget_migration.py b/tests/unit/test_budget_migration.py new file mode 100644 index 0000000..02a4f37 --- /dev/null +++ b/tests/unit/test_budget_migration.py @@ -0,0 +1,101 @@ +"""Upgrade-path coverage for ensure_budget_columns (packages/db/migrate.py). + +create_all never alters existing tables, so an upgraded deployment starts from +a legacy schema: api_keys without spent_microcents and requests_log without the +spend index. These tests build that legacy state by creating the real schema, +seeding rows, then dropping exactly what the pre-budget release lacked — and +pin that the startup migration restores it: (a) the column seeded from +historical request-log spend, (b) the composite index, (c) idempotency across +repeated boots, (d) the ORM (and thus auth) working again. +""" + +from __future__ import annotations + +from sqlalchemy import inspect as sa_inspect +from sqlalchemy import select, text +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.request_log import RequestLog + + +async def _legacy_deploy_engine(tmp_sqlite_url): + """Engine over a DB shaped like the last released schema, with one + budgeted key that already burned 2500 microcents of history.""" + 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) + async with factory() as s: + row = ApiKey( + workspace_id="w1", name="leaked-then-capped", key_hash="h", + key_prefix="p", budget_limit_cents=100, spent_microcents=2500, + ) + s.add(row) + await s.commit() + await s.refresh(row) + s.add(RequestLog( + workspace_id="w1", api_key_id=row.id, + trace_id="t1", model_requested="gpt-4o-mini", + model_resolved="gpt-4o-mini", provider="openai", + routing_strategy="balanced", input_tokens=5, output_tokens=2, + cost_microcents=2500, latency_ms=10, status_code=200, + )) + await s.commit() + # Downgrade to the pre-budget schema: drop the column and the index that + # only the new release's metadata declares. + async with engine.begin() as conn: + await conn.execute(text("DROP INDEX IF EXISTS ix_requests_log_api_key_spend")) + await conn.execute(text("ALTER TABLE api_keys DROP COLUMN spent_microcents")) + await engine.dispose() + return build_engine(tmp_sqlite_url) + + +async def test_ensure_budget_columns_upgrades_legacy_schema(tmp_sqlite_url): + engine = await _legacy_deploy_engine(tmp_sqlite_url) + try: + await ensure_budget_columns(engine) + + async with engine.connect() as conn: + # Seeded from history: a key that already burned 2500 microcents + # must not get a fresh full budget on upgrade. + spent = await conn.scalar( + text("SELECT spent_microcents FROM api_keys WHERE workspace_id = 'w1'") + ) + assert spent == 2500 + + idx_names = { + i["name"] + for i in await conn.run_sync( + lambda sync: sa_inspect(sync).get_indexes("requests_log") + ) + } + assert "ix_requests_log_api_key_spend" in idx_names + + # Idempotent across restarts: a second boot changes nothing. + await ensure_budget_columns(engine) + async with engine.connect() as conn: + assert await conn.scalar( + text("SELECT spent_microcents FROM api_keys WHERE workspace_id = 'w1'") + ) == 2500 + finally: + await engine.dispose() + + +async def test_orm_reads_work_after_upgrade(tmp_sqlite_url): + # The original 503 failure mode: the ORM SELECTs every mapped column, so + # without the migration every authenticated request broke on the upgraded + # database. After ensure_budget_columns the ApiKey select must succeed. + engine = await _legacy_deploy_engine(tmp_sqlite_url) + try: + await ensure_budget_columns(engine) + factory = async_sessionmaker(engine, expire_on_commit=False) + async with factory() as s: + row = (await s.execute(select(ApiKey))).scalar_one() + assert row.spent_microcents == 2500 + assert row.budget_limit_cents == 100 + finally: + await engine.dispose() From 1b3932981fd29a193772a20d11ac72af36774126 Mon Sep 17 00:00:00 2001 From: Hasit Bhatt Date: Thu, 24 Sep 2026 01:52:54 -0700 Subject: [PATCH 11/18] fix(budget): key streaming settlement on usage alone, harden blocking retry A client hang-up AFTER the usage frame charged the whole remaining allowance even though the row already recorded the measured cost. The usage frame is now the only cost-unknown signal: a stream that ends without one still fails closed, a measured one is charged as recorded. The blocking retry also re-added the previous attempt's rolled-back ORM object and never slept between attempts; it now inserts a fresh row per attempt with the same bounded backoff the streaming path uses. --- app/routes/chat.py | 83 ++++++++----- tests/integration/test_budget_enforcement.py | 122 +++++++++++++++++++ 2 files changed, 174 insertions(+), 31 deletions(-) diff --git a/app/routes/chat.py b/app/routes/chat.py index b8f9d81..1ae19f5 100644 --- a/app/routes/chat.py +++ b/app/routes/chat.py @@ -688,16 +688,15 @@ async def sse() -> AsyncGenerator[str, None]: status_code = 200 error_type: str | None = None log_written = False - # True only once a terminal `data: [DONE]` has been emitted, i.e. the - # response was delivered in full. While False, the stream ended early - # (client disconnect / mid-stream upstream error) and the real cost is - # unknown, so the budget claim must be kept (fail-closed) rather than - # released — otherwise a client could stream tokens then hang up before - # the usage frame to bypass the cap. - stream_completed = False - # True once any usage frame has been observed in the stream. A completed - # stream with no usage frame means cost is unknown (client suppressed it - # or the provider omitted it), so the cap must still be enforced. + # The usage frame is the billing signal: True once one has been + # observed. A stream that ends without it — client hung up before + # the usage frame, suppressed it, or the provider omitted it — has + # an unknown cost and is settled fail-closed against the key's + # remaining allowance so the cap cannot be bypassed. With a usage + # frame delivered the cost is known even if the client then + # disconnects, and charging more would over-bill a quantity the + # row already accounts for (and break charged == row.cost, the + # invariant the trace-id idempotence relies on). usage_seen = False async def _finalize() -> None: @@ -799,15 +798,16 @@ async def _already_persisted(s) -> bool: def _settlement_amount() -> int: """Budget charge for this request, in microcents. - When the real cost is unknown — the stream ended without a - terminal [DONE], or a completed stream never delivered a usage - frame (e.g. a client forced include_usage=False or a provider - omitted usage) — charge the full remaining allowance so a client - cannot suppress the usage frame to bypass the cap. + The usage frame is the billing signal. When no usage frame + was ever observed — the stream ended early, the client + suppressed the frame, or the provider omitted it — the real + cost is unknown and the full remaining allowance is charged + (fail-closed) so no client-side choice can bypass the cap. + Once a usage frame was delivered the cost is known — even + if the stream then died — and the recorded cost is charged. """ actual = row_values.get("cost_microcents") or 0 - cost_unknown = (not stream_completed) or (not usage_seen) - if cost_unknown: + if not usage_seen: actual = max( actual, (getattr(kc, "_budget_cap", 0) or 0) @@ -934,7 +934,6 @@ async def _commit_row(*, retry: bool) -> None: agg_model = d["model"] yield f"data: {json.dumps(d, separators=(',', ':'))}\n\n" yield "data: [DONE]\n\n" - stream_completed = True except (asyncio.CancelledError, GeneratorExit): # Client closed the connection (Ctrl+C, tab closed, browser # navigated away, proxy timeout, ...). Two distinct signals @@ -1035,19 +1034,18 @@ async def _commit_row(*, retry: bool) -> None: # is legal; clients reading until [DONE] still get it after # an upstream error. yield "data: [DONE]\n\n" - # The error response was delivered in full (terminal [DONE] sent), - # so settle against the actual cost only — not the full remaining - # allowance. Without this, every mid-stream provider failure would + # Mark the settlement known: the error response was delivered + # in full (terminal [DONE] sent) and its ~0 cost is recorded on + # the log row, so charge the actual cost only — not the full + # remaining allowance. Otherwise every transient mid-stream + # provider failure (rate limit, 5xx, network drop) would # charge (and exhaust) the key's entire remaining budget even - # though the delivered response cost ~0. usage_seen is also set: - # the settlement here is *known* (the log row records the ~0 cost - # of the failed request), so the completed-stream-without-usage - # fail-closed rule does not apply — that rule guards against a - # client suppressing the usage frame, which is not a factor in a - # server-side provider error the client cannot steer. Client - # disconnects still take the GeneratorExit path above and keep - # charging the full remaining allowance. - stream_completed = True + # though the delivered response cost ~0. The remaining-charge + # rule guards against a client suppressing the usage frame — + # not a factor in a server-side error the client cannot steer. + # Client disconnects never reach this branch (GeneratorExit is + # not an Exception); they unwind with usage_seen unchanged, so + # a hang-up before the usage frame still fails closed. usage_seen = True finally: # Same shielding reason as the cancel branch: ensure the @@ -1165,6 +1163,16 @@ async def _commit_row(*, retry: bool) -> None: kc._budget_cap - (getattr(kc, "_budget_spent", 0) or 0), ) + # 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 for attempt in range(1, max_attempts + 1): try: @@ -1174,7 +1182,7 @@ async def _commit_row(*, retry: bool) -> None: ) ) is not None: break # already durable (log + charge committed) - db.add(log) + db.add(RequestLog(**log_values)) await _settle_budget(db, settle_amount, commit=False) await db.commit() break @@ -1187,6 +1195,19 @@ async def _commit_row(*, retry: bool) -> None: logger.warning( "request_log_commit_failed", error=str(commit_err), attempts=attempt, ) + 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 like the arm above. + logger.warning( + "request_log_commit_failed", error=str(commit_err), attempts=attempt, + ) + raise hosted_fallback = _meta_hosted_fallback(response) if isinstance(response, dict) and "_orca_meta" in response: diff --git a/tests/integration/test_budget_enforcement.py b/tests/integration/test_budget_enforcement.py index 9cea6a4..dd0ab06 100644 --- a/tests/integration/test_budget_enforcement.py +++ b/tests/integration/test_budget_enforcement.py @@ -9,6 +9,7 @@ from __future__ import annotations +import asyncio import time from unittest.mock import AsyncMock @@ -463,3 +464,124 @@ async def test_budgeted_blocking_with_usage_charges_actual(budget_env): ) ).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_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 From ff295c763ee5bbe652c1913a57a907b32a087e96 Mon Sep 17 00:00:00 2001 From: Hasit Bhatt Date: Thu, 24 Sep 2026 01:59:01 -0700 Subject: [PATCH 12/18] fix(budget): price an unmeasured mid-stream delivery instead of settling it at zero MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The provider-error branch settled on the row's recorded cost, which is 0 when no usage frame ever arrived — so a capped key could stream free of charge behind a provider that fails mid-generation. What nothing measured is now priced from the characters actually delivered (chars/4, deliberately under-counting so an estimate never over-bills a failure the caller cannot steer), and the same number lands on the row and on the key. --- app/routes/chat.py | 67 ++++++++++++++++--- tests/integration/test_budget_enforcement.py | 53 +++++++++++++++ .../unit/test_unmeasured_stream_settlement.py | 36 ++++++++++ 3 files changed, 145 insertions(+), 11 deletions(-) create mode 100644 tests/unit/test_unmeasured_stream_settlement.py diff --git a/app/routes/chat.py b/app/routes/chat.py index 1ae19f5..1fdecd9 100644 --- a/app/routes/chat.py +++ b/app/routes/chat.py @@ -59,6 +59,42 @@ # 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 `_settle_unmeasured_stream`). +_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 + + +def _settle_unmeasured_stream(agg_usage: dict, agg_output_chars: int, body) -> dict: + """Token counts for a stream that delivered content but reported no usage. + + Reached only when the provider failed mid-generation after content had + already been forwarded: the prompt was billed upstream and the delivered + text is real, so settling that stream at zero would let a flaky provider + be streamed for free against a capped key. Character counts divided by 4 + under-count code and CJK on purpose — an estimate must not over-bill for a + failure the caller cannot steer. + """ + if agg_usage or not agg_output_chars: + return agg_usage + prompt_chars = sum(_text_chars(m.content) for m in body.messages) + return { + "prompt_tokens": max(1, prompt_chars // _CHARS_PER_TOKEN), + "completion_tokens": max(1, agg_output_chars // _CHARS_PER_TOKEN), + } + def _chunk_to_dict(chunk) -> dict: """Normalize a litellm chunk (Pydantic model or dict) into a plain dict. @@ -682,6 +718,10 @@ 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 `_settle_unmeasured_stream`). + 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 @@ -932,6 +972,10 @@ async def _commit_row(*, retry: bool) -> None: usage_seen = True if d.get("model"): agg_model = d["model"] + 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" yield "data: [DONE]\n\n" except (asyncio.CancelledError, GeneratorExit): @@ -1035,17 +1079,18 @@ async def _commit_row(*, retry: bool) -> None: # an upstream error. yield "data: [DONE]\n\n" # Mark the settlement known: the error response was delivered - # in full (terminal [DONE] sent) and its ~0 cost is recorded on - # the log row, so charge the actual cost only — not the full - # remaining allowance. Otherwise every transient mid-stream - # provider failure (rate limit, 5xx, network drop) would - # charge (and exhaust) the key's entire remaining budget even - # though the delivered response cost ~0. The remaining-charge - # rule guards against a client suppressing the usage frame — - # not a factor in a server-side error the client cannot steer. - # Client disconnects never reach this branch (GeneratorExit is - # not an Exception); they unwind with usage_seen unchanged, so - # a hang-up before the usage frame still fails closed. + # in full (terminal [DONE] sent), so charge the recorded cost + # rather than the full remaining allowance. Otherwise every + # transient mid-stream provider failure (rate limit, 5xx, + # network drop) would charge (and exhaust) the key's entire + # remaining budget. What the provider never measured is priced + # from what actually reached the client, so an unmeasured + # partial stream still costs something proportional to the + # delivery instead of nothing — a capped key cannot stream for + # free behind a flaky provider. Client disconnects never reach + # this branch (GeneratorExit is not an Exception); they unwind + # with usage_seen unchanged and keep failing closed. + agg_usage = _settle_unmeasured_stream(agg_usage, agg_output_chars, body) usage_seen = True finally: # Same shielding reason as the cancel branch: ensure the diff --git a/tests/integration/test_budget_enforcement.py b/tests/integration/test_budget_enforcement.py index dd0ab06..fe5cda2 100644 --- a/tests/integration/test_budget_enforcement.py +++ b/tests/integration/test_budget_enforcement.py @@ -379,6 +379,59 @@ async def _gen(): 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_charges_remaining(budget_env): # A budgeted key whose provider ignores the forced include_usage and returns # a usage-less completion has an unknown cost. Mirroring the streaming rule, diff --git a/tests/unit/test_unmeasured_stream_settlement.py b/tests/unit/test_unmeasured_stream_settlement.py new file mode 100644 index 0000000..8212c2b --- /dev/null +++ b/tests/unit/test_unmeasured_stream_settlement.py @@ -0,0 +1,36 @@ +"""`_settle_unmeasured_stream` only prices what nothing else measured.""" + +from __future__ import annotations + +from types import SimpleNamespace + +from app.routes.chat import _settle_unmeasured_stream + + +def _body(prompt: str): + return SimpleNamespace(messages=[SimpleNamespace(content=prompt)]) + + +def test_measured_usage_is_never_replaced_by_an_estimate(): + usage = {"prompt_tokens": 11, "completion_tokens": 22} + got = _settle_unmeasured_stream(usage, 90_000, _body("x" * 400)) + assert got == usage + + +def test_nothing_delivered_stays_unbilled(): + # A failure before the first content chunk delivered no tokens; inventing a + # prompt charge for it would repeat the over-charge this estimate replaces. + assert _settle_unmeasured_stream({}, 0, _body("x" * 400)) == {} + + +def test_estimate_prices_prompt_and_delivery_at_char_quarter(): + got = _settle_unmeasured_stream({}, 4_000, _body("y" * 400)) + assert got == {"prompt_tokens": 100, "completion_tokens": 1000} + + +def test_content_part_lists_count_their_text(): + body = SimpleNamespace(messages=[SimpleNamespace( + content=[{"type": "text", "text": "z" * 800}, {"type": "image_url"}] + )]) + got = _settle_unmeasured_stream({}, 400, body) + assert got == {"prompt_tokens": 200, "completion_tokens": 100} From d4d3e2d7bd7e823c42ce6b8faa6457360b0562a9 Mon Sep 17 00:00:00 2001 From: Hasit Bhatt Date: Thu, 24 Sep 2026 02:28:54 -0700 Subject: [PATCH 13/18] fix(budget): settle an adapter fault on its delivery instead of the whole remaining budget MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ff293c7 priced the unmeasured delivery in the provider-error branch but the sibling `except AdapterError` was left alone: it unwound with usage_seen False, so _settlement_amount treated the cost as unknown and charged a budgeted key its entire remaining lifetime budget for a bug of ours — even when the adapter faulted before the first chunk and nothing was delivered. It now settles the same way: the delivery is priced, nothing settles at 0. --- app/routes/chat.py | 8 ++ .../test_adapter_failure_attribution.py | 91 +++++++++++++++++++ 2 files changed, 99 insertions(+) diff --git a/app/routes/chat.py b/app/routes/chat.py index 1fdecd9..f06076e 100644 --- a/app/routes/chat.py +++ b/app/routes/chat.py @@ -1033,6 +1033,14 @@ async def _commit_row(*, retry: bool) -> None: logger.warning( "chat_completion_stream_adapter_error", served_model=agg_model, ) + # Settle it like the provider-error branch below: an adapter + # fault is our bug, not a choice the caller made, so leaving + # the settlement unknown would charge a budgeted key its entire + # remaining lifetime budget for a failure it cannot steer. The + # delivery is priced from what reached the client; nothing + # delivered settles at the 0 the row already records. + agg_usage = _settle_unmeasured_stream(agg_usage, agg_output_chars, body) + usage_seen = True aclose = getattr(stream_obj, "aclose", None) with anyio.CancelScope(shield=True): if aclose is not None: diff --git a/tests/integration/test_adapter_failure_attribution.py b/tests/integration/test_adapter_failure_attribution.py index 40f2966..a4738f3 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 From 58aec4a9da1de08ddee6b74b5833eb5dd2bce6d9 Mon Sep 17 00:00:00 2001 From: Hasit Bhatt Date: Thu, 24 Sep 2026 03:09:14 -0700 Subject: [PATCH 14/18] fix(budget): keep a given-up settlement counting against the cap The log row and the charge are one transaction, so when every commit attempt fails the delivered cost has no record anywhere: the counter never moved and no reconciliation exists, which leaves a key uncapped for the duration of a write outage. Park the amount instead, keep counting it in the pre-check, and let the next settlement absorb it so the cost becomes durable rather than being billed twice. --- app/routes/chat.py | 26 +++- packages/auth/spend.py | 52 ++++++- tests/integration/test_budget_enforcement.py | 144 +++++++++++++++++++ 3 files changed, 218 insertions(+), 4 deletions(-) diff --git a/app/routes/chat.py b/app/routes/chat.py index f06076e..8d4c1f5 100644 --- a/app/routes/chat.py +++ b/app/routes/chat.py @@ -31,7 +31,13 @@ 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, charge_budget, is_exhausted, read_spent +from packages.auth.spend import ( + MICROCENTS_PER_CENT, + charge_budget, + is_exhausted, + read_spent, + 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 @@ -939,6 +945,16 @@ async def _commit_row(*, retry: bool) -> None: "request_log_commit_failed", error=str(commit_err), attempts=attempt, ) + # Out of retries: this settlement will never be durable, + # and the row that recorded its cost is gone with it. Park + # the amount so the key's cap still counts it — otherwise a + # write outage is a window of free requests (spend.py). + if getattr(kc, "_budget_cap", None) is not None: + record_unsettled_spend( + str(kc.key_id), + recorded_microcents=kc._budget_spent, + microcents=_settlement_amount(), + ) except BaseException: # CancelledError aimed at us, not at the commit — # wait the in-flight attempt out so a row about to @@ -1248,6 +1264,14 @@ async def _commit_row(*, retry: bool) -> None: logger.warning( "request_log_commit_failed", error=str(commit_err), attempts=attempt, ) + # Same last resort as the streaming loop: undurable spend has + # to keep counting against the cap or it is simply lost. + if getattr(kc, "_budget_cap", None) is not None: + record_unsettled_spend( + str(kc.key_id), + recorded_microcents=kc._budget_spent, + microcents=settle_amount, + ) break logger.info( "request_log_commit_retry", error=str(commit_err), attempt=attempt, diff --git a/packages/auth/spend.py b/packages/auth/spend.py index e459308..9379eeb 100644 --- a/packages/auth/spend.py +++ b/packages/auth/spend.py @@ -33,6 +33,35 @@ MICROCENTS_PER_CENT = 10_000 +# 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 +# rolled back and the counter never moved. The amount is parked here, keyed by +# api key -> (counter value it is pending against, microcents), and keeps +# counting against the cap until a settlement that absorbs it becomes durable. +# Process-local by design: the failure it covers is a write outage, which the +# same process is still living through. +_unsettled: dict[str, tuple[int, int]] = {} + + +def record_unsettled_spend( + api_key_id: str, *, recorded_microcents: int, microcents: int +) -> None: + """Keep a settlement that gave up counting against the key's cap.""" + if microcents <= 0: + return + key = str(api_key_id) + parked = _unsettled.get(key) + _unsettled[key] = ( + min(recorded_microcents, parked[0]) if parked else recorded_microcents, + (parked[1] if parked else 0) + microcents, + ) + + +def unsettled_spend(api_key_id: str) -> int: + """Spend delivered for this key but not yet made durable.""" + parked = _unsettled.get(str(api_key_id)) + return parked[1] if parked else 0 + async def read_spent(db: AsyncSession, api_key_id: str) -> int: """Return the key's currently-recorded lifetime spend in microcents.""" @@ -43,8 +72,22 @@ async def read_spent(db: AsyncSession, api_key_id: str) -> int: 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 read_spent(db, api_key_id)) >= cap_microcents + """Fast pre-check: has the key already reached its lifetime cap? + + Includes spend whose settlement gave up, so a write outage cannot be used as + a window of free requests. Parked amounts are dropped once the counter has + moved past where they were parked: ``charge_budget`` is the only writer, so + it moved with them inside it. + """ + key = str(api_key_id) + spent = await read_spent(db, key) + parked = _unsettled.get(key) + if parked is None: + return spent >= cap_microcents + if spent > parked[0]: + del _unsettled[key] + return spent >= cap_microcents + return spent + parked[1] >= cap_microcents async def charge_budget( @@ -67,8 +110,11 @@ async def charge_budget( When ``commit`` is False the UPDATEs are executed but not committed, so the caller can commit them in the same transaction as the request-log write (atomic log + charge — no window where the log lands but the charge is lost). + + Anything parked by a give-up settlement is added to this charge, so that + undeliverable cost becomes durable together with it. """ - actual = actual_microcents or 0 + actual = (actual_microcents or 0) + unsettled_spend(api_key_id) result = await db.execute( update(ApiKey) .where(ApiKey.id == api_key_id, ApiKey.spent_microcents + actual <= cap_microcents) diff --git a/tests/integration/test_budget_enforcement.py b/tests/integration/test_budget_enforcement.py index fe5cda2..31e40b3 100644 --- a/tests/integration/test_budget_enforcement.py +++ b/tests/integration/test_budget_enforcement.py @@ -638,3 +638,147 @@ def _fail_first_log_insert(conn, cursor, statement, parameters, context, execute ).scalars().all() assert len(rows) == 1 # never doubled assert await _get_spent(factory, key_id) == rows[0].cost_microcents + + +class _WriteBlackout: + """Fail every settlement write at the cursor — a sustained "database is locked". + + 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. + """ + + _MATCHES = ("INSERT INTO requests_log", "UPDATE api_keys SET spent_microcents") + + 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 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 fails, so no row and no charge + survives anywhere. When the DB recovers, the next settlement pays for what + was already delivered as well — otherwise the counter understates the key by + everything it spent during the outage and the hard cap is open for its + duration. + """ + from sqlalchemy import 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) + usage = {"prompt_tokens": 10_000, "completion_tokens": 5_000, "total_tokens": 15_000} + fake.acompletion = AsyncMock(side_effect=lambda **kw: _completion("hello", usage=usage)) + + blackout = _WriteBlackout(factory) + blackout.active = True + try: + for i in range(3): + 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 + assert await _get_spent(factory, key_id) == 0 # nothing was recordable + blackout.active = False + + 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 3"}]}, + ) + assert r.status_code == 200, r.text + 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 _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. + """ + 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 + 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" From 072ec73dd452de9280153d71af4962c0d90a3fc9 Mon Sep 17 00:00:00 2001 From: Hasit Bhatt Date: Thu, 24 Sep 2026 03:13:14 -0700 Subject: [PATCH 15/18] fix(budget): park the charge at every exit that reports the row as lost The retries-exhausted arm was not the only one: a cancellation landing during the backoff, or aimed at _finalize while an attempt is in flight, also ends with request_log_commit_failed and no durable settlement. Route all five of those exits through one helper so admitting the loss and closing the cap cannot drift apart again. --- app/routes/chat.py | 58 +++++++++++++++++++--------------------------- 1 file changed, 24 insertions(+), 34 deletions(-) diff --git a/app/routes/chat.py b/app/routes/chat.py index 8d4c1f5..b3579eb 100644 --- a/app/routes/chat.py +++ b/app/routes/chat.py @@ -102,6 +102,23 @@ def _settle_unmeasured_stream(agg_usage: dict, agg_output_chars: int, body) -> d } +def _give_up_settlement(kc: KeyContext, amount: int, attempts: int, error: BaseException) -> 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 against the key's cap — the pre-check keeps + counting it and the next settlement that commits absorbs it (spend.py) — + rather than leaving the cap open for whoever reads the warning. + """ + logger.warning("request_log_commit_failed", error=str(error), attempts=attempts) + if getattr(kc, "_budget_cap", None) is not None: + record_unsettled_spend( + str(kc.key_id), + recorded_microcents=getattr(kc, "_budget_spent", 0) or 0, + microcents=amount, + ) + + def _chunk_to_dict(chunk) -> dict: """Normalize a litellm chunk (Pydantic model or dict) into a plain dict. @@ -935,26 +952,12 @@ 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, + _give_up_settlement( + kc, _settlement_amount(), attempt, commit_err, ) raise continue - logger.warning( - "request_log_commit_failed", - error=str(commit_err), attempts=attempt, - ) - # Out of retries: this settlement will never be durable, - # and the row that recorded its cost is gone with it. Park - # the amount so the key's cap still counts it — otherwise a - # write outage is a window of free requests (spend.py). - if getattr(kc, "_budget_cap", None) is not None: - record_unsettled_spend( - str(kc.key_id), - recorded_microcents=kc._budget_spent, - microcents=_settlement_amount(), - ) + _give_up_settlement(kc, _settlement_amount(), attempt, commit_err) except BaseException: # CancelledError aimed at us, not at the commit — # wait the in-flight attempt out so a row about to @@ -964,9 +967,8 @@ async def _commit_row(*, retry: bool) -> None: try: await commit_task except Exception as commit_err: - logger.warning( - "request_log_commit_failed", - error=str(commit_err), attempts=attempt, + _give_up_settlement( + kc, _settlement_amount(), attempt, commit_err, ) except BaseException: pass @@ -1261,17 +1263,7 @@ async def _commit_row(*, retry: bool) -> None: except Exception: pass if attempt == max_attempts: - logger.warning( - "request_log_commit_failed", error=str(commit_err), attempts=attempt, - ) - # Same last resort as the streaming loop: undurable spend has - # to keep counting against the cap or it is simply lost. - if getattr(kc, "_budget_cap", None) is not None: - record_unsettled_spend( - str(kc.key_id), - recorded_microcents=kc._budget_spent, - microcents=settle_amount, - ) + _give_up_settlement(kc, settle_amount, attempt, commit_err) break logger.info( "request_log_commit_retry", error=str(commit_err), attempt=attempt, @@ -1281,9 +1273,7 @@ async def _commit_row(*, retry: bool) -> None: except BaseException: # Cancelled during the backoff: nothing is in flight and the # row is given up on — say so, then propagate like the arm above. - logger.warning( - "request_log_commit_failed", error=str(commit_err), attempts=attempt, - ) + _give_up_settlement(kc, settle_amount, attempt, commit_err) raise hosted_fallback = _meta_hosted_fallback(response) From d199e6fd3744bb339d5b547dc431d9b967702fae Mon Sep 17 00:00:00 2001 From: Hasit Bhatt Date: Thu, 24 Sep 2026 04:09:27 -0700 Subject: [PATCH 16/18] fix(budget): claim parked spend once, and return it when the claimer dies Park each key's undeliverable cost as one number and hand it to exactly one settlement through a synchronous pop, instead of letting every charge absorb the whole park and expiring it lazily. Two requests that both pass the pre-check over one park used to bill it twice; a park re-recorded against a counter that had already moved past its baseline used to be dropped instead. Whoever claims a park must pay for it or put it back, so the blocked path now also parks on a cancellation that kills the transaction mid-write - otherwise the pop moves the loss window rather than closing it. --- app/routes/chat.py | 47 +++++++++---- packages/auth/spend.py | 69 ++++++++++---------- tests/integration/test_budget_enforcement.py | 50 ++++++++++++++ tests/unit/test_budget_spend.py | 19 ++++++ 4 files changed, 136 insertions(+), 49 deletions(-) diff --git a/app/routes/chat.py b/app/routes/chat.py index b3579eb..e2dace7 100644 --- a/app/routes/chat.py +++ b/app/routes/chat.py @@ -37,6 +37,7 @@ is_exhausted, read_spent, record_unsettled_spend, + take_unsettled_spend, ) from packages.auth.types import KeyContext from packages.db.models.request_log import RequestLog @@ -106,17 +107,13 @@ def _give_up_settlement(kc: KeyContext, amount: int, attempts: int, error: BaseE """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 against the key's cap — the pre-check keeps - counting it and the next settlement that commits absorbs it (spend.py) — - rather than leaving the cap open for whoever reads the warning. + remembers this cost. Park it — along with any spend this settlement had + claimed from a previous give-up — against the key's cap, rather than leaving + the cap open for whoever reads the warning (see `packages.auth.spend`). """ logger.warning("request_log_commit_failed", error=str(error), attempts=attempts) if getattr(kc, "_budget_cap", None) is not None: - record_unsettled_spend( - str(kc.key_id), - recorded_microcents=getattr(kc, "_budget_spent", 0) or 0, - microcents=amount, - ) + record_unsettled_spend(str(kc.key_id), amount + getattr(kc, "_budget_carried", 0)) def _chunk_to_dict(chunk) -> dict: @@ -428,7 +425,9 @@ 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: + async def _settle_budget( + session, actual_microcents: int, *, commit: bool = True, claim: bool = False + ) -> 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 @@ -437,10 +436,18 @@ async def _settle_budget(session, actual_microcents: int, *, commit: bool = True 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. + + `claim` is for the settlement loops, which are the only exits with a + give-up path: they take ownership of spend an earlier give-up parked (one + synchronous pop, so no other in-flight settlement can bill it too) and + carry it into every attempt until it is either durable or parked back. """ cap = getattr(kc, "_budget_cap", None) if cap is None: return + if claim: + kc._budget_carried += take_unsettled_spend(str(kc.key_id)) + actual_microcents += kc._budget_carried await charge_budget(session, str(kc.key_id), cap, actual_microcents, commit=commit) client = await router_cache.get_router(db) @@ -574,6 +581,7 @@ async def _settle_budget(session, actual_microcents: int, *, commit: bool = True ) kc._budget_cap = cap kc._budget_spent = await read_spent(db, str(kc.key_id)) + kc._budget_carried = 0 # parked spend this request's settlement has claimed started_perf = time.perf_counter() completion_kwargs = body.model_dump(exclude_none=True) @@ -887,7 +895,9 @@ async def _commit_row(*, retry: bool) -> None: 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. + retry) and never dropped. Any parked spend claimed for this + attempt travels with every attempt, so it is billed once and + only leaves with a durable commit. """ log = RequestLog(**row_values) if session_mod._session_factory is None: @@ -898,7 +908,9 @@ async def _commit_row(*, retry: bool) -> None: return db.add(log) try: - await _settle_budget(db, _settlement_amount(), commit=False) + await _settle_budget( + db, _settlement_amount(), commit=False, claim=True, + ) await db.commit() except Exception: try: @@ -912,7 +924,9 @@ async def _commit_row(*, retry: bool) -> None: if retry and (await _already_persisted(s)): return s.add(log) - await _settle_budget(s, _settlement_amount(), commit=False) + await _settle_budget( + s, _settlement_amount(), commit=False, claim=True, + ) await s.commit() finally: try: @@ -1254,7 +1268,7 @@ async def _commit_row(*, retry: bool) -> None: ) is not None: break # already durable (log + charge committed) db.add(RequestLog(**log_values)) - await _settle_budget(db, settle_amount, commit=False) + await _settle_budget(db, settle_amount, commit=False, claim=True) await db.commit() break except Exception as commit_err: @@ -1275,6 +1289,13 @@ async def _commit_row(*, retry: bool) -> None: # row is given up on — say so, then propagate like the arm above. _give_up_settlement(kc, settle_amount, attempt, commit_err) 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. Park it, or this request's cost and + # the parked spend it had claimed would vanish together. + _give_up_settlement(kc, settle_amount, attempt, cancel_err) + 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 9379eeb..835a7e9 100644 --- a/packages/auth/spend.py +++ b/packages/auth/spend.py @@ -35,32 +35,38 @@ # 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 -# rolled back and the counter never moved. The amount is parked here, keyed by -# api key -> (counter value it is pending against, microcents), and keeps -# counting against the cap until a settlement that absorbs it becomes durable. -# Process-local by design: the failure it covers is a write outage, which the -# same process is still living through. -_unsettled: dict[str, tuple[int, int]] = {} - - -def record_unsettled_spend( - api_key_id: str, *, recorded_microcents: int, microcents: int -) -> None: - """Keep a settlement that gave up counting against the key's cap.""" +# rolled back and the counter never moved. The amount is parked here and keeps +# counting against the cap until some settlement claims it. +# +# A claim is exclusive: `take_unsettled_spend` is a single synchronous pop, so +# two requests that both passed the pre-check cannot bill the same microcent +# twice. The claimer either commits it (it is now durable) or gives up and +# records it back together with its own cost. Process-local by design: the +# failure it covers is a write outage, which the same process is still living +# through. +_unsettled: dict[str, int] = {} + + +def record_unsettled_spend(api_key_id: str, microcents: int) -> None: + """Keep a settlement that gave up counting against the key's cap. + + Accumulates: replacing would let a second failed settlement shrink the total + and reopen the cap for the rest of the outage. + """ if microcents <= 0: return key = str(api_key_id) - parked = _unsettled.get(key) - _unsettled[key] = ( - min(recorded_microcents, parked[0]) if parked else recorded_microcents, - (parked[1] if parked else 0) + microcents, - ) + _unsettled[key] = _unsettled.get(key, 0) + microcents + + +def take_unsettled_spend(api_key_id: str) -> int: + """Claim this key's parked spend for one settlement; nobody else can claim it.""" + return _unsettled.pop(str(api_key_id), 0) def unsettled_spend(api_key_id: str) -> int: - """Spend delivered for this key but not yet made durable.""" - parked = _unsettled.get(str(api_key_id)) - return parked[1] if parked else 0 + """Parked spend awaiting a claim. Read-only — claiming is `take_*`'s job.""" + return _unsettled.get(str(api_key_id), 0) async def read_spent(db: AsyncSession, api_key_id: str) -> int: @@ -74,20 +80,10 @@ async def read_spent(db: AsyncSession, api_key_id: str) -> int: 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? - Includes spend whose settlement gave up, so a write outage cannot be used as - a window of free requests. Parked amounts are dropped once the counter has - moved past where they were parked: ``charge_budget`` is the only writer, so - it moved with them inside it. + Includes parked spend, so a write outage is not a window of free requests. """ - key = str(api_key_id) - spent = await read_spent(db, key) - parked = _unsettled.get(key) - if parked is None: - return spent >= cap_microcents - if spent > parked[0]: - del _unsettled[key] - return spent >= cap_microcents - return spent + parked[1] >= cap_microcents + spent = await read_spent(db, api_key_id) + return spent + unsettled_spend(api_key_id) >= cap_microcents async def charge_budget( @@ -111,10 +107,11 @@ async def charge_budget( caller can commit them in the same transaction as the request-log write (atomic log + charge — no window where the log lands but the charge is lost). - Anything parked by a give-up settlement is added to this charge, so that - undeliverable cost becomes durable together with it. + ``actual_microcents`` is what this settlement decided to bill, which may + already include spend it claimed via ``take_unsettled_spend``; this function + never reads the park, so the same microcent cannot be billed twice. """ - actual = (actual_microcents or 0) + unsettled_spend(api_key_id) + actual = actual_microcents or 0 result = await db.execute( update(ApiKey) .where(ApiKey.id == api_key_id, ApiKey.spent_microcents + actual <= cap_microcents) diff --git a/tests/integration/test_budget_enforcement.py b/tests/integration/test_budget_enforcement.py index 31e40b3..e689746 100644 --- a/tests/integration/test_budget_enforcement.py +++ b/tests/integration/test_budget_enforcement.py @@ -782,3 +782,53 @@ async def _no_usage(): ) assert r.status_code == 429, r.text assert r.json()["error"]["type"] == "rate_limit_error" + + +async def test_budgeted_concurrent_outage_bills_parked_spend_once(budget_env): + """Two settlements in flight over one park must not both absorb it. + + A parked cost is only ever a debt against the key's cap; whoever claims it + pays for it or puts it back. Two requests that pass the pre-check while it + still counts can each fold it into their own charge, and then the lifetime + counter records a delivery that happened once twice — over-recording spend, + which is the other direction the cap has to be hard in. + """ + from sqlalchemy import 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) + 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) -> 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": f"hi {i}"}]}, + ) + assert r.status_code == 200, r.text + + blackout = _WriteBlackout(factory) + blackout.active = True + try: + await _ask(0) # parks its own cost: nothing can record it + assert await _get_spent(factory, key_id) == 0 + # Both see the park, both fail to settle. Exactly one of them may own it. + await asyncio.gather(_ask(1), _ask(2)) + assert await _get_spent(factory, key_id) == 0 + blackout.active = False + await _ask(3) # durable: pays for itself and for everything still parked + 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 + cost = rows[0].cost_microcents + assert cost > 0 + assert await _get_spent(factory, key_id) == 4 * cost # four deliveries, not five diff --git a/tests/unit/test_budget_spend.py b/tests/unit/test_budget_spend.py index 19d6cee..7e045fd 100644 --- a/tests/unit/test_budget_spend.py +++ b/tests/unit/test_budget_spend.py @@ -9,6 +9,9 @@ charge_budget, is_exhausted, read_spent, + record_unsettled_spend, + take_unsettled_spend, + unsettled_spend, ) @@ -88,3 +91,19 @@ async def test_concurrent_charges_never_exceed_cap(db_session, key): def test_microcent_conversion_constant(): assert MICROCENTS_PER_CENT == 10_000 + + +async def test_parked_spend_is_billed_by_exactly_one_settlement(db_session, key): + """Two settlements in flight at once must not both absorb the same park. + + Both can pass the pre-check while the counter still reads the park's + baseline. The first one claims the park and bills it; the second must + claim nothing, or the same microcents land on the lifetime counter twice. + """ + cap = 10_000 + record_unsettled_spend(key.id, 3_000) + for own in (1_000, 1_000): # B settles, then C + claimed = take_unsettled_spend(key.id) + await charge_budget(db_session, key.id, cap, own + claimed) + assert await read_spent(db_session, key.id) == 5_000 # not 8_000 + assert unsettled_spend(key.id) == 0 From 7ba47e35e4980710f78ba543d61e64af7f51b20e Mon Sep 17 00:00:00 2001 From: Hasit Bhatt Date: Thu, 24 Sep 2026 10:35:27 -0700 Subject: [PATCH 17/18] fix(budget): confirm the settlement is not durable before parking it MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The give-up ran on the exception alone, but the last retry attempt has no follow-up trace_id check to absorb an applied-but-unacked commit: parking that cost on top of the charge already in `spent_microcents` bills one delivery twice, and the next settlement happily claims the park. So the give-up now proves non-durability first — the row is the charge, so its trace_id answers — and reads on its own session, since the failed attempt's is closed or rolled back by then. Unreadable counts as not durable: during a write outage parking is the only thing keeping the cap honest. Covers the cancellation exits too, including the blocking path's, which now releases the pending write's locks before asking. --- app/routes/chat.py | 118 +++++++++++-- tests/integration/test_budget_enforcement.py | 174 ++++++++++++++++++- tests/unit/test_give_up_settlement.py | 70 ++++++++ 3 files changed, 343 insertions(+), 19 deletions(-) create mode 100644 tests/unit/test_give_up_settlement.py diff --git a/app/routes/chat.py b/app/routes/chat.py index e2dace7..2bc686a 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 @@ -103,17 +103,77 @@ def _settle_unmeasured_stream(agg_usage: dict, agg_output_chars: int, body) -> d } -def _give_up_settlement(kc: KeyContext, amount: int, attempts: int, error: BaseException) -> None: +async def _give_up_settlement( + kc: KeyContext, + 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 — along with any spend this settlement had claimed from a previous give-up — 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 not None: - record_unsettled_spend(str(kc.key_id), amount + getattr(kc, "_budget_carried", 0)) + if getattr(kc, "_budget_cap", None) is None: + return + parked = amount + getattr(kc, "_budget_carried", 0) + 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 (and the claimed park + # travelling with it) vanish with the coroutine. + record_unsettled_spend(str(kc.key_id), parked) + raise + if not durable: + record_unsettled_spend(str(kc.key_id), parked) + + +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: @@ -826,7 +886,6 @@ async def _finalize() -> None: # _build_log_row fall back to the measured wall-clock # latency instead of persisting a constant 0. synthetic["_orca_meta"]["latency_ms"] = agg_latency - from sqlalchemy import select from packages.db import session as session_mod @@ -862,9 +921,10 @@ async def _finalize() -> None: row_values.setdefault("id", str(uuid.uuid4())) async def _already_persisted(s) -> bool: - return (await s.scalar( - select(RequestLog.id).where(RequestLog.trace_id == row_values["trace_id"]) - )) is not None + return await _trace_is_persisted(s, row_values["trace_id"]) + + 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. @@ -966,12 +1026,14 @@ 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. - _give_up_settlement( - kc, _settlement_amount(), attempt, commit_err, + await _give_up_settlement( + kc, _settlement_amount(), attempt, commit_err, _durable, ) raise continue - _give_up_settlement(kc, _settlement_amount(), attempt, commit_err) + await _give_up_settlement( + kc, _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 @@ -981,8 +1043,8 @@ async def _commit_row(*, retry: bool) -> None: try: await commit_task except Exception as commit_err: - _give_up_settlement( - kc, _settlement_amount(), attempt, commit_err, + await _give_up_settlement( + kc, _settlement_amount(), attempt, commit_err, _durable, ) except BaseException: pass @@ -1259,6 +1321,12 @@ async def _commit_row(*, retry: bool) -> None: 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 ( @@ -1277,7 +1345,9 @@ async def _commit_row(*, retry: bool) -> None: except Exception: pass if attempt == max_attempts: - _give_up_settlement(kc, settle_amount, attempt, commit_err) + await _give_up_settlement( + kc, settle_amount, attempt, commit_err, _durable, + ) break logger.info( "request_log_commit_retry", error=str(commit_err), attempt=attempt, @@ -1287,14 +1357,26 @@ async def _commit_row(*, retry: bool) -> None: except BaseException: # Cancelled during the backoff: nothing is in flight and the # row is given up on — say so, then propagate like the arm above. - _give_up_settlement(kc, settle_amount, attempt, commit_err) + await _give_up_settlement( + kc, 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. Park it, or this request's cost and - # the parked spend it had claimed would vanish together. - _give_up_settlement(kc, settle_amount, attempt, cancel_err) + # 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 cost + # and the parked spend it had claimed would vanish together. + 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, settle_amount, attempt, cancel_err, _durable, + ) raise hosted_fallback = _meta_hosted_fallback(response) diff --git a/tests/integration/test_budget_enforcement.py b/tests/integration/test_budget_enforcement.py index e689746..b083da3 100644 --- a/tests/integration/test_budget_enforcement.py +++ b/tests/integration/test_budget_enforcement.py @@ -14,6 +14,29 @@ from unittest.mock import AsyncMock import pytest +from sqlalchemy.ext.asyncio import AsyncSession + +from packages.db.models.request_log import RequestLog + + +class _AckLossSession(AsyncSession): + """Commit for real, then report failure as if the ack never came back. + + One-shot, and armed only by the write hook that lets an attempt land, so it + consumes exactly the settlement commit it follows: the row is durable while + its caller still sees an exception — the case the retry loops' trace_id + check exists for. Inert (an ordinary AsyncSession) otherwise. + """ + + drop_ack = False + drops = 0 + + async def commit(self): + await super().commit() + if _AckLossSession.drop_ack: + _AckLossSession.drop_ack = False + _AckLossSession.drops += 1 + raise ConnectionError("connection dropped mid-ack") @pytest.fixture @@ -38,7 +61,9 @@ async def budget_env(tmp_sqlite_url, monkeypatch): from sqlalchemy.ext.asyncio import async_sessionmaker from packages.db import session as session_mod - factory = async_sessionmaker(engine, expire_on_commit=False) + factory = async_sessionmaker( + engine, expire_on_commit=False, class_=_AckLossSession, + ) session_mod._session_factory = factory from app.seed import seed_initial_state @@ -832,3 +857,150 @@ async def _ask(i: int) -> None: cost = rows[0].cost_microcents assert cost > 0 assert await _get_spent(factory, key_id) == 4 * cost # four deliveries, not five + + +class _LastAttemptThenAckLoss: + """Fail every settlement write but the last, and lose that one's ack. + + Deriving the count from the retry budget is what makes this the *final* + attempt — the one with no retry left to run the trace_id check, which is + where a give-up has to decide on the exception alone. + """ + + def __init__(self, factory): + from app.routes.chat import _LOG_COMMIT_BACKOFF_S + + self.last = len(_LOG_COMMIT_BACKOFF_S) + 1 + self.n = 0 + self.drops = 0 + 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 "INSERT INTO requests_log" not in statement: + return + self.n += 1 + if self.n < self.last: + raise RuntimeError("database is locked") + _AckLossSession.drop_ack = self.n == self.last + + def close(self): + from sqlalchemy import event + + event.remove(self._engine, "before_cursor_execute", self._handle) + _AckLossSession.drop_ack = False + self.drops = _AckLossSession.drops + _AckLossSession.drops = 0 + + +async def _ask_blocking(make_client, key: str, content: 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": content}]}, + ) + assert r.status_code == 200, r.text + + +async def _logged_rows(factory, key_id: str): + from sqlalchemy import select + + async with factory() as s: + return ( + await s.execute(select(RequestLog).where(RequestLog.api_key_id == key_id)) + ).scalars().all() + + +async def test_budgeted_blocking_commit_ack_loss_bills_the_delivery_once(budget_env): + """A settlement that is durable must not also be parked. + + The loop's idempotence is the trace_id check at the START of an attempt, so + the final one has no follow-up: it gives up on the exception alone. If that + commit applied and only its ack was lost, parking its cost on top of the + charge already in `spent_microcents` bills one delivery twice, and the next + settlement happily claims the park. + """ + from packages.auth.spend import unsettled_spend + + 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)) + + hook = _LastAttemptThenAckLoss(factory) + try: + await _ask_blocking(make_client, key, "hi") + finally: + hook.close() + # The scenario really ran: every attempt but the last failed outright, and + # the last one raised on a commit the database had already applied. + assert (hook.n, hook.drops) == (hook.last, 1) + + rows = await _logged_rows(factory, key_id) + assert len(rows) == 1 # the row landed, and the retries never doubled it + cost = rows[0].cost_microcents + assert cost > 0 + assert await _get_spent(factory, key_id) == cost + assert unsettled_spend(key_id) == 0 # durable, so there is nothing to park + + await _ask_blocking(make_client, key, "hi again") + assert await _get_spent(factory, key_id) == 2 * cost # not three + + +async def test_budgeted_stream_commit_ack_loss_bills_the_delivery_once(budget_env): + """Same guarantee on the streaming loop, which the give-up also serves.""" + from packages.auth.spend import unsettled_spend + + make_client, fake, factory, _root = budget_env + key, key_id = await _make_budgeted_key(factory, budget_limit_cents=100) + + def _measured_stream(): + async def _gen(): + yield {"choices": [{"delta": {"content": "hi"}, "finish_reason": None}]} + yield { + "usage": { + "prompt_tokens": 10_000, + "completion_tokens": 5_000, + "total_tokens": 15_000, + }, + "choices": [{"delta": {}, "finish_reason": "stop"}], + } + return _gen() + + payload = { + "model": "gpt-4o-mini", + "stream": True, + "messages": [{"role": "user", "content": "hi"}], + } + fake.acompletion = AsyncMock(return_value=_measured_stream()) + + hook = _LastAttemptThenAckLoss(factory) + 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: + hook.close() + # The scenario really ran: every attempt but the last failed outright, and + # the last one raised on a commit the database had already applied. + assert (hook.n, hook.drops) == (hook.last, 1) + + rows = await _logged_rows(factory, key_id) + assert len(rows) == 1 + cost = rows[0].cost_microcents + assert cost > 0 + assert await _get_spent(factory, key_id) == cost + assert unsettled_spend(key_id) == 0 + + fake.acompletion = AsyncMock( + return_value=_completion("hello", usage={ + "prompt_tokens": 10_000, "completion_tokens": 5_000, "total_tokens": 15_000, + }) + ) + await _ask_blocking(make_client, key, "hi again") + assert await _get_spent(factory, key_id) == 2 * cost diff --git a/tests/unit/test_give_up_settlement.py b/tests/unit/test_give_up_settlement.py new file mode 100644 index 0000000..a47e316 --- /dev/null +++ b/tests/unit/test_give_up_settlement.py @@ -0,0 +1,70 @@ +"""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 later +settlement can record it. 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. +""" + +from __future__ import annotations + +import asyncio + +import pytest + +from app.routes.chat import _give_up_settlement +from packages.auth.spend import unsettled_spend + + +class _Kc: + """The three attributes the give-up reads off a KeyContext.""" + + def __init__(self, key_id: str, *, cap: int = 10_000, carried: int = 0): + self.key_id = key_id + self._budget_cap = cap + self._budget_carried = carried + + +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", carried=700) + + async def persisted() -> bool: + return True + + await _give_up_settlement(kc, 900, 3, _failed(), persisted) + assert unsettled_spend(kc.key_id) == 0 + + +async def test_lost_settlement_parks_its_cost_with_what_it_claimed(): + """Nothing is durable: the delivery and the claimed park stay on the cap.""" + kc = _Kc("lost-key", carried=700) + + async def persisted() -> bool: + return False + + await _give_up_settlement(kc, 900, 3, _failed("database is locked"), persisted) + assert unsettled_spend(kc.key_id) == 1_600 + + +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 and the parked spend travelling with it out with the + coroutine. + """ + kc = _Kc("cancelled-probe-key", carried=700) + + async def persisted() -> bool: + raise asyncio.CancelledError + + with pytest.raises(asyncio.CancelledError): + await _give_up_settlement(kc, 900, 1, _failed(), persisted) + assert unsettled_spend(kc.key_id) == 1_600 From 45ce3979b26836d1ff0533db3e8476018b7fed4e Mon Sep 17 00:00:00 2001 From: Hasit Bhatt Date: Thu, 24 Sep 2026 10:35:45 -0700 Subject: [PATCH 18/18] fix(migrations): let a boot that lost the ALTER race go on booting MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Every worker runs ensure_budget_columns in its lifespan, so the first boot after an upgrade starts several of them against one schema none has altered yet. Each inspects, each adds the column, and all but one fail — which took the worker's startup down with it, on exactly the deployment shape the column exists to serve. Each startup statement now runs inside a savepoint that reads "already exists" as success, so Postgres keeps the surrounding transaction alive for the seed and index that follow. The seed is an absolute assignment from the request log, not an increment, so the winner and the loser both landing it is the same number. --- packages/db/migrate.py | 52 +++++++++++++++++++++-------- tests/unit/test_budget_migration.py | 50 +++++++++++++++++++++++++++ 2 files changed, 88 insertions(+), 14 deletions(-) diff --git a/packages/db/migrate.py b/packages/db/migrate.py index 753c723..c5dddce 100644 --- a/packages/db/migrate.py +++ b/packages/db/migrate.py @@ -7,13 +7,39 @@ on every authenticated request — a 503 for the whole API. `ensure_budget_columns` is run once at boot, after `create_all`, and is safe to -call on every start: it inspects the live schema and only acts when the column is -missing. +call on every start: it inspects the live schema, only acts when the change is +missing, and tolerates another process racing it to the same change. """ from __future__ import annotations from sqlalchemy import inspect, text +from sqlalchemy.exc import DBAPIError + + +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 + + +async def _apply_ddl(conn, statement: str) -> 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 + has several processes inspecting a schema none of them has altered yet. Each + then issues the same statement and all but one fail — "column ... already + exists" on Postgres, "duplicate column name" on SQLite — which is success + 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. + """ + try: + async with conn.begin_nested(): + await conn.execute(text(statement)) + except DBAPIError as err: + if not _already_applied(err): + raise async def ensure_budget_columns(engine) -> None: @@ -32,11 +58,10 @@ async def ensure_budget_columns(engine) -> None: is_postgres = engine.dialect.name == "postgresql" if "spent_microcents" not in cols: - await conn.execute( - text( - "ALTER TABLE api_keys ADD COLUMN spent_microcents BIGINT " - "NOT NULL DEFAULT 0" - ) + await _apply_ddl( + conn, + "ALTER TABLE api_keys ADD COLUMN spent_microcents BIGINT " + "NOT NULL DEFAULT 0", ) # Seed lifetime spend from historical request logs so an existing key's # cap is not silently reset to zero (which would re-grant a leaked key @@ -51,8 +76,8 @@ async def ensure_budget_columns(engine) -> None: ) if is_postgres and "budget_limit_cents" in cols: - await conn.execute( - text("ALTER TABLE api_keys ALTER COLUMN budget_limit_cents TYPE BIGINT") + await _apply_ddl( + conn, "ALTER TABLE api_keys ALTER COLUMN budget_limit_cents TYPE BIGINT" ) # The model declares ix_requests_log_api_key_spend (api_key_id, @@ -66,9 +91,8 @@ async def ensure_budget_columns(engine) -> None: ) } if "ix_requests_log_api_key_spend" not in idx: - await conn.execute( - text( - "CREATE INDEX IF NOT EXISTS ix_requests_log_api_key_spend " - "ON requests_log (api_key_id, is_deleted)" - ) + await _apply_ddl( + conn, + "CREATE INDEX IF NOT EXISTS ix_requests_log_api_key_spend " + "ON requests_log (api_key_id, is_deleted)", ) diff --git a/tests/unit/test_budget_migration.py b/tests/unit/test_budget_migration.py index 02a4f37..2200267 100644 --- a/tests/unit/test_budget_migration.py +++ b/tests/unit/test_budget_migration.py @@ -99,3 +99,53 @@ async def test_orm_reads_work_after_upgrade(tmp_sqlite_url): assert row.budget_limit_cents == 100 finally: await engine.dispose() + + +async def test_ensure_budget_columns_survives_a_racing_boot(tmp_sqlite_url, monkeypatch): + """The loser of a concurrent-boot ALTER still boots, and still seeds. + + Every worker runs this at startup, and the first boot after an upgrade + starts several of them at once against one database. Both inspect the + schema before either alters it, so the loser's ALTER meets a column that + appeared in between and the driver rejects it with "duplicate column + name" — which used to escape into the lifespan and keep that worker down. + """ + import packages.db.migrate as migrate + + engine = await _legacy_deploy_engine(tmp_sqlite_url) + try: + # The other boot gets there first: the column is committed, so this + # process's inspection is now stale relative to the schema. + async with engine.begin() as conn: + await conn.execute( + text( + "ALTER TABLE api_keys ADD COLUMN spent_microcents BIGINT " + "NOT NULL DEFAULT 0" + ) + ) + inspector = sa_inspect(engine.sync_engine) + + class _StaleSchema: + def __init__(self, inner): + self._inner = inner + + def get_columns(self, table_name): + return [ + c for c in self._inner.get_columns(table_name) + if c["name"] != "spent_microcents" + ] + + def __getattr__(self, name): + return getattr(self._inner, name) + + monkeypatch.setattr(migrate, "inspect", lambda _sync: _StaleSchema(inspector)) + await ensure_budget_columns(engine) + + async with engine.connect() as conn: + # It went on to seed the column the winner added — before the fix + # the boot died on the ALTER and never reached this. + assert await conn.scalar( + text("SELECT spent_microcents FROM api_keys WHERE workspace_id = 'w1'") + ) == 2500 + finally: + await engine.dispose()