Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
xata.py264 linesDownload Raw Back to vectorstores
1from __future__ import annotations2 3import time4from itertools import repeat5from typing import Any, Dict, Iterable, List, Optional, Tuple, Type6 7from langchain_core.documents import Document8from langchain_core.embeddings import Embeddings9from langchain_core.vectorstores import VectorStore10 11 12class XataVectorStore(VectorStore):13    """`Xata` vector store.14 15    It assumes you have a Xata database16    created with the right schema. See the guide at:17    https://integrations.langchain.com/vectorstores?integration_name=XataVectorStore18 19    """20 21    def __init__(22        self,23        api_key: str,24        db_url: str,25        embedding: Embeddings,26        table_name: str,27    ) -> None:28        """Initialize with Xata client."""29        try:30            from xata.client import XataClient31        except ImportError:32            raise ImportError(33                "Could not import xata python package. "34                "Please install it with `pip install xata`."35            )36        self._client = XataClient(api_key=api_key, db_url=db_url)37        self._embedding: Embeddings = embedding38        self._table_name = table_name or "vectors"39 40    @property41    def embeddings(self) -> Embeddings:42        return self._embedding43 44    def add_vectors(45        self,46        vectors: List[List[float]],47        documents: List[Document],48        ids: Optional[List[str]] = None,49    ) -> List[str]:50        return self._add_vectors(vectors, documents, ids)51 52    def add_texts(53        self,54        texts: Iterable[str],55        metadatas: Optional[List[Dict[Any, Any]]] = None,56        ids: Optional[List[str]] = None,57        **kwargs: Any,58    ) -> List[str]:59        ids = ids60        docs = self._texts_to_documents(texts, metadatas)61 62        vectors = self._embedding.embed_documents(list(texts))63        return self.add_vectors(vectors, docs, ids)64 65    def _add_vectors(66        self,67        vectors: List[List[float]],68        documents: List[Document],69        ids: Optional[List[str]] = None,70    ) -> List[str]:71        """Add vectors to the Xata database."""72 73        rows: List[Dict[str, Any]] = []74        for idx, embedding in enumerate(vectors):75            row = {76                "content": documents[idx].page_content,77                "embedding": embedding,78            }79            if ids:80                row["id"] = ids[idx]81            for key, val in documents[idx].metadata.items():82                if key not in ["id", "content", "embedding"]:83                    row[key] = val84            rows.append(row)85 86        # XXX: I would have liked to use the BulkProcessor here, but it87        # doesn't return the IDs, which we need here. Manual chunking it is.88        chunk_size = 100089        id_list: List[str] = []90        for i in range(0, len(rows), chunk_size):91            chunk = rows[i : i + chunk_size]92 93            r = self._client.records().bulk_insert(self._table_name, {"records": chunk})94            if r.status_code != 200:95                raise Exception(f"Error adding vectors to Xata: {r.status_code} {r}")96            id_list.extend(r["recordIDs"])97        return id_list98 99    @staticmethod100    def _texts_to_documents(101        texts: Iterable[str],102        metadatas: Optional[Iterable[Dict[Any, Any]]] = None,103    ) -> List[Document]:104        """Return list of Documents from list of texts and metadatas."""105        if metadatas is None:106            metadatas = repeat({})107 108        docs = [109            Document(page_content=text, metadata=metadata)110            for text, metadata in zip(texts, metadatas)111        ]112 113        return docs114 115    @classmethod116    def from_texts(117        cls: Type["XataVectorStore"],118        texts: List[str],119        embedding: Embeddings,120        metadatas: Optional[List[dict]] = None,121        api_key: Optional[str] = None,122        db_url: Optional[str] = None,123        table_name: str = "vectors",124        ids: Optional[List[str]] = None,125        **kwargs: Any,126    ) -> "XataVectorStore":127        """Return VectorStore initialized from texts and embeddings."""128 129        if not api_key or not db_url:130            raise ValueError("Xata api_key and db_url must be set.")131 132        embeddings = embedding.embed_documents(texts)133        ids = None  # Xata will generate them for us134        docs = cls._texts_to_documents(texts, metadatas)135 136        vector_db = cls(137            api_key=api_key,138            db_url=db_url,139            embedding=embedding,140            table_name=table_name,141        )142 143        vector_db._add_vectors(embeddings, docs, ids)144        return vector_db145 146    def similarity_search(147        self, query: str, k: int = 4, filter: Optional[dict] = None, **kwargs: Any148    ) -> List[Document]:149        """Return docs most similar to query.150 151        Args:152            query: Text to look up documents similar to.153            k: Number of Documents to return. Defaults to 4.154 155        Returns:156            List of Documents most similar to the query.157        """158        docs_and_scores = self.similarity_search_with_score(query, k, filter=filter)159        documents = [d[0] for d in docs_and_scores]160        return documents161 162    def similarity_search_with_score(163        self, query: str, k: int = 4, filter: Optional[dict] = None, **kwargs: Any164    ) -> List[Tuple[Document, float]]:165        """Run similarity search with Chroma with distance.166 167        Args:168            query (str): Query text to search for.169            k (int): Number of results to return. Defaults to 4.170            filter (Optional[dict]): Filter by metadata. Defaults to None.171 172        Returns:173            List[Tuple[Document, float]]: List of documents most similar to the query174                text with distance in float.175        """176        embedding = self._embedding.embed_query(query)177        payload = {178            "queryVector": embedding,179            "column": "embedding",180            "size": k,181        }182        if filter:183            payload["filter"] = filter184        r = self._client.data().vector_search(self._table_name, payload=payload)185        if r.status_code != 200:186            raise Exception(f"Error running similarity search: {r.status_code} {r}")187        hits = r["records"]188        docs_and_scores = [189            (190                Document(191                    page_content=hit["content"],192                    metadata=self._extractMetadata(hit),193                ),194                hit["xata"]["score"],195            )196            for hit in hits197        ]198        return docs_and_scores199 200    def _extractMetadata(self, record: dict) -> dict:201        """Extract metadata from a record. Filters out known columns."""202        metadata = {}203        for key, val in record.items():204            if key not in ["id", "content", "embedding", "xata"]:205                metadata[key] = val206        return metadata207 208    def delete(209        self,210        ids: Optional[List[str]] = None,211        delete_all: Optional[bool] = None,212        **kwargs: Any,213    ) -> None:214        """Delete by vector IDs.215 216        Args:217            ids: List of ids to delete.218            delete_all: Delete all records in the table.219        """220        if delete_all:221            self._delete_all()222            self.wait_for_indexing(ndocs=0)223        elif ids is not None:224            chunk_size = 500225            for i in range(0, len(ids), chunk_size):226                chunk = ids[i : i + chunk_size]227                operations = [228                    {"delete": {"table": self._table_name, "id": id}} for id in chunk229                ]230                self._client.records().transaction(payload={"operations": operations})231        else:232            raise ValueError("Either ids or delete_all must be set.")233 234    def _delete_all(self) -> None:235        """Delete all records in the table."""236        while True:237            r = self._client.data().query(self._table_name, payload={"columns": ["id"]})238            if r.status_code != 200:239                raise Exception(f"Error running query: {r.status_code} {r}")240            ids = [rec["id"] for rec in r["records"]]241            if len(ids) == 0:242                break243            operations = [244                {"delete": {"table": self._table_name, "id": id}} for id in ids245            ]246            self._client.records().transaction(payload={"operations": operations})247 248    def wait_for_indexing(self, timeout: float = 5, ndocs: int = 1) -> None:249        """Wait for the search index to contain a certain number of250        documents. Useful in tests.251        """252        start = time.time()253        while True:254            r = self._client.data().search_table(255                self._table_name, payload={"query": "", "page": {"size": 0}}256            )257            if r.status_code != 200:258                raise Exception(f"Error running search: {r.status_code} {r}")259            if r["totalCount"] == ndocs:260                break261            if time.time() - start > timeout:262                raise Exception("Timed out waiting for indexing to complete.")263            time.sleep(0.5)264 
codekingpro/portable-devtools · Team Ai