Underground-Digital/Workflow-Engine
0
1import json2import logging3import mimetypes4from collections.abc import Generator5from os import listdir, path6from threading import Lock, Thread7from typing import Any, Optional, Union8 9from configs import dify_config10from core.agent.entities import AgentToolEntity11from core.app.entities.app_invoke_entities import InvokeFrom12from core.helper.module_import_helper import load_single_subclass_from_source13from core.helper.position_helper import is_filtered14from core.model_runtime.utils.encoders import jsonable_encoder15from core.tools.entities.api_entities import UserToolProvider, UserToolProviderTypeLiteral16from core.tools.entities.common_entities import I18nObject17from core.tools.entities.tool_entities import ApiProviderAuthType, ToolInvokeFrom, ToolParameter18from core.tools.errors import ToolProviderNotFoundError19from core.tools.provider.api_tool_provider import ApiToolProviderController20from core.tools.provider.builtin._positions import BuiltinToolProviderSort21from core.tools.provider.builtin_tool_provider import BuiltinToolProviderController22from core.tools.tool.api_tool import ApiTool23from core.tools.tool.builtin_tool import BuiltinTool24from core.tools.tool.tool import Tool25from core.tools.tool_label_manager import ToolLabelManager26from core.tools.utils.configuration import ToolConfigurationManager, ToolParameterConfigurationManager27from extensions.ext_database import db28from models.tools import ApiToolProvider, BuiltinToolProvider, WorkflowToolProvider29from services.tools.tools_transform_service import ToolTransformService30 31logger = logging.getLogger(__name__)32 33 34class ToolManager:35 _builtin_provider_lock = Lock()36 _builtin_providers = {}37 _builtin_providers_loaded = False38 _builtin_tools_labels = {}39 40 @classmethod41 def get_builtin_provider(cls, provider: str) -> BuiltinToolProviderController:42 """43 get the builtin provider44 45 :param provider: the name of the provider46 :return: the provider47 """48 if len(cls._builtin_providers) == 0:49 # init the builtin providers50 cls.load_builtin_providers_cache()51 52 if provider not in cls._builtin_providers:53 raise ToolProviderNotFoundError(f"builtin provider {provider} not found")54 55 return cls._builtin_providers[provider]56 57 @classmethod58 def get_builtin_tool(cls, provider: str, tool_name: str) -> BuiltinTool:59 """60 get the builtin tool61 62 :param provider: the name of the provider63 :param tool_name: the name of the tool64 65 :return: the provider, the tool66 """67 provider_controller = cls.get_builtin_provider(provider)68 tool = provider_controller.get_tool(tool_name)69 70 return tool71 72 @classmethod73 def get_tool(74 cls, provider_type: str, provider_id: str, tool_name: str, tenant_id: Optional[str] = None75 ) -> Union[BuiltinTool, ApiTool]:76 """77 get the tool78 79 :param provider_type: the type of the provider80 :param provider_name: the name of the provider81 :param tool_name: the name of the tool82 83 :return: the tool84 """85 if provider_type == "builtin":86 return cls.get_builtin_tool(provider_id, tool_name)87 elif provider_type == "api":88 if tenant_id is None:89 raise ValueError("tenant id is required for api provider")90 api_provider, _ = cls.get_api_provider_controller(tenant_id, provider_id)91 return api_provider.get_tool(tool_name)92 elif provider_type == "app":93 raise NotImplementedError("app provider not implemented")94 else:95 raise ToolProviderNotFoundError(f"provider type {provider_type} not found")96 97 @classmethod98 def get_tool_runtime(99 cls,100 provider_type: str,101 provider_id: str,102 tool_name: str,103 tenant_id: str,104 invoke_from: InvokeFrom = InvokeFrom.DEBUGGER,105 tool_invoke_from: ToolInvokeFrom = ToolInvokeFrom.AGENT,106 ) -> Union[BuiltinTool, ApiTool]:107 """108 get the tool runtime109 110 :param provider_type: the type of the provider111 :param provider_name: the name of the provider112 :param tool_name: the name of the tool113 114 :return: the tool115 """116 if provider_type == "builtin":117 builtin_tool = cls.get_builtin_tool(provider_id, tool_name)118 119 # check if the builtin tool need credentials120 provider_controller = cls.get_builtin_provider(provider_id)121 if not provider_controller.need_credentials:122 return builtin_tool.fork_tool_runtime(123 runtime={124 "tenant_id": tenant_id,125 "credentials": {},126 "invoke_from": invoke_from,127 "tool_invoke_from": tool_invoke_from,128 }129 )130 131 # get credentials132 builtin_provider: BuiltinToolProvider = (133 db.session.query(BuiltinToolProvider)134 .filter(135 BuiltinToolProvider.tenant_id == tenant_id,136 BuiltinToolProvider.provider == provider_id,137 )138 .first()139 )140 141 if builtin_provider is None:142 raise ToolProviderNotFoundError(f"builtin provider {provider_id} not found")143 144 # decrypt the credentials145 credentials = builtin_provider.credentials146 controller = cls.get_builtin_provider(provider_id)147 tool_configuration = ToolConfigurationManager(tenant_id=tenant_id, provider_controller=controller)148 149 decrypted_credentials = tool_configuration.decrypt_tool_credentials(credentials)150 151 return builtin_tool.fork_tool_runtime(152 runtime={153 "tenant_id": tenant_id,154 "credentials": decrypted_credentials,155 "runtime_parameters": {},156 "invoke_from": invoke_from,157 "tool_invoke_from": tool_invoke_from,158 }159 )160 161 elif provider_type == "api":162 if tenant_id is None:163 raise ValueError("tenant id is required for api provider")164 165 api_provider, credentials = cls.get_api_provider_controller(tenant_id, provider_id)166 167 # decrypt the credentials168 tool_configuration = ToolConfigurationManager(tenant_id=tenant_id, provider_controller=api_provider)169 decrypted_credentials = tool_configuration.decrypt_tool_credentials(credentials)170 171 return api_provider.get_tool(tool_name).fork_tool_runtime(172 runtime={173 "tenant_id": tenant_id,174 "credentials": decrypted_credentials,175 "invoke_from": invoke_from,176 "tool_invoke_from": tool_invoke_from,177 }178 )179 elif provider_type == "workflow":180 workflow_provider = (181 db.session.query(WorkflowToolProvider)182 .filter(WorkflowToolProvider.tenant_id == tenant_id, WorkflowToolProvider.id == provider_id)183 .first()184 )185 186 if workflow_provider is None:187 raise ToolProviderNotFoundError(f"workflow provider {provider_id} not found")188 189 controller = ToolTransformService.workflow_provider_to_controller(db_provider=workflow_provider)190 191 return controller.get_tools(user_id=None, tenant_id=workflow_provider.tenant_id)[0].fork_tool_runtime(192 runtime={193 "tenant_id": tenant_id,194 "credentials": {},195 "invoke_from": invoke_from,196 "tool_invoke_from": tool_invoke_from,197 }198 )199 elif provider_type == "app":200 raise NotImplementedError("app provider not implemented")201 else:202 raise ToolProviderNotFoundError(f"provider type {provider_type} not found")203 204 @classmethod205 def _init_runtime_parameter(cls, parameter_rule: ToolParameter, parameters: dict):206 """207 init runtime parameter208 """209 parameter_value = parameters.get(parameter_rule.name)210 if not parameter_value and parameter_value != 0:211 # get default value212 parameter_value = parameter_rule.default213 if not parameter_value and parameter_rule.required:214 raise ValueError(f"tool parameter {parameter_rule.name} not found in tool config")215 216 if parameter_rule.type == ToolParameter.ToolParameterType.SELECT:217 # check if tool_parameter_config in options218 options = [x.value for x in parameter_rule.options]219 if parameter_value is not None and parameter_value not in options:220 raise ValueError(221 f"tool parameter {parameter_rule.name} value {parameter_value} not in options {options}"222 )223 224 return parameter_rule.type.cast_value(parameter_value)225 226 @classmethod227 def get_agent_tool_runtime(228 cls, tenant_id: str, app_id: str, agent_tool: AgentToolEntity, invoke_from: InvokeFrom = InvokeFrom.DEBUGGER229 ) -> Tool:230 """231 get the agent tool runtime232 """233 tool_entity = cls.get_tool_runtime(234 provider_type=agent_tool.provider_type,235 provider_id=agent_tool.provider_id,236 tool_name=agent_tool.tool_name,237 tenant_id=tenant_id,238 invoke_from=invoke_from,239 tool_invoke_from=ToolInvokeFrom.AGENT,240 )241 runtime_parameters = {}242 parameters = tool_entity.get_all_runtime_parameters()243 for parameter in parameters:244 # check file types245 if (246 parameter.type247 in {248 ToolParameter.ToolParameterType.SYSTEM_FILES,249 ToolParameter.ToolParameterType.FILE,250 ToolParameter.ToolParameterType.FILES,251 }252 and parameter.required253 ):254 raise ValueError(f"file type parameter {parameter.name} not supported in agent")255 256 if parameter.form == ToolParameter.ToolParameterForm.FORM:257 # save tool parameter to tool entity memory258 value = cls._init_runtime_parameter(parameter, agent_tool.tool_parameters)259 runtime_parameters[parameter.name] = value260 261 # decrypt runtime parameters262 encryption_manager = ToolParameterConfigurationManager(263 tenant_id=tenant_id,264 tool_runtime=tool_entity,265 provider_name=agent_tool.provider_id,266 provider_type=agent_tool.provider_type,267 identity_id=f"AGENT.{app_id}",268 )269 runtime_parameters = encryption_manager.decrypt_tool_parameters(runtime_parameters)270 271 tool_entity.runtime.runtime_parameters.update(runtime_parameters)272 return tool_entity273 274 @classmethod275 def get_workflow_tool_runtime(276 cls,277 tenant_id: str,278 app_id: str,279 node_id: str,280 workflow_tool: "ToolEntity",281 invoke_from: InvokeFrom = InvokeFrom.DEBUGGER,282 ) -> Tool:283 """284 get the workflow tool runtime285 """286 tool_entity = cls.get_tool_runtime(287 provider_type=workflow_tool.provider_type,288 provider_id=workflow_tool.provider_id,289 tool_name=workflow_tool.tool_name,290 tenant_id=tenant_id,291 invoke_from=invoke_from,292 tool_invoke_from=ToolInvokeFrom.WORKFLOW,293 )294 runtime_parameters = {}295 parameters = tool_entity.get_all_runtime_parameters()296 297 for parameter in parameters:298 # save tool parameter to tool entity memory299 if parameter.form == ToolParameter.ToolParameterForm.FORM:300 value = cls._init_runtime_parameter(parameter, workflow_tool.tool_configurations)301 runtime_parameters[parameter.name] = value302 303 # decrypt runtime parameters304 encryption_manager = ToolParameterConfigurationManager(305 tenant_id=tenant_id,306 tool_runtime=tool_entity,307 provider_name=workflow_tool.provider_id,308 provider_type=workflow_tool.provider_type,309 identity_id=f"WORKFLOW.{app_id}.{node_id}",310 )311 312 if runtime_parameters:313 runtime_parameters = encryption_manager.decrypt_tool_parameters(runtime_parameters)314 315 tool_entity.runtime.runtime_parameters.update(runtime_parameters)316 return tool_entity317 318 @classmethod319 def get_builtin_provider_icon(cls, provider: str) -> tuple[str, str]:320 """321 get the absolute path of the icon of the builtin provider322 323 :param provider: the name of the provider324 325 :return: the absolute path of the icon, the mime type of the icon326 """327 # get provider328 provider_controller = cls.get_builtin_provider(provider)329 330 absolute_path = path.join(331 path.dirname(path.realpath(__file__)),332 "provider",333 "builtin",334 provider,335 "_assets",336 provider_controller.identity.icon,337 )338 # check if the icon exists339 if not path.exists(absolute_path):340 raise ToolProviderNotFoundError(f"builtin provider {provider} icon not found")341 342 # get the mime type343 mime_type, _ = mimetypes.guess_type(absolute_path)344 mime_type = mime_type or "application/octet-stream"345 346 return absolute_path, mime_type347 348 @classmethod349 def list_builtin_providers(cls) -> Generator[BuiltinToolProviderController, None, None]:350 # use cache first351 if cls._builtin_providers_loaded:352 yield from list(cls._builtin_providers.values())353 return354 355 with cls._builtin_provider_lock:356 if cls._builtin_providers_loaded:357 yield from list(cls._builtin_providers.values())358 return359 360 yield from cls._list_builtin_providers()361 362 @classmethod363 def _list_builtin_providers(cls) -> Generator[BuiltinToolProviderController, None, None]:364 """365 list all the builtin providers366 """367 for provider in listdir(path.join(path.dirname(path.realpath(__file__)), "provider", "builtin")):368 if provider.startswith("__"):369 continue370 371 if path.isdir(path.join(path.dirname(path.realpath(__file__)), "provider", "builtin", provider)):372 if provider.startswith("__"):373 continue374 375 # init provider376 try:377 provider_class = load_single_subclass_from_source(378 module_name=f"core.tools.provider.builtin.{provider}.{provider}",379 script_path=path.join(380 path.dirname(path.realpath(__file__)), "provider", "builtin", provider, f"{provider}.py"381 ),382 parent_type=BuiltinToolProviderController,383 )384 provider: BuiltinToolProviderController = provider_class()385 cls._builtin_providers[provider.identity.name] = provider386 for tool in provider.get_tools():387 cls._builtin_tools_labels[tool.identity.name] = tool.identity.label388 yield provider389 390 except Exception as e:391 logger.error(f"load builtin provider {provider} error: {e}")392 continue393 # set builtin providers loaded394 cls._builtin_providers_loaded = True395 396 @classmethod397 def load_builtin_providers_cache(cls):398 for _ in cls.list_builtin_providers():399 pass400 401 @classmethod402 def clear_builtin_providers_cache(cls):403 cls._builtin_providers = {}404 cls._builtin_providers_loaded = False405 406 @classmethod407 def get_tool_label(cls, tool_name: str) -> Union[I18nObject, None]:408 """409 get the tool label410 411 :param tool_name: the name of the tool412 413 :return: the label of the tool414 """415 if len(cls._builtin_tools_labels) == 0:416 # init the builtin providers417 cls.load_builtin_providers_cache()418 419 if tool_name not in cls._builtin_tools_labels:420 return None421 422 return cls._builtin_tools_labels[tool_name]423 424 @classmethod425 def user_list_providers(426 cls, user_id: str, tenant_id: str, typ: UserToolProviderTypeLiteral427 ) -> list[UserToolProvider]:428 result_providers: dict[str, UserToolProvider] = {}429 430 filters = []431 if not typ:432 filters.extend(["builtin", "api", "workflow"])433 else:434 filters.append(typ)435 436 if "builtin" in filters:437 # get builtin providers438 builtin_providers = cls.list_builtin_providers()439 440 # get db builtin providers441 db_builtin_providers: list[BuiltinToolProvider] = (442 db.session.query(BuiltinToolProvider).filter(BuiltinToolProvider.tenant_id == tenant_id).all()443 )444 445 find_db_builtin_provider = lambda provider: next(446 (x for x in db_builtin_providers if x.provider == provider), None447 )448 449 # append builtin providers450 for provider in builtin_providers:451 # handle include, exclude452 if is_filtered(453 include_set=dify_config.POSITION_TOOL_INCLUDES_SET,454 exclude_set=dify_config.POSITION_TOOL_EXCLUDES_SET,455 data=provider,456 name_func=lambda x: x.identity.name,457 ):458 continue459 460 user_provider = ToolTransformService.builtin_provider_to_user_provider(461 provider_controller=provider,462 db_provider=find_db_builtin_provider(provider.identity.name),463 decrypt_credentials=False,464 )465 466 result_providers[provider.identity.name] = user_provider467 468 # get db api providers469 470 if "api" in filters:471 db_api_providers: list[ApiToolProvider] = (472 db.session.query(ApiToolProvider).filter(ApiToolProvider.tenant_id == tenant_id).all()473 )474 475 api_provider_controllers = [476 {"provider": provider, "controller": ToolTransformService.api_provider_to_controller(provider)}477 for provider in db_api_providers478 ]479 480 # get labels481 labels = ToolLabelManager.get_tools_labels([x["controller"] for x in api_provider_controllers])482 483 for api_provider_controller in api_provider_controllers:484 user_provider = ToolTransformService.api_provider_to_user_provider(485 provider_controller=api_provider_controller["controller"],486 db_provider=api_provider_controller["provider"],487 decrypt_credentials=False,488 labels=labels.get(api_provider_controller["controller"].provider_id, []),489 )490 result_providers[f"api_provider.{user_provider.name}"] = user_provider491 492 if "workflow" in filters:493 # get workflow providers494 workflow_providers: list[WorkflowToolProvider] = (495 db.session.query(WorkflowToolProvider).filter(WorkflowToolProvider.tenant_id == tenant_id).all()496 )497 498 workflow_provider_controllers = []499 for provider in workflow_providers:500 try:501 workflow_provider_controllers.append(502 ToolTransformService.workflow_provider_to_controller(db_provider=provider)503 )504 except Exception as e:505 # app has been deleted506 pass507 508 labels = ToolLabelManager.get_tools_labels(workflow_provider_controllers)509 510 for provider_controller in workflow_provider_controllers:511 user_provider = ToolTransformService.workflow_provider_to_user_provider(512 provider_controller=provider_controller,513 labels=labels.get(provider_controller.provider_id, []),514 )515 result_providers[f"workflow_provider.{user_provider.name}"] = user_provider516 517 return BuiltinToolProviderSort.sort(list(result_providers.values()))518 519 @classmethod520 def get_api_provider_controller(521 cls, tenant_id: str, provider_id: str522 ) -> tuple[ApiToolProviderController, dict[str, Any]]:523 """524 get the api provider525 526 :param provider_name: the name of the provider527 528 :return: the provider controller, the credentials529 """530 provider: ApiToolProvider = (531 db.session.query(ApiToolProvider)532 .filter(533 ApiToolProvider.id == provider_id,534 ApiToolProvider.tenant_id == tenant_id,535 )536 .first()537 )538 539 if provider is None:540 raise ToolProviderNotFoundError(f"api provider {provider_id} not found")541 542 controller = ApiToolProviderController.from_db(543 provider,544 ApiProviderAuthType.API_KEY if provider.credentials["auth_type"] == "api_key" else ApiProviderAuthType.NONE,545 )546 controller.load_bundled_tools(provider.tools)547 548 return controller, provider.credentials549 550 @classmethod551 def user_get_api_provider(cls, provider: str, tenant_id: str) -> dict:552 """553 get api provider554 """555 """556 get tool provider557 """558 provider: ApiToolProvider = (559 db.session.query(ApiToolProvider)560 .filter(561 ApiToolProvider.tenant_id == tenant_id,562 ApiToolProvider.name == provider,563 )564 .first()565 )566 567 if provider is None:568 raise ValueError(f"you have not added provider {provider}")569 570 try:571 credentials = json.loads(provider.credentials_str) or {}572 except:573 credentials = {}574 575 # package tool provider controller576 controller = ApiToolProviderController.from_db(577 provider, ApiProviderAuthType.API_KEY if credentials["auth_type"] == "api_key" else ApiProviderAuthType.NONE578 )579 # init tool configuration580 tool_configuration = ToolConfigurationManager(tenant_id=tenant_id, provider_controller=controller)581 582 decrypted_credentials = tool_configuration.decrypt_tool_credentials(credentials)583 masked_credentials = tool_configuration.mask_tool_credentials(decrypted_credentials)584 585 try:586 icon = json.loads(provider.icon)587 except:588 icon = {"background": "#252525", "content": "\ud83d\ude01"}589 590 # add tool labels591 labels = ToolLabelManager.get_tool_labels(controller)592 593 return jsonable_encoder(594 {595 "schema_type": provider.schema_type,596 "schema": provider.schema,597 "tools": provider.tools,598 "icon": icon,599 "description": provider.description,600 "credentials": masked_credentials,601 "privacy_policy": provider.privacy_policy,602 "custom_disclaimer": provider.custom_disclaimer,603 "labels": labels,604 }605 )606 607 @classmethod608 def get_tool_icon(cls, tenant_id: str, provider_type: str, provider_id: str) -> Union[str, dict]:609 """610 get the tool icon611 612 :param tenant_id: the id of the tenant613 :param provider_type: the type of the provider614 :param provider_id: the id of the provider615 :return:616 """617 provider_type = provider_type618 provider_id = provider_id619 if provider_type == "builtin":620 return (621 dify_config.CONSOLE_API_URL622 + "/console/api/workspaces/current/tool-provider/builtin/"623 + provider_id624 + "/icon"625 )626 elif provider_type == "api":627 try:628 provider: ApiToolProvider = (629 db.session.query(ApiToolProvider)630 .filter(ApiToolProvider.tenant_id == tenant_id, ApiToolProvider.id == provider_id)631 .first()632 )633 return json.loads(provider.icon)634 except:635 return {"background": "#252525", "content": "\ud83d\ude01"}636 elif provider_type == "workflow":637 provider: WorkflowToolProvider = (638 db.session.query(WorkflowToolProvider)639 .filter(WorkflowToolProvider.tenant_id == tenant_id, WorkflowToolProvider.id == provider_id)640 .first()641 )642 if provider is None:643 raise ToolProviderNotFoundError(f"workflow provider {provider_id} not found")644 645 return json.loads(provider.icon)646 else:647 raise ValueError(f"provider type {provider_type} not found")648 649 650# preload builtin tool providers651Thread(target=ToolManager.load_builtin_providers_cache, name="pre_load_builtin_providers_cache", daemon=True).start()652 