diff --git a/.env-template b/.env-template index 5e868f5e..7f3c05a6 100644 --- a/.env-template +++ b/.env-template @@ -109,3 +109,9 @@ MICROSOFT_AUTHORITY=https://{tenantId}.ciamlogin.com/{tenantId} # PAT_MAX_LIFETIME_DAYS=365 # PAT_ALLOW_NON_EXPIRING=false # PAT_MAX_PER_USER=25 + +# Usage quotas (set limits in Admin → Quotas). Usage is counted per calendar +# day, week or month in UTC. Models without a declared price are recorded at $0 +# unless a fallback [input, output] USD rate per 1M tokens is given. +# QUOTA_PERIOD=month +# QUOTA_UNPRICED_RATE_PER_MILLION=[0.5, 1.5] diff --git a/docs/content/Deploying/Access-Control.mdx b/docs/content/Deploying/Access-Control.mdx index 1c452004..dc238349 100644 --- a/docs/content/Deploying/Access-Control.mdx +++ b/docs/content/Deploying/Access-Control.mdx @@ -88,6 +88,7 @@ Admins get a dashboard backed by a REST surface under `/api/admin` (every endpoi | `GET` | `/api/admin/audit` | Authentication/admin audit feed. | | `GET` | `/api/admin/devices/audit` | Remote-device audit feed. | | `GET` | `/api/admin/teams` | Instance-wide oversight of all teams. | +| `GET` `PUT` `DELETE` | `/api/admin/quotas/...` | [Usage quotas](/Deploying/Usage-Quotas) for the instance, teams and users. | Deactivating a user via the dashboard works for any auth type, while OIDC deployments can also offboard through [SCIM](/Deploying/OIDC-SSO#scim-user-provisioning). Both revoke live sessions immediately. @@ -141,9 +142,10 @@ Sharing rules: ## Audit log -Access-control actions are appended to the `auth_events` table alongside the [authentication events](/Deploying/OIDC-SSO#login-auditing). This includes admin actions — `admin_user_activated` / `admin_user_deactivated`, `admin_sessions_revoked`, `role_granted` / `role_revoked` (with `metadata.source` = `manual` or `oidc_group`) — and team events (`team.create`, `team.member_add`, `team.member_role`, `team.member_remove`, `team.share`, `team.unshare`, `team.transfer_owner`, `team.delete`). The acting admin is recorded in the event metadata. +Access-control actions are appended to the `auth_events` table alongside the [authentication events](/Deploying/OIDC-SSO#login-auditing). This includes admin actions — `admin_user_activated` / `admin_user_deactivated`, `admin_sessions_revoked`, `role_granted` / `role_revoked` (with `metadata.source` = `manual` or `oidc_group`), `quota_policy_set` / `quota_policy_deleted` — and team events (`team.create`, `team.member_add`, `team.member_role`, `team.member_remove`, `team.share`, `team.unshare`, `team.transfer_owner`, `team.delete`). The acting admin is recorded in the event metadata. ## Related - [SSO with OIDC](/Deploying/OIDC-SSO) — sign-in, group allowlists, and the `auth_events` table. +- [Usage Quotas](/Deploying/Usage-Quotas) — token and cost limits per user and per team. - [App Configuration](/Deploying/DocsGPT-Settings) — the full settings reference. diff --git a/docs/content/Deploying/Settings-Reference.mdx b/docs/content/Deploying/Settings-Reference.mdx index 33ea3b53..d9f59ec7 100644 --- a/docs/content/Deploying/Settings-Reference.mdx +++ b/docs/content/Deploying/Settings-Reference.mdx @@ -1451,6 +1451,23 @@ Type `int`, default `30`, must be `>= 1`. Days guardrail events are kept before the cleanup task removes them. +## Quotas + +Quota window and the treatment of unpriced models. + +### `QUOTA_PERIOD` + +Type `"day" | "week" | "month"`, default `month`. + +Window every usage quota is measured over. Windows are calendar-aligned in UTC: a day starts at 00:00, a week on Monday, a month on the 1st. + +### `QUOTA_UNPRICED_RATE_PER_MILLION` + +Type `list[float]`, default unset. + +Fallback `[input, output]` USD rates per 1M tokens for models that declare no price, e.g. `[0.5, 1.5]`. Unset, such calls are recorded at $0 and only count toward token quotas. + + ## Scheduler Cadence, quotas and timeouts of scheduled runs. diff --git a/docs/content/Deploying/Usage-Quotas.mdx b/docs/content/Deploying/Usage-Quotas.mdx new file mode 100644 index 00000000..e4ffbaf1 --- /dev/null +++ b/docs/content/Deploying/Usage-Quotas.mdx @@ -0,0 +1,109 @@ +--- +title: Usage Quotas +description: Cap how many tokens or dollars each user may spend per day, week or month, with an instance default, per-team allowances and per-user overrides. +--- + +import { Callout } from 'nextra/components' + +# Usage Quotas + +An instance admin can limit how much each user spends on language models. A quota has two independent budgets: + +- **Tokens** — prompt plus generated tokens. Works for every model, including local ones. +- **Cost (USD)** — tokens priced at the model's catalog rate. Only sees models that declare a price. + +Set either, both or neither. Quotas are managed from **Admin → Quotas**, or through the [API](#api). With no quota set, nothing is limited. + +## Layers + +Limits are set at three layers. For each budget, the first layer that says something wins: + +1. **User override** — one user's own limit. +2. **Team allowance** — what each member of a team gets. +3. **Instance default** — everyone else. + +At each layer a budget is either *not set* (defer to the next layer), a *limit*, or *unlimited*. A limit of `0` blocks the user. The two budgets resolve separately, so a user's token limit can come from their team while their cost limit comes from the instance default. + +### Teams + +A team allowance is **per member**, not a pool the team shares: if the allowance is 2M tokens, each member may use 2M. + +A user in several teams gets the **most generous** allowance among them, and allowances are never added together. Usage is always counted per user, whichever teams they belong to. To hold one person below their team's allowance, give them a user override. + + +Team membership can change without an instance admin — team admins, OIDC group sync and SCIM all add members — so joining a team can only raise a user's allowance to what you granted that team, never lower it. Only instance admins set allowances; team admins cannot. + + +## Windows and enforcement + +Usage is counted over a calendar window in UTC, chosen for the whole instance with [`QUOTA_PERIOD`](/Deploying/Settings-Reference#quotas): `day` (from 00:00), `week` (from Monday) or `month` (from the 1st, the default). Windows are worked out when a request arrives, so there is no reset job to run. + +The quota is checked **before** a request starts. The request that crosses a limit completes; the next one is refused with HTTP `429`: + +```json +{ + "success": false, + "error_code": "quota-exceeded", + "message": "Usage quota reached (1,000,000 of 1,000,000 tokens). It resets at 2026-10-01T00:00:00+00:00.", + "dimension": "tokens", + "unit": "tokens", + "usage": 1000000, + "limit": 1000000, + "bucket": "all", + "source": "instance", + "resets_at": "2026-10-01T00:00:00+00:00" +} +``` + +The response carries a `Retry-After` header. The check covers chat, the agent and OpenAI-compatible APIs, scheduled runs (recorded as `budget_exceeded`) and webhook runs. If the quota check itself fails, the request is allowed. + +Who is charged: + +| Traffic | Charged to | +| --- | --- | +| Chat without an agent | The user | +| A user's own agent, its API key, webhooks and schedules | The agent's owner | +| An agent shared with the user | The user | + +Per-agent token and request limits still apply on top of the owner's quota. + +Users with a quota see their usage and the reset time under **Settings → Analytics**. + +## Pricing + +Cost budgets use the rates in the [model catalog](/Models/cloud-providers), in USD per million tokens: + +```yaml +models: + - id: my-model + input_cost_per_million: 3.0 + output_cost_per_million: 15.0 + cached_input_cost_per_million: 0.3 # optional, prompt-cache reads + cache_write_cost_per_million: 3.75 # optional, prompt-cache writes +``` + +The built-in catalogs ship list prices for hosted models. Override or add rates by dropping a YAML with the same model `id` into `MODELS_CONFIG_DIR`. The cost of each call is stored with its usage row when the call is made, so later price changes do not rewrite history. + + +A model with no declared price is recorded at $0, so a cost budget cannot see it. The Quotas tab lists such models once they have been used. Either limit them with a token budget, declare their rates, or set [`QUOTA_UNPRICED_RATE_PER_MILLION`](/Deploying/Settings-Reference#quotas) to charge a fallback rate. Models a user adds with their own API key are always $0, but their tokens still count. + + +## API + +Every admin endpoint requires the admin role, and every change is written to the [audit log](/Deploying/Access-Control#audit-log) as `quota_policy_set` or `quota_policy_deleted`. + +| Method | Path | Description | +| --- | --- | --- | +| `GET` | `/api/admin/quotas` | All policies by layer, the current window, and used models without a price. | +| `PUT` `DELETE` | `/api/admin/quotas/instance` | The instance default. | +| `GET` `PUT` `DELETE` | `/api/admin/quotas/teams/` | A team's per-member allowance. | +| `GET` `PUT` `DELETE` | `/api/admin/quotas/users/` | A user's override; the user must already exist (SCIM-provisioned, or signed in once), otherwise `404`. `GET` also returns the limits the user ends up with, the layer each came from, and their usage. | +| `GET` | `/api/user/quota` | The caller's own limits, usage and reset time. | + +A `PUT` body sets, per budget, a limit or the unlimited flag; leave both out to defer to the next layer: + +```json +{ "token_limit": 2000000, "cost_unlimited": true, "note": "Research team" } +``` + +`bucket` (default `all`) narrows a policy to `direct` traffic (chat without an agent) or `agent` traffic (anything that runs through an agent, whether or not the agent has an API key). A request must fit both its own bucket and `all`. The dashboard edits `all`. diff --git a/docs/content/Deploying/_meta.js b/docs/content/Deploying/_meta.js index 105b9efd..9afb460d 100644 --- a/docs/content/Deploying/_meta.js +++ b/docs/content/Deploying/_meta.js @@ -15,6 +15,10 @@ export default { "title": "👥 Access Control & Teams", "href": "/Deploying/Access-Control" }, + "Usage-Quotas": { + "title": "📊 Usage Quotas", + "href": "/Deploying/Usage-Quotas" + }, "Docker-Deploying": { "title": "🛳️ Docker Setup", "href": "/Deploying/Docker-Deploying" diff --git a/docsgpt/agents/headless_runner.py b/docsgpt/agents/headless_runner.py index 59c40c5c..c2a4d713 100644 --- a/docsgpt/agents/headless_runner.py +++ b/docsgpt/agents/headless_runner.py @@ -15,6 +15,7 @@ from docsgpt.api.answer.services.prompt_renderer import ( ) from docsgpt.api.answer.services.stream_processor import get_prompt from docsgpt.core.settings import settings +from docsgpt.quotas.service import QuotaExceededError, QuotaService from docsgpt.retriever.retriever_creator import RetrieverCreator from docsgpt.storage.db.repositories.sources import SourcesRepository from docsgpt.storage.db.session import db_readonly @@ -69,7 +70,11 @@ def run_agent_headless( chat_history: Optional[List[Dict[str, Any]]] = None, conversation_id: Optional[str] = None, ) -> Dict[str, Any]: - """Run an agent with no live client; returns a structured outcome dict.""" + """Run an agent with no live client; returns a structured outcome dict. + + Raises: + QuotaExceededError: If the agent owner's usage quota is exhausted. + """ from docsgpt.core.model_utils import ( get_api_key_for_provider, get_default_model_id, @@ -82,6 +87,11 @@ def run_agent_headless( if not owner: raise ValueError("Agent config is missing user_id; cannot run headless.") decoded_token = {"sub": owner} + # An agent run is agent traffic whether or not the agent has a key yet. + is_agent_run = bool(agent_config.get("key") or _resolve_agent_id(agent_config)) + exceeded = QuotaService.check(owner, "agent" if is_agent_run else "direct") + if exceeded is not None: + raise QuotaExceededError(exceeded) retriever_kind = agent_config.get("retriever", "classic") source_id = agent_config.get("source_id") or agent_config.get("source") diff --git a/docsgpt/agents/workflows/workflow_engine.py b/docsgpt/agents/workflows/workflow_engine.py index 54fa6e67..15876c82 100644 --- a/docsgpt/agents/workflows/workflow_engine.py +++ b/docsgpt/agents/workflows/workflow_engine.py @@ -393,6 +393,8 @@ class WorkflowEngine: "prompt": node_prompt, "chat_history": self.agent.chat_history, "decoded_token": self.agent.decoded_token, + # Attributes the node's token usage to the workflow agent. + "agent_id": getattr(self.agent, "agent_id", None), "json_schema": node_json_schema, "retrieved_docs": node_docs, # A template that interpolates the documents itself already carries diff --git a/docsgpt/alembic/versions/0033_quotas.py b/docsgpt/alembic/versions/0033_quotas.py new file mode 100644 index 00000000..165ae021 --- /dev/null +++ b/docsgpt/alembic/versions/0033_quotas.py @@ -0,0 +1,102 @@ +"""0033 quotas — admin-set usage limits and a per-call cost. + +``quota_policies`` holds the limits an instance admin sets at three layers: +the instance default (``subject_id`` NULL), a team's per-member allowance +(``subject_id`` = ``teams.id``) and a single user's override (``subject_id`` = +the auth ``sub``). Each row carries a token budget and a USD budget; per +budget a row either sets a limit (0 blocks), marks it unlimited, or leaves +both empty to defer to the next layer. ``bucket`` narrows a row to chat without +an agent (``direct``) or traffic through an agent (``agent``); ``all`` covers both. + +``subject_id`` is polymorphic, so there is no FK: an AFTER DELETE trigger on +``teams`` scrubs a deleted team's rows, and user rows follow the ``user_roles`` +convention of never blocking user deletion. + +``token_usage.cost`` is the USD cost of the call at write time (see +``docsgpt/pricing.py``); 0 for unpriced and bring-your-own models. +Idempotent both ways. + +Revision ID: 0033_quotas +Revises: 0032_personal_access_tokens +""" + +from typing import Sequence, Union + +from alembic import op + + +revision: str = "0033_quotas" +down_revision: Union[str, None] = "0032_personal_access_tokens" +branch_labels: Union[str, Sequence[str], None] = None +depends_on: Union[str, Sequence[str], None] = None + + +def upgrade() -> None: + op.execute( + "ALTER TABLE token_usage ADD COLUMN IF NOT EXISTS cost NUMERIC(12,8) NOT NULL DEFAULT 0;" + ) + op.execute( + """ + CREATE TABLE IF NOT EXISTS quota_policies ( + id UUID PRIMARY KEY DEFAULT gen_random_uuid(), + scope TEXT NOT NULL CHECK (scope IN ('instance', 'team', 'user')), + subject_id TEXT, + bucket TEXT NOT NULL DEFAULT 'all' + CHECK (bucket IN ('all', 'direct', 'agent')), + token_limit BIGINT CHECK (token_limit >= 0), + token_unlimited BOOLEAN NOT NULL DEFAULT false, + cost_limit_usd NUMERIC(12,4) CHECK (cost_limit_usd >= 0), + cost_unlimited BOOLEAN NOT NULL DEFAULT false, + enabled BOOLEAN NOT NULL DEFAULT true, + note TEXT, + created_by TEXT, + updated_by TEXT, + created_at TIMESTAMPTZ NOT NULL DEFAULT now(), + updated_at TIMESTAMPTZ NOT NULL DEFAULT now(), + CONSTRAINT quota_policies_subject_chk + CHECK ((scope = 'instance') = (subject_id IS NULL)), + CONSTRAINT quota_policies_token_chk + CHECK (NOT (token_unlimited AND token_limit IS NOT NULL)), + CONSTRAINT quota_policies_cost_chk + CHECK (NOT (cost_unlimited AND cost_limit_usd IS NOT NULL)) + ); + """ + ) + # One row per (layer subject, bucket); the instance row's NULL subject folds to ''. + op.execute( + "CREATE UNIQUE INDEX IF NOT EXISTS quota_policies_subject_uidx " + "ON quota_policies (scope, COALESCE(subject_id, ''), bucket);" + ) + op.execute("DROP TRIGGER IF EXISTS quota_policies_set_updated_at ON quota_policies;") + op.execute( + """ + CREATE TRIGGER quota_policies_set_updated_at + BEFORE UPDATE ON quota_policies + FOR EACH ROW EXECUTE FUNCTION set_updated_at(); + """ + ) + op.execute( + """ + CREATE OR REPLACE FUNCTION cleanup_team_quota_policies() RETURNS trigger AS $$ + BEGIN + DELETE FROM quota_policies WHERE scope = 'team' AND subject_id = OLD.id::text; + RETURN OLD; + END; + $$ LANGUAGE plpgsql; + """ + ) + op.execute("DROP TRIGGER IF EXISTS teams_cleanup_quota_policies ON teams;") + op.execute( + """ + CREATE TRIGGER teams_cleanup_quota_policies + AFTER DELETE ON teams + FOR EACH ROW EXECUTE FUNCTION cleanup_team_quota_policies(); + """ + ) + + +def downgrade() -> None: + op.execute("DROP TRIGGER IF EXISTS teams_cleanup_quota_policies ON teams;") + op.execute("DROP FUNCTION IF EXISTS cleanup_team_quota_policies();") + op.execute("DROP TABLE IF EXISTS quota_policies;") + op.execute("ALTER TABLE token_usage DROP COLUMN IF EXISTS cost;") diff --git a/docsgpt/api/admin/__init__.py b/docsgpt/api/admin/__init__.py index 8fa8b444..7bfc3059 100644 --- a/docsgpt/api/admin/__init__.py +++ b/docsgpt/api/admin/__init__.py @@ -1,3 +1,4 @@ from .routes import admin_ns +from . import quotas # noqa: F401 (registers the quota resources on admin_ns) __all__ = ["admin_ns"] diff --git a/docsgpt/api/admin/quotas.py b/docsgpt/api/admin/quotas.py new file mode 100644 index 00000000..0235f36b --- /dev/null +++ b/docsgpt/api/admin/quotas.py @@ -0,0 +1,273 @@ +"""Admin endpoints for usage quotas (RBAC ``admin`` role required). + +Policies are set at three layers: the instance default, a team's per-member +allowance and a single user's override. Every write is audited to +``auth_events`` with the acting admin recorded. +""" + +from __future__ import annotations + +import math +from typing import Any, Optional + +from flask import jsonify, make_response, request +from flask_restx import Resource + +from docsgpt.api.admin.routes import _actor, admin_ns +from docsgpt.api.user.authz import admin_required +from docsgpt.core.settings import settings +from docsgpt.pricing import is_priced +from docsgpt.quotas.service import REQUEST_BUCKETS, QuotaService +from docsgpt.quotas.windows import window_bounds +from docsgpt.storage.db.base_repository import looks_like_uuid +from docsgpt.storage.db.repositories.auth_events import AuthEventsRepository +from docsgpt.storage.db.repositories.quota_policies import BUCKETS, QuotaPoliciesRepository +from docsgpt.storage.db.repositories.teams import TeamsRepository +from docsgpt.storage.db.repositories.token_usage import TokenUsageRepository +from docsgpt.storage.db.repositories.users import UsersRepository +from docsgpt.storage.db.session import db_readonly, db_session + +_MAX_TOKEN_LIMIT = 2**62 +_MAX_COST_LIMIT = 99_999_999.0 +_MAX_NOTE_LENGTH = 500 +_FLAG_DEFAULTS = {"token_unlimited": False, "cost_unlimited": False, "enabled": True} +_BUCKET_MESSAGE = f"bucket must be one of: {', '.join(BUCKETS)}" + + +def _policy_json(row: dict) -> dict: + cost = row.get("cost_limit_usd") + return { + "scope": row["scope"], + "subject_id": row.get("subject_id"), + "bucket": row["bucket"], + "token_limit": row.get("token_limit"), + "token_unlimited": bool(row.get("token_unlimited")), + "cost_limit_usd": float(cost) if cost is not None else None, + "cost_unlimited": bool(row.get("cost_unlimited")), + "enabled": bool(row.get("enabled", True)), + "note": row.get("note"), + "updated_by": row.get("updated_by"), + "updated_at": row.get("updated_at"), + } + + +def _error(message: str, status: int): + return make_response(jsonify({"success": False, "message": message}), status) + + +def _limit_error(value: Any, name: str, whole: bool, maximum: float) -> Optional[str]: + """Return why ``value`` is not a valid limit, or ``None``.""" + if value is None: + return None + number = (int,) if whole else (int, float) + if isinstance(value, bool) or not isinstance(value, number): + return f"{name} must be a {'whole number' if whole else 'number'} or null" + # Range first: ``isfinite`` overflows on an int too large for a float. + if not 0 <= value <= maximum or not math.isfinite(value): + return f"{name} is out of range" + return None + + +def _parse_policy(data: Any) -> tuple[Optional[dict], Optional[str]]: + """Validate a policy body. + + Returns: + ``(fields, None)`` with ``QuotaPoliciesRepository.upsert`` kwargs, or + ``(None, message)`` describing the first problem. + """ + if not isinstance(data, dict): + return None, "Body must be a JSON object" + token_limit, cost_limit = data.get("token_limit"), data.get("cost_limit_usd") + problem = _limit_error(token_limit, "token_limit", True, _MAX_TOKEN_LIMIT) or _limit_error( + cost_limit, "cost_limit_usd", False, _MAX_COST_LIMIT + ) + if problem: + return None, problem + flags = {key: data.get(key, default) for key, default in _FLAG_DEFAULTS.items()} + for key, value in flags.items(): + if not isinstance(value, bool): + return None, f"{key} must be a boolean" + bucket = data.get("bucket", "all") + if bucket not in BUCKETS: + return None, _BUCKET_MESSAGE + note = data.get("note") + if note is not None and not isinstance(note, str): + return None, "note must be a string" + if flags["token_unlimited"] and token_limit is not None: + return None, "Set token_limit or token_unlimited, not both" + if flags["cost_unlimited"] and cost_limit is not None: + return None, "Set cost_limit_usd or cost_unlimited, not both" + if token_limit is None and cost_limit is None and not flags["token_unlimited"] and not flags["cost_unlimited"]: + return None, "Set a limit or mark a budget unlimited; delete the policy to remove it" + return { + "bucket": bucket, + "token_limit": token_limit, + "token_unlimited": flags["token_unlimited"], + "cost_limit_usd": round(float(cost_limit), 4) if cost_limit is not None else None, + "cost_unlimited": flags["cost_unlimited"], + "enabled": flags["enabled"], + "note": (note.strip()[:_MAX_NOTE_LENGTH] or None) if note else None, + }, None + + +def _audit(conn, event: str, scope: str, subject_id: Optional[str], detail: dict) -> None: + actor = _actor() + AuthEventsRepository(conn).insert( + # A user policy is filed under that user; the rest under the acting admin. + subject_id if scope == "user" else (actor or "unknown"), + event, + ip=request.remote_addr, + user_agent=request.headers.get("User-Agent"), + metadata={"by": actor, "via": "admin_api", "scope": scope, "subject_id": subject_id, **detail}, + ) + + +def _put_policy(scope: str, subject_id: Optional[str]): + fields, problem = _parse_policy(request.get_json(silent=True)) + if fields is None: + return _error(problem or "Invalid policy", 400) + with db_session() as conn: + row = QuotaPoliciesRepository(conn).upsert( + scope=scope, subject_id=subject_id, actor=_actor(), **fields + ) + _audit(conn, "quota_policy_set", scope, subject_id, fields) + return make_response(jsonify({"success": True, "policy": _policy_json(row)}), 200) + + +def _delete_policy(scope: str, subject_id: Optional[str]): + bucket = request.args.get("bucket") + if bucket is not None and bucket not in BUCKETS: + return _error(_BUCKET_MESSAGE, 400) + with db_session() as conn: + deleted = QuotaPoliciesRepository(conn).delete(scope, subject_id, bucket) + if deleted: + _audit(conn, "quota_policy_deleted", scope, subject_id, {"bucket": bucket or "*"}) + return make_response(jsonify({"success": True, "deleted": deleted}), 200) + + +def _unpriced_models(conn) -> list[dict]: + """Models used this period whose calls were all recorded at $0 for want of a price.""" + start, _ = window_bounds(settings.QUOTA_PERIOD) + return [ + row + for row in TokenUsageRepository(conn).tokens_by_model(start=start) + # Judged by what was recorded, so a priced model whose provider has since + # been disabled is not listed. BYOM ids are UUIDs and $0 by design; a + # model explicitly priced at $0 is free, not unpriced. + if row["cost"] == 0 and not looks_like_uuid(row["model_id"]) and not is_priced(row["model_id"]) + ] + + +@admin_ns.route("/admin/quotas") +class AdminQuotasResource(Resource): + @admin_required + def get(self): + """Every stored policy, grouped by layer, plus the models cost limits cannot see.""" + start, resets_at = window_bounds(settings.QUOTA_PERIOD) + with db_readonly() as conn: + repo = QuotaPoliciesRepository(conn) + teams = {str(t["id"]): t for t in TeamsRepository(conn).list_all()} + team_policies = [] + for row in repo.list_by_scope("team"): + team = teams.get(str(row["subject_id"]), {}) + team_policies.append( + { + **_policy_json(row), + "team_name": team.get("name"), + "team_slug": team.get("slug"), + "member_count": team.get("member_count"), + } + ) + body = { + "success": True, + "period": settings.QUOTA_PERIOD, + "period_start": start.isoformat(), + "resets_at": resets_at.isoformat(), + "instance": [_policy_json(r) for r in repo.list_by_scope("instance")], + "teams": team_policies, + "users": [_policy_json(r) for r in repo.list_by_scope("user")], + "unpriced_models": _unpriced_models(conn), + } + return make_response(jsonify(body), 200) + + +@admin_ns.route("/admin/quotas/instance") +class AdminInstanceQuotaResource(Resource): + @admin_required + def put(self): + """Set the instance default for one bucket.""" + return _put_policy("instance", None) + + @admin_required + def delete(self): + """Remove the instance default for ``?bucket=``, or for every bucket.""" + return _delete_policy("instance", None) + + +@admin_ns.route("/admin/quotas/teams/") +class AdminTeamQuotaResource(Resource): + @admin_required + def get(self, team_id): + """The per-member allowance of one team.""" + if not looks_like_uuid(team_id): + return _error("Team not found", 404) + with db_readonly() as conn: + if TeamsRepository(conn).get(team_id) is None: + return _error("Team not found", 404) + rows = QuotaPoliciesRepository(conn).list_for_subject("team", team_id) + return make_response( + jsonify({"success": True, "policies": [_policy_json(r) for r in rows]}), 200 + ) + + @admin_required + def put(self, team_id): + """Set the allowance each member of the team gets.""" + if not looks_like_uuid(team_id): + return _error("Team not found", 404) + with db_readonly() as conn: + if TeamsRepository(conn).get(team_id) is None: + return _error("Team not found", 404) + return _put_policy("team", team_id) + + @admin_required + def delete(self, team_id): + """Remove the team's allowance for ``?bucket=``, or for every bucket.""" + if not looks_like_uuid(team_id): + return _error("Team not found", 404) + return _delete_policy("team", team_id) + + +@admin_ns.route("/admin/quotas/users/") +class AdminUserQuotaResource(Resource): + @admin_required + def get(self, user_id): + """A user's overrides and the limits and usage they resolve to.""" + with db_readonly() as conn: + if UsersRepository(conn).get(user_id) is None: + return _error("User not found", 404) + rows = QuotaPoliciesRepository(conn).list_for_subject("user", user_id) + statuses = QuotaService.status(user_id, ("all", *REQUEST_BUCKETS)) + return make_response( + jsonify( + { + "success": True, + "period": settings.QUOTA_PERIOD, + "policies": [_policy_json(r) for r in rows], + "effective": [s.to_dict() for s in statuses], + } + ), + 200, + ) + + @admin_required + def put(self, user_id): + """Set one user's override, which beats team allowances and the default.""" + with db_readonly() as conn: + if UsersRepository(conn).get(user_id) is None: + return _error("User not found", 404) + return _put_policy("user", user_id) + + @admin_required + def delete(self, user_id): + """Remove the user's override for ``?bucket=``, or for every bucket.""" + return _delete_policy("user", user_id) diff --git a/docsgpt/api/answer/routes/answer.py b/docsgpt/api/answer/routes/answer.py index 17fac152..7cdcd722 100644 --- a/docsgpt/api/answer/routes/answer.py +++ b/docsgpt/api/answer/routes/answer.py @@ -103,7 +103,9 @@ class AnswerResource(Resource, BaseAnswerResource): ) if not processor.decoded_token: return make_response({"error": "Unauthorized"}, 401) - if error := self.check_usage(processor.agent_config): + if error := self.check_usage_on_resume( + processor, data["conversation_id"] + ): return error stream = self.complete_stream( question="", @@ -129,7 +131,11 @@ class AnswerResource(Resource, BaseAnswerResource): if not processor.decoded_token: return make_response({"error": "Unauthorized"}, 401) - if error := self.check_usage(processor.agent_config): + if error := self.check_usage( + processor.agent_config, + processor.decoded_token, + agent_id=processor.agent_id, + ): return error should_persist, visibility = resolve_persistence( diff --git a/docsgpt/api/answer/routes/base.py b/docsgpt/api/answer/routes/base.py index 748f0dda..e0cfe4cd 100644 --- a/docsgpt/api/answer/routes/base.py +++ b/docsgpt/api/answer/routes/base.py @@ -23,6 +23,8 @@ from docsgpt.core.model_utils import ( from docsgpt.core.settings import settings from docsgpt.error import sanitize_api_error from docsgpt.llm.llm_creator import LLMCreator +from docsgpt.quotas.http import quota_exceeded_response +from docsgpt.quotas.service import QuotaService from docsgpt.storage.db.repositories.agents import AgentsRepository from docsgpt.storage.db.repositories.conversations import ( HeartbeatState, @@ -112,17 +114,34 @@ class BaseAnswerResource: prepared.append(item) return prepared - def check_usage(self, agent_config: Dict) -> Optional[Response]: - """Check if there is a usage limit and if it is exceeded + def check_usage( + self, + agent_config: Dict, + decoded_token: Optional[Dict] = None, + agent_id: Optional[str] = None, + ) -> Optional[Response]: + """Refuse the request when a usage limit is exhausted. + + The billable user's quota is checked first, for every request; the + agent's own 24h token and request limits then apply to traffic that + runs through an agent. Args: agent_config: The config dict of agent instance + decoded_token: The request's resolved identity; its ``sub`` is the + billable user. + agent_id: The agent the request runs through. A draft agent has no + key, but its usage rows carry the agent id, so it is agent traffic. Returns: None or Response if either of limits exceeded. """ api_key = agent_config.get("user_api_key") + user_id = (decoded_token or {}).get("sub") or agent_config.get("user_id") + exceeded = QuotaService.check(user_id, "agent" if api_key or agent_id else "direct") + if exceeded is not None: + return quota_exceeded_response(exceeded) if not api_key: return None with db_readonly() as conn: @@ -197,6 +216,34 @@ class BaseAnswerResource: ) return None + def check_usage_on_resume(self, processor: Any, conversation_id: Any) -> Optional[Response]: + """Run ``check_usage`` for a tool continuation, releasing its claim on refusal. + + ``resume_from_tool_actions`` has already claimed the paused turn by the + time the limits can be checked (the agent config comes from the claimed + state). A refusal returns before ``complete_stream`` and its cleanup, so + the claim is released here; otherwise retries get a 409 until the stale + claim is reverted. + + Args: + processor: The ``StreamProcessor`` that resumed the turn. + conversation_id: The conversation whose pending state was claimed. + + Returns: + None, or the refusal Response. + """ + error = self.check_usage( + processor.agent_config, processor.decoded_token, agent_id=processor.agent_id + ) + if error is None or not conversation_id: + return error + user = processor.initial_user_id or (processor.decoded_token or {}).get("sub") + try: + ContinuationService().release_claim(str(conversation_id), user) + except Exception: + logger.exception("Failed to release resume claim after a usage refusal") + return error + def complete_stream( self, question: str, diff --git a/docsgpt/api/answer/routes/stream.py b/docsgpt/api/answer/routes/stream.py index f79e0633..cdee274b 100644 --- a/docsgpt/api/answer/routes/stream.py +++ b/docsgpt/api/answer/routes/stream.py @@ -115,7 +115,9 @@ class StreamResource(Resource, BaseAnswerResource): status=401, mimetype="text/event-stream", ) - if error := self.check_usage(processor.agent_config): + if error := self.check_usage_on_resume( + processor, data["conversation_id"] + ): return error return Response( with_sse_keepalive( @@ -151,7 +153,11 @@ class StreamResource(Resource, BaseAnswerResource): mimetype="text/event-stream", ) - if error := self.check_usage(processor.agent_config): + if error := self.check_usage( + processor.agent_config, + processor.decoded_token, + agent_id=processor.agent_id, + ): return error should_persist, visibility = resolve_persistence( visibility_flag=data.get("visibility"), diff --git a/docsgpt/api/pat/rules.py b/docsgpt/api/pat/rules.py index 9b69d1a7..227fe11a 100644 --- a/docsgpt/api/pat/rules.py +++ b/docsgpt/api/pat/rules.py @@ -161,6 +161,7 @@ _CHAT = dict( RULES: dict[tuple[str, str], Rule] = { # Identity and public metadata: any valid token. ("/api/user/me", "GET"): _rule(open=True), + ("/api/user/quota", "GET"): _rule(open=True), ("/api/health", "GET"): _rule(open=True), ("/api/config", "GET"): _rule(open=True), # Agents diff --git a/docsgpt/api/user/me/routes.py b/docsgpt/api/user/me/routes.py index 67e31dfb..a6c70d4b 100644 --- a/docsgpt/api/user/me/routes.py +++ b/docsgpt/api/user/me/routes.py @@ -5,6 +5,9 @@ only from ``request.decoded_token`` (already populated and role-resolved by the auth chokepoint in ``app.py``). Auth-mode-agnostic. ``email``/``name``/ ``picture`` are OIDC-only and optional — they are echoed from the token and are never present for ``simple_jwt``/``session_jwt``/no-auth modes. + +``GET /api/user/quota`` returns the caller's usage against the limits an admin +set for them, without naming the policies behind those limits. """ from __future__ import annotations @@ -13,6 +16,8 @@ from flask import jsonify, make_response, request from flask_restx import Namespace, Resource from docsgpt.api.pat.tokens import is_pat +from docsgpt.core.settings import settings +from docsgpt.quotas.service import REQUEST_BUCKETS, QuotaService me_ns = Namespace("me", description="Current user identity and roles", path="/api") @@ -43,3 +48,34 @@ class MeResource(Resource): "resource_filter": decoded_token.get("resource_filter") or {}, } return make_response(jsonify(body), 200) + + +def _own_budget(budget: dict) -> dict: + return {"limit": budget["limit"], "used": budget["used"]} + + +@me_ns.route("/user/quota") +class MyQuotaResource(Resource): + def get(self): + """Return the caller's limited buckets: ``{bucket, tokens, cost, resets_at}`` each.""" + decoded_token = getattr(request, "decoded_token", None) + user_id = decoded_token.get("sub") if decoded_token else None + if not user_id: + return make_response(jsonify({"success": False}), 401) + statuses = QuotaService.status(user_id, ("all", *REQUEST_BUCKETS)) + buckets = [] + for status in statuses: + if status.limits.unlimited: + continue + data = status.to_dict() + buckets.append( + { + "bucket": data["bucket"], + "tokens": _own_budget(data["tokens"]), + "cost": _own_budget(data["cost"]), + "resets_at": data["resets_at"], + } + ) + return make_response( + jsonify({"success": True, "period": settings.QUOTA_PERIOD, "buckets": buckets}), 200 + ) diff --git a/docsgpt/api/user/scheduler_worker.py b/docsgpt/api/user/scheduler_worker.py index 6207832d..3e085f0e 100644 --- a/docsgpt/api/user/scheduler_worker.py +++ b/docsgpt/api/user/scheduler_worker.py @@ -17,6 +17,7 @@ from sqlalchemy import text as sql_text from docsgpt.agents.headless_runner import run_agent_headless from docsgpt.core.settings import settings from docsgpt.events.publisher import publish_user_event +from docsgpt.quotas.service import QuotaExceededError from docsgpt.storage.db.base_repository import row_to_dict from docsgpt.storage.db.engine import get_engine from docsgpt.storage.db.repositories.conversations import ( @@ -282,6 +283,11 @@ def execute_scheduled_run_body(run_id: str, celery_task_id: Optional[str]) -> Di outcome = {"answer": "", "tool_calls": [], "sources": [], "thought": ""} error_type = "timeout" error_text = "run exceeded soft time limit" + except QuotaExceededError as exc: + # The owner's usage quota is spent; the run never started. + outcome = {"answer": "", "tool_calls": [], "sources": [], "thought": ""} + error_type = "budget_exceeded" + error_text = str(exc) except Exception as exc: outcome = {"answer": "", "tool_calls": [], "sources": [], "thought": ""} error_type = "agent_error" diff --git a/docsgpt/api/v1/routes.py b/docsgpt/api/v1/routes.py index 4137d6dc..fa915d24 100644 --- a/docsgpt/api/v1/routes.py +++ b/docsgpt/api/v1/routes.py @@ -258,6 +258,8 @@ def chat_completions(): try: processor = StreamProcessor(internal_data, decoded_token) + # Set when this request took the resume claim, so a refusal can release it. + claimed_conversation_id = None if internal_data.get("tool_actions"): conversation_id = internal_data.get("conversation_id") @@ -282,6 +284,7 @@ def chat_completions(): claimed_state=pending_state, ) processor.conversation_id = conversation_id + claimed_conversation_id = conversation_id else: # Compatibility fallback for old/completed conversations and # clients that resend the full transcript without resumable @@ -338,7 +341,12 @@ def chat_completions(): ) helper = _V1AnswerHelper() - usage_error = helper.check_usage(processor.agent_config) + if claimed_conversation_id: + usage_error = helper.check_usage_on_resume(processor, claimed_conversation_id) + else: + usage_error = helper.check_usage( + processor.agent_config, processor.decoded_token, agent_id=processor.agent_id + ) if usage_error: return usage_error diff --git a/docsgpt/core/model_settings.py b/docsgpt/core/model_settings.py index d5e8a5c7..6a2a002f 100644 --- a/docsgpt/core/model_settings.py +++ b/docsgpt/core/model_settings.py @@ -32,8 +32,15 @@ class ModelCapabilities: supports_streaming: bool = True supported_attachment_types: List[str] = field(default_factory=list) context_window: int = 128000 - input_cost_per_token: Optional[float] = None - output_cost_per_token: Optional[float] = None + # USD per 1M tokens; consumed by ``docsgpt/pricing.py``. ``None`` means + # "not declared": the call is recorded at $0 unless + # ``QUOTA_UNPRICED_RATE_PER_MILLION`` is set. + input_cost_per_million: Optional[float] = None + output_cost_per_million: Optional[float] = None + # Rates for the prompt-cache sub-bins of the prompt total. ``None`` bills + # those tokens at ``input_cost_per_million``. + cached_input_cost_per_million: Optional[float] = None + cache_write_cost_per_million: Optional[float] = None # OpenAI reasoning-model effort hint (none/minimal/low/medium/high/xhigh; # the accepted subset is model-dependent). Consumed by OpenAILLM — sent # top-level on Chat Completions and nested under ``reasoning`` on the diff --git a/docsgpt/core/model_yaml.py b/docsgpt/core/model_yaml.py index 72f2d1ef..564c362e 100644 --- a/docsgpt/core/model_yaml.py +++ b/docsgpt/core/model_yaml.py @@ -18,7 +18,7 @@ from pathlib import Path from typing import Dict, List, Optional, Sequence import yaml -from pydantic import BaseModel, ConfigDict, Field, field_validator +from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator from docsgpt.core.model_settings import ( AvailableModel, @@ -65,11 +65,33 @@ class _CapabilityFields(BaseModel): supports_streaming: Optional[bool] = None attachments: Optional[List[str]] = None context_window: Optional[int] = None - input_cost_per_token: Optional[float] = None - output_cost_per_token: Optional[float] = None + input_cost_per_million: Optional[float] = Field(default=None, ge=0) + output_cost_per_million: Optional[float] = Field(default=None, ge=0) + cached_input_cost_per_million: Optional[float] = Field(default=None, ge=0) + cache_write_cost_per_million: Optional[float] = Field(default=None, ge=0) reasoning_effort: Optional[str] = None api_flavor: Optional[str] = None + @model_validator(mode="before") + @classmethod + def _per_token_alias(cls, data): + """Accept the deprecated ``*_cost_per_token`` keys, scaled to per-1M.""" + if not isinstance(data, dict): + return data + data = dict(data) + for side in ("input", "output"): + old, new = f"{side}_cost_per_token", f"{side}_cost_per_million" + if old not in data: + continue + value = data.pop(old) + if new in data: + raise ValueError(f"set only one of {old} and {new}") + if isinstance(value, bool) or not isinstance(value, (int, float)): + raise ValueError(f"{old} must be a number") + logger.warning("%s is deprecated; use %s (USD per 1M tokens)", old, new) + data[new] = value * 1_000_000 + return data + @field_validator("reasoning_effort") @classmethod def _valid_reasoning_effort(cls, v: Optional[str]) -> Optional[str]: @@ -237,8 +259,10 @@ def _build_model( supports_streaming=pick("supports_streaming", True), supported_attachment_types=expanded, context_window=pick("context_window", 128000), - input_cost_per_token=pick("input_cost_per_token", None), - output_cost_per_token=pick("output_cost_per_token", None), + input_cost_per_million=pick("input_cost_per_million", None), + output_cost_per_million=pick("output_cost_per_million", None), + cached_input_cost_per_million=pick("cached_input_cost_per_million", None), + cache_write_cost_per_million=pick("cache_write_cost_per_million", None), reasoning_effort=pick("reasoning_effort", None), api_flavor=pick("api_flavor", "chat_completions"), ) diff --git a/docsgpt/core/models/README.md b/docsgpt/core/models/README.md index d67f20fe..5e15a382 100644 --- a/docsgpt/core/models/README.md +++ b/docsgpt/core/models/README.md @@ -108,8 +108,10 @@ defaults: # optional, applied to every model below supports_streaming: bool # default true attachments: [, ...] # default [] context_window: int # default 128000 - input_cost_per_token: float # default null - output_cost_per_token: float # default null + input_cost_per_million: float # USD per 1M prompt tokens; default null (unpriced) + output_cost_per_million: float # USD per 1M generated tokens; default null + cached_input_cost_per_million: float # prompt-cache reads; default: the input rate + cache_write_cost_per_million: float # prompt-cache writes; default: the input rate reasoning_effort: # default null; none|minimal|low|medium|high|xhigh (subset is model-dependent) api_flavor: # chat_completions (default) or responses diff --git a/docsgpt/core/models/anthropic.yaml b/docsgpt/core/models/anthropic.yaml index 4e253396..34a2cedc 100644 --- a/docsgpt/core/models/anthropic.yaml +++ b/docsgpt/core/models/anthropic.yaml @@ -10,14 +10,26 @@ models: description: Most capable Claude model for complex reasoning and agentic coding context_window: 1000000 supports_structured_output: true + input_cost_per_million: 5.0 + output_cost_per_million: 25.0 + cached_input_cost_per_million: 0.5 + cache_write_cost_per_million: 6.25 - id: claude-sonnet-4-6 display_name: Claude Sonnet 4.6 description: Best balance of speed and intelligence with extended thinking context_window: 1000000 supports_structured_output: true + input_cost_per_million: 3.0 + output_cost_per_million: 15.0 + cached_input_cost_per_million: 0.3 + cache_write_cost_per_million: 3.75 - id: claude-haiku-4-5 display_name: Claude Haiku 4.5 description: Fastest Claude model with near-frontier intelligence supports_structured_output: true + input_cost_per_million: 1.0 + output_cost_per_million: 5.0 + cached_input_cost_per_million: 0.1 + cache_write_cost_per_million: 1.25 diff --git a/docsgpt/core/models/deepseek.yaml b/docsgpt/core/models/deepseek.yaml index 017c090a..08271e09 100644 --- a/docsgpt/core/models/deepseek.yaml +++ b/docsgpt/core/models/deepseek.yaml @@ -12,7 +12,11 @@ models: - id: deepseek-v4-flash display_name: DeepSeek V4 Flash description: Cost-efficient 1M-context model with hybrid thinking / non-thinking modes, tool calling and FIM completion + input_cost_per_million: 0.14 + output_cost_per_million: 0.28 - id: deepseek-v4-pro display_name: DeepSeek V4 Pro description: Frontier 1M-context model with hybrid thinking / non-thinking modes for advanced reasoning and agentic coding + input_cost_per_million: 0.435 + output_cost_per_million: 0.87 diff --git a/docsgpt/core/models/docsgpt.yaml b/docsgpt/core/models/docsgpt.yaml index b65b4fbf..4f494434 100644 --- a/docsgpt/core/models/docsgpt.yaml +++ b/docsgpt/core/models/docsgpt.yaml @@ -7,3 +7,6 @@ models: supports_tools: true attachments: [image] context_window: 1048576 + input_cost_per_million: 0.15 + output_cost_per_million: 0.5 + cached_input_cost_per_million: 0.03 diff --git a/docsgpt/core/models/google.yaml b/docsgpt/core/models/google.yaml index a4102d77..3fe7dbc5 100644 --- a/docsgpt/core/models/google.yaml +++ b/docsgpt/core/models/google.yaml @@ -9,9 +9,16 @@ models: - id: gemini-3.1-pro-preview display_name: Gemini 3.1 Pro (preview) description: Most capable Gemini 3 model with advanced reasoning and agentic coding (preview) + # Priced at the >200k-token tier; long prompts are common with attachments. + input_cost_per_million: 4.0 + output_cost_per_million: 18.0 - id: gemini-3.5-flash display_name: Gemini 3.5 Flash description: Frontier-class Flash for sustained performance on agentic and coding tasks + input_cost_per_million: 1.5 + output_cost_per_million: 9 - id: gemini-3.1-flash-lite display_name: Gemini 3.1 Flash-Lite description: Cost-efficient frontier-class multimodal model for high-throughput workloads + input_cost_per_million: 0.25 + output_cost_per_million: 1.5 diff --git a/docsgpt/core/models/groq.yaml b/docsgpt/core/models/groq.yaml index 555951ec..a4d7edfd 100644 --- a/docsgpt/core/models/groq.yaml +++ b/docsgpt/core/models/groq.yaml @@ -8,9 +8,16 @@ models: display_name: GPT-OSS 120B description: OpenAI's open-weight 120B flagship served on Groq's LPU hardware; strong general reasoning with strict structured output support supports_structured_output: true + input_cost_per_million: 0.15 + output_cost_per_million: 0.6 + cached_input_cost_per_million: 0.075 - id: llama-3.3-70b-versatile display_name: Llama 3.3 70B Versatile description: Meta's Llama 3.3 70B for general-purpose chat with parallel tool use + input_cost_per_million: 0.59 + output_cost_per_million: 0.79 - id: llama-3.1-8b-instant display_name: Llama 3.1 8B Instant description: Small, very low-latency Llama model (~560 tok/s) with parallel tool use + input_cost_per_million: 0.05 + output_cost_per_million: 0.08 diff --git a/docsgpt/core/models/novita.yaml b/docsgpt/core/models/novita.yaml index 3fa2e89f..08a2cbd8 100644 --- a/docsgpt/core/models/novita.yaml +++ b/docsgpt/core/models/novita.yaml @@ -8,14 +8,20 @@ models: display_name: DeepSeek V4 Pro description: 1.6T MoE (49B active) with 1M context, hybrid CSA/HCA attention, top-tier reasoning and agentic coding context_window: 1048576 + input_cost_per_million: 1.6 + output_cost_per_million: 3.2 - id: moonshotai/kimi-k2.6 display_name: Kimi K2.6 description: 1T-parameter open-weight MoE with native vision/video, multi-step tool calling, and agentic long-horizon execution attachments: [image] context_window: 262144 + input_cost_per_million: 0.8 + output_cost_per_million: 3.4 - id: zai-org/glm-5 display_name: GLM-5 description: Z.AI 754B-parameter MoE with strong general reasoning, function calling, and structured output context_window: 202800 + input_cost_per_million: 1.0 + output_cost_per_million: 3.2 diff --git a/docsgpt/core/models/openai.yaml b/docsgpt/core/models/openai.yaml index e0f209c5..a4b73a3a 100644 --- a/docsgpt/core/models/openai.yaml +++ b/docsgpt/core/models/openai.yaml @@ -12,9 +12,19 @@ models: context_window: 1050000 api_flavor: responses reasoning_effort: medium + # Short-context rates. Prompts over 272K tokens bill at $10 / $45 (cached $1). + input_cost_per_million: 5.0 + output_cost_per_million: 30.0 + cached_input_cost_per_million: 0.5 - id: gpt-5.4-mini display_name: GPT-5.4 Mini description: Cost-efficient GPT-5.4-class model for high-volume coding, computer use, and subagent workloads + input_cost_per_million: 0.75 + output_cost_per_million: 4.5 + cached_input_cost_per_million: 0.075 - id: gpt-5.4-nano display_name: GPT-5.4 Nano description: Cheapest GPT-5.4-class model, optimized for simple high-volume tasks where speed and cost matter most + input_cost_per_million: 0.2 + output_cost_per_million: 1.25 + cached_input_cost_per_million: 0.02 diff --git a/docsgpt/core/models/openrouter.yaml b/docsgpt/core/models/openrouter.yaml index 0c28dd30..2957fa98 100644 --- a/docsgpt/core/models/openrouter.yaml +++ b/docsgpt/core/models/openrouter.yaml @@ -10,16 +10,25 @@ models: description: Free-tier 480B MoE coder model with strong agentic tool use; rate-limited context_window: 262000 attachments: [] + input_cost_per_million: 0.0 + output_cost_per_million: 0.0 - id: deepseek/deepseek-v3.2 display_name: DeepSeek V3.2 - description: Open-weights reasoning model, very low cost (~$0.25 in / $0.38 out per 1M) + description: Open-weights reasoning model, very low cost (~$0.27 in / $0.40 out per 1M) context_window: 131072 attachments: [] supports_structured_output: true + input_cost_per_million: 0.269 + output_cost_per_million: 0.4 + cached_input_cost_per_million: 0.1345 - id: anthropic/claude-sonnet-4.6 display_name: Claude Sonnet 4.6 (via OpenRouter) description: Frontier Sonnet-class model with 1M context, vision, and extended thinking context_window: 1000000 supports_structured_output: true + input_cost_per_million: 3.0 + output_cost_per_million: 15.0 + cached_input_cost_per_million: 0.3 + cache_write_cost_per_million: 3.75 diff --git a/docsgpt/core/settings/__init__.py b/docsgpt/core/settings/__init__.py index 2af13cea..817281c5 100644 --- a/docsgpt/core/settings/__init__.py +++ b/docsgpt/core/settings/__init__.py @@ -28,6 +28,7 @@ from docsgpt.core.settings.guardrails import GuardrailSettings from docsgpt.core.settings.ingestion import IngestionSettings from docsgpt.core.settings.llm import LLMSettings from docsgpt.core.settings.ocr import OCRSettings +from docsgpt.core.settings.quotas import QuotaSettings from docsgpt.core.settings.retrieval import RetrievalSettings from docsgpt.core.settings.sandbox import SandboxSettings from docsgpt.core.settings.scheduler import SchedulerSettings @@ -54,6 +55,7 @@ SETTINGS_GROUPS: tuple[tuple[str, type[SettingsGroup]], ...] = ( ("Events and devices", EventsSettings), ("Agents", AgentSettings), ("Guardrails", GuardrailSettings), + ("Quotas", QuotaSettings), ("Scheduler", SchedulerSettings), ("Sandbox", SandboxSettings), ("Speech", SpeechSettings), diff --git a/docsgpt/core/settings/quotas.py b/docsgpt/core/settings/quotas.py new file mode 100644 index 00000000..d5d5bb34 --- /dev/null +++ b/docsgpt/core/settings/quotas.py @@ -0,0 +1,36 @@ +"""Admin-set usage quotas and the pricing that feeds their cost budgets.""" + +from __future__ import annotations + +from typing import Literal, Optional + +from pydantic import Field, field_validator + +from docsgpt.core.settings._shared import SettingsGroup + + +class QuotaSettings(SettingsGroup): + """Quota window and the treatment of unpriced models.""" + + QUOTA_PERIOD: Literal["day", "week", "month"] = Field( + default="month", + description=( + "Window every usage quota is measured over. Windows are calendar-aligned in UTC: " + "a day starts at 00:00, a week on Monday, a month on the 1st." + ), + ) + QUOTA_UNPRICED_RATE_PER_MILLION: Optional[list[float]] = Field( + default=None, + description=( + "Fallback `[input, output]` USD rates per 1M tokens for models that declare no price, " + "e.g. `[0.5, 1.5]`. Unset, such calls are recorded at $0 and only count toward token quotas." + ), + ) + @field_validator("QUOTA_UNPRICED_RATE_PER_MILLION") + @classmethod + def _two_non_negative_rates(cls, v: Optional[list[float]]) -> Optional[list[float]]: + if v is None: + return None + if len(v) != 2 or any(rate < 0 for rate in v): + raise ValueError("QUOTA_UNPRICED_RATE_PER_MILLION must be two non-negative numbers") + return v diff --git a/docsgpt/llm/llm_creator.py b/docsgpt/llm/llm_creator.py index 3b3c2a03..995600cb 100644 --- a/docsgpt/llm/llm_creator.py +++ b/docsgpt/llm/llm_creator.py @@ -52,6 +52,7 @@ class LLMCreator: base_url = None upstream_model_id = model_id capabilities = None + model = None if model_id: user_id = model_user_id if user_id is None: @@ -127,4 +128,6 @@ class LLMCreator: # llm.model_id is the upstream name (BYOM resolves it above); stamp # the canonical id (UUID for BYOM) separately for token_usage. llm._canonical_model_id = model_id + # Calls to a user's own model are recorded at $0 (see ``docsgpt/usage.py``). + llm._is_byom = model is not None and getattr(model, "source", "builtin") == "user" return llm diff --git a/docsgpt/pricing.py b/docsgpt/pricing.py new file mode 100644 index 00000000..a05b8c1b --- /dev/null +++ b/docsgpt/pricing.py @@ -0,0 +1,103 @@ +"""USD cost of LLM calls, from the per-model rates in the model catalogs.""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import Optional + +from docsgpt.core.settings import settings + + +@dataclass(frozen=True) +class ModelRates: + """USD-per-1M rates for one model; ``None`` cache rates bill at the prompt rate.""" + + prompt: float + generated: float + cached_input: Optional[float] = None + cache_write: Optional[float] = None + + +def _unpriced_rates() -> Optional[ModelRates]: + """Return the operator's fallback rates for undeclared models, if configured.""" + fallback = settings.QUOTA_UNPRICED_RATE_PER_MILLION + if not fallback: + return None + return ModelRates(prompt=float(fallback[0]), generated=float(fallback[1])) + + +def resolve_model_rates(model: Optional[str]) -> Optional[ModelRates]: + """Return the rates for a registry model id. + + Args: + model: Canonical registry id (catalog id, or the UUID of a BYOM record). + + Returns: + The declared rates, the ``QUOTA_UNPRICED_RATE_PER_MILLION`` fallback when the + model declares none, or ``None`` when there is no fallback either. + """ + # Imported lazily: the registry pulls in the provider plugins, whose LLM + # classes import ``docsgpt.usage`` and, through it, this module. + from docsgpt.core.model_registry import ModelRegistry + + entry = ModelRegistry.get_instance().models.get(str(model)) if model else None + if entry is None: + return _unpriced_rates() + caps = entry.capabilities + if caps.input_cost_per_million is None or caps.output_cost_per_million is None: + return _unpriced_rates() + cached = caps.cached_input_cost_per_million + written = caps.cache_write_cost_per_million + return ModelRates( + prompt=float(caps.input_cost_per_million), + generated=float(caps.output_cost_per_million), + cached_input=float(cached) if cached is not None else None, + cache_write=float(written) if written is not None else None, + ) + + +def is_priced(model: Optional[str]) -> bool: + """Return whether calls to ``model`` are recorded with a cost.""" + return resolve_model_rates(model) is not None + + +def cost_from_rates( + rates: ModelRates, + prompt_tokens: int, + generated_tokens: int, + cached_tokens: Optional[int] = 0, + cache_write_tokens: Optional[int] = 0, +) -> float: + """Return the USD cost of one call at ``rates``. + + ``prompt_tokens`` is the provider's billing total; ``cached_tokens`` and + ``cache_write_tokens`` are the parts of it read from or written to the prompt + cache. The sub-bins are clamped to the prompt total, so a malformed report can + never price a call below "everything cached". + """ + prompt_total = max(int(prompt_tokens or 0), 0) + cached = min(max(int(cached_tokens or 0), 0), prompt_total) + written = min(max(int(cache_write_tokens or 0), 0), prompt_total - cached) + regular = prompt_total - cached - written + cached_rate = rates.cached_input if rates.cached_input is not None else rates.prompt + write_rate = rates.cache_write if rates.cache_write is not None else rates.prompt + return ( + regular * rates.prompt + + cached * cached_rate + + written * write_rate + + max(int(generated_tokens or 0), 0) * rates.generated + ) / 1_000_000.0 + + +def compute_cost_usd( + model: Optional[str], + prompt_tokens: int, + generated_tokens: int, + cached_tokens: Optional[int] = 0, + cache_write_tokens: Optional[int] = 0, +) -> float: + """Return the USD cost of one call to ``model``; ``0.0`` when it has no rates.""" + rates = resolve_model_rates(model) + if rates is None: + return 0.0 + return cost_from_rates(rates, prompt_tokens, generated_tokens, cached_tokens, cache_write_tokens) diff --git a/docsgpt/quotas/__init__.py b/docsgpt/quotas/__init__.py new file mode 100644 index 00000000..9ed71a75 --- /dev/null +++ b/docsgpt/quotas/__init__.py @@ -0,0 +1,25 @@ +"""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, QuotaExceededError, QuotaService +from docsgpt.quotas.windows import window_bounds + +__all__ = [ + "BucketStatus", + "QuotaDefaultsProvider", + "QuotaExceeded", + "QuotaExceededError", + "QuotaService", + "ResolvedLimit", + "ResolvedLimits", + "register_defaults_provider", + "resolve_limits", + "window_bounds", +] diff --git a/docsgpt/quotas/http.py b/docsgpt/quotas/http.py new file mode 100644 index 00000000..c9a1a180 --- /dev/null +++ b/docsgpt/quotas/http.py @@ -0,0 +1,16 @@ +"""HTTP rendering of a quota refusal.""" + +from __future__ import annotations + +from flask import Response, jsonify, make_response + +from docsgpt.quotas.service import QuotaExceeded + + +def quota_exceeded_response(exceeded: QuotaExceeded) -> Response: + """Return the 429 for an exhausted quota, with ``Retry-After`` set to the reset.""" + response = make_response(jsonify(exceeded.to_payload()), 429) + response.headers["Retry-After"] = str(exceeded.retry_after_seconds) + # The reset can be weeks away; the OpenAI SDKs would otherwise retry with backoff. + response.headers["x-should-retry"] = "false" + return response 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..fff7178a --- /dev/null +++ b/docsgpt/quotas/service.py @@ -0,0 +1,183 @@ +"""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: Optional[float], resolved: ResolvedLimit, used: float) -> dict: + return { + "limit": limit, + "used": used, + "source": resolved.source, + "source_id": resolved.source_id, + } + + tokens, cost = self.limits.tokens, self.limits.cost + return { + "bucket": self.bucket, + "tokens": budget(None if tokens.unlimited else int(tokens.limit), tokens, self.tokens_used), + "cost": budget(cost.limit, 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 QuotaExceededError(Exception): + """Raised where a refused request has no HTTP response to carry the refusal.""" + + def __init__(self, exceeded: QuotaExceeded) -> None: + self.exceeded = exceeded + super().__init__(exceeded.to_payload()["message"]) + + +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/models.py b/docsgpt/storage/db/models.py index a0df1a06..1582e32d 100644 --- a/docsgpt/storage/db/models.py +++ b/docsgpt/storage/db/models.py @@ -28,6 +28,7 @@ from sqlalchemy import ( Index, Integer, MetaData, + Numeric, PrimaryKeyConstraint, UniqueConstraint, Table, @@ -241,6 +242,9 @@ token_usage_table = Table( # cache activity, so hit-rate queries stay honest across providers. Column("cached_tokens", Integer), Column("cache_write_tokens", Integer), + # Added in ``0033_quotas``. USD cost of the call at write time; 0 for + # unpriced and bring-your-own models. + Column("cost", Numeric(12, 8), nullable=False, server_default="0"), ) user_logs_table = Table( @@ -1121,3 +1125,42 @@ Index( personal_access_tokens_table.c.user_id, personal_access_tokens_table.c.created_at.desc(), ) + +# --- Usage quotas (migration 0033) ------------------------------------------ +# Admin-set limits at three layers: instance (``subject_id`` NULL), team +# per-member allowance (``teams.id``) and user override (auth ``sub``). Per +# budget a row sets a limit, marks it unlimited, or defers to the next layer. + +quota_policies_table = Table( + "quota_policies", + metadata, + Column("id", UUID(as_uuid=True), primary_key=True, server_default=func.gen_random_uuid()), + Column("scope", Text, nullable=False), + Column("subject_id", Text), + Column("bucket", Text, nullable=False, server_default="all"), + Column("token_limit", BigInteger), + Column("token_unlimited", Boolean, nullable=False, server_default="false"), + Column("cost_limit_usd", Numeric(12, 4)), + Column("cost_unlimited", Boolean, nullable=False, server_default="false"), + Column("enabled", Boolean, nullable=False, server_default="true"), + Column("note", Text), + Column("created_by", Text), + Column("updated_by", Text), + Column("created_at", DateTime(timezone=True), nullable=False, server_default=func.now()), + Column("updated_at", DateTime(timezone=True), nullable=False, server_default=func.now()), + CheckConstraint("scope IN ('instance', 'team', 'user')", name="quota_policies_scope_check"), + CheckConstraint("bucket IN ('all', 'direct', 'agent')", name="quota_policies_bucket_check"), + CheckConstraint("token_limit >= 0", name="quota_policies_token_limit_check"), + CheckConstraint("cost_limit_usd >= 0", name="quota_policies_cost_limit_check"), + CheckConstraint("(scope = 'instance') = (subject_id IS NULL)", name="quota_policies_subject_chk"), + CheckConstraint("NOT (token_unlimited AND token_limit IS NOT NULL)", name="quota_policies_token_chk"), + CheckConstraint("NOT (cost_unlimited AND cost_limit_usd IS NOT NULL)", name="quota_policies_cost_chk"), +) + +Index( + "quota_policies_subject_uidx", + quota_policies_table.c.scope, + func.coalesce(quota_policies_table.c.subject_id, ""), + quota_policies_table.c.bucket, + unique=True, +) 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/docsgpt/storage/db/repositories/token_usage.py b/docsgpt/storage/db/repositories/token_usage.py index 7292dbd1..fb256b55 100644 --- a/docsgpt/storage/db/repositories/token_usage.py +++ b/docsgpt/storage/db/repositories/token_usage.py @@ -37,6 +37,7 @@ class TokenUsageRepository: timestamp: Optional[datetime] = None, cached_tokens: Optional[int] = None, cache_write_tokens: Optional[int] = None, + cost: float = 0.0, ) -> None: # Attribution guard: the ``token_usage_attribution_chk`` CHECK # constraint requires at least one of ``user_id`` / ``api_key`` @@ -62,14 +63,14 @@ class TokenUsageRepository: INSERT INTO token_usage ( user_id, api_key, agent_id, prompt_tokens, generated_tokens, - cached_tokens, cache_write_tokens, + cached_tokens, cache_write_tokens, cost, source, request_id, model_id, timestamp ) VALUES ( :user_id, :api_key, CAST(:agent_id AS uuid), :prompt_tokens, :generated_tokens, - :cached_tokens, :cache_write_tokens, + :cached_tokens, :cache_write_tokens, :cost, :source, :request_id, :model_id, COALESCE(:timestamp, now()) ) """ @@ -82,6 +83,7 @@ class TokenUsageRepository: "generated_tokens": generated_tokens, "cached_tokens": cached_tokens, "cache_write_tokens": cache_write_tokens, + "cost": cost, "source": source, "request_id": request_id, "model_id": model_id, @@ -116,9 +118,13 @@ class TokenUsageRepository: user_id: Optional[str] = None, api_key: Optional[str] = None, ) -> int: - """Total (prompt + generated) tokens in the given time range.""" - clauses = ["timestamp >= :start", "timestamp <= :end"] - params: dict = {"start": start, "end": end} + """Total (prompt + generated) tokens in the given time range. + + Run-level rollup rows (``ROLLUP_SOURCES``) are excluded: their tokens + are already counted on the run's per-call rows. + """ + clauses = ["timestamp >= :start", "timestamp <= :end", "source <> ALL(:rollup_sources)"] + params: dict = {"start": start, "end": end, "rollup_sources": list(self.ROLLUP_SOURCES)} if user_id is not None: clauses.append("user_id = :user_id") params["user_id"] = user_id @@ -132,6 +138,56 @@ class TokenUsageRepository: ) return result.scalar() + def usage_totals(self, *, user_id: str, start: datetime, bucket: str = "all") -> tuple[int, float]: + """Return ``(tokens, cost_usd)`` a user has consumed since ``start``. + + Args: + user_id: The billable user (auth ``sub``). + start: Inclusive window start. + bucket: ``all``, ``agent`` (rows carrying an agent key or an agent + id) or ``direct`` (rows with neither). + + Rollup rows are excluded; side-channel calls count, they are real spend. + """ + clauses = ["user_id = :user_id", "timestamp >= :start", "source <> ALL(:rollup_sources)"] + # Keyless agents and workflow nodes carry an agent id without a key. + if bucket == "direct": + clauses.append("api_key IS NULL AND agent_id IS NULL") + elif bucket == "agent": + clauses.append("(api_key IS NOT NULL OR agent_id IS NOT NULL)") + elif bucket != "all": + raise ValueError(f"unknown usage bucket: {bucket!r}") + row = self._conn.execute( + text( + "SELECT COALESCE(SUM(prompt_tokens + generated_tokens), 0), COALESCE(SUM(cost), 0) " + f"FROM token_usage WHERE {' AND '.join(clauses)}" + ), + {"user_id": user_id, "start": start, "rollup_sources": list(self.ROLLUP_SOURCES)}, + ).one() + return int(row[0]), float(row[1]) + + def tokens_by_model(self, *, start: datetime) -> list[dict]: + """Return ``{model_id, tokens, cost}`` per model since ``start``, busiest first.""" + result = self._conn.execute( + text( + """ + SELECT model_id, + COALESCE(SUM(prompt_tokens + generated_tokens), 0) AS tokens, + COALESCE(SUM(cost), 0) AS cost + FROM token_usage + WHERE timestamp >= :start AND model_id IS NOT NULL + AND source <> ALL(:rollup_sources) + GROUP BY model_id + ORDER BY tokens DESC, model_id + """ + ), + {"start": start, "rollup_sources": list(self.ROLLUP_SOURCES)}, + ) + return [ + {"model_id": row[0], "tokens": int(row[1]), "cost": float(row[2])} + for row in result.fetchall() + ] + # Token usage written outside a user-initiated request (conversation # title generation, history compression, RAG question condensing, # provider fallback). Mirrors the exclusion list in ``count_in_range``. diff --git a/docsgpt/usage.py b/docsgpt/usage.py index 779163a3..82fc3eb8 100644 --- a/docsgpt/usage.py +++ b/docsgpt/usage.py @@ -2,6 +2,7 @@ import logging import time from typing import Any, Dict +from docsgpt.pricing import compute_cost_usd from docsgpt.storage.db.repositories.token_usage import TokenUsageRepository from docsgpt.storage.db.session import db_session from docsgpt.utils import num_tokens_from_object_or_list, num_tokens_from_string @@ -119,6 +120,12 @@ def _persist_call_usage(llm, call_usage): }, ) return + model_id = getattr(llm, "_canonical_model_id", None) + # Bring-your-own models run on the user's own provider key: recorded, never priced. + if getattr(llm, "_is_byom", False): + cost = 0.0 + else: + cost = _call_cost_usd(model_id, call_usage) try: with db_session() as conn: # ``timestamp`` is omitted so Postgres ``server_default @@ -136,16 +143,32 @@ def _persist_call_usage(llm, call_usage): # "0% cache hits". cached_tokens=call_usage.get("cached_tokens"), cache_write_tokens=call_usage.get("cache_write_tokens"), + cost=cost, source=( getattr(llm, "_token_usage_source", None) or "agent_stream" ), request_id=getattr(llm, "_request_id", None), - model_id=getattr(llm, "_canonical_model_id", None), + model_id=model_id, ) except Exception: logger.exception("token_usage persist failed") +def _call_cost_usd(model_id, call_usage) -> float: + """Price one call; a pricing failure records $0 rather than dropping the row.""" + try: + return compute_cost_usd( + model_id, + call_usage["prompt_tokens"], + call_usage["generated_tokens"], + cached_tokens=call_usage.get("cached_tokens"), + cache_write_tokens=call_usage.get("cache_write_tokens"), + ) + except Exception: + logger.exception("token_usage cost computation failed") + return 0.0 + + def _prefer_provider_usage(llm: Any, call_usage: Dict[str, int]) -> Dict[str, int]: """Replace estimates with upstream counts when a provider reported them. diff --git a/docsgpt/worker.py b/docsgpt/worker.py index 3ce7f15e..947f59c2 100755 --- a/docsgpt/worker.py +++ b/docsgpt/worker.py @@ -2122,6 +2122,7 @@ def agent_webhook_worker(self, agent_id, payload): try: # Shared headless path with the scheduler; approval-gated tools auto-deny. from docsgpt.agents.headless_runner import run_agent_headless + from docsgpt.quotas.service import QuotaExceededError outcome = run_agent_headless( agent_config, @@ -2135,6 +2136,12 @@ def agent_webhook_worker(self, agent_id, payload): "tool_calls": outcome.get("tool_calls", []), "thought": outcome.get("thought", ""), } + except QuotaExceededError as e: + # Returned, not raised: retrying cannot succeed before the quota resets. + logging.warning( + f"Webhook skipped for agent {agent_id}: {e}", extra={"agent_id": agent_id} + ) + return {"status": "quota_exceeded", "error": str(e)} except Exception as e: logging.error(f"Error running agent logic: {e}", exc_info=True) raise diff --git a/frontend/src/admin/AdminUI.tsx b/frontend/src/admin/AdminUI.tsx index e62ab82b..2e5be43e 100644 --- a/frontend/src/admin/AdminUI.tsx +++ b/frontend/src/admin/AdminUI.tsx @@ -119,6 +119,8 @@ const EVENT_LABELS: Record = { scim_created: 'Provisioned', scim_deactivated: 'Deactivated (SCIM)', scim_activated: 'Activated (SCIM)', + quota_policy_set: 'Quota set', + quota_policy_deleted: 'Quota removed', }; export function eventLabel(event: string): string { diff --git a/frontend/src/admin/QuotaEditor.tsx b/frontend/src/admin/QuotaEditor.tsx new file mode 100644 index 00000000..7b9de108 --- /dev/null +++ b/frontend/src/admin/QuotaEditor.tsx @@ -0,0 +1,261 @@ +import { useEffect, useState } from 'react'; +import { useSelector } from 'react-redux'; + +import adminService, { type QuotaScope } from '../api/services/adminService'; +import { Button } from '../components/ui/button'; +import { Input } from '../components/ui/input'; +import { + Select, + SelectContent, + SelectItem, + SelectTrigger, + SelectValue, +} from '../components/ui/select'; +import { selectToken } from '../preferences/preferenceSlice'; +import { fmtNumber } from './AdminUI'; +import { + fmtUsd, + formToPolicy, + isEmptyForm, + policyToForm, + usagePercent, + type Budget, + type BudgetMode, + type QuotaForm, + type QuotaPolicy, +} from './quotaUtils'; + +const MODES: { value: BudgetMode; label: string }[] = [ + { value: 'inherit', label: 'Not set here' }, + { value: 'limit', label: 'Limit' }, + { value: 'unlimited', label: 'Unlimited' }, +]; + +export function UsageBar({ + label, + budget, + kind, + caption, +}: { + label: string; + budget: Budget; + kind: 'tokens' | 'cost'; + caption?: string; +}) { + const fmt = (n: number) => (kind === 'cost' ? fmtUsd(n) : fmtNumber(n)); + const percent = usagePercent(budget.used, budget.limit); + const tone = + percent >= 100 + ? 'bg-red-500' + : percent >= 80 + ? 'bg-amber-500' + : 'bg-[#7D54D1]'; + return ( +
+
+ {label} + + {fmt(budget.used)} + {budget.limit === null ? ' · no limit' : ` of ${fmt(budget.limit)}`} + +
+ {budget.limit !== null ? ( +
+
+
+ ) : null} + {caption ? ( +

{caption}

+ ) : null} +
+ ); +} + +function BudgetField({ + label, + hint, + mode, + value, + step, + onMode, + onValue, +}: { + label: string; + hint: string; + mode: BudgetMode; + value: string; + step: string; + onMode: (mode: BudgetMode) => void; + onValue: (value: string) => void; +}) { + return ( +
+

{label}

+
+ + {mode === 'limit' ? ( + onValue(e.target.value)} + className="flex-1" + /> + ) : null} +
+
+ ); +} + +/** + * Edits the ``all``-bucket policy of one subject. Saving a form with neither + * budget set removes the policy, since a policy without an opinion is not stored. + */ +export default function QuotaEditor({ + scope, + subjectId, + policy, + inheritHint, + onSaved, +}: { + scope: QuotaScope; + subjectId: string | null; + policy: QuotaPolicy | null; + inheritHint: string; + onSaved: () => void; +}) { + const token = useSelector(selectToken); + const [form, setForm] = useState(() => policyToForm(policy)); + const [busy, setBusy] = useState(false); + const [error, setError] = useState(null); + + useEffect(() => { + setForm(policyToForm(policy)); + setError(null); + }, [policy, scope, subjectId]); + + const patch = (fields: Partial) => + setForm((prev) => ({ ...prev, ...fields })); + + const submit = async (request: () => Promise) => { + setBusy(true); + setError(null); + try { + const res = await request(); + const json = await res.json().catch(() => ({})); + if (res.ok && json.success !== false) onSaved(); + else setError(json.message || 'Could not save the quota.'); + } catch { + setError('Could not save the quota.'); + } finally { + setBusy(false); + } + }; + + const remove = () => + submit(() => adminService.deleteQuota(scope, subjectId, 'all', token)); + + const save = () => { + if (isEmptyForm(form)) { + if (policy) remove(); + return; + } + const result = formToPolicy(form, policy); + if (!result.ok) { + setError(result.error); + return; + } + submit(() => adminService.setQuota(scope, subjectId, result.policy, token)); + }; + + return ( +
+

{inheritHint}

+ {policy && !policy.enabled ? ( +

+ This policy is disabled and is not enforced. Saving keeps it disabled. +

+ ) : null} + patch({ tokenMode })} + onValue={(tokenLimit) => patch({ tokenLimit })} + /> + patch({ costMode })} + onValue={(costLimit) => patch({ costLimit })} + /> +
+

Note

+ patch({ note: e.target.value })} + className="mt-1" + /> +
+ {error ? ( +

+ {error} +

+ ) : null} +
+ {policy ? ( + + ) : null} + +
+
+ ); +} diff --git a/frontend/src/admin/Quotas.tsx b/frontend/src/admin/Quotas.tsx new file mode 100644 index 00000000..967f4f00 --- /dev/null +++ b/frontend/src/admin/Quotas.tsx @@ -0,0 +1,347 @@ +import { useCallback, useEffect, useMemo, useState } from 'react'; +import { useSelector } from 'react-redux'; + +import adminService, { type QuotaScope } from '../api/services/adminService'; +import teamsService from '../api/services/teamsService'; +import { Button } from '../components/ui/button'; +import { Modal } from '../components/ui/modal'; +import { + Select, + SelectContent, + SelectItem, + SelectTrigger, + SelectValue, +} from '../components/ui/select'; +import { + Table, + TableBody, + TableCell, + TableContainer, + TableHead, + TableHeader, + TableRow, +} from '../components/ui/table'; +import { selectToken } from '../preferences/preferenceSlice'; +import { + LoadError, + Loading, + Pill, + fmtDate, + fmtNumber, + fmtRelative, +} from './AdminUI'; +import QuotaEditor from './QuotaEditor'; +import { describeBudget, type QuotaPolicy } from './quotaUtils'; + +type TeamPolicy = QuotaPolicy & { + team_name?: string | null; + team_slug?: string | null; + member_count?: number | null; +}; + +type Editing = { + scope: QuotaScope; + subjectId: string | null; + title: string; + policy: QuotaPolicy | null; +}; + +const HINTS: Record = { + instance: + 'Applies to every user that no team allowance or user override covers.', + team: 'Each member gets this allowance; it is not a shared pool. A member of several teams gets the most generous one.', + user: 'Overrides team allowances and the instance default for this user.', +}; + +function PolicyCells({ policy }: { policy: QuotaPolicy }) { + return ( + <> + + {describeBudget(policy.token_limit, policy.token_unlimited, 'tokens')} + + + {describeBudget(policy.cost_limit_usd, policy.cost_unlimited, 'cost')} + + + {policy.note || '—'} + + + {fmtRelative(policy.updated_at)} + + + ); +} + +export default function Quotas() { + const token = useSelector(selectToken); + const [data, setData] = useState(null); + const [teams, setTeams] = useState([]); + const [loading, setLoading] = useState(true); + const [editing, setEditing] = useState(null); + const [teamPick, setTeamPick] = useState(''); + + const load = useCallback(async () => { + setLoading(true); + try { + const [quotasRes, teamsJson] = await Promise.all([ + adminService.getQuotas(token), + teamsService.listAll(token).catch(() => ({})), + ]); + setData(await quotasRes.json().catch(() => ({ success: false }))); + setTeams(teamsJson?.teams ?? []); + } catch { + setData({ success: false }); + } finally { + setLoading(false); + } + }, [token]); + + useEffect(() => { + load(); + }, [load]); + + // The editor covers the ``all`` bucket; other buckets are listed read-only. + const isAll = (p: QuotaPolicy) => p.bucket === 'all'; + const instancePolicy: QuotaPolicy | null = + (data?.instance ?? []).find(isAll) ?? null; + const teamPolicies: TeamPolicy[] = data?.teams ?? []; + const userPolicies: QuotaPolicy[] = data?.users ?? []; + const teamsWithoutPolicy = useMemo(() => { + const covered = new Set( + teamPolicies.filter(isAll).map((p) => String(p.subject_id)), + ); + return teams.filter((team) => !covered.has(String(team.id))); + }, [teams, teamPolicies]); + + if (data === null && loading) return ; + if (!data?.success) return ; + + const bucketPill = (policy: QuotaPolicy) => ( + <> + {isAll(policy) ? null : {policy.bucket} traffic} + {policy.enabled ? null : Disabled} + + ); + + return ( +
+

+ Usage is counted per user over each calendar {data.period} (UTC). The + current window resets {fmtDate(data.resets_at)}. A request is refused + once a budget is used up; the request that crosses it still completes. +

+ + {(data.unpriced_models ?? []).length > 0 ? ( +
+

+ Models without a price are invisible to cost limits +

+

+ These were used this {data.period} and recorded at $0:{' '} + {(data.unpriced_models as any[]) + .map((m) => `${m.model_id} (${fmtNumber(m.tokens)} tokens)`) + .join(', ')} + . Use a token limit for them, or declare their rates in the model + catalog. +

+
+ ) : null} + +
+
+

Instance default

+ +
+

+ {instancePolicy + ? `${instancePolicy.enabled ? '' : 'Disabled · '}Tokens: ${describeBudget(instancePolicy.token_limit, instancePolicy.token_unlimited, 'tokens')} · Cost: ${describeBudget(instancePolicy.cost_limit_usd, instancePolicy.cost_unlimited, 'cost')}` + : 'No default: users without a team allowance or override are unlimited.'} +

+
+ +
+
+

Team allowances

+ {teamsWithoutPolicy.length > 0 ? ( +
+ + +
+ ) : null} +
+ {teamPolicies.length === 0 ? ( +

+ No team has an allowance. +

+ ) : ( + + + + + Team + Members + Tokens + Cost + Note + Updated + Actions + + + + {teamPolicies.map((policy) => ( + + + + {policy.team_name ?? policy.subject_id} + + {bucketPill(policy)} + + + {fmtNumber(policy.member_count)} + + + + {isAll(policy) ? ( + + ) : null} + + + ))} + +
+
+ )} +
+ +
+

User overrides

+ {userPolicies.length === 0 ? ( +

+ No user has an override. Add one from a user's menu on the + Users tab. +

+ ) : ( + + + + + User + Tokens + Cost + Note + Updated + Actions + + + + {userPolicies.map((policy) => ( + + + + {policy.subject_id} + + {bucketPill(policy)} + + + + {isAll(policy) ? ( + + ) : null} + + + ))} + +
+
+ )} +
+ + { + if (!open) setEditing(null); + }} + title={editing ? `Quota · ${editing.title}` : 'Quota'} + > + {editing ? ( + { + setEditing(null); + setTeamPick(''); + load(); + }} + /> + ) : null} + +
+ ); +} diff --git a/frontend/src/admin/UserQuotaModal.tsx b/frontend/src/admin/UserQuotaModal.tsx new file mode 100644 index 00000000..241da310 --- /dev/null +++ b/frontend/src/admin/UserQuotaModal.tsx @@ -0,0 +1,122 @@ +import { useCallback, useEffect, useRef, useState } from 'react'; +import { useSelector } from 'react-redux'; + +import adminService from '../api/services/adminService'; +import teamsService from '../api/services/teamsService'; +import { Modal } from '../components/ui/modal'; +import { selectToken } from '../preferences/preferenceSlice'; +import { LoadError, Loading, fmtDate } from './AdminUI'; +import QuotaEditor, { UsageBar } from './QuotaEditor'; +import { + sourceLabel, + type BucketStatus, + type Budget, + type QuotaPolicy, +} from './quotaUtils'; + +/** A user's effective limits and usage, with the editor for their override. */ +export default function UserQuotaModal({ + userId, + onClose, +}: { + userId: string | null; + onClose: () => void; +}) { + const token = useSelector(selectToken); + const [data, setData] = useState(null); + const [teamNames, setTeamNames] = useState>({}); + + // Bumped per request so a slow response for a previous user is discarded + // instead of showing (and letting the editor save) that user's policy. + const requestRef = useRef(0); + + const load = useCallback(async () => { + const request = ++requestRef.current; + setData(null); + if (!userId) return; + try { + const [res, teamsJson] = await Promise.all([ + adminService.getUserQuota(userId, token), + teamsService.listAll(token).catch(() => ({})), + ]); + const json = await res.json().catch(() => ({ success: false })); + if (request !== requestRef.current) return; + setData(json); + setTeamNames( + Object.fromEntries( + (teamsJson?.teams ?? []).map((team: any) => [ + String(team.id), + team.name, + ]), + ), + ); + } catch { + if (request === requestRef.current) setData({ success: false }); + } + }, [userId, token]); + + useEffect(() => { + load(); + }, [load]); + + const overall: BucketStatus | undefined = (data?.effective ?? []).find( + (status: BucketStatus) => status.bucket === 'all', + ); + const override: QuotaPolicy | null = + (data?.policies ?? []).find((p: QuotaPolicy) => p.bucket === 'all') ?? null; + const caption = (budget: Budget) => + sourceLabel( + budget, + budget.source_id ? teamNames[budget.source_id] : undefined, + ); + + return ( + { + if (!open) onClose(); + }} + title={userId ? `Quota · ${userId}` : 'Quota'} + > + {data === null ? ( + + ) : !data.success ? ( + + ) : ( +
+ {overall ? ( +
+ + +

+ Resets {fmtDate(overall.resets_at)} +

+
+ ) : null} +
+

+ User override +

+ +
+
+ )} +
+ ); +} diff --git a/frontend/src/admin/Users.tsx b/frontend/src/admin/Users.tsx index 13e40d13..d7ae7f8e 100644 --- a/frontend/src/admin/Users.tsx +++ b/frontend/src/admin/Users.tsx @@ -1,5 +1,6 @@ import { Eye, + Gauge, LogOut, ShieldCheck, ShieldOff, @@ -40,6 +41,7 @@ import { fmtNumber, fmtRelative, } from './AdminUI'; +import UserQuotaModal from './UserQuotaModal'; type AdminUser = { user_id: string; @@ -70,6 +72,7 @@ export default function Users() { const [busy, setBusy] = useState(null); const [menuUserId, setMenuUserId] = useState(null); const [detail, setDetail] = useState(null); + const [quotaUserId, setQuotaUserId] = useState(null); const [feedback, setFeedback] = useState<{ ok: boolean; message: string; @@ -157,7 +160,14 @@ export default function Users() { isAdmin: boolean, active: boolean, ): Action[] => { - const acts: Action[] = []; + const acts: Action[] = [ + { + key: 'quota', + label: 'Quota', + icon: Gauge, + perform: () => setQuotaUserId(userId), + }, + ]; if (isAdmin) { acts.push({ key: 'revoke', @@ -433,6 +443,11 @@ export default function Users() { /> ) : null} + setQuotaUserId(null)} + /> + { diff --git a/frontend/src/admin/index.tsx b/frontend/src/admin/index.tsx index 6a147f06..5a21c9f4 100644 --- a/frontend/src/admin/index.tsx +++ b/frontend/src/admin/index.tsx @@ -11,6 +11,7 @@ import { Tabs, TabsList, TabsTrigger } from '../components/ui/tabs'; import Admins from './Admins'; import Audit from './Audit'; import Overview from './Overview'; +import Quotas from './Quotas'; import Usage from './Usage'; import Users from './Users'; @@ -19,6 +20,7 @@ const TABS = [ { key: 'users', label: 'Users', path: '/admin/users' }, { key: 'admins', label: 'Admins', path: '/admin/roles' }, { key: 'usage', label: 'Usage', path: '/admin/usage' }, + { key: 'quotas', label: 'Quotas', path: '/admin/quotas' }, { key: 'audit', label: 'Audit', path: '/admin/audit' }, ]; @@ -63,6 +65,7 @@ export default function Admin() { } /> } /> } /> + } /> } /> } /> diff --git a/frontend/src/admin/quotaUtils.test.ts b/frontend/src/admin/quotaUtils.test.ts new file mode 100644 index 00000000..2afad959 --- /dev/null +++ b/frontend/src/admin/quotaUtils.test.ts @@ -0,0 +1,128 @@ +import { describe, expect, it } from 'vitest'; + +import { + describeBudget, + formToPolicy, + isEmptyForm, + policyToForm, + sourceLabel, + usagePercent, + type QuotaPolicy, +} from './quotaUtils'; + +const policy = (fields: Partial): QuotaPolicy => ({ + scope: 'user', + subject_id: 'u1', + bucket: 'all', + token_limit: null, + token_unlimited: false, + cost_limit_usd: null, + cost_unlimited: false, + enabled: true, + ...fields, +}); + +describe('policyToForm', () => { + it('starts a missing policy as inherit', () => { + const form = policyToForm(null); + expect(form.tokenMode).toBe('inherit'); + expect(form.costMode).toBe('inherit'); + expect(isEmptyForm(form)).toBe(true); + }); + + it('keeps zero as a limit, not as inherit', () => { + const form = policyToForm(policy({ token_limit: 0, cost_unlimited: true })); + expect(form.tokenMode).toBe('limit'); + expect(form.tokenLimit).toBe('0'); + expect(form.costMode).toBe('unlimited'); + }); +}); + +describe('formToPolicy', () => { + const base = policyToForm(null); + + it('round-trips limits and trims the note', () => { + const result = formToPolicy({ + ...base, + tokenMode: 'limit', + tokenLimit: ' 5000 ', + costMode: 'limit', + costLimit: '2.5', + note: ' trial ', + }); + expect(result).toEqual({ + ok: true, + policy: { + bucket: 'all', + enabled: true, + token_limit: 5000, + token_unlimited: false, + cost_limit_usd: 2.5, + cost_unlimited: false, + note: 'trial', + }, + }); + }); + + it('keeps a disabled policy disabled', () => { + const form = { ...base, tokenMode: 'limit' as const, tokenLimit: '10' }; + const stored = policy({ enabled: false }); + const result = formToPolicy(form, stored); + expect(result.ok && result.policy.enabled).toBe(false); + }); + + it('sends unlimited without a limit', () => { + const result = formToPolicy({ + ...base, + tokenMode: 'unlimited', + tokenLimit: '99', + }); + expect(result.ok && result.policy.token_limit).toBeNull(); + expect(result.ok && result.policy.token_unlimited).toBe(true); + }); + + it.each(['', '1.5', '-1', 'abc', '1e3'])( + 'rejects token limit %j', + (tokenLimit) => { + expect(formToPolicy({ ...base, tokenMode: 'limit', tokenLimit }).ok).toBe( + false, + ); + }, + ); + + it.each(['', '-0.01', 'abc', 'Infinity'])( + 'rejects cost limit %j', + (costLimit) => { + expect(formToPolicy({ ...base, costMode: 'limit', costLimit }).ok).toBe( + false, + ); + }, + ); +}); + +describe('usagePercent', () => { + it('handles unlimited, zero and overshoot', () => { + expect(usagePercent(50, null)).toBe(0); + expect(usagePercent(0, 0)).toBe(100); + expect(usagePercent(25, 100)).toBe(25); + expect(usagePercent(500, 100)).toBe(100); + }); +}); + +describe('labels', () => { + it('describes budgets', () => { + expect(describeBudget(null, true, 'tokens')).toBe('Unlimited'); + expect(describeBudget(null, false, 'cost')).toBe('—'); + expect(describeBudget(1000, false, 'tokens')).toContain('tokens'); + }); + + it('names the layer a limit came from', () => { + expect(sourceLabel({ limit: null, used: 0 })).toBe('No limit set'); + expect(sourceLabel({ limit: 1, used: 0, source: 'team' }, 'Eng')).toBe( + 'Team: Eng', + ); + expect(sourceLabel({ limit: 1, used: 0, source: 'default' })).toBe( + 'Plan default', + ); + }); +}); diff --git a/frontend/src/admin/quotaUtils.ts b/frontend/src/admin/quotaUtils.ts new file mode 100644 index 00000000..02bc2f6b --- /dev/null +++ b/frontend/src/admin/quotaUtils.ts @@ -0,0 +1,136 @@ +// Pure helpers behind the quota editor and usage bars. + +export type BudgetMode = 'inherit' | 'limit' | 'unlimited'; + +export type QuotaPolicy = { + scope: 'instance' | 'team' | 'user'; + subject_id: string | null; + bucket: string; + token_limit: number | null; + token_unlimited: boolean; + cost_limit_usd: number | null; + cost_unlimited: boolean; + enabled: boolean; + note?: string | null; + updated_by?: string | null; + updated_at?: string | null; +}; + +export type Budget = { + limit: number | null; + used: number; + source?: string | null; + source_id?: string | null; +}; + +export type BucketStatus = { + bucket: string; + tokens: Budget; + cost: Budget; + resets_at: string; +}; + +export type QuotaForm = { + tokenMode: BudgetMode; + tokenLimit: string; + costMode: BudgetMode; + costLimit: string; + note: string; +}; + +const mode = (limit: number | null, unlimited: boolean): BudgetMode => { + if (unlimited) return 'unlimited'; + return limit === null || limit === undefined ? 'inherit' : 'limit'; +}; + +export function policyToForm(policy?: QuotaPolicy | null): QuotaForm { + return { + tokenMode: policy + ? mode(policy.token_limit, policy.token_unlimited) + : 'inherit', + tokenLimit: policy?.token_limit != null ? String(policy.token_limit) : '', + costMode: policy + ? mode(policy.cost_limit_usd, policy.cost_unlimited) + : 'inherit', + costLimit: + policy?.cost_limit_usd != null ? String(policy.cost_limit_usd) : '', + note: policy?.note ?? '', + }; +} + +export type FormResult = + { ok: true; policy: Record } | { ok: false; error: string }; + +// An empty form (both budgets inherited) is not a policy: the caller deletes instead. +export function isEmptyForm(form: QuotaForm): boolean { + return form.tokenMode === 'inherit' && form.costMode === 'inherit'; +} + +// ``existing`` carries the stored ``enabled`` flag through an edit: the form has +// no control for it, and a body without it would switch the policy back on. +export function formToPolicy( + form: QuotaForm, + existing?: QuotaPolicy | null, +): FormResult { + const policy: Record = { + bucket: 'all', + enabled: existing?.enabled ?? true, + token_limit: null, + token_unlimited: form.tokenMode === 'unlimited', + cost_limit_usd: null, + cost_unlimited: form.costMode === 'unlimited', + note: form.note.trim() || null, + }; + if (form.tokenMode === 'limit') { + const raw = form.tokenLimit.trim(); + if (!/^\d+$/.test(raw)) + return { ok: false, error: 'Token limit must be a whole number.' }; + const tokens = Number(raw); + if (!Number.isSafeInteger(tokens)) + return { ok: false, error: 'Token limit is too large.' }; + policy.token_limit = tokens; + } + if (form.costMode === 'limit') { + const raw = form.costLimit.trim(); + const cost = Number(raw); + if (raw === '' || !Number.isFinite(cost) || cost < 0) + return { ok: false, error: 'Cost limit must be a number, 0 or more.' }; + policy.cost_limit_usd = cost; + } + return { ok: true, policy }; +} + +export function usagePercent(used: number, limit: number | null): number { + if (limit === null || limit === undefined) return 0; + if (limit <= 0) return 100; + return Math.min(100, Math.max(0, (used / limit) * 100)); +} + +export function fmtUsd(value?: number | null): string { + return new Intl.NumberFormat(undefined, { + style: 'currency', + currency: 'USD', + maximumFractionDigits: value != null && value < 1 ? 4 : 2, + }).format(value ?? 0); +} + +export function describeBudget( + limit: number | null, + unlimited: boolean, + kind: 'tokens' | 'cost', +): string { + if (unlimited) return 'Unlimited'; + if (limit === null || limit === undefined) return '—'; + return kind === 'cost' + ? fmtUsd(limit) + : `${new Intl.NumberFormat().format(limit)} tokens`; +} + +export function sourceLabel(budget: Budget, teamName?: string): string { + if (!budget.source) return 'No limit set'; + if (budget.source === 'user') return 'User override'; + if (budget.source === 'team') + return teamName ? `Team: ${teamName}` : 'Team allowance'; + if (budget.source === 'instance') return 'Instance default'; + return 'Plan default'; +} diff --git a/frontend/src/api/endpoints.ts b/frontend/src/api/endpoints.ts index 6f5f36e6..89187528 100644 --- a/frontend/src/api/endpoints.ts +++ b/frontend/src/api/endpoints.ts @@ -2,6 +2,7 @@ const endpoints = { USER: { CONFIG: '/api/config', ME: '/api/user/me', + QUOTA: '/api/user/quota', NEW_TOKEN: '/api/generate_token', OIDC_LOGIN: '/api/auth/oidc/login', OIDC_TOKEN: '/api/auth/oidc/token', @@ -173,6 +174,12 @@ const endpoints = { USAGE: '/api/admin/usage', AUDIT: '/api/admin/audit', DEVICE_AUDIT: '/api/admin/devices/audit', + QUOTAS: '/api/admin/quotas', + QUOTA_INSTANCE: '/api/admin/quotas/instance', + QUOTA_TEAM: (id: string) => + `/api/admin/quotas/teams/${encodeURIComponent(id)}`, + QUOTA_USER: (id: string) => + `/api/admin/quotas/users/${encodeURIComponent(id)}`, }, CONVERSATION: { ANSWER: '/api/answer', diff --git a/frontend/src/api/services/adminService.ts b/frontend/src/api/services/adminService.ts index 7359aed4..adf67e71 100644 --- a/frontend/src/api/services/adminService.ts +++ b/frontend/src/api/services/adminService.ts @@ -10,6 +10,14 @@ const qs = (params: Record): string => { return str ? `?${str}` : ''; }; +export type QuotaScope = 'instance' | 'team' | 'user'; + +const quotaUrl = (scope: QuotaScope, subjectId?: string | null): string => { + if (scope === 'team') return endpoints.ADMIN.QUOTA_TEAM(subjectId ?? ''); + if (scope === 'user') return endpoints.ADMIN.QUOTA_USER(subjectId ?? ''); + return endpoints.ADMIN.QUOTA_INSTANCE; +}; + const adminService = { getOverview: (token: string | null): Promise => apiClient.get(endpoints.ADMIN.OVERVIEW, token), @@ -54,6 +62,23 @@ const adminService = { token: string | null, ): Promise => apiClient.get(`${endpoints.ADMIN.DEVICE_AUDIT}${qs(params)}`, token), + getQuotas: (token: string | null): Promise => + apiClient.get(endpoints.ADMIN.QUOTAS, token), + getUserQuota: (userId: string, token: string | null): Promise => + apiClient.get(endpoints.ADMIN.QUOTA_USER(userId), token), + setQuota: ( + scope: QuotaScope, + subjectId: string | null, + policy: Record, + token: string | null, + ): Promise => apiClient.put(quotaUrl(scope, subjectId), policy, token), + deleteQuota: ( + scope: QuotaScope, + subjectId: string | null, + bucket: string, + token: string | null, + ): Promise => + apiClient.delete(`${quotaUrl(scope, subjectId)}${qs({ bucket })}`, token), }; export default adminService; diff --git a/frontend/src/api/services/userService.ts b/frontend/src/api/services/userService.ts index b9de3bc0..fd9d81f6 100644 --- a/frontend/src/api/services/userService.ts +++ b/frontend/src/api/services/userService.ts @@ -7,6 +7,8 @@ const userService = { throttledApiClient.get(endpoints.USER.CONFIG, null), getMe: (token: string | null): Promise => apiClient.get(endpoints.USER.ME, token), + getQuota: (token: string | null): Promise => + apiClient.get(endpoints.USER.QUOTA, token), getNewToken: (): Promise => throttledApiClient.get(endpoints.USER.NEW_TOKEN, null), // Token deliberately null: a stale Authorization header must not be able diff --git a/frontend/src/conversation/conversationHandlers.ts b/frontend/src/conversation/conversationHandlers.ts index 4578c8d2..dd519628 100644 --- a/frontend/src/conversation/conversationHandlers.ts +++ b/frontend/src/conversation/conversationHandlers.ts @@ -1,7 +1,10 @@ +import i18n from 'i18next'; + import { baseURL } from '../api/client'; import conversationService from '../api/services/conversationService'; import { Doc } from '../models/misc'; import { Answer, FEEDBACK, RetrievalPayload } from './conversationModels'; +import { isQuotaError, quotaErrorMessage } from './quotaError'; import { ToolCallsType } from './types'; /** @@ -48,7 +51,9 @@ async function _handlePreStreamHttpError( if (text) { try { const parsed = JSON.parse(text); - if (parsed && typeof parsed === 'object') { + if (isQuotaError(parsed)) { + message = quotaErrorMessage(parsed, i18n.t.bind(i18n), i18n.language); + } else if (parsed && typeof parsed === 'object') { message = (typeof parsed.message === 'string' && parsed.message) || (typeof parsed.error === 'string' && parsed.error) || diff --git a/frontend/src/conversation/quotaError.test.ts b/frontend/src/conversation/quotaError.test.ts new file mode 100644 index 00000000..78220511 --- /dev/null +++ b/frontend/src/conversation/quotaError.test.ts @@ -0,0 +1,52 @@ +import { describe, expect, it } from 'vitest'; + +import { isQuotaError, quotaErrorMessage } from './quotaError'; + +const t = ((key: string, values: Record) => + `${key}|${values.used}|${values.limit}|${values.resetsAt}`) as any; + +describe('isQuotaError', () => { + it('matches only the quota error code', () => { + expect(isQuotaError({ error_code: 'quota-exceeded' })).toBe(true); + expect(isQuotaError({ message: 'Exceeding usage limit' })).toBe(false); + expect(isQuotaError(null)).toBe(false); + expect(isQuotaError('quota-exceeded')).toBe(false); + }); +}); + +describe('quotaErrorMessage', () => { + it('formats token budgets as numbers', () => { + const message = quotaErrorMessage( + { + dimension: 'tokens', + usage: 1200000, + limit: 1000000, + resets_at: '2026-10-01T00:00:00+00:00', + }, + t, + 'en-US', + ); + const [key, used, limit, resetsAt] = message.split('|'); + expect(key).toBe('conversation.quotaExceeded.tokens'); + expect([used, limit]).toEqual(['1,200,000', '1,000,000']); + expect(resetsAt).not.toBe(''); + }); + + it('formats cost budgets as dollars', () => { + const message = quotaErrorMessage( + { dimension: 'cost', usage: 5.25, limit: 5 }, + t, + 'en-US', + ); + expect(message).toBe('conversation.quotaExceeded.cost|$5.25|$5.00|'); + }); + + it('tolerates a malformed reset time', () => { + const message = quotaErrorMessage( + { dimension: 'tokens', usage: 1, limit: 1, resets_at: 'soon' }, + t, + 'en-US', + ); + expect(message.endsWith('|')).toBe(true); + }); +}); diff --git a/frontend/src/conversation/quotaError.ts b/frontend/src/conversation/quotaError.ts new file mode 100644 index 00000000..86bf46dd --- /dev/null +++ b/frontend/src/conversation/quotaError.ts @@ -0,0 +1,47 @@ +import type { TFunction } from 'i18next'; + +export type QuotaErrorBody = { + error_code?: string; + dimension?: string; + usage?: number; + limit?: number; + resets_at?: string; +}; + +export function isQuotaError(body: unknown): body is QuotaErrorBody { + return ( + !!body && + typeof body === 'object' && + (body as QuotaErrorBody).error_code === 'quota-exceeded' + ); +} + +/** The chat message for a 429 ``quota-exceeded`` body, in the user's language. */ +export function quotaErrorMessage( + body: QuotaErrorBody, + t: TFunction, + locale?: string, +): string { + const isCost = body.dimension === 'cost'; + const amount = (value?: number) => + isCost + ? new Intl.NumberFormat(locale, { + style: 'currency', + currency: 'USD', + }).format(value ?? 0) + : new Intl.NumberFormat(locale).format(value ?? 0); + const reset = body.resets_at ? new Date(body.resets_at) : null; + const resetsAt = + reset && !Number.isNaN(reset.getTime()) + ? new Intl.DateTimeFormat(locale, { + dateStyle: 'medium', + timeStyle: 'short', + }).format(reset) + : ''; + return t( + isCost + ? 'conversation.quotaExceeded.cost' + : 'conversation.quotaExceeded.tokens', + { used: amount(body.usage), limit: amount(body.limit), resetsAt }, + ); +} diff --git a/frontend/src/locale/de.json b/frontend/src/locale/de.json index ae7d2630..df6dabfa 100644 --- a/frontend/src/locale/de.json +++ b/frontend/src/locale/de.json @@ -375,6 +375,17 @@ "toolCalls": "Werkzeugaufrufe", "runSuccess": "Erfolgsquote", "feedback": "Feedback" + }, + "quota": { + "title": "Ihr Nutzungskontingent", + "resets": "Wird zurückgesetzt: {{resetsAt}}", + "tokens": "Tokens", + "cost": "Kosten", + "usedOf": "{{used}} von {{limit}}", + "scope": { + "direct": "Chat ohne Agent", + "agent": "Über Agenten" + } } }, "logs": { @@ -1262,6 +1273,10 @@ "running": "Läuft…", "denied": "Vom Benutzer abgelehnt", "failed": "fehlgeschlagen" + }, + "quotaExceeded": { + "tokens": "Sie haben {{used}} von Ihrem Kontingent von {{limit}} Tokens verbraucht. Es wird am {{resetsAt}} zurückgesetzt.", + "cost": "Sie haben {{used}} von Ihrem Nutzungsbudget von {{limit}} verbraucht. Es wird am {{resetsAt}} zurückgesetzt." } }, "agents": { diff --git a/frontend/src/locale/en.json b/frontend/src/locale/en.json index 93674148..cdf03a26 100644 --- a/frontend/src/locale/en.json +++ b/frontend/src/locale/en.json @@ -380,6 +380,17 @@ "toolCalls": "Tool Calls", "runSuccess": "Run Success", "feedback": "Feedback" + }, + "quota": { + "title": "Your usage quota", + "resets": "Resets {{resetsAt}}", + "tokens": "Tokens", + "cost": "Cost", + "usedOf": "{{used}} of {{limit}}", + "scope": { + "direct": "Chat without an agent", + "agent": "Through agents" + } } }, "logs": { @@ -1273,6 +1284,10 @@ "running": "Running…", "denied": "Denied by user", "failed": "failed" + }, + "quotaExceeded": { + "tokens": "You've used {{used}} of your {{limit}} token quota. It resets {{resetsAt}}.", + "cost": "You've used {{used}} of your {{limit}} usage budget. It resets {{resetsAt}}." } }, "agents": { diff --git a/frontend/src/locale/es.json b/frontend/src/locale/es.json index 3f19b13b..76058628 100644 --- a/frontend/src/locale/es.json +++ b/frontend/src/locale/es.json @@ -375,6 +375,17 @@ "toolCalls": "Llamadas a Herramientas", "runSuccess": "Éxito de Ejecución", "feedback": "Retroalimentación" + }, + "quota": { + "title": "Tu cuota de uso", + "resets": "Se restablece el {{resetsAt}}", + "tokens": "Tokens", + "cost": "Coste", + "usedOf": "{{used}} de {{limit}}", + "scope": { + "direct": "Chat sin agente", + "agent": "A través de agentes" + } } }, "logs": { @@ -1262,6 +1273,10 @@ "running": "Ejecutando…", "denied": "Denegado por el usuario", "failed": "falló" + }, + "quotaExceeded": { + "tokens": "Has usado {{used}} de tu cuota de {{limit}} tokens. Se restablece el {{resetsAt}}.", + "cost": "Has usado {{used}} de tu presupuesto de uso de {{limit}}. Se restablece el {{resetsAt}}." } }, "agents": { diff --git a/frontend/src/locale/jp.json b/frontend/src/locale/jp.json index f60227d1..18626d19 100644 --- a/frontend/src/locale/jp.json +++ b/frontend/src/locale/jp.json @@ -375,6 +375,17 @@ "toolCalls": "ツール呼び出し", "runSuccess": "実行成功率", "feedback": "フィードバック" + }, + "quota": { + "title": "利用クォータ", + "resets": "{{resetsAt}} にリセット", + "tokens": "トークン", + "cost": "コスト", + "usedOf": "{{used}} / {{limit}}", + "scope": { + "direct": "エージェントなしのチャット", + "agent": "エージェント経由" + } } }, "logs": { @@ -1262,6 +1273,10 @@ "running": "実行中…", "denied": "ユーザーによって拒否されました", "failed": "失敗" + }, + "quotaExceeded": { + "tokens": "トークンクォータ {{limit}} のうち {{used}} を使用しました。{{resetsAt}} にリセットされます。", + "cost": "利用予算 {{limit}} のうち {{used}} を使用しました。{{resetsAt}} にリセットされます。" } }, "agents": { diff --git a/frontend/src/locale/ru.json b/frontend/src/locale/ru.json index 878451ba..081860c6 100644 --- a/frontend/src/locale/ru.json +++ b/frontend/src/locale/ru.json @@ -375,6 +375,17 @@ "toolCalls": "Вызовы инструментов", "runSuccess": "Успешность запусков", "feedback": "Обратная связь" + }, + "quota": { + "title": "Ваша квота использования", + "resets": "Сброс: {{resetsAt}}", + "tokens": "Токены", + "cost": "Стоимость", + "usedOf": "{{used}} из {{limit}}", + "scope": { + "direct": "Чат без агента", + "agent": "Через агентов" + } } }, "logs": { @@ -1282,6 +1293,10 @@ "running": "Выполняется…", "denied": "Отклонено пользователем", "failed": "не удалось" + }, + "quotaExceeded": { + "tokens": "Вы использовали {{used}} из квоты в {{limit}} токенов. Квота сбросится {{resetsAt}}.", + "cost": "Вы использовали {{used}} из бюджета в {{limit}}. Бюджет сбросится {{resetsAt}}." } }, "agents": { diff --git a/frontend/src/locale/zh-TW.json b/frontend/src/locale/zh-TW.json index 5323803d..7902cf5a 100644 --- a/frontend/src/locale/zh-TW.json +++ b/frontend/src/locale/zh-TW.json @@ -375,6 +375,17 @@ "toolCalls": "工具呼叫", "runSuccess": "執行成功率", "feedback": "回饋" + }, + "quota": { + "title": "您的用量配額", + "resets": "{{resetsAt}} 重設", + "tokens": "權杖", + "cost": "費用", + "usedOf": "{{used}} / {{limit}}", + "scope": { + "direct": "不使用代理的聊天", + "agent": "透過代理" + } } }, "logs": { @@ -1262,6 +1273,10 @@ "running": "執行中…", "denied": "已被使用者拒絕", "failed": "失敗" + }, + "quotaExceeded": { + "tokens": "您已使用 {{limit}} 權杖配額中的 {{used}}。配額將於 {{resetsAt}} 重設。", + "cost": "您已使用 {{limit}} 用量預算中的 {{used}}。預算將於 {{resetsAt}} 重設。" } }, "agents": { diff --git a/frontend/src/locale/zh.json b/frontend/src/locale/zh.json index 93ba0088..7e647c5e 100644 --- a/frontend/src/locale/zh.json +++ b/frontend/src/locale/zh.json @@ -375,6 +375,17 @@ "toolCalls": "工具调用", "runSuccess": "运行成功率", "feedback": "反馈" + }, + "quota": { + "title": "您的用量配额", + "resets": "{{resetsAt}} 重置", + "tokens": "令牌", + "cost": "费用", + "usedOf": "{{used}} / {{limit}}", + "scope": { + "direct": "不使用代理的聊天", + "agent": "通过代理" + } } }, "logs": { @@ -1262,6 +1273,10 @@ "running": "正在运行…", "denied": "已被用户拒绝", "failed": "失败" + }, + "quotaExceeded": { + "tokens": "您已使用 {{limit}} 令牌配额中的 {{used}}。配额将于 {{resetsAt}} 重置。", + "cost": "您已使用 {{limit}} 用量预算中的 {{used}}。预算将于 {{resetsAt}} 重置。" } }, "agents": { diff --git a/frontend/src/settings/Analytics.tsx b/frontend/src/settings/Analytics.tsx index 1183f2af..f67855bd 100644 --- a/frontend/src/settings/Analytics.tsx +++ b/frontend/src/settings/Analytics.tsx @@ -26,6 +26,7 @@ import { useDarkTheme, useLoaderState } from '../hooks'; import { selectToken } from '../preferences/preferenceSlice'; import { htmlLegendPlugin } from '../utils/chartUtils'; import { formatDate } from '../utils/dateTimeUtils'; +import UsageQuota from './components/UsageQuota'; /** * Resolve a CSS custom property on `:root` to a concrete color string. @@ -377,6 +378,7 @@ export default function Analytics({ agentId }: AnalyticsProps) { return (
+ {agentId ? null : }

{t('settings.analytics.subtitle')} diff --git a/frontend/src/settings/components/UsageQuota.tsx b/frontend/src/settings/components/UsageQuota.tsx new file mode 100644 index 00000000..a2147beb --- /dev/null +++ b/frontend/src/settings/components/UsageQuota.tsx @@ -0,0 +1,140 @@ +import { useEffect, useState } from 'react'; +import { useTranslation } from 'react-i18next'; +import { useSelector } from 'react-redux'; + +import { usagePercent } from '../../admin/quotaUtils'; +import userService from '../../api/services/userService'; +import { selectToken } from '../../preferences/preferenceSlice'; + +type Budget = { limit: number | null; used: number }; +type Bucket = { + bucket: string; + tokens: Budget; + cost: Budget; + resets_at: string; +}; + +function Meter({ + label, + budget, + format, +}: { + label: string; + budget: Budget; + format: (value: number) => string; +}) { + const { t } = useTranslation(); + if (budget.limit === null) return null; + const percent = usagePercent(budget.used, budget.limit); + const tone = + percent >= 100 + ? 'bg-red-500' + : percent >= 80 + ? 'bg-amber-500' + : 'bg-[#7D54D1]'; + return ( +

+
+ {label} + + {t('settings.analytics.quota.usedOf', { + used: format(budget.used), + limit: format(budget.limit), + })} + +
+
+
+
+
+ ); +} + +/** The caller's usage against the quota an admin set; renders nothing when unlimited. */ +export default function UsageQuota() { + const { t, i18n } = useTranslation(); + const token = useSelector(selectToken); + const [buckets, setBuckets] = useState([]); + + useEffect(() => { + let cancelled = false; + userService + .getQuota(token) + .then((res: Response) => (res.ok ? res.json() : null)) + .then((json: { buckets?: Bucket[] } | null) => { + if (cancelled) return; + setBuckets(json?.buckets ?? []); + }) + .catch(() => undefined); + return () => { + cancelled = true; + }; + }, [token]); + + if (buckets.length === 0) return null; + + const number = new Intl.NumberFormat(i18n.language); + const usd = new Intl.NumberFormat(i18n.language, { + style: 'currency', + currency: 'USD', + }); + const reset = new Date(buckets[0].resets_at); + const resetsAt = Number.isNaN(reset.getTime()) + ? '' + : new Intl.DateTimeFormat(i18n.language, { + dateStyle: 'medium', + timeStyle: 'short', + }).format(reset); + + // A request must fit its own bucket and ``all``, so each limited one is shown. + const scopeLabel = (name: string) => + name === 'direct' || name === 'agent' + ? t(`settings.analytics.quota.scope.${name}`) + : null; + + return ( +
+
+

+ {t('settings.analytics.quota.title')} +

+ {resetsAt ? ( +

+ {t('settings.analytics.quota.resets', { resetsAt })} +

+ ) : null} +
+ {buckets.map((bucket) => ( +
+ {scopeLabel(bucket.bucket) ? ( +

+ {scopeLabel(bucket.bucket)} +

+ ) : null} +
+ number.format(value)} + /> + usd.format(value)} + /> +
+
+ ))} +
+ ); +} diff --git a/tests/agents/test_workflow_agent_types.py b/tests/agents/test_workflow_agent_types.py index 41bf7cd4..18bf22d3 100644 --- a/tests/agents/test_workflow_agent_types.py +++ b/tests/agents/test_workflow_agent_types.py @@ -214,6 +214,27 @@ class TestWorkflowEngineAgenticNode: assert engine.state["node_agent_agentic_output"] == "agentic answer" assert engine.state["result"] == "agentic answer" + def test_node_usage_is_attributed_to_the_workflow_agent(self, monkeypatch): + engine = create_engine() + engine.agent.agent_id = "11111111-1111-1111-1111-111111111111" + node = create_agent_node(node_id="agent_attr", agent_type="classic") + + captured: Dict[str, Any] = {} + + def capture_create(**kwargs): + captured.update(kwargs) + return StubNodeAgent([{"answer": "ok"}]) + + monkeypatch.setattr(WorkflowNodeAgentFactory, "create", staticmethod(capture_create)) + monkeypatch.setattr( + "docsgpt.core.model_utils.get_api_key_for_provider", + lambda _provider: None, + ) + + list(engine._execute_agent_node(node)) + + assert captured["agent_id"] == "11111111-1111-1111-1111-111111111111" + def test_agentic_node_passes_retriever_config(self, monkeypatch): engine = create_engine() # The node-source authorization gate is exercised separately; diff --git a/tests/api/test_quota_endpoints.py b/tests/api/test_quota_endpoints.py new file mode 100644 index 00000000..69224f9c --- /dev/null +++ b/tests/api/test_quota_endpoints.py @@ -0,0 +1,287 @@ +"""Endpoint tests for the admin quota API and ``GET /api/user/quota``. + +Driven through the real app.py chokepoint against an ephemeral Postgres; only +``handle_auth`` / ``resolve_roles`` are patched. +""" + +from __future__ import annotations + +import json +from contextlib import ExitStack, contextmanager +from unittest.mock import patch + +import pytest +from sqlalchemy import text + +from docsgpt.storage.db.repositories.quota_policies import QuotaPoliciesRepository +from docsgpt.storage.db.repositories.token_usage import TokenUsageRepository +from docsgpt.storage.db.repositories.users import UsersRepository + + +@pytest.fixture +def client(): + from docsgpt.app import app as flask_app + + flask_app.config["TESTING"] = True + return flask_app.test_client() + + +@pytest.fixture +def db(pg_conn): + @contextmanager + def _yield(): + yield pg_conn + + with ExitStack() as stack: + for target in ( + "docsgpt.api.admin.quotas.db_readonly", + "docsgpt.api.admin.quotas.db_session", + "docsgpt.quotas.service.db_readonly", + ): + stack.enter_context(patch(target, _yield)) + yield pg_conn + + +@contextmanager +def _as(sub, *roles): + with patch("docsgpt.app.handle_auth", return_value={"sub": sub}), patch( + "docsgpt.app.resolve_roles", return_value=list(roles) or ["user"] + ): + yield + + +def _admin(): + return _as("admin1", "admin", "user") + + +def _body(resp): + return json.loads(resp.data) + + +def _team(conn, slug="q-team", member=None): + team_id = str( + conn.execute( + text("INSERT INTO teams (name, slug, owner_id) VALUES (:n, :s, 'o') RETURNING id"), + {"n": slug, "s": slug}, + ).scalar() + ) + if member: + conn.execute( + text("INSERT INTO team_members (team_id, user_id, role) VALUES (CAST(:t AS uuid), :u, 'team_member')"), + {"t": team_id, "u": member}, + ) + return team_id + + +ADMIN_ROUTES = [ + ("get", "/api/admin/quotas"), + ("put", "/api/admin/quotas/instance"), + ("delete", "/api/admin/quotas/instance"), + ("get", "/api/admin/quotas/users/u1"), + ("put", "/api/admin/quotas/users/u1"), + ("delete", "/api/admin/quotas/users/u1"), + ("get", "/api/admin/quotas/teams/00000000-0000-0000-0000-000000000000"), + ("put", "/api/admin/quotas/teams/00000000-0000-0000-0000-000000000000"), + ("delete", "/api/admin/quotas/teams/00000000-0000-0000-0000-000000000000"), +] + + +class TestGuard: + @pytest.mark.parametrize("method, path", ADMIN_ROUTES) + def test_non_admin_forbidden(self, client, db, method, path): + with _as("u1"): + assert getattr(client, method)(path, json={"token_limit": 1}).status_code == 403 + + @pytest.mark.parametrize("method, path", ADMIN_ROUTES) + def test_unauthenticated(self, client, method, path): + with patch("docsgpt.app.handle_auth", return_value=None): + assert getattr(client, method)(path, json={"token_limit": 1}).status_code == 401 + + def test_team_admin_cannot_set_their_teams_allowance(self, client, db): + team_id = _team(db) + db.execute( + text("INSERT INTO team_members (team_id, user_id, role) VALUES (CAST(:t AS uuid), 'lead', 'team_admin')"), + {"t": team_id}, + ) + with _as("lead"): + resp = client.put(f"/api/admin/quotas/teams/{team_id}", json={"token_unlimited": True}) + assert resp.status_code == 403 + assert QuotaPoliciesRepository(db).get("team", team_id) is None + + +class TestInstancePolicy: + def test_set_read_delete(self, client, db): + with _admin(): + put = client.put("/api/admin/quotas/instance", json={"token_limit": 1000, "note": " default "}) + assert put.status_code == 200 + policy = _body(put)["policy"] + assert (policy["token_limit"], policy["note"], policy["updated_by"]) == (1000, "default", "admin1") + + client.put("/api/admin/quotas/instance", json={"bucket": "agent", "cost_limit_usd": 2.5}) + overview = _body(client.get("/api/admin/quotas")) + assert [(p["bucket"], p["token_limit"], p["cost_limit_usd"]) for p in overview["instance"]] == [ + ("all", 1000, None), + ("agent", None, 2.5), + ] + assert overview["period"] == "month" + + one_bucket = _body(client.delete("/api/admin/quotas/instance?bucket=agent")) + the_rest = _body(client.delete("/api/admin/quotas/instance")) + remaining = _body(client.get("/api/admin/quotas"))["instance"] + assert (one_bucket["deleted"], the_rest["deleted"], remaining) == (1, 1, []) + + def test_writes_are_audited(self, client, db): + with _admin(): + client.put("/api/admin/quotas/instance", json={"token_limit": 5}) + client.delete("/api/admin/quotas/instance") + client.delete("/api/admin/quotas/instance") + events = db.execute( + text("SELECT user_id, event, metadata FROM auth_events WHERE event LIKE 'quota_policy_%'") + ).fetchall() + by_event = {e[1]: e for e in events} + # The second delete removed nothing, so it left no event. + assert sorted((e[0], e[1]) for e in events) == [ + ("admin1", "quota_policy_deleted"), + ("admin1", "quota_policy_set"), + ] + metadata = by_event["quota_policy_set"][2] + assert metadata["token_limit"] == 5 and metadata["by"] == "admin1" + + @pytest.mark.parametrize( + "body", + [ + None, + [], + {}, + {"note": "only a note"}, + {"token_limit": -1}, + {"token_limit": 1.5}, + {"token_limit": True}, + {"token_limit": "10"}, + {"token_limit": 2**63}, + {"token_limit": 10**400}, + {"cost_limit_usd": 10**400}, + {"cost_limit_usd": -0.01}, + {"cost_limit_usd": "5"}, + {"cost_limit_usd": float("inf")}, + {"cost_limit_usd": 1e12}, + {"token_limit": 1, "token_unlimited": True}, + {"cost_limit_usd": 1, "cost_unlimited": True}, + {"token_unlimited": "yes"}, + {"token_limit": 1, "enabled": "no"}, + {"token_limit": 1, "bucket": "everything"}, + {"token_limit": 1, "note": 7}, + ], + ) + def test_invalid_bodies_rejected(self, client, db, body): + with _admin(): + resp = client.put("/api/admin/quotas/instance", json=body) + assert resp.status_code == 400 + assert QuotaPoliciesRepository(db).list_by_scope("instance") == [] + + def test_unknown_bucket_on_delete(self, client, db): + with _admin(): + resp = client.delete("/api/admin/quotas/instance?bucket=nope") + assert resp.status_code == 400 + + def test_zero_is_accepted_as_a_block(self, client, db): + with _admin(): + resp = client.put("/api/admin/quotas/instance", json={"token_limit": 0, "cost_limit_usd": 0}) + assert resp.status_code == 200 + assert _body(resp)["policy"]["token_limit"] == 0 + + +class TestTeamPolicy: + def test_set_and_list_with_team_details(self, client, db): + team_id = _team(db, "q-eng", member="u1") + with _admin(): + assert client.put(f"/api/admin/quotas/teams/{team_id}", json={"token_limit": 500}).status_code == 200 + (row,) = _body(client.get("/api/admin/quotas"))["teams"] + assert (row["team_slug"], row["token_limit"], row["member_count"]) == ("q-eng", 500, 1) + assert _body(client.get(f"/api/admin/quotas/teams/{team_id}"))["policies"][0]["token_limit"] == 500 + + @pytest.mark.parametrize("team_id", ["not-a-uuid", "00000000-0000-0000-0000-000000000000"]) + def test_unknown_team(self, client, db, team_id): + with _admin(): + assert client.put(f"/api/admin/quotas/teams/{team_id}", json={"token_limit": 1}).status_code == 404 + assert client.get(f"/api/admin/quotas/teams/{team_id}").status_code == 404 + + +class TestUserPolicy: + def test_unknown_user(self, client, db): + with _admin(): + assert client.put("/api/admin/quotas/users/ghost", json={"token_limit": 1}).status_code == 404 + assert client.get("/api/admin/quotas/users/ghost").status_code == 404 + + def test_effective_limits_name_their_source(self, client, db): + UsersRepository(db).upsert("u1") + small, big = _team(db, "q-small", member="u1"), _team(db, "q-big", member="u1") + repo = QuotaPoliciesRepository(db) + repo.upsert(scope="instance", subject_id=None, token_limit=10, cost_limit_usd=1) + repo.upsert(scope="team", subject_id=small, token_limit=100) + repo.upsert(scope="team", subject_id=big, token_limit=900) + TokenUsageRepository(db).insert(user_id="u1", prompt_tokens=40, cost=0.25) + + with _admin(): + body = _body(client.get("/api/admin/quotas/users/u1")) + overall = body["effective"][0] + assert overall["bucket"] == "all" + assert overall["tokens"] == {"limit": 900, "used": 40, "source": "team", "source_id": big} + assert isinstance(overall["tokens"]["limit"], int) + assert overall["cost"] == {"limit": 1.0, "used": 0.25, "source": "instance", "source_id": None} + assert body["policies"] == [] + + with _admin(): + client.put("/api/admin/quotas/users/u1", json={"token_limit": 50}) + body = _body(client.get("/api/admin/quotas/users/u1")) + assert body["effective"][0]["tokens"]["source"] == "user" + assert body["policies"][0]["token_limit"] == 50 + + def test_user_policy_audit_is_filed_under_the_user(self, client, db): + UsersRepository(db).upsert("u1") + with _admin(): + client.put("/api/admin/quotas/users/u1", json={"cost_unlimited": True}) + row = db.execute( + text("SELECT user_id, metadata FROM auth_events WHERE event = 'quota_policy_set'") + ).one() + assert row[0] == "u1" and row[1]["by"] == "admin1" and row[1]["scope"] == "user" + + +class TestUnpricedModels: + def test_lists_models_recorded_at_zero_for_want_of_a_price(self, client, db): + usage = TokenUsageRepository(db) + usage.insert(user_id="u1", prompt_tokens=10, model_id="local-llama") + # Priced when called; its provider may be disabled by now. + usage.insert(user_id="u1", prompt_tokens=5, model_id="retired-priced-model", cost=0.1) + usage.insert(user_id="u1", prompt_tokens=3, model_id="free-model") + usage.insert(user_id="u1", prompt_tokens=7, model_id="7d0c1a52-2f5e-4c53-9a0e-111111111111") + with _admin(), patch("docsgpt.api.admin.quotas.is_priced", lambda m: m == "free-model"): + unpriced = _body(client.get("/api/admin/quotas"))["unpriced_models"] + assert unpriced == [{"model_id": "local-llama", "tokens": 10, "cost": 0.0}] + + +class TestMyQuota: + def test_unlimited_user_sees_no_buckets(self, client, db): + with _as("u1"): + body = _body(client.get("/api/user/quota")) + assert body == {"success": True, "period": "month", "buckets": []} + + def test_limited_user_sees_usage_without_policy_internals(self, client, db): + QuotaPoliciesRepository(db).upsert(scope="user", subject_id="u1", token_limit=100, note="secret note") + TokenUsageRepository(db).insert(user_id="u1", prompt_tokens=30) + with _as("u1"): + body = _body(client.get("/api/user/quota")) + (bucket,) = body["buckets"] + assert bucket["bucket"] == "all" + assert bucket["tokens"] == {"limit": 100, "used": 30} + assert bucket["cost"] == {"limit": None, "used": 0.0} + assert "source" not in json.dumps(body) and "secret" not in json.dumps(body) + + def test_a_user_only_sees_their_own_quota(self, client, db): + QuotaPoliciesRepository(db).upsert(scope="user", subject_id="u1", token_limit=100) + with _as("u2"): + assert _body(client.get("/api/user/quota"))["buckets"] == [] + + def test_unauthenticated(self, client): + with patch("docsgpt.app.handle_auth", return_value=None): + assert client.get("/api/user/quota").status_code == 401 diff --git a/tests/api/user/test_scheduler_worker.py b/tests/api/user/test_scheduler_worker.py index 5b7b55ea..09528cd5 100644 --- a/tests/api/user/test_scheduler_worker.py +++ b/tests/api/user/test_scheduler_worker.py @@ -126,6 +126,31 @@ class TestExecuteScheduledRunBody: assert sched["consecutive_failure_count"] == 1 assert "schedule.run.failed" in {e[0] for e in stub_events} + def test_quota_refusal_marks_budget_exceeded( + self, pg_engine, patched_engine, stub_events, + ): + from datetime import datetime, timezone + + from docsgpt.quotas.service import QuotaExceeded, QuotaExceededError + + exceeded = QuotaExceeded( + user_id="u1", bucket="all", budget="tokens", usage=10, limit=10, + source="instance", source_id=None, + resets_at=datetime(2099, 1, 1, tzinfo=timezone.utc), + ) + with pg_engine.begin() as conn: + _, run, _ = _make_pending_run(conn) + with patch( + "docsgpt.api.user.scheduler_worker.run_agent_headless", + side_effect=QuotaExceededError(exceeded), + ): + result = execute_scheduled_run_body(str(run["id"]), "celery-q") + assert result["status"] == "failed" + with pg_engine.connect() as conn: + row = ScheduleRunsRepository(conn).get_internal(str(run["id"])) + assert row["error_type"] == "budget_exceeded" + assert "Usage quota reached" in row["error"] + def test_autopause_after_threshold( self, pg_engine, patched_engine, stub_events, ): diff --git a/tests/core/test_model_settings.py b/tests/core/test_model_settings.py index 266c9426..288ac6c9 100644 --- a/tests/core/test_model_settings.py +++ b/tests/core/test_model_settings.py @@ -46,8 +46,8 @@ class TestModelCapabilities: assert caps.supports_streaming is True assert caps.supported_attachment_types == [] assert caps.context_window == 128000 - assert caps.input_cost_per_token is None - assert caps.output_cost_per_token is None + assert caps.input_cost_per_million is None + assert caps.output_cost_per_million is None @pytest.mark.unit def test_custom_values(self): @@ -55,7 +55,7 @@ class TestModelCapabilities: supports_tools=True, supports_structured_output=True, context_window=32000, - input_cost_per_token=0.001, + input_cost_per_million=1.0, ) assert caps.supports_tools is True assert caps.context_window == 32000 diff --git a/tests/llm/test_fallback.py b/tests/llm/test_fallback.py index ceba9b65..03974bb5 100644 --- a/tests/llm/test_fallback.py +++ b/tests/llm/test_fallback.py @@ -1015,6 +1015,38 @@ class TestLLMCreatorPassesModelUserId: assert captured["model_user_id"] == "owner-alice" + @pytest.mark.parametrize( + "model_id, source, expected", + [(None, None, False), ("catalog-model", "builtin", False), ("byom-uuid", "user", True)], + ) + def test_byom_flag_follows_the_model_source(self, monkeypatch, model_id, source, expected): + from types import SimpleNamespace + + from docsgpt.llm.llm_creator import LLMCreator + from docsgpt.llm.providers import PROVIDERS_BY_NAME + + class _LLM: + def __init__(self, *args, **kwargs): + pass + + monkeypatch.setattr(PROVIDERS_BY_NAME["openai"], "llm_class", _LLM) + model = SimpleNamespace( + source=source, api_key="own-key", base_url=None, upstream_model_id=None, capabilities=None + ) + registry = SimpleNamespace(get_model=lambda _id, user_id=None: model) + monkeypatch.setattr( + "docsgpt.core.model_registry.ModelRegistry.get_instance", lambda: registry + ) + + llm = LLMCreator.create_llm( + type="openai", api_key="k", user_api_key=None, + decoded_token={"sub": "u1"}, model_id=model_id, + ) + + assert llm._is_byom is expected + assert llm._canonical_model_id == model_id + + # Tests — responding-provider tracking (cross-provider fallback handler fix) 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_enforcement.py b/tests/quotas/test_enforcement.py new file mode 100644 index 00000000..0a54974c --- /dev/null +++ b/tests/quotas/test_enforcement.py @@ -0,0 +1,194 @@ +"""Tests for quota enforcement at the request entry points.""" + +from __future__ import annotations + +import json +from contextlib import contextmanager +from unittest.mock import patch + +import pytest + +from docsgpt.quotas.service import QuotaExceededError +from docsgpt.storage.db.repositories.agents import AgentsRepository +from docsgpt.storage.db.repositories.quota_policies import QuotaPoliciesRepository +from docsgpt.storage.db.repositories.token_usage import TokenUsageRepository + + +@pytest.fixture +def db(pg_conn): + @contextmanager + def _yield(): + yield pg_conn + + with patch("docsgpt.quotas.service.db_readonly", _yield), patch( + "docsgpt.api.answer.routes.base.db_readonly", _yield + ), patch("docsgpt.agents.headless_runner.db_readonly", _yield): + yield pg_conn + + +def _spend(conn, user_id, tokens, api_key=None): + TokenUsageRepository(conn).insert(user_id=user_id, api_key=api_key, prompt_tokens=tokens) + + +def _check(flask_app, agent_config, decoded_token=None, agent_id=None): + from docsgpt.api.answer.routes.base import BaseAnswerResource + + with flask_app.app_context(): + return BaseAnswerResource().check_usage(agent_config, decoded_token, agent_id=agent_id) + + +class TestCheckUsage: + def test_direct_chat_is_limited(self, db, flask_app): + QuotaPoliciesRepository(db).upsert(scope="user", subject_id="u1", token_limit=100) + _spend(db, "u1", 100) + + response = _check(flask_app, {}, {"sub": "u1"}) + + assert response.status_code == 429 + assert int(response.headers["Retry-After"]) >= 1 + body = json.loads(response.data) + assert body["error_code"] == "quota-exceeded" + assert (body["dimension"], body["usage"], body["limit"], body["source"]) == ("tokens", 100, 100, "user") + + def test_direct_chat_under_the_limit_passes(self, db, flask_app): + QuotaPoliciesRepository(db).upsert(scope="user", subject_id="u1", token_limit=100) + _spend(db, "u1", 99) + assert _check(flask_app, {}, {"sub": "u1"}) is None + + def test_no_policies_and_no_identity_pass(self, db, flask_app): + assert _check(flask_app, {}) is None + assert _check(flask_app, {}, {"sub": "u1"}) is None + + def test_agent_traffic_bills_the_resolved_user(self, db, flask_app): + AgentsRepository(db).create("owner", "a", "published", key="k1") + QuotaPoliciesRepository(db).upsert(scope="user", subject_id="owner", token_limit=10) + _spend(db, "owner", 10, api_key="k1") + + config = {"user_api_key": "k1", "user_id": "owner"} + assert _check(flask_app, config, {"sub": "owner"}).status_code == 429 + # A shared agent bills the caller, who has room. + assert _check(flask_app, config, {"sub": "caller"}) is None + # No token resolved: fall back to the agent owner. + assert _check(flask_app, config).status_code == 429 + + def test_quota_runs_before_the_agent_key_lookup(self, db, flask_app): + QuotaPoliciesRepository(db).upsert(scope="user", subject_id="u1", token_limit=0) + response = _check(flask_app, {"user_api_key": "missing"}, {"sub": "u1"}) + assert response.status_code == 429 + + def test_agent_limits_still_apply(self, db, flask_app): + AgentsRepository(db).create( + "owner", "a", "published", key="k2", limited_token_mode=True, token_limit=5 + ) + _spend(db, "owner", 5, api_key="k2") + response = _check(flask_app, {"user_api_key": "k2"}, {"sub": "owner"}) + assert response.status_code == 429 + assert "error_code" not in json.loads(response.data) + + def test_agent_bucket_policy_ignores_direct_chat(self, db, flask_app): + QuotaPoliciesRepository(db).upsert(scope="user", subject_id="u1", bucket="agent", token_limit=10) + _spend(db, "u1", 500) + assert _check(flask_app, {}, {"sub": "u1"}) is None + + def test_a_keyless_agent_chat_is_agent_traffic(self, db, flask_app): + agent_id = str(AgentsRepository(db).create("u1", "draft", "draft")["id"]) + QuotaPoliciesRepository(db).upsert(scope="user", subject_id="u1", bucket="agent", token_limit=10) + TokenUsageRepository(db).insert(user_id="u1", agent_id=agent_id, prompt_tokens=10) + + response = _check(flask_app, {"user_api_key": None}, {"sub": "u1"}, agent_id=agent_id) + + assert response.status_code == 429 + assert json.loads(response.data)["bucket"] == "agent" + # The same spend leaves chat without an agent alone. + assert _check(flask_app, {}, {"sub": "u1"}) is None + + def test_a_refusal_tells_sdk_clients_not_to_retry(self, db, flask_app): + QuotaPoliciesRepository(db).upsert(scope="user", subject_id="u1", token_limit=0) + response = _check(flask_app, {}, {"sub": "u1"}) + assert response.headers["x-should-retry"] == "false" + + +class TestHeadless: + def test_exhausted_owner_is_refused_before_the_run(self, db): + from docsgpt.agents.headless_runner import run_agent_headless + + QuotaPoliciesRepository(db).upsert(scope="instance", subject_id=None, token_limit=10) + _spend(db, "owner", 10) + + with patch("docsgpt.agents.headless_runner.RetrieverCreator") as retriever: + with pytest.raises(QuotaExceededError) as raised: + run_agent_headless({"user_id": "owner", "key": "k"}, "hello") + + retriever.create_retriever.assert_not_called() + assert raised.value.exceeded.source == "instance" + assert "10 of 10 tokens" in str(raised.value) + + + + def test_a_keyless_agent_run_is_agent_traffic(self, db): + from docsgpt.agents.headless_runner import run_agent_headless + + QuotaPoliciesRepository(db).upsert(scope="user", subject_id="owner", bucket="agent", token_limit=0) + config = {"user_id": "owner", "id": "22222222-2222-2222-2222-222222222222"} + + with patch("docsgpt.agents.headless_runner.RetrieverCreator"): + with pytest.raises(QuotaExceededError) as raised: + run_agent_headless(config, "hello") + + assert raised.value.exceeded.bucket == "agent" + + +class TestResumeRefusal: + def _processor(self, user_id="u1"): + from types import SimpleNamespace + + return SimpleNamespace( + agent_config={}, decoded_token={"sub": user_id}, initial_user_id=user_id, agent_id=None + ) + + def test_a_refused_resume_releases_its_claim(self, db, flask_app): + from docsgpt.api.answer.routes.base import BaseAnswerResource + + QuotaPoliciesRepository(db).upsert(scope="user", subject_id="u1", token_limit=0) + with flask_app.app_context(), patch( + "docsgpt.api.answer.routes.base.ContinuationService" + ) as service: + response = BaseAnswerResource().check_usage_on_resume(self._processor(), "conv-1") + + assert response.status_code == 429 + service.return_value.release_claim.assert_called_once_with("conv-1", "u1") + + def test_an_admitted_resume_keeps_its_claim(self, db, flask_app): + from docsgpt.api.answer.routes.base import BaseAnswerResource + + with flask_app.app_context(), patch( + "docsgpt.api.answer.routes.base.ContinuationService" + ) as service: + response = BaseAnswerResource().check_usage_on_resume(self._processor(), "conv-1") + + assert response is None + service.return_value.release_claim.assert_not_called() + + def test_no_claim_means_nothing_to_release(self, db, flask_app): + from docsgpt.api.answer.routes.base import BaseAnswerResource + + QuotaPoliciesRepository(db).upsert(scope="user", subject_id="u1", token_limit=0) + with flask_app.app_context(), patch( + "docsgpt.api.answer.routes.base.ContinuationService" + ) as service: + response = BaseAnswerResource().check_usage_on_resume(self._processor(), None) + + assert response.status_code == 429 + service.return_value.release_claim.assert_not_called() + + def test_a_failed_release_still_returns_the_refusal(self, db, flask_app): + from docsgpt.api.answer.routes.base import BaseAnswerResource + + QuotaPoliciesRepository(db).upsert(scope="user", subject_id="u1", token_limit=0) + with flask_app.app_context(), patch( + "docsgpt.api.answer.routes.base.ContinuationService" + ) as service: + service.return_value.release_claim.side_effect = RuntimeError("db down") + response = BaseAnswerResource().check_usage_on_resume(self._processor(), "conv-1") + + assert response.status_code == 429 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..13975bb3 --- /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, "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..83a9302c --- /dev/null +++ b/tests/storage/db/repositories/test_quota_policies.py @@ -0,0 +1,144 @@ +"""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) + first = repo.delete("user", "u1", "agent") + again = repo.delete("user", "u1", "agent") + rest = repo.delete("user", "u1") + assert (first, again, rest) == (1, 0, 1) + assert repo.get("user", "u2") is not None diff --git a/tests/storage/db/repositories/test_token_usage.py b/tests/storage/db/repositories/test_token_usage.py index 46f59d11..6ce4cb96 100644 --- a/tests/storage/db/repositories/test_token_usage.py +++ b/tests/storage/db/repositories/test_token_usage.py @@ -35,6 +35,63 @@ class TestInsert: ) assert total == 30 + def test_cost_defaults_to_zero_and_round_trips(self, pg_conn): + repo = _repo(pg_conn) + repo.insert(user_id="u-cost", prompt_tokens=1, generated_tokens=1) + repo.insert(user_id="u-cost", prompt_tokens=1, generated_tokens=1, cost=0.00012345) + costs = pg_conn.execute( + text("SELECT cost FROM token_usage WHERE user_id = 'u-cost' ORDER BY id") + ).scalars().all() + assert [float(c) for c in costs] == [0.0, 0.00012345] + + +class TestRollupExclusion: + def test_sum_tokens_ignores_schedule_rollups(self, pg_conn): + repo = _repo(pg_conn) + repo.insert(user_id="u-roll", prompt_tokens=10, generated_tokens=5, source="agent_stream") + repo.insert(user_id="u-roll", prompt_tokens=10, generated_tokens=5, source="schedule") + total = repo.sum_tokens_in_range( + start=_now() - timedelta(minutes=1), end=_now() + timedelta(minutes=1), user_id="u-roll" + ) + assert total == 15 + + +class TestUsageTotals: + def _seed(self, repo): + repo.insert(user_id="u-tot", prompt_tokens=100, generated_tokens=10, cost=0.5) + repo.insert(user_id="u-tot", api_key="k", prompt_tokens=20, generated_tokens=2, cost=0.25) + repo.insert(user_id="u-tot", prompt_tokens=7, generated_tokens=0, cost=0.125, source="title") + # A keyless agent (or workflow node): an agent id without a key. + from docsgpt.storage.db.repositories.agents import AgentsRepository + + agent = AgentsRepository(repo._conn).create("u-tot", "keyless", "draft") + repo.insert( + user_id="u-tot", agent_id=str(agent["id"]), + prompt_tokens=4, generated_tokens=0, cost=0.0625, + ) + repo.insert(user_id="u-tot", prompt_tokens=999, generated_tokens=0, source="schedule") + repo.insert(user_id="u-other", prompt_tokens=999, generated_tokens=0, cost=9) + repo.insert( + user_id="u-tot", prompt_tokens=999, generated_tokens=0, cost=9, + timestamp=_now() - timedelta(days=40), + ) + + @pytest.mark.parametrize( + "bucket, expected", + [("all", (143, 0.9375)), ("direct", (117, 0.625)), ("agent", (26, 0.3125))], + ) + def test_totals_per_bucket(self, pg_conn, bucket, expected): + repo = _repo(pg_conn) + self._seed(repo) + assert repo.usage_totals(user_id="u-tot", start=_now() - timedelta(days=1), bucket=bucket) == expected + + def test_no_usage_is_zero(self, pg_conn): + assert _repo(pg_conn).usage_totals(user_id="nobody", start=_now() - timedelta(days=1)) == (0, 0.0) + + def test_unknown_bucket_rejected(self, pg_conn): + with pytest.raises(ValueError): + _repo(pg_conn).usage_totals(user_id="u", start=_now(), bucket="nope") + class TestReassignApiKey: def test_rewrites_and_preserves_rate_limit_window(self, pg_conn): diff --git a/tests/storage/db/test_migration_0033.py b/tests/storage/db/test_migration_0033.py new file mode 100644 index 00000000..09cb1d4f --- /dev/null +++ b/tests/storage/db/test_migration_0033.py @@ -0,0 +1,145 @@ +"""Migration round-trip test for 0033_quotas.""" + +from __future__ import annotations + +import os +import subprocess +import sys +from pathlib import Path + +import pytest +from sqlalchemy import text +from sqlalchemy.exc import IntegrityError + + +pytestmark = pytest.mark.integration + + +def _alembic_ini() -> Path: + return Path(__file__).resolve().parents[3] / "docsgpt" / "alembic.ini" + + +def _run_alembic(url: str, *args: str) -> None: + subprocess.check_call( + [sys.executable, "-m", "alembic", "-c", str(_alembic_ini()), *args], + timeout=60, + env={**os.environ, "POSTGRES_URI": url}, + ) + + +def _alembic_heads(url: str) -> list[str]: + out = subprocess.check_output( + [sys.executable, "-m", "alembic", "-c", str(_alembic_ini()), "heads"], + timeout=60, + env={**os.environ, "POSTGRES_URI": url}, + text=True, + ) + return [line for line in out.splitlines() if line.strip()] + + +def _alembic_version(conn) -> str: + return conn.execute(text("SELECT version_num FROM alembic_version")).scalar() + + +def _column_exists(conn, table: str, column: str) -> bool: + row = conn.execute( + text( + "SELECT 1 FROM information_schema.columns " + "WHERE table_name = :t AND column_name = :c AND table_schema = 'public'" + ), + {"t": table, "c": column}, + ).fetchone() + return row is not None + + +_0033 = "0033_quotas" +_0032 = "0032_personal_access_tokens" + + +def _table_exists(conn, table: str) -> bool: + return conn.execute(text("SELECT to_regclass(:t)"), {"t": f"public.{table}"}).scalar() is not None + + +def _insert_policy(conn, **values) -> None: + cols = ", ".join(values) + params = ", ".join(f":{k}" for k in values) + conn.execute(text(f"INSERT INTO quota_policies ({cols}) VALUES ({params})"), values) + + +class TestMigration0033RoundTrip: + def test_single_head(self, pg_engine): + url = pg_engine.url.render_as_string(hide_password=False) + assert len(_alembic_heads(url)) == 1 + + def test_head_has_quota_schema(self, pg_engine): + with pg_engine.connect() as conn: + assert _alembic_version(conn) >= _0033 + assert _table_exists(conn, "quota_policies") + assert _column_exists(conn, "token_usage", "cost") + + def test_downgrade_drops_then_upgrade_restores(self, pg_engine): + url = pg_engine.url.render_as_string(hide_password=False) + _run_alembic(url, "downgrade", _0032) + with pg_engine.connect() as conn: + assert _alembic_version(conn) == _0032 + assert not _table_exists(conn, "quota_policies") + assert not _column_exists(conn, "token_usage", "cost") + _run_alembic(url, "upgrade", "head") + with pg_engine.connect() as conn: + assert _table_exists(conn, "quota_policies") + assert _column_exists(conn, "token_usage", "cost") + + def test_upgrade_tolerates_an_existing_cost_column(self, pg_engine): + """A database that already carries ``token_usage.cost`` upgrades cleanly.""" + url = pg_engine.url.render_as_string(hide_password=False) + _run_alembic(url, "downgrade", _0032) + with pg_engine.begin() as conn: + conn.execute(text("ALTER TABLE token_usage ADD COLUMN cost NUMERIC(12,8) NOT NULL DEFAULT 0")) + conn.execute( + text( + "INSERT INTO token_usage (user_id, prompt_tokens, generated_tokens, cost) " + "VALUES ('u-mig33', 10, 1, 0.5)" + ) + ) + _run_alembic(url, "upgrade", "head") + with pg_engine.connect() as conn: + cost = conn.execute(text("SELECT cost FROM token_usage WHERE user_id = 'u-mig33'")).scalar() + assert float(cost) == 0.5 + + +class TestQuotaPolicyConstraints: + def test_one_row_per_subject_and_bucket(self, pg_conn): + _insert_policy(pg_conn, scope="instance", token_limit=10) + with pytest.raises(IntegrityError): + with pg_conn.begin_nested(): + _insert_policy(pg_conn, scope="instance", token_limit=20) + _insert_policy(pg_conn, scope="instance", bucket="agent", token_limit=20) + + @pytest.mark.parametrize( + "values", + [ + {"scope": "instance", "subject_id": "u1"}, + {"scope": "user"}, + {"scope": "user", "subject_id": "u1", "token_limit": 5, "token_unlimited": True}, + {"scope": "user", "subject_id": "u1", "cost_limit_usd": 5, "cost_unlimited": True}, + {"scope": "user", "subject_id": "u1", "token_limit": -1}, + {"scope": "user", "subject_id": "u1", "bucket": "nope"}, + {"scope": "org", "subject_id": "u1"}, + ], + ) + def test_invalid_rows_rejected(self, pg_conn, values): + with pytest.raises(IntegrityError): + with pg_conn.begin_nested(): + _insert_policy(pg_conn, **values) + + def test_deleting_a_team_removes_its_policies(self, pg_conn): + team_id = pg_conn.execute( + text("INSERT INTO teams (name, slug, owner_id) VALUES ('Q', 'q-mig33', 'owner') RETURNING id") + ).scalar() + _insert_policy(pg_conn, scope="team", subject_id=str(team_id), token_limit=10) + _insert_policy(pg_conn, scope="user", subject_id=str(team_id), token_limit=10) + pg_conn.execute(text("DELETE FROM teams WHERE id = :id"), {"id": team_id}) + scopes = pg_conn.execute( + text("SELECT scope FROM quota_policies WHERE subject_id = :id"), {"id": str(team_id)} + ).scalars().all() + assert scopes == ["user"] diff --git a/tests/test_pricing.py b/tests/test_pricing.py new file mode 100644 index 00000000..d8a6ca69 --- /dev/null +++ b/tests/test_pricing.py @@ -0,0 +1,139 @@ +"""Tests for docsgpt/pricing.py and the per-million catalog fields.""" + +from __future__ import annotations + +from types import SimpleNamespace +from unittest.mock import patch + +import pytest + +from docsgpt import pricing +from docsgpt.core.model_settings import ModelCapabilities +from docsgpt.core.model_yaml import ( + BUILTIN_MODELS_DIR, + ModelYAMLError, + load_model_yamls, +) +from docsgpt.pricing import ModelRates, compute_cost_usd, cost_from_rates + + +def _registry(**models): + entries = {k: SimpleNamespace(capabilities=v) for k, v in models.items()} + return SimpleNamespace(models=entries) + + +@pytest.fixture +def priced_registry(): + caps = ModelCapabilities( + input_cost_per_million=2.0, + output_cost_per_million=10.0, + cached_input_cost_per_million=0.2, + cache_write_cost_per_million=2.5, + ) + bare = ModelCapabilities() + with patch( + "docsgpt.core.model_registry.ModelRegistry.get_instance", + return_value=_registry(priced=caps, bare=bare), + ): + yield + + +@pytest.mark.unit +class TestCostFromRates: + def test_prompt_and_generated(self): + rates = ModelRates(prompt=2.0, generated=10.0) + assert cost_from_rates(rates, 1_000_000, 500_000) == pytest.approx(7.0) + + def test_cache_bins_use_their_rates(self): + rates = ModelRates(prompt=2.0, generated=10.0, cached_input=0.2, cache_write=2.5) + cost = cost_from_rates(rates, 1000, 0, cached_tokens=600, cache_write_tokens=100) + assert cost == pytest.approx((300 * 2.0 + 600 * 0.2 + 100 * 2.5) / 1e6) + + def test_missing_cache_rates_bill_at_prompt_rate(self): + rates = ModelRates(prompt=2.0, generated=10.0) + assert cost_from_rates(rates, 1000, 0, cached_tokens=900) == pytest.approx(1000 * 2.0 / 1e6) + + def test_cache_bins_clamped_to_prompt_total(self): + rates = ModelRates(prompt=2.0, generated=0.0, cached_input=0.0, cache_write=0.0) + assert cost_from_rates(rates, 100, 0, cached_tokens=5000, cache_write_tokens=5000) == 0.0 + assert cost_from_rates(rates, 100, 0, cached_tokens=-5) == pytest.approx(100 * 2.0 / 1e6) + + def test_none_and_negative_counts(self): + rates = ModelRates(prompt=2.0, generated=10.0) + assert cost_from_rates(rates, None, -3, None, None) == 0.0 + + +@pytest.mark.unit +class TestComputeCost: + def test_priced_model(self, priced_registry): + assert compute_cost_usd("priced", 1_000_000, 0) == pytest.approx(2.0) + + @pytest.mark.parametrize("model", ["bare", "unknown", None]) + def test_unpriced_model_is_free_without_fallback(self, priced_registry, model): + with patch.object(pricing.settings, "QUOTA_UNPRICED_RATE_PER_MILLION", None): + assert compute_cost_usd(model, 1_000_000, 1_000_000) == 0.0 + assert pricing.is_priced(model) is False + + def test_unpriced_model_uses_fallback(self, priced_registry): + with patch.object(pricing.settings, "QUOTA_UNPRICED_RATE_PER_MILLION", [0.5, 1.5]): + assert compute_cost_usd("bare", 1_000_000, 1_000_000) == pytest.approx(2.0) + assert pricing.is_priced("bare") is True + + +@pytest.mark.unit +class TestCatalogFields: + def _load(self, tmp_path, body): + (tmp_path / "p.yaml").write_text(body) + return load_model_yamls([tmp_path])[0].models[0].capabilities + + def test_per_million_fields(self, tmp_path): + caps = self._load( + tmp_path, + "provider: openai\nmodels:\n - id: m\n input_cost_per_million: 3\n" + " output_cost_per_million: 15\n cached_input_cost_per_million: 0.3\n", + ) + assert (caps.input_cost_per_million, caps.output_cost_per_million) == (3, 15) + assert caps.cached_input_cost_per_million == 0.3 + assert caps.cache_write_cost_per_million is None + + def test_per_token_alias_is_scaled(self, tmp_path): + caps = self._load( + tmp_path, + "provider: openai\ndefaults:\n input_cost_per_token: 0.000003\n" + "models:\n - id: m\n output_cost_per_token: 0.000015\n", + ) + assert caps.input_cost_per_million == pytest.approx(3.0) + assert caps.output_cost_per_million == pytest.approx(15.0) + + def test_both_spellings_rejected(self, tmp_path): + with pytest.raises(ModelYAMLError): + self._load( + tmp_path, + "provider: openai\nmodels:\n - id: m\n input_cost_per_token: 0.1\n" + " input_cost_per_million: 1\n", + ) + + def test_negative_rate_rejected(self, tmp_path): + with pytest.raises(ModelYAMLError): + self._load(tmp_path, "provider: openai\nmodels:\n - id: m\n input_cost_per_million: -1\n") + + def test_hosted_builtin_models_are_priced(self): + hosted = {"anthropic", "deepseek", "docsgpt", "google", "groq", "novita", "openai", "openrouter"} + catalogs = [ + c for c in load_model_yamls([BUILTIN_MODELS_DIR]) if c.source_path.stem in hosted + ] + assert {c.source_path.stem for c in catalogs} == hosted + for catalog in catalogs: + for model in catalog.models: + caps = model.capabilities + assert caps.input_cost_per_million is not None, model.id + assert caps.output_cost_per_million is not None, model.id + + def test_default_docsgpt_model_rates(self): + (model,) = [ + m for c in load_model_yamls([BUILTIN_MODELS_DIR]) for m in c.models if m.id == "docsgpt-local" + ] + caps = model.capabilities + assert (caps.input_cost_per_million, caps.output_cost_per_million) == (0.15, 0.5) + assert caps.cached_input_cost_per_million == 0.03 + assert caps.cache_write_cost_per_million is None diff --git a/tests/test_usage.py b/tests/test_usage.py index b75639d6..67fbd4b0 100644 --- a/tests/test_usage.py +++ b/tests/test_usage.py @@ -712,3 +712,56 @@ def test_decorator_omits_cache_bins_when_provider_reports_none(monkeypatch): assert row["cache_write_tokens"] is None assert llm.emitted[0]["cached_tokens"] is None assert llm.emitted[0]["cache_write_tokens"] is None + + +def _persist_with_cost(monkeypatch, llm, cost_fn): + from docsgpt.usage import _persist_call_usage + + _install_fake_token_repo(monkeypatch) + monkeypatch.setattr("docsgpt.usage.compute_cost_usd", cost_fn) + _persist_call_usage( + llm, {"prompt_tokens": 1000, "generated_tokens": 10, "cached_tokens": 400} + ) + return _FakeTokenUsageRepo.last_instance.inserted[0] + + +class _CostLLM: + decoded_token = {"sub": "user_123"} + user_api_key = None + agent_id = None + _canonical_model_id = "priced-model" + + +@pytest.mark.unit +def test_persist_prices_the_call_by_canonical_model(monkeypatch): + seen = {} + + def cost_fn(model, prompt, generated, cached_tokens=None, cache_write_tokens=None): + seen.update(model=model, prompt=prompt, generated=generated, cached=cached_tokens) + return 0.0123 + + row = _persist_with_cost(monkeypatch, _CostLLM(), cost_fn) + + assert row["cost"] == 0.0123 + assert seen == {"model": "priced-model", "prompt": 1000, "generated": 10, "cached": 400} + + +@pytest.mark.unit +def test_persist_records_byom_calls_at_zero_cost(monkeypatch): + llm = _CostLLM() + llm._is_byom = True + + row = _persist_with_cost(monkeypatch, llm, lambda *a, **k: 9.9) + + assert row["cost"] == 0.0 + assert row["prompt_tokens"] == 1000 + + +@pytest.mark.unit +def test_persist_keeps_the_row_when_pricing_fails(monkeypatch): + def boom(*args, **kwargs): + raise RuntimeError("registry unavailable") + + row = _persist_with_cost(monkeypatch, _CostLLM(), boom) + + assert row["cost"] == 0.0 diff --git a/tests/worker/test_agent_workers.py b/tests/worker/test_agent_workers.py index 8a863f64..033f97d0 100644 --- a/tests/worker/test_agent_workers.py +++ b/tests/worker/test_agent_workers.py @@ -110,6 +110,36 @@ class TestAgentWebhookWorker: with pytest.raises(RuntimeError, match="LLM exploded"): worker.agent_webhook_worker(task_self, agent_id, {"event": "ping"}) + def test_quota_refusal_is_returned_not_retried( + self, pg_conn, patch_worker_db, task_self, monkeypatch + ): + """A spent quota cannot succeed on retry, so the task must not raise.""" + from datetime import datetime, timezone + + from docsgpt import worker + from docsgpt.agents import headless_runner + from docsgpt.quotas.service import QuotaExceeded, QuotaExceededError + + agent = AgentsRepository(pg_conn).create( + user_id="alice", name="hook-agent", status="active", + agent_type="classic", retriever="classic", chunks=2, key="sk-test-q", + ) + exceeded = QuotaExceeded( + user_id="alice", bucket="all", budget="cost", usage=5.0, limit=5.0, + source="user", source_id=None, + resets_at=datetime(2099, 1, 1, tzinfo=timezone.utc), + ) + + def _refuse(*a, **k): + raise QuotaExceededError(exceeded) + + monkeypatch.setattr(headless_runner, "run_agent_headless", _refuse) + + result = worker.agent_webhook_worker(task_self, str(agent["id"]), {"event": "ping"}) + + assert result["status"] == "quota_exceeded" + assert "$5.00 of $5.00" in result["error"] + def test_webhook_journals_headless_denial_for_approval_gated_tool( self, pg_conn, patch_worker_db, task_self, monkeypatch ):