Team Ai
Apppublic

Underground-Digital/Workflow-Engine

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
data_source.py269 linesDownload Raw Back to datasets
1import datetime2import json3 4from flask import request5from flask_login import current_user6from flask_restful import Resource, marshal_with, reqparse7from werkzeug.exceptions import NotFound8 9from controllers.console import api10from controllers.console.wraps import account_initialization_required, setup_required11from core.indexing_runner import IndexingRunner12from core.rag.extractor.entity.extract_setting import ExtractSetting13from core.rag.extractor.notion_extractor import NotionExtractor14from extensions.ext_database import db15from fields.data_source_fields import integrate_list_fields, integrate_notion_info_list_fields16from libs.login import login_required17from models import DataSourceOauthBinding, Document18from services.dataset_service import DatasetService, DocumentService19from tasks.document_indexing_sync_task import document_indexing_sync_task20 21 22class DataSourceApi(Resource):23    @setup_required24    @login_required25    @account_initialization_required26    @marshal_with(integrate_list_fields)27    def get(self):28        # get workspace data source integrates29        data_source_integrates = (30            db.session.query(DataSourceOauthBinding)31            .filter(32                DataSourceOauthBinding.tenant_id == current_user.current_tenant_id,33                DataSourceOauthBinding.disabled == False,34            )35            .all()36        )37 38        base_url = request.url_root.rstrip("/")39        data_source_oauth_base_path = "/console/api/oauth/data-source"40        providers = ["notion"]41 42        integrate_data = []43        for provider in providers:44            # existing_integrate = next((ai for ai in data_source_integrates if ai.provider == provider), None)45            existing_integrates = filter(lambda item: item.provider == provider, data_source_integrates)46            if existing_integrates:47                for existing_integrate in list(existing_integrates):48                    integrate_data.append(49                        {50                            "id": existing_integrate.id,51                            "provider": provider,52                            "created_at": existing_integrate.created_at,53                            "is_bound": True,54                            "disabled": existing_integrate.disabled,55                            "source_info": existing_integrate.source_info,56                            "link": f"{base_url}{data_source_oauth_base_path}/{provider}",57                        }58                    )59            else:60                integrate_data.append(61                    {62                        "id": None,63                        "provider": provider,64                        "created_at": None,65                        "source_info": None,66                        "is_bound": False,67                        "disabled": None,68                        "link": f"{base_url}{data_source_oauth_base_path}/{provider}",69                    }70                )71        return {"data": integrate_data}, 20072 73    @setup_required74    @login_required75    @account_initialization_required76    def patch(self, binding_id, action):77        binding_id = str(binding_id)78        action = str(action)79        data_source_binding = DataSourceOauthBinding.query.filter_by(id=binding_id).first()80        if data_source_binding is None:81            raise NotFound("Data source binding not found.")82        # enable binding83        if action == "enable":84            if data_source_binding.disabled:85                data_source_binding.disabled = False86                data_source_binding.updated_at = datetime.datetime.now(datetime.timezone.utc).replace(tzinfo=None)87                db.session.add(data_source_binding)88                db.session.commit()89            else:90                raise ValueError("Data source is not disabled.")91        # disable binding92        if action == "disable":93            if not data_source_binding.disabled:94                data_source_binding.disabled = True95                data_source_binding.updated_at = datetime.datetime.now(datetime.timezone.utc).replace(tzinfo=None)96                db.session.add(data_source_binding)97                db.session.commit()98            else:99                raise ValueError("Data source is disabled.")100        return {"result": "success"}, 200101 102 103class DataSourceNotionListApi(Resource):104    @setup_required105    @login_required106    @account_initialization_required107    @marshal_with(integrate_notion_info_list_fields)108    def get(self):109        dataset_id = request.args.get("dataset_id", default=None, type=str)110        exist_page_ids = []111        # import notion in the exist dataset112        if dataset_id:113            dataset = DatasetService.get_dataset(dataset_id)114            if not dataset:115                raise NotFound("Dataset not found.")116            if dataset.data_source_type != "notion_import":117                raise ValueError("Dataset is not notion type.")118            documents = Document.query.filter_by(119                dataset_id=dataset_id,120                tenant_id=current_user.current_tenant_id,121                data_source_type="notion_import",122                enabled=True,123            ).all()124            if documents:125                for document in documents:126                    data_source_info = json.loads(document.data_source_info)127                    exist_page_ids.append(data_source_info["notion_page_id"])128        # get all authorized pages129        data_source_bindings = DataSourceOauthBinding.query.filter_by(130            tenant_id=current_user.current_tenant_id, provider="notion", disabled=False131        ).all()132        if not data_source_bindings:133            return {"notion_info": []}, 200134        pre_import_info_list = []135        for data_source_binding in data_source_bindings:136            source_info = data_source_binding.source_info137            pages = source_info["pages"]138            # Filter out already bound pages139            for page in pages:140                if page["page_id"] in exist_page_ids:141                    page["is_bound"] = True142                else:143                    page["is_bound"] = False144            pre_import_info = {145                "workspace_name": source_info["workspace_name"],146                "workspace_icon": source_info["workspace_icon"],147                "workspace_id": source_info["workspace_id"],148                "pages": pages,149            }150            pre_import_info_list.append(pre_import_info)151        return {"notion_info": pre_import_info_list}, 200152 153 154class DataSourceNotionApi(Resource):155    @setup_required156    @login_required157    @account_initialization_required158    def get(self, workspace_id, page_id, page_type):159        workspace_id = str(workspace_id)160        page_id = str(page_id)161        data_source_binding = DataSourceOauthBinding.query.filter(162            db.and_(163                DataSourceOauthBinding.tenant_id == current_user.current_tenant_id,164                DataSourceOauthBinding.provider == "notion",165                DataSourceOauthBinding.disabled == False,166                DataSourceOauthBinding.source_info["workspace_id"] == f'"{workspace_id}"',167            )168        ).first()169        if not data_source_binding:170            raise NotFound("Data source binding not found.")171 172        extractor = NotionExtractor(173            notion_workspace_id=workspace_id,174            notion_obj_id=page_id,175            notion_page_type=page_type,176            notion_access_token=data_source_binding.access_token,177            tenant_id=current_user.current_tenant_id,178        )179 180        text_docs = extractor.extract()181        return {"content": "\n".join([doc.page_content for doc in text_docs])}, 200182 183    @setup_required184    @login_required185    @account_initialization_required186    def post(self):187        parser = reqparse.RequestParser()188        parser.add_argument("notion_info_list", type=list, required=True, nullable=True, location="json")189        parser.add_argument("process_rule", type=dict, required=True, nullable=True, location="json")190        parser.add_argument("doc_form", type=str, default="text_model", required=False, nullable=False, location="json")191        parser.add_argument(192            "doc_language", type=str, default="English", required=False, nullable=False, location="json"193        )194        args = parser.parse_args()195        # validate args196        DocumentService.estimate_args_validate(args)197        notion_info_list = args["notion_info_list"]198        extract_settings = []199        for notion_info in notion_info_list:200            workspace_id = notion_info["workspace_id"]201            for page in notion_info["pages"]:202                extract_setting = ExtractSetting(203                    datasource_type="notion_import",204                    notion_info={205                        "notion_workspace_id": workspace_id,206                        "notion_obj_id": page["page_id"],207                        "notion_page_type": page["type"],208                        "tenant_id": current_user.current_tenant_id,209                    },210                    document_model=args["doc_form"],211                )212                extract_settings.append(extract_setting)213        indexing_runner = IndexingRunner()214        response = indexing_runner.indexing_estimate(215            current_user.current_tenant_id,216            extract_settings,217            args["process_rule"],218            args["doc_form"],219            args["doc_language"],220        )221        return response, 200222 223 224class DataSourceNotionDatasetSyncApi(Resource):225    @setup_required226    @login_required227    @account_initialization_required228    def get(self, dataset_id):229        dataset_id_str = str(dataset_id)230        dataset = DatasetService.get_dataset(dataset_id_str)231        if dataset is None:232            raise NotFound("Dataset not found.")233 234        documents = DocumentService.get_document_by_dataset_id(dataset_id_str)235        for document in documents:236            document_indexing_sync_task.delay(dataset_id_str, document.id)237        return 200238 239 240class DataSourceNotionDocumentSyncApi(Resource):241    @setup_required242    @login_required243    @account_initialization_required244    def get(self, dataset_id, document_id):245        dataset_id_str = str(dataset_id)246        document_id_str = str(document_id)247        dataset = DatasetService.get_dataset(dataset_id_str)248        if dataset is None:249            raise NotFound("Dataset not found.")250 251        document = DocumentService.get_document(dataset_id_str, document_id_str)252        if document is None:253            raise NotFound("Document not found.")254        document_indexing_sync_task.delay(dataset_id_str, document_id_str)255        return 200256 257 258api.add_resource(DataSourceApi, "/data-source/integrates", "/data-source/integrates/<uuid:binding_id>/<string:action>")259api.add_resource(DataSourceNotionListApi, "/notion/pre-import/pages")260api.add_resource(261    DataSourceNotionApi,262    "/notion/workspaces/<uuid:workspace_id>/pages/<uuid:page_id>/<string:page_type>/preview",263    "/datasets/notion-indexing-estimate",264)265api.add_resource(DataSourceNotionDatasetSyncApi, "/datasets/<uuid:dataset_id>/notion/sync")266api.add_resource(267    DataSourceNotionDocumentSyncApi, "/datasets/<uuid:dataset_id>/documents/<uuid:document_id>/notion/sync"268)269