feat(ai): durcissement phase F (#79) - rate limit, redaction, OpenAPI/MCP, E2E

This commit is contained in:
2026-09-11 21:56:50 -04:00
parent 88eecd7671
commit 9dce341bc8
18 changed files with 854 additions and 27 deletions
+25
View File
@@ -17,6 +17,7 @@ from __future__ import annotations
import json
import logging
import os
from collections.abc import Callable
from dataclasses import dataclass, field
from typing import Any
@@ -35,11 +36,14 @@ logger = logging.getLogger("obsigate.agent.loop")
DEFAULT_MAX_ITERATIONS = 10
# Cap the size of a tool result fed back to the model (chars).
MAX_TOOL_RESULT_CHARS = 100_000
# Quota: maximum tool calls executed per agent run (``BOOKSLM_MAX_TOOL_CALLS``).
DEFAULT_MAX_TOOL_CALLS = int(os.environ.get("BOOKSLM_MAX_TOOL_CALLS", "25"))
# Stopping reasons
STOP_DONE = "done"
STOP_MAX_ITERATIONS = "max_iterations"
STOP_CONFIRMATION_REQUIRED = "confirmation_required"
STOP_QUOTA_EXCEEDED = "quota_exceeded"
@dataclass
@@ -158,6 +162,7 @@ async def run_agent(
llm: Callable[..., Any] | None = None,
tools: list[dict[str, Any]] | None = None,
max_iterations: int = DEFAULT_MAX_ITERATIONS,
max_tool_calls: int | None = None,
on_tool_call: Callable[[ToolCallRecord], None] | None = None,
resume_messages: list[dict[str, Any]] | None = None,
confirm_pending: dict[str, Any] | None = None,
@@ -173,6 +178,8 @@ async def run_agent(
tools: Tool schemas to expose. ``None`` exposes all in-app tools;
pass ``[]`` to disable tool calling (plain chat).
max_iterations: Hard cap on LLM round-trips.
max_tool_calls: Hard cap on the total number of executed tool calls
(quota, defaults to ``BOOKSLM_MAX_TOOL_CALLS``).
on_tool_call: Optional callback invoked after each executed tool call.
resume_messages: Conversation snapshot from a paused run (returned with
a ``confirmation`` event). When set, the loop resumes from it.
@@ -187,11 +194,20 @@ async def run_agent(
llm = llm or provider_llm
if tools is None:
tools = get_tool_schemas(scope=ToolScope.IN_APP)
quota = DEFAULT_MAX_TOOL_CALLS if max_tool_calls is None else max_tool_calls
convo = [dict(m) for m in (resume_messages if resume_messages is not None else messages)]
executed: list[ToolCallRecord] = []
if confirm_pending:
if quota is not None and len(executed) >= quota:
return AgentResult(
content="",
messages=convo,
tool_calls=executed,
iterations=0,
stopped=STOP_QUOTA_EXCEEDED,
)
_execute_confirmed(ctx, confirm_pending, convo, executed, on_tool_call)
for iteration in range(1, max_iterations + 1):
@@ -209,6 +225,15 @@ async def run_agent(
convo.append(_assistant_tool_message(response.content, response.tool_calls))
for call in response.tool_calls:
if quota is not None and len(executed) >= quota:
logger.warning(f"Agent reached the tool-call quota ({quota})")
return AgentResult(
content=response.content or "",
messages=convo,
tool_calls=executed,
iterations=iteration,
stopped=STOP_QUOTA_EXCEEDED,
)
try:
result = call_tool(call.name, ctx, call.arguments)
payload = result.data
+2
View File
@@ -63,6 +63,8 @@ def get_current_user(
# Attach vault permissions from the token (snapshot at login time)
user["_token_vaults"] = payload.get("vaults", [])
# Attach the token id for per-token rate limiting (AI tool layer).
user["_token_jti"] = payload.get("jti")
return user
+66
View File
@@ -32,6 +32,7 @@ TAGS_METADATA: list[dict[str, str]] = [
{"name": "Export", "description": "Export notes or whole vaults to HTML, Markdown bundle or ePub."},
{"name": "AI", "description": "AI-powered editor actions, provider status and model discovery."},
{"name": "BooksLM", "description": "Directory-scoped AI chat (NotebookLM-style) over a vault folder."},
{"name": "MCP", "description": "Model Context Protocol server (Streamable HTTP) exposing the shared AI tool layer to external clients (Claude Desktop, Cursor…)."},
{"name": "Sharing", "description": "Create and manage public read-only share links for documents."},
{"name": "Webhooks", "description": "HTTP callbacks signed with HMAC-SHA256 for file events."},
{"name": "Conflicts", "description": "Detect and resolve Syncthing sync-conflict files."},
@@ -79,6 +80,7 @@ _TAG_RULES: list[tuple[re.Pattern[str], str]] = [
(re.compile(r"^/api/push"), "Push"),
(re.compile(r"^/api/ai/bookslm"), "BooksLM"),
(re.compile(r"^/api/ai"), "AI"),
(re.compile(r"^/mcp"), "MCP"),
(re.compile(r"^/api/config/ai-"), "AI"),
(re.compile(r"^/api/share"), "Sharing"),
(re.compile(r"^/api/shares"), "Sharing"),
@@ -141,6 +143,7 @@ _TAG_ALIASES: dict[str, str] = {
"frontend": "Frontend",
"ai": "AI",
"bookslm": "BooksLM",
"mcp": "MCP",
"pdf": "PDF",
"bookmarks": "Bookmarks",
}
@@ -224,6 +227,67 @@ def _is_binary_operation(operation: dict[str, Any]) -> bool:
return any(media.startswith(_BINARY_MEDIA) for media in content)
# ---------------------------------------------------------------------------
# MCP endpoint (not a FastAPI route: custom ASGI mount) — documented manually
# ---------------------------------------------------------------------------
_MCP_DESCRIPTION = (
"**Model Context Protocol** server over Streamable HTTP (JSON-RPC 2.0). "
"Exposes the shared AI tool layer to external MCP clients (Claude Desktop, "
"Cursor…). Authentication uses `Authorization: Bearer <JWT>` (the same "
"token as the REST API).\n\n"
"Primitives: read/search **tools** directly; write/destructive tools as a "
"two-step `propose_<tool>` / `apply_<tool>` pair (signed, single-use "
"confirmation token); **resources** `vault://<name>` and "
"`vault://<name>/<path>` (read-only, secrets redacted); **prompts** "
"`summarize-directory`, `generate-note`, `find-related`.\n\n"
"See `docs/MCP_GUIDE.md` for client setup."
)
def _inject_mcp_path(schema: dict[str, Any]) -> None:
"""Add the MCP Streamable HTTP endpoint to the schema (idempotent)."""
paths = schema.setdefault("paths", {})
if "/mcp" in paths:
return
paths["/mcp"] = {
"post": {
"tags": ["MCP"],
"summary": "MCP Streamable HTTP endpoint (JSON-RPC 2.0)",
"operationId": "mcp_streamable_http",
"description": _MCP_DESCRIPTION,
"requestBody": {
"required": True,
"content": {
"application/json": {
"example": {
"jsonrpc": "2.0",
"id": 1,
"method": "tools/list",
"params": {},
}
}
},
},
"responses": {
"200": {
"description": "JSON-RPC response (or 202 for notifications)",
"content": {
"application/json": {
"example": {
"jsonrpc": "2.0",
"id": 1,
"result": {"tools": []},
}
}
},
}
},
"security": [{"bearerAuth": []}],
}
}
def enrich_openapi_schema(schema: dict[str, Any]) -> dict[str, Any]:
"""Enrich a FastAPI-generated OpenAPI schema in place and return it.
@@ -242,6 +306,8 @@ def enrich_openapi_schema(schema: dict[str, Any]) -> dict[str, Any]:
}
schema["tags"] = TAGS_METADATA
_inject_mcp_path(schema)
components = schema.setdefault("components", {})
security_schemes = components.setdefault("securitySchemes", {})
security_schemes.setdefault("bearerAuth", {
+16
View File
@@ -17,11 +17,22 @@ from backend.tools.context import (
ToolMode,
ToolNotFoundError,
ToolPermissionError,
ToolRateLimitError,
ToolRisk,
ToolScope,
ToolValidationError,
resolve_safe_path,
)
from backend.tools.ratelimit import (
check_and_record as check_tool_rate_limit,
)
from backend.tools.ratelimit import (
get_status as get_tool_rate_limit_status,
)
from backend.tools.ratelimit import (
reset as reset_tool_rate_limit,
)
from backend.tools.redaction import redact_payload
from backend.tools.registry import (
ToolSpec,
call_tool,
@@ -39,15 +50,20 @@ __all__ = [
"ToolMode",
"ToolNotFoundError",
"ToolPermissionError",
"ToolRateLimitError",
"ToolResult",
"ToolRisk",
"ToolScope",
"ToolSpec",
"ToolValidationError",
"call_tool",
"check_tool_rate_limit",
"get_tool",
"get_tool_rate_limit_status",
"get_tool_schemas",
"list_tools",
"redact_payload",
"reset_tool_rate_limit",
"resolve_safe_path",
"tool",
]
+30
View File
@@ -102,6 +102,21 @@ class ToolConfirmationRequired(ToolError):
return payload
class ToolRateLimitError(ToolError):
"""Raised when an identity exceeds its tool-call rate limit (Phase F)."""
code = "rate_limited"
def __init__(self, tool: str, retry_after: int = 0, message: str | None = None):
super().__init__(
message or f"Rate limit exceeded for tool '{tool}'",
code="rate_limited",
details={"tool": tool, "retry_after": retry_after},
)
self.tool = tool
self.retry_after = retry_after
@dataclass
class ToolContext:
"""Carries the caller identity and execution options for a tool call.
@@ -126,6 +141,21 @@ class ToolContext:
def username(self) -> str:
return self.user.get("username", "unknown")
@property
def rate_limit_identity(self) -> str:
"""Identity used by the per-token/per-tool rate limiter.
Prefers the JWT id (``_token_jti``, attached by the auth middleware),
then an explicit ``token_id`` in :attr:`metadata`, then the user id and
finally the username. This yields per-token limiting when a token id is
available and per-account limiting otherwise.
"""
token_id = self.metadata.get("token_id") or self.user.get("_token_jti")
if token_id:
return f"token:{token_id}"
user_id = self.user.get("id") or self.user.get("username")
return f"user:{user_id or 'anonymous'}"
def has_vault_access(self, vault: str) -> bool:
"""Return True if the caller may access *vault*."""
return bool(vault) and check_vault_access(vault, self.user)
+153
View File
@@ -0,0 +1,153 @@
"""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(),
}
+39
View File
@@ -0,0 +1,39 @@
"""Secret redaction for tool results (Phase F).
Tool results are fed back to the LLM (in-app agent loop) or returned to an
external MCP client, so they must never leak credentials. ``read_file`` and
``read_file_raw`` already redact through the shared file service, but other
tools (``diff_backup``, ``search_*``) return content that has not been through
the redactor. This module applies :func:`backend.secret_redactor.redact` to
every string in a tool payload, recursively.
"""
from __future__ import annotations
from typing import Any
from backend.secret_redactor import redact
_MAX_DEPTH = 12
def redact_payload(payload: Any, _depth: int = 0) -> Any:
"""Return a copy of *payload* with every string secret-redacted.
Walks dicts, lists and tuples; scalars are returned unchanged. Strings are
run through :func:`backend.secret_redactor.redact`. A recursion cap avoids
pathological/cyclic structures (tool results are JSON-serialisable, so
cycles should not occur, but the guard keeps this safe).
"""
if _depth > _MAX_DEPTH:
return payload
if isinstance(payload, str):
redacted, count = redact(payload)
return redacted if count else payload
if isinstance(payload, dict):
return {key: redact_payload(value, _depth + 1) for key, value in payload.items()}
if isinstance(payload, list):
return [redact_payload(item, _depth + 1) for item in payload]
if isinstance(payload, tuple):
return tuple(redact_payload(item, _depth + 1) for item in payload)
return payload
+12 -1
View File
@@ -26,10 +26,13 @@ from backend.tools.context import (
ToolError,
ToolNotFoundError,
ToolPermissionError,
ToolRateLimitError,
ToolRisk,
ToolScope,
ToolValidationError,
)
from backend.tools.ratelimit import check_and_record
from backend.tools.redaction import redact_payload
from backend.tools.schemas import ToolResult
logger = logging.getLogger("obsigate.tools.registry")
@@ -186,6 +189,12 @@ def call_tool(
_audit(ctx, spec, arguments, ok=False, error=error.code)
raise error from e
allowed, retry_after = check_and_record(ctx.rate_limit_identity, spec.name)
if not allowed:
rate_error = ToolRateLimitError(spec.name, retry_after)
_audit(ctx, spec, arguments, ok=False, error=rate_error.code)
raise rate_error
vault = getattr(params, "vault", None)
if spec.requires_vault and vault:
try:
@@ -223,4 +232,6 @@ def call_tool(
duration_ms = round((time.perf_counter() - started) * 1000, 2)
logger.debug(f"Tool '{name}' ok in {duration_ms}ms")
_audit(ctx, spec, arguments, ok=True)
return ToolResult(ok=True, data=data)
# Never forward raw secrets to the model (defense in depth: some tools —
# diffs, search snippets — return content that was not pre-redacted).
return ToolResult(ok=True, data=redact_payload(data))
+4 -2
View File
@@ -11,6 +11,7 @@ confirmation mechanism (two-step propose/apply).
from __future__ import annotations
import logging
import os
from typing import Any
from backend.indexer import get_backlinks as _get_backlinks
@@ -90,8 +91,9 @@ from backend.tools.schemas import (
logger = logging.getLogger("obsigate.tools.service")
# Maximum file size returned by ``read_file`` (bytes).
TOOL_MAX_READ_BYTES = 200_000
# Maximum file size returned by ``read_file`` (bytes). Quota configurable via
# ``BOOKSLM_MAX_TOOL_READ_BYTES``.
TOOL_MAX_READ_BYTES = int(os.environ.get("BOOKSLM_MAX_TOOL_READ_BYTES", "200000"))
# ── C1. Vaults / navigation ────────────────────────────────────────────────