Merge branch 'feat/0.4.0.dev3' into feat/0.4.1

# Conflicts:
#	src/backend/bisheng/workflow/nodes/agent/agent.py
This commit is contained in:
GuoQing Zhang
2024-12-19 14:51:28 +08:00
2 changed files with 46 additions and 1 deletions
@@ -42,6 +42,7 @@ class LLMNodeCallbackHandler(BaseCallbackHandler):
output: bool,
output_key: str,
stream: bool = True,
tool_list: Optional[List[Any]] = None,
):
self.callback_manager = callback
self.unique_id = unique_id
@@ -50,6 +51,7 @@ class LLMNodeCallbackHandler(BaseCallbackHandler):
self.output_len = 0
self.output_key = output_key
self.stream = stream
self.tool_list = tool_list
logger.info('on_llm_new_token {} outkey={}', self.output, self.output_key)
async def on_tool_start(self, serialized: Dict[str, Any], input_str: str,
@@ -57,12 +59,26 @@ class LLMNodeCallbackHandler(BaseCallbackHandler):
"""Run when tool starts running."""
logger.debug(
f'on_tool_start serialized={serialized} input_str={input_str} kwargs={kwargs}')
if self.tool_list is not None:
self.tool_list.append({
'type': 'start',
'run_id': kwargs.get('run_id').hex,
'name': serialized['name'],
'input': input_str,
})
if serialized['name'] == 'sql_agent':
self.output = False
async def on_tool_end(self, output: str, **kwargs: Any) -> Any:
"""Run when tool ends running."""
logger.debug(f'on_tool_end output={output} kwargs={kwargs}')
if self.tool_list is not None:
self.tool_list.append({
'type': 'end',
'run_id': kwargs.get('run_id').hex,
'name': kwargs['name'],
'output': output,
})
if kwargs['name'] == 'sql_agent':
self.output = True
@@ -70,6 +86,12 @@ class LLMNodeCallbackHandler(BaseCallbackHandler):
**kwargs: Any) -> Any:
"""Run when tool errors."""
logger.debug(f'on_tool_error error={error} kwargs={kwargs}')
if self.tool_list is not None:
self.tool_list.append({
'type': 'error',
'run_id': kwargs.get('run_id').hex,
'error': str(error),
})
if kwargs['name'] == 'sql_agent':
self.output = True
@@ -37,6 +37,7 @@ class AgentNode(BaseNode):
self._batch_variable_list = {}
self._system_prompt_list = []
self._user_prompt_list = []
self._tool_invoke_list = []
# 聊天消息
self._chat_history_flag = self.node_params['chat_history_flag']['flag']
@@ -213,6 +214,7 @@ class AgentNode(BaseNode):
self._batch_variable_list = {}
self._system_prompt_list = []
self._user_prompt_list = []
self._tool_invoke_list = []
for one in self._system_variables:
variable_map[one] = self.graph_state.get_variable_by_str(one)
@@ -242,6 +244,26 @@ class AgentNode(BaseNode):
'user_prompt': self._user_prompt_list,
'output': result
}
tool_invoke_info = {}
if self._tool_invoke_list:
for one in self._tool_invoke_list:
if one['run_id'] not in tool_invoke_info:
tool_invoke_info[one['run_id']] = {}
if one['type'] == 'start':
tool_invoke_info[one['run_id']].update({
'name': one['name'],
'input': one['input']
})
elif one['type'] == 'end':
tool_invoke_info[one['run_id']].update({
'output': one['output']
})
elif one['type'] == 'error':
tool_invoke_info[one['run_id']].update({
'output': f'Error: {one["error"]}'
})
if tool_invoke_info:
ret['tool_invoke'] = list(tool_invoke_info.values())
if self._batch_variable_list:
ret['batch_variable'] = self._batch_variable_list
return ret
@@ -279,7 +301,8 @@ class AgentNode(BaseNode):
unique_id=unique_id,
node_id=self.id,
output=self._output_user,
output_key=output_key)
output_key=output_key,
tool_list=self._tool_invoke_list)
config = RunnableConfig(callbacks=[llm_callback])
if self._agent_executor_type == 'ReAct':