Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
23 changes: 5 additions & 18 deletions src/agents/agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -58,7 +58,7 @@
)
from .tool_context import ToolContext
from .util import _transforms
from .util._asyncio_tasks import gather_with_cancel
from .util._asyncio_tasks import gather_with_cancel, run_producer_consumer
from .util._types import MaybeAwaitable

if TYPE_CHECKING:
Expand Down Expand Up @@ -912,10 +912,7 @@ async def dispatch_stream_events() -> None:
if is_sentinel:
break

dispatch_task = asyncio.create_task(dispatch_stream_events())
stream_iteration_cancelled = False

try:
async def enqueue_stream_events() -> None:
from .stream_events import AgentUpdatedStreamEvent

current_agent = run_result_streaming.current_agent
Expand All @@ -930,20 +927,10 @@ async def dispatch_stream_events() -> None:
"tool_call": context.tool_call,
}
await event_queue.put(payload)
except asyncio.CancelledError:
stream_iteration_cancelled = True
raise
finally:
if stream_iteration_cancelled:
dispatch_task.cancel()
try:
await dispatch_task
except asyncio.CancelledError:
pass
else:
finally:
await event_queue.put(None)
await event_queue.join()
await dispatch_task

await run_producer_consumer(enqueue_stream_events(), dispatch_stream_events())
run_result = run_result_streaming
else:
run_result = await Runner.run(
Expand Down
126 changes: 65 additions & 61 deletions src/agents/extensions/experimental/codex/codex_tool.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,7 @@
from agents.tool_context import ToolContext
from agents.tracing import SpanError, custom_span
from agents.usage import Usage as AgentsUsage, _make_input_tokens_details
from agents.util._asyncio_tasks import run_producer_consumer
from agents.util._types import MaybeAwaitable

from .codex import Codex
Expand Down Expand Up @@ -1047,78 +1048,81 @@ async def _consume_events(
resolved_thread_id_holder["thread_id"] = resolved_thread_id

event_queue: asyncio.Queue[CodexToolStreamEvent | None] | None = None
dispatch_task: asyncio.Task[None] | None = None

if on_stream is not None:
# Buffer events so user callbacks cannot block the Codex stream loop.
event_queue = asyncio.Queue()

async def _run_handler(payload: CodexToolStreamEvent) -> None:
# Dispatch user callbacks asynchronously to avoid blocking the stream.
async def _run_handler(payload: CodexToolStreamEvent) -> None:
# Dispatch user callbacks asynchronously to avoid blocking the stream.
assert on_stream is not None
try:
maybe_result = on_stream(payload)
if inspect.isawaitable(maybe_result):
await maybe_result
except Exception as exc:
log_model_and_tool_action_error(
logger,
"Error while handling Codex on_stream event",
exc,
)

async def _dispatch() -> None:
assert event_queue is not None
while True:
payload = await event_queue.get()
is_sentinel = payload is None
try:
maybe_result = on_stream(payload)
if inspect.isawaitable(maybe_result):
await maybe_result
except Exception as exc:
log_model_and_tool_action_error(
logger,
"Error while handling Codex on_stream event",
exc,
)
if payload is not None:
await _run_handler(payload)
finally:
event_queue.task_done()
if is_sentinel:
break

async def _dispatch() -> None:
assert event_queue is not None
while True:
payload = await event_queue.get()
is_sentinel = payload is None
try:
if payload is not None:
await _run_handler(payload)
finally:
event_queue.task_done()
if is_sentinel:
break
async def _process_events() -> None:
nonlocal final_response, resolved_thread_id, usage

dispatch_task = asyncio.create_task(_dispatch())
try:
async for raw_event in events:
event = coerce_thread_event(raw_event)
if event_queue is not None:
await event_queue.put(
CodexToolStreamEvent(
event=event,
thread=thread,
tool_call=ctx.tool_call,
)
)

try:
async for raw_event in events:
event = coerce_thread_event(raw_event)
if isinstance(event, ItemStartedEvent):
_handle_item_started(event.item, active_spans, span_data_max_chars)
elif isinstance(event, ItemUpdatedEvent):
_handle_item_updated(event.item, active_spans, span_data_max_chars)
elif isinstance(event, ItemCompletedEvent):
_handle_item_completed(event.item, active_spans, span_data_max_chars)
if is_agent_message_item(event.item):
final_response = event.item.text
elif isinstance(event, TurnCompletedEvent):
usage = event.usage
elif isinstance(event, ThreadStartedEvent):
resolved_thread_id = event.thread_id
if resolved_thread_id_holder is not None:
resolved_thread_id_holder["thread_id"] = resolved_thread_id
elif isinstance(event, TurnFailedEvent):
error = event.error.message
raise UserError(f"Codex turn failed{(': ' + error) if error else ''}")
elif isinstance(event, ThreadErrorEvent):
raise UserError(f"Codex stream error: {event.message}")
finally:
if event_queue is not None:
await event_queue.put(
CodexToolStreamEvent(
event=event,
thread=thread,
tool_call=ctx.tool_call,
)
)
await event_queue.put(None)

if isinstance(event, ItemStartedEvent):
_handle_item_started(event.item, active_spans, span_data_max_chars)
elif isinstance(event, ItemUpdatedEvent):
_handle_item_updated(event.item, active_spans, span_data_max_chars)
elif isinstance(event, ItemCompletedEvent):
_handle_item_completed(event.item, active_spans, span_data_max_chars)
if is_agent_message_item(event.item):
final_response = event.item.text
elif isinstance(event, TurnCompletedEvent):
usage = event.usage
elif isinstance(event, ThreadStartedEvent):
resolved_thread_id = event.thread_id
if resolved_thread_id_holder is not None:
resolved_thread_id_holder["thread_id"] = resolved_thread_id
elif isinstance(event, TurnFailedEvent):
error = event.error.message
raise UserError(f"Codex turn failed{(': ' + error) if error else ''}")
elif isinstance(event, ThreadErrorEvent):
raise UserError(f"Codex stream error: {event.message}")
try:
if on_stream is None:
await _process_events()
else:
await run_producer_consumer(_process_events(), _dispatch())
finally:
if event_queue is not None:
await event_queue.put(None)
await event_queue.join()
if dispatch_task is not None:
await dispatch_task

# Ensure any open spans are closed even on failure.
for span in active_spans.values():
span.finish()
Expand Down
10 changes: 7 additions & 3 deletions src/agents/sandbox/memory/manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -131,11 +131,15 @@ async def flush(self) -> None:
self._ensure_worker()
for rollout_file in rollout_files:
self._queue.put_nowait(rollout_file)
await self._queue.join()
if self._worker_task is not None:
self._queue.put_nowait(_STOP)
await self._worker_task
self._worker_task = None
worker_task = self._worker_task
try:
# The stop marker follows every rollout, so worker completion implies
# that all preceding rollout files were processed.
await worker_task
finally:
self._worker_task = None
await self._run_phase_two()
finally:
_unregister_memory_generation_manager(session=self._session, manager=self)
Expand Down
40 changes: 40 additions & 0 deletions src/agents/util/_asyncio_tasks.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,8 @@
T4 = TypeVar("T4")
T5 = TypeVar("T5")
T6 = TypeVar("T6")
TProducer = TypeVar("TProducer")
TConsumer = TypeVar("TConsumer")


def _consume_future_exception(future: asyncio.Future[Any]) -> None:
Expand Down Expand Up @@ -110,3 +112,41 @@ async def gather_with_cancel(
task.cancel()
await asyncio.gather(*tasks, return_exceptions=True)
raise


async def run_producer_consumer(
producer: Awaitable[TProducer],
consumer: Awaitable[TConsumer],
/,
) -> tuple[TProducer, TConsumer]:
"""Run a producer and consumer with asymmetric failure handling.

The producer must signal completion to the consumer in a ``finally`` block. A producer
failure waits for the consumer to drain before propagating, while a consumer failure or
parent cancellation cancels and drains the sibling task.
"""
producer_task = asyncio.ensure_future(producer)
consumer_task = asyncio.ensure_future(consumer)
tasks = (producer_task, consumer_task)

try:
done, _ = await asyncio.wait(tasks, return_when=asyncio.FIRST_COMPLETED)
if consumer_task in done:
consumer_result = consumer_task.result()
producer_result = await producer_task
return producer_result, consumer_result

try:
producer_result = producer_task.result()
except BaseException:
await consumer_task
raise

consumer_result = await consumer_task
return producer_result, consumer_result
except BaseException:
for task in tasks:
if not task.done():
task.cancel()
await asyncio.gather(*tasks, return_exceptions=True)
raise
Loading