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