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("""\ + + +
+ + ++ Sign in to connect + $client_name +
+ $error_html + + +Missing client_id parameter
", status_code=400 + ) + + client = await asyncio.to_thread(verify_client, client_id) + if not client: + return HTMLResponse( + "Unknown client_id
", status_code=400 + ) + + # Verify redirect_uri is registered + if redirect_uri and redirect_uri not in client["redirect_uris"]: + return HTMLResponse( + "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("Invalid client
", status_code=400) + + if redirect_uri and redirect_uri not in client["redirect_uris"]: + return HTMLResponse( + "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='{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