"""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(), }