mirror of
https://github.com/tiennm99/DocsGPT.git
synced 2026-10-11 03:12:55 +00:00
Merge pull request #2816 from arc53/feat/admin-quotas
feat: admin-set usage quotas per user and per team
This commit is contained in:
83 files changed
+4614
-30
No files matched your search
@@ -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]
|
||||
@@ -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. |
|
||||
|
||||
<Callout type="info" emoji="ℹ️">
|
||||
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.
|
||||
@@ -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.
|
||||
|
||||
@@ -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.
|
||||
|
||||
<Callout type="info" emoji="ℹ️">
|
||||
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.
|
||||
</Callout>
|
||||
|
||||
## 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.
|
||||
|
||||
<Callout type="warning" emoji="⚠️">
|
||||
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.
|
||||
</Callout>
|
||||
|
||||
## 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/<team_id>` | A team's per-member allowance. |
|
||||
| `GET` `PUT` `DELETE` | `/api/admin/quotas/users/<user_id>` | 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`.
|
||||
@@ -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"
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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;")
|
||||
@@ -1,3 +1,4 @@
|
||||
from .routes import admin_ns
|
||||
from . import quotas # noqa: F401 (registers the quota resources on admin_ns)
|
||||
|
||||
__all__ = ["admin_ns"]
|
||||
@@ -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/<string:team_id>")
|
||||
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/<string:user_id>")
|
||||
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)
|
||||
@@ -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(
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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"),
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
)
|
||||
@@ -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"
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"),
|
||||
)
|
||||
|
||||
@@ -108,8 +108,10 @@ defaults: # optional, applied to every model below
|
||||
supports_streaming: bool # default true
|
||||
attachments: [<alias-or-mime>, ...] # 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: <string> # default null; none|minimal|low|medium|high|xhigh (subset is model-dependent)
|
||||
api_flavor: <string> # chat_completions (default) or responses
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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),
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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
|
||||
@@ -0,0 +1,51 @@
|
||||
"""Extension point for limits that do not come from ``quota_policies``.
|
||||
|
||||
A deployment can register a provider that supplies per-user default policies
|
||||
(for example from a subscription plan) and adjusts the quota error payload.
|
||||
Defaults sit below every stored layer: a stored instance, team or user row
|
||||
with an opinion always wins.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Mapping
|
||||
|
||||
|
||||
class QuotaDefaultsProvider:
|
||||
"""Base provider: no defaults, error payload unchanged."""
|
||||
|
||||
def default_policies(self, user_id: str) -> list[dict]:
|
||||
"""Return default policy rows for ``user_id``.
|
||||
|
||||
Each row uses the ``quota_policies`` field names (``bucket``,
|
||||
``token_limit``, ``token_unlimited``, ``cost_limit_usd``,
|
||||
``cost_unlimited``); ``scope`` is set by the caller.
|
||||
"""
|
||||
return []
|
||||
|
||||
def error_payload(self, payload: dict, user_id: str) -> dict:
|
||||
"""Return the payload sent to a client whose quota is exhausted."""
|
||||
return payload
|
||||
|
||||
|
||||
_provider: QuotaDefaultsProvider = QuotaDefaultsProvider()
|
||||
|
||||
|
||||
def register_defaults_provider(provider: QuotaDefaultsProvider) -> None:
|
||||
"""Replace the process-wide defaults provider."""
|
||||
global _provider
|
||||
_provider = provider
|
||||
|
||||
|
||||
def get_defaults_provider() -> QuotaDefaultsProvider:
|
||||
"""Return the registered defaults provider."""
|
||||
return _provider
|
||||
|
||||
|
||||
def default_rows(user_id: str) -> list[dict]:
|
||||
"""Return the provider's defaults for ``user_id`` as ``default``-layer rows."""
|
||||
rows: list[dict] = []
|
||||
for row in get_defaults_provider().default_policies(user_id) or []:
|
||||
if isinstance(row, Mapping):
|
||||
rows.append({"bucket": "all", "enabled": True, **row, "scope": "default", "subject_id": None})
|
||||
return rows
|
||||
@@ -0,0 +1,99 @@
|
||||
"""Resolve the policy rows that apply to a user into effective limits.
|
||||
|
||||
Each budget (tokens, cost) resolves on its own: the user's row wins, then the
|
||||
most generous of the user's team rows, then the instance row, then the
|
||||
registered defaults. A row with neither a limit nor the unlimited flag for a
|
||||
budget has no opinion on it and is skipped.
|
||||
|
||||
Teams resolve to the most generous allowance because team membership is not
|
||||
controlled by the instance admin: under "most restrictive", any team admin
|
||||
could throttle a user by adding them to a low-allowance team.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Iterable, Mapping, Optional
|
||||
|
||||
LAYERS = ("user", "team", "instance", "default")
|
||||
|
||||
_FIELDS = {"tokens": ("token_limit", "token_unlimited"), "cost": ("cost_limit_usd", "cost_unlimited")}
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ResolvedLimit:
|
||||
"""One budget's effective limit and the layer it came from.
|
||||
|
||||
``limit`` is ``None`` when the budget is unlimited. ``source`` is ``None``
|
||||
when no layer had an opinion, otherwise one of ``LAYERS``; ``source_id`` is
|
||||
the team id for a team-sourced limit.
|
||||
"""
|
||||
|
||||
limit: Optional[float] = None
|
||||
source: Optional[str] = None
|
||||
source_id: Optional[str] = None
|
||||
|
||||
@property
|
||||
def unlimited(self) -> bool:
|
||||
return self.limit is None
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ResolvedLimits:
|
||||
"""A user's effective token and cost limits for one bucket."""
|
||||
|
||||
tokens: ResolvedLimit
|
||||
cost: ResolvedLimit
|
||||
|
||||
@property
|
||||
def unlimited(self) -> bool:
|
||||
return self.tokens.unlimited and self.cost.unlimited
|
||||
|
||||
|
||||
def _opinion(row: Mapping, budget: str) -> Optional[tuple[bool, Optional[float]]]:
|
||||
"""Return ``(unlimited, limit)`` for a row's budget, or ``None`` if it defers."""
|
||||
limit_field, unlimited_field = _FIELDS[budget]
|
||||
if row.get(unlimited_field):
|
||||
return True, None
|
||||
value = row.get(limit_field)
|
||||
if value is None:
|
||||
return None
|
||||
return False, float(value)
|
||||
|
||||
|
||||
def _resolve_budget(rows: Iterable[Mapping], budget: str) -> ResolvedLimit:
|
||||
by_layer: dict[str, list[tuple[Mapping, tuple[bool, Optional[float]]]]] = {}
|
||||
for row in rows:
|
||||
opinion = _opinion(row, budget)
|
||||
if opinion is not None:
|
||||
by_layer.setdefault(row["scope"], []).append((row, opinion))
|
||||
for layer in LAYERS:
|
||||
candidates = by_layer.get(layer)
|
||||
if not candidates:
|
||||
continue
|
||||
# Most generous first: unlimited, then the larger limit. Only the team
|
||||
# layer can hold more than one candidate. Ties break on subject id so
|
||||
# the reported source is stable.
|
||||
row, (unlimited, limit) = min(
|
||||
candidates,
|
||||
key=lambda c: (not c[1][0], -(c[1][1] or 0.0), str(c[0].get("subject_id") or "")),
|
||||
)
|
||||
source_id = str(row["subject_id"]) if layer == "team" else None
|
||||
return ResolvedLimit(limit=None if unlimited else limit, source=layer, source_id=source_id)
|
||||
return ResolvedLimit()
|
||||
|
||||
|
||||
def resolve_limits(rows: Iterable[Mapping], bucket: str = "all") -> ResolvedLimits:
|
||||
"""Return the effective limits for ``bucket`` from a user's applicable rows.
|
||||
|
||||
Args:
|
||||
rows: Policy rows that apply to the user (their own, their teams', the
|
||||
instance's and any provider defaults). Disabled rows and rows for
|
||||
other buckets are ignored.
|
||||
bucket: The policy bucket to resolve.
|
||||
"""
|
||||
applicable = [r for r in rows if r.get("bucket", "all") == bucket and r.get("enabled", True)]
|
||||
return ResolvedLimits(
|
||||
tokens=_resolve_budget(applicable, "tokens"),
|
||||
cost=_resolve_budget(applicable, "cost"),
|
||||
)
|
||||
@@ -0,0 +1,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
|
||||
@@ -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}")
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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
|
||||
@@ -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``.
|
||||
|
||||
+24
-1
@@ -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.
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -119,6 +119,8 @@ const EVENT_LABELS: Record<string, string> = {
|
||||
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 {
|
||||
|
||||
@@ -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 (
|
||||
<div>
|
||||
<div className="flex items-baseline justify-between gap-3 text-sm">
|
||||
<span className="text-muted-foreground">{label}</span>
|
||||
<span className="tabular-nums">
|
||||
{fmt(budget.used)}
|
||||
{budget.limit === null ? ' · no limit' : ` of ${fmt(budget.limit)}`}
|
||||
</span>
|
||||
</div>
|
||||
{budget.limit !== null ? (
|
||||
<div
|
||||
className="bg-muted mt-1 h-1.5 w-full overflow-hidden rounded-full"
|
||||
role="progressbar"
|
||||
aria-label={label}
|
||||
aria-valuemin={0}
|
||||
aria-valuemax={100}
|
||||
aria-valuenow={Math.round(percent)}
|
||||
>
|
||||
<div
|
||||
className={`h-full rounded-full ${tone}`}
|
||||
style={{ width: `${percent}%` }}
|
||||
/>
|
||||
</div>
|
||||
) : null}
|
||||
{caption ? (
|
||||
<p className="text-muted-foreground mt-1 text-xs">{caption}</p>
|
||||
) : null}
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
||||
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 (
|
||||
<div>
|
||||
<p className="text-foreground text-sm font-medium">{label}</p>
|
||||
<div className="mt-1 flex items-center gap-2">
|
||||
<Select value={mode} onValueChange={(v) => onMode(v as BudgetMode)}>
|
||||
<SelectTrigger className="w-40" aria-label={`${label} mode`}>
|
||||
<SelectValue />
|
||||
</SelectTrigger>
|
||||
<SelectContent>
|
||||
{MODES.map((m) => (
|
||||
<SelectItem key={m.value} value={m.value}>
|
||||
{m.label}
|
||||
</SelectItem>
|
||||
))}
|
||||
</SelectContent>
|
||||
</Select>
|
||||
{mode === 'limit' ? (
|
||||
<Input
|
||||
type="number"
|
||||
min="0"
|
||||
step={step}
|
||||
inputMode="decimal"
|
||||
value={value}
|
||||
aria-label={label}
|
||||
placeholder={hint}
|
||||
onChange={(e) => onValue(e.target.value)}
|
||||
className="flex-1"
|
||||
/>
|
||||
) : null}
|
||||
</div>
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
||||
/**
|
||||
* 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<QuotaForm>(() => policyToForm(policy));
|
||||
const [busy, setBusy] = useState(false);
|
||||
const [error, setError] = useState<string | null>(null);
|
||||
|
||||
useEffect(() => {
|
||||
setForm(policyToForm(policy));
|
||||
setError(null);
|
||||
}, [policy, scope, subjectId]);
|
||||
|
||||
const patch = (fields: Partial<QuotaForm>) =>
|
||||
setForm((prev) => ({ ...prev, ...fields }));
|
||||
|
||||
const submit = async (request: () => Promise<Response>) => {
|
||||
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 (
|
||||
<div className="space-y-4">
|
||||
<p className="text-muted-foreground text-xs">{inheritHint}</p>
|
||||
{policy && !policy.enabled ? (
|
||||
<p className="text-xs text-amber-600 dark:text-amber-400">
|
||||
This policy is disabled and is not enforced. Saving keeps it disabled.
|
||||
</p>
|
||||
) : null}
|
||||
<BudgetField
|
||||
label="Tokens"
|
||||
hint="e.g. 2000000"
|
||||
step="1"
|
||||
mode={form.tokenMode}
|
||||
value={form.tokenLimit}
|
||||
onMode={(tokenMode) => patch({ tokenMode })}
|
||||
onValue={(tokenLimit) => patch({ tokenLimit })}
|
||||
/>
|
||||
<BudgetField
|
||||
label="Cost (USD)"
|
||||
hint="e.g. 25"
|
||||
step="0.01"
|
||||
mode={form.costMode}
|
||||
value={form.costLimit}
|
||||
onMode={(costMode) => patch({ costMode })}
|
||||
onValue={(costLimit) => patch({ costLimit })}
|
||||
/>
|
||||
<div>
|
||||
<p className="text-foreground text-sm font-medium">Note</p>
|
||||
<Input
|
||||
value={form.note}
|
||||
maxLength={500}
|
||||
aria-label="Note"
|
||||
placeholder="Optional, visible to admins only"
|
||||
onChange={(e) => patch({ note: e.target.value })}
|
||||
className="mt-1"
|
||||
/>
|
||||
</div>
|
||||
{error ? (
|
||||
<p role="alert" className="text-sm text-red-600 dark:text-red-400">
|
||||
{error}
|
||||
</p>
|
||||
) : null}
|
||||
<div className="flex justify-end gap-2">
|
||||
{policy ? (
|
||||
<Button
|
||||
type="button"
|
||||
variant="destructive-outline"
|
||||
size="sm"
|
||||
disabled={busy}
|
||||
onClick={remove}
|
||||
>
|
||||
Remove
|
||||
</Button>
|
||||
) : null}
|
||||
<Button
|
||||
type="button"
|
||||
size="sm"
|
||||
disabled={busy || (isEmptyForm(form) && !policy)}
|
||||
onClick={save}
|
||||
>
|
||||
Save
|
||||
</Button>
|
||||
</div>
|
||||
</div>
|
||||
);
|
||||
}
|
||||
@@ -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<QuotaScope, string> = {
|
||||
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 (
|
||||
<>
|
||||
<TableCell className="tabular-nums">
|
||||
{describeBudget(policy.token_limit, policy.token_unlimited, 'tokens')}
|
||||
</TableCell>
|
||||
<TableCell className="tabular-nums">
|
||||
{describeBudget(policy.cost_limit_usd, policy.cost_unlimited, 'cost')}
|
||||
</TableCell>
|
||||
<TableCell className="text-muted-foreground max-w-56 truncate">
|
||||
{policy.note || '—'}
|
||||
</TableCell>
|
||||
<TableCell className="text-muted-foreground whitespace-nowrap">
|
||||
{fmtRelative(policy.updated_at)}
|
||||
</TableCell>
|
||||
</>
|
||||
);
|
||||
}
|
||||
|
||||
export default function Quotas() {
|
||||
const token = useSelector(selectToken);
|
||||
const [data, setData] = useState<any | null>(null);
|
||||
const [teams, setTeams] = useState<any[]>([]);
|
||||
const [loading, setLoading] = useState(true);
|
||||
const [editing, setEditing] = useState<Editing | null>(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 <Loading />;
|
||||
if (!data?.success) return <LoadError message="Failed to load quotas." />;
|
||||
|
||||
const bucketPill = (policy: QuotaPolicy) => (
|
||||
<>
|
||||
{isAll(policy) ? null : <Pill tone="muted">{policy.bucket} traffic</Pill>}
|
||||
{policy.enabled ? null : <Pill tone="muted">Disabled</Pill>}
|
||||
</>
|
||||
);
|
||||
|
||||
return (
|
||||
<div className="mt-6 space-y-8">
|
||||
<p className="text-muted-foreground text-sm">
|
||||
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.
|
||||
</p>
|
||||
|
||||
{(data.unpriced_models ?? []).length > 0 ? (
|
||||
<div className="rounded-2xl border border-amber-300 bg-amber-50 px-5 py-4 text-sm dark:border-amber-800 dark:bg-amber-950/30">
|
||||
<p className="text-foreground font-medium">
|
||||
Models without a price are invisible to cost limits
|
||||
</p>
|
||||
<p className="text-muted-foreground mt-1">
|
||||
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.
|
||||
</p>
|
||||
</div>
|
||||
) : null}
|
||||
|
||||
<section>
|
||||
<div className="flex items-center justify-between gap-3">
|
||||
<p className="text-foreground font-bold">Instance default</p>
|
||||
<Button
|
||||
variant="outline"
|
||||
size="sm"
|
||||
onClick={() =>
|
||||
setEditing({
|
||||
scope: 'instance',
|
||||
subjectId: null,
|
||||
title: 'Instance default',
|
||||
policy: instancePolicy,
|
||||
})
|
||||
}
|
||||
>
|
||||
{instancePolicy ? 'Edit' : 'Set default'}
|
||||
</Button>
|
||||
</div>
|
||||
<p className="text-muted-foreground mt-1 text-sm">
|
||||
{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.'}
|
||||
</p>
|
||||
</section>
|
||||
|
||||
<section>
|
||||
<div className="flex flex-wrap items-center justify-between gap-3">
|
||||
<p className="text-foreground font-bold">Team allowances</p>
|
||||
{teamsWithoutPolicy.length > 0 ? (
|
||||
<div className="flex items-center gap-2">
|
||||
<Select value={teamPick} onValueChange={setTeamPick}>
|
||||
<SelectTrigger className="w-52" aria-label="Team">
|
||||
<SelectValue placeholder="Choose a team" />
|
||||
</SelectTrigger>
|
||||
<SelectContent>
|
||||
{teamsWithoutPolicy.map((team) => (
|
||||
<SelectItem key={team.id} value={String(team.id)}>
|
||||
{team.name}
|
||||
</SelectItem>
|
||||
))}
|
||||
</SelectContent>
|
||||
</Select>
|
||||
<Button
|
||||
variant="outline"
|
||||
size="sm"
|
||||
disabled={!teamPick}
|
||||
onClick={() => {
|
||||
const team = teams.find((tm) => String(tm.id) === teamPick);
|
||||
setEditing({
|
||||
scope: 'team',
|
||||
subjectId: teamPick,
|
||||
title: team?.name ?? 'Team',
|
||||
policy: null,
|
||||
});
|
||||
}}
|
||||
>
|
||||
Add allowance
|
||||
</Button>
|
||||
</div>
|
||||
) : null}
|
||||
</div>
|
||||
{teamPolicies.length === 0 ? (
|
||||
<p className="text-muted-foreground mt-1 text-sm">
|
||||
No team has an allowance.
|
||||
</p>
|
||||
) : (
|
||||
<TableContainer className="mt-3">
|
||||
<Table>
|
||||
<TableHead>
|
||||
<TableRow>
|
||||
<TableHeader>Team</TableHeader>
|
||||
<TableHeader>Members</TableHeader>
|
||||
<TableHeader>Tokens</TableHeader>
|
||||
<TableHeader>Cost</TableHeader>
|
||||
<TableHeader>Note</TableHeader>
|
||||
<TableHeader>Updated</TableHeader>
|
||||
<TableHeader className="text-right">Actions</TableHeader>
|
||||
</TableRow>
|
||||
</TableHead>
|
||||
<TableBody>
|
||||
{teamPolicies.map((policy) => (
|
||||
<TableRow key={`${policy.subject_id}:${policy.bucket}`}>
|
||||
<TableCell>
|
||||
<span className="mr-2">
|
||||
{policy.team_name ?? policy.subject_id}
|
||||
</span>
|
||||
{bucketPill(policy)}
|
||||
</TableCell>
|
||||
<TableCell className="tabular-nums">
|
||||
{fmtNumber(policy.member_count)}
|
||||
</TableCell>
|
||||
<PolicyCells policy={policy} />
|
||||
<TableCell className="text-right">
|
||||
{isAll(policy) ? (
|
||||
<Button
|
||||
variant="outline"
|
||||
size="sm"
|
||||
onClick={() =>
|
||||
setEditing({
|
||||
scope: 'team',
|
||||
subjectId: policy.subject_id,
|
||||
title: policy.team_name ?? 'Team',
|
||||
policy,
|
||||
})
|
||||
}
|
||||
>
|
||||
Edit
|
||||
</Button>
|
||||
) : null}
|
||||
</TableCell>
|
||||
</TableRow>
|
||||
))}
|
||||
</TableBody>
|
||||
</Table>
|
||||
</TableContainer>
|
||||
)}
|
||||
</section>
|
||||
|
||||
<section>
|
||||
<p className="text-foreground font-bold">User overrides</p>
|
||||
{userPolicies.length === 0 ? (
|
||||
<p className="text-muted-foreground mt-1 text-sm">
|
||||
No user has an override. Add one from a user's menu on the
|
||||
Users tab.
|
||||
</p>
|
||||
) : (
|
||||
<TableContainer className="mt-3">
|
||||
<Table>
|
||||
<TableHead>
|
||||
<TableRow>
|
||||
<TableHeader>User</TableHeader>
|
||||
<TableHeader>Tokens</TableHeader>
|
||||
<TableHeader>Cost</TableHeader>
|
||||
<TableHeader>Note</TableHeader>
|
||||
<TableHeader>Updated</TableHeader>
|
||||
<TableHeader className="text-right">Actions</TableHeader>
|
||||
</TableRow>
|
||||
</TableHead>
|
||||
<TableBody>
|
||||
{userPolicies.map((policy) => (
|
||||
<TableRow key={`${policy.subject_id}:${policy.bucket}`}>
|
||||
<TableCell>
|
||||
<span className="mr-2 break-all">
|
||||
{policy.subject_id}
|
||||
</span>
|
||||
{bucketPill(policy)}
|
||||
</TableCell>
|
||||
<PolicyCells policy={policy} />
|
||||
<TableCell className="text-right">
|
||||
{isAll(policy) ? (
|
||||
<Button
|
||||
variant="outline"
|
||||
size="sm"
|
||||
onClick={() =>
|
||||
setEditing({
|
||||
scope: 'user',
|
||||
subjectId: policy.subject_id,
|
||||
title: policy.subject_id ?? 'User',
|
||||
policy,
|
||||
})
|
||||
}
|
||||
>
|
||||
Edit
|
||||
</Button>
|
||||
) : null}
|
||||
</TableCell>
|
||||
</TableRow>
|
||||
))}
|
||||
</TableBody>
|
||||
</Table>
|
||||
</TableContainer>
|
||||
)}
|
||||
</section>
|
||||
|
||||
<Modal
|
||||
open={editing !== null}
|
||||
onOpenChange={(open) => {
|
||||
if (!open) setEditing(null);
|
||||
}}
|
||||
title={editing ? `Quota · ${editing.title}` : 'Quota'}
|
||||
>
|
||||
{editing ? (
|
||||
<QuotaEditor
|
||||
scope={editing.scope}
|
||||
subjectId={editing.subjectId}
|
||||
policy={editing.policy}
|
||||
inheritHint={HINTS[editing.scope]}
|
||||
onSaved={() => {
|
||||
setEditing(null);
|
||||
setTeamPick('');
|
||||
load();
|
||||
}}
|
||||
/>
|
||||
) : null}
|
||||
</Modal>
|
||||
</div>
|
||||
);
|
||||
}
|
||||
@@ -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<any | null>(null);
|
||||
const [teamNames, setTeamNames] = useState<Record<string, string>>({});
|
||||
|
||||
// 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 (
|
||||
<Modal
|
||||
open={userId !== null}
|
||||
onOpenChange={(open) => {
|
||||
if (!open) onClose();
|
||||
}}
|
||||
title={userId ? `Quota · ${userId}` : 'Quota'}
|
||||
>
|
||||
{data === null ? (
|
||||
<Loading />
|
||||
) : !data.success ? (
|
||||
<LoadError message="Failed to load this user's quota." />
|
||||
) : (
|
||||
<div className="space-y-5">
|
||||
{overall ? (
|
||||
<div className="space-y-3">
|
||||
<UsageBar
|
||||
label={`Tokens this ${data.period}`}
|
||||
kind="tokens"
|
||||
budget={overall.tokens}
|
||||
caption={caption(overall.tokens)}
|
||||
/>
|
||||
<UsageBar
|
||||
label={`Cost this ${data.period}`}
|
||||
kind="cost"
|
||||
budget={overall.cost}
|
||||
caption={caption(overall.cost)}
|
||||
/>
|
||||
<p className="text-muted-foreground text-xs">
|
||||
Resets {fmtDate(overall.resets_at)}
|
||||
</p>
|
||||
</div>
|
||||
) : null}
|
||||
<div className="border-border border-t pt-4">
|
||||
<p className="text-foreground mb-2 text-sm font-bold">
|
||||
User override
|
||||
</p>
|
||||
<QuotaEditor
|
||||
scope="user"
|
||||
subjectId={userId}
|
||||
policy={override}
|
||||
inheritHint="Overrides team allowances and the instance default for this user. Leave a budget as “Not set here” to keep what the user inherits."
|
||||
onSaved={load}
|
||||
/>
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
</Modal>
|
||||
);
|
||||
}
|
||||
@@ -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<string | null>(null);
|
||||
const [menuUserId, setMenuUserId] = useState<string | null>(null);
|
||||
const [detail, setDetail] = useState<any | null>(null);
|
||||
const [quotaUserId, setQuotaUserId] = useState<string | null>(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}
|
||||
|
||||
<UserQuotaModal
|
||||
userId={quotaUserId}
|
||||
onClose={() => setQuotaUserId(null)}
|
||||
/>
|
||||
|
||||
<Modal
|
||||
open={detail !== null}
|
||||
onOpenChange={(open) => {
|
||||
|
||||
@@ -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() {
|
||||
<Route path="users" element={<Users />} />
|
||||
<Route path="roles" element={<Admins />} />
|
||||
<Route path="usage" element={<Usage />} />
|
||||
<Route path="quotas" element={<Quotas />} />
|
||||
<Route path="audit" element={<Audit />} />
|
||||
<Route path="*" element={<Navigate to="/admin" replace />} />
|
||||
</Routes>
|
||||
|
||||
@@ -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>): 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',
|
||||
);
|
||||
});
|
||||
});
|
||||
@@ -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<string, unknown> } | { 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<string, unknown> = {
|
||||
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';
|
||||
}
|
||||
@@ -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',
|
||||
|
||||
@@ -10,6 +10,14 @@ const qs = (params: Record<string, string | number | undefined>): 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<any> =>
|
||||
apiClient.get(endpoints.ADMIN.OVERVIEW, token),
|
||||
@@ -54,6 +62,23 @@ const adminService = {
|
||||
token: string | null,
|
||||
): Promise<any> =>
|
||||
apiClient.get(`${endpoints.ADMIN.DEVICE_AUDIT}${qs(params)}`, token),
|
||||
getQuotas: (token: string | null): Promise<any> =>
|
||||
apiClient.get(endpoints.ADMIN.QUOTAS, token),
|
||||
getUserQuota: (userId: string, token: string | null): Promise<any> =>
|
||||
apiClient.get(endpoints.ADMIN.QUOTA_USER(userId), token),
|
||||
setQuota: (
|
||||
scope: QuotaScope,
|
||||
subjectId: string | null,
|
||||
policy: Record<string, unknown>,
|
||||
token: string | null,
|
||||
): Promise<any> => apiClient.put(quotaUrl(scope, subjectId), policy, token),
|
||||
deleteQuota: (
|
||||
scope: QuotaScope,
|
||||
subjectId: string | null,
|
||||
bucket: string,
|
||||
token: string | null,
|
||||
): Promise<any> =>
|
||||
apiClient.delete(`${quotaUrl(scope, subjectId)}${qs({ bucket })}`, token),
|
||||
};
|
||||
|
||||
export default adminService;
|
||||
@@ -7,6 +7,8 @@ const userService = {
|
||||
throttledApiClient.get(endpoints.USER.CONFIG, null),
|
||||
getMe: (token: string | null): Promise<any> =>
|
||||
apiClient.get(endpoints.USER.ME, token),
|
||||
getQuota: (token: string | null): Promise<any> =>
|
||||
apiClient.get(endpoints.USER.QUOTA, token),
|
||||
getNewToken: (): Promise<any> =>
|
||||
throttledApiClient.get(endpoints.USER.NEW_TOKEN, null),
|
||||
// Token deliberately null: a stale Authorization header must not be able
|
||||
|
||||
@@ -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) ||
|
||||
|
||||
@@ -0,0 +1,52 @@
|
||||
import { describe, expect, it } from 'vitest';
|
||||
|
||||
import { isQuotaError, quotaErrorMessage } from './quotaError';
|
||||
|
||||
const t = ((key: string, values: Record<string, string>) =>
|
||||
`${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);
|
||||
});
|
||||
});
|
||||
@@ -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 },
|
||||
);
|
||||
}
|
||||
@@ -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": {
|
||||
|
||||
@@ -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": {
|
||||
|
||||
@@ -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": {
|
||||
|
||||
@@ -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": {
|
||||
|
||||
@@ -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": {
|
||||
|
||||
@@ -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": {
|
||||
|
||||
@@ -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": {
|
||||
|
||||
@@ -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 (
|
||||
<div className="mt-8">
|
||||
{agentId ? null : <UsageQuota />}
|
||||
<div className="mb-5 flex flex-row flex-wrap items-center justify-between gap-3">
|
||||
<p className="text-muted-foreground text-sm leading-6">
|
||||
{t('settings.analytics.subtitle')}
|
||||
|
||||
@@ -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 (
|
||||
<div className="min-w-48 flex-1">
|
||||
<div className="flex items-baseline justify-between gap-3 text-sm">
|
||||
<span className="text-muted-foreground">{label}</span>
|
||||
<span className="text-foreground tabular-nums">
|
||||
{t('settings.analytics.quota.usedOf', {
|
||||
used: format(budget.used),
|
||||
limit: format(budget.limit),
|
||||
})}
|
||||
</span>
|
||||
</div>
|
||||
<div
|
||||
className="bg-muted mt-1 h-1.5 w-full overflow-hidden rounded-full"
|
||||
role="progressbar"
|
||||
aria-label={label}
|
||||
aria-valuemin={0}
|
||||
aria-valuemax={100}
|
||||
aria-valuenow={Math.round(percent)}
|
||||
>
|
||||
<div
|
||||
className={`h-full rounded-full ${tone}`}
|
||||
style={{ width: `${percent}%` }}
|
||||
/>
|
||||
</div>
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
||||
/** 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<Bucket[]>([]);
|
||||
|
||||
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 (
|
||||
<div className="border-border mb-6 rounded-2xl border px-6 py-5">
|
||||
<div className="flex flex-wrap items-baseline justify-between gap-2">
|
||||
<p className="text-foreground font-bold">
|
||||
{t('settings.analytics.quota.title')}
|
||||
</p>
|
||||
{resetsAt ? (
|
||||
<p className="text-muted-foreground text-xs">
|
||||
{t('settings.analytics.quota.resets', { resetsAt })}
|
||||
</p>
|
||||
) : null}
|
||||
</div>
|
||||
{buckets.map((bucket) => (
|
||||
<div key={bucket.bucket} className="mt-3">
|
||||
{scopeLabel(bucket.bucket) ? (
|
||||
<p className="text-muted-foreground mb-1 text-xs">
|
||||
{scopeLabel(bucket.bucket)}
|
||||
</p>
|
||||
) : null}
|
||||
<div className="flex flex-wrap gap-6">
|
||||
<Meter
|
||||
label={t('settings.analytics.quota.tokens')}
|
||||
budget={bucket.tokens}
|
||||
format={(value) => number.format(value)}
|
||||
/>
|
||||
<Meter
|
||||
label={t('settings.analytics.quota.cost')}
|
||||
budget={bucket.cost}
|
||||
format={(value) => usd.format(value)}
|
||||
/>
|
||||
</div>
|
||||
</div>
|
||||
))}
|
||||
</div>
|
||||
);
|
||||
}
|
||||
@@ -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;
|
||||
|
||||
@@ -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
|
||||
@@ -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,
|
||||
):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
|
||||
Whitespace-only changes.
@@ -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
|
||||
@@ -0,0 +1,47 @@
|
||||
"""Tests for docsgpt/quotas/providers.py."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from docsgpt.quotas import providers
|
||||
from docsgpt.quotas.providers import QuotaDefaultsProvider, default_rows, register_defaults_provider
|
||||
from docsgpt.quotas.resolver import resolve_limits
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _restore_provider():
|
||||
original = providers.get_defaults_provider()
|
||||
yield
|
||||
register_defaults_provider(original)
|
||||
|
||||
|
||||
class _PlanProvider(QuotaDefaultsProvider):
|
||||
def default_policies(self, user_id):
|
||||
return [
|
||||
{"bucket": "agent", "cost_limit_usd": 5.0, "scope": "user"},
|
||||
{"cost_limit_usd": 10.0},
|
||||
"ignored",
|
||||
]
|
||||
|
||||
def error_payload(self, payload, user_id):
|
||||
return {**payload, "error_code": "free-limit-reached"}
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestDefaultsProvider:
|
||||
def test_base_provider_has_no_defaults(self):
|
||||
assert default_rows("u1") == []
|
||||
assert QuotaDefaultsProvider().error_payload({"a": 1}, "u1") == {"a": 1}
|
||||
|
||||
def test_rows_are_forced_into_the_default_layer(self):
|
||||
register_defaults_provider(_PlanProvider())
|
||||
rows = default_rows("u1")
|
||||
assert [r["scope"] for r in rows] == ["default", "default"]
|
||||
assert [r["bucket"] for r in rows] == ["agent", "all"]
|
||||
assert resolve_limits(rows, "agent").cost.source == "default"
|
||||
|
||||
def test_stored_rows_win_over_defaults(self):
|
||||
register_defaults_provider(_PlanProvider())
|
||||
stored = {"scope": "instance", "subject_id": None, "bucket": "all", "enabled": True, "cost_limit_usd": 99}
|
||||
assert resolve_limits(default_rows("u1") + [stored]).cost.limit == 99.0
|
||||
@@ -0,0 +1,124 @@
|
||||
"""Tests for docsgpt/quotas/resolver.py."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from docsgpt.quotas.resolver import ResolvedLimit, resolve_limits
|
||||
|
||||
|
||||
def _row(scope, subject_id=None, **fields):
|
||||
return {"scope": scope, "subject_id": subject_id, "bucket": "all", "enabled": True, **fields}
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestLayers:
|
||||
def test_no_rows_is_unlimited(self):
|
||||
limits = resolve_limits([])
|
||||
assert limits.unlimited
|
||||
assert limits.tokens == ResolvedLimit()
|
||||
|
||||
def test_user_beats_team_beats_instance_beats_default(self):
|
||||
rows = [
|
||||
_row("default", token_limit=1),
|
||||
_row("instance", token_limit=10),
|
||||
_row("team", "t1", token_limit=100),
|
||||
_row("user", "u1", token_limit=5),
|
||||
]
|
||||
assert resolve_limits(rows).tokens == ResolvedLimit(5.0, "user")
|
||||
assert resolve_limits(rows[:3]).tokens == ResolvedLimit(100.0, "team", "t1")
|
||||
assert resolve_limits(rows[:2]).tokens == ResolvedLimit(10.0, "instance")
|
||||
assert resolve_limits(rows[:1]).tokens == ResolvedLimit(1.0, "default")
|
||||
|
||||
def test_user_override_can_be_stricter_than_the_team(self):
|
||||
rows = [_row("team", "t1", token_unlimited=True), _row("user", "u1", token_limit=0)]
|
||||
assert resolve_limits(rows).tokens == ResolvedLimit(0.0, "user")
|
||||
|
||||
def test_user_unlimited_lifts_an_instance_limit(self):
|
||||
rows = [_row("instance", token_limit=10), _row("user", "u1", token_unlimited=True)]
|
||||
resolved = resolve_limits(rows).tokens
|
||||
assert resolved.unlimited and resolved.source == "user"
|
||||
|
||||
def test_budgets_resolve_independently(self):
|
||||
rows = [
|
||||
_row("instance", token_limit=10, cost_limit_usd=1),
|
||||
_row("user", "u1", cost_limit_usd=25),
|
||||
]
|
||||
limits = resolve_limits(rows)
|
||||
assert limits.tokens == ResolvedLimit(10.0, "instance")
|
||||
assert limits.cost == ResolvedLimit(25.0, "user")
|
||||
|
||||
def test_a_row_with_no_opinion_defers(self):
|
||||
rows = [_row("instance", token_limit=10), _row("user", "u1", note="vip")]
|
||||
assert resolve_limits(rows).tokens == ResolvedLimit(10.0, "instance")
|
||||
|
||||
def test_zero_is_a_limit_not_unlimited(self):
|
||||
resolved = resolve_limits([_row("instance", cost_limit_usd=0)]).cost
|
||||
assert resolved.limit == 0.0 and not resolved.unlimited
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestMultipleTeams:
|
||||
def test_most_generous_team_wins(self):
|
||||
rows = [
|
||||
_row("team", "small", token_limit=100),
|
||||
_row("team", "big", token_limit=900),
|
||||
_row("team", "mid", token_limit=500),
|
||||
]
|
||||
assert resolve_limits(rows).tokens == ResolvedLimit(900.0, "team", "big")
|
||||
|
||||
def test_an_unlimited_team_beats_any_limit(self):
|
||||
rows = [_row("team", "big", token_limit=10**12), _row("team", "free", token_unlimited=True)]
|
||||
assert resolve_limits(rows).tokens == ResolvedLimit(None, "team", "free")
|
||||
|
||||
def test_allowances_are_not_added_together(self):
|
||||
rows = [_row("team", "a", token_limit=100), _row("team", "b", token_limit=100)]
|
||||
assert resolve_limits(rows).tokens.limit == 100.0
|
||||
|
||||
def test_equal_teams_report_a_stable_source(self):
|
||||
rows = [_row("team", "b", token_limit=100), _row("team", "a", token_limit=100)]
|
||||
assert resolve_limits(rows).tokens.source_id == "a"
|
||||
assert resolve_limits(list(reversed(rows))).tokens.source_id == "a"
|
||||
|
||||
def test_each_budget_can_come_from_a_different_team(self):
|
||||
rows = [
|
||||
_row("team", "tok", token_limit=900, cost_limit_usd=1),
|
||||
_row("team", "usd", token_limit=100, cost_limit_usd=50),
|
||||
]
|
||||
limits = resolve_limits(rows)
|
||||
assert limits.tokens.source_id == "tok"
|
||||
assert limits.cost.source_id == "usd"
|
||||
|
||||
def test_a_team_without_an_opinion_does_not_lift_the_limit(self):
|
||||
rows = [_row("team", "quiet"), _row("team", "capped", token_limit=100)]
|
||||
assert resolve_limits(rows).tokens == ResolvedLimit(100.0, "team", "capped")
|
||||
|
||||
def test_teams_without_opinions_fall_through_to_instance(self):
|
||||
rows = [_row("team", "quiet", cost_limit_usd=5), _row("instance", token_limit=10)]
|
||||
assert resolve_limits(rows).tokens == ResolvedLimit(10.0, "instance")
|
||||
|
||||
def test_a_zero_team_does_not_block_a_member_of_a_funded_team(self):
|
||||
rows = [_row("team", "blocked", token_limit=0), _row("team", "funded", token_limit=50)]
|
||||
assert resolve_limits(rows).tokens.limit == 50.0
|
||||
|
||||
def test_disabled_team_rows_are_ignored(self):
|
||||
rows = [_row("team", "big", token_limit=900, enabled=False), _row("team", "small", token_limit=100)]
|
||||
assert resolve_limits(rows).tokens == ResolvedLimit(100.0, "team", "small")
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestBuckets:
|
||||
def test_only_rows_of_the_bucket_apply(self):
|
||||
rows = [
|
||||
_row("instance", token_limit=10),
|
||||
{**_row("instance", token_limit=3), "bucket": "agent"},
|
||||
]
|
||||
assert resolve_limits(rows, "all").tokens.limit == 10.0
|
||||
assert resolve_limits(rows, "agent").tokens.limit == 3.0
|
||||
assert resolve_limits(rows, "direct").unlimited
|
||||
|
||||
def test_decimal_limits_become_floats(self):
|
||||
from decimal import Decimal
|
||||
|
||||
resolved = resolve_limits([_row("instance", cost_limit_usd=Decimal("12.5000"))]).cost
|
||||
assert resolved.limit == 12.5 and isinstance(resolved.limit, float)
|
||||
@@ -0,0 +1,237 @@
|
||||
"""Tests for QuotaService against a real Postgres instance."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from contextlib import contextmanager
|
||||
from datetime import datetime, timedelta, timezone
|
||||
|
||||
import pytest
|
||||
from sqlalchemy import text
|
||||
|
||||
from docsgpt.quotas import providers
|
||||
from docsgpt.quotas.providers import QuotaDefaultsProvider, register_defaults_provider
|
||||
from docsgpt.quotas.service import QuotaService
|
||||
from docsgpt.storage.db.repositories.quota_policies import QuotaPoliciesRepository
|
||||
from docsgpt.storage.db.repositories.token_usage import TokenUsageRepository
|
||||
|
||||
NOW = datetime(2026, 9, 23, 12, tzinfo=timezone.utc)
|
||||
THIS_MONTH = NOW - timedelta(days=2)
|
||||
LAST_MONTH = NOW - timedelta(days=40)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def conn(pg_conn, monkeypatch):
|
||||
@contextmanager
|
||||
def _readonly():
|
||||
yield pg_conn
|
||||
|
||||
monkeypatch.setattr("docsgpt.quotas.service.db_readonly", _readonly)
|
||||
monkeypatch.setattr("docsgpt.quotas.service.settings.QUOTA_PERIOD", "month")
|
||||
return pg_conn
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _restore_provider():
|
||||
original = providers.get_defaults_provider()
|
||||
yield
|
||||
register_defaults_provider(original)
|
||||
|
||||
|
||||
def _use(conn, user_id="u1", tokens=0, cost=0.0, api_key=None, when=THIS_MONTH, source="agent_stream"):
|
||||
TokenUsageRepository(conn).insert(
|
||||
user_id=user_id, api_key=api_key, prompt_tokens=tokens, cost=cost, timestamp=when, source=source
|
||||
)
|
||||
|
||||
|
||||
def _policy(conn, scope, subject_id=None, **fields):
|
||||
return QuotaPoliciesRepository(conn).upsert(scope=scope, subject_id=subject_id, **fields)
|
||||
|
||||
|
||||
def _team_with_member(conn, slug, user_id="u1"):
|
||||
team_id = str(
|
||||
conn.execute(
|
||||
text("INSERT INTO teams (name, slug, owner_id) VALUES (:n, :s, 'o') RETURNING id"),
|
||||
{"n": slug, "s": slug},
|
||||
).scalar()
|
||||
)
|
||||
conn.execute(
|
||||
text("INSERT INTO team_members (team_id, user_id, role) VALUES (CAST(:t AS uuid), :u, 'team_member')"),
|
||||
{"t": team_id, "u": user_id},
|
||||
)
|
||||
return team_id
|
||||
|
||||
|
||||
class TestCheck:
|
||||
def test_no_policies_allows(self, conn):
|
||||
_use(conn, tokens=10**9)
|
||||
assert QuotaService.check("u1", now=NOW) is None
|
||||
|
||||
def test_no_user_allows(self, conn):
|
||||
assert QuotaService.check(None, now=NOW) is None
|
||||
|
||||
def test_under_the_limit_allows(self, conn):
|
||||
_policy(conn, "instance", token_limit=100)
|
||||
_use(conn, tokens=99)
|
||||
assert QuotaService.check("u1", now=NOW) is None
|
||||
|
||||
def test_reaching_the_token_limit_blocks(self, conn):
|
||||
_policy(conn, "instance", token_limit=100)
|
||||
_use(conn, tokens=100)
|
||||
exceeded = QuotaService.check("u1", now=NOW)
|
||||
assert (exceeded.budget, exceeded.usage, exceeded.limit, exceeded.source) == ("tokens", 100, 100.0, "instance")
|
||||
assert exceeded.resets_at == datetime(2026, 10, 1, tzinfo=timezone.utc)
|
||||
|
||||
def test_cost_limit_blocks(self, conn):
|
||||
_policy(conn, "user", "u1", cost_limit_usd=1.5)
|
||||
_use(conn, cost=1.0)
|
||||
_use(conn, cost=0.5)
|
||||
exceeded = QuotaService.check("u1", now=NOW)
|
||||
assert (exceeded.budget, exceeded.usage, exceeded.source) == ("cost", 1.5, "user")
|
||||
|
||||
def test_zero_limit_blocks_without_usage(self, conn):
|
||||
_policy(conn, "user", "u1", token_limit=0)
|
||||
assert QuotaService.check("u1", now=NOW).limit == 0
|
||||
|
||||
def test_last_periods_usage_does_not_count(self, conn):
|
||||
_policy(conn, "instance", token_limit=100)
|
||||
_use(conn, tokens=500, when=LAST_MONTH)
|
||||
assert QuotaService.check("u1", now=NOW) is None
|
||||
|
||||
def test_other_users_usage_does_not_count(self, conn):
|
||||
_policy(conn, "instance", token_limit=100)
|
||||
_use(conn, user_id="u2", tokens=500)
|
||||
assert QuotaService.check("u1", now=NOW) is None
|
||||
|
||||
def test_scheduler_rollups_do_not_count(self, conn):
|
||||
_policy(conn, "instance", token_limit=100)
|
||||
_use(conn, tokens=60)
|
||||
_use(conn, tokens=60, source="schedule")
|
||||
assert QuotaService.check("u1", now=NOW) is None
|
||||
|
||||
def test_user_override_lifts_the_instance_limit(self, conn):
|
||||
_policy(conn, "instance", token_limit=100)
|
||||
_policy(conn, "user", "u1", token_unlimited=True)
|
||||
_use(conn, tokens=10**6)
|
||||
assert QuotaService.check("u1", now=NOW) is None
|
||||
|
||||
def test_period_setting_moves_the_window(self, conn, monkeypatch):
|
||||
monkeypatch.setattr("docsgpt.quotas.service.settings.QUOTA_PERIOD", "day")
|
||||
_policy(conn, "instance", token_limit=100)
|
||||
_use(conn, tokens=500, when=NOW - timedelta(days=1))
|
||||
assert QuotaService.check("u1", now=NOW) is None
|
||||
_use(conn, tokens=500, when=NOW - timedelta(hours=1))
|
||||
assert QuotaService.check("u1", now=NOW).resets_at == datetime(2026, 9, 24, tzinfo=timezone.utc)
|
||||
|
||||
def test_unknown_bucket_rejected(self, conn):
|
||||
with pytest.raises(ValueError):
|
||||
QuotaService.check("u1", bucket="all", now=NOW)
|
||||
|
||||
def test_failure_allows_the_request(self, monkeypatch):
|
||||
def boom():
|
||||
raise RuntimeError("db down")
|
||||
|
||||
monkeypatch.setattr("docsgpt.quotas.service.db_readonly", boom)
|
||||
assert QuotaService.check("u1", now=NOW) is None
|
||||
|
||||
|
||||
class TestTeams:
|
||||
def test_the_most_generous_team_sets_the_allowance(self, conn):
|
||||
_policy(conn, "team", _team_with_member(conn, "svc-small"), token_limit=100)
|
||||
big = _team_with_member(conn, "svc-big")
|
||||
_policy(conn, "team", big, token_limit=1000)
|
||||
_use(conn, tokens=500)
|
||||
assert QuotaService.check("u1", now=NOW) is None
|
||||
_use(conn, tokens=500)
|
||||
exceeded = QuotaService.check("u1", now=NOW)
|
||||
assert (exceeded.limit, exceeded.source, exceeded.source_id) == (1000.0, "team", big)
|
||||
|
||||
def test_usage_is_one_total_not_one_per_team(self, conn):
|
||||
for slug in ("svc-a", "svc-b", "svc-c"):
|
||||
_policy(conn, "team", _team_with_member(conn, slug), token_limit=100)
|
||||
_use(conn, tokens=100)
|
||||
assert QuotaService.check("u1", now=NOW).limit == 100.0
|
||||
|
||||
def test_a_team_only_covers_its_members(self, conn):
|
||||
_policy(conn, "instance", token_limit=100)
|
||||
_policy(conn, "team", _team_with_member(conn, "svc-vip", user_id="u2"), token_unlimited=True)
|
||||
_use(conn, tokens=100)
|
||||
assert QuotaService.check("u1", now=NOW).source == "instance"
|
||||
|
||||
|
||||
class TestBuckets:
|
||||
def test_bucket_policies_only_see_their_traffic(self, conn):
|
||||
_policy(conn, "instance", bucket="agent", token_limit=100)
|
||||
_use(conn, tokens=500)
|
||||
_use(conn, tokens=90, api_key="k")
|
||||
assert QuotaService.check("u1", "direct", now=NOW) is None
|
||||
assert QuotaService.check("u1", "agent", now=NOW) is None
|
||||
_use(conn, tokens=10, api_key="k")
|
||||
exceeded = QuotaService.check("u1", "agent", now=NOW)
|
||||
assert (exceeded.bucket, exceeded.usage) == ("agent", 100)
|
||||
assert QuotaService.check("u1", "direct", now=NOW) is None
|
||||
|
||||
def test_the_all_bucket_applies_to_both_kinds_of_traffic(self, conn):
|
||||
_policy(conn, "instance", token_limit=100)
|
||||
_use(conn, tokens=60)
|
||||
_use(conn, tokens=60, api_key="k")
|
||||
assert QuotaService.check("u1", "direct", now=NOW).bucket == "all"
|
||||
assert QuotaService.check("u1", "agent", now=NOW).bucket == "all"
|
||||
|
||||
|
||||
class TestProviderDefaults:
|
||||
class _Plan(QuotaDefaultsProvider):
|
||||
def default_policies(self, user_id):
|
||||
return [{"cost_limit_usd": 5.0}] if user_id == "u1" else []
|
||||
|
||||
def error_payload(self, payload, user_id):
|
||||
return {**payload, "error_code": "free-limit-reached"}
|
||||
|
||||
def test_defaults_apply_without_stored_rows(self, conn):
|
||||
register_defaults_provider(self._Plan())
|
||||
_use(conn, cost=5.0)
|
||||
exceeded = QuotaService.check("u1", now=NOW)
|
||||
assert exceeded.source == "default"
|
||||
assert exceeded.to_payload()["error_code"] == "free-limit-reached"
|
||||
|
||||
def test_a_broken_provider_payload_falls_back(self, conn):
|
||||
class _Broken(self._Plan):
|
||||
def error_payload(self, payload, user_id):
|
||||
raise RuntimeError("nope")
|
||||
|
||||
register_defaults_provider(_Broken())
|
||||
_use(conn, cost=5.0)
|
||||
assert QuotaService.check("u1", now=NOW).to_payload()["error_code"] == "quota-exceeded"
|
||||
|
||||
|
||||
class TestStatusAndPayload:
|
||||
def test_status_reports_limits_and_usage(self, conn):
|
||||
_policy(conn, "instance", token_limit=100, cost_limit_usd=2)
|
||||
_use(conn, tokens=40, cost=0.5)
|
||||
(status,) = QuotaService.status("u1", now=NOW)
|
||||
assert status.to_dict() == {
|
||||
"bucket": "all",
|
||||
"tokens": {"limit": 100, "used": 40, "source": "instance", "source_id": None},
|
||||
"cost": {"limit": 2.0, "used": 0.5, "source": "instance", "source_id": None},
|
||||
"resets_at": "2026-10-01T00:00:00+00:00",
|
||||
}
|
||||
|
||||
def test_unlimited_users_skip_the_usage_query(self, conn, monkeypatch):
|
||||
def fail(*args, **kwargs):
|
||||
raise AssertionError("usage must not be summed for an unlimited user")
|
||||
|
||||
monkeypatch.setattr(TokenUsageRepository, "usage_totals", fail)
|
||||
(status,) = QuotaService.status("u1", now=NOW)
|
||||
assert status.limits.unlimited and status.tokens_used == 0
|
||||
|
||||
def test_payload_shape(self, conn):
|
||||
_policy(conn, "user", "u1", cost_limit_usd=1)
|
||||
_use(conn, cost=1.25)
|
||||
exceeded = QuotaService.check("u1", now=NOW)
|
||||
payload = exceeded.to_payload()
|
||||
assert payload["success"] is False
|
||||
assert payload["error_code"] == "quota-exceeded"
|
||||
assert (payload["dimension"], payload["unit"]) == ("cost", "USD")
|
||||
assert (payload["usage"], payload["limit"]) == (1.25, 1.0)
|
||||
assert payload["resets_at"] == "2026-10-01T00:00:00+00:00"
|
||||
assert "$1.25 of $1.00" in payload["message"]
|
||||
assert exceeded.retry_after_seconds >= 1
|
||||
@@ -0,0 +1,55 @@
|
||||
"""Tests for docsgpt/quotas/windows.py."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timedelta, timezone
|
||||
|
||||
import pytest
|
||||
|
||||
from docsgpt.quotas.windows import window_bounds
|
||||
|
||||
UTC = timezone.utc
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestWindowBounds:
|
||||
@pytest.mark.parametrize(
|
||||
"period, start, end",
|
||||
[
|
||||
("day", datetime(2026, 9, 23, tzinfo=UTC), datetime(2026, 9, 24, tzinfo=UTC)),
|
||||
("week", datetime(2026, 9, 21, tzinfo=UTC), datetime(2026, 9, 28, tzinfo=UTC)),
|
||||
("month", datetime(2026, 9, 1, tzinfo=UTC), datetime(2026, 10, 1, tzinfo=UTC)),
|
||||
],
|
||||
)
|
||||
def test_midweek(self, period, start, end):
|
||||
now = datetime(2026, 9, 23, 15, 30, 12, 99, tzinfo=UTC) # a Wednesday
|
||||
assert window_bounds(period, now) == (start, end)
|
||||
|
||||
def test_week_starts_on_monday_itself(self):
|
||||
monday = datetime(2026, 9, 21, 0, 0, tzinfo=UTC)
|
||||
assert window_bounds("week", monday)[0] == monday
|
||||
|
||||
def test_month_rolls_over_the_year(self):
|
||||
start, end = window_bounds("month", datetime(2026, 12, 31, 23, 59, tzinfo=UTC))
|
||||
assert (start, end) == (datetime(2026, 12, 1, tzinfo=UTC), datetime(2027, 1, 1, tzinfo=UTC))
|
||||
|
||||
def test_leap_february(self):
|
||||
start, end = window_bounds("month", datetime(2028, 2, 29, 12, tzinfo=UTC))
|
||||
assert (end - start).days == 29
|
||||
|
||||
def test_other_timezones_are_read_in_utc(self):
|
||||
tokyo = timezone(timedelta(hours=9))
|
||||
# 2026-10-01 08:00 in Tokyo is still 2026-09-30 in UTC.
|
||||
start, _ = window_bounds("month", datetime(2026, 10, 1, 8, tzinfo=tokyo))
|
||||
assert start == datetime(2026, 9, 1, tzinfo=UTC)
|
||||
|
||||
def test_naive_datetimes_are_utc(self):
|
||||
assert window_bounds("day", datetime(2026, 9, 23, 5))[0] == datetime(2026, 9, 23, tzinfo=UTC)
|
||||
|
||||
def test_defaults_to_now(self):
|
||||
start, end = window_bounds("day")
|
||||
assert start <= datetime.now(UTC) < end
|
||||
|
||||
def test_unknown_period(self):
|
||||
with pytest.raises(ValueError):
|
||||
window_bounds("year")
|
||||
@@ -0,0 +1,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
|
||||
@@ -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):
|
||||
|
||||
@@ -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"]
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
):
|
||||
|
||||
Reference in new issue
Block a user