From 270ffa8488e034837ee43e918c5fef972b7e03a7 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E7=8E=8B=E8=89=AF?= <841369634@qq.com> Date: Tue, 7 Jul 2026 17:00:09 +0800 Subject: [PATCH] fix: Fix the `generate_prompt` method generated a 'None' string --- apps/application/flow/compare/__init__.py | 2 +- apps/application/flow/loop_workflow_manage.py | 39 +++++++++++++- .../ai_chat_step_node/impl/base_chat_node.py | 2 +- .../step_node/mcp_node/impl/base_mcp_node.py | 2 +- .../impl/base_search_document_node.py | 8 +-- .../tool_lib_node/impl/base_tool_lib_node.py | 2 +- .../tool_node/impl/base_tool_node.py | 2 +- .../impl/base_variable_assign_node.py | 2 +- apps/application/flow/workflow_manage.py | 51 ++++++++++++++++++- 9 files changed, 98 insertions(+), 12 deletions(-) diff --git a/apps/application/flow/compare/__init__.py b/apps/application/flow/compare/__init__.py index ce0c430e1ad..0a0fd159d33 100644 --- a/apps/application/flow/compare/__init__.py +++ b/apps/application/flow/compare/__init__.py @@ -64,7 +64,7 @@ def _compare(source_value, compare, target_value): def _assertion(workflow_manage, field_list: List[str], compare: str, value): try: - value = workflow_manage.generate_prompt(value) + value = workflow_manage.generate_field_value(value) except Exception: pass field_value = None diff --git a/apps/application/flow/loop_workflow_manage.py b/apps/application/flow/loop_workflow_manage.py index c236b15dcc5..a1d7daccfdf 100644 --- a/apps/application/flow/loop_workflow_manage.py +++ b/apps/application/flow/loop_workflow_manage.py @@ -179,12 +179,18 @@ def reset_prompt(self, prompt: str): prompt = self.parentWorkflowManage.reset_prompt(prompt) return prompt - def generate_prompt(self, prompt: str): + def generate_prompt(self, prompt: str) -> str: """ 格式化生成提示词 @param prompt: 提示词信息 @return: 格式化后的提示词 """ + if prompt is None: + return '' + if prompt == '': + return '' + if "{{" not in prompt or "}}" not in prompt: + return prompt context = {**self.get_workflow_content(), **self.parentWorkflowManage.get_workflow_content()} prompt = self.reset_prompt(prompt) @@ -192,6 +198,37 @@ def generate_prompt(self, prompt: str): value = prompt_template.format(context=context) return value + def reset_field_value(self, prompt: str): + prompt = super().reset_field_value(prompt) + for field in self.loop_field_list: + chatLabel = f"loop.{field.get('value')}" + chatValue = f"context.get('loop').get('{field.get('value', '')}')" + prompt = prompt.replace(chatLabel, chatValue) + + prompt = self.parentWorkflowManage.reset_field_value(prompt) + return prompt + + def generate_field_value(self, field_value: str): + """ + 格式化生成参数值 + @param field_value: 参数信息 + @return: 格式化后的参数值 + """ + if field_value is None: + return None + if field_value == '': + return '' + if "{{" not in field_value or "}}" not in field_value: + return field_value + + context = {**self.get_workflow_content(), **self.parentWorkflowManage.get_workflow_content()} + field_value = self.reset_field_value(field_value) + prompt_template = PromptTemplate.from_template(field_value, template_format='jinja2') + value = prompt_template.format(context=context) + if value == 'None': + return None + return value + def get_source_type(self): return "APPLICATION" diff --git a/apps/application/flow/step_node/ai_chat_step_node/impl/base_chat_node.py b/apps/application/flow/step_node/ai_chat_step_node/impl/base_chat_node.py index ee2d9cbc311..c6938c01771 100644 --- a/apps/application/flow/step_node/ai_chat_step_node/impl/base_chat_node.py +++ b/apps/application/flow/step_node/ai_chat_step_node/impl/base_chat_node.py @@ -468,7 +468,7 @@ def handle_variables(self, tool_params): # 处理参数中的变量 for k, v in tool_params.items(): if type(v) == str: - tool_params[k] = self.workflow_manage.generate_prompt(tool_params[k]) + tool_params[k] = self.workflow_manage.generate_field_value(v) elif type(v) == dict: self.handle_variables(v) elif (type(v) == list) and len(v) > 0 and (type(v[0]) == str): diff --git a/apps/application/flow/step_node/mcp_node/impl/base_mcp_node.py b/apps/application/flow/step_node/mcp_node/impl/base_mcp_node.py index 5245546b0de..80dbfd6f4ae 100644 --- a/apps/application/flow/step_node/mcp_node/impl/base_mcp_node.py +++ b/apps/application/flow/step_node/mcp_node/impl/base_mcp_node.py @@ -58,7 +58,7 @@ def handle_variables(self, tool_params): # 处理参数中的变量 for k, v in tool_params.items(): if type(v) == str: - tool_params[k] = self.workflow_manage.generate_prompt(tool_params[k]) + tool_params[k] = self.workflow_manage.generate_field_value(v) elif type(v) == dict: self.handle_variables(v) elif (type(v) == list) and len(v) > 0 and (type(v[0]) == str): diff --git a/apps/application/flow/step_node/search_document_node/impl/base_search_document_node.py b/apps/application/flow/step_node/search_document_node/impl/base_search_document_node.py index 1d85cff5331..b90bbb993dc 100644 --- a/apps/application/flow/step_node/search_document_node/impl/base_search_document_node.py +++ b/apps/application/flow/step_node/search_document_node/impl/base_search_document_node.py @@ -125,10 +125,10 @@ def handle_custom_tags(self, document_id_list: List, search_condition_list: list for condition in search_condition_list: tag_key = condition['key'] - field_value = self.workflow_manage.generate_prompt(condition['value']) + field_value = self.workflow_manage.generate_field_value(condition['value']) compare_type = condition['compare'] - if not field_value or field_value == 'None' or len(field_value) == 0: + if not field_value: continue # 构建查询条件 @@ -164,10 +164,10 @@ def handle_custom_tags(self, document_id_list: List, search_condition_list: list for condition in search_condition_list: tag_key = condition['key'] - field_value = self.workflow_manage.generate_prompt(condition['value']) + field_value = self.workflow_manage.generate_field_value(condition['value']) compare_type = condition['compare'] - if not field_value or field_value == 'None' or len(field_value) == 0: + if not field_value: continue if compare_type == 'not_contain': diff --git a/apps/application/flow/step_node/tool_lib_node/impl/base_tool_lib_node.py b/apps/application/flow/step_node/tool_lib_node/impl/base_tool_lib_node.py index 3360544a51e..504173133a5 100644 --- a/apps/application/flow/step_node/tool_lib_node/impl/base_tool_lib_node.py +++ b/apps/application/flow/step_node/tool_lib_node/impl/base_tool_lib_node.py @@ -106,7 +106,7 @@ def convert_value(name: str, value, _type, is_required, source, node): return float(value) return value try: - value = node.workflow_manage.generate_prompt(value) + value = node.workflow_manage.generate_field_value(value) return common_convert_value(_type, value) except Exception as e: raise Exception( diff --git a/apps/application/flow/step_node/tool_node/impl/base_tool_node.py b/apps/application/flow/step_node/tool_node/impl/base_tool_node.py index e269a2b2ba7..a77b14d25cc 100644 --- a/apps/application/flow/step_node/tool_node/impl/base_tool_node.py +++ b/apps/application/flow/step_node/tool_node/impl/base_tool_node.py @@ -81,7 +81,7 @@ def convert_value(name: str, value, _type, is_required, source, node): return float(value) return value try: - value = node.workflow_manage.generate_prompt(value) + value = node.workflow_manage.generate_field_value(value) return common_convert_value(_type, value) except Exception as e: raise Exception( diff --git a/apps/application/flow/step_node/variable_assign_node/impl/base_variable_assign_node.py b/apps/application/flow/step_node/variable_assign_node/impl/base_variable_assign_node.py index b9572805acf..207240125e3 100644 --- a/apps/application/flow/step_node/variable_assign_node/impl/base_variable_assign_node.py +++ b/apps/application/flow/step_node/variable_assign_node/impl/base_variable_assign_node.py @@ -53,7 +53,7 @@ def handle(self, variable, evaluation): result['output_value'] = variable['value'] = val elif variable['type'] == 'string': # 变量解析 例如:{{global.xxx}} - val = self.workflow_manage.generate_prompt(variable['value']) + val = self.workflow_manage.generate_field_value(variable['value']) evaluation(variable, val) result['output_value'] = val else: diff --git a/apps/application/flow/workflow_manage.py b/apps/application/flow/workflow_manage.py index f1323c6d4b7..47de31db2a3 100644 --- a/apps/application/flow/workflow_manage.py +++ b/apps/application/flow/workflow_manage.py @@ -775,18 +775,67 @@ def reset_prompt(self, prompt: str): return prompt - def generate_prompt(self, prompt: str): + def generate_prompt(self, prompt: str) -> str: """ 格式化生成提示词 @param prompt: 提示词信息 @return: 格式化后的提示词 """ + if prompt is None: + return '' + if prompt == '': + return '' + if "{{" not in prompt or "}}" not in prompt: + return prompt + context = self.get_workflow_content() prompt = self.reset_prompt(prompt) prompt_template = PromptTemplate.from_template(prompt, template_format='jinja2') value = prompt_template.format(context=context) return value + def reset_field_value(self, prompt: str): + placeholder = "{}" + for field in self.field_list: + globeLabel = f"{field.get('node_name')}.{field.get('value')}" + globeValue = f"context.get('{field.get('node_id')}',{placeholder}).get('{field.get('value', '')}')" + prompt = prompt.replace(globeLabel, globeValue) + for field in self.global_field_list: + globeLabel = f"全局变量.{field.get('value')}" + globeLabelNew = f"global.{field.get('value')}" + globeValue = f"context.get('global').get('{field.get('value', '')}')" + prompt = prompt.replace(globeLabel, globeValue).replace(globeLabelNew, globeValue) + for field in self.chat_field_list: + chatLabel = f"chat.{field.get('value')}" + chatValue = f"context.get('chat').get('{field.get('value', '')}')" + prompt = prompt.replace(chatLabel, chatValue) + + return prompt + + def generate_field_value(self, field_value: str): + """ + 格式化生成参数值 + @param field_value: 参数信息 + @return: 格式化后的参数值 + """ + if field_value is None: + return None + if not field_value: + return '' + if "{{" not in field_value or "}}" not in field_value: + return field_value + + context = self.get_workflow_content() + field_value = self.reset_field_value(field_value) + field_value_template = PromptTemplate.from_template(field_value, template_format='jinja2') + value = field_value_template.format(context=context) + maxkb_logger.info(f"context: {json.dumps(context, ensure_ascii=False)}") + maxkb_logger.info(f"value: {value}, type: {type(value)}") + maxkb_logger.info(field_value) + if value == 'None': + return None + return value + def get_start_node(self): """ 获取启动节点