feat(quotas): limit resolution and the quota service

docsgpt/quotas resolves a user's effective token and cost limits from the
policy rows that apply to them: their own override, then the most generous
allowance among their teams, then the instance default, then any defaults a
registered provider supplies. Each budget resolves on its own, and a team
counts once however many memberships the user holds in it.

QuotaService compares those limits with the user's token_usage totals over
the current QUOTA_PERIOD window (calendar-aligned, UTC, computed at read
time). It fails open, and skips the usage query for unlimited users.
This commit is contained in:
Alex committed 2026-09-21 12:11:11 +01:00
1 parent 6c42139224
commit 6826313b60
12 files changed
+1175

No files matched your search

+24
View File
@@ -0,0 +1,24 @@
"""Admin-set usage quotas.
Limits live in ``quota_policies`` at three layers (instance default, team
per-member allowance, user override), each with a token budget and a USD
budget. ``QuotaService`` resolves a user's effective limits and compares them
with their ``token_usage`` totals over the current ``QUOTA_PERIOD`` window.
"""
from docsgpt.quotas.providers import QuotaDefaultsProvider, register_defaults_provider
from docsgpt.quotas.resolver import ResolvedLimit, ResolvedLimits, resolve_limits
from docsgpt.quotas.service import BucketStatus, QuotaExceeded, QuotaService
from docsgpt.quotas.windows import window_bounds
__all__ = [
"BucketStatus",
"QuotaDefaultsProvider",
"QuotaExceeded",
"QuotaService",
"ResolvedLimit",
"ResolvedLimits",
"register_defaults_provider",
"resolve_limits",
"window_bounds",
]
+51
View File
@@ -0,0 +1,51 @@
"""Extension point for limits that do not come from ``quota_policies``.
A deployment can register a provider that supplies per-user default policies
(for example from a subscription plan) and adjusts the quota error payload.
Defaults sit below every stored layer: a stored instance, team or user row
with an opinion always wins.
"""
from __future__ import annotations
from typing import Mapping
class QuotaDefaultsProvider:
"""Base provider: no defaults, error payload unchanged."""
def default_policies(self, user_id: str) -> list[dict]:
"""Return default policy rows for ``user_id``.
Each row uses the ``quota_policies`` field names (``bucket``,
``token_limit``, ``token_unlimited``, ``cost_limit_usd``,
``cost_unlimited``); ``scope`` is set by the caller.
"""
return []
def error_payload(self, payload: dict, user_id: str) -> dict:
"""Return the payload sent to a client whose quota is exhausted."""
return payload
_provider: QuotaDefaultsProvider = QuotaDefaultsProvider()
def register_defaults_provider(provider: QuotaDefaultsProvider) -> None:
"""Replace the process-wide defaults provider."""
global _provider
_provider = provider
def get_defaults_provider() -> QuotaDefaultsProvider:
"""Return the registered defaults provider."""
return _provider
def default_rows(user_id: str) -> list[dict]:
"""Return the provider's defaults for ``user_id`` as ``default``-layer rows."""
rows: list[dict] = []
for row in get_defaults_provider().default_policies(user_id) or []:
if isinstance(row, Mapping):
rows.append({"bucket": "all", "enabled": True, **row, "scope": "default", "subject_id": None})
return rows
+99
View File
@@ -0,0 +1,99 @@
"""Resolve the policy rows that apply to a user into effective limits.
Each budget (tokens, cost) resolves on its own: the user's row wins, then the
most generous of the user's team rows, then the instance row, then the
registered defaults. A row with neither a limit nor the unlimited flag for a
budget has no opinion on it and is skipped.
Teams resolve to the most generous allowance because team membership is not
controlled by the instance admin: under "most restrictive", any team admin
could throttle a user by adding them to a low-allowance team.
"""
from __future__ import annotations
from dataclasses import dataclass
from typing import Iterable, Mapping, Optional
LAYERS = ("user", "team", "instance", "default")
_FIELDS = {"tokens": ("token_limit", "token_unlimited"), "cost": ("cost_limit_usd", "cost_unlimited")}
@dataclass(frozen=True)
class ResolvedLimit:
"""One budget's effective limit and the layer it came from.
``limit`` is ``None`` when the budget is unlimited. ``source`` is ``None``
when no layer had an opinion, otherwise one of ``LAYERS``; ``source_id`` is
the team id for a team-sourced limit.
"""
limit: Optional[float] = None
source: Optional[str] = None
source_id: Optional[str] = None
@property
def unlimited(self) -> bool:
return self.limit is None
@dataclass(frozen=True)
class ResolvedLimits:
"""A user's effective token and cost limits for one bucket."""
tokens: ResolvedLimit
cost: ResolvedLimit
@property
def unlimited(self) -> bool:
return self.tokens.unlimited and self.cost.unlimited
def _opinion(row: Mapping, budget: str) -> Optional[tuple[bool, Optional[float]]]:
"""Return ``(unlimited, limit)`` for a row's budget, or ``None`` if it defers."""
limit_field, unlimited_field = _FIELDS[budget]
if row.get(unlimited_field):
return True, None
value = row.get(limit_field)
if value is None:
return None
return False, float(value)
def _resolve_budget(rows: Iterable[Mapping], budget: str) -> ResolvedLimit:
by_layer: dict[str, list[tuple[Mapping, tuple[bool, Optional[float]]]]] = {}
for row in rows:
opinion = _opinion(row, budget)
if opinion is not None:
by_layer.setdefault(row["scope"], []).append((row, opinion))
for layer in LAYERS:
candidates = by_layer.get(layer)
if not candidates:
continue
# Most generous first: unlimited, then the larger limit. Only the team
# layer can hold more than one candidate. Ties break on subject id so
# the reported source is stable.
row, (unlimited, limit) = min(
candidates,
key=lambda c: (not c[1][0], -(c[1][1] or 0.0), str(c[0].get("subject_id") or "")),
)
source_id = str(row["subject_id"]) if layer == "team" else None
return ResolvedLimit(limit=None if unlimited else limit, source=layer, source_id=source_id)
return ResolvedLimit()
def resolve_limits(rows: Iterable[Mapping], bucket: str = "all") -> ResolvedLimits:
"""Return the effective limits for ``bucket`` from a user's applicable rows.
Args:
rows: Policy rows that apply to the user (their own, their teams', the
instance's and any provider defaults). Disabled rows and rows for
other buckets are ignored.
bucket: The policy bucket to resolve.
"""
applicable = [r for r in rows if r.get("bucket", "all") == bucket and r.get("enabled", True)]
return ResolvedLimits(
tokens=_resolve_budget(applicable, "tokens"),
cost=_resolve_budget(applicable, "cost"),
)
+174
View File
@@ -0,0 +1,174 @@
"""Compare a user's usage with their effective limits."""
from __future__ import annotations
import logging
from dataclasses import dataclass
from datetime import datetime
from typing import Optional
from docsgpt.core.settings import settings
from docsgpt.quotas.providers import default_rows, get_defaults_provider
from docsgpt.quotas.resolver import ResolvedLimit, ResolvedLimits, resolve_limits
from docsgpt.quotas.windows import window_bounds
from docsgpt.storage.db.repositories.quota_policies import QuotaPoliciesRepository
from docsgpt.storage.db.repositories.token_usage import TokenUsageRepository
from docsgpt.storage.db.session import db_readonly
logger = logging.getLogger(__name__)
REQUEST_BUCKETS = ("direct", "agent")
_UNITS = {"tokens": "tokens", "cost": "USD"}
@dataclass(frozen=True)
class BucketStatus:
"""A user's limits and usage for one policy bucket in the current window."""
bucket: str
limits: ResolvedLimits
tokens_used: int
cost_used: float
resets_at: datetime
def exceeded_budget(self) -> Optional[str]:
"""Return ``tokens`` or ``cost`` when that budget is used up, else ``None``."""
if not self.limits.tokens.unlimited and self.tokens_used >= self.limits.tokens.limit:
return "tokens"
if not self.limits.cost.unlimited and self.cost_used >= self.limits.cost.limit:
return "cost"
return None
def to_dict(self) -> dict:
"""Return the JSON shape shared by the admin and user quota endpoints."""
def budget(limit: ResolvedLimit, used: float) -> dict:
return {
"limit": limit.limit,
"used": used,
"source": limit.source,
"source_id": limit.source_id,
}
return {
"bucket": self.bucket,
"tokens": budget(self.limits.tokens, self.tokens_used),
"cost": budget(self.limits.cost, round(self.cost_used, 6)),
"resets_at": self.resets_at.isoformat(),
}
@dataclass(frozen=True)
class QuotaExceeded:
"""The exhausted budget that blocks a request."""
user_id: str
bucket: str
budget: str
usage: float
limit: float
source: Optional[str]
source_id: Optional[str]
resets_at: datetime
@property
def retry_after_seconds(self) -> int:
now = datetime.now(self.resets_at.tzinfo)
return max(int((self.resets_at - now).total_seconds()), 1)
def to_payload(self) -> dict:
"""Return the client-facing error body, after the provider's adjustments."""
unit = _UNITS[self.budget]
if self.budget == "cost":
amounts = f"${self.usage:.2f} of ${self.limit:.2f}"
else:
amounts = f"{int(self.usage):,} of {int(self.limit):,} tokens"
payload = {
"success": False,
"error_code": "quota-exceeded",
"message": f"Usage quota reached ({amounts}). It resets at {self.resets_at.isoformat()}.",
"limit_scope": "user_quota",
"dimension": self.budget,
"unit": unit,
"usage": round(self.usage, 6) if self.budget == "cost" else int(self.usage),
"limit": self.limit if self.budget == "cost" else int(self.limit),
"bucket": self.bucket,
"source": self.source,
"resets_at": self.resets_at.isoformat(),
}
try:
return get_defaults_provider().error_payload(payload, self.user_id) or payload
except Exception:
logger.exception("quota defaults provider failed to build the error payload")
return payload
class QuotaService:
"""Resolve limits and measure usage for the current ``QUOTA_PERIOD`` window."""
@staticmethod
def status(
user_id: str,
buckets: tuple[str, ...] = ("all",),
now: Optional[datetime] = None,
) -> list[BucketStatus]:
"""Return the user's status for each of ``buckets``.
Usage is only summed for buckets that carry a limit, so a user with no
applicable policy costs one policy lookup and no usage query.
"""
start, resets_at = window_bounds(settings.QUOTA_PERIOD, now)
statuses: list[BucketStatus] = []
with db_readonly() as conn:
rows = QuotaPoliciesRepository(conn).policies_for_user(user_id) + default_rows(user_id)
usage_repo = TokenUsageRepository(conn)
for bucket in buckets:
limits = resolve_limits(rows, bucket)
tokens_used, cost_used = (0, 0.0)
if not limits.unlimited:
tokens_used, cost_used = usage_repo.usage_totals(user_id=user_id, start=start, bucket=bucket)
statuses.append(BucketStatus(bucket, limits, tokens_used, cost_used, resets_at))
return statuses
@classmethod
def check(
cls, user_id: Optional[str], bucket: str = "direct", now: Optional[datetime] = None
) -> Optional[QuotaExceeded]:
"""Return why ``user_id`` may not start a request, or ``None`` if they may.
Both the ``all`` policies and the request's own bucket must have room.
The check runs before the request, so the call that crosses a limit
completes and the next one is refused. Any failure here allows the
request: a quota outage must not take chat down.
Args:
user_id: The billable user. Requests with no user are not limited.
bucket: ``direct`` for chat without an agent, ``agent`` for traffic through one.
now: Reference instant, for tests.
"""
if not user_id:
return None
if bucket not in REQUEST_BUCKETS:
raise ValueError(f"unknown request bucket: {bucket!r}")
try:
statuses = cls.status(user_id, ("all", bucket), now)
except Exception:
logger.exception("quota check failed; allowing the request", extra={"user_id": user_id})
return None
for status in statuses:
budget = status.exceeded_budget()
if budget is None:
continue
limit = status.limits.tokens if budget == "tokens" else status.limits.cost
return QuotaExceeded(
user_id=user_id,
bucket=status.bucket,
budget=budget,
usage=status.tokens_used if budget == "tokens" else status.cost_used,
limit=limit.limit,
source=limit.source,
source_id=limit.source_id,
resets_at=status.resets_at,
)
return None
+37
View File
@@ -0,0 +1,37 @@
"""Calendar-aligned UTC quota windows, computed at read time (no reset job)."""
from __future__ import annotations
from datetime import datetime, timedelta, timezone
from typing import Optional
PERIODS = ("day", "week", "month")
def window_bounds(period: str, now: Optional[datetime] = None) -> tuple[datetime, datetime]:
"""Return ``(start, resets_at)`` of the window containing ``now``.
Args:
period: ``day`` (from 00:00), ``week`` (from Monday) or ``month`` (from the 1st).
now: Reference instant; defaults to the current time. Naive values are read as UTC.
Raises:
ValueError: If ``period`` is not one of ``PERIODS``.
"""
if now is None:
now = datetime.now(timezone.utc)
elif now.tzinfo is None:
now = now.replace(tzinfo=timezone.utc)
else:
now = now.astimezone(timezone.utc)
midnight = now.replace(hour=0, minute=0, second=0, microsecond=0)
if period == "day":
return midnight, midnight + timedelta(days=1)
if period == "week":
start = midnight - timedelta(days=midnight.weekday())
return start, start + timedelta(days=7)
if period == "month":
start = midnight.replace(day=1)
next_month = (start + timedelta(days=32)).replace(day=1)
return start, next_month
raise ValueError(f"unknown quota period: {period!r}")
@@ -0,0 +1,184 @@
"""Repository for the ``quota_policies`` table.
One row per ``(scope, subject, bucket)``: the instance default (no subject), a
team's per-member allowance (``teams.id``) or a user's override (auth ``sub``).
All methods take a ``Connection`` and do not manage their own transactions.
"""
from __future__ import annotations
from typing import Optional
from sqlalchemy import Connection, text
from docsgpt.storage.db.base_repository import row_to_dict
SCOPES = ("instance", "team", "user")
BUCKETS = ("all", "direct", "agent")
_COLUMNS = (
"id, scope, subject_id, bucket, token_limit, token_unlimited, cost_limit_usd, "
"cost_unlimited, enabled, note, created_by, updated_by, created_at, updated_at"
)
def _validate(scope: str, subject_id: Optional[str], bucket: str) -> None:
if scope not in SCOPES:
raise ValueError(f"unknown quota scope: {scope!r}")
if bucket not in BUCKETS:
raise ValueError(f"unknown quota bucket: {bucket!r}")
if (scope == "instance") != (subject_id is None):
raise ValueError("subject_id is required for team and user policies and must be omitted for instance")
class QuotaPoliciesRepository:
"""Admin-set usage limits."""
def __init__(self, conn: Connection) -> None:
self._conn = conn
# ------------------------------------------------------------------
# Reads
# ------------------------------------------------------------------
def policies_for_user(self, user_id: str) -> list[dict]:
"""Return every enabled row that applies to ``user_id``.
That is the instance rows, the user's own rows, and the rows of each
team the user belongs to (once per team, whatever roles or sources
the membership has).
"""
result = self._conn.execute(
text(
f"""
SELECT {_COLUMNS} FROM quota_policies
WHERE enabled AND (
scope = 'instance'
OR (scope = 'user' AND subject_id = :user_id)
OR (scope = 'team' AND subject_id IN (
SELECT DISTINCT team_id::text FROM team_members WHERE user_id = :user_id
))
)
"""
),
{"user_id": user_id},
)
return [row_to_dict(row) for row in result.fetchall()]
def get(self, scope: str, subject_id: Optional[str], bucket: str = "all") -> Optional[dict]:
"""Return one policy row, or ``None``."""
_validate(scope, subject_id, bucket)
row = self._conn.execute(
text(
f"SELECT {_COLUMNS} FROM quota_policies "
"WHERE scope = :scope AND COALESCE(subject_id, '') = :subject AND bucket = :bucket"
),
{"scope": scope, "subject": subject_id or "", "bucket": bucket},
).fetchone()
return row_to_dict(row) if row is not None else None
def list_for_subject(self, scope: str, subject_id: Optional[str]) -> list[dict]:
"""Return a subject's rows across buckets, ``all`` first."""
_validate(scope, subject_id, "all")
result = self._conn.execute(
text(
f"SELECT {_COLUMNS} FROM quota_policies "
"WHERE scope = :scope AND COALESCE(subject_id, '') = :subject "
"ORDER BY array_position(ARRAY['all', 'direct', 'agent'], bucket)"
),
{"scope": scope, "subject": subject_id or ""},
)
return [row_to_dict(row) for row in result.fetchall()]
def list_by_scope(self, scope: str) -> list[dict]:
"""Return every row of a scope, ordered by subject then bucket."""
if scope not in SCOPES:
raise ValueError(f"unknown quota scope: {scope!r}")
result = self._conn.execute(
text(
f"SELECT {_COLUMNS} FROM quota_policies WHERE scope = :scope "
"ORDER BY subject_id NULLS FIRST, array_position(ARRAY['all', 'direct', 'agent'], bucket)"
),
{"scope": scope},
)
return [row_to_dict(row) for row in result.fetchall()]
# ------------------------------------------------------------------
# Writes
# ------------------------------------------------------------------
def upsert(
self,
*,
scope: str,
subject_id: Optional[str],
bucket: str = "all",
token_limit: Optional[int] = None,
token_unlimited: bool = False,
cost_limit_usd: Optional[float] = None,
cost_unlimited: bool = False,
enabled: bool = True,
note: Optional[str] = None,
actor: Optional[str] = None,
) -> dict:
"""Create or replace the policy for ``(scope, subject_id, bucket)``.
Raises:
ValueError: On an unknown scope or bucket, a subject that does not
match the scope, a negative limit, or a budget that is both
limited and unlimited.
"""
_validate(scope, subject_id, bucket)
if token_unlimited and token_limit is not None:
raise ValueError("token budget cannot be both limited and unlimited")
if cost_unlimited and cost_limit_usd is not None:
raise ValueError("cost budget cannot be both limited and unlimited")
if (token_limit is not None and token_limit < 0) or (cost_limit_usd is not None and cost_limit_usd < 0):
raise ValueError("limits must not be negative")
row = self._conn.execute(
text(
f"""
INSERT INTO quota_policies (
scope, subject_id, bucket, token_limit, token_unlimited,
cost_limit_usd, cost_unlimited, enabled, note, created_by, updated_by
)
VALUES (
:scope, :subject_id, :bucket, :token_limit, :token_unlimited,
:cost_limit_usd, :cost_unlimited, :enabled, :note, :actor, :actor
)
ON CONFLICT (scope, COALESCE(subject_id, ''), bucket) DO UPDATE SET
token_limit = EXCLUDED.token_limit,
token_unlimited = EXCLUDED.token_unlimited,
cost_limit_usd = EXCLUDED.cost_limit_usd,
cost_unlimited = EXCLUDED.cost_unlimited,
enabled = EXCLUDED.enabled,
note = EXCLUDED.note,
updated_by = EXCLUDED.updated_by
RETURNING {_COLUMNS}
"""
),
{
"scope": scope,
"subject_id": subject_id,
"bucket": bucket,
"token_limit": token_limit,
"token_unlimited": token_unlimited,
"cost_limit_usd": cost_limit_usd,
"cost_unlimited": cost_unlimited,
"enabled": enabled,
"note": note,
"actor": actor,
},
).one()
return row_to_dict(row)
def delete(self, scope: str, subject_id: Optional[str], bucket: Optional[str] = None) -> int:
"""Delete a subject's policy for ``bucket``, or all of them; return the count."""
_validate(scope, subject_id, bucket or "all")
clauses = ["scope = :scope", "COALESCE(subject_id, '') = :subject"]
params = {"scope": scope, "subject": subject_id or ""}
if bucket is not None:
clauses.append("bucket = :bucket")
params["bucket"] = bucket
result = self._conn.execute(
text(f"DELETE FROM quota_policies WHERE {' AND '.join(clauses)}"), params
)
return result.rowcount
View File
Whitespace-only changes.
+47
View File
@@ -0,0 +1,47 @@
"""Tests for docsgpt/quotas/providers.py."""
from __future__ import annotations
import pytest
from docsgpt.quotas import providers
from docsgpt.quotas.providers import QuotaDefaultsProvider, default_rows, register_defaults_provider
from docsgpt.quotas.resolver import resolve_limits
@pytest.fixture(autouse=True)
def _restore_provider():
original = providers.get_defaults_provider()
yield
register_defaults_provider(original)
class _PlanProvider(QuotaDefaultsProvider):
def default_policies(self, user_id):
return [
{"bucket": "agent", "cost_limit_usd": 5.0, "scope": "user"},
{"cost_limit_usd": 10.0},
"ignored",
]
def error_payload(self, payload, user_id):
return {**payload, "error_code": "free-limit-reached"}
@pytest.mark.unit
class TestDefaultsProvider:
def test_base_provider_has_no_defaults(self):
assert default_rows("u1") == []
assert QuotaDefaultsProvider().error_payload({"a": 1}, "u1") == {"a": 1}
def test_rows_are_forced_into_the_default_layer(self):
register_defaults_provider(_PlanProvider())
rows = default_rows("u1")
assert [r["scope"] for r in rows] == ["default", "default"]
assert [r["bucket"] for r in rows] == ["agent", "all"]
assert resolve_limits(rows, "agent").cost.source == "default"
def test_stored_rows_win_over_defaults(self):
register_defaults_provider(_PlanProvider())
stored = {"scope": "instance", "subject_id": None, "bucket": "all", "enabled": True, "cost_limit_usd": 99}
assert resolve_limits(default_rows("u1") + [stored]).cost.limit == 99.0
+124
View File
@@ -0,0 +1,124 @@
"""Tests for docsgpt/quotas/resolver.py."""
from __future__ import annotations
import pytest
from docsgpt.quotas.resolver import ResolvedLimit, resolve_limits
def _row(scope, subject_id=None, **fields):
return {"scope": scope, "subject_id": subject_id, "bucket": "all", "enabled": True, **fields}
@pytest.mark.unit
class TestLayers:
def test_no_rows_is_unlimited(self):
limits = resolve_limits([])
assert limits.unlimited
assert limits.tokens == ResolvedLimit()
def test_user_beats_team_beats_instance_beats_default(self):
rows = [
_row("default", token_limit=1),
_row("instance", token_limit=10),
_row("team", "t1", token_limit=100),
_row("user", "u1", token_limit=5),
]
assert resolve_limits(rows).tokens == ResolvedLimit(5.0, "user")
assert resolve_limits(rows[:3]).tokens == ResolvedLimit(100.0, "team", "t1")
assert resolve_limits(rows[:2]).tokens == ResolvedLimit(10.0, "instance")
assert resolve_limits(rows[:1]).tokens == ResolvedLimit(1.0, "default")
def test_user_override_can_be_stricter_than_the_team(self):
rows = [_row("team", "t1", token_unlimited=True), _row("user", "u1", token_limit=0)]
assert resolve_limits(rows).tokens == ResolvedLimit(0.0, "user")
def test_user_unlimited_lifts_an_instance_limit(self):
rows = [_row("instance", token_limit=10), _row("user", "u1", token_unlimited=True)]
resolved = resolve_limits(rows).tokens
assert resolved.unlimited and resolved.source == "user"
def test_budgets_resolve_independently(self):
rows = [
_row("instance", token_limit=10, cost_limit_usd=1),
_row("user", "u1", cost_limit_usd=25),
]
limits = resolve_limits(rows)
assert limits.tokens == ResolvedLimit(10.0, "instance")
assert limits.cost == ResolvedLimit(25.0, "user")
def test_a_row_with_no_opinion_defers(self):
rows = [_row("instance", token_limit=10), _row("user", "u1", note="vip")]
assert resolve_limits(rows).tokens == ResolvedLimit(10.0, "instance")
def test_zero_is_a_limit_not_unlimited(self):
resolved = resolve_limits([_row("instance", cost_limit_usd=0)]).cost
assert resolved.limit == 0.0 and not resolved.unlimited
@pytest.mark.unit
class TestMultipleTeams:
def test_most_generous_team_wins(self):
rows = [
_row("team", "small", token_limit=100),
_row("team", "big", token_limit=900),
_row("team", "mid", token_limit=500),
]
assert resolve_limits(rows).tokens == ResolvedLimit(900.0, "team", "big")
def test_an_unlimited_team_beats_any_limit(self):
rows = [_row("team", "big", token_limit=10**12), _row("team", "free", token_unlimited=True)]
assert resolve_limits(rows).tokens == ResolvedLimit(None, "team", "free")
def test_allowances_are_not_added_together(self):
rows = [_row("team", "a", token_limit=100), _row("team", "b", token_limit=100)]
assert resolve_limits(rows).tokens.limit == 100.0
def test_equal_teams_report_a_stable_source(self):
rows = [_row("team", "b", token_limit=100), _row("team", "a", token_limit=100)]
assert resolve_limits(rows).tokens.source_id == "a"
assert resolve_limits(list(reversed(rows))).tokens.source_id == "a"
def test_each_budget_can_come_from_a_different_team(self):
rows = [
_row("team", "tok", token_limit=900, cost_limit_usd=1),
_row("team", "usd", token_limit=100, cost_limit_usd=50),
]
limits = resolve_limits(rows)
assert limits.tokens.source_id == "tok"
assert limits.cost.source_id == "usd"
def test_a_team_without_an_opinion_does_not_lift_the_limit(self):
rows = [_row("team", "quiet"), _row("team", "capped", token_limit=100)]
assert resolve_limits(rows).tokens == ResolvedLimit(100.0, "team", "capped")
def test_teams_without_opinions_fall_through_to_instance(self):
rows = [_row("team", "quiet", cost_limit_usd=5), _row("instance", token_limit=10)]
assert resolve_limits(rows).tokens == ResolvedLimit(10.0, "instance")
def test_a_zero_team_does_not_block_a_member_of_a_funded_team(self):
rows = [_row("team", "blocked", token_limit=0), _row("team", "funded", token_limit=50)]
assert resolve_limits(rows).tokens.limit == 50.0
def test_disabled_team_rows_are_ignored(self):
rows = [_row("team", "big", token_limit=900, enabled=False), _row("team", "small", token_limit=100)]
assert resolve_limits(rows).tokens == ResolvedLimit(100.0, "team", "small")
@pytest.mark.unit
class TestBuckets:
def test_only_rows_of_the_bucket_apply(self):
rows = [
_row("instance", token_limit=10),
{**_row("instance", token_limit=3), "bucket": "agent"},
]
assert resolve_limits(rows, "all").tokens.limit == 10.0
assert resolve_limits(rows, "agent").tokens.limit == 3.0
assert resolve_limits(rows, "direct").unlimited
def test_decimal_limits_become_floats(self):
from decimal import Decimal
resolved = resolve_limits([_row("instance", cost_limit_usd=Decimal("12.5000"))]).cost
assert resolved.limit == 12.5 and isinstance(resolved.limit, float)
+237
View File
@@ -0,0 +1,237 @@
"""Tests for QuotaService against a real Postgres instance."""
from __future__ import annotations
from contextlib import contextmanager
from datetime import datetime, timedelta, timezone
import pytest
from sqlalchemy import text
from docsgpt.quotas import providers
from docsgpt.quotas.providers import QuotaDefaultsProvider, register_defaults_provider
from docsgpt.quotas.service import QuotaService
from docsgpt.storage.db.repositories.quota_policies import QuotaPoliciesRepository
from docsgpt.storage.db.repositories.token_usage import TokenUsageRepository
NOW = datetime(2026, 9, 23, 12, tzinfo=timezone.utc)
THIS_MONTH = NOW - timedelta(days=2)
LAST_MONTH = NOW - timedelta(days=40)
@pytest.fixture
def conn(pg_conn, monkeypatch):
@contextmanager
def _readonly():
yield pg_conn
monkeypatch.setattr("docsgpt.quotas.service.db_readonly", _readonly)
monkeypatch.setattr("docsgpt.quotas.service.settings.QUOTA_PERIOD", "month")
return pg_conn
@pytest.fixture(autouse=True)
def _restore_provider():
original = providers.get_defaults_provider()
yield
register_defaults_provider(original)
def _use(conn, user_id="u1", tokens=0, cost=0.0, api_key=None, when=THIS_MONTH, source="agent_stream"):
TokenUsageRepository(conn).insert(
user_id=user_id, api_key=api_key, prompt_tokens=tokens, cost=cost, timestamp=when, source=source
)
def _policy(conn, scope, subject_id=None, **fields):
return QuotaPoliciesRepository(conn).upsert(scope=scope, subject_id=subject_id, **fields)
def _team_with_member(conn, slug, user_id="u1"):
team_id = str(
conn.execute(
text("INSERT INTO teams (name, slug, owner_id) VALUES (:n, :s, 'o') RETURNING id"),
{"n": slug, "s": slug},
).scalar()
)
conn.execute(
text("INSERT INTO team_members (team_id, user_id, role) VALUES (CAST(:t AS uuid), :u, 'team_member')"),
{"t": team_id, "u": user_id},
)
return team_id
class TestCheck:
def test_no_policies_allows(self, conn):
_use(conn, tokens=10**9)
assert QuotaService.check("u1", now=NOW) is None
def test_no_user_allows(self, conn):
assert QuotaService.check(None, now=NOW) is None
def test_under_the_limit_allows(self, conn):
_policy(conn, "instance", token_limit=100)
_use(conn, tokens=99)
assert QuotaService.check("u1", now=NOW) is None
def test_reaching_the_token_limit_blocks(self, conn):
_policy(conn, "instance", token_limit=100)
_use(conn, tokens=100)
exceeded = QuotaService.check("u1", now=NOW)
assert (exceeded.budget, exceeded.usage, exceeded.limit, exceeded.source) == ("tokens", 100, 100.0, "instance")
assert exceeded.resets_at == datetime(2026, 10, 1, tzinfo=timezone.utc)
def test_cost_limit_blocks(self, conn):
_policy(conn, "user", "u1", cost_limit_usd=1.5)
_use(conn, cost=1.0)
_use(conn, cost=0.5)
exceeded = QuotaService.check("u1", now=NOW)
assert (exceeded.budget, exceeded.usage, exceeded.source) == ("cost", 1.5, "user")
def test_zero_limit_blocks_without_usage(self, conn):
_policy(conn, "user", "u1", token_limit=0)
assert QuotaService.check("u1", now=NOW).limit == 0
def test_last_periods_usage_does_not_count(self, conn):
_policy(conn, "instance", token_limit=100)
_use(conn, tokens=500, when=LAST_MONTH)
assert QuotaService.check("u1", now=NOW) is None
def test_other_users_usage_does_not_count(self, conn):
_policy(conn, "instance", token_limit=100)
_use(conn, user_id="u2", tokens=500)
assert QuotaService.check("u1", now=NOW) is None
def test_scheduler_rollups_do_not_count(self, conn):
_policy(conn, "instance", token_limit=100)
_use(conn, tokens=60)
_use(conn, tokens=60, source="schedule")
assert QuotaService.check("u1", now=NOW) is None
def test_user_override_lifts_the_instance_limit(self, conn):
_policy(conn, "instance", token_limit=100)
_policy(conn, "user", "u1", token_unlimited=True)
_use(conn, tokens=10**6)
assert QuotaService.check("u1", now=NOW) is None
def test_period_setting_moves_the_window(self, conn, monkeypatch):
monkeypatch.setattr("docsgpt.quotas.service.settings.QUOTA_PERIOD", "day")
_policy(conn, "instance", token_limit=100)
_use(conn, tokens=500, when=NOW - timedelta(days=1))
assert QuotaService.check("u1", now=NOW) is None
_use(conn, tokens=500, when=NOW - timedelta(hours=1))
assert QuotaService.check("u1", now=NOW).resets_at == datetime(2026, 9, 24, tzinfo=timezone.utc)
def test_unknown_bucket_rejected(self, conn):
with pytest.raises(ValueError):
QuotaService.check("u1", bucket="all", now=NOW)
def test_failure_allows_the_request(self, monkeypatch):
def boom():
raise RuntimeError("db down")
monkeypatch.setattr("docsgpt.quotas.service.db_readonly", boom)
assert QuotaService.check("u1", now=NOW) is None
class TestTeams:
def test_the_most_generous_team_sets_the_allowance(self, conn):
_policy(conn, "team", _team_with_member(conn, "svc-small"), token_limit=100)
big = _team_with_member(conn, "svc-big")
_policy(conn, "team", big, token_limit=1000)
_use(conn, tokens=500)
assert QuotaService.check("u1", now=NOW) is None
_use(conn, tokens=500)
exceeded = QuotaService.check("u1", now=NOW)
assert (exceeded.limit, exceeded.source, exceeded.source_id) == (1000.0, "team", big)
def test_usage_is_one_total_not_one_per_team(self, conn):
for slug in ("svc-a", "svc-b", "svc-c"):
_policy(conn, "team", _team_with_member(conn, slug), token_limit=100)
_use(conn, tokens=100)
assert QuotaService.check("u1", now=NOW).limit == 100.0
def test_a_team_only_covers_its_members(self, conn):
_policy(conn, "instance", token_limit=100)
_policy(conn, "team", _team_with_member(conn, "svc-vip", user_id="u2"), token_unlimited=True)
_use(conn, tokens=100)
assert QuotaService.check("u1", now=NOW).source == "instance"
class TestBuckets:
def test_bucket_policies_only_see_their_traffic(self, conn):
_policy(conn, "instance", bucket="agent", token_limit=100)
_use(conn, tokens=500)
_use(conn, tokens=90, api_key="k")
assert QuotaService.check("u1", "direct", now=NOW) is None
assert QuotaService.check("u1", "agent", now=NOW) is None
_use(conn, tokens=10, api_key="k")
exceeded = QuotaService.check("u1", "agent", now=NOW)
assert (exceeded.bucket, exceeded.usage) == ("agent", 100)
assert QuotaService.check("u1", "direct", now=NOW) is None
def test_the_all_bucket_applies_to_both_kinds_of_traffic(self, conn):
_policy(conn, "instance", token_limit=100)
_use(conn, tokens=60)
_use(conn, tokens=60, api_key="k")
assert QuotaService.check("u1", "direct", now=NOW).bucket == "all"
assert QuotaService.check("u1", "agent", now=NOW).bucket == "all"
class TestProviderDefaults:
class _Plan(QuotaDefaultsProvider):
def default_policies(self, user_id):
return [{"cost_limit_usd": 5.0}] if user_id == "u1" else []
def error_payload(self, payload, user_id):
return {**payload, "error_code": "free-limit-reached"}
def test_defaults_apply_without_stored_rows(self, conn):
register_defaults_provider(self._Plan())
_use(conn, cost=5.0)
exceeded = QuotaService.check("u1", now=NOW)
assert exceeded.source == "default"
assert exceeded.to_payload()["error_code"] == "free-limit-reached"
def test_a_broken_provider_payload_falls_back(self, conn):
class _Broken(self._Plan):
def error_payload(self, payload, user_id):
raise RuntimeError("nope")
register_defaults_provider(_Broken())
_use(conn, cost=5.0)
assert QuotaService.check("u1", now=NOW).to_payload()["error_code"] == "quota-exceeded"
class TestStatusAndPayload:
def test_status_reports_limits_and_usage(self, conn):
_policy(conn, "instance", token_limit=100, cost_limit_usd=2)
_use(conn, tokens=40, cost=0.5)
(status,) = QuotaService.status("u1", now=NOW)
assert status.to_dict() == {
"bucket": "all",
"tokens": {"limit": 100.0, "used": 40, "source": "instance", "source_id": None},
"cost": {"limit": 2.0, "used": 0.5, "source": "instance", "source_id": None},
"resets_at": "2026-10-01T00:00:00+00:00",
}
def test_unlimited_users_skip_the_usage_query(self, conn, monkeypatch):
def fail(*args, **kwargs):
raise AssertionError("usage must not be summed for an unlimited user")
monkeypatch.setattr(TokenUsageRepository, "usage_totals", fail)
(status,) = QuotaService.status("u1", now=NOW)
assert status.limits.unlimited and status.tokens_used == 0
def test_payload_shape(self, conn):
_policy(conn, "user", "u1", cost_limit_usd=1)
_use(conn, cost=1.25)
exceeded = QuotaService.check("u1", now=NOW)
payload = exceeded.to_payload()
assert payload["success"] is False
assert payload["error_code"] == "quota-exceeded"
assert (payload["dimension"], payload["unit"]) == ("cost", "USD")
assert (payload["usage"], payload["limit"]) == (1.25, 1.0)
assert payload["resets_at"] == "2026-10-01T00:00:00+00:00"
assert "$1.25 of $1.00" in payload["message"]
assert exceeded.retry_after_seconds >= 1
+55
View File
@@ -0,0 +1,55 @@
"""Tests for docsgpt/quotas/windows.py."""
from __future__ import annotations
from datetime import datetime, timedelta, timezone
import pytest
from docsgpt.quotas.windows import window_bounds
UTC = timezone.utc
@pytest.mark.unit
class TestWindowBounds:
@pytest.mark.parametrize(
"period, start, end",
[
("day", datetime(2026, 9, 23, tzinfo=UTC), datetime(2026, 9, 24, tzinfo=UTC)),
("week", datetime(2026, 9, 21, tzinfo=UTC), datetime(2026, 9, 28, tzinfo=UTC)),
("month", datetime(2026, 9, 1, tzinfo=UTC), datetime(2026, 10, 1, tzinfo=UTC)),
],
)
def test_midweek(self, period, start, end):
now = datetime(2026, 9, 23, 15, 30, 12, 99, tzinfo=UTC) # a Wednesday
assert window_bounds(period, now) == (start, end)
def test_week_starts_on_monday_itself(self):
monday = datetime(2026, 9, 21, 0, 0, tzinfo=UTC)
assert window_bounds("week", monday)[0] == monday
def test_month_rolls_over_the_year(self):
start, end = window_bounds("month", datetime(2026, 12, 31, 23, 59, tzinfo=UTC))
assert (start, end) == (datetime(2026, 12, 1, tzinfo=UTC), datetime(2027, 1, 1, tzinfo=UTC))
def test_leap_february(self):
start, end = window_bounds("month", datetime(2028, 2, 29, 12, tzinfo=UTC))
assert (end - start).days == 29
def test_other_timezones_are_read_in_utc(self):
tokyo = timezone(timedelta(hours=9))
# 2026-10-01 08:00 in Tokyo is still 2026-09-30 in UTC.
start, _ = window_bounds("month", datetime(2026, 10, 1, 8, tzinfo=tokyo))
assert start == datetime(2026, 9, 1, tzinfo=UTC)
def test_naive_datetimes_are_utc(self):
assert window_bounds("day", datetime(2026, 9, 23, 5))[0] == datetime(2026, 9, 23, tzinfo=UTC)
def test_defaults_to_now(self):
start, end = window_bounds("day")
assert start <= datetime.now(UTC) < end
def test_unknown_period(self):
with pytest.raises(ValueError):
window_bounds("year")
@@ -0,0 +1,143 @@
"""Tests for QuotaPoliciesRepository against a real Postgres instance."""
from __future__ import annotations
import pytest
from sqlalchemy import text
from docsgpt.storage.db.repositories.quota_policies import QuotaPoliciesRepository
def _repo(conn) -> QuotaPoliciesRepository:
return QuotaPoliciesRepository(conn)
def _team(conn, slug: str) -> str:
return str(
conn.execute(
text("INSERT INTO teams (name, slug, owner_id) VALUES (:n, :s, 'owner') RETURNING id"),
{"n": slug, "s": slug},
).scalar()
)
def _member(conn, team_id: str, user_id: str, role: str = "team_member", source: str = "manual") -> None:
conn.execute(
text(
"INSERT INTO team_members (team_id, user_id, role, source) "
"VALUES (CAST(:t AS uuid), :u, :r, :s)"
),
{"t": team_id, "u": user_id, "r": role, "s": source},
)
class TestUpsert:
def test_creates_then_replaces(self, pg_conn):
repo = _repo(pg_conn)
created = repo.upsert(scope="user", subject_id="u1", token_limit=100, note="trial", actor="admin1")
assert created["token_limit"] == 100
assert created["created_by"] == created["updated_by"] == "admin1"
replaced = repo.upsert(scope="user", subject_id="u1", cost_limit_usd=2.5, actor="admin2")
assert replaced["id"] == created["id"]
assert replaced["token_limit"] is None
assert float(replaced["cost_limit_usd"]) == 2.5
assert replaced["note"] is None
assert (replaced["created_by"], replaced["updated_by"]) == ("admin1", "admin2")
def test_instance_row_is_a_singleton_per_bucket(self, pg_conn):
repo = _repo(pg_conn)
repo.upsert(scope="instance", subject_id=None, token_limit=1)
repo.upsert(scope="instance", subject_id=None, token_limit=2)
repo.upsert(scope="instance", subject_id=None, bucket="agent", token_limit=3)
rows = repo.list_by_scope("instance")
assert [(r["bucket"], r["token_limit"]) for r in rows] == [("all", 2), ("agent", 3)]
@pytest.mark.parametrize(
"kwargs",
[
{"scope": "org", "subject_id": "x"},
{"scope": "user", "subject_id": None},
{"scope": "instance", "subject_id": "x"},
{"scope": "user", "subject_id": "u", "bucket": "nope"},
{"scope": "user", "subject_id": "u", "token_limit": 1, "token_unlimited": True},
{"scope": "user", "subject_id": "u", "cost_limit_usd": 1, "cost_unlimited": True},
{"scope": "user", "subject_id": "u", "token_limit": -1},
{"scope": "user", "subject_id": "u", "cost_limit_usd": -0.5},
],
)
def test_rejects_invalid_policies(self, pg_conn, kwargs):
with pytest.raises(ValueError):
_repo(pg_conn).upsert(**kwargs)
class TestReads:
def test_get_and_list_for_subject(self, pg_conn):
repo = _repo(pg_conn)
repo.upsert(scope="user", subject_id="u1", bucket="agent", token_limit=5)
repo.upsert(scope="user", subject_id="u1", token_limit=9)
assert repo.get("user", "u1")["token_limit"] == 9
assert repo.get("user", "u1", "direct") is None
assert [r["bucket"] for r in repo.list_for_subject("user", "u1")] == ["all", "agent"]
assert repo.list_for_subject("user", "nobody") == []
def test_list_by_scope_rejects_unknown_scope(self, pg_conn):
with pytest.raises(ValueError):
_repo(pg_conn).list_by_scope("org")
class TestPoliciesForUser:
def test_collects_instance_user_and_team_rows(self, pg_conn):
repo = _repo(pg_conn)
mine, other = _team(pg_conn, "qp-mine"), _team(pg_conn, "qp-other")
_member(pg_conn, mine, "u1")
repo.upsert(scope="instance", subject_id=None, token_limit=1)
repo.upsert(scope="team", subject_id=mine, token_limit=2)
repo.upsert(scope="team", subject_id=other, token_limit=3)
repo.upsert(scope="user", subject_id="u1", token_limit=4)
repo.upsert(scope="user", subject_id="u2", token_limit=5)
limits = sorted(r["token_limit"] for r in repo.policies_for_user("u1"))
assert limits == [1, 2, 4]
def test_a_team_counts_once_however_many_memberships(self, pg_conn):
repo = _repo(pg_conn)
team = _team(pg_conn, "qp-multi")
_member(pg_conn, team, "u1", "team_member", "manual")
_member(pg_conn, team, "u1", "team_admin", "manual")
_member(pg_conn, team, "u1", "team_member", "oidc_group")
repo.upsert(scope="team", subject_id=team, token_limit=7)
assert [r["token_limit"] for r in repo.policies_for_user("u1")] == [7]
def test_every_team_of_the_user_is_included(self, pg_conn):
repo = _repo(pg_conn)
for slug, limit in (("qp-a", 10), ("qp-b", 20), ("qp-c", 30)):
team = _team(pg_conn, slug)
_member(pg_conn, team, "u1")
repo.upsert(scope="team", subject_id=team, token_limit=limit)
assert sorted(r["token_limit"] for r in repo.policies_for_user("u1")) == [10, 20, 30]
def test_disabled_rows_are_left_out(self, pg_conn):
repo = _repo(pg_conn)
repo.upsert(scope="user", subject_id="u1", token_limit=4, enabled=False)
assert repo.policies_for_user("u1") == []
def test_leaving_a_team_drops_its_allowance(self, pg_conn):
repo = _repo(pg_conn)
team = _team(pg_conn, "qp-leave")
_member(pg_conn, team, "u1")
repo.upsert(scope="team", subject_id=team, token_limit=7)
pg_conn.execute(text("DELETE FROM team_members WHERE user_id = 'u1'"))
assert repo.policies_for_user("u1") == []
class TestDelete:
def test_delete_one_bucket_or_all(self, pg_conn):
repo = _repo(pg_conn)
repo.upsert(scope="user", subject_id="u1", token_limit=1)
repo.upsert(scope="user", subject_id="u1", bucket="agent", token_limit=2)
repo.upsert(scope="user", subject_id="u2", token_limit=3)
assert repo.delete("user", "u1", "agent") == 1
assert repo.delete("user", "u1", "agent") == 0
assert repo.delete("user", "u1") == 1
assert repo.get("user", "u2") is not None