Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion apps/application/flow/compare/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
39 changes: 38 additions & 1 deletion apps/application/flow/loop_workflow_manage.py
Original file line number Diff line number Diff line change
Expand Up @@ -179,19 +179,56 @@ 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)
prompt_template = PromptTemplate.from_template(prompt, template_format='jinja2')
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"

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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

# 构建查询条件
Expand Down Expand Up @@ -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':
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
51 changes: 50 additions & 1 deletion apps/application/flow/workflow_manage.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
"""
获取启动节点
Expand Down
Loading