Team Ai
Apppublic

Underground-Digital/Workflow-Engine

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
elasticsearch_vector.py217 linesDownload Raw Back to elasticsearch
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