Spaces:
Sleeping
Sleeping
| from typing import Optional, List | |
| try: | |
| from sqlalchemy.dialects import postgresql | |
| from sqlalchemy.engine import create_engine, Engine | |
| from sqlalchemy.inspection import inspect | |
| from sqlalchemy.orm import sessionmaker, scoped_session | |
| from sqlalchemy.schema import MetaData, Table, Column | |
| from sqlalchemy.sql.expression import text, select, delete | |
| from sqlalchemy.types import DateTime, String | |
| except ImportError: | |
| raise ImportError("`sqlalchemy` not installed") | |
| from phi.memory.db import MemoryDb | |
| from phi.memory.row import MemoryRow | |
| from phi.utils.log import logger | |
| class PgMemoryDb(MemoryDb): | |
| def __init__( | |
| self, | |
| table_name: str, | |
| schema: Optional[str] = "ai", | |
| db_url: Optional[str] = None, | |
| db_engine: Optional[Engine] = None, | |
| ): | |
| """ | |
| This class provides a memory store backed by a postgres table. | |
| The following order is used to determine the database connection: | |
| 1. Use the db_engine if provided | |
| 2. Use the db_url to create the engine | |
| Args: | |
| table_name (str): The name of the table to store memory rows. | |
| schema (Optional[str]): The schema to store the table in. Defaults to "ai". | |
| db_url (Optional[str]): The database URL to connect to. Defaults to None. | |
| db_engine (Optional[Engine]): The database engine to use. Defaults to None. | |
| """ | |
| _engine: Optional[Engine] = db_engine | |
| if _engine is None and db_url is not None: | |
| _engine = create_engine(db_url) | |
| if _engine is None: | |
| raise ValueError("Must provide either db_url or db_engine") | |
| self.table_name: str = table_name | |
| self.schema: Optional[str] = schema | |
| self.db_url: Optional[str] = db_url | |
| self.db_engine: Engine = _engine | |
| self.inspector = inspect(self.db_engine) | |
| self.metadata: MetaData = MetaData(schema=self.schema) | |
| self.Session: scoped_session = scoped_session(sessionmaker(bind=self.db_engine)) | |
| self.table: Table = self.get_table() | |
| def get_table(self) -> Table: | |
| return Table( | |
| self.table_name, | |
| self.metadata, | |
| Column("id", String, primary_key=True), | |
| Column("user_id", String), | |
| Column("memory", postgresql.JSONB, server_default=text("'{}'::jsonb")), | |
| Column("created_at", DateTime(timezone=True), server_default=text("now()")), | |
| Column("updated_at", DateTime(timezone=True), onupdate=text("now()")), | |
| extend_existing=True, | |
| ) | |
| def create(self) -> None: | |
| if not self.table_exists(): | |
| try: | |
| with self.Session() as sess, sess.begin(): | |
| if self.schema is not None: | |
| logger.debug(f"Creating schema: {self.schema}") | |
| sess.execute(text(f"CREATE SCHEMA IF NOT EXISTS {self.schema};")) | |
| logger.debug(f"Creating table: {self.table_name}") | |
| self.table.create(self.db_engine, checkfirst=True) | |
| except Exception as e: | |
| logger.error(f"Error creating table '{self.table.fullname}': {e}") | |
| raise | |
| def memory_exists(self, memory: MemoryRow) -> bool: | |
| columns = [self.table.c.id] | |
| with self.Session() as sess, sess.begin(): | |
| stmt = select(*columns).where(self.table.c.id == memory.id) | |
| result = sess.execute(stmt).first() | |
| return result is not None | |
| def read_memories( | |
| self, user_id: Optional[str] = None, limit: Optional[int] = None, sort: Optional[str] = None | |
| ) -> List[MemoryRow]: | |
| memories: List[MemoryRow] = [] | |
| try: | |
| with self.Session() as sess, sess.begin(): | |
| stmt = select(self.table) | |
| if user_id is not None: | |
| stmt = stmt.where(self.table.c.user_id == user_id) | |
| if limit is not None: | |
| stmt = stmt.limit(limit) | |
| if sort == "asc": | |
| stmt = stmt.order_by(self.table.c.created_at.asc()) | |
| else: | |
| stmt = stmt.order_by(self.table.c.created_at.desc()) | |
| rows = sess.execute(stmt).fetchall() | |
| for row in rows: | |
| if row is not None: | |
| memories.append(MemoryRow.model_validate(row)) | |
| except Exception as e: | |
| logger.debug(f"Exception reading from table: {e}") | |
| logger.debug(f"Table does not exist: {self.table.name}") | |
| logger.debug("Creating table for future transactions") | |
| self.create() | |
| return memories | |
| def upsert_memory(self, memory: MemoryRow, create_and_retry: bool = True) -> None: | |
| """Create a new memory if it does not exist, otherwise update the existing memory""" | |
| try: | |
| with self.Session() as sess, sess.begin(): | |
| # Create an insert statement | |
| stmt = postgresql.insert(self.table).values( | |
| id=memory.id, | |
| user_id=memory.user_id, | |
| memory=memory.memory, | |
| ) | |
| # Define the upsert if the memory already exists | |
| # See: https://docs.sqlalchemy.org/en/20/dialects/postgresql.html#postgresql-insert-on-conflict | |
| stmt = stmt.on_conflict_do_update( | |
| index_elements=["id"], | |
| set_=dict( | |
| user_id=stmt.excluded.user_id, | |
| memory=stmt.excluded.memory, | |
| ), | |
| ) | |
| sess.execute(stmt) | |
| except Exception as e: | |
| logger.debug(f"Exception upserting into table: {e}") | |
| logger.debug(f"Table does not exist: {self.table.name}") | |
| logger.debug("Creating table for future transactions") | |
| self.create() | |
| if create_and_retry: | |
| return self.upsert_memory(memory, create_and_retry=False) | |
| return None | |
| def delete_memory(self, id: str) -> None: | |
| with self.Session() as sess, sess.begin(): | |
| stmt = delete(self.table).where(self.table.c.id == id) | |
| sess.execute(stmt) | |
| def drop_table(self) -> None: | |
| if self.table_exists(): | |
| logger.debug(f"Deleting table: {self.table_name}") | |
| self.table.drop(self.db_engine) | |
| def table_exists(self) -> bool: | |
| logger.debug(f"Checking if table exists: {self.table.name}") | |
| try: | |
| return inspect(self.db_engine).has_table(self.table.name, schema=self.schema) | |
| except Exception as e: | |
| logger.error(e) | |
| return False | |
| def clear(self) -> bool: | |
| with self.Session() as sess, sess.begin(): | |
| stmt = delete(self.table) | |
| sess.execute(stmt) | |
| return True | |
| def __deepcopy__(self, memo): | |
| """ | |
| Create a deep copy of the PgMemoryDb instance, handling unpickleable attributes. | |
| Args: | |
| memo (dict): A dictionary of objects already copied during the current copying pass. | |
| Returns: | |
| PgMemoryDb: A deep-copied instance of PgMemoryDb. | |
| """ | |
| from copy import deepcopy | |
| # Create a new instance without calling __init__ | |
| cls = self.__class__ | |
| copied_obj = cls.__new__(cls) | |
| memo[id(self)] = copied_obj | |
| # Deep copy attributes | |
| for k, v in self.__dict__.items(): | |
| if k in {"metadata", "table"}: | |
| continue | |
| # Reuse db_engine and Session without copying | |
| elif k in {"db_engine", "Session"}: | |
| setattr(copied_obj, k, v) | |
| else: | |
| setattr(copied_obj, k, deepcopy(v, memo)) | |
| # Recreate metadata and table for the copied instance | |
| copied_obj.metadata = MetaData(schema=copied_obj.schema) | |
| copied_obj.table = copied_obj.get_table() | |
| return copied_obj | |