Files
AI-Profile-Router/services/athena-realtime-voice/test_smoke.py
T

244 lines
11 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, RealtimeSession, 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()))
# Two utterances separated by enough time for the first TTS reply.
if 4 <= self.count < 44 or 750 <= self.count < 790:
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_transcript_item_announced_before_completion(self):
async def stt(_):
return web.json_response({"text": "Hallo Athena"})
fake = web.Application()
fake.router.add_post("/audio/transcriptions", stt)
runner, base = await serve(fake)
os.environ["ATHENA_API_BASE_URL"] = base
events = []
class Channel:
readyState = "open"
def send(self, message):
events.append(json.loads(message))
try:
async with ClientSession() as http:
session = RealtimeSession(None, http, None)
session.channel = Channel()
await session.transcribe_and_consult(b"\0\0" * SAMPLE_RATE)
committed = next((index, event) for index, event in enumerate(events)
if event.get("type") == "input_audio_buffer.committed")
completed = next((index, event) for index, event in enumerate(events)
if event.get("type") == "conversation.item.input_audio_transcription.completed")
self.assertLess(committed[0], completed[0])
self.assertEqual(committed[1]["item_id"], completed[1]["item_id"])
finally:
await runner.cleanup()
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")
self.assertEqual(body["chunk_size"], 4)
stream = web.StreamResponse(headers={"Content-Type": "application/octet-stream"})
await stream.prepare(request)
pcm = (6000).to_bytes(2, "little", signed=True) * SAMPLE_RATE
await stream.write(pcm[:SAMPLE_RATE])
await asyncio.wait_for(audio_received.wait(), timeout=3)
await stream.write(pcm[SAMPLE_RATE:])
await stream.write_eof()
return stream
fake = web.Application()
fake.router.add_post("/audio/transcriptions", stt)
fake.router.add_post("/audio/speech/pcm-stream", 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 = []
event_times = []
done = asyncio.Event()
audio_received = asyncio.Event()
completed_responses = 0
function_calls = 0
channel = pc.createDataChannel("oai-events")
microphone = Microphone()
pc.addTrack(microphone)
@pc.on("track")
def on_track(track):
if track.kind == "audio":
async def consume():
for _ in range(600):
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):
nonlocal completed_responses, function_calls
event = json.loads(raw)
events.append(event)
event_times.append((round(time.monotonic() - microphone.start, 1), event.get("type")))
if event.get("type") == "response.done":
output = event.get("response", {}).get("output", [])
if output and output[0].get("type") == "function_call":
function_calls += 1
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":
completed_responses += 1
if completed_responses == 2:
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=25)
except asyncio.TimeoutError:
self.fail(f"No completed voice response; state={pc.connectionState}, "
f"mic_frames={microphone.count}, events: {event_times}")
await asyncio.wait_for(audio_received.wait(), timeout=5)
self.assertEqual(function_calls, 2)
self.assertTrue(any(event.get("type") ==
"conversation.item.input_audio_transcription.completed"
for event in events))
user_items = [event for event in events
if event.get("type") == "input_audio_buffer.committed"]
assistant_items = [event for event in events
if event.get("type") == "conversation.item.created"
and event.get("item", {}).get("role") == "assistant"]
assistant_transcripts = [event for event in events
if event.get("type") == "response.output_audio_transcript.done"]
self.assertEqual(len(user_items), 2)
self.assertEqual(len(assistant_items), 2)
self.assertEqual(len(assistant_transcripts), 2)
self.assertEqual(assistant_items[0]["previous_item_id"], user_items[0]["item_id"])
self.assertEqual(assistant_transcripts[0]["item_id"], assistant_items[0]["item"]["id"])
self.assertEqual(user_items[1]["previous_item_id"], assistant_items[0]["item"]["id"])
self.assertEqual(assistant_items[1]["previous_item_id"], user_items[1]["item_id"])
committed = next((index, event) for index, event in enumerate(events)
if event.get("type") == "input_audio_buffer.committed")
completed = next((index, event) for index, event in enumerate(events)
if event.get("type") == "conversation.item.input_audio_transcription.completed")
self.assertLess(committed[0], completed[0])
self.assertEqual(committed[1]["item_id"], completed[1]["item_id"])
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()