Underground-Digital/Workflow-Engine
0
1import json2import logging3import math4from typing import Any, Optional5from urllib.parse import urlparse6 7import requests8from elasticsearch import Elasticsearch9from flask import current_app10from pydantic import BaseModel, model_validator11 12from core.rag.datasource.vdb.field import Field13from core.rag.datasource.vdb.vector_base import BaseVector14from core.rag.datasource.vdb.vector_factory import AbstractVectorFactory15from core.rag.datasource.vdb.vector_type import VectorType16from core.rag.embedding.embedding_base import Embeddings17from core.rag.models.document import Document18from extensions.ext_redis import redis_client19from models.dataset import Dataset20 21logger = logging.getLogger(__name__)22 23 24class ElasticSearchConfig(BaseModel):25 host: str26 port: int27 username: str28 password: str29 30 @model_validator(mode="before")31 @classmethod32 def validate_config(cls, values: dict) -> dict:33 if not values["host"]:34 raise ValueError("config HOST is required")35 if not values["port"]:36 raise ValueError("config PORT is required")37 if not values["username"]:38 raise ValueError("config USERNAME is required")39 if not values["password"]:40 raise ValueError("config PASSWORD is required")41 return values42 43 44class ElasticSearchVector(BaseVector):45 def __init__(self, index_name: str, config: ElasticSearchConfig, attributes: list):46 super().__init__(index_name.lower())47 self._client = self._init_client(config)48 self._version = self._get_version()49 self._check_version()50 self._attributes = attributes51 52 def _init_client(self, config: ElasticSearchConfig) -> Elasticsearch:53 try:54 parsed_url = urlparse(config.host)55 if parsed_url.scheme in {"http", "https"}:56 hosts = f"{config.host}:{config.port}"57 else:58 hosts = f"http://{config.host}:{config.port}"59 client = Elasticsearch(60 hosts=hosts,61 basic_auth=(config.username, config.password),62 request_timeout=100000,63 retry_on_timeout=True,64 max_retries=10000,65 )66 except requests.exceptions.ConnectionError:67 raise ConnectionError("Vector database connection error")68 69 return client70 71 def _get_version(self) -> str:72 info = self._client.info()73 return info["version"]["number"]74 75 def _check_version(self):76 if self._version < "8.0.0":77 raise ValueError("Elasticsearch vector database version must be greater than 8.0.0")78 79 def get_type(self) -> str:80 return VectorType.ELASTICSEARCH81 82 def add_texts(self, documents: list[Document], embeddings: list[list[float]], **kwargs):83 uuids = self._get_uuids(documents)84 for i in range(len(documents)):85 self._client.index(86 index=self._collection_name,87 id=uuids[i],88 document={89 Field.CONTENT_KEY.value: documents[i].page_content,90 Field.VECTOR.value: embeddings[i] or None,91 Field.METADATA_KEY.value: documents[i].metadata or {},92 },93 )94 self._client.indices.refresh(index=self._collection_name)95 return uuids96 97 def text_exists(self, id: str) -> bool:98 return bool(self._client.exists(index=self._collection_name, id=id))99 100 def delete_by_ids(self, ids: list[str]) -> None:101 for id in ids:102 self._client.delete(index=self._collection_name, id=id)103 104 def delete_by_metadata_field(self, key: str, value: str) -> None:105 query_str = {"query": {"match": {f"metadata.{key}": f"{value}"}}}106 results = self._client.search(index=self._collection_name, body=query_str)107 ids = [hit["_id"] for hit in results["hits"]["hits"]]108 if ids:109 self.delete_by_ids(ids)110 111 def delete(self) -> None:112 self._client.indices.delete(index=self._collection_name)113 114 def search_by_vector(self, query_vector: list[float], **kwargs: Any) -> list[Document]:115 top_k = kwargs.get("top_k", 4)116 num_candidates = math.ceil(top_k * 1.5)117 knn = {"field": Field.VECTOR.value, "query_vector": query_vector, "k": top_k, "num_candidates": num_candidates}118 119 results = self._client.search(index=self._collection_name, knn=knn, size=top_k)120 121 docs_and_scores = []122 for hit in results["hits"]["hits"]:123 docs_and_scores.append(124 (125 Document(126 page_content=hit["_source"][Field.CONTENT_KEY.value],127 vector=hit["_source"][Field.VECTOR.value],128 metadata=hit["_source"][Field.METADATA_KEY.value],129 ),130 hit["_score"],131 )132 )133 134 docs = []135 for doc, score in docs_and_scores:136 score_threshold = float(kwargs.get("score_threshold") or 0.0)137 if score > score_threshold:138 doc.metadata["score"] = score139 docs.append(doc)140 141 return docs142 143 def search_by_full_text(self, query: str, **kwargs: Any) -> list[Document]:144 query_str = {"match": {Field.CONTENT_KEY.value: query}}145 results = self._client.search(index=self._collection_name, query=query_str, size=kwargs.get("top_k", 4))146 docs = []147 for hit in results["hits"]["hits"]:148 docs.append(149 Document(150 page_content=hit["_source"][Field.CONTENT_KEY.value],151 vector=hit["_source"][Field.VECTOR.value],152 metadata=hit["_source"][Field.METADATA_KEY.value],153 )154 )155 156 return docs157 158 def create(self, texts: list[Document], embeddings: list[list[float]], **kwargs):159 metadatas = [d.metadata for d in texts]160 self.create_collection(embeddings, metadatas)161 self.add_texts(texts, embeddings, **kwargs)162 163 def create_collection(164 self, embeddings: list, metadatas: Optional[list[dict]] = None, index_params: Optional[dict] = None165 ):166 lock_name = f"vector_indexing_lock_{self._collection_name}"167 with redis_client.lock(lock_name, timeout=20):168 collection_exist_cache_key = f"vector_indexing_{self._collection_name}"169 if redis_client.get(collection_exist_cache_key):170 logger.info(f"Collection {self._collection_name} already exists.")171 return172 173 if not self._client.indices.exists(index=self._collection_name):174 dim = len(embeddings[0])175 mappings = {176 "properties": {177 Field.CONTENT_KEY.value: {"type": "text"},178 Field.VECTOR.value: { # Make sure the dimension is correct here179 "type": "dense_vector",180 "dims": dim,181 "similarity": "cosine",182 },183 Field.METADATA_KEY.value: {184 "type": "object",185 "properties": {186 "doc_id": {"type": "keyword"} # Map doc_id to keyword type187 },188 },189 }190 }191 self._client.indices.create(index=self._collection_name, mappings=mappings)192 193 redis_client.set(collection_exist_cache_key, 1, ex=3600)194 195 196class ElasticSearchVectorFactory(AbstractVectorFactory):197 def init_vector(self, dataset: Dataset, attributes: list, embeddings: Embeddings) -> ElasticSearchVector:198 if dataset.index_struct_dict:199 class_prefix: str = dataset.index_struct_dict["vector_store"]["class_prefix"]200 collection_name = class_prefix201 else:202 dataset_id = dataset.id203 collection_name = Dataset.gen_collection_name_by_id(dataset_id)204 dataset.index_struct = json.dumps(self.gen_index_struct_dict(VectorType.ELASTICSEARCH, collection_name))205 206 config = current_app.config207 return ElasticSearchVector(208 index_name=collection_name,209 config=ElasticSearchConfig(210 host=config.get("ELASTICSEARCH_HOST"),211 port=config.get("ELASTICSEARCH_PORT"),212 username=config.get("ELASTICSEARCH_USERNAME"),213 password=config.get("ELASTICSEARCH_PASSWORD"),214 ),215 attributes=[],216 )217 