codekingpro/portable-devtools
114k
1"""Retriever tool."""2 3from __future__ import annotations4 5from typing import TYPE_CHECKING, Literal6 7from pydantic import BaseModel, Field8 9# Cannot move Callbacks and Document to TYPE_CHECKING as StructuredTool's10# func/coroutine parameter annotations are evaluated at runtime.11from langchain_core.callbacks import Callbacks # noqa: TC00112from langchain_core.documents import Document # noqa: TC00113from langchain_core.prompts import (14 BasePromptTemplate,15 PromptTemplate,16 aformat_document,17 format_document,18)19from langchain_core.tools.structured import StructuredTool20 21if TYPE_CHECKING:22 from langchain_core.retrievers import BaseRetriever23 24 25class RetrieverInput(BaseModel):26 """Input to the retriever."""27 28 query: str = Field(description="query to look up in retriever")29 30 31def create_retriever_tool(32 retriever: BaseRetriever,33 name: str,34 description: str,35 *,36 document_prompt: BasePromptTemplate | None = None,37 document_separator: str = "\n\n",38 response_format: Literal["content", "content_and_artifact"] = "content",39) -> StructuredTool:40 r"""Create a tool to do retrieval of documents.41 42 Args:43 retriever: The retriever to use for the retrieval44 name: The name for the tool.45 46 This will be passed to the language model, so should be unique and somewhat47 descriptive.48 description: The description for the tool.49 50 This will be passed to the language model, so should be descriptive.51 document_prompt: The prompt to use for the document.52 document_separator: The separator to use between documents.53 response_format: The tool response format.54 55 If `'content'` then the output of the tool is interpreted as the contents of56 a `ToolMessage`. If `'content_and_artifact'` then the output is expected to57 be a two-tuple corresponding to the `(content, artifact)` of a `ToolMessage`58 (artifact being a list of documents in this case).59 60 Returns:61 Tool class to pass to an agent.62 """63 document_prompt_ = document_prompt or PromptTemplate.from_template("{page_content}")64 65 def func(66 query: str, callbacks: Callbacks = None67 ) -> str | tuple[str, list[Document]]:68 docs = retriever.invoke(query, config={"callbacks": callbacks})69 content = document_separator.join(70 format_document(doc, document_prompt_) for doc in docs71 )72 if response_format == "content_and_artifact":73 return (content, docs)74 return content75 76 async def afunc(77 query: str, callbacks: Callbacks = None78 ) -> str | tuple[str, list[Document]]:79 docs = await retriever.ainvoke(query, config={"callbacks": callbacks})80 content = document_separator.join(81 [await aformat_document(doc, document_prompt_) for doc in docs]82 )83 if response_format == "content_and_artifact":84 return (content, docs)85 return content86 87 return StructuredTool(88 name=name,89 description=description,90 func=func,91 coroutine=afunc,92 args_schema=RetrieverInput,93 response_format=response_format,94 )95 