Files

127 lines
4.9 KiB
Python

"""Manual Mac→Athena WebRTC check using a generated test WAV, no microphone."""
import asyncio
import base64
import hashlib
import hmac
import io
import json
import os
import sys
import time
import uuid
import wave
from fractions import Fraction
import numpy as np
from aiohttp import ClientSession
from aiortc import MediaStreamTrack, RTCPeerConnection, RTCSessionDescription
from av import AudioFrame
from server import FRAME_SAMPLES, SAMPLE_RATE
class RecordedMicrophone(MediaStreamTrack):
kind = "audio"
def __init__(self, path):
super().__init__()
with wave.open(path, "rb") as wav:
assert wav.getframerate() == SAMPLE_RATE and wav.getnchannels() == 1
assert wav.getsampwidth() == 2
self.audio = wav.readframes(wav.getnframes())
self.offset = 0
self.pts = 0
self.start = time.monotonic()
async def recv(self):
await asyncio.sleep(max(0, self.start + self.pts / SAMPLE_RATE - time.monotonic()))
chunk = self.audio[self.offset : self.offset + FRAME_SAMPLES * 2]
self.offset += len(chunk)
samples = np.frombuffer(chunk.ljust(FRAME_SAMPLES * 2, b"\0"), dtype="<i2")
frame = AudioFrame.from_ndarray(samples.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
async def main(url, wav_path):
browser_token = os.environ.get("BROWSER_SESSION_TOKEN")
if not browser_token:
secret = os.environ["ATHENA_TALK_REALTIME_SECRET"]
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(
secret.encode(), payload.encode(), hashlib.sha256,
).digest()).rstrip(b"=").decode()
browser_token = f"{payload}.{signature}"
pc = RTCPeerConnection()
pc.addTrack(RecordedMicrophone(wav_path))
channel = pc.createDataChannel("oai-events")
transcribed = asyncio.Event()
response_done = asyncio.Event()
audio_received = asyncio.Event()
transcript = ""
seen = []
@pc.on("track")
def on_track(track):
if track.kind == "audio":
async def consume():
for _ in range(9000):
frame = await track.recv()
if np.max(np.abs(frame.to_ndarray())) > 100:
audio_received.set()
return
asyncio.create_task(consume())
@channel.on("message")
def on_message(raw):
nonlocal transcript
event = json.loads(raw)
seen.append(event.get("type"))
if event.get("type") == "conversation.item.input_audio_transcription.completed":
transcript = event.get("transcript", "")
transcribed.set()
elif event.get("type") == "response.done":
output = event.get("response", {}).get("output", [])
if output and output[0].get("type") == "function_call":
channel.send(json.dumps({"type": "conversation.item.create", "item": {
"type": "function_call_output", "call_id": output[0]["call_id"],
"output": json.dumps({"result": "Hallo zurück!"}),
}}))
channel.send(json.dumps({"type": "response.create"}))
elif output and output[0].get("type") == "message":
response_done.set()
elif event.get("type") == "error":
print("bridge error:", event.get("error", {}).get("message"))
try:
await pc.setLocalDescription(await pc.createOffer())
async with ClientSession() as http:
async with http.post(url, data=pc.localDescription.sdp,
headers={"Authorization": f"Bearer {browser_token}",
"Content-Type": "application/sdp"}) as result:
if result.status != 200:
raise RuntimeError(f"SDP offer failed: HTTP {result.status}")
answer = await result.text()
await pc.setRemoteDescription(RTCSessionDescription(sdp=answer, type="answer"))
try:
await asyncio.wait_for(transcribed.wait(), timeout=90)
await asyncio.wait_for(audio_received.wait(), timeout=90)
await asyncio.wait_for(response_done.wait(), timeout=90)
except asyncio.TimeoutError as exc:
raise RuntimeError(f"Remote voice timeout: state={pc.connectionState}, "
f"transcript={transcript!r}, events={seen}") from exc
print(json.dumps({"webRTC": pc.connectionState, "transcript": transcript,
"spokenResponse": True}, ensure_ascii=False))
finally:
await pc.close()
if __name__ == "__main__":
asyncio.run(main(sys.argv[1], sys.argv[2]))