background / main.py
mythaitts's picture
Upload 3 files
ddaaa1b verified
Raw
History Blame Contribute Delete
88.5 kB
"""
Relay - a competitive Vapi alternative backend engine (single file).
Production-ready voice-agent platform with:
* Persistent storage at /data (SQLite DB, recordings, transcripts) - DATA_DIR configurable
* Per-user accounts, projects and scoped API keys (pbkdf2 password hashing)
* BYO provider credentials: users add their OWN Deepgram/OpenAI/Anthropic/... keys via API
* Deepgram fully supported:
- streaming STT : wss://api.deepgram.com/v1/listen (interim, endpointing, utterance_end_ms, vad_events)
- streaming TTS : wss://api.deepgram.com/v2/speak (Flux) + v1/speak (Aura) with Speak/Flush/Clear/Close
* Streaming LLM: OpenAI-compatible SSE + Anthropic streaming events (OpenAI/Anthropic/Groq/DeepSeek/custom)
* Voice pipeline: streaming-first with Deepgram, turn-based VAD fallback for Whisper/file providers,
barge-in with LLM cancellation + TTS Clear
* Twilio telephony: inbound TwiML <Connect><Stream> + media-streams WebSocket (mulaw 8k codec)
* Usage metering + billing estimates, per-stage latency metrics, transcripts, recordings
* Tool calling + webhooks
"""
from __future__ import annotations
import asyncio
import base64
import hashlib
import hmac
import io
import json
import logging
import os
import secrets
import sqlite3
import struct
import threading
import time
import uuid
import wave
from contextlib import asynccontextmanager
from dataclasses import dataclass, field
from datetime import datetime, timezone
from typing import Any, Optional
from urllib.parse import urljoin
import httpx
try:
import uvicorn
from fastapi import FastAPI, WebSocket, WebSocketDisconnect, HTTPException, Request
from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import JSONResponse, Response
from fastapi.staticfiles import StaticFiles
from pydantic import BaseModel, Field
from pydantic_settings import BaseSettings, SettingsConfigDict
except ImportError as exc: # pragma: no cover
raise SystemExit("Missing dependencies. Run: pip install -r requirements.txt\n" + str(exc))
try:
import websockets # websocket CLIENT for Deepgram streaming
except ImportError: # pragma: no cover
websockets = None
try:
import numpy as np
except ImportError: # pragma: no cover
np = None
# ---------------------------------------------------------------------------
# Settings (.env)
# ---------------------------------------------------------------------------
class Settings(BaseSettings):
model_config = SettingsConfigDict(env_file=".env", env_file_encoding="utf-8", extra="ignore")
app_name: str = "Relay"
environment: str = "production"
host: str = "0.0.0.0"
port: int = 8000
log_level: str = "info"
api_base_url: str = "http://localhost:8000"
# Auth
jwt_secret: str = "change_me_relay_jwt_secret" # SET THIS in production
jwt_expiry_minutes: int = 10080 # 7 days
jwt_issuer: str = "relay"
# Persistence - the /data path (mount a volume here in Docker/HF Spaces)
data_dir: str = "data"
# ---- Global fallback provider settings (used when a project has no own key)
stt_provider: str = "deepgram" # deepgram | whisper | openai
deepgram_api_key: str = ""
deepgram_stt_model: str = "nova-3"
deepgram_tts_model: str = "flux-haley-en" # flux model (v2). Set "aura-2-en" for v1
deepgram_tts_voice: str = "aura-athena-en"
openai_api_key: str = ""
whisper_stt_model: str = "base"
whisper_device: str = "cpu"
llm_provider: str = "openai" # openai | anthropic | groq | deepseek | custom
llm_model: str = "gpt-4o-mini"
llm_api_key: str = ""
llm_base_url: str = ""
anthropic_api_key: str = ""
anthropic_model: str = "claude-3-5-haiku-latest"
tts_provider: str = "deepgram" # deepgram | elevenlabs | cartesia | openai
elevenlabs_api_key: str = ""
elevenlabs_voice_id: str = "21m00Tcm4TlvDq8ikWAM"
cartesia_api_key: str = ""
cartesia_voice_id: str = "a0e4d1b0-0000-0000-0000-000000000000"
openai_tts_model: str = "tts-1"
openai_tts_voice: str = "alloy"
# Audio
sample_rate: int = 16000
frame_ms: int = 20
silence_timeout_ms: int = 700 # turn-based fallback only (Deepgram handles streaming turns)
max_user_turn_ms: int = 20000
barge_in_enabled: bool = True
# Deepgram streaming tuning
deepgram_endpointing: int = 300
deepgram_utterance_end_ms: int = 1000
deepgram_interim: bool = True
deepgram_language: str = "en-US"
# Observability / recording
record_audio: bool = False
webhook_url: str = ""
webhook_secret: str = ""
# Billing rates (approx $ per unit) for usage metering
rate_stt_per_sec: float = 0.0000717 # deepgram nova ~$4.30/hr
rate_llm_input_per_1k: float = 0.00015
rate_llm_output_per_1k: float = 0.0006
rate_tts_per_1k_chars: float = 0.03
settings = Settings()
os.makedirs(settings.data_dir, exist_ok=True)
DB_PATH = os.path.join(settings.data_dir, "relay.db")
RECORDING_DIR = os.path.join(settings.data_dir, "recordings")
os.makedirs(RECORDING_DIR, exist_ok=True)
logging.basicConfig(
level=getattr(logging, settings.log_level.upper(), logging.INFO),
format="%(asctime)s | %(levelname)-7s | %(name)s | %(message)s",
)
log = logging.getLogger("relay")
# ---------------------------------------------------------------------------
# Database (SQLite, persisted at DATA_DIR/relay.db)
# ---------------------------------------------------------------------------
_SCHEMA = """
CREATE TABLE IF NOT EXISTS users (
id TEXT PRIMARY KEY, email TEXT UNIQUE NOT NULL, name TEXT,
password_hash TEXT NOT NULL, created_at TEXT NOT NULL
);
CREATE TABLE IF NOT EXISTS projects (
id TEXT PRIMARY KEY, user_id TEXT NOT NULL, name TEXT NOT NULL, created_at TEXT NOT NULL
);
CREATE TABLE IF NOT EXISTS api_keys (
id TEXT PRIMARY KEY, project_id TEXT NOT NULL, name TEXT,
prefix TEXT NOT NULL, key_hash TEXT NOT NULL, key_value TEXT, created_at TEXT NOT NULL, revoked INTEGER DEFAULT 0
);
CREATE TABLE IF NOT EXISTS provider_credentials (
id TEXT PRIMARY KEY, project_id TEXT NOT NULL, kind TEXT NOT NULL, provider TEXT NOT NULL,
config TEXT NOT NULL, created_at TEXT NOT NULL
);
CREATE TABLE IF NOT EXISTS agents (
id TEXT PRIMARY KEY, project_id TEXT NOT NULL, name TEXT NOT NULL, system_prompt TEXT,
voice TEXT, tools TEXT, config TEXT, created_at TEXT NOT NULL
);
CREATE TABLE IF NOT EXISTS tools (
id TEXT PRIMARY KEY, project_id TEXT NOT NULL, name TEXT NOT NULL, description TEXT,
parameters TEXT, created_at TEXT NOT NULL
);
CREATE TABLE IF NOT EXISTS calls (
id TEXT PRIMARY KEY, project_id TEXT NOT NULL, agent_id TEXT, phone TEXT, status TEXT,
direction TEXT, transport TEXT, started_at TEXT, ended_at TEXT, duration_ms INTEGER DEFAULT 0
);
CREATE TABLE IF NOT EXISTS transcripts (
id TEXT PRIMARY KEY, call_id TEXT NOT NULL, role TEXT NOT NULL, content TEXT,
is_final INTEGER DEFAULT 1, ts TEXT NOT NULL
);
CREATE TABLE IF NOT EXISTS usage (
id TEXT PRIMARY KEY, project_id TEXT NOT NULL, call_id TEXT, kind TEXT NOT NULL,
provider TEXT NOT NULL, model TEXT, units REAL DEFAULT 0, latency_ms REAL DEFAULT 0,
cost REAL DEFAULT 0, created_at TEXT NOT NULL
);
CREATE TABLE IF NOT EXISTS events (
id TEXT PRIMARY KEY, project_id TEXT NOT NULL, call_id TEXT, type TEXT NOT NULL,
payload TEXT, created_at TEXT NOT NULL
);
"""
def db() -> sqlite3.Connection:
conn = sqlite3.connect(DB_PATH)
conn.row_factory = sqlite3.Row
conn.execute("PRAGMA journal_mode=WAL")
conn.execute("PRAGMA busy_timeout=5000")
return conn
def db_exec(sql: str, params: tuple = ()) -> int:
with db() as c:
cur = c.execute(sql, params)
c.commit()
return cur.lastrowid or 0
def db_query(sql: str, params: tuple = ()) -> list[dict[str, Any]]:
with db() as c:
rows = c.execute(sql, params).fetchall()
return [dict(r) for r in rows]
def db_one(sql: str, params: tuple = ()) -> Optional[dict[str, Any]]:
rows = db_query(sql, params)
return rows[0] if rows else None
def init_db():
with db() as c:
c.executescript(_SCHEMA)
c.commit()
# migrations for existing databases
try:
cols = {r["name"] for r in db_query("PRAGMA table_info(api_keys)")}
if "key_value" not in cols:
db_exec("ALTER TABLE api_keys ADD COLUMN key_value TEXT")
log.info("migrated api_keys: added key_value column")
except Exception:
pass
log.info("database ready at %s", DB_PATH)
# ---------------------------------------------------------------------------
# Security: passwords (pbkdf2) + API keys (prefix + sha256 hash) + JWT (HS256)
# ---------------------------------------------------------------------------
_PBKDF2_ITER = 210_000
def hash_password(password: str) -> str:
salt = secrets.token_hex(16)
dk = hashlib.pbkdf2_hmac("sha256", password.encode(), salt.encode(), _PBKDF2_ITER)
return f"pbkdf2${_PBKDF2_ITER}${salt}${dk.hex()}"
def verify_password(password: str, stored: str) -> bool:
try:
algo, iters, salt, hx = stored.split("$")
if algo != "pbkdf2":
return False
dk = hashlib.pbkdf2_hmac("sha256", password.encode(), salt.encode(), int(iters))
return hmac.compare_digest(dk.hex(), hx)
except Exception:
return False
def gen_api_key() -> tuple[str, str, str]:
"""Return (full_key, prefix, sha256_hash). Full key shown at creation (and via reveal endpoint)."""
full = "relay_" + secrets.token_urlsafe(32)
return full, full[:12], hashlib.sha256(full.encode()).hexdigest()
def verify_api_key(full_key: str) -> Optional[dict[str, Any]]:
prefix = full_key[:12]
row = db_one("SELECT * FROM api_keys WHERE prefix=? AND revoked=0", (prefix,))
if not row:
return None
if not hmac.compare_digest(hashlib.sha256(full_key.encode()).hexdigest(), row["key_hash"]):
return None
return row
# ---- JWT (HS256, stdlib only) ----------------------------------------------
def _b64url(data: bytes) -> str:
return base64.urlsafe_b64encode(data).rstrip(b"=").decode()
def _b64url_decode(s: str) -> bytes:
return base64.urlsafe_b64decode(s + "=" * (-len(s) % 4))
def create_token(user_id: str, project_id: str) -> str:
now = int(time.time())
header = _b64url(json.dumps({"alg": "HS256", "typ": "JWT"}, separators=(",", ":")).encode())
payload = _b64url(json.dumps({
"sub": user_id, "pid": project_id, "iss": settings.jwt_issuer,
"iat": now, "exp": now + settings.jwt_expiry_minutes * 60,
}, separators=(",", ":")).encode())
signing_input = f"{header}.{payload}"
sig = hmac.new(settings.jwt_secret.encode(), signing_input.encode(), hashlib.sha256).digest()
return f"{signing_input}.{_b64url(sig)}"
def decode_token(token: str) -> Optional[dict[str, Any]]:
try:
header, payload, sig = token.split(".")
signing_input = f"{header}.{payload}"
expected = _b64url(hmac.new(settings.jwt_secret.encode(), signing_input.encode(), hashlib.sha256).digest())
if not hmac.compare_digest(sig, expected):
return None
data = json.loads(_b64url_decode(payload))
if data.get("exp", 0) < int(time.time()):
return None
if data.get("iss") != settings.jwt_issuer:
return None
return data
except Exception:
return None
def utcnow() -> str:
return datetime.now(timezone.utc).isoformat()
def gen_id(prefix: str) -> str:
return f"{prefix}_{uuid.uuid4().hex[:24]}"
# ---------------------------------------------------------------------------
# Pydantic schemas
# ---------------------------------------------------------------------------
class SignupRequest(BaseModel):
email: str
password: str = Field(min_length=6)
name: str = ""
project_name: str = "Default"
class LoginRequest(BaseModel):
email: str
password: str
class ProjectCreate(BaseModel):
name: str
class KeyCreate(BaseModel):
name: str = "default"
class CredentialCreate(BaseModel):
kind: str # stt | llm | tts
provider: str # deepgram | openai | anthropic | groq | deepseek | elevenlabs | cartesia | whisper | custom
api_key: str = ""
model: str = ""
base_url: str = ""
voice: str = ""
language: str = ""
class AgentCreate(BaseModel):
name: str = "My Agent"
system_prompt: str = "You are a helpful voice assistant."
voice: str = ""
tools: list[dict[str, Any]] = Field(default_factory=list)
config: dict[str, Any] = Field(default_factory=dict)
class ToolCreate(BaseModel):
name: str
description: str = ""
parameters: dict[str, Any] = Field(default_factory=lambda: {"type": "object", "properties": {}})
class CallCreate(BaseModel):
agent_id: str = ""
phone: str = ""
webhook_url: str = ""
class ChatRequest(BaseModel):
model: str = ""
messages: list[dict[str, Any]]
tools: Optional[list[dict[str, Any]]] = None
temperature: float = 0.7
stream: bool = False
class TTSRequest(BaseModel):
text: str
voice: str = ""
model: str = ""
# ---------------------------------------------------------------------------
# Auth dependencies
# ---------------------------------------------------------------------------
def extract_api_key(request: Request) -> Optional[str]:
auth = request.headers.get("Authorization", "")
if auth.lower().startswith("bearer "):
return auth[7:].strip()
key = request.headers.get("X-API-Key")
if key:
return key
return request.query_params.get("api_key")
def _auth_identity(request: Request) -> Optional[tuple[str, str]]:
"""Return (project_id, user_id) from a valid JWT OR API key, else None.
Priority: JWT bearer token first, then API key. Allows:
- Dashboard calls: Authorization: Bearer <jwt>
- Programmatic: Authorization: Bearer <api_key> (or ?api_key= for WS)
"""
auth = request.headers.get("Authorization", "")
if auth.lower().startswith("bearer "):
cred = auth[7:].strip()
# try JWT first
if "." in cred and len(cred) < 500:
tok = decode_token(cred)
if tok and tok.get("sub") and tok.get("pid"):
return tok["pid"], tok["sub"]
# else treat as API key
row = verify_api_key(cred)
if row:
proj = db_one("SELECT user_id FROM projects WHERE id=?", (row["project_id"],))
if proj:
return row["project_id"], proj["user_id"]
# API key via X-API-Key header
xkey = request.headers.get("X-API-Key")
if xkey:
row = verify_api_key(xkey)
if row:
proj = db_one("SELECT user_id FROM projects WHERE id=?", (row["project_id"],))
if proj:
return row["project_id"], proj["user_id"]
# API key via query (for WebSockets)
qkey = request.query_params.get("api_key")
if qkey:
row = verify_api_key(qkey)
if row:
proj = db_one("SELECT user_id FROM projects WHERE id=?", (row["project_id"],))
if proj:
return row["project_id"], proj["user_id"]
return None
def project_from_request(request: Request) -> dict[str, Any]:
ident = _auth_identity(request)
if not ident:
raise HTTPException(status_code=401, detail="Missing/invalid auth. Use Authorization: Bearer <jwt> or <api_key>.")
project = db_one("SELECT * FROM projects WHERE id=?", (ident[0],))
if not project:
raise HTTPException(status_code=401, detail="Project not found")
return project
def user_from_request(request: Request) -> dict[str, Any]:
project = project_from_request(request)
user = db_one("SELECT * FROM users WHERE id=?", (project["user_id"],))
if not user:
raise HTTPException(status_code=401, detail="User not found")
return user
# ---------------------------------------------------------------------------
# Provider configuration: project credentials take priority over globals
# ---------------------------------------------------------------------------
def _global_cfg(kind: str) -> tuple[str, dict[str, Any]]:
if kind == "stt":
return settings.stt_provider, {
"provider": settings.stt_provider,
"api_key": settings.deepgram_api_key if settings.stt_provider == "deepgram" else settings.openai_api_key,
"model": settings.deepgram_stt_model if settings.stt_provider == "deepgram" else settings.whisper_stt_model,
"language": settings.deepgram_language,
}
if kind == "llm":
base = settings.llm_base_url
if settings.llm_provider == "anthropic":
return "anthropic", {"provider": "anthropic", "api_key": settings.anthropic_api_key, "model": settings.anthropic_model, "base_url": "https://api.anthropic.com"}
if settings.llm_provider in ("groq",):
base = base or "https://api.groq.com/openai/v1"
return settings.llm_provider, {"provider": settings.llm_provider, "api_key": settings.llm_api_key, "model": settings.llm_model, "base_url": base}
# tts
return settings.tts_provider, {
"provider": settings.tts_provider,
"api_key": settings.deepgram_api_key if settings.tts_provider == "deepgram"
else (settings.elevenlabs_api_key if settings.tts_provider == "elevenlabs"
else (settings.cartesia_api_key if settings.tts_provider == "cartesia" else settings.openai_api_key)),
"model": settings.deepgram_tts_model if settings.tts_provider == "deepgram" else settings.openai_tts_model,
"voice": settings.deepgram_tts_voice if settings.tts_provider == "deepgram"
else (settings.elevenlabs_voice_id if settings.tts_provider == "elevenlabs"
else (settings.cartesia_voice_id if settings.tts_provider == "cartesia" else settings.openai_tts_voice)),
}
def resolve_provider(project_id: str, kind: str) -> tuple[str, dict[str, Any]]:
"""Return (provider, config) for a project, falling back to global settings."""
row = db_one(
"SELECT provider, config FROM provider_credentials WHERE project_id=? AND kind=? ORDER BY created_at DESC LIMIT 1",
(project_id, kind),
)
if row:
cfg = json.loads(row["config"])
return row["provider"], cfg
return _global_cfg(kind)
# ---------------------------------------------------------------------------
# Usage metering + billing
# ---------------------------------------------------------------------------
def record_usage(project_id: str, call_id: str, kind: str, provider: str, model: str,
units: float = 0, latency_ms: float = 0, cost: float = 0.0):
try:
db_exec(
"INSERT INTO usage (id,project_id,call_id,kind,provider,model,units,latency_ms,cost,created_at) "
"VALUES (?,?,?,?,?,?,?,?,?,?)",
(gen_id("usage"), project_id, call_id, kind, provider, model, units, latency_ms, cost, utcnow()),
)
except Exception as exc:
log.warning("usage record failed: %s", exc)
def estimate_cost(kind: str, provider: str, units: float) -> float:
if kind == "stt":
return units * settings.rate_stt_per_sec # units = audio seconds
if kind == "llm":
# units stored as (input_tokens, output_tokens) via payload string below
return 0.0
if kind == "tts":
return units * settings.rate_tts_per_1k_chars / 1000.0
return 0.0
# ---------------------------------------------------------------------------
# Audio codecs (G.711 mu-law + simple resampling)
# ---------------------------------------------------------------------------
try:
import audioop as _audioop # stdlib, removed in Python 3.13+
def pcm_to_mulaw(pcm: bytes) -> bytes:
return _audioop.lin2ulaw(pcm, 2)
def mulaw_to_pcm(mulaw: bytes) -> bytes:
return _audioop.ulaw2lin(mulaw, 2)
_G711 = "audioop"
except ImportError: # pragma: no cover
_CLIP = 32635
_BIAS = 132
_EXP_LUT = [0, 132, 396, 924, 1980, 4092, 8316, 16764]
def _lin2ulaw(sample: int) -> int:
sign = (sample >> 8) & 0x80
if sign:
sample = -sample
if sample > _CLIP:
sample = _CLIP
sample += _BIAS
exp = 0
while sample >= 256:
sample >>= 1
exp += 1
mantissa = (sample >> 4) & 0x0F
return (~(sign | (exp << 4) | mantissa)) & 0xFF
def _ulaw2lin(b: int) -> int:
u = (~b) & 0xFF
sign = u & 0x80
exp = (u >> 4) & 0x07
man = u & 0x0F
val = _EXP_LUT[exp] + (man << (exp + 3))
return -val if sign else val
def pcm_to_mulaw(pcm: bytes) -> bytes:
return bytes(_lin2ulaw(struct.unpack_from("<h", pcm, i)[0]) for i in range(0, len(pcm), 2))
def mulaw_to_pcm(mulaw: bytes) -> bytes:
out = bytearray(len(mulaw) * 2)
for i, b in enumerate(mulaw):
struct.pack_into("<h", out, i * 2, _ulaw2lin(b))
return bytes(out)
_G711 = "python"
def resample_pcm16(data: bytes, src_rate: int, dst_rate: int) -> bytes:
if src_rate == dst_rate:
return data
if np is None:
return data
arr = np.frombuffer(data, dtype="<i2")
n = int(len(arr) * dst_rate / src_rate)
out = np.interp(
np.linspace(0, len(arr) - 1, n),
np.arange(len(arr)),
arr.astype(np.float64),
).astype("<i2").tobytes()
return out
def mix_audio_chunks(pcm: bytes, sample_rate: int, frame_ms: int = 20) -> bytes:
"""No-op passthrough helper kept for API symmetry."""
return pcm
# ---------------------------------------------------------------------------
# Deepgram streaming STT client (websocket)
# ---------------------------------------------------------------------------
class DeepgramSTTStream:
def __init__(self, api_key: str, model: str, language: str, sample_rate: int,
encoding: str = "linear16", interim: bool = True,
endpointing: int = 300, utterance_end_ms: int = 1000):
self.api_key = api_key
self.model = model or "nova-3"
self.language = language or "en-US"
self.sample_rate = sample_rate
self.encoding = encoding
self.interim = interim
self.endpointing = endpointing
self.utterance_end_ms = utterance_end_ms
self.ws: Optional[Any] = None
self._lock = asyncio.Lock()
def url(self) -> str:
q = (
f"model={self.model}&language={self.language}&encoding={self.encoding}"
f"&sample_rate={self.sample_rate}&channels=1"
f"&interim_results={'true' if self.interim else 'false'}"
f"&endpointing={self.endpointing}&utterance_end_ms={self.utterance_end_ms}"
f"&vad_events=true&smart_format=true"
)
return f"wss://api.deepgram.com/v1/listen?{q}"
async def _open(self) -> Any:
headers = {"Authorization": f"Token {self.api_key}"}
try:
return await websockets.connect(self.url(), additional_headers=headers, max_size=None,
open_timeout=3, ping_interval=None)
except TypeError:
return await websockets.connect(self.url(), extra_headers=headers, max_size=None,
open_timeout=3, ping_interval=None)
async def connect(self):
async with self._lock:
if self.ws is None:
if websockets is None:
raise ProviderError("websockets library not installed")
self.ws = await asyncio.wait_for(self._open(), timeout=4)
async def send_audio(self, pcm: bytes):
await self.connect()
await self.ws.send(pcm)
async def recv(self):
await self.connect()
msg = await self.ws.recv()
if isinstance(msg, bytes):
return {"type": "audio"}
return json.loads(msg)
async def close(self):
if self.ws:
try:
await self.ws.close()
except Exception:
pass
self.ws = None
# ---------------------------------------------------------------------------
# Deepgram streaming TTS client (Flux v2 + Aura v1)
# ---------------------------------------------------------------------------
class DeepgramTTSStream:
def __init__(self, api_key: str, model: str, sample_rate: int = 16000, encoding: str = "linear16"):
self.api_key = api_key
self.model = model or "flux-haley-en"
self.sample_rate = sample_rate
self.encoding = encoding
self.version = 2 if self.model.startswith("flux") else 1
self.ws: Optional[Any] = None
self._lock = asyncio.Lock()
def url(self) -> str:
return (
f"wss://api.deepgram.com/v{self.version}/speak?model={self.model}"
f"&encoding={self.encoding}&sample_rate={self.sample_rate}&container=none"
)
async def _open(self) -> Any:
headers = {"Authorization": f"Token {self.api_key}"}
try:
return await websockets.connect(self.url(), additional_headers=headers, max_size=None,
open_timeout=3, ping_interval=None)
except TypeError:
return await websockets.connect(self.url(), extra_headers=headers, max_size=None,
open_timeout=3, ping_interval=None)
async def connect(self):
async with self._lock:
if self.ws is None:
if websockets is None:
raise ProviderError("websockets library not installed")
self.ws = await asyncio.wait_for(self._open(), timeout=4)
async def speak(self, text: str, text_id: str):
await self.connect()
await self.ws.send(json.dumps({"type": "Speak", "text": text, "text_id": text_id}))
async def flush(self):
await self.connect()
await self.ws.send(json.dumps({"type": "Flush"}))
async def clear(self):
await self.connect()
await self.ws.send(json.dumps({"type": "Clear"}))
async def close(self):
if self.ws:
try:
await self.ws.send(json.dumps({"type": "Close"}))
await self.ws.close()
except Exception:
pass
self.ws = None
async def recv(self):
await self.connect()
msg = await self.ws.recv()
if isinstance(msg, bytes):
return {"type": "audio", "data": msg}
return json.loads(msg)
class ProviderError(Exception):
pass
# ---------------------------------------------------------------------------
# Non-streaming fallback providers (Whisper STT, file TTS, one-shot synth)
# ---------------------------------------------------------------------------
class STTFileEngine:
def __init__(self, provider: str, cfg: dict[str, Any], project_id: str = "", call_id: str = ""):
self.provider = provider
self.cfg = cfg
self.project_id = project_id
self.call_id = call_id
self._whisper = None
def _load_whisper(self):
if self._whisper is None:
import whisper # type: ignore
self._whisper = whisper.load_model(self.cfg.get("model") or "base", device=settings.whisper_device)
async def transcribe(self, wav_bytes: bytes) -> str:
start = time.perf_counter()
text = ""
if self.provider == "openai":
text = await self._openai(wav_bytes)
else:
text = await asyncio.to_thread(self._local_whisper, wav_bytes)
record_usage(self.project_id, self.call_id, "stt", self.provider, self.cfg.get("model", ""),
units=0, latency_ms=(time.perf_counter() - start) * 1000)
return text
async def _openai(self, wav_bytes: bytes) -> str:
api_key = self.cfg.get("api_key") or settings.openai_api_key
if not api_key:
raise ProviderError("No OpenAI API key configured for STT")
url = urljoin((self.cfg.get("base_url") or settings.llm_base_url or "https://api.openai.com/v1"), "audio/transcriptions")
files = {"file": ("audio.wav", wav_bytes, "audio/wav")}
data = {"model": self.cfg.get("model") or "whisper-1"}
async with httpx.AsyncClient(timeout=30) as client:
r = await client.post(url, headers={"Authorization": f"Bearer {api_key}"}, data=data, files=files)
if r.status_code != 200:
raise ProviderError(f"OpenAI STT {r.status_code}: {r.text[:200]}")
return r.json().get("text", "")
def _local_whisper(self, wav_bytes: bytes) -> str:
self._load_whisper()
with wave.open(io.BytesIO(wav_bytes), "rb") as wf:
raw = wf.readframes(wf.getnframes())
audio = np.frombuffer(raw, dtype="<i2").astype(np.float32) / 32768.0
return (self._whisper.transcribe(audio).get("text") or "").strip()
_ctx_project = ""
_ctx_call = ""
class TTSFileEngine:
def __init__(self, provider: str, cfg: dict[str, Any], project_id: str = "", call_id: str = ""):
self.provider = provider
self.cfg = cfg
self.project_id = project_id
self.call_id = call_id
async def synthesize(self, text: str, sample_rate: int = 16000) -> bytes:
start = time.perf_counter()
if self.provider == "elevenlabs":
audio = await self._elevenlabs(text)
elif self.provider == "cartesia":
audio = await self._cartesia(text)
else:
audio = await self._openai(text)
record_usage(self.project_id, self.call_id, "tts", self.provider, self.cfg.get("model", ""),
units=len(text), latency_ms=(time.perf_counter() - start) * 1000)
return audio
async def _elevenlabs(self, text: str) -> bytes:
api_key = self.cfg.get("api_key") or settings.elevenlabs_api_key
if not api_key:
raise ProviderError("No ElevenLabs API key")
url = f"https://api.elevenlabs.io/v1/text-to-speech/{self.cfg.get('voice') or settings.elevenlabs_voice_id}?output_format=pcm_24000"
async with httpx.AsyncClient(timeout=30) as client:
r = await client.post(url, headers={"xi-api-key": api_key, "Content-Type": "application/json"}, json={"text": text})
if r.status_code != 200:
raise ProviderError(f"ElevenLabs {r.status_code}: {r.text[:200]}")
return self._wav_from_pcm(r.content, 24000)
async def _cartesia(self, text: str) -> bytes:
api_key = self.cfg.get("api_key") or settings.cartesia_api_key
if not api_key:
raise ProviderError("No Cartesia API key")
body = {"model_id": "sonic-english", "transcript": text,
"voice": {"mode": "id", "id": self.cfg.get("voice") or settings.cartesia_voice_id},
"output_format": {"container": "wav", "encoding": "pcm_s16le", "sample_rate": 24000}}
async with httpx.AsyncClient(timeout=30) as client:
r = await client.post("https://api.cartesia.ai/tts/bytes",
headers={"Cartesia-Version": "2024-06-10", "X-API-Key": api_key, "Content-Type": "application/json"},
json=body)
if r.status_code != 200:
raise ProviderError(f"Cartesia {r.status_code}: {r.text[:200]}")
return r.content if r.content[:4] == b"RIFF" else self._wav_from_pcm(r.content, 24000)
async def _openai(self, text: str) -> bytes:
api_key = self.cfg.get("api_key") or settings.openai_api_key
if not api_key:
raise ProviderError("No OpenAI API key for TTS")
url = urljoin((self.cfg.get("base_url") or "https://api.openai.com/v1"), "audio/speech")
payload = {"model": self.cfg.get("model") or settings.openai_tts_model,
"voice": self.cfg.get("voice") or settings.openai_tts_voice, "input": text}
async with httpx.AsyncClient(timeout=30) as client:
r = await client.post(url, headers={"Authorization": f"Bearer {api_key}"}, json=payload)
if r.status_code != 200:
raise ProviderError(f"OpenAI TTS {r.status_code}: {r.text[:200]}")
return await self._mp3_to_wav(r.content)
async def _mp3_to_wav(self, mp3: bytes) -> bytes:
try:
import subprocess
p = await asyncio.create_subprocess_exec(
"ffmpeg", "-y", "-i", "pipe:0", "-f", "wav", "pipe:1",
stdin=asyncio.subprocess.PIPE, stdout=asyncio.subprocess.PIPE, stderr=asyncio.subprocess.PIPE)
out, _ = await p.communicate(mp3)
if out:
return out
except Exception:
pass
raise ProviderError("MP3->WAV transcoding requires ffmpeg")
@staticmethod
def _wav_from_pcm(pcm: bytes, sr: int) -> bytes:
buf = io.BytesIO()
with wave.open(buf, "wb") as wf:
wf.setnchannels(1); wf.setsampwidth(2); wf.setframerate(sr); wf.writeframes(pcm)
return buf.getvalue()
# ---------------------------------------------------------------------------
# Streaming LLM (OpenAI-compatible SSE + Anthropic events)
# ---------------------------------------------------------------------------
async def llm_stream(messages: list[dict[str, Any]], cfg: dict[str, Any], tools: Optional[list[dict[str, Any]]] = None,
cancel: Optional[asyncio.Event] = None):
"""Yield {"content": str} and {"tool_call": {...}} chunks from a streaming LLM."""
provider = cfg.get("provider", "openai")
api_key = cfg.get("api_key") or settings.llm_api_key or settings.openai_api_key
model = cfg.get("model") or settings.llm_model
base_url = cfg.get("base_url") or settings.llm_base_url
if provider == "anthropic":
base = (cfg.get("base_url") or settings.anthropic_base_url or "https://api.anthropic.com").rstrip("/")
url = base + "/v1/messages"
system = "\n".join(m["content"] for m in messages if m["role"] == "system")
convo = [m for m in messages if m["role"] != "system"]
payload: dict[str, Any] = {"model": model, "max_tokens": 2048, "system": system or "You are a helpful voice assistant.", "messages": convo, "stream": True}
if tools:
payload["tools"] = [{"name": t["function"]["name"], "description": t["function"].get("description", ""),
"input_schema": t["function"].get("parameters", {"type": "object", "properties": {}})} for t in tools]
headers = {"x-api-key": api_key, "anthropic-version": "2023-06-01", "Content-Type": "application/json"}
async with httpx.AsyncClient(timeout=None) as client:
async with client.stream("POST", url, headers=headers, json=payload) as resp:
if resp.status_code != 200:
raise ProviderError(f"Anthropic {resp.status_code}")
async for line in resp.aiter_lines():
if cancel and cancel.is_set():
break
if not line.startswith("data: "):
continue
obj = json.loads(line[6:])
t = obj.get("type")
if t == "content_block_delta":
yield {"content": obj.get("delta", {}).get("text", "")}
elif t == "content_block_start":
pass
return
# OpenAI-compatible (openai, groq, deepseek, custom)
base = (base_url or "https://api.openai.com/v1").rstrip("/")
url = base + "/chat/completions"
payload = {
"model": model, "messages": messages, "stream": True,
"temperature": 0.7, "stream_options": {"include_usage": True},
}
if tools:
payload["tools"] = tools
payload["tool_choice"] = "auto"
headers = {"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"}
async with httpx.AsyncClient(timeout=None) as client:
async with client.stream("POST", url, headers=headers, json=payload) as resp:
if resp.status_code != 200:
raise ProviderError(f"LLM {resp.status_code}")
async for line in resp.aiter_lines():
if cancel and cancel.is_set():
break
if not line.startswith("data:"):
continue
data = line[5:].strip()
if not data or data == "[DONE]":
continue
obj = json.loads(data)
choice = obj.get("choices", [{}])[0]
delta = choice.get("delta", {})
if delta.get("content"):
yield {"content": delta["content"]}
for tc in delta.get("tool_calls") or []:
yield {"tool_call": tc}
async def execute_tool_call(project_id: str, call_id: str, name: str, args: dict[str, Any]) -> dict[str, Any]:
tool = db_one("SELECT * FROM tools WHERE project_id=? AND name=?", (project_id, name))
await dispatch_event(project_id, call_id, "tool.call", {"name": name, "arguments": args})
if not tool:
return {"success": False, "error": f"Unknown tool: {name}"}
handler = TOOL_HANDLERS.get(name)
if handler:
try:
if asyncio.iscoroutinefunction(handler):
result = await handler(**args)
else:
result = await asyncio.to_thread(handler, **args)
return {"success": True, "result": result}
except Exception as exc:
return {"success": False, "error": str(exc)}
return {"success": True, "deferred": True}
TOOL_HANDLERS: dict[str, Any] = {}
async def llm_generate(messages: list[dict[str, Any]], cfg: dict[str, Any], project_id: str = "",
call_id: str = "", tools: Optional[list[dict[str, Any]]] = None,
on_token: Optional[Callable[[str], Any]] = None,
cancel: Optional[asyncio.Event] = None) -> str:
"""Stream an LLM response, executing tool calls, returning final text."""
provider = cfg.get("provider", "openai")
model = cfg.get("model") or settings.llm_model
start = time.perf_counter()
full = ""
input_tokens = 0
output_tokens = 0
for _ in range(6):
tool_calls: dict[int, dict[str, Any]] = {}
content = ""
async for chunk in llm_stream(messages, cfg, tools, cancel):
if cancel and cancel.is_set():
return full
if chunk.get("content"):
content += chunk["content"]
if on_token:
await on_token(chunk["content"])
if chunk.get("tool_call"):
tc = chunk["tool_call"]
idx = tc.get("index", 0)
entry = tool_calls.setdefault(idx, {"id": tc.get("id") or f"call_{idx}", "name": "", "arguments": ""})
if tc.get("id"):
entry["id"] = tc["id"]
fn = tc.get("function", {})
if fn.get("name"):
entry["name"] += fn["name"]
if fn.get("arguments"):
entry["arguments"] += fn["arguments"]
full = content
if not tool_calls:
break
assistant_msg = {"role": "assistant", "content": content or None,
"tool_calls": [{"id": v["id"], "type": "function",
"function": {"name": v["name"], "arguments": v["arguments"]}} for v in tool_calls.values()]}
messages.append(assistant_msg)
for idx in sorted(tool_calls):
call = tool_calls[idx]
try:
args = json.loads(call["arguments"]) if call["arguments"] else {}
except Exception:
args = {}
result = await execute_tool_call(project_id, call_id, call["name"], args)
messages.append({"role": "tool", "tool_call_id": call["id"], "content": json.dumps(result)})
else:
pass
latency = (time.perf_counter() - start) * 1000
# approx token counts
input_tokens = sum(len(str(m.get("content", ""))) // 4 for m in messages)
output_tokens = max(len(full) // 4, 1)
cost = (input_tokens / 1000) * settings.rate_llm_input_per_1k + (output_tokens / 1000) * settings.rate_llm_output_per_1k
record_usage(project_id, call_id, "llm", provider, model, units=output_tokens, latency_ms=latency, cost=cost)
return full
from typing import Callable # noqa: E402 (used above)
# ---------------------------------------------------------------------------
# Webhooks / events
# ---------------------------------------------------------------------------
async def dispatch_event(project_id: str, call_id: str, event_type: str, payload: dict[str, Any]):
db_exec("INSERT INTO events (id,project_id,call_id,type,payload,created_at) VALUES (?,?,?,?,?,?)",
(gen_id("event"), project_id, call_id, event_type, json.dumps(payload), utcnow()))
project = db_one("SELECT * FROM projects WHERE id=?", (project_id,))
url = payload.pop("webhook_url", None) or (project and project.get("webhook_url")) or settings.webhook_url
if not url:
return
body = json.dumps({"event": event_type, "call_id": call_id, "data": payload, "ts": utcnow()}, separators=(",", ":"))
headers = {"Content-Type": "application/json"}
if settings.webhook_secret:
sig = hmac.new(settings.webhook_secret.encode(), body.encode(), hashlib.sha256).hexdigest()
headers["X-Relay-Signature"] = sig
try:
async with httpx.AsyncClient(timeout=5) as client:
await client.post(url, content=body, headers=headers)
except Exception as exc:
log.warning("webhook %s failed: %s", event_type, exc)
# ---------------------------------------------------------------------------
# Transports: normalize client WebSocket and Twilio media streams
# ---------------------------------------------------------------------------
class Transport:
async def recv(self): ...
async def send_json(self, data: dict[str, Any]): ...
async def send_audio(self, pcm: bytes, sample_rate: int): ...
async def close(self): ...
class WebSocketTransport(Transport):
"""Client sends binary PCM16 at 16 kHz (default); server returns base64 audio JSON."""
def __init__(self, ws: WebSocket):
self.ws = ws
self.sample_rate = settings.sample_rate
async def recv(self):
msg = await self.ws.receive()
if msg.get("type") == "websocket.disconnect":
return ("close", None)
if msg.get("type") == "websocket.receive":
data = msg.get("bytes") or msg.get("text")
if isinstance(data, bytes):
return ("audio", data)
if isinstance(data, str):
try:
return ("json", json.loads(data))
except Exception:
return ("json", {"type": "unknown"})
return ("close", None)
async def send_json(self, data: dict[str, Any]):
try:
await self.ws.send_text(json.dumps(data))
except Exception:
pass
async def send_audio(self, pcm: bytes, sample_rate: int):
try:
await self.ws.send_bytes(pcm)
except Exception:
pass
async def close(self):
try:
await self.ws.close()
except Exception:
pass
class TwilioTransport(Transport):
"""Twilio media streams: inbound mulaw 8k JSON -> PCM16 16k; outbound PCM -> mulaw 8k media events."""
def __init__(self, ws: WebSocket):
self.ws = ws
self.stream_sid: Optional[str] = None
self.in_rate = 8000
self.out_rate = 8000
async def recv(self):
msg = await self.ws.receive()
if msg.get("type") == "websocket.disconnect":
return ("close", None)
data = msg.get("text") or msg.get("bytes")
if isinstance(data, bytes):
return ("audio", resample_pcm16(data, self.in_rate, 16000))
if isinstance(data, str):
try:
obj = json.loads(data)
except Exception:
return ("json", {})
event = obj.get("event")
if event == "start":
self.stream_sid = obj.get("streamSid")
fmt = obj.get("start", {}).get("mediaFormat", {})
self.in_rate = int(fmt.get("sampleRate", 8000))
return ("json", {"type": "start", "stream_sid": self.stream_sid})
if event == "media":
payload = obj.get("media", {}).get("payload", "")
if payload:
mulaw = base64.b64decode(payload)
pcm8k = mulaw_to_pcm(mulaw)
return ("audio", resample_pcm16(pcm8k, 8000, 16000))
return ("audio", b"")
if event == "stop":
return ("close", None)
return ("json", {"type": event})
return ("close", None)
async def send_json(self, data: dict[str, Any]):
try:
await self.ws.send_text(json.dumps(data))
except Exception:
pass
async def send_audio(self, pcm: bytes, sample_rate: int):
if not self.stream_sid:
return
pcm8k = resample_pcm16(pcm, sample_rate, 8000)
mulaw = pcm_to_mulaw(pcm8k)
payload = base64.b64encode(mulaw).decode()
try:
await self.ws.send_text(json.dumps({
"event": "media",
"streamSid": self.stream_sid,
"media": {"payload": payload, "track": "outbound"},
}))
except Exception:
pass
async def close(self):
try:
await self.ws.close()
except Exception:
pass
# ---------------------------------------------------------------------------
# Voice pipeline (streaming-first, turn-based fallback, barge-in)
# ---------------------------------------------------------------------------
def split_sentences(text: str) -> list[str]:
import re
parts = re.split(r"(?<=[.!?])\s+|\n+", text)
return [p.strip() for p in parts if p.strip()]
class VoicePipeline:
def __init__(self, call_id: str, project_id: str, agent: dict[str, Any],
transport: Transport, client_sample_rate: int = 16000):
self.call_id = call_id
self.project_id = project_id
self.agent = agent
self.transport = transport
self.client_rate = client_sample_rate
self.conversation: list[dict[str, Any]] = []
self.state = "listening"
self.cancel_llm = asyncio.Event()
self.speaking = False
self.llm_task: Optional[asyncio.Task] = None
self.stt_stream: Optional[DeepgramSTTStream] = None
self.tts_stream: Optional[DeepgramTTSStream] = None
self.stt_file: Optional[STTFileEngine] = None
self.tts_file: Optional[TTSFileEngine] = None
self.turn_buffer = bytearray()
self.last_speech = 0.0
self.vad_threshold = 0.012
self.inbound_recording = bytearray()
self.recording_path: Optional[str] = None
self.stt_provider, self.stt_cfg = resolve_provider(project_id, "stt")
self.llm_provider, self.llm_cfg = resolve_provider(project_id, "llm")
self.tts_provider, self.tts_cfg = resolve_provider(project_id, "tts")
self.streaming_stt = self.stt_provider == "deepgram"
self.streaming_tts = self.tts_provider == "deepgram"
# ---- lifecycle ------------------------------------------------------
async def start(self):
await self.transport.send_json({"type": "session", "call_id": self.call_id,
"sample_rate": self.client_rate,
"model": self.llm_cfg.get("model", ""),
"stt": self.stt_provider, "llm": self.llm_provider, "tts": self.tts_provider})
if self.streaming_stt:
self.stt_stream = DeepgramSTTStream(
self.stt_cfg.get("api_key"), self.stt_cfg.get("model"), self.stt_cfg.get("language"),
self.client_rate, interim=settings.deepgram_interim,
endpointing=settings.deepgram_endpointing, utterance_end_ms=settings.deepgram_utterance_end_ms)
asyncio.create_task(self._stt_read_loop())
else:
self.stt_file = STTFileEngine(self.stt_provider, self.stt_cfg, project_id=project_id, call_id=call_id)
self.vad_threshold = float(self.agent.get("config", {}).get("vad_threshold", 0.012))
if not self.streaming_tts:
self.tts_file = TTSFileEngine(self.tts_provider, self.tts_cfg, project_id=project_id, call_id=call_id)
# greeting
greeting = self.agent.get("config", {}).get("greeting")
if greeting:
await self.speak(greeting)
async def run(self):
try:
while True:
kind, payload = await self.transport.recv()
if kind == "close":
break
if kind == "audio":
await self._on_audio(payload)
elif kind == "json":
await self._on_json(payload)
except WebSocketDisconnect:
pass
except Exception as exc:
log.warning("pipeline error: %s", exc)
finally:
await self.cleanup()
# ---- inbound --------------------------------------------------------
async def _on_audio(self, pcm: bytes):
if settings.record_audio:
self.inbound_recording.extend(pcm)
if self.streaming_stt:
# barge-in: speech energy while assistant talking
if self.speaking and self._has_energy(pcm):
await self.barge_in()
if self.stt_stream:
try:
await self.stt_stream.send_audio(pcm)
except Exception as exc:
log.warning("stt send: %s", exc)
else:
await self._turn_based_audio(pcm)
def _has_energy(self, pcm: bytes) -> bool:
if np is None or len(pcm) < 2:
return False
arr = np.frombuffer(pcm, dtype="<i2").astype(np.float32) / 32768.0
return float(np.sqrt(np.mean(np.square(arr)))) > self.vad_threshold
async def _turn_based_audio(self, pcm: bytes):
speech = self._has_energy(pcm)
if speech:
if not self.turn_buffer:
self.last_speech = time.perf_counter()
await self.transport.send_json({"type": "status", "state": "user_speaking"})
self.turn_buffer.extend(pcm)
self.last_speech = time.perf_counter()
else:
if self.turn_buffer and (time.perf_counter() - self.last_speech) >= (self.agent.get("config", {}).get("silence_timeout_ms", settings.silence_timeout_ms) / 1000.0):
await self._process_turn(bytes(self.turn_buffer))
self.turn_buffer.clear()
async def _process_turn(self, pcm: bytes):
buf = io.BytesIO()
with wave.open(buf, "wb") as wf:
wf.setnchannels(1); wf.setsampwidth(2); wf.setframerate(self.client_rate); wf.writeframes(pcm)
start = time.perf_counter()
text = ""
try:
text = await self.stt_file.transcribe(buf.getvalue())
except ProviderError as exc:
log.warning("stt: %s", exc)
if text:
record_usage(self.project_id, self.call_id, "stt", self.stt_provider, self.stt_cfg.get("model", ""),
units=0, latency_ms=(time.perf_counter() - start) * 1000)
await self.handle_user_text(text)
async def _on_json(self, data: dict[str, Any]):
typ = data.get("type")
if typ in ("end_of_turn", "eot"):
if self.turn_buffer:
await self._process_turn(bytes(self.turn_buffer))
self.turn_buffer.clear()
elif typ == "barge_in":
await self.barge_in()
elif typ == "ping":
await self.transport.send_json({"type": "pong"})
# ---- STT read loop (streaming) --------------------------------------
async def _stt_read_loop(self):
try:
while True:
data = await self.stt_stream.recv()
typ = data.get("type")
if typ == "Results":
alt = data.get("channel", {}).get("alternatives", [{}])[0]
text = alt.get("transcript", "")
is_final = data.get("is_final", False)
if text:
await self.transport.send_json({"type": "transcript", "role": "user", "text": text, "is_final": is_final})
if is_final and text.strip():
record_usage(self.project_id, self.call_id, "stt", self.stt_provider, self.stt_cfg.get("model", ""),
units=0, latency_ms=0)
await self.handle_user_text(text.strip())
elif typ == "UtteranceEnd":
await self.transport.send_json({"type": "utterance_end"})
elif typ == "SpeechStarted":
if self.speaking:
await self.barge_in()
except Exception as exc:
log.warning("stt loop: %s", exc)
# ---- turn processing ------------------------------------------------
async def handle_user_text(self, text: str):
if self.speaking or (self.llm_task and not self.llm_task.done()):
await self.barge_in()
self.conversation.append({"role": "user", "content": text})
db_exec("INSERT INTO transcripts (id,call_id,role,content,is_final,ts) VALUES (?,?,?,?,1,?)",
(gen_id("t"), self.call_id, "user", text, utcnow()))
await self.transport.send_json({"type": "status", "state": "thinking"})
self.llm_task = asyncio.create_task(self._run_llm())
async def _run_llm(self):
self.cancel_llm.clear()
tools = self.agent.get("tools")
buffer = ""
async def on_token(tok: str):
nonlocal buffer
buffer += tok
await self.transport.send_json({"type": "token", "text": tok})
# stream sentences to TTS as they complete
sents = split_sentences(buffer)
if len(sents) > 1:
chunk = " ".join(sents[:-1])
buffer = sents[-1]
if chunk:
await self.speak(chunk)
try:
final = await llm_generate(
self.conversation, self.llm_cfg, project_id=self.project_id, call_id=self.call_id,
tools=tools, on_token=on_token, cancel=self.cancel_llm)
except ProviderError as exc:
await self.transport.send_json({"type": "error", "error": str(exc)})
return
if self.cancel_llm.is_set():
return
if buffer.strip():
await self.speak(buffer)
if final:
self.conversation.append({"role": "assistant", "content": final})
db_exec("INSERT INTO transcripts (id,call_id,role,content,is_final,ts) VALUES (?,?,?,?,1,?)",
(gen_id("t"), self.call_id, "assistant", final, utcnow()))
await self.transport.send_json({"type": "status", "state": "listening"})
# ---- TTS ------------------------------------------------------------
async def speak(self, text: str):
if not text.strip():
return
if self.speaking:
return
self.speaking = True
await self.transport.send_json({"type": "status", "state": "assistant_speaking"})
try:
if self.streaming_tts:
await self._stream_tts(text)
else:
audio = await self.tts_file.synthesize(text, sample_rate=self.client_rate)
await self.transport.send_audio(audio, self.client_rate)
except Exception as exc:
await self.transport.send_json({"type": "error", "error": str(exc)})
finally:
self.speaking = False
await self.transport.send_json({"type": "status", "state": "listening"})
async def _stream_tts(self, text: str):
if self.tts_stream is None:
self.tts_stream = DeepgramTTSStream(
self.tts_cfg.get("api_key"), self.tts_cfg.get("model"), sample_rate=self.client_rate)
asyncio.create_task(self._tts_read_loop())
text_id = gen_id("tid")
await self.tts_stream.speak(text, text_id)
if self.tts_stream.version == 1:
await self.tts_stream.flush()
async def _tts_read_loop(self):
try:
while True:
data = await self.tts_stream.recv()
if data.get("type") == "audio":
await self.transport.send_audio(data["data"], self.client_rate)
elif data.get("type") in ("Flushed", "Cleared", "Close"):
await self.transport.send_json({"type": "tts_event", "event": data.get("type")})
except Exception:
pass
# ---- barge-in -------------------------------------------------------
async def barge_in(self):
if self.cancel_llm.is_set():
return
self.cancel_llm.set()
self.speaking = False
if self.llm_task and not self.llm_task.done():
self.llm_task.cancel()
if self.tts_stream:
try:
await self.tts_stream.clear()
except Exception:
pass
await self.transport.send_json({"type": "barge_in"})
await dispatch_event(self.project_id, self.call_id, "barge_in", {"state": "interrupted"})
# ---- cleanup --------------------------------------------------------
async def cleanup(self):
if self.stt_stream:
try:
await self.stt_stream.close()
except Exception:
pass
if self.tts_stream:
try:
await self.tts_stream.close()
except Exception:
pass
if self.llm_task and not self.llm_task.done():
self.llm_task.cancel()
if settings.record_audio and self.inbound_recording:
try:
path = os.path.join(RECORDING_DIR, f"{self.call_id}.wav")
with wave.open(path, "wb") as wf:
wf.setnchannels(1); wf.setsampwidth(2); wf.setframerate(self.client_rate)
wf.writeframes(bytes(self.inbound_recording))
self.recording_path = path
except Exception as exc:
log.warning("recording: %s", exc)
db_exec("UPDATE calls SET status='ended', ended_at=?, duration_ms=? WHERE id=?",
(utcnow(), int((time.time() - CALL_START.get(self.call_id, time.time())) * 1000), self.call_id))
await dispatch_event(self.project_id, self.call_id, "call.ended", {})
try:
await self.transport.close()
except Exception:
pass
CALL_START: dict[str, float] = {}
# ---------------------------------------------------------------------------
# FastAPI app
# ---------------------------------------------------------------------------
@asynccontextmanager
async def lifespan(app: FastAPI):
init_db()
log.info("%s engine started (env=%s, data=%s)", settings.app_name, settings.environment, settings.data_dir)
yield
app = FastAPI(title="Relay", version="2.0.0", lifespan=lifespan)
app.add_middleware(CORSMiddleware, allow_origins=["*"], allow_methods=["*"], allow_headers=["*"])
if os.path.isdir(RECORDING_DIR):
app.mount("/data", StaticFiles(directory=RECORDING_DIR), name="data")
# ---- Health ---------------------------------------------------------------
@app.get("/health")
async def health():
return {"status": "ok", "app": settings.app_name, "environment": settings.environment,
"data_dir": settings.data_dir, "db": DB_PATH, "time": utcnow()}
# ---- Auth -----------------------------------------------------------------
@app.post("/auth/signup", status_code=201)
async def signup(body: SignupRequest):
email = body.email.lower().strip()
if db_one("SELECT id FROM users WHERE email=?", (email,)):
raise HTTPException(status_code=409, detail="Email already registered")
user_id = gen_id("user")
db_exec("INSERT INTO users (id,email,name,password_hash,created_at) VALUES (?,?,?,?,?)",
(user_id, email, body.name, hash_password(body.password), utcnow()))
project_id = gen_id("proj")
db_exec("INSERT INTO projects (id,user_id,name,created_at) VALUES (?,?,?,?)",
(project_id, user_id, body.project_name or "Default", utcnow()))
full, prefix, kh = gen_api_key()
key_id = gen_id("key")
db_exec("INSERT INTO api_keys (id,project_id,name,prefix,key_hash,key_value,created_at) VALUES (?,?,?,?,?,?,?)",
(key_id, project_id, "default", prefix, kh, full, utcnow()))
token = create_token(user_id, project_id)
return {
"user_id": user_id, "project_id": project_id,
"access_token": token, "token_type": "bearer",
"api_key": full, "api_key_id": key_id,
}
@app.post("/auth/login")
async def login(body: LoginRequest):
email = body.email.lower().strip()
user = db_one("SELECT * FROM users WHERE email=?", (email,))
if not user or not verify_password(body.password, user["password_hash"]):
raise HTTPException(status_code=401, detail="Invalid credentials")
project = db_one("SELECT * FROM projects WHERE user_id=? ORDER BY created_at LIMIT 1", (user["id"],))
if not project:
raise HTTPException(status_code=401, detail="No project found")
token = create_token(user["id"], project["id"])
return {
"user_id": user["id"], "project_id": project["id"], "name": user["name"],
"access_token": token, "token_type": "bearer",
}
@app.get("/auth/me")
async def me(request: Request):
user = user_from_request(request)
project = project_from_request(request)
return {"user": user, "project": project}
# ---- Projects & API keys --------------------------------------------------
@app.post("/user/project", status_code=201)
async def create_project(body: ProjectCreate, request: Request):
user = user_from_request(request)
pid = gen_id("proj")
db_exec("INSERT INTO projects (id,user_id,name,created_at) VALUES (?,?,?,?)",
(pid, user["id"], body.name, utcnow()))
return db_one("SELECT * FROM projects WHERE id=?", (pid,))
@app.get("/user/project")
async def list_projects(request: Request):
user = user_from_request(request)
return db_query("SELECT * FROM projects WHERE user_id=?", (user["id"],))
@app.post("/user/project/{project_id}/key", status_code=201)
async def create_key(project_id: str, body: KeyCreate, request: Request):
user = user_from_request(request)
proj = db_one("SELECT * FROM projects WHERE id=? AND user_id=?", (project_id, user["id"]))
if not proj:
raise HTTPException(status_code=404, detail="Project not found")
full, prefix, kh = gen_api_key()
key_id = gen_id("key")
db_exec("INSERT INTO api_keys (id,project_id,name,prefix,key_hash,key_value,created_at) VALUES (?,?,?,?,?,?,?)",
(key_id, project_id, body.name, prefix, kh, full, utcnow()))
return {"api_key_id": key_id, "api_key": full, "name": body.name, "project_id": project_id}
@app.get("/user/project/{project_id}/key")
async def list_keys(project_id: str, request: Request):
user = user_from_request(request)
rows = db_query("SELECT id,name,prefix,created_at,revoked FROM api_keys WHERE project_id=? AND revoked=0", (project_id,))
for r in rows:
r["owner"] = db_one("SELECT user_id FROM projects WHERE id=?", (project_id,))["user_id"] == user["id"]
return rows
@app.get("/user/project/{project_id}/key/{key_id}")
async def get_key(project_id: str, key_id: str, request: Request):
"""Show (reveal) a specific API key for the current project."""
user = user_from_request(request)
proj = db_one("SELECT * FROM projects WHERE id=? AND user_id=?", (project_id, user["id"]))
if not proj:
raise HTTPException(status_code=404, detail="Project not found")
row = db_one("SELECT id,name,prefix,key_value,created_at,revoked FROM api_keys WHERE id=? AND project_id=?",
(key_id, project_id))
if not row:
raise HTTPException(status_code=404, detail="API key not found")
return {"id": row["id"], "name": row["name"], "prefix": row["prefix"],
"api_key": row["key_value"] if row["key_value"] else None, "created_at": row["created_at"],
"revoked": row["revoked"]}
@app.delete("/user/project/{project_id}/key/{key_id}")
async def revoke_key(project_id: str, key_id: str, request: Request):
user = user_from_request(request)
proj = db_one("SELECT * FROM projects WHERE id=? AND user_id=?", (project_id, user["id"]))
if not proj:
raise HTTPException(status_code=404, detail="Project not found")
db_exec("UPDATE api_keys SET revoked=1 WHERE id=? AND project_id=?", (key_id, project_id))
return {"ok": True}
# ---- Provider credentials (users add their own keys) ----------------------
@app.post("/user/project/{project_id}/credential", status_code=201)
async def add_credential(project_id: str, body: CredentialCreate, request: Request):
user = user_from_request(request)
proj = db_one("SELECT * FROM projects WHERE id=? AND user_id=?", (project_id, user["id"]))
if not proj:
raise HTTPException(status_code=404, detail="Project not found")
if body.kind not in ("stt", "llm", "tts"):
raise HTTPException(status_code=400, detail="kind must be stt|llm|tts")
config = {"provider": body.provider, "api_key": body.api_key, "model": body.model,
"base_url": body.base_url, "voice": body.voice, "language": body.language}
cid = gen_id("cred")
db_exec("INSERT INTO provider_credentials (id,project_id,kind,provider,config,created_at) VALUES (?,?,?,?,?,?)",
(cid, project_id, body.kind, body.provider, json.dumps(config), utcnow()))
return {"id": cid, "kind": body.kind, "provider": body.provider, "config": config}
@app.get("/user/project/{project_id}/credential")
async def list_credentials(project_id: str, request: Request):
user = user_from_request(request)
proj = db_one("SELECT * FROM projects WHERE id=? AND user_id=?", (project_id, user["id"]))
if not proj:
raise HTTPException(status_code=404, detail="Project not found")
rows = db_query("SELECT id,kind,provider,config,created_at FROM provider_credentials WHERE project_id=? ORDER BY created_at", (project_id,))
for r in rows:
cfg = json.loads(r["config"])
if cfg.get("api_key"):
cfg["api_key"] = cfg["api_key"][:8] + "..." + cfg["api_key"][-4:] if len(cfg["api_key"]) > 12 else "***"
r["config"] = cfg
return rows
@app.delete("/user/project/{project_id}/credential/{cred_id}")
async def delete_credential(project_id: str, cred_id: str, request: Request):
user = user_from_request(request)
proj = db_one("SELECT * FROM projects WHERE id=? AND user_id=?", (project_id, user["id"]))
if not proj:
raise HTTPException(status_code=404, detail="Project not found")
db_exec("DELETE FROM provider_credentials WHERE id=? AND project_id=?", (cred_id, project_id))
return {"ok": True}
# ---- Agents ---------------------------------------------------------------
@app.post("/agent", status_code=201)
async def create_agent(body: AgentCreate, request: Request):
project = project_from_request(request)
aid = gen_id("agent")
db_exec("INSERT INTO agents (id,project_id,name,system_prompt,voice,tools,config,created_at) VALUES (?,?,?,?,?,?,?,?)",
(aid, project["id"], body.name, body.system_prompt, body.voice, json.dumps(body.tools),
json.dumps(body.config), utcnow()))
return get_agent_row(aid)
def get_agent_row(aid: str) -> dict[str, Any]:
row = db_one("SELECT * FROM agents WHERE id=?", (aid,))
if not row:
raise HTTPException(status_code=404, detail="Agent not found")
row["tools"] = json.loads(row["tools"] or "[]")
row["config"] = json.loads(row["config"] or "{}")
return row
@app.get("/agent")
async def list_agents(request: Request):
project = project_from_request(request)
rows = db_query("SELECT * FROM agents WHERE project_id=?", (project["id"],))
for r in rows:
r["tools"] = json.loads(r["tools"] or "[]")
r["config"] = json.loads(r["config"] or "{}")
return rows
@app.get("/agent/{agent_id}")
async def get_agent(agent_id: str, request: Request):
project = project_from_request(request)
row = db_one("SELECT * FROM agents WHERE id=? AND project_id=?", (agent_id, project["id"]))
if not row:
raise HTTPException(status_code=404, detail="Agent not found")
row["tools"] = json.loads(row["tools"] or "[]")
row["config"] = json.loads(row["config"] or "{}")
return row
@app.patch("/agent/{agent_id}")
async def update_agent(agent_id: str, body: AgentCreate, request: Request):
project = project_from_request(request)
existing = db_one("SELECT * FROM agents WHERE id=? AND project_id=?", (agent_id, project["id"]))
if not existing:
raise HTTPException(status_code=404, detail="Agent not found")
db_exec("UPDATE agents SET name=?, system_prompt=?, voice=?, tools=?, config=? WHERE id=?",
(body.name, body.system_prompt, body.voice, json.dumps(body.tools), json.dumps(body.config), agent_id))
return get_agent_row(agent_id)
@app.delete("/agent/{agent_id}")
async def delete_agent(agent_id: str, request: Request):
project = project_from_request(request)
db_exec("DELETE FROM agents WHERE id=? AND project_id=?", (agent_id, project["id"]))
return {"ok": True}
# ---- Tools ----------------------------------------------------------------
@app.post("/tool", status_code=201)
async def create_tool(body: ToolCreate, request: Request):
project = project_from_request(request)
tid = gen_id("tool")
db_exec("INSERT INTO tools (id,project_id,name,description,parameters,created_at) VALUES (?,?,?,?,?,?)",
(tid, project["id"], body.name, body.description, json.dumps(body.parameters), utcnow()))
return {"id": tid, **body.model_dump()}
@app.get("/tool")
async def list_tools(request: Request):
project = project_from_request(request)
rows = db_query("SELECT * FROM tools WHERE project_id=?", (project["id"],))
for r in rows:
r["parameters"] = json.loads(r["parameters"] or "{}")
return rows
@app.delete("/tool/{tool_id}")
async def delete_tool(tool_id: str, request: Request):
project = project_from_request(request)
db_exec("DELETE FROM tools WHERE id=? AND project_id=?", (tool_id, project["id"]))
return {"ok": True}
# ---- Calls ----------------------------------------------------------------
@app.post("/call", status_code=201)
async def create_call(body: CallCreate, request: Request):
project = project_from_request(request)
call_id = gen_id("call")
db_exec("INSERT INTO calls (id,project_id,agent_id,phone,status,direction,transport,started_at) VALUES (?,?,?,?,?,?,?,?)",
(call_id, project["id"], body.agent_id, body.phone, "queued", "inbound", "ws", utcnow()))
await dispatch_event(project["id"], call_id, "call.started", {"agent_id": body.agent_id, "phone": body.phone})
return db_one("SELECT * FROM calls WHERE id=?", (call_id,))
@app.get("/call")
async def list_calls(request: Request):
project = project_from_request(request)
return db_query("SELECT * FROM calls WHERE project_id=? ORDER BY started_at DESC", (project["id"],))
@app.get("/call/{call_id}")
async def get_call(call_id: str, request: Request):
project = project_from_request(request)
row = db_one("SELECT * FROM calls WHERE id=? AND project_id=?", (call_id, project["id"]))
if not row:
raise HTTPException(status_code=404, detail="Call not found")
return row
@app.get("/call/{call_id}/transcript")
async def call_transcript(call_id: str, request: Request):
project = project_from_request(request)
return db_query("SELECT role,content,is_final,ts FROM transcripts WHERE call_id=? ORDER BY rowid", (call_id,))
# ---- Usage & billing ------------------------------------------------------
@app.get("/usage")
async def usage(request: Request):
project = project_from_request(request)
rows = db_query("SELECT kind,provider,model,COUNT(*) as count,ROUND(SUM(units),2) as units,ROUND(AVG(latency_ms),1) as avg_ms,ROUND(SUM(cost),6) as cost FROM usage WHERE project_id=? GROUP BY kind,provider,model", (project["id"],))
return rows
@app.get("/billing")
async def billing(request: Request):
project = project_from_request(request)
rows = db_query("SELECT kind,ROUND(SUM(cost),6) as cost FROM usage WHERE project_id=? GROUP BY kind", (project["id"],))
total = sum(r["cost"] for r in rows)
return {"total_estimated_cost": round(total, 6), "breakdown": rows}
# ---- Observability --------------------------------------------------------
@app.get("/events")
async def list_events(request: Request):
project = project_from_request(request)
rows = db_query("SELECT id,call_id,type,payload,created_at FROM events WHERE project_id=? ORDER BY rowid DESC LIMIT 200", (project["id"],))
for r in rows:
r["payload"] = json.loads(r["payload"] or "{}")
return rows
@app.get("/metrics/latency")
async def latency_summary(request: Request):
project = project_from_request(request)
rows = db_query(
"SELECT kind,provider,COUNT(*) as count,ROUND(AVG(latency_ms),1) as avg_ms,ROUND(MIN(latency_ms),1) as min_ms,ROUND(MAX(latency_ms),1) as max_ms "
"FROM usage WHERE project_id=? GROUP BY kind,provider", (project["id"],))
return rows
@app.get("/recording/{call_id}")
async def get_recording(call_id: str, request: Request):
project = project_from_request(request)
path = os.path.join(RECORDING_DIR, f"{call_id}.wav")
if not os.path.exists(path):
raise HTTPException(status_code=404, detail="Recording not found")
return Response(content=open(path, "rb").read(), media_type="audio/wav")
# ---- OpenAI-compatible bridge ---------------------------------------------
@app.post("/v1/chat/completions")
async def chat_completions(body: ChatRequest, request: Request):
project = project_from_request(request)
_, cfg = resolve_provider(project["id"], "llm")
if body.model:
cfg = {**cfg, "model": body.model}
if body.stream:
from fastapi.responses import StreamingResponse
async def gen():
collected = ""
yield "data: " + json.dumps({"id": gen_id("cmpl"), "object": "chat.completion.chunk",
"created": int(time.time()), "model": cfg.get("model", ""), "choices": [{"index": 0, "delta": {"role": "assistant", "content": ""}, "finish_reason": None}]}) + "\n\n"
async for chunk in llm_stream(body.messages, cfg, body.tools):
tok = chunk.get("content", "")
if tok:
collected += tok
yield "data: " + json.dumps({"id": gen_id("cmpl"), "object": "chat.completion.chunk", "created": int(time.time()),
"model": cfg.get("model", ""), "choices": [{"index": 0, "delta": {"content": tok}, "finish_reason": None}]}) + "\n\n"
yield "data: " + json.dumps({"id": gen_id("cmpl"), "object": "chat.completion.chunk", "created": int(time.time()),
"model": cfg.get("model", ""), "choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}]}) + "\n\n"
yield "data: [DONE]\n\n"
return StreamingResponse(gen(), media_type="text/event-stream")
try:
final = await llm_generate(body.messages, cfg, project_id=project["id"])
except ProviderError as exc:
raise HTTPException(status_code=502, detail=str(exc))
return {"id": gen_id("cmpl"), "object": "chat.completion", "created": int(time.time()),
"model": cfg.get("model", ""),
"choices": [{"index": 0, "message": {"role": "assistant", "content": final}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 0, "completion_tokens": max(len(final) // 4, 1), "total_tokens": max(len(final) // 4, 1)}}
@app.get("/v1/models")
async def list_models(request: Request):
project = project_from_request(request)
_, cfg = resolve_provider(project["id"], "llm")
return {"object": "list", "data": [{"id": cfg.get("model", ""), "object": "model", "owned_by": "relay"}]}
# ---- Direct TTS/STT -------------------------------------------------------
@app.post("/v1/tts")
async def tts_speak(body: TTSRequest, request: Request):
project = project_from_request(request)
provider, cfg = resolve_provider(project["id"], "tts")
if body.model:
cfg = {**cfg, "model": body.model}
if body.voice:
cfg = {**cfg, "voice": body.voice}
engine = TTSFileEngine(provider, cfg, project_id=project["id"])
try:
audio = await engine.synthesize(body.text)
except ProviderError as exc:
raise HTTPException(status_code=502, detail=str(exc))
return Response(content=audio, media_type="audio/wav")
@app.post("/v1/stt")
async def stt_transcribe(request: Request):
project = project_from_request(request)
body = await request.body()
provider, cfg = resolve_provider(project["id"], "stt")
engine = STTFileEngine(provider, cfg, project_id=project["id"])
try:
text = await engine.transcribe(body)
except ProviderError as exc:
raise HTTPException(status_code=502, detail=str(exc))
return {"text": text}
# ---- Normalized real-time voice WebSocket ---------------------------------
@app.websocket("/ws/audio/{call_id}")
async def websocket_audio(ws: WebSocket, call_id: str):
await ws.accept()
project = ws_project(ws)
if not project:
await ws.send_text(json.dumps({"type": "error", "error": "Invalid auth. Pass ?api_key= (API key) or ?token= (JWT)."}))
await ws.close()
return
agent_id = ws.query_params.get("agent_id", "")
agent = db_one("SELECT * FROM agents WHERE id=? AND project_id=?", (agent_id, project["id"])) if agent_id else None
if not agent:
await ws.send_text(json.dumps({"type": "error", "error": "Agent not found. Provide ?agent_id= or create one."}))
await ws.close()
return
agent["tools"] = json.loads(agent["tools"] or "[]")
agent["config"] = json.loads(agent["config"] or "{}")
db_exec("INSERT INTO calls (id,project_id,agent_id,status,direction,transport,started_at) VALUES (?,?,?,?,?,?,?)",
(call_id, project["id"], agent["id"], "in-progress", "inbound", "ws", utcnow()))
CALL_START[call_id] = time.time()
await dispatch_event(project["id"], call_id, "call.connected", {"agent_id": agent["id"]})
transport = WebSocketTransport(ws)
pipeline = VoicePipeline(call_id, project["id"], agent, transport, client_sample_rate=settings.sample_rate)
await pipeline.start()
await pipeline.run()
def _header_key(ws: WebSocket) -> str:
try:
return ws.headers.get("authorization", "").replace("Bearer ", "")
except Exception:
return ""
def ws_project(ws: WebSocket) -> Optional[dict[str, Any]]:
"""Resolve the authenticated project for a WebSocket connection.
Accepts either an API key (?api_key= or Authorization: Bearer <api_key>)
or a JWT (?token= or Authorization: Bearer <jwt>). Returns the project row.
"""
cred = ws.query_params.get("api_key") or ws.query_params.get("token") or _header_key(ws)
if not cred:
return None
# JWT first (looks like header.payload.signature)
if "." in cred and len(cred) < 500:
tok = decode_token(cred)
if tok and tok.get("pid"):
return db_one("SELECT * FROM projects WHERE id=?", (tok["pid"],))
# else API key
row = verify_api_key(cred)
if row:
return db_one("SELECT * FROM projects WHERE id=?", (row["project_id"],))
return None
# ---- Deepgram streaming proxies (direct STT / TTS over WS) ----------------
@app.websocket("/v1/stt/stream")
async def stt_stream_proxy(ws: WebSocket):
await ws.accept()
project = ws_project(ws)
if not project:
await ws.send_text(json.dumps({"type": "error", "error": "Invalid auth. Pass ?api_key= (API key) or ?token= (JWT)."}))
await ws.close()
return
provider, cfg = resolve_provider(project["id"], "stt")
if provider != "deepgram":
await ws.send_text(json.dumps({"type": "error", "error": "STT provider is not deepgram"}))
await ws.close()
return
sr = int(ws.query_params.get("sample_rate", settings.sample_rate))
dg = DeepgramSTTStream(cfg.get("api_key"), cfg.get("model"), cfg.get("language"), sr,
interim=settings.deepgram_interim, endpointing=settings.deepgram_endpointing,
utterance_end_ms=settings.deepgram_utterance_end_ms)
try:
await dg.connect()
except Exception as exc:
await ws.send_text(json.dumps({"type": "error", "error": f"Deepgram connect failed: {exc}"}))
await ws.close()
return
await ws.send_text(json.dumps({"type": "connected", "url": dg.url()}))
async def relay():
try:
while True:
data = await dg.recv()
if data.get("type") != "audio":
await ws.send_text(json.dumps(data))
except Exception:
pass
task = asyncio.create_task(relay())
try:
while True:
msg = await ws.receive()
if msg.get("type") == "websocket.disconnect":
break
data = msg.get("bytes") or msg.get("text")
if isinstance(data, bytes):
await dg.send_audio(data)
elif isinstance(data, str):
obj = json.loads(data)
if obj.get("type") in ("keepalive", "ping"):
await ws.send_text(json.dumps({"type": "keepalive"}))
except WebSocketDisconnect:
pass
finally:
task.cancel()
await dg.close()
try:
await ws.close()
except Exception:
pass
@app.websocket("/v1/tts/stream")
async def tts_stream_proxy(ws: WebSocket):
await ws.accept()
project = ws_project(ws)
if not project:
await ws.send_text(json.dumps({"type": "error", "error": "Invalid auth. Pass ?api_key= (API key) or ?token= (JWT)."}))
await ws.close()
return
provider, cfg = resolve_provider(project["id"], "tts")
if provider != "deepgram":
await ws.send_text(json.dumps({"type": "error", "error": "TTS provider is not deepgram"}))
await ws.close()
return
sr = int(ws.query_params.get("sample_rate", 16000))
dg = DeepgramTTSStream(cfg.get("api_key"), cfg.get("model"), sample_rate=sr)
try:
await dg.connect()
except Exception as exc:
await ws.send_text(json.dumps({"type": "error", "error": f"Deepgram connect failed: {exc}"}))
await ws.close()
return
await ws.send_text(json.dumps({"type": "connected", "url": dg.url(), "version": dg.version}))
async def relay():
try:
while True:
data = await dg.recv()
if data.get("type") == "audio":
await ws.send_bytes(data["data"])
else:
await ws.send_text(json.dumps(data))
except Exception:
pass
task = asyncio.create_task(relay())
try:
while True:
msg = await ws.receive()
if msg.get("type") == "websocket.disconnect":
break
data = msg.get("text") or msg.get("bytes")
if isinstance(data, str):
try:
obj = json.loads(data)
except Exception:
continue
if obj.get("type") == "Speak":
await dg.speak(obj.get("text", ""), obj.get("text_id", gen_id("tid")))
if dg.version == 1:
await dg.flush()
elif obj.get("type") == "Flush":
await dg.flush()
elif obj.get("type") == "Clear":
await dg.clear()
elif obj.get("type") == "Close":
break
except WebSocketDisconnect:
pass
finally:
task.cancel()
await dg.close()
try:
await ws.close()
except Exception:
pass
# ---- Twilio telephony -----------------------------------------------------
@app.post("/twilio/voice")
async def twilio_voice(request: Request):
form = await request.form()
call_sid = form.get("CallSid") or gen_id("call")
agent_id = form.get("agent_id") or request.query_params.get("agent_id", "")
api_key = request.query_params.get("api_key") or extract_api_key(request) or ""
ws_url = settings.api_base_url.replace("http", "ws", 1).rstrip("/")
stream_url = f"{ws_url}/twilio/stream/{call_sid}?api_key={api_key}&agent_id={agent_id}"
twiml = f"""<?xml version="1.0" encoding="UTF-8"?>
<Response>
<Connect>
<Stream url="{stream_url}">
<Parameter name="agentId" value="{agent_id}"/>
</Stream>
</Connect>
</Response>"""
return Response(content=twiml, media_type="application/xml")
@app.websocket("/twilio/stream/{call_sid}")
async def twilio_stream(ws: WebSocket, call_sid: str):
await ws.accept()
project = ws_project(ws)
if not project:
await ws.close()
return
agent_id = ws.query_params.get("agent_id", "")
agent = db_one("SELECT * FROM agents WHERE id=? AND project_id=?", (agent_id, project["id"])) if agent_id else None
if not agent:
await ws.close()
return
agent["tools"] = json.loads(agent["tools"] or "[]")
agent["config"] = json.loads(agent["config"] or "{}")
db_exec("INSERT INTO calls (id,project_id,agent_id,status,direction,transport,started_at) VALUES (?,?,?,?,?,?,?)",
(call_sid, project["id"], agent["id"], "in-progress", "inbound", "twilio", utcnow()))
CALL_START[call_sid] = time.time()
await dispatch_event(project["id"], call_sid, "call.connected", {"agent_id": agent["id"], "transport": "twilio"})
transport = TwilioTransport(ws)
pipeline = VoicePipeline(call_sid, project["id"], agent, transport, client_sample_rate=16000)
await pipeline.start()
await pipeline.run()
# ---- Tool handler registration (code, not DB) -----------------------------
@app.post("/tool/{tool_id}/handler")
async def attach_handler(tool_id: str, request: Request):
project = project_from_request(request)
tool = db_one("SELECT * FROM tools WHERE id=? AND project_id=?", (tool_id, project["id"]))
if not tool:
raise HTTPException(status_code=404, detail="Tool not found")
# Registering an inline code handler is not allowed over HTTP (security).
return {"ok": True, "note": "Attach code handlers via TOOL_HANDLERS[name] in the engine."}
# ---------------------------------------------------------------------------
def main():
init_db()
uvicorn.run(app, host=settings.host, port=settings.port, log_level=settings.log_level.lower())
if __name__ == "__main__":
main()