500 lines
21 KiB
Python
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")))
|