Team Ai
Apppublic

Underground-Digital/Workflow-Engine

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
tidb_vector.py254 linesDownload Raw Back to tidb_vector
1import json2import logging3from typing import Any4 5import sqlalchemy6from pydantic import BaseModel, model_validator7from sqlalchemy import JSON, TEXT, Column, DateTime, String, Table, create_engine, insert8from sqlalchemy import text as sql_text9from sqlalchemy.orm import Session, declarative_base10 11from configs import dify_config12from core.rag.datasource.vdb.vector_base import BaseVector13from core.rag.datasource.vdb.vector_factory import AbstractVectorFactory14from core.rag.datasource.vdb.vector_type import VectorType15from core.rag.embedding.embedding_base import Embeddings16from core.rag.models.document import Document17from extensions.ext_redis import redis_client18from models.dataset import Dataset19 20logger = logging.getLogger(__name__)21 22 23class TiDBVectorConfig(BaseModel):24    host: str25    port: int26    user: str27    password: str28    database: str29    program_name: str30 31    @model_validator(mode="before")32    @classmethod33    def validate_config(cls, values: dict) -> dict:34        if not values["host"]:35            raise ValueError("config TIDB_VECTOR_HOST is required")36        if not values["port"]:37            raise ValueError("config TIDB_VECTOR_PORT is required")38        if not values["user"]:39            raise ValueError("config TIDB_VECTOR_USER is required")40        if not values["password"]:41            raise ValueError("config TIDB_VECTOR_PASSWORD is required")42        if not values["database"]:43            raise ValueError("config TIDB_VECTOR_DATABASE is required")44        if not values["program_name"]:45            raise ValueError("config APPLICATION_NAME is required")46        return values47 48 49class TiDBVector(BaseVector):50    def get_type(self) -> str:51        return VectorType.TIDB_VECTOR52 53    def _table(self, dim: int) -> Table:54        from tidb_vector.sqlalchemy import VectorType55 56        return Table(57            self._collection_name,58            self._orm_base.metadata,59            Column("id", String(36), primary_key=True, nullable=False),60            Column(61                "vector",62                VectorType(dim),63                nullable=False,64                comment="" if self._distance_func is None else f"hnsw(distance={self._distance_func})",65            ),66            Column("text", TEXT, nullable=False),67            Column("meta", JSON, nullable=False),68            Column("create_time", DateTime, server_default=sqlalchemy.text("CURRENT_TIMESTAMP")),69            Column(70                "update_time", DateTime, server_default=sqlalchemy.text("CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP")71            ),72            extend_existing=True,73        )74 75    def __init__(self, collection_name: str, config: TiDBVectorConfig, distance_func: str = "cosine"):76        super().__init__(collection_name)77        self._client_config = config78        self._url = (79            f"mysql+pymysql://{config.user}:{config.password}@{config.host}:{config.port}/{config.database}?"80            f"ssl_verify_cert=true&ssl_verify_identity=true&program_name={config.program_name}"81        )82        self._distance_func = distance_func.lower()83        self._engine = create_engine(self._url)84        self._orm_base = declarative_base()85        self._dimension = 153686 87    def create(self, texts: list[Document], embeddings: list[list[float]], **kwargs):88        logger.info("create collection and add texts, collection_name: " + self._collection_name)89        self._create_collection(len(embeddings[0]))90        self.add_texts(texts, embeddings)91        self._dimension = len(embeddings[0])92        pass93 94    def _create_collection(self, dimension: int):95        logger.info("_create_collection, collection_name " + self._collection_name)96        lock_name = "vector_indexing_lock_{}".format(self._collection_name)97        with redis_client.lock(lock_name, timeout=20):98            collection_exist_cache_key = "vector_indexing_{}".format(self._collection_name)99            if redis_client.get(collection_exist_cache_key):100                return101            with Session(self._engine) as session:102                session.begin()103                create_statement = sql_text(f"""104                    CREATE TABLE IF NOT EXISTS {self._collection_name} (105                        id CHAR(36) PRIMARY KEY,106                        text TEXT NOT NULL,107                        meta JSON NOT NULL,108                        doc_id VARCHAR(64) AS (JSON_UNQUOTE(JSON_EXTRACT(meta, '$.doc_id'))) STORED,109                        KEY (doc_id),110                        vector VECTOR<FLOAT>({dimension}) NOT NULL COMMENT "hnsw(distance={self._distance_func})",111                        create_time DATETIME DEFAULT CURRENT_TIMESTAMP,112                        update_time DATETIME DEFAULT CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP113                    );114                """)115                session.execute(create_statement)116                # tidb vector not support 'CREATE/ADD INDEX' now117                session.commit()118            redis_client.set(collection_exist_cache_key, 1, ex=3600)119 120    def add_texts(self, documents: list[Document], embeddings: list[list[float]], **kwargs):121        table = self._table(len(embeddings[0]))122        ids = self._get_uuids(documents)123        metas = [d.metadata for d in documents]124        texts = [d.page_content for d in documents]125 126        chunks_table_data = []127        with self._engine.connect() as conn, conn.begin():128            for id, text, meta, embedding in zip(ids, texts, metas, embeddings):129                chunks_table_data.append({"id": id, "vector": embedding, "text": text, "meta": meta})130 131                # Execute the batch insert when the batch size is reached132                if len(chunks_table_data) == 500:133                    conn.execute(insert(table).values(chunks_table_data))134                    # Clear the chunks_table_data list for the next batch135                    chunks_table_data.clear()136 137            # Insert any remaining records that didn't make up a full batch138            if chunks_table_data:139                conn.execute(insert(table).values(chunks_table_data))140        return ids141 142    def text_exists(self, id: str) -> bool:143        result = self.get_ids_by_metadata_field("doc_id", id)144        return bool(result)145 146    def delete_by_ids(self, ids: list[str]) -> None:147        with Session(self._engine) as session:148            ids_str = ",".join(f"'{doc_id}'" for doc_id in ids)149            select_statement = sql_text(150                f"""SELECT id FROM {self._collection_name} WHERE meta->>'$.doc_id' in ({ids_str}); """151            )152            result = session.execute(select_statement).fetchall()153        if result:154            ids = [item[0] for item in result]155            self._delete_by_ids(ids)156 157    def _delete_by_ids(self, ids: list[str]) -> bool:158        if ids is None:159            raise ValueError("No ids provided to delete.")160        table = self._table(self._dimension)161        try:162            with self._engine.connect() as conn, conn.begin():163                delete_condition = table.c.id.in_(ids)164                conn.execute(table.delete().where(delete_condition))165                return True166        except Exception as e:167            print("Delete operation failed:", str(e))168            return False169 170    def get_ids_by_metadata_field(self, key: str, value: str):171        with Session(self._engine) as session:172            select_statement = sql_text(173                f"""SELECT id FROM {self._collection_name} WHERE meta->>'$.{key}' = '{value}'; """174            )175            result = session.execute(select_statement).fetchall()176        if result:177            return [item[0] for item in result]178        else:179            return None180 181    def delete_by_metadata_field(self, key: str, value: str) -> None:182        ids = self.get_ids_by_metadata_field(key, value)183        if ids:184            self._delete_by_ids(ids)185 186    def search_by_vector(self, query_vector: list[float], **kwargs: Any) -> list[Document]:187        top_k = kwargs.get("top_k", 4)188        score_threshold = float(kwargs.get("score_threshold") or 0.0)189        filter = kwargs.get("filter")190        distance = 1 - score_threshold191 192        query_vector_str = ", ".join(format(x) for x in query_vector)193        query_vector_str = "[" + query_vector_str + "]"194        logger.debug(195            f"_collection_name: {self._collection_name}, score_threshold: {score_threshold}, distance: {distance}"196        )197 198        docs = []199        if self._distance_func == "l2":200            tidb_func = "Vec_l2_distance"201        elif self._distance_func == "cosine":202            tidb_func = "Vec_Cosine_distance"203        else:204            tidb_func = "Vec_Cosine_distance"205 206        with Session(self._engine) as session:207            select_statement = sql_text(208                f"""SELECT meta, text, distance FROM (209                        SELECT meta, text, {tidb_func}(vector, "{query_vector_str}")  as distance210                        FROM {self._collection_name}211                        ORDER BY distance212                        LIMIT {top_k}213                    ) t WHERE distance < {distance};"""214            )215            res = session.execute(select_statement)216            results = [(row[0], row[1], row[2]) for row in res]217            for meta, text, distance in results:218                metadata = json.loads(meta)219                metadata["score"] = 1 - distance220                docs.append(Document(page_content=text, metadata=metadata))221        return docs222 223    def search_by_full_text(self, query: str, **kwargs: Any) -> list[Document]:224        # tidb doesn't support bm25 search225        return []226 227    def delete(self) -> None:228        with Session(self._engine) as session:229            session.execute(sql_text(f"""DROP TABLE IF EXISTS {self._collection_name};"""))230            session.commit()231 232 233class TiDBVectorFactory(AbstractVectorFactory):234    def init_vector(self, dataset: Dataset, attributes: list, embeddings: Embeddings) -> TiDBVector:235        if dataset.index_struct_dict:236            class_prefix: str = dataset.index_struct_dict["vector_store"]["class_prefix"]237            collection_name = class_prefix.lower()238        else:239            dataset_id = dataset.id240            collection_name = Dataset.gen_collection_name_by_id(dataset_id).lower()241            dataset.index_struct = json.dumps(self.gen_index_struct_dict(VectorType.TIDB_VECTOR, collection_name))242 243        return TiDBVector(244            collection_name=collection_name,245            config=TiDBVectorConfig(246                host=dify_config.TIDB_VECTOR_HOST,247                port=dify_config.TIDB_VECTOR_PORT,248                user=dify_config.TIDB_VECTOR_USER,249                password=dify_config.TIDB_VECTOR_PASSWORD,250                database=dify_config.TIDB_VECTOR_DATABASE,251                program_name=dify_config.APPLICATION_NAME,252            ),253        )254