Underground-Digital/Workflow-Engine
0
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 