Team Ai
Apppublic

Underground-Digital/Workflow-Engine

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
base_agent_runner.py532 linesDownload Raw Back to agent
1import json2import logging3import uuid4from collections.abc import Mapping, Sequence5from datetime import datetime, timezone6from typing import Optional, Union, cast7 8from core.agent.entities import AgentEntity, AgentToolEntity9from core.app.app_config.features.file_upload.manager import FileUploadConfigManager10from core.app.apps.agent_chat.app_config_manager import AgentChatAppConfig11from core.app.apps.base_app_queue_manager import AppQueueManager12from core.app.apps.base_app_runner import AppRunner13from core.app.entities.app_invoke_entities import (14    AgentChatAppGenerateEntity,15    ModelConfigWithCredentialsEntity,16)17from core.callback_handler.agent_tool_callback_handler import DifyAgentCallbackHandler18from core.callback_handler.index_tool_callback_handler import DatasetIndexToolCallbackHandler19from core.file import file_manager20from core.memory.token_buffer_memory import TokenBufferMemory21from core.model_manager import ModelInstance22from core.model_runtime.entities import (23    AssistantPromptMessage,24    LLMUsage,25    PromptMessage,26    PromptMessageContent,27    PromptMessageTool,28    SystemPromptMessage,29    TextPromptMessageContent,30    ToolPromptMessage,31    UserPromptMessage,32)33from core.model_runtime.entities.model_entities import ModelFeature34from core.model_runtime.model_providers.__base.large_language_model import LargeLanguageModel35from core.model_runtime.utils.encoders import jsonable_encoder36from core.prompt.utils.extract_thread_messages import extract_thread_messages37from core.tools.entities.tool_entities import (38    ToolParameter,39    ToolRuntimeVariablePool,40)41from core.tools.tool.dataset_retriever_tool import DatasetRetrieverTool42from core.tools.tool.tool import Tool43from core.tools.tool_manager import ToolManager44from extensions.ext_database import db45from factories import file_factory46from models.model import Conversation, Message, MessageAgentThought, MessageFile47from models.tools import ToolConversationVariables48 49logger = logging.getLogger(__name__)50 51 52class BaseAgentRunner(AppRunner):53    def __init__(54        self,55        tenant_id: str,56        application_generate_entity: AgentChatAppGenerateEntity,57        conversation: Conversation,58        app_config: AgentChatAppConfig,59        model_config: ModelConfigWithCredentialsEntity,60        config: AgentEntity,61        queue_manager: AppQueueManager,62        message: Message,63        user_id: str,64        memory: Optional[TokenBufferMemory] = None,65        prompt_messages: Optional[list[PromptMessage]] = None,66        variables_pool: Optional[ToolRuntimeVariablePool] = None,67        db_variables: Optional[ToolConversationVariables] = None,68        model_instance: ModelInstance = None,69    ) -> None:70        self.tenant_id = tenant_id71        self.application_generate_entity = application_generate_entity72        self.conversation = conversation73        self.app_config = app_config74        self.model_config = model_config75        self.config = config76        self.queue_manager = queue_manager77        self.message = message78        self.user_id = user_id79        self.memory = memory80        self.history_prompt_messages = self.organize_agent_history(prompt_messages=prompt_messages or [])81        self.variables_pool = variables_pool82        self.db_variables_pool = db_variables83        self.model_instance = model_instance84 85        # init callback86        self.agent_callback = DifyAgentCallbackHandler()87        # init dataset tools88        hit_callback = DatasetIndexToolCallbackHandler(89            queue_manager=queue_manager,90            app_id=self.app_config.app_id,91            message_id=message.id,92            user_id=user_id,93            invoke_from=self.application_generate_entity.invoke_from,94        )95        self.dataset_tools = DatasetRetrieverTool.get_dataset_tools(96            tenant_id=tenant_id,97            dataset_ids=app_config.dataset.dataset_ids if app_config.dataset else [],98            retrieve_config=app_config.dataset.retrieve_config if app_config.dataset else None,99            return_resource=app_config.additional_features.show_retrieve_source,100            invoke_from=application_generate_entity.invoke_from,101            hit_callback=hit_callback,102        )103        # get how many agent thoughts have been created104        self.agent_thought_count = (105            db.session.query(MessageAgentThought)106            .filter(107                MessageAgentThought.message_id == self.message.id,108            )109            .count()110        )111        db.session.close()112 113        # check if model supports stream tool call114        llm_model = cast(LargeLanguageModel, model_instance.model_type_instance)115        model_schema = llm_model.get_model_schema(model_instance.model, model_instance.credentials)116        if model_schema and ModelFeature.STREAM_TOOL_CALL in (model_schema.features or []):117            self.stream_tool_call = True118        else:119            self.stream_tool_call = False120 121        # check if model supports vision122        if model_schema and ModelFeature.VISION in (model_schema.features or []):123            self.files = application_generate_entity.files124        else:125            self.files = []126        self.query = None127        self._current_thoughts: list[PromptMessage] = []128 129    def _repack_app_generate_entity(130        self, app_generate_entity: AgentChatAppGenerateEntity131    ) -> AgentChatAppGenerateEntity:132        """133        Repack app generate entity134        """135        if app_generate_entity.app_config.prompt_template.simple_prompt_template is None:136            app_generate_entity.app_config.prompt_template.simple_prompt_template = ""137 138        return app_generate_entity139 140    def _convert_tool_to_prompt_message_tool(self, tool: AgentToolEntity) -> tuple[PromptMessageTool, Tool]:141        """142        convert tool to prompt message tool143        """144        tool_entity = ToolManager.get_agent_tool_runtime(145            tenant_id=self.tenant_id,146            app_id=self.app_config.app_id,147            agent_tool=tool,148            invoke_from=self.application_generate_entity.invoke_from,149        )150        tool_entity.load_variables(self.variables_pool)151 152        message_tool = PromptMessageTool(153            name=tool.tool_name,154            description=tool_entity.description.llm,155            parameters={156                "type": "object",157                "properties": {},158                "required": [],159            },160        )161 162        parameters = tool_entity.get_all_runtime_parameters()163        for parameter in parameters:164            if parameter.form != ToolParameter.ToolParameterForm.LLM:165                continue166 167            parameter_type = parameter.type.as_normal_type()168            if parameter.type in {169                ToolParameter.ToolParameterType.SYSTEM_FILES,170                ToolParameter.ToolParameterType.FILE,171                ToolParameter.ToolParameterType.FILES,172            }:173                continue174            enum = []175            if parameter.type == ToolParameter.ToolParameterType.SELECT:176                enum = [option.value for option in parameter.options]177 178            message_tool.parameters["properties"][parameter.name] = {179                "type": parameter_type,180                "description": parameter.llm_description or "",181            }182 183            if len(enum) > 0:184                message_tool.parameters["properties"][parameter.name]["enum"] = enum185 186            if parameter.required:187                message_tool.parameters["required"].append(parameter.name)188 189        return message_tool, tool_entity190 191    def _convert_dataset_retriever_tool_to_prompt_message_tool(self, tool: DatasetRetrieverTool) -> PromptMessageTool:192        """193        convert dataset retriever tool to prompt message tool194        """195        prompt_tool = PromptMessageTool(196            name=tool.identity.name,197            description=tool.description.llm,198            parameters={199                "type": "object",200                "properties": {},201                "required": [],202            },203        )204 205        for parameter in tool.get_runtime_parameters():206            parameter_type = "string"207 208            prompt_tool.parameters["properties"][parameter.name] = {209                "type": parameter_type,210                "description": parameter.llm_description or "",211            }212 213            if parameter.required:214                if parameter.name not in prompt_tool.parameters["required"]:215                    prompt_tool.parameters["required"].append(parameter.name)216 217        return prompt_tool218 219    def _init_prompt_tools(self) -> tuple[Mapping[str, Tool], Sequence[PromptMessageTool]]:220        """221        Init tools222        """223        tool_instances = {}224        prompt_messages_tools = []225 226        for tool in self.app_config.agent.tools if self.app_config.agent else []:227            try:228                prompt_tool, tool_entity = self._convert_tool_to_prompt_message_tool(tool)229            except Exception:230                # api tool may be deleted231                continue232            # save tool entity233            tool_instances[tool.tool_name] = tool_entity234            # save prompt tool235            prompt_messages_tools.append(prompt_tool)236 237        # convert dataset tools into ModelRuntime Tool format238        for dataset_tool in self.dataset_tools:239            prompt_tool = self._convert_dataset_retriever_tool_to_prompt_message_tool(dataset_tool)240            # save prompt tool241            prompt_messages_tools.append(prompt_tool)242            # save tool entity243            tool_instances[dataset_tool.identity.name] = dataset_tool244 245        return tool_instances, prompt_messages_tools246 247    def update_prompt_message_tool(self, tool: Tool, prompt_tool: PromptMessageTool) -> PromptMessageTool:248        """249        update prompt message tool250        """251        # try to get tool runtime parameters252        tool_runtime_parameters = tool.get_runtime_parameters() or []253 254        for parameter in tool_runtime_parameters:255            if parameter.form != ToolParameter.ToolParameterForm.LLM:256                continue257 258            parameter_type = parameter.type.as_normal_type()259            if parameter.type in {260                ToolParameter.ToolParameterType.SYSTEM_FILES,261                ToolParameter.ToolParameterType.FILE,262                ToolParameter.ToolParameterType.FILES,263            }:264                continue265            enum = []266            if parameter.type == ToolParameter.ToolParameterType.SELECT:267                enum = [option.value for option in parameter.options]268 269            prompt_tool.parameters["properties"][parameter.name] = {270                "type": parameter_type,271                "description": parameter.llm_description or "",272            }273 274            if len(enum) > 0:275                prompt_tool.parameters["properties"][parameter.name]["enum"] = enum276 277            if parameter.required:278                if parameter.name not in prompt_tool.parameters["required"]:279                    prompt_tool.parameters["required"].append(parameter.name)280 281        return prompt_tool282 283    def create_agent_thought(284        self, message_id: str, message: str, tool_name: str, tool_input: str, messages_ids: list[str]285    ) -> MessageAgentThought:286        """287        Create agent thought288        """289        thought = MessageAgentThought(290            message_id=message_id,291            message_chain_id=None,292            thought="",293            tool=tool_name,294            tool_labels_str="{}",295            tool_meta_str="{}",296            tool_input=tool_input,297            message=message,298            message_token=0,299            message_unit_price=0,300            message_price_unit=0,301            message_files=json.dumps(messages_ids) if messages_ids else "",302            answer="",303            observation="",304            answer_token=0,305            answer_unit_price=0,306            answer_price_unit=0,307            tokens=0,308            total_price=0,309            position=self.agent_thought_count + 1,310            currency="USD",311            latency=0,312            created_by_role="account",313            created_by=self.user_id,314        )315 316        db.session.add(thought)317        db.session.commit()318        db.session.refresh(thought)319        db.session.close()320 321        self.agent_thought_count += 1322 323        return thought324 325    def save_agent_thought(326        self,327        agent_thought: MessageAgentThought,328        tool_name: str,329        tool_input: Union[str, dict],330        thought: str,331        observation: Union[str, dict],332        tool_invoke_meta: Union[str, dict],333        answer: str,334        messages_ids: list[str],335        llm_usage: LLMUsage = None,336    ) -> MessageAgentThought:337        """338        Save agent thought339        """340        agent_thought = db.session.query(MessageAgentThought).filter(MessageAgentThought.id == agent_thought.id).first()341 342        if thought is not None:343            agent_thought.thought = thought344 345        if tool_name is not None:346            agent_thought.tool = tool_name347 348        if tool_input is not None:349            if isinstance(tool_input, dict):350                try:351                    tool_input = json.dumps(tool_input, ensure_ascii=False)352                except Exception as e:353                    tool_input = json.dumps(tool_input)354 355            agent_thought.tool_input = tool_input356 357        if observation is not None:358            if isinstance(observation, dict):359                try:360                    observation = json.dumps(observation, ensure_ascii=False)361                except Exception as e:362                    observation = json.dumps(observation)363 364            agent_thought.observation = observation365 366        if answer is not None:367            agent_thought.answer = answer368 369        if messages_ids is not None and len(messages_ids) > 0:370            agent_thought.message_files = json.dumps(messages_ids)371 372        if llm_usage:373            agent_thought.message_token = llm_usage.prompt_tokens374            agent_thought.message_price_unit = llm_usage.prompt_price_unit375            agent_thought.message_unit_price = llm_usage.prompt_unit_price376            agent_thought.answer_token = llm_usage.completion_tokens377            agent_thought.answer_price_unit = llm_usage.completion_price_unit378            agent_thought.answer_unit_price = llm_usage.completion_unit_price379            agent_thought.tokens = llm_usage.total_tokens380            agent_thought.total_price = llm_usage.total_price381 382        # check if tool labels is not empty383        labels = agent_thought.tool_labels or {}384        tools = agent_thought.tool.split(";") if agent_thought.tool else []385        for tool in tools:386            if not tool:387                continue388            if tool not in labels:389                tool_label = ToolManager.get_tool_label(tool)390                if tool_label:391                    labels[tool] = tool_label.to_dict()392                else:393                    labels[tool] = {"en_US": tool, "zh_Hans": tool}394 395        agent_thought.tool_labels_str = json.dumps(labels)396 397        if tool_invoke_meta is not None:398            if isinstance(tool_invoke_meta, dict):399                try:400                    tool_invoke_meta = json.dumps(tool_invoke_meta, ensure_ascii=False)401                except Exception as e:402                    tool_invoke_meta = json.dumps(tool_invoke_meta)403 404            agent_thought.tool_meta_str = tool_invoke_meta405 406        db.session.commit()407        db.session.close()408 409    def update_db_variables(self, tool_variables: ToolRuntimeVariablePool, db_variables: ToolConversationVariables):410        """411        convert tool variables to db variables412        """413        db_variables = (414            db.session.query(ToolConversationVariables)415            .filter(416                ToolConversationVariables.conversation_id == self.message.conversation_id,417            )418            .first()419        )420 421        db_variables.updated_at = datetime.now(timezone.utc).replace(tzinfo=None)422        db_variables.variables_str = json.dumps(jsonable_encoder(tool_variables.pool))423        db.session.commit()424        db.session.close()425 426    def organize_agent_history(self, prompt_messages: list[PromptMessage]) -> list[PromptMessage]:427        """428        Organize agent history429        """430        result = []431        # check if there is a system message in the beginning of the conversation432        for prompt_message in prompt_messages:433            if isinstance(prompt_message, SystemPromptMessage):434                result.append(prompt_message)435 436        messages: list[Message] = (437            db.session.query(Message)438            .filter(439                Message.conversation_id == self.message.conversation_id,440            )441            .order_by(Message.created_at.desc())442            .all()443        )444 445        messages = list(reversed(extract_thread_messages(messages)))446 447        for message in messages:448            if message.id == self.message.id:449                continue450 451            result.append(self.organize_agent_user_prompt(message))452            agent_thoughts: list[MessageAgentThought] = message.agent_thoughts453            if agent_thoughts:454                for agent_thought in agent_thoughts:455                    tools = agent_thought.tool456                    if tools:457                        tools = tools.split(";")458                        tool_calls: list[AssistantPromptMessage.ToolCall] = []459                        tool_call_response: list[ToolPromptMessage] = []460                        try:461                            tool_inputs = json.loads(agent_thought.tool_input)462                        except Exception as e:463                            tool_inputs = {tool: {} for tool in tools}464                        try:465                            tool_responses = json.loads(agent_thought.observation)466                        except Exception as e:467                            tool_responses = dict.fromkeys(tools, agent_thought.observation)468 469                        for tool in tools:470                            # generate a uuid for tool call471                            tool_call_id = str(uuid.uuid4())472                            tool_calls.append(473                                AssistantPromptMessage.ToolCall(474                                    id=tool_call_id,475                                    type="function",476                                    function=AssistantPromptMessage.ToolCall.ToolCallFunction(477                                        name=tool,478                                        arguments=json.dumps(tool_inputs.get(tool, {})),479                                    ),480                                )481                            )482                            tool_call_response.append(483                                ToolPromptMessage(484                                    content=tool_responses.get(tool, agent_thought.observation),485                                    name=tool,486                                    tool_call_id=tool_call_id,487                                )488                            )489 490                        result.extend(491                            [492                                AssistantPromptMessage(493                                    content=agent_thought.thought,494                                    tool_calls=tool_calls,495                                ),496                                *tool_call_response,497                            ]498                        )499                    if not tools:500                        result.append(AssistantPromptMessage(content=agent_thought.thought))501            else:502                if message.answer:503                    result.append(AssistantPromptMessage(content=message.answer))504 505        db.session.close()506 507        return result508 509    def organize_agent_user_prompt(self, message: Message) -> UserPromptMessage:510        files = db.session.query(MessageFile).filter(MessageFile.message_id == message.id).all()511        if files:512            file_extra_config = FileUploadConfigManager.convert(message.app_model_config.to_dict())513 514            if file_extra_config:515                file_objs = file_factory.build_from_message_files(516                    message_files=files, tenant_id=self.tenant_id, config=file_extra_config517                )518            else:519                file_objs = []520 521            if not file_objs:522                return UserPromptMessage(content=message.query)523            else:524                prompt_message_contents: list[PromptMessageContent] = []525                prompt_message_contents.append(TextPromptMessageContent(data=message.query))526                for file_obj in file_objs:527                    prompt_message_contents.append(file_manager.to_prompt_message_content(file_obj))528 529                return UserPromptMessage(content=prompt_message_contents)530        else:531            return UserPromptMessage(content=message.query)532