Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
tool.py140 linesDownload Raw Back to vectorstore
1"""Tools for interacting with vectorstores."""2 3import json4from typing import Any, Dict, Optional5 6from langchain_core.callbacks import (7    AsyncCallbackManagerForToolRun,8    CallbackManagerForToolRun,9)10from langchain_core.language_models import BaseLanguageModel11from langchain_core.tools import BaseTool12from langchain_core.vectorstores import VectorStore13from pydantic import BaseModel, ConfigDict, Field14 15from langchain_community.llms.openai import OpenAI16 17 18class BaseVectorStoreTool(BaseModel):19    """Base class for tools that use a VectorStore."""20 21    vectorstore: VectorStore = Field(exclude=True)22    llm: BaseLanguageModel = Field(default_factory=lambda: OpenAI(temperature=0))23 24    model_config = ConfigDict(25        arbitrary_types_allowed=True,26    )27 28 29def _create_description_from_template(values: Dict[str, Any]) -> Dict[str, Any]:30    values["description"] = values["template"].format(name=values["name"])31    return values32 33 34class VectorStoreQATool(BaseVectorStoreTool, BaseTool):35    """Tool for the VectorDBQA chain. To be initialized with name and chain."""36 37    @staticmethod38    def get_description(name: str, description: str) -> str:39        template: str = (40            "Useful for when you need to answer questions about {name}. "41            "Whenever you need information about {description} "42            "you should ALWAYS use this. "43            "Input should be a fully formed question."44        )45        return template.format(name=name, description=description)46 47    def _run(48        self,49        query: str,50        run_manager: Optional[CallbackManagerForToolRun] = None,51    ) -> str:52        """Use the tool."""53        from langchain_classic.chains.retrieval_qa.base import RetrievalQA54 55        chain = RetrievalQA.from_chain_type(56            self.llm, retriever=self.vectorstore.as_retriever()57        )58        return chain.invoke(59            {chain.input_key: query},60            config={"callbacks": run_manager.get_child() if run_manager else None},61        )[chain.output_key]62 63    async def _arun(64        self,65        query: str,66        run_manager: Optional[AsyncCallbackManagerForToolRun] = None,67    ) -> str:68        """Use the tool asynchronously."""69        from langchain_classic.chains.retrieval_qa.base import RetrievalQA70 71        chain = RetrievalQA.from_chain_type(72            self.llm, retriever=self.vectorstore.as_retriever()73        )74        return (75            await chain.ainvoke(76                {chain.input_key: query},77                config={"callbacks": run_manager.get_child() if run_manager else None},78            )79        )[chain.output_key]80 81 82class VectorStoreQAWithSourcesTool(BaseVectorStoreTool, BaseTool):83    """Tool for the VectorDBQAWithSources chain."""84 85    @staticmethod86    def get_description(name: str, description: str) -> str:87        template: str = (88            "Useful for when you need to answer questions about {name} and the sources "89            "used to construct the answer. "90            "Whenever you need information about {description} "91            "you should ALWAYS use this. "92            " Input should be a fully formed question. "93            "Output is a json serialized dictionary with keys `answer` and `sources`. "94            "Only use this tool if the user explicitly asks for sources."95        )96        return template.format(name=name, description=description)97 98    def _run(99        self,100        query: str,101        run_manager: Optional[CallbackManagerForToolRun] = None,102    ) -> str:103        """Use the tool."""104 105        from langchain_classic.chains.qa_with_sources.retrieval import (106            RetrievalQAWithSourcesChain,107        )108 109        chain = RetrievalQAWithSourcesChain.from_chain_type(110            self.llm, retriever=self.vectorstore.as_retriever()111        )112        return json.dumps(113            chain.invoke(114                {chain.question_key: query},115                return_only_outputs=True,116                config={"callbacks": run_manager.get_child() if run_manager else None},117            )118        )119 120    async def _arun(121        self,122        query: str,123        run_manager: Optional[AsyncCallbackManagerForToolRun] = None,124    ) -> str:125        """Use the tool asynchronously."""126        from langchain_classic.chains.qa_with_sources.retrieval import (127            RetrievalQAWithSourcesChain,128        )129 130        chain = RetrievalQAWithSourcesChain.from_chain_type(131            self.llm, retriever=self.vectorstore.as_retriever()132        )133        return json.dumps(134            await chain.ainvoke(135                {chain.question_key: query},136                return_only_outputs=True,137                config={"callbacks": run_manager.get_child() if run_manager else None},138            )139        )140 
codekingpro/portable-devtools · Team Ai