CI / lint (push) Successful in 1m3s
CI / security (push) Successful in 41s
CI / test (push) Successful in 1m20s
CI / build (push) Successful in 1m19s
CI / e2e (push) Successful in 10m56s
Desktop Build / build-windows (push) Canceled after 0s
Desktop Build / build-linux (push) Canceled after 0s
143 lines
4.8 KiB
Python
143 lines
4.8 KiB
Python
"""Signed, single-use confirmation tokens for MCP mutations (Phase E4).
|
|
|
|
MCP has no "Apply" button, so mutating tools are exposed in two steps:
|
|
``propose_<tool>`` returns a preview plus a **signed token**, and
|
|
``apply_<tool>`` consumes that token to execute the mutation.
|
|
|
|
The token is a short-lived JWT (same secret as the app) carrying the tool name
|
|
and its arguments. Single use is enforced by a persisted JTI blacklist, which
|
|
also protects against replay and TOCTOU (a stale proposal cannot be re-applied).
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import json
|
|
import logging
|
|
import os
|
|
import threading
|
|
import time
|
|
import uuid
|
|
from pathlib import Path
|
|
from typing import Any
|
|
|
|
from jose import JWTError, jwt
|
|
|
|
from backend.auth.jwt_handler import get_secret_key
|
|
|
|
logger = logging.getLogger("obsigate.mcp.confirmations")
|
|
|
|
ALGORITHM = "HS256"
|
|
TOKEN_TYPE = "mcp_confirmation"
|
|
|
|
# Default token lifetime (seconds); override with OBSIGATE_MCP_CONFIRMATION_TTL.
|
|
DEFAULT_TTL = int(os.environ.get("OBSIGATE_MCP_CONFIRMATION_TTL", "300"))
|
|
|
|
_USED_TOKENS_FILE = Path("data/mcp_used_tokens.json")
|
|
_used_lock = threading.RLock()
|
|
_used_loaded = False
|
|
_used_jtis: dict[str, int] = {}
|
|
|
|
|
|
class ConfirmationError(Exception):
|
|
"""Raised when a confirmation token is invalid, expired, or already used."""
|
|
|
|
def __init__(self, message: str, *, code: str = "invalid_confirmation"):
|
|
super().__init__(message)
|
|
self.message = message
|
|
self.code = code
|
|
|
|
|
|
def _load_used() -> None:
|
|
"""Load used JTIs from disk once, dropping expired entries."""
|
|
global _used_loaded, _used_jtis
|
|
if _used_loaded:
|
|
return
|
|
with _used_lock:
|
|
if _used_loaded:
|
|
return
|
|
if _USED_TOKENS_FILE.exists():
|
|
try:
|
|
data = json.loads(_USED_TOKENS_FILE.read_text(encoding="utf-8"))
|
|
now = int(time.time())
|
|
_used_jtis = {jti: exp for jti, exp in data.items() if int(exp) > now}
|
|
except Exception as e: # pragma: no cover - corrupt store
|
|
logger.warning(f"Failed to load used MCP tokens: {e}")
|
|
_used_jtis = {}
|
|
_used_loaded = True
|
|
|
|
|
|
def _save_used() -> None:
|
|
"""Persist used JTIs with their expiry (best-effort)."""
|
|
try:
|
|
_USED_TOKENS_FILE.parent.mkdir(parents=True, exist_ok=True)
|
|
tmp = _USED_TOKENS_FILE.with_suffix(".tmp")
|
|
tmp.write_text(json.dumps(_used_jtis), encoding="utf-8")
|
|
tmp.replace(_USED_TOKENS_FILE)
|
|
except Exception as e: # pragma: no cover - disk error
|
|
logger.warning(f"Failed to persist used MCP tokens: {e}")
|
|
|
|
|
|
def create_confirmation_token(
|
|
tool: str,
|
|
arguments: dict[str, Any],
|
|
username: str,
|
|
*,
|
|
ttl: int | None = None,
|
|
) -> str:
|
|
"""Create a signed confirmation token for a pending mutation."""
|
|
now = int(time.time())
|
|
payload = {
|
|
"type": TOKEN_TYPE,
|
|
"tool": tool,
|
|
"arguments": arguments,
|
|
"sub": username,
|
|
"jti": str(uuid.uuid4()),
|
|
"iat": now,
|
|
"exp": now + (ttl if ttl is not None else DEFAULT_TTL),
|
|
}
|
|
return jwt.encode(payload, get_secret_key(), algorithm=ALGORITHM)
|
|
|
|
|
|
def peek_confirmation_token(token: str) -> dict[str, Any]:
|
|
"""Decode and validate a token without consuming it (signature + expiry)."""
|
|
try:
|
|
payload = jwt.decode(token, get_secret_key(), algorithms=[ALGORITHM])
|
|
except JWTError as e:
|
|
raise ConfirmationError("Invalid or expired confirmation token", code="invalid_confirmation") from e
|
|
|
|
if payload.get("type") != TOKEN_TYPE:
|
|
raise ConfirmationError("Wrong token type", code="invalid_confirmation")
|
|
if not payload.get("tool") or not payload.get("sub"):
|
|
raise ConfirmationError("Malformed confirmation token", code="invalid_confirmation")
|
|
return payload
|
|
|
|
|
|
def consume_confirmation_token(token: str, username: str) -> dict[str, Any]:
|
|
"""Validate a token, enforce single use, and return its payload.
|
|
|
|
Raises:
|
|
ConfirmationError: invalid/expired token, wrong user, or replay.
|
|
"""
|
|
payload = peek_confirmation_token(token)
|
|
|
|
if payload.get("sub") != username:
|
|
raise ConfirmationError("Confirmation token does not belong to this user", code="forbidden")
|
|
|
|
jti = payload.get("jti")
|
|
if not jti:
|
|
raise ConfirmationError("Malformed confirmation token", code="invalid_confirmation")
|
|
|
|
_load_used()
|
|
now = int(time.time())
|
|
with _used_lock:
|
|
if jti in _used_jtis:
|
|
raise ConfirmationError("Confirmation token already used", code="token_reused")
|
|
# Record as used *before* returning so a concurrent replay is rejected.
|
|
_used_jtis[jti] = int(payload.get("exp", now))
|
|
# Opportunistic cleanup of expired JTIs.
|
|
for old_jti in [k for k, exp in _used_jtis.items() if exp <= now]:
|
|
_used_jtis.pop(old_jti, None)
|
|
_save_used()
|
|
|
|
return payload
|