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