From 44c8fe9dd0ddcaa0e9caad743f8207148f1649c7 Mon Sep 17 00:00:00 2001 From: wxg0103 <727495428@qq.com> Date: Fri, 7 Aug 2026 14:46:17 +0800 Subject: [PATCH] feat: add token quota management for chat users with localization updates --- apps/application/serializers/common.py | 5 + apps/chat/serializers/chat.py | 848 ++++++++++-------- apps/locales/en_US/LC_MESSAGES/django.po | 3 + apps/locales/zh_CN/LC_MESSAGES/django.po | 3 + apps/locales/zh_Hant/LC_MESSAGES/django.po | 5 +- .../0008_add_chat_user_token_quota.py | 85 +- .../models/chat_user_token_quota.py | 31 + apps/system_manage/serializers/chat_user.py | 490 +++++----- apps/system_manage/views/system_chat_user.py | 27 +- 9 files changed, 859 insertions(+), 638 deletions(-) diff --git a/apps/application/serializers/common.py b/apps/application/serializers/common.py index 1086e640c9c..311483d7e7d 100644 --- a/apps/application/serializers/common.py +++ b/apps/application/serializers/common.py @@ -18,6 +18,7 @@ from django.db.models import QuerySet from django.utils import timezone from django.utils.translation import gettext_lazy as _ +from system_manage.models.chat_user_token_quota import ChatUserTokenQuota from application.models import Application, ChatRecord, Chat, ApplicationVersion, ChatUserType, ApplicationTypeChoices, \ ExecuteType @@ -397,6 +398,10 @@ def append_chat_record(self, chat_record: ChatRecord): ).save() else: QuerySet(Chat).filter(id=self.chat_id).update(update_time=timezone.now()) + # 记录Token消耗 + total_tokens = (chat_record.message_tokens or 0) + (chat_record.answer_tokens or 0) + if total_tokens > 0: + ChatUserTokenQuota.consume(self.chat_user_id, total_tokens) # 插入会话记录 QuerySet(ChatRecord).update_or_create( id=chat_record.id, diff --git a/apps/chat/serializers/chat.py b/apps/chat/serializers/chat.py index 721f638cffb..02d53545e57 100644 --- a/apps/chat/serializers/chat.py +++ b/apps/chat/serializers/chat.py @@ -1,11 +1,12 @@ # coding=utf-8 """ - @project: MaxKB - @Author:虎虎 - @file: chat.py - @date:2025/6/9 11:23 - @desc: +@project: MaxKB +@Author:虎虎 +@file: chat.py +@date:2025/6/9 11:23 +@desc: """ + import json import os import queue as thread_queue @@ -25,14 +26,23 @@ from application.chat_pipeline.pipeline_manage import PipelineManage from application.chat_pipeline.step.chat_step.i_chat_step import PostResponseHandler from application.chat_pipeline.step.chat_step.impl.base_chat_step import BaseChatStep -from application.chat_pipeline.step.generate_human_message_step.impl.base_generate_human_message_step import \ - BaseGenerateHumanMessageStep +from application.chat_pipeline.step.generate_human_message_step.impl.base_generate_human_message_step import ( + BaseGenerateHumanMessageStep, +) from application.chat_pipeline.step.reset_problem_step.impl.base_reset_problem_step import BaseResetProblemStep from application.chat_pipeline.step.search_dataset_step.impl.base_search_dataset_step import BaseSearchDatasetStep from application.flow.common import Answer from application.flow.tools import to_stream_response_simple -from application.models import Application, ApplicationTypeChoices, \ - ChatUserType, ApplicationChatUserStats, ApplicationAccessToken, ChatRecord, Chat, ApplicationVersion +from application.models import ( + Application, + ApplicationTypeChoices, + ChatUserType, + ApplicationChatUserStats, + ApplicationAccessToken, + ChatRecord, + Chat, + ApplicationVersion, +) from application.serializers.application import ApplicationOperateSerializer from application.serializers.common import ChatInfo from application.workflow.common import WorkflowType, new_instance @@ -58,6 +68,7 @@ from maxkb.conf import PROJECT_DIR from models_provider.models import Model, Status from models_provider.tools import get_model_instance_by_model_workspace_id +from system_manage.models.chat_user_token_quota import ChatUserTokenQuota from system_manage.models.resource_mapping import ResourceMapping @@ -78,65 +89,74 @@ def is_valid(self, *, raise_exception=False): raise AppApiException(400, _("Too many messages")) for index in range(len(messages)): - role = messages[index].get('role') - if role == 'ai' and index % 2 != 1: + role = messages[index].get("role") + if role == "ai" and index % 2 != 1: raise AppApiException(400, _("Authentication failed. Please verify that the parameters are correct.")) - if role == 'user' and index % 2 != 0: + if role == "user" and index % 2 != 0: raise AppApiException(400, _("Authentication failed. Please verify that the parameters are correct.")) - if role not in ['user', 'ai']: + if role not in ["user", "ai"]: raise AppApiException(400, _("Authentication failed. Please verify that the parameters are correct.")) class ChatMessageSerializers(serializers.Serializer): message = serializers.DictField(required=True, label=_("User Questions")) - stream = serializers.BooleanField(required=True, - label=_("Is the answer in streaming mode")) + stream = serializers.BooleanField(required=True, label=_("Is the answer in streaming mode")) re_chat = serializers.BooleanField(required=True, label=_("Do you want to reply again")) - chat_record_id = serializers.UUIDField(required=False, allow_null=True, - label=_("Conversation record id")) + chat_record_id = serializers.UUIDField(required=False, allow_null=True, label=_("Conversation record id")) - node_id = serializers.CharField(required=False, allow_null=True, allow_blank=True, - label=_("Node id")) + node_id = serializers.CharField(required=False, allow_null=True, allow_blank=True, label=_("Node id")) - runtime_node_id = serializers.CharField(required=False, allow_null=True, allow_blank=True, - label=_("Runtime node id")) + runtime_node_id = serializers.CharField( + required=False, allow_null=True, allow_blank=True, label=_("Runtime node id") + ) - node_data = serializers.DictField(required=False, allow_null=True, - label=_("Node parameters")) + node_data = serializers.DictField(required=False, allow_null=True, label=_("Node parameters")) form_data = serializers.DictField(required=False, label=_("Global variables")) - child_node = serializers.DictField(required=False, allow_null=True, - label=_("Child Nodes")) + child_node = serializers.DictField(required=False, allow_null=True, label=_("Child Nodes")) def get_post_handler(chat_info: ChatInfo): class PostHandler(PostResponseHandler): - - def handler(self, - chat_id, - chat_record_id, - paragraph_list: List[Paragraph], - problem_text: str, - answer_text, - manage: PipelineManage, - step: BaseChatStep, - padding_problem_text: str = None, - **kwargs): - answer_list = [[Answer(answer_text, 'ai-chat-node', 'ai-chat-node', 'ai-chat-node', {}, 'ai-chat-node', - kwargs.get('reasoning_content', '')).to_dict()]] - chat_record = ChatRecord(id=chat_record_id, - chat_id=chat_id, - problem_text=problem_text, - answer_text=answer_text, - details=manage.get_details(), - message_tokens=manage.context['message_tokens'], - answer_tokens=manage.context['answer_tokens'], - answer_text_list=answer_list, - run_time=manage.context['run_time'], - index=len(chat_info.chat_record_list) + 1, - ip_address=chat_info.ip_address, - source=chat_info.source - ) + def handler( + self, + chat_id, + chat_record_id, + paragraph_list: List[Paragraph], + problem_text: str, + answer_text, + manage: PipelineManage, + step: BaseChatStep, + padding_problem_text: str = None, + **kwargs, + ): + answer_list = [ + [ + Answer( + answer_text, + "ai-chat-node", + "ai-chat-node", + "ai-chat-node", + {}, + "ai-chat-node", + kwargs.get("reasoning_content", ""), + ).to_dict() + ] + ] + chat_record = ChatRecord( + id=chat_record_id, + chat_id=chat_id, + problem_text=problem_text, + answer_text=answer_text, + details=manage.get_details(), + message_tokens=manage.context["message_tokens"], + answer_tokens=manage.context["answer_tokens"], + answer_text_list=answer_list, + run_time=manage.context["run_time"], + index=len(chat_info.chat_record_list) + 1, + ip_address=chat_info.ip_address, + source=chat_info.source, + ) chat_info.append_chat_record(chat_record) # 重新设置缓存 chat_info.set_cache() @@ -149,74 +169,83 @@ class DebugChatSerializers(serializers.Serializer): def chat(self, instance: dict, base_to_response: BaseToResponse = SystemToResponse()): self.is_valid(raise_exception=True) - chat_id = self.data.get('chat_id') + chat_id = self.data.get("chat_id") chat_info: ChatInfo = ChatInfo.get_cache(chat_id) application = QuerySet(Application).filter(id=chat_info.application_id).first() chat_info.application = application - return ChatSerializers(data={ - 'chat_id': chat_id, "chat_user_id": chat_info.chat_user_id, - "chat_user_type": chat_info.chat_user_type, - "application_id": chat_info.application.id, "debug": True - }).chat(instance, base_to_response) + return ChatSerializers( + data={ + "chat_id": chat_id, + "chat_user_id": chat_info.chat_user_id, + "chat_user_type": chat_info.chat_user_type, + "application_id": chat_info.application.id, + "debug": True, + } + ).chat(instance, base_to_response) -SYSTEM_ROLE = get_file_content(os.path.join(PROJECT_DIR, "apps", "chat", 'template', 'generate_prompt_system')) +SYSTEM_ROLE = get_file_content(os.path.join(PROJECT_DIR, "apps", "chat", "template", "generate_prompt_system")) class PromptGenerateSerializer(serializers.Serializer): - workspace_id = serializers.CharField(required=False, label=_('Workspace ID')) + workspace_id = serializers.CharField(required=False, label=_("Workspace ID")) model_id = serializers.CharField(required=False, allow_blank=True, allow_null=True, label=_("Model")) application_id = serializers.CharField(required=False, allow_blank=True, allow_null=True, label=_("Application")) 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) application = query_set.first() if application is None: - raise AppApiException(500, _('Application id does not exist')) + raise AppApiException(500, _("Application id does not exist")) return application def generate_prompt(self, instance: dict): application = self.is_valid(raise_exception=True) GeneratePromptSerializers(data=instance).is_valid(raise_exception=True) - workspace_id = self.data.get('workspace_id') - model_id = self.data.get('model_id') - prompt = instance.get('prompt') - messages = instance.get('messages') + workspace_id = self.data.get("workspace_id") + model_id = self.data.get("model_id") + prompt = instance.get("prompt") + messages = instance.get("messages") - message = messages[-1]['content'] + message = messages[-1]["content"] q = prompt.replace("{userInput}", message) - messages[-1]['content'] = q + messages[-1]["content"] = q SUPPORTED_MODEL_TYPES = ["LLM", "IMAGE"] - model_exist = QuerySet(Model).filter( - id=model_id, - model_type__in=SUPPORTED_MODEL_TYPES - ).exists() + model_exist = QuerySet(Model).filter(id=model_id, model_type__in=SUPPORTED_MODEL_TYPES).exists() if not model_exist: raise Exception(_("Model does not exists or is not an LLM model")) def process(): - model = get_model_instance_by_model_workspace_id(model_id=model_id, workspace_id=workspace_id, - **application.model_params_setting) + model = get_model_instance_by_model_workspace_id( + model_id=model_id, workspace_id=workspace_id, **application.model_params_setting + ) try: - for r in model.stream([SystemMessage(content=SYSTEM_ROLE), - *[HumanMessage(content=m.get('content')) if m.get( - 'role') == 'user' else AIMessage( - content=m.get('content')) for m in messages]]): - yield 'data: ' + json.dumps({'content': r.content}) + '\n\n' + for r in model.stream( + [ + SystemMessage(content=SYSTEM_ROLE), + *[ + HumanMessage(content=m.get("content")) + if m.get("role") == "user" + else AIMessage(content=m.get("content")) + for m in messages + ], + ] + ): + yield "data: " + json.dumps({"content": r.content}) + "\n\n" except Exception as e: - yield 'data: ' + json.dumps({'error': str(e)}) + '\n\n' + yield "data: " + json.dumps({"error": str(e)}) + "\n\n" return to_stream_response_simple(process()) class OpenAIMessage(serializers.Serializer): - content = serializers.CharField(required=True, label=_('content')) - role = serializers.CharField(required=True, label=_('Role')) + content = serializers.CharField(required=True, label=_("content")) + role = serializers.CharField(required=True, label=_("Role")) class OpenAIInstanceSerializer(serializers.Serializer): @@ -235,26 +264,27 @@ class OpenAIChatSerializer(serializers.Serializer): @staticmethod def get_message(instance): - return instance.get('messages')[-1].get('content') + return instance.get("messages")[-1].get("content") @staticmethod def generate_chat(chat_id, application_id, message, chat_user_id, chat_user_type, ip_address, source): if chat_id is None: chat_id = str(uuid.uuid1()) - chat_info = ChatInfo(chat_id, chat_user_id, chat_user_type, ip_address, source, [], [], - application_id) + chat_info = ChatInfo(chat_id, chat_user_id, chat_user_type, ip_address, source, [], [], application_id) chat_info.set_cache() else: chat_info = ChatInfo.get_cache(chat_id) if chat_info is None: - open_chat = ChatSerializers(data={ - 'chat_id': chat_id, - 'chat_user_id': chat_user_id, - 'chat_user_type': chat_user_type, - 'application_id': application_id, - 'ip_address': ip_address, - 'source': source, - }) + open_chat = ChatSerializers( + data={ + "chat_id": chat_id, + "chat_user_id": chat_user_id, + "chat_user_type": chat_user_type, + "application_id": application_id, + "ip_address": ip_address, + "source": source, + } + ) open_chat.is_valid(raise_exception=True) chat_info = open_chat.re_open_chat(chat_id) chat_info.set_cache() @@ -264,42 +294,45 @@ def chat(self, instance: Dict, with_valid=True): if with_valid: self.is_valid(raise_exception=True) OpenAIInstanceSerializer(data=instance).is_valid(raise_exception=True) - chat_id = instance.get('chat_id') + chat_id = instance.get("chat_id") message = self.get_message(instance) - re_chat = instance.get('re_chat', False) - stream = instance.get('stream', False) - application_id = self.data.get('application_id') - chat_user_id = self.data.get('chat_user_id') - chat_user_type = self.data.get('chat_user_type') - ip_address = self.data.get('ip_address') - source = self.data.get('source') + re_chat = instance.get("re_chat", False) + stream = instance.get("stream", False) + application_id = self.data.get("application_id") + chat_user_id = self.data.get("chat_user_id") + chat_user_type = self.data.get("chat_user_type") + ip_address = self.data.get("ip_address") + source = self.data.get("source") chat_id = self.generate_chat(chat_id, application_id, message, chat_user_id, chat_user_type, ip_address, source) return ChatSerializers( data={ - 'chat_id': chat_id, - 'chat_user_id': chat_user_id, - 'chat_user_type': chat_user_type, - 'application_id': application_id, - 'ip_address': ip_address, - 'source': source, + "chat_id": chat_id, + "chat_user_id": chat_user_id, + "chat_user_type": chat_user_type, + "application_id": application_id, + "ip_address": ip_address, + "source": source, } - ).chat({'message': message, - 're_chat': re_chat, - 'stream': stream, - 'form_data': instance.get('form_data', {}), - 'image_list': instance.get('image_list', []), - 'document_list': instance.get('document_list', []), - 'audio_list': instance.get('audio_list', []), - 'other_list': instance.get('other_list', [])}, - base_to_response=OpenaiToResponse()) + ).chat( + { + "message": message, + "re_chat": re_chat, + "stream": stream, + "form_data": instance.get("form_data", {}), + "image_list": instance.get("image_list", []), + "document_list": instance.get("document_list", []), + "audio_list": instance.get("audio_list", []), + "other_list": instance.get("other_list", []), + }, + base_to_response=OpenaiToResponse(), + ) class ChatSerializers(serializers.Serializer): chat_id = serializers.UUIDField(required=True, label=_("Conversation ID")) chat_user_id = serializers.CharField(required=True, label=_("Client id")) chat_user_type = serializers.CharField(required=True, label=_("Client Type")) - application_id = serializers.UUIDField(required=True, allow_null=True, - label=_("Application ID")) + application_id = serializers.UUIDField(required=True, allow_null=True, label=_("Application ID")) debug = serializers.BooleanField(required=False, label=_("Debug")) ip_address = serializers.CharField(required=False, label=_("IP Address"), allow_null=True, allow_blank=True) source = serializers.JSONField(required=False, label=_("Source")) @@ -308,27 +341,34 @@ def is_valid_application_workflow(self, *, raise_exception=False): self.is_valid_intraday_access_num() def is_valid_chat_id(self, chat_info: ChatInfo): - if self.data.get('application_id') is not None and self.data.get('application_id') != str( - chat_info.application_id): + if self.data.get("application_id") is not None and self.data.get("application_id") != str( + chat_info.application_id + ): raise ChatException(500, _("Conversation does not exist")) def is_valid_intraday_access_num(self): - if not self.data.get('debug') and [ChatUserType.ANONYMOUS_USER.value, - ChatUserType.CHAT_USER.value].__contains__( - self.data.get('chat_user_type')): - access_client = QuerySet(ApplicationChatUserStats).filter(chat_user_id=self.data.get('chat_user_id'), - application_id=self.data.get( - 'application_id')).first() + if not self.data.get("debug") and [ + ChatUserType.ANONYMOUS_USER.value, + ChatUserType.CHAT_USER.value, + ].__contains__(self.data.get("chat_user_type")): + access_client = ( + QuerySet(ApplicationChatUserStats) + .filter(chat_user_id=self.data.get("chat_user_id"), application_id=self.data.get("application_id")) + .first() + ) if access_client is None: - access_client = ApplicationChatUserStats(chat_user_id=self.data.get('chat_user_id'), - chat_user_type=self.data.get('chat_user_type'), - application_id=self.data.get('application_id'), - access_num=0, - intraday_access_num=0) + access_client = ApplicationChatUserStats( + chat_user_id=self.data.get("chat_user_id"), + chat_user_type=self.data.get("chat_user_type"), + application_id=self.data.get("application_id"), + access_num=0, + intraday_access_num=0, + ) access_client.save() - application_access_token = QuerySet(ApplicationAccessToken).filter( - application_id=self.data.get('application_id')).first() + application_access_token = ( + QuerySet(ApplicationAccessToken).filter(application_id=self.data.get("application_id")).first() + ) if application_access_token.access_num <= access_client.intraday_access_num: raise AppChatNumOutOfBoundsFailed(1002, _("The number of visits exceeds today's visits")) @@ -347,52 +387,67 @@ def is_valid_application_simple(self, *, chat_info: ChatInfo, raise_exception=Fa return chat_info def chat_simple(self, chat_info: ChatInfo, instance, base_to_response): - message_dict = instance.get('message') - message = message_dict.get('content', '') if isinstance(message_dict, dict) else message_dict - re_chat = instance.get('re_chat') - stream = instance.get('stream') - chat_user_id = self.data.get('chat_user_id') - chat_user_type = self.data.get('chat_user_type') - ip_address = self.data.get('ip_address') - source = self.data.get('source') + message_dict = instance.get("message") + message = message_dict.get("content", "") if isinstance(message_dict, dict) else message_dict + re_chat = instance.get("re_chat") + stream = instance.get("stream") + chat_user_id = self.data.get("chat_user_id") + chat_user_type = self.data.get("chat_user_type") + ip_address = self.data.get("ip_address") + source = self.data.get("source") form_data = instance.get("form_data") - chat_record_id = instance.get('chat_record_id') + chat_record_id = instance.get("chat_record_id") pipeline_manage_builder = PipelineManage.builder() # 如果开启了问题优化,则添加上问题优化步骤 if chat_info.application.problem_optimization: pipeline_manage_builder.append_step(BaseResetProblemStep) # 构建流水线管理器 - pipeline_message = (pipeline_manage_builder.append_step(BaseSearchDatasetStep) - .append_step(BaseGenerateHumanMessageStep) - .append_step(BaseChatStep) - .add_base_to_response(base_to_response) - .add_debug(self.data.get('debug', False)) - .build()) + pipeline_message = ( + pipeline_manage_builder.append_step(BaseSearchDatasetStep) + .append_step(BaseGenerateHumanMessageStep) + .append_step(BaseChatStep) + .add_base_to_response(base_to_response) + .add_debug(self.data.get("debug", False)) + .build() + ) exclude_paragraph_id_list = [] # 相同问题是否需要排除已经查询到的段落 if re_chat: paragraph_id_list = flat_map( - [[paragraph.get('id') for paragraph in chat_record.details['search_step']['paragraph_list']] for - chat_record in chat_info.chat_record_list if - chat_record.problem_text == message and 'search_step' in chat_record.details and 'paragraph_list' in - chat_record.details['search_step']]) + [ + [paragraph.get("id") for paragraph in chat_record.details["search_step"]["paragraph_list"]] + for chat_record in chat_info.chat_record_list + if chat_record.problem_text == message + and "search_step" in chat_record.details + and "paragraph_list" in chat_record.details["search_step"] + ] + ) exclude_paragraph_id_list = list(set(paragraph_id_list)) # 构建运行参数 - params = chat_info.to_pipeline_manage_params(message, get_post_handler(chat_info), exclude_paragraph_id_list, - chat_user_id, chat_user_type, ip_address, source, stream, - form_data) + params = chat_info.to_pipeline_manage_params( + message, + get_post_handler(chat_info), + exclude_paragraph_id_list, + chat_user_id, + chat_user_type, + ip_address, + source, + stream, + form_data, + ) if chat_record_id: - params['chat_record_id'] = chat_record_id + params["chat_record_id"] = chat_record_id chat_info.set_chat(message) # 运行流水线作业 pipeline_message.run(params) - return pipeline_message.context['chat_result'] + return pipeline_message.context["chat_result"] @staticmethod def get_chat_record(chat_info, chat_record_id): if chat_info is not None: - chat_record_list = [chat_record for chat_record in chat_info.chat_record_list if - str(chat_record.id) == str(chat_record_id)] + chat_record_list = [ + chat_record for chat_record in chat_info.chat_record_list if str(chat_record.id) == str(chat_record_id) + ] if chat_record_list is not None and len(chat_record_list): return chat_record_list[-1] chat_record = QuerySet(ChatRecord).filter(id=chat_record_id, chat_id=chat_info.chat_id).first() @@ -405,25 +460,25 @@ def get_chat_record(chat_info, chat_record_id): def chat_work_flow(self, chat_info: ChatInfo, instance: dict, base_to_response): import queue - message_dict = instance.get('message') - message = message_dict.get('content', '') if isinstance(message_dict, dict) else message_dict - re_chat = instance.get('re_chat') - stream = instance.get('stream') + message_dict = instance.get("message") + message = message_dict.get("content", "") if isinstance(message_dict, dict) else message_dict + re_chat = instance.get("re_chat") + stream = instance.get("stream") chat_user_id = self.data.get("chat_user_id") - chat_user_type = self.data.get('chat_user_type') - ip_address = self.data.get('ip_address') - source = self.data.get('source') - form_data = instance.get('form_data') - image_list = message_dict.get('image_list', []) if isinstance(message_dict, dict) else [] - video_list = message_dict.get('video_list', []) if isinstance(message_dict, dict) else [] - document_list = message_dict.get('document_list', []) if isinstance(message_dict, dict) else [] - audio_list = message_dict.get('audio_list', []) if isinstance(message_dict, dict) else [] - other_list = message_dict.get('other_list', []) if isinstance(message_dict, dict) else [] + chat_user_type = self.data.get("chat_user_type") + ip_address = self.data.get("ip_address") + source = self.data.get("source") + form_data = instance.get("form_data") + image_list = message_dict.get("image_list", []) if isinstance(message_dict, dict) else [] + video_list = message_dict.get("video_list", []) if isinstance(message_dict, dict) else [] + document_list = message_dict.get("document_list", []) if isinstance(message_dict, dict) else [] + audio_list = message_dict.get("audio_list", []) if isinstance(message_dict, dict) else [] + other_list = message_dict.get("other_list", []) if isinstance(message_dict, dict) else [] workspace_id = chat_info.application.workspace_id - chat_record_id = instance.get('chat_record_id') - position = instance.get('position') - chunk_id = instance.get('chunk_id') - debug = self.data.get('debug', False) + chat_record_id = instance.get("chat_record_id") + position = instance.get("position") + chunk_id = instance.get("chunk_id") + debug = self.data.get("debug", False) history_chat_record = chat_info.chat_record_list if chat_record_id is not None: chat_record = self.get_chat_record(chat_info, chat_record_id) @@ -436,106 +491,132 @@ def chat_work_flow(self, chat_info: ChatInfo, instance: dict, base_to_response): chat_record_id_str = str(uuid.uuid7()) if chat_record_id is None else str(chat_record_id) parameters = { - 'history_chat_record': history_chat_record, - 'question': message, - 'chat_id': chat_info.chat_id, - 'chat_record_id': chat_record_id_str, - 'stream': stream, - 're_chat': re_chat, - 'chat_user_id': chat_user_id, - 'chat_user_type': chat_user_type, - 'ip_address': ip_address, - 'source': source, - 'workspace_id': workspace_id, - 'debug': debug, - 'chat_user': chat_info.get_chat_user(), - 'chat_user_group': chat_info.get_chat_user_group(), - 'application_id': str(chat_info.application_id), - 'form_data': form_data or {}, - 'position': position, - 'chunk_id': chunk_id, - 'image_list': image_list or [], - 'document_list': document_list or [], - 'audio_list': audio_list or [], - 'video_list': video_list or [], - 'other_list': other_list or [], + "history_chat_record": history_chat_record, + "question": message, + "chat_id": chat_info.chat_id, + "chat_record_id": chat_record_id_str, + "stream": stream, + "re_chat": re_chat, + "chat_user_id": chat_user_id, + "chat_user_type": chat_user_type, + "ip_address": ip_address, + "source": source, + "workspace_id": workspace_id, + "debug": debug, + "chat_user": chat_info.get_chat_user(), + "chat_user_group": chat_info.get_chat_user_group(), + "application_id": str(chat_info.application_id), + "form_data": form_data or {}, + "position": position, + "chunk_id": chunk_id, + "image_list": image_list or [], + "document_list": document_list or [], + "audio_list": audio_list or [], + "video_list": video_list or [], + "other_list": other_list or [], } result_queue = queue.Queue() aggregation = AggregationManager() - self.save_chat_record(chat_info, chat_info.chat_id, chat_record_id_str, - message_dict) + self.save_chat_record(chat_info, chat_info.chat_id, chat_record_id_str, message_dict) def on_next(wf_manage, content): aggregation.aggregate(content) message_queue = get_message_queue() message_queue.produce(chat_record_id_str, content.to_dict()) if isinstance(content, TextContent): - result_queue.put(('chunk', { - 'content': [{ - 'id': content.id, - 'type': 'TEXT', - 'content': content.content, - }] - })) + result_queue.put( + ( + "chunk", + { + "content": [ + { + "id": content.id, + "type": "TEXT", + "content": content.content, + } + ] + }, + ) + ) elif isinstance(content, ReasoningContent): - result_queue.put(('chunk', { - 'content': [{ - 'id': content.id, - 'type': 'REASONING', - 'content': content.content, - 'status': content.status.value if content.status else None, - }] - })) + result_queue.put( + ( + "chunk", + { + "content": [ + { + "id": content.id, + "type": "REASONING", + "content": content.content, + "status": content.status.value if content.status else None, + } + ] + }, + ) + ) elif isinstance(content, ToolContent): - result_queue.put(('chunk', { - 'content': [{ - 'id': content.id, - 'type': 'TOOL', - 'content': content.content, - 'arguments': content.arguments, - 'result': content.result, - 'status': content.status.value if content.status else None, - }] - })) + result_queue.put( + ( + "chunk", + { + "content": [ + { + "id": content.id, + "type": "TOOL", + "content": content.content, + "arguments": content.arguments, + "result": content.result, + "status": content.status.value if content.status else None, + } + ] + }, + ) + ) elif isinstance(content, FormContent): + def position_to_dict(pos): if pos is None: return None - return { - 'id': pos.id, - 'index': pos.index, - 'children': position_to_dict(pos.children) - } - - result_queue.put(('chunk', { - 'content': [{ - 'id': content.id, - 'type': 'FORM', - 'form_field_list': content.form_field_list, - 'form_content_format': content.form_content_format, - 'is_submit': content.is_submit, - 'form_data': content.form_data, - 'status': content.status.value if content.status else None, - 'position': position_to_dict(content.position), - 'chat_record_id': chat_record_id_str, - }] - })) + return {"id": pos.id, "index": pos.index, "children": position_to_dict(pos.children)} + + result_queue.put( + ( + "chunk", + { + "content": [ + { + "id": content.id, + "type": "FORM", + "form_field_list": content.form_field_list, + "form_content_format": content.form_content_format, + "is_submit": content.is_submit, + "form_data": content.form_data, + "status": content.status.value if content.status else None, + "position": position_to_dict(content.position), + "chat_record_id": chat_record_id_str, + } + ] + }, + ) + ) def on_complete(wf_manage, error): # 注销工作流实例 WorkflowRunRegistry.unregister(chat_record_id_str, str(chat_info.chat_id)) message_queue = get_message_queue() if error: - result_queue.put(('error', error)) - message_queue.produce(chat_record_id_str, - FailureContent(str(uuid_utils.uuid7()), str(error), Status.SUCCESS, - None, None).to_dict()) + result_queue.put(("error", error)) + message_queue.produce( + chat_record_id_str, + FailureContent(str(uuid_utils.uuid7()), str(error), Status.SUCCESS, None, None).to_dict(), + ) QuerySet(ChatRecord).filter(id=chat_record_id).update() - self.update_chat_record(chat_info, chat_info.chat_id, chat_record_id_str, wf_manage.context, - aggregation.get_contents()) - result_queue.put(('done', None)) + self.update_chat_record( + chat_info, chat_info.chat_id, chat_record_id_str, wf_manage.context, aggregation.get_contents() + ) + result_queue.put(("done", None)) message_queue.produce_done(chat_record_id_str) call_back = CallBack(on_next, on_complete) @@ -552,16 +633,18 @@ def get_start_node_fn(wf, wm): parameters=parameters, workflow_type=WorkflowType.APPLICATION, call_back=call_back, - get_start_node=get_start_node_fn + get_start_node=get_start_node_fn, ) if work_flow_manage is None: # 恢复失败,回退到正常流程 - work_flow_manage = WorkflowManage(workflow, parameters, WorkflowType.APPLICATION, - call_back, get_start_node_fn) + work_flow_manage = WorkflowManage( + workflow, parameters, WorkflowType.APPLICATION, call_back, get_start_node_fn + ) else: # 正常创建新的 WorkflowManage - work_flow_manage = WorkflowManage(workflow, parameters, WorkflowType.APPLICATION, - call_back, get_start_node_fn) + work_flow_manage = WorkflowManage( + workflow, parameters, WorkflowType.APPLICATION, call_back, get_start_node_fn + ) work_flow_manage.start_node.workflow_manage = work_flow_manage @@ -571,37 +654,44 @@ def get_start_node_fn(wf, wm): chat_info.set_chat(message) if stream: + def generate(): work_flow_manage.run() while True: msg_type, data = result_queue.get() - if msg_type == 'done': - yield 'data: [DONE]\n\n' + if msg_type == "done": + yield "data: [DONE]\n\n" break - if msg_type == 'error': - yield 'data: ' + json.dumps({ - 'chat_id': str(chat_info.chat_id), - 'chat_record_id': chat_record_id_str, - 'content': [{'type': 'FAILURE', 'content': str(data)}] - }, ensure_ascii=False) + '\n\n' - yield 'data: [DONE]\n\n' + if msg_type == "error": + yield ( + "data: " + + json.dumps( + { + "chat_id": str(chat_info.chat_id), + "chat_record_id": chat_record_id_str, + "content": [{"type": "FAILURE", "content": str(data)}], + }, + ensure_ascii=False, + ) + + "\n\n" + ) + yield "data: [DONE]\n\n" break - if msg_type == 'chunk': - data['chat_id'] = str(chat_info.chat_id) - data['chat_record_id'] = chat_record_id_str - yield 'data: ' + json.dumps(data, ensure_ascii=False) + '\n\n' + if msg_type == "chunk": + data["chat_id"] = str(chat_info.chat_id) + data["chat_record_id"] = chat_record_id_str + yield "data: " + json.dumps(data, ensure_ascii=False) + "\n\n" return to_stream_response_simple(generate()) else: work_flow_manage.run() while True: msg_type, data = result_queue.get() - if msg_type == 'done': + if msg_type == "done": break - if msg_type == 'error': + if msg_type == "error": raise data - return base_to_response.to_block_response( - chat_info.chat_id, chat_record_id_str, '', True, 0, 0) + return base_to_response.to_block_response(chat_info.chat_id, chat_record_id_str, "", True, 0, 0) @staticmethod def save_chat_record(chat_info, chat_id, chat_record_id, question): @@ -620,7 +710,7 @@ def save_chat_record(chat_info, chat_id, chat_record_id, question): source=chat_info.source, workflow_context={}, question=question, - messages=[] + messages=[], ) chat_info.append_chat_record(chat_record) chat_info.set_cache() @@ -628,26 +718,32 @@ def save_chat_record(chat_info, chat_id, chat_record_id, question): @staticmethod def update_chat_record(chat_info, chat_id, chat_record_id, workflow_context, messages): message_tokens = sum( - v.get('message_tokens', 0) for v in workflow_context.values() if - isinstance(v, dict) and 'message_tokens' in v) + v.get("message_tokens", 0) + for v in workflow_context.values() + if isinstance(v, dict) and "message_tokens" in v + ) answer_tokens = sum( - v.get('answer_tokens', 0) for v in workflow_context.values() if - isinstance(v, dict) and 'answer_tokens' in v) + v.get("answer_tokens", 0) for v in workflow_context.values() if isinstance(v, dict) and "answer_tokens" in v + ) + ChatUserTokenQuota.consume(chat_info.chat_user_id, message_tokens + answer_tokens) QuerySet(ChatRecord).filter(id=chat_record_id).update( workflow_context=workflow_context, messages=messages, message_tokens=message_tokens, - answer_tokens=answer_tokens + answer_tokens=answer_tokens, ) def is_valid_chat_user(self): - chat_user_id = self.data.get('chat_user_id') - application_id = self.data.get('application_id') - chat_user_type = self.data.get('chat_user_type') + chat_user_id = self.data.get("chat_user_id") + application_id = self.data.get("application_id") + chat_user_type = self.data.get("chat_user_type") is_auth_chat_user = DatabaseModelManage.get_model("is_auth_chat_user") application_access_token = QuerySet(ApplicationAccessToken).filter(application_id=application_id).first() - if application_access_token and application_access_token.authentication and application_access_token.authentication_value.get( - 'type') == 'login': + if ( + application_access_token + and application_access_token.authentication + and application_access_token.authentication_value.get("type") == "login" + ): if chat_user_type == ChatUserType.ANONYMOUS_USER.value: raise ChatException(500, _("The chat user is not authorized.")) if chat_user_type == ChatUserType.CHAT_USER.value and is_auth_chat_user: @@ -660,10 +756,11 @@ def chat(self, instance: dict, base_to_response: BaseToResponse = SystemToRespon ChatMessageSerializers(data=instance).is_valid(raise_exception=True) chat_info = self.get_chat_info() chat_info.get_application() - chat_info.get_chat_user(asker=(instance.get('form_data') or {}).get('asker')) + chat_info.get_chat_user(asker=(instance.get("form_data") or {}).get("asker")) self.is_valid_chat_id(chat_info) - if not self.data.get('debug'): + if not self.data.get("debug"): self.is_valid_chat_user() + ChatUserTokenQuota.consume(chat_info.chat_user_id, 0) # 触发周期重置 + 配额预校验 if chat_info.application.type == ApplicationTypeChoices.SIMPLE: self.is_valid_application_simple(raise_exception=True, chat_info=chat_info) return self.chat_simple(chat_info, instance, base_to_response) @@ -673,7 +770,7 @@ def chat(self, instance: dict, base_to_response: BaseToResponse = SystemToRespon def get_chat_info(self): self.is_valid(raise_exception=True) - chat_id = self.data.get('chat_id') + chat_id = self.data.get("chat_id") chat_info: ChatInfo = ChatInfo.get_cache(chat_id) if chat_info is None: chat_info: ChatInfo = self.re_open_chat(chat_id) @@ -687,8 +784,9 @@ def re_open_chat(self, chat_id: str): application = QuerySet(Application).filter(id=chat.application_id).first() if application is None: raise ChatException(500, _("Application does not exist")) - application_version = QuerySet(ApplicationVersion).filter(application_id=application.id).order_by( - '-create_time')[0:1].first() + application_version = ( + QuerySet(ApplicationVersion).filter(application_id=application.id).order_by("-create_time")[0:1].first() + ) if application_version is None: raise ChatException(500, _("The application has not been published. Please use it after publishing.")) if application.type == ApplicationTypeChoices.SIMPLE: @@ -697,38 +795,53 @@ def re_open_chat(self, chat_id: str): return self.re_open_chat_work_flow(chat_id, application) def re_open_chat_simple(self, chat_id, application): - if self.data.get('debug'): + if self.data.get("debug"): # 数据集id列表 - knowledge_id_list = [str(row.target_id) for row in - QuerySet(ResourceMapping).filter(source_id=str(application.id), - source_type='APPLICATION', - target_type='KNOWLEDGE')] + knowledge_id_list = [ + str(row.target_id) + for row in QuerySet(ResourceMapping).filter( + source_id=str(application.id), source_type="APPLICATION", target_type="KNOWLEDGE" + ) + ] else: - application_version = QuerySet(ApplicationVersion).filter(application_id=application.id).order_by( - '-create_time')[0:1].first() + application_version = ( + QuerySet(ApplicationVersion).filter(application_id=application.id).order_by("-create_time")[0:1].first() + ) knowledge_id_list = application_version.knowledge_ids # 需要排除的文档 - exclude_document_id_list = [str(document.id) for document in - QuerySet(Document).filter( - knowledge_id__in=knowledge_id_list, - is_active=False)] - chat_info = ChatInfo(chat_id, self.data.get('chat_user_id'), self.data.get('chat_user_type'), - self.data.get('ip_address'), - self.data.get('source'), knowledge_id_list, - exclude_document_id_list, application.id) - chat_record_list = list(QuerySet(ChatRecord).filter(chat_id=chat_id).order_by('-create_time')[0:5]) + exclude_document_id_list = [ + str(document.id) + for document in QuerySet(Document).filter(knowledge_id__in=knowledge_id_list, is_active=False) + ] + chat_info = ChatInfo( + chat_id, + self.data.get("chat_user_id"), + self.data.get("chat_user_type"), + self.data.get("ip_address"), + self.data.get("source"), + knowledge_id_list, + exclude_document_id_list, + application.id, + ) + chat_record_list = list(QuerySet(ChatRecord).filter(chat_id=chat_id).order_by("-create_time")[0:5]) chat_record_list.sort(key=lambda r: r.create_time) for chat_record in chat_record_list: chat_info.chat_record_list.append(chat_record) return chat_info def re_open_chat_work_flow(self, chat_id, application): - chat_info = ChatInfo(chat_id, self.data.get('chat_user_id'), self.data.get('chat_user_type'), - self.data.get('ip_address'), - self.data.get('source'), [], [], - application.id) - chat_record_list = list(QuerySet(ChatRecord).filter(chat_id=chat_id).order_by('-create_time')[0:5]) + chat_info = ChatInfo( + chat_id, + self.data.get("chat_user_id"), + self.data.get("chat_user_type"), + self.data.get("ip_address"), + self.data.get("source"), + [], + [], + application.id, + ) + chat_record_list = list(QuerySet(ChatRecord).filter(chat_id=chat_id).order_by("-create_time")[0:5]) chat_record_list.sort(key=lambda r: r.create_time) for chat_record in chat_record_list: chat_info.chat_record_list.append(chat_record) @@ -749,7 +862,8 @@ def resume(self, request): self.is_valid(raise_exception=True) from application.workflow.message_queue import get_message_queue from application.models import ChatRecord - chat_record_id = self.data.get('chat_record_id') + + chat_record_id = self.data.get("chat_record_id") mq = get_message_queue() start_id = self._resolve_start_id(request) @@ -761,15 +875,15 @@ def resume(self, request): else: chat_record = ChatRecord.objects.filter(id=chat_record_id).first() if not chat_record: - return result.error(_('Chat record not found')) + return result.error(_("Chat record not found")) generator = self._stream_from_db(chat_record, start_id) response = StreamingHttpResponse( generator, - content_type='text/event-stream;charset=utf-8', + content_type="text/event-stream;charset=utf-8", ) - response['Cache-Control'] = 'no-cache' - response['X-Accel-Buffering'] = 'no' + response["Cache-Control"] = "no-cache" + response["X-Accel-Buffering"] = "no" return response @staticmethod @@ -779,12 +893,12 @@ def _resolve_start_id(request: Request) -> str: 兼容 body / query 里显式传的 last_event_id。取不到则从头开始。 """ candidate = ( - request.META.get('HTTP_LAST_EVENT_ID') - or (request.data.get('last_event_id') if hasattr(request, 'data') else None) - or request.query_params.get('last_event_id') + request.META.get("HTTP_LAST_EVENT_ID") + or (request.data.get("last_event_id") if hasattr(request, "data") else None) + or request.query_params.get("last_event_id") ) - candidate = (candidate or '').strip() - return candidate or '0' + candidate = (candidate or "").strip() + return candidate or "0" @staticmethod def _sse(msg_id: str, msg_data: str) -> str: @@ -792,7 +906,7 @@ def _sse(msg_id: str, msg_data: str) -> str: 带 id 字段的 SSE 帧:浏览器会把最后收到的 id 存进 Last-Event-ID, 下次重连自动回传,从而实现断点续传。 """ - return f'id: {msg_id}\ndata: {msg_data}\n\n' + return f"id: {msg_id}\ndata: {msg_data}\n\n" def _stream_from_queue(self, mq, chat_record_id: str, start_id: str): """ @@ -822,9 +936,7 @@ def pump(): except thread_queue.Full: pass - worker = threading.Thread( - target=pump, name=f"resume-{chat_record_id}", daemon=True - ) + worker = threading.Thread(target=pump, name=f"resume-{chat_record_id}", daemon=True) worker.start() try: @@ -833,15 +945,13 @@ def pump(): # 略大于 consume timeout:正常情况下哨兵会先到,这里只防线程异常挂死 item = bridge.get(timeout=_CONSUME_TIMEOUT + 5) except thread_queue.Empty: - maxkb_logger.warning( - f"ResumeStream bridge idle timeout [{chat_record_id}]" - ) + maxkb_logger.warning(f"ResumeStream bridge idle timeout [{chat_record_id}]") break if item is done_sentinel: break msg_id, msg_data = item yield self._sse(msg_id, msg_data) - yield 'data: [DONE]\n\n' + yield "data: [DONE]\n\n" finally: # 客户端提前关闭连接会在 yield 处抛 GeneratorExit,落到这里; # 通知 consume 线程停止,不必再等到 300s 超时 @@ -854,11 +964,11 @@ def _stream_from_db(self, chat_record, start_id: str): """ try: messages = chat_record.messages or [] - resuming = start_id and start_id != '0' + resuming = start_id and start_id != "0" passed = not resuming # 无续传点则全部下发 for msg in messages: - msg_id = str(msg.get('id', '')) if isinstance(msg, dict) else '' + msg_id = str(msg.get("id", "")) if isinstance(msg, dict) else "" if not passed: # 尚未越过续传点:命中该 id 后,从下一条开始发 @@ -871,10 +981,10 @@ def _stream_from_db(self, chat_record, start_id: str): # 续传点在库里没匹配到(比如 id 体系不一致):退化为整段重放,别让客户端收到空流 if not passed: for msg in messages: - msg_id = str(msg.get('id', '')) if isinstance(msg, dict) else '' + msg_id = str(msg.get("id", "")) if isinstance(msg, dict) else "" yield self._sse(msg_id, json.dumps(msg, ensure_ascii=False)) finally: - yield 'data: [DONE]\n\n' + yield "data: [DONE]\n\n" class OpenChatSerializers(serializers.Serializer): @@ -888,25 +998,25 @@ class OpenChatSerializers(serializers.Serializer): def is_valid(self, *, raise_exception=False): super().is_valid(raise_exception=True) - workspace_id = self.data.get('workspace_id') - application_id = self.data.get('application_id') + workspace_id = self.data.get("workspace_id") + application_id = self.data.get("application_id") query_set = QuerySet(Application).filter(id=application_id) if workspace_id: query_set = query_set.filter(workspace_id=workspace_id) if not query_set.exists(): - raise AppApiException(500, gettext('Application does not exist')) + raise AppApiException(500, gettext("Application does not exist")) def open(self): self.is_valid(raise_exception=True) - application_id = self.data.get('application_id') + application_id = self.data.get("application_id") application = QuerySet(Application).get(id=application_id) debug = self.data.get("debug") if not debug: - application_version = QuerySet(ApplicationVersion).filter(application_id=application_id).order_by( - '-create_time')[0:1].first() + application_version = ( + QuerySet(ApplicationVersion).filter(application_id=application_id).order_by("-create_time")[0:1].first() + ) if application_version is None: - raise AppApiException(500, - _("The application has not been published. Please use it after publishing.")) + raise AppApiException(500, _("The application has not been published. Please use it after publishing.")) if application.type == ApplicationTypeChoices.SIMPLE: return self.open_simple(application) else: @@ -914,45 +1024,53 @@ def open(self): def open_work_flow(self, application): self.is_valid(raise_exception=True) - application_id = self.data.get('application_id') + application_id = self.data.get("application_id") chat_user_id = self.data.get("chat_user_id") chat_user_type = self.data.get("chat_user_type") ip_address = self.data.get("ip_address") source = self.data.get("source") debug = self.data.get("debug") chat_id = str(uuid.uuid7()) - chat_info = ChatInfo(chat_id, chat_user_id, chat_user_type, ip_address, source, [], - [], - application_id, debug) + chat_info = ChatInfo(chat_id, chat_user_id, chat_user_type, ip_address, source, [], [], application_id, debug) chat_info.save_chat() chat_info.set_cache() return chat_id def open_simple(self, application): - application_id = self.data.get('application_id') + application_id = self.data.get("application_id") chat_user_id = self.data.get("chat_user_id") chat_user_type = self.data.get("chat_user_type") ip_address = self.data.get("ip_address") source = self.data.get("source") debug = self.data.get("debug") if debug: - knowledge_id_list = [str(row.target_id) for row in - QuerySet(ResourceMapping).filter(source_id=str(application_id), - source_type='APPLICATION', - target_type='KNOWLEDGE')] + knowledge_id_list = [ + str(row.target_id) + for row in QuerySet(ResourceMapping).filter( + source_id=str(application_id), source_type="APPLICATION", target_type="KNOWLEDGE" + ) + ] else: - application_version = QuerySet(ApplicationVersion).filter(application_id=application_id).order_by( - '-create_time')[0:1].first() + application_version = ( + QuerySet(ApplicationVersion).filter(application_id=application_id).order_by("-create_time")[0:1].first() + ) knowledge_id_list = application_version.knowledge_ids chat_id = str(uuid.uuid7()) - chat_info = ChatInfo(chat_id, chat_user_id, chat_user_type, ip_address, source, knowledge_id_list, - [str(document.id) for document in - QuerySet(Document).filter( - knowledge_id__in=knowledge_id_list, - is_active=False)], - application_id, - debug=debug) + chat_info = ChatInfo( + chat_id, + chat_user_id, + chat_user_type, + ip_address, + source, + knowledge_id_list, + [ + str(document.id) + for document in QuerySet(Document).filter(knowledge_id__in=knowledge_id_list, is_active=False) + ], + application_id, + debug=debug, + ) chat_info.save_chat() chat_info.set_cache() return chat_id @@ -963,11 +1081,11 @@ class TextToSpeechSerializers(serializers.Serializer): def text_to_speech(self, instance): self.is_valid(raise_exception=True) - application_id = self.data.get('application_id') + application_id = self.data.get("application_id") application = QuerySet(Application).filter(id=application_id).first() return ApplicationOperateSerializer( - data={'application_id': application_id, - 'user_id': application.user_id}).text_to_speech(instance, False) + data={"application_id": application_id, "user_id": application.user_id} + ).text_to_speech(instance, False) class SpeechToTextSerializers(serializers.Serializer): @@ -975,8 +1093,8 @@ class SpeechToTextSerializers(serializers.Serializer): def speech_to_text(self, instance): self.is_valid(raise_exception=True) - application_id = self.data.get('application_id') + application_id = self.data.get("application_id") application = QuerySet(Application).filter(id=application_id).first() return ApplicationOperateSerializer( - data={'application_id': application_id, - 'user_id': application.user_id}).speech_to_text(instance, False) + data={"application_id": application_id, "user_id": application.user_id} + ).speech_to_text(instance, False) diff --git a/apps/locales/en_US/LC_MESSAGES/django.po b/apps/locales/en_US/LC_MESSAGES/django.po index ac4e37dee23..97b431376d6 100644 --- a/apps/locales/en_US/LC_MESSAGES/django.po +++ b/apps/locales/en_US/LC_MESSAGES/django.po @@ -9591,3 +9591,6 @@ msgstr "Get chat user quota" #: apps/xpack/views/system_chat_user.py:82 msgid "Set chat user quota" msgstr "Set chat user quota" + +msgid "The token quota for the current period has been exhausted. Please contact the administrator." +msgstr "" diff --git a/apps/locales/zh_CN/LC_MESSAGES/django.po b/apps/locales/zh_CN/LC_MESSAGES/django.po index 6d55bb78cb5..c65b4f24ea3 100644 --- a/apps/locales/zh_CN/LC_MESSAGES/django.po +++ b/apps/locales/zh_CN/LC_MESSAGES/django.po @@ -9731,6 +9731,9 @@ msgstr "获取对话用户配额" msgid "Set chat user quota" msgstr "设置对话用户配额" +msgid "The token quota for the current period has been exhausted. Please contact the administrator." +msgstr "当前周期 Tokens 配额已用尽,请联系管理员。" + #: apps/xpack/views/system_chat_user.py:101 msgid "Batch set chat user quota" msgstr "批量设置对话用户配额" \ No newline at end of file diff --git a/apps/locales/zh_Hant/LC_MESSAGES/django.po b/apps/locales/zh_Hant/LC_MESSAGES/django.po index bbce67af6d3..acab840cdae 100644 --- a/apps/locales/zh_Hant/LC_MESSAGES/django.po +++ b/apps/locales/zh_Hant/LC_MESSAGES/django.po @@ -9729,4 +9729,7 @@ msgstr "獲取對話用戶配額" #: apps/xpack/views/system_chat_user.py:82 msgid "Set chat user quota" -msgstr "設置對話用戶配額" \ No newline at end of file +msgstr "設置對話用戶配額" + +msgid "The token quota for the current period has been exhausted. Please contact the administrator." +msgstr "當前週期 Tokens 配額已用盡,請聯繫管理員。" \ No newline at end of file diff --git a/apps/system_manage/migrations/0008_add_chat_user_token_quota.py b/apps/system_manage/migrations/0008_add_chat_user_token_quota.py index 7463225b620..52853ec39ff 100644 --- a/apps/system_manage/migrations/0008_add_chat_user_token_quota.py +++ b/apps/system_manage/migrations/0008_add_chat_user_token_quota.py @@ -4,30 +4,85 @@ from django.db import migrations, models -class Migration(migrations.Migration): +def migrate_historical_tokens(apps, schema_editor): + ChatRecord = apps.get_model('application', 'ChatRecord') + Chat = apps.get_model('application', 'Chat') + ChatUserTokenQuota = apps.get_model('system_manage', 'ChatUserTokenQuota') + + chat_user_map = { + str(c['id']): c['chat_user_id'] + for c in Chat.objects.values('id', 'chat_user_id').iterator() + } + user_totals = {} + queryset = ChatRecord.objects.filter(message_tokens__isnull=False) | ChatRecord.objects.filter( + answer_tokens__isnull=False) + for record in queryset.values('chat_id', 'message_tokens', 'answer_tokens').iterator(chunk_size=5000): + user_id = chat_user_map.get(str(record['chat_id'])) + if not user_id: + continue + tokens = (record['message_tokens'] or 0) + (record['answer_tokens'] or 0) + user_totals[user_id] = user_totals.get(user_id, 0) + tokens + + objs = [ + ChatUserTokenQuota(user_id=uid, total_tokens=total) + for uid, total in user_totals.items() + ] + ChatUserTokenQuota.objects.bulk_create(objs, batch_size=500) + +class Migration(migrations.Migration): dependencies = [ - ('system_manage', '0007_workspaceusergroupresourcepermission'), + ("system_manage", "0007_workspaceusergroupresourcepermission"), ] operations = [ migrations.CreateModel( - name='ChatUserTokenQuota', + name="ChatUserTokenQuota", fields=[ - ('create_time', models.DateTimeField(auto_now_add=True, db_index=True, verbose_name='创建时间')), - ('update_time', models.DateTimeField(auto_now=True, db_index=True, verbose_name='修改时间')), - ('id', models.UUIDField(default=uuid_utils.compat.uuid7, editable=False, primary_key=True, serialize=False, verbose_name='主键id')), - ('user_id', models.UUIDField(db_index=True, verbose_name='用户id')), - ('quota_type', models.CharField(choices=[('UNLIMITED', '不限额'), ('PERIODIC', '按周期限制')], default='UNLIMITED', max_length=20, verbose_name='配额模式')), - ('period_type', models.CharField(blank=True, choices=[('DAY', '天'), ('WEEK', '周'), ('MONTH', '月')], max_length=10, null=True, verbose_name='周期单位')), - ('period_value', models.PositiveIntegerField(blank=True, null=True, verbose_name='周期数量')), - ('token_limit', models.BigIntegerField(blank=True, null=True, verbose_name='Tokens上限')), - ('used_tokens', models.BigIntegerField(default=0, verbose_name='当前周期已使用Tokens')), - ('total_tokens', models.BigIntegerField(default=0, verbose_name='累计Tokens')), - ('period_end', models.DateTimeField(blank=True, null=True, verbose_name='当前周期结束时间')), + ("create_time", models.DateTimeField(auto_now_add=True, db_index=True, verbose_name="创建时间")), + ("update_time", models.DateTimeField(auto_now=True, db_index=True, verbose_name="修改时间")), + ( + "id", + models.UUIDField( + default=uuid_utils.compat.uuid7, + editable=False, + primary_key=True, + serialize=False, + verbose_name="主键id", + ), + ), + ("user_id", models.UUIDField(unique=True, verbose_name="用户id")), + ( + "quota_type", + models.CharField( + choices=[("UNLIMITED", "不限额"), ("PERIODIC", "按周期限制")], + default="UNLIMITED", + max_length=20, + verbose_name="配额模式", + ), + ), + ( + "period_type", + models.CharField( + blank=True, + choices=[("DAY", "天"), ("WEEK", "周"), ("MONTH", "月")], + max_length=10, + null=True, + verbose_name="周期单位", + ), + ), + ("period_value", models.PositiveIntegerField(blank=True, null=True, verbose_name="周期数量")), + ("token_limit", models.BigIntegerField(blank=True, null=True, verbose_name="Tokens上限")), + ("used_tokens", models.BigIntegerField(default=0, verbose_name="当前周期已使用Tokens")), + ("total_tokens", models.BigIntegerField(default=0, verbose_name="累计Tokens")), + ("period_end", models.DateTimeField(blank=True, null=True, verbose_name="当前周期结束时间")), ], options={ - 'db_table': 'chat_user_token_quota', + "db_table": "chat_user_token_quota", }, ), + migrations.RunPython( + migrate_historical_tokens, + migrations.RunPython.noop, + ), ] diff --git a/apps/system_manage/models/chat_user_token_quota.py b/apps/system_manage/models/chat_user_token_quota.py index 9e6e0717ef5..34bb8d61d3b 100644 --- a/apps/system_manage/models/chat_user_token_quota.py +++ b/apps/system_manage/models/chat_user_token_quota.py @@ -7,7 +7,11 @@ import uuid_utils.compat as uuid from django.db import models +from common.exception.app_exception import AppApiException from common.mixins.app_model_mixin import AppModelMixin +from dateutil.relativedelta import relativedelta +from django.utils import timezone +from django.utils.translation import gettext_lazy as _ class QuotaType(models.TextChoices): @@ -47,3 +51,30 @@ class ChatUserTokenQuota(AppModelMixin): class Meta: db_table = "chat_user_token_quota" + + def check_and_reset(self): + if self.quota_type != QuotaType.PERIODIC or self.period_end is None: + return + now = timezone.now() + if now < self.period_end: + return + while self.period_end <= now: + self.period_end += relativedelta( + **{f'{self.period_type.lower()}s': self.period_value} + ) + self.used_tokens = 0 + self.save(update_fields=['used_tokens', 'period_end']) + + @classmethod + def consume(cls, user_id, amount): + if amount <= 0: + return + quota = cls.objects.filter(user_id=user_id).first() + if quota is None or quota.quota_type == QuotaType.UNLIMITED: + return + quota.check_and_reset() + if quota.used_tokens + amount > quota.token_limit: + raise AppApiException(500, _("The token quota for the current period has been exhausted. Please contact the administrator.")) + quota.used_tokens += amount + quota.total_tokens += amount + quota.save(update_fields=['used_tokens', 'total_tokens']) diff --git a/apps/system_manage/serializers/chat_user.py b/apps/system_manage/serializers/chat_user.py index f512fb27218..e6049d931f5 100644 --- a/apps/system_manage/serializers/chat_user.py +++ b/apps/system_manage/serializers/chat_user.py @@ -3,10 +3,13 @@ import re from collections import defaultdict +from dateutil.relativedelta import relativedelta + import uuid_utils.compat as uuid from django.core import validators from django.db import transaction from django.db.models import Q, QuerySet +from django.utils import timezone from django.utils.translation import gettext_lazy as _ from rest_framework import serializers @@ -16,14 +19,14 @@ from common.utils.common import password_encrypt from common.utils.rsa_util import decrypt from system_manage.models import ChatUser, UserGroup, UserGroupRelation +from system_manage.models.chat_user_token_quota import ChatUserTokenQuota from users.serializers.user import PASSWORD_REGEX class ChatUserInstanceSerializer(serializers.ModelSerializer): class Meta: model = ChatUser - fields = ['id', 'username', 'email', 'phone', 'is_active', 'nick_name', 'create_time', 'update_time', - 'source'] + fields = ["id", "username", "email", "phone", "is_active", "nick_name", "create_time", "update_time", "source"] @transaction.atomic @@ -33,12 +36,19 @@ def add_or_edit_user_group_relation(user, user_group_ids): return groups = UserGroup.objects.filter(id__in=user_group_ids) if groups.count() != len(user_group_ids): - raise AppApiException(500, _('Some user groups do not exist')) + raise AppApiException(500, _("Some user groups do not exist")) + + UserGroupRelation.objects.bulk_create([UserGroupRelation(user=user, group=group) for group in groups]) + - UserGroupRelation.objects.bulk_create([ - UserGroupRelation(user=user, group=group) - for group in groups - ]) +def _format_tokens(count): + if count is None: + return "不限" + if count >= 1_000_000: + return f"{count / 1_000_000:.1f}M" + if count >= 1_000: + return f"{count / 1_000:.1f}K" + return str(count) class ChatUserSerializer(serializers.Serializer): @@ -46,12 +56,14 @@ class UserInstance(serializers.Serializer): email = serializers.EmailField( required=False, label=_("Email"), - validators=[validators.EmailValidator( - message=ExceptionCodeConstants.EMAIL_FORMAT_ERROR.value.message, - code=ExceptionCodeConstants.EMAIL_FORMAT_ERROR.value.code - )], + validators=[ + validators.EmailValidator( + message=ExceptionCodeConstants.EMAIL_FORMAT_ERROR.value.message, + code=ExceptionCodeConstants.EMAIL_FORMAT_ERROR.value.code, + ) + ], allow_null=True, - allow_blank=True + allow_blank=True, ) username = serializers.CharField( required=True, @@ -60,10 +72,9 @@ class UserInstance(serializers.Serializer): min_length=4, validators=[ validators.RegexValidator( - regex=re.compile("^.{4,64}$"), - message=_('Username must be 4-64 characters long') + regex=re.compile("^.{4,64}$"), message=_("Username must be 4-64 characters long") ) - ] + ], ) password = serializers.CharField( required=True, @@ -75,9 +86,9 @@ class UserInstance(serializers.Serializer): regex=PASSWORD_REGEX, message=_( "The password must be 6-20 characters long and must be a combination of letters, numbers, and special characters." - ) + ), ) - ] + ], ) nick_name = serializers.CharField( required=True, @@ -85,31 +96,20 @@ class UserInstance(serializers.Serializer): max_length=64, ) phone = serializers.CharField( - required=False, - label=_("Phone"), - max_length=20, - allow_null=True, - allow_blank=True + required=False, label=_("Phone"), max_length=20, allow_null=True, allow_blank=True ) user_group_ids = serializers.ListField( - child=serializers.CharField(required=True), - required=False, - label=_('User Group IDs') - ) - source = serializers.CharField( - required=False, - label=_("Source"), - max_length=20, - default="LOCAL" + child=serializers.CharField(required=True), required=False, label=_("User Group IDs") ) + source = serializers.CharField(required=False, label=_("Source"), max_length=20, default="LOCAL") def is_valid(self, *, raise_exception=True): super().is_valid(raise_exception=True) self._check_unique_username_and_email() def _check_unique_username_and_email(self): - username = self.data.get('username') - nick_name = self.data.get('nick_name') + username = self.data.get("username") + nick_name = self.data.get("nick_name") user = ChatUser.objects.filter(Q(username=username) | Q(nick_name=nick_name)).first() if user: if user.username == username: @@ -118,44 +118,23 @@ def _check_unique_username_and_email(self): raise ExceptionCodeConstants.NICKNAME_IS_EXIST.value.to_app_api_exception() class Query(serializers.Serializer): - username = serializers.CharField( - required=False, - label=_('Username'), - allow_null=True, - allow_blank=True - ) - nick_name = serializers.CharField( - required=False, - label=_('Nickname'), - allow_null=True, - allow_blank=True - ) - source = serializers.CharField( - required=False, - label=_('Source'), - allow_null=True, - allow_blank=True - ) - is_active = serializers.BooleanField( - required=False, - label=_("Is active"), - allow_null=True - ) + username = serializers.CharField(required=False, label=_("Username"), allow_null=True, allow_blank=True) + nick_name = serializers.CharField(required=False, label=_("Nickname"), allow_null=True, allow_blank=True) + source = serializers.CharField(required=False, label=_("Source"), allow_null=True, allow_blank=True) + is_active = serializers.BooleanField(required=False, label=_("Is active"), allow_null=True) def get_query_set(self): - username = self.data.get('username') + username = self.data.get("username") query_set = QuerySet(ChatUser) if username is not None: - query_set = query_set.filter( - Q(username__contains=username)) - nick_name = self.data.get('nick_name') + query_set = query_set.filter(Q(username__contains=username)) + nick_name = self.data.get("nick_name") if nick_name is not None: - query_set = query_set.filter( - Q(nick_name__contains=nick_name)) - source = self.data.get('source') + query_set = query_set.filter(Q(nick_name__contains=nick_name)) + source = self.data.get("source") if source is not None: query_set = query_set.filter(source=source) - is_active = self.data.get('is_active', None) + is_active = self.data.get("is_active", None) if is_active is not None: query_set = query_set.filter(is_active=is_active) query_set = query_set.order_by("-create_time") @@ -164,82 +143,93 @@ def get_query_set(self): def page(self, current_page: int, page_size: int, with_valid=True): if with_valid: self.is_valid(raise_exception=True) - result = page_search(current_page, page_size, - self.get_query_set(), - post_records_handler=lambda u: ChatUserInstanceSerializer(u).data) - user_ids = [user['id'] for user in result['records']] - user_groups = UserGroupRelation.objects.filter( - user__id__in=user_ids - ).select_related('group') + result = page_search( + current_page, + page_size, + self.get_query_set(), + post_records_handler=lambda u: ChatUserInstanceSerializer(u).data, + ) + user_ids = [user["id"] for user in result["records"]] + user_groups = UserGroupRelation.objects.filter(user__id__in=user_ids).select_related("group") - user_groups_map = defaultdict(lambda: {'user_group_ids': [], 'user_group_names': []}) + user_groups_map = defaultdict(lambda: {"user_group_ids": [], "user_group_names": []}) for relation in user_groups: - user_groups_map[str(relation.user_id)]['user_group_ids'].append(str(relation.group_id)) - user_groups_map[str(relation.user_id)]['user_group_names'].append(relation.group.name) - - for user in result['records']: - user.update(user_groups_map.get(str(user['id']), {'user_group_ids': [], 'user_group_names': []})) + user_groups_map[str(relation.user_id)]["user_group_ids"].append(str(relation.group_id)) + user_groups_map[str(relation.user_id)]["user_group_names"].append(relation.group.name) + + for user in result["records"]: + user.update(user_groups_map.get(str(user["id"]), {"user_group_ids": [], "user_group_names": []})) + + # 合并 Token 配额数据 + quotas = ChatUserTokenQuota.objects.filter(user_id__in=user_ids) + now = timezone.now() + quota_map = {} + for q in quotas: + effective_used = q.used_tokens + effective_period_end = q.period_end + if q.quota_type == "PERIODIC" and q.period_end and now >= q.period_end: + effective_used = 0 + effective_period_end = q.period_end + delta_kwargs = {f'{q.period_type.lower()}s': q.period_value} + while effective_period_end <= now: + effective_period_end += relativedelta(**delta_kwargs) + quota_map[str(q.user_id)] = { + "quota_type": q.quota_type, + "used_tokens": effective_used, + "token_limit": q.token_limit, + "total_tokens": q.total_tokens, + "period_end": effective_period_end.isoformat() if effective_period_end else None, + } + for user in result["records"]: + quota = quota_map.get(str(user["id"]), None) + user["token_quota"] = quota return result class BatchDeleteInstance(serializers.Serializer): - ids = serializers.ListField( - child=serializers.UUIDField(required=True), - required=True, - label=_('User IDs') - ) + ids = serializers.ListField(child=serializers.UUIDField(required=True), required=True, label=_("User IDs")) def batch_delete(self): - user_ids = self.data.get('ids') + user_ids = self.data.get("ids") if not user_ids: - raise AppApiException(1004, _('User IDs cannot be empty')) + raise AppApiException(1004, _("User IDs cannot be empty")) ChatUser.objects.filter(id__in=user_ids).delete() return True class BatchAddGroup(serializers.Serializer): - ids = serializers.ListField( - child=serializers.UUIDField(required=True), - required=True, - label=_('User IDs') - ) + ids = serializers.ListField(child=serializers.UUIDField(required=True), required=True, label=_("User IDs")) user_group_ids = serializers.ListField( - child=serializers.CharField(required=True), - required=True, - label=_('User Group IDs') - ) - is_append = serializers.BooleanField( - required=False, - label=_('Is Append'), - default=False + child=serializers.CharField(required=True), required=True, label=_("User Group IDs") ) + is_append = serializers.BooleanField(required=False, label=_("Is Append"), default=False) @transaction.atomic def batch_add_group(self, with_valid=True): if with_valid: self.is_valid(raise_exception=True) - user_ids = self.data.get('ids') - original_group_ids = self.data.get('user_group_ids') - is_append = self.data.get('is_append', False) + user_ids = self.data.get("ids") + original_group_ids = self.data.get("user_group_ids") + is_append = self.data.get("is_append", False) if not user_ids: - raise AppApiException(1004, _('User IDs cannot be empty')) + raise AppApiException(1004, _("User IDs cannot be empty")) if not original_group_ids: - raise AppApiException(1004, _('User Group IDs cannot be empty')) + raise AppApiException(1004, _("User Group IDs cannot be empty")) users = ChatUser.objects.filter(id__in=user_ids) if users.count() != len(user_ids): - raise AppApiException(1004, _('Some users do not exist')) + raise AppApiException(1004, _("Some users do not exist")) groups_count = UserGroup.objects.filter(id__in=original_group_ids).count() if groups_count != len(original_group_ids): - raise AppApiException(1004, _('Some user groups do not exist')) + raise AppApiException(1004, _("Some user groups do not exist")) if is_append: # 获取现有关系 - existing_relations = UserGroupRelation.objects.filter( - user_id__in=user_ids - ).values_list('user_id', 'group_id') + existing_relations = UserGroupRelation.objects.filter(user_id__in=user_ids).values_list( + "user_id", "group_id" + ) existing_groups_map = defaultdict(set) for user_id, group_id in existing_relations: @@ -253,11 +243,7 @@ def batch_add_group(self, with_valid=True): for group_id in new_group_ids: relations_to_create.append( - UserGroupRelation( - id=uuid.uuid7(), - user_id=user_id, - group_id=group_id - ) + UserGroupRelation(id=uuid.uuid7(), user_id=user_id, group_id=group_id) ) # 只创建不存在的关系,不删除现有关系 @@ -269,11 +255,7 @@ def batch_add_group(self, with_valid=True): UserGroupRelation.objects.filter(user_id__in=user_ids).delete() relations_to_create = [ - UserGroupRelation( - id=uuid.uuid7(), - user_id=user_id, - group_id=group_id - ) + UserGroupRelation(id=uuid.uuid7(), user_id=user_id, group_id=group_id) for user_id in user_ids for group_id in original_group_ids ] @@ -284,34 +266,36 @@ def batch_add_group(self, with_valid=True): @transaction.atomic def save(self, instance, with_valid=True): if with_valid: - if instance.get('encrypted'): - instance['password'] = decrypt(instance.get('password')) + if instance.get("encrypted"): + instance["password"] = decrypt(instance.get("password")) self.UserInstance(data=instance).is_valid(raise_exception=True) user = ChatUser( id=uuid.uuid7(), - email=instance.get('email'), - phone=instance.get('phone', ''), - nick_name=instance.get('nick_name', ''), - username=instance.get('username'), - password=password_encrypt(instance.get('password')), - source=instance.get('source', 'LOCAL'), - is_active=True + email=instance.get("email"), + phone=instance.get("phone", ""), + nick_name=instance.get("nick_name", ""), + username=instance.get("username"), + password=password_encrypt(instance.get("password")), + source=instance.get("source", "LOCAL"), + is_active=True, ) user.save() - add_or_edit_user_group_relation(user, instance.get('user_group_ids', [])) + add_or_edit_user_group_relation(user, instance.get("user_group_ids", [])) return ChatUserInstanceSerializer(user).data class UserEditInstance(serializers.Serializer): email = serializers.EmailField( required=False, label=_("Email"), - validators=[validators.EmailValidator( - message=ExceptionCodeConstants.EMAIL_FORMAT_ERROR.value.message, - code=ExceptionCodeConstants.EMAIL_FORMAT_ERROR.value.code - )], + validators=[ + validators.EmailValidator( + message=ExceptionCodeConstants.EMAIL_FORMAT_ERROR.value.message, + code=ExceptionCodeConstants.EMAIL_FORMAT_ERROR.value.code, + ) + ], allow_null=True, - allow_blank=True + allow_blank=True, ) nick_name = serializers.CharField( required=True, @@ -319,20 +303,11 @@ class UserEditInstance(serializers.Serializer): max_length=64, ) phone = serializers.CharField( - required=False, - label=_("Phone"), - max_length=20, - allow_null=True, - allow_blank=True - ) - is_active = serializers.BooleanField( - required=False, - label=_("Is Active") + required=False, label=_("Phone"), max_length=20, allow_null=True, allow_blank=True ) + is_active = serializers.BooleanField(required=False, label=_("Is Active")) user_group_ids = serializers.ListField( - child=serializers.CharField(required=True), - required=False, - label=_('User Group IDs') + child=serializers.CharField(required=True), required=False, label=_("User Group IDs") ) def is_valid(self, *, user_id=None, raise_exception=False): @@ -340,9 +315,9 @@ def is_valid(self, *, user_id=None, raise_exception=False): self._check_unique_nick_name(user_id) def _check_unique_nick_name(self, user_id): - nick_name = self.data.get('nick_name') + nick_name = self.data.get("nick_name") if nick_name and ChatUser.objects.filter(nick_name=nick_name).exclude(id=user_id).exists(): - raise AppApiException(1008, _('Nickname is already in use')) + raise AppApiException(1008, _("Nickname is already in use")) class RePasswordInstance(serializers.Serializer): password = serializers.CharField( @@ -355,9 +330,9 @@ class RePasswordInstance(serializers.Serializer): regex=PASSWORD_REGEX, message=_( "The password must be 6-20 characters long and must be a combination of letters, numbers, and special characters." - ) + ), ) - ] + ], ) re_password = serializers.CharField( required=True, @@ -367,9 +342,9 @@ class RePasswordInstance(serializers.Serializer): regex=PASSWORD_REGEX, message=_( "The confirmation password must be 6-20 characters long and must be a combination of letters, numbers, and special characters." - ) + ), ) - ] + ], ) def is_valid(self, *, raise_exception=False): @@ -377,42 +352,43 @@ def is_valid(self, *, raise_exception=False): self._check_passwords_match() def _check_passwords_match(self): - if self.data.get('password') != self.data.get('re_password'): + if self.data.get("password") != self.data.get("re_password"): raise ExceptionCodeConstants.PASSWORD_NOT_EQ_RE_PASSWORD.value.to_app_api_exception() class Operate(serializers.Serializer): - id = serializers.UUIDField(required=True, label=_('User ID')) + id = serializers.UUIDField(required=True, label=_("User ID")) def is_valid(self, *, raise_exception=False): super().is_valid(raise_exception=True) self._check_user_exists() def _check_user_exists(self): - if not ChatUser.objects.filter(id=self.data.get('id')).exists(): - raise AppApiException(1004, _('User does not exist')) + if not ChatUser.objects.filter(id=self.data.get("id")).exists(): + raise AppApiException(1004, _("User does not exist")) @transaction.atomic def delete(self, with_valid=True): if with_valid: self.is_valid(raise_exception=True) - user_id = self.data.get('id') + user_id = self.data.get("id") ChatUser.objects.filter(id=user_id).delete() return True def edit(self, instance, with_valid=True): if with_valid: self.is_valid(raise_exception=True) - ChatUserSerializer.UserEditInstance(data=instance).is_valid(user_id=self.data.get('id'), - raise_exception=True) - user = ChatUser.objects.filter(id=self.data.get('id')).first() + ChatUserSerializer.UserEditInstance(data=instance).is_valid( + user_id=self.data.get("id"), raise_exception=True + ) + user = ChatUser.objects.filter(id=self.data.get("id")).first() self._update_user_fields(user, instance) user.save() - add_or_edit_user_group_relation(user, instance.get('user_group_ids', [])) + add_or_edit_user_group_relation(user, instance.get("user_group_ids", [])) return ChatUserInstanceSerializer(user).data @staticmethod def _update_user_fields(user, instance): - update_keys = ['email', 'nick_name', 'phone', 'is_active'] + update_keys = ["email", "nick_name", "phone", "is_active"] for key in update_keys: if key in instance and instance.get(key) is not None: setattr(user, key, instance.get(key)) @@ -420,11 +396,12 @@ def _update_user_fields(user, instance): def one(self, with_valid=True): if with_valid: self.is_valid(raise_exception=True) - user = ChatUser.objects.filter(id=self.data.get('id')).first() + user = ChatUser.objects.filter(id=self.data.get("id")).first() user_data = ChatUserInstanceSerializer(user).data # 补充用户组信息 - user_data['user_group_ids'] = list( - UserGroupRelation.objects.filter(user=user).values_list('group_id', flat=True)) + user_data["user_group_ids"] = list( + UserGroupRelation.objects.filter(user=user).values_list("group_id", flat=True) + ) return user_data def re_password(self, instance, with_valid=True): @@ -438,50 +415,50 @@ def re_password(self, instance, with_valid=True): decrypted_data = json.loads(decrypted_raw) if decrypted_raw else {} if isinstance(decrypted_data, dict): instance.update(decrypted_data) - except Exception as e: + except Exception: raise AppApiException(500, _("Invalid encrypted data")) ChatUserSerializer.RePasswordInstance(data=instance).is_valid(raise_exception=True) - user = ChatUser.objects.filter(id=self.data.get('id')).first() - user.password = password_encrypt(instance.get('password')) + user = ChatUser.objects.filter(id=self.data.get("id")).first() + user.password = password_encrypt(instance.get("password")) user.save() return True class GetUserListByGroup(serializers.Serializer): - group_id = serializers.UUIDField(required=True, label=_('Group ID')) + group_id = serializers.UUIDField(required=True, label=_("Group ID")) def is_valid(self, *, raise_exception=False): super().is_valid(raise_exception=True) - group_id = self.data.get('group_id') + group_id = self.data.get("group_id") if not UserGroup.objects.filter(id=group_id).exists(): - raise AppApiException(1004, _('User group does not exist')) + raise AppApiException(1004, _("User group does not exist")) def get_user_list(self, with_valid=True): if with_valid: self.is_valid(raise_exception=True) - group_id = self.data.get('group_id') - user_ids = UserGroupRelation.objects.filter(group_id=group_id).values_list('user_id', flat=True) + group_id = self.data.get("group_id") + user_ids = UserGroupRelation.objects.filter(group_id=group_id).values_list("user_id", flat=True) users = ChatUser.objects.exclude(id__in=user_ids) return ChatUserInstanceSerializer(users, many=True).data @classmethod def list(cls): - users = ChatUser.objects.all().order_by('-create_time') + users = ChatUser.objects.all().order_by("-create_time") return ChatUserInstanceSerializer(users, many=True).data class UserGroupModelSerializer(serializers.ModelSerializer): class Meta: model = UserGroup - fields = ['id', 'name'] + fields = ["id", "name"] class UserGroupCreateSerializer(serializers.Serializer): - id = serializers.CharField(required=False, label='ID') - name = serializers.CharField(required=True, label='User Group Name') + id = serializers.CharField(required=False, label="ID") + name = serializers.CharField(required=True, label="User Group Name") def validate(self, data): - id = data.get('id') - name = data.get('name') + id = data.get("id") + name = data.get("name") if id: group = UserGroup.objects.filter(id=id).first() if not group: @@ -498,57 +475,50 @@ def validate(self, data): def create_or_update_group(self, with_valid=True): if with_valid: self.is_valid(raise_exception=True) - id = self.data.get('id') - name = self.data.get('name') + id = self.data.get("id") + name = self.data.get("name") if id: group = UserGroup.objects.get(id=id) group.name = name group.save() else: - group = UserGroup.objects.create( - id=uuid.uuid7(), - name=name - ) + group = UserGroup.objects.create(id=uuid.uuid7(), name=name) group.save() return UserGroupModelSerializer(group).data def get_user_group_list(self): - groups = UserGroup.objects.all().order_by('name') + groups = UserGroup.objects.all().order_by("name") return UserGroupModelSerializer(groups, many=True).data class UserGroupDeleteSerializer(serializers.Serializer): - id = serializers.CharField(required=True, label='ID') + id = serializers.CharField(required=True, label="ID") def validate(self, data): - id = data.get('id') + id = data.get("id") group = UserGroup.objects.filter(id=id).first() if not group: raise AppApiException(500, _("User group does not exist")) - if group.id == 'default': + if group.id == "default": raise AppApiException(500, _("Default user group cannot be deleted")) return data def delete(self, with_valid=True): if with_valid: self.is_valid(raise_exception=True) - id = self.data.get('id') + id = self.data.get("id") UserGroupRelation.objects.filter(group_id=id).delete() UserGroup.objects.filter(id=id).delete() return True class UserGroupAddMemberSerializer(serializers.Serializer): - id = serializers.CharField(required=True, label='ID') - user_ids = serializers.ListField( - child=serializers.CharField(required=True), - required=True, - label=_('User IDs') - ) + id = serializers.CharField(required=True, label="ID") + user_ids = serializers.ListField(child=serializers.CharField(required=True), required=True, label=_("User IDs")) def validate(self, data): - id = data.get('id') - user_ids = data.get('user_ids') + id = data.get("id") + user_ids = data.get("user_ids") group = UserGroup.objects.filter(id=id).first() if not group: raise AppApiException(500, _("User group does not exist")) @@ -559,35 +529,33 @@ def validate(self, data): def add_member(self, with_valid=True): if with_valid: self.is_valid(raise_exception=True) - user_ids = self.data.get('user_ids') + user_ids = self.data.get("user_ids") current_user_group_ids = set( - str(user_id) for user_id in - UserGroupRelation.objects.filter(group__id=self.data.get('id')).values_list('user_id', flat=True) + str(user_id) + for user_id in UserGroupRelation.objects.filter(group__id=self.data.get("id")).values_list( + "user_id", flat=True + ) ) to_add = set(user_ids).difference(current_user_group_ids) if to_add: - UserGroupRelation.objects.bulk_create([ - UserGroupRelation( - id=uuid.uuid7(), - user_id=user_id, - group_id=self.data.get('id') - ) - for user_id in to_add - ]) + UserGroupRelation.objects.bulk_create( + [ + UserGroupRelation(id=uuid.uuid7(), user_id=user_id, group_id=self.data.get("id")) + for user_id in to_add + ] + ) return True class UserGroupRemoveMemberSerializer(serializers.Serializer): - id = serializers.CharField(required=True, label='ID') + id = serializers.CharField(required=True, label="ID") group_relation_ids = serializers.ListField( - child=serializers.CharField(required=True), - required=True, - label=_('User group relation IDs') + child=serializers.CharField(required=True), required=True, label=_("User group relation IDs") ) def validate(self, data): - id = data.get('id') - user_ids = data.get('group_relation_ids') + id = data.get("id") + user_ids = data.get("group_relation_ids") if UserGroup.objects.filter(id=id).count() == 0: raise AppApiException(500, _("User group does not exist")) if not user_ids: @@ -597,21 +565,21 @@ def validate(self, data): def remove_member(self, with_valid=True): if with_valid: self.is_valid(raise_exception=True) - group_relation_ids = self.data.get('group_relation_ids') + group_relation_ids = self.data.get("group_relation_ids") UserGroupRelation.objects.filter(id__in=group_relation_ids).delete() return True class UserGroupListPageSerializer(serializers.Serializer): class Query(serializers.Serializer): - group_id = serializers.CharField(required=True, label=_('Group ID')) - username = serializers.CharField(required=False, label=_('Username'), allow_null=True) - nick_name = serializers.CharField(required=False, label=_('Nick Name'), allow_null=True) - source = serializers.CharField(required=False, label=_('Source'), allow_null=True) + group_id = serializers.CharField(required=True, label=_("Group ID")) + username = serializers.CharField(required=False, label=_("Username"), allow_null=True) + nick_name = serializers.CharField(required=False, label=_("Nick Name"), allow_null=True) + source = serializers.CharField(required=False, label=_("Source"), allow_null=True) def is_valid(self, *, raise_exception=False): super().is_valid(raise_exception=raise_exception) - group_id = self.data.get('group_id') + group_id = self.data.get("group_id") if not UserGroup.objects.filter(id=group_id).exists(): raise AppApiException(500, _("User group does not exist")) @@ -624,18 +592,18 @@ def page(self, current_page, page_size): query_set, post_records_handler=lambda relation: { **ChatUserInstanceSerializer(relation.user).data, - 'user_group_relation_id': relation.id - } + "user_group_relation_id": relation.id, + }, ) return result def get_query_set(self): - group_id = self.data.get('group_id') + group_id = self.data.get("group_id") - username = self.data.get('username') - nick_name = self.data.get('nick_name') - source = self.data.get('source') - query_set = UserGroupRelation.objects.filter(group_id=group_id).select_related('user') + username = self.data.get("username") + nick_name = self.data.get("nick_name") + source = self.data.get("source") + query_set = UserGroupRelation.objects.filter(group_id=group_id).select_related("user") if username is not None: query_set = query_set.filter(user__username__contains=username) @@ -643,34 +611,53 @@ def get_query_set(self): query_set = query_set.filter(user__nick_name__contains=nick_name) if source is not None: query_set = query_set.filter(user__source=source) - return query_set.order_by('-user__create_time') + return query_set.order_by("-user__create_time") class RePasswordSerializer(serializers.Serializer): - password = serializers.CharField(required=True, label=_("Password"), - validators=[validators.RegexValidator(regex=re.compile( - "^(?![a-zA-Z]+$)(?![A-Z0-9]+$)(?![A-Z_!@#$%^&*`~.()-+=]+$)(?![a-z0-9]+$)(?![a-z_!@#$%^&*`~()-+=]+$)" - "(?![0-9_!@#$%^&*`~()-+=]+$)[a-zA-Z0-9_!@#$%^&*`~.()-+=]{6,20}$") - , message=_( - "The confirmation password must be 6-20 characters long and must be a combination of letters, numbers, and special characters."))]) - - re_password = serializers.CharField(required=True, label=_("Confirm Password"), - validators=[validators.RegexValidator(regex=re.compile( - "^(?![a-zA-Z]+$)(?![A-Z0-9]+$)(?![A-Z_!@#$%^&*`~.()-+=]+$)(?![a-z0-9]+$)(?![a-z_!@#$%^&*`~()-+=]+$)" - "(?![0-9_!@#$%^&*`~()-+=]+$)[a-zA-Z0-9_!@#$%^&*`~.()-+=]{6,20}$") - , message=_( - "The confirmation password must be 6-20 characters long and must be a combination of letters, numbers, and special characters."))] - ) + password = serializers.CharField( + required=True, + label=_("Password"), + validators=[ + validators.RegexValidator( + regex=re.compile( + "^(?![a-zA-Z]+$)(?![A-Z0-9]+$)(?![A-Z_!@#$%^&*`~.()-+=]+$)(?![a-z0-9]+$)(?![a-z_!@#$%^&*`~()-+=]+$)" + "(?![0-9_!@#$%^&*`~()-+=]+$)[a-zA-Z0-9_!@#$%^&*`~.()-+=]{6,20}$" + ), + message=_( + "The confirmation password must be 6-20 characters long and must be a combination of letters, numbers, and special characters." + ), + ) + ], + ) + + re_password = serializers.CharField( + required=True, + label=_("Confirm Password"), + validators=[ + validators.RegexValidator( + regex=re.compile( + "^(?![a-zA-Z]+$)(?![A-Z0-9]+$)(?![A-Z_!@#$%^&*`~.()-+=]+$)(?![a-z0-9]+$)(?![a-z_!@#$%^&*`~()-+=]+$)" + "(?![0-9_!@#$%^&*`~()-+=]+$)[a-zA-Z0-9_!@#$%^&*`~.()-+=]{6,20}$" + ), + message=_( + "The confirmation password must be 6-20 characters long and must be a combination of letters, numbers, and special characters." + ), + ) + ], + ) class Meta: model = ChatUser - fields = '__all__' + fields = "__all__" def is_valid(self, *, raise_exception=False): super().is_valid(raise_exception=True) - if self.data.get('password') != self.data.get('re_password'): - raise AppApiException(ExceptionCodeConstants.PASSWORD_NOT_EQ_RE_PASSWORD.value.code, - ExceptionCodeConstants.PASSWORD_NOT_EQ_RE_PASSWORD.value.message) + if self.data.get("password") != self.data.get("re_password"): + raise AppApiException( + ExceptionCodeConstants.PASSWORD_NOT_EQ_RE_PASSWORD.value.code, + ExceptionCodeConstants.PASSWORD_NOT_EQ_RE_PASSWORD.value.message, + ) return True def reset_password(self, user_id): @@ -679,8 +666,7 @@ def reset_password(self, user_id): :return: 是否成功 """ if self.is_valid(): - QuerySet(ChatUser).filter(id=user_id).update( - password=password_encrypt(self.data.get('password'))) + QuerySet(ChatUser).filter(id=user_id).update(password=password_encrypt(self.data.get("password"))) return True @@ -695,9 +681,9 @@ def profile(user: ChatUser): if not user: return {} return { - 'id': user.id, - 'username': user.username, - 'nick_name': user.nick_name, - 'email': user.email, - 'source': user.source, + "id": user.id, + "username": user.username, + "nick_name": user.nick_name, + "email": user.email, + "source": user.source, } diff --git a/apps/system_manage/views/system_chat_user.py b/apps/system_manage/views/system_chat_user.py index 7cac22f5660..bc8bfd11000 100644 --- a/apps/system_manage/views/system_chat_user.py +++ b/apps/system_manage/views/system_chat_user.py @@ -12,8 +12,13 @@ from common.result import result from models_provider.api.model import DefaultModelResponse from system_manage.api.chat_user import BatchAddGroupApi, ChatUserAPI, ChatUserPageApi, EditUserApi -from system_manage.api.user_group import AddMemberApi, CreateUserGroupApi, DeleteUserGroupApi, RemoveMemberApi, \ - UserGroupListApi +from system_manage.api.user_group import ( + AddMemberApi, + CreateUserGroupApi, + DeleteUserGroupApi, + RemoveMemberApi, + UserGroupListApi, +) from system_manage.models import ChatUser, UserGroup from system_manage.serializers.chat_user import ( ChatUserSerializer, @@ -78,7 +83,12 @@ class List(APIView): tags=[_("System/Chat user")], # type: ignore responses=ChatUserPageApi.get_response(), ) - @has_permissions(PermissionConstants.CHAT_USER_READ, PermissionConstants.USER_GROUP_READ, RoleConstants.ADMIN, RoleConstants.WORKSPACE_MANAGE) + @has_permissions( + PermissionConstants.CHAT_USER_READ, + PermissionConstants.USER_GROUP_READ, + RoleConstants.ADMIN, + RoleConstants.WORKSPACE_MANAGE, + ) def get(self, request: Request): return result.success(ChatUserSerializer.list()) @@ -257,7 +267,12 @@ class SystemChatUserGroupView(APIView): responses=CreateUserGroupApi.get_response(), tags=[_("System/User Group")], # type: ignore ) # type: ignore - @has_permissions(PermissionConstants.USER_GROUP_CREATE, PermissionConstants.USER_GROUP_EDIT, RoleConstants.ADMIN, RoleConstants.WORKSPACE_MANAGE) + @has_permissions( + PermissionConstants.USER_GROUP_CREATE, + PermissionConstants.USER_GROUP_EDIT, + RoleConstants.ADMIN, + RoleConstants.WORKSPACE_MANAGE, + ) @log( menu="User group", operate="Create or update user group", @@ -342,7 +357,9 @@ class RemoveMember(APIView): responses=DefaultModelResponse, tags=[_("System/User Group")], # type: ignore ) - @has_permissions(PermissionConstants.USER_GROUP_REMOVE_MEMBER, RoleConstants.ADMIN, RoleConstants.WORKSPACE_MANAGE) + @has_permissions( + PermissionConstants.USER_GROUP_REMOVE_MEMBER, RoleConstants.ADMIN, RoleConstants.WORKSPACE_MANAGE + ) @log( menu="User group", operate="Remove member from user group",