mirror of
https://github.com/dataelement/bisheng.git
synced 2026-09-24 23:19:52 +08:00
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:
@@ -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':
|
||||
|
||||
Reference in New Issue
Block a user