Keep realtime conversation items ordered across turns
This commit is contained in:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user