codekingpro/portable-devtools
114k
1from __future__ import annotations2 3import json4import logging5from hashlib import sha16from threading import Thread7from typing import Any, Dict, Iterable, List, Mapping, Optional, Tuple, Union8 9import numpy as np10from langchain_core.documents import Document11from langchain_core.embeddings import Embeddings12from langchain_core.vectorstores import VectorStore13from pydantic_settings import BaseSettings, SettingsConfigDict14from typing_extensions import TypedDict15 16from langchain_community.vectorstores.utils import maximal_marginal_relevance17 18logger = logging.getLogger()19DEBUG = False20 21Metadata = Mapping[str, Union[str, int, float, bool]]22 23 24class QueryResult(TypedDict):25 ids: List[List[str]]26 embeddings: List[Any]27 documents: List[Document]28 metadatas: Optional[List[Metadata]]29 distances: Optional[List[float]]30 31 32class ApacheDorisSettings(BaseSettings):33 """Apache Doris client configuration.34 35 Attributes:36 apache_doris_host (str) : An URL to connect to frontend.37 Defaults to 'localhost'.38 apache_doris_port (int) : URL port to connect with HTTP. Defaults to 9030.39 username (str) : Username to login. Defaults to 'root'.40 password (str) : Password to login. Defaults to None.41 database (str) : Database name to find the table. Defaults to 'default'.42 table (str) : Table name to operate on.43 Defaults to 'langchain'.44 45 column_map (Dict) : Column type map to project column name onto langchain46 semantics. Must have keys: `text`, `id`, `vector`,47 must be same size to number of columns. For example:48 .. code-block:: python49 50 {51 'id': 'text_id',52 'embedding': 'text_embedding',53 'document': 'text_plain',54 'metadata': 'metadata_dictionary_in_json',55 }56 57 Defaults to identity map.58 """59 60 host: str = "localhost"61 port: int = 903062 username: str = "root"63 password: str = ""64 65 column_map: Dict[str, str] = {66 "id": "id",67 "document": "document",68 "embedding": "embedding",69 "metadata": "metadata",70 }71 72 database: str = "default"73 table: str = "langchain"74 75 def __getitem__(self, item: str) -> Any:76 return getattr(self, item)77 78 model_config = SettingsConfigDict(79 env_file=".env",80 env_file_encoding="utf-8",81 env_prefix="apache_doris_",82 extra="ignore",83 )84 85 86class ApacheDoris(VectorStore):87 """`Apache Doris` vector store.88 89 You need a `pymysql` python package, and a valid account90 to connect to Apache Doris.91 92 For more information, please visit93 [Apache Doris official site](https://doris.apache.org/)94 [Apache Doris github](https://github.com/apache/doris)95 """96 97 def __init__(98 self,99 embedding: Embeddings,100 *,101 config: Optional[ApacheDorisSettings] = None,102 **kwargs: Any,103 ) -> None:104 """Constructor for Apache Doris.105 106 Args:107 embedding (Embeddings): Text embedding model.108 config (ApacheDorisSettings): Apache Doris client configuration information.109 """110 try:111 import pymysql # type: ignore[import-untyped]112 except ImportError:113 raise ImportError(114 "Could not import pymysql python package. "115 "Please install it with `pip install pymysql`."116 )117 try:118 from tqdm import tqdm119 120 self.pgbar = tqdm121 except ImportError:122 # Just in case if tqdm is not installed123 self.pgbar = lambda x, **kwargs: x124 super().__init__()125 if config is not None:126 self.config = config127 else:128 self.config = ApacheDorisSettings()129 assert self.config130 assert self.config.host and self.config.port131 assert self.config.column_map and self.config.database and self.config.table132 for k in ["id", "embedding", "document", "metadata"]:133 assert k in self.config.column_map134 135 # initialize the schema136 dim = len(embedding.embed_query("test"))137 138 self.schema = f"""\139CREATE TABLE IF NOT EXISTS {self.config.database}.{self.config.table}( 140 {self.config.column_map["id"]} varchar(50),141 {self.config.column_map["document"]} string,142 {self.config.column_map["embedding"]} array<float>,143 {self.config.column_map["metadata"]} string144) ENGINE = OLAP UNIQUE KEY(id) DISTRIBUTED BY HASH(id) \145 PROPERTIES ("replication_allocation" = "tag.location.default: 1")\146"""147 self.dim = dim148 self.BS = "\\"149 self.must_escape = ("\\", "'")150 self._embedding = embedding151 self.dist_order = "DESC"152 _debug_output(self.config)153 154 # Create a connection to Apache Doris155 self.connection = pymysql.connect(156 host=self.config.host,157 port=self.config.port,158 user=self.config.username,159 password=self.config.password,160 database=self.config.database,161 **kwargs,162 )163 164 _debug_output(self.schema)165 _get_named_result(self.connection, self.schema)166 167 def escape_str(self, value: str) -> str:168 return "".join(f"{self.BS}{c}" if c in self.must_escape else c for c in value)169 170 @property171 def embeddings(self) -> Embeddings:172 return self._embedding173 174 def _build_insert_sql(self, transac: Iterable, column_names: Iterable[str]) -> str:175 ks = ",".join(column_names)176 embed_tuple_index = tuple(column_names).index(177 self.config.column_map["embedding"]178 )179 _data = []180 for n in transac:181 n = ",".join(182 [183 (184 f"'{self.escape_str(str(_n))}'"185 if idx != embed_tuple_index186 else f"{str(_n)}"187 )188 for (idx, _n) in enumerate(n)189 ]190 )191 _data.append(f"({n})")192 i_str = f"""193 INSERT INTO194 {self.config.database}.{self.config.table}({ks})195 VALUES196 {",".join(_data)}197 """198 return i_str199 200 def _insert(self, transac: Iterable, column_names: Iterable[str]) -> None:201 _insert_query = self._build_insert_sql(transac, column_names)202 _debug_output(_insert_query)203 _get_named_result(self.connection, _insert_query)204 205 def add_texts(206 self,207 texts: Iterable[str],208 metadatas: Optional[List[dict]] = None,209 batch_size: int = 32,210 ids: Optional[Iterable[str]] = None,211 **kwargs: Any,212 ) -> List[str]:213 """Insert more texts through the embeddings and add to the VectorStore.214 215 Args:216 texts: Iterable of strings to add to the VectorStore.217 ids: Optional list of ids to associate with the texts.218 batch_size: Batch size of insertion219 metadata: Optional column data to be inserted220 221 Returns:222 List of ids from adding the texts into the VectorStore.223 224 """225 # Embed and create the documents226 ids = ids or [sha1(t.encode("utf-8")).hexdigest() for t in texts]227 colmap_ = self.config.column_map228 transac = []229 column_names = {230 colmap_["id"]: ids,231 colmap_["document"]: texts,232 colmap_["embedding"]: self._embedding.embed_documents(list(texts)),233 }234 metadatas = metadatas or [{} for _ in texts]235 column_names[colmap_["metadata"]] = map(json.dumps, metadatas)236 assert len(set(colmap_) - set(column_names)) >= 0237 keys, values = zip(*column_names.items())238 try:239 t = None240 for v in self.pgbar(241 zip(*values), desc="Inserting data...", total=len(metadatas)242 ):243 assert (244 len(v[keys.index(self.config.column_map["embedding"])]) == self.dim245 )246 transac.append(v)247 if len(transac) == batch_size:248 if t:249 t.join()250 t = Thread(target=self._insert, args=[transac, keys])251 t.start()252 transac = []253 if len(transac) > 0:254 if t:255 t.join()256 self._insert(transac, keys)257 return [i for i in ids]258 except Exception as e:259 logger.error(f"\033[91m\033[1m{type(e)}\033[0m \033[95m{str(e)}\033[0m")260 return []261 262 @classmethod263 def from_texts(264 cls,265 texts: List[str],266 embedding: Embeddings,267 metadatas: Optional[List[Dict[Any, Any]]] = None,268 config: Optional[ApacheDorisSettings] = None,269 text_ids: Optional[Iterable[str]] = None,270 batch_size: int = 32,271 **kwargs: Any,272 ) -> ApacheDoris:273 """Create Apache Doris wrapper with existing texts274 275 Args:276 embedding_function (Embeddings): Function to extract text embedding277 texts (Iterable[str]): List or tuple of strings to be added278 config (ApacheDorisSettings, Optional): Apache Doris configuration279 text_ids (Optional[Iterable], optional): IDs for the texts.280 Defaults to None.281 batch_size (int, optional): BatchSize when transmitting data to Apache282 Doris. Defaults to 32.283 metadata (List[dict], optional): metadata to texts. Defaults to None.284 Returns:285 Apache Doris Index286 """287 ctx = cls(embedding, config=config, **kwargs)288 ctx.add_texts(texts, ids=text_ids, batch_size=batch_size, metadatas=metadatas)289 return ctx290 291 def __repr__(self) -> str:292 """Text representation for Apache Doris Vector Store, prints frontends, username293 and schemas. Easy to use with `str(ApacheDoris())`294 295 Returns:296 repr: string to show connection info and data schema297 """298 _repr = f"\033[92m\033[1m{self.config.database}.{self.config.table} @ "299 _repr += f"{self.config.host}:{self.config.port}\033[0m\n\n"300 _repr += f"\033[1musername: {self.config.username}\033[0m\n\nTable Schema:\n"301 width = 25302 fields = 3303 _repr += "-" * (width * fields + 1) + "\n"304 columns = ["name", "type", "key"]305 _repr += f"|\033[94m{columns[0]:24s}\033[0m|\033[96m{columns[1]:24s}"306 _repr += f"\033[0m|\033[96m{columns[2]:24s}\033[0m|\n"307 _repr += "-" * (width * fields + 1) + "\n"308 q_str = f"DESC {self.config.database}.{self.config.table}"309 _debug_output(q_str)310 rs = _get_named_result(self.connection, q_str)311 for r in rs:312 _repr += f"|\033[94m{r['Field']:24s}\033[0m|\033[96m{r['Type']:24s}"313 _repr += f"\033[0m|\033[96m{r['Key']:24s}\033[0m|\n"314 _repr += "-" * (width * fields + 1) + "\n"315 return _repr316 317 def _build_query_sql(318 self, q_emb: List[float], topk: int, where_str: Optional[str] = None319 ) -> str:320 q_emb_str = ",".join(map(str, q_emb))321 if where_str:322 where_str = f"WHERE {where_str}"323 else:324 where_str = ""325 326 q_str = f"""327 SELECT 328 id as id,329 {self.config.column_map["document"]} as document, 330 {self.config.column_map["metadata"]} as metadata, 331 cosine_distance(array<float>[{q_emb_str}],332 {self.config.column_map["embedding"]}) as dist,333 {self.config.column_map["embedding"]} as embedding334 FROM {self.config.database}.{self.config.table}335 {where_str}336 ORDER BY dist {self.dist_order}337 LIMIT {topk}338 """339 340 _debug_output(q_str)341 return q_str342 343 def similarity_search(344 self, query: str, k: int = 4, where_str: Optional[str] = None, **kwargs: Any345 ) -> List[Document]:346 """Perform a similarity search with Apache Doris347 348 Args:349 query (str): query string350 k (int, optional): Top K neighbors to retrieve. Defaults to 4.351 where_str (Optional[str], optional): where condition string.352 Defaults to None.353 354 NOTE: Please do not let end-user to fill this and always be aware355 of SQL injection. When dealing with metadatas, remember to356 use `{self.metadata_column}.attribute` instead of `attribute`357 alone. The default name for it is `metadata`.358 359 Returns:360 List[Document]: List of Documents361 """362 return self.similarity_search_by_vector(363 self._embedding.embed_query(query), k, where_str, **kwargs364 )365 366 def similarity_search_by_vector(367 self,368 embedding: List[float],369 k: int = 4,370 where_str: Optional[str] = None,371 **kwargs: Any,372 ) -> List[Document]:373 """Perform a similarity search with Apache Doris by vectors374 375 Args:376 query (str): query string377 k (int, optional): Top K neighbors to retrieve. Defaults to 4.378 where_str (Optional[str], optional): where condition string.379 Defaults to None.380 381 NOTE: Please do not let end-user to fill this and always be aware382 of SQL injection. When dealing with metadatas, remember to383 use `{self.metadata_column}.attribute` instead of `attribute`384 alone. The default name for it is `metadata`.385 386 Returns:387 List[Document]: List of (Document, similarity)388 """389 q_str = self._build_query_sql(embedding, k, where_str)390 try:391 q_r = _get_named_result(self.connection, q_str)392 return [393 Document(394 page_content=r[self.config.column_map["document"]],395 metadata=json.loads(r[self.config.column_map["metadata"]]),396 )397 for r in q_r398 ]399 except Exception as e:400 logger.error(f"\033[91m\033[1m{type(e)}\033[0m \033[95m{str(e)}\033[0m")401 return []402 403 def similarity_search_with_relevance_scores(404 self, query: str, k: int = 4, where_str: Optional[str] = None, **kwargs: Any405 ) -> List[Tuple[Document, float]]:406 """Perform a similarity search with Apache Doris407 408 Args:409 query (str): query string410 k (int, optional): Top K neighbors to retrieve. Defaults to 4.411 where_str (Optional[str], optional): where condition string.412 Defaults to None.413 414 NOTE: Please do not let end-user to fill this and always be aware415 of SQL injection. When dealing with metadatas, remember to416 use `{self.metadata_column}.attribute` instead of `attribute`417 alone. The default name for it is `metadata`.418 419 Returns:420 List[Document]: List of documents421 """422 q_str = self._build_query_sql(self._embedding.embed_query(query), k, where_str)423 try:424 return [425 (426 Document(427 page_content=r[self.config.column_map["document"]],428 metadata=json.loads(r[self.config.column_map["metadata"]]),429 ),430 r["dist"],431 )432 for r in _get_named_result(self.connection, q_str)433 ]434 except Exception as e:435 logger.error(f"\033[91m\033[1m{type(e)}\033[0m \033[95m{str(e)}\033[0m")436 return []437 438 def drop(self) -> None:439 """440 Helper function: Drop data441 """442 _get_named_result(443 self.connection,444 f"DROP TABLE IF EXISTS {self.config.database}.{self.config.table}",445 )446 447 @property448 def metadata_column(self) -> str:449 return self.config.column_map["metadata"]450 451 def max_marginal_relevance_search_by_vector(452 self,453 embedding: list[float],454 k: int = 4,455 fetch_k: int = 20,456 lambda_mult: float = 0.5,457 **kwargs: Any,458 ) -> list[Document]:459 q_str = self._build_query_sql(embedding, fetch_k, None)460 q_r = _get_named_result(self.connection, q_str)461 results = QueryResult(462 ids=[r["id"] for r in q_r],463 embeddings=[464 json.loads(r[self.config.column_map["embedding"]]) for r in q_r465 ],466 documents=[r[self.config.column_map["document"]] for r in q_r],467 metadatas=[json.loads(r[self.config.column_map["metadata"]]) for r in q_r],468 distances=[r["dist"] for r in q_r],469 )470 471 mmr_selected = maximal_marginal_relevance(472 np.array(embedding, dtype=np.float32),473 results["embeddings"],474 k=k,475 lambda_mult=lambda_mult,476 )477 478 candidates = _results_to_docs(results)479 480 selected_results = [r for i, r in enumerate(candidates) if i in mmr_selected]481 return selected_results482 483 def max_marginal_relevance_search(484 self,485 query: str,486 k: int = 5,487 fetch_k: int = 20,488 lambda_mult: float = 0.5,489 filter: Optional[Dict[str, str]] = None,490 where_document: Optional[Dict[str, str]] = None,491 **kwargs: Any,492 ) -> List[Document]:493 if self.embeddings is None:494 raise ValueError(495 "For MMR search, you must specify an embedding function oncreation."496 )497 498 embedding = self.embeddings.embed_query(query)499 return self.max_marginal_relevance_search_by_vector(500 embedding,501 k,502 fetch_k,503 lambda_mult=lambda_mult,504 filter=filter,505 where_document=where_document,506 )507 508 509def _has_mul_sub_str(s: str, *args: Any) -> bool:510 """Check if a string has multiple substrings.511 512 Args:513 s: The string to check514 *args: The substrings to check for in the string515 516 Returns:517 bool: True if all substrings are present in the string, False otherwise518 """519 for a in args:520 if a not in s:521 return False522 return True523 524 525def _debug_output(s: Any) -> None:526 """Print a debug message if DEBUG is True.527 528 Args:529 s: The message to print530 """531 if DEBUG:532 print(s) # noqa: T201533 534 535def _get_named_result(connection: Any, query: str) -> List[dict[str, Any]]:536 """Get a named result from a query.537 538 Args:539 connection: The connection to the database540 query: The query to execute541 542 Returns:543 List[dict[str, Any]]: The result of the query544 """545 cursor = connection.cursor()546 cursor.execute(query)547 columns = cursor.description548 result = []549 for value in cursor.fetchall():550 r = {}551 for idx, datum in enumerate(value):552 k = columns[idx][0]553 r[k] = datum554 result.append(r)555 _debug_output(result)556 cursor.close()557 return result558 559 560def _results_to_docs(results: Any) -> List[Document]:561 return [doc for doc, _ in _results_to_docs_and_scores(results)]562 563 564def _results_to_docs_and_scores(results: Any) -> List[Tuple[Document, float]]:565 return [566 (Document(page_content=result[0], metadata=result[1] or {}), result[2])567 for result in zip(568 results["documents"],569 results["metadatas"],570 results["distances"],571 )572 ]573 