Add Athena WebRTC bridge for OpenClaw Talk
This commit is contained in:
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()
|
||||
Reference in new issue
Block a user