Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
tool_retry.py404 linesDownload Raw Back to middleware
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 
codekingpro/portable-devtools · Team Ai