Underground-Digital/Workflow-Engine
0
1import array2import json3import re4import uuid5from contextlib import contextmanager6from typing import Any7 8import jieba.posseg as pseg9import nltk10import numpy11import oracledb12from nltk.corpus import stopwords13from pydantic import BaseModel, model_validator14 15from configs import dify_config16from core.rag.datasource.vdb.vector_base import BaseVector17from core.rag.datasource.vdb.vector_factory import AbstractVectorFactory18from core.rag.datasource.vdb.vector_type import VectorType19from core.rag.embedding.embedding_base import Embeddings20from core.rag.models.document import Document21from extensions.ext_redis import redis_client22from models.dataset import Dataset23 24oracledb.defaults.fetch_lobs = False25 26 27class OracleVectorConfig(BaseModel):28 host: str29 port: int30 user: str31 password: str32 database: str33 34 @model_validator(mode="before")35 @classmethod36 def validate_config(cls, values: dict) -> dict:37 if not values["host"]:38 raise ValueError("config ORACLE_HOST is required")39 if not values["port"]:40 raise ValueError("config ORACLE_PORT is required")41 if not values["user"]:42 raise ValueError("config ORACLE_USER is required")43 if not values["password"]:44 raise ValueError("config ORACLE_PASSWORD is required")45 if not values["database"]:46 raise ValueError("config ORACLE_DB is required")47 return values48 49 50SQL_CREATE_TABLE = """51CREATE TABLE IF NOT EXISTS {table_name} (52 id varchar2(100)53 ,text CLOB NOT NULL54 ,meta JSON55 ,embedding vector NOT NULL56)57"""58SQL_CREATE_INDEX = """59CREATE INDEX IF NOT EXISTS idx_docs_{table_name} ON {table_name}(text) 60INDEXTYPE IS CTXSYS.CONTEXT PARAMETERS 61('FILTER CTXSYS.NULL_FILTER SECTION GROUP CTXSYS.HTML_SECTION_GROUP LEXER sys.my_chinese_vgram_lexer')62"""63 64 65class OracleVector(BaseVector):66 def __init__(self, collection_name: str, config: OracleVectorConfig):67 super().__init__(collection_name)68 self.pool = self._create_connection_pool(config)69 self.table_name = f"embedding_{collection_name}"70 71 def get_type(self) -> str:72 return VectorType.ORACLE73 74 def numpy_converter_in(self, value):75 if value.dtype == numpy.float64:76 dtype = "d"77 elif value.dtype == numpy.float32:78 dtype = "f"79 else:80 dtype = "b"81 return array.array(dtype, value)82 83 def input_type_handler(self, cursor, value, arraysize):84 if isinstance(value, numpy.ndarray):85 return cursor.var(86 oracledb.DB_TYPE_VECTOR,87 arraysize=arraysize,88 inconverter=self.numpy_converter_in,89 )90 91 def numpy_converter_out(self, value):92 if value.typecode == "b":93 dtype = numpy.int894 elif value.typecode == "f":95 dtype = numpy.float3296 else:97 dtype = numpy.float6498 return numpy.array(value, copy=False, dtype=dtype)99 100 def output_type_handler(self, cursor, metadata):101 if metadata.type_code is oracledb.DB_TYPE_VECTOR:102 return cursor.var(103 metadata.type_code,104 arraysize=cursor.arraysize,105 outconverter=self.numpy_converter_out,106 )107 108 def _create_connection_pool(self, config: OracleVectorConfig):109 return oracledb.create_pool(110 user=config.user,111 password=config.password,112 dsn="{}:{}/{}".format(config.host, config.port, config.database),113 min=1,114 max=50,115 increment=1,116 )117 118 @contextmanager119 def _get_cursor(self):120 conn = self.pool.acquire()121 conn.inputtypehandler = self.input_type_handler122 conn.outputtypehandler = self.output_type_handler123 cur = conn.cursor()124 try:125 yield cur126 finally:127 cur.close()128 conn.commit()129 conn.close()130 131 def create(self, texts: list[Document], embeddings: list[list[float]], **kwargs):132 dimension = len(embeddings[0])133 self._create_collection(dimension)134 return self.add_texts(texts, embeddings)135 136 def add_texts(self, documents: list[Document], embeddings: list[list[float]], **kwargs):137 values = []138 pks = []139 for i, doc in enumerate(documents):140 doc_id = doc.metadata.get("doc_id", str(uuid.uuid4()))141 pks.append(doc_id)142 values.append(143 (144 doc_id,145 doc.page_content,146 json.dumps(doc.metadata),147 # array.array("f", embeddings[i]),148 numpy.array(embeddings[i]),149 )150 )151 # print(f"INSERT INTO {self.table_name} (id, text, meta, embedding) VALUES (:1, :2, :3, :4)")152 with self._get_cursor() as cur:153 cur.executemany(154 f"INSERT INTO {self.table_name} (id, text, meta, embedding) VALUES (:1, :2, :3, :4)", values155 )156 return pks157 158 def text_exists(self, id: str) -> bool:159 with self._get_cursor() as cur:160 cur.execute(f"SELECT id FROM {self.table_name} WHERE id = '%s'" % (id,))161 return cur.fetchone() is not None162 163 def get_by_ids(self, ids: list[str]) -> list[Document]:164 with self._get_cursor() as cur:165 cur.execute(f"SELECT meta, text FROM {self.table_name} WHERE id IN %s", (tuple(ids),))166 docs = []167 for record in cur:168 docs.append(Document(page_content=record[1], metadata=record[0]))169 return docs170 171 def delete_by_ids(self, ids: list[str]) -> None:172 with self._get_cursor() as cur:173 cur.execute(f"DELETE FROM {self.table_name} WHERE id IN %s" % (tuple(ids),))174 175 def delete_by_metadata_field(self, key: str, value: str) -> None:176 with self._get_cursor() as cur:177 cur.execute(f"DELETE FROM {self.table_name} WHERE meta->>%s = %s", (key, value))178 179 def search_by_vector(self, query_vector: list[float], **kwargs: Any) -> list[Document]:180 """181 Search the nearest neighbors to a vector.182 183 :param query_vector: The input vector to search for similar items.184 :param top_k: The number of nearest neighbors to return, default is 5.185 :return: List of Documents that are nearest to the query vector.186 """187 top_k = kwargs.get("top_k", 4)188 with self._get_cursor() as cur:189 cur.execute(190 f"SELECT meta, text, vector_distance(embedding,:1) AS distance FROM {self.table_name}"191 f" ORDER BY distance fetch first {top_k} rows only",192 [numpy.array(query_vector)],193 )194 docs = []195 score_threshold = float(kwargs.get("score_threshold") or 0.0)196 for record in cur:197 metadata, text, distance = record198 score = 1 - distance199 metadata["score"] = score200 if score > score_threshold:201 docs.append(Document(page_content=text, metadata=metadata))202 return docs203 204 def search_by_full_text(self, query: str, **kwargs: Any) -> list[Document]:205 top_k = kwargs.get("top_k", 5)206 # just not implement fetch by score_threshold now, may be later207 score_threshold = float(kwargs.get("score_threshold") or 0.0)208 if len(query) > 0:209 # Check which language the query is in210 zh_pattern = re.compile("[\u4e00-\u9fa5]+")211 match = zh_pattern.search(query)212 entities = []213 # match: query condition maybe is a chinese sentence, so using Jieba split,else using nltk split214 if match:215 words = pseg.cut(query)216 current_entity = ""217 for word, pos in words:218 if pos in {"nr", "Ng", "eng", "nz", "n", "ORG", "v"}: # nr: 人名, ns: 地名, nt: 机构名219 current_entity += word220 else:221 if current_entity:222 entities.append(current_entity)223 current_entity = ""224 if current_entity:225 entities.append(current_entity)226 else:227 try:228 nltk.data.find("tokenizers/punkt")229 nltk.data.find("corpora/stopwords")230 except LookupError:231 nltk.download("punkt")232 nltk.download("stopwords")233 print("run download")234 e_str = re.sub(r"[^\w ]", "", query)235 all_tokens = nltk.word_tokenize(e_str)236 stop_words = stopwords.words("english")237 for token in all_tokens:238 if token not in stop_words:239 entities.append(token)240 with self._get_cursor() as cur:241 cur.execute(242 f"select meta, text, embedding FROM {self.table_name}"243 f" WHERE CONTAINS(text, :1, 1) > 0 order by score(1) desc fetch first {top_k} rows only",244 [" ACCUM ".join(entities)],245 )246 docs = []247 for record in cur:248 metadata, text, embedding = record249 docs.append(Document(page_content=text, vector=embedding, metadata=metadata))250 return docs251 else:252 return [Document(page_content="", metadata={})]253 return []254 255 def delete(self) -> None:256 with self._get_cursor() as cur:257 cur.execute(f"DROP TABLE IF EXISTS {self.table_name} cascade constraints")258 259 def _create_collection(self, dimension: int):260 cache_key = f"vector_indexing_{self._collection_name}"261 lock_name = f"{cache_key}_lock"262 with redis_client.lock(lock_name, timeout=20):263 collection_exist_cache_key = f"vector_indexing_{self._collection_name}"264 if redis_client.get(collection_exist_cache_key):265 return266 267 with self._get_cursor() as cur:268 cur.execute(SQL_CREATE_TABLE.format(table_name=self.table_name))269 redis_client.set(collection_exist_cache_key, 1, ex=3600)270 with self._get_cursor() as cur:271 cur.execute(SQL_CREATE_INDEX.format(table_name=self.table_name))272 273 274class OracleVectorFactory(AbstractVectorFactory):275 def init_vector(self, dataset: Dataset, attributes: list, embeddings: Embeddings) -> OracleVector:276 if dataset.index_struct_dict:277 class_prefix: str = dataset.index_struct_dict["vector_store"]["class_prefix"]278 collection_name = class_prefix279 else:280 dataset_id = dataset.id281 collection_name = Dataset.gen_collection_name_by_id(dataset_id)282 dataset.index_struct = json.dumps(self.gen_index_struct_dict(VectorType.ORACLE, collection_name))283 284 return OracleVector(285 collection_name=collection_name,286 config=OracleVectorConfig(287 host=dify_config.ORACLE_HOST,288 port=dify_config.ORACLE_PORT,289 user=dify_config.ORACLE_USER,290 password=dify_config.ORACLE_PASSWORD,291 database=dify_config.ORACLE_DATABASE,292 ),293 )294 