mirror of
https://github.com/tiennm99/DocsGPT.git
synced 2026-10-04 16:13:23 +00:00
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:
1 parent
6c42139224
commit
6826313b60
12 files changed
+1175
No files matched your search
@@ -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",
|
||||
]
|
||||
@@ -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
|
||||
@@ -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"),
|
||||
)
|
||||
@@ -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
|
||||
@@ -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
|
||||
Whitespace-only changes.
@@ -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
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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
|
||||
Reference in new issue
Block a user