Add Athena WebRTC bridge for OpenClaw Talk
This commit is contained in:
@@ -0,0 +1,441 @@
|
||||
"""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
|
||||
self.emit({"type": "conversation.item.input_audio_transcription.completed",
|
||||
"item_id": f"item_{uuid.uuid4().hex}", "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")))
|
||||
Reference in New Issue
Block a user