|
|
|
@@ -12,6 +12,8 @@ import asyncio
|
|
|
|
|
import queue
|
|
|
|
|
import threading
|
|
|
|
|
import time
|
|
|
|
|
from pathlib import Path
|
|
|
|
|
from types import SimpleNamespace
|
|
|
|
|
|
|
|
|
|
import pytest
|
|
|
|
|
|
|
|
|
@@ -205,6 +207,154 @@ def _run_test(coro_factory, timeout=10.0):
|
|
|
|
|
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)
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|