fix(gateway): flush spoken acknowledgments at tool boundaries

This commit is contained in:
Franci Penov
2026-09-08 16:22:10 -07:00
committed by Teknium
parent fc36288832
commit c50bf92f44
4 changed files with 162 additions and 4 deletions
+1
View File
@@ -0,0 +1 @@
francip
+7 -1
View File
@@ -876,7 +876,7 @@ class TurnRunner:
delta_sinks = [sc for sc in ((stream_consumer if want_stream_deltas else None), stts) if sc is not None] delta_sinks = [sc for sc in ((stream_consumer if want_stream_deltas else None), stts) if sc is not None]
stream_delta_cb = None stream_delta_cb = None
if delta_sinks: if delta_sinks:
def stream_delta_cb(text: str) -> None: def stream_delta_cb(text: Optional[str]) -> None:
if ctx._run_still_current(): if ctx._run_still_current():
for sink in delta_sinks: for sink in delta_sinks:
sink.on_delta(text) sink.on_delta(text)
@@ -884,6 +884,12 @@ class TurnRunner:
def interim_assistant_cb(text: str, *, already_streamed: bool = False) -> None: def interim_assistant_cb(text: str, *, already_streamed: bool = False) -> None:
if not ctx._run_still_current(): if not ctx._run_still_current():
return return
if stts is not None:
# Flush accepted deltas; completed commentary is a separate speech segment.
stts.on_delta(None)
if not already_streamed:
stts.on_delta(text)
stts.on_delta(None)
if stream_consumer is not None: if stream_consumer is not None:
stream_consumer.on_segment_break() if already_streamed else stream_consumer.on_commentary(text) stream_consumer.on_segment_break() if already_streamed else stream_consumer.on_commentary(text)
elif not already_streamed and ctx._status_adapter and str(text or "").strip(): elif not already_streamed and ctx._status_adapter and str(text or "").strip():
+4 -3
View File
@@ -66,11 +66,12 @@ class StreamingTTSConsumer:
if log_errors: if log_errors:
logger.debug("streaming TTS on_delta error", exc_info=True) logger.debug("streaming TTS on_delta error", exc_info=True)
def on_delta(self, text: str) -> None: def on_delta(self, text: Optional[str]) -> None:
"""Receive a text delta from the agent. Non-blocking.""" """Receive text, or flush a ``None`` segment boundary without ending audio. Non-blocking."""
if self._aborted or not self.active or self._finished: if self._aborted or not self.active or self._finished:
return return
self._enqueue_clauses(self._chunker.feed(text), "streaming TTS queue full, dropping clause", clauses = self._chunker.flush() if text is None else self._chunker.feed(text)
self._enqueue_clauses(clauses, "streaming TTS queue full, dropping clause",
log_errors=True) log_errors=True)
def finish(self) -> None: def finish(self) -> None:
@@ -12,6 +12,8 @@ import asyncio
import queue import queue
import threading import threading
import time import time
from pathlib import Path
from types import SimpleNamespace
import pytest import pytest
@@ -205,6 +207,154 @@ def _run_test(coro_factory, timeout=10.0):
loop.close() loop.close()
@pytest.fixture
def gateway_tts_turn(monkeypatch, tmp_path):
"""Real agent delivery and gateway consumers, with only the audio transport faked."""
from agent.agent_runtime_helpers import strip_think_blocks
from agent.stream_delivery import StreamDeliveryMixin
from gateway.config import StreamingConfig
from gateway.run_turn_runner import TurnRunner
from gateway.stream_consumer import StreamConsumerConfig
from gateway.turn_context import TurnContext
monkeypatch.setattr(Path, "home", lambda: tmp_path)
monkeypatch.setenv("HERMES_HOME", str(tmp_path / ".hermes"))
class Agent(StreamDeliveryMixin):
_strip_think_blocks = strip_think_blocks
class Streamer(FakeStreamer):
def __init__(self):
super().__init__()
self.clauses = []
def stream(self, text):
self.clauses.append(text)
yield b"\x01\x00" * 480
class Adapter(FakeVoiceAdapter):
def __init__(self):
super().__init__()
self.heard = asyncio.Queue()
async def write_streaming_tts(self, handle, chunk):
await super().write_streaming_tts(handle, chunk)
self.heard.put_nowait(chunk)
def make(loop, *, text_streaming=True, interim_enabled=True, active=True):
streamer, adapter = Streamer(), Adapter()
monkeypatch.setattr("tools.tts_streaming.resolve_streaming_provider", lambda cfg: streamer if active else None)
tts = StreamingTTSConsumer(adapter, "voice", {}, loop)
current = [True]
ctx = TurnContext(
streaming_tts_consumer_holder=[tts], user_config={},
resolve_display_setting=lambda *args: text_streaming,
interim_assistant_messages_enabled=interim_enabled,
source=SimpleNamespace(platform=SimpleNamespace(value="realtime"), chat_id="voice"),
_run_still_current=lambda: current[0],
)
runner = SimpleNamespace(
config=SimpleNamespace(streaming=StreamingConfig()),
_adapter_for_source=lambda source: adapter,
_build_stream_consumer_config=lambda *args, **kwargs: (StreamConsumerConfig(), None),
)
_, delta, interim, want_interim = TurnRunner(runner, ctx)._setup_stream_consumer("realtime")
agent = Agent()
agent.stream_delta_callback, agent._stream_callback = delta, None
agent.interim_assistant_callback = interim if want_interim else None
return agent, tts, adapter, streamer, current
return make
@pytest.mark.parametrize("text_streaming", [False, True])
@pytest.mark.parametrize("interim_enabled", [False, True])
@pytest.mark.parametrize("source", ["streamed", "commentary", "commentary_after_delta", "interim"])
def test_gateway_speaks_acknowledgment_before_tool_result(
gateway_tts_turn, text_streaming, interim_enabled, source,
):
async def run():
agent, tts, adapter, streamer, _ = gateway_tts_turn(
asyncio.get_running_loop(), text_streaming=text_streaming, interim_enabled=interim_enabled,
)
acknowledgment, result = "I will check that.", "The result is available."
leading = "One moment." if source == "commentary_after_delta" else ""
def before_tool():
message = {"role": "assistant", "content": acknowledgment}
if leading:
assert agent._deliver_to_stream_callbacks(leading)
agent._record_streamed_assistant_text(leading)
if source == "streamed":
assert agent._deliver_to_stream_callbacks(acknowledgment)
agent._record_streamed_assistant_text(acknowledgment)
elif source in ("commentary", "commentary_after_delta"):
agent._fire_streamed_codex_commentary(acknowledgment)
message = {"role": "assistant", "content": "", "codex_message_items": [{
"type": "message", "phase": "commentary",
"content": [{"type": "output_text", "text": acknowledgment}],
}]}
agent._emit_interim_assistant_message(message)
# A tool boundary must flush deltas even when optional commentary is off.
if source == "streamed" or not interim_enabled:
agent.stream_delta_callback(None)
tts.start()
try:
await asyncio.to_thread(before_tool)
should_acknowledge = source == "streamed" or interim_enabled
before_result = ([leading] if leading else []) + ([acknowledgment] if should_acknowledge else [])
if before_result:
# The tool result cannot be sent until the acknowledgment reaches PCM.
for _ in before_result:
await asyncio.wait_for(adapter.heard.get(), timeout=3)
assert streamer.clauses == before_result
assert adapter.finish_count == 0
assert not tts.done
await asyncio.to_thread(agent.stream_delta_callback, None)
await asyncio.to_thread(agent.stream_delta_callback, result)
tts.finish()
assert await tts.wait_complete(timeout=3)
assert streamer.clauses == before_result + [result]
assert adapter.begin_count == adapter.finish_count == 1
assert adapter.abort_count == 0
assert tts.suppress_whole_file
finally:
tts.finish()
await tts.wait_complete(timeout=3)
asyncio.run(run())
@pytest.mark.parametrize("state", ["stale", "finished", "aborted", "inactive"])
def test_gateway_segment_callbacks_preserve_terminal_guards(gateway_tts_turn, state):
async def run():
agent, tts, adapter, streamer, current = gateway_tts_turn(
asyncio.get_running_loop(), active=state != "inactive",
)
if state == "stale":
current[0] = False
elif state == "finished":
tts.finish()
elif state == "aborted":
tts.abort()
def late_callbacks():
agent.stream_delta_callback("Late streamed text.")
agent.stream_delta_callback(None)
agent._fire_streamed_codex_commentary("Late commentary.")
await asyncio.to_thread(late_callbacks)
tts.finish()
tts.start()
await tts.wait_complete(timeout=3)
assert tts.done
assert streamer.clauses == []
assert adapter.written_chunks == []
asyncio.run(run())
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# Adapter contract defaults (BasePlatformAdapter) # Adapter contract defaults (BasePlatformAdapter)
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------