Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
entity.py630 linesDownload Raw Back to memory
1"""Deprecated as of LangChain v0.3.4 and will be removed in LangChain v1.0.0."""2 3import logging4from abc import ABC, abstractmethod5from collections.abc import Iterable6from itertools import islice7from typing import TYPE_CHECKING, Any8 9from langchain_core._api import deprecated10from langchain_core.language_models import BaseLanguageModel11from langchain_core.messages import BaseMessage, get_buffer_string12from langchain_core.prompts import BasePromptTemplate13from pydantic import BaseModel, ConfigDict, Field14from typing_extensions import override15 16from langchain_classic.chains.llm import LLMChain17from langchain_classic.memory.chat_memory import BaseChatMemory18from langchain_classic.memory.prompt import (19    ENTITY_EXTRACTION_PROMPT,20    ENTITY_SUMMARIZATION_PROMPT,21)22from langchain_classic.memory.utils import get_prompt_input_key23 24if TYPE_CHECKING:25    import sqlite326 27logger = logging.getLogger(__name__)28 29 30@deprecated(31    since="0.3.1",32    removal="2.0.0",33    alternative="langchain.agents.create_agent",34    addendum=(35        "For agents that need to remember prior interactions, use "36        "`create_agent` with checkpointing or the `Store` API. See "37        "https://docs.langchain.com/oss/python/langchain/short-term-memory and "38        "https://docs.langchain.com/oss/python/langchain/long-term-memory"39    ),40)41class BaseEntityStore(BaseModel, ABC):42    """Abstract base class for Entity store."""43 44    @abstractmethod45    def get(self, key: str, default: str | None = None) -> str | None:46        """Get entity value from store."""47 48    @abstractmethod49    def set(self, key: str, value: str | None) -> None:50        """Set entity value in store."""51 52    @abstractmethod53    def delete(self, key: str) -> None:54        """Delete entity value from store."""55 56    @abstractmethod57    def exists(self, key: str) -> bool:58        """Check if entity exists in store."""59 60    @abstractmethod61    def clear(self) -> None:62        """Delete all entities from store."""63 64 65@deprecated(66    since="0.3.1",67    removal="2.0.0",68    alternative="langchain.agents.create_agent",69    addendum=(70        "For agents that need to remember prior interactions, use "71        "`create_agent` with checkpointing or the `Store` API. See "72        "https://docs.langchain.com/oss/python/langchain/short-term-memory and "73        "https://docs.langchain.com/oss/python/langchain/long-term-memory"74    ),75)76class InMemoryEntityStore(BaseEntityStore):77    """In-memory Entity store."""78 79    store: dict[str, str | None] = {}80 81    @override82    def get(self, key: str, default: str | None = None) -> str | None:83        return self.store.get(key, default)84 85    @override86    def set(self, key: str, value: str | None) -> None:87        self.store[key] = value88 89    @override90    def delete(self, key: str) -> None:91        del self.store[key]92 93    @override94    def exists(self, key: str) -> bool:95        return key in self.store96 97    @override98    def clear(self) -> None:99        return self.store.clear()100 101 102@deprecated(103    since="0.3.1",104    removal="2.0.0",105    alternative="langchain.agents.create_agent",106    addendum=(107        "For agents that need to remember prior interactions, use "108        "`create_agent` with checkpointing or the `Store` API. See "109        "https://docs.langchain.com/oss/python/langchain/short-term-memory and "110        "https://docs.langchain.com/oss/python/langchain/long-term-memory"111    ),112)113class UpstashRedisEntityStore(BaseEntityStore):114    """Upstash Redis backed Entity store.115 116    Entities get a TTL of 1 day by default, and117    that TTL is extended by 3 days every time the entity is read back.118    """119 120    def __init__(121        self,122        session_id: str = "default",123        url: str = "",124        token: str = "",125        key_prefix: str = "memory_store",126        ttl: int | None = 60 * 60 * 24,127        recall_ttl: int | None = 60 * 60 * 24 * 3,128        *args: Any,129        **kwargs: Any,130    ):131        """Initializes the RedisEntityStore.132 133        Args:134            session_id: Unique identifier for the session.135            url: URL of the Redis server.136            token: Authentication token for the Redis server.137            key_prefix: Prefix for keys in the Redis store.138            ttl: Time-to-live for keys in seconds (default 1 day).139            recall_ttl: Time-to-live extension for keys when recalled (default 3 days).140            *args: Additional positional arguments.141            **kwargs: Additional keyword arguments.142        """143        try:144            from upstash_redis import Redis145        except ImportError as e:146            msg = (147                "Could not import upstash_redis python package. "148                "Please install it with `pip install upstash_redis`."149            )150            raise ImportError(msg) from e151 152        super().__init__(*args, **kwargs)153 154        try:155            self.redis_client = Redis(url=url, token=token)156        except Exception as exc:157            error_msg = "Upstash Redis instance could not be initiated"158            logger.exception(error_msg)159            raise RuntimeError(error_msg) from exc160 161        self.session_id = session_id162        self.key_prefix = key_prefix163        self.ttl = ttl164        self.recall_ttl = recall_ttl or ttl165 166    @property167    def full_key_prefix(self) -> str:168        """Returns the full key prefix with session ID."""169        return f"{self.key_prefix}:{self.session_id}"170 171    @override172    def get(self, key: str, default: str | None = None) -> str | None:173        res = (174            self.redis_client.getex(f"{self.full_key_prefix}:{key}", ex=self.recall_ttl)175            or default176            or ""177        )178        logger.debug(179            "Upstash Redis MEM get '%s:%s': '%s'", self.full_key_prefix, key, res180        )181        return res182 183    @override184    def set(self, key: str, value: str | None) -> None:185        if not value:186            return self.delete(key)187        self.redis_client.set(f"{self.full_key_prefix}:{key}", value, ex=self.ttl)188        logger.debug(189            "Redis MEM set '%s:%s': '%s' EX %s",190            self.full_key_prefix,191            key,192            value,193            self.ttl,194        )195        return None196 197    @override198    def delete(self, key: str) -> None:199        self.redis_client.delete(f"{self.full_key_prefix}:{key}")200 201    @override202    def exists(self, key: str) -> bool:203        return self.redis_client.exists(f"{self.full_key_prefix}:{key}") == 1204 205    @override206    def clear(self) -> None:207        def scan_and_delete(cursor: int) -> int:208            cursor, keys_to_delete = self.redis_client.scan(209                cursor,210                f"{self.full_key_prefix}:*",211            )212            self.redis_client.delete(*keys_to_delete)213            return cursor214 215        cursor = scan_and_delete(0)216        while cursor != 0:217            scan_and_delete(cursor)218 219 220@deprecated(221    since="0.3.1",222    removal="2.0.0",223    alternative="langchain.agents.create_agent",224    addendum=(225        "For agents that need to remember prior interactions, use "226        "`create_agent` with checkpointing or the `Store` API. See "227        "https://docs.langchain.com/oss/python/langchain/short-term-memory and "228        "https://docs.langchain.com/oss/python/langchain/long-term-memory"229    ),230)231class RedisEntityStore(BaseEntityStore):232    """Redis-backed Entity store.233 234    Entities get a TTL of 1 day by default, and235    that TTL is extended by 3 days every time the entity is read back.236    """237 238    redis_client: Any239    session_id: str = "default"240    key_prefix: str = "memory_store"241    ttl: int | None = 60 * 60 * 24242    recall_ttl: int | None = 60 * 60 * 24 * 3243 244    def __init__(245        self,246        session_id: str = "default",247        url: str = "redis://localhost:6379/0",248        key_prefix: str = "memory_store",249        ttl: int | None = 60 * 60 * 24,250        recall_ttl: int | None = 60 * 60 * 24 * 3,251        *args: Any,252        **kwargs: Any,253    ):254        """Initializes the RedisEntityStore.255 256        Args:257            session_id: Unique identifier for the session.258            url: URL of the Redis server.259            key_prefix: Prefix for keys in the Redis store.260            ttl: Time-to-live for keys in seconds (default 1 day).261            recall_ttl: Time-to-live extension for keys when recalled (default 3 days).262            *args: Additional positional arguments.263            **kwargs: Additional keyword arguments.264        """265        try:266            import redis267        except ImportError as e:268            msg = (269                "Could not import redis python package. "270                "Please install it with `pip install redis`."271            )272            raise ImportError(msg) from e273 274        super().__init__(*args, **kwargs)275 276        try:277            from langchain_community.utilities.redis import get_client278        except ImportError as e:279            msg = (280                "Could not import langchain_community.utilities.redis.get_client. "281                "Please install it with `pip install langchain-community`."282            )283            raise ImportError(msg) from e284 285        try:286            self.redis_client = get_client(redis_url=url, decode_responses=True)287        except redis.exceptions.ConnectionError:288            logger.exception("Redis client could not connect")289 290        self.session_id = session_id291        self.key_prefix = key_prefix292        self.ttl = ttl293        self.recall_ttl = recall_ttl or ttl294 295    @property296    def full_key_prefix(self) -> str:297        """Returns the full key prefix with session ID."""298        return f"{self.key_prefix}:{self.session_id}"299 300    @override301    def get(self, key: str, default: str | None = None) -> str | None:302        res = (303            self.redis_client.getex(f"{self.full_key_prefix}:{key}", ex=self.recall_ttl)304            or default305            or ""306        )307        logger.debug("REDIS MEM get '%s:%s': '%s'", self.full_key_prefix, key, res)308        return res309 310    @override311    def set(self, key: str, value: str | None) -> None:312        if not value:313            return self.delete(key)314        self.redis_client.set(f"{self.full_key_prefix}:{key}", value, ex=self.ttl)315        logger.debug(316            "REDIS MEM set '%s:%s': '%s' EX %s",317            self.full_key_prefix,318            key,319            value,320            self.ttl,321        )322        return None323 324    @override325    def delete(self, key: str) -> None:326        self.redis_client.delete(f"{self.full_key_prefix}:{key}")327 328    @override329    def exists(self, key: str) -> bool:330        return self.redis_client.exists(f"{self.full_key_prefix}:{key}") == 1331 332    @override333    def clear(self) -> None:334        # iterate a list in batches of size batch_size335        def batched(iterable: Iterable[Any], batch_size: int) -> Iterable[Any]:336            iterator = iter(iterable)337            while batch := list(islice(iterator, batch_size)):338                yield batch339 340        for keybatch in batched(341            self.redis_client.scan_iter(f"{self.full_key_prefix}:*"),342            500,343        ):344            self.redis_client.delete(*keybatch)345 346 347@deprecated(348    since="0.3.1",349    removal="2.0.0",350    alternative="langchain.agents.create_agent",351    addendum=(352        "For agents that need to remember prior interactions, use "353        "`create_agent` with checkpointing or the `Store` API. See "354        "https://docs.langchain.com/oss/python/langchain/short-term-memory and "355        "https://docs.langchain.com/oss/python/langchain/long-term-memory"356    ),357)358class SQLiteEntityStore(BaseEntityStore):359    """SQLite-backed Entity store with safe query construction."""360 361    session_id: str = "default"362    table_name: str = "memory_store"363    conn: Any = None364 365    model_config = ConfigDict(366        arbitrary_types_allowed=True,367    )368 369    def __init__(370        self,371        session_id: str = "default",372        db_file: str = "entities.db",373        table_name: str = "memory_store",374        *args: Any,375        **kwargs: Any,376    ):377        """Initializes the SQLiteEntityStore.378 379        Args:380            session_id: Unique identifier for the session.381            db_file: Path to the SQLite database file.382            table_name: Name of the table to store entities.383            *args: Additional positional arguments.384            **kwargs: Additional keyword arguments.385        """386        super().__init__(*args, **kwargs)387        try:388            import sqlite3389        except ImportError as e:390            msg = (391                "Could not import sqlite3 python package. "392                "Please install it with `pip install sqlite3`."393            )394            raise ImportError(msg) from e395 396        # Basic validation to prevent obviously malicious table/session names397        if not table_name.isidentifier() or not session_id.isidentifier():398            # Since we validate here, we can safely suppress the S608 bandit warning399            msg = "Table name and session ID must be valid Python identifiers."400            raise ValueError(msg)401 402        self.conn = sqlite3.connect(db_file)403        self.session_id = session_id404        self.table_name = table_name405        self._create_table_if_not_exists()406 407    @property408    def full_table_name(self) -> str:409        """Returns the full table name with session ID."""410        return f"{self.table_name}_{self.session_id}"411 412    def _execute_query(self, query: str, params: tuple = ()) -> "sqlite3.Cursor":413        """Executes a query with proper connection handling."""414        with self.conn:415            return self.conn.execute(query, params)416 417    def _create_table_if_not_exists(self) -> None:418        """Creates the entity table if it doesn't exist, using safe quoting."""419        # Use standard SQL double quotes for the table name identifier420        create_table_query = f"""421            CREATE TABLE IF NOT EXISTS "{self.full_table_name}" (422                key TEXT PRIMARY KEY,423                value TEXT424            )425        """426        self._execute_query(create_table_query)427 428    def get(self, key: str, default: str | None = None) -> str | None:429        """Retrieves a value, safely quoting the table name."""430        # `?` placeholder is used for the value to prevent SQL injection431        # Ignore S608 since we validate for malicious table/session names in `__init__`432        query = f'SELECT value FROM "{self.full_table_name}" WHERE key = ?'  # noqa: S608433        cursor = self._execute_query(query, (key,))434        result = cursor.fetchone()435        return result[0] if result is not None else default436 437    def set(self, key: str, value: str | None) -> None:438        """Inserts or replaces a value, safely quoting the table name."""439        if not value:440            return self.delete(key)441        # Ignore S608 since we validate for malicious table/session names in `__init__`442        query = (443            "INSERT OR REPLACE INTO "  # noqa: S608444            f'"{self.full_table_name}" (key, value) VALUES (?, ?)'445        )446        self._execute_query(query, (key, value))447        return None448 449    def delete(self, key: str) -> None:450        """Deletes a key-value pair, safely quoting the table name."""451        # Ignore S608 since we validate for malicious table/session names in `__init__`452        query = f'DELETE FROM "{self.full_table_name}" WHERE key = ?'  # noqa: S608453        self._execute_query(query, (key,))454 455    def exists(self, key: str) -> bool:456        """Checks for the existence of a key, safely quoting the table name."""457        # Ignore S608 since we validate for malicious table/session names in `__init__`458        query = f'SELECT 1 FROM "{self.full_table_name}" WHERE key = ? LIMIT 1'  # noqa: S608459        cursor = self._execute_query(query, (key,))460        return cursor.fetchone() is not None461 462    @override463    def clear(self) -> None:464        # Ignore S608 since we validate for malicious table/session names in `__init__`465        query = f"""466            DELETE FROM {self.full_table_name}467        """  # noqa: S608468        with self.conn:469            self.conn.execute(query)470 471 472@deprecated(473    since="0.3.1",474    removal="2.0.0",475    alternative="langchain.agents.create_agent",476    addendum=(477        "For agents that need to remember prior interactions, use "478        "`create_agent` with checkpointing or the `Store` API. See "479        "https://docs.langchain.com/oss/python/langchain/short-term-memory and "480        "https://docs.langchain.com/oss/python/langchain/long-term-memory"481    ),482)483class ConversationEntityMemory(BaseChatMemory):484    """Entity extractor & summarizer memory.485 486    Extracts named entities from the recent chat history and generates summaries.487    With a swappable entity store, persisting entities across conversations.488    Defaults to an in-memory entity store, and can be swapped out for a Redis,489    SQLite, or other entity store.490    """491 492    human_prefix: str = "Human"493    ai_prefix: str = "AI"494    llm: BaseLanguageModel495    entity_extraction_prompt: BasePromptTemplate = ENTITY_EXTRACTION_PROMPT496    entity_summarization_prompt: BasePromptTemplate = ENTITY_SUMMARIZATION_PROMPT497 498    # Cache of recently detected entity names, if any499    # It is updated when load_memory_variables is called:500    entity_cache: list[str] = []501 502    # Number of recent message pairs to consider when updating entities:503    k: int = 3504 505    chat_history_key: str = "history"506 507    # Store to manage entity-related data:508    entity_store: BaseEntityStore = Field(default_factory=InMemoryEntityStore)509 510    @property511    def buffer(self) -> list[BaseMessage]:512        """Access chat memory messages."""513        return self.chat_memory.messages514 515    @property516    def memory_variables(self) -> list[str]:517        """Will always return list of memory variables."""518        return ["entities", self.chat_history_key]519 520    def load_memory_variables(self, inputs: dict[str, Any]) -> dict[str, Any]:521        """Load memory variables.522 523        Returns chat history and all generated entities with summaries if available,524        and updates or clears the recent entity cache.525 526        New entity name can be found when calling this method, before the entity527        summaries are generated, so the entity cache values may be empty if no entity528        descriptions are generated yet.529        """530        # Create an LLMChain for predicting entity names from the recent chat history:531        chain = LLMChain(llm=self.llm, prompt=self.entity_extraction_prompt)532 533        if self.input_key is None:534            prompt_input_key = get_prompt_input_key(inputs, self.memory_variables)535        else:536            prompt_input_key = self.input_key537 538        # Extract an arbitrary window of the last message pairs from539        # the chat history, where the hyperparameter k is the540        # number of message pairs:541        buffer_string = get_buffer_string(542            self.buffer[-self.k * 2 :],543            human_prefix=self.human_prefix,544            ai_prefix=self.ai_prefix,545        )546 547        # Generates a comma-separated list of named entities,548        # e.g. "Jane, White House, UFO"549        # or "NONE" if no named entities are extracted:550        output = chain.predict(551            history=buffer_string,552            input=inputs[prompt_input_key],553        )554 555        # If no named entities are extracted, assigns an empty list.556        if output.strip() == "NONE":557            entities = []558        else:559            # Make a list of the extracted entities:560            entities = [w.strip() for w in output.split(",")]561 562        # Make a dictionary of entities with summary if exists:563        entity_summaries = {}564 565        for entity in entities:566            entity_summaries[entity] = self.entity_store.get(entity, "")567 568        # Replaces the entity name cache with the most recently discussed entities,569        # or if no entities were extracted, clears the cache:570        self.entity_cache = entities571 572        # Should we return as message objects or as a string?573        if self.return_messages:574            # Get last `k` pair of chat messages:575            buffer: Any = self.buffer[-self.k * 2 :]576        else:577            # Reuse the string we made earlier:578            buffer = buffer_string579 580        return {581            self.chat_history_key: buffer,582            "entities": entity_summaries,583        }584 585    def save_context(self, inputs: dict[str, Any], outputs: dict[str, str]) -> None:586        """Save context from this conversation history to the entity store.587 588        Generates a summary for each entity in the entity cache by prompting589        the model, and saves these summaries to the entity store.590        """591        super().save_context(inputs, outputs)592 593        if self.input_key is None:594            prompt_input_key = get_prompt_input_key(inputs, self.memory_variables)595        else:596            prompt_input_key = self.input_key597 598        # Extract an arbitrary window of the last message pairs from599        # the chat history, where the hyperparameter k is the600        # number of message pairs:601        buffer_string = get_buffer_string(602            self.buffer[-self.k * 2 :],603            human_prefix=self.human_prefix,604            ai_prefix=self.ai_prefix,605        )606 607        input_data = inputs[prompt_input_key]608 609        # Create an LLMChain for predicting entity summarization from the context610        chain = LLMChain(llm=self.llm, prompt=self.entity_summarization_prompt)611 612        # Generate new summaries for entities and save them in the entity store613        for entity in self.entity_cache:614            # Get existing summary if it exists615            existing_summary = self.entity_store.get(entity, "")616            output = chain.predict(617                summary=existing_summary,618                entity=entity,619                history=buffer_string,620                input=input_data,621            )622            # Save the updated summary to the entity store623            self.entity_store.set(entity, output.strip())624 625    def clear(self) -> None:626        """Clear memory contents."""627        self.chat_memory.clear()628        self.entity_cache.clear()629        self.entity_store.clear()630 
codekingpro/portable-devtools · Team Ai