mirror of
https://github.com/AletheiaVox/signal_bridge_remote.git
synced 2026-10-07 03:18:17 +08:00
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:
60
.env.example
Normal file
60
.env.example
Normal 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
3
requirements-phone.txt
Normal file
@@ -0,0 +1,3 @@
|
||||
# Signal Bridge Remote — Phone Relay Client Dependencies
|
||||
buttplug>=1.0.0,<2.0.0
|
||||
websockets>=12.0
|
||||
103
server/app.py
103
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",
|
||||
},
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"))
|
||||
|
||||
|
||||
228
server/governor.py
Normal file
228
server/governor.py
Normal 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()
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
439
server/oauth.py
Normal file
439
server/oauth.py
Normal 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
432
server/oauth_routes.py
Normal 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("&", "&")
|
||||
.replace("<", "<")
|
||||
.replace(">", ">")
|
||||
.replace('"', """)
|
||||
.replace("'", "'")
|
||||
)
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user