From 5406550bca8d25550b8ac00664dc3bbe6a5925a4 Mon Sep 17 00:00:00 2001 From: Alex Date: Mon, 21 Sep 2026 11:51:23 +0100 Subject: [PATCH] feat(quotas): enforce user quotas on chat, agent, scheduled and webhook runs check_usage now checks the billable user's quota on every request, before the per-agent 24h limits, which keep applying to traffic through an agent. Until now a request without an agent key skipped every limit. A refusal is a 429 with Retry-After and a body naming the budget, usage, limit, the layer the limit came from and when it resets. Headless runs check the agent owner's quota before starting. A refused scheduled run is recorded as budget_exceeded; a refused webhook run returns a quota_exceeded result instead of raising, so Celery does not retry it. --- docsgpt/agents/headless_runner.py | 10 ++- docsgpt/api/answer/routes/answer.py | 8 +- docsgpt/api/answer/routes/base.py | 18 +++- docsgpt/api/answer/routes/stream.py | 8 +- docsgpt/api/user/scheduler_worker.py | 6 ++ docsgpt/api/v1/routes.py | 2 +- docsgpt/quotas/__init__.py | 3 +- docsgpt/quotas/http.py | 14 +++ docsgpt/quotas/service.py | 8 ++ docsgpt/worker.py | 7 ++ tests/api/user/test_scheduler_worker.py | 25 ++++++ tests/quotas/test_enforcement.py | 108 ++++++++++++++++++++++++ tests/worker/test_agent_workers.py | 30 +++++++ 13 files changed, 238 insertions(+), 9 deletions(-) create mode 100644 docsgpt/quotas/http.py create mode 100644 tests/quotas/test_enforcement.py diff --git a/docsgpt/agents/headless_runner.py b/docsgpt/agents/headless_runner.py index 59c40c5c..344e5fe3 100644 --- a/docsgpt/agents/headless_runner.py +++ b/docsgpt/agents/headless_runner.py @@ -15,6 +15,7 @@ from docsgpt.api.answer.services.prompt_renderer import ( ) from docsgpt.api.answer.services.stream_processor import get_prompt from docsgpt.core.settings import settings +from docsgpt.quotas.service import QuotaExceededError, QuotaService from docsgpt.retriever.retriever_creator import RetrieverCreator from docsgpt.storage.db.repositories.sources import SourcesRepository from docsgpt.storage.db.session import db_readonly @@ -69,7 +70,11 @@ def run_agent_headless( chat_history: Optional[List[Dict[str, Any]]] = None, conversation_id: Optional[str] = None, ) -> Dict[str, Any]: - """Run an agent with no live client; returns a structured outcome dict.""" + """Run an agent with no live client; returns a structured outcome dict. + + Raises: + QuotaExceededError: If the agent owner's usage quota is exhausted. + """ from docsgpt.core.model_utils import ( get_api_key_for_provider, get_default_model_id, @@ -82,6 +87,9 @@ def run_agent_headless( if not owner: raise ValueError("Agent config is missing user_id; cannot run headless.") decoded_token = {"sub": owner} + exceeded = QuotaService.check(owner, "agent" if agent_config.get("key") else "direct") + if exceeded is not None: + raise QuotaExceededError(exceeded) retriever_kind = agent_config.get("retriever", "classic") source_id = agent_config.get("source_id") or agent_config.get("source") diff --git a/docsgpt/api/answer/routes/answer.py b/docsgpt/api/answer/routes/answer.py index 17fac152..cdc63cca 100644 --- a/docsgpt/api/answer/routes/answer.py +++ b/docsgpt/api/answer/routes/answer.py @@ -103,7 +103,9 @@ class AnswerResource(Resource, BaseAnswerResource): ) if not processor.decoded_token: return make_response({"error": "Unauthorized"}, 401) - if error := self.check_usage(processor.agent_config): + if error := self.check_usage( + processor.agent_config, processor.decoded_token + ): return error stream = self.complete_stream( question="", @@ -129,7 +131,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( + processor.agent_config, processor.decoded_token + ): return error should_persist, visibility = resolve_persistence( diff --git a/docsgpt/api/answer/routes/base.py b/docsgpt/api/answer/routes/base.py index 748f0dda..c61cfd48 100644 --- a/docsgpt/api/answer/routes/base.py +++ b/docsgpt/api/answer/routes/base.py @@ -23,6 +23,8 @@ from docsgpt.core.model_utils import ( from docsgpt.core.settings import settings from docsgpt.error import sanitize_api_error from docsgpt.llm.llm_creator import LLMCreator +from docsgpt.quotas.http import quota_exceeded_response +from docsgpt.quotas.service import QuotaService from docsgpt.storage.db.repositories.agents import AgentsRepository from docsgpt.storage.db.repositories.conversations import ( HeartbeatState, @@ -112,17 +114,29 @@ 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 + ) -> 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. 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 else "direct") + if exceeded is not None: + return quota_exceeded_response(exceeded) if not api_key: return None with db_readonly() as conn: diff --git a/docsgpt/api/answer/routes/stream.py b/docsgpt/api/answer/routes/stream.py index f79e0633..fdc5a00b 100644 --- a/docsgpt/api/answer/routes/stream.py +++ b/docsgpt/api/answer/routes/stream.py @@ -115,7 +115,9 @@ class StreamResource(Resource, BaseAnswerResource): status=401, mimetype="text/event-stream", ) - if error := self.check_usage(processor.agent_config): + if error := self.check_usage( + processor.agent_config, processor.decoded_token + ): return error return Response( with_sse_keepalive( @@ -151,7 +153,9 @@ 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 + ): return error should_persist, visibility = resolve_persistence( visibility_flag=data.get("visibility"), diff --git a/docsgpt/api/user/scheduler_worker.py b/docsgpt/api/user/scheduler_worker.py index 6207832d..3e085f0e 100644 --- a/docsgpt/api/user/scheduler_worker.py +++ b/docsgpt/api/user/scheduler_worker.py @@ -17,6 +17,7 @@ from sqlalchemy import text as sql_text from docsgpt.agents.headless_runner import run_agent_headless from docsgpt.core.settings import settings from docsgpt.events.publisher import publish_user_event +from docsgpt.quotas.service import QuotaExceededError from docsgpt.storage.db.base_repository import row_to_dict from docsgpt.storage.db.engine import get_engine from docsgpt.storage.db.repositories.conversations import ( @@ -282,6 +283,11 @@ def execute_scheduled_run_body(run_id: str, celery_task_id: Optional[str]) -> Di outcome = {"answer": "", "tool_calls": [], "sources": [], "thought": ""} error_type = "timeout" error_text = "run exceeded soft time limit" + except QuotaExceededError as exc: + # The owner's usage quota is spent; the run never started. + outcome = {"answer": "", "tool_calls": [], "sources": [], "thought": ""} + error_type = "budget_exceeded" + error_text = str(exc) except Exception as exc: outcome = {"answer": "", "tool_calls": [], "sources": [], "thought": ""} error_type = "agent_error" diff --git a/docsgpt/api/v1/routes.py b/docsgpt/api/v1/routes.py index 4137d6dc..cd208e84 100644 --- a/docsgpt/api/v1/routes.py +++ b/docsgpt/api/v1/routes.py @@ -338,7 +338,7 @@ def chat_completions(): ) helper = _V1AnswerHelper() - usage_error = helper.check_usage(processor.agent_config) + usage_error = helper.check_usage(processor.agent_config, processor.decoded_token) if usage_error: return usage_error diff --git a/docsgpt/quotas/__init__.py b/docsgpt/quotas/__init__.py index ae1b5f4f..9ed71a75 100644 --- a/docsgpt/quotas/__init__.py +++ b/docsgpt/quotas/__init__.py @@ -8,13 +8,14 @@ with their ``token_usage`` totals over the current ``QUOTA_PERIOD`` window. from docsgpt.quotas.providers import QuotaDefaultsProvider, register_defaults_provider from docsgpt.quotas.resolver import ResolvedLimit, ResolvedLimits, resolve_limits -from docsgpt.quotas.service import BucketStatus, QuotaExceeded, QuotaService +from docsgpt.quotas.service import BucketStatus, QuotaExceeded, QuotaExceededError, QuotaService from docsgpt.quotas.windows import window_bounds __all__ = [ "BucketStatus", "QuotaDefaultsProvider", "QuotaExceeded", + "QuotaExceededError", "QuotaService", "ResolvedLimit", "ResolvedLimits", diff --git a/docsgpt/quotas/http.py b/docsgpt/quotas/http.py new file mode 100644 index 00000000..494b87cf --- /dev/null +++ b/docsgpt/quotas/http.py @@ -0,0 +1,14 @@ +"""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) + return response diff --git a/docsgpt/quotas/service.py b/docsgpt/quotas/service.py index 775f3109..7659fe1d 100644 --- a/docsgpt/quotas/service.py +++ b/docsgpt/quotas/service.py @@ -104,6 +104,14 @@ class QuotaExceeded: 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.""" diff --git a/docsgpt/worker.py b/docsgpt/worker.py index 3ce7f15e..947f59c2 100755 --- a/docsgpt/worker.py +++ b/docsgpt/worker.py @@ -2122,6 +2122,7 @@ def agent_webhook_worker(self, agent_id, payload): try: # Shared headless path with the scheduler; approval-gated tools auto-deny. from docsgpt.agents.headless_runner import run_agent_headless + from docsgpt.quotas.service import QuotaExceededError outcome = run_agent_headless( agent_config, @@ -2135,6 +2136,12 @@ def agent_webhook_worker(self, agent_id, payload): "tool_calls": outcome.get("tool_calls", []), "thought": outcome.get("thought", ""), } + except QuotaExceededError as e: + # Returned, not raised: retrying cannot succeed before the quota resets. + logging.warning( + f"Webhook skipped for agent {agent_id}: {e}", extra={"agent_id": agent_id} + ) + return {"status": "quota_exceeded", "error": str(e)} except Exception as e: logging.error(f"Error running agent logic: {e}", exc_info=True) raise diff --git a/tests/api/user/test_scheduler_worker.py b/tests/api/user/test_scheduler_worker.py index 5b7b55ea..09528cd5 100644 --- a/tests/api/user/test_scheduler_worker.py +++ b/tests/api/user/test_scheduler_worker.py @@ -126,6 +126,31 @@ class TestExecuteScheduledRunBody: assert sched["consecutive_failure_count"] == 1 assert "schedule.run.failed" in {e[0] for e in stub_events} + def test_quota_refusal_marks_budget_exceeded( + self, pg_engine, patched_engine, stub_events, + ): + from datetime import datetime, timezone + + from docsgpt.quotas.service import QuotaExceeded, QuotaExceededError + + exceeded = QuotaExceeded( + user_id="u1", bucket="all", budget="tokens", usage=10, limit=10, + source="instance", source_id=None, + resets_at=datetime(2099, 1, 1, tzinfo=timezone.utc), + ) + with pg_engine.begin() as conn: + _, run, _ = _make_pending_run(conn) + with patch( + "docsgpt.api.user.scheduler_worker.run_agent_headless", + side_effect=QuotaExceededError(exceeded), + ): + result = execute_scheduled_run_body(str(run["id"]), "celery-q") + assert result["status"] == "failed" + with pg_engine.connect() as conn: + row = ScheduleRunsRepository(conn).get_internal(str(run["id"])) + assert row["error_type"] == "budget_exceeded" + assert "Usage quota reached" in row["error"] + def test_autopause_after_threshold( self, pg_engine, patched_engine, stub_events, ): diff --git a/tests/quotas/test_enforcement.py b/tests/quotas/test_enforcement.py new file mode 100644 index 00000000..b02dbae1 --- /dev/null +++ b/tests/quotas/test_enforcement.py @@ -0,0 +1,108 @@ +"""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): + from docsgpt.api.answer.routes.base import BaseAnswerResource + + with flask_app.app_context(): + return BaseAnswerResource().check_usage(agent_config, decoded_token) + + +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 + + +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) + diff --git a/tests/worker/test_agent_workers.py b/tests/worker/test_agent_workers.py index 8a863f64..033f97d0 100644 --- a/tests/worker/test_agent_workers.py +++ b/tests/worker/test_agent_workers.py @@ -110,6 +110,36 @@ class TestAgentWebhookWorker: with pytest.raises(RuntimeError, match="LLM exploded"): worker.agent_webhook_worker(task_self, agent_id, {"event": "ping"}) + def test_quota_refusal_is_returned_not_retried( + self, pg_conn, patch_worker_db, task_self, monkeypatch + ): + """A spent quota cannot succeed on retry, so the task must not raise.""" + from datetime import datetime, timezone + + from docsgpt import worker + from docsgpt.agents import headless_runner + from docsgpt.quotas.service import QuotaExceeded, QuotaExceededError + + agent = AgentsRepository(pg_conn).create( + user_id="alice", name="hook-agent", status="active", + agent_type="classic", retriever="classic", chunks=2, key="sk-test-q", + ) + exceeded = QuotaExceeded( + user_id="alice", bucket="all", budget="cost", usage=5.0, limit=5.0, + source="user", source_id=None, + resets_at=datetime(2099, 1, 1, tzinfo=timezone.utc), + ) + + def _refuse(*a, **k): + raise QuotaExceededError(exceeded) + + monkeypatch.setattr(headless_runner, "run_agent_headless", _refuse) + + result = worker.agent_webhook_worker(task_self, str(agent["id"]), {"event": "ping"}) + + assert result["status"] == "quota_exceeded" + assert "$5.00 of $5.00" in result["error"] + def test_webhook_journals_headless_denial_for_approval_gated_tool( self, pg_conn, patch_worker_db, task_self, monkeypatch ):