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