codekingpro/portable-devtools
114k
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 