Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
pgembedding.py532 linesDownload Raw Back to vectorstores
1from __future__ import annotations2 3import logging4import uuid5from typing import Any, Dict, Iterable, List, Optional, Tuple, Type6 7import sqlalchemy8from sqlalchemy import func9from sqlalchemy.dialects.postgresql import JSON, UUID10from sqlalchemy.orm import Session, relationship11 12try:13    from sqlalchemy.orm import declarative_base14except ImportError:15    from sqlalchemy.ext.declarative import declarative_base16 17from langchain_core.documents import Document18from langchain_core.embeddings import Embeddings19from langchain_core.utils import get_from_dict_or_env20from langchain_core.vectorstores import VectorStore21 22Base = declarative_base()  # type: Any23 24 25ADA_TOKEN_COUNT = 153626_LANGCHAIN_DEFAULT_COLLECTION_NAME = "langchain"27 28 29class BaseModel(Base):30    """Base model for all SQL stores."""31 32    __abstract__ = True33    uuid = sqlalchemy.Column(UUID(as_uuid=True), primary_key=True, default=uuid.uuid4)34 35 36class CollectionStore(BaseModel):37    """Collection store."""38 39    __tablename__ = "langchain_pg_collection"40 41    name = sqlalchemy.Column(sqlalchemy.String)42    cmetadata = sqlalchemy.Column(JSON)43 44    embeddings = relationship(45        "EmbeddingStore",46        back_populates="collection",47        passive_deletes=True,48    )49 50    @classmethod51    def get_by_name(cls, session: Session, name: str) -> Optional["CollectionStore"]:52        return session.query(cls).filter(cls.name == name).first()53 54    @classmethod55    def get_or_create(56        cls,57        session: Session,58        name: str,59        cmetadata: Optional[dict] = None,60    ) -> Tuple["CollectionStore", bool]:61        """62        Get or create a collection.63        Returns [Collection, bool] where the bool is True if the collection was created.64        """65        created = False66        collection = cls.get_by_name(session, name)67        if collection:68            return collection, created69 70        collection = cls(name=name, cmetadata=cmetadata)71        session.add(collection)72        session.commit()73        created = True74        return collection, created75 76 77class EmbeddingStore(BaseModel):78    """Embedding store."""79 80    __tablename__ = "langchain_pg_embedding"81 82    collection_id = sqlalchemy.Column(83        UUID(as_uuid=True),84        sqlalchemy.ForeignKey(85            f"{CollectionStore.__tablename__}.uuid",86            ondelete="CASCADE",87        ),88    )89    collection = relationship(CollectionStore, back_populates="embeddings")90 91    embedding = sqlalchemy.Column(sqlalchemy.ARRAY(sqlalchemy.REAL))  # type: ignore[var-annotated]92    document = sqlalchemy.Column(sqlalchemy.String, nullable=True)93    cmetadata = sqlalchemy.Column(JSON, nullable=True)94 95    # custom_id : any user defined id96    custom_id = sqlalchemy.Column(sqlalchemy.String, nullable=True)97 98 99class QueryResult:100    """Result from a query."""101 102    EmbeddingStore: EmbeddingStore103    distance: float104 105 106class PGEmbedding(VectorStore):107    """`Postgres` with the `pg_embedding` extension as a vector store.108 109    pg_embedding uses sequential scan by default. but you can create a HNSW index110    using the create_hnsw_index method.111    - `connection_string` is a postgres connection string.112    - `embedding_function` any embedding function implementing113        `langchain.embeddings.base.Embeddings` interface.114    - `collection_name` is the name of the collection to use. (default: langchain)115        - NOTE: This is not the name of the table, but the name of the collection.116            The tables will be created when initializing the store (if not exists)117            So, make sure the user has the right permissions to create tables.118    - `distance_strategy` is the distance strategy to use. (default: EUCLIDEAN)119        - `EUCLIDEAN` is the euclidean distance.120    - `pre_delete_collection` if True, will delete the collection if it exists.121        (default: False)122        - Useful for testing.123    """124 125    def __init__(126        self,127        connection_string: str,128        embedding_function: Embeddings,129        collection_name: str = _LANGCHAIN_DEFAULT_COLLECTION_NAME,130        collection_metadata: Optional[dict] = None,131        pre_delete_collection: bool = False,132        logger: Optional[logging.Logger] = None,133    ) -> None:134        self.connection_string = connection_string135        self.embedding_function = embedding_function136        self.collection_name = collection_name137        self.collection_metadata = collection_metadata138        self.pre_delete_collection = pre_delete_collection139        self.logger = logger or logging.getLogger(__name__)140        self.__post_init__()141 142    def __post_init__(143        self,144    ) -> None:145        self._conn = self.connect()146        self.create_hnsw_extension()147        self.create_tables_if_not_exists()148        self.create_collection()149 150    @property151    def embeddings(self) -> Embeddings:152        return self.embedding_function153 154    def connect(self) -> sqlalchemy.engine.Connection:155        engine = sqlalchemy.create_engine(self.connection_string)156        conn = engine.connect()157        return conn158 159    def create_hnsw_extension(self) -> None:160        try:161            with Session(self._conn) as session:162                statement = sqlalchemy.text("CREATE EXTENSION IF NOT EXISTS embedding")163                session.execute(statement)164                session.commit()165        except Exception as e:166            self.logger.exception(e)167 168    def create_tables_if_not_exists(self) -> None:169        with self._conn.begin():170            Base.metadata.create_all(self._conn)171 172    def drop_tables(self) -> None:173        with self._conn.begin():174            Base.metadata.drop_all(self._conn)175 176    def create_collection(self) -> None:177        if self.pre_delete_collection:178            self.delete_collection()179        with Session(self._conn) as session:180            CollectionStore.get_or_create(181                session, self.collection_name, cmetadata=self.collection_metadata182            )183 184    def create_hnsw_index(185        self,186        max_elements: int = 10000,187        dims: int = ADA_TOKEN_COUNT,188        m: int = 8,189        ef_construction: int = 16,190        ef_search: int = 16,191    ) -> None:192        create_index_query = sqlalchemy.text(193            "CREATE INDEX IF NOT EXISTS langchain_pg_embedding_idx "194            "ON langchain_pg_embedding USING hnsw (embedding) "195            "WITH ("196            "maxelements = {}, "197            "dims = {}, "198            "m = {}, "199            "efconstruction = {}, "200            "efsearch = {}"201            ");".format(max_elements, dims, m, ef_construction, ef_search)202        )203 204        # Execute the queries205        try:206            with Session(self._conn) as session:207                # Create the HNSW index208                session.execute(create_index_query)209                session.commit()210            print("HNSW extension and index created successfully.")  # noqa: T201211        except Exception as e:212            print(f"Failed to create HNSW extension or index: {e}")  # noqa: T201213 214    def delete_collection(self) -> None:215        self.logger.debug("Trying to delete collection")216        with Session(self._conn) as session:217            collection = self.get_collection(session)218            if not collection:219                self.logger.warning("Collection not found")220                return221            session.delete(collection)222            session.commit()223 224    def get_collection(self, session: Session) -> Optional["CollectionStore"]:225        return CollectionStore.get_by_name(session, self.collection_name)226 227    @classmethod228    def _initialize_from_embeddings(229        cls,230        texts: List[str],231        embeddings: List[List[float]],232        embedding: Embeddings,233        metadatas: Optional[List[dict]] = None,234        ids: Optional[List[str]] = None,235        collection_name: str = _LANGCHAIN_DEFAULT_COLLECTION_NAME,236        pre_delete_collection: bool = False,237        **kwargs: Any,238    ) -> PGEmbedding:239        if ids is None:240            ids = [str(uuid.uuid4()) for _ in texts]241 242        if not metadatas:243            metadatas = [{} for _ in texts]244 245        connection_string = cls.get_connection_string(kwargs)246 247        store = cls(248            connection_string=connection_string,249            collection_name=collection_name,250            embedding_function=embedding,251            pre_delete_collection=pre_delete_collection,252        )253 254        store.add_embeddings(255            texts=texts, embeddings=embeddings, metadatas=metadatas, ids=ids, **kwargs256        )257 258        return store259 260    def add_embeddings(261        self,262        texts: List[str],263        embeddings: List[List[float]],264        metadatas: List[dict],265        ids: List[str],266        **kwargs: Any,267    ) -> None:268        with Session(self._conn) as session:269            collection = self.get_collection(session)270            if not collection:271                raise ValueError("Collection not found")272            for text, metadata, embedding, id in zip(texts, metadatas, embeddings, ids):273                embedding_store = EmbeddingStore(274                    embedding=embedding,275                    document=text,276                    cmetadata=metadata,277                    custom_id=id,278                )279                collection.embeddings.append(embedding_store)280                session.add(embedding_store)281            session.commit()282 283    def add_texts(284        self,285        texts: Iterable[str],286        metadatas: Optional[List[dict]] = None,287        ids: Optional[List[str]] = None,288        **kwargs: Any,289    ) -> List[str]:290        if ids is None:291            ids = [str(uuid.uuid4()) for _ in texts]292 293        embeddings = self.embedding_function.embed_documents(list(texts))294 295        if not metadatas:296            metadatas = [{} for _ in texts]297 298        with Session(self._conn) as session:299            collection = self.get_collection(session)300            if not collection:301                raise ValueError("Collection not found")302            for text, metadata, embedding, id in zip(texts, metadatas, embeddings, ids):303                embedding_store = EmbeddingStore(304                    embedding=embedding,305                    document=text,306                    cmetadata=metadata,307                    custom_id=id,308                )309                collection.embeddings.append(embedding_store)310                session.add(embedding_store)311            session.commit()312 313        return ids314 315    def similarity_search(316        self,317        query: str,318        k: int = 4,319        filter: Optional[dict] = None,320        **kwargs: Any,321    ) -> List[Document]:322        embedding = self.embedding_function.embed_query(text=query)323        return self.similarity_search_by_vector(324            embedding=embedding,325            k=k,326            filter=filter,327        )328 329    def similarity_search_with_score(330        self,331        query: str,332        k: int = 4,333        filter: Optional[dict] = None,334    ) -> List[Tuple[Document, float]]:335        embedding = self.embedding_function.embed_query(query)336        docs = self.similarity_search_with_score_by_vector(337            embedding=embedding, k=k, filter=filter338        )339        return docs340 341    def similarity_search_with_score_by_vector(342        self,343        embedding: List[float],344        k: int = 4,345        filter: Optional[dict] = None,346    ) -> List[Tuple[Document, float]]:347        with Session(self._conn) as session:348            collection = self.get_collection(session)349            set_enable_seqscan_stmt = sqlalchemy.text("SET enable_seqscan = off")350            session.execute(set_enable_seqscan_stmt)351            if not collection:352                raise ValueError("Collection not found")353 354            filter_by = EmbeddingStore.collection_id == collection.uuid355 356            if filter is not None:357                filter_clauses = []358                for key, value in filter.items():359                    IN = "in"360                    if isinstance(value, dict) and IN in map(str.lower, value):361                        value_case_insensitive = {362                            k.lower(): v for k, v in value.items()363                        }364                        filter_by_metadata = EmbeddingStore.cmetadata[key].astext.in_(365                            value_case_insensitive[IN]366                        )367                        filter_clauses.append(filter_by_metadata)368                    elif isinstance(value, dict) and "substring" in map(369                        str.lower, value370                    ):371                        filter_by_metadata = EmbeddingStore.cmetadata[key].astext.ilike(372                            f"%{value['substring']}%"373                        )374                        filter_clauses.append(filter_by_metadata)375                    else:376                        filter_by_metadata = EmbeddingStore.cmetadata[377                            key378                        ].astext == str(value)379                        filter_clauses.append(filter_by_metadata)380 381                filter_by = sqlalchemy.and_(filter_by, *filter_clauses)382 383            results: List[QueryResult] = (384                session.query(385                    EmbeddingStore,386                    func.abs(EmbeddingStore.embedding.op("<->")(embedding)).label(387                        "distance"388                    ),389                )  # Specify the columns you need here, e.g., EmbeddingStore.embedding390                .filter(filter_by)391                .order_by(392                    func.abs(EmbeddingStore.embedding.op("<->")(embedding)).asc()393                )  # Using PostgreSQL specific operator with the correct column name394                .limit(k)395                .all()396            )397 398        docs = [399            (400                Document(401                    page_content=result.EmbeddingStore.document,  # type: ignore[arg-type]402                    metadata=result.EmbeddingStore.cmetadata,403                ),404                result.distance if self.embedding_function is not None else 0.0,405            )406            for result in results407        ]408        return docs409 410    def similarity_search_by_vector(411        self,412        embedding: List[float],413        k: int = 4,414        filter: Optional[dict] = None,415        **kwargs: Any,416    ) -> List[Document]:417        docs_and_scores = self.similarity_search_with_score_by_vector(418            embedding=embedding, k=k, filter=filter419        )420        return [doc for doc, _ in docs_and_scores]421 422    @classmethod423    def from_texts(424        cls: Type[PGEmbedding],425        texts: List[str],426        embedding: Embeddings,427        metadatas: Optional[List[dict]] = None,428        collection_name: str = _LANGCHAIN_DEFAULT_COLLECTION_NAME,429        ids: Optional[List[str]] = None,430        pre_delete_collection: bool = False,431        **kwargs: Any,432    ) -> PGEmbedding:433        embeddings = embedding.embed_documents(list(texts))434 435        return cls._initialize_from_embeddings(436            texts,437            embeddings,438            embedding,439            metadatas=metadatas,440            ids=ids,441            collection_name=collection_name,442            pre_delete_collection=pre_delete_collection,443            **kwargs,444        )445 446    @classmethod447    def from_embeddings(448        cls,449        text_embeddings: List[Tuple[str, List[float]]],450        embedding: Embeddings,451        metadatas: Optional[List[dict]] = None,452        collection_name: str = _LANGCHAIN_DEFAULT_COLLECTION_NAME,453        ids: Optional[List[str]] = None,454        pre_delete_collection: bool = False,455        **kwargs: Any,456    ) -> PGEmbedding:457        texts = [t[0] for t in text_embeddings]458        embeddings = [t[1] for t in text_embeddings]459 460        return cls._initialize_from_embeddings(461            texts,462            embeddings,463            embedding,464            metadatas=metadatas,465            ids=ids,466            collection_name=collection_name,467            pre_delete_collection=pre_delete_collection,468            **kwargs,469        )470 471    @classmethod472    def from_existing_index(473        cls: Type[PGEmbedding],474        embedding: Embeddings,475        collection_name: str = _LANGCHAIN_DEFAULT_COLLECTION_NAME,476        pre_delete_collection: bool = False,477        **kwargs: Any,478    ) -> PGEmbedding:479        connection_string = cls.get_connection_string(kwargs)480 481        store = cls(482            connection_string=connection_string,483            collection_name=collection_name,484            embedding_function=embedding,485            pre_delete_collection=pre_delete_collection,486        )487 488        return store489 490    @classmethod491    def get_connection_string(cls, kwargs: Dict[str, Any]) -> str:492        connection_string: str = get_from_dict_or_env(493            data=kwargs,494            key="connection_string",495            env_key="POSTGRES_CONNECTION_STRING",496        )497 498        if not connection_string:499            raise ValueError(500                "Postgres connection string is required"501                "Either pass it as a parameter"502                "or set the POSTGRES_CONNECTION_STRING environment variable."503            )504 505        return connection_string506 507    @classmethod508    def from_documents(509        cls: Type[PGEmbedding],510        documents: List[Document],511        embedding: Embeddings,512        collection_name: str = _LANGCHAIN_DEFAULT_COLLECTION_NAME,513        ids: Optional[List[str]] = None,514        pre_delete_collection: bool = False,515        **kwargs: Any,516    ) -> PGEmbedding:517        texts = [d.page_content for d in documents]518        metadatas = [d.metadata for d in documents]519        connection_string = cls.get_connection_string(kwargs)520 521        kwargs["connection_string"] = connection_string522 523        return cls.from_texts(524            texts=texts,525            pre_delete_collection=pre_delete_collection,526            embedding=embedding,527            metadatas=metadatas,528            ids=ids,529            collection_name=collection_name,530            **kwargs,531        )532 
codekingpro/portable-devtools · Team Ai