Files
ObsiGate/backend/tools/ratelimit.py
T

154 lines
5.0 KiB
Python

"""Rate limiting for the AI tool layer (Phase F).
Complements the IP-based login limiter (``backend.ratelimit``) with a
per-identity, per-tool sliding-window limiter applied to every tool call,
whether it originates from the in-app assistant or the MCP server.
Two counters are maintained per identity:
* a **global** counter (all tools combined), capped by
``OBSIGATE_TOOL_RATE_LIMIT``;
* a **per-tool** counter, capped by ``OBSIGATE_TOOL_RATE_LIMIT_PER_TOOL``
(defaults to the global cap) so a single expensive tool cannot consume the
whole budget.
The window length is ``OBSIGATE_TOOL_RATE_WINDOW`` seconds (default 60).
The identity is the token JTI when available (``_token_jti`` attached by the
auth middleware), otherwise the user id/username. This gives "per token / per
tool" limiting while remaining meaningful for anonymous/disabled-auth mode.
"""
from __future__ import annotations
import logging
import os
import time
from collections import deque
from typing import Any
logger = logging.getLogger("obsigate.tools.ratelimit")
# --- Configuration (read at call time so tests can monkeypatch env) ---
DEFAULT_RATE_LIMIT = 60
DEFAULT_RATE_WINDOW = 60
def _env_int(name: str, default: int) -> int:
try:
return int(os.environ.get(name, str(default)))
except (TypeError, ValueError):
return default
def _global_limit() -> int:
return _env_int("OBSIGATE_TOOL_RATE_LIMIT", DEFAULT_RATE_LIMIT)
def _per_tool_limit() -> int:
return _env_int("OBSIGATE_TOOL_RATE_LIMIT_PER_TOOL", _global_limit())
def _window() -> int:
return max(1, _env_int("OBSIGATE_TOOL_RATE_WINDOW", DEFAULT_RATE_WINDOW))
# --- In-memory store: {counter_key: deque[timestamp]} ---
_calls: dict[str, deque] = {}
def _identity_key(identity: str | None) -> str:
return identity or "anonymous"
def _counter_key(identity: str | None, tool: str | None) -> str:
base = _identity_key(identity)
return f"{base}\x00{tool}" if tool else base
def _prune(entries: deque, now: float, window: int) -> None:
cutoff = now - window
while entries and entries[0] <= cutoff:
entries.popleft()
def check_and_record(identity: str | None, tool: str, *, now: float | None = None) -> tuple[bool, int]:
"""Record one call and report whether it is allowed.
Args:
identity: Rate-limit identity (token JTI or username).
tool: Tool name (used for the per-tool counter).
now: Optional timestamp override (tests).
Returns:
``(allowed, retry_after)``. When ``allowed`` is False, ``retry_after``
is the number of seconds until the oldest call leaves the window.
"""
now = time.time() if now is None else now
window = _window()
global_limit = _global_limit()
tool_limit = _per_tool_limit()
global_entries = _calls.setdefault(_counter_key(identity, None), deque())
tool_entries = _calls.setdefault(_counter_key(identity, tool), deque())
_prune(global_entries, now, window)
_prune(tool_entries, now, window)
if len(global_entries) >= global_limit or len(tool_entries) >= tool_limit:
oldest = min(
global_entries[0] if len(global_entries) >= global_limit else now,
tool_entries[0] if len(tool_entries) >= tool_limit else now,
)
retry_after = max(1, int(oldest + window - now) + 1)
logger.warning(f"Tool rate limit exceeded for '{_identity_key(identity)}' on '{tool}'")
return False, retry_after
global_entries.append(now)
tool_entries.append(now)
return True, 0
def remaining(identity: str | None, tool: str | None = None) -> int:
"""Return the number of calls still allowed in the current window."""
now = time.time()
window = _window()
if tool:
entries = _calls.get(_counter_key(identity, tool))
if entries is None:
return _per_tool_limit()
_prune(entries, now, window)
return max(0, _per_tool_limit() - len(entries))
entries = _calls.get(_counter_key(identity, None))
if entries is None:
return _global_limit()
_prune(entries, now, window)
return max(0, _global_limit() - len(entries))
def reset(identity: str | None = None) -> None:
"""Clear rate-limit state (all identities, or a single one)."""
if identity is None:
_calls.clear()
return
prefix = f"{_identity_key(identity)}\x00"
for key in [k for k in _calls if k == _identity_key(identity) or k.startswith(prefix)]:
del _calls[key]
def get_status(identity: str | None = None) -> dict[str, Any]:
"""Diagnostic snapshot of the limiter (for tests / admin tooling)."""
if identity is None:
return {
"tracked_identities": len({k.split("\x00", 1)[0] for k in _calls}),
"global_limit": _global_limit(),
"per_tool_limit": _per_tool_limit(),
"window_seconds": _window(),
}
return {
"identity": _identity_key(identity),
"remaining_global": remaining(identity),
"global_limit": _global_limit(),
"per_tool_limit": _per_tool_limit(),
"window_seconds": _window(),
}