Team Ai
Apppublic

Underground-Digital/Workflow-Engine

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
weaviate_vector.py284 linesDownload Raw Back to weaviate
1import datetime2import json3from typing import Any, Optional4 5import requests6import weaviate7from pydantic import BaseModel, model_validator8 9from configs import dify_config10from core.rag.datasource.vdb.field import Field11from core.rag.datasource.vdb.vector_base import BaseVector12from core.rag.datasource.vdb.vector_factory import AbstractVectorFactory13from core.rag.datasource.vdb.vector_type import VectorType14from core.rag.embedding.embedding_base import Embeddings15from core.rag.models.document import Document16from extensions.ext_redis import redis_client17from models.dataset import Dataset18 19 20class WeaviateConfig(BaseModel):21    endpoint: str22    api_key: Optional[str] = None23    batch_size: int = 10024 25    @model_validator(mode="before")26    @classmethod27    def validate_config(cls, values: dict) -> dict:28        if not values["endpoint"]:29            raise ValueError("config WEAVIATE_ENDPOINT is required")30        return values31 32 33class WeaviateVector(BaseVector):34    def __init__(self, collection_name: str, config: WeaviateConfig, attributes: list):35        super().__init__(collection_name)36        self._client = self._init_client(config)37        self._attributes = attributes38 39    def _init_client(self, config: WeaviateConfig) -> weaviate.Client:40        auth_config = weaviate.auth.AuthApiKey(api_key=config.api_key)41 42        weaviate.connect.connection.has_grpc = False43 44        try:45            client = weaviate.Client(46                url=config.endpoint, auth_client_secret=auth_config, timeout_config=(5, 60), startup_period=None47            )48        except requests.exceptions.ConnectionError:49            raise ConnectionError("Vector database connection error")50 51        client.batch.configure(52            # `batch_size` takes an `int` value to enable auto-batching53            # (`None` is used for manual batching)54            batch_size=config.batch_size,55            # dynamically update the `batch_size` based on import speed56            dynamic=True,57            # `timeout_retries` takes an `int` value to retry on time outs58            timeout_retries=3,59        )60 61        return client62 63    def get_type(self) -> str:64        return VectorType.WEAVIATE65 66    def get_collection_name(self, dataset: Dataset) -> str:67        if dataset.index_struct_dict:68            class_prefix: str = dataset.index_struct_dict["vector_store"]["class_prefix"]69            if not class_prefix.endswith("_Node"):70                # original class_prefix71                class_prefix += "_Node"72 73            return class_prefix74 75        dataset_id = dataset.id76        return Dataset.gen_collection_name_by_id(dataset_id)77 78    def to_index_struct(self) -> dict:79        return {"type": self.get_type(), "vector_store": {"class_prefix": self._collection_name}}80 81    def create(self, texts: list[Document], embeddings: list[list[float]], **kwargs):82        # create collection83        self._create_collection()84        # create vector85        self.add_texts(texts, embeddings)86 87    def _create_collection(self):88        lock_name = "vector_indexing_lock_{}".format(self._collection_name)89        with redis_client.lock(lock_name, timeout=20):90            collection_exist_cache_key = "vector_indexing_{}".format(self._collection_name)91            if redis_client.get(collection_exist_cache_key):92                return93            schema = self._default_schema(self._collection_name)94            if not self._client.schema.contains(schema):95                # create collection96                self._client.schema.create_class(schema)97            redis_client.set(collection_exist_cache_key, 1, ex=3600)98 99    def add_texts(self, documents: list[Document], embeddings: list[list[float]], **kwargs):100        uuids = self._get_uuids(documents)101        texts = [d.page_content for d in documents]102        metadatas = [d.metadata for d in documents]103 104        ids = []105 106        with self._client.batch as batch:107            for i, text in enumerate(texts):108                data_properties = {Field.TEXT_KEY.value: text}109                if metadatas is not None:110                    for key, val in metadatas[i].items():111                        data_properties[key] = self._json_serializable(val)112 113                batch.add_data_object(114                    data_object=data_properties,115                    class_name=self._collection_name,116                    uuid=uuids[i],117                    vector=embeddings[i] if embeddings else None,118                )119                ids.append(uuids[i])120        return ids121 122    def delete_by_metadata_field(self, key: str, value: str):123        # check whether the index already exists124        schema = self._default_schema(self._collection_name)125        if self._client.schema.contains(schema):126            where_filter = {"operator": "Equal", "path": [key], "valueText": value}127 128            self._client.batch.delete_objects(class_name=self._collection_name, where=where_filter, output="minimal")129 130    def delete(self):131        # check whether the index already exists132        schema = self._default_schema(self._collection_name)133        if self._client.schema.contains(schema):134            self._client.schema.delete_class(self._collection_name)135 136    def text_exists(self, id: str) -> bool:137        collection_name = self._collection_name138        schema = self._default_schema(self._collection_name)139 140        # check whether the index already exists141        if not self._client.schema.contains(schema):142            return False143        result = (144            self._client.query.get(collection_name)145            .with_additional(["id"])146            .with_where(147                {148                    "path": ["doc_id"],149                    "operator": "Equal",150                    "valueText": id,151                }152            )153            .with_limit(1)154            .do()155        )156 157        if "errors" in result:158            raise ValueError(f"Error during query: {result['errors']}")159 160        entries = result["data"]["Get"][collection_name]161        if len(entries) == 0:162            return False163 164        return True165 166    def delete_by_ids(self, ids: list[str]) -> None:167        # check whether the index already exists168        schema = self._default_schema(self._collection_name)169        if self._client.schema.contains(schema):170            for uuid in ids:171                try:172                    self._client.data_object.delete(173                        class_name=self._collection_name,174                        uuid=uuid,175                    )176                except weaviate.UnexpectedStatusCodeException as e:177                    # tolerate not found error178                    if e.status_code != 404:179                        raise e180 181    def search_by_vector(self, query_vector: list[float], **kwargs: Any) -> list[Document]:182        """Look up similar documents by embedding vector in Weaviate."""183        collection_name = self._collection_name184        properties = self._attributes185        properties.append(Field.TEXT_KEY.value)186        query_obj = self._client.query.get(collection_name, properties)187 188        vector = {"vector": query_vector}189        if kwargs.get("where_filter"):190            query_obj = query_obj.with_where(kwargs.get("where_filter"))191        result = (192            query_obj.with_near_vector(vector)193            .with_limit(kwargs.get("top_k", 4))194            .with_additional(["vector", "distance"])195            .do()196        )197        if "errors" in result:198            raise ValueError(f"Error during query: {result['errors']}")199 200        docs_and_scores = []201        for res in result["data"]["Get"][collection_name]:202            text = res.pop(Field.TEXT_KEY.value)203            score = 1 - res["_additional"]["distance"]204            docs_and_scores.append((Document(page_content=text, metadata=res), score))205 206        docs = []207        for doc, score in docs_and_scores:208            score_threshold = float(kwargs.get("score_threshold") or 0.0)209            # check score threshold210            if score > score_threshold:211                doc.metadata["score"] = score212                docs.append(doc)213        # Sort the documents by score in descending order214        docs = sorted(docs, key=lambda x: x.metadata["score"], reverse=True)215        return docs216 217    def search_by_full_text(self, query: str, **kwargs: Any) -> list[Document]:218        """Return docs using BM25F.219 220        Args:221            query: Text to look up documents similar to.222            k: Number of Documents to return. Defaults to 4.223 224        Returns:225            List of Documents most similar to the query.226        """227        collection_name = self._collection_name228        content: dict[str, Any] = {"concepts": [query]}229        properties = self._attributes230        properties.append(Field.TEXT_KEY.value)231        if kwargs.get("search_distance"):232            content["certainty"] = kwargs.get("search_distance")233        query_obj = self._client.query.get(collection_name, properties)234        if kwargs.get("where_filter"):235            query_obj = query_obj.with_where(kwargs.get("where_filter"))236        query_obj = query_obj.with_additional(["vector"])237        properties = ["text"]238        result = query_obj.with_bm25(query=query, properties=properties).with_limit(kwargs.get("top_k", 4)).do()239        if "errors" in result:240            raise ValueError(f"Error during query: {result['errors']}")241        docs = []242        for res in result["data"]["Get"][collection_name]:243            text = res.pop(Field.TEXT_KEY.value)244            additional = res.pop("_additional")245            docs.append(Document(page_content=text, vector=additional["vector"], metadata=res))246        return docs247 248    def _default_schema(self, index_name: str) -> dict:249        return {250            "class": index_name,251            "properties": [252                {253                    "name": "text",254                    "dataType": ["text"],255                }256            ],257        }258 259    def _json_serializable(self, value: Any) -> Any:260        if isinstance(value, datetime.datetime):261            return value.isoformat()262        return value263 264 265class WeaviateVectorFactory(AbstractVectorFactory):266    def init_vector(self, dataset: Dataset, attributes: list, embeddings: Embeddings) -> WeaviateVector:267        if dataset.index_struct_dict:268            class_prefix: str = dataset.index_struct_dict["vector_store"]["class_prefix"]269            collection_name = class_prefix270        else:271            dataset_id = dataset.id272            collection_name = Dataset.gen_collection_name_by_id(dataset_id)273            dataset.index_struct = json.dumps(self.gen_index_struct_dict(VectorType.WEAVIATE, collection_name))274 275        return WeaviateVector(276            collection_name=collection_name,277            config=WeaviateConfig(278                endpoint=dify_config.WEAVIATE_ENDPOINT,279                api_key=dify_config.WEAVIATE_API_KEY,280                batch_size=dify_config.WEAVIATE_BATCH_SIZE,281            ),282            attributes=attributes,283        )284