codekingpro/portable-devtools
114k
1"""Tool retry middleware for agents."""2 3from __future__ import annotations4 5import asyncio6import time7import warnings8from typing import TYPE_CHECKING, Any9 10from langchain_core.messages import ToolMessage11 12from langchain.agents.middleware._retry import (13 OnFailure,14 RetryOn,15 calculate_delay,16 should_retry_exception,17 validate_retry_params,18)19from langchain.agents.middleware.types import AgentMiddleware, AgentState, ContextT, ResponseT20 21if TYPE_CHECKING:22 from collections.abc import Awaitable, Callable23 24 from langgraph.types import Command25 26 from langchain.agents.middleware.types import ToolCallRequest27 from langchain.tools import BaseTool28 29 30class ToolRetryMiddleware(AgentMiddleware[AgentState[ResponseT], ContextT, ResponseT]):31 """Middleware that automatically retries failed tool calls with configurable backoff.32 33 Supports retrying on specific exceptions and exponential backoff.34 35 Examples:36 !!! example "Basic usage with default settings (2 retries, exponential backoff)"37 38 ```python39 from langchain.agents import create_agent40 from langchain.agents.middleware import ToolRetryMiddleware41 42 agent = create_agent(model, tools=[search_tool], middleware=[ToolRetryMiddleware()])43 ```44 45 !!! example "Retry specific exceptions only"46 47 ```python48 from requests.exceptions import RequestException, Timeout49 50 retry = ToolRetryMiddleware(51 max_retries=4,52 retry_on=(RequestException, Timeout),53 backoff_factor=1.5,54 )55 ```56 57 !!! example "Custom exception filtering"58 59 ```python60 from requests.exceptions import HTTPError61 62 63 def should_retry(exc: Exception) -> bool:64 # Only retry on 5xx errors65 if isinstance(exc, HTTPError):66 return 500 <= exc.status_code < 60067 return False68 69 70 retry = ToolRetryMiddleware(71 max_retries=3,72 retry_on=should_retry,73 )74 ```75 76 !!! example "Apply to specific tools with custom error handling"77 78 ```python79 def format_error(exc: Exception) -> str:80 return "Database temporarily unavailable. Please try again later."81 82 83 retry = ToolRetryMiddleware(84 max_retries=4,85 tools=["search_database"],86 on_failure=format_error,87 )88 ```89 90 !!! example "Apply to specific tools using `BaseTool` instances"91 92 ```python93 from langchain_core.tools import tool94 95 96 @tool97 def search_database(query: str) -> str:98 '''Search the database.'''99 return results100 101 102 retry = ToolRetryMiddleware(103 max_retries=4,104 tools=[search_database], # Pass BaseTool instance105 )106 ```107 108 !!! example "Constant backoff (no exponential growth)"109 110 ```python111 retry = ToolRetryMiddleware(112 max_retries=5,113 backoff_factor=0.0, # No exponential growth114 initial_delay=2.0, # Always wait 2 seconds115 )116 ```117 118 !!! example "Raise exception on failure"119 120 ```python121 retry = ToolRetryMiddleware(122 max_retries=2,123 on_failure="error", # Re-raise exception instead of returning message124 )125 ```126 """127 128 def __init__(129 self,130 *,131 max_retries: int = 2,132 tools: list[BaseTool | str] | None = None,133 retry_on: RetryOn = (Exception,),134 on_failure: OnFailure = "continue",135 backoff_factor: float = 2.0,136 initial_delay: float = 1.0,137 max_delay: float = 60.0,138 jitter: bool = True,139 ) -> None:140 """Initialize `ToolRetryMiddleware`.141 142 Args:143 max_retries: Maximum number of retry attempts after the initial call.144 145 Must be `>= 0`.146 tools: Optional list of tools or tool names to apply retry logic to.147 148 Can be a list of `BaseTool` instances or tool name strings.149 150 If `None`, applies to all tools.151 retry_on: Either a tuple of exception types to retry on, or a callable152 that takes an exception and returns `True` if it should be retried.153 154 Default is to retry on all exceptions.155 on_failure: Behavior when all retries are exhausted.156 157 Options:158 159 - `'continue'`: Return a `ToolMessage` with error details,160 allowing the LLM to handle the failure and potentially recover.161 - `'error'`: Re-raise the exception, stopping agent execution.162 - **Custom callable:** Function that takes the exception and returns a163 string for the `ToolMessage` content, allowing custom error164 formatting.165 166 **Deprecated values** (for backwards compatibility):167 168 - `'return_message'`: Use `'continue'` instead.169 - `'raise'`: Use `'error'` instead.170 backoff_factor: Multiplier for exponential backoff.171 172 Each retry waits `initial_delay * (backoff_factor ** retry_number)`173 seconds.174 175 Set to `0.0` for constant delay.176 initial_delay: Initial delay in seconds before first retry.177 max_delay: Maximum delay in seconds between retries.178 179 Caps exponential backoff growth.180 jitter: Whether to add random jitter (`±25%`) to delay to avoid thundering herd.181 182 Raises:183 ValueError: If `max_retries < 0` or delays are negative.184 """185 super().__init__()186 187 # Validate parameters188 validate_retry_params(max_retries, initial_delay, max_delay, backoff_factor)189 190 # Handle backwards compatibility for deprecated on_failure values191 if on_failure == "raise": # type: ignore[comparison-overlap]192 msg = ( # type: ignore[unreachable]193 "on_failure='raise' is deprecated and will be removed in a future version. "194 "Use on_failure='error' instead."195 )196 warnings.warn(msg, DeprecationWarning, stacklevel=2)197 on_failure = "error"198 elif on_failure == "return_message": # type: ignore[comparison-overlap]199 msg = ( # type: ignore[unreachable]200 "on_failure='return_message' is deprecated and will be removed "201 "in a future version. Use on_failure='continue' instead."202 )203 warnings.warn(msg, DeprecationWarning, stacklevel=2)204 on_failure = "continue"205 206 self.max_retries = max_retries207 208 # Extract tool names from BaseTool instances or strings209 self._tool_filter: list[str] | None210 if tools is not None:211 self._tool_filter = [tool.name if not isinstance(tool, str) else tool for tool in tools]212 else:213 self._tool_filter = None214 215 self.tools = [] # No additional tools registered by this middleware216 self.retry_on = retry_on217 self.on_failure = on_failure218 self.backoff_factor = backoff_factor219 self.initial_delay = initial_delay220 self.max_delay = max_delay221 self.jitter = jitter222 223 def _should_retry_tool(self, tool_name: str) -> bool:224 """Check if retry logic should apply to this tool.225 226 Args:227 tool_name: Name of the tool being called.228 229 Returns:230 `True` if retry logic should apply, `False` otherwise.231 """232 if self._tool_filter is None:233 return True234 return tool_name in self._tool_filter235 236 @staticmethod237 def _format_failure_message(tool_name: str, exc: Exception, attempts_made: int) -> str:238 """Format the failure message when retries are exhausted.239 240 Args:241 tool_name: Name of the tool that failed.242 exc: The exception that caused the failure.243 attempts_made: Number of attempts actually made.244 245 Returns:246 Formatted error message string.247 """248 exc_type = type(exc).__name__249 exc_msg = str(exc)250 attempt_word = "attempt" if attempts_made == 1 else "attempts"251 return (252 f"Tool '{tool_name}' failed after {attempts_made} {attempt_word} "253 f"with {exc_type}: {exc_msg}. Please try again."254 )255 256 def _handle_failure(257 self, tool_name: str, tool_call_id: str | None, exc: Exception, attempts_made: int258 ) -> ToolMessage:259 """Handle failure when all retries are exhausted.260 261 Args:262 tool_name: Name of the tool that failed.263 tool_call_id: ID of the tool call (may be `None`).264 exc: The exception that caused the failure.265 attempts_made: Number of attempts actually made.266 267 Returns:268 `ToolMessage` with error details.269 270 Raises:271 Exception: If `on_failure` is `'error'`, re-raises the exception.272 """273 if self.on_failure == "error":274 raise exc275 276 if callable(self.on_failure):277 content = self.on_failure(exc)278 else:279 content = self._format_failure_message(tool_name, exc, attempts_made)280 281 return ToolMessage(282 content=content,283 tool_call_id=tool_call_id,284 name=tool_name,285 status="error",286 )287 288 def wrap_tool_call(289 self,290 request: ToolCallRequest,291 handler: Callable[[ToolCallRequest], ToolMessage | Command[Any]],292 ) -> ToolMessage | Command[Any]:293 """Intercept tool execution and retry on failure.294 295 Args:296 request: Tool call request with call dict, `BaseTool`, state, and runtime.297 handler: Callable to execute the tool (can be called multiple times).298 299 Returns:300 `ToolMessage` or `Command` (the final result).301 302 Raises:303 RuntimeError: If the retry loop completes without returning. This should not happen.304 """305 tool_name = request.tool.name if request.tool else request.tool_call["name"]306 307 # Check if retry should apply to this tool308 if not self._should_retry_tool(tool_name):309 return handler(request)310 311 tool_call_id = request.tool_call["id"]312 313 # Initial attempt + retries314 for attempt in range(self.max_retries + 1):315 try:316 return handler(request)317 except Exception as exc:318 attempts_made = attempt + 1 # attempt is 0-indexed319 320 # Check if we should retry this exception321 if not should_retry_exception(exc, self.retry_on):322 # Exception is not retryable, handle failure immediately323 return self._handle_failure(tool_name, tool_call_id, exc, attempts_made)324 325 # Check if we have more retries left326 if attempt < self.max_retries:327 # Calculate and apply backoff delay328 delay = calculate_delay(329 attempt,330 backoff_factor=self.backoff_factor,331 initial_delay=self.initial_delay,332 max_delay=self.max_delay,333 jitter=self.jitter,334 )335 if delay > 0:336 time.sleep(delay)337 # Continue to next retry338 else:339 # No more retries, handle failure340 return self._handle_failure(tool_name, tool_call_id, exc, attempts_made)341 342 # Unreachable: loop always returns via handler success or _handle_failure343 msg = "Unexpected: retry loop completed without returning"344 raise RuntimeError(msg)345 346 async def awrap_tool_call(347 self,348 request: ToolCallRequest,349 handler: Callable[[ToolCallRequest], Awaitable[ToolMessage | Command[Any]]],350 ) -> ToolMessage | Command[Any]:351 """Intercept and control async tool execution with retry logic.352 353 Args:354 request: Tool call request with call `dict`, `BaseTool`, state, and runtime.355 handler: Async callable to execute the tool and returns `ToolMessage` or356 `Command`.357 358 Returns:359 `ToolMessage` or `Command` (the final result).360 361 Raises:362 RuntimeError: If the retry loop completes without returning. This should not happen.363 """364 tool_name = request.tool.name if request.tool else request.tool_call["name"]365 366 # Check if retry should apply to this tool367 if not self._should_retry_tool(tool_name):368 return await handler(request)369 370 tool_call_id = request.tool_call["id"]371 372 # Initial attempt + retries373 for attempt in range(self.max_retries + 1):374 try:375 return await handler(request)376 except Exception as exc:377 attempts_made = attempt + 1 # attempt is 0-indexed378 379 # Check if we should retry this exception380 if not should_retry_exception(exc, self.retry_on):381 # Exception is not retryable, handle failure immediately382 return self._handle_failure(tool_name, tool_call_id, exc, attempts_made)383 384 # Check if we have more retries left385 if attempt < self.max_retries:386 # Calculate and apply backoff delay387 delay = calculate_delay(388 attempt,389 backoff_factor=self.backoff_factor,390 initial_delay=self.initial_delay,391 max_delay=self.max_delay,392 jitter=self.jitter,393 )394 if delay > 0:395 await asyncio.sleep(delay)396 # Continue to next retry397 else:398 # No more retries, handle failure399 return self._handle_failure(tool_name, tool_call_id, exc, attempts_made)400 401 # Unreachable: loop always returns via handler success or _handle_failure402 msg = "Unexpected: retry loop completed without returning"403 raise RuntimeError(msg)404 