Team Ai
Datasetpublic

codekingpro/portable-devtools

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