mirror of
https://github.com/tiennm99/DocsGPT.git
synced 2026-10-11 03:12:55 +00:00
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.
This commit is contained in:
1 parent
6826313b60
commit
5406550bca
13 files changed
+238
-9
No files matched your search
@@ -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")
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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"),
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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
|
||||
@@ -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."""
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
):
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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