From 1a3eec3f18f2e6324d2596a1ee2b32cde2846bbd Mon Sep 17 00:00:00 2001
From: wxg0103 <727495428@qq.com>
Date: Wed, 2 Sep 2026 11:18:54 +0800
Subject: [PATCH] refactor: improve code formatting and enhance XLSX validation
logic
---
.../serializers/application_chat.py | 388 +++++++++++-------
apps/chat/serializers/chat_authentication.py | 14 +-
apps/common/handle/impl/common_handle.py | 91 ++--
.../handle/impl/qa/xlsx_parse_qa_handle.py | 56 +--
.../impl/table/xlsx_parse_table_handle.py | 42 +-
.../handle/impl/text/xlsx_split_handle.py | 95 +++--
.../serializers/knowledge_workflow.py | 12 +
apps/knowledge/serializers/problem.py | 28 +-
apps/oss/views/file.py | 104 +++--
apps/users/serializers/user.py | 11 +-
10 files changed, 522 insertions(+), 319 deletions(-)
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''
- paragraph_list.append({'title': title[0:255],
- 'content': content[0:102400],
- 'problem_list': problem_list})
- return {'name': file_name, 'paragraphs': paragraph_list}
+ content = f""
+ 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''
+ cell_value = f""
# 使用标题作为键,单元格的值作为值存入字典
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''
- return cell_value.replace('\n', '
').replace('|', '|')
+ return f""
+ 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''
+ cell_value = f""
# 使用标题作为键,单元格的值作为值存入字典
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