feat(bidi): proactive session reconnect for bidi session (#3874)

This commit is contained in:
mehtarac
2026-08-28 08:15:02 -04:00
committed by GitHub
parent 45bcea8845
commit 23d39acacc
16 changed files with 1747 additions and 92 deletions
@@ -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: