refactor(bidi): correct docstrings and drop unused barge-in and connection stop reasons (#4728)

Co-authored-by: Claude Opus 5.5 <noreply@anthropic.com>
This commit is contained in:
mehtarac
2026-09-30 10:30:42 -04:00
committed by GitHub
co-authored by Claude Opus 5.5
parent 4dfeca8c74
commit 2da3a696f0
26 changed files with 127 additions and 181 deletions
@@ -80,7 +80,7 @@ async def main():
async for event in agent.receive():
if isinstance(event, BidiBargeInEvent):
print(f"Barge-in: {event.reason}")
print("Barge-in detected")
# Custom handling:
# - Update UI to show barge-in
# - Log analytics
@@ -95,8 +95,7 @@ asyncio.run(main())
### Key Events
**BidiBargeInEvent** - Emitted when barge-in detected:
- `reason`: `"user_speech"` (most common) or `"error"`
**BidiBargeInEvent** - Emitted when a barge-in is detected. It carries no fields beyond `type`.
## Barge-in Hooks
@@ -114,7 +113,7 @@ class BargeInTracker:
async def on_barge_in(self, event: BidiBargeInHookEvent):
self.barge_in_count += 1
print(f"Barge-in #{self.barge_in_count}: {event.reason}")
print(f"Barge-in #{self.barge_in_count}")
# Log to analytics
# Update UI
@@ -219,12 +219,7 @@ Emitted when the streaming connection is closed.
**Properties:**
- `connection_id`: Unique identifier for this streaming connection
- `reason`: Why the connection closed
- `"client_disconnect"`: Client disconnected
- `"timeout"`: Connection timed out
- `"error"`: Error occurred
- `"complete"`: Conversation completed normally
- `"user_request"`: Cancellation requested through `agent.cancel()`
- `reason`: Why the connection closed. Always `"user_request"`, emitted once a cancellation requested through `agent.cancel()` takes effect.
### Response Lifecycle Events
@@ -469,22 +464,15 @@ Signals a barge-in that stops response generation or playback, typically when th
```python
{
"type": "bidi_barge_in",
"reason": "user_speech"
"type": "bidi_barge_in"
}
```
**Properties:**
- `reason`: Why response output should stop
- `"user_speech"`: User started speaking (most common)
- `"error"`: Error stopped output
**Usage:**
```python
async for event in agent.receive():
if event["type"] == "bidi_barge_in":
print(f"Barge-in: {event['reason']}")
print("Barge-in detected")
# Audio output automatically cleared
# Model ready for new input
```
@@ -204,7 +204,7 @@ This section contains practical hook implementations for common use cases.
### Tracking barge-ins
Count barge-ins and record their reasons:
Count barge-ins:
```python
from strands.bidi.agent import BidiAgent
@@ -221,7 +221,7 @@ class BargeInTracker:
async def on_barge_in(self, event: BidiBargeInEvent) -> None:
self.barge_in_count += 1
print(f"Barge-in #{self.barge_in_count}: {event.reason}")
print(f"Barge-in #{self.barge_in_count}")
tracker = BargeInTracker()
@@ -165,8 +165,7 @@ agents. See [Tool-Level Attributes](../observability-evaluation/traces.md#tool-l
### Barge-in events
A barge-in is recorded as a `bidi_barge_in` event on the session span with a
`barge_in.reason` attribute.
A barge-in is recorded as a `bidi_barge_in` event on the session span.
To count barge-ins, query the session span's `bidi_barge_in` events. Keeping the
event on the session also captures barge-ins between response spans. For detection
@@ -211,9 +210,7 @@ The console exporter is the fastest way to confirm spans are being produced. A s
{
"name": "bidi_barge_in",
"timestamp": "2026-08-10T13:27:28.100000Z",
"attributes": {
"barge_in.reason": "user_speech"
}
"attributes": {}
}
],
"links": []
@@ -35,7 +35,7 @@ class _TaskGroup:
return self
async def __aexit__(self, *_: Any) -> None:
"""Execute tasks in group.
"""Wait for all tasks in the group, cancelling the rest if one fails.
The following execution rules are enforced:
- The context stops executing all tasks if at least one task raises an Exception or the context is cancelled.
@@ -24,6 +24,9 @@ class _TaskPool:
Adds a clean up callback to run after task completes.
Args:
coro: Coroutine to run as a task in the pool.
Returns:
The created task.
"""
+2 -3
View File
@@ -190,12 +190,11 @@ def end_connection_span(tracer: Tracer, span: Span, error: Exception | None = No
tracer._end_span(span, error=error)
def add_barge_in_event(span: Span, reason: str) -> None:
def add_barge_in_event(span: Span) -> None:
"""Record a barge-in as a span event on the session span.
Args:
span: The session span to add the event to.
reason: Reason for the barge-in.
"""
if span and span.is_recording():
span.add_event("bidi_barge_in", attributes={"barge_in.reason": reason})
span.add_event("bidi_barge_in")
+8 -1
View File
@@ -434,10 +434,17 @@ class BidiAgent(LocalAgent):
Yields:
Model output events processed by background tasks including audio output,
text responses, tool calls, and connection updates.
text responses, tool calls, and connection updates. The agent also yields:
- A completed ``BidiTextBlockEvent``, ``BidiReasoningBlockEvent``, or
``BidiTranscriptBlockEvent`` after each text, reasoning, or transcript stop event.
- Tool execution events such as ``ToolStreamEvent`` and ``ToolResultEvent``, and a
``ToolResultMessageEvent`` once a tool group's results are recorded.
Raises:
RuntimeError: If start has not been called.
ConnectionTimeoutError: If the model connection times out and automatic restart is
disabled.
"""
if not self._started:
raise RuntimeError("agent not started | call start before receiving")
+12 -7
View File
@@ -205,7 +205,12 @@ class _AgentLoop:
self._arm_restart_timer()
async def stop(self) -> None:
"""Stop the agent loop."""
"""Stop the agent loop.
Closes the send gate and cancels the restart timer, then cancels background tasks (the
model reader and running tools) and stops the model. The session span is ended and
``BidiAgentStopEvent`` fires even if teardown fails.
"""
logger.debug("agent loop stopping")
self._started = False
@@ -308,8 +313,8 @@ class _AgentLoop:
restart_event=restart_event,
)
except Exception:
# The restart event was queued before the failing swap. Surface it before
# preserving the existing behavior of raising the restart failure.
# The restart event was queued before the failing swap. Surface it, then
# raise the restart failure to the caller.
yield self._event_queue.get_nowait()
raise
continue
@@ -348,7 +353,7 @@ class _AgentLoop:
async def _on_restart_warning(self, time_left_s: int) -> None:
"""Timer callback: surface an approaching-restart warning to the receiver."""
logger.debug("time_left_s=<%.1f> | emitting connection warning", time_left_s)
logger.debug("time_left_s=<%s> | emitting connection warning", time_left_s)
await self._event_queue.put(BidiConnectionWarningEvent(time_left_s=time_left_s))
async def _on_restart_deadline(self) -> None:
@@ -384,7 +389,7 @@ class _AgentLoop:
await asyncio.wait_for(self._turn_complete.wait(), timeout=_MODEL_RESTART_TURN_TIMEOUT_S)
except asyncio.TimeoutError:
logger.debug(
"no turn boundary within %.1fs | forcing restart",
"timeout_s=<%s> | no turn boundary, forcing restart",
_MODEL_RESTART_TURN_TIMEOUT_S,
)
@@ -684,12 +689,12 @@ class _AgentLoop:
elif isinstance(event, BidiBargeInEvent):
if self._session_span:
_telemetry.add_barge_in_event(self._session_span, event["reason"])
_telemetry.add_barge_in_event(self._session_span)
# A barge-in ends the current response; the user's next turn owes a reply.
self._response_active = False
self._update_turn_state()
await self._agent.hooks.invoke_callbacks_async(BidiBargeInHookEvent(self._agent, event["reason"]))
await self._agent.hooks.invoke_callbacks_async(BidiBargeInHookEvent(self._agent))
elif isinstance(event, BidiResponseStopEvent):
if response_span:
+2 -7
View File
@@ -60,17 +60,12 @@ class BidiBargeInEvent(_HookEvent):
"""Event triggered to stop current response generation or playback.
This event is fired when the user barges in (e.g., by speaking during the
assistant's response) or when an error stops output. This is
specific to a response and does not pause the bidirectional session.
assistant's response). This is specific to a response and does not pause
the bidirectional session.
Hook providers can use this event to log barge-ins, stop playback, or trigger cleanup.
Attributes:
reason: Why response output should stop ("user_speech" or "error").
"""
reason: Literal["user_speech", "error"]
@dataclass
class BidiBeforeConnectionRestartEvent(_HookEvent):
+1 -1
View File
@@ -262,7 +262,7 @@ class _AudioOutputStream(OutputStream):
logger.debug("audio_bytes=<%d> | audio chunk buffered for playback", len(data))
elif isinstance(event, BidiBargeInEvent):
logger.debug("reason=<%s> | clearing audio buffer due to barge-in", event["reason"])
logger.debug("clearing audio buffer due to barge-in")
self._buffer.clear()
if self._audio_processor is not None:
self._audio_processor.clear_far_data()
@@ -7,7 +7,7 @@ InvokeModelWithBidirectionalStream protocol.
Nova Sonic specifics:
- Hierarchical event sequences: connectionStart → promptStart → content streaming
- Base64-encoded audio format with hex encoding
- Base64-encoded audio
- Tool execution with content containers and identifier tracking
- 8-minute connection limits with proper cleanup sequences
- Barge-in detection through stopReason events
@@ -890,7 +890,7 @@ class BedrockNovaSonicModel(BidiModel, AudioCapable):
if stop_reason == "INTERRUPTED":
# The user holds the turn until Nova answers, even if the response already ended.
response_state.idle.clear()
events.append(BidiBargeInEvent("user_speech"))
events.append(BidiBargeInEvent())
if response_state.response_id is not None:
events.extend(self._complete_response(response_state))
return events
+3 -11
View File
@@ -2,14 +2,6 @@
Implements the BidiModel interface for Google's Gemini Live API using the
official Google GenAI SDK for simplified and robust WebSocket communication.
Key improvements over custom WebSocket implementation:
- Uses official google-genai SDK with native Live API support
- Simplified session management with client.aio.live.connect()
- Built-in tool integration and event handling
- Automatic WebSocket connection management and error handling
- Native support for audio/text streaming and barge-in
"""
import base64
@@ -463,7 +455,7 @@ class GoogleGeminiLiveModel(BidiModel, AudioCapable):
events: list[BidiOutputEvent] = []
if server_content.interrupted:
events.append(BidiBargeInEvent(reason="user_speech"))
events.append(BidiBargeInEvent())
input_transcript = server_content.input_transcription
if input_transcript and input_transcript.text:
@@ -753,7 +745,7 @@ class GoogleGeminiLiveModel(BidiModel, AudioCapable):
"input_audio_transcription": {},
# Sliding-window context compression removes the ~15-min audio-only session cap, so a
# session resumed across proactive restarts can continue indefinitely rather than
# dying at the cap (gemini_session.md).
# dying at the cap.
"context_window_compression": {"sliding_window": {}},
}
@@ -761,7 +753,7 @@ class GoogleGeminiLiveModel(BidiModel, AudioCapable):
config_dict["session_resumption"] = {"handle": live_session_handle}
# Enables send_client_content for initial history seeding before realtime mode.
# Not supported on Vertex AI; HistoryConfig requires google-genai>=1.67 (floor bump tracked separately).
# Not supported on Vertex AI.
has_messages = kwargs.get("has_messages", False)
if has_messages and getattr(self._client, "vertexai", False) is not True:
config_dict["history_config"] = {"initial_history_in_client_content": True}
+2 -15
View File
@@ -1,17 +1,4 @@
"""Bidirectional streaming model interface.
Defines the abstract interface for models that support real-time bidirectional
communication with persistent connections. Unlike traditional request-response
models, bidirectional models maintain an open connection for streaming audio,
text, and tool interactions.
Features:
- Persistent connection management with connect/close lifecycle
- Real-time bidirectional communication (send and receive simultaneously)
- Provider-agnostic event normalization
- Support for audio, text, image, and tool result streaming
"""
"""Bidirectional streaming model interface: start a persistent connection, send and receive concurrently, then stop."""
import abc
import logging
@@ -113,7 +100,7 @@ class BidiModel(Model, abc.ABC):
Terminates the active bidirectional connection and cleans up any associated
resources such as network connections, buffers, or background tasks. After
calling close(), the model instance cannot be used until start() is called again.
calling stop(), the model instance cannot be used until start() is called again.
"""
pass
+4 -4
View File
@@ -60,8 +60,6 @@ from .model import AudioCapable, BidiModel, ConnectionTimeoutError
logger = logging.getLogger(__name__)
# Test idle_timeout_ms
# OpenAI Realtime API configuration
OPENAI_MAX_TIMEOUT_S = 3000 # 50 minutes
"""Max timeout before closing connection.
@@ -209,7 +207,9 @@ class OpenAIRealtimeModel(BidiModel, AudioCapable):
api_key: OpenAI API key. Defaults to ``OPENAI_API_KEY``.
organization: OpenAI organization. Defaults to ``OPENAI_ORGANIZATION``.
project: OpenAI project. Defaults to ``OPENAI_PROJECT``.
timeout_s: Maximum connection duration in seconds.
timeout_s: Maximum connection duration in seconds. Unless ``connection.restart_after_s`` is
set, the agent restarts the connection 5 minutes before this limit, so a value of 300 or
less disables the proactive restart.
voice: Output voice identifier. Defaults to ``alloy``.
**model_config: Model configuration.
@@ -594,7 +594,7 @@ class OpenAIRealtimeModel(BidiModel, AudioCapable):
state = state if state is not None else self._session_state
if event_type == "input_audio_buffer.speech_started":
events: list[BidiOutputEvent] = [BidiBargeInEvent(reason="user_speech")]
events: list[BidiOutputEvent] = [BidiBargeInEvent()]
if state.transcription_enabled:
events.extend(state.start_transcript("user", openai_event["item_id"]))
return events
+38 -62
View File
@@ -1,24 +1,10 @@
"""Bidirectional streaming types for real-time audio/text conversations.
"""Output event types for bidirectional streaming.
Type definitions for bidirectional streaming that extends Strands' existing streaming
capabilities with real-time audio and persistent connection support.
Key features:
- Audio output events with standardized formats
- Barge-in detection and handling
- Connection lifecycle management
- Provider-agnostic event types
- Type-safe discriminated unions with TypedEvent
- JSON-serializable output events (audio stored as base64 strings)
Audio format normalization:
- Supports PCM, WAV, Opus, and MP3 formats
- Describes sample rates in Hz
- Normalizes channel configurations (mono/stereo)
- Abstracts provider-specific encodings
- Audio output stored as base64-encoded strings for JSON compatibility
Defines the provider-agnostic events produced by bidirectional models and ``BidiAgent``:
connection lifecycle (start, restart, warning, stop), response start and stop, audio, text,
reasoning, and transcript streams (start, delta, stop, and the completed block), barge-in,
token usage, and tool-use groups. Also defines the ``AudioChannel``, ``AudioFormat``, and
``Role`` literals and the ``BidiOutputEvent`` union.
"""
import logging
@@ -39,7 +25,11 @@ AudioChannel = Literal[1, 2]
- Stereo: 2
"""
AudioFormat = Literal["pcm", "wav", "opus", "mp3"]
"""Audio encoding format."""
"""Audio encoding format of model audio output and ``AudioStreamConfig``.
Distinct from ``strands.types.media.AudioFormat``, the wider set of formats that types
``AudioDelta.format`` on audio input.
"""
Role = Literal["user", "assistant"]
"""Role of a message sender.
@@ -81,7 +71,7 @@ def _normalize_role(role: Any, default: Role = "user") -> Role:
class BidiConnectionStartEvent(TypedEvent):
"""Streaming connection established and ready for interaction.
Parameters:
Args:
connection_id: Unique identifier for this streaming connection.
model: Model identifier (e.g., "gpt-realtime-2.1", "gemini-3.8-live").
"""
@@ -113,13 +103,13 @@ class BidiConnectionRestartEvent(TypedEvent):
Emitted on both restart paths: reactively after the model reports a timeout, and
proactively when the restart timer fires ahead of the provider's limit.
Parameters:
Args:
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.
turn_interrupted: True if the restart cut off an in-progress assistant response or a
user turn that had not been answered yet. The new connection receives the history
as context, so that turn is not answered on its own; an app can re-prompt or notify
the user when this is set.
"""
def __init__(
@@ -139,9 +129,9 @@ class BidiConnectionRestartEvent(TypedEvent):
)
@property
def reason(self) -> str:
def reason(self) -> Literal["timeout", "scheduled"]:
"""What triggered the restart ("timeout" or "scheduled")."""
return cast(str, self["reason"])
return cast(Literal["timeout", "scheduled"], self["reason"])
@property
def timeout_error(self) -> "ConnectionTimeoutError | None":
@@ -150,7 +140,7 @@ class BidiConnectionRestartEvent(TypedEvent):
@property
def turn_interrupted(self) -> bool:
"""True if the restart cut an in-progress or owed turn that will not be answered."""
"""True if the restart cut off an in-progress response or an unanswered user turn."""
return cast(bool, self["turn_interrupted"])
@@ -159,7 +149,7 @@ class BidiConnectionWarningEvent(TypedEvent):
Emitted by the proactive restart timer before a restart; informational only.
Parameters:
Args:
time_left_s: Approximate seconds until the scheduled restart.
"""
@@ -181,7 +171,7 @@ class BidiConnectionWarningEvent(TypedEvent):
class BidiResponseStartEvent(TypedEvent):
"""Start of a model response.
Parameters:
Args:
response_id: Unique identifier for this response (used in BidiResponseStopEvent).
"""
@@ -211,7 +201,7 @@ class BidiAudioStartEvent(TypedEvent):
class BidiAudioDeltaEvent(TypedEvent):
"""Incremental audio output from the model.
Parameters:
Args:
audio: Base64-encoded audio chunk.
format: Audio encoding format.
sample_rate: Number of audio samples per second in Hz.
@@ -405,7 +395,7 @@ class BidiReasoningBlockEvent(TypedEvent):
class BidiTranscriptStartEvent(TypedEvent):
"""Beginning of a user or assistant transcript, before its text arrives.
Parameters:
Args:
role: Who is speaking ("user" or "assistant").
content_id: Unique identifier shared by this transcript's events.
"""
@@ -434,7 +424,7 @@ class BidiTranscriptStartEvent(TypedEvent):
class BidiTranscriptDeltaEvent(TypedEvent):
"""Incremental transcription of user or assistant speech.
Parameters:
Args:
delta: The incremental transcript text.
role: Who is speaking ("user" or "assistant").
content_id: Unique identifier shared by this transcript's events.
@@ -470,7 +460,7 @@ class BidiTranscriptDeltaEvent(TypedEvent):
class BidiTranscriptStopEvent(TypedEvent):
"""End of a transcript stream, before its completed block is emitted.
Parameters:
Args:
role: Who spoke ("user" or "assistant").
content_id: Unique identifier shared by this transcript's events.
"""
@@ -499,7 +489,7 @@ class BidiTranscriptStopEvent(TypedEvent):
class BidiTranscriptBlockEvent(TypedEvent):
"""Complete transcript, emitted after its stop event by the agent.
Parameters:
Args:
transcript: The final transcript text.
role: Who spoke ("user" or "assistant").
content_id: Unique identifier shared by this transcript's events.
@@ -533,31 +523,17 @@ class BidiTranscriptBlockEvent(TypedEvent):
class BidiBargeInEvent(TypedEvent):
"""Stop current response generation or playback while the session continues.
"""Stop current response generation or playback while the session continues."""
Parameters:
reason: Why response output should stop.
"""
def __init__(self, reason: Literal["user_speech", "error"]):
def __init__(self) -> None:
"""Initialize barge-in event."""
super().__init__(
{
"type": "bidi_barge_in",
"reason": reason,
}
)
@property
def reason(self) -> str:
"""Why response output should stop."""
return cast(str, self["reason"])
super().__init__({"type": "bidi_barge_in"})
class BidiResponseStopEvent(TypedEvent):
"""Response output ended. User transcription may still be pending.
Parameters:
Args:
response_id: ID of the response that ended (matches BidiResponseStartEvent).
"""
@@ -596,7 +572,7 @@ class BidiUsageEvent(TypedEvent):
Tracks token consumption across different modalities (audio, text, images)
during bidirectional streaming sessions.
Parameters:
Args:
input_tokens: Total tokens used for all input modalities.
output_tokens: Total tokens used for all output modalities.
total_tokens: Sum of input and output tokens.
@@ -663,7 +639,7 @@ class BidiUsageEvent(TypedEvent):
class BidiToolUseBlocksEvent(TypedEvent):
"""A complete group of tool calls requested by the model.
Parameters:
Args:
tool_uses: Tool calls to execute together.
"""
@@ -680,15 +656,15 @@ class BidiToolUseBlocksEvent(TypedEvent):
class BidiConnectionStopEvent(TypedEvent):
"""Streaming connection closed.
Parameters:
Args:
connection_id: Unique identifier for this streaming connection (matches BidiConnectionStartEvent).
reason: Why the connection was closed.
reason: Why the connection was closed. ``"user_request"`` after ``agent.cancel()`` takes effect.
"""
def __init__(
self,
connection_id: str,
reason: Literal["client_disconnect", "timeout", "error", "complete", "user_request"],
reason: Literal["user_request"],
):
"""Initialize connection stop event."""
super().__init__(
@@ -705,9 +681,9 @@ class BidiConnectionStopEvent(TypedEvent):
return cast(str, self["connection_id"])
@property
def reason(self) -> str:
def reason(self) -> Literal["user_request"]:
"""Why the connection was closed."""
return cast(str, self["reason"])
return cast(Literal["user_request"], self["reason"])
# ============================================================================
@@ -80,7 +80,7 @@ class MockBidiModel(BidiModel):
yield event
# Yield connection end event
yield BidiConnectionStopEvent(connection_id=self._connection_id, reason="complete")
yield BidiConnectionStopEvent(connection_id=self._connection_id, reason="user_request")
def set_events(self, events):
"""Helper to set events this mock model will yield."""
@@ -629,7 +629,7 @@ async def test_bidi_agent_receive_events_from_model(agent, events):
exp_events = [
BidiConnectionStartEvent(connection_id=unittest.mock.ANY, model=unittest.mock.ANY),
*events,
BidiConnectionStopEvent(connection_id=unittest.mock.ANY, reason="complete"),
BidiConnectionStopEvent(connection_id=unittest.mock.ANY, reason="user_request"),
]
await agent.start()
@@ -351,7 +351,7 @@ async def test_receive_barge_in_does_not_wait_for_transcription(agent, agenerato
complete_a = BidiResponseStopEvent("a")
start_b = BidiResponseStartEvent("b")
audio_b = BidiAudioDeltaEvent("cancelled", "pcm", 24000, 1, content_id="audio")
barge_in = BidiBargeInEvent("user_speech")
barge_in = BidiBargeInEvent()
complete_b = BidiResponseStopEvent("b")
transcript = BidiTranscriptBlockEvent("Earlier question.", "user", content_id="speech-a")
native_events = [
@@ -374,7 +374,7 @@ async def test_receive_barge_in_does_not_wait_for_transcription(agent, agenerato
exp_events = [*native_events[:4], transcript, *native_events[4:]]
tru_events = [await asyncio.wait_for(anext(reader), 1) for _ in exp_events]
assert tru_events == exp_events
assert hooks.events_received == [BidiBargeInHookEvent(agent=agent, reason="user_speech")]
assert hooks.events_received == [BidiBargeInHookEvent(agent=agent)]
finally:
await reader.aclose()
await agent.stop()
@@ -573,7 +573,7 @@ async def test_model_processes_transcripts_before_consumer_reads(loop, agent, ag
"stream_event,hook_type",
[
(BidiResponseStopEvent(response_id="r1"), BidiResponseStopHookEvent),
(BidiBargeInEvent(reason="user_speech"), BidiBargeInHookEvent),
(BidiBargeInEvent(), BidiBargeInHookEvent),
(BidiTranscriptStartEvent(role="assistant", content_id="assistant-transcript"), MessageAddedEvent),
],
)
@@ -706,7 +706,7 @@ async def test_agent_stop_hook(agent, agenerator, cleanup_fails):
@pytest.mark.asyncio
async def test_bidi_agent_loop_receive_restart_connection(loop, agent, agenerator):
timeout_error = ConnectionTimeoutError("test timeout", test_restart_config=1)
close_event = BidiConnectionStopEvent(connection_id="test", reason="complete")
close_event = BidiConnectionStopEvent(connection_id="test", reason="user_request")
agent.model.receive = unittest.mock.Mock(side_effect=[timeout_error, agenerator([close_event])])
@@ -762,7 +762,7 @@ async def test_bidi_agent_loop_auto_reconnect_default_on(loop, agent, agenerator
# An empty connection config uses the default restart behavior.
agent.model.get_connection_config.return_value = {}
timeout_error = ConnectionTimeoutError("test timeout")
close_event = BidiConnectionStopEvent(connection_id="test", reason="complete")
close_event = BidiConnectionStopEvent(connection_id="test", reason="user_request")
agent.model.receive = unittest.mock.Mock(side_effect=[timeout_error, agenerator([close_event])])
await loop.start()
@@ -992,7 +992,7 @@ async def test_restart_fences_superseded_reader_stream_close_error():
await loop.start()
first = BidiConnectionStopEvent(connection_id="first", reason="complete")
first = BidiConnectionStopEvent(connection_id="first", reason="user_request")
await model.emit(first)
assert await loop._event_queue.get() is first
@@ -1001,7 +1001,7 @@ async def test_restart_fences_superseded_reader_stream_close_error():
await loop._restart_connection(None, loop._generation)
assert model.restart_calls == 1
second = BidiConnectionStopEvent(connection_id="second", reason="complete")
second = BidiConnectionStopEvent(connection_id="second", reason="user_request")
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
@@ -1062,7 +1062,7 @@ async def test_stale_reader_error_is_dropped_not_raised(loop, agent, agenerator)
# An error raised on a superseded (older) generation.
await loop._event_queue.put(_ReaderError(loop._generation - 1, OSError("stale connection error")))
sentinel = BidiConnectionStopEvent(connection_id="after-stale-error", reason="complete")
sentinel = BidiConnectionStopEvent(connection_id="after-stale-error", reason="user_request")
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)
@@ -1102,7 +1102,7 @@ async def test_stale_reactive_timeout_dropped_after_proactive_swap(loop, agent,
# A timeout tagged with the pre-swap generation is now stale; receive() must drop it.
await loop._event_queue.put(_ReaderError(stale_generation, ConnectionTimeoutError("stale timeout")))
sentinel = BidiConnectionStopEvent(connection_id="after-stale-timeout", reason="complete")
sentinel = BidiConnectionStopEvent(connection_id="after-stale-timeout", reason="user_request")
feed = asyncio.create_task(_feed_after_drain(loop, sentinel))
result = await asyncio.wait_for(loop.receive().__anext__(), timeout=2.0)
assert result is sentinel
@@ -309,7 +309,7 @@ async def test_tool_call_span_closed_on_error(loop, agent, agenerator, otel_setu
async def test_connection_restart_span(loop, agent, agenerator, otel_setup):
"""Connection restart creates a span with error message."""
timeout_error = ConnectionTimeoutError("8 minute timeout")
close_event = BidiConnectionStopEvent(connection_id="test", reason="complete")
close_event = BidiConnectionStopEvent(connection_id="test", reason="user_request")
agent.model.receive = unittest.mock.Mock(side_effect=[timeout_error, agenerator([close_event])])
@@ -380,7 +380,7 @@ async def test_barge_in_event_recorded_on_session_span(loop, agent, agenerator,
"""Barge-in events are added to the session span."""
events = [
BidiResponseStartEvent(response_id="resp-3"),
BidiBargeInEvent(reason="user_speech"),
BidiBargeInEvent(),
BidiResponseStopEvent(response_id="resp-3"),
]
agent.model.receive = unittest.mock.Mock(return_value=agenerator(events))
@@ -398,9 +398,7 @@ async def test_barge_in_event_recorded_on_session_span(loop, agent, agenerator,
assert len(session_spans) == 1
span_events = session_spans[0].events
assert any(
event.name == "bidi_barge_in" and event.attributes["barge_in.reason"] == "user_speech" for event in span_events
)
assert any(event.name == "bidi_barge_in" for event in span_events)
@pytest.mark.asyncio
@@ -63,7 +63,7 @@ def response_stop_event(agent):
@pytest.fixture
def barge_in_event(agent):
return BidiBargeInEvent(agent=agent, reason="user_speech")
return BidiBargeInEvent(agent=agent)
def test_event_should_reverse_callbacks(agent_stop_event, response_stop_event, barge_in_event):
@@ -74,10 +74,10 @@ def test_event_should_reverse_callbacks(agent_stop_event, response_stop_event, b
def test_barge_in_event_fields(agent):
event = BidiBargeInEvent(agent=agent, reason="error")
event = BidiBargeInEvent(agent=agent)
tru_event = {field.name: getattr(event, field.name) for field in fields(event)}
exp_event = {"agent": agent, "reason": "error"}
exp_event = {"agent": agent}
assert tru_event == exp_event
@@ -179,7 +179,7 @@ async def test_output_connection_stop_flushes_partial_content(console, output_st
await output_stream(BidiTextStartEvent("text"))
await output_stream(BidiTextDeltaEvent("Partial", "text"))
await output_stream(BidiConnectionStopEvent("connection", reason="complete"))
await output_stream(BidiConnectionStopEvent("connection", reason="user_request"))
assert capsys.readouterr().out.strip() == "Partial"
assert not console._display.blocks
@@ -204,7 +204,7 @@ async def test_audio_io_output_barge_in(audio_output):
content_id="audio",
)
await audio_output(audio_event)
barge_in_event = BidiBargeInEvent(reason="user_speech")
barge_in_event = BidiBargeInEvent()
await audio_output(barge_in_event)
tru_data, _ = audio_output._callback(None, frame_count=1)
@@ -642,7 +642,7 @@ async def test_output_clears_reference_on_barge_in(py_audio, aec_agent, mock_aud
)
output._callback(None, frame_count=2)
await output(BidiBargeInEvent(reason="user_speech"))
await output(BidiBargeInEvent())
assert audio_io._audio_processor._get_far_data() == b""
await input_.stop()
@@ -662,7 +662,7 @@ def test_barge_in_closes_response_before_next_turn(nova_model, role):
tru_events = nova_model._convert_nova_event({"contentEnd": {"type": "TEXT", "stopReason": "INTERRUPTED"}}, state)
exp_events = [
BidiBargeInEvent("user_speech"),
BidiBargeInEvent(),
BidiTranscriptStopEvent(role, "t1"),
BidiResponseStopEvent("r1"),
]
@@ -756,7 +756,7 @@ def test_response_after_barge_in_finishes_before_next_user_transcript(nova_model
assert response_state.generation_stage == "FINAL"
elif native_event.get("contentEnd", {}).get("contentId") == "control":
assert events == [
BidiBargeInEvent("user_speech"),
BidiBargeInEvent(),
*([BidiAudioStopEvent(content_id=ANY)] if final_fragments else []),
BidiTranscriptStopEvent("assistant", content_id="t1"),
BidiResponseStopEvent("r1"),
@@ -767,7 +767,7 @@ def test_response_after_barge_in_finishes_before_next_user_transcript(nova_model
assert events == []
exp_events = [
*([BidiAudioStartEvent(content_id=ANY)] if final_fragments else []),
BidiBargeInEvent("user_speech"),
BidiBargeInEvent(),
*([BidiAudioStopEvent(content_id=ANY)] if final_fragments else []),
BidiTranscriptStopEvent("assistant", content_id="t1"),
BidiResponseStopEvent("r1"),
@@ -787,7 +787,7 @@ def test_barge_in_after_response_stop_only_stops_playback(nova_model):
tru_events = nova_model._convert_nova_event(
{"contentEnd": {"type": "TEXT", "stopReason": "INTERRUPTED"}}, response_state
)
exp_events = [BidiBargeInEvent("user_speech")]
exp_events = [BidiBargeInEvent()]
assert tru_events == exp_events
assert response_state == _ResponseState()
assert not response_state.idle.is_set()
@@ -1499,7 +1499,7 @@ def test_audio_stream_preserves_content_id(nova_model, interrupted):
BidiAudioStartEvent(content_id),
BidiAudioDeltaEvent("YQ==", "pcm", 16000, 1, content_id),
BidiAudioDeltaEvent("Yg==", "pcm", 16000, 1, content_id),
*([BidiBargeInEvent("user_speech")] if interrupted else []),
*([BidiBargeInEvent()] if interrupted else []),
BidiAudioStopEvent(content_id),
BidiResponseStopEvent(ANY),
]
@@ -923,7 +923,7 @@ async def test_event_conversion(mock_genai_client, model, live_message, server_c
barge_in_events = model._convert_gemini_live_event(mock_barge_in, turn_state)
assert barge_in_events == [
BidiBargeInEvent(reason="user_speech"),
BidiBargeInEvent(),
BidiAudioStopEvent(content_id=unittest.mock.ANY),
]
@@ -1049,7 +1049,7 @@ def test_barge_in_emitted_alongside_other_server_content(model, complete_with_ou
tru_events = [event for message in messages for event in model._convert_gemini_live_event(message, turn_state)]
exp_events = [
BidiResponseStartEvent(unittest.mock.ANY),
BidiBargeInEvent("user_speech"),
BidiBargeInEvent(),
BidiTranscriptStartEvent("assistant", content_id=unittest.mock.ANY),
BidiTranscriptDeltaEvent("partial reply", "assistant", content_id=unittest.mock.ANY),
BidiTranscriptStopEvent("assistant", content_id=unittest.mock.ANY),
@@ -1098,7 +1098,7 @@ async def test_barge_in_preserves_user_transcription_already_in_progress(
turn_state,
)
exp_events = [
BidiBargeInEvent("user_speech"),
BidiBargeInEvent(),
BidiTranscriptDeltaEvent(" second", "user", started[0].content_id),
]
assert tru_events == exp_events
@@ -1150,7 +1150,7 @@ def test_transcription_fragments_complete_at_turn_boundary(model):
True,
[{"interrupted": True}, {"turn_complete": True}],
[
BidiBargeInEvent(reason="user_speech"),
BidiBargeInEvent(),
BidiTranscriptStopEvent("user", content_id=unittest.mock.ANY),
BidiResponseStopEvent("r1"),
],
@@ -1160,7 +1160,7 @@ def test_transcription_fragments_complete_at_turn_boundary(model):
True,
[{"interrupted": True, "turn_complete": True}],
[
BidiBargeInEvent(reason="user_speech"),
BidiBargeInEvent(),
BidiTranscriptStopEvent("user", content_id=unittest.mock.ANY),
BidiResponseStopEvent("r1"),
],
@@ -1353,7 +1353,7 @@ def test_audio_stops_once_at_generation_boundary(model, live_message, server_con
BidiAudioDeltaEvent("Zmlyc3Q=", format="pcm", sample_rate=24000, channels=1, content_id=content_id),
]
if ending == "interrupted":
exp_events.append(BidiBargeInEvent(reason="user_speech"))
exp_events.append(BidiBargeInEvent())
exp_events.extend(
[
BidiAudioDeltaEvent("bGFzdA==", format="pcm", sample_rate=24000, channels=1, content_id=content_id),
@@ -1390,7 +1390,7 @@ async def test_turn_complete_without_open_response_emits_nothing(
tru_events.extend(
model._convert_gemini_live_event(live_message(server_content=server_content(turn_complete=True)), turn_state)
)
exp_events = [BidiBargeInEvent("user_speech")] if interrupted else []
exp_events = [BidiBargeInEvent()] if interrupted else []
assert tru_events == exp_events
@@ -1413,7 +1413,7 @@ async def test_barge_in_completes_at_turn_boundary(
events = model._convert_gemini_live_event(live_message(server_content=server_content(interrupted=True)), turn_state)
assert events == [BidiBargeInEvent(reason="user_speech")]
assert events == [BidiBargeInEvent()]
assert turn_state.response_id is not None
tru_events = model._convert_gemini_live_event(
@@ -848,7 +848,7 @@ async def test_event_conversion(model):
speech_started = {"type": "input_audio_buffer.speech_started", "item_id": "speech"}
tru_events = model._convert_openai_event(speech_started)
exp_events = [
BidiBargeInEvent("user_speech"),
BidiBargeInEvent(),
BidiTranscriptStartEvent("user", "speech"),
]
assert tru_events == exp_events
@@ -1251,7 +1251,7 @@ async def test_disabled_transcription_does_not_associate_audio_with_missing_tran
]
tru_events = [event for native in native_events for event in model._convert_openai_event(native) or []]
exp_events = [
BidiBargeInEvent("user_speech"),
BidiBargeInEvent(),
BidiResponseStartEvent("a"),
BidiResponseStartEvent("b"),
]
@@ -2000,7 +2000,7 @@ async def test_native_acknowledgments_correlate_inputs_and_late_transcripts(mode
if event == BidiResponseStopEvent("b"):
break
assert tru_events == [
BidiBargeInEvent("user_speech"),
BidiBargeInEvent(),
BidiTranscriptStartEvent("user", content_id="speech"),
BidiResponseStartEvent("a"),
BidiResponseStopEvent("a"),
@@ -2046,7 +2046,7 @@ async def test_receive_defers_response_during_speech(model, mock_websocket, bloc
reader = model.receive()
try:
await anext(reader)
assert await anext(reader) == BidiBargeInEvent("user_speech")
assert await anext(reader) == BidiBargeInEvent()
await model.send(BidiMessage(content=[block]))
mock_websocket.send.assert_awaited_once()
@@ -86,7 +86,7 @@ from strands.bidi.types.events import _normalize_role
{"transcript": "Hello", "role": "assistant", "content_id": "t1"},
"bidi_transcript_block",
),
(BidiBargeInEvent, {"reason": "user_speech"}, "bidi_barge_in"),
(BidiBargeInEvent, {}, "bidi_barge_in"),
(
BidiResponseStopEvent,
{"response_id": "r1"},
@@ -99,7 +99,7 @@ from strands.bidi.types.events import _normalize_role
),
(
BidiConnectionStopEvent,
{"connection_id": "c1", "reason": "complete"},
{"connection_id": "c1", "reason": "user_request"},
"bidi_connection_stop",
),
],