154 lines
5.0 KiB
Python
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(),
|
|
}
|