Team Ai
Apppublic

Underground-Digital/Workflow-Engine

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
cot_agent_runner.py422 linesDownload Raw Back to agent
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