Files
AletheiaVox 7bb9d8473b fix(security): require auth on every MCP request; stop stale refreshes banning clients
- Remove the authless "sole connected phone" fallback and SB_REQUIRE_MCP_AUTH:
  an unauthenticated request no longer reaches whichever phone is online alone.
- Mcp-Session-Id is no longer a credential; the Bearer token is checked on
  every request (MCP auth spec).
- Refresh tokens live 90 days (was 30, equal to the access token, so they
  were always dead when first needed).
- A rejected refresh no longer counts toward the IP ban: a client with an
  expired token was retrying into a self-renewing ban on its own IP.
- 401s carry the RFC 9728 WWW-Authenticate discovery header.
2026-09-27 12:55:16 +02:00

589 lines
22 KiB
Python

"""
Signal Bridge Remote — Main Server Application
Single FastAPI app that serves three roles:
1. OAuth-style auth (register, login, token refresh)
2. MCP endpoint (Streamable HTTP — tool calls from Claude)
3. WebSocket relay hub (persistent phone connections)
Plus rate limiting, IP banning, and the dead man's switch.
"""
from __future__ import annotations
import asyncio
import json
import logging
import os
import sys
import uuid
from contextlib import asynccontextmanager
from fastapi import FastAPI, Request, Response, WebSocket, WebSocketDisconnect
from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import JSONResponse
from . import config
from .auth import (
init_db, create_user, verify_user, create_token, verify_token,
extract_token, ip_tracker, rate_limiter,
get_safety_config, set_safety_config,
)
from .mcp_tools import TOOLS, HANDLERS, current_user_id
from .oauth import init_oauth_db
from .oauth_routes import router as oauth_router, _base_url
from .relay_hub import check_ws_ip_limit, release_ws_ip_slot, get_ip_from_headers
from .session_registry import registry
from .governor import governor
from .safety import dead_man_switch
# ── Logging ─────────────────────────────────────────────────────────────
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s [%(name)s] %(levelname)s: %(message)s",
datefmt="%H:%M:%S",
)
log = logging.getLogger("signal_bridge")
# ── Lifespan ────────────────────────────────────────────────────────────
@asynccontextmanager
async def lifespan(app: FastAPI):
config.validate()
init_db()
init_oauth_db()
await dead_man_switch.start()
log.info(f"Signal Bridge Remote started on {config.HOST}:{config.PORT}")
log.info(f"Registration {'OPEN' if config.REGISTRATION_OPEN else 'CLOSED'}")
yield
await dead_man_switch.stop()
log.info("Signal Bridge Remote shutting down")
app = FastAPI(
title="Signal Bridge Remote",
version="1.0.0",
lifespan=lifespan,
)
app.add_middleware(
CORSMiddleware,
allow_origins=config.CORS_ORIGINS,
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
# Mount OAuth routes (metadata, registration, authorize, token)
app.include_router(oauth_router)
# ════════════════════════════════════════════════════════════════════════
# Auth helpers
# ════════════════════════════════════════════════════════════════════════
def _get_ip(request: Request) -> str:
forwarded = request.headers.get("X-Forwarded-For", "")
if forwarded:
return forwarded.split(",")[0].strip()
return request.client.host if request.client else "unknown"
async def _require_auth(request: Request) -> dict | None:
"""Validate Bearer token. Returns user dict or None."""
token = extract_token(request.headers.get("Authorization", ""))
if not token:
return None
return verify_token(token)
# ════════════════════════════════════════════════════════════════════════
# Auth Endpoints
# ════════════════════════════════════════════════════════════════════════
@app.post("/auth/register")
async def register(request: Request):
"""Register a new user account."""
ip = _get_ip(request)
if await ip_tracker.is_banned(ip):
return JSONResponse({"error": "Temporarily banned"}, status_code=429)
if not await rate_limiter.check(f"auth:{ip}", config.RATE_LIMIT_AUTH):
return JSONResponse({"error": "Too many attempts"}, status_code=429)
if not config.REGISTRATION_OPEN:
return JSONResponse({"error": "Registration is closed"}, status_code=403)
body = await request.json()
username = body.get("username", "").strip()
password = body.get("password", "")
try:
user = await asyncio.to_thread(create_user, username, password)
except ValueError as e:
# Don't count validation errors (short username, weak password) toward IP ban.
# Only actual auth failures (wrong credentials) should inflate the ban counter.
return JSONResponse({"error": str(e)}, status_code=400)
token = create_token(user["user_id"], user["username"])
await ip_tracker.clear_failures(ip)
return {"user_id": user["user_id"], "username": user["username"], "token": token}
@app.post("/auth/login")
async def login(request: Request):
"""Authenticate and receive a JWT."""
ip = _get_ip(request)
if await ip_tracker.is_banned(ip):
return JSONResponse({"error": "Temporarily banned"}, status_code=429)
if not await rate_limiter.check(f"auth:{ip}", config.RATE_LIMIT_AUTH):
return JSONResponse({"error": "Too many attempts"}, status_code=429)
body = await request.json()
username = body.get("username", "")
password = body.get("password", "")
user = await asyncio.to_thread(verify_user, username, password)
if not user:
await ip_tracker.record_failure(ip)
return JSONResponse({"error": "Invalid credentials"}, status_code=401)
token = create_token(user["user_id"], user["username"])
await ip_tracker.clear_failures(ip)
return {"user_id": user["user_id"], "username": user["username"], "token": token}
# ════════════════════════════════════════════════════════════════════════
# MCP Endpoint — Streamable HTTP (JSON-RPC over POST + GET)
#
# Implements the MCP Streamable HTTP transport spec:
# - POST: JSON-RPC requests from client
# - GET: SSE stream for server-to-client notifications (kept open)
# - Mcp-Session-Id header for session tracking
# - Bearer token (OAuth or login JWT) required on every request
# ════════════════════════════════════════════════════════════════════════
# In-memory MCP session tracking (maps session_id → user_id)
_mcp_sessions: dict[str, str] = {}
async def _resolve_mcp_user(request: Request) -> dict | None:
"""
Resolve the user for an MCP request: a valid Bearer token, on every
request, or nobody.
There used to be two more paths here — an Mcp-Session-Id lookup and a
"sole connected phone" fallback for unauthenticated requests. Both are
gone: the fallback handed an anonymous caller whichever single phone was
online, and a session ID must not outlive or replace the token that
opened it (MCP auth spec: authorization on every HTTP request).
"""
return await _require_auth(request)
@app.post("/mcp")
async def mcp_endpoint(request: Request):
"""
MCP Streamable HTTP endpoint (POST).
Accepts JSON-RPC requests, routes tool calls to the authenticated
user's phone via the session registry.
"""
ip = _get_ip(request)
if await ip_tracker.is_banned(ip):
return JSONResponse({"error": "Temporarily banned"}, status_code=429)
if not await rate_limiter.check(f"global:{ip}", config.RATE_LIMIT_GLOBAL):
return JSONResponse({"error": "Rate limit exceeded"}, status_code=429)
# Parse JSON-RPC first (we need to check if it's an initialize request)
try:
body = await request.json()
except Exception:
return _jsonrpc_error(None, -32700, "Parse error: invalid JSON")
method = body.get("method", "")
params = body.get("params", {})
req_id = body.get("id")
# Notifications (no id) — spec requires a bare 202 ACK, not a JSON-RPC error.
if req_id is None:
return Response(status_code=202)
# Resolve user
user = await _resolve_mcp_user(request)
if not user:
# RFC 9728: the challenge header is how MCP clients discover where
# to start the OAuth flow — a bare 401 leaves them stranded.
base = _base_url(request)
return JSONResponse(
{"jsonrpc": "2.0", "error": {"code": -32000, "message": "Authentication required: send a Bearer token (OAuth or login JWT)"}},
status_code=401,
headers={
"WWW-Authenticate": f'Bearer resource_metadata="{base}/.well-known/oauth-protected-resource"'
},
)
# Rate limit per user for commands
if not await rate_limiter.check(
f"cmd:{user['user_id']}", config.RATE_LIMIT_COMMANDS
):
return JSONResponse(
{"jsonrpc": "2.0", "error": {"code": -32000, "message": "Command rate limit exceeded"}},
status_code=429,
)
# Set user context for tool handlers
current_user_id.set(user["user_id"])
# ── Route by method ─────────────────────────────────────────────
if method == "initialize":
# Generate a session ID and bind it to this user
session_id = str(uuid.uuid4())
_mcp_sessions[session_id] = user["user_id"]
log.info(f"MCP session created: {session_id[:8]}... for user {user['user_id']}")
client_ver = (params or {}).get("protocolVersion", "2025-03-26")
supported = {"2024-11-05", "2025-03-26", "2025-06-18"}
result = {
"protocolVersion": client_ver if client_ver in supported else "2025-03-26",
"capabilities": {"tools": {}},
"serverInfo": {"name": "Signal Bridge Remote", "version": "1.0.0"},
}
response = JSONResponse({"jsonrpc": "2.0", "id": req_id, "result": result})
response.headers["Mcp-Session-Id"] = session_id
return response
elif method == "tools/list":
log.info(f"tools/list hit — user={user['user_id']} ua={request.headers.get('user-agent','?')}")
return _jsonrpc_result(req_id, {"tools": TOOLS})
elif method == "tools/call":
tool_name = params.get("name", "")
tool_args = params.get("arguments", {})
handler = HANDLERS.get(tool_name)
if not handler:
return _jsonrpc_error(req_id, -32601, f"Unknown tool: {tool_name}")
try:
result_text = await handler(**tool_args)
return _jsonrpc_result(req_id, {
"content": [{"type": "text", "text": result_text}],
})
except Exception as e:
log.error(f"Tool {tool_name} error: {e}")
return _jsonrpc_result(req_id, {
"content": [{"type": "text", "text": f"Error: {e}"}],
"isError": True,
})
elif method == "ping":
return _jsonrpc_result(req_id, {})
elif method == "resources/list":
return _jsonrpc_result(req_id, {"resources": []})
elif method == "prompts/list":
return _jsonrpc_result(req_id, {"prompts": []})
elif method.startswith("notifications/"):
# MCP notifications (e.g. notifications/initialized) are fire-and-forget.
# Return empty success — no error, no noise.
return _jsonrpc_result(req_id, {})
else:
return _jsonrpc_error(req_id, -32601, f"Unknown method: {method}")
@app.get("/mcp")
async def mcp_sse_endpoint(request: Request):
"""
MCP Streamable HTTP endpoint (GET).
Opens an SSE stream for server-to-client notifications.
We don't currently use server-initiated notifications,
so this just stays open to satisfy the spec.
"""
from starlette.responses import StreamingResponse
async def event_stream():
# Send a keep-alive comment, then hold the connection open
yield ": connected\n\n"
try:
while True:
await asyncio.sleep(30)
yield ": keepalive\n\n"
except asyncio.CancelledError:
pass
return StreamingResponse(
event_stream(),
media_type="text/event-stream",
headers={
"Cache-Control": "no-cache",
"Connection": "keep-alive",
},
)
def _jsonrpc_result(req_id, result):
return JSONResponse({"jsonrpc": "2.0", "id": req_id, "result": result})
def _jsonrpc_error(req_id, code, message):
return JSONResponse(
{"jsonrpc": "2.0", "id": req_id, "error": {"code": code, "message": message}}
)
# ════════════════════════════════════════════════════════════════════════
# WebSocket Relay — Phone connections
# ════════════════════════════════════════════════════════════════════════
@app.websocket("/ws/phone")
async def websocket_phone(websocket: WebSocket):
"""
WebSocket endpoint for phone relay clients.
The phone connects here, authenticates with its JWT,
and maintains a persistent connection for receiving device commands.
"""
await websocket.accept()
await _handle_phone_ws(websocket)
async def _handle_phone_ws(ws: WebSocket):
"""
Full phone WebSocket lifecycle: auth → register → message loop → cleanup.
"""
from .models import CommandAck
ip = get_ip_from_headers(
ws.client.host if ws.client else None,
dict(ws.headers) if ws.headers else None,
)
# IP-level rate limiting
rejection = await check_ws_ip_limit(ip)
if rejection:
await ws.close(4003, rejection)
return
user_id = None
try:
# Wait for auth message
raw = await asyncio.wait_for(ws.receive_text(), timeout=10.0)
msg = json.loads(raw)
if msg.get("type") != "phone_auth" or "token" not in msg:
await ws.close(4001, "First message must be phone_auth")
await ip_tracker.record_failure(ip)
return
user = verify_token(msg["token"])
if not user:
await ws.close(4001, "Invalid token")
await ip_tracker.record_failure(ip)
return
user_id = user["user_id"]
await ip_tracker.clear_failures(ip)
await ws.send_json({
"type": "auth_ok",
"user_id": user_id,
"message": "Connected to Signal Bridge relay",
})
log.info(f"Phone connected: user={user['username']} ip={ip}")
# Create a wrapper that looks like a websockets ServerConnection
wrapper = _FastAPIWSWrapper(ws)
session = await registry.register(user_id, wrapper)
# Load per-user governor config from database
effective_config = _effective_safety_config(user_id)
governor.apply_user_config(user_id, effective_config)
# Request device list (phone also sends proactively, but this is a backup)
log.info(f"Requesting device scan from phone: user={user_id}")
await ws.send_json({"type": "scan"})
# Message loop
while True:
try:
raw = await ws.receive_text()
msg = json.loads(raw)
msg_type = msg.get("type")
if msg_type == "heartbeat_pong":
await registry.update_heartbeat(user_id)
elif msg_type == "command_ack":
ack = CommandAck(
success=msg.get("success", True),
message=msg.get("message", ""),
request_id=msg.get("request_id"),
data=msg.get("data"),
)
if ack.request_id:
session.resolve_ack(ack.request_id, ack)
elif msg_type == "phone_emergency_stop":
# Phone-initiated emergency stop (volume keys, etc.)
# Tell the governor so heat stops accumulating
governor.record_stop(user_id)
log.warning(f"Phone emergency stop: user={user_id}")
elif msg_type == "device_list":
await registry.update_devices(user_id, msg.get("devices", []))
log.info(f"Devices updated: user={user_id}, count={len(msg.get('devices', []))}")
except WebSocketDisconnect:
break
except json.JSONDecodeError:
continue
except asyncio.TimeoutError:
await ws.close(4001, "Auth timeout")
except WebSocketDisconnect:
pass
except Exception as e:
log.error(f"Phone WS error: {e}")
finally:
if user_id:
await registry.unregister(user_id)
governor.remove_user(user_id)
log.info(f"Phone disconnected: user={user_id}")
await release_ws_ip_slot(ip)
class _FastAPIWSWrapper:
"""
Minimal wrapper to make a FastAPI WebSocket look enough like a
websockets ServerConnection for the session registry and safety module.
"""
def __init__(self, ws: WebSocket):
self._ws = ws
async def send(self, data: str):
await self._ws.send_text(data)
async def close(self, code: int = 1000, reason: str = ""):
await self._ws.close(code, reason)
@property
def transport(self):
return self # duck typing for _get_ip fallback
def get_extra_info(self, key):
if key == "peername" and self._ws.client:
return (self._ws.client.host, self._ws.client.port)
return None
# ════════════════════════════════════════════════════════════════════════
# Safety Config (per-user governor settings)
# ════════════════════════════════════════════════════════════════════════
def _effective_safety_config(user_id: str) -> dict:
"""Merge per-user overrides with server defaults."""
defaults = {
"governor_enabled": config.GOVERNOR_ENABLED,
"heat_rate": config.GOVERNOR_HEAT_RATE,
"cool_rate": config.GOVERNOR_COOL_RATE,
"cooldown_threshold": config.GOVERNOR_COOLDOWN_THRESHOLD,
"cooldown_exit": config.GOVERNOR_COOLDOWN_EXIT,
"cooldown_duration": config.GOVERNOR_COOLDOWN_DURATION,
}
overrides = get_safety_config(user_id)
merged = {**defaults, **overrides}
return merged
@app.get("/safety/config")
async def get_safety_config_endpoint(request: Request):
"""Get the effective safety config for the authenticated user."""
token = extract_token(request.headers.get("authorization", ""))
if not token:
return JSONResponse({"error": "Not authenticated"}, status_code=401)
user = verify_token(token)
if not user:
return JSONResponse({"error": "Invalid token"}, status_code=401)
return _effective_safety_config(user["user_id"])
@app.post("/safety/config")
async def set_safety_config_endpoint(request: Request):
"""Update per-user safety config overrides."""
token = extract_token(request.headers.get("authorization", ""))
if not token:
return JSONResponse({"error": "Not authenticated"}, status_code=401)
user = verify_token(token)
if not user:
return JSONResponse({"error": "Invalid token"}, status_code=401)
try:
body = await request.json()
except Exception:
return JSONResponse({"error": "Invalid JSON"}, status_code=400)
overrides = set_safety_config(user["user_id"], body)
effective = _effective_safety_config(user["user_id"])
# Update the live governor with new config
governor.apply_user_config(user["user_id"], effective)
log.info(f"Safety config updated for user {user['user_id']}: {overrides}")
return effective
@app.get("/safety/status")
async def safety_status(request: Request):
"""Get current governor state for the authenticated user."""
token = extract_token(request.headers.get("authorization", ""))
if not token:
return JSONResponse({"error": "Not authenticated"}, status_code=401)
user = verify_token(token)
if not user:
return JSONResponse({"error": "Invalid token"}, status_code=401)
state = governor.get_state(user["user_id"])
state["config"] = _effective_safety_config(user["user_id"])
return state
# ════════════════════════════════════════════════════════════════════════
# Health & Status
# ════════════════════════════════════════════════════════════════════════
@app.get("/health")
async def health():
return {
"status": "ok",
"active_phones": registry.active_count,
"banned_ips": ip_tracker.banned_count,
}
# ════════════════════════════════════════════════════════════════════════
# Init module
# ════════════════════════════════════════════════════════════════════════
@app.get("/")
async def root():
return {
"service": "Signal Bridge Remote",
"version": "1.0.0",
"endpoints": {
"auth": "/auth/register, /auth/login",
"oauth": "/.well-known/oauth-authorization-server, /oauth/register, /oauth/authorize, /oauth/token",
"mcp": "/mcp (POST, JSON-RPC)",
"phone_relay": "/ws/phone (WebSocket)",
"safety": "/safety/config (GET, POST), /safety/status (GET)",
"health": "/health",
},
}