Stream realtime TTS audio as it is synthesized

This commit is contained in:
Mikei386
2026-09-16 14:55:12 +02:00
parent 97c8072b3d
commit 3f684b7325
3 changed files with 60 additions and 31 deletions
+43 -28
View File
@@ -137,27 +137,30 @@ class SpeechTrack(MediaStreamTrack):
def __init__(self):
super().__init__()
self._queue: asyncio.Queue[bytes] = asyncio.Queue(maxsize=2000)
self._pending = bytearray()
self._frames_enqueued = 0
self._frames_played = 0
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 enqueue(self, pcm: bytes) -> int:
if len(pcm) != FRAME_SAMPLES * 2:
raise ValueError("SpeechTrack expects one complete audio frame")
await self._queue.put(pcm)
self._frames_enqueued += 1
return self._frames_enqueued
async def wait_played(self, target: int) -> None:
while self._frames_played < target:
await asyncio.sleep(0.02)
await asyncio.sleep(0.2)
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")
try:
chunk = self._queue.get_nowait()
self._frames_played += 1
except asyncio.QueueEmpty:
chunk = b"\x00" * FRAME_SAMPLES * 2
frame = AudioFrame.from_ndarray(
np.frombuffer(chunk, dtype="<i2").reshape(1, FRAME_SAMPLES),
format="s16", layout="mono",
@@ -331,23 +334,35 @@ class RealtimeSession:
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",
os.environ["ATHENA_API_BASE_URL"].rstrip("/") + "/audio/speech/pcm-stream",
headers={"Content-Type": "application/json", **self.auth_headers()},
json={"input": text, "voice": os.environ.get("TTS_VOICE", "alloy"),
"response_format": "wav"}, timeout=120,
"chunk_size": 4}, timeout=120,
) as response:
response.raise_for_status()
pcm = read_wav(await response.read())
self.track.enqueue(pcm)
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)
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})
pending = bytearray()
target = 0
frame_bytes = FRAME_SAMPLES * 2
async for chunk in response.content.iter_chunked(16384):
pending.extend(chunk)
while len(pending) >= frame_bytes:
target = await self.track.enqueue(bytes(pending[:frame_bytes]))
del pending[:frame_bytes]
if len(pending) % 2:
raise ValueError("Athena TTS returned incomplete PCM samples")
if pending:
target = await self.track.enqueue(bytes(pending).ljust(frame_bytes, b"\0"))
if not target:
raise ValueError("Athena TTS returned no audio")
# Keep half-duplex active until the streamed audio has played.
await self.track.wait_played(target)
self.emit({"type": "conversation.item.done", "item": item})
self.emit({"type": "response.done", "response": {
"id": response_id, "status": "completed", "output": [item],