diff --git a/apps/application/serializers/application_chat.py b/apps/application/serializers/application_chat.py index dd9464a9ba5..0236b178eb5 100644 --- a/apps/application/serializers/application_chat.py +++ b/apps/application/serializers/application_chat.py @@ -1,11 +1,12 @@ # coding=utf-8 """ - @project: MaxKB - @Author:虎虎 - @file: application_chat.py - @date:2025/6/10 11:06 - @desc: +@project: MaxKB +@Author:虎虎 +@file: application_chat.py +@date:2025/6/10 11:06 +@desc: """ + import datetime import os import re @@ -45,80 +46,90 @@ class ApplicationChatResponseSerializers(serializers.Serializer): class ApplicationChatRecordExportRequest(serializers.Serializer): - select_ids = serializers.ListField(required=True, label=_("Chat ID List"), - child=serializers.UUIDField(required=True, label=_("Chat ID"))) + select_ids = serializers.ListField( + required=True, label=_("Chat ID List"), child=serializers.UUIDField(required=True, label=_("Chat ID")) + ) class ApplicationChatQuerySerializers(serializers.Serializer): workspace_id = serializers.CharField(required=False, allow_null=True, allow_blank=True, label=_("Workspace ID")) abstract = serializers.CharField(required=False, allow_blank=True, allow_null=True, label=_("summary")) username = serializers.CharField(required=False, allow_blank=True, allow_null=True, label=_("username")) - start_time = serializers.DateField(format='%Y-%m-%d', label=_("Start time")) - end_time = serializers.DateField(format='%Y-%m-%d', label=_("End time")) + start_time = serializers.DateField(format="%Y-%m-%d", label=_("Start time")) + end_time = serializers.DateField(format="%Y-%m-%d", label=_("End time")) application_id = serializers.UUIDField(required=True, label=_("Application ID")) - min_star = serializers.IntegerField(required=False, min_value=0, - label=_("Minimum number of likes")) - min_trample = serializers.IntegerField(required=False, min_value=0, - label=_("Minimum number of clicks")) - comparer = serializers.CharField(required=False, label=_("Comparator"), validators=[ - validators.RegexValidator(regex=re.compile("^and|or$"), - message=_("Only supports and|or"), code=500) - ]) + min_star = serializers.IntegerField(required=False, min_value=0, label=_("Minimum number of likes")) + min_trample = serializers.IntegerField(required=False, min_value=0, label=_("Minimum number of clicks")) + comparer = serializers.CharField( + required=False, + label=_("Comparator"), + validators=[ + validators.RegexValidator(regex=re.compile("^and|or$"), message=_("Only supports and|or"), code=500) + ], + ) def is_valid(self, *, raise_exception=False): super().is_valid(raise_exception=True) - workspace_id = self.data.get('workspace_id') - query_set = QuerySet(Application).filter(id=self.data.get('application_id')) + workspace_id = self.data.get("workspace_id") + query_set = QuerySet(Application).filter(id=self.data.get("application_id")) if workspace_id: query_set = query_set.filter(workspace_id=workspace_id) if not query_set.exists(): - raise AppApiException(500, _('Application id does not exist')) + raise AppApiException(500, _("Application id does not exist")) def get_end_time(self): - d = datetime.datetime.strptime(self.data.get('end_time'), '%Y-%m-%d').date() + d = datetime.datetime.strptime(self.data.get("end_time"), "%Y-%m-%d").date() naive = datetime.datetime.combine(d, datetime.time.max) return timezone.make_aware(naive, timezone.get_default_timezone()) def get_start_time(self): - d = datetime.datetime.strptime(self.data.get('start_time'), '%Y-%m-%d').date() + d = datetime.datetime.strptime(self.data.get("start_time"), "%Y-%m-%d").date() naive = datetime.datetime.combine(d, datetime.time.min) return timezone.make_aware(naive, timezone.get_default_timezone()) def get_query_set(self, select_ids=None): end_time = self.get_end_time() start_time = self.get_start_time() - query_set = QuerySet(model=get_dynamics_model( - {'application_chat.application_id': models.CharField(), - 'application_chat.abstract': models.CharField(), - 'application_chat.asker': models.JSONField(), - "star_num": models.IntegerField(), - 'trample_num': models.IntegerField(), - 'comparer': models.CharField(), - 'application_chat.update_time': models.DateTimeField(), - 'application_chat.id': models.UUIDField(), - 'application_chat_record_temp.id': models.UUIDField()})) - - base_query_dict = {'application_chat.application_id': self.data.get("application_id"), - 'application_chat.update_time__gte': start_time, - 'application_chat.update_time__lte': end_time, - } - if 'abstract' in self.data and self.data.get('abstract') is not None: - base_query_dict['application_chat.abstract__icontains'] = self.data.get('abstract') - if 'username' in self.data and self.data.get('username') is not None: - base_query_dict['application_chat.asker__username__icontains'] = self.data.get('username') - + query_set = QuerySet( + model=get_dynamics_model( + { + "application_chat.application_id": models.CharField(), + "application_chat.abstract": models.CharField(), + "application_chat.asker": models.JSONField(), + "star_num": models.IntegerField(), + "trample_num": models.IntegerField(), + "comparer": models.CharField(), + "application_chat.update_time": models.DateTimeField(), + "application_chat.id": models.UUIDField(), + "application_chat_record_temp.id": models.UUIDField(), + } + ) + ) + + base_query_dict = { + "application_chat.application_id": self.data.get("application_id"), + "application_chat.update_time__gte": start_time, + "application_chat.update_time__lte": end_time, + } + if "abstract" in self.data and self.data.get("abstract") is not None: + base_query_dict["application_chat.abstract__icontains"] = self.data.get("abstract") if select_ids is not None and len(select_ids) > 0: - base_query_dict['application_chat.id__in'] = select_ids + base_query_dict["application_chat.id__in"] = select_ids base_condition = Q(**base_query_dict) + if "username" in self.data and self.data.get("username") is not None: + username = self.data.get("username") + base_condition = base_condition & ( + Q(**{"application_chat.asker__username__icontains": username}) + | Q(**{"application_chat.asker__nick_name__icontains": username}) + ) min_star_query = None min_trample_query = None - if 'min_star' in self.data and self.data.get('min_star') is not None: - min_star_query = Q(star_num__gte=self.data.get('min_star')) - if 'min_trample' in self.data and self.data.get('min_trample') is not None: - min_trample_query = Q(trample_num__gte=self.data.get('min_trample')) + if "min_star" in self.data and self.data.get("min_star") is not None: + min_star_query = Q(star_num__gte=self.data.get("min_star")) + if "min_trample" in self.data and self.data.get("min_trample") is not None: + min_trample_query = Q(trample_num__gte=self.data.get("min_trample")) if min_star_query is not None and min_trample_query is not None: - if self.data.get( - 'comparer') is not None and self.data.get('comparer') == 'or': + if self.data.get("comparer") is not None and self.data.get("comparer") == "or": condition = base_condition & (min_star_query | min_trample_query) else: condition = base_condition & (min_star_query & min_trample_query) @@ -129,71 +140,113 @@ def get_query_set(self, select_ids=None): else: condition = base_condition - return { - 'default_queryset': query_set.filter(condition).order_by("-application_chat.update_time") - } + return {"default_queryset": query_set.filter(condition).order_by("-application_chat.update_time")} def list(self, with_valid=True): if with_valid: self.is_valid(raise_exception=True) - return native_search(self.get_query_set(), select_string=get_file_content( - os.path.join(PROJECT_DIR, "apps", "application", 'sql', - ('list_application_chat_ee.sql' if ['PE', 'EE'].__contains__( - edition) else 'list_application_chat.sql'))), - with_table_name=False) + return native_search( + self.get_query_set(), + select_string=get_file_content( + os.path.join( + PROJECT_DIR, + "apps", + "application", + "sql", + ( + "list_application_chat_ee.sql" + if ["PE", "EE"].__contains__(edition) + else "list_application_chat.sql" + ), + ) + ), + with_table_name=False, + ) @staticmethod def paragraph_list_to_string(paragraph_list): return "\n**********\n".join( - [f"{paragraph.get('title')}:\n{paragraph.get('content')}" for paragraph in - paragraph_list] if paragraph_list is not None else '') + [f"{paragraph.get('title')}:\n{paragraph.get('content')}" for paragraph in paragraph_list] + if paragraph_list is not None + else "" + ) @staticmethod def to_row(row: Dict): - details = row.get('details') or {} - padding_problem_text = ' '.join((node.get("answer", "") or "") for key, node in details.items() if - node.get("type") == 'question-node') - search_dataset_node_list = [(key, node) for key, node in details.items() if - node.get("type") == 'search-dataset-node' or node.get( - "step_type") == 'search_step' or node.get("type") == 'search-knowledge-node'] - reference_paragraph_len = '\n'.join([str(len(node.get('paragraph_list', - []))) if key == 'search_step' else node.get( - 'name') + ':' + str( - len(node.get('paragraph_list', [])) if node.get('paragraph_list', []) is not None else '0') for - key, node in search_dataset_node_list]) - reference_paragraph = '\n----------\n'.join( - [ApplicationChatQuerySerializers.paragraph_list_to_string(node.get('paragraph_list', - [])) if key == 'search_step' else node.get( - 'name') + ':\n' + ApplicationChatQuerySerializers.paragraph_list_to_string(node.get('paragraph_list', - [])) for - key, node in search_dataset_node_list]) - improve_paragraph_list = row.get('improve_paragraph_list') or [] - vote_status_map = {'-1': '未投票', '0': '赞同', '1': '反对'} - vote_reason_map = {'accurate': gettext('accurate'), 'complete': gettext('complete'), - 'inaccurate': gettext('inaccurate'), 'incomplete': gettext('incomplete'), - 'other': gettext('Other'), } - return [str(row.get('chat_id')), row.get('abstract'), row.get('problem_text'), padding_problem_text, - row.get('answer_text'), vote_status_map.get(row.get('vote_status')), - vote_reason_map.get(row.get('vote_reason')), - row.get('vote_other_content'), - reference_paragraph_len, - reference_paragraph, - "\n".join([ + details = row.get("details") or {} + padding_problem_text = " ".join( + (node.get("answer", "") or "") for key, node in details.items() if node.get("type") == "question-node" + ) + search_dataset_node_list = [ + (key, node) + for key, node in details.items() + if node.get("type") == "search-dataset-node" + or node.get("step_type") == "search_step" + or node.get("type") == "search-knowledge-node" + ] + reference_paragraph_len = "\n".join( + [ + str(len(node.get("paragraph_list", []))) + if key == "search_step" + else node.get("name") + + ":" + + str(len(node.get("paragraph_list", [])) if node.get("paragraph_list", []) is not None else "0") + for key, node in search_dataset_node_list + ] + ) + reference_paragraph = "\n----------\n".join( + [ + ApplicationChatQuerySerializers.paragraph_list_to_string(node.get("paragraph_list", [])) + if key == "search_step" + else node.get("name") + + ":\n" + + ApplicationChatQuerySerializers.paragraph_list_to_string(node.get("paragraph_list", [])) + for key, node in search_dataset_node_list + ] + ) + improve_paragraph_list = row.get("improve_paragraph_list") or [] + vote_status_map = {"-1": "未投票", "0": "赞同", "1": "反对"} + vote_reason_map = { + "accurate": gettext("accurate"), + "complete": gettext("complete"), + "inaccurate": gettext("inaccurate"), + "incomplete": gettext("incomplete"), + "other": gettext("Other"), + } + return [ + str(row.get("chat_id")), + row.get("abstract"), + row.get("problem_text"), + padding_problem_text, + row.get("answer_text"), + vote_status_map.get(row.get("vote_status")), + vote_reason_map.get(row.get("vote_reason")), + row.get("vote_other_content"), + reference_paragraph_len, + reference_paragraph, + "\n".join( + [ f"{improve_paragraph_list[index].get('title')}\n{improve_paragraph_list[index].get('content')}" - for index in range(len(improve_paragraph_list))]), - row.get('asker').get('username'), - (row.get('message_tokens') or 0) + (row.get('answer_tokens') or 0), - row.get('ip_address') or '-', - get_source_display(row.get('source')), - row.get('run_time'), - str(row.get('create_time').astimezone(pytz.timezone(TIME_ZONE)).strftime('%Y-%m-%d %H:%M:%S') - if row.get('create_time') is not None else None)] + for index in range(len(improve_paragraph_list)) + ] + ), + row.get("asker").get("username"), + (row.get("message_tokens") or 0) + (row.get("answer_tokens") or 0), + row.get("ip_address") or "-", + get_source_display(row.get("source")), + row.get("run_time"), + str( + row.get("create_time").astimezone(pytz.timezone(TIME_ZONE)).strftime("%Y-%m-%d %H:%M:%S") + if row.get("create_time") is not None + else None + ), + ] @staticmethod def reset_value(value): if isinstance(value, str): - value = re.sub(ILLEGAL_CHARACTERS_RE, '', value) - if value.startswith(('=', '+', '-', '@')): + value = re.sub(ILLEGAL_CHARACTERS_RE, "", value) + if value.startswith(("=", "+", "-", "@")): value = "'" + value if isinstance(value, datetime.datetime): eastern = pytz.timezone(TIME_ZONE) @@ -208,31 +261,50 @@ def export(self, data, with_valid=True): def stream_response(): workbook = openpyxl.Workbook(write_only=True) - worksheet = workbook.create_sheet(title='Sheet1') + worksheet = workbook.create_sheet(title="Sheet1") current_page = 1 page_size = 500 - headers = [gettext('Conversation ID'), gettext('summary'), gettext('User Questions'), - gettext('Problem after optimization'), - gettext('answer'), gettext('User feedback'), gettext('Feedback reason'), - gettext('Other reason content'), - gettext('Reference segment number'), - gettext('Section title + content'), - gettext('Annotation'), gettext('User'), gettext('Consuming tokens'), - gettext('Ip Address'), gettext('source'), - gettext('Time consumed (s)'), - gettext('Question Time')] + headers = [ + gettext("Conversation ID"), + gettext("summary"), + gettext("User Questions"), + gettext("Problem after optimization"), + gettext("answer"), + gettext("User feedback"), + gettext("Feedback reason"), + gettext("Other reason content"), + gettext("Reference segment number"), + gettext("Section title + content"), + gettext("Annotation"), + gettext("User"), + gettext("Consuming tokens"), + gettext("Ip Address"), + gettext("source"), + gettext("Time consumed (s)"), + gettext("Question Time"), + ] worksheet.append(headers) - for data_list in native_page_handler(page_size, self.get_query_set(data.get('select_ids')), - primary_key='application_chat_record_temp.id', - primary_queryset='default_queryset', - get_primary_value=lambda item: item.get('id'), - select_string=get_file_content( - os.path.join(PROJECT_DIR, "apps", "application", 'sql', - ('export_application_chat_ee.sql' if ['PE', - 'EE'].__contains__( - edition) else 'export_application_chat.sql'))), - with_table_name=False): - + for data_list in native_page_handler( + page_size, + self.get_query_set(data.get("select_ids")), + primary_key="application_chat_record_temp.id", + primary_queryset="default_queryset", + get_primary_value=lambda item: item.get("id"), + select_string=get_file_content( + os.path.join( + PROJECT_DIR, + "apps", + "application", + "sql", + ( + "export_application_chat_ee.sql" + if ["PE", "EE"].__contains__(edition) + else "export_application_chat.sql" + ), + ) + ), + with_table_name=False, + ): for item in data_list: row = [self.reset_value(v) for v in self.to_row(item)] worksheet.append(row) @@ -244,56 +316,74 @@ def stream_response(): output.close() workbook.close() - response = StreamingHttpResponse(stream_response(), - content_type='application/vnd.open.xmlformats-officedocument.spreadsheetml.sheet') - response['Content-Disposition'] = 'attachment; filename="data.xlsx"' + response = StreamingHttpResponse( + stream_response(), content_type="application/vnd.open.xmlformats-officedocument.spreadsheetml.sheet" + ) + response["Content-Disposition"] = 'attachment; filename="data.xlsx"' return response def page(self, current_page: int, page_size: int, with_valid=True): if with_valid: self.is_valid(raise_exception=True) - return native_page_search(current_page, page_size, self.get_query_set(), select_string=get_file_content( - os.path.join(PROJECT_DIR, "apps", "application", 'sql', - ('list_application_chat_ee.sql' if ['PE', 'EE'].__contains__( - edition) else 'list_application_chat.sql'))), - with_table_name=False) + return native_page_search( + current_page, + page_size, + self.get_query_set(), + select_string=get_file_content( + os.path.join( + PROJECT_DIR, + "apps", + "application", + "sql", + ( + "list_application_chat_ee.sql" + if ["PE", "EE"].__contains__(edition) + else "list_application_chat.sql" + ), + ) + ), + with_table_name=False, + ) class ChatCountSerializer(serializers.Serializer): chat_id = serializers.UUIDField(required=True, label=_("Conversation ID")) def get_query_set(self): - return QuerySet(ChatRecord).filter(chat_id=self.data.get('chat_id')) + return QuerySet(ChatRecord).filter(chat_id=self.data.get("chat_id")) def update_chat(self): self.is_valid(raise_exception=True) - count_chat_record = native_search(self.get_query_set(), get_file_content( - os.path.join(PROJECT_DIR, "apps", "application", 'sql', 'count_chat_record.sql')), with_search_one=True) - QuerySet(Chat).filter(id=self.data.get('chat_id')).update(star_num=count_chat_record.get('star_num', 0) or 0, - trample_num=count_chat_record.get('trample_num', - 0) or 0, - chat_record_count=count_chat_record.get( - 'chat_record_count', 0) or 0, - mark_sum=count_chat_record.get('mark_sum', 0) or 0) + count_chat_record = native_search( + self.get_query_set(), + get_file_content(os.path.join(PROJECT_DIR, "apps", "application", "sql", "count_chat_record.sql")), + with_search_one=True, + ) + QuerySet(Chat).filter(id=self.data.get("chat_id")).update( + star_num=count_chat_record.get("star_num", 0) or 0, + trample_num=count_chat_record.get("trample_num", 0) or 0, + chat_record_count=count_chat_record.get("chat_record_count", 0) or 0, + mark_sum=count_chat_record.get("mark_sum", 0) or 0, + ) return True def get_source_display(source): - if not source or not isinstance(source, dict) or 'type' not in source: - return '-' - source_type = source.get('type') + if not source or not isinstance(source, dict) or "type" not in source: + return "-" + source_type = source.get("type") # 定义映射关系 source_mapping = { - ChatSourceChoices.ONLINE.value: gettext('Online Usage'), - ChatSourceChoices.API_CALL.value: gettext('API Call'), - ChatSourceChoices.ENTERPRISE_WECHAT.value: gettext('Enterprise WeChat'), - ChatSourceChoices.WECHAT_PUBLIC_ACCOUNT.value: gettext('WeChat Public Account'), - ChatSourceChoices.LARK.value: gettext('Lark'), - ChatSourceChoices.DINGTALK.value: gettext('DingTalk'), - ChatSourceChoices.ENTERPRISE_WECHAT_ROBOT.value: gettext('Enterprise WeChat Robot'), - ChatSourceChoices.TRIGGER.value: gettext('Trigger'), - ChatSourceChoices.SLACK.value: gettext('Slack'), + ChatSourceChoices.ONLINE.value: gettext("Online Usage"), + ChatSourceChoices.API_CALL.value: gettext("API Call"), + ChatSourceChoices.ENTERPRISE_WECHAT.value: gettext("Enterprise WeChat"), + ChatSourceChoices.WECHAT_PUBLIC_ACCOUNT.value: gettext("WeChat Public Account"), + ChatSourceChoices.LARK.value: gettext("Lark"), + ChatSourceChoices.DINGTALK.value: gettext("DingTalk"), + ChatSourceChoices.ENTERPRISE_WECHAT_ROBOT.value: gettext("Enterprise WeChat Robot"), + ChatSourceChoices.TRIGGER.value: gettext("Trigger"), + ChatSourceChoices.SLACK.value: gettext("Slack"), } return source_mapping.get(source_type, str(source_type)) diff --git a/apps/chat/serializers/chat_authentication.py b/apps/chat/serializers/chat_authentication.py index 51ec2e65be2..e3a45bf6381 100644 --- a/apps/chat/serializers/chat_authentication.py +++ b/apps/chat/serializers/chat_authentication.py @@ -15,7 +15,6 @@ from rest_framework import serializers from application.models import ApplicationAccessToken, Application, ApplicationVersion -from application.serializers.application import ApplicationSerializerModel from common.auth.common import FileToken, ChatToken from common.auth.constants.operate_constants import Operate from common.constants.authentication_type import AuthenticationType @@ -38,7 +37,7 @@ def auth(self, request): # 校验token if token is not None: token_details = signing.loads(token[7:]) - except Exception as e: + except Exception: pass chat_user_id = token_details.get("id") or str(uuid.uuid7()) _type = AuthenticationType.CHAT_USER @@ -72,7 +71,7 @@ def auth(self, request, with_valid=True): # 校验token if token is not None: token_details = signing.loads(token[7:]) - except Exception as e: + except Exception: pass if with_valid: self.is_valid(raise_exception=True) @@ -222,7 +221,14 @@ def profile(self, with_valid=True): node for node in ((application.work_flow or {}).get("nodes", []) or []) if node.get("id") == "base-node" ] return { - **ApplicationSerializerModel(application).data, + "id": application.id, + "name": application.name, + "desc": application.desc, + "prologue": application.prologue, + "icon": application.icon, + "type": application.type, + "dialogue_number": application.dialogue_number, + "problem_optimization": application.problem_optimization, "stt_model_id": application.stt_model_id, "tts_model_id": application.tts_model_id, "stt_model_enable": application.stt_model_enable, diff --git a/apps/common/handle/impl/common_handle.py b/apps/common/handle/impl/common_handle.py index 16c647a9626..99cab6f63c6 100644 --- a/apps/common/handle/impl/common_handle.py +++ b/apps/common/handle/impl/common_handle.py @@ -1,14 +1,14 @@ # coding=utf-8 """ - @project: MaxKB - @Author:虎 - @file: tools.py - @date:2024/9/11 16:41 - @desc: +@project: MaxKB +@Author:虎 +@file: tools.py +@date:2024/9/11 16:41 +@desc: """ + import io import traceback -from functools import reduce from io import BytesIO from xml.etree.ElementTree import fromstring from zipfile import ZipFile @@ -23,8 +23,42 @@ from knowledge.models import File from PIL import ImageFile + ImageFile.LOAD_TRUNCATED_IMAGES = True -PILImage.MAX_IMAGE_PIXELS = None + +# 全局图片解码像素上限(不再禁用 Pillow 的解压炸弹保护)。 +# 超过该上限 Pillow 会告警,超过 2 倍会直接抛错,避免超大图片耗尽 worker 内存。 +PILImage.MAX_IMAGE_PIXELS = 50_000_000 + +# 内嵌图片解码保护(防解压炸弹 / 超大尺寸图片导致共享 worker OOM)。 +MAX_EMBED_IMAGE_PIXELS = 16_000_000 +MAX_EMBED_IMAGE_AGGREGATE_PIXELS = 64_000_000 + +# XLSX(zip) 压缩包防护,限制成员数 / 解压后总大小 / 解压膨胀比。 +MAX_EMBED_ARCHIVE_MEMBERS = 10_000 +MAX_EMBED_ARCHIVE_UNCOMPRESSED_BYTES = 1024 * 1024 * 1024 +MAX_EMBED_ARCHIVE_EXPANSION_RATIO = 50 + + +def validate_xlsx_archive(archive: ZipFile): + infolist = archive.infolist() + if len(infolist) > MAX_EMBED_ARCHIVE_MEMBERS: + raise ValueError(f"XLSX archive member count exceeds limit: {len(infolist)}") + total_uncompressed = sum(info.file_size for info in infolist) + total_compressed = sum(info.compress_size for info in infolist) + if total_uncompressed > MAX_EMBED_ARCHIVE_UNCOMPRESSED_BYTES: + raise ValueError("XLSX archive uncompressed size exceeds limit") + if total_compressed > 0 and total_uncompressed > total_compressed * MAX_EMBED_ARCHIVE_EXPANSION_RATIO: + raise ValueError("XLSX archive expansion ratio exceeds limit") + + +def validate_xlsx_buffer(buffer): + archive = ZipFile(buffer) + try: + validate_xlsx_archive(archive) + finally: + archive.close() + def parse_element(element) -> {}: data = {} @@ -87,15 +121,16 @@ def handle_images(deps, archive: ZipFile) -> []: def xlsx_embed_cells_images(buffer) -> {}: archive = ZipFile(buffer) + validate_xlsx_archive(archive) # 解析cellImage.xml文件 deps = get_dependents(archive, get_rels_path("xl/cellimages.xml")) image_rel = handle_images(deps=deps, archive=archive) # 工作表及其中图片ID sheet_list = {} for item in archive.namelist(): - if not item.startswith('xl/worksheets/sheet'): + if not item.startswith("xl/worksheets/sheet"): continue - key = item.split('/')[-1].split('.')[0].split('sheet')[-1] + key = item.split("/")[-1].split(".")[0].split("sheet")[-1] sheet_list[key] = parse_element_sheet_xml(fromstring(archive.read(item))) cell_images_xml = parse_element(fromstring(archive.read("xl/cellimages.xml"))) cell_images_rel = {} @@ -104,18 +139,11 @@ def xlsx_embed_cells_images(buffer) -> {}: for cnv, embed in cell_images_xml.items(): cell_images_xml[cnv] = cell_images_rel.get(embed) result = {} + total_pixels = 0 for key, img in cell_images_xml.items(): - all_cells = [ - cell - for _sheet_id, sheet in sheet_list.items() - if sheet is not None - for cell in sheet or [] - ] - - image_excel_id_list = [ - cell for cell in all_cells - if isinstance(cell, str) and key in cell - ] + all_cells = [cell for _sheet_id, sheet in sheet_list.items() if sheet is not None for cell in sheet or []] + + image_excel_id_list = [cell for cell in all_cells if isinstance(cell, str) and key in cell] # print(key, img) if img is None: continue @@ -123,9 +151,24 @@ def xlsx_embed_cells_images(buffer) -> {}: image_excel_id = image_excel_id_list[-1] f = archive.open(img.target) img_byte = io.BytesIO() - im = PILImage.open(f).convert('RGB') - im.save(img_byte, format='JPEG') - image = File(id=uuid.uuid7(), file_name=img.path, meta={'debug': False, 'content': img_byte.getvalue()}) - result['=' + image_excel_id] = image + try: + with PILImage.open(f) as im: + width, height = im.size + pixels = width * height + if pixels > MAX_EMBED_IMAGE_PIXELS: + maxkb_logger.warning( + f"Skip oversized embedded image {img.path}: {width}x{height} pixels exceeds limit" + ) + continue + total_pixels += pixels + if total_pixels > MAX_EMBED_IMAGE_AGGREGATE_PIXELS: + maxkb_logger.warning("Skip embedded images in archive: aggregate pixels exceed limit") + break + im.convert("RGB").save(img_byte, format="JPEG") + except Exception as e: + maxkb_logger.error(f"Error decoding image {img.target}: {e}, {traceback.format_exc()}") + continue + image = File(id=uuid.uuid7(), file_name=img.path, meta={"debug": False, "content": img_byte.getvalue()}) + result["=" + image_excel_id] = image archive.close() return result diff --git a/apps/common/handle/impl/qa/xlsx_parse_qa_handle.py b/apps/common/handle/impl/qa/xlsx_parse_qa_handle.py index 71332b1b9e1..ff5987763f6 100644 --- a/apps/common/handle/impl/qa/xlsx_parse_qa_handle.py +++ b/apps/common/handle/impl/qa/xlsx_parse_qa_handle.py @@ -1,18 +1,19 @@ # coding=utf-8 """ - @project: maxkb - @Author:虎 - @file: xlsx_parse_qa_handle.py - @date:2024/5/21 14:59 - @desc: +@project: maxkb +@Author:虎 +@file: xlsx_parse_qa_handle.py +@date:2024/5/21 14:59 +@desc: """ + import io import traceback import openpyxl from common.handle.base_parse_qa_handle import BaseParseQAHandle, get_title_row_index_dict, get_row_value -from common.handle.impl.common_handle import xlsx_embed_cells_images +from common.handle.impl.common_handle import xlsx_embed_cells_images, validate_xlsx_buffer from common.utils.logger import maxkb_logger @@ -22,28 +23,26 @@ def handle_sheet(file_name, sheet, image_dict): title_row_list = next(rows) title_row_list = [row.value for row in title_row_list] except Exception as e: - return {'name': file_name, 'paragraphs': []} + return {"name": file_name, "paragraphs": []} if len(title_row_list) == 0: - return {'name': file_name, 'paragraphs': []} + return {"name": file_name, "paragraphs": []} title_row_index_dict = get_title_row_index_dict(title_row_list) paragraph_list = [] for row in rows: - content = get_row_value(row, title_row_index_dict, 'content') + content = get_row_value(row, title_row_index_dict, "content") if content is None or content.value is None: continue - problem = get_row_value(row, title_row_index_dict, 'problem_list') - problem = str(problem.value) if problem is not None and problem.value is not None else '' - problem_list = [{'content': p[0:255]} for p in problem.split('\n') if len(p.strip()) > 0] - title = get_row_value(row, title_row_index_dict, 'title') - title = str(title.value) if title is not None and title.value is not None else '' + problem = get_row_value(row, title_row_index_dict, "problem_list") + problem = str(problem.value) if problem is not None and problem.value is not None else "" + problem_list = [{"content": p[0:255]} for p in problem.split("\n") if len(p.strip()) > 0] + title = get_row_value(row, title_row_index_dict, "title") + title = str(title.value) if title is not None and title.value is not None else "" content = str(content.value) image = image_dict.get(content, None) if image is not None: - content = f'![](./oss/file/{image.id})' - paragraph_list.append({'title': title[0:255], - 'content': content[0:102400], - 'problem_list': problem_list}) - return {'name': file_name, 'paragraphs': paragraph_list} + content = f"![](./oss/file/{image.id})" + paragraph_list.append({"title": title[0:255], "content": content[0:102400], "problem_list": problem_list}) + return {"name": file_name, "paragraphs": paragraph_list} class XlsxParseQAHandle(BaseParseQAHandle): @@ -56,6 +55,7 @@ def support(self, file, get_buffer): def handle(self, file, get_buffer, save_image): buffer = get_buffer(file) try: + validate_xlsx_buffer(io.BytesIO(buffer)) workbook = openpyxl.load_workbook(io.BytesIO(buffer)) try: image_dict: dict = xlsx_embed_cells_images(io.BytesIO(buffer)) @@ -64,12 +64,16 @@ def handle(self, file, get_buffer, save_image): image_dict = {} worksheets = workbook.worksheets worksheets_size = len(worksheets) - return [row for row in - [handle_sheet(file.name, - sheet, - image_dict) if worksheets_size == 1 and sheet.title == 'Sheet1' else handle_sheet( - sheet.title, sheet, image_dict) for sheet - in worksheets] if row is not None] + return [ + row + for row in [ + handle_sheet(file.name, sheet, image_dict) + if worksheets_size == 1 and sheet.title == "Sheet1" + else handle_sheet(sheet.title, sheet, image_dict) + for sheet in worksheets + ] + if row is not None + ] except Exception as e: maxkb_logger.error(f"Error processing XLSX file {file.name}: {e}, {traceback.format_exc()}") - return [{'name': file.name, 'paragraphs': []}] + return [{"name": file.name, "paragraphs": []}] diff --git a/apps/common/handle/impl/table/xlsx_parse_table_handle.py b/apps/common/handle/impl/table/xlsx_parse_table_handle.py index cf2bf68cf8d..fb76678b282 100644 --- a/apps/common/handle/impl/table/xlsx_parse_table_handle.py +++ b/apps/common/handle/impl/table/xlsx_parse_table_handle.py @@ -6,7 +6,7 @@ from openpyxl import load_workbook from common.handle.base_parse_table_handle import BaseParseTableHandle -from common.handle.impl.common_handle import xlsx_embed_cells_images +from common.handle.impl.common_handle import xlsx_embed_cells_images, validate_xlsx_buffer from common.handle.impl.xlsx_utils import iter_sheet_content_rows from common.utils.logger import maxkb_logger @@ -14,7 +14,7 @@ class XlsxParseTableHandle(BaseParseTableHandle): def support(self, file, get_buffer): file_name: str = file.name.lower() - if file_name.endswith('.xlsx'): + if file_name.endswith(".xlsx"): return True return False @@ -30,7 +30,7 @@ def fill_merged_cells(self, sheet, image_dict): return data for idx, cell in enumerate(title_row): if cell.value is None: - headers.append(' ' * (idx + 1)) + headers.append(" " * (idx + 1)) else: headers.append(cell.value) @@ -47,10 +47,10 @@ def fill_merged_cells(self, sheet, image_dict): cell_value = sheet[merged_range.min_row][merged_range.min_col - 1].value break if cell_value is None: - cell_value = '' + cell_value = "" image = image_dict.get(cell_value, None) if image is not None: - cell_value = f'![](./oss/file/{image.id})' + cell_value = f"![](./oss/file/{image.id})" # 使用标题作为键,单元格的值作为值存入字典 row_data[headers[col_idx]] = cell_value @@ -61,6 +61,7 @@ def fill_merged_cells(self, sheet, image_dict): def handle(self, file, get_buffer, save_image): buffer = get_buffer(file) try: + validate_xlsx_buffer(io.BytesIO(buffer)) wb = load_workbook(io.BytesIO(buffer)) try: image_dict: dict = xlsx_embed_cells_images(io.BytesIO(buffer)) @@ -76,13 +77,13 @@ def handle(self, file, get_buffer, save_image): for row in data: row_output = "; ".join([f"{key}: {value}" for key, value in row.items()]) # print(row_output) - paragraphs.append({'title': '', 'content': row_output}) + paragraphs.append({"title": "", "content": row_output}) - result.append({'name': sheetname, 'paragraphs': paragraphs}) + result.append({"name": sheetname, "paragraphs": paragraphs}) except BaseException as e: maxkb_logger.error(f"Error processing XLSX file {file.name}: {e}, {traceback.format_exc()}") - return [{'name': file.name, 'paragraphs': []}] + return [{"name": file.name, "paragraphs": []}] return result def get_content(self, file, save_image): @@ -94,9 +95,9 @@ def get_content(self, file, save_image): if len(image_dict) > 0: save_image(image_dict.values()) except Exception as e: - maxkb_logger.error(f'Exception: {e}') + maxkb_logger.error(f"Exception: {e}") image_dict = {} - md_tables = '' + md_tables = "" # 遍历所有工作表 for sheetname in workbook.sheetnames: sheet = workbook[sheetname] @@ -105,22 +106,25 @@ def get_content(self, file, save_image): continue # 添加 sheet 名称作为标题 - md_tables += f'## {sheetname}\n\n' + md_tables += f"## {sheetname}\n\n" # 提取表头和内容 headers = [f"{key}" for key, value in rows[0].items()] # 构建 Markdown 表格 - md_table = '| ' + ' | '.join(headers) + ' |\n' - md_table += '| ' + ' | '.join(['---'] * len(headers)) + ' |\n' + md_table = "| " + " | ".join(headers) + " |\n" + md_table += "| " + " | ".join(["---"] * len(headers)) + " |\n" for row in rows: - r = [f'{value}' for key, value in row.items()] - md_table += '| ' + ' | '.join( - [str(cell).replace('\n', '
') if cell is not None else '' for cell in r]) + ' |\n' + r = [f"{value}" for key, value in row.items()] + md_table += ( + "| " + + " | ".join([str(cell).replace("\n", "
") if cell is not None else "" for cell in r]) + + " |\n" + ) - md_tables += md_table + '\n\n' + md_tables += md_table + "\n\n" return md_tables except Exception as e: - maxkb_logger.error(f'excel split handle error: {e}') - return f'error: {e}' + maxkb_logger.error(f"excel split handle error: {e}") + return f"error: {e}" diff --git a/apps/common/handle/impl/text/xlsx_split_handle.py b/apps/common/handle/impl/text/xlsx_split_handle.py index c4d1ebb0fdf..4da581c0e93 100644 --- a/apps/common/handle/impl/text/xlsx_split_handle.py +++ b/apps/common/handle/impl/text/xlsx_split_handle.py @@ -1,11 +1,12 @@ # coding=utf-8 """ - @project: maxkb - @Author:虎 - @file: xlsx_parse_qa_handle.py - @date:2024/5/21 14:59 - @desc: +@project: maxkb +@Author:虎 +@file: xlsx_parse_qa_handle.py +@date:2024/5/21 14:59 +@desc: """ + import io import traceback from typing import List @@ -14,40 +15,46 @@ from openpyxl import load_workbook from common.handle.base_split_handle import BaseSplitHandle -from common.handle.impl.common_handle import xlsx_embed_cells_images +from common.handle.impl.common_handle import xlsx_embed_cells_images, validate_xlsx_buffer from common.handle.impl.xlsx_utils import iter_sheet_content_rows from common.utils.logger import maxkb_logger -splitter = '\n`-----------------------------------`\n' +splitter = "\n`-----------------------------------`\n" def post_cell(image_dict, cell_value): image = image_dict.get(cell_value, None) if image is not None: - return f'![](./oss/file/{image.id})' - return cell_value.replace('\n', '
').replace('|', '|') + return f"![](./oss/file/{image.id})" + return cell_value.replace("\n", "
").replace("|", "|") def row_to_md(row, image_dict): - return '| ' + ' | '.join( - [post_cell(image_dict, str(cell.value if cell.value is not None else '')) if cell is not None else '' for cell - in row]) + ' |\n' + return ( + "| " + + " | ".join( + [ + post_cell(image_dict, str(cell.value if cell.value is not None else "")) if cell is not None else "" + for cell in row + ] + ) + + " |\n" + ) def handle_sheet(file_name, sheet, image_dict, limit: int): rows = iter_sheet_content_rows(sheet) paragraphs = [] - result = {'name': file_name, 'content': paragraphs} + result = {"name": file_name, "content": paragraphs} try: title_row_list = next(rows) title_md_content = row_to_md(title_row_list, image_dict) - title_md_content += '| ' + ' | '.join( - ['---' if cell is not None else '' for cell in title_row_list]) + ' |\n' + title_md_content += "| " + " | ".join(["---" if cell is not None else "" for cell in title_row_list]) + " |\n" except Exception as e: return result if len(title_row_list) == 0: return result - result_item_content = '' + result_item_content = "" for row in rows: next_md_content = row_to_md(row, image_dict) next_md_content_len = len(next_md_content) @@ -59,10 +66,10 @@ def handle_sheet(file_name, sheet, image_dict, limit: int): if result_item_content_len + next_md_content_len < limit: result_item_content += next_md_content else: - paragraphs.append({'content': result_item_content, 'title': ''}) + paragraphs.append({"content": result_item_content, "title": ""}) result_item_content = title_md_content + next_md_content if len(result_item_content) > 0: - paragraphs.append({'content': result_item_content, 'title': ''}) + paragraphs.append({"content": result_item_content, "title": ""}) return result @@ -79,7 +86,7 @@ def fill_merged_cells(self, sheet, image_dict): return data for idx, cell in enumerate(title_row): if cell.value is None: - headers.append(' ' * (idx + 1)) + headers.append(" " * (idx + 1)) else: headers.append(cell.value) @@ -98,7 +105,7 @@ def fill_merged_cells(self, sheet, image_dict): image = image_dict.get(cell_value, None) if image is not None: - cell_value = f'![](./oss/file/{image.id})' + cell_value = f"![](./oss/file/{image.id})" # 使用标题作为键,单元格的值作为值存入字典 row_data[headers[col_idx]] = cell_value @@ -109,6 +116,7 @@ def fill_merged_cells(self, sheet, image_dict): def handle(self, file, pattern_list: List, with_filter: bool, limit: int, get_buffer, save_image): buffer = get_buffer(file) try: + validate_xlsx_buffer(io.BytesIO(buffer)) if type(limit) is str: limit = int(limit) workbook = openpyxl.load_workbook(io.BytesIO(buffer)) @@ -119,16 +127,19 @@ def handle(self, file, pattern_list: List, with_filter: bool, limit: int, get_bu image_dict = {} worksheets = workbook.worksheets worksheets_size = len(worksheets) - return [row for row in - [handle_sheet(file.name, - sheet, - image_dict, - limit) if worksheets_size == 1 and sheet.title == 'Sheet1' else handle_sheet( - sheet.title, sheet, image_dict, limit) for sheet - in worksheets] if row is not None] + return [ + row + for row in [ + handle_sheet(file.name, sheet, image_dict, limit) + if worksheets_size == 1 and sheet.title == "Sheet1" + else handle_sheet(sheet.title, sheet, image_dict, limit) + for sheet in worksheets + ] + if row is not None + ] except Exception as e: maxkb_logger.error(f"Error processing XLSX file {file.name}: {e}, {traceback.format_exc()}") - return [{'name': file.name, 'content': []}] + return [{"name": file.name, "content": []}] def get_content(self, file, save_image): try: @@ -139,9 +150,9 @@ def get_content(self, file, save_image): if len(image_dict) > 0: save_image(image_dict.values()) except Exception as e: - maxkb_logger.error(f'Exception: {e}') + maxkb_logger.error(f"Exception: {e}") image_dict = {} - md_tables = '' + md_tables = "" # 遍历所有工作表 for sheetname in workbook.sheetnames: sheet = workbook[sheetname] @@ -150,41 +161,41 @@ def get_content(self, file, save_image): continue # 添加 sheet 名称作为标题 - md_tables += f'## {sheetname}\n\n' + md_tables += f"## {sheetname}\n\n" # 提取表头和内容 headers = [f"{key}" for key, value in rows[0].items()] # 构建 Markdown 表格 - md_table = '| ' + ' | '.join(headers) + ' |\n' - md_table += '| ' + ' | '.join(['---'] * len(headers)) + ' |\n' + md_table = "| " + " | ".join(headers) + " |\n" + md_table += "| " + " | ".join(["---"] * len(headers)) + " |\n" for row in rows: r = [self._escape_cell_content(value) for key, value in row.items()] - md_table += '| ' + ' | '.join(r) + ' |\n' + md_table += "| " + " | ".join(r) + " |\n" - md_tables += md_table + '\n\n' + md_tables += md_table + "\n\n" return md_tables except Exception as e: - maxkb_logger.error(f'excel split handle error: {e}') - return f'error: {e}' + maxkb_logger.error(f"excel split handle error: {e}") + return f"error: {e}" def _escape_cell_content(self, cell_value): """转义单元格内容,避免破坏 Markdown 表格结构""" if cell_value is None: - return '' + return "" cell_str = str(cell_value) # 替换换行符为
- cell_str = cell_str.replace('\n', '
') + cell_str = cell_str.replace("\n", "
") # 转义管道符 | 为 HTML 实体 - cell_str = cell_str.replace('|', '|') + cell_str = cell_str.replace("|", "|") # 如果内容包含反引号,需要转义 - if '`' in cell_str: - cell_str = cell_str.replace('`', '`') + if "`" in cell_str: + cell_str = cell_str.replace("`", "`") return cell_str diff --git a/apps/knowledge/serializers/knowledge_workflow.py b/apps/knowledge/serializers/knowledge_workflow.py index 0b3fc74e007..49a6612bae9 100644 --- a/apps/knowledge/serializers/knowledge_workflow.py +++ b/apps/knowledge/serializers/knowledge_workflow.py @@ -219,6 +219,18 @@ def action(self, instance: Dict, user, with_valid=True, sync_log_id=None): ), ) work_flow_manage.run() + # 需要把文件改成永久文件 + data_source = instance.get("data_source") or {} + file_ids = [item.get("file_id") for item in data_source.get("file_list") or [] if item.get("file_id")] + knowledge_id = str(self.data.get("knowledge_id")) + file_list = list(QuerySet(File).filter(id__in=file_ids)) + for file in file_list: + meta = dict(file.meta or {}) + meta.update(debug=False, knowledge_id=knowledge_id) + file.source_type = FileSourceType.KNOWLEDGE.value + file.source_id = knowledge_id + file.meta = meta + QuerySet(File).bulk_update(file_list, ["source_type", "source_id", "meta"]) return { "id": knowledge_action_id, "knowledge_id": self.data.get("knowledge_id"), diff --git a/apps/knowledge/serializers/problem.py b/apps/knowledge/serializers/problem.py index 2455aba206a..1c6e1560775 100644 --- a/apps/knowledge/serializers/problem.py +++ b/apps/knowledge/serializers/problem.py @@ -20,15 +20,12 @@ class ProblemSerializer(serializers.ModelSerializer): class Meta: model = Problem - fields = ["id", "content", "knowledge_id", "create_time", "update_time", "hit_num", "last_hit_time"] - read_only_fields = ["hit_num", "last_hit_time"] + fields = ["id", "content", "knowledge_id", "create_time", "update_time"] class ProblemInstanceSerializer(serializers.Serializer): id = serializers.CharField(required=False, label=_("problem id")) content = serializers.CharField(required=True, max_length=256, label=_("content")) - hit_num = serializers.IntegerField(read_only=True, label=_("recall count")) - last_hit_time = serializers.DateTimeField(read_only=True, allow_null=True, label=_("last recall time")) class ProblemEditSerializer(serializers.Serializer): @@ -98,12 +95,22 @@ def association(self, instance: Dict, with_valid=True): self.is_valid(raise_exception=True) BatchAssociation(data=instance).is_valid(raise_exception=True) knowledge_id = self.data.get("knowledge_id") - paragraph_list = instance.get("paragraph_list") - problem_id_list = instance.get("problem_id_list") - problem_list = QuerySet(Problem).filter(id__in=problem_id_list) + paragraph_list = instance.get("paragraph_list") or [] + problem_id_list = instance.get("problem_id_list") or [] + paragraph_id_list = [p.get("paragraph_id") for p in paragraph_list] + + # 校验目标段落都属于当前知识库, 防止跨知识库关联并回读他人内容 + if QuerySet(Paragraph).filter(id__in=paragraph_id_list, knowledge_id=knowledge_id).count() != len( + set(paragraph_id_list) + ): + raise AppApiException(500, _("Paragraph does not exist")) + # 仅允许关联当前知识库下的问题 + problem_list = QuerySet(Problem).filter(id__in=problem_id_list, knowledge_id=knowledge_id) + if problem_list.count() != len(set(problem_id_list)): + raise AppApiException(500, _("Problem does not exist")) exits_problem_paragraph_mapping = QuerySet(ProblemParagraphMapping).filter( - problem_id__in=problem_id_list, paragraph_id__in=[p.get("paragraph_id") for p in paragraph_list] + problem_id__in=problem_id_list, paragraph_id__in=paragraph_id_list ) problem_paragraph_mapping_list = [ @@ -166,7 +173,10 @@ def list_paragraph(self, with_valid=True): if problem_paragraph_mapping is None or len(problem_paragraph_mapping) == 0: return [] return native_search( - QuerySet(Paragraph).filter(id__in=[row.paragraph_id for row in problem_paragraph_mapping]), + QuerySet(Paragraph).filter( + knowledge_id=self.data.get("knowledge_id"), + id__in=[row.paragraph_id for row in problem_paragraph_mapping], + ), select_string=get_file_content( os.path.join(PROJECT_DIR, "apps", "knowledge", "sql", "list_paragraph.sql") ), diff --git a/apps/oss/views/file.py b/apps/oss/views/file.py index 6c2641839d4..7e6a1df2bfc 100644 --- a/apps/oss/views/file.py +++ b/apps/oss/views/file.py @@ -9,27 +9,32 @@ from common.auth.authentication import has_permissions from common.auth.common import ChatAuthentication from common.auth.constants.role_constants import RoleConstants +from common.exception.app_exception import AppUnauthorizedFailed from common.log.log import log from common.result import result from knowledge.api.file import FileGetAPI, FileUploadAPI, GetUrlContentAPI +from knowledge.models import FileSourceType +from maxkb.const import CONFIG from oss.serializers.file import FileSerializer, get_url_content class FileRetrievalView(APIView): @extend_schema( - methods=['GET'], - summary=_('Get file'), - description=_('Get file'), - operation_id=_('Get file'), # type: ignore + methods=["GET"], + summary=_("Get file"), + description=_("Get file"), + operation_id=_("Get file"), # type: ignore parameters=FileGetAPI.get_parameters(), responses=FileGetAPI.get_response(), - tags=[_('File')] # type: ignore + tags=[_("File")], # type: ignore ) def get(self, request: Request, file_id: str): - return FileSerializer.Operate(data={ - 'id': file_id, - 'http_range': request.headers.get('Range', ''), - }).get(mk_file_auth=request.COOKIES.get('mk_file_auth')) + return FileSerializer.Operate( + data={ + "id": file_id, + "http_range": request.headers.get("Range", ""), + } + ).get(mk_file_auth=request.COOKIES.get("mk_file_auth")) class FileView(APIView): @@ -37,61 +42,78 @@ class FileView(APIView): parser_classes = [MultiPartParser] @extend_schema( - methods=['POST'], - summary=_('Upload file'), - description=_('Upload file'), - operation_id=_('Upload file'), # type: ignore + methods=["POST"], + summary=_("Upload file"), + description=_("Upload file"), + operation_id=_("Upload file"), # type: ignore parameters=FileUploadAPI.get_parameters(), request=FileUploadAPI.get_request(), responses=FileUploadAPI.get_response(), - tags=[_('File')] # type: ignore + tags=[_("File")], # type: ignore ) - @log(menu='file', operate='Upload file') + @log(menu="file", operate="Upload file") def post(self, request: Request): - return result.success(FileSerializer(data={ - 'file': request.FILES.get('file'), - 'source_id': request.data.get('source_id'), - 'source_type': request.data.get('source_type'), - }).upload(user_id=(str(request.user.id) if request.user else request.auth.chat_user_id))) + source_id = request.data.get("source_id") + source_type = request.data.get("source_type") or FileSourceType.TEMPORARY_120_MINUTE.value + # 聊天路径(/chat/...)或匿名会话下只能上传聊天文件,禁止将文件归属到 + # Application/Knowledge 等其他受保护资源,无论调用者是否登录。 + is_chat_path = request.path.startswith(CONFIG.get_chat_path()) + if request.user is None or is_chat_path: + if source_type != FileSourceType.CHAT.value: + raise AppUnauthorizedFailed(403, _("No permission")) + return result.success( + FileSerializer( + data={ + "file": request.FILES.get("file"), + "source_id": source_id, + "source_type": source_type, + } + ).upload(user_id=(str(request.user.id) if request.user else request.auth.chat_user_id)) + ) class Operate(APIView): authentication_classes = [TokenAuth] @extend_schema( - methods=['DELETE'], - summary=_('Delete file'), - description=_('Delete file'), - operation_id=_('Delete file'), # type: ignore + methods=["DELETE"], + summary=_("Delete file"), + description=_("Delete file"), + operation_id=_("Delete file"), # type: ignore parameters=FileGetAPI.get_parameters(), responses=FileGetAPI.get_response(), - tags=[_('File')] # type: ignore + tags=[_("File")], # type: ignore ) - @log(menu='file', operate='Delete file') + @log(menu="file", operate="Delete file") @has_permissions(RoleConstants.ADMIN, RoleConstants.WORKSPACE_MANAGE, RoleConstants.USER) def delete(self, request: Request, file_id: str): - return result.success(FileSerializer.Operate( - data={ - "id": file_id, - "http_range": request.headers.get("Range", ""), - } - ).delete(mk_file_auth=request.COOKIES.get("mk_file_auth"))) + return result.success( + FileSerializer.Operate( + data={ + "id": file_id, + "http_range": request.headers.get("Range", ""), + } + ).delete(mk_file_auth=request.COOKIES.get("mk_file_auth")) + ) class GetUrlView(APIView): authentication_classes = [AllTokenAuth] @extend_schema( - methods=['GET'], - summary=_('Get url'), + methods=["GET"], + summary=_("Get url"), parameters=GetUrlContentAPI.get_parameters(), - description=_('Get url'), - operation_id=_('Get url'), # type: ignore - tags=[_('Chat')] # type: ignore + description=_("Get url"), + operation_id=_("Get url"), # type: ignore + tags=[_("Chat")], # type: ignore ) def get(self, request: Request, application_id: str): - if isinstance(request.auth, ChatAuthentication) and request.auth.application_id and str( - request.auth.application_id) != application_id: - return result.error(_('No permission')) - url = request.query_params.get('url') + if ( + isinstance(request.auth, ChatAuthentication) + and request.auth.application_id + and str(request.auth.application_id) != application_id + ): + return result.error(_("No permission")) + url = request.query_params.get("url") result_data = get_url_content(url, application_id) return result.success(result_data) diff --git a/apps/users/serializers/user.py b/apps/users/serializers/user.py index a57e418d1d7..40954d3cee8 100644 --- a/apps/users/serializers/user.py +++ b/apps/users/serializers/user.py @@ -7,6 +7,7 @@ @desc: """ +import datetime import json import os import random @@ -28,7 +29,6 @@ from common.auth.struct.auth import Auth from common.constants.cache_version import Cache_Version from common.constants.exception_code_constants import ExceptionCodeConstants -from common.constants.resource_permission_constants import ResourcePermissionConstants from common.database_model_manage.database_model_manage import DatabaseModelManage from common.db.search import page_search from common.exception.app_exception import AppApiException @@ -36,7 +36,7 @@ from common.utils.rsa_util import decrypt from maxkb.conf import PROJECT_DIR from maxkb.const import CONFIG -from system_manage.models import AuthTargetType, SettingType, SystemSetting, WorkspaceUserResourcePermission +from system_manage.models import SettingType, SystemSetting from users.models import User from users.models.user_group import SystemUserGroup, SystemUserGroupRelation @@ -990,10 +990,11 @@ def is_valid(self, *, raise_exception=False): super().is_valid(raise_exception=raise_exception) code_cache_key = self.data.get("email") + ":" + self.data.get("type") code_cache_key_lock = code_cache_key + "_lock" - ttl = cache.ttl(get_key(code_cache_key_lock), version=version) - if ttl is not None and ttl > 0: + ttl = cache.ttl(code_cache_key_lock, version=version) + seconds = ttl.total_seconds() if isinstance(ttl, datetime.timedelta) else ttl + if seconds is not None and seconds > 0: raise AppApiException( - 500, _("Do not send emails again within {seconds} seconds").format(seconds=int(ttl.total_seconds())) + 500, _("Do not send emails again within {seconds} seconds").format(seconds=int(seconds)) ) return True