Files
AI-Profile-Router/services/athena-realtime-voice/server.py
T

500 lines
21 KiB
Python

"""OpenAI Realtime WebRTC subset backed by Athena STT/TTS and OpenClaw tools.
The browser's OpenAI-compatible data channel delegates reasoning to OpenClaw's
``openclaw_agent_consult`` tool. This server never calls the LLM directly.
"""
from __future__ import annotations
import asyncio
import base64
import hashlib
import hmac
import io
import json
import logging
import os
import time
import uuid
import wave
from fractions import Fraction
import numpy as np
from aiohttp import ClientSession, FormData, web
from aiortc import MediaStreamTrack, RTCPeerConnection, RTCSessionDescription
from aiortc.mediastreams import MediaStreamError
from av import AudioFrame, AudioResampler
from cryptography.exceptions import InvalidSignature
from cryptography.hazmat.primitives.serialization import load_pem_public_key
SAMPLE_RATE = 24000
FRAME_SAMPLES = 480
MAX_OFFER_BYTES = 64 * 1024
SECRET = web.AppKey("secret", str)
ORIGIN = web.AppKey("origin", str)
USED_TOKENS = web.AppKey("used_tokens", dict)
PEERS = web.AppKey("peers", set)
HTTP = web.AppKey("http", ClientSession)
PUBLIC_KEY_URL = web.AppKey("public_key_url", str)
logger = logging.getLogger("athena_realtime_voice")
logger.setLevel(logging.INFO)
logger.propagate = False
_handler = logging.StreamHandler()
_handler.setFormatter(logging.Formatter("%(asctime)s %(levelname)s %(message)s"))
logger.addHandler(_handler)
def decode_payload(payload: str, timestamp: int) -> dict:
try:
body = json.loads(base64.urlsafe_b64decode(payload + "=" * (-len(payload) % 4)))
if not isinstance(body, dict) or not isinstance(body.get("jti"), str):
raise ValueError("invalid payload")
if not timestamp < int(body["exp"]) <= timestamp + 65:
raise ValueError("expired token")
return body
except (TypeError, KeyError, ValueError, UnicodeError, json.JSONDecodeError) as exc:
raise web.HTTPUnauthorized(text="Invalid realtime session token") from exc
def decode_token(token: str, secret: str, now: int | None = None) -> dict:
"""Validate the short-lived, signed browser session without exposing a router key."""
try:
payload, signature = token.split(".", 1)
expected = base64.urlsafe_b64encode(
hmac.new(secret.encode(), payload.encode(), hashlib.sha256).digest()
).rstrip(b"=").decode()
if not hmac.compare_digest(signature, expected):
raise ValueError("invalid signature")
timestamp = int(time.time()) if now is None else now
return decode_payload(payload, timestamp)
except (TypeError, KeyError, ValueError, UnicodeError, json.JSONDecodeError) as exc:
raise web.HTTPUnauthorized(text="Invalid realtime session token") from exc
async def decode_public_key_token(token: str, http: ClientSession, url: str) -> dict:
try:
payload, signature = token.split(".", 1)
if "." in signature:
raise ValueError("invalid token")
async with http.get(url, timeout=5) as response:
response.raise_for_status()
key_pem = await response.read()
if len(key_pem) > 4096:
raise ValueError("public key too large")
public_key = load_pem_public_key(key_pem)
public_key.verify(
base64.urlsafe_b64decode(signature + "=" * (-len(signature) % 4)),
payload.encode(),
)
return decode_payload(payload, int(time.time()))
except (ValueError, InvalidSignature, TypeError) as exc:
raise web.HTTPUnauthorized(text="Invalid realtime session token") from exc
except web.HTTPException:
raise
except Exception as exc:
raise web.HTTPServiceUnavailable(text="OpenClaw public key unavailable") from exc
def pcm_rms(data: bytes) -> float:
samples = np.frombuffer(data[: len(data) // 2 * 2], dtype="<i2")
if not samples.size:
return 0.0
normalized = samples.astype(np.float32) / 32768.0
return float(np.sqrt(np.mean(normalized * normalized)))
def make_wav(data: bytes) -> bytes:
output = io.BytesIO()
with wave.open(output, "wb") as wav:
wav.setnchannels(1)
wav.setsampwidth(2)
wav.setframerate(SAMPLE_RATE)
wav.writeframes(data)
return output.getvalue()
def read_wav(data: bytes) -> bytes:
with wave.open(io.BytesIO(data), "rb") as wav:
if wav.getsampwidth() != 2:
raise ValueError("Athena TTS must return 16-bit WAV")
channels, rate = wav.getnchannels(), wav.getframerate()
samples = np.frombuffer(wav.readframes(wav.getnframes()), dtype="<i2")
if channels > 1:
samples = samples.reshape(-1, channels)[:, 0]
if rate != SAMPLE_RATE:
count = max(1, round(samples.size * SAMPLE_RATE / rate))
samples = np.interp(
np.arange(count) * rate / SAMPLE_RATE,
np.arange(samples.size),
samples,
).astype("<i2")
return samples.astype("<i2").tobytes()
class SpeechTrack(MediaStreamTrack):
kind = "audio"
def __init__(self):
super().__init__()
self._queue: asyncio.Queue[bytes] = asyncio.Queue(maxsize=2000)
self._frames_enqueued = 0
self._frames_played = 0
self._pts = 0
self._start = time.monotonic()
async def enqueue(self, pcm: bytes) -> int:
if len(pcm) != FRAME_SAMPLES * 2:
raise ValueError("SpeechTrack expects one complete audio frame")
await self._queue.put(pcm)
self._frames_enqueued += 1
return self._frames_enqueued
async def wait_played(self, target: int) -> None:
while self._frames_played < target:
await asyncio.sleep(0.02)
await asyncio.sleep(0.2)
async def recv(self) -> AudioFrame:
await asyncio.sleep(max(0, self._start + self._pts / SAMPLE_RATE - time.monotonic()))
try:
chunk = self._queue.get_nowait()
self._frames_played += 1
except asyncio.QueueEmpty:
chunk = b"\x00" * FRAME_SAMPLES * 2
frame = AudioFrame.from_ndarray(
np.frombuffer(chunk, dtype="<i2").reshape(1, FRAME_SAMPLES),
format="s16", layout="mono",
)
frame.sample_rate = SAMPLE_RATE
frame.pts = self._pts
frame.time_base = Fraction(1, SAMPLE_RATE)
self._pts += FRAME_SAMPLES
return frame
class RealtimeSession:
def __init__(self, pc: RTCPeerConnection, http: ClientSession, track: SpeechTrack):
self.pc, self.http, self.track = pc, http, track
self.channel = None
self.busy = False
self.speaking = False
self.speech = bytearray()
self.silence_since: float | None = None
self.call_id: str | None = None
self.last_item_id: str | None = None
self.pending_reply: str | None = None
self.turn_task: asyncio.Task | None = None
self.dropped_speech_while_busy = False
self.vad_threshold = float(os.environ.get("VAD_THRESHOLD", "0.018"))
self.silence_ms = int(os.environ.get("SILENCE_MS", "750"))
def emit(self, event: dict) -> None:
if self.channel and self.channel.readyState == "open":
self.channel.send(json.dumps(event, ensure_ascii=False))
def on_message(self, message: str | bytes) -> None:
if not isinstance(message, str) or len(message) > 262144:
return
try:
event = json.loads(message)
except json.JSONDecodeError:
return
kind = event.get("type")
if kind in {"conversation.item.create", "response.create", "response.cancel"}:
logger.info("client event=%s item_type=%s call_matches=%s busy=%s pending_reply=%s",
kind, (event.get("item") or {}).get("type"),
(event.get("item") or {}).get("call_id") == self.call_id,
self.busy, self.pending_reply is not None)
if kind == "session.update":
self.emit({"type": "session.updated", "session": {"id": "athena-local"}})
elif kind == "conversation.item.create":
item = event.get("item") or {}
if item.get("type") == "function_call_output" and item.get("call_id") == self.call_id:
try:
answer = json.loads(item.get("output") or "{}")
self.pending_reply = str(answer.get("result") or answer.get("error") or "")
except (ValueError, TypeError):
self.pending_reply = "Die Antwort konnte nicht gelesen werden."
elif item.get("type") == "message" and item.get("role") == "user":
text = " ".join(str(part.get("text") or "") for part in item.get("content", []))
if text.strip() and not self.busy:
self.turn_task = asyncio.create_task(self.consult(text.strip()))
elif kind == "response.create" and self.pending_reply is not None:
text, self.pending_reply = self.pending_reply, None
self.turn_task = asyncio.create_task(self.speak(text))
elif kind == "response.cancel":
if self.turn_task:
self.turn_task.cancel()
async def consume_audio(self, remote: MediaStreamTrack) -> None:
resampler = AudioResampler(format="s16", layout="mono", rate=SAMPLE_RATE)
try:
while True:
frame = await remote.recv()
for converted in resampler.resample(frame):
pcm = bytes(converted.planes[0])[: converted.samples * 2]
await self.on_audio(pcm)
except asyncio.CancelledError:
raise
except MediaStreamError:
return
except Exception:
logger.exception("incoming audio stream stopped unexpectedly")
return
async def on_audio(self, pcm: bytes) -> None:
if not pcm:
return
if self.busy:
if not self.dropped_speech_while_busy and pcm_rms(pcm) >= self.vad_threshold:
self.dropped_speech_while_busy = True
logger.info("microphone speech ignored while response is busy")
return
self.dropped_speech_while_busy = False
voiced = pcm_rms(pcm) >= self.vad_threshold
now = time.monotonic()
if voiced and not self.speaking:
self.speaking = True
self.speech.clear()
self.silence_since = None
self.emit({"type": "input_audio_buffer.speech_started"})
logger.info("speech started")
if not self.speaking:
return
self.speech.extend(pcm)
if voiced:
self.silence_since = None
elif self.silence_since is None:
self.silence_since = now
elapsed_silence = (now - self.silence_since) * 1000 if self.silence_since else 0
if elapsed_silence >= self.silence_ms or len(self.speech) >= SAMPLE_RATE * 2 * 45:
self.speaking = False
self.emit({"type": "input_audio_buffer.speech_stopped"})
logger.info("speech stopped bytes=%s", len(self.speech))
audio = bytes(self.speech)
self.speech.clear()
if len(audio) >= SAMPLE_RATE * 2 // 4:
self.busy = True
self.turn_task = asyncio.create_task(self.transcribe_and_consult(audio))
async def transcribe_and_consult(self, pcm: bytes) -> None:
try:
form = FormData()
form.add_field("file", make_wav(pcm), filename="talk.wav", content_type="audio/wav")
form.add_field("model", "qwen3-asr")
form.add_field("language", os.environ.get("STT_LANGUAGE", "de"))
async with self.http.post(
os.environ["ATHENA_API_BASE_URL"].rstrip("/") + "/audio/transcriptions",
headers=self.auth_headers(), data=form, timeout=90,
) as response:
response.raise_for_status()
text = str((await response.json()).get("text") or "").strip()
if not text:
self.busy = False
return
item_id = f"item_{uuid.uuid4().hex}"
# OpenClaw tracks transcript items before accepting their text.
self.emit({"type": "input_audio_buffer.committed", "item_id": item_id,
"previous_item_id": self.last_item_id})
self.last_item_id = item_id
self.emit({"type": "conversation.item.input_audio_transcription.completed",
"item_id": item_id, "transcript": text})
logger.info("transcript completed item_id=%s", item_id)
await self.consult(text)
except asyncio.CancelledError:
self.busy = False
raise
except Exception:
logger.exception("speech transcription failed")
self.busy = False
self.emit({"type": "error", "error": {"message": "Athena STT failed"}})
async def consult(self, text: str) -> None:
self.busy = True
self.call_id = f"call_{uuid.uuid4().hex}"
response_id = f"resp_{uuid.uuid4().hex}"
self.emit({"type": "response.created", "response": {"id": response_id}})
self.emit({"type": "response.done", "response": {
"id": response_id, "status": "completed", "output": [{
"id": f"item_{uuid.uuid4().hex}", "type": "function_call",
"status": "completed", "call_id": self.call_id,
"name": "openclaw_agent_consult",
"arguments": json.dumps({"prompt": text}, ensure_ascii=False),
}],
}})
logger.info("consult requested call_id=%s response_id=%s", self.call_id, response_id)
async def speak(self, text: str) -> None:
if not text.strip():
self.busy = False
return
response_id = f"resp_{uuid.uuid4().hex}"
item_id = f"item_{uuid.uuid4().hex}"
try:
logger.info("speaking response_id=%s", response_id)
self.emit({"type": "response.created", "response": {"id": response_id}})
async with self.http.post(
os.environ["ATHENA_API_BASE_URL"].rstrip("/") + "/audio/speech/pcm-stream",
headers={"Content-Type": "application/json", **self.auth_headers()},
json={"input": text, "voice": os.environ.get("TTS_VOICE", "alloy"),
"chunk_size": 4}, timeout=120,
) as response:
response.raise_for_status()
item = {"id": item_id, "type": "message", "role": "assistant",
"status": "completed"}
self.emit({"type": "conversation.item.created", "item": item,
"previous_item_id": self.last_item_id})
self.last_item_id = item_id
self.emit({"type": "response.output_audio_transcript.done",
"item_id": item_id, "transcript": text})
pending = bytearray()
target = 0
frame_bytes = FRAME_SAMPLES * 2
async for chunk in response.content.iter_chunked(16384):
pending.extend(chunk)
while len(pending) >= frame_bytes:
target = await self.track.enqueue(bytes(pending[:frame_bytes]))
del pending[:frame_bytes]
if len(pending) % 2:
raise ValueError("Athena TTS returned incomplete PCM samples")
if pending:
target = await self.track.enqueue(bytes(pending).ljust(frame_bytes, b"\0"))
if not target:
raise ValueError("Athena TTS returned no audio")
# Keep half-duplex active until the streamed audio has played.
await self.track.wait_played(target)
self.emit({"type": "conversation.item.done", "item": item})
self.emit({"type": "response.done", "response": {
"id": response_id, "status": "completed", "output": [item],
}})
logger.info("speech response complete response_id=%s", response_id)
except asyncio.CancelledError:
self.emit({"type": "response.cancelled", "response": {"id": response_id}})
raise
except Exception:
logger.exception("speech synthesis failed")
self.emit({"type": "error", "error": {"message": "Athena TTS failed"}})
finally:
self.busy = False
@staticmethod
def auth_headers() -> dict[str, str]:
key = os.environ.get("ATHENA_API_KEY", "")
return {"Authorization": f"Bearer {key}"} if key else {}
async def offer(request: web.Request) -> web.Response:
if request.headers.get("Origin") not in (None, request.app[ORIGIN]):
raise web.HTTPForbidden(text="Origin not allowed")
token = request.headers.get("Authorization", "").removeprefix("Bearer ")
if request.app[PUBLIC_KEY_URL]:
claims = await decode_public_key_token(token, request.app[HTTP], request.app[PUBLIC_KEY_URL])
else:
claims = decode_token(token, request.app[SECRET])
used = request.app[USED_TOKENS]
now = time.time()
for key, expiry in list(used.items()):
if expiry < now:
del used[key]
if claims["jti"] in used:
raise web.HTTPUnauthorized(text="Realtime token already used")
if len(request.app[PEERS]) >= int(os.environ.get("MAX_SESSIONS", "4")):
raise web.HTTPServiceUnavailable(text="Too many realtime sessions")
used[claims["jti"]] = claims["exp"]
if request.content_length is not None and request.content_length > MAX_OFFER_BYTES:
raise web.HTTPRequestEntityTooLarge(max_size=MAX_OFFER_BYTES, actual_size=request.content_length)
sdp = await request.text()
if len(sdp.encode()) > MAX_OFFER_BYTES:
raise web.HTTPRequestEntityTooLarge(max_size=MAX_OFFER_BYTES, actual_size=len(sdp))
pc = RTCPeerConnection()
speech = SpeechTrack()
pc.addTrack(speech)
session = RealtimeSession(pc, request.app[HTTP], speech)
request.app[PEERS].add(pc)
@pc.on("datachannel")
def on_channel(channel):
session.channel = channel
@channel.on("open")
def on_open():
session.emit({"type": "session.created", "session": {"id": "athena-local"}})
@channel.on("message")
def on_message(message):
session.on_message(message)
@pc.on("track")
def on_track(track):
if track.kind == "audio":
asyncio.create_task(session.consume_audio(track))
@pc.on("connectionstatechange")
async def on_state_change():
if pc.connectionState in {"closed", "failed"}:
if session.turn_task:
session.turn_task.cancel()
await pc.close()
request.app[PEERS].discard(pc)
try:
await pc.setRemoteDescription(RTCSessionDescription(sdp=sdp, type="offer"))
await pc.setLocalDescription(await pc.createAnswer())
return web.Response(text=pc.localDescription.sdp, content_type="application/sdp",
headers={"Access-Control-Allow-Origin": request.app[ORIGIN]})
except Exception:
await pc.close()
request.app[PEERS].discard(pc)
raise web.HTTPBadRequest(text="Invalid WebRTC offer")
async def preflight(request: web.Request) -> web.Response:
if request.headers.get("Origin") != request.app[ORIGIN]:
raise web.HTTPForbidden(text="Origin not allowed")
return web.Response(headers={
"Access-Control-Allow-Origin": request.app[ORIGIN],
"Access-Control-Allow-Methods": "POST, OPTIONS",
"Access-Control-Allow-Headers": "Authorization, Content-Type",
"Access-Control-Max-Age": "600",
})
async def lifecycle(app: web.Application):
app[HTTP] = ClientSession()
yield
await asyncio.gather(*(pc.close() for pc in app[PEERS]))
await app[HTTP].close()
async def health(_: web.Request) -> web.Response:
return web.json_response({"ok": True})
def create_app() -> web.Application:
secret = os.environ.get("ATHENA_TALK_REALTIME_SECRET", "")
public_key_url = os.environ.get("OPENCLAW_PUBLIC_KEY_URL", "")
if public_key_url and not public_key_url.startswith("https://"):
raise RuntimeError("OPENCLAW_PUBLIC_KEY_URL must use HTTPS")
if not public_key_url and len(secret) < 32:
raise RuntimeError("Set OPENCLAW_PUBLIC_KEY_URL or a 32+ character secret")
if not os.environ.get("ATHENA_API_BASE_URL"):
raise RuntimeError("ATHENA_API_BASE_URL is required")
origin = os.environ.get("OPENCLAW_ORIGIN", "")
if not origin.startswith("https://"):
raise RuntimeError("OPENCLAW_ORIGIN must be an HTTPS origin")
app = web.Application(client_max_size=MAX_OFFER_BYTES)
app[SECRET] = secret
app[PUBLIC_KEY_URL] = public_key_url
app[ORIGIN] = origin
app[USED_TOKENS] = {}
app[PEERS] = set()
app.cleanup_ctx.append(lifecycle)
app.router.add_route("OPTIONS", "/v1/realtime/calls", preflight)
app.router.add_post("/v1/realtime/calls", offer)
app.router.add_get("/health", health)
return app
if __name__ == "__main__":
web.run_app(create_app(), host="0.0.0.0", port=int(os.environ.get("PORT", "8090")))