Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
model_call_limit.py268 linesDownload Raw Back to middleware
1"""Call tracking middleware for agents."""2 3from __future__ import annotations4 5from typing import TYPE_CHECKING, Annotated, Any, Literal6 7from langchain_core.messages import AIMessage8from langgraph.channels.untracked_value import UntrackedValue9from typing_extensions import NotRequired, override10 11from langchain.agents.middleware.types import (12    AgentMiddleware,13    AgentState,14    ContextT,15    PrivateStateAttr,16    ResponseT,17    hook_config,18)19 20if TYPE_CHECKING:21    from langgraph.runtime import Runtime22 23 24class ModelCallLimitState(AgentState[ResponseT]):25    """State schema for `ModelCallLimitMiddleware`.26 27    Extends `AgentState` with model call tracking fields.28 29    Type Parameters:30        ResponseT: The type of the structured response. Defaults to `Any`.31    """32 33    thread_model_call_count: NotRequired[Annotated[int, PrivateStateAttr]]34    run_model_call_count: NotRequired[Annotated[int, UntrackedValue, PrivateStateAttr]]35 36 37def _build_limit_exceeded_message(38    thread_count: int,39    run_count: int,40    thread_limit: int | None,41    run_limit: int | None,42) -> str:43    """Build a message indicating which limits were exceeded.44 45    Args:46        thread_count: Current thread model call count.47        run_count: Current run model call count.48        thread_limit: Thread model call limit (if set).49        run_limit: Run model call limit (if set).50 51    Returns:52        A formatted message describing which limits were exceeded.53    """54    exceeded_limits = []55    if thread_limit is not None and thread_count >= thread_limit:56        exceeded_limits.append(f"thread limit ({thread_count}/{thread_limit})")57    if run_limit is not None and run_count >= run_limit:58        exceeded_limits.append(f"run limit ({run_count}/{run_limit})")59 60    return f"Model call limits exceeded: {', '.join(exceeded_limits)}"61 62 63class ModelCallLimitExceededError(Exception):64    """Exception raised when model call limits are exceeded.65 66    This exception is raised when the configured exit behavior is `'error'` and either67    the thread or run model call limit has been exceeded.68    """69 70    def __init__(71        self,72        thread_count: int,73        run_count: int,74        thread_limit: int | None,75        run_limit: int | None,76    ) -> None:77        """Initialize the exception with call count information.78 79        Args:80            thread_count: Current thread model call count.81            run_count: Current run model call count.82            thread_limit: Thread model call limit (if set).83            run_limit: Run model call limit (if set).84        """85        self.thread_count = thread_count86        self.run_count = run_count87        self.thread_limit = thread_limit88        self.run_limit = run_limit89 90        msg = _build_limit_exceeded_message(thread_count, run_count, thread_limit, run_limit)91        super().__init__(msg)92 93 94class ModelCallLimitMiddleware(95    AgentMiddleware[ModelCallLimitState[ResponseT], ContextT, ResponseT]96):97    """Tracks model call counts and enforces limits.98 99    This middleware monitors the number of model calls made during agent execution100    and can terminate the agent when specified limits are reached. It supports101    both thread-level and run-level call counting with configurable exit behaviors.102 103    Thread-level: The middleware tracks the number of model calls and persists104    call count across multiple runs (invocations) of the agent.105 106    Run-level: The middleware tracks the number of model calls made during a single107    run (invocation) of the agent.108 109    Example:110        ```python111        from langchain.agents.middleware import ModelCallLimitMiddleware112        from langchain.agents import create_agent113 114        # Create middleware with limits115        call_tracker = ModelCallLimitMiddleware(thread_limit=10, run_limit=5, exit_behavior="end")116 117        agent = create_agent("openai:gpt-4o", middleware=[call_tracker])118 119        # Agent will automatically jump to end when limits are exceeded120        result = await agent.invoke({"messages": [HumanMessage("Help me with a task")]})121        ```122    """123 124    state_schema = ModelCallLimitState  # type: ignore[assignment]125 126    def __init__(127        self,128        *,129        thread_limit: int | None = None,130        run_limit: int | None = None,131        exit_behavior: Literal["end", "error"] = "end",132    ) -> None:133        """Initialize the call tracking middleware.134 135        Args:136            thread_limit: Maximum number of model calls allowed per thread.137 138                `None` means no limit.139            run_limit: Maximum number of model calls allowed per run.140 141                `None` means no limit.142            exit_behavior: What to do when limits are exceeded.143 144                - `'end'`: Jump to the end of the agent execution and145                    inject an artificial AI message indicating that the limit was146                    exceeded.147                - `'error'`: Raise a `ModelCallLimitExceededError`148 149        Raises:150            ValueError: If both limits are `None` or if `exit_behavior` is invalid.151        """152        super().__init__()153 154        if thread_limit is None and run_limit is None:155            msg = "At least one limit must be specified (thread_limit or run_limit)"156            raise ValueError(msg)157 158        if exit_behavior not in {"end", "error"}:159            msg = f"Invalid exit_behavior: {exit_behavior}. Must be 'end' or 'error'"160            raise ValueError(msg)161 162        self.thread_limit = thread_limit163        self.run_limit = run_limit164        self.exit_behavior = exit_behavior165 166    @hook_config(can_jump_to=["end"])167    @override168    def before_model(169        self, state: ModelCallLimitState[ResponseT], runtime: Runtime[ContextT]170    ) -> dict[str, Any] | None:171        """Check model call limits before making a model call.172 173        Args:174            state: The current agent state containing call counts.175            runtime: The langgraph runtime.176 177        Returns:178            If limits are exceeded and exit_behavior is `'end'`, returns179                a `Command` to jump to the end with a limit exceeded message. Otherwise180                returns `None`.181 182        Raises:183            ModelCallLimitExceededError: If limits are exceeded and `exit_behavior`184                is `'error'`.185        """186        thread_count = state.get("thread_model_call_count", 0)187        run_count = state.get("run_model_call_count", 0)188 189        # Check if any limits will be exceeded after the next call190        thread_limit_exceeded = self.thread_limit is not None and thread_count >= self.thread_limit191        run_limit_exceeded = self.run_limit is not None and run_count >= self.run_limit192 193        if thread_limit_exceeded or run_limit_exceeded:194            if self.exit_behavior == "error":195                raise ModelCallLimitExceededError(196                    thread_count=thread_count,197                    run_count=run_count,198                    thread_limit=self.thread_limit,199                    run_limit=self.run_limit,200                )201            if self.exit_behavior == "end":202                # Create a message indicating the limit was exceeded203                limit_message = _build_limit_exceeded_message(204                    thread_count, run_count, self.thread_limit, self.run_limit205                )206                limit_ai_message = AIMessage(content=limit_message)207 208                return {"jump_to": "end", "messages": [limit_ai_message]}209 210        return None211 212    @hook_config(can_jump_to=["end"])213    async def abefore_model(214        self,215        state: ModelCallLimitState[ResponseT],216        runtime: Runtime[ContextT],217    ) -> dict[str, Any] | None:218        """Async check model call limits before making a model call.219 220        Args:221            state: The current agent state containing call counts.222            runtime: The langgraph runtime.223 224        Returns:225            If limits are exceeded and exit_behavior is `'end'`, returns226                a `Command` to jump to the end with a limit exceeded message. Otherwise227                returns `None`.228 229        Raises:230            ModelCallLimitExceededError: If limits are exceeded and `exit_behavior`231                is `'error'`.232        """233        return self.before_model(state, runtime)234 235    @override236    def after_model(237        self, state: ModelCallLimitState[ResponseT], runtime: Runtime[ContextT]238    ) -> dict[str, Any] | None:239        """Increment model call counts after a model call.240 241        Args:242            state: The current agent state.243            runtime: The langgraph runtime.244 245        Returns:246            State updates with incremented call counts.247        """248        return {249            "thread_model_call_count": state.get("thread_model_call_count", 0) + 1,250            "run_model_call_count": state.get("run_model_call_count", 0) + 1,251        }252 253    async def aafter_model(254        self,255        state: ModelCallLimitState[ResponseT],256        runtime: Runtime[ContextT],257    ) -> dict[str, Any] | None:258        """Async increment model call counts after a model call.259 260        Args:261            state: The current agent state.262            runtime: The langgraph runtime.263 264        Returns:265            State updates with incremented call counts.266        """267        return self.after_model(state, runtime)268 
codekingpro/portable-devtools · Team Ai