Skip to content

Commit bedd7a9

Browse files
committed
fix(realtime): address Codex review for input guardrails
Snapshot the active agent and its input guardrails when the transcription event is handled so a concurrent update_agent()/handoff cannot run a different agent's guardrails, and run the input guardrails concurrently so a slow guardrail cannot delay the forced response cancel.
1 parent 6b3646c commit bedd7a9

1 file changed

Lines changed: 40 additions & 24 deletions

File tree

src/agents/realtime/session.py

Lines changed: 40 additions & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -17,6 +17,7 @@
1717
)
1818
from ..agent import Agent
1919
from ..exceptions import ToolInputGuardrailTripwireTriggered, UserError
20+
from ..guardrail import InputGuardrail, InputGuardrailResult
2021
from ..handoffs import Handoff
2122
from ..items import ToolApprovalItem
2223
from ..logger import logger
@@ -1270,37 +1271,32 @@ async def _run_output_guardrails(self, text: str, response_id: str) -> bool:
12701271

12711272
return False
12721273

1273-
async def _run_input_guardrails(self, text: str, item_id: str) -> bool:
1274+
async def _run_input_guardrails(
1275+
self,
1276+
text: str,
1277+
item_id: str,
1278+
agent: RealtimeAgent,
1279+
input_guardrails: list[InputGuardrail[Any]],
1280+
) -> bool:
12741281
"""Run input guardrails on the user's transcribed input. Returns True if any guardrail was
12751282
triggered.
1276-
"""
1277-
combined_guardrails = self._current_agent.input_guardrails + self._run_config.get(
1278-
"input_guardrails", []
1279-
)
1280-
seen_ids: set[int] = set()
1281-
input_guardrails = []
1282-
for guardrail in combined_guardrails:
1283-
guardrail_id = id(guardrail)
1284-
if guardrail_id not in seen_ids:
1285-
input_guardrails.append(guardrail)
1286-
seen_ids.add(guardrail_id)
12871283
1284+
``agent`` and ``input_guardrails`` are snapshotted when the transcription event is handled
1285+
so that a concurrent ``update_agent()`` or handoff cannot swap in a different agent's
1286+
guardrails before this background task runs.
1287+
"""
12881288
# If we've already interrupted the response for this user item, skip.
12891289
if not input_guardrails or item_id in self._interrupted_input_item_ids:
12901290
return False
12911291

1292-
triggered_results = []
1293-
1294-
for guardrail in input_guardrails:
1292+
async def _run_one(guardrail: InputGuardrail[Any]) -> InputGuardrailResult | None:
12951293
try:
1296-
result = await guardrail.run(
1294+
return await guardrail.run(
12971295
# TODO (rm) Remove this cast, it's wrong
1298-
cast(Agent[Any], self._current_agent),
1296+
cast(Agent[Any], agent),
12991297
text,
13001298
self._context_wrapper,
13011299
)
1302-
if result.output.tripwire_triggered:
1303-
triggered_results.append(result)
13041300
except Exception as exc:
13051301
logger.warning(
13061302
"Input guardrail %r raised %s: %s; skipping it.",
@@ -1309,7 +1305,14 @@ async def _run_input_guardrails(self, text: str, item_id: str) -> bool:
13091305
exc,
13101306
)
13111307
logger.debug("Input guardrail failure details.", exc_info=True)
1312-
continue
1308+
return None
1309+
1310+
# Run the guardrails concurrently so a slow guardrail cannot delay the forced cancel behind
1311+
# unrelated guardrails, which would let the unsafe turn keep generating.
1312+
results = await asyncio.gather(*(_run_one(guardrail) for guardrail in input_guardrails))
1313+
triggered_results = [
1314+
result for result in results if result is not None and result.output.tripwire_triggered
1315+
]
13131316

13141317
if triggered_results:
13151318
# Double-check: bail if already interrupted for this user item.
@@ -1353,14 +1356,27 @@ def _enqueue_guardrail_task(self, text: str, response_id: str) -> None:
13531356
task.add_done_callback(self._on_guardrail_task_done)
13541357

13551358
def _enqueue_input_guardrail_task(self, text: str, item_id: str) -> None:
1359+
# Snapshot the active agent and its guardrails now; a later update_agent()/handoff must not
1360+
# change which guardrails run against this transcript.
1361+
agent = self._current_agent
1362+
combined_guardrails = agent.input_guardrails + self._run_config.get("input_guardrails", [])
1363+
1364+
seen_ids: set[int] = set()
1365+
input_guardrails: list[InputGuardrail[Any]] = []
1366+
for guardrail in combined_guardrails:
1367+
guardrail_id = id(guardrail)
1368+
if guardrail_id not in seen_ids:
1369+
input_guardrails.append(guardrail)
1370+
seen_ids.add(guardrail_id)
1371+
13561372
# Skip creating a no-op task when no input guardrails are configured.
1357-
if not self._current_agent.input_guardrails and not self._run_config.get(
1358-
"input_guardrails"
1359-
):
1373+
if not input_guardrails:
13601374
return
13611375

13621376
# Runs the input guardrails in a separate task to avoid blocking the main loop.
1363-
task = asyncio.create_task(self._run_input_guardrails(text, item_id))
1377+
task = asyncio.create_task(
1378+
self._run_input_guardrails(text, item_id, agent, input_guardrails)
1379+
)
13641380
# Reuse the shared guardrail task set + done callback so completed tasks are removed,
13651381
# exceptions surface as events, and close() cancels any still-running task.
13661382
self._guardrail_tasks.add(task)

0 commit comments

Comments
 (0)