Underground-Digital/Workflow-Engine
0
1import json2from abc import ABC, abstractmethod3from collections.abc import Generator4from typing import Optional, Union5 6from core.agent.base_agent_runner import BaseAgentRunner7from core.agent.entities import AgentScratchpadUnit8from core.agent.output_parser.cot_output_parser import CotAgentOutputParser9from core.app.apps.base_app_queue_manager import PublishFrom10from core.app.entities.queue_entities import QueueAgentThoughtEvent, QueueMessageEndEvent, QueueMessageFileEvent11from core.model_runtime.entities.llm_entities import LLMResult, LLMResultChunk, LLMResultChunkDelta, LLMUsage12from core.model_runtime.entities.message_entities import (13 AssistantPromptMessage,14 PromptMessage,15 ToolPromptMessage,16 UserPromptMessage,17)18from core.ops.ops_trace_manager import TraceQueueManager19from core.prompt.agent_history_prompt_transform import AgentHistoryPromptTransform20from core.tools.entities.tool_entities import ToolInvokeMeta21from core.tools.tool.tool import Tool22from core.tools.tool_engine import ToolEngine23from models.model import Message24 25 26class CotAgentRunner(BaseAgentRunner, ABC):27 _is_first_iteration = True28 _ignore_observation_providers = ["wenxin"]29 _historic_prompt_messages: list[PromptMessage] = None30 _agent_scratchpad: list[AgentScratchpadUnit] = None31 _instruction: str = None32 _query: str = None33 _prompt_messages_tools: list[PromptMessage] = None34 35 def run(36 self,37 message: Message,38 query: str,39 inputs: dict[str, str],40 ) -> Union[Generator, LLMResult]:41 """42 Run Cot agent application43 """44 app_generate_entity = self.application_generate_entity45 self._repack_app_generate_entity(app_generate_entity)46 self._init_react_state(query)47 48 trace_manager = app_generate_entity.trace_manager49 50 # check model mode51 if "Observation" not in app_generate_entity.model_conf.stop:52 if app_generate_entity.model_conf.provider not in self._ignore_observation_providers:53 app_generate_entity.model_conf.stop.append("Observation")54 55 app_config = self.app_config56 57 # init instruction58 inputs = inputs or {}59 instruction = app_config.prompt_template.simple_prompt_template60 self._instruction = self._fill_in_inputs_from_external_data_tools(instruction, inputs)61 62 iteration_step = 163 max_iteration_steps = min(app_config.agent.max_iteration, 5) + 164 65 # convert tools into ModelRuntime Tool format66 tool_instances, self._prompt_messages_tools = self._init_prompt_tools()67 68 function_call_state = True69 llm_usage = {"usage": None}70 final_answer = ""71 72 def increase_usage(final_llm_usage_dict: dict[str, LLMUsage], usage: LLMUsage):73 if not final_llm_usage_dict["usage"]:74 final_llm_usage_dict["usage"] = usage75 else:76 llm_usage = final_llm_usage_dict["usage"]77 llm_usage.prompt_tokens += usage.prompt_tokens78 llm_usage.completion_tokens += usage.completion_tokens79 llm_usage.prompt_price += usage.prompt_price80 llm_usage.completion_price += usage.completion_price81 llm_usage.total_price += usage.total_price82 83 model_instance = self.model_instance84 85 while function_call_state and iteration_step <= max_iteration_steps:86 # continue to run until there is not any tool call87 function_call_state = False88 89 if iteration_step == max_iteration_steps:90 # the last iteration, remove all tools91 self._prompt_messages_tools = []92 93 message_file_ids = []94 95 agent_thought = self.create_agent_thought(96 message_id=message.id, message="", tool_name="", tool_input="", messages_ids=message_file_ids97 )98 99 if iteration_step > 1:100 self.queue_manager.publish(101 QueueAgentThoughtEvent(agent_thought_id=agent_thought.id), PublishFrom.APPLICATION_MANAGER102 )103 104 # recalc llm max tokens105 prompt_messages = self._organize_prompt_messages()106 self.recalc_llm_max_tokens(self.model_config, prompt_messages)107 # invoke model108 chunks: Generator[LLMResultChunk, None, None] = model_instance.invoke_llm(109 prompt_messages=prompt_messages,110 model_parameters=app_generate_entity.model_conf.parameters,111 tools=[],112 stop=app_generate_entity.model_conf.stop,113 stream=True,114 user=self.user_id,115 callbacks=[],116 )117 118 # check llm result119 if not chunks:120 raise ValueError("failed to invoke llm")121 122 usage_dict = {}123 react_chunks = CotAgentOutputParser.handle_react_stream_output(chunks, usage_dict)124 scratchpad = AgentScratchpadUnit(125 agent_response="",126 thought="",127 action_str="",128 observation="",129 action=None,130 )131 132 # publish agent thought if it's first iteration133 if iteration_step == 1:134 self.queue_manager.publish(135 QueueAgentThoughtEvent(agent_thought_id=agent_thought.id), PublishFrom.APPLICATION_MANAGER136 )137 138 for chunk in react_chunks:139 if isinstance(chunk, AgentScratchpadUnit.Action):140 action = chunk141 # detect action142 scratchpad.agent_response += json.dumps(chunk.model_dump())143 scratchpad.action_str = json.dumps(chunk.model_dump())144 scratchpad.action = action145 else:146 scratchpad.agent_response += chunk147 scratchpad.thought += chunk148 yield LLMResultChunk(149 model=self.model_config.model,150 prompt_messages=prompt_messages,151 system_fingerprint="",152 delta=LLMResultChunkDelta(index=0, message=AssistantPromptMessage(content=chunk), usage=None),153 )154 155 scratchpad.thought = scratchpad.thought.strip() or "I am thinking about how to help you"156 self._agent_scratchpad.append(scratchpad)157 158 # get llm usage159 if "usage" in usage_dict:160 increase_usage(llm_usage, usage_dict["usage"])161 else:162 usage_dict["usage"] = LLMUsage.empty_usage()163 164 self.save_agent_thought(165 agent_thought=agent_thought,166 tool_name=scratchpad.action.action_name if scratchpad.action else "",167 tool_input={scratchpad.action.action_name: scratchpad.action.action_input} if scratchpad.action else {},168 tool_invoke_meta={},169 thought=scratchpad.thought,170 observation="",171 answer=scratchpad.agent_response,172 messages_ids=[],173 llm_usage=usage_dict["usage"],174 )175 176 if not scratchpad.is_final():177 self.queue_manager.publish(178 QueueAgentThoughtEvent(agent_thought_id=agent_thought.id), PublishFrom.APPLICATION_MANAGER179 )180 181 if not scratchpad.action:182 # failed to extract action, return final answer directly183 final_answer = ""184 else:185 if scratchpad.action.action_name.lower() == "final answer":186 # action is final answer, return final answer directly187 try:188 if isinstance(scratchpad.action.action_input, dict):189 final_answer = json.dumps(scratchpad.action.action_input)190 elif isinstance(scratchpad.action.action_input, str):191 final_answer = scratchpad.action.action_input192 else:193 final_answer = f"{scratchpad.action.action_input}"194 except json.JSONDecodeError:195 final_answer = f"{scratchpad.action.action_input}"196 else:197 function_call_state = True198 # action is tool call, invoke tool199 tool_invoke_response, tool_invoke_meta = self._handle_invoke_action(200 action=scratchpad.action,201 tool_instances=tool_instances,202 message_file_ids=message_file_ids,203 trace_manager=trace_manager,204 )205 scratchpad.observation = tool_invoke_response206 scratchpad.agent_response = tool_invoke_response207 208 self.save_agent_thought(209 agent_thought=agent_thought,210 tool_name=scratchpad.action.action_name,211 tool_input={scratchpad.action.action_name: scratchpad.action.action_input},212 thought=scratchpad.thought,213 observation={scratchpad.action.action_name: tool_invoke_response},214 tool_invoke_meta={scratchpad.action.action_name: tool_invoke_meta.to_dict()},215 answer=scratchpad.agent_response,216 messages_ids=message_file_ids,217 llm_usage=usage_dict["usage"],218 )219 220 self.queue_manager.publish(221 QueueAgentThoughtEvent(agent_thought_id=agent_thought.id), PublishFrom.APPLICATION_MANAGER222 )223 224 # update prompt tool message225 for prompt_tool in self._prompt_messages_tools:226 self.update_prompt_message_tool(tool_instances[prompt_tool.name], prompt_tool)227 228 iteration_step += 1229 230 yield LLMResultChunk(231 model=model_instance.model,232 prompt_messages=prompt_messages,233 delta=LLMResultChunkDelta(234 index=0, message=AssistantPromptMessage(content=final_answer), usage=llm_usage["usage"]235 ),236 system_fingerprint="",237 )238 239 # save agent thought240 self.save_agent_thought(241 agent_thought=agent_thought,242 tool_name="",243 tool_input={},244 tool_invoke_meta={},245 thought=final_answer,246 observation={},247 answer=final_answer,248 messages_ids=[],249 )250 251 self.update_db_variables(self.variables_pool, self.db_variables_pool)252 # publish end event253 self.queue_manager.publish(254 QueueMessageEndEvent(255 llm_result=LLMResult(256 model=model_instance.model,257 prompt_messages=prompt_messages,258 message=AssistantPromptMessage(content=final_answer),259 usage=llm_usage["usage"] or LLMUsage.empty_usage(),260 system_fingerprint="",261 )262 ),263 PublishFrom.APPLICATION_MANAGER,264 )265 266 def _handle_invoke_action(267 self,268 action: AgentScratchpadUnit.Action,269 tool_instances: dict[str, Tool],270 message_file_ids: list[str],271 trace_manager: Optional[TraceQueueManager] = None,272 ) -> tuple[str, ToolInvokeMeta]:273 """274 handle invoke action275 :param action: action276 :param tool_instances: tool instances277 :param message_file_ids: message file ids278 :param trace_manager: trace manager279 :return: observation, meta280 """281 # action is tool call, invoke tool282 tool_call_name = action.action_name283 tool_call_args = action.action_input284 tool_instance = tool_instances.get(tool_call_name)285 286 if not tool_instance:287 answer = f"there is not a tool named {tool_call_name}"288 return answer, ToolInvokeMeta.error_instance(answer)289 290 if isinstance(tool_call_args, str):291 try:292 tool_call_args = json.loads(tool_call_args)293 except json.JSONDecodeError:294 pass295 296 # invoke tool297 tool_invoke_response, message_files, tool_invoke_meta = ToolEngine.agent_invoke(298 tool=tool_instance,299 tool_parameters=tool_call_args,300 user_id=self.user_id,301 tenant_id=self.tenant_id,302 message=self.message,303 invoke_from=self.application_generate_entity.invoke_from,304 agent_tool_callback=self.agent_callback,305 trace_manager=trace_manager,306 )307 308 # publish files309 for message_file_id, save_as in message_files:310 if save_as:311 self.variables_pool.set_file(tool_name=tool_call_name, value=message_file_id, name=save_as)312 313 # publish message file314 self.queue_manager.publish(315 QueueMessageFileEvent(message_file_id=message_file_id), PublishFrom.APPLICATION_MANAGER316 )317 # add message file ids318 message_file_ids.append(message_file_id)319 320 return tool_invoke_response, tool_invoke_meta321 322 def _convert_dict_to_action(self, action: dict) -> AgentScratchpadUnit.Action:323 """324 convert dict to action325 """326 return AgentScratchpadUnit.Action(action_name=action["action"], action_input=action["action_input"])327 328 def _fill_in_inputs_from_external_data_tools(self, instruction: str, inputs: dict) -> str:329 """330 fill in inputs from external data tools331 """332 for key, value in inputs.items():333 try:334 instruction = instruction.replace(f"{{{{{key}}}}}", str(value))335 except Exception as e:336 continue337 338 return instruction339 340 def _init_react_state(self, query) -> None:341 """342 init agent scratchpad343 """344 self._query = query345 self._agent_scratchpad = []346 self._historic_prompt_messages = self._organize_historic_prompt_messages()347 348 @abstractmethod349 def _organize_prompt_messages(self) -> list[PromptMessage]:350 """351 organize prompt messages352 """353 354 def _format_assistant_message(self, agent_scratchpad: list[AgentScratchpadUnit]) -> str:355 """356 format assistant message357 """358 message = ""359 for scratchpad in agent_scratchpad:360 if scratchpad.is_final():361 message += f"Final Answer: {scratchpad.agent_response}"362 else:363 message += f"Thought: {scratchpad.thought}\n\n"364 if scratchpad.action_str:365 message += f"Action: {scratchpad.action_str}\n\n"366 if scratchpad.observation:367 message += f"Observation: {scratchpad.observation}\n\n"368 369 return message370 371 def _organize_historic_prompt_messages(372 self, current_session_messages: Optional[list[PromptMessage]] = None373 ) -> list[PromptMessage]:374 """375 organize historic prompt messages376 """377 result: list[PromptMessage] = []378 scratchpads: list[AgentScratchpadUnit] = []379 current_scratchpad: AgentScratchpadUnit = None380 381 for message in self.history_prompt_messages:382 if isinstance(message, AssistantPromptMessage):383 if not current_scratchpad:384 current_scratchpad = AgentScratchpadUnit(385 agent_response=message.content,386 thought=message.content or "I am thinking about how to help you",387 action_str="",388 action=None,389 observation=None,390 )391 scratchpads.append(current_scratchpad)392 if message.tool_calls:393 try:394 current_scratchpad.action = AgentScratchpadUnit.Action(395 action_name=message.tool_calls[0].function.name,396 action_input=json.loads(message.tool_calls[0].function.arguments),397 )398 current_scratchpad.action_str = json.dumps(current_scratchpad.action.to_dict())399 except:400 pass401 elif isinstance(message, ToolPromptMessage):402 if current_scratchpad:403 current_scratchpad.observation = message.content404 elif isinstance(message, UserPromptMessage):405 if scratchpads:406 result.append(AssistantPromptMessage(content=self._format_assistant_message(scratchpads)))407 scratchpads = []408 current_scratchpad = None409 410 result.append(message)411 412 if scratchpads:413 result.append(AssistantPromptMessage(content=self._format_assistant_message(scratchpads)))414 415 historic_prompts = AgentHistoryPromptTransform(416 model_config=self.model_config,417 prompt_messages=current_session_messages or [],418 history_messages=result,419 memory=self.memory,420 ).get_prompt()421 return historic_prompts422 