Add Athena WebRTC bridge for OpenClaw Talk

This commit is contained in:
Mikei386 committed 2026-09-16 14:10:39 +02:00
1 parent 276df0eae5
commit e86c905968
14 files changed
+1223 -13

No files matched your search

@@ -0,0 +1,175 @@
"""Local end-to-end WebRTC smoke test with fake Athena STT and TTS endpoints."""
import asyncio
import base64
import hashlib
import hmac
import json
import os
import time
import unittest
import uuid
from fractions import Fraction
import numpy as np
from aiohttp import ClientSession, web
from aiortc import MediaStreamTrack, RTCPeerConnection, RTCSessionDescription
from av import AudioFrame
from cryptography.hazmat.primitives.asymmetric.ed25519 import Ed25519PrivateKey
from cryptography.hazmat.primitives.serialization import Encoding, PublicFormat
from server import FRAME_SAMPLES, SAMPLE_RATE, create_app, decode_public_key_token, make_wav
class Microphone(MediaStreamTrack):
kind = "audio"
def __init__(self):
super().__init__()
self.count = 0
self.start = time.monotonic()
async def recv(self):
await asyncio.sleep(max(0, self.start + self.count * .02 - time.monotonic()))
# 0.6 s speech followed by silence to trigger server VAD.
if 4 <= self.count < 44:
index = np.arange(FRAME_SAMPLES) + self.count * FRAME_SAMPLES
samples = (8000 * np.sin(2 * np.pi * 440 * index / SAMPLE_RATE)).astype("<i2")
else:
samples = np.zeros(FRAME_SAMPLES, dtype="<i2")
samples = samples.reshape(1, FRAME_SAMPLES)
frame = AudioFrame.from_ndarray(samples, format="s16", layout="mono")
frame.sample_rate = SAMPLE_RATE
frame.pts = self.count * FRAME_SAMPLES
frame.time_base = Fraction(1, SAMPLE_RATE)
self.count += 1
return frame
async def serve(app):
runner = web.AppRunner(app)
await runner.setup()
site = web.TCPSite(runner, "127.0.0.1", 0)
await site.start()
port = site._server.sockets[0].getsockname()[1]
return runner, f"http://127.0.0.1:{port}"
class VoiceSmokeTest(unittest.IsolatedAsyncioTestCase):
async def test_public_key_session_token(self):
private = Ed25519PrivateKey.generate()
pem = private.public_key().public_bytes(Encoding.PEM, PublicFormat.SubjectPublicKeyInfo)
key_server = web.Application()
async def key_handler(_):
return web.Response(body=pem)
key_server.router.add_get("/key", key_handler)
runner, base = await serve(key_server)
payload = base64.urlsafe_b64encode(json.dumps({
"exp": int(time.time()) + 60, "jti": str(uuid.uuid4()),
}).encode()).rstrip(b"=").decode()
signature = base64.urlsafe_b64encode(private.sign(payload.encode())).rstrip(b"=").decode()
try:
async with ClientSession() as http:
claims = await decode_public_key_token(f"{payload}.{signature}", http, base + "/key")
self.assertIn("jti", claims)
with self.assertRaises(web.HTTPUnauthorized):
await decode_public_key_token(f"{payload}.{signature}x", http, base + "/key")
finally:
await runner.cleanup()
async def test_browser_session_and_audio_round_trip(self):
async def stt(request):
data = await request.post()
self.assertIn("file", data)
return web.json_response({"text": "Hallo Athena"})
async def tts(request):
body = await request.json()
self.assertEqual(body["input"], "Hallo zurück")
return web.Response(body=make_wav((6000).to_bytes(2, "little", signed=True) * SAMPLE_RATE),
content_type="audio/wav")
fake = web.Application()
fake.router.add_post("/audio/transcriptions", stt)
fake.router.add_post("/audio/speech", tts)
fake_runner, fake_url = await serve(fake)
os.environ.update({
"ATHENA_TALK_REALTIME_SECRET": "test-secret-" * 4,
"ATHENA_API_BASE_URL": fake_url,
"OPENCLAW_ORIGIN": "https://oc.example.test",
"SILENCE_MS": "200",
})
voice_runner, voice_url = await serve(create_app())
pc = RTCPeerConnection()
events = []
done = asyncio.Event()
audio_received = asyncio.Event()
channel = pc.createDataChannel("oai-events")
pc.addTrack(Microphone())
@pc.on("track")
def on_track(track):
if track.kind == "audio":
async def consume():
for _ in range(150):
frame = await track.recv()
if np.max(np.abs(frame.to_ndarray())) > 100:
audio_received.set()
return
asyncio.create_task(consume())
@channel.on("open")
def on_open():
channel.send(json.dumps({"type": "session.update", "session": {}}))
@channel.on("message")
def on_message(raw):
event = json.loads(raw)
events.append(event)
if event.get("type") == "response.done":
output = event.get("response", {}).get("output", [])
if output and output[0].get("type") == "function_call":
call = output[0]
channel.send(json.dumps({"type": "conversation.item.create", "item": {
"type": "function_call_output", "call_id": call["call_id"],
"output": json.dumps({"result": "Hallo zurück"}),
}}))
channel.send(json.dumps({"type": "response.create"}))
elif output and output[0].get("type") == "message":
done.set()
try:
await pc.setLocalDescription(await pc.createOffer())
payload = base64.urlsafe_b64encode(json.dumps({
"exp": int(time.time()) + 60, "jti": str(uuid.uuid4()),
}).encode()).rstrip(b"=").decode()
signature = base64.urlsafe_b64encode(hmac.new(
os.environ["ATHENA_TALK_REALTIME_SECRET"].encode(),
payload.encode(), hashlib.sha256,
).digest()).rstrip(b"=").decode()
async with ClientSession() as http:
async with http.post(voice_url + "/v1/realtime/calls",
data=pc.localDescription.sdp,
headers={"Authorization": f"Bearer {payload}.{signature}",
"Content-Type": "application/sdp"}) as response:
self.assertEqual(response.status, 200, await response.text())
answer = await response.text()
await pc.setRemoteDescription(RTCSessionDescription(sdp=answer, type="answer"))
try:
await asyncio.wait_for(done.wait(), timeout=20)
except asyncio.TimeoutError:
self.fail(f"No completed voice response; events: {[e.get('type') for e in events]}")
await asyncio.wait_for(audio_received.wait(), timeout=5)
self.assertTrue(any(event.get("type") ==
"conversation.item.input_audio_transcription.completed"
for event in events))
self.assertTrue(any(event.get("type") ==
"response.output_audio_transcript.done" for event in events))
finally:
await pc.close()
await voice_runner.cleanup()
await fake_runner.cleanup()
if __name__ == "__main__":
unittest.main()