Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
vectorstore.py126 linesDownload Raw Back to memory
1"""Class for a VectorStore-backed memory object."""2 3from collections.abc import Sequence4from typing import Any5 6from langchain_core._api import deprecated7from langchain_core.documents import Document8from langchain_core.vectorstores import VectorStoreRetriever9from pydantic import Field10 11from langchain_classic.base_memory import BaseMemory12from langchain_classic.memory.utils import get_prompt_input_key13 14 15@deprecated(16    since="0.3.1",17    removal="2.0.0",18    alternative="langchain.agents.create_agent",19    addendum=(20        "For agents that need to remember prior interactions, use "21        "`create_agent` with checkpointing or the `Store` API. See "22        "https://docs.langchain.com/oss/python/langchain/short-term-memory and "23        "https://docs.langchain.com/oss/python/langchain/long-term-memory"24    ),25)26class VectorStoreRetrieverMemory(BaseMemory):27    """Vector Store Retriever Memory.28 29    Store the conversation history in a vector store and retrieves the relevant30    parts of past conversation based on the input.31    """32 33    retriever: VectorStoreRetriever = Field(exclude=True)34    """VectorStoreRetriever object to connect to."""35 36    memory_key: str = "history"37    """Key name to locate the memories in the result of load_memory_variables."""38 39    input_key: str | None = None40    """Key name to index the inputs to load_memory_variables."""41 42    return_docs: bool = False43    """Whether or not to return the result of querying the database directly."""44 45    exclude_input_keys: Sequence[str] = Field(default_factory=tuple)46    """Input keys to exclude in addition to memory key when constructing the document"""47 48    @property49    def memory_variables(self) -> list[str]:50        """The list of keys emitted from the load_memory_variables method."""51        return [self.memory_key]52 53    def _get_prompt_input_key(self, inputs: dict[str, Any]) -> str:54        """Get the input key for the prompt."""55        if self.input_key is None:56            return get_prompt_input_key(inputs, self.memory_variables)57        return self.input_key58 59    def _documents_to_memory_variables(60        self,61        docs: list[Document],62    ) -> dict[str, list[Document] | str]:63        result: list[Document] | str64        if not self.return_docs:65            result = "\n".join([doc.page_content for doc in docs])66        else:67            result = docs68        return {self.memory_key: result}69 70    def load_memory_variables(71        self,72        inputs: dict[str, Any],73    ) -> dict[str, list[Document] | str]:74        """Return history buffer."""75        input_key = self._get_prompt_input_key(inputs)76        query = inputs[input_key]77        docs = self.retriever.invoke(query)78        return self._documents_to_memory_variables(docs)79 80    async def aload_memory_variables(81        self,82        inputs: dict[str, Any],83    ) -> dict[str, list[Document] | str]:84        """Return history buffer."""85        input_key = self._get_prompt_input_key(inputs)86        query = inputs[input_key]87        docs = await self.retriever.ainvoke(query)88        return self._documents_to_memory_variables(docs)89 90    def _form_documents(91        self,92        inputs: dict[str, Any],93        outputs: dict[str, str],94    ) -> list[Document]:95        """Format context from this conversation to buffer."""96        # Each document should only include the current turn, not the chat history97        exclude = set(self.exclude_input_keys)98        exclude.add(self.memory_key)99        filtered_inputs = {k: v for k, v in inputs.items() if k not in exclude}100        texts = [101            f"{k}: {v}"102            for k, v in list(filtered_inputs.items()) + list(outputs.items())103        ]104        page_content = "\n".join(texts)105        return [Document(page_content=page_content)]106 107    def save_context(self, inputs: dict[str, Any], outputs: dict[str, str]) -> None:108        """Save context from this conversation to buffer."""109        documents = self._form_documents(inputs, outputs)110        self.retriever.add_documents(documents)111 112    async def asave_context(113        self,114        inputs: dict[str, Any],115        outputs: dict[str, str],116    ) -> None:117        """Save context from this conversation to buffer."""118        documents = self._form_documents(inputs, outputs)119        await self.retriever.aadd_documents(documents)120 121    def clear(self) -> None:122        """Nothing to clear."""123 124    async def aclear(self) -> None:125        """Nothing to clear."""126 
codekingpro/portable-devtools · Team Ai