Underground-Digital/Workflow-Engine
0
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 