codekingpro/portable-devtools
114k
1"""2Pebblo Retrieval Chain with Identity & Semantic Enforcement for question-answering3against a vector database.4"""5 6import datetime7import inspect8import logging9from importlib.metadata import version10from typing import Any, Dict, List, Optional11 12from langchain_classic.chains.base import Chain13from langchain_classic.chains.combine_documents.base import BaseCombineDocumentsChain14from langchain_core.callbacks import (15 AsyncCallbackManagerForChainRun,16 CallbackManagerForChainRun,17)18from langchain_core.documents import Document19from langchain_core.language_models import BaseLanguageModel20from langchain_core.vectorstores import VectorStoreRetriever21from pydantic import ConfigDict, Field, validator22 23from langchain_community.chains.pebblo_retrieval.enforcement_filters import (24 SUPPORTED_VECTORSTORES,25 set_enforcement_filters,26)27from langchain_community.chains.pebblo_retrieval.models import (28 App,29 AuthContext,30 ChainInfo,31 Framework,32 Model,33 SemanticContext,34 VectorDB,35)36from langchain_community.chains.pebblo_retrieval.utilities import (37 PLUGIN_VERSION,38 PebbloRetrievalAPIWrapper,39 get_runtime,40)41 42logger = logging.getLogger(__name__)43 44 45class PebbloRetrievalQA(Chain):46 """47 Retrieval Chain with Identity & Semantic Enforcement for question-answering48 against a vector database.49 """50 51 combine_documents_chain: BaseCombineDocumentsChain52 """Chain to use to combine the documents."""53 input_key: str = "query" #: :meta private:54 output_key: str = "result" #: :meta private:55 return_source_documents: bool = False56 """Return the source documents or not."""57 58 retriever: VectorStoreRetriever = Field(exclude=True)59 """VectorStore to use for retrieval."""60 auth_context_key: str = "auth_context" #: :meta private:61 """Authentication context for identity enforcement."""62 semantic_context_key: str = "semantic_context" #: :meta private:63 """Semantic context for semantic enforcement."""64 app_name: str #: :meta private:65 """App name."""66 owner: str #: :meta private:67 """Owner of app."""68 description: str #: :meta private:69 """Description of app."""70 api_key: Optional[str] = None #: :meta private:71 """Pebblo cloud API key for app."""72 classifier_url: Optional[str] = None #: :meta private:73 """Classifier endpoint."""74 classifier_location: str = "local" #: :meta private:75 """Classifier location. It could be either of 'local' or 'pebblo-cloud'."""76 _discover_sent: bool = False #: :meta private:77 """Flag to check if discover payload has been sent."""78 enable_prompt_gov: bool = True #: :meta private:79 """Flag to check if prompt governance is enabled or not"""80 pb_client: PebbloRetrievalAPIWrapper = Field(81 default_factory=PebbloRetrievalAPIWrapper82 )83 """Pebblo Retrieval API client"""84 85 def _call(86 self,87 inputs: Dict[str, Any],88 run_manager: Optional[CallbackManagerForChainRun] = None,89 ) -> Dict[str, Any]:90 """Run get_relevant_text and llm on input query.91 92 If chain has 'return_source_documents' as 'True', returns93 the retrieved documents as well under the key 'source_documents'.94 95 Example:96 .. code-block:: python97 98 res = indexqa({'query': 'This is my query'})99 answer, docs = res['result'], res['source_documents']100 """101 prompt_time = datetime.datetime.now().isoformat()102 _run_manager = run_manager or CallbackManagerForChainRun.get_noop_manager()103 question = inputs[self.input_key]104 auth_context = inputs.get(self.auth_context_key)105 semantic_context = inputs.get(self.semantic_context_key)106 _, prompt_entities = self.pb_client.check_prompt_validity(question)107 108 accepts_run_manager = (109 "run_manager" in inspect.signature(self._get_docs).parameters110 )111 if accepts_run_manager:112 docs = self._get_docs(113 question, auth_context, semantic_context, run_manager=_run_manager114 )115 else:116 docs = self._get_docs(question, auth_context, semantic_context) # type: ignore[call-arg]117 answer = self.combine_documents_chain.run(118 input_documents=docs, question=question, callbacks=_run_manager.get_child()119 )120 121 self.pb_client.send_prompt(122 self.app_name,123 self.retriever,124 question,125 answer,126 auth_context,127 docs,128 prompt_entities,129 prompt_time,130 self.enable_prompt_gov,131 )132 133 if self.return_source_documents:134 return {self.output_key: answer, "source_documents": docs}135 else:136 return {self.output_key: answer}137 138 async def _acall(139 self,140 inputs: Dict[str, Any],141 run_manager: Optional[AsyncCallbackManagerForChainRun] = None,142 ) -> Dict[str, Any]:143 """Run get_relevant_text and llm on input query.144 145 If chain has 'return_source_documents' as 'True', returns146 the retrieved documents as well under the key 'source_documents'.147 148 Example:149 .. code-block:: python150 151 res = indexqa({'query': 'This is my query'})152 answer, docs = res['result'], res['source_documents']153 """154 prompt_time = datetime.datetime.now().isoformat()155 _run_manager = run_manager or AsyncCallbackManagerForChainRun.get_noop_manager()156 question = inputs[self.input_key]157 auth_context = inputs.get(self.auth_context_key)158 semantic_context = inputs.get(self.semantic_context_key)159 accepts_run_manager = (160 "run_manager" in inspect.signature(self._aget_docs).parameters161 )162 163 _, prompt_entities = await self.pb_client.acheck_prompt_validity(question)164 165 if accepts_run_manager:166 docs = await self._aget_docs(167 question, auth_context, semantic_context, run_manager=_run_manager168 )169 else:170 docs = await self._aget_docs(question, auth_context, semantic_context) # type: ignore[call-arg]171 answer = await self.combine_documents_chain.arun(172 input_documents=docs, question=question, callbacks=_run_manager.get_child()173 )174 175 await self.pb_client.asend_prompt(176 self.app_name,177 self.retriever,178 question,179 answer,180 auth_context,181 docs,182 prompt_entities,183 prompt_time,184 self.enable_prompt_gov,185 )186 187 if self.return_source_documents:188 return {self.output_key: answer, "source_documents": docs}189 else:190 return {self.output_key: answer}191 192 model_config = ConfigDict(193 populate_by_name=True,194 arbitrary_types_allowed=True,195 extra="forbid",196 )197 198 @property199 def input_keys(self) -> List[str]:200 """Input keys.201 202 :meta private:203 """204 return [self.input_key, self.auth_context_key, self.semantic_context_key]205 206 @property207 def output_keys(self) -> List[str]:208 """Output keys.209 210 :meta private:211 """212 _output_keys = [self.output_key]213 if self.return_source_documents:214 _output_keys += ["source_documents"]215 return _output_keys216 217 @property218 def _chain_type(self) -> str:219 """Return the chain type."""220 return "pebblo_retrieval_qa"221 222 @classmethod223 def from_chain_type(224 cls,225 llm: BaseLanguageModel,226 app_name: str,227 description: str,228 owner: str,229 chain_type: str = "stuff",230 chain_type_kwargs: Optional[dict] = None,231 api_key: Optional[str] = None,232 classifier_url: Optional[str] = None,233 classifier_location: str = "local",234 **kwargs: Any,235 ) -> "PebbloRetrievalQA":236 """Load chain from chain type."""237 from langchain_classic.chains.question_answering import load_qa_chain238 239 _chain_type_kwargs = chain_type_kwargs or {}240 combine_documents_chain = load_qa_chain(241 llm, chain_type=chain_type, **_chain_type_kwargs242 )243 244 # generate app245 app: App = PebbloRetrievalQA._get_app_details(246 app_name=app_name,247 description=description,248 owner=owner,249 llm=llm,250 **kwargs,251 )252 # initialize Pebblo API client253 pb_client = PebbloRetrievalAPIWrapper(254 api_key=api_key,255 classifier_location=classifier_location,256 classifier_url=classifier_url,257 )258 # send app discovery request259 pb_client.send_app_discover(app)260 return cls(261 combine_documents_chain=combine_documents_chain,262 app_name=app_name,263 owner=owner,264 description=description,265 api_key=api_key,266 classifier_url=classifier_url,267 classifier_location=classifier_location,268 pb_client=pb_client,269 **kwargs,270 )271 272 @validator("retriever", pre=True, always=True)273 def validate_vectorstore(274 cls, retriever: VectorStoreRetriever275 ) -> VectorStoreRetriever:276 """277 Validate that the vectorstore of the retriever is supported vectorstores.278 """279 if retriever.vectorstore.__class__.__name__ not in SUPPORTED_VECTORSTORES:280 raise ValueError(281 f"Vectorstore must be an instance of one of the supported "282 f"vectorstores: {SUPPORTED_VECTORSTORES}. "283 f"Got '{retriever.vectorstore.__class__.__name__}' instead."284 )285 return retriever286 287 def _get_docs(288 self,289 question: str,290 auth_context: Optional[AuthContext],291 semantic_context: Optional[SemanticContext],292 *,293 run_manager: CallbackManagerForChainRun,294 ) -> List[Document]:295 """Get docs."""296 set_enforcement_filters(self.retriever, auth_context, semantic_context)297 return self.retriever.invoke(298 question, config={"callbacks": run_manager.get_child()}299 )300 301 async def _aget_docs(302 self,303 question: str,304 auth_context: Optional[AuthContext],305 semantic_context: Optional[SemanticContext],306 *,307 run_manager: AsyncCallbackManagerForChainRun,308 ) -> List[Document]:309 """Get docs."""310 set_enforcement_filters(self.retriever, auth_context, semantic_context)311 return await self.retriever.ainvoke(312 question, config={"callbacks": run_manager.get_child()}313 )314 315 @staticmethod316 def _get_app_details(317 app_name: str,318 owner: str,319 description: str,320 llm: BaseLanguageModel,321 **kwargs: Any,322 ) -> App:323 """Fetch app details. Internal method.324 Returns:325 App: App details.326 """327 framework, runtime = get_runtime()328 chains = PebbloRetrievalQA.get_chain_details(llm, **kwargs)329 app = App(330 name=app_name,331 owner=owner,332 description=description,333 runtime=runtime,334 framework=framework,335 chains=chains,336 plugin_version=PLUGIN_VERSION,337 client_version=Framework(338 name="langchain_community",339 version=version("langchain_community"),340 ),341 )342 return app343 344 @classmethod345 def set_discover_sent(cls) -> None:346 cls._discover_sent = True347 348 @classmethod349 def get_chain_details(350 cls, llm: BaseLanguageModel, **kwargs: Any351 ) -> List[ChainInfo]:352 """353 Get chain details.354 355 Args:356 llm (BaseLanguageModel): Language model instance.357 **kwargs: Additional keyword arguments.358 359 Returns:360 List[ChainInfo]: Chain details.361 """362 llm_dict = llm.__dict__363 chains = [364 ChainInfo(365 name=cls.__name__,366 model=Model(367 name=llm_dict.get("model_name", llm_dict.get("model")),368 vendor=llm.__class__.__name__,369 ),370 vector_dbs=[371 VectorDB(372 name=kwargs["retriever"].vectorstore.__class__.__name__,373 embedding_model=str(374 kwargs["retriever"].vectorstore._embeddings.model375 )376 if hasattr(kwargs["retriever"].vectorstore, "_embeddings")377 else (378 str(kwargs["retriever"].vectorstore._embedding.model)379 if hasattr(kwargs["retriever"].vectorstore, "_embedding")380 else None381 ),382 )383 ],384 ),385 ]386 return chains387 