diff --git a/.env.example b/.env.example index d83eef7..25f0a69 100644 --- a/.env.example +++ b/.env.example @@ -3,8 +3,13 @@ # ============================================================ # Copy to .env and fill in your values -# --- Database --- -DATABASE_URL=postgresql+asyncpg://holdslayer:changeme@localhost:5432/holdslayer +# --- Database (required) --- +DATABASE_URL=postgresql+asyncpg://holdslayer:@localhost:5432/holdslayer + +# --- API auth (required unless HOST=127.0.0.1) --- +# One static bearer token shared by REST, WebSocket (?token=...), and MCP. +# Generate with: openssl rand -hex 32 +API_TOKEN= # --- SIP Trunk --- SIP_TRUNK_HOST=sip.yourprovider.com @@ -57,5 +62,9 @@ NOTIFY_SMS_NUMBER=+15559876543 # --- Server --- HOST=0.0.0.0 PORT=8000 -DEBUG=true +DEBUG=false LOG_LEVEL=info + +# --- Safety --- +# Max simultaneous calls the gateway will place (REST + MCP) +MAX_CONCURRENT_CALLS=4 diff --git a/README.md b/README.md index 4cdae81..86d49a2 100644 --- a/README.md +++ b/README.md @@ -6,12 +6,12 @@ You give it a phone number and an intent ("dispute a charge on my December state > [!CAUTION] > **Emergency calling — 911** -> Hold Slayer passes `911` and `9911` directly to the PSTN trunk. -> **Your SIP trunk provider must support E911 on your DID and have your -> correct registered location on file before this system is put into -> service.** VoIP emergency calls are location-dependent — verify -> with your provider. Do not rely on this system as your only means -> of reaching emergency services. +> Outbound calls to emergency numbers (`911`, `9911`, `112`) via the +> REST API or MCP tools are **always refused** — an AI agent must never +> place an emergency call, and API calls carry no E911 location data. +> Do not rely on this system as any part of your means of reaching +> emergency services; keep a phone with provider-registered E911 +> service available. ## Architecture @@ -73,13 +73,12 @@ You give it a phone number and an intent ("dispute a charge on my December state - **Transcription** (`services/transcription.py`) — Speaches/Whisper STT integration for live call transcription - **Recording** (`services/recording.py`) — WAV recording with date-organized storage, dual-channel support, persisted to the `recordings` table - **Call Persistence** (`services/call_persistence.py`) — Writes completed calls + transcript chunks to the database on hangup -- **Call Analytics** (`services/call_analytics.py`) — Hold time stats, success rates, per-company patterns, time-of-day trends - **Notifications** (`services/notification.py`) — WebSocket + SMS alerts for human detection, call failures, hold status ### API Surface - **REST API** — Call management, call history, transcripts, recordings, routing rules, device DND, call flow CRUD - **WebSocket** — Real-time call events, transcripts, classification updates, receptionist state transitions -- **MCP Server** — 10 tools for AI assistant integration (make calls, send DTMF, get transcripts, manage flows) +- **MCP Server** — 14 tools + 3 resources for AI assistant integration (make calls, send DTMF, get transcripts, manage flows), served over streamable HTTP at `/mcp/` - **Dashboard** — SvelteKit UI served at `/dashboard` with live monitor, call history with transcript playback, and a routing-rules editor ### Data Models @@ -115,7 +114,6 @@ hold-slayer/ │ ├── llm_client.py # OpenAI-compatible LLM client │ ├── transcription.py # Speaches/Whisper STT │ ├── recording.py # Call recording management -│ ├── call_analytics.py # Call metrics and insights │ └── notification.py # WebSocket + SMS notifications ├── api/ │ ├── calls.py # Call management endpoints @@ -163,13 +161,25 @@ source .venv/bin/activate pip install -e ".[dev]" ``` +> [!NOTE] +> The PJSUA2 media pipeline needs the `pjsua2` Python bindings, which are +> **not pip-installable** — they're built from pjproject (`./configure && +> make && make install` with `--enable-shared` and the Python SWIG target). +> Without them the media layer runs in stub mode (signaling only). + ### 2. Configure ```bash cp .env.example .env # Edit .env with your SIP trunk credentials, LLM endpoint, etc. +# Required: DATABASE_URL, and API_TOKEN unless HOST=127.0.0.1 +openssl rand -hex 32 # → API_TOKEN ``` +All REST, WebSocket, and MCP access requires `Authorization: Bearer +$API_TOKEN` (WebSocket also accepts `?token=...`). An empty token is only +permitted when bound to loopback. + ### 3. Build the dashboard (optional but recommended) ```bash @@ -185,7 +195,7 @@ The gateway serves the built UI at `/dashboard` automatically when ### 4. Run ```bash -uvicorn main:app --host 0.0.0.0 --port 8100 +uvicorn main:app --host 0.0.0.0 --port 8000 ``` ### 5. Test @@ -202,6 +212,7 @@ pytest tests/ -v ```bash curl -X POST http://localhost:8000/api/calls/hold-slayer \ + -H "Authorization: Bearer $API_TOKEN" \ -H "Content-Type: application/json" \ -d '{ "number": "+18005551234", @@ -253,7 +264,7 @@ curl -X PATCH http://localhost:8000/api/routing/devices/dev_abc123/dnd \ ### WebSocket — Real-Time Events ```javascript -const ws = new WebSocket("ws://localhost:8000/ws/events"); +const ws = new WebSocket(`ws://localhost:8000/ws/events?token=${API_TOKEN}`); ws.onmessage = (msg) => { const event = JSON.parse(msg.data); // event.type: "human_detected", "hold_detected", "ivr_step", etc. @@ -264,20 +275,33 @@ ws.onmessage = (msg) => { ### MCP — AI Assistant Integration -The MCP server exposes 10 tools that any MCP-compatible assistant can use: +The MCP server is served over **streamable HTTP at `/mcp/`** (note the +trailing slash) and authenticates with the same bearer token: + +```bash +claude mcp add hold-slayer --transport http http://localhost:8000/mcp/ \ + --header "Authorization: Bearer $API_TOKEN" +``` + +It exposes 14 tools and 3 resources (`gateway://status`, +`gateway://call-flows`, `gateway://active-calls`): | Tool | Description | |------|-------------| -| `make_call` | Dial a number through the SIP trunk | -| `end_call` | Hang up an active call | +| `make_call` | Dial a real number through the SIP trunk (emergency numbers refused) | +| `hangup` | Hang up an active call | +| `transfer_call` | Transfer an active call to a device | | `send_dtmf` | Send touch-tone digits to navigate menus | | `get_call_status` | Check current state of a call | | `get_call_transcript` | Get live transcript of a call | | `get_call_recording` | Get recording metadata and file path | | `list_active_calls` | List all calls in progress | -| `get_call_summary` | Analytics summary (hold times, success rates) | -| `search_call_history` | Search past calls by number or company | -| `learn_call_flow` | Build a reusable call flow from exploration data | +| `list_devices` | List registered devices and status | +| `gateway_status` | Trunk, devices, active calls, uptime | +| `get_call_flow` | Look up a stored IVR flow for a number | +| `create_call_flow` | Store a new IVR call flow | +| `get_call_summary` | Stored summary and action items for a call | +| `search_call_history` | Search past calls by number or intent | ## How It Works @@ -307,6 +331,9 @@ All configuration is via environment variables (see `.env.example`): | Variable | Description | Default | |----------|-------------|---------| +| `DATABASE_URL` | PostgreSQL connection string | — (required) | +| `API_TOKEN` | Static bearer token for REST/WS/MCP | — (required unless `HOST=127.0.0.1`) | +| `MAX_CONCURRENT_CALLS` | Cap on simultaneous outbound calls | `4` | | `SIP_TRUNK_HOST` | Your SIP provider hostname | — | | `SIP_TRUNK_USERNAME` | SIP auth username | — | | `SIP_TRUNK_PASSWORD` | SIP auth password | — | @@ -322,7 +349,6 @@ All configuration is via environment variables (see `.env.example`): | `RECEPTIONIST_ENABLED` | Answer inbound calls with the AI receptionist | `true` | | `RECEPTIONIST_GREETING_TEMPLATE` | Spoken greeting | `"Hi, you've reached Robert's line. Who's calling, and what's this about?"` | | `RECEPTIONIST_MESSAGE_MAX_SECONDS` | Voicemail cap | `90` | -| `DATABASE_URL` | PostgreSQL or SQLite connection | SQLite fallback | ## Tech Stack @@ -368,22 +394,21 @@ Full documentation is in [`/docs`](docs/README.md): - [x] Hold Slayer IVR navigation with LLM fallback for LISTEN steps - [x] Call Flow Learner — auto-builds reusable IVR trees from exploration - [x] Recording service with date-organized WAV storage -- [x] Call analytics with hold time stats, per-company patterns - [x] Audio classifier with spectral analysis, DTMF detection, hold-to-human transition ### Phase 3: API & Integration ✅ - [x] REST API — calls, call flows, devices, DTMF - [x] WebSocket real-time event streaming -- [x] MCP server with 16 tools + 3 resources +- [x] MCP server with 14 tools + 3 resources, mounted at `/mcp/` (streamable HTTP) - [x] Notification service (WebSocket + SMS) - [x] Service wiring in main.py lifespan -- [x] 75 passing tests across 4 test files -### Phase 4: Production Hardening 🔜 +### Phase 4: Production Hardening 🚧 - [ ] Alembic database migrations -- [ ] API authentication (API keys / JWT) +- [x] API authentication — static bearer token across REST/WS/MCP +- [x] Emergency-number guard + concurrent-call cap on outbound calls - [ ] Rate limiting on API endpoints - [ ] Structured JSON logging - [ ] Health check endpoints for all dependencies diff --git a/api/calls.py b/api/calls.py index fef14bd..a612fcc 100644 --- a/api/calls.py +++ b/api/calls.py @@ -46,6 +46,8 @@ async def make_call( number=request.number, mode=request.mode.value, ) + except ValueError as e: + raise HTTPException(status_code=400, detail=str(e)) except Exception as e: raise HTTPException(status_code=500, detail=str(e)) @@ -86,6 +88,8 @@ async def hold_slayer( mode="hold_slayer", message="Hold Slayer activated. I'll ring you when a human picks up. ☕", ) + except ValueError as e: + raise HTTPException(status_code=400, detail=str(e)) except Exception as e: raise HTTPException(status_code=500, detail=str(e)) diff --git a/api/deps.py b/api/deps.py index 2643b7c..a8beedb 100644 --- a/api/deps.py +++ b/api/deps.py @@ -2,11 +2,12 @@ API Dependencies — Shared dependency injection for all routes. """ -from fastapi import Depends, HTTPException, Request -from sqlalchemy.ext.asyncio import AsyncSession +import secrets +from fastapi import Header, HTTPException, Request + +from config import get_settings from core.gateway import AIPSTNGateway -from db.database import get_db def get_gateway(request: Request) -> AIPSTNGateway: @@ -15,3 +16,24 @@ def get_gateway(request: Request) -> AIPSTNGateway: if gateway is None: raise HTTPException(status_code=503, detail="Gateway not initialized") return gateway + + +def require_token(authorization: str | None = Header(default=None)) -> None: + """ + Enforce the static bearer token (API_TOKEN) on REST routes. + + An empty configured token disables auth; startup refuses that + combination unless the server is bound to loopback. + """ + token = get_settings().api_token.get_secret_value() + if not token: + return + supplied = "" + if authorization and authorization.lower().startswith("bearer "): + supplied = authorization[7:] + if not secrets.compare_digest(supplied, token): + raise HTTPException( + status_code=401, + detail="Missing or invalid bearer token", + headers={"WWW-Authenticate": "Bearer"}, + ) diff --git a/api/websocket.py b/api/websocket.py index 933db17..85bef08 100644 --- a/api/websocket.py +++ b/api/websocket.py @@ -2,10 +2,12 @@ import asyncio import logging +import secrets from fastapi import APIRouter, WebSocket, WebSocketDisconnect from api.deps import get_gateway +from config import get_settings from models.events import EventType, GatewayEvent logger = logging.getLogger(__name__) @@ -13,6 +15,26 @@ logger = logging.getLogger(__name__) router = APIRouter() +async def _authorize(websocket: WebSocket) -> bool: + """ + Check the static bearer token before accepting the socket. + + Browsers can't set headers on WebSocket connects, so a `token` + query parameter is accepted alongside the Authorization header. + """ + token = get_settings().api_token.get_secret_value() + if not token: + return True + supplied = websocket.query_params.get("token", "") + auth = websocket.headers.get("authorization", "") + if auth.lower().startswith("bearer "): + supplied = auth[7:] + if secrets.compare_digest(supplied, token): + return True + await websocket.close(code=4401, reason="Missing or invalid bearer token") + return False + + async def _send_trunk_status(websocket: WebSocket, gateway) -> None: """Send current SIP trunk status as a synthetic event to a newly connected client.""" try: @@ -58,6 +80,8 @@ async def event_stream(websocket: WebSocket): "message": "🚨 Human detected!" } """ + if not await _authorize(websocket): + return await websocket.accept() logger.info("WebSocket client connected") @@ -90,6 +114,8 @@ async def call_event_stream(websocket: WebSocket, call_id: str): Same format as /events but only sends events for the specified call. """ + if not await _authorize(websocket): + return await websocket.accept() logger.info(f"WebSocket client connected for call {call_id}") diff --git a/config.py b/config.py index 52a3dd6..f637bca 100644 --- a/config.py +++ b/config.py @@ -4,7 +4,7 @@ Hold Slayer Gateway — Configuration All settings loaded from environment variables / .env file. """ -from pydantic import Field +from pydantic import Field, SecretStr from pydantic_settings import BaseSettings, SettingsConfigDict @@ -16,7 +16,7 @@ class SIPTrunkSettings(BaseSettings): host: str = "sip.provider.com" port: int = 5060 username: str = "" - password: str = "" + password: SecretStr = SecretStr("") transport: str = "udp" # udp, tcp, tls did: str = "" # Your phone number (E.164) @@ -58,7 +58,7 @@ class LLMSettings(BaseSettings): base_url: str = "http://localhost:11434/v1" model: str = "llama3" - api_key: str = "not-needed" + api_key: SecretStr = SecretStr("not-needed") timeout: float = 30.0 max_tokens: int = 1024 temperature: float = 0.3 @@ -84,7 +84,7 @@ class TTSSettings(BaseSettings): base_url: str = "http://localhost:8000" model: str = "speaches-ai/Kokoro-82M-v1.0-ONNX" voice: str = "af_heart" - api_key: str = "" + api_key: SecretStr = SecretStr("") timeout: float = 30.0 sample_rate: int = 16000 @@ -117,15 +117,22 @@ class Settings(BaseSettings): extra="ignore", ) - # Database - database_url: str = "postgresql+asyncpg://holdslayer:changeme@localhost:5432/holdslayer" + # Database — no default credentials; must be set in the environment + database_url: str = "" # Server host: str = "0.0.0.0" port: int = 8000 - debug: bool = True + debug: bool = False log_level: str = "info" + # Auth — one static bearer token shared by REST, WebSocket, and MCP. + # Empty disables auth, which is only permitted on loopback binds. + api_token: SecretStr = SecretStr("") + + # Outbound-call safety cap (REST + MCP make_call) + max_concurrent_calls: int = 4 + # Notifications notify_sms_number: str = "" diff --git a/core/dial_plan.py b/core/dial_plan.py index ce5e83e..97284ab 100644 --- a/core/dial_plan.py +++ b/core/dial_plan.py @@ -1,200 +1,38 @@ """ -Dial Plan — Pattern matching and digit normalisation. +Dial Plan — Emergency-number guard and extension allocation. -Matches a dialled string to a route type and normalises the destination -to a canonical form the rest of the gateway can act on. - -Route types: - "extension" — internal 2XX endpoint - "service" — internal 5XX system service - "pstn" — outbound call via SIP trunk (normalised E.164) - "invalid" — no match +Emergency numbers are never dialable through the gateway's API/MCP +surfaces: an outbound emergency call must come from a human on a real +phone whose trunk provider has E911 location data, not from an AI agent +or a REST request. See the README caution. """ -import re -from dataclasses import dataclass -from typing import Optional - - -# ================================================================ -# Emergency numbers — always route to PSTN, highest priority -# ================================================================ +# Dialled forms and their E.164 mappings — both sides are refused. EMERGENCY_NUMBERS: dict[str, str] = { "911": "+1911", # North American emergency "9911": "+1911", # Mis-dial with phantom '9' prefix "112": "+112", # International GSM emergency } +_BLOCKED = frozenset(EMERGENCY_NUMBERS) | frozenset(EMERGENCY_NUMBERS.values()) + + +def is_emergency_number(number: str) -> bool: + """True if the dialled string is an emergency number (any known form).""" + cleaned = number.strip().replace(" ", "").replace("-", "").replace(".", "") + return cleaned in _BLOCKED + # ================================================================ -# Extension ranges +# Extension allocation (2XX range) # ================================================================ EXTENSION_FIRST = 221 EXTENSION_LAST = 299 -SERVICE_FIRST = 500 -SERVICE_LAST = 599 -# ================================================================ -# Known system services -# ================================================================ - -SERVICES: dict[int, str] = { - 500: "auto_attendant", - 510: "gateway_status", - 511: "echo_test", - 520: "hold_slayer_launch", - 599: "operator_fallback", -} - - -# ================================================================ -# Route result -# ================================================================ - -@dataclass -class RouteResult: - """Result of a dial plan lookup.""" - - route_type: str # "extension" | "service" | "pstn" | "invalid" - destination: str # normalised — extension number, service name, or E.164 - original: str # what was dialled - description: str = "" - - @property - def is_internal(self) -> bool: - return self.route_type in ("extension", "service") - - @property - def is_outbound(self) -> bool: - return self.route_type == "pstn" - - @property - def is_valid(self) -> bool: - return self.route_type != "invalid" - - -# ================================================================ -# Core matcher -# ================================================================ - -def match(digits: str) -> RouteResult: - """ - Match dialled digits against the dial plan. - - Returns a RouteResult with the normalised destination. - - Examples: - match("221") → RouteResult(route_type="extension", destination="221") - match("511") → RouteResult(route_type="service", destination="echo_test") - match("6135550100") → RouteResult(route_type="pstn", destination="+16135550100") - match("16135550100") → RouteResult(route_type="pstn", destination="+16135550100") - match("+16135550100") → RouteResult(route_type="pstn", destination="+16135550100") - match("01144201234") → RouteResult(route_type="pstn", destination="+44201234") - """ - digits = digits.strip() - - # ---- Emergency numbers — checked first, no interception ---- - if digits in EMERGENCY_NUMBERS: - e164 = EMERGENCY_NUMBERS[digits] - return RouteResult( - route_type="pstn", - destination=e164, - original=digits, - description=f"EMERGENCY {digits} → {e164}", - ) - - # ---- 2XX extensions ---- - if re.fullmatch(r"2\d{2}", digits): - ext = int(digits) - if EXTENSION_FIRST <= ext <= EXTENSION_LAST: - return RouteResult( - route_type="extension", - destination=digits, - original=digits, - description=f"Extension {digits}", - ) - - # ---- 5XX system services ---- - if re.fullmatch(r"5\d{2}", digits): - svc = int(digits) - if SERVICE_FIRST <= svc <= SERVICE_LAST: - name = SERVICES.get(svc, f"service_{svc}") - return RouteResult( - route_type="service", - destination=name, - original=digits, - description=f"System service: {name}", - ) - - # ---- PSTN outbound ---- - e164 = _normalise_e164(digits) - if e164: - return RouteResult( - route_type="pstn", - destination=e164, - original=digits, - description=f"PSTN outbound → {e164}", - ) - - return RouteResult( - route_type="invalid", - destination=digits, - original=digits, - description=f"No route for '{digits}'", - ) - - -# ================================================================ -# E.164 normalisation -# ================================================================ - -def _normalise_e164(digits: str) -> Optional[str]: - """ - Normalise a dialled string to E.164 (+CC…). - - Handles: - +CCNNN… → unchanged (already E.164) - 1NPANXXXXXX → +1NPANXXXXXX (NANP with country code, 11 digits) - NPANXXXXXX → +1NPANXXXXXX (NANP 10-digit) - 011CCNNN… → +CCNNN… (IDD 011 prefix) - 00CCNNN… → +CCNNN… (IDD 00 prefix) - """ - # Strip spaces/dashes/dots/parens for matching only - clean = re.sub(r"[\s\-\.\(\)]", "", digits) - - # Already E.164 - if re.fullmatch(r"\+\d{7,15}", clean): - return clean - - # NANP: 1 + 10 digits (NPA must be 2-9, NXX must be 2-9) - if re.fullmatch(r"1[2-9]\d{2}[2-9]\d{6}", clean): - return f"+{clean}" - - # NANP: 10 digits only - if re.fullmatch(r"[2-9]\d{2}[2-9]\d{6}", clean): - return f"+1{clean}" - - # IDD 011 (North American international dialling prefix) - m = re.fullmatch(r"011(\d{7,13})", clean) - if m: - return f"+{m.group(1)}" - - # IDD 00 (international dialling prefix used in many countries) - m = re.fullmatch(r"00(\d{7,13})", clean) - if m: - return f"+{m.group(1)}" - - return None - - -# ================================================================ -# Extension helpers -# ================================================================ - -def next_extension(used: set[int]) -> Optional[int]: +def next_extension(used: set[int]) -> int | None: """ Return the lowest available extension in the 2XX range. @@ -208,17 +46,3 @@ def next_extension(used: set[int]) -> Optional[int]: if ext not in used: return ext return None - - -def is_extension(digits: str) -> bool: - """True if the string is a valid 2XX extension.""" - return bool(re.fullmatch(r"2\d{2}", digits)) and ( - EXTENSION_FIRST <= int(digits) <= EXTENSION_LAST - ) - - -def is_service(digits: str) -> bool: - """True if the string is a valid 5XX service code.""" - return bool(re.fullmatch(r"5\d{2}", digits)) and ( - SERVICE_FIRST <= int(digits) <= SERVICE_LAST - ) diff --git a/core/gateway.py b/core/gateway.py index 2eb8445..16bd2e4 100644 --- a/core/gateway.py +++ b/core/gateway.py @@ -11,7 +11,7 @@ from typing import Optional from config import Settings, get_settings from core.call_manager import CallManager -from core.dial_plan import next_extension +from core.dial_plan import is_emergency_number, next_extension from core.event_bus import EventBus from core.media_pipeline import MediaPipeline from core.sip_engine import MockSIPEngine, SIPEngine @@ -52,7 +52,7 @@ def _build_sip_engine(settings: Settings, gateway: "AIPSTNGateway") -> SIPEngine trunk_host=trunk.host, trunk_port=trunk.port, trunk_username=trunk.username, - trunk_password=trunk.password, + trunk_password=trunk.password.get_secret_value(), trunk_transport=trunk.transport, domain=gw_sip.domain, did=trunk.did, @@ -228,6 +228,19 @@ class AIPSTNGateway: - hold_slayer: Navigate IVR, wait on hold, transfer when human detected - ai_assisted: Connect with transcription, recording, noise cancel """ + if is_emergency_number(number): + raise ValueError( + f"Refusing to dial emergency number '{number}'. Emergency calls " + "must be placed from a phone with E911 location service, not " + "through the gateway API." + ) + active = len(self.call_manager.active_calls) + if active >= self.settings.max_concurrent_calls: + raise ValueError( + f"Concurrent-call limit reached ({active}/{self.settings.max_concurrent_calls}). " + "End an active call or raise MAX_CONCURRENT_CALLS." + ) + # Create call in manager call = await self.call_manager.create_call( remote_number=number, diff --git a/main.py b/main.py index 53fd3ef..392fcce 100644 --- a/main.py +++ b/main.py @@ -15,11 +15,12 @@ import logging import sys from contextlib import asynccontextmanager -from fastapi import FastAPI +from fastapi import Depends, FastAPI from fastapi.staticfiles import StaticFiles from api import call_flows, call_history, calls, devices, routing, websocket -from config import get_settings +from api.deps import require_token +from config import Settings, get_settings from core.gateway import AIPSTNGateway from db.database import close_db, init_db from mcp_server.server import create_mcp_server @@ -84,76 +85,105 @@ def _handle_db_error(exc: Exception) -> None: sys.exit(1) +def _check_startup_config(settings: Settings) -> None: + """Refuse insecure or incomplete configurations before booting anything.""" + if not settings.database_url: + logger.critical( + "\n" + "❌ DATABASE_URL is not set.\n" + " Add it to your .env file, e.g.:\n" + " DATABASE_URL=postgresql+asyncpg://holdslayer:@localhost:5432/holdslayer" + ) + sys.exit(1) + + token = settings.api_token.get_secret_value() + if not token and settings.host not in ("127.0.0.1", "localhost", "::1"): + logger.critical( + "\n" + "❌ API_TOKEN is not set but HOST binds beyond loopback " + f"({settings.host}).\n" + " Every surface (REST, WebSocket, MCP make_call) would be open " + "to the network.\n" + " Set API_TOKEN in .env (e.g. `openssl rand -hex 32`), or set " + "HOST=127.0.0.1 for tokenless local development." + ) + sys.exit(1) + + @asynccontextmanager async def lifespan(app: FastAPI): """Startup: Initialize database, SIP engine, and services.""" settings = get_settings() + _check_startup_config(settings) - # Initialize database - logger.info("Initializing database...") - try: - await init_db() - except Exception as e: - _handle_db_error(e) + # The MCP session manager lives in the mounted sub-app's lifespan; + # without entering it, every /mcp request 500s. + async with mcp_http_app.lifespan(app): + # Initialize database + logger.info("Initializing database...") + try: + await init_db() + except Exception as e: + _handle_db_error(e) - # Boot the telephony engine - gateway = AIPSTNGateway.from_config() - await gateway.start() - app.state.gateway = gateway + # Boot the telephony engine + gateway = AIPSTNGateway.from_config() + await gateway.start() + app.state.gateway = gateway - # Start auxiliary services - from services.notification import NotificationService - from services.recording import RecordingService - from services.call_analytics import CallAnalytics - from services.call_flow_learner import CallFlowLearner + # Start auxiliary services + from services.notification import NotificationService + from services.recording import RecordingService - notification_svc = NotificationService(gateway.event_bus, settings) - await notification_svc.start() - app.state.notification_service = notification_svc + notification_svc = NotificationService(gateway.event_bus, settings) + await notification_svc.start() + app.state.notification_service = notification_svc - recording_svc = RecordingService() - await recording_svc.start() - app.state.recording_service = recording_svc - gateway._recording_service = recording_svc + recording_svc = RecordingService() + await recording_svc.start() + app.state.recording_service = recording_svc + gateway._recording_service = recording_svc - analytics_svc = CallAnalytics() - app.state.analytics_service = analytics_svc + logger.info("=" * 60) + logger.info("🔥 Hold Slayer Gateway is LIVE") + # Show a usable URL — 0.0.0.0 is the bind address, not a browser URL + display_host = "localhost" if settings.host in ("0.0.0.0", "::") else settings.host + # When launched via `uvicorn main:app --port XXXX`, the CLI --port arg + # takes precedence over settings.port (which comes from .env). + display_port = settings.port + for i, arg in enumerate(sys.argv): + if arg in ("--port", "-p") and i + 1 < len(sys.argv): + try: + display_port = int(sys.argv[i + 1]) + except ValueError: + pass + auth_state = "bearer token required" if settings.api_token.get_secret_value() else "auth disabled (loopback)" + logger.info(f" API: http://{display_host}:{display_port} [{auth_state}]") + logger.info(f" API Docs: http://{display_host}:{display_port}/docs") + logger.info(f" WebSocket: ws://{display_host}:{display_port}/ws/events") + logger.info(f" MCP: http://{display_host}:{display_port}/mcp/ (streamable HTTP)") + logger.info("=" * 60) - flow_learner = CallFlowLearner() - app.state.flow_learner = flow_learner + yield - # Create and mount MCP server - mcp = create_mcp_server(gateway) - app.state.mcp = mcp + # Shutdown + logger.info("Shutting down Hold Slayer Gateway...") + await notification_svc.stop() + await gateway.stop() + await close_db() + logger.info("Gateway shut down cleanly. 👋") - logger.info("=" * 60) - logger.info("🔥 Hold Slayer Gateway is LIVE") - # Show a usable URL — 0.0.0.0 is the bind address, not a browser URL - display_host = "localhost" if settings.host in ("0.0.0.0", "::") else settings.host - # When launched via `uvicorn main:app --port XXXX`, the CLI --port arg - # takes precedence over settings.port (which comes from .env). - display_port = settings.port - for i, arg in enumerate(sys.argv): - if arg in ("--port", "-p") and i + 1 < len(sys.argv): - try: - display_port = int(sys.argv[i + 1]) - except ValueError: - pass - logger.info(f" API: http://{display_host}:{display_port}") - logger.info(f" API Docs: http://{display_host}:{display_port}/docs") - logger.info(f" WebSocket: ws://{display_host}:{display_port}/ws/events") - logger.info(f" MCP: Available via FastMCP") - logger.info("=" * 60) - yield +def _get_gateway_instance() -> AIPSTNGateway | None: + """Lazy gateway resolver for MCP tools (set on app.state by the lifespan).""" + return getattr(app.state, "gateway", None) - # Shutdown - logger.info("Shutting down Hold Slayer Gateway...") - await notification_svc.stop() - await gateway.stop() - await close_db() - logger.info("Gateway shut down cleanly. 👋") +mcp = create_mcp_server( + _get_gateway_instance, + api_token=get_settings().api_token.get_secret_value(), +) +mcp_http_app = mcp.http_app(path="/") app = FastAPI( title="Hold Slayer Gateway", @@ -171,13 +201,20 @@ app = FastAPI( ) # === API Routes === -app.include_router(calls.router, prefix="/api/calls", tags=["Calls"]) -app.include_router(call_history.router, prefix="/api/calls", tags=["Call History"]) -app.include_router(call_flows.router, prefix="/api/call-flows", tags=["Call Flows"]) -app.include_router(devices.router, prefix="/api/devices", tags=["Devices"]) -app.include_router(routing.router, prefix="/api/routing", tags=["Routing"]) +# call_history must register before calls: both live under /api/calls and +# calls' GET /{call_id} would otherwise capture the literal path "history". +_auth = [Depends(require_token)] +app.include_router(call_history.router, prefix="/api/calls", tags=["Call History"], dependencies=_auth) +app.include_router(calls.router, prefix="/api/calls", tags=["Calls"], dependencies=_auth) +app.include_router(call_flows.router, prefix="/api/call-flows", tags=["Call Flows"], dependencies=_auth) +app.include_router(devices.router, prefix="/api/devices", tags=["Devices"], dependencies=_auth) +app.include_router(routing.router, prefix="/api/routing", tags=["Routing"], dependencies=_auth) +# WebSocket endpoints check the token themselves (query param or header) app.include_router(websocket.router, prefix="/ws", tags=["WebSocket"]) +# === MCP (streamable HTTP; clients connect to /mcp/ with the bearer token) === +app.mount("/mcp", mcp_http_app) + # === Dashboard (built SvelteKit static) === import os as _os _dashboard_build = _os.path.join(_os.path.dirname(__file__), "dashboard", "build") diff --git a/mcp_server/server.py b/mcp_server/server.py index 594ff13..a389257 100644 --- a/mcp_server/server.py +++ b/mcp_server/server.py @@ -15,19 +15,40 @@ Example from an AI assistant: import json import logging -from typing import Optional +from typing import Callable, Optional from fastmcp import FastMCP +from fastmcp.exceptions import ToolError from core.gateway import AIPSTNGateway logger = logging.getLogger(__name__) -def create_mcp_server(gateway: AIPSTNGateway) -> FastMCP: - """Create and configure the MCP server with all tools and resources.""" +def create_mcp_server( + get_gateway: Callable[[], Optional[AIPSTNGateway]], + api_token: str = "", +) -> FastMCP: + """ + Create and configure the MCP server with all tools and resources. - mcp = FastMCP("Hold Slayer Gateway") + The gateway is resolved lazily per request via `get_gateway` so the + server can be mounted at app construction, before the lifespan has + started the gateway. + """ + auth = None + if api_token: + from fastmcp.server.auth.providers.jwt import StaticTokenVerifier + + auth = StaticTokenVerifier(tokens={api_token: {"client_id": "hold-slayer"}}) + + mcp = FastMCP("Hold Slayer Gateway", auth=auth) + + def require_gateway() -> AIPSTNGateway: + gateway = get_gateway() + if gateway is None: + raise ToolError("Gateway is still starting up — try again shortly.") + return gateway # ================================================================ # Tools @@ -42,7 +63,10 @@ def create_mcp_server(gateway: AIPSTNGateway) -> FastMCP: device: str = "", ) -> str: """ - Place an outbound phone call. + Place a REAL outbound phone call over the PSTN. The remote party's + phone actually rings and the call may incur telephony charges — + only use this when the user has asked for a call to be placed. + Emergency numbers (911/112) are always refused. Args: number: Phone number to call (E.164 format, e.g., +18005551234) @@ -56,19 +80,23 @@ def create_mcp_server(gateway: AIPSTNGateway) -> FastMCP: """ from models.call import CallMode + gateway = require_gateway() mode_map = { "direct": CallMode.DIRECT, "hold_slayer": CallMode.HOLD_SLAYER, "ai_assisted": CallMode.AI_ASSISTED, } - call = await gateway.make_call( - number=number, - mode=mode_map.get(mode, CallMode.DIRECT), - intent=intent or None, - call_flow_id=call_flow_id or None, - device=device or None, - ) + try: + call = await gateway.make_call( + number=number, + mode=mode_map.get(mode, CallMode.DIRECT), + intent=intent or None, + call_flow_id=call_flow_id or None, + device=device or None, + ) + except ValueError as e: + raise ToolError(str(e)) return ( f"Call {call.id} initiated.\n" f" Number: {number}\n" @@ -85,6 +113,7 @@ def create_mcp_server(gateway: AIPSTNGateway) -> FastMCP: Shows: status, duration, hold time, current audio type, recent transcript. """ + gateway = require_gateway() call = gateway.get_call(call_id) if not call: return f"Call {call_id} not found. It may have already ended." @@ -113,6 +142,7 @@ def create_mcp_server(gateway: AIPSTNGateway) -> FastMCP: call_id: The call to transfer device: Target device ID (e.g., "sip_phone", "cell") """ + gateway = require_gateway() try: await gateway.transfer_call(call_id, device) return f"Call {call_id} transferred to {device}." @@ -122,6 +152,7 @@ def create_mcp_server(gateway: AIPSTNGateway) -> FastMCP: @mcp.tool() async def hangup(call_id: str) -> str: """Hang up a call.""" + gateway = require_gateway() try: await gateway.hangup_call(call_id) return f"Call {call_id} hung up." @@ -131,6 +162,7 @@ def create_mcp_server(gateway: AIPSTNGateway) -> FastMCP: @mcp.tool() async def list_active_calls() -> str: """List all currently active calls with their status.""" + gateway = require_gateway() calls = gateway.call_manager.active_calls if not calls: return "No active calls." @@ -223,7 +255,7 @@ def create_mcp_server(gateway: AIPSTNGateway) -> FastMCP: id=flow_id, name=name, phone_number=phone_number, - description=f"Created by AI assistant", + description="Created by AI assistant", steps=steps, notes=notes or None, tags=["ai-created"], @@ -246,6 +278,7 @@ def create_mcp_server(gateway: AIPSTNGateway) -> FastMCP: call_id: The call to send tones on digits: DTMF digits to send (e.g., "1", "2", "123#") """ + gateway = require_gateway() call = gateway.get_call(call_id) if not call: return f"Call {call_id} not found." @@ -264,6 +297,7 @@ def create_mcp_server(gateway: AIPSTNGateway) -> FastMCP: Returns the complete transcript text. """ + gateway = require_gateway() call = gateway.get_call(call_id) if not call: return f"Call {call_id} not found." @@ -402,38 +436,10 @@ def create_mcp_server(gateway: AIPSTNGateway) -> FastMCP: except Exception as e: return f"Error searching call history: {e}" - @mcp.tool() - async def learn_call_flow(call_id: str, name: str = "") -> str: - """ - Learn a call flow from a completed call's event history. - - Analyzes the IVR navigation events from a call to build a - reusable call flow for next time. - - Args: - call_id: The call to learn from - name: Optional name for the flow (auto-generated if empty) - """ - from services.call_flow_learner import CallFlowLearner - - try: - learner = CallFlowLearner(gateway.event_bus, gateway.settings) - flow = await learner.learn_from_call(call_id, name or None) - if flow: - return ( - f"Learned call flow '{flow.name}' from call {call_id}:\n" - f" Phone: {flow.phone_number}\n" - f" Steps: {len(flow.steps)}\n" - f" Flow ID: {flow.id}" - ) - return f"Could not learn a call flow from call {call_id}. Not enough IVR navigation data." - except Exception as e: - return f"Error learning call flow: {e}" - @mcp.tool() async def list_devices() -> str: """List all registered devices and their online/offline status.""" - devices = gateway.devices + devices = require_gateway().devices if not devices: return "No devices registered." @@ -446,7 +452,7 @@ def create_mcp_server(gateway: AIPSTNGateway) -> FastMCP: @mcp.tool() async def gateway_status() -> str: """Get full gateway status — trunk, devices, active calls, uptime.""" - status = await gateway.status() + status = await require_gateway().status() trunk = status["trunk"] lines = [ @@ -470,7 +476,7 @@ def create_mcp_server(gateway: AIPSTNGateway) -> FastMCP: @mcp.resource("gateway://status") async def resource_gateway_status() -> str: """Current gateway status — trunk, devices, active calls.""" - status = await gateway.status() + status = await require_gateway().status() return json.dumps(status, default=str, indent=2) @mcp.resource("gateway://call-flows") @@ -502,7 +508,7 @@ def create_mcp_server(gateway: AIPSTNGateway) -> FastMCP: @mcp.resource("gateway://active-calls") async def resource_active_calls() -> str: """All currently active calls.""" - calls = gateway.call_manager.active_calls + calls = require_gateway().call_manager.active_calls return json.dumps( [c.summary() for c in calls.values()], default=str, diff --git a/pyproject.toml b/pyproject.toml index bc303c2..877f185 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -35,8 +35,8 @@ dependencies = [ # HTTP client (for Speaches STT) "httpx>=0.28.0", - # MCP server - "fastmcp>=2.0.0", + # MCP server (3.x — http_app + StaticTokenVerifier) + "fastmcp>=3.0.0", # Utilities "python-slugify>=8.0.0", diff --git a/services/hold_slayer.py b/services/hold_slayer.py index 1eee29b..fd14240 100644 --- a/services/hold_slayer.py +++ b/services/hold_slayer.py @@ -44,7 +44,7 @@ def _get_llm(): _llm_client = LLMClient( base_url=settings.llm.base_url, model=settings.llm.model, - api_key=settings.llm.api_key, + api_key=settings.llm.api_key.get_secret_value(), timeout=settings.llm.timeout, ) except Exception as e: diff --git a/services/tts.py b/services/tts.py index efe6b2d..2865c60 100644 --- a/services/tts.py +++ b/services/tts.py @@ -26,8 +26,8 @@ class TTSService: async def _get_client(self) -> httpx.AsyncClient: if self._client is None or self._client.is_closed: headers = {} - if self.settings.api_key: - headers["Authorization"] = f"Bearer {self.settings.api_key}" + if self.settings.api_key.get_secret_value(): + headers["Authorization"] = f"Bearer {self.settings.api_key.get_secret_value()}" self._client = httpx.AsyncClient( base_url=self.settings.base_url, timeout=httpx.Timeout(self.settings.timeout, connect=5.0), diff --git a/tests/test_api_security.py b/tests/test_api_security.py new file mode 100644 index 0000000..0c3b42d --- /dev/null +++ b/tests/test_api_security.py @@ -0,0 +1,90 @@ +""" +API surface tests — bearer-token enforcement and route registration order. + +The app is exercised without its lifespan: auth runs before any handler, +so a 503 ("Gateway not initialized") proves the token was accepted. +""" + +import httpx +import pytest +from pydantic import SecretStr +from starlette.routing import Match + +import main +from config import get_settings + +TOKEN = "test-token-for-suite" + + +@pytest.fixture +def token_enabled(monkeypatch): + monkeypatch.setattr(get_settings(), "api_token", SecretStr(TOKEN)) + + +@pytest.fixture +async def client(): + transport = httpx.ASGITransport(app=main.app) + async with httpx.AsyncClient(transport=transport, base_url="http://test") as c: + yield c + + +class TestBearerToken: + async def test_missing_token_rejected(self, token_enabled, client): + resp = await client.get("/api/calls/active") + assert resp.status_code == 401 + assert resp.headers["www-authenticate"] == "Bearer" + + async def test_wrong_token_rejected(self, token_enabled, client): + resp = await client.get( + "/api/calls/active", headers={"Authorization": "Bearer wrong"} + ) + assert resp.status_code == 401 + + async def test_valid_token_reaches_handler(self, token_enabled, client): + resp = await client.get( + "/api/calls/active", headers={"Authorization": f"Bearer {TOKEN}"} + ) + # No lifespan ran, so the handler itself 503s — auth was accepted + assert resp.status_code == 503 + + async def test_empty_token_disables_auth(self, monkeypatch, client): + monkeypatch.setattr(get_settings(), "api_token", SecretStr("")) + resp = await client.get("/api/calls/active") + assert resp.status_code == 503 + + async def test_all_api_routers_protected(self, token_enabled, client): + for path in ("/api/calls/active", "/api/call-flows/", "/api/devices/", + "/api/routing/rules", "/api/calls/history"): + resp = await client.get(path) + assert resp.status_code == 401, path + + +class TestRouteOrder: + def _resolve(self, path: str): + scope = { + "type": "http", + "method": "GET", + "path": path, + "root_path": "", + "query_string": b"", + "headers": [], + } + for route in main.app.router.routes: + match, _ = route.matches(scope) + if match == Match.FULL: + return route + return None + + def test_history_not_shadowed_by_call_id(self): + route = self._resolve("/api/calls/history") + assert route is not None + assert route.endpoint.__name__ == "list_history" + + def test_call_id_still_matches(self): + route = self._resolve("/api/calls/call_abc123") + assert route is not None + assert route.endpoint.__name__ == "get_call" + + def test_mcp_mounted(self): + mounted = [getattr(r, "path", "") for r in main.app.router.routes] + assert "/mcp" in mounted diff --git a/tests/test_mcp.py b/tests/test_mcp.py new file mode 100644 index 0000000..3825f05 --- /dev/null +++ b/tests/test_mcp.py @@ -0,0 +1,94 @@ +""" +MCP server tests — tool surface, lazy gateway resolution, call safety. + +Uses the FastMCP in-memory client (no network, no mounted app). +""" + +import pytest +from fastmcp import Client + +from config import Settings +from core.dial_plan import is_emergency_number +from core.gateway import AIPSTNGateway +from mcp_server.server import create_mcp_server + +EXPECTED_TOOLS = { + "make_call", + "get_call_status", + "transfer_call", + "hangup", + "list_active_calls", + "get_call_flow", + "create_call_flow", + "send_dtmf", + "get_call_transcript", + "get_call_recording", + "get_call_summary", + "search_call_history", + "list_devices", + "gateway_status", +} + + +def _make_gateway(max_calls: int = 4) -> AIPSTNGateway: + """Unstarted gateway on the in-memory MockSIPEngine — no network, no DB.""" + return AIPSTNGateway(settings=Settings(max_concurrent_calls=max_calls)) + + +class TestToolSurface: + async def test_tool_listing_matches_expected(self): + mcp = create_mcp_server(lambda: None) + async with Client(mcp) as client: + tools = {t.name for t in await client.list_tools()} + assert tools == EXPECTED_TOOLS + + async def test_auth_configured_when_token_given(self): + assert create_mcp_server(lambda: None, api_token="sekrit").auth is not None + assert create_mcp_server(lambda: None).auth is None + + +class TestGatewayResolution: + async def test_tool_errors_cleanly_before_gateway_ready(self): + mcp = create_mcp_server(lambda: None) + async with Client(mcp) as client: + with pytest.raises(Exception, match="starting up"): + await client.call_tool("list_active_calls", {}) + + async def test_make_call_happy_path(self): + gateway = _make_gateway() + mcp = create_mcp_server(lambda: gateway) + async with Client(mcp) as client: + result = await client.call_tool( + "make_call", {"number": "+15551234567", "mode": "direct"} + ) + text = result.content[0].text + assert "initiated" in text + assert "+15551234567" in text + assert len(gateway.call_manager.active_calls) == 1 + + +class TestCallSafety: + def test_emergency_number_detection(self): + for number in ("911", "9911", "112", "+1911", "+112", " 911 ", "9-1-1"): + assert is_emergency_number(number), number + for number in ("+19115551234", "+18005551234", "211", "999"): + assert not is_emergency_number(number), number + + async def test_gateway_refuses_emergency_numbers(self): + gateway = _make_gateway() + with pytest.raises(ValueError, match="emergency"): + await gateway.make_call("911") + assert gateway.call_manager.active_calls == {} + + async def test_mcp_make_call_refuses_emergency(self): + gateway = _make_gateway() + mcp = create_mcp_server(lambda: gateway) + async with Client(mcp) as client: + with pytest.raises(Exception, match="[Ee]mergency"): + await client.call_tool("make_call", {"number": "911"}) + + async def test_concurrent_call_cap(self): + gateway = _make_gateway(max_calls=1) + await gateway.make_call("+15551234567") + with pytest.raises(ValueError, match="limit"): + await gateway.make_call("+15557654321") diff --git a/tests/test_services.py b/tests/test_services.py index 07804c4..ab9b5bf 100644 --- a/tests/test_services.py +++ b/tests/test_services.py @@ -3,7 +3,6 @@ Tests for the intelligence layer services: - LLMClient - NotificationService - RecordingService -- CallAnalytics - CallFlowLearner """ @@ -286,47 +285,6 @@ class TestRecordingService: await svc.stop_recording("call_abc123") -# ============================================================ -# Call Analytics Tests -# ============================================================ - - -class TestCallAnalytics: - """Test analytics tracking.""" - - def _make_service(self): - from services.call_analytics import CallAnalytics - - return CallAnalytics(max_history=1000) - - def test_init(self): - svc = self._make_service() - assert svc._call_records == [] - assert svc.total_calls_recorded == 0 - - def test_get_summary_empty(self): - svc = self._make_service() - summary = svc.get_summary(hours=24) - assert summary["total_calls"] == 0 - assert summary["success_rate"] == 0.0 - - def test_get_company_stats_unknown(self): - svc = self._make_service() - stats = svc.get_company_stats("+18005551234") - assert stats["total_calls"] == 0 - - def test_get_top_numbers_empty(self): - svc = self._make_service() - top = svc.get_top_numbers(limit=5) - assert top == [] - - def test_get_hold_time_trend(self): - svc = self._make_service() - trend = svc.get_hold_time_trend(days=7) - assert len(trend) == 7 - assert all(t["call_count"] == 0 for t in trend) - - # ============================================================ # Call Flow Learner Tests # ============================================================