diff --git a/src/agents/run_config.py b/src/agents/run_config.py index fcc9b01315..5a150ee2dd 100644 --- a/src/agents/run_config.py +++ b/src/agents/run_config.py @@ -329,6 +329,14 @@ class RunConfig: the run continue. """ + max_model_retries: int = 0 + """Maximum number of automatic retries when the model produces a malformed response that + triggers a ``ModelBehaviorError`` (e.g. invalid tool call JSON, nonexistent tool, etc.). + + On each retry, the error message is fed back to the model as a synthetic user message so + the model can self-correct. Defaults to 0 (no retry). + """ + class RunOptions(TypedDict, Generic[TContext]): """Arguments for ``AgentRunner`` methods.""" diff --git a/src/agents/run_internal/run_loop.py b/src/agents/run_internal/run_loop.py index 45f09c0fa0..42530384ec 100644 --- a/src/agents/run_internal/run_loop.py +++ b/src/agents/run_internal/run_loop.py @@ -1345,8 +1345,6 @@ def _tool_search_fingerprint(raw_item: Any) -> str: model_settings = get_model_settings(execution_agent, run_config) model_settings = maybe_reset_tool_choice(public_agent, tool_use_tracker, model_settings) - final_response: ModelResponse | None = None - streamed_response_output: list[ResponseOutputItem] = [] if server_conversation_tracker is not None: items_for_input = ( @@ -1455,212 +1453,230 @@ async def rewind_model_request() -> None: if server_conversation_tracker is not None: server_conversation_tracker.rewind_input(filtered.input) - stream_failed_retry_attempts: list[int] = [0] + max_retries = run_config.max_model_retries + for attempt in range(max_retries + 1): + final_response = None + streamed_response_output = [] + stream_failed_retry_attempts: list[int] = [0] - retry_stream = stream_response_with_retry( - get_stream=lambda: model.stream_response( - filtered.instructions, - filtered.input, - model_settings, - all_tools, - output_schema, - handoffs, - get_model_tracing_impl( - run_config.tracing_disabled, run_config.trace_include_sensitive_data + retry_stream = stream_response_with_retry( + get_stream=lambda: model.stream_response( + filtered.instructions, + filtered.input, + model_settings, + all_tools, + output_schema, + handoffs, + get_model_tracing_impl( + run_config.tracing_disabled, run_config.trace_include_sensitive_data + ), + previous_response_id=previous_response_id, + conversation_id=conversation_id, + prompt=prompt_config, ), + rewind=rewind_model_request, + retry_settings=model_settings.retry, + get_retry_advice=model.get_retry_advice, previous_response_id=previous_response_id, conversation_id=conversation_id, - prompt=prompt_config, - ), - rewind=rewind_model_request, - retry_settings=model_settings.retry, - get_retry_advice=model.get_retry_advice, - previous_response_id=previous_response_id, - conversation_id=conversation_id, - failed_retry_attempts_out=stream_failed_retry_attempts, - ) + failed_retry_attempts_out=stream_failed_retry_attempts, + ) - async for event in retry_stream: - streamed_result._event_queue.put_nowait(RawResponsesStreamEvent(data=event)) - - terminal_response: Response | None = None - is_completed_event = False - if isinstance(event, ResponseCompletedEvent): - is_completed_event = True - terminal_response = event.response - elif getattr(event, "type", None) in {"response.incomplete", "response.failed"}: - event_type = cast(str, event.type) - maybe_response = getattr(event, "response", None) - raise response_terminal_failure_error( - event_type, - maybe_response if isinstance(maybe_response, Response) else None, - ) - elif getattr(event, "type", None) in {"error", "response.error"}: - raise response_error_event_failure_error(cast(str, event.type), event) - - if terminal_response is not None: - if is_completed_event and not terminal_response.output and streamed_response_output: - # Some streaming backends emit output items during item.done events while leaving - # the terminal response output empty. Preserve those items so the runner can - # resolve the completed step correctly. - terminal_response.output = list(streamed_response_output) - usage = ( - apply_retry_attempt_usage( - Usage( - requests=1, - input_tokens=terminal_response.usage.input_tokens, - output_tokens=terminal_response.usage.output_tokens, - total_tokens=terminal_response.usage.total_tokens, - input_tokens_details=terminal_response.usage.input_tokens_details, - output_tokens_details=terminal_response.usage.output_tokens_details, - ), - stream_failed_retry_attempts[0], + async for event in retry_stream: + streamed_result._event_queue.put_nowait(RawResponsesStreamEvent(data=event)) + + terminal_response: Response | None = None + is_completed_event = False + if isinstance(event, ResponseCompletedEvent): + is_completed_event = True + terminal_response = event.response + elif getattr(event, "type", None) in {"response.incomplete", "response.failed"}: + event_type = cast(str, event.type) + maybe_response = getattr(event, "response", None) + raise response_terminal_failure_error( + event_type, + maybe_response if isinstance(maybe_response, Response) else None, ) - if terminal_response.usage - else Usage() - ) - final_response = ModelResponse( - output=terminal_response.output, - usage=usage, - response_id=terminal_response.id, - request_id=getattr(terminal_response, "_request_id", None), - ) - - if isinstance(event, ResponseOutputItemDoneEvent): - output_item = event.item - streamed_response_output.append(output_item) - output_item_type = getattr(output_item, "type", None) - - if output_item_type == "tool_search_call": - emitted_tool_search_fingerprints.add(_tool_search_fingerprint(output_item)) - streamed_result._event_queue.put_nowait( - RunItemStreamEvent( - item=ToolSearchCallItem( - raw_item=coerce_tool_search_call_raw_item(output_item), - agent=public_agent, + elif getattr(event, "type", None) in {"error", "response.error"}: + raise response_error_event_failure_error(cast(str, event.type), event) + + if terminal_response is not None: + if is_completed_event and not terminal_response.output and streamed_response_output: + # Some streaming backends emit output items during item.done events while leaving + # the terminal response output empty. Preserve those items so the runner can + # resolve the completed step correctly. + terminal_response.output = list(streamed_response_output) + usage = ( + apply_retry_attempt_usage( + Usage( + requests=1, + input_tokens=terminal_response.usage.input_tokens, + output_tokens=terminal_response.usage.output_tokens, + total_tokens=terminal_response.usage.total_tokens, + input_tokens_details=terminal_response.usage.input_tokens_details, + output_tokens_details=terminal_response.usage.output_tokens_details, ), - name="tool_search_called", + stream_failed_retry_attempts[0], ) + if terminal_response.usage + else Usage() + ) + final_response = ModelResponse( + output=terminal_response.output, + usage=usage, + response_id=terminal_response.id, + request_id=getattr(terminal_response, "_request_id", None), ) - elif output_item_type == "tool_search_output": - emitted_tool_search_fingerprints.add(_tool_search_fingerprint(output_item)) - streamed_result._event_queue.put_nowait( - RunItemStreamEvent( - item=ToolSearchOutputItem( - raw_item=coerce_tool_search_output_raw_item(output_item), - agent=public_agent, - ), - name="tool_search_output_created", + if isinstance(event, ResponseOutputItemDoneEvent): + output_item = event.item + streamed_response_output.append(output_item) + output_item_type = getattr(output_item, "type", None) + + if output_item_type == "tool_search_call": + emitted_tool_search_fingerprints.add(_tool_search_fingerprint(output_item)) + streamed_result._event_queue.put_nowait( + RunItemStreamEvent( + item=ToolSearchCallItem( + raw_item=coerce_tool_search_call_raw_item(output_item), + agent=public_agent, + ), + name="tool_search_called", + ) ) - ) - elif isinstance(output_item, McpListTools): - hosted_mcp_tool_metadata.update(collect_mcp_list_tools_metadata([output_item])) + elif output_item_type == "tool_search_output": + emitted_tool_search_fingerprints.add(_tool_search_fingerprint(output_item)) + streamed_result._event_queue.put_nowait( + RunItemStreamEvent( + item=ToolSearchOutputItem( + raw_item=coerce_tool_search_output_raw_item(output_item), + agent=public_agent, + ), + name="tool_search_output_created", + ) + ) - elif isinstance(output_item, TOOL_CALL_TYPES): - output_call_id: str | None = getattr( - output_item, "call_id", getattr(output_item, "id", None) - ) + elif isinstance(output_item, McpListTools): + hosted_mcp_tool_metadata.update(collect_mcp_list_tools_metadata([output_item])) - if ( - output_call_id - and isinstance(output_call_id, str) - and output_call_id not in emitted_tool_call_ids - ): - emitted_tool_call_ids.add(output_call_id) - - # Look up tool description from precomputed map ("last wins" matches - # execution behavior in process_model_response). - tool_lookup_key = get_function_tool_lookup_key_for_call(output_item) - matched_tool = ( - tool_map.get(tool_lookup_key) if tool_lookup_key is not None else None + elif isinstance(output_item, TOOL_CALL_TYPES): + output_call_id: str | None = getattr( + output_item, "call_id", getattr(output_item, "id", None) ) + if ( - matched_tool is None - and output_schema is not None - and isinstance(output_item, ResponseFunctionToolCall) - and output_item.name == "json_tool_call" + output_call_id + and isinstance(output_call_id, str) + and output_call_id not in emitted_tool_call_ids ): - matched_tool = build_litellm_json_tool_call(output_item) - tool_description: str | None = None - tool_title: str | None = None - tool_origin = None - if isinstance(output_item, McpCall): - metadata = hosted_mcp_tool_metadata.get( - (output_item.server_label, output_item.name) + emitted_tool_call_ids.add(output_call_id) + + # Look up tool description from precomputed map ("last wins" matches + # execution behavior in process_model_response). + tool_lookup_key = get_function_tool_lookup_key_for_call(output_item) + matched_tool = ( + tool_map.get(tool_lookup_key) if tool_lookup_key is not None else None ) - if metadata is not None: - tool_description = metadata.description - tool_title = metadata.title - tool_origin = ToolOrigin( - type=ToolOriginType.MCP, - mcp_server_name=output_item.server_label, + if ( + matched_tool is None + and output_schema is not None + and isinstance(output_item, ResponseFunctionToolCall) + and output_item.name == "json_tool_call" + ): + matched_tool = build_litellm_json_tool_call(output_item) + tool_description: str | None = None + tool_title: str | None = None + tool_origin = None + if isinstance(output_item, McpCall): + metadata = hosted_mcp_tool_metadata.get( + (output_item.server_label, output_item.name) + ) + if metadata is not None: + tool_description = metadata.description + tool_title = metadata.title + tool_origin = ToolOrigin( + type=ToolOriginType.MCP, + mcp_server_name=output_item.server_label, + ) + elif matched_tool is not None: + tool_description = getattr(matched_tool, "description", None) + tool_title = getattr(matched_tool, "_mcp_title", None) + tool_origin = get_function_tool_origin(matched_tool) + + tool_item = ToolCallItem( + raw_item=cast(ToolCallItemTypes, output_item), + agent=public_agent, + description=tool_description, + title=tool_title, + tool_origin=tool_origin, + ) + streamed_result._event_queue.put_nowait( + RunItemStreamEvent(item=tool_item, name="tool_called") ) - elif matched_tool is not None: - tool_description = getattr(matched_tool, "description", None) - tool_title = getattr(matched_tool, "_mcp_title", None) - tool_origin = get_function_tool_origin(matched_tool) - - tool_item = ToolCallItem( - raw_item=cast(ToolCallItemTypes, output_item), - agent=public_agent, - description=tool_description, - title=tool_title, - tool_origin=tool_origin, - ) - streamed_result._event_queue.put_nowait( - RunItemStreamEvent(item=tool_item, name="tool_called") - ) - elif isinstance(output_item, ResponseReasoningItem): - reasoning_id: str | None = getattr(output_item, "id", None) + elif isinstance(output_item, ResponseReasoningItem): + reasoning_id: str | None = getattr(output_item, "id", None) - if reasoning_id and reasoning_id not in emitted_reasoning_item_ids: - emitted_reasoning_item_ids.add(reasoning_id) + if reasoning_id and reasoning_id not in emitted_reasoning_item_ids: + emitted_reasoning_item_ids.add(reasoning_id) - reasoning_item = ReasoningItem(raw_item=output_item, agent=public_agent) - streamed_result._event_queue.put_nowait( - RunItemStreamEvent(item=reasoning_item, name="reasoning_item_created") - ) + reasoning_item = ReasoningItem(raw_item=output_item, agent=public_agent) + streamed_result._event_queue.put_nowait( + RunItemStreamEvent(item=reasoning_item, name="reasoning_item_created") + ) - if final_response is not None: - context_wrapper.usage.add(final_response.usage) - await asyncio.gather( - ( - public_agent.hooks.on_llm_end(context_wrapper, public_agent, final_response) - if public_agent.hooks - else _coro.noop_coroutine() - ), - hooks.on_llm_end(context_wrapper, public_agent, final_response), - ) + if final_response is not None: + context_wrapper.usage.add(final_response.usage) + await asyncio.gather( + ( + public_agent.hooks.on_llm_end(context_wrapper, public_agent, final_response) + if public_agent.hooks + else _coro.noop_coroutine() + ), + hooks.on_llm_end(context_wrapper, public_agent, final_response), + ) - if not final_response: - raise ModelBehaviorError("Model did not produce a final response!") + if not final_response: + raise ModelBehaviorError("Model did not produce a final response!") - if server_conversation_tracker is not None: - # Streaming uses the same rewind helper, so a successful retry must restore delivered - # input tracking before the next turn computes server-managed deltas. - server_conversation_tracker.mark_input_as_sent(filtered.input) - server_conversation_tracker.track_server_items(final_response) - - single_step_result = await get_single_step_result_from_response( - bindings=bindings, - original_input=streamed_result.input, - pre_step_items=streamed_result._model_input_items, - new_response=final_response, - output_schema=output_schema, - all_tools=all_tools, - handoffs=handoffs, - hooks=hooks, - context_wrapper=context_wrapper, - run_config=run_config, - error_handlers=error_handlers, - tool_use_tracker=tool_use_tracker, - server_manages_conversation=server_conversation_tracker is not None, - event_queue=streamed_result._event_queue, - before_side_effects=raise_if_input_guardrail_tripwire_known, - ) + if server_conversation_tracker is not None: + # Streaming uses the same rewind helper, so a successful retry must restore delivered + # input tracking before the next turn computes server-managed deltas. + server_conversation_tracker.mark_input_as_sent(filtered.input) + server_conversation_tracker.track_server_items(final_response) + + try: + single_step_result = await get_single_step_result_from_response( + bindings=bindings, + original_input=streamed_result.input, + pre_step_items=streamed_result._model_input_items, + new_response=final_response, + output_schema=output_schema, + all_tools=all_tools, + handoffs=handoffs, + hooks=hooks, + context_wrapper=context_wrapper, + run_config=run_config, + error_handlers=error_handlers, + tool_use_tracker=tool_use_tracker, + server_manages_conversation=server_conversation_tracker is not None, + event_queue=streamed_result._event_queue, + before_side_effects=raise_if_input_guardrail_tripwire_known, + ) + break + except ModelBehaviorError as e: + if attempt >= max_retries: + raise + logger.debug( + "ModelBehaviorError on attempt %s/%s: %s. Retrying with feedback.", + attempt + 1, + max_retries + 1, + e.message, + ) + filtered.input = list(filtered.input) + [ + {"role": "user", "content": f"Your previous response was invalid: {e.message}"} + ] items_to_filter = session_items_for_turn(single_step_result) @@ -1760,39 +1776,57 @@ async def run_single_turn( else: input = _prepare_turn_input_items(original_input, generated_items, reasoning_item_id_policy) - new_response = await get_new_response( - bindings, - system_prompt, - input, - output_schema, - all_tools, - handoffs, - hooks, - context_wrapper, - run_config, - tool_use_tracker, - server_conversation_tracker, - prompt_config, - session=session, - session_items_to_rewind=session_items_to_rewind, - prompt_cache_key_resolver=prompt_cache_key_resolver, - ) + max_retries = run_config.max_model_retries + retry_input = input + for attempt in range(max_retries + 1): + new_response = await get_new_response( + bindings, + system_prompt, + retry_input, + output_schema, + all_tools, + handoffs, + hooks, + context_wrapper, + run_config, + tool_use_tracker, + server_conversation_tracker + if attempt == 0 + else None, + prompt_config, + session=session, + session_items_to_rewind=session_items_to_rewind if attempt == 0 else None, + prompt_cache_key_resolver=prompt_cache_key_resolver if attempt == 0 else None, + ) - return await get_single_step_result_from_response( - bindings=bindings, - original_input=original_input, - pre_step_items=generated_items, - new_response=new_response, - output_schema=output_schema, - all_tools=all_tools, - handoffs=handoffs, - hooks=hooks, - context_wrapper=context_wrapper, - run_config=run_config, - error_handlers=error_handlers, - tool_use_tracker=tool_use_tracker, - server_manages_conversation=server_conversation_tracker is not None, - ) + try: + return await get_single_step_result_from_response( + bindings=bindings, + original_input=original_input, + pre_step_items=generated_items, + new_response=new_response, + output_schema=output_schema, + all_tools=all_tools, + handoffs=handoffs, + hooks=hooks, + context_wrapper=context_wrapper, + run_config=run_config, + error_handlers=error_handlers, + tool_use_tracker=tool_use_tracker, + server_manages_conversation=server_conversation_tracker is not None, + ) + except ModelBehaviorError as e: + if attempt >= max_retries: + raise + logger.debug( + "ModelBehaviorError on attempt %s/%s: %s. Retrying with feedback.", + attempt + 1, + max_retries + 1, + e.message, + ) + retry_input = list(retry_input) + [ + {"role": "user", "content": f"Your previous response was invalid: {e.message}"} + ] async def get_new_response(