codekingpro/portable-devtools
114k
1"""LLM-based tool selector middleware."""2 3from __future__ import annotations4 5import logging6from dataclasses import dataclass7from typing import TYPE_CHECKING, Annotated, Any, Literal, Union8 9from langchain_core.language_models.chat_models import BaseChatModel10from langchain_core.messages import AIMessage, HumanMessage11from pydantic import Field, TypeAdapter12from typing_extensions import TypedDict13 14from langchain.agents.middleware.types import (15 AgentMiddleware,16 AgentState,17 ContextT,18 ModelRequest,19 ModelResponse,20 ResponseT,21)22from langchain.chat_models.base import init_chat_model23 24if TYPE_CHECKING:25 from collections.abc import Awaitable, Callable26 27 from langchain.tools import BaseTool28 29logger = logging.getLogger(__name__)30 31DEFAULT_SYSTEM_PROMPT = (32 "Your goal is to select the most relevant tools for answering the user's query."33)34 35 36@dataclass37class _SelectionRequest:38 """Prepared inputs for tool selection."""39 40 available_tools: list[BaseTool]41 system_message: str42 last_user_message: HumanMessage43 model: BaseChatModel44 valid_tool_names: list[str]45 46 47def _create_tool_selection_response(tools: list[BaseTool]) -> TypeAdapter[Any]:48 """Create a structured output schema for tool selection.49 50 Args:51 tools: Available tools to include in the schema.52 53 Returns:54 `TypeAdapter` for a schema where each tool name is a `Literal` with its55 description.56 57 Raises:58 AssertionError: If `tools` is empty.59 """60 if not tools:61 msg = "Invalid usage: tools must be non-empty"62 raise AssertionError(msg)63 64 # Create a Union of Annotated Literal types for each tool name with description65 # For instance: Union[Annotated[Literal["tool1"], Field(description="...")], ...]66 literals = [67 Annotated[Literal[tool.name], Field(description=tool.description)] for tool in tools68 ]69 selected_tool_type = Union[tuple(literals)] # type: ignore[valid-type] # noqa: UP00770 71 description = "Tools to use. Place the most relevant tools first."72 73 class ToolSelectionResponse(TypedDict):74 """Use to select relevant tools."""75 76 tools: Annotated[list[selected_tool_type], Field(description=description)] # type: ignore[valid-type]77 78 return TypeAdapter(ToolSelectionResponse)79 80 81def _render_tool_list(tools: list[BaseTool]) -> str:82 """Format tools as markdown list.83 84 Args:85 tools: Tools to format.86 87 Returns:88 Markdown string with each tool on a new line.89 """90 return "\n".join(f"- {tool.name}: {tool.description}" for tool in tools)91 92 93class LLMToolSelectorMiddleware(AgentMiddleware[AgentState[ResponseT], ContextT, ResponseT]):94 """Uses an LLM to select relevant tools before calling the main model.95 96 When an agent has many tools available, this middleware filters them down97 to only the most relevant ones for the user's query. This reduces token usage98 and helps the main model focus on the right tools.99 100 Examples:101 !!! example "Limit to 3 tools"102 103 ```python104 from langchain.agents.middleware import LLMToolSelectorMiddleware105 106 middleware = LLMToolSelectorMiddleware(max_tools=3)107 108 agent = create_agent(109 model="openai:gpt-4o",110 tools=[tool1, tool2, tool3, tool4, tool5],111 middleware=[middleware],112 )113 ```114 115 !!! example "Use a smaller model for selection"116 117 ```python118 middleware = LLMToolSelectorMiddleware(model="openai:gpt-4o-mini", max_tools=2)119 ```120 """121 122 def __init__(123 self,124 *,125 model: str | BaseChatModel | None = None,126 system_prompt: str = DEFAULT_SYSTEM_PROMPT,127 max_tools: int | None = None,128 always_include: list[str] | None = None,129 ) -> None:130 """Initialize the tool selector.131 132 Args:133 model: Model to use for selection.134 135 If not provided, uses the agent's main model.136 137 Can be a model identifier string or `BaseChatModel` instance.138 system_prompt: Instructions for the selection model.139 max_tools: Maximum number of tools to select.140 141 If the model selects more, only the first `max_tools` will be used.142 143 If not specified, there is no limit.144 always_include: Tool names to always include regardless of selection.145 146 These do not count against the `max_tools` limit.147 """148 super().__init__()149 self.system_prompt = system_prompt150 self.max_tools = max_tools151 self.always_include = always_include or []152 153 if isinstance(model, (BaseChatModel, type(None))):154 self.model: BaseChatModel | None = model155 else:156 self.model = init_chat_model(model)157 158 def _prepare_selection_request(159 self, request: ModelRequest[ContextT]160 ) -> _SelectionRequest | None:161 """Prepare inputs for tool selection.162 163 Args:164 request: the model request.165 166 Returns:167 `SelectionRequest` with prepared inputs, or `None` if no selection is168 needed.169 170 Raises:171 ValueError: If tools in `always_include` are not found in the request.172 AssertionError: If no user message is found in the request messages.173 """174 # If no tools available, return None175 if not request.tools or len(request.tools) == 0:176 return None177 178 # Filter to only BaseTool instances (exclude provider-specific tool dicts)179 base_tools = [tool for tool in request.tools if not isinstance(tool, dict)]180 181 # Validate that always_include tools exist182 if self.always_include:183 available_tool_names = {tool.name for tool in base_tools}184 missing_tools = [185 name for name in self.always_include if name not in available_tool_names186 ]187 if missing_tools:188 msg = (189 f"Tools in always_include not found in request: {missing_tools}. "190 f"Available tools: {sorted(available_tool_names)}"191 )192 raise ValueError(msg)193 194 # Separate tools that are always included from those available for selection195 available_tools = [tool for tool in base_tools if tool.name not in self.always_include]196 197 # If no tools available for selection, return None198 if not available_tools:199 return None200 201 system_message = self.system_prompt202 # If there's a max_tools limit, append instructions to the system prompt203 if self.max_tools is not None:204 system_message += (205 f"\nIMPORTANT: List the tool names in order of relevance, "206 f"with the most relevant first. "207 f"If you exceed the maximum number of tools, "208 f"only the first {self.max_tools} will be used."209 )210 211 # Get the last user message from the conversation history212 last_user_message: HumanMessage213 for message in reversed(request.messages):214 if isinstance(message, HumanMessage):215 last_user_message = message216 break217 else:218 msg = "No user message found in request messages"219 raise AssertionError(msg)220 221 model = self.model or request.model222 valid_tool_names = [tool.name for tool in available_tools]223 224 return _SelectionRequest(225 available_tools=available_tools,226 system_message=system_message,227 last_user_message=last_user_message,228 model=model,229 valid_tool_names=valid_tool_names,230 )231 232 def _process_selection_response(233 self,234 response: dict[str, Any],235 available_tools: list[BaseTool],236 valid_tool_names: list[str],237 request: ModelRequest[ContextT],238 ) -> ModelRequest[ContextT]:239 """Process the selection response and return filtered `ModelRequest`."""240 selected_tool_names: list[str] = []241 invalid_tool_selections = []242 243 for tool_name in response["tools"]:244 if tool_name not in valid_tool_names:245 invalid_tool_selections.append(tool_name)246 continue247 248 # Only add if not already selected and within max_tools limit249 if tool_name not in selected_tool_names and (250 self.max_tools is None or len(selected_tool_names) < self.max_tools251 ):252 selected_tool_names.append(tool_name)253 254 if invalid_tool_selections:255 msg = f"Model selected invalid tools: {invalid_tool_selections}"256 raise ValueError(msg)257 258 # Filter tools based on selection and append always-included tools259 selected_tools: list[BaseTool] = [260 tool for tool in available_tools if tool.name in selected_tool_names261 ]262 always_included_tools: list[BaseTool] = [263 tool264 for tool in request.tools265 if not isinstance(tool, dict) and tool.name in self.always_include266 ]267 selected_tools.extend(always_included_tools)268 269 # Also preserve any provider-specific tool dicts from the original request270 provider_tools = [tool for tool in request.tools if isinstance(tool, dict)]271 272 return request.override(tools=[*selected_tools, *provider_tools])273 274 def wrap_model_call(275 self,276 request: ModelRequest[ContextT],277 handler: Callable[[ModelRequest[ContextT]], ModelResponse[ResponseT]],278 ) -> ModelResponse[ResponseT] | AIMessage:279 """Filter tools based on LLM selection before invoking the model via handler.280 281 Args:282 request: Model request to execute (includes state and runtime).283 handler: Async callback that executes the model request and returns284 `ModelResponse`.285 286 Returns:287 The model call result.288 289 Raises:290 AssertionError: If the selection model response is not a dict.291 """292 selection_request = self._prepare_selection_request(request)293 if selection_request is None:294 return handler(request)295 296 # Create dynamic response model with Literal enum of available tool names297 type_adapter = _create_tool_selection_response(selection_request.available_tools)298 schema = type_adapter.json_schema()299 structured_model = selection_request.model.with_structured_output(schema)300 301 response = structured_model.invoke(302 [303 {"role": "system", "content": selection_request.system_message},304 selection_request.last_user_message,305 ]306 )307 308 # Response should be a dict since we're passing a schema (not a Pydantic model class)309 if not isinstance(response, dict):310 msg = f"Expected dict response, got {type(response)}"311 raise AssertionError(msg) # noqa: TRY004312 modified_request = self._process_selection_response(313 response, selection_request.available_tools, selection_request.valid_tool_names, request314 )315 return handler(modified_request)316 317 async def awrap_model_call(318 self,319 request: ModelRequest[ContextT],320 handler: Callable[[ModelRequest[ContextT]], Awaitable[ModelResponse[ResponseT]]],321 ) -> ModelResponse[ResponseT] | AIMessage:322 """Filter tools based on LLM selection before invoking the model via handler.323 324 Args:325 request: Model request to execute (includes state and runtime).326 handler: Async callback that executes the model request and returns327 `ModelResponse`.328 329 Returns:330 The model call result.331 332 Raises:333 AssertionError: If the selection model response is not a dict.334 """335 selection_request = self._prepare_selection_request(request)336 if selection_request is None:337 return await handler(request)338 339 # Create dynamic response model with Literal enum of available tool names340 type_adapter = _create_tool_selection_response(selection_request.available_tools)341 schema = type_adapter.json_schema()342 structured_model = selection_request.model.with_structured_output(schema)343 344 response = await structured_model.ainvoke(345 [346 {"role": "system", "content": selection_request.system_message},347 selection_request.last_user_message,348 ]349 )350 351 # Response should be a dict since we're passing a schema (not a Pydantic model class)352 if not isinstance(response, dict):353 msg = f"Expected dict response, got {type(response)}"354 raise AssertionError(msg) # noqa: TRY004355 modified_request = self._process_selection_response(356 response, selection_request.available_tools, selection_request.valid_tool_names, request357 )358 return await handler(modified_request)359 