Add Athena WebRTC bridge for OpenClaw Talk
This commit is contained in:
@@ -0,0 +1,11 @@
|
||||
FROM python:3.11-slim
|
||||
|
||||
ENV PYTHONDONTWRITEBYTECODE=1 PYTHONUNBUFFERED=1 PORT=8090
|
||||
WORKDIR /app
|
||||
COPY requirements.txt .
|
||||
RUN pip install --no-cache-dir -r requirements.txt && useradd --uid 10001 --create-home athenavoice
|
||||
COPY server.py .
|
||||
USER athenavoice
|
||||
EXPOSE 8090
|
||||
HEALTHCHECK --interval=30s --timeout=5s --start-period=15s CMD python -c "import urllib.request; urllib.request.urlopen('http://127.0.0.1:8090/health', timeout=3)" || exit 1
|
||||
CMD ["python", "server.py"]
|
||||
@@ -0,0 +1,74 @@
|
||||
# Athena realtime voice bridge (experimental)
|
||||
|
||||
This independent service lets OpenClaw's existing browser Talk UI use Athena
|
||||
Whisper, OpenClaw's agent, and Athena Qwen3-TTS through the browser's supported
|
||||
OpenAI-style WebRTC transport. It does not modify OpenClaw or switch an Athena
|
||||
profile. The existing `gateway-relay` path remains available.
|
||||
|
||||
Flow: OpenClaw's `athena-talk` plugin signs a single-use 60-second browser token;
|
||||
the browser posts an SDP offer to this service; microphone audio is sent over
|
||||
WebRTC; Athena transcribes it; the service asks OpenClaw to run
|
||||
`openclaw_agent_consult`; Athena speaks the returned text over the same WebRTC
|
||||
connection. This is a half-duplex prototype. Barge-in and remote-network TURN
|
||||
support are not implemented.
|
||||
|
||||
Required service environment:
|
||||
|
||||
| Name | Purpose |
|
||||
| --- | --- |
|
||||
| `OPENCLAW_PUBLIC_KEY_URL` | Existing OpenClaw HTTPS origin plus `/plugins/athena-talk/realtime/public-key`. The service fetches the plugin's public verification key. |
|
||||
| `ATHENA_API_BASE_URL` | Router API base, `http://router:8081/v1` in the Compose deployment. |
|
||||
| `ATHENA_API_KEY` | Router key, if required. |
|
||||
| `OPENCLAW_ORIGIN` | Exact HTTPS origin of the UI, e.g. `https://oc.casaderoll.de`. |
|
||||
| `PORT` | HTTP listen port; default 8090. The Compose service shares the existing WireGuard gateway network namespace. |
|
||||
|
||||
The plugin provides an HTTPS offer route on the existing OpenClaw origin and
|
||||
forwards SDP to this service. Configure
|
||||
`talk.realtime.providers.athena-talk.realtimeUpstreamUrl` with Athena's internal
|
||||
HTTP URL ending in `/v1/realtime/calls`. Set the service's
|
||||
`OPENCLAW_PUBLIC_KEY_URL` to the OpenClaw HTTPS public-key route above. No new
|
||||
secret is needed in OpenClaw: its plugin keeps a private signing key in memory,
|
||||
and the service receives only the public key. Select
|
||||
`talk.realtime.transport: "webrtc"` only after the service is reachable. Keep
|
||||
`gateway-relay` as a rollback option.
|
||||
|
||||
The WebRTC media connection needs a route from the client device to Athena's
|
||||
ICE candidate addresses. The service runs directly inside Athena's existing
|
||||
private WireGuard network namespace, so home/VPN clients can reach it at
|
||||
`192.168.1.212:8090`; the OpenClaw HTTPS proxy covers signaling only. Clients
|
||||
outside the private network need TURN support, which is not implemented.
|
||||
|
||||
The deployed service is `mike-ai-realtime-voice` in `compose.yaml`. It uses no
|
||||
GPU and does not switch Athena's active model profile. Deploy it with the
|
||||
stack's normal environment:
|
||||
|
||||
```sh
|
||||
cd /opt/mike-ai/stack
|
||||
docker compose --env-file /etc/mike-ai/stack.env up -d --no-deps --build realtime-voice
|
||||
```
|
||||
|
||||
OpenClaw 2026.9.4 uses the `athena-talk` plugin version 1.1.0 with
|
||||
`talk.realtime.transport` set to `webrtc` and
|
||||
`talk.realtime.providers.athena-talk.realtimeUpstreamUrl` set to
|
||||
`http://192.168.1.212:8090/v1/realtime/calls`. On the Mac, turn off
|
||||
“Echtzeitweiterleitung über Gateway verwenden”, which applies only to the
|
||||
older `gateway-relay` mode. A Gateway restart after plugin installation was
|
||||
needed to register the two new HTTPS routes. The plugin signs short-lived
|
||||
session tokens with an in-memory Ed25519 key; the service fetches only its
|
||||
public key. A Gateway restart rotates the key automatically.
|
||||
|
||||
Local isolated test (fake STT/TTS, synthetic microphone and speaker audio):
|
||||
|
||||
```sh
|
||||
python3.11 -m venv .venv
|
||||
.venv/bin/pip install -r requirements.txt
|
||||
.venv/bin/python -m unittest -v test_smoke.py
|
||||
```
|
||||
|
||||
The synthetic microphone test passed over the actual OpenClaw HTTPS offer
|
||||
route and WireGuard media path with production Whisper and Qwen3-TTS on
|
||||
2026-09-16. It supplied a fake agent result; it does not prove that the Mac
|
||||
or browser UI completes a real agent-consult turn. Keep this integration
|
||||
experimental until that UI round trip has been observed.
|
||||
For isolated tests, `ATHENA_TALK_REALTIME_SECRET` can replace
|
||||
`OPENCLAW_PUBLIC_KEY_URL`; do not use that test mode in the deployed setup.
|
||||
@@ -0,0 +1,126 @@
|
||||
"""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]))
|
||||
@@ -0,0 +1,5 @@
|
||||
aiohttp==3.14.3
|
||||
aiortc==1.15.0
|
||||
av==17.1.0
|
||||
cryptography==50.0.1
|
||||
numpy==2.4.6
|
||||
@@ -0,0 +1,441 @@
|
||||
"""OpenAI Realtime WebRTC subset backed by Athena STT/TTS and OpenClaw tools.
|
||||
|
||||
The browser's OpenAI-compatible data channel delegates reasoning to OpenClaw's
|
||||
``openclaw_agent_consult`` tool. This server never calls the LLM directly.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import base64
|
||||
import hashlib
|
||||
import hmac
|
||||
import io
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
import uuid
|
||||
import wave
|
||||
from fractions import Fraction
|
||||
|
||||
import numpy as np
|
||||
from aiohttp import ClientSession, FormData, web
|
||||
from aiortc import MediaStreamTrack, RTCPeerConnection, RTCSessionDescription
|
||||
from av import AudioFrame, AudioResampler
|
||||
from cryptography.exceptions import InvalidSignature
|
||||
from cryptography.hazmat.primitives.serialization import load_pem_public_key
|
||||
|
||||
SAMPLE_RATE = 24000
|
||||
FRAME_SAMPLES = 480
|
||||
MAX_OFFER_BYTES = 64 * 1024
|
||||
SECRET = web.AppKey("secret", str)
|
||||
ORIGIN = web.AppKey("origin", str)
|
||||
USED_TOKENS = web.AppKey("used_tokens", dict)
|
||||
PEERS = web.AppKey("peers", set)
|
||||
HTTP = web.AppKey("http", ClientSession)
|
||||
PUBLIC_KEY_URL = web.AppKey("public_key_url", str)
|
||||
|
||||
|
||||
def decode_payload(payload: str, timestamp: int) -> dict:
|
||||
try:
|
||||
body = json.loads(base64.urlsafe_b64decode(payload + "=" * (-len(payload) % 4)))
|
||||
if not isinstance(body, dict) or not isinstance(body.get("jti"), str):
|
||||
raise ValueError("invalid payload")
|
||||
if not timestamp < int(body["exp"]) <= timestamp + 65:
|
||||
raise ValueError("expired token")
|
||||
return body
|
||||
except (TypeError, KeyError, ValueError, UnicodeError, json.JSONDecodeError) as exc:
|
||||
raise web.HTTPUnauthorized(text="Invalid realtime session token") from exc
|
||||
|
||||
|
||||
def decode_token(token: str, secret: str, now: int | None = None) -> dict:
|
||||
"""Validate the short-lived, signed browser session without exposing a router key."""
|
||||
try:
|
||||
payload, signature = token.split(".", 1)
|
||||
expected = base64.urlsafe_b64encode(
|
||||
hmac.new(secret.encode(), payload.encode(), hashlib.sha256).digest()
|
||||
).rstrip(b"=").decode()
|
||||
if not hmac.compare_digest(signature, expected):
|
||||
raise ValueError("invalid signature")
|
||||
timestamp = int(time.time()) if now is None else now
|
||||
return decode_payload(payload, timestamp)
|
||||
except (TypeError, KeyError, ValueError, UnicodeError, json.JSONDecodeError) as exc:
|
||||
raise web.HTTPUnauthorized(text="Invalid realtime session token") from exc
|
||||
|
||||
|
||||
async def decode_public_key_token(token: str, http: ClientSession, url: str) -> dict:
|
||||
try:
|
||||
payload, signature = token.split(".", 1)
|
||||
if "." in signature:
|
||||
raise ValueError("invalid token")
|
||||
async with http.get(url, timeout=5) as response:
|
||||
response.raise_for_status()
|
||||
key_pem = await response.read()
|
||||
if len(key_pem) > 4096:
|
||||
raise ValueError("public key too large")
|
||||
public_key = load_pem_public_key(key_pem)
|
||||
public_key.verify(
|
||||
base64.urlsafe_b64decode(signature + "=" * (-len(signature) % 4)),
|
||||
payload.encode(),
|
||||
)
|
||||
return decode_payload(payload, int(time.time()))
|
||||
except (ValueError, InvalidSignature, TypeError) as exc:
|
||||
raise web.HTTPUnauthorized(text="Invalid realtime session token") from exc
|
||||
except web.HTTPException:
|
||||
raise
|
||||
except Exception as exc:
|
||||
raise web.HTTPServiceUnavailable(text="OpenClaw public key unavailable") from exc
|
||||
|
||||
|
||||
def pcm_rms(data: bytes) -> float:
|
||||
samples = np.frombuffer(data[: len(data) // 2 * 2], dtype="<i2")
|
||||
if not samples.size:
|
||||
return 0.0
|
||||
normalized = samples.astype(np.float32) / 32768.0
|
||||
return float(np.sqrt(np.mean(normalized * normalized)))
|
||||
|
||||
|
||||
def make_wav(data: bytes) -> bytes:
|
||||
output = io.BytesIO()
|
||||
with wave.open(output, "wb") as wav:
|
||||
wav.setnchannels(1)
|
||||
wav.setsampwidth(2)
|
||||
wav.setframerate(SAMPLE_RATE)
|
||||
wav.writeframes(data)
|
||||
return output.getvalue()
|
||||
|
||||
|
||||
def read_wav(data: bytes) -> bytes:
|
||||
with wave.open(io.BytesIO(data), "rb") as wav:
|
||||
if wav.getsampwidth() != 2:
|
||||
raise ValueError("Athena TTS must return 16-bit WAV")
|
||||
channels, rate = wav.getnchannels(), wav.getframerate()
|
||||
samples = np.frombuffer(wav.readframes(wav.getnframes()), dtype="<i2")
|
||||
if channels > 1:
|
||||
samples = samples.reshape(-1, channels)[:, 0]
|
||||
if rate != SAMPLE_RATE:
|
||||
count = max(1, round(samples.size * SAMPLE_RATE / rate))
|
||||
samples = np.interp(
|
||||
np.arange(count) * rate / SAMPLE_RATE,
|
||||
np.arange(samples.size),
|
||||
samples,
|
||||
).astype("<i2")
|
||||
return samples.astype("<i2").tobytes()
|
||||
|
||||
|
||||
class SpeechTrack(MediaStreamTrack):
|
||||
kind = "audio"
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self._queue: asyncio.Queue[bytes] = asyncio.Queue(maxsize=2000)
|
||||
self._pending = bytearray()
|
||||
self._pts = 0
|
||||
self._start = time.monotonic()
|
||||
|
||||
def enqueue(self, pcm: bytes) -> None:
|
||||
for offset in range(0, len(pcm), FRAME_SAMPLES * 2):
|
||||
try:
|
||||
self._queue.put_nowait(pcm[offset : offset + FRAME_SAMPLES * 2])
|
||||
except asyncio.QueueFull:
|
||||
break
|
||||
|
||||
async def recv(self) -> AudioFrame:
|
||||
await asyncio.sleep(max(0, self._start + self._pts / SAMPLE_RATE - time.monotonic()))
|
||||
if not self._pending:
|
||||
try:
|
||||
self._pending.extend(self._queue.get_nowait())
|
||||
except asyncio.QueueEmpty:
|
||||
pass
|
||||
chunk = bytes(self._pending[: FRAME_SAMPLES * 2])
|
||||
del self._pending[: FRAME_SAMPLES * 2]
|
||||
chunk = chunk.ljust(FRAME_SAMPLES * 2, b"\x00")
|
||||
frame = AudioFrame.from_ndarray(
|
||||
np.frombuffer(chunk, dtype="<i2").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
|
||||
|
||||
|
||||
class RealtimeSession:
|
||||
def __init__(self, pc: RTCPeerConnection, http: ClientSession, track: SpeechTrack):
|
||||
self.pc, self.http, self.track = pc, http, track
|
||||
self.channel = None
|
||||
self.busy = False
|
||||
self.speaking = False
|
||||
self.speech = bytearray()
|
||||
self.silence_since: float | None = None
|
||||
self.call_id: str | None = None
|
||||
self.pending_reply: str | None = None
|
||||
self.turn_task: asyncio.Task | None = None
|
||||
self.vad_threshold = float(os.environ.get("VAD_THRESHOLD", "0.018"))
|
||||
self.silence_ms = int(os.environ.get("SILENCE_MS", "750"))
|
||||
|
||||
def emit(self, event: dict) -> None:
|
||||
if self.channel and self.channel.readyState == "open":
|
||||
self.channel.send(json.dumps(event, ensure_ascii=False))
|
||||
|
||||
def on_message(self, message: str | bytes) -> None:
|
||||
if not isinstance(message, str) or len(message) > 262144:
|
||||
return
|
||||
try:
|
||||
event = json.loads(message)
|
||||
except json.JSONDecodeError:
|
||||
return
|
||||
kind = event.get("type")
|
||||
if kind == "session.update":
|
||||
self.emit({"type": "session.updated", "session": {"id": "athena-local"}})
|
||||
elif kind == "conversation.item.create":
|
||||
item = event.get("item") or {}
|
||||
if item.get("type") == "function_call_output" and item.get("call_id") == self.call_id:
|
||||
try:
|
||||
answer = json.loads(item.get("output") or "{}")
|
||||
self.pending_reply = str(answer.get("result") or answer.get("error") or "")
|
||||
except (ValueError, TypeError):
|
||||
self.pending_reply = "Die Antwort konnte nicht gelesen werden."
|
||||
elif item.get("type") == "message" and item.get("role") == "user":
|
||||
text = " ".join(str(part.get("text") or "") for part in item.get("content", []))
|
||||
if text.strip() and not self.busy:
|
||||
self.turn_task = asyncio.create_task(self.consult(text.strip()))
|
||||
elif kind == "response.create" and self.pending_reply is not None:
|
||||
text, self.pending_reply = self.pending_reply, None
|
||||
self.turn_task = asyncio.create_task(self.speak(text))
|
||||
elif kind == "response.cancel":
|
||||
if self.turn_task:
|
||||
self.turn_task.cancel()
|
||||
|
||||
async def consume_audio(self, remote: MediaStreamTrack) -> None:
|
||||
resampler = AudioResampler(format="s16", layout="mono", rate=SAMPLE_RATE)
|
||||
try:
|
||||
while True:
|
||||
frame = await remote.recv()
|
||||
for converted in resampler.resample(frame):
|
||||
pcm = bytes(converted.planes[0])[: converted.samples * 2]
|
||||
await self.on_audio(pcm)
|
||||
except asyncio.CancelledError:
|
||||
raise
|
||||
except Exception:
|
||||
return
|
||||
|
||||
async def on_audio(self, pcm: bytes) -> None:
|
||||
if self.busy or not pcm:
|
||||
return
|
||||
voiced = pcm_rms(pcm) >= self.vad_threshold
|
||||
now = time.monotonic()
|
||||
if voiced and not self.speaking:
|
||||
self.speaking = True
|
||||
self.speech.clear()
|
||||
self.silence_since = None
|
||||
self.emit({"type": "input_audio_buffer.speech_started"})
|
||||
if not self.speaking:
|
||||
return
|
||||
self.speech.extend(pcm)
|
||||
if voiced:
|
||||
self.silence_since = None
|
||||
elif self.silence_since is None:
|
||||
self.silence_since = now
|
||||
elapsed_silence = (now - self.silence_since) * 1000 if self.silence_since else 0
|
||||
if elapsed_silence >= self.silence_ms or len(self.speech) >= SAMPLE_RATE * 2 * 45:
|
||||
self.speaking = False
|
||||
self.emit({"type": "input_audio_buffer.speech_stopped"})
|
||||
audio = bytes(self.speech)
|
||||
self.speech.clear()
|
||||
if len(audio) >= SAMPLE_RATE * 2 // 4:
|
||||
self.busy = True
|
||||
self.turn_task = asyncio.create_task(self.transcribe_and_consult(audio))
|
||||
|
||||
async def transcribe_and_consult(self, pcm: bytes) -> None:
|
||||
try:
|
||||
form = FormData()
|
||||
form.add_field("file", make_wav(pcm), filename="talk.wav", content_type="audio/wav")
|
||||
form.add_field("model", "whisper-1")
|
||||
form.add_field("language", os.environ.get("STT_LANGUAGE", "de"))
|
||||
async with self.http.post(
|
||||
os.environ["ATHENA_API_BASE_URL"].rstrip("/") + "/audio/transcriptions",
|
||||
headers=self.auth_headers(), data=form, timeout=90,
|
||||
) as response:
|
||||
response.raise_for_status()
|
||||
text = str((await response.json()).get("text") or "").strip()
|
||||
if not text:
|
||||
self.busy = False
|
||||
return
|
||||
self.emit({"type": "conversation.item.input_audio_transcription.completed",
|
||||
"item_id": f"item_{uuid.uuid4().hex}", "transcript": text})
|
||||
await self.consult(text)
|
||||
except asyncio.CancelledError:
|
||||
self.busy = False
|
||||
raise
|
||||
except Exception:
|
||||
self.busy = False
|
||||
self.emit({"type": "error", "error": {"message": "Athena STT failed"}})
|
||||
|
||||
async def consult(self, text: str) -> None:
|
||||
self.busy = True
|
||||
self.call_id = f"call_{uuid.uuid4().hex}"
|
||||
response_id = f"resp_{uuid.uuid4().hex}"
|
||||
self.emit({"type": "response.created", "response": {"id": response_id}})
|
||||
self.emit({"type": "response.done", "response": {
|
||||
"id": response_id, "status": "completed", "output": [{
|
||||
"id": f"item_{uuid.uuid4().hex}", "type": "function_call",
|
||||
"status": "completed", "call_id": self.call_id,
|
||||
"name": "openclaw_agent_consult",
|
||||
"arguments": json.dumps({"prompt": text}, ensure_ascii=False),
|
||||
}],
|
||||
}})
|
||||
|
||||
async def speak(self, text: str) -> None:
|
||||
if not text.strip():
|
||||
self.busy = False
|
||||
return
|
||||
response_id = f"resp_{uuid.uuid4().hex}"
|
||||
try:
|
||||
self.emit({"type": "response.created", "response": {"id": response_id}})
|
||||
async with self.http.post(
|
||||
os.environ["ATHENA_API_BASE_URL"].rstrip("/") + "/audio/speech",
|
||||
headers={"Content-Type": "application/json", **self.auth_headers()},
|
||||
json={"input": text, "voice": os.environ.get("TTS_VOICE", "alloy"),
|
||||
"response_format": "wav"}, timeout=120,
|
||||
) as response:
|
||||
response.raise_for_status()
|
||||
pcm = read_wav(await response.read())
|
||||
self.track.enqueue(pcm)
|
||||
self.emit({"type": "response.output_audio_transcript.done", "transcript": text})
|
||||
# Keep half-duplex active until the queued audio has actually played.
|
||||
await asyncio.sleep(len(pcm) / (SAMPLE_RATE * 2) + 0.2)
|
||||
self.emit({"type": "response.done", "response": {
|
||||
"id": response_id, "status": "completed", "output": [
|
||||
{"id": f"item_{uuid.uuid4().hex}", "type": "message", "role": "assistant", "status": "completed"}
|
||||
],
|
||||
}})
|
||||
except asyncio.CancelledError:
|
||||
self.emit({"type": "response.cancelled", "response": {"id": response_id}})
|
||||
raise
|
||||
except Exception:
|
||||
self.emit({"type": "error", "error": {"message": "Athena TTS failed"}})
|
||||
finally:
|
||||
self.busy = False
|
||||
|
||||
@staticmethod
|
||||
def auth_headers() -> dict[str, str]:
|
||||
key = os.environ.get("ATHENA_API_KEY", "")
|
||||
return {"Authorization": f"Bearer {key}"} if key else {}
|
||||
|
||||
|
||||
async def offer(request: web.Request) -> web.Response:
|
||||
if request.headers.get("Origin") not in (None, request.app[ORIGIN]):
|
||||
raise web.HTTPForbidden(text="Origin not allowed")
|
||||
token = request.headers.get("Authorization", "").removeprefix("Bearer ")
|
||||
if request.app[PUBLIC_KEY_URL]:
|
||||
claims = await decode_public_key_token(token, request.app[HTTP], request.app[PUBLIC_KEY_URL])
|
||||
else:
|
||||
claims = decode_token(token, request.app[SECRET])
|
||||
used = request.app[USED_TOKENS]
|
||||
now = time.time()
|
||||
for key, expiry in list(used.items()):
|
||||
if expiry < now:
|
||||
del used[key]
|
||||
if claims["jti"] in used:
|
||||
raise web.HTTPUnauthorized(text="Realtime token already used")
|
||||
if len(request.app[PEERS]) >= int(os.environ.get("MAX_SESSIONS", "4")):
|
||||
raise web.HTTPServiceUnavailable(text="Too many realtime sessions")
|
||||
used[claims["jti"]] = claims["exp"]
|
||||
if request.content_length is not None and request.content_length > MAX_OFFER_BYTES:
|
||||
raise web.HTTPRequestEntityTooLarge(max_size=MAX_OFFER_BYTES, actual_size=request.content_length)
|
||||
sdp = await request.text()
|
||||
if len(sdp.encode()) > MAX_OFFER_BYTES:
|
||||
raise web.HTTPRequestEntityTooLarge(max_size=MAX_OFFER_BYTES, actual_size=len(sdp))
|
||||
pc = RTCPeerConnection()
|
||||
speech = SpeechTrack()
|
||||
pc.addTrack(speech)
|
||||
session = RealtimeSession(pc, request.app[HTTP], speech)
|
||||
request.app[PEERS].add(pc)
|
||||
|
||||
@pc.on("datachannel")
|
||||
def on_channel(channel):
|
||||
session.channel = channel
|
||||
|
||||
@channel.on("open")
|
||||
def on_open():
|
||||
session.emit({"type": "session.created", "session": {"id": "athena-local"}})
|
||||
|
||||
@channel.on("message")
|
||||
def on_message(message):
|
||||
session.on_message(message)
|
||||
|
||||
@pc.on("track")
|
||||
def on_track(track):
|
||||
if track.kind == "audio":
|
||||
asyncio.create_task(session.consume_audio(track))
|
||||
|
||||
@pc.on("connectionstatechange")
|
||||
async def on_state_change():
|
||||
if pc.connectionState in {"closed", "failed"}:
|
||||
if session.turn_task:
|
||||
session.turn_task.cancel()
|
||||
await pc.close()
|
||||
request.app[PEERS].discard(pc)
|
||||
|
||||
try:
|
||||
await pc.setRemoteDescription(RTCSessionDescription(sdp=sdp, type="offer"))
|
||||
await pc.setLocalDescription(await pc.createAnswer())
|
||||
return web.Response(text=pc.localDescription.sdp, content_type="application/sdp",
|
||||
headers={"Access-Control-Allow-Origin": request.app[ORIGIN]})
|
||||
except Exception:
|
||||
await pc.close()
|
||||
request.app[PEERS].discard(pc)
|
||||
raise web.HTTPBadRequest(text="Invalid WebRTC offer")
|
||||
|
||||
|
||||
async def preflight(request: web.Request) -> web.Response:
|
||||
if request.headers.get("Origin") != request.app[ORIGIN]:
|
||||
raise web.HTTPForbidden(text="Origin not allowed")
|
||||
return web.Response(headers={
|
||||
"Access-Control-Allow-Origin": request.app[ORIGIN],
|
||||
"Access-Control-Allow-Methods": "POST, OPTIONS",
|
||||
"Access-Control-Allow-Headers": "Authorization, Content-Type",
|
||||
"Access-Control-Max-Age": "600",
|
||||
})
|
||||
|
||||
|
||||
async def lifecycle(app: web.Application):
|
||||
app[HTTP] = ClientSession()
|
||||
yield
|
||||
await asyncio.gather(*(pc.close() for pc in app[PEERS]))
|
||||
await app[HTTP].close()
|
||||
|
||||
|
||||
async def health(_: web.Request) -> web.Response:
|
||||
return web.json_response({"ok": True})
|
||||
|
||||
|
||||
def create_app() -> web.Application:
|
||||
secret = os.environ.get("ATHENA_TALK_REALTIME_SECRET", "")
|
||||
public_key_url = os.environ.get("OPENCLAW_PUBLIC_KEY_URL", "")
|
||||
if public_key_url and not public_key_url.startswith("https://"):
|
||||
raise RuntimeError("OPENCLAW_PUBLIC_KEY_URL must use HTTPS")
|
||||
if not public_key_url and len(secret) < 32:
|
||||
raise RuntimeError("Set OPENCLAW_PUBLIC_KEY_URL or a 32+ character secret")
|
||||
if not os.environ.get("ATHENA_API_BASE_URL"):
|
||||
raise RuntimeError("ATHENA_API_BASE_URL is required")
|
||||
origin = os.environ.get("OPENCLAW_ORIGIN", "")
|
||||
if not origin.startswith("https://"):
|
||||
raise RuntimeError("OPENCLAW_ORIGIN must be an HTTPS origin")
|
||||
app = web.Application(client_max_size=MAX_OFFER_BYTES)
|
||||
app[SECRET] = secret
|
||||
app[PUBLIC_KEY_URL] = public_key_url
|
||||
app[ORIGIN] = origin
|
||||
app[USED_TOKENS] = {}
|
||||
app[PEERS] = set()
|
||||
app.cleanup_ctx.append(lifecycle)
|
||||
app.router.add_route("OPTIONS", "/v1/realtime/calls", preflight)
|
||||
app.router.add_post("/v1/realtime/calls", offer)
|
||||
app.router.add_get("/health", health)
|
||||
return app
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
web.run_app(create_app(), host="0.0.0.0", port=int(os.environ.get("PORT", "8090")))
|
||||
@@ -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