Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
epsilla.py379 linesDownload Raw Back to vectorstores
1"""Wrapper around Epsilla vector database."""2 3from __future__ import annotations4 5import logging6import uuid7from typing import TYPE_CHECKING, Any, Iterable, List, Optional, Type8 9from langchain_core.documents import Document10from langchain_core.embeddings import Embeddings11from langchain_core.vectorstores import VectorStore12 13if TYPE_CHECKING:14    from pyepsilla import vectordb15 16logger = logging.getLogger()17 18 19class Epsilla(VectorStore):20    """21    Wrapper around Epsilla vector database.22 23    As a prerequisite, you need to install ``pyepsilla`` package24    and have a running Epsilla vector database (for example, through our docker image)25    See the following documentation for how to run an Epsilla vector database:26    https://epsilla-inc.gitbook.io/epsilladb/quick-start27 28    Args:29        client (Any): Epsilla client to connect to.30        embeddings (Embeddings): Function used to embed the texts.31        db_path (Optional[str]): The path where the database will be persisted.32                                 Defaults to "/tmp/langchain-epsilla".33        db_name (Optional[str]): Give a name to the loaded database.34                                 Defaults to "langchain_store".35    Example:36        .. code-block:: python37 38            from langchain_community.vectorstores import Epsilla39            from pyepsilla import vectordb40 41            client = vectordb.Client()42            embeddings = OpenAIEmbeddings()43            db_path = "/tmp/vectorstore"44            db_name = "langchain_store"45            epsilla = Epsilla(client, embeddings, db_path, db_name)46    """47 48    _LANGCHAIN_DEFAULT_DB_NAME: str = "langchain_store"49    _LANGCHAIN_DEFAULT_DB_PATH: str = "/tmp/langchain-epsilla"50    _LANGCHAIN_DEFAULT_TABLE_NAME: str = "langchain_collection"51 52    def __init__(53        self,54        client: Any,55        embeddings: Embeddings,56        db_path: Optional[str] = _LANGCHAIN_DEFAULT_DB_PATH,57        db_name: Optional[str] = _LANGCHAIN_DEFAULT_DB_NAME,58    ):59        """Initialize with necessary components."""60        try:61            import pyepsilla62        except ImportError as e:63            raise ImportError(64                "Could not import pyepsilla python package. "65                "Please install pyepsilla package with `pip install pyepsilla`."66            ) from e67 68        if not isinstance(69            client, (pyepsilla.vectordb.Client, pyepsilla.cloud.client.Vectordb)70        ):71            raise TypeError(72                "client should be an instance of pyepsilla.vectordb.Client or "73                f"pyepsilla.cloud.client.Vectordb, got {type(client)}"74            )75 76        self._client: vectordb.Client = client77        self._db_name = db_name78        self._embeddings = embeddings79        self._collection_name = Epsilla._LANGCHAIN_DEFAULT_TABLE_NAME80        self._client.load_db(db_name=db_name, db_path=db_path)81        self._client.use_db(db_name=db_name)82 83    @property84    def embeddings(self) -> Optional[Embeddings]:85        return self._embeddings86 87    def use_collection(self, collection_name: str) -> None:88        """89        Set default collection to use.90 91        Args:92            collection_name (str): The name of the collection.93        """94        self._collection_name = collection_name95 96    def clear_data(self, collection_name: str = "") -> None:97        """98        Clear data in a collection.99 100        Args:101            collection_name (Optional[str]): The name of the collection.102                If not provided, the default collection will be used.103        """104        if not collection_name:105            collection_name = self._collection_name106        self._client.drop_table(collection_name)107 108    def get(109        self, collection_name: str = "", response_fields: Optional[List[str]] = None110    ) -> List[dict]:111        """Get the collection.112 113        Args:114            collection_name (Optional[str]): The name of the collection115                to retrieve data from.116                If not provided, the default collection will be used.117            response_fields (Optional[List[str]]): List of field names in the result.118                If not specified, all available fields will be responded.119 120        Returns:121            A list of the retrieved data.122        """123        if not collection_name:124            collection_name = self._collection_name125        status_code, response = self._client.get(126            table_name=collection_name, response_fields=response_fields127        )128        if status_code != 200:129            logger.error(f"Failed to get records: {response['message']}")130            raise Exception("Error: {}.".format(response["message"]))131        return response["result"]132 133    def _create_collection(134        self, table_name: str, embeddings: list, metadatas: Optional[list[dict]] = None135    ) -> None:136        if not embeddings:137            raise ValueError("Embeddings list is empty.")138 139        dim = len(embeddings[0])140        fields: List[dict] = [141            {"name": "id", "dataType": "INT"},142            {"name": "text", "dataType": "STRING"},143            {"name": "embeddings", "dataType": "VECTOR_FLOAT", "dimensions": dim},144        ]145        if metadatas is not None:146            field_names = [field["name"] for field in fields]147            for metadata in metadatas:148                for key, value in metadata.items():149                    if key in field_names:150                        continue151                    d_type: str152                    if isinstance(value, str):153                        d_type = "STRING"154                    elif isinstance(value, int):155                        d_type = "INT"156                    elif isinstance(value, float):157                        d_type = "FLOAT"158                    elif isinstance(value, bool):159                        d_type = "BOOL"160                    else:161                        raise ValueError(f"Unsupported data type for {key}.")162                    fields.append({"name": key, "dataType": d_type})163                    field_names.append(key)164 165        status_code, response = self._client.create_table(166            table_name, table_fields=fields167        )168        if status_code != 200:169            if status_code == 409:170                logger.info(f"Continuing with the existing table {table_name}.")171            else:172                logger.error(173                    f"Failed to create collection {table_name}: {response['message']}"174                )175                raise Exception("Error: {}.".format(response["message"]))176 177    def add_texts(178        self,179        texts: Iterable[str],180        metadatas: Optional[List[dict]] = None,181        collection_name: Optional[str] = "",182        drop_old: Optional[bool] = False,183        **kwargs: Any,184    ) -> List[str]:185        """186        Embed texts and add them to the database.187 188        Args:189            texts (Iterable[str]): The texts to embed.190            metadatas (Optional[List[dict]]): Metadata dicts191                        attached to each of the texts. Defaults to None.192            collection_name (Optional[str]): Which collection to use.193                        Defaults to "langchain_collection".194                        If provided, default collection name will be set as well.195            drop_old (Optional[bool]): Whether to drop the previous collection196                        and create a new one. Defaults to False.197 198        Returns:199            List of ids of the added texts.200        """201        if not collection_name:202            collection_name = self._collection_name203        else:204            self._collection_name = collection_name205 206        if drop_old:207            self._client.drop_db(db_name=collection_name)208 209        texts = list(texts)210        try:211            embeddings = self._embeddings.embed_documents(texts)212        except NotImplementedError:213            embeddings = [self._embeddings.embed_query(x) for x in texts]214 215        if len(embeddings) == 0:216            logger.debug("Nothing to insert, skipping.")217            return []218 219        self._create_collection(220            table_name=collection_name, embeddings=embeddings, metadatas=metadatas221        )222 223        ids = [hash(uuid.uuid4()) for _ in texts]224        records = []225        for index, id in enumerate(ids):226            record = {227                "id": id,228                "text": texts[index],229                "embeddings": embeddings[index],230            }231            if metadatas is not None:232                metadata = metadatas[index].items()233                for key, value in metadata:234                    record[key] = value235            records.append(record)236 237        status_code, response = self._client.insert(238            table_name=collection_name, records=records239        )240        if status_code != 200:241            logger.error(242                f"Failed to add records to {collection_name}: {response['message']}"243            )244            raise Exception("Error: {}.".format(response["message"]))245        return [str(id) for id in ids]246 247    def similarity_search(248        self, query: str, k: int = 4, collection_name: str = "", **kwargs: Any249    ) -> List[Document]:250        """251        Return the documents that are semantically most relevant to the query.252 253        Args:254            query (str): String to query the vectorstore with.255            k (Optional[int]): Number of documents to return. Defaults to 4.256            collection_name (Optional[str]): Collection to use.257                Defaults to "langchain_store" or the one provided before.258        Returns:259            List of documents that are semantically most relevant to the query260        """261        if not collection_name:262            collection_name = self._collection_name263        query_vector = self._embeddings.embed_query(query)264        status_code, response = self._client.query(265            table_name=collection_name,266            query_field="embeddings",267            query_vector=query_vector,268            limit=k,269        )270        if status_code != 200:271            logger.error(f"Search failed: {response['message']}.")272            raise Exception("Error: {}.".format(response["message"]))273 274        exclude_keys = ["id", "text", "embeddings"]275        return list(276            map(277                lambda item: Document(278                    page_content=item["text"],279                    metadata={280                        key: item[key] for key in item if key not in exclude_keys281                    },282                ),283                response["result"],284            )285        )286 287    @classmethod288    def from_texts(289        cls: Type[Epsilla],290        texts: List[str],291        embedding: Embeddings,292        metadatas: Optional[List[dict]] = None,293        client: Any = None,294        db_path: Optional[str] = _LANGCHAIN_DEFAULT_DB_PATH,295        db_name: Optional[str] = _LANGCHAIN_DEFAULT_DB_NAME,296        collection_name: Optional[str] = _LANGCHAIN_DEFAULT_TABLE_NAME,297        drop_old: Optional[bool] = False,298        **kwargs: Any,299    ) -> Epsilla:300        """Create an Epsilla vectorstore from raw documents.301 302        Args:303            texts (List[str]): List of text data to be inserted.304            embeddings (Embeddings): Embedding function.305            client (pyepsilla.vectordb.Client): Epsilla client to connect to.306            metadatas (Optional[List[dict]]): Metadata for each text.307                    Defaults to None.308            db_path (Optional[str]): The path where the database will be persisted.309                    Defaults to "/tmp/langchain-epsilla".310            db_name (Optional[str]): Give a name to the loaded database.311                    Defaults to "langchain_store".312            collection_name (Optional[str]): Which collection to use.313                    Defaults to "langchain_collection".314                    If provided, default collection name will be set as well.315            drop_old (Optional[bool]): Whether to drop the previous collection316                    and create a new one. Defaults to False.317 318        Returns:319            Epsilla: Epsilla vector store.320        """321        instance = Epsilla(client, embedding, db_path=db_path, db_name=db_name)322        instance.add_texts(323            texts,324            metadatas=metadatas,325            collection_name=collection_name,326            drop_old=drop_old,327            **kwargs,328        )329 330        return instance331 332    @classmethod333    def from_documents(334        cls: Type[Epsilla],335        documents: List[Document],336        embedding: Embeddings,337        client: Any = None,338        db_path: Optional[str] = _LANGCHAIN_DEFAULT_DB_PATH,339        db_name: Optional[str] = _LANGCHAIN_DEFAULT_DB_NAME,340        collection_name: Optional[str] = _LANGCHAIN_DEFAULT_TABLE_NAME,341        drop_old: Optional[bool] = False,342        **kwargs: Any,343    ) -> Epsilla:344        """Create an Epsilla vectorstore from a list of documents.345 346        Args:347            texts (List[str]): List of text data to be inserted.348            embeddings (Embeddings): Embedding function.349            client (pyepsilla.vectordb.Client): Epsilla client to connect to.350            metadatas (Optional[List[dict]]): Metadata for each text.351                    Defaults to None.352            db_path (Optional[str]): The path where the database will be persisted.353                    Defaults to "/tmp/langchain-epsilla".354            db_name (Optional[str]): Give a name to the loaded database.355                    Defaults to "langchain_store".356            collection_name (Optional[str]): Which collection to use.357                    Defaults to "langchain_collection".358                    If provided, default collection name will be set as well.359            drop_old (Optional[bool]): Whether to drop the previous collection360                    and create a new one. Defaults to False.361 362        Returns:363            Epsilla: Epsilla vector store.364        """365        texts = [doc.page_content for doc in documents]366        metadatas = [doc.metadata for doc in documents]367 368        return cls.from_texts(369            texts,370            embedding,371            metadatas=metadatas,372            client=client,373            db_path=db_path,374            db_name=db_name,375            collection_name=collection_name,376            drop_old=drop_old,377            **kwargs,378        )379 
codekingpro/portable-devtools · Team Ai