From 37c4046c9087dca1d3b2ffc6392683a526744aa7 Mon Sep 17 00:00:00 2001 From: GuoQing Zhang Date: Thu, 19 Dec 2024 14:46:46 +0800 Subject: [PATCH] feat: agent node add tool invoke log --- .../bisheng/chat/clients/llm_callback.py | 22 +++++++++++++++ .../bisheng/workflow/nodes/agent/agent.py | 28 +++++++++++++++++-- 2 files changed, 48 insertions(+), 2 deletions(-) diff --git a/src/backend/bisheng/chat/clients/llm_callback.py b/src/backend/bisheng/chat/clients/llm_callback.py index c3ec765fa..e931dff2c 100644 --- a/src/backend/bisheng/chat/clients/llm_callback.py +++ b/src/backend/bisheng/chat/clients/llm_callback.py @@ -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 diff --git a/src/backend/bisheng/workflow/nodes/agent/agent.py b/src/backend/bisheng/workflow/nodes/agent/agent.py index 53dd3a9d2..54274dd27 100644 --- a/src/backend/bisheng/workflow/nodes/agent/agent.py +++ b/src/backend/bisheng/workflow/nodes/agent/agent.py @@ -34,9 +34,11 @@ class AgentNode(BaseNode): self._user_prompt = PromptTemplateParser(template=self.node_params['user_prompt']) self._user_variables = self._user_prompt.extract() - self.batch_variable_list = [] + # log data + 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'] @@ -210,6 +212,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) @@ -239,6 +242,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 @@ -275,7 +298,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':