Keep realtime conversation items ordered across turns
This commit is contained in:
@@ -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; omitting the item caused “Realtime transcript refers to an unknown
|
||||
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:
|
||||
|
||||
|
||||
@@ -12,6 +12,7 @@ import hashlib
|
||||
import hmac
|
||||
import io
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
import uuid
|
||||
@@ -21,6 +22,7 @@ from fractions import Fraction
|
||||
import numpy as np
|
||||
from aiohttp import ClientSession, FormData, web
|
||||
from aiortc import MediaStreamTrack, RTCPeerConnection, RTCSessionDescription
|
||||
from aiortc.mediastreams import MediaStreamError
|
||||
from av import AudioFrame, AudioResampler
|
||||
from cryptography.exceptions import InvalidSignature
|
||||
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)
|
||||
HTTP = web.AppKey("http", ClientSession)
|
||||
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:
|
||||
@@ -170,8 +178,10 @@ class RealtimeSession:
|
||||
self.speech = bytearray()
|
||||
self.silence_since: float | None = None
|
||||
self.call_id: str | None = None
|
||||
self.last_item_id: str | None = None
|
||||
self.pending_reply: str | 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.silence_ms = int(os.environ.get("SILENCE_MS", "750"))
|
||||
|
||||
@@ -187,6 +197,11 @@ class RealtimeSession:
|
||||
except json.JSONDecodeError:
|
||||
return
|
||||
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":
|
||||
self.emit({"type": "session.updated", "session": {"id": "athena-local"}})
|
||||
elif kind == "conversation.item.create":
|
||||
@@ -218,12 +233,21 @@ class RealtimeSession:
|
||||
await self.on_audio(pcm)
|
||||
except asyncio.CancelledError:
|
||||
raise
|
||||
except MediaStreamError:
|
||||
return
|
||||
except Exception:
|
||||
logger.exception("incoming audio stream stopped unexpectedly")
|
||||
return
|
||||
|
||||
async def on_audio(self, pcm: bytes) -> None:
|
||||
if self.busy or not pcm:
|
||||
if not pcm:
|
||||
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
|
||||
now = time.monotonic()
|
||||
if voiced and not self.speaking:
|
||||
@@ -231,6 +255,7 @@ class RealtimeSession:
|
||||
self.speech.clear()
|
||||
self.silence_since = None
|
||||
self.emit({"type": "input_audio_buffer.speech_started"})
|
||||
logger.info("speech started")
|
||||
if not self.speaking:
|
||||
return
|
||||
self.speech.extend(pcm)
|
||||
@@ -242,6 +267,7 @@ class RealtimeSession:
|
||||
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"})
|
||||
logger.info("speech stopped bytes=%s", len(self.speech))
|
||||
audio = bytes(self.speech)
|
||||
self.speech.clear()
|
||||
if len(audio) >= SAMPLE_RATE * 2 // 4:
|
||||
@@ -265,14 +291,18 @@ class RealtimeSession:
|
||||
return
|
||||
item_id = f"item_{uuid.uuid4().hex}"
|
||||
# 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",
|
||||
"item_id": item_id, "transcript": text})
|
||||
logger.info("transcript completed item_id=%s", item_id)
|
||||
await self.consult(text)
|
||||
except asyncio.CancelledError:
|
||||
self.busy = False
|
||||
raise
|
||||
except Exception:
|
||||
logger.exception("speech transcription failed")
|
||||
self.busy = False
|
||||
self.emit({"type": "error", "error": {"message": "Athena STT failed"}})
|
||||
|
||||
@@ -289,13 +319,16 @@ class RealtimeSession:
|
||||
"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:
|
||||
if not text.strip():
|
||||
self.busy = False
|
||||
return
|
||||
response_id = f"resp_{uuid.uuid4().hex}"
|
||||
item_id = f"item_{uuid.uuid4().hex}"
|
||||
try:
|
||||
logger.info("speaking response_id=%s", response_id)
|
||||
self.emit({"type": "response.created", "response": {"id": response_id}})
|
||||
async with self.http.post(
|
||||
os.environ["ATHENA_API_BASE_URL"].rstrip("/") + "/audio/speech",
|
||||
@@ -306,18 +339,25 @@ class RealtimeSession:
|
||||
response.raise_for_status()
|
||||
pcm = read_wav(await response.read())
|
||||
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.
|
||||
await asyncio.sleep(len(pcm) / (SAMPLE_RATE * 2) + 0.2)
|
||||
self.emit({"type": "conversation.item.done", "item": item})
|
||||
self.emit({"type": "response.done", "response": {
|
||||
"id": response_id, "status": "completed", "output": [
|
||||
{"id": f"item_{uuid.uuid4().hex}", "type": "message", "role": "assistant", "status": "completed"}
|
||||
],
|
||||
"id": response_id, "status": "completed", "output": [item],
|
||||
}})
|
||||
logger.info("speech response complete response_id=%s", response_id)
|
||||
except asyncio.CancelledError:
|
||||
self.emit({"type": "response.cancelled", "response": {"id": response_id}})
|
||||
raise
|
||||
except Exception:
|
||||
logger.exception("speech synthesis failed")
|
||||
self.emit({"type": "error", "error": {"message": "Athena TTS failed"}})
|
||||
finally:
|
||||
self.busy = False
|
||||
|
||||
@@ -31,8 +31,8 @@ class Microphone(MediaStreamTrack):
|
||||
|
||||
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:
|
||||
# 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:
|
||||
@@ -132,16 +132,20 @@ class VoiceSmokeTest(unittest.IsolatedAsyncioTestCase):
|
||||
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")
|
||||
pc.addTrack(Microphone())
|
||||
microphone = Microphone()
|
||||
pc.addTrack(microphone)
|
||||
|
||||
@pc.on("track")
|
||||
def on_track(track):
|
||||
if track.kind == "audio":
|
||||
async def consume():
|
||||
for _ in range(150):
|
||||
for _ in range(600):
|
||||
frame = await track.recv()
|
||||
if np.max(np.abs(frame.to_ndarray())) > 100:
|
||||
audio_received.set()
|
||||
@@ -154,11 +158,14 @@ class VoiceSmokeTest(unittest.IsolatedAsyncioTestCase):
|
||||
|
||||
@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"],
|
||||
@@ -166,6 +173,8 @@ class VoiceSmokeTest(unittest.IsolatedAsyncioTestCase):
|
||||
}}))
|
||||
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:
|
||||
@@ -186,13 +195,29 @@ class VoiceSmokeTest(unittest.IsolatedAsyncioTestCase):
|
||||
answer = await response.text()
|
||||
await pc.setRemoteDescription(RTCSessionDescription(sdp=answer, type="answer"))
|
||||
try:
|
||||
await asyncio.wait_for(done.wait(), timeout=20)
|
||||
await asyncio.wait_for(done.wait(), timeout=25)
|
||||
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)
|
||||
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)
|
||||
|
||||
Reference in New Issue
Block a user