codekingpro/portable-devtools
114k
1"""Chain for question-answering against a vector database."""2 3from __future__ import annotations4 5import inspect6from abc import abstractmethod7from typing import Any8 9from langchain_core._api import deprecated10from langchain_core.callbacks import (11 AsyncCallbackManagerForChainRun,12 CallbackManagerForChainRun,13 Callbacks,14)15from langchain_core.documents import Document16from langchain_core.language_models import BaseLanguageModel17from langchain_core.prompts import PromptTemplate18from langchain_core.retrievers import BaseRetriever19from langchain_core.vectorstores import VectorStore20from pydantic import ConfigDict, Field, model_validator21from typing_extensions import override22 23from langchain_classic.chains.base import Chain24from langchain_classic.chains.combine_documents.base import BaseCombineDocumentsChain25from langchain_classic.chains.combine_documents.stuff import StuffDocumentsChain26from langchain_classic.chains.llm import LLMChain27from langchain_classic.chains.question_answering import load_qa_chain28from langchain_classic.chains.question_answering.stuff_prompt import PROMPT_SELECTOR29 30 31@deprecated(32 since="0.2.13",33 removal="2.0.0",34 alternative="langchain.agents.create_agent",35 addendum=(36 "Build new RAG flows with `create_agent` and a retrieval tool. See "37 "https://docs.langchain.com/oss/python/langchain/rag"38 ),39)40class BaseRetrievalQA(Chain):41 """Base class for question-answering chains."""42 43 combine_documents_chain: BaseCombineDocumentsChain44 """Chain to use to combine the documents."""45 input_key: str = "query"46 output_key: str = "result"47 return_source_documents: bool = False48 """Return the source documents or not."""49 50 model_config = ConfigDict(51 populate_by_name=True,52 arbitrary_types_allowed=True,53 extra="forbid",54 )55 56 @property57 def input_keys(self) -> list[str]:58 """Input keys."""59 return [self.input_key]60 61 @property62 def output_keys(self) -> list[str]:63 """Output keys."""64 _output_keys = [self.output_key]65 if self.return_source_documents:66 _output_keys = [*_output_keys, "source_documents"]67 return _output_keys68 69 @classmethod70 def from_llm(71 cls,72 llm: BaseLanguageModel,73 prompt: PromptTemplate | None = None,74 callbacks: Callbacks = None,75 llm_chain_kwargs: dict | None = None,76 **kwargs: Any,77 ) -> BaseRetrievalQA:78 """Initialize from LLM."""79 _prompt = prompt or PROMPT_SELECTOR.get_prompt(llm)80 llm_chain = LLMChain(81 llm=llm,82 prompt=_prompt,83 callbacks=callbacks,84 **(llm_chain_kwargs or {}),85 )86 document_prompt = PromptTemplate(87 input_variables=["page_content"],88 template="Context:\n{page_content}",89 )90 combine_documents_chain = StuffDocumentsChain(91 llm_chain=llm_chain,92 document_variable_name="context",93 document_prompt=document_prompt,94 callbacks=callbacks,95 )96 97 return cls(98 combine_documents_chain=combine_documents_chain,99 callbacks=callbacks,100 **kwargs,101 )102 103 @classmethod104 def from_chain_type(105 cls,106 llm: BaseLanguageModel,107 chain_type: str = "stuff",108 chain_type_kwargs: dict | None = None,109 **kwargs: Any,110 ) -> BaseRetrievalQA:111 """Load chain from chain type."""112 _chain_type_kwargs = chain_type_kwargs or {}113 combine_documents_chain = load_qa_chain(114 llm,115 chain_type=chain_type,116 **_chain_type_kwargs,117 )118 return cls(combine_documents_chain=combine_documents_chain, **kwargs)119 120 @abstractmethod121 def _get_docs(122 self,123 question: str,124 *,125 run_manager: CallbackManagerForChainRun,126 ) -> list[Document]:127 """Get documents to do question answering over."""128 129 def _call(130 self,131 inputs: dict[str, Any],132 run_manager: CallbackManagerForChainRun | None = None,133 ) -> dict[str, Any]:134 """Run get_relevant_text and llm on input query.135 136 If chain has 'return_source_documents' as 'True', returns137 the retrieved documents as well under the key 'source_documents'.138 139 Example:140 ```python141 res = indexqa({"query": "This is my query"})142 answer, docs = res["result"], res["source_documents"]143 ```144 """145 _run_manager = run_manager or CallbackManagerForChainRun.get_noop_manager()146 question = inputs[self.input_key]147 accepts_run_manager = (148 "run_manager" in inspect.signature(self._get_docs).parameters149 )150 if accepts_run_manager:151 docs = self._get_docs(question, run_manager=_run_manager)152 else:153 docs = self._get_docs(question) # type: ignore[call-arg]154 answer = self.combine_documents_chain.run(155 input_documents=docs,156 question=question,157 callbacks=_run_manager.get_child(),158 )159 160 if self.return_source_documents:161 return {self.output_key: answer, "source_documents": docs}162 return {self.output_key: answer}163 164 @abstractmethod165 async def _aget_docs(166 self,167 question: str,168 *,169 run_manager: AsyncCallbackManagerForChainRun,170 ) -> list[Document]:171 """Get documents to do question answering over."""172 173 async def _acall(174 self,175 inputs: dict[str, Any],176 run_manager: AsyncCallbackManagerForChainRun | None = None,177 ) -> dict[str, Any]:178 """Run get_relevant_text and llm on input query.179 180 If chain has 'return_source_documents' as 'True', returns181 the retrieved documents as well under the key 'source_documents'.182 183 Example:184 ```python185 res = indexqa({"query": "This is my query"})186 answer, docs = res["result"], res["source_documents"]187 ```188 """189 _run_manager = run_manager or AsyncCallbackManagerForChainRun.get_noop_manager()190 question = inputs[self.input_key]191 accepts_run_manager = (192 "run_manager" in inspect.signature(self._aget_docs).parameters193 )194 if accepts_run_manager:195 docs = await self._aget_docs(question, run_manager=_run_manager)196 else:197 docs = await self._aget_docs(question) # type: ignore[call-arg]198 answer = await self.combine_documents_chain.arun(199 input_documents=docs,200 question=question,201 callbacks=_run_manager.get_child(),202 )203 204 if self.return_source_documents:205 return {self.output_key: answer, "source_documents": docs}206 return {self.output_key: answer}207 208 209@deprecated(210 since="0.1.17",211 removal="2.0.0",212 alternative="langchain.agents.create_agent",213 addendum=(214 "Build new RAG flows with `create_agent` and a retrieval tool. See "215 "https://docs.langchain.com/oss/python/langchain/rag"216 ),217)218class RetrievalQA(BaseRetrievalQA):219 """Chain for question-answering against an index.220 221 This class is deprecated. See below for an example implementation using222 `create_retrieval_chain`:223 224 ```python225 from langchain_classic.chains import create_retrieval_chain226 from langchain_classic.chains.combine_documents import (227 create_stuff_documents_chain,228 )229 from langchain_core.prompts import ChatPromptTemplate230 from langchain_openai import ChatOpenAI231 232 233 retriever = ... # Your retriever234 model = ChatOpenAI()235 236 system_prompt = (237 "Use the given context to answer the question. "238 "If you don't know the answer, say you don't know. "239 "Use three sentence maximum and keep the answer concise. "240 "Context: {context}"241 )242 prompt = ChatPromptTemplate.from_messages(243 [244 ("system", system_prompt),245 ("human", "{input}"),246 ]247 )248 question_answer_chain = create_stuff_documents_chain(model, prompt)249 chain = create_retrieval_chain(retriever, question_answer_chain)250 251 chain.invoke({"input": query})252 ```253 254 Example:255 ```python256 from langchain_openai import OpenAI257 from langchain_classic.chains import RetrievalQA258 from langchain_community.vectorstores import FAISS259 from langchain_core.vectorstores import VectorStoreRetriever260 261 retriever = VectorStoreRetriever(vectorstore=FAISS(...))262 retrievalQA = RetrievalQA.from_llm(llm=OpenAI(), retriever=retriever)263 ```264 """265 266 retriever: BaseRetriever = Field(exclude=True)267 268 def _get_docs(269 self,270 question: str,271 *,272 run_manager: CallbackManagerForChainRun,273 ) -> list[Document]:274 """Get docs."""275 return self.retriever.invoke(276 question,277 config={"callbacks": run_manager.get_child()},278 )279 280 async def _aget_docs(281 self,282 question: str,283 *,284 run_manager: AsyncCallbackManagerForChainRun,285 ) -> list[Document]:286 """Get docs."""287 return await self.retriever.ainvoke(288 question,289 config={"callbacks": run_manager.get_child()},290 )291 292 @property293 def _chain_type(self) -> str:294 """Return the chain type."""295 return "retrieval_qa"296 297 298@deprecated(299 since="0.2.13",300 removal="2.0.0",301 alternative="langchain.agents.create_agent",302 addendum=(303 "Build new RAG flows with `create_agent` and a retrieval tool. See "304 "https://docs.langchain.com/oss/python/langchain/rag"305 ),306)307class VectorDBQA(BaseRetrievalQA):308 """Chain for question-answering against a vector database."""309 310 vectorstore: VectorStore = Field(exclude=True, alias="vectorstore")311 """Vector Database to connect to."""312 k: int = 4313 """Number of documents to query for."""314 search_type: str = "similarity"315 """Search type to use over vectorstore. `similarity` or `mmr`."""316 search_kwargs: dict[str, Any] = Field(default_factory=dict)317 """Extra search args."""318 319 @model_validator(mode="before")320 @classmethod321 def validate_search_type(cls, values: dict) -> Any:322 """Validate search type."""323 if "search_type" in values:324 search_type = values["search_type"]325 if search_type not in ("similarity", "mmr"):326 msg = f"search_type of {search_type} not allowed."327 raise ValueError(msg)328 return values329 330 @override331 def _get_docs(332 self,333 question: str,334 *,335 run_manager: CallbackManagerForChainRun,336 ) -> list[Document]:337 """Get docs."""338 if self.search_type == "similarity":339 docs = self.vectorstore.similarity_search(340 question,341 k=self.k,342 **self.search_kwargs,343 )344 elif self.search_type == "mmr":345 docs = self.vectorstore.max_marginal_relevance_search(346 question,347 k=self.k,348 **self.search_kwargs,349 )350 else:351 msg = f"search_type of {self.search_type} not allowed."352 raise ValueError(msg)353 return docs354 355 async def _aget_docs(356 self,357 question: str,358 *,359 run_manager: AsyncCallbackManagerForChainRun,360 ) -> list[Document]:361 """Get docs."""362 msg = "VectorDBQA does not support async"363 raise NotImplementedError(msg)364 365 @property366 def _chain_type(self) -> str:367 """Return the chain type."""368 return "vector_db_qa"369 