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 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:
|
||||||
|
|
||||||
|
|||||||
@@ -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,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,7 +173,9 @@ 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":
|
||||||
done.set()
|
completed_responses += 1
|
||||||
|
if completed_responses == 2:
|
||||||
|
done.set()
|
||||||
|
|
||||||
try:
|
try:
|
||||||
await pc.setLocalDescription(await pc.createOffer())
|
await pc.setLocalDescription(await pc.createOffer())
|
||||||
@@ -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)
|
||||||
|
|||||||
Reference in New Issue
Block a user