Team Ai
Apppublic

Underground-Digital/Workflow-Engine

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