Team Ai
Apppublic

Underground-Digital/Workflow-Engine

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
oraclevector.py294 linesDownload Raw Back to oracle
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