Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
base.py387 linesDownload Raw Back to pebblo_retrieval
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 
codekingpro/portable-devtools · Team Ai