Underground-Digital/Workflow-Engine
0
1import json2import logging3from collections.abc import Generator4from copy import deepcopy5from typing import Any, Optional, Union6 7from core.agent.base_agent_runner import BaseAgentRunner8from core.app.apps.base_app_queue_manager import PublishFrom9from core.app.entities.queue_entities import QueueAgentThoughtEvent, QueueMessageEndEvent, QueueMessageFileEvent10from core.file import file_manager11from core.model_runtime.entities import (12 AssistantPromptMessage,13 LLMResult,14 LLMResultChunk,15 LLMResultChunkDelta,16 LLMUsage,17 PromptMessage,18 PromptMessageContent,19 PromptMessageContentType,20 SystemPromptMessage,21 TextPromptMessageContent,22 ToolPromptMessage,23 UserPromptMessage,24)25from core.prompt.agent_history_prompt_transform import AgentHistoryPromptTransform26from core.tools.entities.tool_entities import ToolInvokeMeta27from core.tools.tool_engine import ToolEngine28from models.model import Message29 30logger = logging.getLogger(__name__)31 32 33class FunctionCallAgentRunner(BaseAgentRunner):34 def run(self, message: Message, query: str, **kwargs: Any) -> Generator[LLMResultChunk, None, None]:35 """36 Run FunctionCall agent application37 """38 self.query = query39 app_generate_entity = self.application_generate_entity40 41 app_config = self.app_config42 43 # convert tools into ModelRuntime Tool format44 tool_instances, prompt_messages_tools = self._init_prompt_tools()45 46 iteration_step = 147 max_iteration_steps = min(app_config.agent.max_iteration, 5) + 148 49 # continue to run until there is not any tool call50 function_call_state = True51 llm_usage = {"usage": None}52 final_answer = ""53 54 # get tracing instance55 trace_manager = app_generate_entity.trace_manager56 57 def increase_usage(final_llm_usage_dict: dict[str, LLMUsage], usage: LLMUsage):58 if not final_llm_usage_dict["usage"]:59 final_llm_usage_dict["usage"] = usage60 else:61 llm_usage = final_llm_usage_dict["usage"]62 llm_usage.prompt_tokens += usage.prompt_tokens63 llm_usage.completion_tokens += usage.completion_tokens64 llm_usage.prompt_price += usage.prompt_price65 llm_usage.completion_price += usage.completion_price66 llm_usage.total_price += usage.total_price67 68 model_instance = self.model_instance69 70 while function_call_state and iteration_step <= max_iteration_steps:71 function_call_state = False72 73 if iteration_step == max_iteration_steps:74 # the last iteration, remove all tools75 prompt_messages_tools = []76 77 message_file_ids = []78 agent_thought = self.create_agent_thought(79 message_id=message.id, message="", tool_name="", tool_input="", messages_ids=message_file_ids80 )81 82 # recalc llm max tokens83 prompt_messages = self._organize_prompt_messages()84 self.recalc_llm_max_tokens(self.model_config, prompt_messages)85 # invoke model86 chunks: Union[Generator[LLMResultChunk, None, None], LLMResult] = model_instance.invoke_llm(87 prompt_messages=prompt_messages,88 model_parameters=app_generate_entity.model_conf.parameters,89 tools=prompt_messages_tools,90 stop=app_generate_entity.model_conf.stop,91 stream=self.stream_tool_call,92 user=self.user_id,93 callbacks=[],94 )95 96 tool_calls: list[tuple[str, str, dict[str, Any]]] = []97 98 # save full response99 response = ""100 101 # save tool call names and inputs102 tool_call_names = ""103 tool_call_inputs = ""104 105 current_llm_usage = None106 107 if self.stream_tool_call:108 is_first_chunk = True109 for chunk in chunks:110 if is_first_chunk:111 self.queue_manager.publish(112 QueueAgentThoughtEvent(agent_thought_id=agent_thought.id), PublishFrom.APPLICATION_MANAGER113 )114 is_first_chunk = False115 # check if there is any tool call116 if self.check_tool_calls(chunk):117 function_call_state = True118 tool_calls.extend(self.extract_tool_calls(chunk))119 tool_call_names = ";".join([tool_call[1] for tool_call in tool_calls])120 try:121 tool_call_inputs = json.dumps(122 {tool_call[1]: tool_call[2] for tool_call in tool_calls}, ensure_ascii=False123 )124 except json.JSONDecodeError as e:125 # ensure ascii to avoid encoding error126 tool_call_inputs = json.dumps({tool_call[1]: tool_call[2] for tool_call in tool_calls})127 128 if chunk.delta.message and chunk.delta.message.content:129 if isinstance(chunk.delta.message.content, list):130 for content in chunk.delta.message.content:131 response += content.data132 else:133 response += chunk.delta.message.content134 135 if chunk.delta.usage:136 increase_usage(llm_usage, chunk.delta.usage)137 current_llm_usage = chunk.delta.usage138 139 yield chunk140 else:141 result: LLMResult = chunks142 # check if there is any tool call143 if self.check_blocking_tool_calls(result):144 function_call_state = True145 tool_calls.extend(self.extract_blocking_tool_calls(result))146 tool_call_names = ";".join([tool_call[1] for tool_call in tool_calls])147 try:148 tool_call_inputs = json.dumps(149 {tool_call[1]: tool_call[2] for tool_call in tool_calls}, ensure_ascii=False150 )151 except json.JSONDecodeError as e:152 # ensure ascii to avoid encoding error153 tool_call_inputs = json.dumps({tool_call[1]: tool_call[2] for tool_call in tool_calls})154 155 if result.usage:156 increase_usage(llm_usage, result.usage)157 current_llm_usage = result.usage158 159 if result.message and result.message.content:160 if isinstance(result.message.content, list):161 for content in result.message.content:162 response += content.data163 else:164 response += result.message.content165 166 if not result.message.content:167 result.message.content = ""168 169 self.queue_manager.publish(170 QueueAgentThoughtEvent(agent_thought_id=agent_thought.id), PublishFrom.APPLICATION_MANAGER171 )172 173 yield LLMResultChunk(174 model=model_instance.model,175 prompt_messages=result.prompt_messages,176 system_fingerprint=result.system_fingerprint,177 delta=LLMResultChunkDelta(178 index=0,179 message=result.message,180 usage=result.usage,181 ),182 )183 184 assistant_message = AssistantPromptMessage(content="", tool_calls=[])185 if tool_calls:186 assistant_message.tool_calls = [187 AssistantPromptMessage.ToolCall(188 id=tool_call[0],189 type="function",190 function=AssistantPromptMessage.ToolCall.ToolCallFunction(191 name=tool_call[1], arguments=json.dumps(tool_call[2], ensure_ascii=False)192 ),193 )194 for tool_call in tool_calls195 ]196 else:197 assistant_message.content = response198 199 self._current_thoughts.append(assistant_message)200 201 # save thought202 self.save_agent_thought(203 agent_thought=agent_thought,204 tool_name=tool_call_names,205 tool_input=tool_call_inputs,206 thought=response,207 tool_invoke_meta=None,208 observation=None,209 answer=response,210 messages_ids=[],211 llm_usage=current_llm_usage,212 )213 self.queue_manager.publish(214 QueueAgentThoughtEvent(agent_thought_id=agent_thought.id), PublishFrom.APPLICATION_MANAGER215 )216 217 final_answer += response + "\n"218 219 # call tools220 tool_responses = []221 for tool_call_id, tool_call_name, tool_call_args in tool_calls:222 tool_instance = tool_instances.get(tool_call_name)223 if not tool_instance:224 tool_response = {225 "tool_call_id": tool_call_id,226 "tool_call_name": tool_call_name,227 "tool_response": f"there is not a tool named {tool_call_name}",228 "meta": ToolInvokeMeta.error_instance(f"there is not a tool named {tool_call_name}").to_dict(),229 }230 else:231 # invoke tool232 tool_invoke_response, message_files, tool_invoke_meta = ToolEngine.agent_invoke(233 tool=tool_instance,234 tool_parameters=tool_call_args,235 user_id=self.user_id,236 tenant_id=self.tenant_id,237 message=self.message,238 invoke_from=self.application_generate_entity.invoke_from,239 agent_tool_callback=self.agent_callback,240 trace_manager=trace_manager,241 )242 # publish files243 for message_file_id, save_as in message_files:244 if save_as:245 self.variables_pool.set_file(tool_name=tool_call_name, value=message_file_id, name=save_as)246 247 # publish message file248 self.queue_manager.publish(249 QueueMessageFileEvent(message_file_id=message_file_id), PublishFrom.APPLICATION_MANAGER250 )251 # add message file ids252 message_file_ids.append(message_file_id)253 254 tool_response = {255 "tool_call_id": tool_call_id,256 "tool_call_name": tool_call_name,257 "tool_response": tool_invoke_response,258 "meta": tool_invoke_meta.to_dict(),259 }260 261 tool_responses.append(tool_response)262 if tool_response["tool_response"] is not None:263 self._current_thoughts.append(264 ToolPromptMessage(265 content=tool_response["tool_response"],266 tool_call_id=tool_call_id,267 name=tool_call_name,268 )269 )270 271 if len(tool_responses) > 0:272 # save agent thought273 self.save_agent_thought(274 agent_thought=agent_thought,275 tool_name=None,276 tool_input=None,277 thought=None,278 tool_invoke_meta={279 tool_response["tool_call_name"]: tool_response["meta"] for tool_response in tool_responses280 },281 observation={282 tool_response["tool_call_name"]: tool_response["tool_response"]283 for tool_response in tool_responses284 },285 answer=None,286 messages_ids=message_file_ids,287 )288 self.queue_manager.publish(289 QueueAgentThoughtEvent(agent_thought_id=agent_thought.id), PublishFrom.APPLICATION_MANAGER290 )291 292 # update prompt tool293 for prompt_tool in prompt_messages_tools:294 self.update_prompt_message_tool(tool_instances[prompt_tool.name], prompt_tool)295 296 iteration_step += 1297 298 self.update_db_variables(self.variables_pool, self.db_variables_pool)299 # publish end event300 self.queue_manager.publish(301 QueueMessageEndEvent(302 llm_result=LLMResult(303 model=model_instance.model,304 prompt_messages=prompt_messages,305 message=AssistantPromptMessage(content=final_answer),306 usage=llm_usage["usage"] or LLMUsage.empty_usage(),307 system_fingerprint="",308 )309 ),310 PublishFrom.APPLICATION_MANAGER,311 )312 313 def check_tool_calls(self, llm_result_chunk: LLMResultChunk) -> bool:314 """315 Check if there is any tool call in llm result chunk316 """317 if llm_result_chunk.delta.message.tool_calls:318 return True319 return False320 321 def check_blocking_tool_calls(self, llm_result: LLMResult) -> bool:322 """323 Check if there is any blocking tool call in llm result324 """325 if llm_result.message.tool_calls:326 return True327 return False328 329 def extract_tool_calls(330 self, llm_result_chunk: LLMResultChunk331 ) -> Union[None, list[tuple[str, str, dict[str, Any]]]]:332 """333 Extract tool calls from llm result chunk334 335 Returns:336 List[Tuple[str, str, Dict[str, Any]]]: [(tool_call_id, tool_call_name, tool_call_args)]337 """338 tool_calls = []339 for prompt_message in llm_result_chunk.delta.message.tool_calls:340 args = {}341 if prompt_message.function.arguments != "":342 args = json.loads(prompt_message.function.arguments)343 344 tool_calls.append(345 (346 prompt_message.id,347 prompt_message.function.name,348 args,349 )350 )351 352 return tool_calls353 354 def extract_blocking_tool_calls(self, llm_result: LLMResult) -> Union[None, list[tuple[str, str, dict[str, Any]]]]:355 """356 Extract blocking tool calls from llm result357 358 Returns:359 List[Tuple[str, str, Dict[str, Any]]]: [(tool_call_id, tool_call_name, tool_call_args)]360 """361 tool_calls = []362 for prompt_message in llm_result.message.tool_calls:363 args = {}364 if prompt_message.function.arguments != "":365 args = json.loads(prompt_message.function.arguments)366 367 tool_calls.append(368 (369 prompt_message.id,370 prompt_message.function.name,371 args,372 )373 )374 375 return tool_calls376 377 def _init_system_message(378 self, prompt_template: str, prompt_messages: Optional[list[PromptMessage]] = None379 ) -> list[PromptMessage]:380 """381 Initialize system message382 """383 if not prompt_messages and prompt_template:384 return [385 SystemPromptMessage(content=prompt_template),386 ]387 388 if prompt_messages and not isinstance(prompt_messages[0], SystemPromptMessage) and prompt_template:389 prompt_messages.insert(0, SystemPromptMessage(content=prompt_template))390 391 return prompt_messages392 393 def _organize_user_query(self, query, prompt_messages: list[PromptMessage]) -> list[PromptMessage]:394 """395 Organize user query396 """397 if self.files:398 prompt_message_contents: list[PromptMessageContent] = []399 prompt_message_contents.append(TextPromptMessageContent(data=query))400 for file_obj in self.files:401 prompt_message_contents.append(file_manager.to_prompt_message_content(file_obj))402 403 prompt_messages.append(UserPromptMessage(content=prompt_message_contents))404 else:405 prompt_messages.append(UserPromptMessage(content=query))406 407 return prompt_messages408 409 def _clear_user_prompt_image_messages(self, prompt_messages: list[PromptMessage]) -> list[PromptMessage]:410 """411 As for now, gpt supports both fc and vision at the first iteration.412 We need to remove the image messages from the prompt messages at the first iteration.413 """414 prompt_messages = deepcopy(prompt_messages)415 416 for prompt_message in prompt_messages:417 if isinstance(prompt_message, UserPromptMessage):418 if isinstance(prompt_message.content, list):419 prompt_message.content = "\n".join(420 [421 content.data422 if content.type == PromptMessageContentType.TEXT423 else "[image]"424 if content.type == PromptMessageContentType.IMAGE425 else "[file]"426 for content in prompt_message.content427 ]428 )429 430 return prompt_messages431 432 def _organize_prompt_messages(self):433 prompt_template = self.app_config.prompt_template.simple_prompt_template or ""434 self.history_prompt_messages = self._init_system_message(prompt_template, self.history_prompt_messages)435 query_prompt_messages = self._organize_user_query(self.query, [])436 437 self.history_prompt_messages = AgentHistoryPromptTransform(438 model_config=self.model_config,439 prompt_messages=[*query_prompt_messages, *self._current_thoughts],440 history_messages=self.history_prompt_messages,441 memory=self.memory,442 ).get_prompt()443 444 prompt_messages = [*self.history_prompt_messages, *query_prompt_messages, *self._current_thoughts]445 if len(self._current_thoughts) != 0:446 # clear messages after the first iteration447 prompt_messages = self._clear_user_prompt_image_messages(prompt_messages)448 return prompt_messages449 