From 6a8bc353c5feb6be9272ae4af145b38ebaecc1b0 Mon Sep 17 00:00:00 2001 From: Aletheia Date: Tue, 7 Jul 2026 20:05:45 +0200 Subject: [PATCH] feat: OAuth 2.0 support, server-side safety governor, multi-user auth mode Publishes server work that shipped in the Android edition but never made it to this repo: - Full OAuth 2.0 flow (discovery metadata, dynamic client registration, authorize + token endpoints) so claude.ai remote connectors and the Android app can authenticate per-user instead of relying on the sole-phone fallback. - Safety governor: server-side heat model (intensity x time) with automatic cooldown, per-user overrides via GET/POST /safety/config, and governor state piggybacked on heartbeat pings so relay clients can display it. - SB_REQUIRE_MCP_AUTH env flag for multi-user deployments (disables the unauthenticated sole-phone fallback). - requirements-phone.txt and .env.example documenting the new knobs. Co-Authored-By: Claude Fable 5 --- .env.example | 60 ++++++ requirements-phone.txt | 3 + server/app.py | 103 +++++++++- server/auth.py | 71 ++++++- server/config.py | 11 ++ server/governor.py | 228 +++++++++++++++++++++ server/mcp_tools.py | 52 ++++- server/oauth.py | 439 +++++++++++++++++++++++++++++++++++++++++ server/oauth_routes.py | 432 ++++++++++++++++++++++++++++++++++++++++ server/safety.py | 15 +- 10 files changed, 1398 insertions(+), 16 deletions(-) create mode 100644 .env.example create mode 100644 requirements-phone.txt create mode 100644 server/governor.py create mode 100644 server/oauth.py create mode 100644 server/oauth_routes.py diff --git a/.env.example b/.env.example new file mode 100644 index 0000000..2bd66d7 --- /dev/null +++ b/.env.example @@ -0,0 +1,60 @@ +# Signal Bridge Remote — Environment Configuration +# Copy this to .env and fill in your values. + +# ═══ REQUIRED ═══════════════════════════════════════════════════════════ +# Generate with: python -c "import secrets; print(secrets.token_hex(32))" +SB_SECRET_KEY=your-secret-key-here + +# ═══ Server ═════════════════════════════════════════════════════════════ +SB_HOST=0.0.0.0 +SB_PORT=8420 + +# ═══ Auth ═══════════════════════════════════════════════════════════════ +# Set to "false" to lock out new registrations after your users are set up +SB_REGISTRATION_OPEN=true +# How long login tokens last (hours) +SB_TOKEN_EXPIRY_HOURS=168 +# Set to "true" to require OAuth/token auth on every MCP request (multi-user +# servers). Default "false" keeps the single-user convenience fallback: an +# unauthenticated MCP request is routed to the sole connected phone session. +SB_REQUIRE_MCP_AUTH=false + +# ═══ Rate Limiting (anti-harassment) ════════════════════════════════════ +# Auth endpoint: strict to prevent credential stuffing +SB_RATE_LIMIT_AUTH=5/minute +# Command endpoint: generous for normal use, blocks floods +SB_RATE_LIMIT_COMMANDS=120/minute +# Global per-IP limit +SB_RATE_LIMIT_GLOBAL=300/minute +# Max WebSocket connections per IP +SB_MAX_WS_PER_IP=3 +# Auto-ban after N failed auth attempts (per IP, 1-hour window) +SB_BAN_THRESHOLD=20 +# How long bans last (minutes) +SB_BAN_DURATION_MINUTES=30 + +# ═══ Safety ═════════════════════════════════════════════════════════════ +# How often to ping phones (seconds) +SB_HEARTBEAT_INTERVAL=2.0 +# How long before declaring a phone dead (seconds) +SB_HEARTBEAT_TIMEOUT=6.0 + +# ═══ Governor (session intensity limiter) ═══════════════════════════════ +# Heat accumulates from intensity × time and dissipates when idle; when it +# hits the threshold the governor forces a cooldown. Server-wide defaults — +# each user can override via GET/POST /safety/config. +SB_GOVERNOR_ENABLED=true +# Heat units/second at intensity 1.0 +SB_GOVERNOR_HEAT_RATE=3.0 +# Heat units/second dissipated when idle +SB_GOVERNOR_COOL_RATE=2.0 +# Heat % that triggers cooldown +SB_GOVERNOR_COOLDOWN_ENTER=90.0 +# Heat % at which cooldown may end +SB_GOVERNOR_COOLDOWN_EXIT=30.0 +# Minimum seconds a cooldown lasts +SB_GOVERNOR_COOLDOWN_DURATION=30.0 + +# ═══ Database ═══════════════════════════════════════════════════════════ +# SQLite file path (auto-created) +SB_DB_PATH=./signal_bridge.db diff --git a/requirements-phone.txt b/requirements-phone.txt new file mode 100644 index 0000000..00fbaf7 --- /dev/null +++ b/requirements-phone.txt @@ -0,0 +1,3 @@ +# Signal Bridge Remote — Phone Relay Client Dependencies +buttplug>=1.0.0,<2.0.0 +websockets>=12.0 diff --git a/server/app.py b/server/app.py index f84990b..f42cf91 100644 --- a/server/app.py +++ b/server/app.py @@ -25,10 +25,14 @@ 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 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 ───────────────────────────────────────────────────────────── @@ -47,9 +51,11 @@ log = logging.getLogger("signal_bridge") 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'}") + log.info(f"MCP auth {'REQUIRED' if config.REQUIRE_MCP_AUTH else 'optional (sole-phone fallback enabled)'}") yield await dead_man_switch.stop() log.info("Signal Bridge Remote shutting down") @@ -69,6 +75,9 @@ app.add_middleware( allow_headers=["*"], ) +# Mount OAuth routes (metadata, registration, authorize, token) +app.include_router(oauth_router) + # ════════════════════════════════════════════════════════════════════════ # Auth helpers @@ -180,10 +189,12 @@ async def _resolve_mcp_user(request: Request) -> dict | None: return {"user_id": _mcp_sessions[session_id]} # 3. Fall back to sole active phone session (authless / claude.ai init) - fallback_user_id = await registry.get_sole_user_id() - if fallback_user_id: - log.info(f"MCP request without auth — using active session: {fallback_user_id}") - return {"user_id": fallback_user_id} + # Disabled when SB_REQUIRE_MCP_AUTH=true (multi-user mode). + if not config.REQUIRE_MCP_AUTH: + fallback_user_id = await registry.get_sole_user_id() + if fallback_user_id: + log.info(f"MCP request without auth — using active session: {fallback_user_id}") + return {"user_id": fallback_user_id} return None @@ -395,6 +406,10 @@ async def _handle_phone_ws(ws: WebSocket): 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"}) @@ -417,6 +432,11 @@ async def _handle_phone_ws(ws: WebSocket): ) 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', []))}") @@ -435,6 +455,7 @@ async def _handle_phone_ws(ws: WebSocket): 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) @@ -463,6 +484,78 @@ class _FastAPIWSWrapper: 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 # ════════════════════════════════════════════════════════════════════════ @@ -487,8 +580,10 @@ async def root(): "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", }, } diff --git a/server/auth.py b/server/auth.py index 1980943..b279c5c 100644 --- a/server/auth.py +++ b/server/auth.py @@ -24,7 +24,7 @@ from . import config # ════════════════════════════════════════════════════════════════════════ def init_db(): - """Create user table if it doesn't exist.""" + """Create user and safety_config tables if they don't exist.""" conn = sqlite3.connect(config.DB_PATH) conn.execute(""" CREATE TABLE IF NOT EXISTS users ( @@ -35,10 +35,79 @@ def init_db(): is_active INTEGER DEFAULT 1 ) """) + conn.execute(""" + CREATE TABLE IF NOT EXISTS safety_config ( + user_id TEXT PRIMARY KEY REFERENCES users(id), + governor_enabled INTEGER DEFAULT 1, + heat_rate REAL DEFAULT NULL, + cool_rate REAL DEFAULT NULL, + cooldown_threshold REAL DEFAULT NULL, + cooldown_exit REAL DEFAULT NULL, + cooldown_duration REAL DEFAULT NULL, + updated_at TEXT NOT NULL + ) + """) conn.commit() conn.close() +def get_safety_config(user_id: str) -> dict: + """Get per-user safety config. Returns overrides only (NULLs omitted).""" + conn = _get_conn() + row = conn.execute( + "SELECT * FROM safety_config WHERE user_id = ?", (user_id,) + ).fetchone() + conn.close() + + if not row: + return {} + + result = {} + for key in ("governor_enabled", "heat_rate", "cool_rate", + "cooldown_threshold", "cooldown_exit", "cooldown_duration"): + if row[key] is not None: + result[key] = row[key] + return result + + +def set_safety_config(user_id: str, overrides: dict) -> dict: + """Set per-user safety config overrides. Returns the merged config.""" + allowed_keys = { + "governor_enabled", "heat_rate", "cool_rate", + "cooldown_threshold", "cooldown_exit", "cooldown_duration", + } + filtered = {k: v for k, v in overrides.items() if k in allowed_keys} + + conn = _get_conn() + existing = conn.execute( + "SELECT user_id FROM safety_config WHERE user_id = ?", (user_id,) + ).fetchone() + + now = datetime.now(timezone.utc).isoformat() + + if existing: + # Update existing overrides + sets = ", ".join(f"{k} = ?" for k in filtered) + if sets: + conn.execute( + f"UPDATE safety_config SET {sets}, updated_at = ? WHERE user_id = ?", + (*filtered.values(), now, user_id), + ) + else: + # Insert new row + cols = ", ".join(["user_id", "updated_at"] + list(filtered.keys())) + placeholders = ", ".join(["?"] * (2 + len(filtered))) + conn.execute( + f"INSERT INTO safety_config ({cols}) VALUES ({placeholders})", + (user_id, now, *filtered.values()), + ) + + conn.commit() + conn.close() + + return get_safety_config(user_id) + + def _get_conn(): conn = sqlite3.connect(config.DB_PATH) conn.row_factory = sqlite3.Row diff --git a/server/config.py b/server/config.py index 1f96497..9cc0b03 100644 --- a/server/config.py +++ b/server/config.py @@ -19,6 +19,7 @@ CORS_ORIGINS = os.getenv("SB_CORS_ORIGINS", "*").split(",") # ── Auth ──────────────────────────────────────────────────────────────── TOKEN_EXPIRY_HOURS = int(os.getenv("SB_TOKEN_EXPIRY_HOURS", "168")) # 1 week REGISTRATION_OPEN = os.getenv("SB_REGISTRATION_OPEN", "true").lower() == "true" +REQUIRE_MCP_AUTH = os.getenv("SB_REQUIRE_MCP_AUTH", "false").lower() == "true" # ── Rate Limiting ─────────────────────────────────────────────────────── # Format: "count/period" — e.g. "5/minute", "100/hour" @@ -33,6 +34,16 @@ BAN_DURATION_MINUTES = int(os.getenv("SB_BAN_DURATION_MINUTES", "30")) HEARTBEAT_INTERVAL_S = float(os.getenv("SB_HEARTBEAT_INTERVAL", "2.0")) HEARTBEAT_TIMEOUT_S = float(os.getenv("SB_HEARTBEAT_TIMEOUT", "6.0")) +# ── Governor (session intensity limiter) ─────────────────────────────── +# Heat accumulates based on intensity × time, dissipates when idle. +# Cooldown triggers when heat reaches threshold, exits at the floor. +GOVERNOR_ENABLED = os.getenv("SB_GOVERNOR_ENABLED", "true").lower() == "true" +GOVERNOR_HEAT_RATE = float(os.getenv("SB_GOVERNOR_HEAT_RATE", "3.0")) # heat units/sec at intensity=1.0 +GOVERNOR_COOL_RATE = float(os.getenv("SB_GOVERNOR_COOL_RATE", "2.0")) # heat units/sec dissipation when idle +GOVERNOR_COOLDOWN_THRESHOLD = float(os.getenv("SB_GOVERNOR_COOLDOWN_ENTER", "90.0")) # heat% to trigger cooldown +GOVERNOR_COOLDOWN_EXIT = float(os.getenv("SB_GOVERNOR_COOLDOWN_EXIT", "30.0")) # heat% to exit cooldown +GOVERNOR_COOLDOWN_DURATION = float(os.getenv("SB_GOVERNOR_COOLDOWN_DURATION", "30.0")) # min seconds in cooldown + # ── Database ──────────────────────────────────────────────────────────── DB_PATH = os.getenv("SB_DB_PATH", str(Path(__file__).parent / "signal_bridge.db")) diff --git a/server/governor.py b/server/governor.py new file mode 100644 index 0000000..a834731 --- /dev/null +++ b/server/governor.py @@ -0,0 +1,228 @@ +""" +Signal Bridge Remote — Session Intensity Governor + +Tracks cumulative session intensity ("heat") per user and enforces +cooldown periods when the threshold is reached. + +Heat model: + - Accumulates at: current_intensity × HEAT_RATE per second + - Dissipates at: COOL_RATE per second when intensity = 0 + - Partial dissipation when running at low intensity: + net_rate = (intensity × HEAT_RATE) - COOL_RATE + - Cooldown triggers at COOLDOWN_THRESHOLD (default 90%) + - Cooldown exits when heat falls to COOLDOWN_EXIT (default 30%) + AND at least COOLDOWN_DURATION seconds have passed + +The governor state is piggybacked on heartbeat pings so the phone +can display heat level and cooldown countdown in real time. + +This is a soft safety layer — the hard safety is the dead man's switch. +The governor exists to pace sessions, not to prevent hardware damage. +""" +from __future__ import annotations +import logging +import time +from dataclasses import dataclass, field + +from . import config + +log = logging.getLogger("signal_bridge.governor") + + +@dataclass +class GovernorConfig: + """Per-user governor config (overrides server defaults).""" + heat_rate: float = config.GOVERNOR_HEAT_RATE + cool_rate: float = config.GOVERNOR_COOL_RATE + cooldown_threshold: float = config.GOVERNOR_COOLDOWN_THRESHOLD + cooldown_exit: float = config.GOVERNOR_COOLDOWN_EXIT + cooldown_duration: float = config.GOVERNOR_COOLDOWN_DURATION + enabled: bool = config.GOVERNOR_ENABLED + + +@dataclass +class GovernorState: + """Per-user heat tracking state.""" + heat: float = 0.0 # 0..100 + current_intensity: float = 0.0 # last known intensity (0..1) + in_cooldown: bool = False + cooldown_entered_at: float = 0.0 + cooldown_count: int = 0 # total cooldowns this session + last_tick: float = field(default_factory=time.time) + cfg: GovernorConfig = field(default_factory=GovernorConfig) + + def tick(self) -> None: + """Update heat based on elapsed time since last tick.""" + now = time.time() + dt = now - self.last_tick + self.last_tick = now + + if dt <= 0 or dt > 10: + # Sanity: skip huge jumps (e.g., system clock change) + return + + if self.in_cooldown: + # During cooldown: always dissipate, intensity is forced to 0 + self.heat -= self.cfg.cool_rate * dt + self.heat = max(0.0, self.heat) + + # Check cooldown exit conditions + elapsed = now - self.cooldown_entered_at + if (self.heat <= self.cfg.cooldown_exit + and elapsed >= self.cfg.cooldown_duration): + self.in_cooldown = False + log.info( + f"Cooldown ended: heat={self.heat:.1f}% " + f"after {elapsed:.0f}s" + ) + else: + # Normal operation: accumulate or dissipate + net = (self.current_intensity * self.cfg.heat_rate + - self.cfg.cool_rate) + # Only dissipate if there's actual heat to lose + if net < 0 and self.heat <= 0: + return + self.heat += net * dt + self.heat = max(0.0, min(100.0, self.heat)) + + # Check cooldown trigger + if self.heat >= self.cfg.cooldown_threshold: + self.in_cooldown = True + self.cooldown_entered_at = now + self.cooldown_count += 1 + self.current_intensity = 0.0 + log.warning( + f"Cooldown triggered: heat={self.heat:.1f}% " + f"(cooldown #{self.cooldown_count})" + ) + + def record_command(self, intensity: float) -> None: + """Record that a command was sent at a given intensity.""" + self.current_intensity = max(0.0, min(1.0, intensity)) + + def record_stop(self) -> None: + """Record that devices were stopped.""" + self.current_intensity = 0.0 + + @property + def cooldown_remaining(self) -> int: + """Seconds remaining in cooldown (0 if not in cooldown).""" + if not self.in_cooldown: + return 0 + elapsed = time.time() - self.cooldown_entered_at + # Time-based minimum + time_remaining = max(0, self.cfg.cooldown_duration - elapsed) + # Heat-based: estimate time to reach exit threshold + if self.cfg.cool_rate > 0: + heat_remaining = max( + 0, + (self.heat - self.cfg.cooldown_exit) + / self.cfg.cool_rate, + ) + else: + heat_remaining = 0 + return int(max(time_remaining, heat_remaining)) + + @property + def predicted_seconds(self) -> int | None: + """ + At current intensity, how many seconds until cooldown triggers? + Returns None if intensity is 0 or heat is dissipating. + """ + if self.in_cooldown or self.current_intensity <= 0: + return None + net = (self.current_intensity * self.cfg.heat_rate + - self.cfg.cool_rate) + if net <= 0: + return None # heat is stable or dissipating + remaining_heat = self.cfg.cooldown_threshold - self.heat + if remaining_heat <= 0: + return 0 + return int(remaining_heat / net) + + def to_dict(self) -> dict: + """Serialize for piggybacking on heartbeat pings.""" + return { + "heat_pct": round(self.heat, 1), + "in_cooldown": self.in_cooldown, + "cooldown_remaining": self.cooldown_remaining, + "cooldown_count": self.cooldown_count, + "predicted_seconds": self.predicted_seconds, + } + + +class Governor: + """ + Central governor managing per-user heat state. + + Usage: + - Call tick(user_id) on every heartbeat to update heat + - Call check(user_id) before sending commands — returns (allowed, reason) + - Call record_command(user_id, intensity) after sending a command + - Call record_stop(user_id) when devices are stopped + - Call get_state(user_id) to get current state for heartbeat piggyback + """ + + def __init__(self): + self._states: dict[str, GovernorState] = {} + + def _get(self, user_id: str) -> GovernorState: + if user_id not in self._states: + self._states[user_id] = GovernorState() + return self._states[user_id] + + def tick(self, user_id: str) -> None: + """Advance the heat model. Call on every heartbeat.""" + self._get(user_id).tick() + + def check(self, user_id: str) -> tuple[bool, str]: + """ + Check if a command is allowed. + Returns (allowed: bool, reason: str). + """ + state = self._get(user_id) + if not state.cfg.enabled: + return True, "" + + if state.in_cooldown: + remaining = state.cooldown_remaining + return False, ( + f"Cooldown active ({remaining}s remaining). " + f"Session heat reached {state.cfg.cooldown_threshold:.0f}%. " + f"Wait for cooldown to complete before sending more commands." + ) + return True, "" + + def apply_user_config(self, user_id: str, effective: dict) -> None: + """Apply per-user config overrides from the database.""" + state = self._get(user_id) + state.cfg = GovernorConfig( + enabled=bool(effective.get("governor_enabled", True)), + heat_rate=float(effective.get("heat_rate", config.GOVERNOR_HEAT_RATE)), + cool_rate=float(effective.get("cool_rate", config.GOVERNOR_COOL_RATE)), + cooldown_threshold=float(effective.get("cooldown_threshold", config.GOVERNOR_COOLDOWN_THRESHOLD)), + cooldown_exit=float(effective.get("cooldown_exit", config.GOVERNOR_COOLDOWN_EXIT)), + cooldown_duration=float(effective.get("cooldown_duration", config.GOVERNOR_COOLDOWN_DURATION)), + ) + log.info(f"Applied user config for {user_id}: heat_rate={state.cfg.heat_rate}, " + f"cool_rate={state.cfg.cool_rate}, threshold={state.cfg.cooldown_threshold}") + + def record_command(self, user_id: str, intensity: float) -> None: + """Record that a command was dispatched.""" + self._get(user_id).record_command(intensity) + + def record_stop(self, user_id: str) -> None: + """Record that devices were stopped.""" + self._get(user_id).record_stop() + + def get_state(self, user_id: str) -> dict: + """Get governor state dict for heartbeat piggyback.""" + return self._get(user_id).to_dict() + + def remove_user(self, user_id: str) -> None: + """Clean up when a user disconnects.""" + self._states.pop(user_id, None) + + +# Singleton +governor = Governor() diff --git a/server/mcp_tools.py b/server/mcp_tools.py index 28bdcd1..605674f 100644 --- a/server/mcp_tools.py +++ b/server/mcp_tools.py @@ -24,6 +24,7 @@ from .models import ( DeviceCommand, PatternCommand, StopCommand, ScanCommand, ReadSensorCommand, CommandAck, ) +from .governor import governor from .session_registry import registry # Set by auth middleware before each MCP request @@ -66,10 +67,30 @@ def _register_tool(name: str, description: str, params: dict, required: list[str # Helper # ════════════════════════════════════════════════════════════════════════ -async def _send(command: dict) -> str: - """Route a command to the current user's phone and return result text.""" +async def _send(command: dict, intensity: float = 0.0) -> str: + """ + Route a command to the current user's phone and return result text. + + If intensity > 0, the governor checks if the command is allowed + and records the intensity for heat tracking. + """ user_id = current_user_id.get() + + # Governor check (skip for stop commands and scans) + cmd_type = command.get("type", "") + if cmd_type not in ("stop", "scan") and intensity > 0: + allowed, reason = governor.check(user_id) + if not allowed: + return f"Blocked by governor: {reason}" + ack = await registry.send_to_user(user_id, command) + + # Record intensity for heat tracking + if ack.success and intensity > 0: + governor.record_command(user_id, intensity) + elif ack.success and cmd_type == "stop": + governor.record_stop(user_id) + if ack.success: return ack.message or "OK" else: @@ -126,6 +147,17 @@ async def list_devices(**kwargs) -> str: + (f" | floor: {floor}" if floor > 0 else "") + (f" | {notes}" if notes else "") ) + + # Append governor state so Claude knows the session budget + gov = governor.get_state(user_id) + heat = gov["heat_pct"] + if gov["in_cooldown"]: + lines.append(f"\n⚠ Governor: COOLDOWN ({gov['cooldown_remaining']}s remaining)") + elif heat > 0: + lines.append(f"\nGovernor: {heat:.0f}% heat" + + (f" (~{gov['predicted_seconds']}s to cooldown)" + if gov["predicted_seconds"] is not None else "")) + return "\n".join(lines) @@ -166,13 +198,14 @@ def _make_output_handler(output_type: OutputType): async def handler( device: str = "all", intensity: float = 0.5, duration: float = 0, **kw ) -> str: + clamped = max(0.0, min(1.0, float(intensity))) cmd = DeviceCommand( action=output_type, device=device, - intensity=max(0.0, min(1.0, intensity)), - duration=max(0.0, duration), + intensity=clamped, + duration=max(0.0, float(duration)), ) - return await _send(cmd.model_dump()) + return await _send(cmd.model_dump(), intensity=clamped) return handler @@ -290,15 +323,16 @@ def _make_pattern_handler(pattern_name: str): hold_seconds: float = 0, **kw, ) -> str: + clamped = max(0.0, min(1.0, float(intensity))) cmd = PatternCommand( pattern=pattern_name, output_type=OutputType(output_type), device=device, - intensity=max(0.0, min(1.0, intensity)), - duration=max(0.0, duration), - hold_seconds=max(0.0, hold_seconds), + intensity=clamped, + duration=max(0.0, float(duration)), + hold_seconds=max(0.0, float(hold_seconds)), ) - return await _send(cmd.model_dump()) + return await _send(cmd.model_dump(), intensity=clamped) return handler diff --git a/server/oauth.py b/server/oauth.py new file mode 100644 index 0000000..63c0978 --- /dev/null +++ b/server/oauth.py @@ -0,0 +1,439 @@ +""" +Signal Bridge Remote — OAuth 2.0 Authorization Server + +Implements the OAuth flow required by MCP Streamable HTTP transport so +that claude.ai custom connectors can authenticate users without a +pre-shared Bearer token. + +Endpoints: + GET /.well-known/oauth-authorization-server RFC 8414 metadata + POST /oauth/register RFC 7591 dynamic client registration + GET /oauth/authorize Authorization endpoint (login page) + POST /oauth/authorize Login form submission + POST /oauth/token Token endpoint (code + refresh) + +Supports PKCE (RFC 7636) with S256 method. +""" +from __future__ import annotations + +import asyncio +import hashlib +import base64 +import json +import logging +import secrets +import sqlite3 +import string +import time +import urllib.parse +from datetime import datetime, timezone + +import bcrypt + +from . import config +from .auth import verify_user, create_token, verify_token, _get_conn + +log = logging.getLogger("signal_bridge.oauth") + +# ════════════════════════════════════════════════════════════════════════ +# Constants +# ════════════════════════════════════════════════════════════════════════ + +AUTH_CODE_EXPIRY_S = 300 # 5 minutes — per OAuth spec recommendation +REFRESH_TOKEN_EXPIRY_S = 86400 * 30 # 30 days + + +# ════════════════════════════════════════════════════════════════════════ +# Database — OAuth tables +# ════════════════════════════════════════════════════════════════════════ + +def init_oauth_db(): + """Create OAuth tables if they don't exist. Called from app lifespan.""" + conn = sqlite3.connect(config.DB_PATH) + conn.execute(""" + CREATE TABLE IF NOT EXISTS oauth_clients ( + client_id TEXT PRIMARY KEY, + client_secret_hash TEXT, + client_name TEXT NOT NULL, + redirect_uris TEXT NOT NULL, + created_at TEXT NOT NULL + ) + """) + conn.execute(""" + CREATE TABLE IF NOT EXISTS oauth_codes ( + code TEXT PRIMARY KEY, + client_id TEXT NOT NULL, + user_id TEXT NOT NULL, + redirect_uri TEXT NOT NULL, + code_challenge TEXT, + code_challenge_method TEXT, + expires_at REAL NOT NULL, + used INTEGER DEFAULT 0 + ) + """) + conn.execute(""" + CREATE TABLE IF NOT EXISTS oauth_refresh_tokens ( + token TEXT PRIMARY KEY, + client_id TEXT NOT NULL, + user_id TEXT NOT NULL, + expires_at REAL NOT NULL, + revoked INTEGER DEFAULT 0 + ) + """) + conn.commit() + conn.close() + log.info("OAuth database tables ready") + + +# ════════════════════════════════════════════════════════════════════════ +# Client Registration (RFC 7591) +# ════════════════════════════════════════════════════════════════════════ + +def register_client(client_name: str, redirect_uris: list[str]) -> dict: + """ + Register a new OAuth client. Returns client_id and client_secret. + The secret is returned in plaintext exactly once; we store the hash. + """ + client_id = secrets.token_urlsafe(24) + client_secret = secrets.token_urlsafe(48) + client_secret_hash = bcrypt.hashpw( + client_secret.encode(), bcrypt.gensalt() + ).decode() + + conn = _get_conn() + conn.execute( + "INSERT INTO oauth_clients (client_id, client_secret_hash, client_name, redirect_uris, created_at) " + "VALUES (?, ?, ?, ?, ?)", + ( + client_id, + client_secret_hash, + client_name, + json.dumps(redirect_uris), + datetime.now(timezone.utc).isoformat(), + ), + ) + conn.commit() + conn.close() + + log.info(f"OAuth client registered: {client_name} ({client_id[:12]}...)") + return { + "client_id": client_id, + "client_secret": client_secret, + "client_name": client_name, + "redirect_uris": redirect_uris, + } + + +def verify_client(client_id: str, client_secret: str | None = None) -> dict | None: + """ + Look up a client by ID. If client_secret is provided, verify it. + Returns client dict or None. + """ + conn = _get_conn() + row = conn.execute( + "SELECT * FROM oauth_clients WHERE client_id = ?", (client_id,) + ).fetchone() + conn.close() + + if not row: + return None + + if client_secret is not None: + if not row["client_secret_hash"]: + return None + if not bcrypt.checkpw(client_secret.encode(), row["client_secret_hash"].encode()): + return None + + return { + "client_id": row["client_id"], + "client_name": row["client_name"], + "redirect_uris": json.loads(row["redirect_uris"]), + } + + +# ════════════════════════════════════════════════════════════════════════ +# Authorization Codes +# ════════════════════════════════════════════════════════════════════════ + +def create_auth_code( + client_id: str, + user_id: str, + redirect_uri: str, + code_challenge: str | None = None, + code_challenge_method: str | None = None, +) -> str: + """Issue a short-lived authorization code.""" + code = secrets.token_urlsafe(48) + + conn = _get_conn() + conn.execute( + "INSERT INTO oauth_codes " + "(code, client_id, user_id, redirect_uri, code_challenge, code_challenge_method, expires_at) " + "VALUES (?, ?, ?, ?, ?, ?, ?)", + ( + code, + client_id, + user_id, + redirect_uri, + code_challenge, + code_challenge_method, + time.time() + AUTH_CODE_EXPIRY_S, + ), + ) + conn.commit() + conn.close() + return code + + +def consume_auth_code(code: str, client_id: str, code_verifier: str | None = None) -> dict | None: + """ + Validate and consume an authorization code. Returns user info or None. + Each code can only be used once. + """ + conn = _get_conn() + row = conn.execute( + "SELECT * FROM oauth_codes WHERE code = ? AND client_id = ? AND used = 0", + (code, client_id), + ).fetchone() + + if not row: + conn.close() + return None + + # Check expiry + if time.time() > row["expires_at"]: + conn.execute("DELETE FROM oauth_codes WHERE code = ?", (code,)) + conn.commit() + conn.close() + return None + + # PKCE verification + if row["code_challenge"]: + if not code_verifier: + conn.close() + return None + + method = row["code_challenge_method"] or "plain" + if method == "S256": + digest = hashlib.sha256(code_verifier.encode("ascii")).digest() + computed = base64.urlsafe_b64encode(digest).rstrip(b"=").decode("ascii") + else: + computed = code_verifier + + if computed != row["code_challenge"]: + conn.close() + return None + + # Mark as used + conn.execute("UPDATE oauth_codes SET used = 1 WHERE code = ?", (code,)) + conn.commit() + conn.close() + + return {"user_id": row["user_id"], "redirect_uri": row["redirect_uri"]} + + +# ════════════════════════════════════════════════════════════════════════ +# Refresh Tokens +# ════════════════════════════════════════════════════════════════════════ + +def create_refresh_token(client_id: str, user_id: str) -> str: + """Issue a long-lived refresh token.""" + token = secrets.token_urlsafe(64) + + conn = _get_conn() + conn.execute( + "INSERT INTO oauth_refresh_tokens (token, client_id, user_id, expires_at) " + "VALUES (?, ?, ?, ?)", + (token, client_id, user_id, time.time() + REFRESH_TOKEN_EXPIRY_S), + ) + conn.commit() + conn.close() + return token + + +def consume_refresh_token(token: str, client_id: str) -> dict | None: + """ + Validate a refresh token. Revokes the old one (rotation). + Returns user info or None. + """ + conn = _get_conn() + row = conn.execute( + "SELECT * FROM oauth_refresh_tokens " + "WHERE token = ? AND client_id = ? AND revoked = 0", + (token, client_id), + ).fetchone() + + if not row: + conn.close() + return None + + if time.time() > row["expires_at"]: + conn.execute("DELETE FROM oauth_refresh_tokens WHERE token = ?", (token,)) + conn.commit() + conn.close() + return None + + # Revoke old token (rotation) + conn.execute("UPDATE oauth_refresh_tokens SET revoked = 1 WHERE token = ?", (token,)) + conn.commit() + conn.close() + + return {"user_id": row["user_id"]} + + +def cleanup_expired(): + """Purge expired codes and revoked/expired refresh tokens.""" + now = time.time() + conn = _get_conn() + conn.execute("DELETE FROM oauth_codes WHERE expires_at < ? OR used = 1", (now,)) + conn.execute( + "DELETE FROM oauth_refresh_tokens WHERE expires_at < ? OR revoked = 1", + (now,), + ) + conn.commit() + conn.close() + + +# ════════════════════════════════════════════════════════════════════════ +# PKCE Helpers +# ════════════════════════════════════════════════════════════════════════ + +def _verify_pkce(challenge: str, method: str, verifier: str) -> bool: + """Verify a PKCE code_verifier against the stored challenge.""" + if method == "S256": + digest = hashlib.sha256(verifier.encode("ascii")).digest() + computed = base64.urlsafe_b64encode(digest).rstrip(b"=").decode("ascii") + else: + computed = verifier + return computed == challenge + + +# ════════════════════════════════════════════════════════════════════════ +# HTML Login Page +# ════════════════════════════════════════════════════════════════════════ + +_LOGIN_PAGE_TEMPLATE = string.Template("""\ + + + + + +Signal Bridge — Sign In + + + +
+

Signal Bridge

+

+ Sign in to connect + $client_name +

+ $error_html +
+ + + + + + + + + + + +
+ +
+ + +""") + + +def render_login_page(**kwargs) -> str: + """Render the login page with safe $-substitution (no CSS brace conflicts).""" + return _LOGIN_PAGE_TEMPLATE.safe_substitute(**kwargs) diff --git a/server/oauth_routes.py b/server/oauth_routes.py new file mode 100644 index 0000000..db87aa8 --- /dev/null +++ b/server/oauth_routes.py @@ -0,0 +1,432 @@ +""" +Signal Bridge Remote — OAuth 2.0 Route Handlers + +FastAPI routes for the OAuth authorization server. +Mounted in app.py via include_router(). +""" +from __future__ import annotations + +import asyncio +import json +import logging +import urllib.parse + +from fastapi import APIRouter, Request +from fastapi.responses import HTMLResponse, JSONResponse, RedirectResponse + +from . import config +from .auth import verify_user, create_token, ip_tracker, rate_limiter +from .oauth import ( + init_oauth_db, + register_client, + verify_client, + create_auth_code, + consume_auth_code, + create_refresh_token, + consume_refresh_token, + cleanup_expired, + render_login_page, +) + +log = logging.getLogger("signal_bridge.oauth") + +router = APIRouter() + + +# ════════════════════════════════════════════════════════════════════════ +# Helper +# ════════════════════════════════════════════════════════════════════════ + +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" + + +def _base_url(request: Request) -> str: + """Derive the external base URL from the request.""" + # Respect X-Forwarded-Proto / X-Forwarded-Host if behind a reverse proxy + proto = request.headers.get("X-Forwarded-Proto", request.url.scheme) + host = request.headers.get("X-Forwarded-Host", request.headers.get("Host", request.url.netloc)) + return f"{proto}://{host}" + + +# ════════════════════════════════════════════════════════════════════════ +# RFC 8414 — OAuth Authorization Server Metadata +# ════════════════════════════════════════════════════════════════════════ + +@router.get("/.well-known/oauth-authorization-server") +async def oauth_metadata(request: Request): + """ + Discovery endpoint. MCP clients fetch this to learn where to + authorize, exchange tokens, and register. + """ + base = _base_url(request) + return JSONResponse({ + "issuer": base, + "authorization_endpoint": f"{base}/oauth/authorize", + "token_endpoint": f"{base}/oauth/token", + "registration_endpoint": f"{base}/oauth/register", + "response_types_supported": ["code"], + "grant_types_supported": ["authorization_code", "refresh_token"], + "code_challenge_methods_supported": ["S256", "plain"], + "token_endpoint_auth_methods_supported": ["client_secret_post"], + "scopes_supported": ["signal_bridge"], + }) + + +# ════════════════════════════════════════════════════════════════════════ +# RFC 7591 — Dynamic Client Registration +# ════════════════════════════════════════════════════════════════════════ + +@router.post("/oauth/register") +async def oauth_register_client(request: Request): + """ + Dynamic client registration. MCP clients call this once to obtain + a client_id and client_secret before starting the auth flow. + """ + 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"oauth_reg:{ip}", "10/hour"): + return JSONResponse({"error": "rate_limit_exceeded"}, status_code=429) + + try: + body = await request.json() + except Exception: + return JSONResponse({"error": "invalid_request"}, status_code=400) + + client_name = body.get("client_name", "Unknown MCP Client") + redirect_uris = body.get("redirect_uris", []) + + if not redirect_uris or not isinstance(redirect_uris, list): + return JSONResponse( + {"error": "invalid_client_metadata", + "error_description": "redirect_uris is required and must be a non-empty array"}, + status_code=400, + ) + + # Validate redirect URIs (must be valid URLs) + for uri in redirect_uris: + parsed = urllib.parse.urlparse(uri) + if not parsed.scheme or not parsed.netloc: + # Allow localhost without scheme validation for dev + if "localhost" not in uri and "127.0.0.1" not in uri: + return JSONResponse( + {"error": "invalid_redirect_uri", + "error_description": f"Invalid redirect_uri: {uri}"}, + status_code=400, + ) + + result = await asyncio.to_thread(register_client, client_name, redirect_uris) + + # RFC 7591 response format + return JSONResponse({ + "client_id": result["client_id"], + "client_secret": result["client_secret"], + "client_name": result["client_name"], + "redirect_uris": result["redirect_uris"], + "grant_types": ["authorization_code", "refresh_token"], + "response_types": ["code"], + "token_endpoint_auth_method": "client_secret_post", + }, status_code=201) + + +# ════════════════════════════════════════════════════════════════════════ +# Authorization Endpoint +# ════════════════════════════════════════════════════════════════════════ + +@router.get("/oauth/authorize") +async def oauth_authorize_get(request: Request): + """ + Authorization endpoint (GET). Renders the login page. + Query params: client_id, redirect_uri, response_type, state, + code_challenge, code_challenge_method + """ + params = request.query_params + client_id = params.get("client_id", "") + redirect_uri = params.get("redirect_uri", "") + response_type = params.get("response_type", "") + state = params.get("state", "") + code_challenge = params.get("code_challenge", "") + code_challenge_method = params.get("code_challenge_method", "") + + # Validate + if response_type != "code": + return _authorize_error( + redirect_uri, state, "unsupported_response_type", + "Only response_type=code is supported" + ) + + if not client_id: + return HTMLResponse( + "

Error

Missing client_id parameter

", status_code=400 + ) + + client = await asyncio.to_thread(verify_client, client_id) + if not client: + return HTMLResponse( + "

Error

Unknown client_id

", status_code=400 + ) + + # Verify redirect_uri is registered + if redirect_uri and redirect_uri not in client["redirect_uris"]: + return HTMLResponse( + "

Error

redirect_uri not registered for this client

", + status_code=400, + ) + + # Use first registered URI if none specified + if not redirect_uri: + redirect_uri = client["redirect_uris"][0] + + # Render login page + html = render_login_page( + client_name=_escape_html(client["client_name"]), + client_id=_escape_html(client_id), + redirect_uri=_escape_html(redirect_uri), + state=_escape_html(state), + code_challenge=_escape_html(code_challenge), + code_challenge_method=_escape_html(code_challenge_method), + error_html="", + ) + return HTMLResponse(html) + + +@router.post("/oauth/authorize") +async def oauth_authorize_post(request: Request): + """ + Authorization endpoint (POST). Handles login form submission. + On success, redirects to redirect_uri with authorization code. + """ + ip = _get_ip(request) + + if await ip_tracker.is_banned(ip): + return HTMLResponse("

Temporarily banned

", status_code=429) + + if not await rate_limiter.check(f"auth:{ip}", config.RATE_LIMIT_AUTH): + return HTMLResponse("

Too many attempts

", status_code=429) + + form = await request.form() + username = form.get("username", "") + password = form.get("password", "") + client_id = form.get("client_id", "") + redirect_uri = form.get("redirect_uri", "") + state = form.get("state", "") + code_challenge = form.get("code_challenge", "") + code_challenge_method = form.get("code_challenge_method", "") + + # Verify client + client = await asyncio.to_thread(verify_client, client_id) + if not client: + return HTMLResponse("

Error

Invalid client

", status_code=400) + + if redirect_uri and redirect_uri not in client["redirect_uris"]: + return HTMLResponse( + "

Error

Invalid redirect_uri

", status_code=400 + ) + if not redirect_uri: + redirect_uri = client["redirect_uris"][0] + + # Verify credentials + user = await asyncio.to_thread(verify_user, username, password) + if not user: + await ip_tracker.record_failure(ip) + html = render_login_page( + client_name=_escape_html(client["client_name"]), + client_id=_escape_html(client_id), + redirect_uri=_escape_html(redirect_uri), + state=_escape_html(state), + code_challenge=_escape_html(code_challenge), + code_challenge_method=_escape_html(code_challenge_method), + error_html='
Invalid username or password.
', + ) + return HTMLResponse(html) + + await ip_tracker.clear_failures(ip) + + # Issue authorization code + code = await asyncio.to_thread( + create_auth_code, + client_id, + user["user_id"], + redirect_uri, + code_challenge or None, + code_challenge_method or None, + ) + + # Redirect back to client + params = {"code": code} + if state: + params["state"] = state + + separator = "&" if "?" in redirect_uri else "?" + target = redirect_uri + separator + urllib.parse.urlencode(params) + + log.info(f"OAuth code issued for user={user['username']} client={client_id[:12]}...") + return RedirectResponse(target, status_code=302) + + +# ════════════════════════════════════════════════════════════════════════ +# Token Endpoint +# ════════════════════════════════════════════════════════════════════════ + +@router.post("/oauth/token") +async def oauth_token(request: Request): + """ + Token endpoint. Supports: + - grant_type=authorization_code (exchange code for tokens) + - grant_type=refresh_token (rotate refresh token) + + Client authenticates via client_secret_post (credentials in body). + """ + 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"oauth_token:{ip}", "30/minute"): + return JSONResponse({"error": "rate_limit_exceeded"}, status_code=429) + + # Accept both form-encoded and JSON + content_type = request.headers.get("content-type", "") + if "application/x-www-form-urlencoded" in content_type: + form = await request.form() + body = dict(form) + else: + try: + body = await request.json() + except Exception: + return JSONResponse({"error": "invalid_request"}, status_code=400) + + grant_type = body.get("grant_type", "") + client_id = body.get("client_id", "") + client_secret = body.get("client_secret", "") + + # Verify client credentials + client = await asyncio.to_thread(verify_client, client_id, client_secret) + if not client: + await ip_tracker.record_failure(ip) + return JSONResponse({"error": "invalid_client"}, status_code=401) + + if grant_type == "authorization_code": + return await _handle_authorization_code(body, client_id, ip) + elif grant_type == "refresh_token": + return await _handle_refresh_token(body, client_id, ip) + else: + return JSONResponse({"error": "unsupported_grant_type"}, status_code=400) + + +async def _handle_authorization_code(body: dict, client_id: str, ip: str): + """Exchange an authorization code for access + refresh tokens.""" + code = body.get("code", "") + code_verifier = body.get("code_verifier") + + if not code: + return JSONResponse( + {"error": "invalid_request", "error_description": "Missing code"}, + status_code=400, + ) + + result = await asyncio.to_thread(consume_auth_code, code, client_id, code_verifier) + if not result: + await ip_tracker.record_failure(ip) + return JSONResponse({"error": "invalid_grant"}, status_code=400) + + # Look up user to get username for token + from .auth import _get_conn as get_conn + conn = get_conn() + user_row = conn.execute( + "SELECT username FROM users WHERE id = ?", (result["user_id"],) + ).fetchone() + conn.close() + + if not user_row: + return JSONResponse({"error": "invalid_grant"}, status_code=400) + + # Issue tokens — reuse the existing JWT infrastructure + access_token = create_token(result["user_id"], user_row["username"]) + refresh_token = await asyncio.to_thread( + create_refresh_token, client_id, result["user_id"] + ) + + log.info(f"OAuth tokens issued for user={result['user_id']} via auth code") + + return JSONResponse({ + "access_token": access_token, + "token_type": "Bearer", + "expires_in": config.TOKEN_EXPIRY_HOURS * 3600, + "refresh_token": refresh_token, + }) + + +async def _handle_refresh_token(body: dict, client_id: str, ip: str): + """Exchange a refresh token for a new access + refresh token pair.""" + token = body.get("refresh_token", "") + if not token: + return JSONResponse( + {"error": "invalid_request", "error_description": "Missing refresh_token"}, + status_code=400, + ) + + result = await asyncio.to_thread(consume_refresh_token, token, client_id) + if not result: + await ip_tracker.record_failure(ip) + return JSONResponse({"error": "invalid_grant"}, status_code=400) + + # Look up user + from .auth import _get_conn as get_conn + conn = get_conn() + user_row = conn.execute( + "SELECT username FROM users WHERE id = ?", (result["user_id"],) + ).fetchone() + conn.close() + + if not user_row: + return JSONResponse({"error": "invalid_grant"}, status_code=400) + + access_token = create_token(result["user_id"], user_row["username"]) + new_refresh = await asyncio.to_thread( + create_refresh_token, client_id, result["user_id"] + ) + + log.info(f"OAuth tokens refreshed for user={result['user_id']}") + + return JSONResponse({ + "access_token": access_token, + "token_type": "Bearer", + "expires_in": config.TOKEN_EXPIRY_HOURS * 3600, + "refresh_token": new_refresh, + }) + + +# ════════════════════════════════════════════════════════════════════════ +# Error Helpers +# ════════════════════════════════════════════════════════════════════════ + +def _authorize_error(redirect_uri: str, state: str, error: str, description: str): + """Redirect back with an OAuth error if we have a valid redirect_uri.""" + if not redirect_uri: + return HTMLResponse( + f"

Error

{description}

", status_code=400 + ) + params = {"error": error, "error_description": description} + if state: + params["state"] = state + separator = "&" if "?" in redirect_uri else "?" + target = redirect_uri + separator + urllib.parse.urlencode(params) + return RedirectResponse(target, status_code=302) + + +def _escape_html(s: str) -> str: + """Minimal HTML escaping for template interpolation.""" + return ( + s.replace("&", "&") + .replace("<", "<") + .replace(">", ">") + .replace('"', """) + .replace("'", "'") + ) diff --git a/server/safety.py b/server/safety.py index 2cd9140..b9761c7 100644 --- a/server/safety.py +++ b/server/safety.py @@ -14,6 +14,7 @@ import logging import time from . import config +from .governor import governor from .session_registry import registry log = logging.getLogger("signal_bridge.safety") @@ -70,8 +71,15 @@ class DeadManSwitch: sessions = await registry.get_all_sessions() for user_id, session in sessions.items(): - # Send ping - ping = {"type": "heartbeat_ping", "timestamp": now} + # Tick the governor (advances heat model) + governor.tick(user_id) + + # Send ping with governor state piggybacked + ping = { + "type": "heartbeat_ping", + "timestamp": now, + **governor.get_state(user_id), + } try: await session.websocket.send(json.dumps(ping)) except Exception: @@ -92,6 +100,8 @@ class DeadManSwitch: async def _emergency_stop(self, user_id: str, session): """Send stop-all and disconnect the session.""" log.critical(f"EMERGENCY STOP for user {user_id} — all devices halted") + governor.record_stop(user_id) + try: stop_cmd = {"type": "stop", "device": "all", "emergency": True} await session.websocket.send(json.dumps(stop_cmd)) @@ -104,6 +114,7 @@ class DeadManSwitch: pass await registry.unregister(user_id) + governor.remove_user(user_id) # Singleton