fix(gateway): flush spoken acknowledgments at tool boundaries
This commit is contained in:
@@ -0,0 +1 @@
|
|||||||
|
francip
|
||||||
@@ -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():
|
||||||
|
|||||||
@@ -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)
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|||||||
Reference in New Issue
Block a user