diff --git a/.agents/references/realtime-session-lifecycle.md b/.agents/references/realtime-session-lifecycle.md index 4a1f9c68a4..e60a251e07 100644 --- a/.agents/references/realtime-session-lifecycle.md +++ b/.agents/references/realtime-session-lifecycle.md @@ -42,7 +42,9 @@ Do not add a new side effect before a failure point without defining who release ## Guardrails and Response Ordering - Realtime output guardrails inspect accumulated transcript text at configured debounce thresholds, not each token and not a final `Runner` output object. They emit `guardrail_tripped` instead of raising a normal Runner tripwire exception. -- A tripped output guardrail marks the response interrupted before awaiting transport work, emits one trip event per response, forces response cancellation, and sends safe follow-up input naming the guardrail. Concurrent guardrail tasks must not interrupt or message the same response twice. +- A tripped output guardrail marks the response interrupted before awaiting transport work, emits one trip event per response, forces response cancellation, and sends safe follow-up input naming the guardrail. Correlate that follow-up's `response.create` request with the exact response that starts; do not classify the next arbitrary turn as recovery. Custom models that cannot echo the request correlation ID must preserve response request ordering so the session can use its ordered fallback for both successful starts and create failures. If the correlated recovery response also trips, reject it normally without starting another automatic recovery; a later user turn starts a new recovery chain. Concurrent guardrail tasks must not interrupt or message the same response twice. +- Bind each response to the agent that was active when the turn started. Later guardrail debounce thresholds must use that source-agent snapshot even if `update_agent()` or a handoff changes the current agent first. +- Text-only guardrail trips cancel the identified response without truncating playback state from an earlier audio response. Audio transcript trips still use the playback interrupt path so consumers receive `audio_interrupted`. - Guardrail callbacks can run after audio has already been buffered or played. Consumers must treat `audio_interrupted` as the signal to stop local playback; text rejection alone cannot retract audio already delivered. - An exception from one output guardrail is logged and skipped so it does not silently terminate the live session. Exceptions that escape the background guardrail task must become a `RealtimeError` event rather than disappearing. - Realtime function-tool input guardrails follow the same optional pre-approval and mandatory post-approval ordering as standard function tools, but their rejection is returned through Realtime tool output and events. diff --git a/examples/realtime/app/README.md b/examples/realtime/app/README.md index 7d003a6f39..b35f58ab5e 100644 --- a/examples/realtime/app/README.md +++ b/examples/realtime/app/README.md @@ -30,6 +30,26 @@ cd examples/realtime/app && LOG_LEVEL=DEBUG uv run python server.py The debug logs include concise summaries for server, model, session, history, tool, handoff, error, and usage events. Audio frames and high-volume delta events are omitted, and transcript content is not logged. Uvicorn and WebSocket protocol logging remain at INFO so `LOG_LEVEL=DEBUG` does not dump wire payloads. +### Testing delayed output guardrails + +Enable the manual delayed guardrail when validating response lifecycle behavior: + +```bash +cd examples/realtime/app +REALTIME_GUARDRAIL_TEST=1 \ +REALTIME_GUARDRAIL_TEST_DELAY_SECONDS=1 \ +REALTIME_GUARDRAIL_TEST_TRIGGER_PHRASE="red balloon" \ +REALTIME_GUARDRAIL_TEST_DEBOUNCE_TEXT_LENGTH=1 \ +LOG_LEVEL=DEBUG \ +uv run python server.py +``` + +Ask the agent to say "red balloon" in a short response. Test mode lowers the debounce threshold to 1 by default so even this short output is checked. The delay makes it likely that the response finishes before the guardrail result is ready. The app should still receive a `guardrail_tripped` event and generate the safe follow-up without cancelling a newer response. + +The guardrail rejects every response containing the trigger phrase. If the automatic safe follow-up also contains it, the follow-up should be rejected with another `guardrail_tripped` event, but the SDK should not generate a third response automatically. A new user turn should still be checked normally and can trigger the guardrail again. + +To exercise the overlapping-response case, speak again immediately after the triggering response finishes but before the delayed guardrail returns. The newer response should continue without being interrupted. + ## Customization To use the same UI with your own agents, edit `agent.py` and ensure get_starting_agent() returns the right starting agent for your use case. diff --git a/examples/realtime/app/agent.py b/examples/realtime/app/agent.py index e83564e279..cdf84c339e 100644 --- a/examples/realtime/app/agent.py +++ b/examples/realtime/app/agent.py @@ -1,6 +1,8 @@ import asyncio +import os +from typing import Any -from agents import function_tool +from agents import GuardrailFunctionOutput, OutputGuardrail, function_tool from agents.extensions.handoff_prompt import RECOMMENDED_PROMPT_PREFIX from agents.realtime import RealtimeAgent, realtime_handoff @@ -8,6 +10,41 @@ When running the UI example locally, you can edit this file to change the setup. THe server will use the agent returned from get_starting_agent() as the starting agent.""" + +async def delayed_manual_test_guardrail( + _context: Any, + _agent: Any, + output: str, +) -> GuardrailFunctionOutput: + """Delay output validation so response lifecycle races are easy to test manually.""" + delay_seconds = float(os.getenv("REALTIME_GUARDRAIL_TEST_DELAY_SECONDS", "1.0")) + trigger_phrase = os.getenv("REALTIME_GUARDRAIL_TEST_TRIGGER_PHRASE", "red balloon") + await asyncio.sleep(max(delay_seconds, 0.0)) + return GuardrailFunctionOutput( + output_info={"trigger_phrase": trigger_phrase, "delay_seconds": delay_seconds}, + tripwire_triggered=bool(trigger_phrase) and trigger_phrase.casefold() in output.casefold(), + ) + + +def get_manual_guardrail_test_debounce_text_length() -> int | None: + """Return the short manual-test debounce threshold when the guardrail is enabled.""" + if os.getenv("REALTIME_GUARDRAIL_TEST", "").casefold() not in {"1", "true", "yes", "on"}: + return None + configured = int(os.getenv("REALTIME_GUARDRAIL_TEST_DEBOUNCE_TEXT_LENGTH", "1")) + return max(configured, 1) + + +manual_test_output_guardrails = ( + [ + OutputGuardrail( + guardrail_function=delayed_manual_test_guardrail, + name="delayed_manual_test_guardrail", + ) + ] + if get_manual_guardrail_test_debounce_text_length() is not None + else [] +) + ### TOOLS @@ -65,6 +102,7 @@ def get_weather(city: str) -> str: 2. Use the faq lookup tool to answer the question. Do not rely on your own knowledge. 3. If you cannot answer the question, transfer back to the triage agent.""", tools=[faq_lookup_tool], + output_guardrails=manual_test_output_guardrails, ) seat_booking_agent = RealtimeAgent( @@ -79,6 +117,7 @@ def get_weather(city: str) -> str: 3. Use the update seat tool to update the seat on the flight. If the customer asks a question that is not related to the routine, transfer back to the triage agent. """, tools=[update_seat], + output_guardrails=manual_test_output_guardrails, ) triage_agent = RealtimeAgent( @@ -89,6 +128,7 @@ def get_weather(city: str) -> str: "You are a helpful triaging agent. You can use your tools to delegate questions to other appropriate agents." ), tools=[get_weather], + output_guardrails=manual_test_output_guardrails, handoffs=[ realtime_handoff(faq_agent, tool_name_override="transfer_to_faq_agent"), realtime_handoff(seat_booking_agent, tool_name_override="transfer_to_seat_booking_agent"), diff --git a/examples/realtime/app/server.py b/examples/realtime/app/server.py index d92d8b0353..31ebee9e47 100644 --- a/examples/realtime/app/server.py +++ b/examples/realtime/app/server.py @@ -14,7 +14,7 @@ from typing_extensions import assert_never from agents.realtime import RealtimeRunner, RealtimeSession, RealtimeSessionEvent -from agents.realtime.config import RealtimeUserInputMessage +from agents.realtime.config import RealtimeRunConfig, RealtimeUserInputMessage from agents.realtime.items import RealtimeItem from agents.realtime.model import RealtimeModelConfig from agents.realtime.model_events import ( @@ -27,15 +27,15 @@ # Import TwilioHandler class - handle both module and package use cases if TYPE_CHECKING: # For type checking, use the relative import - from .agent import get_starting_agent + from .agent import get_manual_guardrail_test_debounce_text_length, get_starting_agent else: # At runtime, try both import styles try: # Try relative import first (when used as a package) - from .agent import get_starting_agent + from .agent import get_manual_guardrail_test_debounce_text_length, get_starting_agent except ImportError: # Fall back to direct import (when run as a script) - from agent import get_starting_agent + from agent import get_manual_guardrail_test_debounce_text_length, get_starting_agent _requested_log_level = os.getenv("LOG_LEVEL", "INFO").upper() @@ -45,6 +45,13 @@ logger.setLevel(_log_level) +def _get_runner_config() -> RealtimeRunConfig | None: + debounce_text_length = get_manual_guardrail_test_debounce_text_length() + if debounce_text_length is None: + return None + return {"guardrails_settings": {"debounce_text_length": debounce_text_length}} + + class RealtimeWebSocketManager: def __init__(self): self.active_sessions: dict[str, RealtimeSession] = {} @@ -56,7 +63,7 @@ async def connect(self, websocket: WebSocket, session_id: str): self.websockets[session_id] = websocket agent = get_starting_agent() - runner = RealtimeRunner(agent) + runner = RealtimeRunner(agent, config=_get_runner_config()) # If you want to customize the runner behavior, you can pass options: # runner_config = RealtimeRunConfig(async_tool_calls=False) # runner = RealtimeRunner(agent, config=runner_config) @@ -205,7 +212,7 @@ def _log_debug_event(self, session_id: str, event: RealtimeSessionEvent) -> None def _log_debug_model_event(self, session_id: str, event: Any) -> None: model_event = event.data - if model_event.type in {"audio", "transcript_delta"}: + if model_event.type in {"audio", "output_text_delta", "transcript_delta"}: return if isinstance(model_event, RealtimeModelRawServerEvent): diff --git a/src/agents/realtime/__init__.py b/src/agents/realtime/__init__.py index 5310d0bd5a..d3999b9602 100644 --- a/src/agents/realtime/__init__.py +++ b/src/agents/realtime/__init__.py @@ -70,6 +70,7 @@ RealtimeModelItemDeletedEvent, RealtimeModelItemUpdatedEvent, RealtimeModelOtherEvent, + RealtimeModelOutputTextDeltaEvent, RealtimeModelOutputTokensDetails, RealtimeModelToolCallEvent, RealtimeModelTranscriptDeltaEvent, @@ -173,6 +174,7 @@ "RealtimeModelItemDeletedEvent", "RealtimeModelItemUpdatedEvent", "RealtimeModelOtherEvent", + "RealtimeModelOutputTextDeltaEvent", "RealtimeModelOutputTokensDetails", "RealtimeModelToolCallEvent", "RealtimeModelTranscriptDeltaEvent", diff --git a/src/agents/realtime/model_events.py b/src/agents/realtime/model_events.py index 2716d32026..54bbb6803c 100644 --- a/src/agents/realtime/model_events.py +++ b/src/agents/realtime/model_events.py @@ -16,6 +16,10 @@ class RealtimeModelErrorEvent: error: Any type: Literal["error"] = "error" + response_create_id: str | None = None + """Correlation ID of the failed response.create request, when available.""" + is_guardrail_recovery: bool | None = None + """Whether the failed response.create was requested for guardrail recovery, when known.""" @dataclass @@ -106,6 +110,17 @@ class RealtimeModelTranscriptDeltaEvent: type: Literal["transcript_delta"] = "transcript_delta" +@dataclass +class RealtimeModelOutputTextDeltaEvent: + """Partial text output update.""" + + item_id: str + delta: str + response_id: str + + type: Literal["output_text_delta"] = "output_text_delta" + + @dataclass class RealtimeModelItemUpdatedEvent: """Item added to the history or updated.""" @@ -138,6 +153,12 @@ class RealtimeModelTurnStartedEvent: """Triggered when the model starts generating a response for a turn.""" type: Literal["turn_started"] = "turn_started" + response_id: str | None = None + """Provider response ID for this turn, when available.""" + response_create_id: str | None = None + """Correlation ID copied from the response.create request, when available.""" + is_guardrail_recovery: bool | None = None + """Whether this turn was requested for guardrail recovery, when known.""" @dataclass @@ -189,6 +210,8 @@ class RealtimeModelTurnEndedEvent: """Triggered when the model finishes generating a response for a turn.""" type: Literal["turn_ended"] = "turn_ended" + response_id: str | None = None + """Provider response ID for this turn, when available.""" @dataclass @@ -208,6 +231,10 @@ class RealtimeModelExceptionEvent: context: str | None = None type: Literal["exception"] = "exception" + response_create_id: str | None = None + """Correlation ID of the failed response.create request, when available.""" + is_guardrail_recovery: bool | None = None + """Whether the failed response.create was requested for guardrail recovery, when known.""" @dataclass @@ -228,6 +255,7 @@ class RealtimeModelRawServerEvent: | RealtimeModelInputAudioTimeoutTriggeredEvent | RealtimeModelInputAudioTranscriptionCompletedEvent | RealtimeModelTranscriptDeltaEvent + | RealtimeModelOutputTextDeltaEvent | RealtimeModelItemUpdatedEvent | RealtimeModelItemDeletedEvent | RealtimeModelConnectionStatusEvent diff --git a/src/agents/realtime/model_inputs.py b/src/agents/realtime/model_inputs.py index c167ce34f8..1d219b34dc 100644 --- a/src/agents/realtime/model_inputs.py +++ b/src/agents/realtime/model_inputs.py @@ -68,6 +68,12 @@ class RealtimeModelSendUserInput: user_input: RealtimeModelUserInput """The user input to send.""" + response_create_id: str | None = None + """Optional correlation ID for the response.create triggered by this input. + + Custom models should echo this value on the matching turn_started event when possible. + """ + @dataclass class RealtimeModelSendAudio: @@ -98,6 +104,12 @@ class RealtimeModelSendInterrupt: force_response_cancel: bool = False """Force sending a response.cancel event even if automatic cancellation is enabled.""" + response_id: str | None = None + """Limit the interrupt to the response that triggered it, when supported.""" + + cancel_response_only: bool = False + """Cancel the response without interrupting or truncating audio playback.""" + @dataclass class RealtimeModelSendSessionUpdate: diff --git a/src/agents/realtime/openai_realtime.py b/src/agents/realtime/openai_realtime.py index f36aafffa4..531faf25f7 100644 --- a/src/agents/realtime/openai_realtime.py +++ b/src/agents/realtime/openai_realtime.py @@ -57,6 +57,9 @@ from openai.types.realtime.realtime_function_tool import ( RealtimeFunctionTool as OpenAISessionFunction, ) +from openai.types.realtime.realtime_response_create_params import ( + RealtimeResponseCreateParams as OpenAIRealtimeResponseCreateParams, +) from openai.types.realtime.realtime_response_usage import RealtimeResponseUsage from openai.types.realtime.realtime_server_event import ( RealtimeServerEvent as OpenAIRealtimeServerEvent, @@ -133,6 +136,7 @@ RealtimeModelInputTokensDetails, RealtimeModelItemDeletedEvent, RealtimeModelItemUpdatedEvent, + RealtimeModelOutputTextDeltaEvent, RealtimeModelOutputTokensDetails, RealtimeModelRawServerEvent, RealtimeModelToolCallEvent, @@ -152,6 +156,7 @@ ) FormatInput: TypeAlias = str | AudioPCM | AudioPCMU | AudioPCMA | Mapping[str, Any] | None +_RESPONSE_CREATE_EVENT_ID_METADATA_KEY = "openai_agents_response_create_event_id" # Avoid direct imports of non-exported names by referencing via module @@ -233,6 +238,7 @@ class _PendingResponseCreate: request_version: int target_version: int is_manual: bool + response_create_id: str | None class _ResponseCreateSequencer: @@ -292,18 +298,35 @@ async def set_response_control( self._response_control = control self._condition.notify_all() - async def mark_response_created(self) -> None: + async def mark_response_created( + self, + response_create_event_id: str | None, + ) -> tuple[str | None, bool]: async with self._condition: + pending = self._pending_response_create + matches_pending = pending is not None and ( + pending.response_create_id is None + or ( + response_create_event_id is not None + and response_create_event_id == pending.event_id + ) + ) + response_create_id = ( + pending.response_create_id if pending is not None and matches_pending else None + ) + is_guardrail_recovery = response_create_id is not None self._ongoing_response = True - self._pending_response_create = None - self._response_control = "free" + if matches_pending: + self._pending_response_create = None + self._response_control = "free" self._condition.notify_all() + return response_create_id, is_guardrail_recovery async def mark_response_done(self) -> None: async with self._condition: self._ongoing_response = False - self._pending_response_create = None - self._response_control = "free" + if self._pending_response_create is None: + self._response_control = "free" self._condition.notify_all() async def release_waiters(self) -> None: @@ -327,29 +350,35 @@ async def reserve_response_create_request(self, *, manual: bool = False) -> int: self._condition.notify_all() return request_version - async def clear_pending_response_create(self, event_id: str | None = None) -> bool: + async def clear_pending_response_create( + self, event_id: str | None = None + ) -> _PendingResponseCreate | None: async with self._condition: if ( self._response_control != "create_requested" or self._pending_response_create is None ): - return False + return None if event_id is not None and self._pending_response_create.event_id != event_id: - return False + return None # The caller only uses the no-event-id path for response.create-like # server errors, so clearing here won't release unrelated requests. - self._pending_request_versions.discard(self._pending_response_create.request_version) - if self._pending_response_create.is_manual: - self._manual_response_create_versions.discard( - self._pending_response_create.request_version - ) + pending = self._pending_response_create + self._pending_request_versions.discard(pending.request_version) + if pending.is_manual: + self._manual_response_create_versions.discard(pending.request_version) self._pending_response_create = None self._response_control = "free" self._condition.notify_all() - return True + return pending async def wait_for_response_create_slot( - self, request_version: int, *, manual: bool = False, event_id: str | None = None + self, + request_version: int, + *, + manual: bool = False, + event_id: str | None = None, + response_create_id: str | None = None, ) -> _PendingResponseCreate | None: while True: async with self._condition: @@ -381,6 +410,7 @@ async def wait_for_response_create_slot( request_version=request_version, target_version=target_version, is_manual=manual, + response_create_id=response_create_id, ) self._pending_response_create = pending return pending @@ -751,8 +781,11 @@ async def _set_response_control( ) -> None: await self._response_create_sequencer.set_response_control(control) - async def _mark_response_created(self) -> None: - await self._response_create_sequencer.mark_response_created() + async def _mark_response_created( + self, + response_create_event_id: str | None = None, + ) -> tuple[str | None, bool]: + return await self._response_create_sequencer.mark_response_created(response_create_event_id) async def _mark_response_done(self) -> None: await self._response_create_sequencer.mark_response_done() @@ -765,7 +798,9 @@ async def _release_response_waiters(self) -> None: async def _reserve_response_create_request(self, *, manual: bool = False) -> int: return await self._response_create_sequencer.reserve_response_create_request(manual=manual) - async def _clear_pending_response_create(self, event_id: str | None = None) -> bool: + async def _clear_pending_response_create( + self, event_id: str | None = None + ) -> _PendingResponseCreate | None: return await self._response_create_sequencer.clear_pending_response_create(event_id) async def _send_response_create_when_idle( @@ -774,21 +809,46 @@ async def _send_response_create_when_idle( *, response_create: OpenAIResponseCreateEvent | None = None, manual: bool = False, + response_create_id: str | None = None, ) -> None: pending = await self._response_create_sequencer.wait_for_response_create_slot( request_version, manual=manual, event_id=response_create.event_id if response_create is not None else None, + response_create_id=response_create_id, ) if pending is None: return try: - response_create_event = ( - response_create.model_copy(update={"event_id": pending.event_id}) - if response_create is not None - else OpenAIResponseCreateEvent(type="response.create", event_id=pending.event_id) - ) + response_params: OpenAIRealtimeResponseCreateParams | None + if pending.response_create_id is not None: + response_params = ( + response_create.response + if response_create is not None and response_create.response is not None + else OpenAIRealtimeResponseCreateParams() + ) + response_metadata = dict(response_params.metadata or {}) + response_metadata[_RESPONSE_CREATE_EVENT_ID_METADATA_KEY] = pending.event_id + response_params = response_params.model_copy(update={"metadata": response_metadata}) + else: + response_params = response_create.response if response_create is not None else None + if response_create is not None: + response_update: dict[str, Any] = {"event_id": pending.event_id} + if response_params is not None: + response_update["response"] = response_params + response_create_event = response_create.model_copy(update=response_update) + elif response_params is not None: + response_create_event = OpenAIResponseCreateEvent( + type="response.create", + event_id=pending.event_id, + response=response_params, + ) + else: + response_create_event = OpenAIResponseCreateEvent( + type="response.create", + event_id=pending.event_id, + ) await self._send_raw_message(response_create_event) except BaseException: await self._clear_pending_response_create(pending.event_id) @@ -802,28 +862,46 @@ async def _send_response_create_in_background( *, response_create: OpenAIResponseCreateEvent | None = None, manual: bool = False, + response_create_id: str | None = None, ) -> None: try: await self._send_response_create_when_idle( request_version, response_create=response_create, manual=manual, + response_create_id=response_create_id, ) except asyncio.CancelledError: logger.debug("Deferred response.create task was cancelled") except AssertionError as exc: - if str(exc) != "Not connected": + if str(exc) != "Not connected" or response_create_id is not None: + await self._emit_event( + RealtimeModelExceptionEvent( + exception=exc, + context="Error sending deferred response.create", + response_create_id=response_create_id, + is_guardrail_recovery=response_create_id is not None, + ) + ) + except websockets.exceptions.ConnectionClosed as exc: + if response_create_id is None: + logger.debug("Skipping deferred response.create because the websocket is closed") + else: await self._emit_event( RealtimeModelExceptionEvent( - exception=exc, context="Error sending deferred response.create" + exception=exc, + context="Error sending deferred response.create", + response_create_id=response_create_id, + is_guardrail_recovery=response_create_id is not None, ) ) - except websockets.exceptions.ConnectionClosed: - logger.debug("Skipping deferred response.create because the websocket is closed") except Exception as exc: await self._emit_event( RealtimeModelExceptionEvent( - exception=exc, context="Error sending deferred response.create" + exception=exc, + context="Error sending deferred response.create", + response_create_id=response_create_id, + is_guardrail_recovery=response_create_id is not None, ) ) @@ -833,12 +911,14 @@ def _start_response_create( *, response_create: OpenAIResponseCreateEvent | None = None, manual: bool = False, + response_create_id: str | None = None, ) -> None: task = asyncio.create_task( self._send_response_create_in_background( request_version, response_create=response_create, manual=manual, + response_create_id=response_create_id, ) ) self._response_create_tasks.add(task) @@ -861,8 +941,13 @@ async def _cancel_response_create_tasks(self) -> None: async def _send_user_input(self, event: RealtimeModelSendUserInput) -> None: converted = _ConversionHelper.convert_user_input_to_item_create(event) await self._send_raw_message(converted) - request_version = await self._reserve_response_create_request() - self._start_response_create(request_version) + manual = event.response_create_id is not None + request_version = await self._reserve_response_create_request(manual=manual) + self._start_response_create( + request_version, + manual=manual, + response_create_id=event.response_create_id, + ) async def _send_audio(self, event: RealtimeModelSendAudio) -> None: converted = _ConversionHelper.convert_audio_to_input_audio_buffer_append(event) @@ -921,6 +1006,22 @@ def _get_audio_limits(self, item_id: str, item_content_index: int) -> tuple[floa return audio_state.audio_length_ms, max_audio_ms async def _send_interrupt(self, event: RealtimeModelSendInterrupt) -> None: + if event.cancel_response_only: + session = self._created_session + automatic_response_cancellation_enabled = ( + session + and session.audio is not None + and session.audio.input is not None + and session.audio.input.turn_detection is not None + and session.audio.input.turn_detection.interrupt_response is True + ) + should_cancel_response = event.force_response_cancel or ( + not automatic_response_cancellation_enabled + ) + if should_cancel_response: + await self._cancel_response(response_id=event.response_id) + return + playback_state = self._get_playback_state() current_item_id = playback_state.get("current_item_id") current_item_content_index = playback_state.get("current_item_content_index") @@ -975,7 +1076,7 @@ async def _send_interrupt(self, event: RealtimeModelSendInterrupt) -> None: not automatic_response_cancellation_enabled ) if should_cancel_response: - await self._cancel_response() + await self._cancel_response(response_id=event.response_id) if current_item_id is not None and elapsed_ms is not None: self._audio_state_tracker.on_interrupted() @@ -1067,12 +1168,17 @@ async def close(self) -> None: else: await self._release_response_waiters() - async def _cancel_response(self) -> None: + async def _cancel_response(self, *, response_id: str | None = None) -> None: if not await self._response_create_sequencer.begin_cancel_response(): return try: - await self._send_raw_message(OpenAIResponseCancelEvent(type="response.cancel")) + cancel_event = ( + OpenAIResponseCancelEvent(type="response.cancel", response_id=response_id) + if response_id is not None + else OpenAIResponseCancelEvent(type="response.cancel") + ) + await self._send_raw_message(cancel_event) except Exception: await self._set_response_control("free") raise @@ -1229,28 +1335,57 @@ async def _handle_ws_event(self, event: dict[str, Any]): if not automatic_response_cancellation_enabled: await self._cancel_response() elif parsed.type == "response.created": - await self._mark_response_created() - await self._emit_event(RealtimeModelTurnStartedEvent()) + response = getattr(parsed, "response", None) + metadata = getattr(response, "metadata", None) + response_create_event_id = ( + metadata.get(_RESPONSE_CREATE_EVENT_ID_METADATA_KEY) + if isinstance(metadata, Mapping) + else None + ) + response_create_id, is_guardrail_recovery = await self._mark_response_created( + response_create_event_id + ) + response_id = getattr(response, "id", None) + await self._emit_event( + RealtimeModelTurnStartedEvent( + response_id=response_id, + response_create_id=response_create_id, + is_guardrail_recovery=is_guardrail_recovery, + ) + ) elif parsed.type == "response.done": await self._mark_response_done() if parsed.response.usage is not None: await self._emit_event( _ConversionHelper.convert_response_usage(parsed.response.usage) ) - await self._emit_event(RealtimeModelTurnEndedEvent()) + await self._emit_event( + RealtimeModelTurnEndedEvent(response_id=getattr(parsed.response, "id", None)) + ) elif parsed.type == "session.created": await self._send_tracing_config(self._tracing_config) self._update_created_session(parsed.session) elif parsed.type == "session.updated": self._update_created_session(parsed.session) elif parsed.type == "error": + response_create_id = None + is_guardrail_recovery = None if ( not self._ongoing_response and self._response_control == "create_requested" and self._error_matches_pending_response_create(parsed.error) ): - await self._clear_pending_response_create(parsed.error.event_id) - await self._emit_event(RealtimeModelErrorEvent(error=parsed.error)) + pending = await self._clear_pending_response_create(parsed.error.event_id) + if pending is not None: + response_create_id = pending.response_create_id + is_guardrail_recovery = response_create_id is not None + await self._emit_event( + RealtimeModelErrorEvent( + error=parsed.error, + response_create_id=response_create_id, + is_guardrail_recovery=is_guardrail_recovery, + ) + ) elif parsed.type == "conversation.item.deleted": await self._emit_event(RealtimeModelItemDeletedEvent(item_id=parsed.item_id)) elif ( @@ -1286,9 +1421,16 @@ async def _handle_ws_event(self, event: dict[str, Any]): item_id=parsed.item_id, delta=parsed.delta, response_id=parsed.response_id ) ) + elif parsed.type == "response.output_text.delta": + await self._emit_event( + RealtimeModelOutputTextDeltaEvent( + item_id=parsed.item_id, + delta=parsed.delta, + response_id=parsed.response_id, + ) + ) elif ( parsed.type == "conversation.item.input_audio_transcription.delta" - or parsed.type == "response.output_text.delta" or parsed.type == "response.function_call_arguments.delta" ): # No support for partials yet diff --git a/src/agents/realtime/session.py b/src/agents/realtime/session.py index 0a985c0387..df177b150b 100644 --- a/src/agents/realtime/session.py +++ b/src/agents/realtime/session.py @@ -4,6 +4,7 @@ import dataclasses import inspect import json +from collections import deque from collections.abc import AsyncIterator, Sequence from functools import partial from typing import Any, cast @@ -71,6 +72,7 @@ from .model_events import ( RealtimeModelEvent, RealtimeModelInputAudioTranscriptionCompletedEvent, + RealtimeModelOutputTextDeltaEvent, RealtimeModelToolCallEvent, RealtimeModelUsageEvent, ) @@ -225,8 +227,17 @@ def __init__( # Guardrails state tracking self._interrupted_response_ids: set[str] = set() + self._active_output_response_id: str | None = None + self._response_agent_snapshots: dict[str, RealtimeAgent] = {} + self._unscoped_response_agent_snapshot: RealtimeAgent | None = None self._item_transcripts: dict[str, str] = {} # item_id -> accumulated transcript self._item_guardrail_run_counts: dict[str, int] = {} # item_id -> run count + self._pending_guardrail_recovery_ids: set[str] = set() + self._pending_guardrail_recovery_order: deque[str] = deque() + self._active_guardrail_recovery_response_ids: set[str] = set() + self._legacy_guardrail_recovery_turn_active = False + self._legacy_guardrail_recovery_response_id: str | None = None + self._guardrail_recovery_request_counter = 0 self._debounce_text_length = self._run_config.get("guardrails_settings", {}).get( "debounce_text_length", 100 ) @@ -369,6 +380,10 @@ async def on_event(self, event: RealtimeModelEvent) -> None: return if event.type == "error": + self._discard_failed_guardrail_recovery( + event.response_create_id, + event.is_guardrail_recovery, + ) await self._put_event(RealtimeError(info=self._event_info, error=event.error)) elif event.type == "function_call": agent_snapshot = self._current_agent @@ -422,13 +437,13 @@ async def on_event(self, event: RealtimeModelEvent) -> None: ) ) elif event.type == "transcript_delta": - # Accumulate transcript text for guardrail debouncing per item_id item_id = event.item_id - if item_id not in self._item_transcripts: - self._item_transcripts[item_id] = "" - self._item_guardrail_run_counts[item_id] = 0 - - self._item_transcripts[item_id] += event.delta + self._record_output_guardrail_delta( + item_id, + event.delta, + event.response_id, + is_audio_output=True, + ) self._history = self._get_new_history( self._history, AssistantMessageItem( @@ -436,16 +451,14 @@ async def on_event(self, event: RealtimeModelEvent) -> None: content=[AssistantAudio(transcript=self._item_transcripts[item_id])], ), ) - - # Check if we should run guardrails based on debounce threshold - current_length = len(self._item_transcripts[item_id]) - threshold = self._debounce_text_length - next_run_threshold = (self._item_guardrail_run_counts[item_id] + 1) * threshold - - if current_length >= next_run_threshold: - self._item_guardrail_run_counts[item_id] += 1 - # Pass response_id so we can ensure only a single interrupt per response - self._enqueue_guardrail_task(self._item_transcripts[item_id], event.response_id) + elif event.type == "output_text_delta": + assert isinstance(event, RealtimeModelOutputTextDeltaEvent) + self._record_output_guardrail_delta( + event.item_id, + event.delta, + event.response_id, + is_audio_output=False, + ) elif event.type == "item_updated": is_new = not any(item.item_id == event.item.item_id for item in self._history) @@ -518,6 +531,25 @@ async def on_event(self, event: RealtimeModelEvent) -> None: elif event.type == "connection_status": pass elif event.type == "turn_started": + if event.response_id is None: + self._unscoped_response_agent_snapshot = self._current_agent + else: + self._response_agent_snapshots[event.response_id] = self._current_agent + is_guardrail_recovery = self._consume_guardrail_recovery( + event.response_create_id, + event.is_guardrail_recovery, + ) + if is_guardrail_recovery: + if event.response_id is not None: + self._active_guardrail_recovery_response_ids.add(event.response_id) + self._legacy_guardrail_recovery_turn_active = False + self._legacy_guardrail_recovery_response_id = None + else: + self._legacy_guardrail_recovery_turn_active = True + self._legacy_guardrail_recovery_response_id = None + else: + self._legacy_guardrail_recovery_turn_active = False + self._legacy_guardrail_recovery_response_id = None await self._put_event( RealtimeAgentStartEvent( agent=self._current_agent, @@ -529,8 +561,20 @@ async def on_event(self, event: RealtimeModelEvent) -> None: self._context_wrapper.usage.add(event.usage) elif event.type == "turn_ended": # Clear guardrail state for next turn + self._active_output_response_id = None + if event.response_id is None: + self._response_agent_snapshots.clear() + else: + self._response_agent_snapshots.pop(event.response_id, None) + self._unscoped_response_agent_snapshot = None self._item_transcripts.clear() self._item_guardrail_run_counts.clear() + if event.response_id is None: + self._active_guardrail_recovery_response_ids.clear() + else: + self._active_guardrail_recovery_response_ids.discard(event.response_id) + self._legacy_guardrail_recovery_turn_active = False + self._legacy_guardrail_recovery_response_id = None await self._put_event( RealtimeAgentEndEvent( @@ -539,6 +583,10 @@ async def on_event(self, event: RealtimeModelEvent) -> None: ) ) elif event.type == "exception": + self._discard_failed_guardrail_recovery( + event.response_create_id, + event.is_guardrail_recovery, + ) # Store the exception to be raised in __aiter__ self._stored_exception = event.exception elif event.type == "other": @@ -1299,12 +1347,19 @@ def _image_url_str(val: object) -> str | None: # Otherwise, add it to the end return old_history + [event] - async def _run_output_guardrails(self, text: str, response_id: str) -> bool: + async def _run_output_guardrails( + self, + text: str, + response_id: str, + agent_snapshot: RealtimeAgent, + is_guardrail_recovery: bool, + is_audio_output: bool, + ) -> bool: """Run output guardrails on the given text. Returns True if any guardrail was triggered.""" if self._closing or self._closed: return False - combined_guardrails = self._current_agent.output_guardrails + self._run_config.get( + combined_guardrails = agent_snapshot.output_guardrails + self._run_config.get( "output_guardrails", [] ) seen_ids: set[int] = set() @@ -1326,7 +1381,7 @@ async def _run_output_guardrails(self, text: str, response_id: str) -> bool: result = await guardrail.run( # TODO (rm) Remove this cast, it's wrong self._context_wrapper, - cast(Agent[Any], self._current_agent), + cast(Agent[Any], agent_snapshot), text, ) if self._closing or self._closed: @@ -1360,31 +1415,156 @@ async def _run_output_guardrails(self, text: str, response_id: str) -> bool: ): return False - # Interrupt the model - if self._closing or self._closed: - return False - await self._model.send_event(RealtimeModelSendInterrupt(force_response_cancel=True)) + # Interrupt only while the response that produced this output remains active. + if self._is_guardrail_response_active(response_id): + await self._model.send_event( + RealtimeModelSendInterrupt( + force_response_cancel=True, + response_id=response_id, + cancel_response_only=not is_audio_output, + ) + ) - # Send guardrail triggered message - if self._closing or self._closed: - return False - guardrail_names = [result.guardrail.get_name() for result in triggered_results] - await self._model.send_event( - RealtimeModelSendUserInput( - user_input=f"guardrail triggered: {', '.join(guardrail_names)}" + # Start one automatic recovery response for each ordinary response. If that recovery + # response also trips a guardrail, reject it without extending the same recovery chain. + if not is_guardrail_recovery: + if self._closing or self._closed: + return False + guardrail_names = [result.guardrail.get_name() for result in triggered_results] + self._guardrail_recovery_request_counter += 1 + response_create_id = ( + f"guardrail_recovery_{self._guardrail_recovery_request_counter}" ) - ) + self._pending_guardrail_recovery_ids.add(response_create_id) + self._pending_guardrail_recovery_order.append(response_create_id) + try: + await self._model.send_event( + RealtimeModelSendUserInput( + user_input=f"guardrail triggered: {', '.join(guardrail_names)}", + response_create_id=response_create_id, + ) + ) + except BaseException: + self._discard_guardrail_recovery(response_create_id) + raise return True return False - def _enqueue_guardrail_task(self, text: str, response_id: str) -> None: + def _discard_guardrail_recovery(self, response_create_id: str) -> None: + self._pending_guardrail_recovery_ids.discard(response_create_id) + try: + self._pending_guardrail_recovery_order.remove(response_create_id) + except ValueError: + pass + + def _discard_failed_guardrail_recovery( + self, + response_create_id: str | None, + is_guardrail_recovery: bool | None, + ) -> None: + if response_create_id is not None: + self._discard_guardrail_recovery(response_create_id) + return + + if is_guardrail_recovery is not True or not self._pending_guardrail_recovery_order: + return + + recovery_id = self._pending_guardrail_recovery_order.popleft() + self._pending_guardrail_recovery_ids.discard(recovery_id) + + def _consume_guardrail_recovery( + self, + response_create_id: str | None, + is_guardrail_recovery: bool | None, + ) -> bool: + if response_create_id is not None: + was_pending = response_create_id in self._pending_guardrail_recovery_ids + if was_pending: + self._discard_guardrail_recovery(response_create_id) + return was_pending or is_guardrail_recovery is True + + if is_guardrail_recovery is False: + return False + if not self._pending_guardrail_recovery_order: + return is_guardrail_recovery is True + + recovery_id = self._pending_guardrail_recovery_order.popleft() + self._pending_guardrail_recovery_ids.discard(recovery_id) + return True + + def _is_guardrail_response_active(self, response_id: str) -> bool: + return ( + not self._closing + and not self._closed + and self._active_output_response_id == response_id + ) + + def _record_output_guardrail_delta( + self, + item_id: str, + delta: str, + response_id: str, + *, + is_audio_output: bool, + ) -> None: + """Accumulate text or audio transcript deltas using the same guardrail debounce.""" + self._active_output_response_id = response_id + agent_snapshot = self._response_agent_snapshots.get(response_id) + if agent_snapshot is None: + agent_snapshot = self._unscoped_response_agent_snapshot or self._current_agent + self._response_agent_snapshots[response_id] = agent_snapshot + if ( + self._legacy_guardrail_recovery_turn_active + and self._legacy_guardrail_recovery_response_id is None + ): + self._legacy_guardrail_recovery_response_id = response_id + is_guardrail_recovery = ( + response_id in self._active_guardrail_recovery_response_ids + or self._legacy_guardrail_recovery_response_id == response_id + ) + if item_id not in self._item_transcripts: + self._item_transcripts[item_id] = "" + self._item_guardrail_run_counts[item_id] = 0 + + self._item_transcripts[item_id] += delta + current_length = len(self._item_transcripts[item_id]) + next_run_threshold = ( + self._item_guardrail_run_counts[item_id] + 1 + ) * self._debounce_text_length + if current_length >= next_run_threshold: + self._item_guardrail_run_counts[item_id] += 1 + self._enqueue_guardrail_task( + self._item_transcripts[item_id], + response_id, + agent_snapshot=agent_snapshot, + is_guardrail_recovery=is_guardrail_recovery, + is_audio_output=is_audio_output, + ) + + def _enqueue_guardrail_task( + self, + text: str, + response_id: str, + *, + agent_snapshot: RealtimeAgent, + is_guardrail_recovery: bool, + is_audio_output: bool, + ) -> None: # Runs the guardrails in a separate task to avoid blocking the main loop if self._closing or self._closed: return - task = asyncio.create_task(self._run_output_guardrails(text, response_id)) + task = asyncio.create_task( + self._run_output_guardrails( + text, + response_id, + agent_snapshot, + is_guardrail_recovery, + is_audio_output, + ) + ) self._guardrail_tasks.add(task) # Add callback to remove completed tasks and handle exceptions diff --git a/src/agents/result.py b/src/agents/result.py index 7bccccec91..1093978468 100644 --- a/src/agents/result.py +++ b/src/agents/result.py @@ -579,6 +579,10 @@ class RunResultStreaming(RunResultBase): interruptions: list[ToolApprovalItem] = field(default_factory=list) """Pending tool approval requests (interruptions) for this run.""" _waiting_on_event_queue: bool = field(default=False, repr=False) + _active_stream_consumers: int = field(default=0, init=False, repr=False) + _stream_consumers_stopped: asyncio.Event = field( + default_factory=asyncio.Event, init=False, repr=False + ) _current_turn_persisted_item_count: int = 0 """Number of items from new_items already persisted to session for the @@ -775,6 +779,22 @@ def cancel(self, mode: Literal["immediate", "after_turn"] = "immediate") -> None # Don't call _cleanup_tasks() or clear queues yet pass + async def _wait_for_turn_event_consumption(self) -> None: + """Wait for active consumers to finish processing the current turn's events.""" + if self._active_stream_consumers == 0: + return + + queue_drained = asyncio.create_task(self._event_queue.join()) + consumers_stopped = asyncio.create_task(self._stream_consumers_stopped.wait()) + tasks = {queue_drained, consumers_stopped} + try: + await asyncio.wait(tasks, return_when=asyncio.FIRST_COMPLETED) + finally: + for task in tasks: + if not task.done(): + task.cancel() + await asyncio.gather(*tasks, return_exceptions=True) + async def stream_events(self) -> AsyncIterator[StreamEvent]: """Stream deltas for new items as they are generated. We're using the types from the OpenAI Responses API, so these are semantic events: each event has a `type` field that @@ -784,9 +804,45 @@ async def stream_events(self) -> AsyncIterator[StreamEvent]: - A MaxTurnsExceeded exception if the agent exceeds the max_turns limit. - A GuardrailTripwireTriggered exception if a guardrail is tripped. """ + registered_consumer_task: asyncio.Task[Any] | None = None + item_acknowledgement_pending = False + + def acknowledge_item() -> None: + nonlocal item_acknowledgement_pending + if not item_acknowledgement_pending: + return + item_acknowledgement_pending = False + self._event_queue.task_done() + + def unregister_consumer(task: asyncio.Task[Any]) -> None: + nonlocal registered_consumer_task + if registered_consumer_task is not task: + return + acknowledge_item() + registered_consumer_task = None + self._active_stream_consumers -= 1 + if self._active_stream_consumers == 0: + self._stream_consumers_stopped.set() + + def register_current_consumer() -> None: + nonlocal registered_consumer_task + current_task = asyncio.current_task() + if current_task is None or registered_consumer_task is current_task: + return + if registered_consumer_task is not None: + previous_task = registered_consumer_task + previous_task.remove_done_callback(unregister_consumer) + unregister_consumer(previous_task) + registered_consumer_task = current_task + self._active_stream_consumers += 1 + self._stream_consumers_stopped.clear() + current_task.add_done_callback(unregister_consumer) + + register_current_consumer() cancelled = False try: while True: + register_current_consumer() self._check_errors() should_drain_queued_events = isinstance( self._stored_exception, MaxTurnsExceeded @@ -826,9 +882,16 @@ async def stream_events(self) -> AsyncIterator[StreamEvent]: self._check_errors() break - yield item - self._event_queue.task_done() + item_acknowledgement_pending = True + try: + yield item + finally: + acknowledge_item() finally: + if registered_consumer_task is not None: + consumer_task = registered_consumer_task + consumer_task.remove_done_callback(unregister_consumer) + unregister_consumer(consumer_task) try: if cancelled: # Cancellation should return promptly, so avoid waiting on long-running tasks. diff --git a/src/agents/run_internal/run_loop.py b/src/agents/run_internal/run_loop.py index 4873ed542d..a94e740551 100644 --- a/src/agents/run_internal/run_loop.py +++ b/src/agents/run_internal/run_loop.py @@ -325,6 +325,28 @@ def _complete_stream_interruption( streamed_result._event_queue.put_nowait(QueueCompleteSentinel()) +async def _wait_for_streamed_turn_events_and_stop_if_cancelled( + streamed_result: RunResultStreaming, +) -> bool: + """Let consumers process the completed turn before starting another one.""" + await streamed_result._wait_for_turn_event_consumption() + if streamed_result._cancel_mode != "after_turn": + return False + + streamed_result.is_complete = True + streamed_result._event_queue.put_nowait(QueueCompleteSentinel()) + return True + + +def _publish_streamed_result_agent( + streamed_result: RunResultStreaming, + agent: Agent[Any], +) -> None: + """Publish an agent transition before cancellation can complete the streamed run.""" + streamed_result.current_agent = agent + streamed_result._current_agent_output_schema = get_output_schema(agent) + + async def _save_resumed_stream_items( *, session: Session | None, @@ -872,6 +894,7 @@ async def _save_stream_items_without_count( current_agent = turn_result.next_step.new_agent if run_state is not None: run_state._current_agent = current_agent + _publish_streamed_result_agent(streamed_result, current_agent) if current_span: current_span.finish(reset_current=True) current_span = None @@ -880,6 +903,10 @@ async def _save_stream_items_without_count( AgentUpdatedStreamEvent(new_agent=current_agent) ) run_state._current_step = NextStepRunAgain() # type: ignore[assignment] + if await _wait_for_streamed_turn_events_and_stop_if_cancelled( + streamed_result + ): + break continue if isinstance(turn_result.next_step, NextStepFinalOutput): @@ -903,6 +930,10 @@ async def _save_stream_items_without_count( store_setting, ) run_state._current_step = NextStepRunAgain() # type: ignore[assignment] + if await _wait_for_streamed_turn_events_and_stop_if_cancelled( + streamed_result + ): + break continue run_state._current_step = None @@ -1181,6 +1212,7 @@ async def _save_stream_items_without_count( current_agent = turn_result.next_step.new_agent if run_state is not None: run_state._current_agent = current_agent + _publish_streamed_result_agent(streamed_result, current_agent) current_span.finish(reset_current=True) current_span = None should_run_agent_start_hooks = True @@ -1190,9 +1222,7 @@ async def _save_stream_items_without_count( if streamed_result._state is not None: streamed_result._state._current_step = NextStepRunAgain() - if streamed_result._cancel_mode == "after_turn": # type: ignore[comparison-overlap] - streamed_result.is_complete = True - streamed_result._event_queue.put_nowait(QueueCompleteSentinel()) + if await _wait_for_streamed_turn_events_and_stop_if_cancelled(streamed_result): break elif isinstance(turn_result.next_step, NextStepFinalOutput): await _finalize_streamed_final_output( @@ -1242,9 +1272,7 @@ async def _save_stream_items_without_count( store_setting, ) - if streamed_result._cancel_mode == "after_turn": # type: ignore[comparison-overlap] - streamed_result.is_complete = True - streamed_result._event_queue.put_nowait(QueueCompleteSentinel()) + if await _wait_for_streamed_turn_events_and_stop_if_cancelled(streamed_result): break except Exception as e: if current_span and _should_attach_generic_agent_error(e): diff --git a/tests/realtime/test_app_server_debug.py b/tests/realtime/test_app_server_debug.py index 847a13b8a7..494779cf05 100644 --- a/tests/realtime/test_app_server_debug.py +++ b/tests/realtime/test_app_server_debug.py @@ -9,7 +9,10 @@ from agents.realtime.events import RealtimeEventInfo, RealtimeRawModelEvent from agents.realtime.items import InputText, UserMessageItem -from agents.realtime.model_events import RealtimeModelItemUpdatedEvent +from agents.realtime.model_events import ( + RealtimeModelItemUpdatedEvent, + RealtimeModelOutputTextDeltaEvent, +) from agents.run_context import RunContextWrapper @@ -18,6 +21,8 @@ def app_server(monkeypatch: pytest.MonkeyPatch) -> ModuleType: app_dir = Path(__file__).parents[2] / "examples" / "realtime" / "app" monkeypatch.chdir(app_dir) monkeypatch.setenv("LOG_LEVEL", "DEBUG") + monkeypatch.delenv("REALTIME_GUARDRAIL_TEST", raising=False) + monkeypatch.delenv("REALTIME_GUARDRAIL_TEST_DEBOUNCE_TEXT_LENGTH", raising=False) module = importlib.import_module("examples.realtime.app.server") return importlib.reload(module) @@ -42,3 +47,35 @@ def test_item_updated_debug_summary_uses_concrete_event_type( assert "item-1" in caplog.text assert "input_text" in caplog.text assert "sensitive transcript" not in caplog.text + + +def test_output_text_delta_is_omitted_from_debug_logs( + app_server: ModuleType, + caplog: pytest.LogCaptureFixture, +) -> None: + event = RealtimeRawModelEvent( + data=RealtimeModelOutputTextDeltaEvent( + item_id="item-1", + response_id="response-1", + delta="partial text", + ), + info=RealtimeEventInfo(context=RunContextWrapper(None)), + ) + + with caplog.at_level(logging.DEBUG, logger=app_server.__name__): + app_server.manager._log_debug_event("session-1", event) + + assert caplog.records == [] + + +def test_manual_guardrail_test_uses_short_debounce_threshold( + app_server: ModuleType, + monkeypatch: pytest.MonkeyPatch, +) -> None: + assert app_server._get_runner_config() is None + + monkeypatch.setenv("REALTIME_GUARDRAIL_TEST", "1") + assert app_server._get_runner_config() == {"guardrails_settings": {"debounce_text_length": 1}} + + monkeypatch.setenv("REALTIME_GUARDRAIL_TEST_DEBOUNCE_TEXT_LENGTH", "7") + assert app_server._get_runner_config() == {"guardrails_settings": {"debounce_text_length": 7}} diff --git a/tests/realtime/test_model_events.py b/tests/realtime/test_model_events.py index 031567b632..42efa77720 100644 --- a/tests/realtime/test_model_events.py +++ b/tests/realtime/test_model_events.py @@ -27,6 +27,17 @@ def test_usage_event_types_are_publicly_exported() -> None: assert getattr(realtime, name) is not None +def test_output_text_delta_event_is_publicly_exported() -> None: + event = realtime.RealtimeModelOutputTextDeltaEvent( + item_id="item_1", + delta="hello", + response_id="response_1", + ) + + assert "RealtimeModelOutputTextDeltaEvent" in realtime.__all__ + assert event.type == "output_text_delta" + + def test_custom_model_can_construct_typed_usage_without_openai_types() -> None: event = realtime.RealtimeModelUsageEvent( usage=Usage(requests=1, input_tokens=8, output_tokens=5, total_tokens=13), diff --git a/tests/realtime/test_openai_realtime.py b/tests/realtime/test_openai_realtime.py index aeeb58081b..6732c6d973 100644 --- a/tests/realtime/test_openai_realtime.py +++ b/tests/realtime/test_openai_realtime.py @@ -16,7 +16,10 @@ from agents.realtime.model_events import ( RealtimeModelAudioEvent, RealtimeModelErrorEvent, + RealtimeModelExceptionEvent, + RealtimeModelOutputTextDeltaEvent, RealtimeModelToolCallEvent, + RealtimeModelTurnStartedEvent, RealtimeModelUsageEvent, ) from agents.realtime.model_inputs import ( @@ -519,7 +522,13 @@ def validate_python(self, event): self._string_adapter.validate_python(voice) if event["type"] == "response.done": return SimpleNamespace(type=event["type"], response=SimpleNamespace(usage=None)) - return SimpleNamespace(type=event["type"]) + return SimpleNamespace( + type=event["type"], + response=SimpleNamespace( + id=None, + metadata=event.get("response", {}).get("metadata"), + ), + ) monkeypatch.setattr(model, "_send_raw_message", fake_send_raw) model._server_event_type_adapter = CustomVoiceRejectingAdapter() @@ -534,7 +543,9 @@ def validate_python(self, event): response_with_custom_voice = { "type": "response.created", - "response": {"audio": {"output": {"voice": {"id": "voice_test"}}}}, + "response": { + "audio": {"output": {"voice": {"id": "voice_test"}}}, + }, } await model._handle_ws_event(response_with_custom_voice) @@ -964,6 +975,9 @@ async def test_send_event_dispatch(self, model, monkeypatch): await model.send_event(RealtimeModelSendUserInput(user_input="hi")) await asyncio.sleep(0) + pending_event_id = model._pending_response_create_event_id + assert pending_event_id is not None + await model._mark_response_created(pending_event_id) await model._mark_response_done() await model.send_event(RealtimeModelSendAudio(audio=b"a", commit=False)) await model.send_event(RealtimeModelSendAudio(audio=b"a", commit=True)) @@ -1002,11 +1016,22 @@ async def test_interrupt_force_cancel_overrides_auto_cancellation(self, model, m monkeypatch.setattr(model, "_send_raw_message", send_raw) monkeypatch.setattr(model, "_emit_event", emit_event) - await model._send_interrupt(RealtimeModelSendInterrupt(force_response_cancel=True)) + await model._send_interrupt( + RealtimeModelSendInterrupt( + force_response_cancel=True, + response_id="response_1", + ) + ) assert send_raw.await_count == 2 payload_types = {call.args[0].type for call in send_raw.call_args_list} assert payload_types == {"conversation.item.truncate", "response.cancel"} + cancel_event = next( + call.args[0] + for call in send_raw.call_args_list + if call.args[0].type == "response.cancel" + ) + assert cancel_event.response_id == "response_1" assert model._ongoing_response is True assert model._response_control == "cancel_requested" @@ -1015,6 +1040,38 @@ async def test_interrupt_force_cancel_overrides_auto_cancellation(self, model, m assert model._response_control == "free" assert model._audio_state_tracker.get_last_audio_item() is None + @pytest.mark.asyncio + async def test_cancel_response_only_does_not_truncate_stale_audio(self, model, monkeypatch): + model._audio_state_tracker.set_audio_format("pcm16") + model._audio_state_tracker.on_audio_delta("stale_audio_item", 0, b"\x00" * 4800) + await model._mark_response_created() + model._created_session = SimpleNamespace( + audio=SimpleNamespace( + input=SimpleNamespace(turn_detection=SimpleNamespace(interrupt_response=True)) + ) + ) + + send_raw = AsyncMock() + emit_event = AsyncMock() + monkeypatch.setattr(model, "_send_raw_message", send_raw) + monkeypatch.setattr(model, "_emit_event", emit_event) + + await model._send_interrupt( + RealtimeModelSendInterrupt( + force_response_cancel=True, + response_id="text_response", + cancel_response_only=True, + ) + ) + + assert send_raw.await_count == 1 + assert send_raw.await_args is not None + cancel_event = send_raw.await_args.args[0] + assert cancel_event.type == "response.cancel" + assert cancel_event.response_id == "text_response" + emit_event.assert_not_awaited() + assert model._audio_state_tracker.get_last_audio_item() == ("stale_audio_item", 0) + @pytest.mark.asyncio async def test_interrupt_respects_auto_cancellation_when_not_forced(self, model, monkeypatch): """Interrupt should avoid sending response.cancel when relying on automatic cancellation.""" @@ -1039,6 +1096,221 @@ async def test_interrupt_respects_auto_cancellation_when_not_forced(self, model, assert all(call.args[0].type != "response.cancel" for call in send_raw.call_args_list) assert model._ongoing_response is True + @pytest.mark.asyncio + async def test_output_text_delta_emits_provider_neutral_event(self, model): + listener = AsyncMock() + model.add_listener(listener) + + await model._handle_ws_event( + { + "type": "response.output_text.delta", + "event_id": "event_1", + "item_id": "item_1", + "response_id": "response_1", + "output_index": 0, + "content_index": 0, + "delta": "hello", + } + ) + + emitted = [call.args[0] for call in listener.on_event.call_args_list] + assert [event.type for event in emitted] == ["raw_server_event", "output_text_delta"] + text_event = emitted[1] + assert isinstance(text_event, RealtimeModelOutputTextDeltaEvent) + assert text_event.item_id == "item_1" + assert text_event.response_id == "response_1" + assert text_event.delta == "hello" + + @pytest.mark.asyncio + async def test_response_create_correlation_skips_an_already_requested_user_turn( + self, model, monkeypatch + ): + payload_types: list[str] = [] + response_create_events: list[Any] = [] + + async def fake_send_raw(event): + payload_types.append(event.type) + if event.type == "response.create": + response_create_events.append(event) + + class ResponseLifecycleAdapter: + def validate_python(self, event): + return SimpleNamespace( + type=event["type"], + response=SimpleNamespace( + id=event["response"]["id"], + metadata=event["response"].get("metadata"), + usage=None, + ), + ) + + monkeypatch.setattr(model, "_send_raw_message", fake_send_raw) + model._server_event_type_adapter = ResponseLifecycleAdapter() + listener = AsyncMock() + model.add_listener(listener) + + await model._send_user_input(RealtimeModelSendUserInput(user_input="user input")) + await asyncio.sleep(0) + await model._send_user_input( + RealtimeModelSendUserInput( + user_input="guardrail recovery", + response_create_id="recovery_1", + ) + ) + await asyncio.sleep(0) + + assert payload_types == [ + "conversation.item.create", + "response.create", + "conversation.item.create", + ] + assert len(response_create_events) == 1 + + await model._handle_ws_event( + { + "type": "response.created", + "response": {"id": "response_user"}, + } + ) + await model._handle_ws_event({"type": "response.done", "response": {"id": "response_user"}}) + await asyncio.sleep(0) + + assert payload_types[-1] == "response.create" + assert len(response_create_events) == 2 + assert response_create_events[1].response is not None + await model._handle_ws_event( + { + "type": "response.created", + "response": { + "id": "response_recovery", + "metadata": response_create_events[1].response.metadata, + }, + } + ) + + turn_started_events = [ + call.args[0] + for call in listener.on_event.call_args_list + if isinstance(call.args[0], RealtimeModelTurnStartedEvent) + ] + assert [ + ( + event.response_id, + event.response_create_id, + event.is_guardrail_recovery, + ) + for event in turn_started_events + ] == [ + ("response_user", None, False), + ("response_recovery", "recovery_1", True), + ] + + @pytest.mark.asyncio + async def test_server_vad_turn_does_not_consume_pending_guardrail_recovery( + self, model, monkeypatch + ): + response_create_events: list[Any] = [] + + async def fake_send_raw(event): + if event.type == "response.create": + response_create_events.append(event) + + class ResponseLifecycleAdapter: + def validate_python(self, event): + return SimpleNamespace( + type=event["type"], + response=SimpleNamespace( + id=event["response"]["id"], + metadata=event["response"].get("metadata"), + usage=None, + ), + ) + + monkeypatch.setattr(model, "_send_raw_message", fake_send_raw) + model._server_event_type_adapter = ResponseLifecycleAdapter() + listener = AsyncMock() + model.add_listener(listener) + + await model._send_user_input( + RealtimeModelSendUserInput( + user_input="guardrail recovery", + response_create_id="recovery_1", + ) + ) + await asyncio.sleep(0) + assert len(response_create_events) == 1 + pending_event_id = model._pending_response_create_event_id + assert pending_event_id is not None + + await model._handle_ws_event( + {"type": "response.created", "response": {"id": "response_vad"}} + ) + assert model._pending_response_create_event_id == pending_event_id + assert model._response_control == "create_requested" + + await model._handle_ws_event({"type": "response.done", "response": {"id": "response_vad"}}) + assert model._pending_response_create_event_id == pending_event_id + assert model._response_control == "create_requested" + + await model._handle_ws_event( + { + "type": "response.created", + "response": { + "id": "response_recovery", + "metadata": response_create_events[0].response.metadata, + }, + } + ) + + turn_started_events = [ + call.args[0] + for call in listener.on_event.call_args_list + if isinstance(call.args[0], RealtimeModelTurnStartedEvent) + ] + assert [ + ( + event.response_id, + event.response_create_id, + event.is_guardrail_recovery, + ) + for event in turn_started_events + ] == [ + ("response_vad", None, False), + ("response_recovery", "recovery_1", True), + ] + assert model._pending_response_create_event_id is None + assert model._response_control == "free" + + @pytest.mark.asyncio + async def test_correlated_response_create_send_failure_emits_correlation( + self, model, monkeypatch + ): + async def fake_send_raw(event): + if event.type == "response.create": + raise RuntimeError("response.create failed") + + monkeypatch.setattr(model, "_send_raw_message", fake_send_raw) + listener = AsyncMock() + model.add_listener(listener) + + await model._send_user_input( + RealtimeModelSendUserInput( + user_input="guardrail recovery", + response_create_id="recovery_1", + ) + ) + await asyncio.sleep(0) + await asyncio.sleep(0) + + exception_events = [ + call.args[0] + for call in listener.on_event.call_args_list + if isinstance(call.args[0], RealtimeModelExceptionEvent) + ] + assert len(exception_events) == 1 + assert exception_events[0].response_create_id == "recovery_1" + assert model._response_control == "free" + @pytest.mark.asyncio async def test_send_user_input_defers_response_create_without_blocking_caller( self, model, monkeypatch @@ -1141,7 +1413,9 @@ async def fake_send_raw(event): await asyncio.sleep(0) assert payload_types == ["conversation.item.create", "response.create"] - await model._mark_response_created() + pending_event_id = model._pending_response_create_event_id + assert pending_event_id is not None + await model._mark_response_created(pending_event_id) second_task = asyncio.create_task( model._send_user_input(RealtimeModelSendUserInput(user_input="second")) @@ -1195,7 +1469,9 @@ async def fake_send_raw(event): assert payload_types.count("response.create") == 1 - await model._mark_response_created() + pending_event_id = model._pending_response_create_event_id + assert pending_event_id is not None + await model._mark_response_created(pending_event_id) await asyncio.sleep(0) await model._mark_response_done() @@ -1306,6 +1582,43 @@ async def fake_send_raw(event): "response.create", ] + @pytest.mark.asyncio + async def test_response_create_server_error_emits_correlation(self, model, monkeypatch): + async def fake_send_raw(_event): + pass + + emit_event = AsyncMock() + monkeypatch.setattr(model, "_send_raw_message", fake_send_raw) + monkeypatch.setattr(model, "_emit_event", emit_event) + + await model._send_user_input( + RealtimeModelSendUserInput( + user_input="guardrail recovery", + response_create_id="recovery_1", + ) + ) + await asyncio.sleep(0) + + pending_event_id = model._pending_response_create_event_id + assert pending_event_id is not None + + await model._handle_ws_event( + { + "type": "error", + "event_id": "event_err_1", + "error": { + "type": "invalid_request_error", + "code": "bad_response_create", + "message": "bad response.create", + "event_id": pending_event_id, + }, + } + ) + + error_event = emit_event.call_args_list[-1].args[0] + assert isinstance(error_event, RealtimeModelErrorEvent) + assert error_event.response_create_id == "recovery_1" + @pytest.mark.asyncio async def test_missing_unrelated_error_event_id_does_not_release_in_flight_response_create( self, model, monkeypatch @@ -1643,7 +1956,9 @@ async def fake_send_raw(event): assert payload_types.count("response.create") == 1 - await model._mark_response_created() + pending_event_id = model._pending_response_create_event_id + assert pending_event_id is not None + await model._mark_response_created(pending_event_id) await model._mark_response_done() await asyncio.sleep(0) @@ -1693,7 +2008,9 @@ async def fake_send_raw(event): assert payload_types.count("response.create") == 1 - await model._mark_response_created() + pending_event_id = model._pending_response_create_event_id + assert pending_event_id is not None + await model._mark_response_created(pending_event_id) await model._mark_response_done() await asyncio.sleep(0) diff --git a/tests/realtime/test_session.py b/tests/realtime/test_session.py index 8d69916505..1578d99504 100644 --- a/tests/realtime/test_session.py +++ b/tests/realtime/test_session.py @@ -50,6 +50,8 @@ RealtimeModelItemDeletedEvent, RealtimeModelItemUpdatedEvent, RealtimeModelOtherEvent, + RealtimeModelOutputTextDeltaEvent, + RealtimeModelRawServerEvent, RealtimeModelToolCallEvent, RealtimeModelTranscriptDeltaEvent, RealtimeModelTurnEndedEvent, @@ -66,6 +68,7 @@ from agents.realtime.session import ( REJECTION_MESSAGE, RealtimeSession, + _PendingToolOutput, _PendingToolOutputSendError, _serialize_tool_output, ) @@ -3548,7 +3551,13 @@ async def failing_guardrail(context, agent, output): monkeypatch.setattr(_debug, "DONT_LOG_TOOL_DATA", tool_redacted) with caplog.at_level(logging.DEBUG, logger="openai.agents"): - triggered = await session._run_output_guardrails("model text", "response-id") + triggered = await session._run_output_guardrails( + "model text", + "response-id", + agent, + is_guardrail_recovery=False, + is_audio_output=False, + ) assert triggered is False records = [ @@ -3595,7 +3604,13 @@ async def __call__(self, context, agent, output): monkeypatch.setattr(_debug, "DONT_LOG_TOOL_DATA", False) with caplog.at_level(logging.WARNING, logger="openai.agents"): - triggered = await session._run_output_guardrails("model text", "response-id") + triggered = await session._run_output_guardrails( + "model text", + "response-id", + agent, + is_guardrail_recovery=False, + is_audio_output=False, + ) assert triggered is False records = [ @@ -3638,6 +3653,7 @@ async def test_transcript_delta_triggers_guardrail_at_threshold( if isinstance(event, RealtimeModelSendInterrupt) ) assert interrupt_event.force_response_cancel is True + assert interrupt_event.cancel_response_only is False assert len(mock_model.sent_messages) == 1 assert mock_model.sent_messages[0] == "guardrail triggered: triggered_guardrail" @@ -3650,6 +3666,246 @@ async def test_transcript_delta_triggers_guardrail_at_threshold( assert len(guardrail_events) == 1 assert guardrail_events[0].message == "this is more than ten characters" + @pytest.mark.asyncio + async def test_text_output_deltas_trigger_guardrails_at_threshold( + self, mock_model, mock_agent, triggered_guardrail + ): + session = RealtimeSession( + mock_model, + mock_agent, + None, + run_config={ + "output_guardrails": [triggered_guardrail], + "guardrails_settings": {"debounce_text_length": 10}, + }, + ) + + await session.on_event( + RealtimeModelOutputTextDeltaEvent( + item_id="item_text", + response_id="resp_text", + delta="blocked ", + ) + ) + assert mock_model.interrupts_called == 0 + + await session.on_event( + RealtimeModelOutputTextDeltaEvent( + item_id="item_text", + response_id="resp_text", + delta="content", + ) + ) + await self._wait_for_guardrail_tasks(session) + + assert mock_model.interrupts_called == 1 + interrupt_event = next( + event + for event in mock_model.sent_events + if isinstance(event, RealtimeModelSendInterrupt) + ) + assert interrupt_event.cancel_response_only is True + assert mock_model.sent_messages == ["guardrail triggered: triggered_guardrail"] + assert session._item_transcripts["item_text"] == "blocked content" + events = [] + while not session._event_queue.empty(): + events.append(await session._event_queue.get()) + guardrail_events = [ + event for event in events if isinstance(event, RealtimeGuardrailTripped) + ] + assert len(guardrail_events) == 1 + assert guardrail_events[0].message == "blocked content" + + @pytest.mark.asyncio + @pytest.mark.parametrize( + "invalid_event", + [ + {"type": "response.output_text.delta", "delta": "blocked"}, + { + "type": "response.output_text.delta", + "item_id": "item_text", + "response_id": "resp_text", + "delta": "", + }, + { + "type": "response.output_audio_transcript.delta", + "item_id": "item_text", + "response_id": "resp_text", + "delta": "blocked", + }, + ], + ids=["missing-identifiers", "empty-delta", "audio-transcript-handled-separately"], + ) + async def test_unrelated_raw_server_events_do_not_schedule_text_guardrails( + self, mock_model, mock_agent, triggered_guardrail, invalid_event + ): + session = RealtimeSession( + mock_model, + mock_agent, + None, + run_config={ + "output_guardrails": [triggered_guardrail], + "guardrails_settings": {"debounce_text_length": 1}, + }, + ) + + await session.on_event(RealtimeModelRawServerEvent(data=invalid_event)) + await self._wait_for_guardrail_tasks(session) + + assert mock_model.interrupts_called == 0 + assert session._item_transcripts == {} + + @pytest.mark.asyncio + @pytest.mark.parametrize("start_next_response", [False, True]) + async def test_late_guardrail_result_preserves_result_without_interrupting_new_response( + self, mock_model, mock_agent, start_next_response + ): + guardrail_started = asyncio.Event() + release_guardrail = asyncio.Event() + + async def delayed_guardrail(_context, _agent, _output): + guardrail_started.set() + await release_guardrail.wait() + return GuardrailFunctionOutput(output_info={}, tripwire_triggered=True) + + session = RealtimeSession( + mock_model, + mock_agent, + None, + run_config={ + "output_guardrails": [ + OutputGuardrail( + guardrail_function=delayed_guardrail, + name="delayed_guardrail", + ) + ], + "guardrails_settings": {"debounce_text_length": 5}, + }, + ) + + await session.on_event( + RealtimeModelOutputTextDeltaEvent( + item_id="item_1", + response_id="response_1", + delta="blocked", + ) + ) + await asyncio.wait_for(guardrail_started.wait(), timeout=1) + + await session.on_event(RealtimeModelTurnEndedEvent()) + if start_next_response: + await session.on_event( + RealtimeModelOutputTextDeltaEvent( + item_id="item_2", + response_id="response_2", + delta="safe", + ) + ) + release_guardrail.set() + await self._wait_for_guardrail_tasks(session) + + assert mock_model.interrupts_called == 0 + assert mock_model.sent_messages == ["guardrail triggered: delayed_guardrail"] + queued_events = [] + while not session._event_queue.empty(): + queued_events.append(await session._event_queue.get()) + guardrail_events = [ + event for event in queued_events if isinstance(event, RealtimeGuardrailTripped) + ] + assert len(guardrail_events) == 1 + assert guardrail_events[0].message == "blocked" + + @pytest.mark.asyncio + @pytest.mark.parametrize("output_kind", ["text", "audio"]) + async def test_guardrail_runs_when_turn_ends_before_task_starts( + self, mock_model, mock_agent, triggered_guardrail, output_kind + ): + session = RealtimeSession( + mock_model, + mock_agent, + None, + run_config={ + "output_guardrails": [triggered_guardrail], + "guardrails_settings": {"debounce_text_length": 5}, + }, + ) + + delta_event: RealtimeModelOutputTextDeltaEvent | RealtimeModelTranscriptDeltaEvent + if output_kind == "text": + delta_event = RealtimeModelOutputTextDeltaEvent( + item_id="item_1", + response_id="response_1", + delta="blocked", + ) + else: + delta_event = RealtimeModelTranscriptDeltaEvent( + item_id="item_1", + response_id="response_1", + delta="blocked", + ) + + await session.on_event(delta_event) + await session.on_event(RealtimeModelTurnEndedEvent()) + await self._wait_for_guardrail_tasks(session) + + assert mock_model.interrupts_called == 0 + assert mock_model.sent_messages == ["guardrail triggered: triggered_guardrail"] + queued_events = [] + while not session._event_queue.empty(): + queued_events.append(await session._event_queue.get()) + guardrail_events = [ + event for event in queued_events if isinstance(event, RealtimeGuardrailTripped) + ] + assert len(guardrail_events) == 1 + assert guardrail_events[0].message == "blocked" + + @pytest.mark.asyncio + async def test_guardrail_uses_agent_that_produced_the_response(self, mock_model): + guardrail_agents = [] + + def source_guardrail(_context, agent, _output): + guardrail_agents.append(agent) + return GuardrailFunctionOutput(output_info={}, tripwire_triggered=True) + + source_agent = RealtimeAgent( + name="source", + output_guardrails=[ + OutputGuardrail( + guardrail_function=source_guardrail, + name="source_guardrail", + ) + ], + ) + successor_agent = RealtimeAgent(name="successor") + session = RealtimeSession( + mock_model, + source_agent, + None, + run_config={"guardrails_settings": {"debounce_text_length": 5}}, + ) + + await session.on_event(RealtimeModelTurnStartedEvent(response_id="response_1")) + await session.on_event( + RealtimeModelOutputTextDeltaEvent( + item_id="item_1", + response_id="response_1", + delta="bl", + ) + ) + await session.update_agent(successor_agent) + await session.on_event( + RealtimeModelOutputTextDeltaEvent( + item_id="item_1", + response_id="response_1", + delta="ocked", + ) + ) + await self._wait_for_guardrail_tasks(session) + + assert guardrail_agents == [source_agent] + assert mock_model.interrupts_called == 1 + assert mock_model.sent_messages == ["guardrail triggered: source_guardrail"] + @pytest.mark.asyncio async def test_agent_and_run_config_guardrails_not_run_twice(self, mock_model): """Guardrails shared by agent and run config should execute once.""" @@ -3855,6 +4111,514 @@ async def test_agent_output_guardrails_triggered(self, mock_model, triggered_gua assert len(guardrail_events) == 1 assert guardrail_events[0].message == "this is more than ten characters" + @pytest.mark.asyncio + async def test_guardrail_recovery_is_rejected_without_starting_another_recovery( + self, mock_model, mock_agent, triggered_guardrail + ): + session = RealtimeSession( + mock_model, + mock_agent, + None, + run_config={ + "output_guardrails": [triggered_guardrail], + "guardrails_settings": {"debounce_text_length": 5}, + }, + ) + + await session.on_event( + RealtimeModelOutputTextDeltaEvent( + item_id="item_1", + response_id="response_1", + delta="blocked", + ) + ) + await self._wait_for_guardrail_tasks(session) + + assert mock_model.interrupts_called == 1 + assert mock_model.sent_messages == ["guardrail triggered: triggered_guardrail"] + + await session.on_event(RealtimeModelTurnEndedEvent()) + await session.on_event(RealtimeModelTurnStartedEvent()) + await session.on_event( + RealtimeModelOutputTextDeltaEvent( + item_id="item_2", + response_id="response_2", + delta="blocked again", + ) + ) + await self._wait_for_guardrail_tasks(session) + + assert mock_model.interrupts_called == 2 + assert mock_model.sent_messages == ["guardrail triggered: triggered_guardrail"] + + await session.on_event(RealtimeModelTurnEndedEvent()) + await session.on_event(RealtimeModelTurnStartedEvent()) + await session.on_event( + RealtimeModelOutputTextDeltaEvent( + item_id="item_3", + response_id="response_3", + delta="blocked once more", + ) + ) + await self._wait_for_guardrail_tasks(session) + + assert mock_model.interrupts_called == 3 + assert mock_model.sent_messages == [ + "guardrail triggered: triggered_guardrail", + "guardrail triggered: triggered_guardrail", + ] + queued_events = [] + while not session._event_queue.empty(): + queued_events.append(await session._event_queue.get()) + assert ( + len([event for event in queued_events if isinstance(event, RealtimeGuardrailTripped)]) + == 3 + ) + + @pytest.mark.asyncio + async def test_guardrail_recovery_is_correlated_with_its_response(self, mock_model, mock_agent): + release_delayed_guardrail = asyncio.Event() + + async def guardrail_function(context, agent, output): + if output == "blocked first response": + await release_delayed_guardrail.wait() + return GuardrailFunctionOutput(output_info=None, tripwire_triggered=True) + + guardrail = OutputGuardrail( + guardrail_function=guardrail_function, + name="triggered_guardrail", + ) + session = RealtimeSession( + mock_model, + mock_agent, + None, + run_config={ + "output_guardrails": [guardrail], + "guardrails_settings": {"debounce_text_length": 5}, + }, + ) + + await session.on_event( + RealtimeModelOutputTextDeltaEvent( + item_id="item_1", + response_id="response_1", + delta="blocked first response", + ) + ) + await session.on_event(RealtimeModelTurnEndedEvent(response_id="response_1")) + + # The user's next response.create is sent before the delayed guardrail completes. + await session.send_message("next user turn") + release_delayed_guardrail.set() + await self._wait_for_guardrail_tasks(session) + + recovery_events = [ + event + for event in mock_model.sent_events + if isinstance(event, RealtimeModelSendUserInput) + and event.response_create_id is not None + ] + assert len(recovery_events) == 1 + first_recovery_id = recovery_events[0].response_create_id + assert first_recovery_id is not None + + # The user's already-requested turn starts first and must remain an ordinary turn. + await session.on_event( + RealtimeModelTurnStartedEvent( + response_id="response_2", + is_guardrail_recovery=False, + ) + ) + await session.on_event( + RealtimeModelOutputTextDeltaEvent( + item_id="item_2", + response_id="response_2", + delta="blocked user response", + ) + ) + await self._wait_for_guardrail_tasks(session) + + recovery_events = [ + event + for event in mock_model.sent_events + if isinstance(event, RealtimeModelSendUserInput) + and event.response_create_id is not None + ] + assert len(recovery_events) == 2 + + # The correlated replacement remains a recovery even though another turn started first. + await session.on_event(RealtimeModelTurnEndedEvent(response_id="response_2")) + await session.on_event( + RealtimeModelTurnStartedEvent( + response_id="response_3", + response_create_id=first_recovery_id, + ) + ) + await session.on_event( + RealtimeModelOutputTextDeltaEvent( + item_id="item_3", + response_id="response_3", + delta="blocked recovery response", + ) + ) + await self._wait_for_guardrail_tasks(session) + + recovery_events = [ + event + for event in mock_model.sent_events + if isinstance(event, RealtimeModelSendUserInput) + and event.response_create_id is not None + ] + assert len(recovery_events) == 2 + + @pytest.mark.asyncio + async def test_guardrail_recovery_uses_request_order_without_create_correlation( + self, mock_model, mock_agent + ): + release_delayed_guardrail = asyncio.Event() + + async def guardrail_function(context, agent, output): + if output == "blocked first response": + await release_delayed_guardrail.wait() + return GuardrailFunctionOutput(output_info=None, tripwire_triggered=True) + + guardrail = OutputGuardrail( + guardrail_function=guardrail_function, + name="triggered_guardrail", + ) + session = RealtimeSession( + mock_model, + mock_agent, + None, + run_config={ + "output_guardrails": [guardrail], + "guardrails_settings": {"debounce_text_length": 5}, + }, + ) + + await session.on_event( + RealtimeModelOutputTextDeltaEvent( + item_id="item_1", + response_id="response_1", + delta="blocked first response", + ) + ) + await session.on_event(RealtimeModelTurnEndedEvent(response_id="response_1")) + + # An ordinary user request is already queued when the delayed guardrail creates recovery. + await session.send_message("next user turn") + release_delayed_guardrail.set() + await self._wait_for_guardrail_tasks(session) + + await session.on_event( + RealtimeModelTurnStartedEvent( + response_id="response_2", + is_guardrail_recovery=False, + ) + ) + await session.on_event(RealtimeModelTurnEndedEvent(response_id="response_2")) + + # A custom model can provide the response ID while omitting response_create_id. + await session.on_event(RealtimeModelTurnStartedEvent(response_id="response_3")) + await session.on_event( + RealtimeModelOutputTextDeltaEvent( + item_id="item_3", + response_id="response_3", + delta="blocked recovery response", + ) + ) + await self._wait_for_guardrail_tasks(session) + + recovery_events = [ + event + for event in mock_model.sent_events + if isinstance(event, RealtimeModelSendUserInput) + and event.response_create_id is not None + ] + assert len(recovery_events) == 1 + + @pytest.mark.asyncio + async def test_guardrail_recovery_create_failure_releases_pending_request( + self, mock_model, mock_agent, triggered_guardrail + ): + session = RealtimeSession( + mock_model, + mock_agent, + None, + run_config={ + "output_guardrails": [triggered_guardrail], + "guardrails_settings": {"debounce_text_length": 5}, + }, + ) + + await session.on_event( + RealtimeModelOutputTextDeltaEvent( + item_id="item_1", + response_id="response_1", + delta="blocked", + ) + ) + await self._wait_for_guardrail_tasks(session) + + recovery_event = next( + event + for event in mock_model.sent_events + if isinstance(event, RealtimeModelSendUserInput) + and event.response_create_id is not None + ) + assert recovery_event.response_create_id in session._pending_guardrail_recovery_ids + + await session.on_event( + RealtimeModelErrorEvent( + error={"message": "response.create failed"}, + response_create_id=recovery_event.response_create_id, + ) + ) + + assert recovery_event.response_create_id not in session._pending_guardrail_recovery_ids + assert not session._pending_guardrail_recovery_order + + @pytest.mark.asyncio + async def test_unrelated_error_preserves_pending_uncorrelated_recovery( + self, mock_model, mock_agent, triggered_guardrail + ): + session = RealtimeSession( + mock_model, + mock_agent, + None, + run_config={ + "output_guardrails": [triggered_guardrail], + "guardrails_settings": {"debounce_text_length": 5}, + }, + ) + + await session.on_event( + RealtimeModelOutputTextDeltaEvent( + item_id="item_1", + response_id="response_1", + delta="blocked", + ) + ) + await self._wait_for_guardrail_tasks(session) + + recovery_event = next( + event + for event in mock_model.sent_events + if isinstance(event, RealtimeModelSendUserInput) + and event.response_create_id is not None + ) + assert recovery_event.response_create_id in session._pending_guardrail_recovery_ids + + await session.on_event(RealtimeModelErrorEvent(error={"message": "bad item"})) + + assert recovery_event.response_create_id in session._pending_guardrail_recovery_ids + assert list(session._pending_guardrail_recovery_order) == [ + recovery_event.response_create_id + ] + + await session.on_event(RealtimeModelTurnEndedEvent(response_id="response_1")) + await session.on_event( + RealtimeModelTurnStartedEvent( + response_id="response_2", + is_guardrail_recovery=True, + ) + ) + await session.on_event( + RealtimeModelOutputTextDeltaEvent( + item_id="item_2", + response_id="response_2", + delta="blocked again", + ) + ) + await self._wait_for_guardrail_tasks(session) + + recovery_events = [ + event + for event in mock_model.sent_events + if isinstance(event, RealtimeModelSendUserInput) + and event.response_create_id is not None + ] + assert len(recovery_events) == 1 + + @pytest.mark.asyncio + async def test_uncorrelated_recovery_create_failure_releases_pending_recovery( + self, mock_model, mock_agent, triggered_guardrail + ): + session = RealtimeSession( + mock_model, + mock_agent, + None, + run_config={ + "output_guardrails": [triggered_guardrail], + "guardrails_settings": {"debounce_text_length": 5}, + }, + ) + + await session.on_event( + RealtimeModelOutputTextDeltaEvent( + item_id="item_1", + response_id="response_1", + delta="blocked", + ) + ) + await self._wait_for_guardrail_tasks(session) + + recovery_event = next( + event + for event in mock_model.sent_events + if isinstance(event, RealtimeModelSendUserInput) + and event.response_create_id is not None + ) + assert recovery_event.response_create_id in session._pending_guardrail_recovery_ids + + await session.on_event( + RealtimeModelErrorEvent( + error={"message": "response.create failed"}, + is_guardrail_recovery=True, + ) + ) + + assert recovery_event.response_create_id not in session._pending_guardrail_recovery_ids + assert not session._pending_guardrail_recovery_order + + await session.on_event(RealtimeModelTurnEndedEvent(response_id="response_1")) + await session.on_event( + RealtimeModelTurnStartedEvent( + response_id="response_2", + is_guardrail_recovery=False, + ) + ) + await session.on_event( + RealtimeModelOutputTextDeltaEvent( + item_id="item_2", + response_id="response_2", + delta="blocked again", + ) + ) + await self._wait_for_guardrail_tasks(session) + + recovery_events = [ + event + for event in mock_model.sent_events + if isinstance(event, RealtimeModelSendUserInput) + and event.response_create_id is not None + ] + assert len(recovery_events) == 2 + + @pytest.mark.asyncio + async def test_server_vad_turn_does_not_consume_pending_recovery( + self, mock_model, mock_agent, triggered_guardrail + ): + session = RealtimeSession( + mock_model, + mock_agent, + None, + run_config={ + "output_guardrails": [triggered_guardrail], + "guardrails_settings": {"debounce_text_length": 5}, + }, + ) + + await session.on_event( + RealtimeModelOutputTextDeltaEvent( + item_id="item_1", + response_id="response_1", + delta="blocked", + ) + ) + await self._wait_for_guardrail_tasks(session) + + first_recovery = next( + event + for event in mock_model.sent_events + if isinstance(event, RealtimeModelSendUserInput) + and event.response_create_id is not None + ) + await session.on_event(RealtimeModelTurnEndedEvent(response_id="response_1")) + + await session.on_event( + RealtimeModelTurnStartedEvent( + response_id="response_vad", + is_guardrail_recovery=False, + ) + ) + await session.on_event( + RealtimeModelOutputTextDeltaEvent( + item_id="item_vad", + response_id="response_vad", + delta="blocked again", + ) + ) + await self._wait_for_guardrail_tasks(session) + + recovery_events = [ + event + for event in mock_model.sent_events + if isinstance(event, RealtimeModelSendUserInput) + and event.response_create_id is not None + ] + assert len(recovery_events) == 2 + assert first_recovery.response_create_id in session._pending_guardrail_recovery_ids + + @pytest.mark.asyncio + async def test_coalesced_tool_outputs_do_not_precede_pending_recovery( + self, mock_model, mock_agent, triggered_guardrail + ): + session = RealtimeSession( + mock_model, + mock_agent, + None, + run_config={ + "output_guardrails": [triggered_guardrail], + "guardrails_settings": {"debounce_text_length": 5}, + }, + ) + + for call_id in ("call_1", "call_2"): + await session._send_pending_tool_output( + _PendingToolOutput( + tool_call=RealtimeModelToolCallEvent( + name="tool", + call_id=call_id, + arguments="{}", + ), + output="ok", + start_response=True, + ) + ) + + await session.on_event( + RealtimeModelOutputTextDeltaEvent( + item_id="item_1", + response_id="response_1", + delta="blocked", + ) + ) + await self._wait_for_guardrail_tasks(session) + + await session.on_event( + RealtimeModelTurnStartedEvent( + response_id="response_recovery", + is_guardrail_recovery=True, + ) + ) + await session.on_event( + RealtimeModelOutputTextDeltaEvent( + item_id="item_recovery", + response_id="response_recovery", + delta="blocked again", + ) + ) + await self._wait_for_guardrail_tasks(session) + + recovery_events = [ + event + for event in mock_model.sent_events + if isinstance(event, RealtimeModelSendUserInput) + and event.response_create_id is not None + ] + assert len(recovery_events) == 1 + @pytest.mark.asyncio async def test_concurrent_guardrail_tasks_interrupt_once_per_response(self, mock_model): """Even if multiple guardrail tasks trigger concurrently for the same response_id, diff --git a/tests/test_run_impl_resume_paths.py b/tests/test_run_impl_resume_paths.py index 22cf1c0768..1bfb118f20 100644 --- a/tests/test_run_impl_resume_paths.py +++ b/tests/test_run_impl_resume_paths.py @@ -1,3 +1,4 @@ +import asyncio import json from typing import Any, cast @@ -7,6 +8,7 @@ import agents.run as run_module from agents import Agent, Runner, function_tool from agents.agent import ToolsToFinalOutputResult +from agents.agent_output import AgentOutputSchema from agents.items import ( MessageOutputItem, ModelResponse, @@ -21,6 +23,7 @@ from agents.run_internal.agent_bindings import bind_public_agent from agents.run_internal.run_loop import ( NextStepFinalOutput, + NextStepHandoff, NextStepInterruption, NextStepRunAgain, ProcessedResponse, @@ -232,6 +235,122 @@ async def fake_run_single_turn(**_kwargs): assert "function_call" in saved_types +@pytest.mark.asyncio +@pytest.mark.parametrize("continuation", ["run_again", "handoff"]) +async def test_resumed_stream_waits_for_event_consumption_before_continuing( + monkeypatch: pytest.MonkeyPatch, + continuation: str, +) -> None: + agent = Agent(name="resume-agent") + delegate = Agent(name="delegate", output_type=int) + state: RunState[dict[str, str]] = RunState( + context=RunContextWrapper(context={}), + original_input="input", + starting_agent=agent, + max_turns=2, + ) + state._current_step = NextStepInterruption(interruptions=[]) + state._model_responses = [ + ModelResponse(output=[], usage=Usage(), response_id="resp_1"), + ] + state._last_processed_response = ProcessedResponse( + new_items=[], + handoffs=[], + functions=[], + computer_actions=[], + local_shell_calls=[], + shell_calls=[], + apply_patch_calls=[], + tools_used=[], + mcp_approval_requests=[], + interruptions=[], + ) + + tool_output_item = ToolCallOutputItem( + agent=agent, + raw_item={ + "type": "function_call_output", + "call_id": "call-resume", + "output": "ok", + }, + output="ok", + ) + next_step = NextStepHandoff(delegate) if continuation == "handoff" else NextStepRunAgain() + allow_resume_resolution = asyncio.Event() + + async def fake_resolve_interrupted_turn(**_kwargs: object) -> SingleStepResult: + await allow_resume_resolution.wait() + return SingleStepResult( + original_input="input", + model_response=ModelResponse(output=[], usage=Usage(), response_id="resp_resume"), + pre_step_items=[], + new_step_items=[tool_output_item], + next_step=next_step, + tool_input_guardrail_results=[], + tool_output_guardrail_results=[], + ) + + next_model_turn_started = asyncio.Event() + allow_model_turn_to_finish = asyncio.Event() + + async def fake_run_single_turn_streamed(*_args: object, **_kwargs: object) -> SingleStepResult: + next_model_turn_started.set() + await allow_model_turn_to_finish.wait() + return SingleStepResult( + original_input="input", + model_response=ModelResponse(output=[], usage=Usage(), response_id="unexpected"), + pre_step_items=[], + new_step_items=[], + next_step=NextStepFinalOutput("unexpected"), + tool_input_guardrail_results=[], + tool_output_guardrail_results=[], + ) + + monkeypatch.setattr(run_loop, "resolve_interrupted_turn", fake_resolve_interrupted_turn) + monkeypatch.setattr(run_loop, "run_single_turn_streamed", fake_run_single_turn_streamed) + + result = Runner.run_streamed(agent, state) + consumer_active = asyncio.Event() + consumer_suspended = asyncio.Event() + release_consumer = asyncio.Event() + cancel_called = asyncio.Event() + + async def consume_events() -> None: + async for event in result.stream_events(): + if event.type == "agent_updated_stream_event": + consumer_active.set() + if event.type == "run_item_stream_event" and event.name == "tool_output": + consumer_suspended.set() + await release_consumer.wait() + result.cancel(mode="after_turn") + cancel_called.set() + + consumer_task = asyncio.create_task(consume_events()) + await asyncio.wait_for(consumer_active.wait(), timeout=1) + allow_resume_resolution.set() + await asyncio.wait_for(consumer_suspended.wait(), timeout=1) + await asyncio.sleep(0) + await asyncio.sleep(0) + + assert not next_model_turn_started.is_set() + + release_consumer.set() + await asyncio.wait_for(cancel_called.wait(), timeout=1) + allow_model_turn_to_finish.set() + await asyncio.wait_for(consumer_task, timeout=1) + + assert not next_model_turn_started.is_set() + assert result.final_output is None + expected_agent = delegate if continuation == "handoff" else agent + assert result.current_agent is expected_agent + assert result.last_agent is expected_agent + assert result.to_state()._current_agent is expected_agent + if continuation == "handoff": + assert result._current_agent_output_schema is not None + assert isinstance(result._current_agent_output_schema, AgentOutputSchema) + assert result._current_agent_output_schema.output_type is int + + @pytest.mark.parametrize( ("conversation_id", "previous_response_id", "auto_previous_response_id"), [ diff --git a/tests/test_soft_cancel.py b/tests/test_soft_cancel.py index ddb51f8f17..fe94368dcb 100644 --- a/tests/test_soft_cancel.py +++ b/tests/test_soft_cancel.py @@ -1,13 +1,23 @@ """Tests for soft cancel (after_turn mode) functionality.""" +import asyncio import json +from collections.abc import AsyncGenerator +from typing import cast import pytest from agents import Agent, Runner, SQLiteSession +from agents.agent_output import AgentOutputSchema +from agents.stream_events import StreamEvent from .fake_model import FakeModel -from .test_responses import get_function_tool, get_function_tool_call, get_text_message +from .test_responses import ( + get_function_tool, + get_function_tool_call, + get_handoff_tool_call, + get_text_message, +) @pytest.mark.asyncio @@ -140,7 +150,8 @@ async def test_soft_cancel_tracks_usage(): @pytest.mark.asyncio -async def test_soft_cancel_stops_next_turn(): +@pytest.mark.parametrize("consumer_suspensions", [0, 1, 3]) +async def test_soft_cancel_stops_next_turn(consumer_suspensions: int): """Verify soft cancel prevents next turn from starting.""" model = FakeModel() agent = Agent( @@ -165,9 +176,137 @@ async def test_soft_cancel_stops_next_turn(): if event.type == "run_item_stream_event" and event.name == "tool_output": turns_completed += 1 if turns_completed == 1: + for _ in range(consumer_suspensions): + await asyncio.sleep(0) result.cancel(mode="after_turn") assert turns_completed == 1, "Should complete exactly 1 turn" + assert result.final_output is None + assert result.context_wrapper.usage.requests == 1 + + +@pytest.mark.asyncio +async def test_streamed_run_completes_without_an_event_consumer(): + """Turn acknowledgement must not block a run whose events are not consumed.""" + model = FakeModel() + model.add_multiple_turn_outputs( + [ + [get_function_tool_call("tool1", "{}")], + [get_text_message("Turn 2")], + ] + ) + agent = Agent( + name="Assistant", + model=model, + tools=[get_function_tool("tool1", "result1")], + ) + + result = Runner.run_streamed(agent, input="Hello") + assert result.run_loop_task is not None + await asyncio.wait_for(result.run_loop_task, timeout=1) + + assert result.final_output == "Turn 2" + assert result.context_wrapper.usage.requests == 2 + + +@pytest.mark.asyncio +async def test_closing_stream_consumer_releases_turn_acknowledgement(): + """Closing an iterator must not deadlock while a turn awaits its consumer.""" + model = FakeModel() + model.add_multiple_turn_outputs( + [ + [get_function_tool_call("tool1", "{}")], + [get_text_message("Turn 2")], + ] + ) + agent = Agent( + name="Assistant", + model=model, + tools=[get_function_tool("tool1", "result1")], + ) + + result = Runner.run_streamed(agent, input="Hello") + events = cast(AsyncGenerator[StreamEvent, None], result.stream_events()) + while True: + event = await anext(events) + if event.type == "run_item_stream_event" and event.name == "tool_output": + break + + await asyncio.wait_for(events.aclose(), timeout=1) + + assert result.final_output == "Turn 2" + assert result.context_wrapper.usage.requests == 2 + + +@pytest.mark.asyncio +async def test_cancelled_stream_consumer_releases_turn_acknowledgement(): + """Cancelling a consumer suspended after yield must release the completed turn.""" + model = FakeModel() + model.add_multiple_turn_outputs( + [ + [get_function_tool_call("tool1", "{}")], + [get_text_message("Turn 2")], + ] + ) + agent = Agent( + name="Assistant", + model=model, + tools=[get_function_tool("tool1", "result1")], + ) + + result = Runner.run_streamed(agent, input="Hello") + events = cast(AsyncGenerator[StreamEvent, None], result.stream_events()) + consumer_suspended = asyncio.Event() + keep_consumer_suspended = asyncio.Event() + + async def consume_events() -> None: + async for event in events: + if event.type == "run_item_stream_event" and event.name == "tool_output": + consumer_suspended.set() + await keep_consumer_suspended.wait() + + consumer_task = asyncio.create_task(consume_events()) + await asyncio.wait_for(consumer_suspended.wait(), timeout=1) + + consumer_task.cancel() + with pytest.raises(asyncio.CancelledError): + await consumer_task + + assert result.run_loop_task is not None + await asyncio.wait_for(result.run_loop_task, timeout=1) + + assert result.final_output == "Turn 2" + assert result.context_wrapper.usage.requests == 2 + assert result._active_stream_consumers == 0 + + await asyncio.wait_for(events.aclose(), timeout=1) + + +@pytest.mark.asyncio +async def test_immediate_cancel_releases_turn_acknowledgement(): + """Immediate cancellation must cancel a run waiting for streamed event acknowledgement.""" + model = FakeModel() + model.add_multiple_turn_outputs( + [ + [get_function_tool_call("tool1", "{}")], + [get_text_message("Turn 2")], + ] + ) + agent = Agent( + name="Assistant", + model=model, + tools=[get_function_tool("tool1", "result1")], + ) + + result = Runner.run_streamed(agent, input="Hello") + async for event in result.stream_events(): + if event.type == "run_item_stream_event" and event.name == "tool_output": + await asyncio.sleep(0) + result.cancel(mode="immediate") + + assert result.is_complete + assert result.final_output is None + assert result.context_wrapper.usage.requests == 1 @pytest.mark.asyncio @@ -436,6 +575,64 @@ async def on_invoke_handoff(context, data): await session.clear_session() +@pytest.mark.asyncio +async def test_soft_cancel_waits_for_handoff_event_consumption_before_next_turn(): + """A suspended handoff consumer can stop the run before the delegate model starts.""" + second_request_started = asyncio.Event() + + class HandoffModel(FakeModel): + def __init__(self) -> None: + super().__init__() + self.request_count = 0 + + async def stream_response(self, *args, **kwargs): + self.request_count += 1 + if self.request_count == 2: + second_request_started.set() + async for event in super().stream_response(*args, **kwargs): + yield event + + model = HandoffModel() + delegate = Agent(name="Delegate", model=model, output_type=int) + triage = Agent(name="Triage", model=model, handoffs=[delegate]) + model.add_multiple_turn_outputs( + [ + [get_handoff_tool_call(delegate)], + [get_text_message("Delegate response")], + ] + ) + + result = Runner.run_streamed(triage, input="Route this request") + consumer_suspended = asyncio.Event() + release_consumer = asyncio.Event() + + async def consume_events() -> None: + async for event in result.stream_events(): + if event.type == "run_item_stream_event" and event.name == "handoff_requested": + consumer_suspended.set() + await release_consumer.wait() + result.cancel(mode="after_turn") + + consumer_task = asyncio.create_task(consume_events()) + await asyncio.wait_for(consumer_suspended.wait(), timeout=1) + await asyncio.sleep(0) + await asyncio.sleep(0) + + assert not second_request_started.is_set() + + release_consumer.set() + await asyncio.wait_for(consumer_task, timeout=1) + + assert result.final_output is None + assert result.context_wrapper.usage.requests == 1 + assert result.current_agent is delegate + assert result.last_agent is delegate + assert result.to_state()._current_agent is delegate + assert result._current_agent_output_schema is not None + assert isinstance(result._current_agent_output_schema, AgentOutputSchema) + assert result._current_agent_output_schema.output_type is int + + @pytest.mark.asyncio async def test_soft_cancel_with_session_and_multiple_turns(): """Verify soft cancel with session across multiple turns."""