feat(quotas): quota_policies table and a per-call cost on token_usage

Migration 0033 adds quota_policies (instance default, team per-member
allowance, user override; a token budget and a USD budget per row) and
token_usage.cost. The column is added IF NOT EXISTS so a database that
already carries it upgrades cleanly.

Every usage row now records the call's USD cost from the model catalog;
bring-your-own models are recorded at $0.
This commit is contained in:
Alex committed 2026-09-21 12:11:11 +01:00
1 parent b77561288d
commit 43fad2a865
9 files changed
+415 -3

No files matched your search

+102
View File
@@ -0,0 +1,102 @@
"""0033 quotas — admin-set usage limits and a per-call cost.
``quota_policies`` holds the limits an instance admin sets at three layers:
the instance default (``subject_id`` NULL), a team's per-member allowance
(``subject_id`` = ``teams.id``) and a single user's override (``subject_id`` =
the auth ``sub``). Each row carries a token budget and a USD budget; per
budget a row either sets a limit (0 blocks), marks it unlimited, or leaves
both empty to defer to the next layer. ``bucket`` narrows a row to chat without
an agent (``direct``) or traffic through an agent (``agent``); ``all`` covers both.
``subject_id`` is polymorphic, so there is no FK: an AFTER DELETE trigger on
``teams`` scrubs a deleted team's rows, and user rows follow the ``user_roles``
convention of never blocking user deletion.
``token_usage.cost`` is the USD cost of the call at write time (see
``docsgpt/pricing.py``); 0 for unpriced and bring-your-own models.
Idempotent both ways.
Revision ID: 0033_quotas
Revises: 0032_personal_access_tokens
"""
from typing import Sequence, Union
from alembic import op
revision: str = "0033_quotas"
down_revision: Union[str, None] = "0032_personal_access_tokens"
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def upgrade() -> None:
op.execute(
"ALTER TABLE token_usage ADD COLUMN IF NOT EXISTS cost NUMERIC(12,8) NOT NULL DEFAULT 0;"
)
op.execute(
"""
CREATE TABLE IF NOT EXISTS quota_policies (
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
scope TEXT NOT NULL CHECK (scope IN ('instance', 'team', 'user')),
subject_id TEXT,
bucket TEXT NOT NULL DEFAULT 'all'
CHECK (bucket IN ('all', 'direct', 'agent')),
token_limit BIGINT CHECK (token_limit >= 0),
token_unlimited BOOLEAN NOT NULL DEFAULT false,
cost_limit_usd NUMERIC(12,4) CHECK (cost_limit_usd >= 0),
cost_unlimited BOOLEAN NOT NULL DEFAULT false,
enabled BOOLEAN NOT NULL DEFAULT true,
note TEXT,
created_by TEXT,
updated_by TEXT,
created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
updated_at TIMESTAMPTZ NOT NULL DEFAULT now(),
CONSTRAINT quota_policies_subject_chk
CHECK ((scope = 'instance') = (subject_id IS NULL)),
CONSTRAINT quota_policies_token_chk
CHECK (NOT (token_unlimited AND token_limit IS NOT NULL)),
CONSTRAINT quota_policies_cost_chk
CHECK (NOT (cost_unlimited AND cost_limit_usd IS NOT NULL))
);
"""
)
# One row per (layer subject, bucket); the instance row's NULL subject folds to ''.
op.execute(
"CREATE UNIQUE INDEX IF NOT EXISTS quota_policies_subject_uidx "
"ON quota_policies (scope, COALESCE(subject_id, ''), bucket);"
)
op.execute("DROP TRIGGER IF EXISTS quota_policies_set_updated_at ON quota_policies;")
op.execute(
"""
CREATE TRIGGER quota_policies_set_updated_at
BEFORE UPDATE ON quota_policies
FOR EACH ROW EXECUTE FUNCTION set_updated_at();
"""
)
op.execute(
"""
CREATE OR REPLACE FUNCTION cleanup_team_quota_policies() RETURNS trigger AS $$
BEGIN
DELETE FROM quota_policies WHERE scope = 'team' AND subject_id = OLD.id::text;
RETURN OLD;
END;
$$ LANGUAGE plpgsql;
"""
)
op.execute("DROP TRIGGER IF EXISTS teams_cleanup_quota_policies ON teams;")
op.execute(
"""
CREATE TRIGGER teams_cleanup_quota_policies
AFTER DELETE ON teams
FOR EACH ROW EXECUTE FUNCTION cleanup_team_quota_policies();
"""
)
def downgrade() -> None:
op.execute("DROP TRIGGER IF EXISTS teams_cleanup_quota_policies ON teams;")
op.execute("DROP FUNCTION IF EXISTS cleanup_team_quota_policies();")
op.execute("DROP TABLE IF EXISTS quota_policies;")
op.execute("ALTER TABLE token_usage DROP COLUMN IF EXISTS cost;")
+3
View File
@@ -52,6 +52,7 @@ class LLMCreator:
base_url = None
upstream_model_id = model_id
capabilities = None
model = None
if model_id:
user_id = model_user_id
if user_id is None:
@@ -127,4 +128,6 @@ class LLMCreator:
# llm.model_id is the upstream name (BYOM resolves it above); stamp
# the canonical id (UUID for BYOM) separately for token_usage.
llm._canonical_model_id = model_id
# Calls to a user's own model are recorded at $0 (see ``docsgpt/usage.py``).
llm._is_byom = model is not None and getattr(model, "source", "builtin") == "user"
return llm
+43
View File
@@ -28,6 +28,7 @@ from sqlalchemy import (
Index,
Integer,
MetaData,
Numeric,
PrimaryKeyConstraint,
UniqueConstraint,
Table,
@@ -241,6 +242,9 @@ token_usage_table = Table(
# cache activity, so hit-rate queries stay honest across providers.
Column("cached_tokens", Integer),
Column("cache_write_tokens", Integer),
# Added in ``0033_quotas``. USD cost of the call at write time; 0 for
# unpriced and bring-your-own models.
Column("cost", Numeric(12, 8), nullable=False, server_default="0"),
)
user_logs_table = Table(
@@ -1121,3 +1125,42 @@ Index(
personal_access_tokens_table.c.user_id,
personal_access_tokens_table.c.created_at.desc(),
)
# --- Usage quotas (migration 0033) ------------------------------------------
# Admin-set limits at three layers: instance (``subject_id`` NULL), team
# per-member allowance (``teams.id``) and user override (auth ``sub``). Per
# budget a row sets a limit, marks it unlimited, or defers to the next layer.
quota_policies_table = Table(
"quota_policies",
metadata,
Column("id", UUID(as_uuid=True), primary_key=True, server_default=func.gen_random_uuid()),
Column("scope", Text, nullable=False),
Column("subject_id", Text),
Column("bucket", Text, nullable=False, server_default="all"),
Column("token_limit", BigInteger),
Column("token_unlimited", Boolean, nullable=False, server_default="false"),
Column("cost_limit_usd", Numeric(12, 4)),
Column("cost_unlimited", Boolean, nullable=False, server_default="false"),
Column("enabled", Boolean, nullable=False, server_default="true"),
Column("note", Text),
Column("created_by", Text),
Column("updated_by", Text),
Column("created_at", DateTime(timezone=True), nullable=False, server_default=func.now()),
Column("updated_at", DateTime(timezone=True), nullable=False, server_default=func.now()),
CheckConstraint("scope IN ('instance', 'team', 'user')", name="quota_policies_scope_check"),
CheckConstraint("bucket IN ('all', 'direct', 'agent')", name="quota_policies_bucket_check"),
CheckConstraint("token_limit >= 0", name="quota_policies_token_limit_check"),
CheckConstraint("cost_limit_usd >= 0", name="quota_policies_cost_limit_check"),
CheckConstraint("(scope = 'instance') = (subject_id IS NULL)", name="quota_policies_subject_chk"),
CheckConstraint("NOT (token_unlimited AND token_limit IS NOT NULL)", name="quota_policies_token_chk"),
CheckConstraint("NOT (cost_unlimited AND cost_limit_usd IS NOT NULL)", name="quota_policies_cost_chk"),
)
Index(
"quota_policies_subject_uidx",
quota_policies_table.c.scope,
func.coalesce(quota_policies_table.c.subject_id, ""),
quota_policies_table.c.bucket,
unique=True,
)
@@ -37,6 +37,7 @@ class TokenUsageRepository:
timestamp: Optional[datetime] = None,
cached_tokens: Optional[int] = None,
cache_write_tokens: Optional[int] = None,
cost: float = 0.0,
) -> None:
# Attribution guard: the ``token_usage_attribution_chk`` CHECK
# constraint requires at least one of ``user_id`` / ``api_key``
@@ -62,14 +63,14 @@ class TokenUsageRepository:
INSERT INTO token_usage (
user_id, api_key, agent_id,
prompt_tokens, generated_tokens,
cached_tokens, cache_write_tokens,
cached_tokens, cache_write_tokens, cost,
source, request_id, model_id, timestamp
)
VALUES (
:user_id, :api_key,
CAST(:agent_id AS uuid),
:prompt_tokens, :generated_tokens,
:cached_tokens, :cache_write_tokens,
:cached_tokens, :cache_write_tokens, :cost,
:source, :request_id, :model_id, COALESCE(:timestamp, now())
)
"""
@@ -82,6 +83,7 @@ class TokenUsageRepository:
"generated_tokens": generated_tokens,
"cached_tokens": cached_tokens,
"cache_write_tokens": cache_write_tokens,
"cost": cost,
"source": source,
"request_id": request_id,
"model_id": model_id,
+24 -1
View File
@@ -2,6 +2,7 @@ import logging
import time
from typing import Any, Dict
from docsgpt.pricing import compute_cost_usd
from docsgpt.storage.db.repositories.token_usage import TokenUsageRepository
from docsgpt.storage.db.session import db_session
from docsgpt.utils import num_tokens_from_object_or_list, num_tokens_from_string
@@ -119,6 +120,12 @@ def _persist_call_usage(llm, call_usage):
},
)
return
model_id = getattr(llm, "_canonical_model_id", None)
# Bring-your-own models run on the user's own provider key: recorded, never priced.
if getattr(llm, "_is_byom", False):
cost = 0.0
else:
cost = _call_cost_usd(model_id, call_usage)
try:
with db_session() as conn:
# ``timestamp`` is omitted so Postgres ``server_default
@@ -136,16 +143,32 @@ def _persist_call_usage(llm, call_usage):
# "0% cache hits".
cached_tokens=call_usage.get("cached_tokens"),
cache_write_tokens=call_usage.get("cache_write_tokens"),
cost=cost,
source=(
getattr(llm, "_token_usage_source", None) or "agent_stream"
),
request_id=getattr(llm, "_request_id", None),
model_id=getattr(llm, "_canonical_model_id", None),
model_id=model_id,
)
except Exception:
logger.exception("token_usage persist failed")
def _call_cost_usd(model_id, call_usage) -> float:
"""Price one call; a pricing failure records $0 rather than dropping the row."""
try:
return compute_cost_usd(
model_id,
call_usage["prompt_tokens"],
call_usage["generated_tokens"],
cached_tokens=call_usage.get("cached_tokens"),
cache_write_tokens=call_usage.get("cache_write_tokens"),
)
except Exception:
logger.exception("token_usage cost computation failed")
return 0.0
def _prefer_provider_usage(llm: Any, call_usage: Dict[str, int]) -> Dict[str, int]:
"""Replace estimates with upstream counts when a provider reported them.
+32
View File
@@ -1015,6 +1015,38 @@ class TestLLMCreatorPassesModelUserId:
assert captured["model_user_id"] == "owner-alice"
@pytest.mark.parametrize(
"model_id, source, expected",
[(None, None, False), ("catalog-model", "builtin", False), ("byom-uuid", "user", True)],
)
def test_byom_flag_follows_the_model_source(self, monkeypatch, model_id, source, expected):
from types import SimpleNamespace
from docsgpt.llm.llm_creator import LLMCreator
from docsgpt.llm.providers import PROVIDERS_BY_NAME
class _LLM:
def __init__(self, *args, **kwargs):
pass
monkeypatch.setattr(PROVIDERS_BY_NAME["openai"], "llm_class", _LLM)
model = SimpleNamespace(
source=source, api_key="own-key", base_url=None, upstream_model_id=None, capabilities=None
)
registry = SimpleNamespace(get_model=lambda _id, user_id=None: model)
monkeypatch.setattr(
"docsgpt.core.model_registry.ModelRegistry.get_instance", lambda: registry
)
llm = LLMCreator.create_llm(
type="openai", api_key="k", user_api_key=None,
decoded_token={"sub": "u1"}, model_id=model_id,
)
assert llm._is_byom is expected
assert llm._canonical_model_id == model_id
# Tests — responding-provider tracking (cross-provider fallback handler fix)
@@ -35,6 +35,15 @@ class TestInsert:
)
assert total == 30
def test_cost_defaults_to_zero_and_round_trips(self, pg_conn):
repo = _repo(pg_conn)
repo.insert(user_id="u-cost", prompt_tokens=1, generated_tokens=1)
repo.insert(user_id="u-cost", prompt_tokens=1, generated_tokens=1, cost=0.00012345)
costs = pg_conn.execute(
text("SELECT cost FROM token_usage WHERE user_id = 'u-cost' ORDER BY id")
).scalars().all()
assert [float(c) for c in costs] == [0.0, 0.00012345]
class TestReassignApiKey:
def test_rewrites_and_preserves_rate_limit_window(self, pg_conn):
+145
View File
@@ -0,0 +1,145 @@
"""Migration round-trip test for 0033_quotas."""
from __future__ import annotations
import os
import subprocess
import sys
from pathlib import Path
import pytest
from sqlalchemy import text
from sqlalchemy.exc import IntegrityError
pytestmark = pytest.mark.integration
def _alembic_ini() -> Path:
return Path(__file__).resolve().parents[3] / "docsgpt" / "alembic.ini"
def _run_alembic(url: str, *args: str) -> None:
subprocess.check_call(
[sys.executable, "-m", "alembic", "-c", str(_alembic_ini()), *args],
timeout=60,
env={**os.environ, "POSTGRES_URI": url},
)
def _alembic_heads(url: str) -> list[str]:
out = subprocess.check_output(
[sys.executable, "-m", "alembic", "-c", str(_alembic_ini()), "heads"],
timeout=60,
env={**os.environ, "POSTGRES_URI": url},
text=True,
)
return [line for line in out.splitlines() if line.strip()]
def _alembic_version(conn) -> str:
return conn.execute(text("SELECT version_num FROM alembic_version")).scalar()
def _column_exists(conn, table: str, column: str) -> bool:
row = conn.execute(
text(
"SELECT 1 FROM information_schema.columns "
"WHERE table_name = :t AND column_name = :c AND table_schema = 'public'"
),
{"t": table, "c": column},
).fetchone()
return row is not None
_0033 = "0033_quotas"
_0032 = "0032_personal_access_tokens"
def _table_exists(conn, table: str) -> bool:
return conn.execute(text("SELECT to_regclass(:t)"), {"t": f"public.{table}"}).scalar() is not None
def _insert_policy(conn, **values) -> None:
cols = ", ".join(values)
params = ", ".join(f":{k}" for k in values)
conn.execute(text(f"INSERT INTO quota_policies ({cols}) VALUES ({params})"), values)
class TestMigration0033RoundTrip:
def test_single_head(self, pg_engine):
url = pg_engine.url.render_as_string(hide_password=False)
assert len(_alembic_heads(url)) == 1
def test_head_has_quota_schema(self, pg_engine):
with pg_engine.connect() as conn:
assert _alembic_version(conn) >= _0033
assert _table_exists(conn, "quota_policies")
assert _column_exists(conn, "token_usage", "cost")
def test_downgrade_drops_then_upgrade_restores(self, pg_engine):
url = pg_engine.url.render_as_string(hide_password=False)
_run_alembic(url, "downgrade", _0032)
with pg_engine.connect() as conn:
assert _alembic_version(conn) == _0032
assert not _table_exists(conn, "quota_policies")
assert not _column_exists(conn, "token_usage", "cost")
_run_alembic(url, "upgrade", "head")
with pg_engine.connect() as conn:
assert _table_exists(conn, "quota_policies")
assert _column_exists(conn, "token_usage", "cost")
def test_upgrade_tolerates_an_existing_cost_column(self, pg_engine):
"""A database that already carries ``token_usage.cost`` upgrades cleanly."""
url = pg_engine.url.render_as_string(hide_password=False)
_run_alembic(url, "downgrade", _0032)
with pg_engine.begin() as conn:
conn.execute(text("ALTER TABLE token_usage ADD COLUMN cost NUMERIC(12,8) NOT NULL DEFAULT 0"))
conn.execute(
text(
"INSERT INTO token_usage (user_id, prompt_tokens, generated_tokens, cost) "
"VALUES ('u-mig33', 10, 1, 0.5)"
)
)
_run_alembic(url, "upgrade", "head")
with pg_engine.connect() as conn:
cost = conn.execute(text("SELECT cost FROM token_usage WHERE user_id = 'u-mig33'")).scalar()
assert float(cost) == 0.5
class TestQuotaPolicyConstraints:
def test_one_row_per_subject_and_bucket(self, pg_conn):
_insert_policy(pg_conn, scope="instance", token_limit=10)
with pytest.raises(IntegrityError):
with pg_conn.begin_nested():
_insert_policy(pg_conn, scope="instance", token_limit=20)
_insert_policy(pg_conn, scope="instance", bucket="agent", token_limit=20)
@pytest.mark.parametrize(
"values",
[
{"scope": "instance", "subject_id": "u1"},
{"scope": "user"},
{"scope": "user", "subject_id": "u1", "token_limit": 5, "token_unlimited": True},
{"scope": "user", "subject_id": "u1", "cost_limit_usd": 5, "cost_unlimited": True},
{"scope": "user", "subject_id": "u1", "token_limit": -1},
{"scope": "user", "subject_id": "u1", "bucket": "nope"},
{"scope": "org", "subject_id": "u1"},
],
)
def test_invalid_rows_rejected(self, pg_conn, values):
with pytest.raises(IntegrityError):
with pg_conn.begin_nested():
_insert_policy(pg_conn, **values)
def test_deleting_a_team_removes_its_policies(self, pg_conn):
team_id = pg_conn.execute(
text("INSERT INTO teams (name, slug, owner_id) VALUES ('Q', 'q-mig33', 'owner') RETURNING id")
).scalar()
_insert_policy(pg_conn, scope="team", subject_id=str(team_id), token_limit=10)
_insert_policy(pg_conn, scope="user", subject_id=str(team_id), token_limit=10)
pg_conn.execute(text("DELETE FROM teams WHERE id = :id"), {"id": team_id})
scopes = pg_conn.execute(
text("SELECT scope FROM quota_policies WHERE subject_id = :id"), {"id": str(team_id)}
).scalars().all()
assert scopes == ["user"]
+53
View File
@@ -712,3 +712,56 @@ def test_decorator_omits_cache_bins_when_provider_reports_none(monkeypatch):
assert row["cache_write_tokens"] is None
assert llm.emitted[0]["cached_tokens"] is None
assert llm.emitted[0]["cache_write_tokens"] is None
def _persist_with_cost(monkeypatch, llm, cost_fn):
from docsgpt.usage import _persist_call_usage
_install_fake_token_repo(monkeypatch)
monkeypatch.setattr("docsgpt.usage.compute_cost_usd", cost_fn)
_persist_call_usage(
llm, {"prompt_tokens": 1000, "generated_tokens": 10, "cached_tokens": 400}
)
return _FakeTokenUsageRepo.last_instance.inserted[0]
class _CostLLM:
decoded_token = {"sub": "user_123"}
user_api_key = None
agent_id = None
_canonical_model_id = "priced-model"
@pytest.mark.unit
def test_persist_prices_the_call_by_canonical_model(monkeypatch):
seen = {}
def cost_fn(model, prompt, generated, cached_tokens=None, cache_write_tokens=None):
seen.update(model=model, prompt=prompt, generated=generated, cached=cached_tokens)
return 0.0123
row = _persist_with_cost(monkeypatch, _CostLLM(), cost_fn)
assert row["cost"] == 0.0123
assert seen == {"model": "priced-model", "prompt": 1000, "generated": 10, "cached": 400}
@pytest.mark.unit
def test_persist_records_byom_calls_at_zero_cost(monkeypatch):
llm = _CostLLM()
llm._is_byom = True
row = _persist_with_cost(monkeypatch, llm, lambda *a, **k: 9.9)
assert row["cost"] == 0.0
assert row["prompt_tokens"] == 1000
@pytest.mark.unit
def test_persist_keeps_the_row_when_pricing_fails(monkeypatch):
def boom(*args, **kwargs):
raise RuntimeError("registry unavailable")
row = _persist_with_cost(monkeypatch, _CostLLM(), boom)
assert row["cost"] == 0.0