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

445 lines
18 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 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 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)
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._pending = bytearray()
self._pts = 0
self._start = time.monotonic()
def enqueue(self, pcm: bytes) -> None:
for offset in range(0, len(pcm), FRAME_SAMPLES * 2):
try:
self._queue.put_nowait(pcm[offset : offset + FRAME_SAMPLES * 2])
except asyncio.QueueFull:
break
async def recv(self) -> AudioFrame:
await asyncio.sleep(max(0, self._start + self._pts / SAMPLE_RATE - time.monotonic()))
if not self._pending:
try:
self._pending.extend(self._queue.get_nowait())
except asyncio.QueueEmpty:
pass
chunk = bytes(self._pending[: FRAME_SAMPLES * 2])
del self._pending[: FRAME_SAMPLES * 2]
chunk = chunk.ljust(FRAME_SAMPLES * 2, b"\x00")
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.pending_reply: str | None = None
self.turn_task: asyncio.Task | None = None
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 == "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 Exception:
return
async def on_audio(self, pcm: bytes) -> None:
if self.busy or not pcm:
return
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"})
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"})
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", "whisper-1")
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})
self.emit({"type": "conversation.item.input_audio_transcription.completed",
"item_id": item_id, "transcript": text})
await self.consult(text)
except asyncio.CancelledError:
self.busy = False
raise
except Exception:
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),
}],
}})
async def speak(self, text: str) -> None:
if not text.strip():
self.busy = False
return
response_id = f"resp_{uuid.uuid4().hex}"
try:
self.emit({"type": "response.created", "response": {"id": response_id}})
async with self.http.post(
os.environ["ATHENA_API_BASE_URL"].rstrip("/") + "/audio/speech",
headers={"Content-Type": "application/json", **self.auth_headers()},
json={"input": text, "voice": os.environ.get("TTS_VOICE", "alloy"),
"response_format": "wav"}, timeout=120,
) as response:
response.raise_for_status()
pcm = read_wav(await response.read())
self.track.enqueue(pcm)
self.emit({"type": "response.output_audio_transcript.done", "transcript": text})
# Keep half-duplex active until the queued audio has actually played.
await asyncio.sleep(len(pcm) / (SAMPLE_RATE * 2) + 0.2)
self.emit({"type": "response.done", "response": {
"id": response_id, "status": "completed", "output": [
{"id": f"item_{uuid.uuid4().hex}", "type": "message", "role": "assistant", "status": "completed"}
],
}})
except asyncio.CancelledError:
self.emit({"type": "response.cancelled", "response": {"id": response_id}})
raise
except Exception:
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")))