Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
apache_doris.py573 linesDownload Raw Back to vectorstores
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 
codekingpro/portable-devtools · Team Ai