Underground-Digital/Workflow-Engine
0
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 