Underground-Digital/Workflow-Engine
0
1import base642import enum3import hashlib4import hmac5import json6import logging7import os8import pickle9import re10import time11from json import JSONDecodeError12 13from sqlalchemy import func14from sqlalchemy.dialects.postgresql import JSONB15 16from configs import dify_config17from core.rag.retrieval.retrieval_methods import RetrievalMethod18from extensions.ext_database import db19from extensions.ext_storage import storage20 21from .account import Account22from .model import App, Tag, TagBinding, UploadFile23from .types import StringUUID24 25 26class DatasetPermissionEnum(str, enum.Enum):27 ONLY_ME = "only_me"28 ALL_TEAM = "all_team_members"29 PARTIAL_TEAM = "partial_members"30 31 32class Dataset(db.Model):33 __tablename__ = "datasets"34 __table_args__ = (35 db.PrimaryKeyConstraint("id", name="dataset_pkey"),36 db.Index("dataset_tenant_idx", "tenant_id"),37 db.Index("retrieval_model_idx", "retrieval_model", postgresql_using="gin"),38 )39 40 INDEXING_TECHNIQUE_LIST = ["high_quality", "economy", None]41 PROVIDER_LIST = ["vendor", "external", None]42 43 id = db.Column(StringUUID, server_default=db.text("uuid_generate_v4()"))44 tenant_id = db.Column(StringUUID, nullable=False)45 name = db.Column(db.String(255), nullable=False)46 description = db.Column(db.Text, nullable=True)47 provider = db.Column(db.String(255), nullable=False, server_default=db.text("'vendor'::character varying"))48 permission = db.Column(db.String(255), nullable=False, server_default=db.text("'only_me'::character varying"))49 data_source_type = db.Column(db.String(255))50 indexing_technique = db.Column(db.String(255), nullable=True)51 index_struct = db.Column(db.Text, nullable=True)52 created_by = db.Column(StringUUID, nullable=False)53 created_at = db.Column(db.DateTime, nullable=False, server_default=db.text("CURRENT_TIMESTAMP(0)"))54 updated_by = db.Column(StringUUID, nullable=True)55 updated_at = db.Column(db.DateTime, nullable=False, server_default=db.text("CURRENT_TIMESTAMP(0)"))56 embedding_model = db.Column(db.String(255), nullable=True)57 embedding_model_provider = db.Column(db.String(255), nullable=True)58 collection_binding_id = db.Column(StringUUID, nullable=True)59 retrieval_model = db.Column(JSONB, nullable=True)60 61 @property62 def dataset_keyword_table(self):63 dataset_keyword_table = (64 db.session.query(DatasetKeywordTable).filter(DatasetKeywordTable.dataset_id == self.id).first()65 )66 if dataset_keyword_table:67 return dataset_keyword_table68 69 return None70 71 @property72 def index_struct_dict(self):73 return json.loads(self.index_struct) if self.index_struct else None74 75 @property76 def external_retrieval_model(self):77 default_retrieval_model = {78 "top_k": 2,79 "score_threshold": 0.0,80 }81 return self.retrieval_model or default_retrieval_model82 83 @property84 def created_by_account(self):85 return db.session.get(Account, self.created_by)86 87 @property88 def latest_process_rule(self):89 return (90 DatasetProcessRule.query.filter(DatasetProcessRule.dataset_id == self.id)91 .order_by(DatasetProcessRule.created_at.desc())92 .first()93 )94 95 @property96 def app_count(self):97 return (98 db.session.query(func.count(AppDatasetJoin.id))99 .filter(AppDatasetJoin.dataset_id == self.id, App.id == AppDatasetJoin.app_id)100 .scalar()101 )102 103 @property104 def document_count(self):105 return db.session.query(func.count(Document.id)).filter(Document.dataset_id == self.id).scalar()106 107 @property108 def available_document_count(self):109 return (110 db.session.query(func.count(Document.id))111 .filter(112 Document.dataset_id == self.id,113 Document.indexing_status == "completed",114 Document.enabled == True,115 Document.archived == False,116 )117 .scalar()118 )119 120 @property121 def available_segment_count(self):122 return (123 db.session.query(func.count(DocumentSegment.id))124 .filter(125 DocumentSegment.dataset_id == self.id,126 DocumentSegment.status == "completed",127 DocumentSegment.enabled == True,128 )129 .scalar()130 )131 132 @property133 def word_count(self):134 return (135 Document.query.with_entities(func.coalesce(func.sum(Document.word_count)))136 .filter(Document.dataset_id == self.id)137 .scalar()138 )139 140 @property141 def doc_form(self):142 document = db.session.query(Document).filter(Document.dataset_id == self.id).first()143 if document:144 return document.doc_form145 return None146 147 @property148 def retrieval_model_dict(self):149 default_retrieval_model = {150 "search_method": RetrievalMethod.SEMANTIC_SEARCH.value,151 "reranking_enable": False,152 "reranking_model": {"reranking_provider_name": "", "reranking_model_name": ""},153 "top_k": 2,154 "score_threshold_enabled": False,155 }156 return self.retrieval_model or default_retrieval_model157 158 @property159 def tags(self):160 tags = (161 db.session.query(Tag)162 .join(TagBinding, Tag.id == TagBinding.tag_id)163 .filter(164 TagBinding.target_id == self.id,165 TagBinding.tenant_id == self.tenant_id,166 Tag.tenant_id == self.tenant_id,167 Tag.type == "knowledge",168 )169 .all()170 )171 172 return tags or []173 174 @property175 def external_knowledge_info(self):176 if self.provider != "external":177 return None178 external_knowledge_binding = (179 db.session.query(ExternalKnowledgeBindings).filter(ExternalKnowledgeBindings.dataset_id == self.id).first()180 )181 if not external_knowledge_binding:182 return None183 external_knowledge_api = (184 db.session.query(ExternalKnowledgeApis)185 .filter(ExternalKnowledgeApis.id == external_knowledge_binding.external_knowledge_api_id)186 .first()187 )188 if not external_knowledge_api:189 return None190 return {191 "external_knowledge_id": external_knowledge_binding.external_knowledge_id,192 "external_knowledge_api_id": external_knowledge_api.id,193 "external_knowledge_api_name": external_knowledge_api.name,194 "external_knowledge_api_endpoint": json.loads(external_knowledge_api.settings).get("endpoint", ""),195 }196 197 @staticmethod198 def gen_collection_name_by_id(dataset_id: str) -> str:199 normalized_dataset_id = dataset_id.replace("-", "_")200 return f"Vector_index_{normalized_dataset_id}_Node"201 202 203class DatasetProcessRule(db.Model):204 __tablename__ = "dataset_process_rules"205 __table_args__ = (206 db.PrimaryKeyConstraint("id", name="dataset_process_rule_pkey"),207 db.Index("dataset_process_rule_dataset_id_idx", "dataset_id"),208 )209 210 id = db.Column(StringUUID, nullable=False, server_default=db.text("uuid_generate_v4()"))211 dataset_id = db.Column(StringUUID, nullable=False)212 mode = db.Column(db.String(255), nullable=False, server_default=db.text("'automatic'::character varying"))213 rules = db.Column(db.Text, nullable=True)214 created_by = db.Column(StringUUID, nullable=False)215 created_at = db.Column(db.DateTime, nullable=False, server_default=db.text("CURRENT_TIMESTAMP(0)"))216 217 MODES = ["automatic", "custom"]218 PRE_PROCESSING_RULES = ["remove_stopwords", "remove_extra_spaces", "remove_urls_emails"]219 AUTOMATIC_RULES = {220 "pre_processing_rules": [221 {"id": "remove_extra_spaces", "enabled": True},222 {"id": "remove_urls_emails", "enabled": False},223 ],224 "segmentation": {"delimiter": "\n", "max_tokens": 500, "chunk_overlap": 50},225 }226 227 def to_dict(self):228 return {229 "id": self.id,230 "dataset_id": self.dataset_id,231 "mode": self.mode,232 "rules": self.rules_dict,233 "created_by": self.created_by,234 "created_at": self.created_at,235 }236 237 @property238 def rules_dict(self):239 try:240 return json.loads(self.rules) if self.rules else None241 except JSONDecodeError:242 return None243 244 245class Document(db.Model):246 __tablename__ = "documents"247 __table_args__ = (248 db.PrimaryKeyConstraint("id", name="document_pkey"),249 db.Index("document_dataset_id_idx", "dataset_id"),250 db.Index("document_is_paused_idx", "is_paused"),251 db.Index("document_tenant_idx", "tenant_id"),252 )253 254 # initial fields255 id = db.Column(StringUUID, nullable=False, server_default=db.text("uuid_generate_v4()"))256 tenant_id = db.Column(StringUUID, nullable=False)257 dataset_id = db.Column(StringUUID, nullable=False)258 position = db.Column(db.Integer, nullable=False)259 data_source_type = db.Column(db.String(255), nullable=False)260 data_source_info = db.Column(db.Text, nullable=True)261 dataset_process_rule_id = db.Column(StringUUID, nullable=True)262 batch = db.Column(db.String(255), nullable=False)263 name = db.Column(db.String(255), nullable=False)264 created_from = db.Column(db.String(255), nullable=False)265 created_by = db.Column(StringUUID, nullable=False)266 created_api_request_id = db.Column(StringUUID, nullable=True)267 created_at = db.Column(db.DateTime, nullable=False, server_default=db.text("CURRENT_TIMESTAMP(0)"))268 269 # start processing270 processing_started_at = db.Column(db.DateTime, nullable=True)271 272 # parsing273 file_id = db.Column(db.Text, nullable=True)274 word_count = db.Column(db.Integer, nullable=True)275 parsing_completed_at = db.Column(db.DateTime, nullable=True)276 277 # cleaning278 cleaning_completed_at = db.Column(db.DateTime, nullable=True)279 280 # split281 splitting_completed_at = db.Column(db.DateTime, nullable=True)282 283 # indexing284 tokens = db.Column(db.Integer, nullable=True)285 indexing_latency = db.Column(db.Float, nullable=True)286 completed_at = db.Column(db.DateTime, nullable=True)287 288 # pause289 is_paused = db.Column(db.Boolean, nullable=True, server_default=db.text("false"))290 paused_by = db.Column(StringUUID, nullable=True)291 paused_at = db.Column(db.DateTime, nullable=True)292 293 # error294 error = db.Column(db.Text, nullable=True)295 stopped_at = db.Column(db.DateTime, nullable=True)296 297 # basic fields298 indexing_status = db.Column(db.String(255), nullable=False, server_default=db.text("'waiting'::character varying"))299 enabled = db.Column(db.Boolean, nullable=False, server_default=db.text("true"))300 disabled_at = db.Column(db.DateTime, nullable=True)301 disabled_by = db.Column(StringUUID, nullable=True)302 archived = db.Column(db.Boolean, nullable=False, server_default=db.text("false"))303 archived_reason = db.Column(db.String(255), nullable=True)304 archived_by = db.Column(StringUUID, nullable=True)305 archived_at = db.Column(db.DateTime, nullable=True)306 updated_at = db.Column(db.DateTime, nullable=False, server_default=db.text("CURRENT_TIMESTAMP(0)"))307 doc_type = db.Column(db.String(40), nullable=True)308 doc_metadata = db.Column(db.JSON, nullable=True)309 doc_form = db.Column(db.String(255), nullable=False, server_default=db.text("'text_model'::character varying"))310 doc_language = db.Column(db.String(255), nullable=True)311 312 DATA_SOURCES = ["upload_file", "notion_import", "website_crawl"]313 314 @property315 def display_status(self):316 status = None317 if self.indexing_status == "waiting":318 status = "queuing"319 elif self.indexing_status not in {"completed", "error", "waiting"} and self.is_paused:320 status = "paused"321 elif self.indexing_status in {"parsing", "cleaning", "splitting", "indexing"}:322 status = "indexing"323 elif self.indexing_status == "error":324 status = "error"325 elif self.indexing_status == "completed" and not self.archived and self.enabled:326 status = "available"327 elif self.indexing_status == "completed" and not self.archived and not self.enabled:328 status = "disabled"329 elif self.indexing_status == "completed" and self.archived:330 status = "archived"331 return status332 333 @property334 def data_source_info_dict(self):335 if self.data_source_info:336 try:337 data_source_info_dict = json.loads(self.data_source_info)338 except JSONDecodeError:339 data_source_info_dict = {}340 341 return data_source_info_dict342 return None343 344 @property345 def data_source_detail_dict(self):346 if self.data_source_info:347 if self.data_source_type == "upload_file":348 data_source_info_dict = json.loads(self.data_source_info)349 file_detail = (350 db.session.query(UploadFile)351 .filter(UploadFile.id == data_source_info_dict["upload_file_id"])352 .one_or_none()353 )354 if file_detail:355 return {356 "upload_file": {357 "id": file_detail.id,358 "name": file_detail.name,359 "size": file_detail.size,360 "extension": file_detail.extension,361 "mime_type": file_detail.mime_type,362 "created_by": file_detail.created_by,363 "created_at": file_detail.created_at.timestamp(),364 }365 }366 elif self.data_source_type in {"notion_import", "website_crawl"}:367 return json.loads(self.data_source_info)368 return {}369 370 @property371 def average_segment_length(self):372 if self.word_count and self.word_count != 0 and self.segment_count and self.segment_count != 0:373 return self.word_count // self.segment_count374 return 0375 376 @property377 def dataset_process_rule(self):378 if self.dataset_process_rule_id:379 return db.session.get(DatasetProcessRule, self.dataset_process_rule_id)380 return None381 382 @property383 def dataset(self):384 return db.session.query(Dataset).filter(Dataset.id == self.dataset_id).one_or_none()385 386 @property387 def segment_count(self):388 return DocumentSegment.query.filter(DocumentSegment.document_id == self.id).count()389 390 @property391 def hit_count(self):392 return (393 DocumentSegment.query.with_entities(func.coalesce(func.sum(DocumentSegment.hit_count)))394 .filter(DocumentSegment.document_id == self.id)395 .scalar()396 )397 398 def to_dict(self):399 return {400 "id": self.id,401 "tenant_id": self.tenant_id,402 "dataset_id": self.dataset_id,403 "position": self.position,404 "data_source_type": self.data_source_type,405 "data_source_info": self.data_source_info,406 "dataset_process_rule_id": self.dataset_process_rule_id,407 "batch": self.batch,408 "name": self.name,409 "created_from": self.created_from,410 "created_by": self.created_by,411 "created_api_request_id": self.created_api_request_id,412 "created_at": self.created_at,413 "processing_started_at": self.processing_started_at,414 "file_id": self.file_id,415 "word_count": self.word_count,416 "parsing_completed_at": self.parsing_completed_at,417 "cleaning_completed_at": self.cleaning_completed_at,418 "splitting_completed_at": self.splitting_completed_at,419 "tokens": self.tokens,420 "indexing_latency": self.indexing_latency,421 "completed_at": self.completed_at,422 "is_paused": self.is_paused,423 "paused_by": self.paused_by,424 "paused_at": self.paused_at,425 "error": self.error,426 "stopped_at": self.stopped_at,427 "indexing_status": self.indexing_status,428 "enabled": self.enabled,429 "disabled_at": self.disabled_at,430 "disabled_by": self.disabled_by,431 "archived": self.archived,432 "archived_reason": self.archived_reason,433 "archived_by": self.archived_by,434 "archived_at": self.archived_at,435 "updated_at": self.updated_at,436 "doc_type": self.doc_type,437 "doc_metadata": self.doc_metadata,438 "doc_form": self.doc_form,439 "doc_language": self.doc_language,440 "display_status": self.display_status,441 "data_source_info_dict": self.data_source_info_dict,442 "average_segment_length": self.average_segment_length,443 "dataset_process_rule": self.dataset_process_rule.to_dict() if self.dataset_process_rule else None,444 "dataset": self.dataset.to_dict() if self.dataset else None,445 "segment_count": self.segment_count,446 "hit_count": self.hit_count,447 }448 449 @classmethod450 def from_dict(cls, data: dict):451 return cls(452 id=data.get("id"),453 tenant_id=data.get("tenant_id"),454 dataset_id=data.get("dataset_id"),455 position=data.get("position"),456 data_source_type=data.get("data_source_type"),457 data_source_info=data.get("data_source_info"),458 dataset_process_rule_id=data.get("dataset_process_rule_id"),459 batch=data.get("batch"),460 name=data.get("name"),461 created_from=data.get("created_from"),462 created_by=data.get("created_by"),463 created_api_request_id=data.get("created_api_request_id"),464 created_at=data.get("created_at"),465 processing_started_at=data.get("processing_started_at"),466 file_id=data.get("file_id"),467 word_count=data.get("word_count"),468 parsing_completed_at=data.get("parsing_completed_at"),469 cleaning_completed_at=data.get("cleaning_completed_at"),470 splitting_completed_at=data.get("splitting_completed_at"),471 tokens=data.get("tokens"),472 indexing_latency=data.get("indexing_latency"),473 completed_at=data.get("completed_at"),474 is_paused=data.get("is_paused"),475 paused_by=data.get("paused_by"),476 paused_at=data.get("paused_at"),477 error=data.get("error"),478 stopped_at=data.get("stopped_at"),479 indexing_status=data.get("indexing_status"),480 enabled=data.get("enabled"),481 disabled_at=data.get("disabled_at"),482 disabled_by=data.get("disabled_by"),483 archived=data.get("archived"),484 archived_reason=data.get("archived_reason"),485 archived_by=data.get("archived_by"),486 archived_at=data.get("archived_at"),487 updated_at=data.get("updated_at"),488 doc_type=data.get("doc_type"),489 doc_metadata=data.get("doc_metadata"),490 doc_form=data.get("doc_form"),491 doc_language=data.get("doc_language"),492 )493 494 495class DocumentSegment(db.Model):496 __tablename__ = "document_segments"497 __table_args__ = (498 db.PrimaryKeyConstraint("id", name="document_segment_pkey"),499 db.Index("document_segment_dataset_id_idx", "dataset_id"),500 db.Index("document_segment_document_id_idx", "document_id"),501 db.Index("document_segment_tenant_dataset_idx", "dataset_id", "tenant_id"),502 db.Index("document_segment_tenant_document_idx", "document_id", "tenant_id"),503 db.Index("document_segment_dataset_node_idx", "dataset_id", "index_node_id"),504 db.Index("document_segment_tenant_idx", "tenant_id"),505 )506 507 # initial fields508 id = db.Column(StringUUID, nullable=False, server_default=db.text("uuid_generate_v4()"))509 tenant_id = db.Column(StringUUID, nullable=False)510 dataset_id = db.Column(StringUUID, nullable=False)511 document_id = db.Column(StringUUID, nullable=False)512 position = db.Column(db.Integer, nullable=False)513 content = db.Column(db.Text, nullable=False)514 answer = db.Column(db.Text, nullable=True)515 word_count = db.Column(db.Integer, nullable=False)516 tokens = db.Column(db.Integer, nullable=False)517 518 # indexing fields519 keywords = db.Column(db.JSON, nullable=True)520 index_node_id = db.Column(db.String(255), nullable=True)521 index_node_hash = db.Column(db.String(255), nullable=True)522 523 # basic fields524 hit_count = db.Column(db.Integer, nullable=False, default=0)525 enabled = db.Column(db.Boolean, nullable=False, server_default=db.text("true"))526 disabled_at = db.Column(db.DateTime, nullable=True)527 disabled_by = db.Column(StringUUID, nullable=True)528 status = db.Column(db.String(255), nullable=False, server_default=db.text("'waiting'::character varying"))529 created_by = db.Column(StringUUID, nullable=False)530 created_at = db.Column(db.DateTime, nullable=False, server_default=db.text("CURRENT_TIMESTAMP(0)"))531 updated_by = db.Column(StringUUID, nullable=True)532 updated_at = db.Column(db.DateTime, nullable=False, server_default=db.text("CURRENT_TIMESTAMP(0)"))533 indexing_at = db.Column(db.DateTime, nullable=True)534 completed_at = db.Column(db.DateTime, nullable=True)535 error = db.Column(db.Text, nullable=True)536 stopped_at = db.Column(db.DateTime, nullable=True)537 538 @property539 def dataset(self):540 return db.session.query(Dataset).filter(Dataset.id == self.dataset_id).first()541 542 @property543 def document(self):544 return db.session.query(Document).filter(Document.id == self.document_id).first()545 546 @property547 def previous_segment(self):548 return (549 db.session.query(DocumentSegment)550 .filter(DocumentSegment.document_id == self.document_id, DocumentSegment.position == self.position - 1)551 .first()552 )553 554 @property555 def next_segment(self):556 return (557 db.session.query(DocumentSegment)558 .filter(DocumentSegment.document_id == self.document_id, DocumentSegment.position == self.position + 1)559 .first()560 )561 562 def get_sign_content(self):563 signed_urls = []564 text = self.content565 566 # For data before v0.10.0567 pattern = r"/files/([a-f0-9\-]+)/image-preview"568 matches = re.finditer(pattern, text)569 for match in matches:570 upload_file_id = match.group(1)571 nonce = os.urandom(16).hex()572 timestamp = str(int(time.time()))573 data_to_sign = f"image-preview|{upload_file_id}|{timestamp}|{nonce}"574 secret_key = dify_config.SECRET_KEY.encode() if dify_config.SECRET_KEY else b""575 sign = hmac.new(secret_key, data_to_sign.encode(), hashlib.sha256).digest()576 encoded_sign = base64.urlsafe_b64encode(sign).decode()577 578 params = f"timestamp={timestamp}&nonce={nonce}&sign={encoded_sign}"579 signed_url = f"{match.group(0)}?{params}"580 signed_urls.append((match.start(), match.end(), signed_url))581 582 # For data after v0.10.0583 pattern = r"/files/([a-f0-9\-]+)/file-preview"584 matches = re.finditer(pattern, text)585 for match in matches:586 upload_file_id = match.group(1)587 nonce = os.urandom(16).hex()588 timestamp = str(int(time.time()))589 data_to_sign = f"file-preview|{upload_file_id}|{timestamp}|{nonce}"590 secret_key = dify_config.SECRET_KEY.encode() if dify_config.SECRET_KEY else b""591 sign = hmac.new(secret_key, data_to_sign.encode(), hashlib.sha256).digest()592 encoded_sign = base64.urlsafe_b64encode(sign).decode()593 594 params = f"timestamp={timestamp}&nonce={nonce}&sign={encoded_sign}"595 signed_url = f"{match.group(0)}?{params}"596 signed_urls.append((match.start(), match.end(), signed_url))597 598 # Reconstruct the text with signed URLs599 offset = 0600 for start, end, signed_url in signed_urls:601 text = text[: start + offset] + signed_url + text[end + offset :]602 offset += len(signed_url) - (end - start)603 604 return text605 606 607class AppDatasetJoin(db.Model):608 __tablename__ = "app_dataset_joins"609 __table_args__ = (610 db.PrimaryKeyConstraint("id", name="app_dataset_join_pkey"),611 db.Index("app_dataset_join_app_dataset_idx", "dataset_id", "app_id"),612 )613 614 id = db.Column(StringUUID, primary_key=True, nullable=False, server_default=db.text("uuid_generate_v4()"))615 app_id = db.Column(StringUUID, nullable=False)616 dataset_id = db.Column(StringUUID, nullable=False)617 created_at = db.Column(db.DateTime, nullable=False, server_default=db.func.current_timestamp())618 619 @property620 def app(self):621 return db.session.get(App, self.app_id)622 623 624class DatasetQuery(db.Model):625 __tablename__ = "dataset_queries"626 __table_args__ = (627 db.PrimaryKeyConstraint("id", name="dataset_query_pkey"),628 db.Index("dataset_query_dataset_id_idx", "dataset_id"),629 )630 631 id = db.Column(StringUUID, primary_key=True, nullable=False, server_default=db.text("uuid_generate_v4()"))632 dataset_id = db.Column(StringUUID, nullable=False)633 content = db.Column(db.Text, nullable=False)634 source = db.Column(db.String(255), nullable=False)635 source_app_id = db.Column(StringUUID, nullable=True)636 created_by_role = db.Column(db.String, nullable=False)637 created_by = db.Column(StringUUID, nullable=False)638 created_at = db.Column(db.DateTime, nullable=False, server_default=db.func.current_timestamp())639 640 641class DatasetKeywordTable(db.Model):642 __tablename__ = "dataset_keyword_tables"643 __table_args__ = (644 db.PrimaryKeyConstraint("id", name="dataset_keyword_table_pkey"),645 db.Index("dataset_keyword_table_dataset_id_idx", "dataset_id"),646 )647 648 id = db.Column(StringUUID, primary_key=True, server_default=db.text("uuid_generate_v4()"))649 dataset_id = db.Column(StringUUID, nullable=False, unique=True)650 keyword_table = db.Column(db.Text, nullable=False)651 data_source_type = db.Column(652 db.String(255), nullable=False, server_default=db.text("'database'::character varying")653 )654 655 @property656 def keyword_table_dict(self):657 class SetDecoder(json.JSONDecoder):658 def __init__(self, *args, **kwargs):659 super().__init__(object_hook=self.object_hook, *args, **kwargs)660 661 def object_hook(self, dct):662 if isinstance(dct, dict):663 for keyword, node_idxs in dct.items():664 if isinstance(node_idxs, list):665 dct[keyword] = set(node_idxs)666 return dct667 668 # get dataset669 dataset = Dataset.query.filter_by(id=self.dataset_id).first()670 if not dataset:671 return None672 if self.data_source_type == "database":673 return json.loads(self.keyword_table, cls=SetDecoder) if self.keyword_table else None674 else:675 file_key = "keyword_files/" + dataset.tenant_id + "/" + self.dataset_id + ".txt"676 try:677 keyword_table_text = storage.load_once(file_key)678 if keyword_table_text:679 return json.loads(keyword_table_text.decode("utf-8"), cls=SetDecoder)680 return None681 except Exception as e:682 logging.exception(str(e))683 return None684 685 686class Embedding(db.Model):687 __tablename__ = "embeddings"688 __table_args__ = (689 db.PrimaryKeyConstraint("id", name="embedding_pkey"),690 db.UniqueConstraint("model_name", "hash", "provider_name", name="embedding_hash_idx"),691 db.Index("created_at_idx", "created_at"),692 )693 694 id = db.Column(StringUUID, primary_key=True, server_default=db.text("uuid_generate_v4()"))695 model_name = db.Column(696 db.String(255), nullable=False, server_default=db.text("'text-embedding-ada-002'::character varying")697 )698 hash = db.Column(db.String(64), nullable=False)699 embedding = db.Column(db.LargeBinary, nullable=False)700 created_at = db.Column(db.DateTime, nullable=False, server_default=db.text("CURRENT_TIMESTAMP(0)"))701 provider_name = db.Column(db.String(255), nullable=False, server_default=db.text("''::character varying"))702 703 def set_embedding(self, embedding_data: list[float]):704 self.embedding = pickle.dumps(embedding_data, protocol=pickle.HIGHEST_PROTOCOL)705 706 def get_embedding(self) -> list[float]:707 return pickle.loads(self.embedding)708 709 710class DatasetCollectionBinding(db.Model):711 __tablename__ = "dataset_collection_bindings"712 __table_args__ = (713 db.PrimaryKeyConstraint("id", name="dataset_collection_bindings_pkey"),714 db.Index("provider_model_name_idx", "provider_name", "model_name"),715 )716 717 id = db.Column(StringUUID, primary_key=True, server_default=db.text("uuid_generate_v4()"))718 provider_name = db.Column(db.String(40), nullable=False)719 model_name = db.Column(db.String(255), nullable=False)720 type = db.Column(db.String(40), server_default=db.text("'dataset'::character varying"), nullable=False)721 collection_name = db.Column(db.String(64), nullable=False)722 created_at = db.Column(db.DateTime, nullable=False, server_default=db.text("CURRENT_TIMESTAMP(0)"))723 724 725class TidbAuthBinding(db.Model):726 __tablename__ = "tidb_auth_bindings"727 __table_args__ = (728 db.PrimaryKeyConstraint("id", name="tidb_auth_bindings_pkey"),729 db.Index("tidb_auth_bindings_tenant_idx", "tenant_id"),730 db.Index("tidb_auth_bindings_active_idx", "active"),731 db.Index("tidb_auth_bindings_created_at_idx", "created_at"),732 db.Index("tidb_auth_bindings_status_idx", "status"),733 )734 id = db.Column(StringUUID, primary_key=True, server_default=db.text("uuid_generate_v4()"))735 tenant_id = db.Column(StringUUID, nullable=True)736 cluster_id = db.Column(db.String(255), nullable=False)737 cluster_name = db.Column(db.String(255), nullable=False)738 active = db.Column(db.Boolean, nullable=False, server_default=db.text("false"))739 status = db.Column(db.String(255), nullable=False, server_default=db.text("CREATING"))740 account = db.Column(db.String(255), nullable=False)741 password = db.Column(db.String(255), nullable=False)742 created_at = db.Column(db.DateTime, nullable=False, server_default=db.text("CURRENT_TIMESTAMP(0)"))743 744 745class Whitelist(db.Model):746 __tablename__ = "whitelists"747 __table_args__ = (748 db.PrimaryKeyConstraint("id", name="whitelists_pkey"),749 db.Index("whitelists_tenant_idx", "tenant_id"),750 )751 id = db.Column(StringUUID, primary_key=True, server_default=db.text("uuid_generate_v4()"))752 tenant_id = db.Column(StringUUID, nullable=True)753 category = db.Column(db.String(255), nullable=False)754 created_at = db.Column(db.DateTime, nullable=False, server_default=db.text("CURRENT_TIMESTAMP(0)"))755 756 757class DatasetPermission(db.Model):758 __tablename__ = "dataset_permissions"759 __table_args__ = (760 db.PrimaryKeyConstraint("id", name="dataset_permission_pkey"),761 db.Index("idx_dataset_permissions_dataset_id", "dataset_id"),762 db.Index("idx_dataset_permissions_account_id", "account_id"),763 db.Index("idx_dataset_permissions_tenant_id", "tenant_id"),764 )765 766 id = db.Column(StringUUID, server_default=db.text("uuid_generate_v4()"), primary_key=True)767 dataset_id = db.Column(StringUUID, nullable=False)768 account_id = db.Column(StringUUID, nullable=False)769 tenant_id = db.Column(StringUUID, nullable=False)770 has_permission = db.Column(db.Boolean, nullable=False, server_default=db.text("true"))771 created_at = db.Column(db.DateTime, nullable=False, server_default=db.text("CURRENT_TIMESTAMP(0)"))772 773 774class ExternalKnowledgeApis(db.Model):775 __tablename__ = "external_knowledge_apis"776 __table_args__ = (777 db.PrimaryKeyConstraint("id", name="external_knowledge_apis_pkey"),778 db.Index("external_knowledge_apis_tenant_idx", "tenant_id"),779 db.Index("external_knowledge_apis_name_idx", "name"),780 )781 782 id = db.Column(StringUUID, nullable=False, server_default=db.text("uuid_generate_v4()"))783 name = db.Column(db.String(255), nullable=False)784 description = db.Column(db.String(255), nullable=False)785 tenant_id = db.Column(StringUUID, nullable=False)786 settings = db.Column(db.Text, nullable=True)787 created_by = db.Column(StringUUID, nullable=False)788 created_at = db.Column(db.DateTime, nullable=False, server_default=db.text("CURRENT_TIMESTAMP(0)"))789 updated_by = db.Column(StringUUID, nullable=True)790 updated_at = db.Column(db.DateTime, nullable=False, server_default=db.text("CURRENT_TIMESTAMP(0)"))791 792 def to_dict(self):793 return {794 "id": self.id,795 "tenant_id": self.tenant_id,796 "name": self.name,797 "description": self.description,798 "settings": self.settings_dict,799 "dataset_bindings": self.dataset_bindings,800 "created_by": self.created_by,801 "created_at": self.created_at.isoformat(),802 }803 804 @property805 def settings_dict(self):806 try:807 return json.loads(self.settings) if self.settings else None808 except JSONDecodeError:809 return None810 811 @property812 def dataset_bindings(self):813 external_knowledge_bindings = (814 db.session.query(ExternalKnowledgeBindings)815 .filter(ExternalKnowledgeBindings.external_knowledge_api_id == self.id)816 .all()817 )818 dataset_ids = [binding.dataset_id for binding in external_knowledge_bindings]819 datasets = db.session.query(Dataset).filter(Dataset.id.in_(dataset_ids)).all()820 dataset_bindings = []821 for dataset in datasets:822 dataset_bindings.append({"id": dataset.id, "name": dataset.name})823 824 return dataset_bindings825 826 827class ExternalKnowledgeBindings(db.Model):828 __tablename__ = "external_knowledge_bindings"829 __table_args__ = (830 db.PrimaryKeyConstraint("id", name="external_knowledge_bindings_pkey"),831 db.Index("external_knowledge_bindings_tenant_idx", "tenant_id"),832 db.Index("external_knowledge_bindings_dataset_idx", "dataset_id"),833 db.Index("external_knowledge_bindings_external_knowledge_idx", "external_knowledge_id"),834 db.Index("external_knowledge_bindings_external_knowledge_api_idx", "external_knowledge_api_id"),835 )836 837 id = db.Column(StringUUID, nullable=False, server_default=db.text("uuid_generate_v4()"))838 tenant_id = db.Column(StringUUID, nullable=False)839 external_knowledge_api_id = db.Column(StringUUID, nullable=False)840 dataset_id = db.Column(StringUUID, nullable=False)841 external_knowledge_id = db.Column(db.Text, nullable=False)842 created_by = db.Column(StringUUID, nullable=False)843 created_at = db.Column(db.DateTime, nullable=False, server_default=db.text("CURRENT_TIMESTAMP(0)"))844 updated_by = db.Column(StringUUID, nullable=True)845 updated_at = db.Column(db.DateTime, nullable=False, server_default=db.text("CURRENT_TIMESTAMP(0)"))846 