diff --git a/apps/application/migrations/0014_applicationversion_knowledge_ids.py b/apps/application/migrations/0014_applicationversion_knowledge_ids.py new file mode 100644 index 00000000000..676b0cb0d58 --- /dev/null +++ b/apps/application/migrations/0014_applicationversion_knowledge_ids.py @@ -0,0 +1,61 @@ +# Generated by Django 5.2.15 on 2026-07-29 02:42 + +from django.db import migrations, models + + +def forwards(apps, schema_editor): + Application = apps.get_model("application", "Application") + ResourceMapping = apps.get_model("system_manage", "ResourceMapping") + ApplicationVersion = apps.get_model("application", "ApplicationVersion") + + APPLICATION = "APPLICATION" + KNOWLEDGE = "KNOWLEDGE" + SIMPLE = "SIMPLE" + db_alias = schema_editor.connection.alias + simple_application_ids = { + str(app_id) + for app_id in Application.objects.using(db_alias) + .filter(type=SIMPLE) + .values_list("id", flat=True) + } + mapping = {} + qs = ( + ResourceMapping.objects.using(db_alias) + .filter(source_type=APPLICATION, target_type=KNOWLEDGE) + .values_list("source_id", "target_id") + ) + for source_id, target_id in qs.iterator(): + if source_id in simple_application_ids: + mapping.setdefault(source_id, []).append(target_id) + mapping = {k: list(dict.fromkeys(v)) for k, v in mapping.items()} + + updates = [] + for obj in ApplicationVersion.objects.using(db_alias).iterator(): + app_id = str(obj.application_id) + if app_id not in simple_application_ids: + continue + knowledge_ids = mapping.get(app_id) + if knowledge_ids: + obj.knowledge_ids = knowledge_ids + updates.append(obj) + if updates: + ApplicationVersion.objects.using(db_alias).bulk_update( + updates, ["knowledge_ids"], batch_size=500 + ) + + +class Migration(migrations.Migration): + + dependencies = [ + ('application', '0013_application_long_term_enable_and_more'), + ('system_manage', '0005_resourcemapping'), + ] + + operations = [ + migrations.AddField( + model_name='applicationversion', + name='knowledge_ids', + field=models.JSONField(default=list, verbose_name='数据集id列表'), + ), + migrations.RunPython(forwards, migrations.RunPython.noop), + ] diff --git a/apps/application/models/application.py b/apps/application/models/application.py index 34824b29eee..71d345164cc 100644 --- a/apps/application/models/application.py +++ b/apps/application/models/application.py @@ -192,6 +192,7 @@ class ApplicationVersion(AppModelMixin): long_term_model_params_setting = models.JSONField(verbose_name="长期记忆模型参数相关设置", default=dict) long_term_trigger_type = models.CharField(verbose_name='长期记忆触发类型', default='ROUND') long_term_trigger_setting = models.JSONField(verbose_name='长期记忆触发配置', default=dict) + knowledge_ids = models.JSONField(verbose_name="数据集id列表", default=list) class Meta: db_table = "application_version" diff --git a/apps/application/serializers/application.py b/apps/application/serializers/application.py index d5c6f9a8261..c3a74a9c8a1 100644 --- a/apps/application/serializers/application.py +++ b/apps/application/serializers/application.py @@ -86,7 +86,7 @@ def get_bound_tool_ids(instance: Dict) -> List[str]: """ tool_ids = set() for key in ("tool_ids", "skill_tool_ids", "mcp_tool_ids"): - for tool_id in (instance.get(key) or []): + for tool_id in instance.get(key) or []: tool_ids.add(str(tool_id)) if instance.get("mcp_tool_id"): tool_ids.add(str(instance.get("mcp_tool_id"))) @@ -96,7 +96,7 @@ def collect(node_data): if node_data.get(key): tool_ids.add(str(node_data.get(key))) for key in ("mcp_tool_ids", "tool_ids", "skill_tool_ids"): - for tool_id in (node_data.get(key) or []): + for tool_id in node_data.get(key) or []: tool_ids.add(str(tool_id)) _walk_workflow_nodes(instance.get("work_flow"), collect) @@ -109,11 +109,11 @@ def get_bound_application_ids(instance: Dict) -> List[str]: ai-chat-node 的 node_data 包含 application_ids 列表。 """ application_ids = set() - for app_id in (instance.get("application_ids") or []): + for app_id in instance.get("application_ids") or []: application_ids.add(str(app_id)) def collect(node_data): - for app_id in (node_data.get("application_ids") or []): + for app_id in node_data.get("application_ids") or []: application_ids.add(str(app_id)) _walk_workflow_nodes(instance.get("work_flow"), collect) @@ -186,7 +186,9 @@ def validate_bound_tool_permissions(user_id: str, workspace_id: str, instance: D application_ids = get_bound_application_ids(instance) if application_ids: authorized_application_ids = set(get_authorized_application_ids(user_id, workspace_id, application_ids)) - unauthorized_application_ids = [app_id for app_id in application_ids if app_id not in authorized_application_ids] + unauthorized_application_ids = [ + app_id for app_id in application_ids if app_id not in authorized_application_ids + ] if unauthorized_application_ids: message = lazy_format( _("No permission to use application(s): {application_ids}"), @@ -1268,6 +1270,14 @@ def publish(self, instance, with_valid=True): workspace_id=workspace_id, ) self.reset_application_version(work_flow_version, application) + # 如果是简易应用 需要存入 knowledge_ids + if application.type == ApplicationTypeChoices.SIMPLE: + work_flow_version.knowledge_ids = [ + str(row.target_id) + for row in QuerySet(ResourceMapping).filter( + source_id=str(application.id), source_type="APPLICATION", target_type="KNOWLEDGE" + ) + ] work_flow_version.save() access_token = hashlib.md5(str(uuid.uuid7()).encode()).hexdigest()[8:24] application_access_token = QuerySet(ApplicationAccessToken).filter(application_id=application.id).first() @@ -1781,9 +1791,7 @@ def batch_delete(self, instance: Dict, with_valid=True): id_list = instance.get("id_list") workspace_id = self.data.get("workspace_id") id_list = list( - QuerySet(Application) - .filter(id__in=id_list, workspace_id=workspace_id) - .values_list("id", flat=True) + QuerySet(Application).filter(id__in=id_list, workspace_id=workspace_id).values_list("id", flat=True) ) QuerySet(ApplicationVersion).filter(application_id__in=id_list).delete() @@ -1842,9 +1850,7 @@ def batch_clean_time(self, instance: Dict, with_valid=True): class BatchCleanTimeSerializer(BatchSerializer): clean_time = serializers.IntegerField(required=True, min_value=1, max_value=100000, label=_("Clean time")) - file_clean_time = serializers.IntegerField( - required=True, min_value=1, max_value=100000, label=_("File clean time") - ) + file_clean_time = serializers.IntegerField(required=True, min_value=1, max_value=100000, label=_("File clean time")) def is_valid(self, *, model=None, raise_exception=False): super().is_valid(model=model, raise_exception=True) diff --git a/apps/application/serializers/common.py b/apps/application/serializers/common.py index 2773c8fd7f6..dfd340b9d2a 100644 --- a/apps/application/serializers/common.py +++ b/apps/application/serializers/common.py @@ -1,23 +1,23 @@ # coding=utf-8 """ - @project: MaxKB - @Author:虎虎 - @file: common.py - @date:2025/6/9 13:42 - @desc: +@project: MaxKB +@Author:虎虎 +@file: common.py +@date:2025/6/9 13:42 +@desc: """ -from typing import List -from django.core.cache import cache -from django.db.models import QuerySet -from django.utils import timezone -from django.utils.translation import gettext_lazy as _ +from typing import List -from application.models import Application, ChatRecord, Chat, ApplicationVersion, ChatUserType, ApplicationTypeChoices +from application.models import Application, ApplicationTypeChoices, ApplicationVersion, Chat, ChatRecord, ChatUserType from application.serializers.application_chat import ChatCountSerializer from common.constants.cache_version import Cache_Version from common.database_model_manage.database_model_manage import DatabaseModelManage from common.exception.app_exception import ChatException +from django.core.cache import cache +from django.db.models import QuerySet +from django.utils import timezone +from django.utils.translation import gettext_lazy as _ from knowledge.models import Document from models_provider.models import Model from models_provider.tools import get_model_credential @@ -26,12 +26,7 @@ class ToolExecute: - def __init__(self, tool_id: str, - tool_record_id: str, - workspace_id: str, - source_type, - source_id, - debug=False): + def __init__(self, tool_id: str, tool_record_id: str, workspace_id: str, source_type, source_id, debug=False): self.tool_id = tool_id self.workspace_id = workspace_id self.source_type = source_type @@ -42,8 +37,12 @@ def __init__(self, tool_id: str, def get_record(self): if self.tool_record_id: if self.debug: - return self.to_record(cache.get(Cache_Version.TOOL_WORKFLOW_EXECUTE.get_key(key=self.tool_record_id), - version=Cache_Version.TOOL_WORKFLOW_EXECUTE.get_version())) + return self.to_record( + cache.get( + Cache_Version.TOOL_WORKFLOW_EXECUTE.get_key(key=self.tool_record_id), + version=Cache_Version.TOOL_WORKFLOW_EXECUTE.get_version(), + ) + ) else: return QuerySet(ToolRecord).filter(tool_id=self.tool_id, id=self.tool_record_id).first() return None @@ -51,61 +50,74 @@ def get_record(self): def to_record(self, tool_record_dict): if tool_record_dict is None: return None - return ToolRecord(id=tool_record_dict.get('id'), - tool_id=tool_record_dict.get('tool_id'), - workspace_id=tool_record_dict.get('workspace_id'), - source_type=tool_record_dict.get('source_type'), - source_id=tool_record_dict.get('source_id'), - meta=tool_record_dict.get('meta'), - state=tool_record_dict.get('state'), - run_time=tool_record_dict.get('run_time')) + return ToolRecord( + id=tool_record_dict.get("id"), + tool_id=tool_record_dict.get("tool_id"), + workspace_id=tool_record_dict.get("workspace_id"), + source_type=tool_record_dict.get("source_type"), + source_id=tool_record_dict.get("source_id"), + meta=tool_record_dict.get("meta"), + state=tool_record_dict.get("state"), + run_time=tool_record_dict.get("run_time"), + ) def to_dict(self, tool_record): - return {'id': tool_record.id, - 'tool_id': tool_record.tool_id, - 'workspace_id': tool_record.workspace_id, - 'source_type': tool_record.source_type, - 'source_id': tool_record.source_id, - 'meta': tool_record.meta, - 'state': tool_record.state, - 'run_time': tool_record.run_time} + return { + "id": tool_record.id, + "tool_id": tool_record.tool_id, + "workspace_id": tool_record.workspace_id, + "source_type": tool_record.source_type, + "source_id": tool_record.source_id, + "meta": tool_record.meta, + "state": tool_record.state, + "run_time": tool_record.run_time, + } def set_record(self, tool_record): - cache.set(Cache_Version.TOOL_WORKFLOW_EXECUTE.get_key(key=self.tool_record_id), self.to_dict(tool_record), - version=Cache_Version.TOOL_WORKFLOW_EXECUTE.get_version(), - timeout=60 * 30) + cache.set( + Cache_Version.TOOL_WORKFLOW_EXECUTE.get_key(key=self.tool_record_id), + self.to_dict(tool_record), + version=Cache_Version.TOOL_WORKFLOW_EXECUTE.get_version(), + timeout=60 * 30, + ) if not self.debug: - QuerySet(ToolRecord).update_or_create(id=tool_record.id, - create_defaults={'id': tool_record.id, - 'tool_id': tool_record.tool_id, - 'state': tool_record.state, - 'workspace_id': tool_record.workspace_id, - "source_type": tool_record.source_type, - 'source_id': tool_record.source_id, - 'meta': tool_record.meta, - 'run_time': tool_record.run_time}, - defaults={ - 'workspace_id': tool_record.workspace_id, - 'tool_id': tool_record.tool_id, - "source_type": tool_record.source_type, - 'source_id': tool_record.source_id, - 'state': tool_record.state, - 'meta': tool_record.meta, - 'run_time': tool_record.run_time - }) + QuerySet(ToolRecord).update_or_create( + id=tool_record.id, + create_defaults={ + "id": tool_record.id, + "tool_id": tool_record.tool_id, + "state": tool_record.state, + "workspace_id": tool_record.workspace_id, + "source_type": tool_record.source_type, + "source_id": tool_record.source_id, + "meta": tool_record.meta, + "run_time": tool_record.run_time, + }, + defaults={ + "workspace_id": tool_record.workspace_id, + "tool_id": tool_record.tool_id, + "source_type": tool_record.source_type, + "source_id": tool_record.source_id, + "state": tool_record.state, + "meta": tool_record.meta, + "run_time": tool_record.run_time, + }, + ) class ChatInfo: - def __init__(self, - chat_id: str, - chat_user_id: str, - chat_user_type: str, - ip_address: str, - source: {}, - knowledge_id_list: List[str], - exclude_document_id_list: list[str], - application_id: str, - debug=False): + def __init__( + self, + chat_id: str, + chat_user_id: str, + chat_user_type: str, + ip_address: str, + source: {}, + knowledge_id_list: List[str], + exclude_document_id_list: list[str], + application_id: str, + debug=False, + ): """ :param chat_id: 对话id :param chat_user_id 对话用户id @@ -133,36 +145,44 @@ def __init__(self, @staticmethod def get_no_references_setting(knowledge_setting, model_setting): no_references_setting = knowledge_setting.get( - 'no_references_setting', { - 'status': 'ai_questioning', - 'value': '{question}'}) - if no_references_setting.get('status') == 'ai_questioning': - no_references_prompt = model_setting.get('no_references_prompt', '{question}') - no_references_setting['value'] = no_references_prompt if len(no_references_prompt) > 0 else "{question}" + "no_references_setting", {"status": "ai_questioning", "value": "{question}"} + ) + if no_references_setting.get("status") == "ai_questioning": + no_references_prompt = model_setting.get("no_references_prompt", "{question}") + no_references_setting["value"] = no_references_prompt if len(no_references_prompt) > 0 else "{question}" return no_references_setting def get_application(self): if self.debug: application = QuerySet(Application).filter(id=self.application_id).first() if not application: - raise ChatException(500, _('The application does not exist')) + raise ChatException(500, _("The application does not exist")) else: - application = QuerySet(ApplicationVersion).filter(application_id=self.application_id).order_by( - '-create_time')[0:1].first() + application = ( + QuerySet(ApplicationVersion) + .filter(application_id=self.application_id) + .order_by("-create_time")[0:1] + .first() + ) if not application: raise ChatException(500, _("The application has not been published. Please use it after publishing.")) if application.type == ApplicationTypeChoices.SIMPLE.value: - # 数据集id列表 - knowledge_id_list = [str(row.target_id) for row in - QuerySet(ResourceMapping).filter(source_id=self.application_id, - source_type='APPLICATION', - target_type='KNOWLEDGE')] + # 数据集id列表 这里需要从application中获取知识库 不能从关联表获取 + if self.debug: + knowledge_id_list = [ + str(row.target_id) + for row in QuerySet(ResourceMapping).filter( + source_id=self.application_id, source_type="APPLICATION", target_type="KNOWLEDGE" + ) + ] + else: + knowledge_id_list = application.knowledge_ids # 需要排除的文档 - exclude_document_id_list = [str(document.id) for document in - QuerySet(Document).filter( - knowledge_id__in=knowledge_id_list, - is_active=False)] + exclude_document_id_list = [ + str(document.id) + for document in QuerySet(Document).filter(knowledge_id__in=knowledge_id_list, is_active=False) + ] self.knowledge_id_list = knowledge_id_list self.exclude_document_id_list = exclude_document_id_list self.application = application @@ -175,37 +195,38 @@ def get_chat_user(self, asker=None): if self.chat_user_type == ChatUserType.CHAT_USER.value and chat_user_model: chat_user = QuerySet(chat_user_model).filter(id=self.chat_user_id).first() return { - 'id': str(chat_user.id), - 'email': chat_user.email, - 'phone': chat_user.phone, - 'nick_name': chat_user.nick_name, - 'username': chat_user.username, - 'source': chat_user.source + "id": str(chat_user.id), + "email": chat_user.email, + "phone": chat_user.phone, + "nick_name": chat_user.nick_name, + "username": chat_user.username, + "source": chat_user.source, } else: if asker: if isinstance(asker, dict): self.chat_user = asker else: - self.chat_user = {'username': asker} + self.chat_user = {"username": asker} else: - self.chat_user = {'username': '游客'} + self.chat_user = {"username": "游客"} return self.chat_user def get_chat_user_group(self, asker=None): chat_user = self.get_chat_user(asker=asker) - chat_user_id = chat_user.get('id') + chat_user_id = chat_user.get("id") if not chat_user_id: return [] user_group_relation_model = DatabaseModelManage.get_model("user_group_relation") if user_group_relation_model: - return [{ - 'id': user_group_relation.group_id, - 'name': user_group_relation.group.name - } for user_group_relation in - QuerySet(user_group_relation_model).select_related('group').filter(user_id=chat_user_id)] + return [ + {"id": user_group_relation.group_id, "name": user_group_relation.group.name} + for user_group_relation in QuerySet(user_group_relation_model) + .select_related("group") + .filter(user_id=chat_user_id) + ] return [] def to_base_pipeline_manage_params(self): @@ -222,63 +243,89 @@ def to_base_pipeline_manage_params(self): credential = get_model_credential(model.provider, model.model_type, model.model_name) model_params_setting = credential.get_model_params_setting_form(model.model_name).get_default_form_data() return { - 'knowledge_id_list': self.knowledge_id_list, - 'exclude_document_id_list': self.exclude_document_id_list, - 'exclude_paragraph_id_list': [], - 'top_n': 3 if knowledge_setting.get('top_n') is None else knowledge_setting.get('top_n'), - 'similarity': 0.6 if knowledge_setting.get('similarity') is None else knowledge_setting.get('similarity'), - 'max_paragraph_char_number': knowledge_setting.get('max_paragraph_char_number') or 5000, - 'history_chat_record': self.chat_record_list, - 'chat_id': self.chat_id, - 'dialogue_number': self.application.dialogue_number, - 'problem_optimization_prompt': self.application.problem_optimization_prompt if self.application.problem_optimization_prompt is not None and len( - self.application.problem_optimization_prompt) > 0 else _( - "() contains the user's question. Answer the guessed user's question based on the context ({question}) Requirement: Output a complete question and put it in the tag"), - 'prompt': model_setting.get( - 'prompt') if 'prompt' in model_setting and len(model_setting.get( - 'prompt')) > 0 else Application.get_default_model_prompt(), - 'system': model_setting.get( - 'system', None), - 'model_id': model_id, - 'problem_optimization': self.application.problem_optimization, - 'stream': True, - 'model_setting': model_setting, - 'model_params_setting': model_params_setting if self.application.model_params_setting is None or len( - self.application.model_params_setting.keys()) == 0 else self.application.model_params_setting, - 'search_mode': self.application.knowledge_setting.get('search_mode') or 'embedding', - 'no_references_setting': self.get_no_references_setting(self.application.knowledge_setting, model_setting), - 'workspace_id': self.application.workspace_id, - 'application_id': self.application_id, - 'mcp_enable': self.application.mcp_enable, - 'mcp_tool_ids': self.application.mcp_tool_ids, - 'mcp_servers': self.application.mcp_servers, - 'mcp_source': self.application.mcp_source, - 'tool_enable': self.application.tool_enable, - 'tool_ids': self.application.tool_ids, - 'application_enable': self.application.application_enable, - 'application_ids': self.application.application_ids, - 'skill_tool_ids': self.application.skill_tool_ids, - 'mcp_output_enable': self.application.mcp_output_enable, + "knowledge_id_list": self.knowledge_id_list, + "exclude_document_id_list": self.exclude_document_id_list, + "exclude_paragraph_id_list": [], + "top_n": 3 if knowledge_setting.get("top_n") is None else knowledge_setting.get("top_n"), + "similarity": 0.6 if knowledge_setting.get("similarity") is None else knowledge_setting.get("similarity"), + "max_paragraph_char_number": knowledge_setting.get("max_paragraph_char_number") or 5000, + "history_chat_record": self.chat_record_list, + "chat_id": self.chat_id, + "dialogue_number": self.application.dialogue_number, + "problem_optimization_prompt": self.application.problem_optimization_prompt + if self.application.problem_optimization_prompt is not None + and len(self.application.problem_optimization_prompt) > 0 + else _( + "() contains the user's question. Answer the guessed user's question based on the context ({question}) Requirement: Output a complete question and put it in the tag" + ), + "prompt": model_setting.get("prompt") + if "prompt" in model_setting and len(model_setting.get("prompt")) > 0 + else Application.get_default_model_prompt(), + "system": model_setting.get("system", None), + "model_id": model_id, + "problem_optimization": self.application.problem_optimization, + "stream": True, + "model_setting": model_setting, + "model_params_setting": model_params_setting + if self.application.model_params_setting is None or len(self.application.model_params_setting.keys()) == 0 + else self.application.model_params_setting, + "search_mode": self.application.knowledge_setting.get("search_mode") or "embedding", + "no_references_setting": self.get_no_references_setting(self.application.knowledge_setting, model_setting), + "workspace_id": self.application.workspace_id, + "application_id": self.application_id, + "mcp_enable": self.application.mcp_enable, + "mcp_tool_ids": self.application.mcp_tool_ids, + "mcp_servers": self.application.mcp_servers, + "mcp_source": self.application.mcp_source, + "tool_enable": self.application.tool_enable, + "tool_ids": self.application.tool_ids, + "application_enable": self.application.application_enable, + "application_ids": self.application.application_ids, + "skill_tool_ids": self.application.skill_tool_ids, + "mcp_output_enable": self.application.mcp_output_enable, } - def to_pipeline_manage_params(self, problem_text: str, post_response_handler, - exclude_paragraph_id_list, chat_user_id: str, chat_user_type, ip_address, source, - stream=True, - form_data=None): + def to_pipeline_manage_params( + self, + problem_text: str, + post_response_handler, + exclude_paragraph_id_list, + chat_user_id: str, + chat_user_type, + ip_address, + source, + stream=True, + form_data=None, + ): if form_data is None: form_data = {} params = self.to_base_pipeline_manage_params() - return {**params, 'problem_text': problem_text, 'post_response_handler': post_response_handler, - 'exclude_paragraph_id_list': exclude_paragraph_id_list, 'stream': stream, 'chat_user_id': chat_user_id, - 'chat_user_type': chat_user_type, 'ip_address': ip_address, 'source': source, 'form_data': form_data} + return { + **params, + "problem_text": problem_text, + "post_response_handler": post_response_handler, + "exclude_paragraph_id_list": exclude_paragraph_id_list, + "stream": stream, + "chat_user_id": chat_user_id, + "chat_user_type": chat_user_type, + "ip_address": ip_address, + "source": source, + "form_data": form_data, + } def set_chat(self, question): if not self.debug: if not QuerySet(Chat).filter(id=self.chat_id).exists(): - Chat(id=self.chat_id, application_id=self.application_id, abstract=question[0:1024], - chat_user_id=self.chat_user_id, chat_user_type=self.chat_user_type, - ip_address=self.ip_address, source=self.source, - asker=self.get_chat_user()).save() + Chat( + id=self.chat_id, + application_id=self.application_id, + abstract=question[0:1024], + chat_user_id=self.chat_user_id, + chat_user_type=self.chat_user_type, + ip_address=self.ip_address, + source=self.source, + asker=self.get_chat_user(), + ).save() def set_chat_variable(self, chat_context): if not self.debug: @@ -287,9 +334,12 @@ def set_chat_variable(self, chat_context): chat.meta = {**(chat.meta if isinstance(chat.meta, dict) else {}), **chat_context} chat.save() else: - cache.set(Cache_Version.CHAT_VARIABLE.get_key(key=self.chat_id), chat_context, - version=Cache_Version.CHAT_VARIABLE.get_version(), - timeout=60 * 30) + cache.set( + Cache_Version.CHAT_VARIABLE.get_key(key=self.chat_id), + chat_context, + version=Cache_Version.CHAT_VARIABLE.get_version(), + timeout=60 * 30, + ) def get_chat_variable(self): if not self.debug: @@ -298,8 +348,13 @@ def get_chat_variable(self): return chat.meta return {} else: - return cache.get(Cache_Version.CHAT_VARIABLE.get_key(key=self.chat_id), - version=Cache_Version.CHAT_VARIABLE.get_version()) or {} + return ( + cache.get( + Cache_Version.CHAT_VARIABLE.get_key(key=self.chat_id), + version=Cache_Version.CHAT_VARIABLE.get_version(), + ) + or {} + ) def append_chat_record(self, chat_record: ChatRecord): chat_record.problem_text = chat_record.problem_text[0:10240] if chat_record.problem_text is not None else "" @@ -316,135 +371,167 @@ def append_chat_record(self, chat_record: ChatRecord): self.chat_record_list.append(chat_record) if not self.debug: if not QuerySet(Chat).filter(id=self.chat_id).exists(): - Chat(id=self.chat_id, application_id=self.application_id, abstract=chat_record.problem_text[0:1024], - chat_user_id=self.chat_user_id, chat_user_type=self.chat_user_type, - ip_address=self.ip_address, source=self.source, - asker=self.get_chat_user()).save() + Chat( + id=self.chat_id, + application_id=self.application_id, + abstract=chat_record.problem_text[0:1024], + chat_user_id=self.chat_user_id, + chat_user_type=self.chat_user_type, + ip_address=self.ip_address, + source=self.source, + asker=self.get_chat_user(), + ).save() else: QuerySet(Chat).filter(id=self.chat_id).update(update_time=timezone.now()) # 插入会话记录 - QuerySet(ChatRecord).update_or_create(id=chat_record.id, - create_defaults={'id': chat_record.id, - 'chat_id': chat_record.chat_id, - "vote_status": chat_record.vote_status, - 'problem_text': chat_record.problem_text, - 'answer_text': chat_record.answer_text, - 'answer_text_list': chat_record.answer_text_list, - 'message_tokens': chat_record.message_tokens, - 'answer_tokens': chat_record.answer_tokens, - 'const': chat_record.const, - 'details': chat_record.details, - 'improve_paragraph_id_list': chat_record.improve_paragraph_id_list, - 'run_time': chat_record.run_time, - 'source': chat_record.source, - 'ip_address': chat_record.ip_address or '', - 'index': chat_record.index}, - defaults={ - "vote_status": chat_record.vote_status, - 'problem_text': chat_record.problem_text, - 'answer_text': chat_record.answer_text, - 'answer_text_list': chat_record.answer_text_list, - 'message_tokens': chat_record.message_tokens, - 'answer_tokens': chat_record.answer_tokens, - 'const': chat_record.const, - 'details': chat_record.details, - 'improve_paragraph_id_list': chat_record.improve_paragraph_id_list, - 'run_time': chat_record.run_time, - 'index': chat_record.index, - 'source': chat_record.source, - 'ip_address': chat_record.ip_address or '', - }) - ChatCountSerializer(data={'chat_id': self.chat_id}).update_chat() + QuerySet(ChatRecord).update_or_create( + id=chat_record.id, + create_defaults={ + "id": chat_record.id, + "chat_id": chat_record.chat_id, + "vote_status": chat_record.vote_status, + "problem_text": chat_record.problem_text, + "answer_text": chat_record.answer_text, + "answer_text_list": chat_record.answer_text_list, + "message_tokens": chat_record.message_tokens, + "answer_tokens": chat_record.answer_tokens, + "const": chat_record.const, + "details": chat_record.details, + "improve_paragraph_id_list": chat_record.improve_paragraph_id_list, + "run_time": chat_record.run_time, + "source": chat_record.source, + "ip_address": chat_record.ip_address or "", + "index": chat_record.index, + }, + defaults={ + "vote_status": chat_record.vote_status, + "problem_text": chat_record.problem_text, + "answer_text": chat_record.answer_text, + "answer_text_list": chat_record.answer_text_list, + "message_tokens": chat_record.message_tokens, + "answer_tokens": chat_record.answer_tokens, + "const": chat_record.const, + "details": chat_record.details, + "improve_paragraph_id_list": chat_record.improve_paragraph_id_list, + "run_time": chat_record.run_time, + "index": chat_record.index, + "source": chat_record.source, + "ip_address": chat_record.ip_address or "", + }, + ) + ChatCountSerializer(data={"chat_id": self.chat_id}).update_chat() def to_dict(self): return { - 'chat_id': self.chat_id, - 'chat_user_id': self.chat_user_id, - 'chat_user_type': self.chat_user_type, - 'ip_address': self.ip_address, - 'source': self.source, - 'knowledge_id_list': self.knowledge_id_list, - 'exclude_document_id_list': self.exclude_document_id_list, - 'application_id': self.application_id, - 'chat_record_list': [self.chat_record_to_map(c) for c in self.chat_record_list][-20:], - 'debug': self.debug + "chat_id": self.chat_id, + "chat_user_id": self.chat_user_id, + "chat_user_type": self.chat_user_type, + "ip_address": self.ip_address, + "source": self.source, + "knowledge_id_list": self.knowledge_id_list, + "exclude_document_id_list": self.exclude_document_id_list, + "application_id": self.application_id, + "chat_record_list": [self.chat_record_to_map(c) for c in self.chat_record_list][-20:], + "debug": self.debug, } def chat_record_to_map(self, chat_record): - return {'id': chat_record.id, - 'chat_id': chat_record.chat_id, - 'vote_status': chat_record.vote_status, - 'problem_text': chat_record.problem_text, - 'answer_text': chat_record.answer_text, - 'answer_text_list': chat_record.answer_text_list, - 'message_tokens': chat_record.message_tokens, - 'answer_tokens': chat_record.answer_tokens, - 'const': chat_record.const, - 'details': chat_record.details, - 'improve_paragraph_id_list': chat_record.improve_paragraph_id_list, - 'run_time': chat_record.run_time, - 'source': chat_record.source, - 'ip_address': chat_record.ip_address, - 'index': chat_record.index} + return { + "id": chat_record.id, + "chat_id": chat_record.chat_id, + "vote_status": chat_record.vote_status, + "problem_text": chat_record.problem_text, + "answer_text": chat_record.answer_text, + "answer_text_list": chat_record.answer_text_list, + "message_tokens": chat_record.message_tokens, + "answer_tokens": chat_record.answer_tokens, + "const": chat_record.const, + "details": chat_record.details, + "improve_paragraph_id_list": chat_record.improve_paragraph_id_list, + "run_time": chat_record.run_time, + "source": chat_record.source, + "ip_address": chat_record.ip_address, + "index": chat_record.index, + } @staticmethod def map_to_chat_record(chat_record_dict): - return ChatRecord(id=chat_record_dict.get('id'), - chat_id=chat_record_dict.get('chat_id'), - vote_status=chat_record_dict.get('vote_status'), - problem_text=chat_record_dict.get('problem_text'), - answer_text=chat_record_dict.get('answer_text'), - answer_text_list=chat_record_dict.get('answer_text_list'), - message_tokens=chat_record_dict.get('message_tokens'), - answer_tokens=chat_record_dict.get('answer_tokens'), - const=chat_record_dict.get('const'), - details=chat_record_dict.get('details'), - improve_paragraph_id_list=chat_record_dict.get('improve_paragraph_id_list'), - run_time=chat_record_dict.get('run_time'), - index=chat_record_dict.get('index'), - source=chat_record_dict.get('source'), - ip_address=chat_record_dict.get('ip_address')) + return ChatRecord( + id=chat_record_dict.get("id"), + chat_id=chat_record_dict.get("chat_id"), + vote_status=chat_record_dict.get("vote_status"), + problem_text=chat_record_dict.get("problem_text"), + answer_text=chat_record_dict.get("answer_text"), + answer_text_list=chat_record_dict.get("answer_text_list"), + message_tokens=chat_record_dict.get("message_tokens"), + answer_tokens=chat_record_dict.get("answer_tokens"), + const=chat_record_dict.get("const"), + details=chat_record_dict.get("details"), + improve_paragraph_id_list=chat_record_dict.get("improve_paragraph_id_list"), + run_time=chat_record_dict.get("run_time"), + index=chat_record_dict.get("index"), + source=chat_record_dict.get("source"), + ip_address=chat_record_dict.get("ip_address"), + ) def set_cache(self): - cache.set(Cache_Version.CHAT.get_key(key=self.chat_id), self.to_dict(), - version=Cache_Version.CHAT_INFO.get_version(), - timeout=60 * 30) + cache.set( + Cache_Version.CHAT.get_key(key=self.chat_id), + self.to_dict(), + version=Cache_Version.CHAT_INFO.get_version(), + timeout=60 * 30, + ) @staticmethod def map_to_chat_info(chat_info_dict): - c = ChatInfo(chat_info_dict.get('chat_id'), chat_info_dict.get('chat_user_id'), - chat_info_dict.get('chat_user_type'), chat_info_dict.get('ip_address'), - chat_info_dict.get('source'), - chat_info_dict.get('knowledge_id_list'), - chat_info_dict.get('exclude_document_id_list'), - chat_info_dict.get('application_id'), - debug=chat_info_dict.get('debug')) - c.chat_record_list = [ChatInfo.map_to_chat_record(c_r) for c_r in chat_info_dict.get('chat_record_list')] + c = ChatInfo( + chat_info_dict.get("chat_id"), + chat_info_dict.get("chat_user_id"), + chat_info_dict.get("chat_user_type"), + chat_info_dict.get("ip_address"), + chat_info_dict.get("source"), + chat_info_dict.get("knowledge_id_list"), + chat_info_dict.get("exclude_document_id_list"), + chat_info_dict.get("application_id"), + debug=chat_info_dict.get("debug"), + ) + c.chat_record_list = [ChatInfo.map_to_chat_record(c_r) for c_r in chat_info_dict.get("chat_record_list")] return c @staticmethod def get_cache(chat_id): - chat_info_dict = cache.get(Cache_Version.CHAT.get_key(key=chat_id), - version=Cache_Version.CHAT_INFO.get_version()) + chat_info_dict = cache.get( + Cache_Version.CHAT.get_key(key=chat_id), version=Cache_Version.CHAT_INFO.get_version() + ) if chat_info_dict: return ChatInfo.map_to_chat_info(chat_info_dict) return None def update_resource_mapping_by_application(application_id: str, other_resource_mapping=None): - from application.flow.tools import get_instance_resource, save_workflow_mapping, \ - application_instance_field_call_dict + from application.flow.tools import ( + application_instance_field_call_dict, + get_instance_resource, + save_workflow_mapping, + ) from system_manage.models.resource_mapping import ResourceType + if other_resource_mapping is None: other_resource_mapping = [] application = QuerySet(Application).filter(id=application_id).first() - instance_mapping = get_instance_resource(application, ResourceType.APPLICATION, str(application.id), - application_instance_field_call_dict) - if application.type == 'WORK_FLOW': - save_workflow_mapping(application.work_flow, ResourceType.APPLICATION, str(application_id), - instance_mapping + other_resource_mapping) + instance_mapping = get_instance_resource( + application, ResourceType.APPLICATION, str(application.id), application_instance_field_call_dict + ) + if application.type == "WORK_FLOW": + save_workflow_mapping( + application.work_flow, + ResourceType.APPLICATION, + str(application_id), + instance_mapping + other_resource_mapping, + ) return else: - save_workflow_mapping({}, ResourceType.APPLICATION, str(application_id), - instance_mapping + other_resource_mapping) + save_workflow_mapping( + {}, ResourceType.APPLICATION, str(application_id), instance_mapping + other_resource_mapping + ) diff --git a/apps/chat/serializers/chat.py b/apps/chat/serializers/chat.py index 4514dc457bc..5990ee85683 100644 --- a/apps/chat/serializers/chat.py +++ b/apps/chat/serializers/chat.py @@ -495,11 +495,16 @@ 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): - # 数据集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')] + 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')] + else: + 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 @@ -584,10 +589,15 @@ def open_simple(self, application): ip_address = self.data.get("ip_address") source = self.data.get("source") debug = self.data.get("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')] + 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')] + else: + 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()) ChatInfo(chat_id, chat_user_id, chat_user_type, ip_address, source, knowledge_id_list,