mirror of
https://github.com/AletheiaVox/signal_bridge_remote.git
synced 2026-10-07 03:18:17 +08:00
Add files via upload
This commit is contained in:
1
server/__init__.py
Normal file
1
server/__init__.py
Normal file
@@ -0,0 +1 @@
|
||||
# Signal Bridge Remote — Server Package
|
||||
BIN
server/__pycache__/__init__.cpython-310.pyc
Normal file
BIN
server/__pycache__/__init__.cpython-310.pyc
Normal file
Binary file not shown.
BIN
server/__pycache__/app.cpython-310.pyc
Normal file
BIN
server/__pycache__/app.cpython-310.pyc
Normal file
Binary file not shown.
BIN
server/__pycache__/auth.cpython-310.pyc
Normal file
BIN
server/__pycache__/auth.cpython-310.pyc
Normal file
Binary file not shown.
BIN
server/__pycache__/config.cpython-310.pyc
Normal file
BIN
server/__pycache__/config.cpython-310.pyc
Normal file
Binary file not shown.
BIN
server/__pycache__/mcp_tools.cpython-310.pyc
Normal file
BIN
server/__pycache__/mcp_tools.cpython-310.pyc
Normal file
Binary file not shown.
BIN
server/__pycache__/models.cpython-310.pyc
Normal file
BIN
server/__pycache__/models.cpython-310.pyc
Normal file
Binary file not shown.
BIN
server/__pycache__/relay_hub.cpython-310.pyc
Normal file
BIN
server/__pycache__/relay_hub.cpython-310.pyc
Normal file
Binary file not shown.
BIN
server/__pycache__/safety.cpython-310.pyc
Normal file
BIN
server/__pycache__/safety.cpython-310.pyc
Normal file
Binary file not shown.
BIN
server/__pycache__/session_registry.cpython-310.pyc
Normal file
BIN
server/__pycache__/session_registry.cpython-310.pyc
Normal file
Binary file not shown.
494
server/app.py
Normal file
494
server/app.py
Normal 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
226
server/auth.py
Normal 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
46
server/config.py
Normal 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
373
server/mcp_tools.py
Normal 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
123
server/models.py
Normal 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
53
server/relay_hub.py
Normal 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
110
server/safety.py
Normal 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
171
server/session_registry.py
Normal 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()
|
||||
Reference in New Issue
Block a user