Add files via upload

This commit is contained in:
Aletheia
2026-03-14 19:40:48 +01:00
committed by GitHub
commit 142e83d512
25 changed files with 3477 additions and 0 deletions

1
server/__init__.py Normal file
View File

@@ -0,0 +1 @@
# Signal Bridge Remote — Server Package

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

494
server/app.py Normal file
View File

@@ -0,0 +1,494 @@
"""
Signal Bridge Remote — Main Server Application
Single FastAPI app that serves three roles:
1. OAuth-style auth (register, login, token refresh)
2. MCP endpoint (Streamable HTTP — tool calls from Claude)
3. WebSocket relay hub (persistent phone connections)
Plus rate limiting, IP banning, and the dead man's switch.
"""
from __future__ import annotations
import asyncio
import json
import logging
import os
import sys
import uuid
from contextlib import asynccontextmanager
from fastapi import FastAPI, Request, WebSocket, WebSocketDisconnect
from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import JSONResponse
from . import config
from .auth import (
init_db, create_user, verify_user, create_token, verify_token,
extract_token, ip_tracker, rate_limiter,
)
from .mcp_tools import TOOLS, HANDLERS, current_user_id
from .relay_hub import check_ws_ip_limit, release_ws_ip_slot, get_ip_from_headers
from .session_registry import registry
from .safety import dead_man_switch
# ── Logging ─────────────────────────────────────────────────────────────
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s [%(name)s] %(levelname)s: %(message)s",
datefmt="%H:%M:%S",
)
log = logging.getLogger("signal_bridge")
# ── Lifespan ────────────────────────────────────────────────────────────
@asynccontextmanager
async def lifespan(app: FastAPI):
config.validate()
init_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'}")
yield
await dead_man_switch.stop()
log.info("Signal Bridge Remote shutting down")
app = FastAPI(
title="Signal Bridge Remote",
version="1.0.0",
lifespan=lifespan,
)
app.add_middleware(
CORSMiddleware,
allow_origins=config.CORS_ORIGINS,
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
# ════════════════════════════════════════════════════════════════════════
# Auth helpers
# ════════════════════════════════════════════════════════════════════════
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"
async def _require_auth(request: Request) -> dict | None:
"""Validate Bearer token. Returns user dict or None."""
token = extract_token(request.headers.get("Authorization", ""))
if not token:
return None
return verify_token(token)
# ════════════════════════════════════════════════════════════════════════
# Auth Endpoints
# ════════════════════════════════════════════════════════════════════════
@app.post("/auth/register")
async def register(request: Request):
"""Register a new user account."""
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"auth:{ip}", config.RATE_LIMIT_AUTH):
return JSONResponse({"error": "Too many attempts"}, status_code=429)
if not config.REGISTRATION_OPEN:
return JSONResponse({"error": "Registration is closed"}, status_code=403)
body = await request.json()
username = body.get("username", "").strip()
password = body.get("password", "")
try:
user = await asyncio.to_thread(create_user, username, password)
except ValueError as e:
# Don't count validation errors (short username, weak password) toward IP ban.
# Only actual auth failures (wrong credentials) should inflate the ban counter.
return JSONResponse({"error": str(e)}, status_code=400)
token = create_token(user["user_id"], user["username"])
await ip_tracker.clear_failures(ip)
return {"user_id": user["user_id"], "username": user["username"], "token": token}
@app.post("/auth/login")
async def login(request: Request):
"""Authenticate and receive a JWT."""
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"auth:{ip}", config.RATE_LIMIT_AUTH):
return JSONResponse({"error": "Too many attempts"}, status_code=429)
body = await request.json()
username = body.get("username", "")
password = body.get("password", "")
user = await asyncio.to_thread(verify_user, username, password)
if not user:
await ip_tracker.record_failure(ip)
return JSONResponse({"error": "Invalid credentials"}, status_code=401)
token = create_token(user["user_id"], user["username"])
await ip_tracker.clear_failures(ip)
return {"user_id": user["user_id"], "username": user["username"], "token": token}
# ════════════════════════════════════════════════════════════════════════
# MCP Endpoint — Streamable HTTP (JSON-RPC over POST + GET)
#
# Implements the MCP Streamable HTTP transport spec:
# - POST: JSON-RPC requests from client
# - GET: SSE stream for server-to-client notifications (kept open)
# - Mcp-Session-Id header for session tracking
# - Authless mode for claude.ai connector, Bearer token for Claude Desktop
# ════════════════════════════════════════════════════════════════════════
# In-memory MCP session tracking (maps session_id → user_id)
_mcp_sessions: dict[str, str] = {}
async def _resolve_mcp_user(request: Request) -> dict | None:
"""
Resolve the user for an MCP request.
Priority: Bearer token > Mcp-Session-Id lookup > sole active phone session.
"""
# 1. Try Bearer token auth (Claude Desktop)
user = await _require_auth(request)
if user:
return user
# 2. Try Mcp-Session-Id (subsequent requests from claude.ai)
session_id = request.headers.get("mcp-session-id", "")
if session_id and session_id in _mcp_sessions:
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}
return None
@app.post("/mcp")
async def mcp_endpoint(request: Request):
"""
MCP Streamable HTTP endpoint (POST).
Accepts JSON-RPC requests, routes tool calls to the authenticated
user's phone via the session registry.
"""
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"global:{ip}", config.RATE_LIMIT_GLOBAL):
return JSONResponse({"error": "Rate limit exceeded"}, status_code=429)
# Parse JSON-RPC first (we need to check if it's an initialize request)
try:
body = await request.json()
except Exception:
return _jsonrpc_error(None, -32700, "Parse error: invalid JSON")
method = body.get("method", "")
params = body.get("params", {})
req_id = body.get("id")
# Resolve user
user = await _resolve_mcp_user(request)
if not user:
return JSONResponse(
{"jsonrpc": "2.0", "error": {"code": -32000, "message": "No auth token and no active phone session"}},
status_code=401,
)
# Rate limit per user for commands
if not await rate_limiter.check(
f"cmd:{user['user_id']}", config.RATE_LIMIT_COMMANDS
):
return JSONResponse(
{"jsonrpc": "2.0", "error": {"code": -32000, "message": "Command rate limit exceeded"}},
status_code=429,
)
# Set user context for tool handlers
current_user_id.set(user["user_id"])
# ── Route by method ─────────────────────────────────────────────
if method == "initialize":
# Generate a session ID and bind it to this user
session_id = str(uuid.uuid4())
_mcp_sessions[session_id] = user["user_id"]
log.info(f"MCP session created: {session_id[:8]}... for user {user['user_id']}")
result = {
"protocolVersion": "2025-03-26",
"capabilities": {"tools": {}},
"serverInfo": {"name": "Signal Bridge Remote", "version": "1.0.0"},
}
response = JSONResponse({"jsonrpc": "2.0", "id": req_id, "result": result})
response.headers["Mcp-Session-Id"] = session_id
return response
elif method == "tools/list":
return _jsonrpc_result(req_id, {"tools": TOOLS})
elif method == "tools/call":
tool_name = params.get("name", "")
tool_args = params.get("arguments", {})
handler = HANDLERS.get(tool_name)
if not handler:
return _jsonrpc_error(req_id, -32601, f"Unknown tool: {tool_name}")
try:
result_text = await handler(**tool_args)
return _jsonrpc_result(req_id, {
"content": [{"type": "text", "text": result_text}],
})
except Exception as e:
log.error(f"Tool {tool_name} error: {e}")
return _jsonrpc_result(req_id, {
"content": [{"type": "text", "text": f"Error: {e}"}],
"isError": True,
})
elif method == "ping":
return _jsonrpc_result(req_id, {})
elif method == "resources/list":
return _jsonrpc_result(req_id, {"resources": []})
elif method == "prompts/list":
return _jsonrpc_result(req_id, {"prompts": []})
elif method.startswith("notifications/"):
# MCP notifications (e.g. notifications/initialized) are fire-and-forget.
# Return empty success — no error, no noise.
return _jsonrpc_result(req_id, {})
else:
return _jsonrpc_error(req_id, -32601, f"Unknown method: {method}")
@app.get("/mcp")
async def mcp_sse_endpoint(request: Request):
"""
MCP Streamable HTTP endpoint (GET).
Opens an SSE stream for server-to-client notifications.
We don't currently use server-initiated notifications,
so this just stays open to satisfy the spec.
"""
from starlette.responses import StreamingResponse
async def event_stream():
# Send a keep-alive comment, then hold the connection open
yield ": connected\n\n"
try:
while True:
await asyncio.sleep(30)
yield ": keepalive\n\n"
except asyncio.CancelledError:
pass
return StreamingResponse(
event_stream(),
media_type="text/event-stream",
headers={
"Cache-Control": "no-cache",
"Connection": "keep-alive",
},
)
def _jsonrpc_result(req_id, result):
return JSONResponse({"jsonrpc": "2.0", "id": req_id, "result": result})
def _jsonrpc_error(req_id, code, message):
return JSONResponse(
{"jsonrpc": "2.0", "id": req_id, "error": {"code": code, "message": message}}
)
# ════════════════════════════════════════════════════════════════════════
# WebSocket Relay — Phone connections
# ════════════════════════════════════════════════════════════════════════
@app.websocket("/ws/phone")
async def websocket_phone(websocket: WebSocket):
"""
WebSocket endpoint for phone relay clients.
The phone connects here, authenticates with its JWT,
and maintains a persistent connection for receiving device commands.
"""
await websocket.accept()
await _handle_phone_ws(websocket)
async def _handle_phone_ws(ws: WebSocket):
"""
Full phone WebSocket lifecycle: auth → register → message loop → cleanup.
"""
from .models import CommandAck
ip = get_ip_from_headers(
ws.client.host if ws.client else None,
dict(ws.headers) if ws.headers else None,
)
# IP-level rate limiting
rejection = await check_ws_ip_limit(ip)
if rejection:
await ws.close(4003, rejection)
return
user_id = None
try:
# Wait for auth message
raw = await asyncio.wait_for(ws.receive_text(), timeout=10.0)
msg = json.loads(raw)
if msg.get("type") != "phone_auth" or "token" not in msg:
await ws.close(4001, "First message must be phone_auth")
await ip_tracker.record_failure(ip)
return
user = verify_token(msg["token"])
if not user:
await ws.close(4001, "Invalid token")
await ip_tracker.record_failure(ip)
return
user_id = user["user_id"]
await ip_tracker.clear_failures(ip)
await ws.send_json({
"type": "auth_ok",
"user_id": user_id,
"message": "Connected to Signal Bridge relay",
})
log.info(f"Phone connected: user={user['username']} ip={ip}")
# Create a wrapper that looks like a websockets ServerConnection
wrapper = _FastAPIWSWrapper(ws)
session = await registry.register(user_id, wrapper)
# 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"})
# Message loop
while True:
try:
raw = await ws.receive_text()
msg = json.loads(raw)
msg_type = msg.get("type")
if msg_type == "heartbeat_pong":
await registry.update_heartbeat(user_id)
elif msg_type == "command_ack":
ack = CommandAck(
success=msg.get("success", True),
message=msg.get("message", ""),
request_id=msg.get("request_id"),
data=msg.get("data"),
)
if ack.request_id:
session.resolve_ack(ack.request_id, ack)
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', []))}")
except WebSocketDisconnect:
break
except json.JSONDecodeError:
continue
except asyncio.TimeoutError:
await ws.close(4001, "Auth timeout")
except WebSocketDisconnect:
pass
except Exception as e:
log.error(f"Phone WS error: {e}")
finally:
if user_id:
await registry.unregister(user_id)
log.info(f"Phone disconnected: user={user_id}")
await release_ws_ip_slot(ip)
class _FastAPIWSWrapper:
"""
Minimal wrapper to make a FastAPI WebSocket look enough like a
websockets ServerConnection for the session registry and safety module.
"""
def __init__(self, ws: WebSocket):
self._ws = ws
async def send(self, data: str):
await self._ws.send_text(data)
async def close(self, code: int = 1000, reason: str = ""):
await self._ws.close(code, reason)
@property
def transport(self):
return self # duck typing for _get_ip fallback
def get_extra_info(self, key):
if key == "peername" and self._ws.client:
return (self._ws.client.host, self._ws.client.port)
return None
# ════════════════════════════════════════════════════════════════════════
# Health & Status
# ════════════════════════════════════════════════════════════════════════
@app.get("/health")
async def health():
return {
"status": "ok",
"active_phones": registry.active_count,
"banned_ips": ip_tracker.banned_count,
}
# ════════════════════════════════════════════════════════════════════════
# Init module
# ════════════════════════════════════════════════════════════════════════
@app.get("/")
async def root():
return {
"service": "Signal Bridge Remote",
"version": "1.0.0",
"endpoints": {
"auth": "/auth/register, /auth/login",
"mcp": "/mcp (POST, JSON-RPC)",
"phone_relay": "/ws/phone (WebSocket)",
"health": "/health",
},
}

226
server/auth.py Normal file
View File

@@ -0,0 +1,226 @@
"""
Signal Bridge Remote — Authentication & Rate Limiting
JWT-based auth with bcrypt password hashing, SQLite user store,
progressive IP banning, and per-endpoint rate limiting.
"""
from __future__ import annotations
import asyncio
import sqlite3
import time
import uuid
from collections import defaultdict
from datetime import datetime, timedelta, timezone
from typing import Optional
import bcrypt
import jwt
from . import config
# ════════════════════════════════════════════════════════════════════════
# Database
# ════════════════════════════════════════════════════════════════════════
def init_db():
"""Create user table if it doesn't exist."""
conn = sqlite3.connect(config.DB_PATH)
conn.execute("""
CREATE TABLE IF NOT EXISTS users (
id TEXT PRIMARY KEY,
username TEXT UNIQUE NOT NULL,
password_hash TEXT NOT NULL,
created_at TEXT NOT NULL,
is_active INTEGER DEFAULT 1
)
""")
conn.commit()
conn.close()
def _get_conn():
conn = sqlite3.connect(config.DB_PATH)
conn.row_factory = sqlite3.Row
return conn
# ════════════════════════════════════════════════════════════════════════
# User Management
# ════════════════════════════════════════════════════════════════════════
def create_user(username: str, password: str) -> dict:
"""Register a new user. Returns user dict or raises ValueError."""
if len(username) < 3 or len(username) > 32:
raise ValueError("Username must be 3-32 characters")
if len(password) < 8:
raise ValueError("Password must be at least 8 characters")
user_id = str(uuid.uuid4())
password_hash = bcrypt.hashpw(password.encode(), bcrypt.gensalt()).decode()
conn = _get_conn()
try:
conn.execute(
"INSERT INTO users (id, username, password_hash, created_at) VALUES (?, ?, ?, ?)",
(user_id, username.lower().strip(), password_hash,
datetime.now(timezone.utc).isoformat()),
)
conn.commit()
except sqlite3.IntegrityError:
raise ValueError("Username already taken")
finally:
conn.close()
return {"user_id": user_id, "username": username.lower().strip()}
def verify_user(username: str, password: str) -> Optional[dict]:
"""Check credentials. Returns user dict or None."""
conn = _get_conn()
row = conn.execute(
"SELECT * FROM users WHERE username = ? AND is_active = 1",
(username.lower().strip(),),
).fetchone()
conn.close()
if not row:
return None
if not bcrypt.checkpw(password.encode(), row["password_hash"].encode()):
return None
return {"user_id": row["id"], "username": row["username"]}
# ════════════════════════════════════════════════════════════════════════
# JWT Tokens
# ════════════════════════════════════════════════════════════════════════
def create_token(user_id: str, username: str) -> str:
"""Issue a JWT token."""
payload = {
"sub": user_id,
"username": username,
"iat": datetime.now(timezone.utc),
"exp": datetime.now(timezone.utc) + timedelta(hours=config.TOKEN_EXPIRY_HOURS),
}
return jwt.encode(payload, config.SECRET_KEY, algorithm="HS256")
def verify_token(token: str) -> Optional[dict]:
"""Validate a JWT. Returns payload dict or None."""
try:
payload = jwt.decode(token, config.SECRET_KEY, algorithms=["HS256"])
return {"user_id": payload["sub"], "username": payload["username"]}
except (jwt.ExpiredSignatureError, jwt.InvalidTokenError):
return None
def extract_token(authorization: str) -> Optional[str]:
"""Pull the token from an Authorization header value."""
if not authorization:
return None
parts = authorization.split()
if len(parts) == 2 and parts[0].lower() == "bearer":
return parts[1]
return None
# ════════════════════════════════════════════════════════════════════════
# IP Ban Tracker
# ════════════════════════════════════════════════════════════════════════
class IPBanTracker:
"""
Tracks failed auth attempts per IP and issues temporary bans
after threshold is exceeded. Designed to frustrate targeted
harassment without affecting legitimate users.
"""
def __init__(self):
self._failures: dict[str, list[float]] = defaultdict(list)
self._bans: dict[str, float] = {} # ip → ban_expires_at timestamp
self._lock = asyncio.Lock()
async def record_failure(self, ip: str):
"""Record a failed auth attempt. May trigger a ban."""
async with self._lock:
now = time.time()
window = now - 3600 # 1-hour sliding window
self._failures[ip] = [t for t in self._failures[ip] if t > window]
self._failures[ip].append(now)
if len(self._failures[ip]) >= config.BAN_THRESHOLD:
self._bans[ip] = now + (config.BAN_DURATION_MINUTES * 60)
self._failures[ip] = []
async def is_banned(self, ip: str) -> bool:
"""Check if an IP is currently banned."""
async with self._lock:
if ip not in self._bans:
return False
if time.time() > self._bans[ip]:
del self._bans[ip]
return False
return True
async def clear_failures(self, ip: str):
"""Reset failure counter on successful auth."""
async with self._lock:
self._failures.pop(ip, None)
@property
def banned_count(self) -> int:
"""Number of currently banned IPs."""
now = time.time()
return sum(1 for exp in self._bans.values() if exp > now)
# Singleton
ip_tracker = IPBanTracker()
# ════════════════════════════════════════════════════════════════════════
# Rate Limiter (simple token bucket per key)
# ════════════════════════════════════════════════════════════════════════
class RateLimiter:
"""
Simple in-memory rate limiter using sliding window counters.
Parse rate strings like "5/minute", "100/hour".
"""
PERIODS = {
"second": 1, "minute": 60, "hour": 3600, "day": 86400,
}
def __init__(self):
self._windows: dict[str, list[float]] = defaultdict(list)
self._lock = asyncio.Lock()
@staticmethod
def _parse_rate(rate_str: str) -> tuple[int, int]:
"""Parse '5/minute' → (5, 60)."""
count_str, period_str = rate_str.split("/")
return int(count_str), RateLimiter.PERIODS[period_str]
async def check(self, key: str, rate_str: str) -> bool:
"""
Check if request is allowed. Returns True if allowed.
Automatically records the attempt if allowed.
"""
max_count, period = self._parse_rate(rate_str)
async with self._lock:
now = time.time()
window_start = now - period
self._windows[key] = [t for t in self._windows[key] if t > window_start]
if len(self._windows[key]) >= max_count:
return False
self._windows[key].append(now)
return True
# Singleton
rate_limiter = RateLimiter()

46
server/config.py Normal file
View File

@@ -0,0 +1,46 @@
"""
Signal Bridge Remote — Server Configuration
All settings are loaded from environment variables with sensible defaults.
In production, set SB_SECRET_KEY to a random 64-char string.
"""
import os
from pathlib import Path
from dotenv import load_dotenv
load_dotenv() # Load .env file before reading any env vars
# ── Server ──────────────────────────────────────────────────────────────
HOST = os.getenv("SB_HOST", "0.0.0.0")
PORT = int(os.getenv("SB_PORT", "8420"))
SECRET_KEY = os.getenv("SB_SECRET_KEY", "") # MUST be set in production
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"
# ── Rate Limiting ───────────────────────────────────────────────────────
# Format: "count/period" — e.g. "5/minute", "100/hour"
RATE_LIMIT_AUTH = os.getenv("SB_RATE_LIMIT_AUTH", "5/minute")
RATE_LIMIT_COMMANDS = os.getenv("SB_RATE_LIMIT_COMMANDS", "120/minute")
RATE_LIMIT_GLOBAL = os.getenv("SB_RATE_LIMIT_GLOBAL", "300/minute")
MAX_WS_PER_IP = int(os.getenv("SB_MAX_WS_PER_IP", "3"))
BAN_THRESHOLD = int(os.getenv("SB_BAN_THRESHOLD", "20"))
BAN_DURATION_MINUTES = int(os.getenv("SB_BAN_DURATION_MINUTES", "30"))
# ── Safety ──────────────────────────────────────────────────────────────
HEARTBEAT_INTERVAL_S = float(os.getenv("SB_HEARTBEAT_INTERVAL", "2.0"))
HEARTBEAT_TIMEOUT_S = float(os.getenv("SB_HEARTBEAT_TIMEOUT", "6.0"))
# ── Database ────────────────────────────────────────────────────────────
DB_PATH = os.getenv("SB_DB_PATH", str(Path(__file__).parent / "signal_bridge.db"))
def validate():
"""Check that critical config is set. Call on startup."""
if not SECRET_KEY:
raise RuntimeError(
"SB_SECRET_KEY is not set. Generate one with: "
"python -c \"import secrets; print(secrets.token_hex(32))\""
)

373
server/mcp_tools.py Normal file
View File

@@ -0,0 +1,373 @@
"""
Signal Bridge Remote — MCP Tool Definitions
All tools that Claude can call to control devices. Each tool:
1. Validates input
2. Builds a command message
3. Routes it through the session registry to the user's phone
4. Returns the result to Claude
Expanded to support ALL Buttplug output types:
vibrate, rotate, oscillate, constrict, temperature, led, position, spray
And sensor input types:
battery, rssi, pressure, button, depth, position
"""
from __future__ import annotations
import asyncio
import contextvars
import json
from typing import Optional
from .models import (
OutputType, InputType,
DeviceCommand, PatternCommand, StopCommand, ScanCommand, ReadSensorCommand,
CommandAck,
)
from .session_registry import registry
# Set by auth middleware before each MCP request
current_user_id: contextvars.ContextVar[str] = contextvars.ContextVar("current_user_id")
# ════════════════════════════════════════════════════════════════════════
# Tool registry — built at import time, consumed by the MCP endpoint
# ════════════════════════════════════════════════════════════════════════
TOOLS: list[dict] = [] # MCP tool definitions (schema)
HANDLERS: dict[str, callable] = {} # tool_name → async handler function
def _register_tool(name: str, description: str, params: dict, required: list[str] = None):
"""Decorator factory for registering MCP tools."""
def decorator(fn):
schema = {
"type": "object",
"properties": params,
}
# Infer required fields: any param without a "default" key is required
if required is not None:
schema["required"] = required
else:
inferred = [k for k, v in params.items() if "default" not in v]
if inferred:
schema["required"] = inferred
TOOLS.append({
"name": name,
"description": description,
"inputSchema": schema,
})
HANDLERS[name] = fn
return fn
return decorator
# ════════════════════════════════════════════════════════════════════════
# Helper
# ════════════════════════════════════════════════════════════════════════
async def _send(command: dict) -> str:
"""Route a command to the current user's phone and return result text."""
user_id = current_user_id.get()
ack = await registry.send_to_user(user_id, command)
if ack.success:
return ack.message or "OK"
else:
return f"Error: {ack.message}"
# ════════════════════════════════════════════════════════════════════════
# Device Discovery
# ════════════════════════════════════════════════════════════════════════
@_register_tool(
"list_devices",
"List all connected devices with their capabilities, intensity floors, and notes.",
{},
)
async def list_devices(**kwargs) -> str:
user_id = current_user_id.get()
devices = await registry.get_devices(user_id)
# If cache is empty but phone is connected, try requesting a fresh scan
if not devices:
session = await registry.get_session(user_id)
if session:
# Phone is connected but device list is empty — request a scan
try:
scan_ack = await session.send_command({"type": "scan"}, timeout=15.0)
if scan_ack.success:
# Give a moment for the device_list message to arrive and be processed
await asyncio.sleep(0.5)
devices = await registry.get_devices(user_id)
except Exception:
pass
if not devices:
# Check if there's even a session
session = await registry.get_session(user_id)
if not session:
return (
"No phone connected. Start the relay client on your phone/PC "
"and connect it to the server."
)
return (
"Phone is connected but no devices found. Make sure Intiface Central "
"is running and devices are turned on."
)
lines = []
for d in devices:
caps = ", ".join(d.get("capabilities", {}).keys())
notes = d.get("notes", "")
floor = d.get("intensity_floor", 0)
lines.append(
f"• {d.get('short_name', '?')} — capabilities: [{caps}]"
+ (f" | floor: {floor}" if floor > 0 else "")
+ (f" | {notes}" if notes else "")
)
return "\n".join(lines)
@_register_tool(
"scan_devices",
"Rescan for new or reconnected Bluetooth devices.",
{},
)
async def scan_devices(**kwargs) -> str:
return await _send(ScanCommand().model_dump())
# ════════════════════════════════════════════════════════════════════════
# Output Commands — one tool per output type
# ════════════════════════════════════════════════════════════════════════
_OUTPUT_PARAMS = {
"device": {
"type": "string",
"description": "Device short name (e.g. 'ferri', 'lush', 'gravity') or 'all'",
"default": "all",
},
"intensity": {
"type": "number",
"description": "Intensity from 0.0 (off) to 1.0 (maximum)",
"default": 0.5,
},
"duration": {
"type": "number",
"description": "Duration in seconds. 0 = stay on until stop command.",
"default": 0,
},
}
def _make_output_handler(output_type: OutputType):
"""Factory for output command handlers."""
async def handler(
device: str = "all", intensity: float = 0.5, duration: float = 0, **kw
) -> str:
cmd = DeviceCommand(
action=output_type,
device=device,
intensity=max(0.0, min(1.0, intensity)),
duration=max(0.0, duration),
)
return await _send(cmd.model_dump())
return handler
# Standard outputs (available on most devices)
_register_tool(
"vibrate",
"Send vibration to a device. Most common output type.",
_OUTPUT_PARAMS,
)(_make_output_handler(OutputType.VIBRATE))
_register_tool(
"rotate",
"Send rotation/sonic pulse output. Device-specific — some devices use this "
"for sonic clitoral stimulation rather than physical rotation.",
_OUTPUT_PARAMS,
)(_make_output_handler(OutputType.ROTATE))
_register_tool(
"oscillate",
"Send oscillation/thrusting output. Device-specific — typically linear "
"thrusting motion.",
_OUTPUT_PARAMS,
)(_make_output_handler(OutputType.OSCILLATE))
# Extended outputs (device-specific, may not be available on all hardware)
_register_tool(
"constrict",
"Send constriction/compression output. Device-specific — available on "
"devices with squeeze or compression mechanisms.",
_OUTPUT_PARAMS,
)(_make_output_handler(OutputType.CONSTRICT))
_register_tool(
"temperature",
"Set temperature output. Device-specific — available on devices with "
"heating or cooling elements. Intensity maps to temperature range.",
_OUTPUT_PARAMS,
)(_make_output_handler(OutputType.TEMPERATURE))
_register_tool(
"led",
"Control LED light output. Device-specific — intensity controls brightness.",
_OUTPUT_PARAMS,
)(_make_output_handler(OutputType.LED))
_register_tool(
"position",
"Set linear position. Device-specific — intensity maps to position "
"along the device's range of motion (0.0 = retracted, 1.0 = extended).",
_OUTPUT_PARAMS,
)(_make_output_handler(OutputType.POSITION))
_register_tool(
"spray",
"Trigger spray/liquid output. Device-specific.",
_OUTPUT_PARAMS,
)(_make_output_handler(OutputType.SPRAY))
# ════════════════════════════════════════════════════════════════════════
# Stop
# ════════════════════════════════════════════════════════════════════════
@_register_tool(
"stop",
"Immediately stop all output on a device (or all devices). "
"Also cancels any running patterns.",
{
"device": {
"type": "string",
"description": "Device short name or 'all'",
"default": "all",
},
},
)
async def stop(device: str = "all", **kwargs) -> str:
return await _send(StopCommand(device=device).model_dump())
# ════════════════════════════════════════════════════════════════════════
# Patterns — work with ANY output type
# ════════════════════════════════════════════════════════════════════════
_PATTERN_PARAMS = {
"device": {
"type": "string",
"description": "Device short name or 'all'",
"default": "all",
},
"output_type": {
"type": "string",
"description": "Which output to modulate: vibrate, rotate, oscillate, "
"constrict, temperature, led, position, spray",
"default": "vibrate",
},
"intensity": {
"type": "number",
"description": "Peak intensity (0.0–1.0)",
"default": 0.6,
},
"duration": {
"type": "number",
"description": "Duration in seconds",
"default": 10,
},
}
def _make_pattern_handler(pattern_name: str):
async def handler(
device: str = "all",
output_type: str = "vibrate",
intensity: float = 0.6,
duration: float = 10,
hold_seconds: float = 0,
**kw,
) -> str:
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),
)
return await _send(cmd.model_dump())
return handler
_register_tool(
"pulse",
"Rhythmic on/off pattern. 0.5s on at intensity, 0.3s off, repeating. "
"Works with any output type (default: vibrate).",
_PATTERN_PARAMS,
)(_make_pattern_handler("pulse"))
_register_tool(
"wave",
"Smooth sine-wave intensity modulation. Rises and falls continuously. "
"Works with any output type (default: vibrate).",
_PATTERN_PARAMS,
)(_make_pattern_handler("wave"))
_register_tool(
"escalate",
"Gradual ramp from 0% to peak intensity over the duration, then hold at peak. "
"Use hold_seconds to auto-stop after holding (0 = hold indefinitely until stop command). "
"Works with any output type (default: vibrate).",
{k: v for k, v in _PATTERN_PARAMS.items() if k != "intensity"}
| {
"intensity": {"type": "number", "description": "Peak intensity to ramp up to", "default": 1.0},
"hold_seconds": {
"type": "number",
"description": "Seconds to hold at peak after ramp completes. 0 = hold indefinitely until explicit stop.",
"default": 0,
},
},
)(_make_pattern_handler("escalate"))
# ════════════════════════════════════════════════════════════════════════
# Sensor Inputs — read data FROM the device
# ════════════════════════════════════════════════════════════════════════
@_register_tool(
"read_battery",
"Read battery level from a device. Returns percentage (0-100).",
{
"device": {
"type": "string",
"description": "Device short name",
},
},
)
async def read_battery(device: str, **kwargs) -> str:
cmd = ReadSensorCommand(sensor=InputType.BATTERY, device=device)
return await _send(cmd.model_dump())
@_register_tool(
"read_sensor",
"Read a sensor value from a device. Available sensors depend on hardware: "
"battery, rssi (signal strength), pressure, button, depth, position. "
"Not all devices support all sensors.",
{
"device": {
"type": "string",
"description": "Device short name",
},
"sensor": {
"type": "string",
"description": "Sensor type: battery, rssi, pressure, button, depth, position",
},
},
)
async def read_sensor(device: str, sensor: str, **kwargs) -> str:
cmd = ReadSensorCommand(sensor=InputType(sensor), device=device)
return await _send(cmd.model_dump())

123
server/models.py Normal file
View File

@@ -0,0 +1,123 @@
"""
Signal Bridge Remote — Shared Models & Command Protocol
Defines the JSON message format between all three tiers:
Claude ←(MCP)→ VPS Server ←(WebSocket)→ Phone
"""
from __future__ import annotations
from pydantic import BaseModel, Field
from typing import Optional, Any
from enum import Enum
# ── Output Types (commands TO the device) ───────────────────────────────
class OutputType(str, Enum):
VIBRATE = "vibrate"
ROTATE = "rotate"
OSCILLATE = "oscillate"
CONSTRICT = "constrict" # compression / squeeze
TEMPERATURE = "temperature" # heating / cooling
LED = "led" # light control
POSITION = "position" # linear positioning
SPRAY = "spray" # liquid / spray
# ── Input Types (readings FROM the device) ──────────────────────────────
class InputType(str, Enum):
BATTERY = "battery"
RSSI = "rssi" # signal strength
PRESSURE = "pressure"
BUTTON = "button"
DEPTH = "depth"
POSITION = "position"
# ════════════════════════════════════════════════════════════════════════
# Server → Phone messages
# ════════════════════════════════════════════════════════════════════════
class DeviceCommand(BaseModel):
"""Direct output command to a device."""
type: str = "command"
action: OutputType
device: str = "all"
intensity: float = Field(0.5, ge=0.0, le=1.0)
duration: float = Field(0.0, ge=0.0) # 0 = indefinite
class PatternCommand(BaseModel):
"""Run a named pattern on a device."""
type: str = "pattern"
pattern: str # "pulse", "wave", "escalate"
output_type: OutputType = OutputType.VIBRATE
device: str = "all"
intensity: float = Field(0.6, ge=0.0, le=1.0)
duration: float = Field(10.0, ge=0.0)
hold_seconds: float = Field(0.0, ge=0.0) # escalate only: 0 = hold at peak indefinitely
class StopCommand(BaseModel):
type: str = "stop"
device: str = "all"
class ScanCommand(BaseModel):
type: str = "scan"
class ReadSensorCommand(BaseModel):
type: str = "read_sensor"
sensor: InputType
device: str
class HeartbeatPing(BaseModel):
type: str = "heartbeat_ping"
timestamp: float
# ════════════════════════════════════════════════════════════════════════
# Phone → Server messages
# ════════════════════════════════════════════════════════════════════════
class HeartbeatPong(BaseModel):
type: str = "heartbeat_pong"
timestamp: float
class DeviceListReport(BaseModel):
type: str = "device_list"
devices: list[dict[str, Any]] = []
class CommandAck(BaseModel):
type: str = "command_ack"
success: bool = True
message: str = ""
request_id: Optional[str] = None
data: Optional[dict[str, Any]] = None # sensor readings, etc.
class PhoneAuth(BaseModel):
type: str = "phone_auth"
token: str
# ════════════════════════════════════════════════════════════════════════
# MCP JSON-RPC models
# ════════════════════════════════════════════════════════════════════════
class MCPRequest(BaseModel):
jsonrpc: str = "2.0"
id: Optional[str | int] = None
method: str
params: Optional[dict[str, Any]] = None
class MCPResponse(BaseModel):
jsonrpc: str = "2.0"
id: Optional[str | int] = None
result: Optional[Any] = None
error: Optional[dict[str, Any]] = None

53
server/relay_hub.py Normal file
View File

@@ -0,0 +1,53 @@
"""
Signal Bridge Remote — WebSocket Relay Hub (utilities)
IP tracking and rate limiting for phone WebSocket connections.
The actual WebSocket handling lives in app.py using FastAPI's native WebSocket.
This module provides the shared state and helpers used by the app endpoint.
"""
from __future__ import annotations
import asyncio
import logging
from collections import defaultdict
from . import config
from .auth import ip_tracker
log = logging.getLogger("signal_bridge.relay")
# Track WebSocket connections per IP for rate limiting
ws_count_by_ip: dict[str, int] = defaultdict(int)
ws_lock = asyncio.Lock()
async def check_ws_ip_limit(ip: str) -> str | None:
"""
Check if an IP is allowed to open a new WebSocket connection.
Returns an error reason string if rejected, None if allowed.
Automatically increments the counter if allowed.
"""
if await ip_tracker.is_banned(ip):
return "Temporarily banned"
async with ws_lock:
if ws_count_by_ip[ip] >= config.MAX_WS_PER_IP:
return "Too many connections from this IP"
ws_count_by_ip[ip] += 1
return None
async def release_ws_ip_slot(ip: str):
"""Decrement the connection counter for an IP when a WebSocket disconnects."""
async with ws_lock:
ws_count_by_ip[ip] = max(0, ws_count_by_ip[ip] - 1)
def get_ip_from_headers(host: str | None, headers: dict | None = None) -> str:
"""Extract real IP, checking X-Forwarded-For for reverse proxy setups."""
if headers:
forwarded = headers.get("X-Forwarded-For", headers.get("x-forwarded-for", ""))
if forwarded:
return forwarded.split(",")[0].strip()
return host or "unknown"

110
server/safety.py Normal file
View File

@@ -0,0 +1,110 @@
"""
Signal Bridge Remote — Safety Systems
Dead Man's Switch: Monitors phone connections via heartbeat.
If a phone stops responding, all its devices are stopped immediately.
This is non-negotiable safety infrastructure. Hardware must NEVER
be left running unattended after a connection failure.
"""
from __future__ import annotations
import asyncio
import json
import logging
import time
from . import config
from .session_registry import registry
log = logging.getLogger("signal_bridge.safety")
class DeadManSwitch:
"""
Periodic heartbeat monitor for all active phone sessions.
Every HEARTBEAT_INTERVAL_S seconds:
1. Send a heartbeat_ping to each phone
2. Check if any phones missed their last heartbeat by > HEARTBEAT_TIMEOUT_S
3. If so, send emergency stop and disconnect
The phone relay client responds to pings with pongs.
The session registry tracks last_heartbeat timestamps.
"""
def __init__(self):
self._task: asyncio.Task | None = None
self._running = False
async def start(self):
"""Start the heartbeat monitor loop."""
if self._running:
return
self._running = True
self._task = asyncio.create_task(self._monitor_loop())
log.info(
f"Dead man's switch active: ping every {config.HEARTBEAT_INTERVAL_S}s, "
f"timeout after {config.HEARTBEAT_TIMEOUT_S}s"
)
async def stop(self):
"""Stop the monitor."""
self._running = False
if self._task:
self._task.cancel()
try:
await self._task
except asyncio.CancelledError:
pass
async def _monitor_loop(self):
while self._running:
try:
await self._check_all()
except Exception as e:
log.error(f"Heartbeat monitor error: {e}")
await asyncio.sleep(config.HEARTBEAT_INTERVAL_S)
async def _check_all(self):
now = time.time()
sessions = await registry.get_all_sessions()
for user_id, session in sessions.items():
# Send ping
ping = {"type": "heartbeat_ping", "timestamp": now}
try:
await session.websocket.send(json.dumps(ping))
except Exception:
# Can't even send — connection dead
log.warning(f"DEAD MAN'S SWITCH: Cannot reach phone for user {user_id}")
await self._emergency_stop(user_id, session)
continue
# Check if last pong is too old
elapsed = now - session.last_heartbeat
if elapsed > config.HEARTBEAT_TIMEOUT_S:
log.warning(
f"DEAD MAN'S SWITCH: Phone heartbeat timeout for user {user_id} "
f"({elapsed:.1f}s since last pong)"
)
await self._emergency_stop(user_id, session)
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")
try:
stop_cmd = {"type": "stop", "device": "all", "emergency": True}
await session.websocket.send(json.dumps(stop_cmd))
except Exception:
pass # best effort — the phone client also has its own local failsafe
try:
await session.websocket.close(1001, "Heartbeat timeout — emergency stop")
except Exception:
pass
await registry.unregister(user_id)
# Singleton
dead_man_switch = DeadManSwitch()

171
server/session_registry.py Normal file
View File

@@ -0,0 +1,171 @@
"""
Signal Bridge Remote — Session Registry
Maps authenticated users to their active phone WebSocket connections.
Handles routing commands from MCP tool calls to the correct phone.
"""
from __future__ import annotations
import asyncio
import json
import logging
import time
import uuid
from dataclasses import dataclass, field
from typing import Optional, Any, Protocol, runtime_checkable
from .models import CommandAck
log = logging.getLogger("signal_bridge.sessions")
# ════════════════════════════════════════════════════════════════════════
# WebSocket Protocol — works with any WebSocket implementation
# (FastAPI wrapper, websockets library, etc.)
# ════════════════════════════════════════════════════════════════════════
@runtime_checkable
class WebSocketLike(Protocol):
"""Minimal interface for a WebSocket connection."""
async def send(self, data: str) -> None: ...
async def close(self, code: int = 1000, reason: str = "") -> None: ...
@dataclass
class PhoneSession:
"""An active connection from a user's phone."""
user_id: str
websocket: WebSocketLike
connected_at: float = field(default_factory=time.time)
last_heartbeat: float = field(default_factory=time.time)
devices: list[dict[str, Any]] = field(default_factory=list)
# Pending command acknowledgments: request_id → Future
_pending: dict[str, asyncio.Future] = field(default_factory=dict)
async def send_command(self, command: dict, timeout: float = 10.0) -> CommandAck:
"""Send a command and wait for acknowledgment."""
request_id = str(uuid.uuid4())[:8]
command["request_id"] = request_id
loop = asyncio.get_running_loop()
future: asyncio.Future[CommandAck] = loop.create_future()
self._pending[request_id] = future
try:
await self.websocket.send(json.dumps(command))
ack = await asyncio.wait_for(future, timeout=timeout)
return ack
except asyncio.TimeoutError:
return CommandAck(success=False, message="Phone did not respond in time")
finally:
self._pending.pop(request_id, None)
def resolve_ack(self, request_id: str, ack: CommandAck):
"""Called when the phone sends a command_ack."""
future = self._pending.get(request_id)
if future and not future.done():
future.set_result(ack)
async def send_fire_and_forget(self, command: dict):
"""Send without waiting for ack (used for heartbeats, stops)."""
try:
await self.websocket.send(json.dumps(command))
except Exception:
pass # connection probably dead, heartbeat will catch it
class SessionRegistry:
"""
Central registry mapping users to their active phone sessions.
Thread-safe via asyncio locks.
"""
def __init__(self):
self._sessions: dict[str, PhoneSession] = {} # user_id → session
self._lock = asyncio.Lock()
async def register(self, user_id: str, websocket: WebSocketLike) -> PhoneSession:
"""Register a new phone connection for a user."""
async with self._lock:
# Close existing session if any (phone reconnected)
old = self._sessions.get(user_id)
if old:
log.info(f"Replacing existing session for user {user_id}")
try:
await old.websocket.close(1000, "Replaced by new connection")
except Exception:
pass
session = PhoneSession(user_id=user_id, websocket=websocket)
self._sessions[user_id] = session
log.info(f"Phone connected: user={user_id}")
return session
async def unregister(self, user_id: str):
"""Remove a phone session."""
async with self._lock:
session = self._sessions.pop(user_id, None)
if session:
log.info(f"Phone disconnected: user={user_id}")
async def get_session(self, user_id: str) -> Optional[PhoneSession]:
"""Get the active phone session for a user."""
async with self._lock:
return self._sessions.get(user_id)
async def send_to_user(
self, user_id: str, command: dict, wait_ack: bool = True
) -> CommandAck:
"""Route a command to a user's phone. Returns ack."""
session = await self.get_session(user_id)
if not session:
return CommandAck(
success=False,
message="No phone connected. Open Intiface and connect to the relay.",
)
if wait_ack:
return await session.send_command(command)
else:
await session.send_fire_and_forget(command)
return CommandAck(success=True, message="Sent (no ack requested)")
async def update_heartbeat(self, user_id: str):
"""Mark that a heartbeat pong was received."""
async with self._lock:
session = self._sessions.get(user_id)
if session:
session.last_heartbeat = time.time()
async def update_devices(self, user_id: str, devices: list[dict]):
"""Update the device list for a user's session."""
async with self._lock:
session = self._sessions.get(user_id)
if session:
session.devices = devices
async def get_devices(self, user_id: str) -> list[dict]:
"""Get device list for a user."""
session = await self.get_session(user_id)
return session.devices if session else []
async def get_all_sessions(self) -> dict[str, PhoneSession]:
"""Get snapshot of all sessions (for heartbeat monitor)."""
async with self._lock:
return dict(self._sessions)
async def get_sole_user_id(self) -> Optional[str]:
"""If exactly one phone session is active, return its user_id.
Used for authless MCP access (e.g. claude.ai connector)."""
async with self._lock:
if len(self._sessions) == 1:
return next(iter(self._sessions))
return None
@property
def active_count(self) -> int:
return len(self._sessions)
# Singleton
registry = SessionRegistry()