176 lines
7.3 KiB
Python
176 lines
7.3 KiB
Python
"""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()
|