diff --git a/apps/application/workflow/message/aggregator/aggregation_manager.py b/apps/application/workflow/message/aggregator/aggregation_manager.py index 149345c9788..5f3c1c7441e 100644 --- a/apps/application/workflow/message/aggregator/aggregation_manager.py +++ b/apps/application/workflow/message/aggregator/aggregation_manager.py @@ -1,10 +1,11 @@ # coding=utf-8 """ - @project: MaxKB - @file: aggregation_manager.py - @date:2026/7/22 16:24 - @desc: 聚合管理器 +@project: MaxKB +@file: aggregation_manager.py +@date:2026/7/22 16:24 +@desc: 聚合管理器 """ + from typing import Dict, List from application.workflow.message.struct.content import Content @@ -29,9 +30,11 @@ def contents(self) -> List[Content]: def aggregate(self, chunk: Content) -> None: """ 聚合内容块 - + @param chunk: 内容块 """ + if not AggregatorFactory.is_aggregatable(chunk): + return key = f"{chunk.id}_{chunk.type.value if hasattr(chunk.type, 'value') else chunk.type}" idx = self._key_to_index.get(key) @@ -53,7 +56,7 @@ def clear(self) -> None: def get_contents(self) -> List[Dict]: """ 获取所有聚合后的内容(字典格式) - + @return: 内容字典列表 """ return [content.to_dict() for content in self._contents] diff --git a/apps/application/workflow/message/aggregator/aggregator_factory.py b/apps/application/workflow/message/aggregator/aggregator_factory.py index fbd4488efdb..7f79d6ee1fc 100644 --- a/apps/application/workflow/message/aggregator/aggregator_factory.py +++ b/apps/application/workflow/message/aggregator/aggregator_factory.py @@ -8,16 +8,16 @@ from typing import Dict, Type, Optional -from application.workflow.message.aggregator.impl import ProgressAggregator -from application.workflow.message.struct.content import Content -from application.workflow.message.struct.progress_content import ProgressContent -from application.workflow.message.struct.text_content import TextContent -from application.workflow.message.struct.reasoning_content import ReasoningContent -from application.workflow.message.struct.tool_content import ToolContent from application.workflow.message.aggregator.content_aggregator import ContentAggregator -from application.workflow.message.aggregator.impl.text_aggregator import TextAggregator +from application.workflow.message.aggregator.impl import FormAggregator from application.workflow.message.aggregator.impl.reasoning_aggregator import ReasoningAggregator +from application.workflow.message.aggregator.impl.text_aggregator import TextAggregator from application.workflow.message.aggregator.impl.tool_aggregator import ToolAggregator +from application.workflow.message.struct.content import Content +from application.workflow.message.struct.form_content import FormContent +from application.workflow.message.struct.reasoning_content import ReasoningContent +from application.workflow.message.struct.text_content import TextContent +from application.workflow.message.struct.tool_content import ToolContent class AggregatorFactory: @@ -30,7 +30,7 @@ class AggregatorFactory: TextContent: TextAggregator(), ReasoningContent: ReasoningAggregator(), ToolContent: ToolAggregator(), - ProgressContent: ProgressAggregator(), + FormContent: FormAggregator(), } @classmethod @@ -56,3 +56,15 @@ def get_aggregator_optional(cls, content_class: Type[Content]) -> Optional[Conte @return: 聚合器实例或None """ return cls._aggregators.get(content_class) + + @classmethod + def is_aggregatable(cls, chunk: Content) -> bool: + """ + 判断给定的内容对象是否可以被聚合 + + @param chunk: 内容对象实例 + @return: True 表示有对应的聚合器,False 表示没有 + """ + if chunk is None: + return False + return type(chunk) in cls._aggregators diff --git a/apps/application/workflow/message/aggregator/impl/__init__.py b/apps/application/workflow/message/aggregator/impl/__init__.py index 257f446f989..34a83e07003 100644 --- a/apps/application/workflow/message/aggregator/impl/__init__.py +++ b/apps/application/workflow/message/aggregator/impl/__init__.py @@ -9,6 +9,6 @@ from application.workflow.message.aggregator.impl.text_aggregator import TextAggregator from application.workflow.message.aggregator.impl.reasoning_aggregator import ReasoningAggregator from application.workflow.message.aggregator.impl.tool_aggregator import ToolAggregator -from application.workflow.message.aggregator.impl.progress_aggregator import ProgressAggregator +from application.workflow.message.aggregator.impl.form_aggregator import FormAggregator -__all__ = ["TextAggregator", "ReasoningAggregator", "ToolAggregator", "ProgressAggregator"] +__all__ = ["TextAggregator", "ReasoningAggregator", "ToolAggregator", "FormAggregator"] diff --git a/apps/application/workflow/message/aggregator/impl/form_aggregator.py b/apps/application/workflow/message/aggregator/impl/form_aggregator.py new file mode 100644 index 00000000000..663fef26bd2 --- /dev/null +++ b/apps/application/workflow/message/aggregator/impl/form_aggregator.py @@ -0,0 +1,49 @@ +""" +@project: MaxKB +@file: form_aggregator.py +@date:2026/9/16 +@desc: +""" + +from application.workflow.message.aggregator.content_aggregator import ContentAggregator +from application.workflow.message.struct.form_content import FormContent + + +class FormAggregator(ContentAggregator[FormContent]): + """ + 推理内容聚合器 + 用于合并流式推理内容块 + """ + + def aggregate(self, prev: FormContent, chunk: FormContent) -> FormContent: + """ + 聚合推理内容 + + @param prev: 之前的内容 + @param chunk: 新的内容块 + @return: 合并后的内容 + """ + if prev is None: + return chunk + + # 合并 status: 优先使用 chunk 的,否则使用 prev 的 + merged_status = chunk.status if chunk.status else prev.status + form_field_list = chunk.form_field_list if chunk.form_field_list else prev.form_field_list + form_content_format = chunk.form_content_format if chunk.form_content_format else prev.form_content_format + is_submit = chunk.is_submit if chunk.is_submit else prev.is_submit + form_data = chunk.form_data if chunk.form_data else prev.form_data + # 合并基础字段 + merged_id = chunk.id if chunk.id else prev.id + merged_node_info = chunk.node_info if chunk.node_info else prev.node_info + merged_position = chunk.position if chunk.position else prev.position + result = FormContent( + merged_id, + form_field_list, + form_content_format, + is_submit, + merged_status, + merged_node_info, + merged_position, + form_data, + ) + return result diff --git a/apps/application/workflow/message/aggregator/impl/progress_aggregator.py b/apps/application/workflow/message/aggregator/impl/progress_aggregator.py deleted file mode 100644 index 3d7e6239b63..00000000000 --- a/apps/application/workflow/message/aggregator/impl/progress_aggregator.py +++ /dev/null @@ -1,38 +0,0 @@ -# coding=utf-8 -""" -@project: MaxKB -@file: progress_aggregator.py -@date:2026/7/22 16:24 -@desc: ReasoningContent 聚合器 -""" - -from application.workflow.message.aggregator.content_aggregator import ContentAggregator -from application.workflow.message.struct.progress_content import ProgressContent - - -class ProgressAggregator(ContentAggregator[ProgressContent]): - """ - 推理内容聚合器 - 用于合并流式推理内容块 - """ - - def aggregate(self, prev: ProgressContent, chunk: ProgressContent) -> ProgressContent: - """ - 聚合推理内容 - - @param prev: 之前的内容 - @param chunk: 新的内容块 - @return: 合并后的内容 - """ - if prev is None: - return chunk - - # 合并 status: 优先使用 chunk 的,否则使用 prev 的 - merged_status = chunk.status if chunk.status else prev.status - # 合并基础字段 - merged_id = chunk.id if chunk.id else prev.id - merged_node_info = chunk.node_info if chunk.node_info else prev.node_info - merged_position = chunk.position if chunk.position else prev.position - result = ProgressContent(merged_id, merged_status, merged_node_info, merged_position) - - return result diff --git a/apps/chat/serializers/chat.py b/apps/chat/serializers/chat.py index da2d71ff17b..f7861d4168b 100644 --- a/apps/chat/serializers/chat.py +++ b/apps/chat/serializers/chat.py @@ -337,7 +337,7 @@ def on_complete(wf_manage, error): if chat_record: old_details = chat_record.details if position and chat_record.messages: - messages = [*chat_record.messages, *messages] + messages = list({m.get("id"): m for m in [*chat_record.messages, *messages]}.values()) details = wf_manage.get_details(position=position, old_details=old_details) self.update_chat_record(chat_user_id, chat_record_id_str, wf_manage.context, messages, details) ChatCountSerializer(data={"chat_id": chat_id}).update_chat()