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 (
|
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",
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
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,
|
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
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
|
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
|
||||||
|
|||||||
Reference in New Issue
Block a user