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 <noreply@anthropic.com>
This commit is contained in:
Aletheia
2026-07-07 20:05:45 +02:00
parent f8a9245f90
commit 6a8bc353c5
10 changed files with 1398 additions and 16 deletions

60
.env.example Normal file
View File

@@ -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

3
requirements-phone.txt Normal file
View File

@@ -0,0 +1,3 @@
# Signal Bridge Remote — Phone Relay Client Dependencies
buttplug>=1.0.0,<2.0.0
websockets>=12.0

View File

@@ -25,10 +25,14 @@ from . import config
from .auth import ( from .auth import (
init_db, create_user, verify_user, create_token, verify_token, init_db, create_user, verify_user, create_token, verify_token,
extract_token, ip_tracker, rate_limiter, extract_token, ip_tracker, rate_limiter,
get_safety_config, set_safety_config,
) )
from .mcp_tools import TOOLS, HANDLERS, current_user_id 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 .relay_hub import check_ws_ip_limit, release_ws_ip_slot, get_ip_from_headers
from .session_registry import registry from .session_registry import registry
from .governor import governor
from .safety import dead_man_switch from .safety import dead_man_switch
# ── Logging ───────────────────────────────────────────────────────────── # ── Logging ─────────────────────────────────────────────────────────────
@@ -47,9 +51,11 @@ log = logging.getLogger("signal_bridge")
async def lifespan(app: FastAPI): async def lifespan(app: FastAPI):
config.validate() config.validate()
init_db() init_db()
init_oauth_db()
await dead_man_switch.start() await dead_man_switch.start()
log.info(f"Signal Bridge Remote started on {config.HOST}:{config.PORT}") 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"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 yield
await dead_man_switch.stop() await dead_man_switch.stop()
log.info("Signal Bridge Remote shutting down") log.info("Signal Bridge Remote shutting down")
@@ -69,6 +75,9 @@ app.add_middleware(
allow_headers=["*"], allow_headers=["*"],
) )
# Mount OAuth routes (metadata, registration, authorize, token)
app.include_router(oauth_router)
# ════════════════════════════════════════════════════════════════════════ # ════════════════════════════════════════════════════════════════════════
# Auth helpers # Auth helpers
@@ -180,10 +189,12 @@ async def _resolve_mcp_user(request: Request) -> dict | None:
return {"user_id": _mcp_sessions[session_id]} return {"user_id": _mcp_sessions[session_id]}
# 3. Fall back to sole active phone session (authless / claude.ai init) # 3. Fall back to sole active phone session (authless / claude.ai init)
fallback_user_id = await registry.get_sole_user_id() # Disabled when SB_REQUIRE_MCP_AUTH=true (multi-user mode).
if fallback_user_id: if not config.REQUIRE_MCP_AUTH:
log.info(f"MCP request without auth — using active session: {fallback_user_id}") fallback_user_id = await registry.get_sole_user_id()
return {"user_id": fallback_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 return None
@@ -395,6 +406,10 @@ async def _handle_phone_ws(ws: WebSocket):
wrapper = _FastAPIWSWrapper(ws) wrapper = _FastAPIWSWrapper(ws)
session = await registry.register(user_id, wrapper) 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) # Request device list (phone also sends proactively, but this is a backup)
log.info(f"Requesting device scan from phone: user={user_id}") log.info(f"Requesting device scan from phone: user={user_id}")
await ws.send_json({"type": "scan"}) await ws.send_json({"type": "scan"})
@@ -417,6 +432,11 @@ async def _handle_phone_ws(ws: WebSocket):
) )
if ack.request_id: if ack.request_id:
session.resolve_ack(ack.request_id, ack) 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": elif msg_type == "device_list":
await registry.update_devices(user_id, msg.get("devices", [])) await registry.update_devices(user_id, msg.get("devices", []))
log.info(f"Devices updated: user={user_id}, count={len(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: finally:
if user_id: if user_id:
await registry.unregister(user_id) await registry.unregister(user_id)
governor.remove_user(user_id)
log.info(f"Phone disconnected: user={user_id}") log.info(f"Phone disconnected: user={user_id}")
await release_ws_ip_slot(ip) await release_ws_ip_slot(ip)
@@ -463,6 +484,78 @@ class _FastAPIWSWrapper:
return None 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 # Health & Status
# ════════════════════════════════════════════════════════════════════════ # ════════════════════════════════════════════════════════════════════════
@@ -487,8 +580,10 @@ async def root():
"version": "1.0.0", "version": "1.0.0",
"endpoints": { "endpoints": {
"auth": "/auth/register, /auth/login", "auth": "/auth/register, /auth/login",
"oauth": "/.well-known/oauth-authorization-server, /oauth/register, /oauth/authorize, /oauth/token",
"mcp": "/mcp (POST, JSON-RPC)", "mcp": "/mcp (POST, JSON-RPC)",
"phone_relay": "/ws/phone (WebSocket)", "phone_relay": "/ws/phone (WebSocket)",
"safety": "/safety/config (GET, POST), /safety/status (GET)",
"health": "/health", "health": "/health",
}, },
} }

View File

@@ -24,7 +24,7 @@ from . import config
# ════════════════════════════════════════════════════════════════════════ # ════════════════════════════════════════════════════════════════════════
def init_db(): 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 = sqlite3.connect(config.DB_PATH)
conn.execute(""" conn.execute("""
CREATE TABLE IF NOT EXISTS users ( CREATE TABLE IF NOT EXISTS users (
@@ -35,10 +35,79 @@ def init_db():
is_active INTEGER DEFAULT 1 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.commit()
conn.close() 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(): def _get_conn():
conn = sqlite3.connect(config.DB_PATH) conn = sqlite3.connect(config.DB_PATH)
conn.row_factory = sqlite3.Row conn.row_factory = sqlite3.Row

View File

@@ -19,6 +19,7 @@ CORS_ORIGINS = os.getenv("SB_CORS_ORIGINS", "*").split(",")
# ── Auth ──────────────────────────────────────────────────────────────── # ── Auth ────────────────────────────────────────────────────────────────
TOKEN_EXPIRY_HOURS = int(os.getenv("SB_TOKEN_EXPIRY_HOURS", "168")) # 1 week TOKEN_EXPIRY_HOURS = int(os.getenv("SB_TOKEN_EXPIRY_HOURS", "168")) # 1 week
REGISTRATION_OPEN = os.getenv("SB_REGISTRATION_OPEN", "true").lower() == "true" REGISTRATION_OPEN = os.getenv("SB_REGISTRATION_OPEN", "true").lower() == "true"
REQUIRE_MCP_AUTH = os.getenv("SB_REQUIRE_MCP_AUTH", "false").lower() == "true"
# ── Rate Limiting ─────────────────────────────────────────────────────── # ── Rate Limiting ───────────────────────────────────────────────────────
# Format: "count/period" — e.g. "5/minute", "100/hour" # 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_INTERVAL_S = float(os.getenv("SB_HEARTBEAT_INTERVAL", "2.0"))
HEARTBEAT_TIMEOUT_S = float(os.getenv("SB_HEARTBEAT_TIMEOUT", "6.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 ──────────────────────────────────────────────────────────── # ── Database ────────────────────────────────────────────────────────────
DB_PATH = os.getenv("SB_DB_PATH", str(Path(__file__).parent / "signal_bridge.db")) DB_PATH = os.getenv("SB_DB_PATH", str(Path(__file__).parent / "signal_bridge.db"))

228
server/governor.py Normal file
View File

@@ -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()

View File

@@ -24,6 +24,7 @@ from .models import (
DeviceCommand, PatternCommand, StopCommand, ScanCommand, ReadSensorCommand, DeviceCommand, PatternCommand, StopCommand, ScanCommand, ReadSensorCommand,
CommandAck, CommandAck,
) )
from .governor import governor
from .session_registry import registry from .session_registry import registry
# Set by auth middleware before each MCP request # 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 # Helper
# ════════════════════════════════════════════════════════════════════════ # ════════════════════════════════════════════════════════════════════════
async def _send(command: dict) -> str: async def _send(command: dict, intensity: float = 0.0) -> str:
"""Route a command to the current user's phone and return result text.""" """
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() 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) 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: if ack.success:
return ack.message or "OK" return ack.message or "OK"
else: else:
@@ -126,6 +147,17 @@ async def list_devices(**kwargs) -> str:
+ (f" | floor: {floor}" if floor > 0 else "") + (f" | floor: {floor}" if floor > 0 else "")
+ (f" | {notes}" if notes 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) return "\n".join(lines)
@@ -166,13 +198,14 @@ def _make_output_handler(output_type: OutputType):
async def handler( async def handler(
device: str = "all", intensity: float = 0.5, duration: float = 0, **kw device: str = "all", intensity: float = 0.5, duration: float = 0, **kw
) -> str: ) -> str:
clamped = max(0.0, min(1.0, float(intensity)))
cmd = DeviceCommand( cmd = DeviceCommand(
action=output_type, action=output_type,
device=device, device=device,
intensity=max(0.0, min(1.0, intensity)), intensity=clamped,
duration=max(0.0, duration), duration=max(0.0, float(duration)),
) )
return await _send(cmd.model_dump()) return await _send(cmd.model_dump(), intensity=clamped)
return handler return handler
@@ -290,15 +323,16 @@ def _make_pattern_handler(pattern_name: str):
hold_seconds: float = 0, hold_seconds: float = 0,
**kw, **kw,
) -> str: ) -> str:
clamped = max(0.0, min(1.0, float(intensity)))
cmd = PatternCommand( cmd = PatternCommand(
pattern=pattern_name, pattern=pattern_name,
output_type=OutputType(output_type), output_type=OutputType(output_type),
device=device, device=device,
intensity=max(0.0, min(1.0, intensity)), intensity=clamped,
duration=max(0.0, duration), duration=max(0.0, float(duration)),
hold_seconds=max(0.0, hold_seconds), hold_seconds=max(0.0, float(hold_seconds)),
) )
return await _send(cmd.model_dump()) return await _send(cmd.model_dump(), intensity=clamped)
return handler return handler

439
server/oauth.py Normal file
View File

@@ -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("""\
<!DOCTYPE html>
<html lang="en">
<head>
<meta charset="utf-8">
<meta name="viewport" content="width=device-width, initial-scale=1">
<title>Signal Bridge — Sign In</title>
<style>
*, *::before, *::after { box-sizing: border-box; margin: 0; padding: 0; }
body {
font-family: -apple-system, BlinkMacSystemFont, "Segoe UI", Roboto, sans-serif;
background: #1a1a2e;
color: #e0e0e0;
display: flex;
align-items: center;
justify-content: center;
min-height: 100vh;
padding: 1rem;
}
.card {
background: #16213e;
border: 1px solid #0f3460;
border-radius: 12px;
padding: 2rem;
width: 100%;
max-width: 380px;
box-shadow: 0 8px 32px rgba(0,0,0,0.4);
}
h1 {
font-size: 1.3rem;
font-weight: 600;
margin-bottom: 0.3rem;
color: #e94560;
}
.subtitle {
font-size: 0.85rem;
color: #8892a4;
margin-bottom: 1.5rem;
}
.client-name {
color: #53c0f0;
font-weight: 500;
}
label {
display: block;
font-size: 0.8rem;
color: #8892a4;
margin-bottom: 0.3rem;
text-transform: uppercase;
letter-spacing: 0.05em;
}
input[type="text"], input[type="password"] {
width: 100%;
padding: 0.7rem 0.9rem;
border: 1px solid #0f3460;
border-radius: 8px;
background: #1a1a2e;
color: #e0e0e0;
font-size: 0.95rem;
margin-bottom: 1rem;
outline: none;
transition: border-color 0.2s;
}
input:focus { border-color: #e94560; }
button {
width: 100%;
padding: 0.75rem;
border: none;
border-radius: 8px;
background: #e94560;
color: #fff;
font-size: 1rem;
font-weight: 600;
cursor: pointer;
transition: background 0.2s;
}
button:hover { background: #c73652; }
.error {
background: rgba(233,69,96,0.15);
border: 1px solid #e94560;
border-radius: 8px;
padding: 0.7rem 0.9rem;
margin-bottom: 1rem;
font-size: 0.85rem;
color: #e94560;
}
.footer {
text-align: center;
margin-top: 1.2rem;
font-size: 0.75rem;
color: #555;
}
</style>
</head>
<body>
<div class="card">
<h1>Signal Bridge</h1>
<p class="subtitle">
Sign in to connect
<span class="client-name">$client_name</span>
</p>
$error_html
<form method="POST" action="/oauth/authorize">
<label for="username">Username</label>
<input type="text" id="username" name="username" required autocomplete="username" autofocus>
<label for="password">Password</label>
<input type="password" id="password" name="password" required autocomplete="current-password">
<input type="hidden" name="client_id" value="$client_id">
<input type="hidden" name="redirect_uri" value="$redirect_uri">
<input type="hidden" name="state" value="$state">
<input type="hidden" name="code_challenge" value="$code_challenge">
<input type="hidden" name="code_challenge_method" value="$code_challenge_method">
<input type="hidden" name="response_type" value="code">
<button type="submit">Sign In</button>
</form>
<p class="footer">Your credentials are verified by this server only.</p>
</div>
</body>
</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)

432
server/oauth_routes.py Normal file
View File

@@ -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(
"<h1>Error</h1><p>Missing client_id parameter</p>", status_code=400
)
client = await asyncio.to_thread(verify_client, client_id)
if not client:
return HTMLResponse(
"<h1>Error</h1><p>Unknown client_id</p>", status_code=400
)
# Verify redirect_uri is registered
if redirect_uri and redirect_uri not in client["redirect_uris"]:
return HTMLResponse(
"<h1>Error</h1><p>redirect_uri not registered for this client</p>",
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("<h1>Temporarily banned</h1>", status_code=429)
if not await rate_limiter.check(f"auth:{ip}", config.RATE_LIMIT_AUTH):
return HTMLResponse("<h1>Too many attempts</h1>", 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("<h1>Error</h1><p>Invalid client</p>", status_code=400)
if redirect_uri and redirect_uri not in client["redirect_uris"]:
return HTMLResponse(
"<h1>Error</h1><p>Invalid redirect_uri</p>", 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='<div class="error">Invalid username or password.</div>',
)
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"<h1>Error</h1><p>{description}</p>", 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("&", "&amp;")
.replace("<", "&lt;")
.replace(">", "&gt;")
.replace('"', "&quot;")
.replace("'", "&#x27;")
)

View File

@@ -14,6 +14,7 @@ import logging
import time import time
from . import config from . import config
from .governor import governor
from .session_registry import registry from .session_registry import registry
log = logging.getLogger("signal_bridge.safety") log = logging.getLogger("signal_bridge.safety")
@@ -70,8 +71,15 @@ class DeadManSwitch:
sessions = await registry.get_all_sessions() sessions = await registry.get_all_sessions()
for user_id, session in sessions.items(): for user_id, session in sessions.items():
# Send ping # Tick the governor (advances heat model)
ping = {"type": "heartbeat_ping", "timestamp": now} governor.tick(user_id)
# Send ping with governor state piggybacked
ping = {
"type": "heartbeat_ping",
"timestamp": now,
**governor.get_state(user_id),
}
try: try:
await session.websocket.send(json.dumps(ping)) await session.websocket.send(json.dumps(ping))
except Exception: except Exception:
@@ -92,6 +100,8 @@ class DeadManSwitch:
async def _emergency_stop(self, user_id: str, session): async def _emergency_stop(self, user_id: str, session):
"""Send stop-all and disconnect the session.""" """Send stop-all and disconnect the session."""
log.critical(f"EMERGENCY STOP for user {user_id} — all devices halted") log.critical(f"EMERGENCY STOP for user {user_id} — all devices halted")
governor.record_stop(user_id)
try: try:
stop_cmd = {"type": "stop", "device": "all", "emergency": True} stop_cmd = {"type": "stop", "device": "all", "emergency": True}
await session.websocket.send(json.dumps(stop_cmd)) await session.websocket.send(json.dumps(stop_cmd))
@@ -104,6 +114,7 @@ class DeadManSwitch:
pass pass
await registry.unregister(user_id) await registry.unregister(user_id)
governor.remove_user(user_id)
# Singleton # Singleton