mirror of
https://github.com/strands-agents/harness-sdk.git
synced 2026-10-02 02:44:48 +08:00
feat(bidi): proactive session reconnect for bidi session (#3874)
This commit is contained in:
@@ -24,6 +24,7 @@ from .types.events import (
|
||||
BidiConnectionCloseEvent,
|
||||
BidiConnectionRestartEvent,
|
||||
BidiConnectionStartEvent,
|
||||
BidiConnectionWarningEvent,
|
||||
BidiErrorEvent,
|
||||
BidiImageInputEvent,
|
||||
BidiInputEvent,
|
||||
@@ -37,6 +38,9 @@ from .types.events import (
|
||||
ModalityUsage,
|
||||
)
|
||||
|
||||
# Reconnect configuration (declared by providers, tunable via provider_config)
|
||||
from .types.model import BidiConnectionConfig
|
||||
|
||||
__all__ = [
|
||||
# Main interface
|
||||
"BidiAgent",
|
||||
@@ -48,6 +52,7 @@ __all__ = [
|
||||
# Output Event types
|
||||
"BidiConnectionStartEvent",
|
||||
"BidiConnectionRestartEvent",
|
||||
"BidiConnectionWarningEvent",
|
||||
"BidiConnectionCloseEvent",
|
||||
"BidiResponseStartEvent",
|
||||
"BidiResponseCompleteEvent",
|
||||
@@ -58,6 +63,8 @@ __all__ = [
|
||||
"ModalityUsage",
|
||||
"BidiErrorEvent",
|
||||
"BidiOutputEvent",
|
||||
# Reconnect configuration
|
||||
"BidiConnectionConfig",
|
||||
# Tool Event types (reused from standard agent)
|
||||
"ToolUseStreamEvent",
|
||||
"ToolResultEvent",
|
||||
|
||||
@@ -125,21 +125,31 @@ def end_response_span(
|
||||
tracer._end_span(span, attributes=attributes, error=error)
|
||||
|
||||
|
||||
def start_restart_span(tracer: Tracer, parent_span: Span | None = None, error_message: str | None = None) -> Span:
|
||||
def start_restart_span(
|
||||
tracer: Tracer,
|
||||
parent_span: Span | None = None,
|
||||
reason: str = "timeout",
|
||||
error_message: str | None = None,
|
||||
) -> Span:
|
||||
"""Start a span for a connection restart.
|
||||
|
||||
Args:
|
||||
tracer: Tracer instance.
|
||||
parent_span: Parent session span.
|
||||
error_message: The timeout error message that triggered the restart.
|
||||
reason: What triggered the restart ("timeout" reactively, "scheduled" proactively).
|
||||
error_message: The timeout error message that triggered a reactive restart, if any.
|
||||
|
||||
Returns:
|
||||
The restart span.
|
||||
"""
|
||||
attributes: dict[str, AttributeValue] = tracer._get_common_attributes(operation_name="bidi_connection_restart")
|
||||
attributes["gen_ai.bidi.restart_reason"] = reason
|
||||
|
||||
# The message of the timeout that triggered a reactive restart. Namespaced under
|
||||
# gen_ai.bidi.* (not gen_ai.error.*) so it does not read as a failure of this span,
|
||||
# which may end successfully; a failed restart is recorded via end_restart_span.
|
||||
if error_message:
|
||||
attributes["gen_ai.error.message"] = error_message
|
||||
attributes["gen_ai.bidi.restart_error_message"] = error_message
|
||||
|
||||
return tracer._start_span("bidi_connection_restart", parent_span, attributes=attributes)
|
||||
|
||||
|
||||
@@ -0,0 +1,97 @@
|
||||
"""Proactive reconnect timer for bidirectional streaming.
|
||||
|
||||
``_BidiReconnectTimer`` fires a warning then a deadline callback at caller-supplied offsets;
|
||||
it holds no reconnect policy. ``resolve_deadline_s`` reads the deadline from a provider's
|
||||
declared ``BidiConnectionConfig``.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
from collections.abc import Awaitable, Callable
|
||||
|
||||
from ..types.model import BidiConnectionConfig
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def resolve_deadline_s(connection_config: BidiConnectionConfig) -> float | None:
|
||||
"""Resolve the proactive reconnect deadline in seconds from a connection config.
|
||||
|
||||
Args:
|
||||
connection_config: Provider-declared reconnect timing.
|
||||
|
||||
Returns:
|
||||
``restart_after_s`` if declared and positive, else ``None`` (no proactive timer).
|
||||
"""
|
||||
restart_after_s = connection_config.get("restart_after_s")
|
||||
if restart_after_s is None or restart_after_s <= 0:
|
||||
return None
|
||||
return restart_after_s
|
||||
|
||||
|
||||
class _BidiReconnectTimer:
|
||||
"""Fire a warning then a deadline callback ahead of a provider's connection limit.
|
||||
|
||||
The clock is injectable so tests can drive timing without wall time.
|
||||
|
||||
Attributes:
|
||||
_sleep: Injectable async sleep, defaults to ``asyncio.sleep``.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
on_warning: Callable[[float], Awaitable[None]],
|
||||
on_deadline: Callable[[], Awaitable[None]],
|
||||
sleep: Callable[[float], Awaitable[None]] | None = None,
|
||||
) -> None:
|
||||
"""Initialize the timer.
|
||||
|
||||
Args:
|
||||
on_warning: Awaitable called with seconds-left when the warning lead elapses.
|
||||
on_deadline: Awaitable called when the reconnect deadline elapses.
|
||||
sleep: Injectable async sleep (for tests). Defaults to ``asyncio.sleep``.
|
||||
"""
|
||||
self._on_warning = on_warning
|
||||
self._on_deadline = on_deadline
|
||||
self._sleep = sleep or asyncio.sleep
|
||||
self._task: asyncio.Task | None = None
|
||||
|
||||
def arm(self, deadline_s: float, warning_lead_s: float) -> None:
|
||||
"""Arm the warning and deadline timers, cancelling any previously armed cycle.
|
||||
|
||||
Args:
|
||||
deadline_s: Seconds from now until the deadline callback fires.
|
||||
warning_lead_s: Seconds before the deadline to fire the warning callback.
|
||||
"""
|
||||
self.cancel()
|
||||
self._task = asyncio.create_task(self._run(deadline_s, warning_lead_s))
|
||||
logger.debug(
|
||||
"deadline_s=<%.1f>, warning_lead_s=<%.1f> | proactive reconnect timer armed",
|
||||
deadline_s,
|
||||
warning_lead_s,
|
||||
)
|
||||
|
||||
def cancel(self) -> None:
|
||||
"""Cancel the armed timer, if any. Safe to call when idle."""
|
||||
if self._task is not None:
|
||||
self._task.cancel()
|
||||
self._task = None
|
||||
|
||||
async def _run(self, deadline_s: float, warning_lead_s: float) -> None:
|
||||
"""Sleep until the warning lead, fire the warning, then fire the deadline.
|
||||
|
||||
The warning fires ``warning_lead_s`` before the deadline. When the lead is zero
|
||||
or exceeds the deadline, the warning is emitted immediately and the remaining
|
||||
wait runs down to the deadline.
|
||||
"""
|
||||
warning_at_s = max(deadline_s - warning_lead_s, 0.0)
|
||||
|
||||
await self._sleep(warning_at_s)
|
||||
time_left_s = deadline_s - warning_at_s
|
||||
await self._on_warning(time_left_s)
|
||||
|
||||
await self._sleep(deadline_s - warning_at_s)
|
||||
# Detach before the callback re-arms this timer; cancelling a live self-reference
|
||||
# would abort the reconnect the callback runs.
|
||||
self._task = None
|
||||
await self._on_deadline()
|
||||
@@ -8,7 +8,8 @@ import logging
|
||||
import time
|
||||
import warnings
|
||||
from collections.abc import AsyncGenerator
|
||||
from typing import TYPE_CHECKING, Any, cast
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Any, Literal, cast
|
||||
|
||||
from opentelemetry.trace import Span
|
||||
|
||||
@@ -27,11 +28,12 @@ from ...hooks.events import (
|
||||
)
|
||||
from .. import _telemetry
|
||||
from .._async import _TaskPool, stop_all
|
||||
from ..models import BidiModelTimeoutError
|
||||
from ..models import BidiModel, BidiModelTimeoutError
|
||||
from ..types.events import (
|
||||
BidiAudioStreamEvent,
|
||||
BidiConnectionCloseEvent,
|
||||
BidiConnectionRestartEvent,
|
||||
BidiConnectionWarningEvent,
|
||||
BidiInputEvent,
|
||||
BidiInterruptionEvent,
|
||||
BidiOutputEvent,
|
||||
@@ -41,12 +43,35 @@ from ..types.events import (
|
||||
BidiTranscriptStreamEvent,
|
||||
BidiUsageEvent,
|
||||
)
|
||||
from ..types.model import BidiConnectionConfig
|
||||
from ._reconnect_timer import _BidiReconnectTimer, resolve_deadline_s
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .agent import BidiAgent
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Bound on awaiting a superseded reader after its stream is closed before cancelling it.
|
||||
_READER_REAP_TIMEOUT_S = 2.0
|
||||
|
||||
# Fixed advance notice, in seconds before a scheduled reconnect, for the warning event.
|
||||
_WARNING_LEAD_S = 10.0
|
||||
|
||||
# Max seconds a proactive reconnect waits for a turn boundary before forcing the swap.
|
||||
_TURN_ALIGN_WAIT_S = 10.0
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _ReaderError:
|
||||
"""A model-reader error tagged with the connection generation it was raised on.
|
||||
|
||||
receive() uses the tag to drop the error if the connection was superseded before it was
|
||||
consumed.
|
||||
"""
|
||||
|
||||
generation: int
|
||||
error: Exception
|
||||
|
||||
|
||||
class _BidiAgentLoop:
|
||||
"""Agent loop.
|
||||
@@ -60,7 +85,7 @@ class _BidiAgentLoop:
|
||||
This allows passing custom data (user_id, session_id, database connections, etc.)
|
||||
that tools can access via their invocation_state parameter.
|
||||
_send_gate: Gate the sending of events to the model.
|
||||
Blocks when agent is resetting the model connection after timeout.
|
||||
Blocks while the agent is reconnecting the model connection.
|
||||
"""
|
||||
|
||||
def __init__(self, agent: "BidiAgent") -> None:
|
||||
@@ -75,16 +100,47 @@ class _BidiAgentLoop:
|
||||
self._started = False
|
||||
self._task_pool = _TaskPool()
|
||||
self._event_queue: asyncio.Queue
|
||||
# Connection-lifecycle events (warning, restart); unbounded and drained by receive() with
|
||||
# priority over the data queue, so a busy consumer neither drops them nor stalls the timer.
|
||||
self._lifecycle_queue: asyncio.Queue
|
||||
# Holds a data event pulled alongside a lifecycle event, for the next receive() iteration.
|
||||
self._pending_data: BidiOutputEvent | Exception | None = None
|
||||
self._invocation_state: dict[str, Any]
|
||||
self._model_task: asyncio.Task | None = None
|
||||
|
||||
self._send_gate = asyncio.Event()
|
||||
|
||||
self._tracer = get_tracer()
|
||||
self._session_span: Span | None = None
|
||||
self._accumulated_input_tokens: int
|
||||
self._accumulated_output_tokens: int
|
||||
self._accumulated_total_tokens: int
|
||||
self._accumulated_cache_read_tokens: int
|
||||
|
||||
# Session totals = baseline (finished connections) + current (this connection).
|
||||
self._current_input_tokens = 0
|
||||
self._current_output_tokens = 0
|
||||
self._current_total_tokens = 0
|
||||
self._current_cache_read_tokens = 0
|
||||
self._baseline_input_tokens = 0
|
||||
self._baseline_output_tokens = 0
|
||||
self._baseline_total_tokens = 0
|
||||
self._baseline_cache_read_tokens = 0
|
||||
|
||||
self._reconnect_timer = _BidiReconnectTimer(
|
||||
on_warning=self._on_reconnect_warning,
|
||||
on_deadline=self._on_reconnect_deadline,
|
||||
)
|
||||
# Guards _restart_connection against concurrent reactive + proactive entry.
|
||||
self._reconnecting = False
|
||||
# Incremented per reconnect so a superseded reader's events (and its stream-close
|
||||
# error) are dropped rather than forwarded after the swap.
|
||||
self._generation = 0
|
||||
|
||||
# Turn-boundary tracking, so a proactive reconnect waits for the current turn to
|
||||
# finish rather than cutting off a response or dropping an unanswered user turn.
|
||||
# A provider that emits neither response nor transcript events never leaves the
|
||||
# boundary state, so the aligned wait is a no-op (reconnect fires immediately).
|
||||
self._response_active = False
|
||||
self._awaiting_response = False
|
||||
self._turn_complete = asyncio.Event()
|
||||
self._turn_complete.set()
|
||||
|
||||
async def start(self, invocation_state: dict[str, Any] | None = None) -> None:
|
||||
"""Start the agent loop.
|
||||
@@ -130,26 +186,32 @@ class _BidiAgentLoop:
|
||||
self._session_span = None
|
||||
raise
|
||||
_telemetry.end_connection_span(self._tracer, connection_span)
|
||||
self._accumulated_input_tokens = 0
|
||||
self._accumulated_output_tokens = 0
|
||||
self._accumulated_total_tokens = 0
|
||||
self._accumulated_cache_read_tokens = 0
|
||||
self._reset_token_tracking()
|
||||
self._reset_turn_state()
|
||||
|
||||
self._event_queue = asyncio.Queue(maxsize=1)
|
||||
self._lifecycle_queue = asyncio.Queue()
|
||||
self._pending_data = None
|
||||
|
||||
self._task_pool = _TaskPool()
|
||||
self._task_pool.create(self._run_model())
|
||||
self._model_task = self._task_pool.create(self._run_model(self._generation))
|
||||
|
||||
self._invocation_state = invocation_state or {}
|
||||
self._send_gate.set()
|
||||
self._started = True
|
||||
|
||||
self._arm_reconnect_timer()
|
||||
|
||||
async def stop(self) -> None:
|
||||
"""Stop the agent loop."""
|
||||
logger.debug("agent loop stopping")
|
||||
|
||||
self._started = False
|
||||
self._send_gate.clear()
|
||||
self._reconnect_timer.cancel()
|
||||
# Unblock a deadline callback waiting on a turn boundary (it is past the timer's cancel);
|
||||
# once released it re-checks _started and no-ops.
|
||||
self._turn_complete.set()
|
||||
self._invocation_state = {}
|
||||
|
||||
async def stop_tasks() -> None:
|
||||
@@ -195,13 +257,19 @@ class _BidiAgentLoop:
|
||||
if isinstance(event, BidiTextInputEvent):
|
||||
message: Message = {"role": event.role, "content": [{"text": event.text}]}
|
||||
await self._agent._append_messages(message)
|
||||
if event.role == "user":
|
||||
# A user text turn owes a response, same as a finished audio turn. Mark it so a
|
||||
# proactive reconnect waits for the reply instead of swapping mid-turn; without
|
||||
# this, a text-driven session always looks idle and the turn can be cut.
|
||||
self._awaiting_response = True
|
||||
self._update_turn_state()
|
||||
|
||||
await self._agent.model.send(event)
|
||||
|
||||
async def receive(self) -> AsyncGenerator[BidiOutputEvent, None]:
|
||||
"""Receive model and tool call events.
|
||||
|
||||
Returns:
|
||||
Yields:
|
||||
Model and tool call events.
|
||||
|
||||
Raises:
|
||||
@@ -211,12 +279,29 @@ class _BidiAgentLoop:
|
||||
raise RuntimeError("loop not started | call start before receiving")
|
||||
|
||||
while True:
|
||||
event = await self._event_queue.get()
|
||||
if isinstance(event, BidiModelTimeoutError):
|
||||
logger.debug("model timeout error received")
|
||||
yield BidiConnectionRestartEvent(event)
|
||||
await self._restart_connection(event)
|
||||
continue
|
||||
event = await self._next_event()
|
||||
if isinstance(event, _ReaderError):
|
||||
if event.generation != self._generation:
|
||||
# Superseded reader: its connection was already replaced, so drop the error
|
||||
# rather than surface or restart on it.
|
||||
logger.debug("dropping stale reader error from a superseded connection")
|
||||
continue
|
||||
error = event.error
|
||||
if isinstance(error, BidiModelTimeoutError):
|
||||
logger.debug("model timeout error received")
|
||||
if not self._auto_reconnect_enabled():
|
||||
logger.debug("auto_reconnect disabled | surfacing timeout to caller")
|
||||
raise error
|
||||
yield BidiConnectionRestartEvent(
|
||||
reason="timeout",
|
||||
timeout_error=error,
|
||||
turn_interrupted=not self._turn_complete.is_set(),
|
||||
)
|
||||
# Raise-time generation: a swap during the yield makes _restart_connection
|
||||
# decline this now-stale trigger.
|
||||
await self._restart_connection(error, event.generation)
|
||||
continue
|
||||
raise error
|
||||
|
||||
if isinstance(event, Exception):
|
||||
raise event
|
||||
@@ -228,55 +313,311 @@ class _BidiAgentLoop:
|
||||
|
||||
yield event
|
||||
|
||||
async def _restart_connection(self, timeout_error: BidiModelTimeoutError) -> None:
|
||||
"""Restart the model connection after timeout.
|
||||
async def _next_event(self) -> Any:
|
||||
"""Return the next event, draining the lifecycle queue with priority over data events.
|
||||
|
||||
When both queues are idle, waits on whichever produces first; a data event pulled
|
||||
alongside a lifecycle event is held in ``_pending_data`` for the next call.
|
||||
"""
|
||||
if not self._lifecycle_queue.empty():
|
||||
return self._lifecycle_queue.get_nowait()
|
||||
if self._pending_data is not None:
|
||||
event, self._pending_data = self._pending_data, None
|
||||
return event
|
||||
if not self._event_queue.empty():
|
||||
return self._event_queue.get_nowait()
|
||||
|
||||
get_lifecycle = asyncio.ensure_future(self._lifecycle_queue.get())
|
||||
get_data = asyncio.ensure_future(self._event_queue.get())
|
||||
done, pending = await asyncio.wait({get_lifecycle, get_data}, return_when=asyncio.FIRST_COMPLETED)
|
||||
for task in pending:
|
||||
task.cancel()
|
||||
if get_lifecycle in done:
|
||||
if get_data in done:
|
||||
self._pending_data = get_data.result()
|
||||
return get_lifecycle.result()
|
||||
return get_data.result()
|
||||
|
||||
def _connection_config(self) -> BidiConnectionConfig:
|
||||
"""Return the model's declared connection config, or an empty config if none.
|
||||
|
||||
Providers are not required to declare ``connection_config``; a missing or empty
|
||||
config means reactive-only reconnect with no proactive timer.
|
||||
"""
|
||||
connection_config = getattr(self._agent.model, "connection_config", None)
|
||||
return cast(BidiConnectionConfig, connection_config) if connection_config else {}
|
||||
|
||||
def _auto_reconnect_enabled(self) -> bool:
|
||||
"""Whether the agent reconnects automatically.
|
||||
|
||||
Automatic reconnect is the default: a provider is opted in unless it explicitly
|
||||
declares ``auto_reconnect: False`` in its ``connection_config``. A provider that
|
||||
declares no ``connection_config`` at all is treated as opted in.
|
||||
"""
|
||||
return self._connection_config().get("auto_reconnect", True)
|
||||
|
||||
def _arm_reconnect_timer(self) -> None:
|
||||
"""Arm the proactive reconnect timer when the model opts in with a declared deadline.
|
||||
|
||||
Owns the arming policy (auto_reconnect + a declared ``restart_after_s``); the timer
|
||||
itself is a pure mechanism. A no-op when reconnect is disabled or none is declared.
|
||||
"""
|
||||
if not self._auto_reconnect_enabled():
|
||||
return
|
||||
deadline_s = resolve_deadline_s(self._connection_config())
|
||||
if deadline_s is None:
|
||||
return
|
||||
self._reconnect_timer.arm(deadline_s, _WARNING_LEAD_S)
|
||||
|
||||
async def _on_reconnect_warning(self, time_left_s: float) -> None:
|
||||
"""Timer callback: surface an approaching-reconnect warning to the receiver.
|
||||
|
||||
Emitted on the lifecycle queue: non-blocking (so it never stalls the timer's run to the
|
||||
deadline) and unbounded (so a busy consumer never drops it).
|
||||
"""
|
||||
logger.debug("time_left_s=<%.1f> | emitting connection warning", time_left_s)
|
||||
self._lifecycle_queue.put_nowait(BidiConnectionWarningEvent(time_left_s=time_left_s))
|
||||
|
||||
async def _on_reconnect_deadline(self) -> None:
|
||||
"""Timer callback: align to a turn boundary, then reconnect proactively.
|
||||
|
||||
Waits (bounded) for the current turn to finish so the swap does not cut off a
|
||||
response or drop an unanswered user turn; surfaces any failure on the event queue.
|
||||
"""
|
||||
logger.debug("proactive reconnect deadline reached")
|
||||
# Capture before the wait so _restart_connection can decline if the loop stopped or a
|
||||
# reactive swap ran while we waited.
|
||||
generation = self._generation
|
||||
await self._await_turn_boundary()
|
||||
# A forced swap (the wait timed out) leaves _turn_complete clear: an in-progress or owed
|
||||
# turn is being cut and won't be answered on replay. Capture before the swap resets it.
|
||||
turn_interrupted = not self._turn_complete.is_set()
|
||||
try:
|
||||
reconnected = await self._restart_connection(None, generation)
|
||||
except Exception as error:
|
||||
await self._event_queue.put(error)
|
||||
return
|
||||
# Announce an actual swap on the lifecycle queue: non-blocking, and priority-drained by
|
||||
# receive() so it precedes any output from the new connection.
|
||||
if reconnected:
|
||||
self._lifecycle_queue.put_nowait(
|
||||
BidiConnectionRestartEvent(reason="scheduled", turn_interrupted=turn_interrupted)
|
||||
)
|
||||
|
||||
async def _await_turn_boundary(self) -> None:
|
||||
"""Wait for the current turn to finish, bounded so the reconnect beats the limit.
|
||||
|
||||
Returns immediately at a turn boundary (including for a provider that emits no turn
|
||||
events). Otherwise waits up to ``_TURN_ALIGN_WAIT_S`` for the turn to complete, then
|
||||
proceeds so the swap does not overrun the headroom below the provider limit.
|
||||
"""
|
||||
if self._turn_complete.is_set():
|
||||
return
|
||||
try:
|
||||
await asyncio.wait_for(self._turn_complete.wait(), timeout=_TURN_ALIGN_WAIT_S)
|
||||
except asyncio.TimeoutError:
|
||||
logger.debug("no turn boundary within %.1fs | forcing reconnect", _TURN_ALIGN_WAIT_S)
|
||||
|
||||
def _reset_turn_state(self) -> None:
|
||||
"""Reset turn tracking to the idle boundary state."""
|
||||
self._response_active = False
|
||||
self._awaiting_response = False
|
||||
self._turn_complete.set()
|
||||
|
||||
def _update_turn_state(self) -> None:
|
||||
"""Mark the turn complete when idle, or in-progress while a response is owed/active."""
|
||||
if self._response_active or self._awaiting_response:
|
||||
self._turn_complete.clear()
|
||||
else:
|
||||
self._turn_complete.set()
|
||||
|
||||
async def _restart_connection(self, timeout_error: BidiModelTimeoutError | None, generation: int) -> bool:
|
||||
"""Restart the model connection, reactively (after timeout) or proactively (timer).
|
||||
|
||||
The single guard point for both paths: declines when the loop has stopped, when the
|
||||
connection was already swapped since the trigger, or when a restart is in flight.
|
||||
|
||||
Args:
|
||||
timeout_error: Timeout error reported by the model.
|
||||
timeout_error: Timeout error on the reactive path, or ``None`` when proactive.
|
||||
generation: Connection generation the trigger was raised for; the restart is declined
|
||||
as stale if the connection has since been swapped.
|
||||
|
||||
Returns:
|
||||
``True`` if this call performed the swap, ``False`` if it declined. Raises if the
|
||||
swap itself fails.
|
||||
"""
|
||||
logger.debug("resetting model connection")
|
||||
if not self._started:
|
||||
logger.debug("loop stopped | ignoring reconnect trigger")
|
||||
return False
|
||||
if generation != self._generation:
|
||||
logger.debug(
|
||||
"trigger_generation=<%d>, current_generation=<%d> | connection already swapped | ignoring restart",
|
||||
generation,
|
||||
self._generation,
|
||||
)
|
||||
return False
|
||||
if self._reconnecting:
|
||||
logger.debug("reconnect already in progress | ignoring duplicate trigger")
|
||||
return False
|
||||
self._reconnecting = True
|
||||
self._reconnect_timer.cancel()
|
||||
|
||||
self._send_gate.clear()
|
||||
reason: Literal["timeout", "scheduled"] = "timeout" if timeout_error is not None else "scheduled"
|
||||
logger.debug("reason=<%s> | resetting model connection", reason)
|
||||
|
||||
# The before-restart hook runs before the restart begins; per the hook contract its
|
||||
# exceptions propagate to the caller (receive()), leaving the send gate closed.
|
||||
await self._agent.hooks.invoke_callbacks_async(BidiBeforeConnectionRestartEvent(self._agent, timeout_error))
|
||||
try:
|
||||
self._send_gate.clear()
|
||||
self._fold_token_baseline()
|
||||
|
||||
# A raising before-restart hook propagates out with the send gate left closed.
|
||||
await self._agent.hooks.invoke_callbacks_async(
|
||||
BidiBeforeConnectionRestartEvent(self._agent, reason=reason, timeout_error=timeout_error)
|
||||
)
|
||||
await self._swap_connection(reason, timeout_error)
|
||||
|
||||
self._reset_turn_state()
|
||||
self._arm_reconnect_timer()
|
||||
self._send_gate.set()
|
||||
finally:
|
||||
self._reconnecting = False
|
||||
|
||||
return True
|
||||
|
||||
async def _swap_connection(
|
||||
self, reason: Literal["timeout", "scheduled"], timeout_error: BidiModelTimeoutError | None
|
||||
) -> None:
|
||||
"""Swap to a new connection under a restart span, firing the after-restart hook.
|
||||
|
||||
Supersedes the current reader (generation bump) so its stream-close error is fenced
|
||||
rather than forwarded, then reconnects and starts the new reader. A failed swap is
|
||||
re-raised after telemetry and the after-restart hook report it, leaving the gate closed.
|
||||
"""
|
||||
restart_span = _telemetry.start_restart_span(
|
||||
self._tracer, parent_span=self._session_span, error_message=str(timeout_error)
|
||||
self._tracer,
|
||||
parent_span=self._session_span,
|
||||
reason=reason,
|
||||
error_message=str(timeout_error) if timeout_error is not None else None,
|
||||
)
|
||||
|
||||
restart_exception = None
|
||||
restart_kwargs = timeout_error.restart_config if timeout_error is not None else {}
|
||||
restart_exception: Exception | None = None
|
||||
try:
|
||||
await self._agent.model.stop()
|
||||
await self._agent.model.start(
|
||||
self._agent.system_prompt,
|
||||
self._agent.tool_registry.get_all_tool_specs(),
|
||||
self._agent.messages,
|
||||
**timeout_error.restart_config,
|
||||
)
|
||||
self._task_pool.create(self._run_model())
|
||||
previous_reader = self._model_task
|
||||
self._generation += 1
|
||||
await self._reconnect_model(restart_kwargs)
|
||||
await self._await_superseded_reader(previous_reader)
|
||||
self._model_task = self._task_pool.create(self._run_model(self._generation))
|
||||
except Exception as exception:
|
||||
restart_exception = exception
|
||||
finally:
|
||||
_telemetry.end_restart_span(self._tracer, restart_span, error=restart_exception)
|
||||
|
||||
await self._agent.hooks.invoke_callbacks_async(
|
||||
BidiAfterConnectionRestartEvent(self._agent, restart_exception)
|
||||
BidiAfterConnectionRestartEvent(self._agent, reason=reason, exception=restart_exception)
|
||||
)
|
||||
|
||||
# A failed restart leaves no running model task and a closed send gate, so surface it
|
||||
# to the receive() consumer rather than idling silently. Telemetry and the after-restart
|
||||
# hook have already reported the error above.
|
||||
if restart_exception is not None:
|
||||
raise restart_exception
|
||||
|
||||
self._send_gate.set()
|
||||
async def _reconnect_model(self, restart_kwargs: dict[str, Any]) -> None:
|
||||
"""Reconnect via the provider's ``reconnect()``, or ``stop()`` then ``start()``.
|
||||
|
||||
async def _run_model(self) -> None:
|
||||
The fallback is transitional: providers that have not implemented ``reconnect()``
|
||||
inherit the protocol no-op, so route them through stop/start until they adopt it.
|
||||
"""
|
||||
model = self._agent.model
|
||||
system_prompt = self._agent.system_prompt
|
||||
tools = self._agent.tool_registry.get_all_tool_specs()
|
||||
messages = self._agent.messages
|
||||
|
||||
# "Provider didn't override reconnect()" — it still resolves to the protocol's no-op
|
||||
# default, so fall back to stop/start. getattr on the type (not the instance) tolerates
|
||||
# a model whose class does not expose reconnect at all (e.g. a mock).
|
||||
if getattr(type(model), "reconnect", None) is BidiModel.reconnect:
|
||||
await model.stop()
|
||||
await model.start(system_prompt, tools, messages, **restart_kwargs)
|
||||
return
|
||||
|
||||
await model.reconnect(system_prompt, tools, messages, **restart_kwargs)
|
||||
|
||||
async def _await_superseded_reader(self, task: asyncio.Task | None) -> None:
|
||||
"""Await a superseded reader after its stream is closed; cancel only as a backstop.
|
||||
|
||||
The reader is expected to fall out of receive() when its connection is closed. It is
|
||||
awaited (not force-cancelled) so a live provider read is never interrupted; the
|
||||
bounded cancel handles a provider whose receive() does not unblock on close.
|
||||
"""
|
||||
if task is None or task.done():
|
||||
return
|
||||
try:
|
||||
await asyncio.wait_for(task, timeout=_READER_REAP_TIMEOUT_S)
|
||||
except Exception as error:
|
||||
# Expected: the reader's stream-close error, or the reap timeout. Logged, not
|
||||
# forwarded — the current reader is the only one that surfaces errors to the consumer.
|
||||
logger.debug("error=<%s> | superseded reader reaped", error)
|
||||
|
||||
@property
|
||||
def _accumulated_input_tokens(self) -> int:
|
||||
return self._baseline_input_tokens + self._current_input_tokens
|
||||
|
||||
@property
|
||||
def _accumulated_output_tokens(self) -> int:
|
||||
return self._baseline_output_tokens + self._current_output_tokens
|
||||
|
||||
@property
|
||||
def _accumulated_total_tokens(self) -> int:
|
||||
return self._baseline_total_tokens + self._current_total_tokens
|
||||
|
||||
@property
|
||||
def _accumulated_cache_read_tokens(self) -> int:
|
||||
return self._baseline_cache_read_tokens + self._current_cache_read_tokens
|
||||
|
||||
def _reset_token_tracking(self) -> None:
|
||||
"""Reset per-connection and baseline token tracking at session start."""
|
||||
self._current_input_tokens = 0
|
||||
self._current_output_tokens = 0
|
||||
self._current_total_tokens = 0
|
||||
self._current_cache_read_tokens = 0
|
||||
self._baseline_input_tokens = 0
|
||||
self._baseline_output_tokens = 0
|
||||
self._baseline_total_tokens = 0
|
||||
self._baseline_cache_read_tokens = 0
|
||||
|
||||
def _fold_token_baseline(self) -> None:
|
||||
"""Fold the current connection's token totals into the baseline before reconnect."""
|
||||
self._baseline_input_tokens += self._current_input_tokens
|
||||
self._baseline_output_tokens += self._current_output_tokens
|
||||
self._baseline_total_tokens += self._current_total_tokens
|
||||
self._baseline_cache_read_tokens += self._current_cache_read_tokens
|
||||
self._current_input_tokens = 0
|
||||
self._current_output_tokens = 0
|
||||
self._current_total_tokens = 0
|
||||
self._current_cache_read_tokens = 0
|
||||
|
||||
def _record_usage(self, event: BidiUsageEvent) -> None:
|
||||
"""Update the current connection's token counts from a usage event.
|
||||
|
||||
Cumulative providers report a running total (replace); delta providers report
|
||||
per-response counts (add).
|
||||
"""
|
||||
cache_read = event.cache_read_input_tokens or 0
|
||||
|
||||
if getattr(self._agent.model, "usage_is_cumulative", False):
|
||||
self._current_input_tokens = event.input_tokens
|
||||
self._current_output_tokens = event.output_tokens
|
||||
self._current_total_tokens = event.total_tokens
|
||||
self._current_cache_read_tokens = cache_read
|
||||
else:
|
||||
self._current_input_tokens += event.input_tokens
|
||||
self._current_output_tokens += event.output_tokens
|
||||
self._current_total_tokens += event.total_tokens
|
||||
self._current_cache_read_tokens += cache_read
|
||||
|
||||
async def _run_model(self, generation: int) -> None:
|
||||
"""Task for running the model.
|
||||
|
||||
Events are streamed through the event queue.
|
||||
Events are streamed through the event queue. Once superseded by a reconnect
|
||||
(``generation`` no longer current), the stream-close error and any further handling
|
||||
are dropped, so a closed old connection cannot mutate the new connection's state.
|
||||
"""
|
||||
logger.debug("model task starting")
|
||||
|
||||
@@ -287,7 +628,13 @@ class _BidiAgentLoop:
|
||||
|
||||
try:
|
||||
async for event in self._agent.model.receive():
|
||||
if generation != self._generation:
|
||||
return
|
||||
await self._event_queue.put(event)
|
||||
# The put can suspend on the full queue across a reconnect; re-check so a stale
|
||||
# event from the closed connection is not applied to the new connection's state.
|
||||
if generation != self._generation:
|
||||
return
|
||||
|
||||
if isinstance(event, BidiResponseStartEvent):
|
||||
if response_span:
|
||||
@@ -302,6 +649,9 @@ class _BidiAgentLoop:
|
||||
)
|
||||
response_start_time = time.perf_counter()
|
||||
time_to_first_audio_ms = None
|
||||
self._response_active = True
|
||||
self._awaiting_response = False
|
||||
self._update_turn_state()
|
||||
|
||||
elif isinstance(event, BidiAudioStreamEvent):
|
||||
if response_start_time is not None and time_to_first_audio_ms is None:
|
||||
@@ -316,15 +666,21 @@ class _BidiAgentLoop:
|
||||
time_to_first_audio_ms=time_to_first_audio_ms,
|
||||
)
|
||||
response_span = None
|
||||
self._response_active = False
|
||||
self._update_turn_state()
|
||||
|
||||
elif isinstance(event, BidiTranscriptStreamEvent):
|
||||
if event["is_final"]:
|
||||
message: Message = {"role": event["role"], "content": [{"text": event["text"]}]}
|
||||
await self._agent._append_messages(message)
|
||||
if event["role"] == "user":
|
||||
# A finished user turn owes a response; hold reconnect until it lands.
|
||||
self._awaiting_response = True
|
||||
self._update_turn_state()
|
||||
|
||||
elif isinstance(event, ToolUseStreamEvent):
|
||||
tool_use = event["current_tool_use"]
|
||||
self._task_pool.create(self._run_tool(tool_use))
|
||||
self._task_pool.create(self._run_tool(tool_use, generation))
|
||||
|
||||
elif isinstance(event, BidiInterruptionEvent):
|
||||
if self._session_span:
|
||||
@@ -337,16 +693,19 @@ class _BidiAgentLoop:
|
||||
interrupted_response_id=event.get("interrupted_response_id"),
|
||||
)
|
||||
)
|
||||
# A barge-in ends the current response; the user's next turn owes a reply.
|
||||
self._response_active = False
|
||||
self._update_turn_state()
|
||||
|
||||
elif isinstance(event, BidiUsageEvent):
|
||||
self._accumulated_input_tokens += event.input_tokens
|
||||
self._accumulated_output_tokens += event.output_tokens
|
||||
self._accumulated_total_tokens += event.total_tokens
|
||||
self._accumulated_cache_read_tokens += event.cache_read_input_tokens or 0
|
||||
self._record_usage(event)
|
||||
|
||||
except Exception as error:
|
||||
model_error = error
|
||||
await self._event_queue.put(error)
|
||||
# Tag with this reader's generation so receive() drops it if superseded. The put can
|
||||
# suspend on a full queue across a swap, which the pre-put check alone can't fence.
|
||||
if generation == self._generation:
|
||||
await self._event_queue.put(_ReaderError(generation, error))
|
||||
finally:
|
||||
if response_span:
|
||||
stop_reason = "error" if model_error else "incomplete"
|
||||
@@ -359,11 +718,15 @@ class _BidiAgentLoop:
|
||||
)
|
||||
response_span = None
|
||||
|
||||
async def _run_tool(self, tool_use: ToolUse) -> None:
|
||||
async def _run_tool(self, tool_use: ToolUse, generation: int) -> None:
|
||||
"""Task for running tool requested by the model using the tool executor.
|
||||
|
||||
Args:
|
||||
tool_use: Tool use request from model.
|
||||
generation: Connection generation that issued the tool use. If a reconnect
|
||||
advances the generation before the tool finishes, the result is recorded
|
||||
in history but not sent, since the new connection never issued this
|
||||
tool_use_id and would reject the result.
|
||||
"""
|
||||
logger.debug("tool_name=<%s> | tool execution starting", tool_use["name"])
|
||||
|
||||
@@ -434,6 +797,18 @@ class _BidiAgentLoop:
|
||||
)
|
||||
return # Skip sending result to model
|
||||
|
||||
# Wait out any in-flight reconnect (send() gates on the swap), then re-check: a tool
|
||||
# that finished across a swap must not send its result to the new connection, which
|
||||
# never issued this tool_use_id and would reject it. The exchange is already recorded
|
||||
# in messages above for the provider's reconnect replay.
|
||||
await self._send_gate.wait()
|
||||
if generation != self._generation:
|
||||
logger.warning(
|
||||
"tool_use_id=<%s> | tool completed across reconnect | result recorded, not sent to new connection",
|
||||
tool_use["toolUseId"],
|
||||
)
|
||||
return
|
||||
|
||||
# Send result to model
|
||||
await self.send(tool_result_event)
|
||||
|
||||
|
||||
@@ -19,7 +19,12 @@ from typing import TYPE_CHECKING, Any
|
||||
|
||||
import pyaudio
|
||||
|
||||
from ..types.events import BidiAudioInputEvent, BidiAudioStreamEvent, BidiInterruptionEvent, BidiOutputEvent
|
||||
from ..types.events import (
|
||||
BidiAudioInputEvent,
|
||||
BidiAudioStreamEvent,
|
||||
BidiInterruptionEvent,
|
||||
BidiOutputEvent,
|
||||
)
|
||||
from ..types.io import BidiInput, BidiOutput
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
||||
@@ -24,6 +24,7 @@ from ..types.events import (
|
||||
BidiInputEvent,
|
||||
BidiOutputEvent,
|
||||
)
|
||||
from ..types.model import BidiConnectionConfig
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -38,9 +39,17 @@ class BidiModel(Protocol):
|
||||
|
||||
Attributes:
|
||||
config: Configuration dictionary with provider-specific settings.
|
||||
connection_config: Declared connection limit and reconnect timing. Providers that
|
||||
support proactive reconnect populate this; an empty config means reactive-only
|
||||
behavior.
|
||||
usage_is_cumulative: Whether the provider reports cumulative connection token totals
|
||||
(True) rather than per-response deltas (False, the default when absent). Providers
|
||||
reporting deltas may omit it.
|
||||
"""
|
||||
|
||||
config: dict[str, Any]
|
||||
connection_config: BidiConnectionConfig
|
||||
usage_is_cumulative: bool
|
||||
|
||||
async def start(
|
||||
self,
|
||||
@@ -115,6 +124,26 @@ class BidiModel(Protocol):
|
||||
"""
|
||||
...
|
||||
|
||||
async def reconnect(
|
||||
self,
|
||||
system_prompt: str | None = None,
|
||||
tools: list[ToolSpec] | None = None,
|
||||
messages: Messages | None = None,
|
||||
**restart_kwargs: Any,
|
||||
) -> None:
|
||||
"""Close the current connection and establish a new one, preserving context.
|
||||
|
||||
Equivalent to ``stop()`` then ``start()``, but implemented by the provider so it
|
||||
can apply its own resume mechanism (e.g. a session handle).
|
||||
|
||||
Args:
|
||||
system_prompt: System instructions to configure model behavior.
|
||||
tools: Tool specifications that the model can invoke during the conversation.
|
||||
messages: Conversation history to replay for providers that resume via replay.
|
||||
**restart_kwargs: Provider-specific restart options (e.g. from a timeout error).
|
||||
"""
|
||||
...
|
||||
|
||||
|
||||
class BidiModelTimeoutError(Exception):
|
||||
"""Model timeout error.
|
||||
@@ -131,6 +160,6 @@ class BidiModelTimeoutError(Exception):
|
||||
message: Timeout message from model.
|
||||
**restart_config: Configure restart specific behaviors in the call to model start.
|
||||
"""
|
||||
super().__init__(self, message)
|
||||
super().__init__(message)
|
||||
|
||||
self.restart_config = restart_config
|
||||
|
||||
@@ -63,7 +63,7 @@ from ..types.events import (
|
||||
BidiTranscriptStreamEvent,
|
||||
BidiUsageEvent,
|
||||
)
|
||||
from ..types.model import AudioConfig
|
||||
from ..types.model import AudioConfig, BidiConnectionConfig
|
||||
from .model import BidiModel, BidiModelTimeoutError
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -137,6 +137,9 @@ class BidiNovaSonicModel(BidiModel):
|
||||
- inference: Model inference settings (max_tokens, temperature, top_p)
|
||||
- turn_detection: Turn detection configuration (v2 only feature)
|
||||
- endpointingSensitivity: "HIGH" | "MEDIUM" | "LOW" (optional)
|
||||
- connection: Reconnect overrides merged over the provider defaults
|
||||
(e.g. restart_after_s, auto_reconnect); see
|
||||
BidiConnectionConfig.
|
||||
client_config: AWS authentication (boto_session OR region, not both)
|
||||
**kwargs: Reserved for future parameters.
|
||||
|
||||
@@ -148,8 +151,21 @@ class BidiNovaSonicModel(BidiModel):
|
||||
# Store model ID
|
||||
self.model_id = model_id
|
||||
|
||||
# Validate turn_detection configuration
|
||||
# Nova caps a connection at ~8 min; reconnect at 7 min, leaving headroom below the cap.
|
||||
# It also reports cumulative usage totals.
|
||||
self.connection_config: BidiConnectionConfig = {"restart_after_s": 420}
|
||||
self.usage_is_cumulative = True
|
||||
|
||||
provider_config = provider_config or {}
|
||||
|
||||
# Merge any caller-supplied connection overrides over the provider defaults, so
|
||||
# reconnect behavior can be tuned (or opted out via auto_reconnect) through
|
||||
# provider_config without replacing the whole config.
|
||||
self.connection_config = cast(
|
||||
BidiConnectionConfig, {**self.connection_config, **provider_config.get("connection", {})}
|
||||
)
|
||||
|
||||
# Validate turn_detection configuration
|
||||
if "turn_detection" in provider_config and provider_config["turn_detection"]:
|
||||
if model_id == NOVA_SONIC_V1_MODEL_ID:
|
||||
raise ValueError(
|
||||
@@ -592,6 +608,26 @@ class BidiNovaSonicModel(BidiModel):
|
||||
|
||||
logger.debug("nova connection closed")
|
||||
|
||||
async def reconnect(
|
||||
self,
|
||||
system_prompt: str | None = None,
|
||||
tools: list[ToolSpec] | None = None,
|
||||
messages: Messages | None = None,
|
||||
**restart_kwargs: Any,
|
||||
) -> None:
|
||||
"""Reconnect by closing the connection and starting a new one, replaying messages.
|
||||
|
||||
Args:
|
||||
system_prompt: System instructions for the new connection.
|
||||
tools: Tool specifications for the new connection.
|
||||
messages: Conversation history to replay into the new connection.
|
||||
**restart_kwargs: Reserved for provider-specific restart options.
|
||||
"""
|
||||
logger.debug("nova reconnect starting")
|
||||
await self.stop()
|
||||
await self.start(system_prompt, tools, messages, **restart_kwargs)
|
||||
logger.debug("connection_id=<%s> | nova reconnect complete", self._connection_id)
|
||||
|
||||
def _convert_nova_event(self, nova_event: dict[str, Any]) -> BidiOutputEvent | None:
|
||||
"""Convert Nova Sonic events to TypedEvent format."""
|
||||
# Handle completion start - track completionId
|
||||
@@ -601,20 +637,12 @@ class BidiNovaSonicModel(BidiModel):
|
||||
logger.debug("completion_id=<%s> | nova completion started", self._current_completion_id)
|
||||
return None
|
||||
|
||||
# Handle completion end
|
||||
# completionEnd brackets the whole prompt/session, not a turn (its completionId is
|
||||
# constant across turns). Per-turn boundaries come from contentEnd stopReason below,
|
||||
# so only clear completion tracking here.
|
||||
if "completionEnd" in nova_event:
|
||||
completion_data = nova_event["completionEnd"]
|
||||
completion_id = completion_data.get("completionId", self._current_completion_id)
|
||||
stop_reason = completion_data.get("stopReason", "END_TURN")
|
||||
|
||||
event = BidiResponseCompleteEvent(
|
||||
response_id=completion_id or str(uuid.uuid4()), # Fallback to UUID if missing
|
||||
stop_reason="interrupted" if stop_reason == "INTERRUPTED" else "complete",
|
||||
)
|
||||
|
||||
# Clear completion tracking
|
||||
self._current_completion_id = None
|
||||
return event
|
||||
return None
|
||||
|
||||
# Handle audio output
|
||||
if "audioOutput" in nova_event:
|
||||
@@ -693,7 +721,19 @@ class BidiNovaSonicModel(BidiModel):
|
||||
)
|
||||
|
||||
if "contentEnd" in nova_event:
|
||||
content_end = nova_event["contentEnd"]
|
||||
stop_reason = content_end.get("stopReason")
|
||||
# Nova ends a turn after its FINAL assistant text block (which follows the audio).
|
||||
# Both that text block and the preceding audio block carry END_TURN, so gate on
|
||||
# the FINAL text to emit exactly one per-turn complete, after that text is in
|
||||
# history. INTERRUPTED (barge-in) ends the turn regardless of block.
|
||||
is_final_text = content_end.get("type") == "TEXT" and self._generation_stage == "FINAL"
|
||||
self._generation_stage = None
|
||||
if stop_reason == "INTERRUPTED" or (stop_reason == "END_TURN" and is_final_text):
|
||||
return BidiResponseCompleteEvent(
|
||||
response_id=self._current_completion_id or str(uuid.uuid4()),
|
||||
stop_reason="interrupted" if stop_reason == "INTERRUPTED" else "complete",
|
||||
)
|
||||
|
||||
# Ignore all other events
|
||||
return None
|
||||
|
||||
@@ -7,6 +7,7 @@ from .events import (
|
||||
BidiConnectionCloseEvent,
|
||||
BidiConnectionRestartEvent,
|
||||
BidiConnectionStartEvent,
|
||||
BidiConnectionWarningEvent,
|
||||
BidiErrorEvent,
|
||||
BidiImageInputEvent,
|
||||
BidiInputEvent,
|
||||
@@ -20,6 +21,7 @@ from .events import (
|
||||
ModalityUsage,
|
||||
)
|
||||
from .io import BidiInput, BidiOutput
|
||||
from .model import BidiConnectionConfig
|
||||
|
||||
__all__ = [
|
||||
"BidiInput",
|
||||
@@ -33,6 +35,7 @@ __all__ = [
|
||||
# Output Events
|
||||
"BidiConnectionStartEvent",
|
||||
"BidiConnectionRestartEvent",
|
||||
"BidiConnectionWarningEvent",
|
||||
"BidiConnectionCloseEvent",
|
||||
"BidiResponseStartEvent",
|
||||
"BidiResponseCompleteEvent",
|
||||
@@ -43,4 +46,6 @@ __all__ = [
|
||||
"ModalityUsage",
|
||||
"BidiErrorEvent",
|
||||
"BidiOutputEvent",
|
||||
# Reconnect configuration
|
||||
"BidiConnectionConfig",
|
||||
]
|
||||
|
||||
@@ -246,25 +246,74 @@ class BidiConnectionStartEvent(TypedEvent):
|
||||
|
||||
|
||||
class BidiConnectionRestartEvent(TypedEvent):
|
||||
"""Agent is restarting the model connection after timeout."""
|
||||
"""Agent is restarting the model connection.
|
||||
|
||||
def __init__(self, timeout_error: "BidiModelTimeoutError"):
|
||||
"""Initialize.
|
||||
Emitted on both reconnect paths: reactively after the model reports a timeout, and
|
||||
proactively when the reconnect timer fires ahead of the provider's limit.
|
||||
|
||||
Args:
|
||||
timeout_error: Timeout error reported by the model.
|
||||
"""
|
||||
Parameters:
|
||||
reason: What triggered the restart ("timeout" reactively, "scheduled" proactively).
|
||||
timeout_error: The model's timeout error on the reactive path; None when scheduled.
|
||||
turn_interrupted: True if the restart cut an in-progress or owed turn (the alignment
|
||||
wait could not complete it before the deadline, or a timeout struck mid-turn). The
|
||||
provider replays history as context, so that turn will not be answered on its own —
|
||||
an app can re-prompt or notify the user when this is set.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
reason: Literal["timeout", "scheduled"],
|
||||
timeout_error: "BidiModelTimeoutError | None" = None,
|
||||
turn_interrupted: bool = False,
|
||||
):
|
||||
"""Initialize connection restart event."""
|
||||
super().__init__(
|
||||
{
|
||||
"type": "bidi_connection_restart",
|
||||
"reason": reason,
|
||||
"timeout_error": timeout_error,
|
||||
"turn_interrupted": turn_interrupted,
|
||||
}
|
||||
)
|
||||
|
||||
@property
|
||||
def timeout_error(self) -> "BidiModelTimeoutError":
|
||||
"""Model timeout error."""
|
||||
return cast("BidiModelTimeoutError", self["timeout_error"])
|
||||
def reason(self) -> str:
|
||||
"""What triggered the restart ("timeout" or "scheduled")."""
|
||||
return cast(str, self["reason"])
|
||||
|
||||
@property
|
||||
def timeout_error(self) -> "BidiModelTimeoutError | None":
|
||||
"""Model timeout error on the reactive path; None when scheduled."""
|
||||
return cast("BidiModelTimeoutError | None", self["timeout_error"])
|
||||
|
||||
@property
|
||||
def turn_interrupted(self) -> bool:
|
||||
"""True if the restart cut an in-progress or owed turn that will not be answered."""
|
||||
return cast(bool, self["turn_interrupted"])
|
||||
|
||||
|
||||
class BidiConnectionWarningEvent(TypedEvent):
|
||||
"""Agent is approaching a proactive reconnect.
|
||||
|
||||
Emitted by the proactive reconnect timer before a reconnect; informational only.
|
||||
|
||||
Parameters:
|
||||
time_left_s: Approximate seconds until the scheduled reconnect.
|
||||
"""
|
||||
|
||||
def __init__(self, time_left_s: float):
|
||||
"""Initialize connection warning event."""
|
||||
super().__init__(
|
||||
{
|
||||
"type": "bidi_connection_warning",
|
||||
"time_left_s": time_left_s,
|
||||
}
|
||||
)
|
||||
|
||||
@property
|
||||
def time_left_s(self) -> float:
|
||||
"""Approximate seconds until the scheduled reconnect."""
|
||||
return cast(float, self["time_left_s"])
|
||||
|
||||
|
||||
class BidiResponseStartEvent(TypedEvent):
|
||||
@@ -632,6 +681,7 @@ BidiInputEvent = BidiTextInputEvent | BidiAudioInputEvent | BidiImageInputEvent
|
||||
BidiOutputEvent = (
|
||||
BidiConnectionStartEvent
|
||||
| BidiConnectionRestartEvent
|
||||
| BidiConnectionWarningEvent
|
||||
| BidiResponseStartEvent
|
||||
| BidiAudioStreamEvent
|
||||
| BidiTranscriptStreamEvent
|
||||
|
||||
@@ -34,3 +34,25 @@ class AudioConfig(TypedDict, total=False):
|
||||
channels: AudioChannel
|
||||
format: AudioFormat
|
||||
voice: str
|
||||
|
||||
|
||||
class BidiConnectionConfig(TypedDict, total=False):
|
||||
"""Declared reconnect timing for a bidirectional model.
|
||||
|
||||
Providers declare this so the agent loop can reconnect proactively, before the provider
|
||||
terminates the connection on its own limit. A provider that declares nothing (empty config)
|
||||
keeps reactive-only behavior: no proactive timer, reconnect only after the provider reports
|
||||
a timeout.
|
||||
|
||||
All fields are optional. The proactive timer arms only when ``restart_after_s`` is declared.
|
||||
|
||||
Attributes:
|
||||
restart_after_s: Seconds after a connection is established at which to proactively
|
||||
reconnect. Set it at least ~10s below the provider's own connection limit: the
|
||||
reconnect may wait briefly for the current turn to finish (aligning the swap to a
|
||||
turn boundary), and that wait plus the swap must complete before the provider's limit.
|
||||
auto_reconnect: Whether the loop reconnects automatically (default True).
|
||||
"""
|
||||
|
||||
restart_after_s: int
|
||||
auto_reconnect: bool
|
||||
|
||||
@@ -209,22 +209,28 @@ class BidiInterruptionEvent(BidiHookEvent):
|
||||
|
||||
@dataclass
|
||||
class BidiBeforeConnectionRestartEvent(BidiHookEvent):
|
||||
"""Event emitted before agent attempts to restart model connection after timeout.
|
||||
"""Event emitted before the agent restarts the model connection.
|
||||
|
||||
A restart is triggered either reactively, after the model reports a timeout, or
|
||||
proactively, when the reconnect timer fires ahead of the provider's limit.
|
||||
|
||||
Attributes:
|
||||
timeout_error: Timeout error reported by the model.
|
||||
reason: What triggered the restart ("timeout" reactively, "scheduled" proactively).
|
||||
timeout_error: The model's timeout error on the reactive path; None when scheduled.
|
||||
"""
|
||||
|
||||
timeout_error: "BidiModelTimeoutError"
|
||||
reason: Literal["timeout", "scheduled"]
|
||||
timeout_error: "BidiModelTimeoutError | None" = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class BidiAfterConnectionRestartEvent(BidiHookEvent):
|
||||
"""Event emitted after agent attempts to restart model connection after timeout.
|
||||
"""Event emitted after the agent attempts to restart the model connection.
|
||||
|
||||
Attribtues:
|
||||
exception: Populated if exception was raised during connection restart.
|
||||
None value means the restart was successful.
|
||||
Attributes:
|
||||
reason: What triggered the restart ("timeout" reactively, "scheduled" proactively).
|
||||
exception: Populated if an exception was raised during the restart. None means success.
|
||||
"""
|
||||
|
||||
reason: Literal["timeout", "scheduled"]
|
||||
exception: Exception | None = None
|
||||
|
||||
@@ -23,6 +23,8 @@ class MockBidiModel:
|
||||
|
||||
def __init__(self, config=None, model_id="mock-model"):
|
||||
self.config = config or {"audio": {"input_rate": 16000, "output_rate": 24000, "channels": 1}}
|
||||
self.connection_config = {}
|
||||
self.usage_is_cumulative = False
|
||||
self.model_id = model_id
|
||||
self._connection_id = None
|
||||
self._started = False
|
||||
@@ -39,6 +41,10 @@ class MockBidiModel:
|
||||
self._started = False
|
||||
self._connection_id = None
|
||||
|
||||
async def reconnect(self, system_prompt=None, tools=None, messages=None, **restart_kwargs):
|
||||
await self.stop()
|
||||
await self.start(system_prompt, tools, messages, **restart_kwargs)
|
||||
|
||||
async def send(self, content):
|
||||
if not self._started:
|
||||
raise RuntimeError("model not started | call start before sending")
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import asyncio
|
||||
import unittest.mock
|
||||
import warnings
|
||||
|
||||
@@ -6,12 +7,16 @@ import pytest_asyncio
|
||||
|
||||
from strands import tool
|
||||
from strands.experimental.bidi import BidiAgent
|
||||
from strands.experimental.bidi.agent.loop import _ReaderError
|
||||
from strands.experimental.bidi.models import BidiModel, BidiModelTimeoutError
|
||||
from strands.experimental.bidi.types.events import (
|
||||
BidiConnectionCloseEvent,
|
||||
BidiConnectionRestartEvent,
|
||||
BidiConnectionWarningEvent,
|
||||
BidiTextInputEvent,
|
||||
BidiUsageEvent,
|
||||
)
|
||||
from strands.experimental.hooks.events import BidiBeforeConnectionRestartEvent
|
||||
from strands.types._events import ToolResultEvent, ToolResultMessageEvent, ToolUseStreamEvent
|
||||
|
||||
|
||||
@@ -50,14 +55,16 @@ async def test_bidi_agent_loop_receive_restart_connection(loop, agent, agenerato
|
||||
break
|
||||
|
||||
exp_events = [
|
||||
BidiConnectionRestartEvent(timeout_error),
|
||||
BidiConnectionRestartEvent(reason="timeout", timeout_error=timeout_error),
|
||||
text_event,
|
||||
]
|
||||
assert tru_events == exp_events
|
||||
|
||||
agent.model.stop.assert_called_once()
|
||||
assert agent.model.start.call_count == 2
|
||||
agent.model.start.assert_called_with(
|
||||
# The reactive path reconnects through the single reconnect() method. start() is
|
||||
# called once (at loop.start()); the restart goes through reconnect() with the
|
||||
# timeout's restart_config forwarded.
|
||||
assert agent.model.start.call_count == 1
|
||||
agent.model.reconnect.assert_called_once_with(
|
||||
agent.system_prompt,
|
||||
agent.tool_registry.get_all_tool_specs(),
|
||||
agent.messages,
|
||||
@@ -65,6 +72,675 @@ async def test_bidi_agent_loop_receive_restart_connection(loop, agent, agenerato
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bidi_agent_loop_auto_reconnect_default_on(loop, agent, agenerator):
|
||||
"""Auto reconnect is the default: a timeout triggers reconnect without any opt-in."""
|
||||
# AsyncMock(spec=BidiModel) generates connection_config as a Mock; force the realistic
|
||||
# "provider declared nothing" case so the loop falls back to its default (reconnect).
|
||||
agent.model.connection_config = {}
|
||||
timeout_error = BidiModelTimeoutError("test timeout")
|
||||
text_event = BidiTextInputEvent(text="after restart")
|
||||
agent.model.receive = unittest.mock.Mock(side_effect=[timeout_error, agenerator([text_event])])
|
||||
|
||||
await loop.start()
|
||||
|
||||
received = []
|
||||
async for event in loop.receive():
|
||||
received.append(event)
|
||||
if len(received) >= 2:
|
||||
break
|
||||
|
||||
agent.model.reconnect.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bidi_agent_loop_auto_reconnect_opt_out_surfaces_timeout(loop, agent, agenerator):
|
||||
"""A provider opting out with auto_reconnect=False surfaces the timeout instead of reconnecting."""
|
||||
agent.model.connection_config = {"auto_reconnect": False}
|
||||
timeout_error = BidiModelTimeoutError("test timeout")
|
||||
agent.model.receive = unittest.mock.Mock(side_effect=[timeout_error, agenerator([])])
|
||||
|
||||
await loop.start()
|
||||
|
||||
with pytest.raises(BidiModelTimeoutError):
|
||||
async for _ in loop.receive():
|
||||
pass
|
||||
|
||||
agent.model.reconnect.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bidi_agent_loop_proactive_reconnect_before_deadline(loop, agent, agenerator):
|
||||
"""A declared limit arms the timer, which emits a warning and reconnects proactively."""
|
||||
agent.model.connection_config = {"restart_after_s": 5}
|
||||
agent.model.receive = unittest.mock.Mock(return_value=agenerator([]))
|
||||
|
||||
# Drive timing without wall time: the first cycle's sleeps return immediately; the re-armed
|
||||
# cycle after the swap parks, so exactly one proactive reconnect fires.
|
||||
sleep_count = 0
|
||||
|
||||
async def fake_sleep(_seconds):
|
||||
nonlocal sleep_count
|
||||
sleep_count += 1
|
||||
if sleep_count > 2:
|
||||
await asyncio.Event().wait()
|
||||
await asyncio.sleep(0)
|
||||
|
||||
loop._reconnect_timer._sleep = fake_sleep
|
||||
|
||||
await loop.start()
|
||||
|
||||
# The proactive timer enqueues a warning, then reconnects, then enqueues the scheduled event.
|
||||
warning = await loop._lifecycle_queue.get()
|
||||
assert isinstance(warning, BidiConnectionWarningEvent)
|
||||
|
||||
restart = await loop._lifecycle_queue.get()
|
||||
assert isinstance(restart, BidiConnectionRestartEvent)
|
||||
assert restart.reason == "scheduled"
|
||||
assert restart.turn_interrupted is False # swapped at an idle boundary, no turn cut
|
||||
|
||||
agent.model.reconnect.assert_called()
|
||||
|
||||
await loop.stop()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_proactive_timer_when_restart_after_not_positive(loop, agent, agenerator):
|
||||
"""A non-positive restart_after_s must not arm a zero-deadline hot reconnect loop."""
|
||||
agent.model.connection_config = {"restart_after_s": 0}
|
||||
agent.model.receive = unittest.mock.Mock(return_value=agenerator([]))
|
||||
|
||||
await loop.start()
|
||||
|
||||
assert loop._reconnect_timer._task is None # proactive disabled; reactive path remains
|
||||
|
||||
await loop.stop()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bidi_agent_loop_no_timer_without_declared_limit(loop, agent, agenerator):
|
||||
"""A provider that declares no limit arms no proactive timer; reconnect stays reactive-only."""
|
||||
agent.model.connection_config = {}
|
||||
agent.model.receive = unittest.mock.Mock(return_value=agenerator([]))
|
||||
|
||||
await loop.start()
|
||||
|
||||
assert loop._reconnect_timer._task is None
|
||||
|
||||
await loop.stop()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bidi_agent_loop_no_timer_when_auto_reconnect_disabled(loop, agent, agenerator):
|
||||
"""auto_reconnect=False is the only opt-out: no proactive timer arms."""
|
||||
agent.model.connection_config = {"restart_after_s": 420, "auto_reconnect": False}
|
||||
agent.model.receive = unittest.mock.Mock(return_value=agenerator([]))
|
||||
|
||||
await loop.start()
|
||||
|
||||
assert loop._reconnect_timer._task is None
|
||||
|
||||
await loop.stop()
|
||||
|
||||
|
||||
class _NoReconnectModel(BidiModel):
|
||||
"""A provider that inherits the protocol's no-op reconnect(), like Gemini/OpenAI today."""
|
||||
|
||||
def __init__(self):
|
||||
self.config = {}
|
||||
self.connection_config = {}
|
||||
self.started: list = []
|
||||
self.stopped = 0
|
||||
|
||||
async def start(self, system_prompt=None, tools=None, messages=None, **kwargs):
|
||||
self.started.append(system_prompt)
|
||||
|
||||
async def stop(self):
|
||||
self.stopped += 1
|
||||
|
||||
def receive(self): ...
|
||||
|
||||
async def send(self, content): ...
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_reconnect_falls_back_to_stop_start_when_provider_lacks_reconnect():
|
||||
"""A provider that has not implemented reconnect() is reconnected via stop() + start()."""
|
||||
model = _NoReconnectModel()
|
||||
agent = BidiAgent(model=model, system_prompt="hi")
|
||||
|
||||
await agent._loop._reconnect_model({})
|
||||
|
||||
assert model.stopped == 1
|
||||
assert model.started == ["hi"] # start() called once with the agent's system prompt
|
||||
|
||||
|
||||
class _StreamModel(BidiModel):
|
||||
"""Reader blocks on a live 'stream' and raises when stop() closes it, like Nova/awscrt.
|
||||
|
||||
The reader is terminated by the stream closing (an OSError), not by a force-cancel, so
|
||||
a reconnect must fence that error instead of forwarding it to the consumer.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
self.config = {}
|
||||
self.connection_config = {}
|
||||
self.reconnect_calls = 0
|
||||
self._closed = asyncio.Event()
|
||||
self._inbox: asyncio.Queue = asyncio.Queue()
|
||||
|
||||
async def start(self, system_prompt=None, tools=None, messages=None, **kwargs):
|
||||
self._closed = asyncio.Event()
|
||||
self._inbox = asyncio.Queue()
|
||||
|
||||
async def stop(self):
|
||||
self._closed.set()
|
||||
|
||||
async def reconnect(self, system_prompt=None, tools=None, messages=None, **kwargs):
|
||||
self.reconnect_calls += 1
|
||||
await self.stop()
|
||||
await self.start(system_prompt, tools, messages, **kwargs)
|
||||
|
||||
async def send(self, content):
|
||||
return None
|
||||
|
||||
async def emit(self, event):
|
||||
await self._inbox.put(event)
|
||||
|
||||
async def receive(self):
|
||||
closed, inbox = self._closed, self._inbox
|
||||
while True:
|
||||
getter = asyncio.ensure_future(inbox.get())
|
||||
waiter = asyncio.ensure_future(closed.wait())
|
||||
done, pending = await asyncio.wait({getter, waiter}, return_when=asyncio.FIRST_COMPLETED)
|
||||
for task in pending:
|
||||
task.cancel()
|
||||
if waiter in done:
|
||||
getter.cancel()
|
||||
raise OSError("stream closed")
|
||||
yield getter.result()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_reconnect_fences_superseded_reader_stream_close_error():
|
||||
"""Reconnect closes the old stream (reader raises); that error must not leak to the consumer."""
|
||||
model = _StreamModel()
|
||||
agent = BidiAgent(model=model, system_prompt="hi")
|
||||
loop = agent._loop
|
||||
|
||||
await loop.start()
|
||||
|
||||
first = BidiTextInputEvent(text="first")
|
||||
await model.emit(first)
|
||||
assert await loop._event_queue.get() is first
|
||||
|
||||
# Proactive-style reconnect: reconnect() -> stop() closes the old stream, so the old
|
||||
# reader raises OSError. It is superseded, so that error must be dropped, not queued.
|
||||
await loop._restart_connection(None, loop._generation)
|
||||
assert model.reconnect_calls == 1
|
||||
|
||||
second = BidiTextInputEvent(text="second")
|
||||
await model.emit(second)
|
||||
# The new connection's event arrives; a leaked OSError would have surfaced here instead.
|
||||
assert await loop._event_queue.get() is second
|
||||
|
||||
await loop.stop()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stale_reader_event_does_not_corrupt_state_across_reconnect():
|
||||
"""A reader suspended in the full-queue put across a swap must not record its stale event.
|
||||
|
||||
Guards the generation re-check after the queue put: without it, a cumulative-usage provider
|
||||
double-counts the old connection's running total onto the already-folded baseline.
|
||||
"""
|
||||
model = _StreamModel()
|
||||
model.usage_is_cumulative = True # like Nova: usage events report a running total
|
||||
agent = BidiAgent(model=model, system_prompt="hi")
|
||||
loop = agent._loop
|
||||
|
||||
await loop.start()
|
||||
loop._reconnect_timer.cancel()
|
||||
|
||||
await model.emit(BidiUsageEvent(input_tokens=60, output_tokens=40, total_tokens=100))
|
||||
await model.emit(BidiUsageEvent(input_tokens=90, output_tokens=60, total_tokens=150))
|
||||
for _ in range(30):
|
||||
await asyncio.sleep(0)
|
||||
# usage1 recorded; the reader is now suspended inside put(usage2) on the full queue.
|
||||
assert loop._accumulated_total_tokens == 100
|
||||
|
||||
swap = asyncio.create_task(loop._restart_connection(None, loop._generation))
|
||||
for _ in range(30):
|
||||
await asyncio.sleep(0)
|
||||
await loop._event_queue.get() # drain, unblocking the old reader's put(usage2)
|
||||
await swap
|
||||
for _ in range(30):
|
||||
await asyncio.sleep(0)
|
||||
|
||||
# The stale usage2 must not be recorded onto the new connection (no cumulative double count).
|
||||
assert loop._accumulated_total_tokens == 100
|
||||
|
||||
await loop.stop()
|
||||
|
||||
|
||||
async def _feed_after_drain(loop, event):
|
||||
"""Put ``event`` once the queue has drained (so a maxsize-1 put does not block)."""
|
||||
while loop._event_queue.qsize() > 0:
|
||||
await asyncio.sleep(0)
|
||||
await loop._event_queue.put(event)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stale_reader_error_is_dropped_not_raised(loop, agent, agenerator):
|
||||
"""A generic error from a superseded reader must be dropped, not surfaced into the new connection.
|
||||
|
||||
Without the generation tag, a stale error re-raised by receive() kills the healthy, just-swapped
|
||||
session.
|
||||
"""
|
||||
agent.model.receive = unittest.mock.Mock(return_value=agenerator([]))
|
||||
await loop.start()
|
||||
loop._reconnect_timer.cancel()
|
||||
|
||||
# An error raised on a superseded (older) generation.
|
||||
await loop._event_queue.put(_ReaderError(loop._generation - 1, OSError("stale connection error")))
|
||||
|
||||
sentinel = BidiTextInputEvent(text="after stale error")
|
||||
feed = asyncio.create_task(_feed_after_drain(loop, sentinel))
|
||||
# receive() must drop the stale error and go on to the next event, not raise it.
|
||||
result = await asyncio.wait_for(loop.receive().__anext__(), timeout=2.0)
|
||||
assert result is sentinel
|
||||
await feed
|
||||
|
||||
await loop.stop()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_current_reader_error_is_surfaced(loop, agent, agenerator):
|
||||
"""A genuine error from the current reader must still surface to the consumer."""
|
||||
agent.model.receive = unittest.mock.Mock(return_value=agenerator([]))
|
||||
await loop.start()
|
||||
loop._reconnect_timer.cancel()
|
||||
|
||||
await loop._event_queue.put(_ReaderError(loop._generation, OSError("live connection error")))
|
||||
|
||||
with pytest.raises(OSError, match="live connection error"):
|
||||
await asyncio.wait_for(loop.receive().__anext__(), timeout=2.0)
|
||||
|
||||
await loop.stop()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stale_reactive_timeout_dropped_after_proactive_swap(loop, agent, agenerator):
|
||||
"""A timeout raised on an old generation, dequeued after a proactive swap, must not reconnect again."""
|
||||
agent.model.connection_config = {}
|
||||
agent.model.receive = unittest.mock.Mock(return_value=agenerator([]))
|
||||
await loop.start()
|
||||
loop._reconnect_timer.cancel()
|
||||
|
||||
stale_generation = loop._generation
|
||||
await loop._restart_connection(None, loop._generation) # a proactive swap advances the generation
|
||||
reconnects = agent.model.reconnect.call_count
|
||||
|
||||
# A timeout tagged with the pre-swap generation is now stale; receive() must drop it.
|
||||
await loop._event_queue.put(_ReaderError(stale_generation, BidiModelTimeoutError("stale timeout")))
|
||||
|
||||
sentinel = BidiTextInputEvent(text="after stale timeout")
|
||||
feed = asyncio.create_task(_feed_after_drain(loop, sentinel))
|
||||
result = await asyncio.wait_for(loop.receive().__anext__(), timeout=2.0)
|
||||
assert result is sentinel
|
||||
await feed
|
||||
assert agent.model.reconnect.call_count == reconnects # no second reconnect from the stale timeout
|
||||
|
||||
await loop.stop()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_lifecycle_events_take_priority_over_data_in_receive(loop, agent, agenerator):
|
||||
"""A queued lifecycle event is delivered before an older data event (in-order, not dropped)."""
|
||||
agent.model.receive = unittest.mock.Mock(return_value=agenerator([]))
|
||||
await loop.start()
|
||||
loop._reconnect_timer.cancel()
|
||||
|
||||
data = BidiTextInputEvent(text="new-connection output")
|
||||
warning = BidiConnectionWarningEvent(time_left_s=10.0)
|
||||
await loop._event_queue.put(data) # data queued first...
|
||||
loop._lifecycle_queue.put_nowait(warning) # ...lifecycle second, but wins
|
||||
|
||||
first = await asyncio.wait_for(loop.receive().__anext__(), timeout=2.0)
|
||||
assert first is warning
|
||||
|
||||
await loop.stop()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_lifecycle_event_delivered_while_consumer_idle(loop, agent, agenerator):
|
||||
"""A lifecycle event emitted while both queues are empty still wakes receive()."""
|
||||
agent.model.receive = unittest.mock.Mock(return_value=agenerator([]))
|
||||
await loop.start()
|
||||
loop._reconnect_timer.cancel()
|
||||
|
||||
consumer = loop.receive()
|
||||
warning = BidiConnectionWarningEvent(time_left_s=10.0)
|
||||
|
||||
async def emit():
|
||||
await asyncio.sleep(0) # let the consumer reach the both-queues-empty wait
|
||||
loop._lifecycle_queue.put_nowait(warning)
|
||||
|
||||
asyncio.create_task(emit())
|
||||
first = await asyncio.wait_for(consumer.__anext__(), timeout=2.0)
|
||||
assert first is warning
|
||||
|
||||
await loop.stop()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tool_result_not_sent_when_completed_during_reconnect(agenerator):
|
||||
"""A tool completing inside the reconnect window must not deliver its result to the new connection.
|
||||
|
||||
The gen re-check after the send gate reopens guards this; the window is opened by a
|
||||
suspending before-restart hook (a public extension point).
|
||||
"""
|
||||
order = []
|
||||
release_tool = asyncio.Event()
|
||||
|
||||
@tool
|
||||
async def slow_tool():
|
||||
await release_tool.wait()
|
||||
return "result"
|
||||
|
||||
model = unittest.mock.AsyncMock(spec=BidiModel)
|
||||
model.connection_config = {}
|
||||
model.send.side_effect = lambda event: order.append("send")
|
||||
model.reconnect.side_effect = lambda *a, **k: order.append("reconnect")
|
||||
model.receive = unittest.mock.Mock(return_value=agenerator([]))
|
||||
|
||||
agent = BidiAgent(model=model, tools=[slow_tool], system_prompt="hi")
|
||||
loop = agent._loop
|
||||
await loop.start()
|
||||
loop._reconnect_timer.cancel()
|
||||
|
||||
async def drain():
|
||||
while True:
|
||||
await loop._event_queue.get()
|
||||
|
||||
drain_task = asyncio.create_task(drain())
|
||||
tool_use = {"toolUseId": "t1", "name": "slow_tool", "input": {}}
|
||||
tool_task = asyncio.create_task(loop._run_tool(tool_use, loop._generation))
|
||||
for _ in range(10):
|
||||
await asyncio.sleep(0)
|
||||
|
||||
async def before_restart_hook(event):
|
||||
# Release the tool mid-reconnect: the gate is closed but the generation not yet bumped.
|
||||
release_tool.set()
|
||||
for _ in range(50):
|
||||
await asyncio.sleep(0)
|
||||
|
||||
agent.hooks.add_callback(BidiBeforeConnectionRestartEvent, before_restart_hook)
|
||||
|
||||
await loop._restart_connection(None, loop._generation)
|
||||
await asyncio.wait_for(tool_task, timeout=2)
|
||||
drain_task.cancel()
|
||||
|
||||
assert "send" not in order, f"stale tool result sent to new connection: {order}"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stale_reactive_restart_ignored_after_proactive_swap(agent, agenerator):
|
||||
"""A stale timeout restart (raised for an old generation) must not tear down the new connection."""
|
||||
agent.model.connection_config = {}
|
||||
agent.model.receive = unittest.mock.Mock(return_value=agenerator([]))
|
||||
|
||||
loop = agent._loop
|
||||
await loop.start()
|
||||
loop._reconnect_timer.cancel()
|
||||
|
||||
stale_generation = loop._generation
|
||||
await loop._restart_connection(None, loop._generation) # a proactive swap advances the generation
|
||||
assert loop._generation == stale_generation + 1
|
||||
reconnects = agent.model.reconnect.call_count
|
||||
|
||||
await loop._restart_connection(BidiModelTimeoutError("stale"), stale_generation)
|
||||
assert agent.model.reconnect.call_count == reconnects # stale trigger ignored
|
||||
|
||||
await loop.stop()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_deadline_callback_does_not_reconnect_after_stop(agent, agenerator):
|
||||
"""A proactive deadline callback in flight during stop() must not reconnect the model."""
|
||||
agent.model.connection_config = {"restart_after_s": 415}
|
||||
agent.model.receive = unittest.mock.Mock(return_value=agenerator([]))
|
||||
|
||||
loop = agent._loop
|
||||
await loop.start()
|
||||
loop._reconnect_timer.cancel()
|
||||
|
||||
loop._response_active = True # mid-turn: the callback waits for the boundary
|
||||
loop._update_turn_state()
|
||||
|
||||
deadline_task = asyncio.create_task(loop._on_reconnect_deadline())
|
||||
for _ in range(10):
|
||||
await asyncio.sleep(0)
|
||||
|
||||
await loop.stop() # stop() releases the boundary wait; the callback no-ops on _started
|
||||
await asyncio.wait_for(deadline_task, timeout=2)
|
||||
|
||||
agent.model.reconnect.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_send_user_text_marks_turn_awaiting_response(loop, agent, agenerator):
|
||||
"""A user text turn owes a reply, so it holds the turn boundary like a finished audio turn."""
|
||||
agent.model.receive = unittest.mock.Mock(return_value=agenerator([]))
|
||||
await loop.start()
|
||||
|
||||
await loop.send(BidiTextInputEvent(text="hello", role="user"))
|
||||
assert loop._awaiting_response is True
|
||||
assert not loop._turn_complete.is_set() # a proactive reconnect would now wait
|
||||
|
||||
await loop.stop()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_send_assistant_text_does_not_mark_awaiting_response(loop, agent, agenerator):
|
||||
"""Injected assistant context is not an owed user turn and must not hold the boundary."""
|
||||
agent.model.receive = unittest.mock.Mock(return_value=agenerator([]))
|
||||
await loop.start()
|
||||
|
||||
await loop.send(BidiTextInputEvent(text="injected context", role="assistant"))
|
||||
assert loop._awaiting_response is False
|
||||
assert loop._turn_complete.is_set()
|
||||
|
||||
await loop.stop()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_forced_swap_flags_interrupted_turn(agent, agenerator):
|
||||
"""A swap forced while a turn is owed sets turn_interrupted so the app can re-prompt."""
|
||||
agent.model.connection_config = {"restart_after_s": 415}
|
||||
agent.model.receive = unittest.mock.Mock(return_value=agenerator([]))
|
||||
loop = agent._loop
|
||||
await loop.start()
|
||||
loop._reconnect_timer.cancel()
|
||||
|
||||
loop._response_active = True # a turn is in progress and will not complete in time
|
||||
loop._update_turn_state()
|
||||
|
||||
# Force the turn-alignment wait to time out immediately (no wall-clock wait).
|
||||
with unittest.mock.patch("strands.experimental.bidi.agent.loop._TURN_ALIGN_WAIT_S", 0.0):
|
||||
await loop._on_reconnect_deadline()
|
||||
|
||||
restart = await loop._lifecycle_queue.get()
|
||||
assert isinstance(restart, BidiConnectionRestartEvent)
|
||||
assert restart.turn_interrupted is True
|
||||
|
||||
await loop.stop()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_proactive_reconnect_waits_for_turn_boundary(loop, agent, agenerator):
|
||||
"""A proactive reconnect defers until the in-progress turn completes (turn alignment)."""
|
||||
# The real timer is cancelled so the deadline is driven manually; the turn state is set
|
||||
# directly, and _await_turn_boundary waits up to _TURN_ALIGN_WAIT_S for the boundary.
|
||||
agent.model.connection_config = {"restart_after_s": 60}
|
||||
agent.model.receive = unittest.mock.Mock(return_value=agenerator([]))
|
||||
await loop.start()
|
||||
loop._reconnect_timer.cancel()
|
||||
|
||||
# Mid-response: not at a turn boundary.
|
||||
loop._response_active = True
|
||||
loop._update_turn_state()
|
||||
|
||||
deadline = asyncio.create_task(loop._on_reconnect_deadline())
|
||||
for _ in range(10):
|
||||
await asyncio.sleep(0)
|
||||
assert not agent.model.reconnect.called # held: the turn has not finished
|
||||
|
||||
# Turn completes -> boundary reached -> reconnect proceeds.
|
||||
loop._response_active = False
|
||||
loop._update_turn_state()
|
||||
await deadline
|
||||
assert agent.model.reconnect.called
|
||||
|
||||
await loop.stop()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bidi_agent_loop_restart_hook_reports_reason(loop, agent, agenerator):
|
||||
"""The reactive path reports reason='timeout' with the error; proactive reports 'scheduled' with None."""
|
||||
from strands.experimental.hooks.events import BidiBeforeConnectionRestartEvent
|
||||
|
||||
before_events = []
|
||||
agent.hooks.add_callback(
|
||||
BidiBeforeConnectionRestartEvent, lambda event: before_events.append((event.reason, event.timeout_error))
|
||||
)
|
||||
agent.model.connection_config = {}
|
||||
agent.model.receive = unittest.mock.Mock(return_value=agenerator([]))
|
||||
|
||||
await loop.start()
|
||||
|
||||
timeout_error = BidiModelTimeoutError("boom")
|
||||
await loop._restart_connection(timeout_error, loop._generation)
|
||||
await loop._restart_connection(None, loop._generation)
|
||||
|
||||
assert before_events[0] == ("timeout", timeout_error)
|
||||
assert before_events[1] == ("scheduled", None)
|
||||
|
||||
await loop.stop()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bidi_agent_loop_reconnect_is_reentrancy_guarded(loop, agent, agenerator):
|
||||
"""A second trigger arriving while a reconnect is in flight is a no-op, not a racing duplicate."""
|
||||
agent.model.connection_config = {}
|
||||
agent.model.receive = unittest.mock.Mock(return_value=agenerator([]))
|
||||
|
||||
# Block the reconnect so the first call holds the guard while the second is attempted.
|
||||
release = asyncio.Event()
|
||||
reconnect_calls = 0
|
||||
|
||||
async def blocking_reconnect(*_args, **_kwargs):
|
||||
nonlocal reconnect_calls
|
||||
reconnect_calls += 1
|
||||
await release.wait()
|
||||
|
||||
agent.model.reconnect = blocking_reconnect
|
||||
|
||||
await loop.start()
|
||||
|
||||
first = asyncio.create_task(loop._restart_connection(None, loop._generation))
|
||||
for _ in range(10):
|
||||
await asyncio.sleep(0)
|
||||
if reconnect_calls == 1:
|
||||
break
|
||||
|
||||
# First reconnect is now suspended mid-flight, still holding the guard.
|
||||
await loop._restart_connection(None, loop._generation)
|
||||
assert reconnect_calls == 1
|
||||
|
||||
release.set()
|
||||
await first
|
||||
assert reconnect_calls == 1
|
||||
|
||||
await loop.stop()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bidi_agent_loop_proactive_reconnect_completes_when_reconnect_suspends(loop, agent, agenerator):
|
||||
"""The proactive reconnect runs on the timer's task, so it must not cancel itself mid-flight.
|
||||
|
||||
Guards against the timer cancelling the very task running its deadline callback: with a
|
||||
reconnect that actually suspends, a self-cancel would abort the swap and leave the gate closed.
|
||||
"""
|
||||
agent.model.connection_config = {"restart_after_s": 5}
|
||||
agent.model.receive = unittest.mock.Mock(return_value=agenerator([]))
|
||||
|
||||
reconnect_done = False
|
||||
|
||||
async def suspending_reconnect(*_args, **_kwargs):
|
||||
nonlocal reconnect_done
|
||||
await asyncio.sleep(0) # genuine suspension after the timer fires its deadline
|
||||
reconnect_done = True
|
||||
|
||||
agent.model.reconnect = suspending_reconnect
|
||||
|
||||
# Drive timing without wall time: the first cycle fires immediately, the re-armed cycle parks.
|
||||
sleep_count = 0
|
||||
|
||||
async def fake_sleep(_seconds):
|
||||
nonlocal sleep_count
|
||||
sleep_count += 1
|
||||
if sleep_count > 2:
|
||||
await asyncio.Event().wait()
|
||||
await asyncio.sleep(0)
|
||||
|
||||
loop._reconnect_timer._sleep = fake_sleep
|
||||
|
||||
await loop.start()
|
||||
|
||||
# Drain notification events like a real consumer, so the proactive path is not blocked
|
||||
# enqueuing the warning/restart events on the size-1 queue before it reconnects.
|
||||
for _ in range(50):
|
||||
await asyncio.sleep(0)
|
||||
while not loop._event_queue.empty():
|
||||
loop._event_queue.get_nowait()
|
||||
if reconnect_done:
|
||||
break
|
||||
|
||||
assert reconnect_done
|
||||
assert loop._send_gate.is_set()
|
||||
|
||||
await loop.stop()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bidi_agent_loop_cumulative_usage_not_double_counted(loop, agent, agenerator):
|
||||
"""Cumulative providers replace running counts rather than summing successive totals."""
|
||||
from strands.experimental.bidi.types.events import BidiUsageEvent
|
||||
|
||||
agent.model.usage_is_cumulative = True
|
||||
events = [
|
||||
BidiUsageEvent(input_tokens=100, output_tokens=50, total_tokens=150),
|
||||
BidiUsageEvent(input_tokens=250, output_tokens=120, total_tokens=370),
|
||||
]
|
||||
agent.model.receive = unittest.mock.Mock(return_value=agenerator(events))
|
||||
|
||||
await loop.start()
|
||||
|
||||
received = []
|
||||
async for event in loop.receive():
|
||||
received.append(event)
|
||||
if len(received) >= 2:
|
||||
break
|
||||
|
||||
# Latest cumulative total wins (370), not the sum of the two events (520).
|
||||
assert loop._accumulated_input_tokens == 250
|
||||
assert loop._accumulated_output_tokens == 120
|
||||
assert loop._accumulated_total_tokens == 370
|
||||
|
||||
await loop.stop()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bidi_agent_loop_receive_tool_use(loop, agent, agenerator):
|
||||
tool_use = {"toolUseId": "t1", "name": "time_tool", "input": {}}
|
||||
@@ -103,6 +779,44 @@ async def test_bidi_agent_loop_receive_tool_use(loop, agent, agenerator):
|
||||
agent.model.send.assert_called_with(tool_result_event)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bidi_agent_loop_tool_result_not_sent_after_reconnect(loop, agent, agenerator):
|
||||
"""A tool completing after a reconnect records its result but does not send it.
|
||||
|
||||
The tool_use_id is scoped to the connection that issued the call; sending the result to
|
||||
the reconnected connection would be rejected by the provider (e.g. Nova
|
||||
"Not expecting a tool result") and end the session.
|
||||
"""
|
||||
tool_use = {"toolUseId": "t1", "name": "time_tool", "input": {}}
|
||||
|
||||
agent.model.receive = unittest.mock.Mock(return_value=agenerator([]))
|
||||
await loop.start()
|
||||
|
||||
# A reconnect during tool execution advances the connection generation.
|
||||
issuing_generation = loop._generation
|
||||
loop._generation += 1
|
||||
|
||||
# Drain the event queue (maxsize=1) so _run_tool's puts do not block.
|
||||
async def drain():
|
||||
while True:
|
||||
await loop._event_queue.get()
|
||||
|
||||
drain_task = asyncio.create_task(drain())
|
||||
try:
|
||||
await loop._run_tool(tool_use, issuing_generation)
|
||||
await asyncio.sleep(0)
|
||||
finally:
|
||||
drain_task.cancel()
|
||||
|
||||
# The completed exchange is recorded for the provider's reconnect replay...
|
||||
assert len(agent.messages) == 2
|
||||
assert agent.messages[0]["role"] == "assistant"
|
||||
assert agent.messages[0]["content"] == [{"toolUse": tool_use}]
|
||||
assert agent.messages[1]["content"][0]["toolResult"]["toolUseId"] == "t1"
|
||||
# ...but the stale result is not sent to the reconnected connection.
|
||||
agent.model.send.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bidi_agent_loop_request_state_initialized_for_tools(loop, agent, agenerator):
|
||||
"""Test that request_state is initialized in invocation_state before tool execution.
|
||||
|
||||
@@ -316,6 +316,8 @@ async def test_connection_restart_span(loop, agent, agenerator, otel_setup):
|
||||
spans = otel_setup.get_finished_spans()
|
||||
restart_spans = [s for s in spans if "bidi_connection_restart" in s.name]
|
||||
assert len(restart_spans) == 1
|
||||
assert restart_spans[0].attributes["gen_ai.bidi.restart_reason"] == "timeout"
|
||||
assert restart_spans[0].attributes["gen_ai.bidi.restart_error_message"] == "8 minute timeout"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -345,7 +347,7 @@ async def test_restart_failure_propagates_and_reports(loop, agent, agenerator):
|
||||
"""A failed restart surfaces to receive(), keeps the gate closed, and fires the after-restart hook."""
|
||||
timeout_error = BidiModelTimeoutError("8 minute timeout")
|
||||
agent.model.receive = unittest.mock.Mock(side_effect=[timeout_error, agenerator([])])
|
||||
agent.model.start = unittest.mock.AsyncMock(side_effect=[None, ConnectionError("restart failed")])
|
||||
agent.model.reconnect = unittest.mock.AsyncMock(side_effect=ConnectionError("restart failed"))
|
||||
|
||||
after_errors = []
|
||||
agent.hooks.add_callback(BidiAfterConnectionRestartEvent, lambda event: after_errors.append(event.exception))
|
||||
|
||||
@@ -0,0 +1,109 @@
|
||||
"""Unit tests for the proactive reconnect timer.
|
||||
|
||||
The timer is exercised with an injected fake clock so timing is deterministic and does
|
||||
not depend on wall time or a running provider.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
|
||||
import pytest
|
||||
|
||||
from strands.experimental.bidi.agent._reconnect_timer import _BidiReconnectTimer, resolve_deadline_s
|
||||
|
||||
# resolve_deadline_s
|
||||
|
||||
|
||||
def test_resolve_deadline_none_when_not_declared():
|
||||
"""No declared restart_after_s means no proactive timer; reconnect stays reactive-only."""
|
||||
assert resolve_deadline_s({}) is None
|
||||
assert resolve_deadline_s({"auto_reconnect": True}) is None
|
||||
|
||||
|
||||
def test_resolve_deadline_is_restart_after_s():
|
||||
"""Deadline is the declared restart_after_s."""
|
||||
assert resolve_deadline_s({"restart_after_s": 540}) == 540
|
||||
|
||||
|
||||
def test_resolve_deadline_none_when_not_positive():
|
||||
"""A non-positive restart_after_s declares no usable deadline; no proactive timer arms."""
|
||||
assert resolve_deadline_s({"restart_after_s": 0}) is None
|
||||
assert resolve_deadline_s({"restart_after_s": -5}) is None
|
||||
|
||||
|
||||
# _BidiReconnectTimer
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_timer_fires_warning_then_deadline():
|
||||
"""The timer fires the warning at the lead offset, then the deadline."""
|
||||
warnings, deadlines = [], []
|
||||
sleeps = []
|
||||
|
||||
async def fake_sleep(seconds):
|
||||
sleeps.append(seconds)
|
||||
|
||||
timer = _BidiReconnectTimer(
|
||||
on_warning=lambda t: _record(warnings, t),
|
||||
on_deadline=lambda: _record(deadlines, None),
|
||||
sleep=fake_sleep,
|
||||
)
|
||||
|
||||
# deadline 420, warning 30 before it => sleep 390 then 30.
|
||||
timer.arm(deadline_s=420.0, warning_lead_s=30.0)
|
||||
|
||||
await timer._task
|
||||
|
||||
assert sleeps == [390.0, 30.0]
|
||||
assert warnings == [30.0] # time_left_s at warning == warning_lead_s
|
||||
assert deadlines == [None]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_timer_cancel_is_safe_when_idle():
|
||||
"""cancel() before arming does not raise."""
|
||||
timer = _BidiReconnectTimer(on_warning=_noop_arg, on_deadline=_noop)
|
||||
timer.cancel() # should not raise
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_timer_rearm_cancels_previous():
|
||||
"""Re-arming cancels the prior cycle so only one deadline fires."""
|
||||
deadlines = []
|
||||
|
||||
started = asyncio.Event()
|
||||
|
||||
async def slow_sleep(seconds):
|
||||
started.set()
|
||||
await asyncio.sleep(3600)
|
||||
|
||||
timer = _BidiReconnectTimer(
|
||||
on_warning=_noop_arg,
|
||||
on_deadline=lambda: _record(deadlines, None),
|
||||
sleep=slow_sleep,
|
||||
)
|
||||
|
||||
timer.arm(deadline_s=420.0, warning_lead_s=30.0)
|
||||
await started.wait()
|
||||
first_task = timer._task
|
||||
|
||||
timer.arm(deadline_s=420.0, warning_lead_s=30.0)
|
||||
|
||||
await asyncio.sleep(0)
|
||||
assert first_task.cancelled() or first_task.done()
|
||||
|
||||
timer.cancel()
|
||||
|
||||
|
||||
# Helpers
|
||||
|
||||
|
||||
async def _record(sink, value):
|
||||
sink.append(value)
|
||||
|
||||
|
||||
async def _noop():
|
||||
return None
|
||||
|
||||
|
||||
async def _noop_arg(_value):
|
||||
return None
|
||||
@@ -31,6 +31,7 @@ from strands.experimental.bidi.types.events import (
|
||||
BidiAudioStreamEvent,
|
||||
BidiImageInputEvent,
|
||||
BidiInterruptionEvent,
|
||||
BidiResponseCompleteEvent,
|
||||
BidiResponseStartEvent,
|
||||
BidiTextInputEvent,
|
||||
BidiTranscriptStreamEvent,
|
||||
@@ -220,6 +221,183 @@ async def test_model_stop_alone(nova_model):
|
||||
await nova_model.stop() # Should not raise
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stop_is_idempotent(nova_model, mock_stream):
|
||||
"""Calling stop() twice on a started model does not re-close the stream or raise."""
|
||||
await nova_model.start()
|
||||
await nova_model.stop()
|
||||
assert mock_stream.close.call_count == 1
|
||||
|
||||
# Second stop must be a no-op: the stream reference is cleared on first stop, so
|
||||
# close() is not called again and no AttributeError is raised.
|
||||
await nova_model.stop()
|
||||
assert mock_stream.close.call_count == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_content_end_end_turn_emits_response_complete(nova_model):
|
||||
"""A per-turn boundary (contentEnd END_TURN) emits a response-complete event."""
|
||||
nova_model._current_completion_id = "c1"
|
||||
|
||||
# Intermediate blocks are not a turn boundary.
|
||||
assert nova_model._convert_nova_event({"contentEnd": {"type": "TEXT", "stopReason": "PARTIAL_TURN"}}) is None
|
||||
|
||||
# The audio block's END_TURN is deduped away; only the FINAL assistant text block emits
|
||||
# the per-turn complete (so it fires once, after that text is in history).
|
||||
assert nova_model._convert_nova_event({"contentEnd": {"type": "AUDIO", "stopReason": "END_TURN"}}) is None
|
||||
|
||||
nova_model._generation_stage = "FINAL"
|
||||
end = nova_model._convert_nova_event({"contentEnd": {"type": "TEXT", "stopReason": "END_TURN"}})
|
||||
assert isinstance(end, BidiResponseCompleteEvent)
|
||||
assert end.stop_reason == "complete"
|
||||
|
||||
# A barge-in ends the turn regardless of block/stage.
|
||||
interrupted = nova_model._convert_nova_event({"contentEnd": {"type": "AUDIO", "stopReason": "INTERRUPTED"}})
|
||||
assert isinstance(interrupted, BidiResponseCompleteEvent)
|
||||
assert interrupted.stop_reason == "interrupted"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_completion_end_is_not_a_turn_boundary(nova_model):
|
||||
"""completionEnd brackets the whole session, so it is not a per-turn response-complete."""
|
||||
nova_model._current_completion_id = "c1"
|
||||
result = nova_model._convert_nova_event({"completionEnd": {"stopReason": "END_TURN"}})
|
||||
assert result is None
|
||||
assert nova_model._current_completion_id is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_connection_config_declared(nova_model):
|
||||
"""Nova declares its reconnect deadline and cumulative usage semantics."""
|
||||
assert nova_model.connection_config["restart_after_s"] == 420
|
||||
assert nova_model.usage_is_cumulative is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_connection_config_overrides_merge_over_defaults(model_id, boto_session):
|
||||
"""provider_config['connection'] tunes individual fields without dropping the defaults."""
|
||||
model = BidiNovaSonicModel(
|
||||
model_id=model_id,
|
||||
client_config={"boto_session": boto_session},
|
||||
provider_config={"connection": {"auto_reconnect": False}},
|
||||
)
|
||||
|
||||
# Overridden field takes the caller's value.
|
||||
assert model.connection_config["auto_reconnect"] is False
|
||||
# Untouched default is preserved.
|
||||
assert model.connection_config["restart_after_s"] == 420
|
||||
# usage_is_cumulative is a separate provider trait, unaffected by connection overrides.
|
||||
assert model.usage_is_cumulative is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_reconnect_replays_history_through_start_path(nova_model, mock_stream):
|
||||
"""reconnect() stops the old connection and re-initializes with the same context."""
|
||||
tools = [
|
||||
{
|
||||
"name": "get_weather",
|
||||
"description": "Get weather information",
|
||||
"inputSchema": {"json": json.dumps({"type": "object", "properties": {}})},
|
||||
}
|
||||
]
|
||||
messages = [
|
||||
{"role": "user", "content": [{"text": "What's the weather?"}]},
|
||||
{"role": "assistant", "content": [{"text": "It's sunny and 72 degrees."}]},
|
||||
]
|
||||
|
||||
await nova_model.start(system_prompt="You are helpful", tools=tools, messages=messages)
|
||||
first_connection_id = nova_model._connection_id
|
||||
mock_stream.input_stream.send.reset_mock()
|
||||
|
||||
await nova_model.reconnect(system_prompt="You are helpful", tools=tools, messages=messages)
|
||||
|
||||
# Old stream was closed and a fresh connection established with a new id.
|
||||
assert mock_stream.close.called
|
||||
assert nova_model._connection_id is not None
|
||||
assert nova_model._connection_id != first_connection_id
|
||||
|
||||
# History was replayed through the same initialization path start() uses:
|
||||
# sessionStart + promptStart + system prompt (3) + 2 text messages (3 events each).
|
||||
sent_events = [call.args[0].value.bytes_.decode("utf-8") for call in mock_stream.input_stream.send.call_args_list]
|
||||
user_events = [e for e in sent_events if '"role": "USER"' in e]
|
||||
assistant_events = [e for e in sent_events if '"role": "ASSISTANT"' in e]
|
||||
assert len(user_events) >= 1
|
||||
assert len(assistant_events) >= 1
|
||||
|
||||
await nova_model.stop()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_reconnect_twice_does_not_raise(nova_model):
|
||||
"""Two reconnects in succession are safe (relies on idempotent stop())."""
|
||||
await nova_model.start(system_prompt="You are helpful")
|
||||
await nova_model.reconnect(system_prompt="You are helpful")
|
||||
await nova_model.reconnect(system_prompt="You are helpful")
|
||||
assert nova_model._connection_id is not None
|
||||
await nova_model.stop()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_proactive_reconnect_end_to_end_through_agent(model_id, boto_session, mock_client, mock_stream):
|
||||
"""End-to-end: BidiAgent + real Nova model proactively reconnects before the deadline.
|
||||
|
||||
Drives the full chain against the real BidiNovaSonicModel (mocked Bedrock transport):
|
||||
the loop reads Nova's connection_config, arms the proactive timer, emits a warning,
|
||||
and reconnects through Nova's own reconnect() before the session deadline, replaying
|
||||
history via Nova's initialization path. No live AWS calls are made.
|
||||
"""
|
||||
from strands.experimental.bidi.agent.agent import BidiAgent
|
||||
from strands.experimental.bidi.types.events import BidiConnectionWarningEvent
|
||||
|
||||
# Nova never emits events on its own here; await_output blocks so the model task idles
|
||||
# while the proactive timer drives the reconnect.
|
||||
output = AsyncMock()
|
||||
never = asyncio.Event()
|
||||
|
||||
async def blocking_receive():
|
||||
await never.wait()
|
||||
|
||||
output.receive = AsyncMock(side_effect=blocking_receive)
|
||||
mock_stream.await_output = AsyncMock(return_value=(None, output))
|
||||
|
||||
model = BidiNovaSonicModel(model_id=model_id, client_config={"boto_session": boto_session})
|
||||
# A small deadline; the injected clock below fires it without wall time.
|
||||
model.connection_config = {"restart_after_s": 1}
|
||||
|
||||
agent = BidiAgent(model=model, system_prompt="You are helpful")
|
||||
|
||||
# Drive the timer without wall time: the first cycle's sleeps return immediately, the re-armed
|
||||
# cycle after the swap parks, so exactly one proactive reconnect fires.
|
||||
sleep_count = 0
|
||||
|
||||
async def fake_sleep(_seconds):
|
||||
nonlocal sleep_count
|
||||
sleep_count += 1
|
||||
if sleep_count > 2:
|
||||
await asyncio.Event().wait()
|
||||
await asyncio.sleep(0)
|
||||
|
||||
agent._loop._reconnect_timer._sleep = fake_sleep
|
||||
|
||||
await agent.start()
|
||||
|
||||
first_connection_id = model._connection_id
|
||||
|
||||
warning_seen = False
|
||||
async for event in agent.receive():
|
||||
if isinstance(event, BidiConnectionWarningEvent):
|
||||
warning_seen = True
|
||||
# Once a reconnect has produced a new connection id, the proactive cycle completed.
|
||||
if model._connection_id is not None and model._connection_id != first_connection_id:
|
||||
break
|
||||
|
||||
assert warning_seen
|
||||
assert model._connection_id != first_connection_id
|
||||
assert mock_stream.close.called # old connection was torn down via reconnect -> stop
|
||||
|
||||
await agent.stop()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_model_stop_after_start_failure(model_id, boto_session):
|
||||
with patch("strands.experimental.bidi.models.nova_sonic.AsyncBedrockRuntimeClient") as mock_cls:
|
||||
|
||||
Reference in New Issue
Block a user