Keep realtime conversation items ordered across turns

This commit is contained in:
Mikei386
2026-09-16 14:42:30 +02:00
parent b9dcadedd5
commit 97c8072b3d
3 changed files with 82 additions and 13 deletions
+4
View File
@@ -16,6 +16,10 @@ The bridge announces each committed user audio item before sending its completed
transcript with the same item ID. OpenClaw requires that sequence to persist the transcript with the same item ID. OpenClaw requires that sequence to persist the
transcript; omitting the item caused “Realtime transcript refers to an unknown transcript; omitting the item caused “Realtime transcript refers to an unknown
speech item” in the browser. speech item” in the browser.
Spoken assistant replies are also announced as conversation items with matching
transcript IDs and explicit predecessor IDs. Without them, OpenClaw's ordered
transcript storage can wait on a missing assistant item when the next user turn
arrives, leaving later questions unanswered.
Required service environment: Required service environment:
+46 -6
View File
@@ -12,6 +12,7 @@ import hashlib
import hmac import hmac
import io import io
import json import json
import logging
import os import os
import time import time
import uuid import uuid
@@ -21,6 +22,7 @@ from fractions import Fraction
import numpy as np import numpy as np
from aiohttp import ClientSession, FormData, web from aiohttp import ClientSession, FormData, web
from aiortc import MediaStreamTrack, RTCPeerConnection, RTCSessionDescription from aiortc import MediaStreamTrack, RTCPeerConnection, RTCSessionDescription
from aiortc.mediastreams import MediaStreamError
from av import AudioFrame, AudioResampler from av import AudioFrame, AudioResampler
from cryptography.exceptions import InvalidSignature from cryptography.exceptions import InvalidSignature
from cryptography.hazmat.primitives.serialization import load_pem_public_key from cryptography.hazmat.primitives.serialization import load_pem_public_key
@@ -34,6 +36,12 @@ USED_TOKENS = web.AppKey("used_tokens", dict)
PEERS = web.AppKey("peers", set) PEERS = web.AppKey("peers", set)
HTTP = web.AppKey("http", ClientSession) HTTP = web.AppKey("http", ClientSession)
PUBLIC_KEY_URL = web.AppKey("public_key_url", str) PUBLIC_KEY_URL = web.AppKey("public_key_url", str)
logger = logging.getLogger("athena_realtime_voice")
logger.setLevel(logging.INFO)
logger.propagate = False
_handler = logging.StreamHandler()
_handler.setFormatter(logging.Formatter("%(asctime)s %(levelname)s %(message)s"))
logger.addHandler(_handler)
def decode_payload(payload: str, timestamp: int) -> dict: def decode_payload(payload: str, timestamp: int) -> dict:
@@ -170,8 +178,10 @@ class RealtimeSession:
self.speech = bytearray() self.speech = bytearray()
self.silence_since: float | None = None self.silence_since: float | None = None
self.call_id: str | None = None self.call_id: str | None = None
self.last_item_id: str | None = None
self.pending_reply: str | None = None self.pending_reply: str | None = None
self.turn_task: asyncio.Task | None = None self.turn_task: asyncio.Task | None = None
self.dropped_speech_while_busy = False
self.vad_threshold = float(os.environ.get("VAD_THRESHOLD", "0.018")) self.vad_threshold = float(os.environ.get("VAD_THRESHOLD", "0.018"))
self.silence_ms = int(os.environ.get("SILENCE_MS", "750")) self.silence_ms = int(os.environ.get("SILENCE_MS", "750"))
@@ -187,6 +197,11 @@ class RealtimeSession:
except json.JSONDecodeError: except json.JSONDecodeError:
return return
kind = event.get("type") kind = event.get("type")
if kind in {"conversation.item.create", "response.create", "response.cancel"}:
logger.info("client event=%s item_type=%s call_matches=%s busy=%s pending_reply=%s",
kind, (event.get("item") or {}).get("type"),
(event.get("item") or {}).get("call_id") == self.call_id,
self.busy, self.pending_reply is not None)
if kind == "session.update": if kind == "session.update":
self.emit({"type": "session.updated", "session": {"id": "athena-local"}}) self.emit({"type": "session.updated", "session": {"id": "athena-local"}})
elif kind == "conversation.item.create": elif kind == "conversation.item.create":
@@ -218,12 +233,21 @@ class RealtimeSession:
await self.on_audio(pcm) await self.on_audio(pcm)
except asyncio.CancelledError: except asyncio.CancelledError:
raise raise
except MediaStreamError:
return
except Exception: except Exception:
logger.exception("incoming audio stream stopped unexpectedly")
return return
async def on_audio(self, pcm: bytes) -> None: async def on_audio(self, pcm: bytes) -> None:
if self.busy or not pcm: if not pcm:
return return
if self.busy:
if not self.dropped_speech_while_busy and pcm_rms(pcm) >= self.vad_threshold:
self.dropped_speech_while_busy = True
logger.info("microphone speech ignored while response is busy")
return
self.dropped_speech_while_busy = False
voiced = pcm_rms(pcm) >= self.vad_threshold voiced = pcm_rms(pcm) >= self.vad_threshold
now = time.monotonic() now = time.monotonic()
if voiced and not self.speaking: if voiced and not self.speaking:
@@ -231,6 +255,7 @@ class RealtimeSession:
self.speech.clear() self.speech.clear()
self.silence_since = None self.silence_since = None
self.emit({"type": "input_audio_buffer.speech_started"}) self.emit({"type": "input_audio_buffer.speech_started"})
logger.info("speech started")
if not self.speaking: if not self.speaking:
return return
self.speech.extend(pcm) self.speech.extend(pcm)
@@ -242,6 +267,7 @@ class RealtimeSession:
if elapsed_silence >= self.silence_ms or len(self.speech) >= SAMPLE_RATE * 2 * 45: if elapsed_silence >= self.silence_ms or len(self.speech) >= SAMPLE_RATE * 2 * 45:
self.speaking = False self.speaking = False
self.emit({"type": "input_audio_buffer.speech_stopped"}) self.emit({"type": "input_audio_buffer.speech_stopped"})
logger.info("speech stopped bytes=%s", len(self.speech))
audio = bytes(self.speech) audio = bytes(self.speech)
self.speech.clear() self.speech.clear()
if len(audio) >= SAMPLE_RATE * 2 // 4: if len(audio) >= SAMPLE_RATE * 2 // 4:
@@ -265,14 +291,18 @@ class RealtimeSession:
return return
item_id = f"item_{uuid.uuid4().hex}" item_id = f"item_{uuid.uuid4().hex}"
# OpenClaw tracks transcript items before accepting their text. # OpenClaw tracks transcript items before accepting their text.
self.emit({"type": "input_audio_buffer.committed", "item_id": item_id}) self.emit({"type": "input_audio_buffer.committed", "item_id": item_id,
"previous_item_id": self.last_item_id})
self.last_item_id = item_id
self.emit({"type": "conversation.item.input_audio_transcription.completed", self.emit({"type": "conversation.item.input_audio_transcription.completed",
"item_id": item_id, "transcript": text}) "item_id": item_id, "transcript": text})
logger.info("transcript completed item_id=%s", item_id)
await self.consult(text) await self.consult(text)
except asyncio.CancelledError: except asyncio.CancelledError:
self.busy = False self.busy = False
raise raise
except Exception: except Exception:
logger.exception("speech transcription failed")
self.busy = False self.busy = False
self.emit({"type": "error", "error": {"message": "Athena STT failed"}}) self.emit({"type": "error", "error": {"message": "Athena STT failed"}})
@@ -289,13 +319,16 @@ class RealtimeSession:
"arguments": json.dumps({"prompt": text}, ensure_ascii=False), "arguments": json.dumps({"prompt": text}, ensure_ascii=False),
}], }],
}}) }})
logger.info("consult requested call_id=%s response_id=%s", self.call_id, response_id)
async def speak(self, text: str) -> None: async def speak(self, text: str) -> None:
if not text.strip(): if not text.strip():
self.busy = False self.busy = False
return return
response_id = f"resp_{uuid.uuid4().hex}" response_id = f"resp_{uuid.uuid4().hex}"
item_id = f"item_{uuid.uuid4().hex}"
try: try:
logger.info("speaking response_id=%s", response_id)
self.emit({"type": "response.created", "response": {"id": response_id}}) self.emit({"type": "response.created", "response": {"id": response_id}})
async with self.http.post( async with self.http.post(
os.environ["ATHENA_API_BASE_URL"].rstrip("/") + "/audio/speech", os.environ["ATHENA_API_BASE_URL"].rstrip("/") + "/audio/speech",
@@ -306,18 +339,25 @@ class RealtimeSession:
response.raise_for_status() response.raise_for_status()
pcm = read_wav(await response.read()) pcm = read_wav(await response.read())
self.track.enqueue(pcm) self.track.enqueue(pcm)
self.emit({"type": "response.output_audio_transcript.done", "transcript": text}) item = {"id": item_id, "type": "message", "role": "assistant",
"status": "completed"}
self.emit({"type": "conversation.item.created", "item": item,
"previous_item_id": self.last_item_id})
self.last_item_id = item_id
self.emit({"type": "response.output_audio_transcript.done",
"item_id": item_id, "transcript": text})
# Keep half-duplex active until the queued audio has actually played. # Keep half-duplex active until the queued audio has actually played.
await asyncio.sleep(len(pcm) / (SAMPLE_RATE * 2) + 0.2) await asyncio.sleep(len(pcm) / (SAMPLE_RATE * 2) + 0.2)
self.emit({"type": "conversation.item.done", "item": item})
self.emit({"type": "response.done", "response": { self.emit({"type": "response.done", "response": {
"id": response_id, "status": "completed", "output": [ "id": response_id, "status": "completed", "output": [item],
{"id": f"item_{uuid.uuid4().hex}", "type": "message", "role": "assistant", "status": "completed"}
],
}}) }})
logger.info("speech response complete response_id=%s", response_id)
except asyncio.CancelledError: except asyncio.CancelledError:
self.emit({"type": "response.cancelled", "response": {"id": response_id}}) self.emit({"type": "response.cancelled", "response": {"id": response_id}})
raise raise
except Exception: except Exception:
logger.exception("speech synthesis failed")
self.emit({"type": "error", "error": {"message": "Athena TTS failed"}}) self.emit({"type": "error", "error": {"message": "Athena TTS failed"}})
finally: finally:
self.busy = False self.busy = False
+31 -6
View File
@@ -31,8 +31,8 @@ class Microphone(MediaStreamTrack):
async def recv(self): async def recv(self):
await asyncio.sleep(max(0, self.start + self.count * .02 - time.monotonic())) await asyncio.sleep(max(0, self.start + self.count * .02 - time.monotonic()))
# 0.6 s speech followed by silence to trigger server VAD. # Two utterances separated by enough time for the first TTS reply.
if 4 <= self.count < 44: if 4 <= self.count < 44 or 750 <= self.count < 790:
index = np.arange(FRAME_SAMPLES) + self.count * FRAME_SAMPLES index = np.arange(FRAME_SAMPLES) + self.count * FRAME_SAMPLES
samples = (8000 * np.sin(2 * np.pi * 440 * index / SAMPLE_RATE)).astype("<i2") samples = (8000 * np.sin(2 * np.pi * 440 * index / SAMPLE_RATE)).astype("<i2")
else: else:
@@ -132,16 +132,20 @@ class VoiceSmokeTest(unittest.IsolatedAsyncioTestCase):
voice_runner, voice_url = await serve(create_app()) voice_runner, voice_url = await serve(create_app())
pc = RTCPeerConnection() pc = RTCPeerConnection()
events = [] events = []
event_times = []
done = asyncio.Event() done = asyncio.Event()
audio_received = asyncio.Event() audio_received = asyncio.Event()
completed_responses = 0
function_calls = 0
channel = pc.createDataChannel("oai-events") channel = pc.createDataChannel("oai-events")
pc.addTrack(Microphone()) microphone = Microphone()
pc.addTrack(microphone)
@pc.on("track") @pc.on("track")
def on_track(track): def on_track(track):
if track.kind == "audio": if track.kind == "audio":
async def consume(): async def consume():
for _ in range(150): for _ in range(600):
frame = await track.recv() frame = await track.recv()
if np.max(np.abs(frame.to_ndarray())) > 100: if np.max(np.abs(frame.to_ndarray())) > 100:
audio_received.set() audio_received.set()
@@ -154,11 +158,14 @@ class VoiceSmokeTest(unittest.IsolatedAsyncioTestCase):
@channel.on("message") @channel.on("message")
def on_message(raw): def on_message(raw):
nonlocal completed_responses, function_calls
event = json.loads(raw) event = json.loads(raw)
events.append(event) events.append(event)
event_times.append((round(time.monotonic() - microphone.start, 1), event.get("type")))
if event.get("type") == "response.done": if event.get("type") == "response.done":
output = event.get("response", {}).get("output", []) output = event.get("response", {}).get("output", [])
if output and output[0].get("type") == "function_call": if output and output[0].get("type") == "function_call":
function_calls += 1
call = output[0] call = output[0]
channel.send(json.dumps({"type": "conversation.item.create", "item": { channel.send(json.dumps({"type": "conversation.item.create", "item": {
"type": "function_call_output", "call_id": call["call_id"], "type": "function_call_output", "call_id": call["call_id"],
@@ -166,6 +173,8 @@ class VoiceSmokeTest(unittest.IsolatedAsyncioTestCase):
}})) }}))
channel.send(json.dumps({"type": "response.create"})) channel.send(json.dumps({"type": "response.create"}))
elif output and output[0].get("type") == "message": elif output and output[0].get("type") == "message":
completed_responses += 1
if completed_responses == 2:
done.set() done.set()
try: try:
@@ -186,13 +195,29 @@ class VoiceSmokeTest(unittest.IsolatedAsyncioTestCase):
answer = await response.text() answer = await response.text()
await pc.setRemoteDescription(RTCSessionDescription(sdp=answer, type="answer")) await pc.setRemoteDescription(RTCSessionDescription(sdp=answer, type="answer"))
try: try:
await asyncio.wait_for(done.wait(), timeout=20) await asyncio.wait_for(done.wait(), timeout=25)
except asyncio.TimeoutError: except asyncio.TimeoutError:
self.fail(f"No completed voice response; events: {[e.get('type') for e in events]}") 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) await asyncio.wait_for(audio_received.wait(), timeout=5)
self.assertEqual(function_calls, 2)
self.assertTrue(any(event.get("type") == self.assertTrue(any(event.get("type") ==
"conversation.item.input_audio_transcription.completed" "conversation.item.input_audio_transcription.completed"
for event in events)) 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) committed = next((index, event) for index, event in enumerate(events)
if event.get("type") == "input_audio_buffer.committed") if event.get("type") == "input_audio_buffer.committed")
completed = next((index, event) for index, event in enumerate(events) completed = next((index, event) for index, event in enumerate(events)