From 79ce5270f8cc929333bd4af6fedb840a5bdcac26 Mon Sep 17 00:00:00 2001 From: wujin Date: Thu, 30 Jul 2026 09:47:28 +0800 Subject: [PATCH] [agentserver] Add Voice Live Bridge to Invocations Add a typed Voice Live Bridge Protocol 1.0 submodule at azure.ai.agentserver.invocations.voice over the existing invocations_ws transport. Include the complete protocol runtime, immutable models, response and session helpers, bounded concurrency and cleanup, DTMF, handoff, history mutation, tests, docs, samples, APIView, and package metadata. Keep tracing aligned with raw invocations_ws: emit structured close diagnostics without framework-owned connection or turn spans, and expose only content-free Voice protocol metrics through OpenTelemetry. Preserve application-selected WebSocket close codes, package identity, MyPy-clean samples, and per-test telemetry isolation without changing agentserver-core. --- .../CHANGELOG.md | 20 + .../MANIFEST.in | 2 +- .../README.md | 58 +- .../azure-ai-agentserver-invocations/api.md | 608 ++++++ .../api.metadata.yml | 4 +- .../ai/agentserver/invocations/_invocation.py | 5 + .../agentserver/invocations/_invocation_ws.py | 27 +- .../agentserver/invocations/voice/__init__.py | 67 + .../ai/agentserver/invocations/voice/_host.py | 1708 +++++++++++++++++ .../agentserver/invocations/voice/_models.py | 272 +++ .../invocations/voice/_protocol.py | 500 +++++ .../agentserver/invocations/voice/_runtime.py | 814 ++++++++ .../dev_requirements.txt | 3 +- .../docs/voice-live-bridge-sdk-design.md | 365 ++++ .../docs/voice-live-bridge.md | 211 ++ .../pyproject.toml | 10 +- .../basic_voice_agent/basic_voice_agent.py | 28 + .../basic_voice_agent/requirements.txt | 1 + .../samples/resilient_langgraph/agent.py | 8 +- .../samples/resilient_langgraph/app.py | 12 +- .../samples/resilient_multiturn/app.py | 8 +- .../samples/resilient_research/agent.py | 8 +- .../samples/resilient_research/app.py | 8 +- .../ws_bidirectional_streaming_agent.py | 27 +- .../tests/conftest.py | 11 +- .../tests/test_ws_close_event.py | 20 + .../tests/voice/test_voice_host.py | 1662 ++++++++++++++++ .../tests/voice/test_voice_protocol.py | 303 +++ sdk/agentserver/cspell.yaml | 23 + shared_requirements.txt | 1 + 30 files changed, 6753 insertions(+), 41 deletions(-) create mode 100644 sdk/agentserver/azure-ai-agentserver-invocations/azure/ai/agentserver/invocations/voice/__init__.py create mode 100644 sdk/agentserver/azure-ai-agentserver-invocations/azure/ai/agentserver/invocations/voice/_host.py create mode 100644 sdk/agentserver/azure-ai-agentserver-invocations/azure/ai/agentserver/invocations/voice/_models.py create mode 100644 sdk/agentserver/azure-ai-agentserver-invocations/azure/ai/agentserver/invocations/voice/_protocol.py create mode 100644 sdk/agentserver/azure-ai-agentserver-invocations/azure/ai/agentserver/invocations/voice/_runtime.py create mode 100644 sdk/agentserver/azure-ai-agentserver-invocations/docs/voice-live-bridge-sdk-design.md create mode 100644 sdk/agentserver/azure-ai-agentserver-invocations/docs/voice-live-bridge.md create mode 100644 sdk/agentserver/azure-ai-agentserver-invocations/samples/basic_voice_agent/basic_voice_agent.py create mode 100644 sdk/agentserver/azure-ai-agentserver-invocations/samples/basic_voice_agent/requirements.txt create mode 100644 sdk/agentserver/azure-ai-agentserver-invocations/tests/voice/test_voice_host.py create mode 100644 sdk/agentserver/azure-ai-agentserver-invocations/tests/voice/test_voice_protocol.py diff --git a/sdk/agentserver/azure-ai-agentserver-invocations/CHANGELOG.md b/sdk/agentserver/azure-ai-agentserver-invocations/CHANGELOG.md index 3d2dcd5d0504..26fdf96dcb2a 100644 --- a/sdk/agentserver/azure-ai-agentserver-invocations/CHANGELOG.md +++ b/sdk/agentserver/azure-ai-agentserver-invocations/CHANGELOG.md @@ -2,6 +2,19 @@ ## 1.0.0b8 (Unreleased) +### Features Added + +- Added Invocations package identity to the combined `x-platform-server` value. +- Added the preview `azure.ai.agentserver.invocations.voice` submodule, a typed + implementation of Voice Live Bridge Protocol `1.0` over the existing + `invocations_ws` transport. +- Added `VoiceAgentServerHost`, immutable Voice events, ordered multi-item text + output, proactive admission, cancellation and terminal arbitration, DTMF, + handoff, history mutation, and session controls without exposing wire frames. +- Added exact-message deduplication, bounded callback coordination and cleanup, + cooperative cancellation, content-free protocol metrics, strict protocol + validation, and per-connection replay-free state. + ### Samples - Added samples showing how to build crash-resilient invocation agents on top of the new core resilient-task primitive: `resilient_multiturn` (suspend/resume conversation), `resilient_langgraph` (real-time streaming LangGraph integration with crash recovery + steering), and `resilient_research` (multi-stage research loop with checkpointing). See the [Resilient Task Developer Guide](https://github.com/Azure/azure-sdk-for-python/blob/main/sdk/agentserver/azure-ai-agentserver-core/docs/tasks-guide.md) for the underlying API. @@ -9,10 +22,17 @@ ### Bugs Fixed - The cancel (`POST /invocations/{id}/cancel`) and get (`GET /invocations/{id}`) endpoints now resolve the session id consistently with the invoke endpoint, so custom cancel/get handlers can reliably look up per-session state. +- Preserved an application-selected WebSocket close code in structured close + diagnostics instead of recording a normal `1000` after the handler returned. ### Other Changes - Bumped the minimum `azure-ai-agentserver-core` dependency to `>=2.0.0b9`. +- Voice now ships in the Invocations distribution and shares its package version + and release artifact; no separate Voice package or server identity is required. +- Voice follows the existing `invocations_ws` tracing behavior: the transport + emits structured close diagnostics but creates no framework-owned connection + or turn spans. ## 1.0.0b7 (2026-07-22) diff --git a/sdk/agentserver/azure-ai-agentserver-invocations/MANIFEST.in b/sdk/agentserver/azure-ai-agentserver-invocations/MANIFEST.in index cd83a6c13bfa..3bf9c8e2ce09 100644 --- a/sdk/agentserver/azure-ai-agentserver-invocations/MANIFEST.in +++ b/sdk/agentserver/azure-ai-agentserver-invocations/MANIFEST.in @@ -1,7 +1,7 @@ include *.md include LICENSE recursive-include tests *.py -recursive-include samples *.py *.md +recursive-include samples *.py *.md *.txt include azure/__init__.py include azure/ai/__init__.py include azure/ai/agentserver/__init__.py diff --git a/sdk/agentserver/azure-ai-agentserver-invocations/README.md b/sdk/agentserver/azure-ai-agentserver-invocations/README.md index 6769c2b2ca16..06c8bcc2f3f2 100644 --- a/sdk/agentserver/azure-ai-agentserver-invocations/README.md +++ b/sdk/agentserver/azure-ai-agentserver-invocations/README.md @@ -5,6 +5,10 @@ The `azure-ai-agentserver-invocations` package provides the invocation protocol - **HTTP** (`invocations` protocol) — `POST /invocations`, `GET /invocations/{id}`, `POST /invocations/{id}/cancel`, `GET /invocations/docs/openapi.json`, `GET /invocations/docs/asyncapi.{json,yaml}`. - **WebSocket** (`invocations_ws` protocol) — full-duplex streaming at `/invocations_ws`, registered with `@app.ws_handler`. +The package also includes the preview +`azure.ai.agentserver.invocations.voice` submodule: a typed implementation of +Voice Live Bridge Protocol `1.0` on the existing `/invocations_ws` transport. + ## Getting started ### Install the package @@ -30,6 +34,15 @@ This automatically installs `azure-ai-agentserver-core` as a dependency. - `@app.cancel_invocation_handler` — Optional. Handles `POST /invocations/{id}/cancel`. - `@app.ws_handler` — Optional. Handles WebSocket connections at `/invocations_ws`. +### Voice Live Bridge submodule + +`VoiceAgentServerHost` derives from `InvocationAgentServerHost` and owns the +`/invocations_ws` route for exact Voice Live Bridge Protocol `1.0`. It exposes +typed async callbacks, immutable inbound events, response and item helpers, +terminal arbitration, DTMF, handoff, history mutation, and session controls. +Voice Live continues to own audio, speech recognition, synthesis, voice +activity detection, turn-taking, and barge-in. + ### Protocol endpoints | Method | Route | Required | Description | @@ -296,12 +309,16 @@ app.run() - Calls `await websocket.accept()` before invoking your handler. - Runs WebSocket Ping/Pong keep-alive in the background — disabled by default; enable by setting the `WS_KEEPALIVE_INTERVAL` environment variable (auto-injected by AgentService into hosted-agent containers). Set the value to `0` to disable. Frames are sent at the WebSocket protocol layer (RFC 6455 opcode `0x9`/`0xA`) by the underlying Hypercorn server, which keeps the connection alive across upstream proxy / load-balancer idle timeouts without any extra application traffic. - Closes the connection cleanly on handler return (close code `1000`) or maps an uncaught handler exception to close code `1011`. -- Emits a structured close-event log line carrying `azure.ai.agentserver.invocations_ws.session_id`, `azure.ai.agentserver.invocations_ws.close_code`, and `azure.ai.agentserver.invocations_ws.duration_ms`. The same fields are recorded as OpenTelemetry span attributes so the connection lifetime is visible end-to-end. -- Inherits `/readiness`, OpenTelemetry export, graceful shutdown, and the `x-platform-server` identity header from `azure-ai-agentserver-core`. +- Emits a structured close-event log line carrying `azure.ai.agentserver.invocations_ws.session_id`, `azure.ai.agentserver.invocations_ws.close_code`, and `azure.ai.agentserver.invocations_ws.duration_ms`. +- Inherits `/readiness`, OpenTelemetry export configuration, and graceful shutdown from `azure-ai-agentserver-core`. ### Per-connection tracing -A WebSocket connection is wrapped by the SDK in a single connection-scoped `websocket_session` OpenTelemetry span. The span carries the GenAI semantic-convention attributes plus `azure.ai.agentserver.invocations_ws.session_id`, `close_code`, and `duration_ms`. Any child spans your handler opens — e.g. via `opentelemetry.trace.get_tracer(...).start_as_current_span(...)` — are automatically parented to the connection span. +`invocations_ws` does not create a framework-owned connection span. Application +protocols and handlers own any spans they need, while the transport reports its +connection outcome through the structured close-event log described above. The +typed Voice submodule follows this same tracing behavior and does not add +connection or turn spans. ### Handler signature @@ -323,6 +340,40 @@ The handler receives a Starlette [`WebSocket`][starlette-ws] and returns `None`. | `1011` | Handler raised an unhandled exception (mapped by the SDK). | | `4000`-`4999` | Application-defined codes (set by the handler via `await websocket.close(code=...)` — surfaced unchanged to the client). | +## Typed Voice Live Bridge (preview) + +No additional distribution is required. Import the typed protocol from the +Invocations child namespace: + +```python +from azure.ai.agentserver.invocations.voice import ( + UserMessageEvent, + VoiceAgentServerHost, + VoiceResponse, + VoiceSession, +) + +app = VoiceAgentServerHost() + + +@app.on_user_message +async def answer( + session: VoiceSession, + event: UserMessageEvent, + response: VoiceResponse, +) -> None: + del session + await response.send_text(f"You said: {event.text}") + + +app.run() +``` + +The host owns Bridge framing, IDs, ordering, callback coordination, bounded +connection state, and terminal races. Application code remains text-in and +text-out. See the [Voice Live Bridge guide](https://github.com/Azure/azure-sdk-for-python/blob/main/sdk/agentserver/azure-ai-agentserver-invocations/docs/voice-live-bridge.md) for +streaming, DTMF, handoff, proactive responses, privacy, and troubleshooting. + ## Troubleshooting ### Reporting issues @@ -339,6 +390,7 @@ Visit the [Samples](https://github.com/Azure/azure-sdk-for-python/tree/main/sdk/ | [async_invoke_agent](https://github.com/Azure/azure-sdk-for-python/tree/main/sdk/agentserver/azure-ai-agentserver-invocations/samples/async_invoke_agent/) | Long-running operations with polling and cancellation | | [ws_invoke_agent](https://github.com/Azure/azure-sdk-for-python/tree/main/sdk/agentserver/azure-ai-agentserver-invocations/samples/ws_invoke_agent/) | Combined `POST /invocations` (HTTP) and `/invocations_ws` (WebSocket) host | | [ws_bidirectional_streaming_agent](https://github.com/Azure/azure-sdk-for-python/tree/main/sdk/agentserver/azure-ai-agentserver-invocations/samples/ws_bidirectional_streaming_agent/) | Full-duplex `/invocations_ws` agent: concurrent token streams + mid-flight cancel (relies on the SDK's WS protocol Ping/Pong keep-alive, not application-level heartbeats) | +| [basic_voice_agent](https://github.com/Azure/azure-sdk-for-python/tree/main/sdk/agentserver/azure-ai-agentserver-invocations/samples/basic_voice_agent/) | Typed Voice Live Bridge `1.0` text-in/text-out agent | ## Contributing diff --git a/sdk/agentserver/azure-ai-agentserver-invocations/api.md b/sdk/agentserver/azure-ai-agentserver-invocations/api.md index 8593e6499ec0..b53ee3f0947f 100644 --- a/sdk/agentserver/azure-ai-agentserver-invocations/api.md +++ b/sdk/agentserver/azure-ai-agentserver-invocations/api.md @@ -29,4 +29,612 @@ namespace azure.ai.agentserver.invocations def ws_handler(self, fn: WSHandler) -> WSHandler: ... +namespace azure.ai.agentserver.invocations.voice + + @dataclass(frozen=True) + class azure.ai.agentserver.invocations.voice.BargeInEvent: + heard_text: str + item_id: Optional[str] + response_id: str + + def __delattr__() -> None: ... + + def __eq__() -> None: ... + + def __hash__() -> None: ... + + def __init__( + response_id: str, + heard_text: str, + item_id: str | None + ) -> None: ... + + def __repr__() -> None: ... + + def __setattr__() -> None: ... + + + @dataclass(frozen=True) + class azure.ai.agentserver.invocations.voice.ConversationHistoryItem: + content: tuple[Union[InputTextPart, InputImagePart], Ellipsis] + item_id: str + role: Literal["user"] = user + + def __delattr__() -> None: ... + + def __eq__() -> None: ... + + def __hash__() -> None: ... + + def __init__(item_id: str, content: tuple) -> None: ... + + def __repr__() -> None: ... + + def __setattr__() -> None: ... + + + @dataclass(frozen=True) + class azure.ai.agentserver.invocations.voice.ConversationItemCreateEvent: + item: ConversationHistoryItem + previous_item_id: Optional[str] + request_id: str + + def __delattr__() -> None: ... + + def __eq__() -> None: ... + + def __hash__() -> None: ... + + def __init__( + request_id: str, + item: ConversationHistoryItem, + previous_item_id: str | None + ) -> None: ... + + def __repr__() -> None: ... + + def __setattr__() -> None: ... + + + @dataclass(frozen=True) + class azure.ai.agentserver.invocations.voice.ConversationItemDeleteEvent: + item_id: str + request_id: str + + def __delattr__() -> None: ... + + def __eq__() -> None: ... + + def __hash__() -> None: ... + + def __init__(request_id: str, item_id: str) -> None: ... + + def __repr__() -> None: ... + + def __setattr__() -> None: ... + + + @dataclass(frozen=True) + class azure.ai.agentserver.invocations.voice.DtmfCollectedEvent: + collection_id: str + completion_reason: str + digits: str + item_id: str + + def __delattr__() -> None: ... + + def __eq__() -> None: ... + + def __hash__() -> None: ... + + def __init__( + item_id: str, + collection_id: str, + digits: str, + completion_reason: str + ) -> None: ... + + def __repr__() -> None: ... + + def __setattr__() -> None: ... + + + @dataclass(frozen=True) + class azure.ai.agentserver.invocations.voice.DtmfCollectionCancelledEvent: + collection_id: str + reason: str + + def __delattr__() -> None: ... + + def __eq__() -> None: ... + + def __hash__() -> None: ... + + def __init__(collection_id: str, reason: str) -> None: ... + + def __repr__() -> None: ... + + def __setattr__() -> None: ... + + + @dataclass(frozen=True) + class azure.ai.agentserver.invocations.voice.DtmfCollectionRejectedEvent: + collection_id: str + reason: str + + def __delattr__() -> None: ... + + def __eq__() -> None: ... + + def __hash__() -> None: ... + + def __init__(collection_id: str, reason: str) -> None: ... + + def __repr__() -> None: ... + + def __setattr__() -> None: ... + + + @dataclass(frozen=True) + class azure.ai.agentserver.invocations.voice.DtmfKeyEvent: + digit: str + + def __delattr__() -> None: ... + + def __eq__() -> None: ... + + def __hash__() -> None: ... + + def __init__(digit: str) -> None: ... + + def __repr__() -> None: ... + + def __setattr__() -> None: ... + + + @dataclass(frozen=True) + class azure.ai.agentserver.invocations.voice.HandoffFailedEvent: + code: str + item_id: str + message: Optional[str] + target: str + + def __delattr__() -> None: ... + + def __eq__() -> None: ... + + def __hash__() -> None: ... + + def __init__( + item_id: str, + target: str, + code: str, + message: str | None + ) -> None: ... + + def __repr__() -> None: ... + + def __setattr__() -> None: ... + + + @dataclass(frozen=True) + class azure.ai.agentserver.invocations.voice.InputImagePart: + alt: Optional[str] + image_ref: str + mime_type: str + type: Literal["input_image"] = input_image + + def __delattr__() -> None: ... + + def __eq__() -> None: ... + + def __hash__() -> None: ... + + def __init__( + image_ref: str, + mime_type: str, + alt: str | None + ) -> None: ... + + def __repr__() -> None: ... + + def __setattr__() -> None: ... + + + @dataclass(frozen=True) + class azure.ai.agentserver.invocations.voice.InputTextPart: + text: str + type: Literal["input_text"] = input_text + + def __delattr__() -> None: ... + + def __eq__() -> None: ... + + def __hash__() -> None: ... + + def __init__(text: str) -> None: ... + + def __repr__() -> None: ... + + def __setattr__() -> None: ... + + + @dataclass(frozen=True) + class azure.ai.agentserver.invocations.voice.ResponseCancellationOutcome: + heard_text: str + item_id: Optional[str] + kind: Literal["cancelled", "barge_in"] + response_id: str + + def __delattr__() -> None: ... + + def __eq__() -> None: ... + + def __hash__() -> None: ... + + def __init__( + response_id: str, + kind: Literal, + heard_text: str, + item_id: str | None + ) -> None: ... + + def __repr__() -> None: ... + + def __setattr__() -> None: ... + + + @dataclass(frozen=True) + class azure.ai.agentserver.invocations.voice.ResponseTimeoutEvent: + item_ids: Optional[tuple[str, Ellipsis]] + response_id: Optional[str] + stage: str + + def __delattr__() -> None: ... + + def __eq__() -> None: ... + + def __hash__() -> None: ... + + def __init__( + stage: str, + response_id: str | None, + item_ids: tuple + ) -> None: ... + + def __repr__() -> None: ... + + def __setattr__() -> None: ... + + + @dataclass(frozen=True) + class azure.ai.agentserver.invocations.voice.ResponseTimeouts: + first_output_ms: int + idle_ms: int + max_duration_ms: int + + def __delattr__() -> None: ... + + def __eq__() -> None: ... + + def __hash__() -> None: ... + + def __init__( + first_output_ms: int, + idle_ms: int, + max_duration_ms: int + ) -> None: ... + + def __repr__() -> None: ... + + def __setattr__() -> None: ... + + + @dataclass(frozen=True) + class azure.ai.agentserver.invocations.voice.SessionEndEvent: + reason: str + + def __delattr__() -> None: ... + + def __eq__() -> None: ... + + def __hash__() -> None: ... + + def __init__(reason: str) -> None: ... + + def __repr__() -> None: ... + + def __setattr__() -> None: ... + + + @dataclass(frozen=True) + class azure.ai.agentserver.invocations.voice.SessionStartEvent: + caller: Optional[Mapping[str, Any]] + greeting: Optional[str] + no_input_timeout_ms: Optional[int] + protocol_version: str + reconnect: bool + response_timeouts: ResponseTimeouts + + def __delattr__() -> None: ... + + def __eq__() -> None: ... + + def __hash__() -> None: ... + + def __init__( + protocol_version: str, + reconnect: bool, + response_timeouts: ResponseTimeouts, + greeting: str | None, + no_input_timeout_ms: int | None, + caller: Mapping + ) -> None: ... + + def __repr__() -> None: ... + + def __setattr__() -> None: ... + + + @dataclass(frozen=True) + class azure.ai.agentserver.invocations.voice.UserMessageEvent: + property text: str # Read-only + content: tuple[Union[InputTextPart, InputImagePart], Ellipsis] + item_id: str + + def __delattr__() -> None: ... + + def __eq__() -> None: ... + + def __hash__() -> None: ... + + def __init__(item_id: str, content: tuple) -> None: ... + + def __repr__() -> None: ... + + def __setattr__() -> None: ... + + + @dataclass(frozen=True) + class azure.ai.agentserver.invocations.voice.UserNoInputEvent: + count: int + item_id: str + + def __delattr__() -> None: ... + + def __eq__() -> None: ... + + def __hash__() -> None: ... + + def __init__(item_id: str, count: int) -> None: ... + + def __repr__() -> None: ... + + def __setattr__() -> None: ... + + + @dataclass(frozen=True) + class azure.ai.agentserver.invocations.voice.UserSpeechStartedEvent: + + def __delattr__() -> None: ... + + def __eq__() -> None: ... + + def __hash__() -> None: ... + + def __init__() -> None: ... + + def __repr__() -> None: ... + + def __setattr__() -> None: ... + + + class azure.ai.agentserver.invocations.voice.VoiceAgentServerHost(InvocationAgentServerHost): + property routes: list[BaseRoute] # Read-only + property ws_ping_interval: float # Read-only + + def __init__(self, **kwargs: Any) -> None: ... + + def cancel_invocation_handler(self, fn: Callable[[Request], Awaitable[Response]]) -> Callable[[Request], Awaitable[Response]]: ... + + def get_asyncapi_spec_json(self) -> Optional[dict[str, Any]]: ... + + def get_asyncapi_spec_yaml(self) -> Optional[str]: ... + + def get_invocation_handler(self, fn: Callable[[Request], Awaitable[Response]]) -> Callable[[Request], Awaitable[Response]]: ... + + def get_openapi_spec(self) -> Optional[dict[str, Any]]: ... + + def invoke_handler(self, fn: Callable[[Request], Awaitable[Response]]) -> Callable[[Request], Awaitable[Response]]: ... + + def on_barge_in(self, fn: BargeInCallback) -> BargeInCallback: ... + + def on_conversation_item_create(self, fn: ConversationItemCreateCallback) -> ConversationItemCreateCallback: ... + + def on_conversation_item_delete(self, fn: ConversationItemDeleteCallback) -> ConversationItemDeleteCallback: ... + + def on_dtmf_collected(self, fn: DtmfCollectedCallback) -> DtmfCollectedCallback: ... + + def on_dtmf_collection_cancelled(self, fn: DtmfCollectionCancelledCallback) -> DtmfCollectionCancelledCallback: ... + + def on_dtmf_collection_rejected(self, fn: DtmfCollectionRejectedCallback) -> DtmfCollectionRejectedCallback: ... + + def on_dtmf_key(self, fn: DtmfKeyCallback) -> DtmfKeyCallback: ... + + def on_handoff_failed(self, fn: HandoffFailedCallback) -> HandoffFailedCallback: ... + + def on_response_timeout(self, fn: ResponseTimeoutCallback) -> ResponseTimeoutCallback: ... + + def on_session_end(self, fn: SessionEndCallback) -> SessionEndCallback: ... + + def on_session_start(self, fn: SessionStartCallback) -> SessionStartCallback: ... + + def on_user_message(self, fn: UserMessageCallback) -> UserMessageCallback: ... + + def on_user_no_input(self, fn: UserNoInputCallback) -> UserNoInputCallback: ... + + def on_user_speech_started(self, fn: UserSpeechStartedCallback) -> UserSpeechStartedCallback: ... + + def ws_handler(self, fn: Any) -> Any: ... + + + class azure.ai.agentserver.invocations.voice.VoiceBridgeConnectionClosedError(RuntimeError): + + + class azure.ai.agentserver.invocations.voice.VoiceBridgeProtocolError(ValueError): + + def __init__( + self, + message: str, + *, + close_code: int = 1002 + ) -> None: ... + + + class azure.ai.agentserver.invocations.voice.VoiceCancellationToken: + property is_cancelled: bool # Read-only + + def __init__(self) -> None: ... + + async def wait(self) -> None: ... + + + class azure.ai.agentserver.invocations.voice.VoiceProactiveResponseDroppedError(RuntimeError): + + def __init__( + self, + response_id: str, + reason: str + ) -> None: ... + + + class azure.ai.agentserver.invocations.voice.VoiceResponse: + property cancellation: VoiceCancellationToken # Read-only + property in_reply_to: tuple[str, ] | None # Read-only + property is_cancel_pending: bool # Read-only + property is_terminal: bool # Read-only + property is_wire_opened: bool # Read-only + property response_id: str # Read-only + + def __init__(self) -> None: ... + + async def cancel( + self, + *, + reason: str | None = ... + ) -> ResponseCancellationOutcome: ... + + async def collect_dtmf( + self, + *, + initial_timeout_ms: int, + inter_digit_timeout_ms: int, + max_digits: int, + terminator: str | None = ... + ) -> str: ... + + async def decline( + self, + *, + reason: str | None = ... + ) -> None: ... + + async def done(self) -> None: ... + + async def fail( + self, + *, + code: str, + message: str + ) -> None: ... + + async def handoff( + self, + *, + message: str | None = ..., + target: str + ) -> None: ... + + def new_text_item(self) -> VoiceTextItem: ... + + async def send_text( + self, + text: str, + *, + voice: Mapping[str, Any] | None = ... + ) -> None: ... + + async def send_text_delta( + self, + delta: str, + *, + voice: Mapping[str, Any] | None = ... + ) -> None: ... + + async def send_text_done( + self, + *, + voice: Mapping[str, Any] | None = ... + ) -> None: ... + + + class azure.ai.agentserver.invocations.voice.VoiceSession: + property caller: Mapping[str, Any] | None # Read-only + property greeting: str | None # Read-only + property no_input_timeout_ms: int | None # Read-only + property reconnect: bool # Read-only + property response_timeouts: ResponseTimeouts # Read-only + + def __init__(self) -> None: ... + + async def cancel_dtmf_collection(self, collection_id: str) -> None: ... + + async def end_call( + self, + *, + mode: Literal["drain", "immediate"] = "drain", + reason: str + ) -> None: ... + + async def report_error( + self, + *, + code: str, + message: str + ) -> None: ... + + async def start_proactive_response( + self, + *, + admission_timeout_ms: int = 60000, + supersede_key: str | None = ... + ) -> VoiceResponse: ... + + + class azure.ai.agentserver.invocations.voice.VoiceTextItem: + property item_id: str # Read-only + + def __init__(self) -> None: ... + + async def send_text( + self, + text: str, + *, + voice: Mapping[str, Any] | None = ... + ) -> None: ... + + async def send_text_delta( + self, + delta: str, + *, + voice: Mapping[str, Any] | None = ... + ) -> None: ... + + async def send_text_done( + self, + *, + voice: Mapping[str, Any] | None = ... + ) -> None: ... + + ``` \ No newline at end of file diff --git a/sdk/agentserver/azure-ai-agentserver-invocations/api.metadata.yml b/sdk/agentserver/azure-ai-agentserver-invocations/api.metadata.yml index 570a1d2dada6..1ea5deacac9e 100644 --- a/sdk/agentserver/azure-ai-agentserver-invocations/api.metadata.yml +++ b/sdk/agentserver/azure-ai-agentserver-invocations/api.metadata.yml @@ -1,3 +1,3 @@ -apiMdSha256: 28f40971c3ed93df127abaf7faf714a62065ef4440aaf511b9269c584811520c +apiMdSha256: f7024c9da5f751d754d7c6da7c25d84fe15c6565d08ad50ba26aa5a2aa2db9da parserVersion: 0.3.30 -pythonVersion: 3.11.9 +pythonVersion: 3.13.12 diff --git a/sdk/agentserver/azure-ai-agentserver-invocations/azure/ai/agentserver/invocations/_invocation.py b/sdk/agentserver/azure-ai-agentserver-invocations/azure/ai/agentserver/invocations/_invocation.py index e14184a6e2a1..909af0489f85 100644 --- a/sdk/agentserver/azure-ai-agentserver-invocations/azure/ai/agentserver/invocations/_invocation.py +++ b/sdk/agentserver/azure-ai-agentserver-invocations/azure/ai/agentserver/invocations/_invocation.py @@ -23,6 +23,7 @@ from azure.ai.agentserver.core import ( # pylint: disable=no-name-in-module AgentServerHost, FoundryAgentRequestContext, + build_server_version, create_error_response, reset_request_context, set_request_context, @@ -38,6 +39,7 @@ from ._constants import InvocationConstants from ._invocation_ws import _WSHandlerMixin +from ._version import VERSION as _INVOCATIONS_VERSION logger = logging.getLogger("azure.ai.agentserver") @@ -273,6 +275,9 @@ def __init__( # Merge with any routes from sibling mixins via cooperative init existing = list(kwargs.pop("routes", None) or []) super().__init__(routes=existing + invocation_routes, **kwargs) + self.register_server_version( + build_server_version("azure-ai-agentserver-invocations", _INVOCATIONS_VERSION) + ) # --- Invocations startup configuration logging --- logger.info( diff --git a/sdk/agentserver/azure-ai-agentserver-invocations/azure/ai/agentserver/invocations/_invocation_ws.py b/sdk/agentserver/azure-ai-agentserver-invocations/azure/ai/agentserver/invocations/_invocation_ws.py index ffdaad5ad2e8..62f0d0b579da 100644 --- a/sdk/agentserver/azure-ai-agentserver-invocations/azure/ai/agentserver/invocations/_invocation_ws.py +++ b/sdk/agentserver/azure-ai-agentserver-invocations/azure/ai/agentserver/invocations/_invocation_ws.py @@ -24,7 +24,7 @@ import logging import time import uuid -from collections.abc import Awaitable, Callable +from collections.abc import Awaitable, Callable, MutableMapping from typing import TYPE_CHECKING, Any, Optional from starlette.websockets import WebSocket, WebSocketDisconnect, WebSocketState @@ -42,12 +42,16 @@ # ``InvocationAgentServerHost`` MRO actually inherits ``AgentServerHost``, # which keeps the diamond out of the runtime class graph. if TYPE_CHECKING: - _MixinBase = AgentServerHost + class _MixinBase(AgentServerHost): + """Static-checking base that exposes AgentServerHost attributes.""" else: - _MixinBase = object + class _MixinBase: + """Runtime base that keeps AgentServerHost out of the mixin MRO.""" logger = logging.getLogger("azure.ai.agentserver") +_APPLICATION_CLOSE_CODE = "azure.ai.agentserver.invocations_ws.application_close_code" + WSHandler = Callable[[WebSocket], Awaitable[None]] @@ -221,6 +225,16 @@ async def _ws_endpoint(self, websocket: WebSocket) -> None: close_code: int = InvocationsWSConstants.CLOSE_NORMAL handler_exc: Optional[BaseException] = None + original_send = websocket.send + + async def _tracked_send(message: MutableMapping[str, Any]) -> None: + if message.get("type") == "websocket.close": + websocket.scope[_APPLICATION_CLOSE_CODE] = int( + message.get("code", InvocationsWSConstants.CLOSE_NORMAL) + ) + await original_send(message) + + websocket.send = _tracked_send # type: ignore[method-assign] try: close_code, handler_exc = await self._invoke_user_handler(websocket, session_id) except BaseException as exc: # pylint: disable=broad-exception-caught @@ -275,7 +289,12 @@ async def _invoke_user_handler( raise RuntimeError("_invoke_user_handler called with no registered ws_handler") try: await ws_fn(websocket) - return InvocationsWSConstants.CLOSE_NORMAL, None + return int( + websocket.scope.get( + _APPLICATION_CLOSE_CODE, + InvocationsWSConstants.CLOSE_NORMAL, + ) + ), None except WebSocketDisconnect as exc: # Client (or proxy) closed first — surface their code, not 1011. return ( diff --git a/sdk/agentserver/azure-ai-agentserver-invocations/azure/ai/agentserver/invocations/voice/__init__.py b/sdk/agentserver/azure-ai-agentserver-invocations/azure/ai/agentserver/invocations/voice/__init__.py new file mode 100644 index 000000000000..5a2b8965893b --- /dev/null +++ b/sdk/agentserver/azure-ai-agentserver-invocations/azure/ai/agentserver/invocations/voice/__init__.py @@ -0,0 +1,67 @@ +# --------------------------------------------------------- +# Copyright (c) Microsoft Corporation. All rights reserved. +# --------------------------------------------------------- +"""Typed Voice Live Bridge Protocol host on the Invocations WebSocket transport.""" + +from ._host import VoiceAgentServerHost +from ._models import ( + BargeInEvent, + ConversationHistoryItem, + ConversationItemCreateEvent, + ConversationItemDeleteEvent, + DtmfCollectedEvent, + DtmfCollectionCancelledEvent, + DtmfCollectionRejectedEvent, + DtmfKeyEvent, + HandoffFailedEvent, + InputImagePart, + InputTextPart, + ResponseCancellationOutcome, + ResponseTimeoutEvent, + ResponseTimeouts, + SessionEndEvent, + SessionStartEvent, + UserContentPart, + UserMessageEvent, + UserNoInputEvent, + UserSpeechStartedEvent, +) +from ._protocol import ( + VoiceBridgeConnectionClosedError, + VoiceBridgeProtocolError, + VoiceProactiveResponseDroppedError, +) +from ._runtime import VoiceCancellationToken, VoiceResponse, VoiceSession, VoiceTextItem +from .._version import VERSION + +__all__ = [ + "BargeInEvent", + "ConversationHistoryItem", + "ConversationItemCreateEvent", + "ConversationItemDeleteEvent", + "DtmfCollectedEvent", + "DtmfCollectionCancelledEvent", + "DtmfCollectionRejectedEvent", + "DtmfKeyEvent", + "HandoffFailedEvent", + "InputImagePart", + "InputTextPart", + "ResponseCancellationOutcome", + "ResponseTimeoutEvent", + "ResponseTimeouts", + "SessionEndEvent", + "SessionStartEvent", + "UserContentPart", + "UserMessageEvent", + "UserNoInputEvent", + "UserSpeechStartedEvent", + "VoiceAgentServerHost", + "VoiceBridgeConnectionClosedError", + "VoiceBridgeProtocolError", + "VoiceCancellationToken", + "VoiceProactiveResponseDroppedError", + "VoiceResponse", + "VoiceSession", + "VoiceTextItem", +] +__version__ = VERSION diff --git a/sdk/agentserver/azure-ai-agentserver-invocations/azure/ai/agentserver/invocations/voice/_host.py b/sdk/agentserver/azure-ai-agentserver-invocations/azure/ai/agentserver/invocations/voice/_host.py new file mode 100644 index 000000000000..fec9fd33ed48 --- /dev/null +++ b/sdk/agentserver/azure-ai-agentserver-invocations/azure/ai/agentserver/invocations/voice/_host.py @@ -0,0 +1,1708 @@ +# --------------------------------------------------------- +# Copyright (c) Microsoft Corporation. All rights reserved. +# --------------------------------------------------------- +"""Typed Voice Live bridge host built on the invocations_ws transport.""" + +from __future__ import annotations + +import asyncio # pylint: disable=do-not-import-asyncio +import hashlib +import inspect +import logging +import time +from collections import OrderedDict +from collections.abc import Awaitable, Callable +from dataclasses import dataclass +from typing import Any, Optional + +from opentelemetry import metrics + +from .._invocation import InvocationAgentServerHost +from .._version import VERSION + +from ._models import ( + BargeInEvent, + ConversationItemCreateEvent, + ConversationItemDeleteEvent, + DtmfCollectedEvent, + DtmfCollectionCancelledEvent, + DtmfCollectionRejectedEvent, + DtmfKeyEvent, + HandoffFailedEvent, + ResponseCancellationOutcome, + ResponseTimeoutEvent, + SessionEndEvent, + SessionStartEvent, + UserMessageEvent, + UserNoInputEvent, + UserSpeechStartedEvent, +) +from ._protocol import ( + PROTOCOL_VERSION, + VoiceBridgeConnectionClosedError, + VoiceBridgeProtocolError, + VoiceProactiveResponseDroppedError, + canonical_payload, + decode_frame, + encode_frame, + new_id, + optional_string, + parse_conversation_item_create, + parse_conversation_item_delete, + parse_dtmf, + parse_dtmf_collection_cancelled, + parse_dtmf_collection_rejected, + parse_handoff_failed, + parse_response_timeout, + parse_session_start, + parse_user_message, + require_positive_int, + require_prefixed_id, + require_string, + safe_code, +) +from ._runtime import VoiceResponse, VoiceSession + +logger = logging.getLogger("azure.ai.agentserver") +_METER = metrics.get_meter("Azure.AI.AgentServer.Invocations.Voice", VERSION) +_ACTIVATION_COUNTER = _METER.create_counter("azure.ai.agentserver.invocations.voice.activations") +_CALLBACK_DURATION = _METER.create_histogram( + "azure.ai.agentserver.invocations.voice.callback.duration", + unit="ms", +) +_CALLBACK_ERROR_COUNTER = _METER.create_counter("azure.ai.agentserver.invocations.voice.callback.errors") +_FIRST_OUTPUT_DURATION = _METER.create_histogram( + "azure.ai.agentserver.invocations.voice.first_output.duration", + unit="ms", +) +_TERMINAL_COUNTER = _METER.create_counter("azure.ai.agentserver.invocations.voice.response.terminals") +_PROTOCOL_VIOLATION_COUNTER = _METER.create_counter("azure.ai.agentserver.invocations.voice.protocol.violations") +_ACTIVE_CONNECTIONS = _METER.create_up_down_counter("azure.ai.agentserver.invocations.voice.active_connections") +_CLOSE_CODE_COUNTER = _METER.create_counter("azure.ai.agentserver.invocations.voice.close_codes") + +# VoiceResponse and _VoiceConnection are the public/internal halves of one +# runtime and intentionally drive each other's private terminal hooks. +# pylint: disable=protected-access + +SessionStartCallback = Callable[[VoiceSession, SessionStartEvent], Awaitable[None]] +UserMessageCallback = Callable[[VoiceSession, UserMessageEvent, VoiceResponse], Awaitable[None]] +UserNoInputCallback = Callable[[VoiceSession, UserNoInputEvent, VoiceResponse], Awaitable[None]] +UserSpeechStartedCallback = Callable[[VoiceSession, UserSpeechStartedEvent], Awaitable[None]] +DtmfKeyCallback = Callable[[VoiceSession, DtmfKeyEvent], Awaitable[None]] +DtmfCollectedCallback = Callable[[VoiceSession, DtmfCollectedEvent, VoiceResponse], Awaitable[None]] +DtmfCollectionRejectedCallback = Callable[[VoiceSession, DtmfCollectionRejectedEvent], Awaitable[None]] +DtmfCollectionCancelledCallback = Callable[[VoiceSession, DtmfCollectionCancelledEvent], Awaitable[None]] +HandoffFailedCallback = Callable[[VoiceSession, HandoffFailedEvent, VoiceResponse], Awaitable[None]] +ConversationItemCreateCallback = Callable[[VoiceSession, ConversationItemCreateEvent], Awaitable[None]] +ConversationItemDeleteCallback = Callable[[VoiceSession, ConversationItemDeleteEvent], Awaitable[None]] +BargeInCallback = Callable[[VoiceSession, BargeInEvent], Awaitable[None]] +ResponseTimeoutCallback = Callable[[VoiceSession, ResponseTimeoutEvent], Awaitable[None]] +SessionEndCallback = Callable[[VoiceSession, SessionEndEvent], Awaitable[None]] + +_MAX_CALLBACK_QUEUE = 128 +_MAX_SEEN_MESSAGES = 4096 +_MAX_RECENT_RESPONSES = 64 +_MAX_RESOLVED_PREFIXES = 64 +_MAX_PENDING_PROACTIVE = 16 +_CLEANUP_TIMEOUT_SECONDS = 5.0 +_AGENT_TO_BRIDGE_TYPES = { + "session.ready", + "session.rejected", + "conversation.item.created", + "conversation.item.deleted", + "conversation.item.failed", + "response.created", + "response.none", + "response.output_text.delta", + "response.output_text.done", + "response.done", + "response.cancel", + "handoff", + "end_call", + "dtmf.collect", + "dtmf.collect.cancel", + "error", +} + + +@dataclass(frozen=True) +class _CallbackWork: + kind: str + event: Any + callback: Callable[..., Awaitable[None]] | None + response: VoiceResponse | None = None + item_id: str | None = None + request_id: str | None = None + success_type: str | None = None + + +class VoiceAgentServerHost(InvocationAgentServerHost): # pylint: disable=too-many-instance-attributes + """AgentServer host implementing Voice Live Bridge Protocol 1.0.""" + + def __init__(self, **kwargs: Any) -> None: + self._on_session_start: Optional[SessionStartCallback] = None + self._on_user_message: Optional[UserMessageCallback] = None + self._on_user_no_input: Optional[UserNoInputCallback] = None + self._on_user_speech_started: Optional[UserSpeechStartedCallback] = None + self._on_dtmf_key: Optional[DtmfKeyCallback] = None + self._on_dtmf_collected: Optional[DtmfCollectedCallback] = None + self._on_dtmf_collection_rejected: Optional[DtmfCollectionRejectedCallback] = None + self._on_dtmf_collection_cancelled: Optional[DtmfCollectionCancelledCallback] = None + self._on_handoff_failed: Optional[HandoffFailedCallback] = None + self._on_conversation_item_create: Optional[ConversationItemCreateCallback] = None + self._on_conversation_item_delete: Optional[ConversationItemDeleteCallback] = None + self._on_barge_in: Optional[BargeInCallback] = None + self._on_response_timeout: Optional[ResponseTimeoutCallback] = None + self._on_session_end: Optional[SessionEndCallback] = None + super().__init__(**kwargs) + InvocationAgentServerHost.ws_handler(self, self._handle_voice_websocket) + + def ws_handler(self, fn: Any) -> Any: + """Reject raw handler replacement on the typed host. + + :param fn: Raw handler that cannot be registered. + :type fn: Any + :raises RuntimeError: Always. + """ + del fn + raise RuntimeError( + "VoiceAgentServerHost owns /invocations_ws. " + "Use typed voice callbacks, or InvocationAgentServerHost for a custom protocol." + ) + + def on_session_start(self, fn: SessionStartCallback) -> SessionStartCallback: + """Register the optional application-start callback. + + :param fn: Async callback invoked before readiness. + :type fn: Callable[[VoiceSession, SessionStartEvent], Awaitable[None]] + :return: Registered callback. + :rtype: Callable[[VoiceSession, SessionStartEvent], Awaitable[None]] + """ + self._on_session_start = self._register_once("on_session_start", self._on_session_start, fn) + return fn + + def on_user_message(self, fn: UserMessageCallback) -> UserMessageCallback: + """Register the required completed-user-turn callback. + + :param fn: Async callback receiving ordered content and a lazy response. + :type fn: Callable[[VoiceSession, UserMessageEvent, VoiceResponse], Awaitable[None]] + :return: Registered callback. + :rtype: Callable[[VoiceSession, UserMessageEvent, VoiceResponse], Awaitable[None]] + """ + self._on_user_message = self._register_once("on_user_message", self._on_user_message, fn) + return fn + + def on_user_no_input(self, fn: UserNoInputCallback) -> UserNoInputCallback: + """Register the optional bridge-generated silence-turn callback. + + :param fn: Async no-input callback. + :type fn: Callable[[VoiceSession, UserNoInputEvent, VoiceResponse], Awaitable[None]] + :return: Registered callback. + :rtype: Callable[[VoiceSession, UserNoInputEvent, VoiceResponse], Awaitable[None]] + """ + self._on_user_no_input = self._register_once("on_user_no_input", self._on_user_no_input, fn) + return fn + + def on_user_speech_started(self, fn: UserSpeechStartedCallback) -> UserSpeechStartedCallback: + """Register the optional advisory speech-start callback. + + :param fn: Async advisory callback. + :type fn: Callable[[VoiceSession, UserSpeechStartedEvent], Awaitable[None]] + :return: Registered callback. + :rtype: Callable[[VoiceSession, UserSpeechStartedEvent], Awaitable[None]] + """ + self._on_user_speech_started = self._register_once("on_user_speech_started", self._on_user_speech_started, fn) + return fn + + def on_dtmf_key(self, fn: DtmfKeyCallback) -> DtmfKeyCallback: + """Register the optional raw DTMF key callback. + + :param fn: Async session-signal callback. + :type fn: Callable[[VoiceSession, DtmfKeyEvent], Awaitable[None]] + :return: Registered callback. + :rtype: Callable[[VoiceSession, DtmfKeyEvent], Awaitable[None]] + """ + self._on_dtmf_key = self._register_once("on_dtmf_key", self._on_dtmf_key, fn) + return fn + + def on_dtmf_collected(self, fn: DtmfCollectedCallback) -> DtmfCollectedCallback: + """Register the optional completed DTMF collection turn callback. + + :param fn: Async response-producing callback. + :type fn: Callable[[VoiceSession, DtmfCollectedEvent, VoiceResponse], Awaitable[None]] + :return: Registered callback. + :rtype: Callable[[VoiceSession, DtmfCollectedEvent, VoiceResponse], Awaitable[None]] + """ + self._on_dtmf_collected = self._register_once("on_dtmf_collected", self._on_dtmf_collected, fn) + return fn + + def on_dtmf_collection_rejected(self, fn: DtmfCollectionRejectedCallback) -> DtmfCollectionRejectedCallback: + """Register the optional DTMF collection rejection callback. + + :param fn: Async collection-control callback. + :type fn: Callable[[VoiceSession, DtmfCollectionRejectedEvent], Awaitable[None]] + :return: Registered callback. + :rtype: Callable[[VoiceSession, DtmfCollectionRejectedEvent], Awaitable[None]] + """ + self._on_dtmf_collection_rejected = self._register_once( + "on_dtmf_collection_rejected", self._on_dtmf_collection_rejected, fn + ) + return fn + + def on_dtmf_collection_cancelled(self, fn: DtmfCollectionCancelledCallback) -> DtmfCollectionCancelledCallback: + """Register the optional DTMF collection cancellation callback. + + :param fn: Async collection-control callback. + :type fn: Callable[[VoiceSession, DtmfCollectionCancelledEvent], Awaitable[None]] + :return: Registered callback. + :rtype: Callable[[VoiceSession, DtmfCollectionCancelledEvent], Awaitable[None]] + """ + self._on_dtmf_collection_cancelled = self._register_once( + "on_dtmf_collection_cancelled", self._on_dtmf_collection_cancelled, fn + ) + return fn + + def on_handoff_failed(self, fn: HandoffFailedCallback) -> HandoffFailedCallback: + """Register the optional handoff recovery-turn callback. + + :param fn: Async response-producing recovery callback. + :type fn: Callable[[VoiceSession, HandoffFailedEvent, VoiceResponse], Awaitable[None]] + :return: Registered callback. + :rtype: Callable[[VoiceSession, HandoffFailedEvent, VoiceResponse], Awaitable[None]] + """ + self._on_handoff_failed = self._register_once("on_handoff_failed", self._on_handoff_failed, fn) + return fn + + def on_conversation_item_create(self, fn: ConversationItemCreateCallback) -> ConversationItemCreateCallback: + """Register the optional durable history-create callback. + + :param fn: Async history mutation callback. + :type fn: Callable[[VoiceSession, ConversationItemCreateEvent], Awaitable[None]] + :return: Registered callback. + :rtype: Callable[[VoiceSession, ConversationItemCreateEvent], Awaitable[None]] + """ + self._on_conversation_item_create = self._register_once( + "on_conversation_item_create", self._on_conversation_item_create, fn + ) + return fn + + def on_conversation_item_delete(self, fn: ConversationItemDeleteCallback) -> ConversationItemDeleteCallback: + """Register the optional durable history-delete callback. + + :param fn: Async history mutation callback. + :type fn: Callable[[VoiceSession, ConversationItemDeleteEvent], Awaitable[None]] + :return: Registered callback. + :rtype: Callable[[VoiceSession, ConversationItemDeleteEvent], Awaitable[None]] + """ + self._on_conversation_item_delete = self._register_once( + "on_conversation_item_delete", self._on_conversation_item_delete, fn + ) + return fn + + def on_barge_in(self, fn: BargeInCallback) -> BargeInCallback: + """Register the optional caller-interruption callback. + + :param fn: Async callback invoked after response cancellation. + :type fn: Callable[[VoiceSession, BargeInEvent], Awaitable[None]] + :return: Registered callback. + :rtype: Callable[[VoiceSession, BargeInEvent], Awaitable[None]] + """ + self._on_barge_in = self._register_once("on_barge_in", self._on_barge_in, fn) + return fn + + def on_response_timeout(self, fn: ResponseTimeoutCallback) -> ResponseTimeoutCallback: + """Register the optional response-timeout callback. + + :param fn: Async callback invoked after local work is tombstoned. + :type fn: Callable[[VoiceSession, ResponseTimeoutEvent], Awaitable[None]] + :return: Registered callback. + :rtype: Callable[[VoiceSession, ResponseTimeoutEvent], Awaitable[None]] + """ + self._on_response_timeout = self._register_once("on_response_timeout", self._on_response_timeout, fn) + return fn + + def on_session_end(self, fn: SessionEndCallback) -> SessionEndCallback: + """Register the optional session-end callback. + + :param fn: Async callback invoked during bounded teardown. + :type fn: Callable[[VoiceSession, SessionEndEvent], Awaitable[None]] + :return: Registered callback. + :rtype: Callable[[VoiceSession, SessionEndEvent], Awaitable[None]] + """ + self._on_session_end = self._register_once("on_session_end", self._on_session_end, fn) + return fn + + @staticmethod + def _register_once(name: str, current: Any, fn: Any) -> Any: + if current is not None: + raise RuntimeError(f"{name} callback is already registered") + if not inspect.iscoroutinefunction(fn): + raise TypeError(f"{name} expects an async function") + return fn + + async def _handle_voice_websocket(self, websocket: Any) -> None: + connection = _VoiceConnection( + websocket=websocket, + on_session_start=self._on_session_start, + on_user_message=self._on_user_message, + on_user_no_input=self._on_user_no_input, + on_user_speech_started=self._on_user_speech_started, + on_dtmf_key=self._on_dtmf_key, + on_dtmf_collected=self._on_dtmf_collected, + on_dtmf_collection_rejected=self._on_dtmf_collection_rejected, + on_dtmf_collection_cancelled=self._on_dtmf_collection_cancelled, + on_handoff_failed=self._on_handoff_failed, + on_conversation_item_create=self._on_conversation_item_create, + on_conversation_item_delete=self._on_conversation_item_delete, + on_barge_in=self._on_barge_in, + on_response_timeout=self._on_response_timeout, + on_session_end=self._on_session_end, + ) + _ACTIVE_CONNECTIONS.add(1) + try: + await connection.run() + finally: + _ACTIVE_CONNECTIONS.add(-1) + + +class _VoiceConnection: # pylint: disable=too-many-instance-attributes,too-many-public-methods + """Per-WebSocket protocol runtime.""" + + def __init__( + self, + *, + websocket: Any, + on_session_start: Optional[SessionStartCallback], + on_user_message: Optional[UserMessageCallback], + on_user_no_input: Optional[UserNoInputCallback], + on_user_speech_started: Optional[UserSpeechStartedCallback], + on_dtmf_key: Optional[DtmfKeyCallback], + on_dtmf_collected: Optional[DtmfCollectedCallback], + on_dtmf_collection_rejected: Optional[DtmfCollectionRejectedCallback], + on_dtmf_collection_cancelled: Optional[DtmfCollectionCancelledCallback], + on_handoff_failed: Optional[HandoffFailedCallback], + on_conversation_item_create: Optional[ConversationItemCreateCallback], + on_conversation_item_delete: Optional[ConversationItemDeleteCallback], + on_barge_in: Optional[BargeInCallback], + on_response_timeout: Optional[ResponseTimeoutCallback], + on_session_end: Optional[SessionEndCallback], + ) -> None: + self._websocket = websocket + self._on_session_start = on_session_start + self._on_user_message = on_user_message + self._on_user_no_input = on_user_no_input + self._on_user_speech_started = on_user_speech_started + self._on_dtmf_key = on_dtmf_key + self._on_dtmf_collected = on_dtmf_collected + self._on_dtmf_collection_rejected = on_dtmf_collection_rejected + self._on_dtmf_collection_cancelled = on_dtmf_collection_cancelled + self._on_handoff_failed = on_handoff_failed + self._on_conversation_item_create = on_conversation_item_create + self._on_conversation_item_delete = on_conversation_item_delete + self._on_barge_in = on_barge_in + self._on_response_timeout = on_response_timeout + self._on_session_end = on_session_end + self._send_lock = asyncio.Lock() + self._state_lock = asyncio.Lock() + self._callback_queue: asyncio.Queue[_CallbackWork | None] = asyncio.Queue(maxsize=_MAX_CALLBACK_QUEUE) + self._callback_worker: asyncio.Task[None] | None = None + self._active_customer_task: asyncio.Task[None] | None = None + self._active_release: asyncio.Event | None = None + self._cleanup_tasks: set[asyncio.Task[None]] = set() + self._session: VoiceSession | None = None + self._active_response: VoiceResponse | None = None + self._pending_turns: OrderedDict[str, VoiceResponse] = OrderedDict() + self._resolved_input_prefixes: OrderedDict[tuple[str, ...], tuple[VoiceResponse, bool]] = OrderedDict() + self._recent_responses: OrderedDict[str, VoiceResponse] = OrderedDict() + self._seen_response_ids: set[str] = set() + self._terminal_response_ids: set[str] = set() + self._seen_messages: OrderedDict[str, str] = OrderedDict() + self._seen_input_ids: set[str] = set() + self._playback_outcomes: set[str] = set() + self._abandoned_proactive_cancels: set[str] = set() + self._cancel_waiters: dict[str, asyncio.Future[ResponseCancellationOutcome]] = {} + self._pending_proactive: OrderedDict[ + str, + tuple[VoiceResponse, asyncio.Future[tuple[bool, str]]], + ] = OrderedDict() + self._dtmf_collections: dict[str, str] = {} + self._dtmf_cancel_pending: set[str] = set() + self._recent_dtmf_cancel_races: OrderedDict[str, None] = OrderedDict() + self._response_start_ns: dict[str, int] = {} + self._first_output_recorded: set[str] = set() + self._activation_recorded = False + self._close_recorded = False + self._ready = False + self._closed = False + self._ending = False + + @property + def ending(self) -> bool: + return self._ending or self._closed + + async def run(self) -> None: + """Activate the protocol and run the sole receive pump.""" + graceful_end = False + if not await self._activate(): + await self._shutdown_runtime(drain_callbacks=False) + return + try: + while not self._closed: + payload = await self._receive_with_worker_supervision() + if payload is None: + break + if not await self._dispatch(payload): + graceful_end = True + break + except VoiceBridgeProtocolError as exc: + _PROTOCOL_VIOLATION_COUNTER.add(1, {"close_code": exc.close_code}) + logger.warning("Voice bridge protocol violation: %s", exc) + await self._close(code=exc.close_code, reason="Protocol error") + except Exception as exc: # pylint: disable=broad-exception-caught + logger.error("Voice bridge runtime failed: %s", type(exc).__name__) + await self._close(code=1011, reason="Internal server error") + finally: + await self._shutdown_runtime(drain_callbacks=graceful_end) + + async def _receive_with_worker_supervision(self) -> dict[str, Any] | None: + worker = self._callback_worker + assert worker is not None + receive_task = asyncio.create_task( + self._receive_payload(), + name="voice_receive", + ) + try: + done, _ = await asyncio.wait((receive_task, worker), return_when=asyncio.FIRST_COMPLETED) + except BaseException: + receive_task.cancel() + await asyncio.gather(receive_task, return_exceptions=True) + raise + if worker in done: + receive_task.cancel() + await asyncio.gather(receive_task, return_exceptions=True) + if worker.cancelled(): + raise RuntimeError("Voice callback coordinator was cancelled unexpectedly") + error = worker.exception() + if error is not None: + raise error + raise RuntimeError("Voice callback coordinator stopped unexpectedly") + return receive_task.result() + + async def send(self, message_type: str, **fields: Any) -> None: + """Serialize one SDK-owned application frame. + + :param message_type: Wire message discriminator. + :type message_type: str + """ + if self._closed: + raise VoiceBridgeConnectionClosedError("The voice connection is closed") + async with self._send_lock: + if self._closed: + raise VoiceBridgeConnectionClosedError("The voice connection is closed") + response_id = fields.get("response_id") + if isinstance(response_id, str): + async with self._state_lock: + if response_id in self._terminal_response_ids: + raise VoiceBridgeConnectionClosedError("The voice response is terminal") + # Perform the WebSocket write OUTSIDE _state_lock. Holding the state + # lock across a send that stalls on outbound backpressure would block + # the receive pump from acquiring _state_lock to process barge_in, + # response.timeout, or session termination — defeating the full-duplex + # cancellation path. _send_lock still serializes writes; the terminal + # check above rejects the common case, and the bridge drops any single + # frame that races past a terminal it has already recorded. + await self._websocket.send_text(encode_frame(message_type, **fields)) + self._record_first_output(message_type, fields) + + def _remember_resolved_prefix_locked( + self, prefix: tuple[str, ...], value: tuple[VoiceResponse, bool] + ) -> None: + # Bounded LRU: these entries only reconcile a late response.timeout that + # references a prefix the SDK already resolved before the bridge saw the + # response.created — a single-round-trip window. Retaining every resolved + # prefix (each holding a full VoiceResponse with accumulated output text) + # for the life of the call would grow without bound; cap it and evict the + # oldest, since a timeout arriving that far out of order is stale. + self._resolved_input_prefixes.pop(prefix, None) + self._resolved_input_prefixes[prefix] = value + while len(self._resolved_input_prefixes) > _MAX_RESOLVED_PREFIXES: + self._resolved_input_prefixes.popitem(last=False) + + async def open_response(self, response_id: str, in_reply_to: tuple[str, ...] | None) -> bool: + """Consume one explicit input prefix and announce a reply response. + + :param response_id: SDK-allocated response identifier. + :type response_id: str + :param in_reply_to: Non-empty ordered input prefix. + :type in_reply_to: tuple[str, ...] or None + :return: ``False`` when a bridge terminal already won; otherwise ``True``. + :rtype: bool + """ + self._ensure_ready() + if not in_reply_to: + raise RuntimeError("Reply response requires a non-empty in_reply_to prefix") + async with self._state_lock: + if response_id in self._terminal_response_ids: + return False + active = self._active_response + if active is None or active.response_id != response_id: + raise VoiceBridgeConnectionClosedError("The response is no longer active") + pending_ids = tuple(self._pending_turns) + if pending_ids[: len(in_reply_to)] != in_reply_to: + raise RuntimeError("in_reply_to must be an ordered prefix of pending inputs") + self._remember_resolved_prefix_locked(in_reply_to, (active, True)) + for item_id in in_reply_to: + self._pending_turns.pop(item_id, None) + await self.send( + "response.created", + response_id=response_id, + in_reply_to=list(in_reply_to), + ) + return True + + async def decline_response(self, in_reply_to: tuple[str, ...], reason: str | None) -> None: + """Consume one explicit input prefix without opening a response. + + :param in_reply_to: Non-empty ordered input prefix. + :type in_reply_to: tuple[str, ...] + :param reason: Optional open-enum decline reason. + :type reason: str or None + """ + self._ensure_ready() + response_id: str | None = None + async with self._state_lock: + active = self._active_response + if ( + active is not None + and active.in_reply_to == in_reply_to + and active.response_id in self._terminal_response_ids + ): + return + pending_ids = tuple(self._pending_turns) + if pending_ids[: len(in_reply_to)] != in_reply_to: + raise RuntimeError("in_reply_to must be an ordered prefix of pending inputs") + if in_reply_to: + response_id = self._pending_turns[in_reply_to[0]].response_id + self._remember_resolved_prefix_locked( + in_reply_to, (self._pending_turns[in_reply_to[0]], False) + ) + for item_id in in_reply_to: + self._pending_turns.pop(item_id, None) + # Claim the terminal for this input prefix BEFORE emitting response.none. + # response.none carries no response_id on the wire, so send()'s + # response-scoped terminal check cannot guard it. Without claiming here, a + # response.timeout processed by the receive pump while this send is + # suspended (send lock or transport backpressure) would let response.none + # reach the wire after the bridge terminal already won. Serializing the + # claim with timeout arbitration ensures only the winning terminal is emitted. + if response_id is not None: + async with self._state_lock: + if response_id in self._terminal_response_ids: + return + self._terminal_response_ids.add(response_id) + fields: dict[str, Any] = {"in_reply_to": list(in_reply_to)} + if reason is not None: + fields["reason"] = reason + await self.send("response.none", **fields) + if response_id is not None: + self._record_first_output_for_response(response_id) + self._record_terminal(response_id, "none") + + async def begin_cancel(self, response_id: str, reason: str | None) -> asyncio.Future[ResponseCancellationOutcome]: + """Register cancellation arbitration before sending ``response.cancel``. + + :param response_id: Open response to cancel. + :type response_id: str + :param reason: Optional open-enum cancellation reason. + :type reason: str or None + :return: Future resolved by the winning playback terminal. + :rtype: asyncio.Future[ResponseCancellationOutcome] + """ + self._ensure_ready() + async with self._state_lock: + response = self._find_response_locked(response_id) + if response is None or not response.is_wire_opened: + raise VoiceBridgeConnectionClosedError("The response is not open") + if response_id in self._cancel_waiters: + raise RuntimeError("Response cancellation is already pending") + future: asyncio.Future[ResponseCancellationOutcome] = asyncio.get_running_loop().create_future() + self._cancel_waiters[response_id] = future + fields: dict[str, Any] = {"response_id": response_id} + if reason is not None: + fields["reason"] = reason + try: + await self.send("response.cancel", **fields) + except BaseException: + async with self._state_lock: + self._cancel_waiters.pop(response_id, None) + raise + return future + + async def response_completed(self, response_id: str, terminal_kind: str = "done") -> None: + """Move a response into bounded playback-reconciliation state. + + :param response_id: Completed response identifier. + :type response_id: str + :param terminal_kind: Low-cardinality terminal classification. + :type terminal_kind: str + """ + async with self._state_lock: + response = self._find_response_locked(response_id) + if response is None: + return + local_terminal_won = response_id not in self._terminal_response_ids + self._terminal_response_ids.add(response_id) + self._remember_response_locked(response) + if self._active_response is response: + self._active_response = None + if local_terminal_won: + self._record_terminal(response_id, terminal_kind) + + async def register_dtmf_collection( + self, + *, + response_id: str, + collection_id: str, + max_digits: int, + terminator: str | None, + initial_timeout_ms: int, + inter_digit_timeout_ms: int, + ) -> None: + """Register and emit one response-scoped DTMF collection request. + + :keyword response_id: Open source response identifier. + :paramtype response_id: str + :keyword collection_id: SDK-allocated collection identifier. + :paramtype collection_id: str + :keyword max_digits: Positive maximum returned digit count. + :paramtype max_digits: int + :keyword terminator: Optional single DTMF terminator. + :paramtype terminator: str or None + :keyword initial_timeout_ms: Positive first-key timeout. + :paramtype initial_timeout_ms: int + :keyword inter_digit_timeout_ms: Positive inter-key timeout. + :paramtype inter_digit_timeout_ms: int + """ + self._ensure_ready() + async with self._state_lock: + if self._dtmf_collections: + raise RuntimeError("Only one DTMF collection may be pending or active") + response = self._find_response_locked(response_id) + if response is None or response.is_terminal: + raise VoiceBridgeConnectionClosedError("The source response is not open") + self._dtmf_collections[collection_id] = response_id + fields: dict[str, Any] = { + "response_id": response_id, + "collection_id": collection_id, + "max_digits": max_digits, + "initial_timeout_ms": initial_timeout_ms, + "inter_digit_timeout_ms": inter_digit_timeout_ms, + } + if terminator is not None: + fields["terminator"] = terminator + try: + await self.send("dtmf.collect", **fields) + except BaseException: + async with self._state_lock: + self._dtmf_collections.pop(collection_id, None) + raise + + async def cancel_dtmf_collection(self, collection_id: str) -> None: + """Emit explicit cancellation for one known DTMF collection. + + :param collection_id: SDK-allocated collection identifier. + :type collection_id: str + """ + self._ensure_ready() + async with self._state_lock: + if collection_id not in self._dtmf_collections: + raise RuntimeError("Unknown or completed DTMF collection_id") + if collection_id in self._dtmf_cancel_pending: + raise RuntimeError("DTMF collection cancellation is already pending") + self._dtmf_cancel_pending.add(collection_id) + try: + await self.send("dtmf.collect.cancel", collection_id=collection_id) + except BaseException: + async with self._state_lock: + self._dtmf_cancel_pending.discard(collection_id) + raise + + async def end_call(self, reason: str, mode: str) -> None: + """Emit one call terminal and seal active work. + + :param reason: Non-empty open-enum termination reason. + :type reason: str + :param mode: Closed ``drain`` or ``immediate`` mode. + :type mode: str + """ + self._ensure_ready() + async with self._state_lock: + if self._ending: + return + self._ending = True + response = self._active_response + if response is not None: + self._terminal_response_ids.add(response.response_id) + release = self._active_release + active_task = self._active_customer_task + await self.send("end_call", reason=reason, mode=mode) + if response is not None: + await response._mark_terminal() + self._record_terminal(response.response_id, "end_call") + if release is not None and active_task is not asyncio.current_task(): + release.set() + + async def start_proactive_response( + self, + *, + admission_timeout_ms: int, + supersede_key: str | None, + ) -> VoiceResponse: + """Request admission and return only after ``response.accepted``. + + :keyword admission_timeout_ms: Positive admission deadline. + :paramtype admission_timeout_ms: int + :keyword supersede_key: Optional non-empty supersession key. + :paramtype supersede_key: str or None + :return: Accepted proactive response. + :rtype: VoiceResponse + """ + self._ensure_ready() + response = VoiceResponse._create( + self, + response_id=new_id("r"), + in_reply_to=None, + wire_opened=True, + accepted=False, + ) + future: asyncio.Future[tuple[bool, str]] = asyncio.get_running_loop().create_future() + async with self._state_lock: + if len(self._pending_proactive) >= _MAX_PENDING_PROACTIVE: + raise RuntimeError("Too many proactive admission outcomes are pending") + if response.response_id in self._seen_response_ids: + raise RuntimeError("Generated proactive response_id was already used") + self._seen_response_ids.add(response.response_id) + self._pending_proactive[response.response_id] = (response, future) + fields: dict[str, Any] = { + "response_id": response.response_id, + "admission_timeout_ms": admission_timeout_ms, + } + if supersede_key is not None: + fields["supersede_key"] = supersede_key + try: + await self.send("response.created", **fields) + except BaseException: + async with self._state_lock: + self._pending_proactive.pop(response.response_id, None) + raise + try: + accepted, reason = await future + except asyncio.CancelledError: + async with self._state_lock: + self._abandoned_proactive_cancels.add(response.response_id) + try: + await self.send( + "response.cancel", + response_id=response.response_id, + reason="cancelled_by_agent", + ) + except VoiceBridgeConnectionClosedError: + pass + raise + if not accepted: + await response._mark_terminal() + raise VoiceProactiveResponseDroppedError(response.response_id, reason) + return response + + async def report_session_error(self, code: str, message: str) -> None: + """Emit a session-scoped terminal error. + + :param code: Bounded machine-readable code. + :type code: str + :param message: Bounded diagnostic message. + :type message: str + """ + self._ensure_ready() + async with self._state_lock: + if self._ending: + return + self._ending = True + response = self._active_response + if response is not None: + self._terminal_response_ids.add(response.response_id) + release = self._active_release + await self.send("error", code=code, message=message) + if response is not None: + await response._mark_terminal() + self._record_terminal(response.response_id, "session_error") + if release is not None: + release.set() + + async def _activate(self) -> bool: # pylint: disable=too-many-return-statements + try: + payload = await self._receive_payload() + except VoiceBridgeProtocolError as exc: + await self._reject("invalid_session_start", close_code=exc.close_code) + return False + if payload is None: + self._record_activation("closed") + return False + if payload.get("type") != "session.start": + await self._reject("invalid_session_start", close_code=1002) + return False + try: + event = parse_session_start(payload) + except VoiceBridgeProtocolError: + code = ( + "protocol_mismatch" if payload.get("protocol_version") != PROTOCOL_VERSION else "invalid_session_start" + ) + await self._reject(code, close_code=1002) + return False + if self._on_user_message is None: + await self._reject("startup_failed", close_code=1011) + return False + session = VoiceSession._create(self, event) + self._session = session + startup_started_ns = time.monotonic_ns() + try: + if not await self._run_session_start_callback(session, event): + return False + except Exception as exc: # pylint: disable=broad-exception-caught + logger.error("Voice session-start callback failed: %s", type(exc).__name__) + _CALLBACK_ERROR_COUNTER.add(1, {"kind": "session.start"}) + await self._reject("startup_failed", close_code=1011) + return False + finally: + if self._on_session_start is not None: + _CALLBACK_DURATION.record( + (time.monotonic_ns() - startup_started_ns) / 1_000_000, + {"kind": "session.start"}, + ) + await self.send("session.ready") + self._record_activation("ready") + self._ready = True + self._callback_worker = asyncio.create_task(self._callback_worker_loop(), name="voice_callback_coordinator") + return True + + async def _run_session_start_callback(self, session: VoiceSession, event: SessionStartEvent) -> bool: + async def _invoke_customer_callback() -> None: + if self._on_session_start is not None: + await self._on_session_start(session, event) + + receive_task = asyncio.create_task( + self._receive_payload(), + name="voice_activation_receive", + ) + customer_task = asyncio.create_task( + _invoke_customer_callback(), + name="voice_session_start", + ) + try: + done, _ = await asyncio.wait((customer_task, receive_task), return_when=asyncio.FIRST_COMPLETED) + except BaseException: + receive_task.cancel() + await asyncio.gather(receive_task, return_exceptions=True) + if not customer_task.done(): + self._schedule_customer_cleanup(customer_task) + raise + + if receive_task in done: + if not customer_task.done(): + self._schedule_customer_cleanup(customer_task) + else: + await asyncio.gather(customer_task, return_exceptions=True) + try: + early_payload = receive_task.result() + except VoiceBridgeProtocolError as exc: + await self._reject("protocol_mismatch", close_code=exc.close_code) + except Exception as exc: # pylint: disable=broad-exception-caught + logger.error("Voice activation receive failed: %s", type(exc).__name__) + self._record_activation("startup_failed") + await self._close(code=1011, reason="Internal server error") + else: + if early_payload is None: + self._record_activation("closed") + else: + await self._reject("protocol_mismatch", close_code=1008) + return False + + receive_task.cancel() + await asyncio.gather(receive_task, return_exceptions=True) + if customer_task.cancelled(): + raise RuntimeError("Voice session-start callback was cancelled") + error = customer_task.exception() + if error is not None: + raise error + return True + + async def _reject(self, code: str, *, close_code: int) -> None: + if self._closed: + return + self._record_activation(code) + if code in ("invalid_session_start", "protocol_mismatch"): + _PROTOCOL_VIOLATION_COUNTER.add(1, {"close_code": close_code}) + try: + await self.send( + "session.rejected", + code=safe_code(code, "startup_failed"), + retriable=False, + ) + finally: + await self._close(code=close_code, reason="Session rejected") + + async def _dispatch(self, payload: dict[str, Any]) -> bool: # pylint: disable=too-many-branches + message_type = payload["type"] + if message_type == "user.message": + message_event = parse_user_message(payload) + await self._enqueue_turn( + message_event.item_id, + message_event, + self._on_user_message, + "user.message", + ) + elif message_type == "user.no_input": + no_input_event = UserNoInputEvent( + item_id=require_prefixed_id(payload, "item_id", "in_"), + count=require_positive_int(payload, "count"), + ) + await self._enqueue_turn( + no_input_event.item_id, + no_input_event, + self._on_user_no_input, + "user.no_input", + ) + elif message_type == "user.speech_started": + await self._enqueue_signal(UserSpeechStartedEvent(), self._on_user_speech_started, "user.speech_started") + elif message_type == "conversation.item.create": + create_event = parse_conversation_item_create(payload) + self._enqueue_history( + create_event, + self._on_conversation_item_create, + "conversation.item.create", + "conversation.item.created", + ) + elif message_type == "conversation.item.delete": + delete_event = parse_conversation_item_delete(payload) + self._enqueue_history( + delete_event, + self._on_conversation_item_delete, + "conversation.item.delete", + "conversation.item.deleted", + ) + elif message_type == "dtmf": + dtmf_event = parse_dtmf(payload) + if isinstance(dtmf_event, DtmfKeyEvent): + await self._enqueue_signal(dtmf_event, self._on_dtmf_key, "dtmf.key") + else: + await self._consume_dtmf_collection( + dtmf_event.collection_id, + preserve_cancel_race=True, + ) + await self._enqueue_turn( + dtmf_event.item_id, + dtmf_event, + self._on_dtmf_collected, + "dtmf.collected", + ) + elif message_type == "dtmf.collect.rejected": + rejected_event = parse_dtmf_collection_rejected(payload) + await self._consume_dtmf_collection( + rejected_event.collection_id, + allow_late_cancel_rejection=rejected_event.reason == "collection_not_found", + preserve_cancel_race=rejected_event.reason != "collection_not_found", + ) + await self._enqueue_signal( + rejected_event, + self._on_dtmf_collection_rejected, + "dtmf.collect.rejected", + ) + elif message_type == "dtmf.collect.cancelled": + cancelled_event = parse_dtmf_collection_cancelled(payload) + await self._consume_dtmf_collection( + cancelled_event.collection_id, + preserve_cancel_race=cancelled_event.reason != "cancelled_by_agent", + ) + await self._enqueue_signal( + cancelled_event, + self._on_dtmf_collection_cancelled, + "dtmf.collect.cancelled", + ) + elif message_type == "handoff.failed": + handoff_event = parse_handoff_failed(payload) + await self._enqueue_turn( + handoff_event.item_id, + handoff_event, + self._on_handoff_failed, + "handoff.failed", + ) + elif message_type == "barge_in": + await self._handle_playback_terminal(payload, kind="barge_in") + elif message_type == "response.cancelled": + await self._handle_playback_terminal(payload, kind="cancelled") + elif message_type == "response.timeout": + await self._handle_response_timeout(parse_response_timeout(payload)) + elif message_type == "response.accepted": + await self._handle_response_accepted(payload) + elif message_type == "response.dropped": + await self._handle_response_dropped(payload) + elif message_type == "session.end": + await self._handle_session_end(payload) + return False + else: + if message_type == "session.start" or message_type in _AGENT_TO_BRIDGE_TYPES: + raise VoiceBridgeProtocolError( + f"{message_type} is not valid from the bridge after readiness", + close_code=1008, + ) + logger.debug("Ignoring unknown post-readiness voice message type: %s", message_type) + return True + + async def _enqueue_turn( + self, + item_id: str, + event: Any, + callback: Callable[..., Awaitable[None]] | None, + kind: str, + ) -> None: + response = VoiceResponse._create(self, in_reply_to=(item_id,)) + async with self._state_lock: + if self._ending: + raise VoiceBridgeProtocolError(f"{kind} arrived after session terminal", close_code=1008) + if item_id in self._seen_input_ids: + raise VoiceBridgeProtocolError("Input item_id was reused", close_code=1008) + self._seen_input_ids.add(item_id) + if response.response_id in self._seen_response_ids: + raise RuntimeError("Generated response_id was already used") + self._seen_response_ids.add(response.response_id) + self._pending_turns[item_id] = response + self._response_start_ns[response.response_id] = time.monotonic_ns() + self._put_work( + _CallbackWork( + kind=kind, + event=event, + callback=callback, + response=response, + item_id=item_id, + ) + ) + + async def _enqueue_signal( + self, + event: Any, + callback: Callable[..., Awaitable[None]] | None, + kind: str, + ) -> None: + if callback is not None: + self._put_work(_CallbackWork(kind=kind, event=event, callback=callback)) + + def _enqueue_history( + self, + event: ConversationItemCreateEvent | ConversationItemDeleteEvent, + callback: Callable[..., Awaitable[None]] | None, + kind: str, + success_type: str, + ) -> None: + self._put_work( + _CallbackWork( + kind=kind, + event=event, + callback=callback, + request_id=event.request_id, + success_type=success_type, + ) + ) + + async def _consume_dtmf_collection( + self, + collection_id: str, + *, + allow_late_cancel_rejection: bool = False, + preserve_cancel_race: bool = False, + ) -> None: + async with self._state_lock: + source_response_id = self._dtmf_collections.pop(collection_id, None) + if source_response_id is None: + if allow_late_cancel_rejection and collection_id in self._recent_dtmf_cancel_races: + self._recent_dtmf_cancel_races.pop(collection_id, None) + return + raise VoiceBridgeProtocolError("Unknown DTMF collection_id", close_code=1008) + cancel_pending = collection_id in self._dtmf_cancel_pending + self._dtmf_cancel_pending.discard(collection_id) + if cancel_pending and preserve_cancel_race: + self._recent_dtmf_cancel_races[collection_id] = None + while len(self._recent_dtmf_cancel_races) > _MAX_RECENT_RESPONSES: + self._recent_dtmf_cancel_races.popitem(last=False) + + def _put_work(self, work: _CallbackWork) -> None: + try: + self._callback_queue.put_nowait(work) + except asyncio.QueueFull as exc: + raise VoiceBridgeProtocolError("Voice callback queue limit exceeded", close_code=1008) from exc + + async def _callback_worker_loop(self) -> None: + while True: + work = await self._callback_queue.get() + try: + if work is None: + return + if work.response is not None: + await self._process_turn_work(work) + else: + await self._process_signal_work(work) + finally: + self._callback_queue.task_done() + + # pylint: disable=too-many-statements,too-many-branches + async def _process_turn_work(self, work: _CallbackWork) -> None: + response = work.response + assert response is not None + assert self._session is not None + if response.is_terminal: + return + + release = asyncio.Event() + async with self._state_lock: + if response.is_terminal or self._ending: + return + self._active_response = response + self._active_release = release + + if work.callback is None: + try: + await response._fail_callback() + finally: + async with self._state_lock: + if self._active_response is response: + self._active_response = None + if self._active_release is release: + self._active_release = None + if work.item_id is not None: + self._pending_turns.pop(work.item_id, None) + if response.is_wire_opened: + self._remember_response_locked(response) + return + + callback_started_ns = time.monotonic_ns() + + async def _invoke_customer_callback() -> None: + assert work.callback is not None + await work.callback(self._session, work.event, response) + + customer_task: asyncio.Task[None] = asyncio.create_task( + _invoke_customer_callback(), + name=f"voice_{work.kind}", + ) + release_task = asyncio.create_task(release.wait(), name="voice_turn_release") + async with self._state_lock: + self._active_customer_task = customer_task + try: + done, _ = await asyncio.wait((customer_task, release_task), return_when=asyncio.FIRST_COMPLETED) + if release_task in done and not customer_task.done(): + await self._schedule_customer_cleanup(customer_task) + elif customer_task.cancelled(): + if not response.is_terminal and not self.ending: + await response._fail_callback() + else: + error = customer_task.exception() + if error is not None: + terminal_race = isinstance(error, VoiceBridgeConnectionClosedError) and ( + response.is_terminal or response.response_id in self._terminal_response_ids + ) + if not terminal_race: + logger.error("Voice callback failed: %s", type(error).__name__) + _CALLBACK_ERROR_COUNTER.add(1, {"kind": work.kind}) + await response._fail_callback() + else: + await response._complete_callback() + except asyncio.CancelledError: + if not customer_task.done(): + self._schedule_customer_cleanup(customer_task) + raise + finally: + _CALLBACK_DURATION.record( + (time.monotonic_ns() - callback_started_ns) / 1_000_000, + {"kind": work.kind}, + ) + release_task.cancel() + await asyncio.gather(release_task, return_exceptions=True) + async with self._state_lock: + if self._active_customer_task is customer_task: + self._active_customer_task = None + if self._active_response is response: + self._active_response = None + if self._active_release is release: + self._active_release = None + if work.item_id is not None: + self._pending_turns.pop(work.item_id, None) + if response.is_wire_opened: + self._remember_response_locked(response) + + async def _process_signal_work(self, work: _CallbackWork) -> None: + if self._session is None: + return + if work.success_type is not None: + assert work.request_id is not None + if work.callback is None: + await self.send( + "conversation.item.failed", + request_id=work.request_id, + code="mutation_failed", + message="No history mutation callback is registered", + ) + return + callback_started_ns = time.monotonic_ns() + try: + await self._await_signal_callback(work) + except Exception as exc: # pylint: disable=broad-exception-caught + logger.error("Voice history callback failed: %s", type(exc).__name__) + _CALLBACK_ERROR_COUNTER.add(1, {"kind": work.kind}) + await self.send( + "conversation.item.failed", + request_id=work.request_id, + code="mutation_failed", + message="History mutation callback failed", + ) + else: + await self.send(work.success_type, request_id=work.request_id) + finally: + _CALLBACK_DURATION.record( + (time.monotonic_ns() - callback_started_ns) / 1_000_000, + {"kind": work.kind}, + ) + return + if work.callback is None: + return + callback_started_ns = time.monotonic_ns() + try: + await self._await_signal_callback(work) + except Exception as exc: # pylint: disable=broad-exception-caught + logger.error("Voice signal callback failed: %s", type(exc).__name__) + _CALLBACK_ERROR_COUNTER.add(1, {"kind": work.kind}) + finally: + _CALLBACK_DURATION.record( + (time.monotonic_ns() - callback_started_ns) / 1_000_000, + {"kind": work.kind}, + ) + + async def _await_signal_callback(self, work: _CallbackWork) -> None: + assert self._session is not None + assert work.callback is not None + + async def _invoke_customer_callback() -> None: + assert self._session is not None + assert work.callback is not None + await work.callback(self._session, work.event) + + customer_task: asyncio.Task[None] = asyncio.create_task( + _invoke_customer_callback(), + name=f"voice_{work.kind}", + ) + try: + await asyncio.shield(customer_task) + except asyncio.CancelledError as exc: + if not customer_task.done(): + self._schedule_customer_cleanup(customer_task) + raise + current_task = asyncio.current_task() + cancelling = getattr(current_task, "cancelling", lambda: 0) + if customer_task.cancelled() and cancelling() == 0: + raise RuntimeError("Voice signal callback was cancelled") from exc + raise + + def _schedule_customer_cleanup(self, task: asyncio.Task[None]) -> asyncio.Task[None]: + task.cancel() + + async def _bounded_cleanup() -> None: + try: + await asyncio.wait_for(asyncio.shield(task), timeout=_CLEANUP_TIMEOUT_SECONDS) + except asyncio.TimeoutError: + logger.warning("Voice callback ignored cancellation beyond cleanup deadline") + except asyncio.CancelledError: + pass + except Exception: # pylint: disable=broad-exception-caught + pass + + cleanup = asyncio.create_task(_bounded_cleanup(), name="voice_callback_cleanup") + self._cleanup_tasks.add(cleanup) + cleanup.add_done_callback(self._cleanup_tasks.discard) + return cleanup + + async def _handle_playback_terminal(self, payload: dict[str, Any], *, kind: str) -> None: + response_id = require_prefixed_id(payload, "response_id", "r_") + heard_text = require_string(payload, "heard_text") + item_id = optional_string(payload, "item_id") + if item_id is not None and (not item_id.startswith("it_") or len(item_id) <= 3): + raise VoiceBridgeProtocolError("Playback item_id must start with it_", close_code=1008) + async with self._state_lock: + if response_id in self._playback_outcomes: + return + response = self._find_response_locked(response_id) + if response is None: + if response_id in self._pending_proactive: + raise VoiceBridgeProtocolError( + f"{kind} is invalid before proactive response.accepted", + close_code=1008, + ) + if response_id in self._seen_response_ids: + return + raise VoiceBridgeProtocolError("Unknown playback response_id", close_code=1008) + if item_id is not None and not response._owns_item_id(item_id): + raise VoiceBridgeProtocolError("Playback item_id does not belong to response_id", close_code=1008) + self._playback_outcomes.add(response_id) + waiter = self._cancel_waiters.pop(response_id, None) + abandoned = response_id in self._abandoned_proactive_cancels + self._abandoned_proactive_cancels.discard(response_id) + if kind == "cancelled" and waiter is None and not abandoned: + self._playback_outcomes.remove(response_id) + raise VoiceBridgeProtocolError( + "response.cancelled requires a pending response.cancel", + close_code=1008, + ) + active = self._active_response is response + release = self._active_release if active and not response.is_cancel_pending else None + playback_terminal_won = response_id not in self._terminal_response_ids + self._terminal_response_ids.add(response_id) + await response._mark_terminal() + outcome = ResponseCancellationOutcome( + response_id=response_id, + kind="barge_in" if kind == "barge_in" else "cancelled", + heard_text=heard_text, + item_id=item_id, + ) + if waiter is not None and not waiter.done(): + waiter.set_result(outcome) + if playback_terminal_won: + self._record_terminal(response_id, kind) + if release is not None: + release.set() + if kind == "barge_in" and self._on_barge_in is not None: + self._put_work( + _CallbackWork( + kind="barge_in", + event=BargeInEvent( + response_id=response_id, + heard_text=heard_text, + item_id=item_id, + ), + callback=self._on_barge_in, + ) + ) + + async def _handle_response_timeout(self, event: ResponseTimeoutEvent) -> None: + responses: list[VoiceResponse] = [] + timeout_metric_winners: set[str] = set() + release: asyncio.Event | None = None + cancel_waiter: asyncio.Future[ResponseCancellationOutcome] | None = None + async with self._state_lock: + if event.response_id is not None: + response = self._find_response_locked(event.response_id) + if event.response_id in self._playback_outcomes: + return + if response is None: + if event.response_id in self._seen_response_ids: + return + raise VoiceBridgeProtocolError("Unknown response.timeout response_id", close_code=1008) + responses.append(response) + self._playback_outcomes.add(event.response_id) + if event.response_id not in self._terminal_response_ids: + timeout_metric_winners.add(event.response_id) + self._terminal_response_ids.add(event.response_id) + if self._active_response is response: + release = self._active_release + cancel_waiter = self._cancel_waiters.pop(event.response_id, None) + else: + assert event.item_ids is not None + pending_ids = tuple(self._pending_turns) + if pending_ids[: len(event.item_ids)] == event.item_ids: + for item_id in event.item_ids: + response = self._pending_turns.pop(item_id) + if response.is_wire_opened: + raise VoiceBridgeProtocolError( + "response.timeout item_ids referenced an open response", + close_code=1008, + ) + responses.append(response) + if response.response_id not in self._terminal_response_ids: + timeout_metric_winners.add(response.response_id) + self._terminal_response_ids.add(response.response_id) + if self._active_response is response: + release = self._active_release + else: + remaining = event.item_ids + while remaining: + resolved = next( + ( + (prefix, response, opened_response) + for prefix, (response, opened_response) in self._resolved_input_prefixes.items() + if remaining[: len(prefix)] == prefix + ), + None, + ) + if resolved is not None: + prefix, response, opened_response = resolved + self._resolved_input_prefixes.pop(prefix, None) + if response not in responses: + responses.append(response) + if response.response_id not in self._terminal_response_ids: + timeout_metric_winners.add(response.response_id) + self._terminal_response_ids.add(response.response_id) + if self._active_response is response: + release = self._active_release + if opened_response: + self._playback_outcomes.add(response.response_id) + cancel_waiter = self._cancel_waiters.pop(response.response_id, cancel_waiter) + remaining = remaining[len(prefix) :] + continue + + pending_ids = tuple(self._pending_turns) + if pending_ids[: len(remaining)] != remaining: + raise VoiceBridgeProtocolError( + "response.timeout item_ids do not match the pending or just-resolved prefix", + close_code=1008, + ) + for item_id in remaining: + response = self._pending_turns.pop(item_id) + if response.is_wire_opened: + raise VoiceBridgeProtocolError( + "response.timeout item_ids referenced an open response", + close_code=1008, + ) + if response not in responses: + responses.append(response) + if response.response_id not in self._terminal_response_ids: + timeout_metric_winners.add(response.response_id) + self._terminal_response_ids.add(response.response_id) + if self._active_response is response: + release = self._active_release + remaining = () + for response in responses: + await response._mark_terminal() + if response.response_id in timeout_metric_winners: + self._record_terminal(response.response_id, "timeout") + if cancel_waiter is not None and not cancel_waiter.done(): + cancel_waiter.set_exception(VoiceBridgeConnectionClosedError("Response terminated by timeout")) + if release is not None: + release.set() + if self._on_response_timeout is not None: + self._put_work( + _CallbackWork( + kind="response.timeout", + event=event, + callback=self._on_response_timeout, + ) + ) + + async def _handle_response_accepted(self, payload: dict[str, Any]) -> None: + response_id = require_prefixed_id(payload, "response_id", "r_") + async with self._state_lock: + pending = self._pending_proactive.get(response_id) + if pending is None: + raise VoiceBridgeProtocolError("Unknown proactive response_id", close_code=1008) + response, future = pending + if self._active_response is not None and not self._active_response.is_terminal: + raise VoiceBridgeProtocolError( + "Proactive response accepted while another response is active", + close_code=1008, + ) + self._pending_proactive.pop(response_id, None) + self._active_response = response + self._response_start_ns[response_id] = time.monotonic_ns() + await response._mark_accepted() + if not future.done(): + future.set_result((True, "")) + + async def _handle_response_dropped(self, payload: dict[str, Any]) -> None: + response_id = require_prefixed_id(payload, "response_id", "r_") + reason = safe_code(require_string(payload, "reason", non_empty=True), "dropped") + async with self._state_lock: + pending = self._pending_proactive.pop(response_id, None) + if pending is None: + raise VoiceBridgeProtocolError("Unknown proactive response_id", close_code=1008) + response, future = pending + self._abandoned_proactive_cancels.discard(response_id) + self._terminal_response_ids.add(response_id) + await response._mark_terminal() + self._record_terminal(response_id, "dropped") + if not future.done(): + future.set_result((False, reason)) + + async def _handle_session_end(self, payload: dict[str, Any]) -> None: + event = SessionEndEvent(reason=require_string(payload, "reason", non_empty=True)) + async with self._state_lock: + self._ending = True + responses = list(self._pending_turns.values()) + self._pending_turns.clear() + if self._active_response is not None and self._active_response not in responses: + responses.append(self._active_response) + self._terminal_response_ids.update(response.response_id for response in responses) + release = self._active_release + for response in responses: + was_terminal = response.is_terminal + await response._mark_terminal() + if not was_terminal: + self._record_terminal(response.response_id, "session_end") + self._fail_helper_waiters("Voice session ended") + if release is not None: + release.set() + if self._on_session_end is not None: + self._put_work(_CallbackWork(kind="session.end", event=event, callback=self._on_session_end)) + + async def _receive_payload(self) -> dict[str, Any] | None: + while True: + message = await self._websocket.receive() + message_type = message.get("type") + if message_type == "websocket.disconnect": + self._record_close(int(message.get("code") or 1006)) + self._closed = True + return None + if message_type != "websocket.receive": + raise VoiceBridgeProtocolError("Unexpected ASGI WebSocket event") + if message.get("bytes") is not None: + _PROTOCOL_VIOLATION_COUNTER.add(1, {"close_code": 1003}) + await self._close(code=1003, reason="Binary data is unsupported") + return None + frame = message.get("text") + if not isinstance(frame, str): + raise VoiceBridgeProtocolError("Voice bridge requires JSON text frames") + payload = decode_frame(frame) + message_id = require_string(payload, "id", non_empty=True) + # Store a fixed-size digest, not the full canonical JSON: with a 1 MB + # frame limit and up to _MAX_SEEN_MESSAGES entries, retaining the + # complete canonical text per id would let a peer pin gigabytes of + # memory (even via unique ignored message types) before the entry + # count cap is reached. A sha256 digest bounds each entry to 64 chars. + digest = hashlib.sha256(canonical_payload(payload).encode("utf-8")).hexdigest() + previous = self._seen_messages.get(message_id) + if previous is not None: + if previous != digest: + raise VoiceBridgeProtocolError("Message id was reused with different content", close_code=1008) + continue + if len(self._seen_messages) >= _MAX_SEEN_MESSAGES: + raise VoiceBridgeProtocolError("Message dedupe limit exceeded", close_code=1008) + self._seen_messages[message_id] = digest + return payload + + async def _shutdown_runtime(self, *, drain_callbacks: bool) -> None: + self._ending = True + async with self._state_lock: + responses = list(self._pending_turns.values()) + self._pending_turns.clear() + if self._active_response is not None and self._active_response not in responses: + responses.append(self._active_response) + self._terminal_response_ids.update(response.response_id for response in responses) + release = self._active_release + for response in responses: + was_terminal = response.is_terminal + await response._mark_terminal() + if not was_terminal: + self._record_terminal(response.response_id, "connection_closed") + self._fail_helper_waiters("Voice connection closed") + if release is not None: + release.set() + + worker = self._callback_worker + if worker is not None: + if not drain_callbacks: + worker.cancel() + await asyncio.gather(worker, return_exceptions=True) + else: + try: + await asyncio.wait_for(self._callback_queue.join(), timeout=_CLEANUP_TIMEOUT_SECONDS) + except asyncio.TimeoutError: + logger.warning("Voice callback drain exceeded cleanup deadline") + if drain_callbacks and not worker.done(): + try: + self._callback_queue.put_nowait(None) + except asyncio.QueueFull: + worker.cancel() + try: + await asyncio.wait_for(worker, timeout=_CLEANUP_TIMEOUT_SECONDS) + except asyncio.TimeoutError: + worker.cancel() + await asyncio.gather(worker, return_exceptions=True) + except asyncio.CancelledError: + pass + + self._closed = True + self._record_close(1000) + if self._cleanup_tasks: + try: + await asyncio.wait_for( + asyncio.gather(*tuple(self._cleanup_tasks), return_exceptions=True), + timeout=_CLEANUP_TIMEOUT_SECONDS, + ) + except asyncio.TimeoutError: + for task in tuple(self._cleanup_tasks): + task.cancel() + + def _fail_helper_waiters(self, message: str) -> None: + for future in tuple(self._cancel_waiters.values()): + if not future.done(): + future.set_exception(VoiceBridgeConnectionClosedError(message)) + self._cancel_waiters.clear() + self._abandoned_proactive_cancels.clear() + self._dtmf_collections.clear() + self._dtmf_cancel_pending.clear() + self._recent_dtmf_cancel_races.clear() + for _, proactive_future in self._pending_proactive.values(): + if not proactive_future.done(): + proactive_future.set_exception(VoiceBridgeConnectionClosedError(message)) + self._pending_proactive.clear() + + def _find_response_locked(self, response_id: str) -> VoiceResponse | None: + if self._active_response is not None and self._active_response.response_id == response_id: + return self._active_response + return self._recent_responses.get(response_id) + + def _remember_response_locked(self, response: VoiceResponse) -> None: + self._recent_responses[response.response_id] = response + self._recent_responses.move_to_end(response.response_id) + while len(self._recent_responses) > _MAX_RECENT_RESPONSES: + self._recent_responses.popitem(last=False) + + def _record_first_output(self, message_type: str, fields: dict[str, Any]) -> None: + if message_type not in ("response.output_text.delta", "response.output_text.done"): + return + response_id = fields.get("response_id") + if isinstance(response_id, str): + self._record_first_output_for_response(response_id) + + def _record_first_output_for_response(self, response_id: str) -> None: + if response_id in self._first_output_recorded: + return + started = self._response_start_ns.get(response_id) + if started is None: + return + self._first_output_recorded.add(response_id) + duration_ms = (time.monotonic_ns() - started) / 1_000_000 + _FIRST_OUTPUT_DURATION.record(duration_ms) + + def _record_terminal(self, response_id: str, terminal_kind: str) -> None: + _TERMINAL_COUNTER.add(1, {"kind": terminal_kind}) + self._response_start_ns.pop(response_id, None) + self._first_output_recorded.discard(response_id) + + def _record_activation(self, result: str) -> None: + if self._activation_recorded: + return + self._activation_recorded = True + _ACTIVATION_COUNTER.add(1, {"result": result}) + + def _record_close(self, code: int) -> None: + if self._close_recorded: + return + self._close_recorded = True + _CLOSE_CODE_COUNTER.add(1, {"code": code}) + + def _ensure_ready(self) -> None: + if not self._ready or self._closed: + raise VoiceBridgeConnectionClosedError("The voice connection is not ready") + if self._ending: + raise VoiceBridgeConnectionClosedError("The voice session is ending") + + async def _close(self, *, code: int, reason: str) -> None: + if self._closed: + return + self._record_close(code) + self._closed = True + await self._websocket.close(code=code, reason=reason) diff --git a/sdk/agentserver/azure-ai-agentserver-invocations/azure/ai/agentserver/invocations/voice/_models.py b/sdk/agentserver/azure-ai-agentserver-invocations/azure/ai/agentserver/invocations/voice/_models.py new file mode 100644 index 000000000000..c7a97fe01e77 --- /dev/null +++ b/sdk/agentserver/azure-ai-agentserver-invocations/azure/ai/agentserver/invocations/voice/_models.py @@ -0,0 +1,272 @@ +# --------------------------------------------------------- +# Copyright (c) Microsoft Corporation. All rights reserved. +# --------------------------------------------------------- +"""Public models for the typed Voice Live bridge host.""" + +from __future__ import annotations + +from collections.abc import Mapping +from dataclasses import dataclass, field +from typing import Any, Literal, Union + + +@dataclass(frozen=True) +class ResponseTimeouts: + """Effective response deadlines advertised by the bridge. + + :param first_output_ms: Maximum time to first output or explicit decline. + :param idle_ms: Maximum time between output progress messages. + :param max_duration_ms: Absolute maximum response duration. + """ + + first_output_ms: int + idle_ms: int + max_duration_ms: int + + +@dataclass(frozen=True) +class InputTextPart: + """One text part from an ordered ``user.message`` payload. + + :param text: Final recognized or application-supplied text. + """ + + text: str + type: Literal["input_text"] = field(default="input_text", init=False) + + +@dataclass(frozen=True) +class InputImagePart: + """One reference-only image part from ``user.message``. + + :param image_ref: Short-lived fetchable image reference. + :param mime_type: Image media type. + :param alt: Optional untrusted caller-provided caption. + """ + + image_ref: str + mime_type: str + alt: str | None = None + type: Literal["input_image"] = field(default="input_image", init=False) + + +UserContentPart = Union[InputTextPart, InputImagePart] + + +@dataclass(frozen=True) +class ConversationHistoryItem: + """One caller-app supplied user-role history item. + + :param item_id: Caller/bridge-allocated ``hi_`` identifier. + :param content: Supported content parts in original order. + """ + + item_id: str + content: tuple[UserContentPart, ...] + role: Literal["user"] = field(default="user", init=False) + + +@dataclass(frozen=True) +class ConversationItemCreateEvent: + """Non-response-producing history create request. + + :param request_id: Inbound envelope identifier used for result correlation. + :param item: User-role history item to persist. + :param previous_item_id: Insertion predecessor, ``root``, or ``None`` to append. + """ + + request_id: str + item: ConversationHistoryItem + previous_item_id: str | None = None + + +@dataclass(frozen=True) +class ConversationItemDeleteEvent: + """Non-response-producing history delete request. + + :param request_id: Inbound envelope identifier used for result correlation. + :param item_id: Existing history item to delete. + """ + + request_id: str + item_id: str + + +@dataclass(frozen=True) +class SessionStartEvent: + """Validated application-start event delivered before ``session.ready``. + + :param protocol_version: Exact accepted bridge protocol version. + :param reconnect: Whether this transport reattaches the logical session. + :param response_timeouts: Effective response deadlines. + :param greeting: Optional bridge-owned greeting; absent on reconnect. + :param no_input_timeout_ms: Optional user-silence threshold. + :param caller: Deeply read-only, untrusted caller metadata. + """ + + protocol_version: str + reconnect: bool + response_timeouts: ResponseTimeouts + greeting: str | None = None + no_input_timeout_ms: int | None = None + caller: Mapping[str, Any] | None = None + + +@dataclass(frozen=True) +class UserMessageEvent: + """Completed user turn with ordered content parts. + + :param item_id: Bridge-allocated ``in_`` input item identifier. + :param content: Supported content parts in original wire order. + """ + + item_id: str + content: tuple[UserContentPart, ...] + + @property + def text(self) -> str: + """Return all text parts joined with one space. + + :return: Convenience text projection preserving text-part order. + :rtype: str + """ + return " ".join(part.text for part in self.content if isinstance(part, InputTextPart)) + + +@dataclass(frozen=True) +class UserNoInputEvent: + """Bridge-generated silence turn. + + :param item_id: Bridge-allocated ``in_`` input item identifier. + :param count: Consecutive no-input count. + """ + + item_id: str + count: int + + +@dataclass(frozen=True) +class UserSpeechStartedEvent: + """Advisory signal that caller speech began while no response was open.""" + + +@dataclass(frozen=True) +class DtmfKeyEvent: + """One raw session-scoped DTMF key. + + :param digit: Exactly one of ``0``–``9``, ``*``, or ``#``. + """ + + digit: str + + +@dataclass(frozen=True) +class DtmfCollectedEvent: + """Completed DTMF collection delivered as a new response turn. + + :param item_id: Bridge-allocated ``in_`` input item identifier. + :param collection_id: SDK-allocated ``dc_`` collection identifier. + :param digits: Collected digits, excluding the terminator. + :param completion_reason: Open-enum completion reason. + """ + + item_id: str + collection_id: str + digits: str + completion_reason: str + + +@dataclass(frozen=True) +class DtmfCollectionRejectedEvent: + """DTMF collection request that the bridge did not start. + + :param collection_id: Rejected ``dc_`` collection identifier. + :param reason: Open-enum rejection reason. + """ + + collection_id: str + reason: str + + +@dataclass(frozen=True) +class DtmfCollectionCancelledEvent: + """Pending or active DTMF collection that ended without a result turn. + + :param collection_id: Cancelled ``dc_`` collection identifier. + :param reason: Open-enum cancellation reason. + """ + + collection_id: str + reason: str + + +@dataclass(frozen=True) +class HandoffFailedEvent: + """Bridge-generated recovery turn after target activation failed. + + :param item_id: Bridge-allocated ``in_`` recovery item identifier. + :param target: Same-project target agent name. + :param code: Open-enum failure code. + :param message: Optional sanitized diagnostic detail. + """ + + item_id: str + target: str + code: str + message: str | None = None + + +@dataclass(frozen=True) +class BargeInEvent: + """Caller interruption and playback outcome. + + :param response_id: Agent-owned response that was interrupted. + :param heard_text: Approximate text played before the cut. + :param item_id: Output item playing at the cut, when one existed. + """ + + response_id: str + heard_text: str + item_id: str | None = None + + +@dataclass(frozen=True) +class ResponseTimeoutEvent: + """Terminal response-deadline notification. + + Exactly one of ``response_id`` and ``item_ids`` is populated. + + :param stage: Open-enum timeout stage. + :param response_id: Open response terminated by the bridge. + :param item_ids: Pending ordered input batch terminated before response open. + """ + + stage: str + response_id: str | None = None + item_ids: tuple[str, ...] | None = None + + +@dataclass(frozen=True) +class ResponseCancellationOutcome: + """Winning playback outcome returned by ``VoiceResponse.cancel``. + + :param response_id: Response whose cancellation was requested. + :param kind: Winning terminal, ``cancelled`` or caller ``barge_in``. + :param heard_text: Approximate text played before the terminal. + :param item_id: Output item playing at the terminal, when one existed. + """ + + response_id: str + kind: Literal["cancelled", "barge_in"] + heard_text: str + item_id: str | None = None + + +@dataclass(frozen=True) +class SessionEndEvent: + """Bridge-initiated session termination. + + :param reason: Open-enum termination reason. + """ + + reason: str diff --git a/sdk/agentserver/azure-ai-agentserver-invocations/azure/ai/agentserver/invocations/voice/_protocol.py b/sdk/agentserver/azure-ai-agentserver-invocations/azure/ai/agentserver/invocations/voice/_protocol.py new file mode 100644 index 000000000000..376bec52a96d --- /dev/null +++ b/sdk/agentserver/azure-ai-agentserver-invocations/azure/ai/agentserver/invocations/voice/_protocol.py @@ -0,0 +1,500 @@ +# --------------------------------------------------------- +# Copyright (c) Microsoft Corporation. All rights reserved. +# --------------------------------------------------------- +"""Internal codec for Voice Live Bridge Protocol 1.0.""" + +# Internal codec helpers are documented at the public host/model layer. +# pylint: disable=docstring-missing-param,docstring-missing-return,docstring-missing-rtype + +from __future__ import annotations + +import datetime +import json +import re +import uuid +from collections.abc import Mapping, Sequence +from types import MappingProxyType +from typing import Any, cast + +from ._models import ( + ConversationHistoryItem, + ConversationItemCreateEvent, + ConversationItemDeleteEvent, + DtmfCollectedEvent, + DtmfCollectionCancelledEvent, + DtmfCollectionRejectedEvent, + DtmfKeyEvent, + HandoffFailedEvent, + InputImagePart, + InputTextPart, + ResponseTimeoutEvent, + ResponseTimeouts, + SessionStartEvent, + UserMessageEvent, +) + +PROTOCOL_VERSION = "1.0" +MAX_ERROR_MESSAGE_LENGTH = 1024 +_SAFE_CODE = re.compile(r"^[A-Za-z0-9._-]{1,64}$") +_VOICE_TYPE_ALIASES = {"azure-platform": "azure-standard", "custom": "azure-custom"} +_VOICE_TYPES = { + "openai", + "azure-standard", + "azure-custom", + "azure-personal", + "avatar-voice-sync", + "azure-realtime-native", +} +_VOICE_REQUIRED_STRING_FIELDS = { + "name", + "endpoint_id", + "model", +} +_VOICE_NULLABLE_STRING_FIELDS = { + "locale", + "style", + "pitch", + "rate", + "volume", + "custom_lexicon_url", + "custom_text_normalization_url", + "multi_talker_speaker_name", +} +_VOICE_FIELDS = { + "type", + "temperature", + "prefer_locales", + *_VOICE_REQUIRED_STRING_FIELDS, + *_VOICE_NULLABLE_STRING_FIELDS, +} +_AZURE_VOICE_OPTIONAL_FIELDS = { + "temperature", + "custom_lexicon_url", + "custom_text_normalization_url", + "prefer_locales", + "locale", + "style", + "pitch", + "rate", + "volume", +} +_VOICE_VARIANT_FIELDS = { + "openai": {"type", "name"}, + "azure-realtime-native": {"type", "name"}, + "azure-standard": {"type", "name", "multi_talker_speaker_name", *_AZURE_VOICE_OPTIONAL_FIELDS}, + "azure-custom": {"type", "name", "endpoint_id", *_AZURE_VOICE_OPTIONAL_FIELDS}, + "azure-personal": {"type", "name", "model", *_AZURE_VOICE_OPTIONAL_FIELDS}, + "avatar-voice-sync": {"type", "model", *_AZURE_VOICE_OPTIONAL_FIELDS}, +} +_RFC3339 = re.compile( + r"^(?P\d{4}-\d{2}-\d{2})T(?P