diff --git a/docsgpt/quotas/__init__.py b/docsgpt/quotas/__init__.py new file mode 100644 index 00000000..ae1b5f4f --- /dev/null +++ b/docsgpt/quotas/__init__.py @@ -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", +] diff --git a/docsgpt/quotas/providers.py b/docsgpt/quotas/providers.py new file mode 100644 index 00000000..8b38ef97 --- /dev/null +++ b/docsgpt/quotas/providers.py @@ -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 diff --git a/docsgpt/quotas/resolver.py b/docsgpt/quotas/resolver.py new file mode 100644 index 00000000..0f3ebb2e --- /dev/null +++ b/docsgpt/quotas/resolver.py @@ -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"), + ) diff --git a/docsgpt/quotas/service.py b/docsgpt/quotas/service.py new file mode 100644 index 00000000..775f3109 --- /dev/null +++ b/docsgpt/quotas/service.py @@ -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 diff --git a/docsgpt/quotas/windows.py b/docsgpt/quotas/windows.py new file mode 100644 index 00000000..58f5e67c --- /dev/null +++ b/docsgpt/quotas/windows.py @@ -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}") diff --git a/docsgpt/storage/db/repositories/quota_policies.py b/docsgpt/storage/db/repositories/quota_policies.py new file mode 100644 index 00000000..3713fc48 --- /dev/null +++ b/docsgpt/storage/db/repositories/quota_policies.py @@ -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 diff --git a/tests/quotas/__init__.py b/tests/quotas/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/tests/quotas/test_providers.py b/tests/quotas/test_providers.py new file mode 100644 index 00000000..aa51383c --- /dev/null +++ b/tests/quotas/test_providers.py @@ -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 diff --git a/tests/quotas/test_resolver.py b/tests/quotas/test_resolver.py new file mode 100644 index 00000000..4401f0b7 --- /dev/null +++ b/tests/quotas/test_resolver.py @@ -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) diff --git a/tests/quotas/test_service.py b/tests/quotas/test_service.py new file mode 100644 index 00000000..0a365799 --- /dev/null +++ b/tests/quotas/test_service.py @@ -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 diff --git a/tests/quotas/test_windows.py b/tests/quotas/test_windows.py new file mode 100644 index 00000000..2a243ff2 --- /dev/null +++ b/tests/quotas/test_windows.py @@ -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") diff --git a/tests/storage/db/repositories/test_quota_policies.py b/tests/storage/db/repositories/test_quota_policies.py new file mode 100644 index 00000000..25396210 --- /dev/null +++ b/tests/storage/db/repositories/test_quota_policies.py @@ -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