588 lines
21 KiB
Python
588 lines
21 KiB
Python
#!/usr/bin/env python3
|
|
"""The Collector — a smooth-operator conversational voice AI.
|
|
|
|
Pipeline (latency-optimized):
|
|
STT faster-whisper small.en (CPU, int8) ~1-2s / short utterance
|
|
LLM ornith-1.5:9b-64k @ .222 (warm) ~600ms first token, streamed
|
|
TTS Piper lessac-high (CPU, 22.05k->16k) ~400ms / sentence, overlapped w/ LLM
|
|
|
|
Interfaces:
|
|
WS /ws browser full-duplex voice (PCM16 16 kHz)
|
|
WS /signalwire/stream phone media stream (mulaw 8 kHz)
|
|
POST /api/call {to, task} outbound phone call on request
|
|
POST /api/text {text} text chat (returns text)
|
|
GET /api/health
|
|
"""
|
|
import asyncio
|
|
import base64
|
|
import json
|
|
import logging
|
|
import os
|
|
import re
|
|
|
|
import audioop
|
|
import httpx
|
|
import numpy as np
|
|
import yaml
|
|
from fastapi import FastAPI, Request, WebSocket, WebSocketDisconnect
|
|
from fastapi.responses import FileResponse, JSONResponse, PlainTextResponse
|
|
|
|
logging.basicConfig(level=logging.INFO,
|
|
format="%(asctime)s %(name)s %(levelname)s %(message)s")
|
|
log = logging.getLogger("collector")
|
|
|
|
BASE = "/opt/collector"
|
|
|
|
|
|
def load_config() -> dict:
|
|
with open(os.path.join(BASE, "config.yaml")) as f:
|
|
return yaml.safe_load(f)
|
|
|
|
|
|
CFG = load_config()
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# audio helpers (stdlib audioop)
|
|
# ---------------------------------------------------------------------------
|
|
def ulaw_to_pcm(data: bytes) -> bytes:
|
|
return audioop.ulaw2lin(data, 2)
|
|
|
|
|
|
def pcm_to_ulaw(pcm: bytes) -> bytes:
|
|
return audioop.lin2ulaw(pcm, 2)
|
|
|
|
|
|
def resample(pcm: bytes, src: int, dst: int) -> bytes:
|
|
if src == dst or not pcm:
|
|
return pcm
|
|
out, _ = audioop.ratecv(pcm, 2, 1, src, dst, None)
|
|
return out
|
|
|
|
|
|
def rms(pcm: bytes) -> int:
|
|
return audioop.rms(pcm, 2) if pcm else 0
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# STT
|
|
# ---------------------------------------------------------------------------
|
|
class STT:
|
|
def __init__(self, cfg):
|
|
self.cfg = cfg
|
|
self.model = None
|
|
self.remote_url = (cfg.get("remote_url") or "").rstrip("/")
|
|
|
|
def load(self):
|
|
# Remote GPU STT (nightmare) is primary; local whisper is a lazy fallback.
|
|
return self
|
|
|
|
def transcribe(self, pcm16k: bytes) -> str:
|
|
if self.remote_url:
|
|
try:
|
|
return self._remote(pcm16k)
|
|
except Exception as e:
|
|
log.warning("remote STT failed (%s); falling back to local", e)
|
|
return self._local(pcm16k)
|
|
|
|
def _remote(self, pcm16k: bytes) -> str:
|
|
import requests
|
|
r = requests.post(self.remote_url + "/transcribe", data=pcm16k, timeout=30)
|
|
r.raise_for_status()
|
|
return (r.json().get("text") or "").strip()
|
|
|
|
def _local(self, pcm16k: bytes) -> str:
|
|
if self.model is None:
|
|
from faster_whisper import WhisperModel
|
|
c = self.cfg
|
|
log.info("Loading local faster-whisper %s (fallback)...", c["model"])
|
|
self.model = WhisperModel(c["model"], device=c["device"],
|
|
compute_type=c["compute_type"],
|
|
cpu_threads=c.get("cpu_threads", 2))
|
|
audio = np.frombuffer(pcm16k, dtype=np.int16).astype(np.float32) / 32768.0
|
|
segs, _ = self.model.transcribe(
|
|
audio, beam_size=1, language="en", vad_filter=True,
|
|
condition_on_previous_text=False, no_speech_threshold=0.6)
|
|
return " ".join(s.text.strip() for s in segs).strip()
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# TTS
|
|
# ---------------------------------------------------------------------------
|
|
class TTS:
|
|
def __init__(self, cfg):
|
|
self.cfg = cfg
|
|
self.voice_path = cfg.get("piper_voice") or \
|
|
"/opt/collector/voices/en_US-lessac-high.onnx"
|
|
self._piper = None
|
|
|
|
def load(self):
|
|
if self._piper is None:
|
|
# onnxruntime pins threads to CPUs 0..N-1; the unprivileged LXC
|
|
# cpuset (e.g. 8,17) rejects that. Setting an explicit thread count
|
|
# disables the affinity pin (which is otherwise just log noise).
|
|
import onnxruntime
|
|
_orig = onnxruntime.InferenceSession.__init__
|
|
|
|
def _patched(self, *a, **k):
|
|
so = k.get("sess_options") or onnxruntime.SessionOptions()
|
|
so.intra_op_num_threads = 2
|
|
k["sess_options"] = so
|
|
_orig(self, *a, **k)
|
|
|
|
onnxruntime.InferenceSession.__init__ = _patched
|
|
|
|
from piper import PiperVoice
|
|
log.info("Loading Piper voice %s ...", self.voice_path)
|
|
self._piper = PiperVoice.load(self.voice_path)
|
|
log.info("TTS ready")
|
|
return self._piper
|
|
|
|
def synthesize_16k(self, text: str) -> bytes:
|
|
"""Return PCM16 mono 16 kHz for `text` (piper yields AudioChunk iterators)."""
|
|
if not text.strip():
|
|
return b""
|
|
speed = float(self.cfg.get("speed", 1.05))
|
|
from piper import SynthesisConfig
|
|
syn = SynthesisConfig(length_scale=1.0 / speed)
|
|
chunks = []
|
|
rate = 22050
|
|
for chunk in self._piper.synthesize(text, syn_config=syn):
|
|
audio = (chunk.audio_float_array * 32767).astype(np.int16).tobytes()
|
|
rate = chunk.sample_rate
|
|
chunks.append(audio)
|
|
pcm = b"".join(chunks)
|
|
return resample(pcm, rate, 16000)
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# LLM (ornith streaming, /api/chat, think:false)
|
|
# ---------------------------------------------------------------------------
|
|
async def stream_llm(messages, cfg):
|
|
payload = {
|
|
"model": cfg["ollama"]["model"],
|
|
"messages": messages,
|
|
"stream": True,
|
|
"think": False,
|
|
"keep_alive": "30m",
|
|
"options": {
|
|
"temperature": cfg["ollama"].get("temperature", 0.5),
|
|
"num_predict": cfg["ollama"].get("num_predict", 256),
|
|
},
|
|
}
|
|
urls = [cfg["ollama"]["base_url"], cfg["ollama"].get("fallback_url", "")]
|
|
for base in urls:
|
|
if not base:
|
|
continue
|
|
url = base.rstrip("/") + "/api/chat"
|
|
try:
|
|
async with httpx.AsyncClient(
|
|
timeout=httpx.Timeout(10.0, read=60.0)) as client:
|
|
async with client.stream("POST", url, json=payload) as r:
|
|
r.raise_for_status()
|
|
async for line in r.aiter_lines():
|
|
if not line.strip():
|
|
continue
|
|
try:
|
|
obj = json.loads(line)
|
|
except Exception:
|
|
continue
|
|
if obj.get("done"):
|
|
return
|
|
msg = obj.get("message", {})
|
|
content = msg.get("content") or msg.get("thinking") or ""
|
|
if content:
|
|
yield content
|
|
return
|
|
except Exception as e:
|
|
log.warning("LLM host %s failed: %s", base, e)
|
|
log.error("All LLM hosts failed")
|
|
|
|
|
|
async def sentence_stream(chunks):
|
|
"""Yield sentences as they complete, enabling TTS/LLM overlap."""
|
|
buf = ""
|
|
async for chunk in chunks:
|
|
buf += chunk
|
|
while True:
|
|
m = re.search(r"[.!?\n]", buf)
|
|
if not m:
|
|
break
|
|
idx = m.end()
|
|
sentence = buf[:idx].strip()
|
|
buf = buf[idx:].strip()
|
|
if sentence:
|
|
yield sentence
|
|
if buf.strip():
|
|
yield buf.strip()
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Conversation session (shared by browser WS + phone stream)
|
|
# ---------------------------------------------------------------------------
|
|
SPEECH_RMS = 400 # 16-bit PCM: voice onset
|
|
BARGE_RMS = 900 # sustained level that interrupts playback
|
|
SILENCE_END_MS = 700 # ms of quiet after speech before we turn
|
|
FRAME_MS = 20
|
|
MAX_UTT_BYTES = 18 * 16000 * 2 # 18s cap
|
|
|
|
|
|
class Session:
|
|
def __init__(self, cfg, system_prompt):
|
|
self.cfg = cfg
|
|
self.stt = STT(cfg["stt"])
|
|
self.tts = TTS(cfg["tts"])
|
|
self.messages = [{"role": "system", "content": system_prompt}]
|
|
self.buf = bytearray()
|
|
self.talking = False
|
|
self.silence_ms = 0
|
|
self._barge = False
|
|
self._processing = False
|
|
self._queue = asyncio.Queue()
|
|
self._worker = None
|
|
|
|
async def warm(self):
|
|
await asyncio.to_thread(self.stt.load)
|
|
await asyncio.to_thread(self.tts.load)
|
|
try:
|
|
async for _ in stream_llm(
|
|
[{"role": "user", "content": "hi"}], self.cfg):
|
|
break
|
|
except Exception as e:
|
|
log.warning("LLM warmup: %s", e)
|
|
log.info("Collector warmed (STT + TTS + LLM)")
|
|
|
|
def start_worker(self, send_audio, send_text):
|
|
"""Begin the continuous conversation loop for this connection."""
|
|
self._worker = asyncio.create_task(
|
|
self._process_loop(send_audio, send_text))
|
|
|
|
async def feed_audio(self, pcm16k: bytes):
|
|
"""Continuous VAD. Call for every 20ms frame. Never drops speech —
|
|
even during a reply, overlapping audio is buffered and processed next
|
|
(barge-in)."""
|
|
level = rms(pcm16k)
|
|
if level > SPEECH_RMS:
|
|
self.talking = True
|
|
self.silence_ms = 0
|
|
if self._processing:
|
|
self._barge = True # user started talking over the reply
|
|
elif self.talking:
|
|
self.silence_ms += FRAME_MS
|
|
self.buf.extend(pcm16k)
|
|
if self.talking and (self.silence_ms >= SILENCE_END_MS or
|
|
len(self.buf) > MAX_UTT_BYTES):
|
|
self.talking = False
|
|
self.silence_ms = 0
|
|
pcm = bytes(self.buf)
|
|
self.buf.clear()
|
|
self._queue.put_nowait(pcm)
|
|
|
|
async def _process_loop(self, send_audio, send_text):
|
|
while True:
|
|
pcm = await self._queue.get()
|
|
self._processing = True
|
|
self._barge = False
|
|
text = await asyncio.to_thread(self.stt.transcribe, pcm)
|
|
if text:
|
|
await send_text(text, "user")
|
|
self.messages.append({"role": "user", "content": text})
|
|
await self._reply(send_audio, send_text)
|
|
self._processing = False
|
|
|
|
async def speak(self, text: str, send_audio, send_text):
|
|
"""Speak an injected message (outbound preamble), then keep listening."""
|
|
self.messages.append({"role": "user", "content": text})
|
|
self._processing = True
|
|
self._barge = False
|
|
await self._reply(send_audio, send_text)
|
|
self._processing = False
|
|
|
|
async def _reply(self, send_audio, send_text):
|
|
full = []
|
|
async for sentence in sentence_stream(stream_llm(self.messages, self.cfg)):
|
|
if self._barge:
|
|
log.info("barge-in: stopping reply")
|
|
break
|
|
full.append(sentence)
|
|
pcm = await asyncio.to_thread(self.tts.synthesize_16k, sentence)
|
|
if pcm:
|
|
await send_audio(pcm)
|
|
if full:
|
|
reply = " ".join(full)
|
|
await send_text(reply, "assistant")
|
|
self.messages.append({"role": "assistant", "content": reply})
|
|
if len(self.messages) > 24:
|
|
self.messages = [self.messages[0]] + self.messages[-23:]
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# SignalWire client
|
|
# ---------------------------------------------------------------------------
|
|
class Phone:
|
|
def __init__(self, cfg):
|
|
t = cfg.get("signalwire", {})
|
|
self.sid = t.get("account_sid", "")
|
|
self.token = t.get("auth_token", "")
|
|
self.number = t.get("phone_number", "")
|
|
self.client = None
|
|
self.tasks = {} # call_sid -> outbound task
|
|
if self.sid and self.token:
|
|
from signalwire.rest import Client as SWClient
|
|
self.client = SWClient(self.sid, self.token,
|
|
signalwire_space_url="templeofdoom.signalwire.com")
|
|
|
|
@property
|
|
def ready(self):
|
|
return self.client is not None and bool(self.number)
|
|
|
|
def create_call(self, to: str, url: str):
|
|
return self.client.calls.create(to=to, from_=self.number, url=url)
|
|
|
|
|
|
PHONE = Phone(CFG)
|
|
app = FastAPI(title="The Collector")
|
|
|
|
|
|
def system_prompt() -> str:
|
|
return CFG["agent"].get("system_prompt", "You are The Collector.")
|
|
|
|
|
|
def stream_twiml(direction: str, public_base: str) -> str:
|
|
# SignalWire <Stream> requires wss://, not https://
|
|
wss_base = public_base.replace("https://", "wss://", 1)
|
|
stream_url = f"{wss_base}/signalwire/stream?direction={direction}"
|
|
return (f'<?xml version="1.0" encoding="UTF-8"?><Response><Connect>'
|
|
f'<Stream url="{stream_url}"/></Connect></Response>')
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Health
|
|
# ---------------------------------------------------------------------------
|
|
@app.get("/")
|
|
async def index():
|
|
return FileResponse(os.path.join(BASE, "static", "index.html"))
|
|
|
|
|
|
@app.get("/api/health")
|
|
async def health():
|
|
return {
|
|
"ok": True,
|
|
"name": CFG["agent"]["name"],
|
|
"model": CFG["ollama"]["model"],
|
|
"llm": CFG["ollama"]["base_url"],
|
|
"phone": PHONE.ready,
|
|
"number": PHONE.number if PHONE.ready else None,
|
|
}
|
|
|
|
|
|
@app.post("/api/config")
|
|
async def api_config(request: Request):
|
|
"""Accept runtime config (tunnel URL, temperature, tts speed) + persist."""
|
|
body = await request.json()
|
|
changed = {}
|
|
if "public_base_url" in body:
|
|
url = (body["public_base_url"] or "").strip().rstrip("/")
|
|
CFG.setdefault("server", {})["public_base_url"] = url
|
|
changed["public_base_url"] = url
|
|
if "temperature" in body:
|
|
t = float(body["temperature"])
|
|
CFG.setdefault("ollama", {})["temperature"] = max(0.0, min(1.5, t))
|
|
changed["temperature"] = CFG["ollama"]["temperature"]
|
|
if "speed" in body:
|
|
s = float(body["speed"])
|
|
CFG.setdefault("tts", {})["speed"] = max(0.5, min(2.0, s))
|
|
changed["speed"] = CFG["tts"]["speed"]
|
|
if changed:
|
|
try:
|
|
with open(os.path.join(BASE, "config.yaml")) as f:
|
|
cfg = yaml.safe_load(f)
|
|
if "public_base_url" in changed:
|
|
cfg.setdefault("server", {})["public_base_url"] = changed["public_base_url"]
|
|
if "temperature" in changed:
|
|
cfg.setdefault("ollama", {})["temperature"] = changed["temperature"]
|
|
if "speed" in changed:
|
|
cfg.setdefault("tts", {})["speed"] = changed["speed"]
|
|
with open(os.path.join(BASE, "config.yaml"), "w") as f:
|
|
yaml.safe_dump(cfg, f, sort_keys=False)
|
|
except Exception as e:
|
|
log.warning("config persist failed: %s", e)
|
|
return {"ok": True, **changed}
|
|
return JSONResponse({"ok": False, "error": "nothing to set"}, 400)
|
|
|
|
|
|
@app.get("/api/settings")
|
|
async def api_settings():
|
|
return {
|
|
"temperature": CFG["ollama"].get("temperature", 0.5),
|
|
"speed": CFG["tts"].get("speed", 1.05),
|
|
"model": CFG["ollama"]["model"],
|
|
"voice": "lessac-high",
|
|
"llm": CFG["ollama"]["base_url"],
|
|
"stt": "GPU (nightmare)",
|
|
}
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Text chat
|
|
# ---------------------------------------------------------------------------
|
|
@app.post("/api/text")
|
|
async def api_text(request: Request):
|
|
body = await request.json()
|
|
text = (body.get("text") or "").strip()
|
|
if not text:
|
|
return JSONResponse({"ok": False, "error": "empty text"}, 400)
|
|
msgs = [{"role": "system", "content": system_prompt()},
|
|
{"role": "user", "content": text}]
|
|
out = []
|
|
async for chunk in stream_llm(msgs, CFG):
|
|
out.append(chunk)
|
|
return {"ok": True, "reply": "".join(out).strip()}
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Outbound call
|
|
# ---------------------------------------------------------------------------
|
|
@app.post("/api/call")
|
|
async def api_call(request: Request):
|
|
if not PHONE.ready:
|
|
return JSONResponse({"ok": False, "error": "SignalWire not configured"}, 400)
|
|
body = await request.json()
|
|
to = (body.get("to") or "").strip()
|
|
task = (body.get("task") or "").strip()
|
|
if not to:
|
|
return JSONResponse({"ok": False, "error": "missing 'to'"}, 400)
|
|
public = CFG["server"].get("public_base_url", "")
|
|
if not public:
|
|
return JSONResponse({"ok": False, "error": "public_base_url not set"}, 500)
|
|
call = await asyncio.to_thread(
|
|
PHONE.create_call, to, f"{public}/signalwire/outbound")
|
|
PHONE.tasks[call.sid] = task
|
|
return {"ok": True, "call_sid": call.sid, "to": to, "task": task}
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# SignalWire webhooks
|
|
# ---------------------------------------------------------------------------
|
|
@app.post("/signalwire/voice")
|
|
async def sw_inbound():
|
|
public = CFG["server"].get("public_base_url", "")
|
|
return PlainTextResponse(stream_twiml("inbound", public),
|
|
media_type="application/xml")
|
|
|
|
|
|
@app.post("/signalwire/outbound")
|
|
async def sw_outbound():
|
|
public = CFG["server"].get("public_base_url", "")
|
|
return PlainTextResponse(stream_twiml("outbound", public),
|
|
media_type="application/xml")
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Browser full-duplex voice
|
|
# ---------------------------------------------------------------------------
|
|
@app.websocket("/ws")
|
|
async def browser_ws(ws: WebSocket):
|
|
await ws.accept()
|
|
session = Session(CFG, system_prompt())
|
|
|
|
async def send_audio(pcm):
|
|
await ws.send_bytes(pcm)
|
|
|
|
async def send_text(text, role):
|
|
await ws.send_text(json.dumps({"type": "transcript", "role": role,
|
|
"text": text}))
|
|
|
|
await session.warm()
|
|
session.start_worker(send_audio, send_text)
|
|
try:
|
|
while True:
|
|
data = await ws.receive()
|
|
if data["type"] == "websocket.receive":
|
|
if data.get("bytes"):
|
|
await session.feed_audio(data["bytes"])
|
|
elif data["type"] == "websocket.disconnect":
|
|
break
|
|
except WebSocketDisconnect:
|
|
pass
|
|
finally:
|
|
if session._worker:
|
|
session._worker.cancel()
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# SignalWire media stream (phone)
|
|
# ---------------------------------------------------------------------------
|
|
async def send_phone_audio(ws, stream_sid, pcm16k):
|
|
pcm8 = resample(pcm16k, 16000, 8000)
|
|
ulaw = pcm_to_ulaw(pcm8)
|
|
frame = 160 # 20ms @ 8kHz
|
|
for i in range(0, len(ulaw), frame):
|
|
chunk = ulaw[i:i + frame]
|
|
await ws.send_text(json.dumps({
|
|
"event": "media", "streamSid": stream_sid,
|
|
"media": {"payload": base64.b64encode(chunk).decode()}}))
|
|
await asyncio.sleep(0.02)
|
|
|
|
|
|
@app.websocket("/signalwire/stream")
|
|
async def sw_stream(ws: WebSocket, direction: str = "inbound"):
|
|
await ws.accept()
|
|
session = Session(CFG, system_prompt())
|
|
stream_sid = None
|
|
call_sid = None
|
|
first = True
|
|
|
|
async def send_audio(pcm16k):
|
|
await send_phone_audio(ws, stream_sid, pcm16k)
|
|
|
|
async def send_text(text, role):
|
|
pass # no transcript channel on phone; audio only
|
|
|
|
await session.warm()
|
|
session.start_worker(send_audio, send_text)
|
|
try:
|
|
while True:
|
|
msg = await ws.receive_text()
|
|
try:
|
|
obj = json.loads(msg)
|
|
except Exception:
|
|
continue
|
|
ev = obj.get("event")
|
|
|
|
if ev == "start":
|
|
start = obj.get("start", {})
|
|
stream_sid = start.get("streamSid", stream_sid)
|
|
call_sid = start.get("callSid", call_sid)
|
|
if direction == "outbound" and first:
|
|
first = False
|
|
task = PHONE.tasks.pop(call_sid, "")
|
|
if task:
|
|
preamble = CFG["agent"].get("outbound_preamble", "") + task
|
|
await session.speak(preamble, send_audio, send_text)
|
|
|
|
elif ev == "media":
|
|
payload = obj.get("media", {}).get("payload", "")
|
|
if not payload:
|
|
continue
|
|
raw = base64.b64decode(payload)
|
|
pcm8 = ulaw_to_pcm(raw)
|
|
pcm16k = resample(pcm8, 8000, 16000)
|
|
await session.feed_audio(pcm16k)
|
|
|
|
elif ev == "stop":
|
|
break
|
|
except WebSocketDisconnect:
|
|
pass
|
|
finally:
|
|
if session._worker:
|
|
session._worker.cancel()
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Entrypoint
|
|
# ---------------------------------------------------------------------------
|
|
if __name__ == "__main__":
|
|
import uvicorn
|
|
s = CFG["server"]
|
|
uvicorn.run(app, host=s.get("host", "0.0.0.0"), port=s.get("port", 8766))
|