codekingpro/portable-devtools
114k
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 