Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
relyt.py519 linesDownload Raw Back to vectorstores
1from __future__ import annotations2 3import logging4import uuid5from typing import Any, Callable, Dict, Iterable, List, Optional, Sequence, Tuple, Type6 7from sqlalchemy import Column, String, Table, create_engine, insert, text8from sqlalchemy.dialects.postgresql import JSON, TEXT9 10try:11    from sqlalchemy.orm import declarative_base12except ImportError:13    from sqlalchemy.ext.declarative import declarative_base14 15from langchain_core.documents import Document16from langchain_core.embeddings import Embeddings17from langchain_core.utils import get_from_dict_or_env18from langchain_core.vectorstores import VectorStore19 20_LANGCHAIN_DEFAULT_EMBEDDING_DIM = 153621_LANGCHAIN_DEFAULT_COLLECTION_NAME = "langchain_document"22 23Base = declarative_base()  # type: Any24 25 26class Relyt(VectorStore):27    """`Relyt` (distributed PostgreSQL) vector store.28 29    Relyt is a distributed full postgresql syntax cloud-native database.30    - `connection_string` is a postgres connection string.31    - `embedding_function` any embedding function implementing32        `langchain.embeddings.base.Embeddings` interface.33    - `collection_name` is the name of the collection to use. (default: langchain)34        - NOTE: This is not the name of the table, but the name of the collection.35            The tables will be created when initializing the store (if not exists)36            So, make sure the user has the right permissions to create tables.37    - `pre_delete_collection` if True, will delete the collection if it exists.38        (default: False)39        - Useful for testing.40 41    """42 43    def __init__(44        self,45        connection_string: str,46        embedding_function: Embeddings,47        embedding_dimension: int = _LANGCHAIN_DEFAULT_EMBEDDING_DIM,48        collection_name: str = _LANGCHAIN_DEFAULT_COLLECTION_NAME,49        pre_delete_collection: bool = False,50        logger: Optional[logging.Logger] = None,51        engine_args: Optional[dict] = None,52    ) -> None:53        """Initialize a PGVecto_rs vectorstore.54 55        Args:56            embedding: Embeddings to use.57            dimension: Dimension of the embeddings.58            db_url: Database URL.59            collection_name: Name of the collection.60            new_table: Whether to create a new table or connect to an existing one.61            If true, the table will be dropped if exists, then recreated.62            Defaults to False.63        """64        try:65            from pgvecto_rs.sdk import PGVectoRs66 67            PGVectoRs(68                db_url=connection_string,69                collection_name=collection_name,70                dimension=embedding_dimension,71                recreate=pre_delete_collection,72            )73        except ImportError as e:74            raise ImportError(75                "Unable to import pgvector_rs.sdk , please install with "76                '`pip install "pgvecto_rs[sdk]"`.'77            ) from e78 79        self.connection_string = connection_string80        self.embedding_function = embedding_function81        self.embedding_dimension = embedding_dimension82        self.collection_name = collection_name83        self.pre_delete_collection = pre_delete_collection84        self.logger = logger or logging.getLogger(__name__)85        self.__post_init__(engine_args)86 87    def __post_init__(88        self,89        engine_args: Optional[dict] = None,90    ) -> None:91        """92        Initialize the store.93        """94 95        _engine_args = engine_args or {}96 97        if (98            "pool_recycle" not in _engine_args99        ):  # Check if pool_recycle is not in _engine_args100            _engine_args["pool_recycle"] = (101                3600  # Set pool_recycle to 3600s if not present102            )103 104        self.engine = create_engine(self.connection_string, **_engine_args)105        self.create_collection()106 107    @property108    def embeddings(self) -> Embeddings:109        return self.embedding_function110 111    def _select_relevance_score_fn(self) -> Callable[[float], float]:112        return self._euclidean_relevance_score_fn113 114    def create_table_if_not_exists(self) -> None:115        # Define the dynamic table116        """117        Table(118            self.collection_name,119            Base.metadata,120            Column("id", TEXT, primary_key=True, default=uuid.uuid4),121            Column("embedding", Vector(self.embedding_dimension)),122            Column("document", String, nullable=True),123            Column("metadata", JSON, nullable=True),124            extend_existing=True,125        )126        """127        with self.engine.connect() as conn:128            with conn.begin():129                # create vectors130                conn.execute(text("CREATE EXTENSION IF NOT EXISTS vectors"))131                conn.execute(text('CREATE EXTENSION IF NOT EXISTS "uuid-ossp"'))132 133                # Create the table134                # Base.metadata.create_all(conn)135                table_name = f"{self.collection_name}"136                table_query = text(137                    f"""138                    SELECT 1139                    FROM pg_class140                    WHERE relname = '{table_name}';141                """142                )143                result = conn.execute(table_query).scalar()144                if not result:145                    table_statement = text(146                        f"""147                            CREATE TABLE {table_name} (148                                id TEXT PRIMARY KEY DEFAULT uuid_generate_v4(),149                                embedding vector({self.embedding_dimension}),150                                document TEXT,151                                metadata JSON152                            ) USING heap;153                        """154                    )155                    conn.execute(table_statement)156 157                # Check if the index exists158                index_name = f"{self.collection_name}_embedding_idx"159                index_query = text(160                    f"""161                    SELECT 1162                    FROM pg_indexes163                    WHERE indexname = '{index_name}';164                """165                )166                result = conn.execute(index_query).scalar()167 168                # Create the index if it doesn't exist169                if not result:170                    index_statement = text(171                        f"""172                        CREATE INDEX {index_name}173                        ON {self.collection_name}174                        USING vectors (embedding vector_l2_ops)175                        WITH (options = $$176                        optimizing.optimizing_threads = 30177                        segment.max_growing_segment_size = 600178                        segment.max_sealed_segment_size = 30000000179                        [indexing.hnsw]180                        m=30181                        ef_construction=500182                        $$);183                    """184                    )185                    conn.execute(index_statement)186 187    def create_collection(self) -> None:188        if self.pre_delete_collection:189            self.delete_collection()190        self.create_table_if_not_exists()191 192    def delete_collection(self) -> None:193        self.logger.debug("Trying to delete collection")194        drop_statement = text(f"DROP TABLE IF EXISTS {self.collection_name};")195        with self.engine.connect() as conn:196            with conn.begin():197                conn.execute(drop_statement)198 199    def add_texts(200        self,201        texts: Iterable[str],202        metadatas: Optional[List[dict]] = None,203        ids: Optional[List[str]] = None,204        batch_size: int = 500,205        **kwargs: Any,206    ) -> List[str]:207        """Run more texts through the embeddings and add to the vectorstore.208 209        Args:210            texts: Iterable of strings to add to the vectorstore.211            metadatas: Optional list of metadatas associated with the texts.212            kwargs: vectorstore specific parameters213 214        Returns:215            List of ids from adding the texts into the vectorstore.216        """217        from pgvecto_rs.sqlalchemy import Vector218 219        if ids is None:220            ids = [str(uuid.uuid1()) for _ in texts]221 222        embeddings = self.embedding_function.embed_documents(list(texts))223 224        if not metadatas:225            metadatas = [{} for _ in texts]226 227        # Define the table schema228        chunks_table = Table(229            self.collection_name,230            Base.metadata,231            Column("id", TEXT, primary_key=True),232            Column("embedding", Vector(self.embedding_dimension)),233            Column("document", String, nullable=True),234            Column("metadata", JSON, nullable=True),235            extend_existing=True,236        )237 238        chunks_table_data = []239        with self.engine.connect() as conn:240            with conn.begin():241                for document, metadata, chunk_id, embedding in zip(242                    texts, metadatas, ids, embeddings243                ):244                    chunks_table_data.append(245                        {246                            "id": chunk_id,247                            "embedding": embedding,248                            "document": document,249                            "metadata": metadata,250                        }251                    )252 253                    # Execute the batch insert when the batch size is reached254                    if len(chunks_table_data) == batch_size:255                        conn.execute(insert(chunks_table).values(chunks_table_data))256                        # Clear the chunks_table_data list for the next batch257                        chunks_table_data.clear()258 259                # Insert any remaining records that didn't make up a full batch260                if chunks_table_data:261                    conn.execute(insert(chunks_table).values(chunks_table_data))262 263        return ids264 265    def similarity_search(266        self,267        query: str,268        k: int = 4,269        filter: Optional[dict] = None,270        **kwargs: Any,271    ) -> List[Document]:272        """Run similarity search with AnalyticDB with distance.273 274        Args:275            query (str): Query text to search for.276            k (int): Number of results to return. Defaults to 4.277            filter (Optional[Dict[str, str]]): Filter by metadata. Defaults to None.278 279        Returns:280            List of Documents most similar to the query.281        """282        embedding = self.embedding_function.embed_query(text=query)283        return self.similarity_search_by_vector(284            embedding=embedding,285            k=k,286            filter=filter,287        )288 289    def similarity_search_with_score(290        self,291        query: str,292        k: int = 4,293        filter: Optional[dict] = None,294    ) -> List[Tuple[Document, float]]:295        """Return docs most similar to query.296 297        Args:298            query: Text to look up documents similar to.299            k: Number of Documents to return. Defaults to 4.300            filter (Optional[Dict[str, str]]): Filter by metadata. Defaults to None.301 302        Returns:303            List of Documents most similar to the query and score for each304        """305        embedding = self.embedding_function.embed_query(query)306        docs = self.similarity_search_with_score_by_vector(307            embedding=embedding, k=k, filter=filter308        )309        return docs310 311    def similarity_search_with_score_by_vector(312        self,313        embedding: List[float],314        k: int = 4,315        filter: Optional[dict] = None,316    ) -> List[Tuple[Document, float]]:317        # Add the filter if provided318        try:319            from sqlalchemy.engine import Row320        except ImportError:321            raise ImportError(322                "Could not import Row from sqlalchemy.engine. "323                "Please 'pip install sqlalchemy>=1.4'."324            )325 326        filter_condition = ""327        if filter is not None:328            conditions = [329                f"metadata->>{key!r} = {value!r}" for key, value in filter.items()330            ]331            filter_condition = f"WHERE {' AND '.join(conditions)}"332 333        # Define the base query334        sql_query = f"""335            set vectors.enable_search_growing = on;336            set vectors.enable_search_write = on;337            SELECT document, metadata, embedding <-> :embedding as distance338            FROM {self.collection_name}339            {filter_condition}340            ORDER BY embedding <-> :embedding341            LIMIT :k342        """343 344        # Set up the query parameters345        embedding_str = ", ".join(format(x) for x in embedding)346        embedding_str = "[" + embedding_str + "]"347        params = {"embedding": embedding_str, "k": k}348 349        # Execute the query and fetch the results350        with self.engine.connect() as conn:351            results: Sequence[Row] = conn.execute(text(sql_query), params).fetchall()352 353        documents_with_scores = [354            (355                Document(356                    page_content=result.document,357                    metadata=result.metadata,358                ),359                result.distance if self.embedding_function is not None else None,360            )361            for result in results362        ]363        return documents_with_scores364 365    def similarity_search_by_vector(366        self,367        embedding: List[float],368        k: int = 4,369        filter: Optional[dict] = None,370        **kwargs: Any,371    ) -> List[Document]:372        """Return docs most similar to embedding vector.373 374        Args:375            embedding: Embedding to look up documents similar to.376            k: Number of Documents to return. Defaults to 4.377            filter (Optional[Dict[str, str]]): Filter by metadata. Defaults to None.378 379        Returns:380            List of Documents most similar to the query vector.381        """382        docs_and_scores = self.similarity_search_with_score_by_vector(383            embedding=embedding, k=k, filter=filter384        )385        return [doc for doc, _ in docs_and_scores]386 387    def delete(self, ids: Optional[List[str]] = None, **kwargs: Any) -> Optional[bool]:388        """Delete by vector IDs.389 390        Args:391            ids: List of ids to delete.392        """393        from pgvecto_rs.sqlalchemy import Vector394 395        if ids is None:396            raise ValueError("No ids provided to delete.")397 398        # Define the table schema399        chunks_table = Table(400            self.collection_name,401            Base.metadata,402            Column("id", TEXT, primary_key=True),403            Column("embedding", Vector(self.embedding_dimension)),404            Column("document", String, nullable=True),405            Column("metadata", JSON, nullable=True),406            extend_existing=True,407        )408 409        try:410            with self.engine.connect() as conn:411                with conn.begin():412                    delete_condition = chunks_table.c.id.in_(ids)413                    conn.execute(chunks_table.delete().where(delete_condition))414                    return True415        except Exception as e:416            print("Delete operation failed:", str(e))  # noqa: T201417            return False418 419    @classmethod420    def from_texts(421        cls: Type[Relyt],422        texts: List[str],423        embedding: Embeddings,424        metadatas: Optional[List[dict]] = None,425        embedding_dimension: int = _LANGCHAIN_DEFAULT_EMBEDDING_DIM,426        collection_name: str = _LANGCHAIN_DEFAULT_COLLECTION_NAME,427        ids: Optional[List[str]] = None,428        pre_delete_collection: bool = False,429        engine_args: Optional[dict] = None,430        **kwargs: Any,431    ) -> Relyt:432        """433        Return VectorStore initialized from texts and embeddings.434        Postgres Connection string is required435        Either pass it as a parameter436        or set the PG_CONNECTION_STRING environment variable.437        """438 439        connection_string = cls.get_connection_string(kwargs)440 441        store = cls(442            connection_string=connection_string,443            collection_name=collection_name,444            embedding_function=embedding,445            embedding_dimension=embedding_dimension,446            pre_delete_collection=pre_delete_collection,447            engine_args=engine_args,448        )449 450        store.add_texts(texts=texts, metadatas=metadatas, ids=ids, **kwargs)451        return store452 453    @classmethod454    def get_connection_string(cls, kwargs: Dict[str, Any]) -> str:455        connection_string: str = get_from_dict_or_env(456            data=kwargs,457            key="connection_string",458            env_key="PG_CONNECTION_STRING",459        )460 461        if not connection_string:462            raise ValueError(463                "Postgres connection string is required"464                "Either pass it as a parameter"465                "or set the PG_CONNECTION_STRING environment variable."466            )467 468        return connection_string469 470    @classmethod471    def from_documents(472        cls: Type[Relyt],473        documents: List[Document],474        embedding: Embeddings,475        embedding_dimension: int = _LANGCHAIN_DEFAULT_EMBEDDING_DIM,476        collection_name: str = _LANGCHAIN_DEFAULT_COLLECTION_NAME,477        ids: Optional[List[str]] = None,478        pre_delete_collection: bool = False,479        engine_args: Optional[dict] = None,480        **kwargs: Any,481    ) -> Relyt:482        """483        Return VectorStore initialized from documents and embeddings.484        Postgres Connection string is required485        Either pass it as a parameter486        or set the PG_CONNECTION_STRING environment variable.487        """488 489        texts = [d.page_content for d in documents]490        metadatas = [d.metadata for d in documents]491        connection_string = cls.get_connection_string(kwargs)492 493        kwargs["connection_string"] = connection_string494 495        return cls.from_texts(496            texts=texts,497            pre_delete_collection=pre_delete_collection,498            embedding=embedding,499            embedding_dimension=embedding_dimension,500            metadatas=metadatas,501            ids=ids,502            collection_name=collection_name,503            engine_args=engine_args,504            **kwargs,505        )506 507    @classmethod508    def connection_string_from_db_params(509        cls,510        driver: str,511        host: str,512        port: int,513        database: str,514        user: str,515        password: str,516    ) -> str:517        """Return connection string from database parameters."""518        return f"postgresql+{driver}://{user}:{password}@{host}:{port}/{database}"519 
codekingpro/portable-devtools · Team Ai