| import sqlalchemy |
| from langchain.vectorstores.pgvector import PGVector as Base,CollectionStore |
| from sqlalchemy.orm import Session |
|
|
| class PGVector(Base): |
|
|
| def __init__( |
| self, |
| pre_delete_embeddings: bool = False, |
| **kwargs |
| ) -> None: |
| self.pre_delete_embeddings = pre_delete_embeddings |
| Base.__init__(self, **kwargs) |
|
|
| def connect(self) -> sqlalchemy.engine.Connection: |
| engine = sqlalchemy.create_engine(self.connection_string, echo=True) |
| conn = engine.connect() |
| return conn |
| |
| def delete_embeddings(self) -> None: |
| self.logger.debug("Trying to delete embeddings") |
| with Session(self._conn) as session: |
| collection = self.get_collection(session) |
| if not collection: |
| self.logger.error("Collection not found") |
| return |
| session.delete(collection.embeddings) |
| session.commit() |
|
|
| def create_collection(self) -> None: |
| if self.pre_delete_collection: |
| self.delete_collection() |
| with Session(self._conn) as session: |
| CollectionStore.get_or_create( |
| session, self.collection_name, cmetadata=self.collection_metadata |
| ) |
| if self.pre_delete_embeddings: |
| self.delete_embeddings() |