Team Ai
Apppublic

Underground-Digital/Workflow-Engine

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
datasets_segments.py427 linesDownload Raw Back to datasets
1import uuid2from datetime import datetime, timezone3 4import pandas as pd5from flask import request6from flask_login import current_user7from flask_restful import Resource, marshal, reqparse8from werkzeug.exceptions import Forbidden, NotFound9 10import services11from controllers.console import api12from controllers.console.app.error import ProviderNotInitializeError13from controllers.console.datasets.error import InvalidActionError, NoFileUploadedError, TooManyFilesError14from controllers.console.wraps import (15    account_initialization_required,16    cloud_edition_billing_knowledge_limit_check,17    cloud_edition_billing_resource_check,18    setup_required,19)20from core.errors.error import LLMBadRequestError, ProviderTokenNotInitError21from core.model_manager import ModelManager22from core.model_runtime.entities.model_entities import ModelType23from extensions.ext_database import db24from extensions.ext_redis import redis_client25from fields.segment_fields import segment_fields26from libs.login import login_required27from models import DocumentSegment28from services.dataset_service import DatasetService, DocumentService, SegmentService29from tasks.batch_create_segment_to_index_task import batch_create_segment_to_index_task30from tasks.disable_segment_from_index_task import disable_segment_from_index_task31from tasks.enable_segment_to_index_task import enable_segment_to_index_task32 33 34class DatasetDocumentSegmentListApi(Resource):35    @setup_required36    @login_required37    @account_initialization_required38    def get(self, dataset_id, document_id):39        dataset_id = str(dataset_id)40        document_id = str(document_id)41        dataset = DatasetService.get_dataset(dataset_id)42        if not dataset:43            raise NotFound("Dataset not found.")44 45        try:46            DatasetService.check_dataset_permission(dataset, current_user)47        except services.errors.account.NoPermissionError as e:48            raise Forbidden(str(e))49 50        document = DocumentService.get_document(dataset_id, document_id)51 52        if not document:53            raise NotFound("Document not found.")54 55        parser = reqparse.RequestParser()56        parser.add_argument("last_id", type=str, default=None, location="args")57        parser.add_argument("limit", type=int, default=20, location="args")58        parser.add_argument("status", type=str, action="append", default=[], location="args")59        parser.add_argument("hit_count_gte", type=int, default=None, location="args")60        parser.add_argument("enabled", type=str, default="all", location="args")61        parser.add_argument("keyword", type=str, default=None, location="args")62        args = parser.parse_args()63 64        last_id = args["last_id"]65        limit = min(args["limit"], 100)66        status_list = args["status"]67        hit_count_gte = args["hit_count_gte"]68        keyword = args["keyword"]69 70        query = DocumentSegment.query.filter(71            DocumentSegment.document_id == str(document_id), DocumentSegment.tenant_id == current_user.current_tenant_id72        )73 74        if last_id is not None:75            last_segment = db.session.get(DocumentSegment, str(last_id))76            if last_segment:77                query = query.filter(DocumentSegment.position > last_segment.position)78            else:79                return {"data": [], "has_more": False, "limit": limit}, 20080 81        if status_list:82            query = query.filter(DocumentSegment.status.in_(status_list))83 84        if hit_count_gte is not None:85            query = query.filter(DocumentSegment.hit_count >= hit_count_gte)86 87        if keyword:88            query = query.where(DocumentSegment.content.ilike(f"%{keyword}%"))89 90        if args["enabled"].lower() != "all":91            if args["enabled"].lower() == "true":92                query = query.filter(DocumentSegment.enabled == True)93            elif args["enabled"].lower() == "false":94                query = query.filter(DocumentSegment.enabled == False)95 96        total = query.count()97        segments = query.order_by(DocumentSegment.position).limit(limit + 1).all()98 99        has_more = False100        if len(segments) > limit:101            has_more = True102            segments = segments[:-1]103 104        return {105            "data": marshal(segments, segment_fields),106            "doc_form": document.doc_form,107            "has_more": has_more,108            "limit": limit,109            "total": total,110        }, 200111 112 113class DatasetDocumentSegmentApi(Resource):114    @setup_required115    @login_required116    @account_initialization_required117    @cloud_edition_billing_resource_check("vector_space")118    def patch(self, dataset_id, segment_id, action):119        dataset_id = str(dataset_id)120        dataset = DatasetService.get_dataset(dataset_id)121        if not dataset:122            raise NotFound("Dataset not found.")123        # check user's model setting124        DatasetService.check_dataset_model_setting(dataset)125        # The role of the current user in the ta table must be admin, owner, or editor126        if not current_user.is_editor:127            raise Forbidden()128 129        try:130            DatasetService.check_dataset_permission(dataset, current_user)131        except services.errors.account.NoPermissionError as e:132            raise Forbidden(str(e))133        if dataset.indexing_technique == "high_quality":134            # check embedding model setting135            try:136                model_manager = ModelManager()137                model_manager.get_model_instance(138                    tenant_id=current_user.current_tenant_id,139                    provider=dataset.embedding_model_provider,140                    model_type=ModelType.TEXT_EMBEDDING,141                    model=dataset.embedding_model,142                )143            except LLMBadRequestError:144                raise ProviderNotInitializeError(145                    "No Embedding Model available. Please configure a valid provider "146                    "in the Settings -> Model Provider."147                )148            except ProviderTokenNotInitError as ex:149                raise ProviderNotInitializeError(ex.description)150 151        segment = DocumentSegment.query.filter(152            DocumentSegment.id == str(segment_id), DocumentSegment.tenant_id == current_user.current_tenant_id153        ).first()154 155        if not segment:156            raise NotFound("Segment not found.")157 158        if segment.status != "completed":159            raise NotFound("Segment is not completed, enable or disable function is not allowed")160 161        document_indexing_cache_key = "document_{}_indexing".format(segment.document_id)162        cache_result = redis_client.get(document_indexing_cache_key)163        if cache_result is not None:164            raise InvalidActionError("Document is being indexed, please try again later")165 166        indexing_cache_key = "segment_{}_indexing".format(segment.id)167        cache_result = redis_client.get(indexing_cache_key)168        if cache_result is not None:169            raise InvalidActionError("Segment is being indexed, please try again later")170 171        if action == "enable":172            if segment.enabled:173                raise InvalidActionError("Segment is already enabled.")174 175            segment.enabled = True176            segment.disabled_at = None177            segment.disabled_by = None178            db.session.commit()179 180            # Set cache to prevent indexing the same segment multiple times181            redis_client.setex(indexing_cache_key, 600, 1)182 183            enable_segment_to_index_task.delay(segment.id)184 185            return {"result": "success"}, 200186        elif action == "disable":187            if not segment.enabled:188                raise InvalidActionError("Segment is already disabled.")189 190            segment.enabled = False191            segment.disabled_at = datetime.now(timezone.utc).replace(tzinfo=None)192            segment.disabled_by = current_user.id193            db.session.commit()194 195            # Set cache to prevent indexing the same segment multiple times196            redis_client.setex(indexing_cache_key, 600, 1)197 198            disable_segment_from_index_task.delay(segment.id)199 200            return {"result": "success"}, 200201        else:202            raise InvalidActionError()203 204 205class DatasetDocumentSegmentAddApi(Resource):206    @setup_required207    @login_required208    @account_initialization_required209    @cloud_edition_billing_resource_check("vector_space")210    @cloud_edition_billing_knowledge_limit_check("add_segment")211    def post(self, dataset_id, document_id):212        # check dataset213        dataset_id = str(dataset_id)214        dataset = DatasetService.get_dataset(dataset_id)215        if not dataset:216            raise NotFound("Dataset not found.")217        # check document218        document_id = str(document_id)219        document = DocumentService.get_document(dataset_id, document_id)220        if not document:221            raise NotFound("Document not found.")222        if not current_user.is_editor:223            raise Forbidden()224        # check embedding model setting225        if dataset.indexing_technique == "high_quality":226            try:227                model_manager = ModelManager()228                model_manager.get_model_instance(229                    tenant_id=current_user.current_tenant_id,230                    provider=dataset.embedding_model_provider,231                    model_type=ModelType.TEXT_EMBEDDING,232                    model=dataset.embedding_model,233                )234            except LLMBadRequestError:235                raise ProviderNotInitializeError(236                    "No Embedding Model available. Please configure a valid provider "237                    "in the Settings -> Model Provider."238                )239            except ProviderTokenNotInitError as ex:240                raise ProviderNotInitializeError(ex.description)241        try:242            DatasetService.check_dataset_permission(dataset, current_user)243        except services.errors.account.NoPermissionError as e:244            raise Forbidden(str(e))245        # validate args246        parser = reqparse.RequestParser()247        parser.add_argument("content", type=str, required=True, nullable=False, location="json")248        parser.add_argument("answer", type=str, required=False, nullable=True, location="json")249        parser.add_argument("keywords", type=list, required=False, nullable=True, location="json")250        args = parser.parse_args()251        SegmentService.segment_create_args_validate(args, document)252        segment = SegmentService.create_segment(args, document, dataset)253        return {"data": marshal(segment, segment_fields), "doc_form": document.doc_form}, 200254 255 256class DatasetDocumentSegmentUpdateApi(Resource):257    @setup_required258    @login_required259    @account_initialization_required260    @cloud_edition_billing_resource_check("vector_space")261    def patch(self, dataset_id, document_id, segment_id):262        # check dataset263        dataset_id = str(dataset_id)264        dataset = DatasetService.get_dataset(dataset_id)265        if not dataset:266            raise NotFound("Dataset not found.")267        # check user's model setting268        DatasetService.check_dataset_model_setting(dataset)269        # check document270        document_id = str(document_id)271        document = DocumentService.get_document(dataset_id, document_id)272        if not document:273            raise NotFound("Document not found.")274        if dataset.indexing_technique == "high_quality":275            # check embedding model setting276            try:277                model_manager = ModelManager()278                model_manager.get_model_instance(279                    tenant_id=current_user.current_tenant_id,280                    provider=dataset.embedding_model_provider,281                    model_type=ModelType.TEXT_EMBEDDING,282                    model=dataset.embedding_model,283                )284            except LLMBadRequestError:285                raise ProviderNotInitializeError(286                    "No Embedding Model available. Please configure a valid provider "287                    "in the Settings -> Model Provider."288                )289            except ProviderTokenNotInitError as ex:290                raise ProviderNotInitializeError(ex.description)291            # check segment292        segment_id = str(segment_id)293        segment = DocumentSegment.query.filter(294            DocumentSegment.id == str(segment_id), DocumentSegment.tenant_id == current_user.current_tenant_id295        ).first()296        if not segment:297            raise NotFound("Segment not found.")298        # The role of the current user in the ta table must be admin, owner, or editor299        if not current_user.is_editor:300            raise Forbidden()301        try:302            DatasetService.check_dataset_permission(dataset, current_user)303        except services.errors.account.NoPermissionError as e:304            raise Forbidden(str(e))305        # validate args306        parser = reqparse.RequestParser()307        parser.add_argument("content", type=str, required=True, nullable=False, location="json")308        parser.add_argument("answer", type=str, required=False, nullable=True, location="json")309        parser.add_argument("keywords", type=list, required=False, nullable=True, location="json")310        args = parser.parse_args()311        SegmentService.segment_create_args_validate(args, document)312        segment = SegmentService.update_segment(args, segment, document, dataset)313        return {"data": marshal(segment, segment_fields), "doc_form": document.doc_form}, 200314 315    @setup_required316    @login_required317    @account_initialization_required318    def delete(self, dataset_id, document_id, segment_id):319        # check dataset320        dataset_id = str(dataset_id)321        dataset = DatasetService.get_dataset(dataset_id)322        if not dataset:323            raise NotFound("Dataset not found.")324        # check user's model setting325        DatasetService.check_dataset_model_setting(dataset)326        # check document327        document_id = str(document_id)328        document = DocumentService.get_document(dataset_id, document_id)329        if not document:330            raise NotFound("Document not found.")331        # check segment332        segment_id = str(segment_id)333        segment = DocumentSegment.query.filter(334            DocumentSegment.id == str(segment_id), DocumentSegment.tenant_id == current_user.current_tenant_id335        ).first()336        if not segment:337            raise NotFound("Segment not found.")338        # The role of the current user in the ta table must be admin or owner339        if not current_user.is_editor:340            raise Forbidden()341        try:342            DatasetService.check_dataset_permission(dataset, current_user)343        except services.errors.account.NoPermissionError as e:344            raise Forbidden(str(e))345        SegmentService.delete_segment(segment, document, dataset)346        return {"result": "success"}, 200347 348 349class DatasetDocumentSegmentBatchImportApi(Resource):350    @setup_required351    @login_required352    @account_initialization_required353    @cloud_edition_billing_resource_check("vector_space")354    @cloud_edition_billing_knowledge_limit_check("add_segment")355    def post(self, dataset_id, document_id):356        # check dataset357        dataset_id = str(dataset_id)358        dataset = DatasetService.get_dataset(dataset_id)359        if not dataset:360            raise NotFound("Dataset not found.")361        # check document362        document_id = str(document_id)363        document = DocumentService.get_document(dataset_id, document_id)364        if not document:365            raise NotFound("Document not found.")366        # get file from request367        file = request.files["file"]368        # check file369        if "file" not in request.files:370            raise NoFileUploadedError()371 372        if len(request.files) > 1:373            raise TooManyFilesError()374        # check file type375        if not file.filename.endswith(".csv"):376            raise ValueError("Invalid file type. Only CSV files are allowed")377 378        try:379            # Skip the first row380            df = pd.read_csv(file)381            result = []382            for index, row in df.iterrows():383                if document.doc_form == "qa_model":384                    data = {"content": row[0], "answer": row[1]}385                else:386                    data = {"content": row[0]}387                result.append(data)388            if len(result) == 0:389                raise ValueError("The CSV file is empty.")390            # async job391            job_id = str(uuid.uuid4())392            indexing_cache_key = "segment_batch_import_{}".format(str(job_id))393            # send batch add segments task394            redis_client.setnx(indexing_cache_key, "waiting")395            batch_create_segment_to_index_task.delay(396                str(job_id), result, dataset_id, document_id, current_user.current_tenant_id, current_user.id397            )398        except Exception as e:399            return {"error": str(e)}, 500400        return {"job_id": job_id, "job_status": "waiting"}, 200401 402    @setup_required403    @login_required404    @account_initialization_required405    def get(self, job_id):406        job_id = str(job_id)407        indexing_cache_key = "segment_batch_import_{}".format(job_id)408        cache_result = redis_client.get(indexing_cache_key)409        if cache_result is None:410            raise ValueError("The job is not exist.")411 412        return {"job_id": job_id, "job_status": cache_result.decode()}, 200413 414 415api.add_resource(DatasetDocumentSegmentListApi, "/datasets/<uuid:dataset_id>/documents/<uuid:document_id>/segments")416api.add_resource(DatasetDocumentSegmentApi, "/datasets/<uuid:dataset_id>/segments/<uuid:segment_id>/<string:action>")417api.add_resource(DatasetDocumentSegmentAddApi, "/datasets/<uuid:dataset_id>/documents/<uuid:document_id>/segment")418api.add_resource(419    DatasetDocumentSegmentUpdateApi,420    "/datasets/<uuid:dataset_id>/documents/<uuid:document_id>/segments/<uuid:segment_id>",421)422api.add_resource(423    DatasetDocumentSegmentBatchImportApi,424    "/datasets/<uuid:dataset_id>/documents/<uuid:document_id>/segments/batch_import",425    "/datasets/batch_import_status/<uuid:job_id>",426)427