Team Ai
Apppublic

Underground-Digital/Workflow-Engine

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
dataset.py846 linesDownload Raw Back to models
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