codekingpro/portable-devtools
114k
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 