dify/api/core/tools/__base/tool.py

197 lines
6.3 KiB
Python
Raw Normal View History

2024-02-01 18:11:57 +08:00
from abc import ABC, abstractmethod
2024-08-29 14:09:47 +08:00
from collections.abc import Generator
from copy import deepcopy
2024-09-20 23:48:48 +08:00
from typing import TYPE_CHECKING, Any, Optional
2024-09-20 23:48:48 +08:00
from core.tools.__base.tool_runtime import ToolRuntime
from core.tools.entities.tool_entities import (
2024-09-20 23:48:48 +08:00
ToolEntity,
ToolInvokeMessage,
ToolParameter,
ToolProviderType,
)
from core.tools.utils.tool_parameter_converter import ToolParameterConverter
2024-02-01 18:11:57 +08:00
if TYPE_CHECKING:
from core.file.file_obj import FileVar
2024-09-20 23:48:48 +08:00
class Tool(ABC):
"""
The base class of a tool
"""
2024-09-20 23:48:48 +08:00
entity: ToolEntity
runtime: ToolRuntime
2024-09-20 23:48:48 +08:00
def __init__(self, entity: ToolEntity, runtime: ToolRuntime) -> None:
self.entity = entity
self.runtime = runtime
2024-09-20 23:48:48 +08:00
def fork_tool_runtime(self, runtime: ToolRuntime) -> "Tool":
"""
2024-07-29 16:40:04 +08:00
fork a new tool with meta data
2024-07-29 16:40:04 +08:00
:param meta: the meta data of a tool call processing, tenant_id is required
:return: the new tool
"""
return self.__class__(
2024-09-20 23:48:48 +08:00
entity=self.entity.model_copy(),
runtime=runtime,
)
2024-07-29 16:40:04 +08:00
@abstractmethod
def tool_provider_type(self) -> ToolProviderType:
"""
2024-07-29 16:40:04 +08:00
get the tool provider type
2024-07-29 16:40:04 +08:00
:return: the tool provider type
"""
2024-07-29 16:40:04 +08:00
2024-08-29 14:06:10 +08:00
def invoke(self, user_id: str, tool_parameters: dict[str, Any]) -> Generator[ToolInvokeMessage]:
if self.runtime and self.runtime.runtime_parameters:
2024-01-31 11:58:07 +08:00
tool_parameters.update(self.runtime.runtime_parameters)
# try parse tool parameters into the correct type
tool_parameters = self._transform_tool_parameters_type(tool_parameters)
result = self._invoke(
user_id=user_id,
tool_parameters=tool_parameters,
)
2024-08-30 18:11:38 +08:00
if isinstance(result, ToolInvokeMessage):
2024-08-30 18:11:38 +08:00
def single_generator():
yield result
2024-09-14 02:47:01 +08:00
2024-08-30 18:11:38 +08:00
return single_generator()
elif isinstance(result, list):
2024-09-14 02:47:01 +08:00
2024-08-30 18:11:38 +08:00
def generator():
yield from result
2024-09-14 02:47:01 +08:00
2024-08-30 18:11:38 +08:00
return generator()
else:
return result
2024-08-29 14:06:10 +08:00
def _transform_tool_parameters_type(self, tool_parameters: dict[str, Any]) -> dict[str, Any]:
"""
Transform tool parameters type
"""
# Temp fix for the issue that the tool parameters will be converted to empty while validating the credentials
result = deepcopy(tool_parameters)
2024-09-20 23:48:48 +08:00
for parameter in self.entity.parameters:
if parameter.name in tool_parameters:
2024-07-29 16:40:04 +08:00
result[parameter.name] = ToolParameterConverter.cast_parameter_by_type(
tool_parameters[parameter.name], parameter.type
)
return result
@abstractmethod
def _invoke(
self, user_id: str, tool_parameters: dict[str, Any]
2024-09-14 02:47:01 +08:00
) -> ToolInvokeMessage | list[ToolInvokeMessage] | Generator[ToolInvokeMessage, None, None]:
pass
2024-07-29 16:40:04 +08:00
def get_runtime_parameters(self) -> list[ToolParameter]:
"""
2024-07-29 16:40:04 +08:00
get the runtime parameters
2024-07-29 16:40:04 +08:00
interface for developer to dynamic change the parameters of a tool depends on the variables pool
2024-07-29 16:40:04 +08:00
:return: the runtime parameters
"""
2024-09-20 23:48:48 +08:00
return self.entity.parameters
2024-07-29 16:40:04 +08:00
2024-03-08 20:31:13 +08:00
def get_all_runtime_parameters(self) -> list[ToolParameter]:
"""
2024-07-29 16:40:04 +08:00
get all runtime parameters
2024-03-08 20:31:13 +08:00
2024-07-29 16:40:04 +08:00
:return: all runtime parameters
2024-03-08 20:31:13 +08:00
"""
2024-09-20 23:48:48 +08:00
parameters = self.entity.parameters
2024-03-08 20:31:13 +08:00
parameters = parameters.copy()
user_parameters = self.get_runtime_parameters() or []
user_parameters = user_parameters.copy()
# override parameters
for parameter in user_parameters:
# check if parameter in tool parameters
for tool_parameter in parameters:
if tool_parameter.name == parameter.name:
2024-09-20 23:48:48 +08:00
# override parameter
tool_parameter.type = parameter.type
tool_parameter.form = parameter.form
tool_parameter.required = parameter.required
tool_parameter.default = parameter.default
tool_parameter.options = parameter.options
tool_parameter.llm_description = parameter.llm_description
2024-03-08 20:31:13 +08:00
break
else:
# add new parameter
parameters.append(parameter)
return parameters
2024-07-29 16:40:04 +08:00
def create_image_message(self, image: str, save_as: str = "") -> ToolInvokeMessage:
"""
2024-07-29 16:40:04 +08:00
create an image message
2024-07-29 16:40:04 +08:00
:param image: the url of the image
:return: the image message
"""
2024-09-14 02:47:01 +08:00
return ToolInvokeMessage(
type=ToolInvokeMessage.MessageType.IMAGE, message=ToolInvokeMessage.TextMessage(text=image), save_as=save_as
)
2024-07-29 16:40:04 +08:00
def create_file_var_message(self, file_var: "FileVar") -> ToolInvokeMessage:
return ToolInvokeMessage(
2024-09-14 02:47:01 +08:00
type=ToolInvokeMessage.MessageType.FILE_VAR, message=None, meta={"file_var": file_var}, save_as=""
)
2024-07-29 16:40:04 +08:00
def create_link_message(self, link: str, save_as: str = "") -> ToolInvokeMessage:
"""
2024-07-29 16:40:04 +08:00
create a link message
2024-07-29 16:40:04 +08:00
:param link: the url of the link
:return: the link message
"""
2024-09-14 02:47:01 +08:00
return ToolInvokeMessage(
type=ToolInvokeMessage.MessageType.LINK, message=ToolInvokeMessage.TextMessage(text=link), save_as=save_as
)
2024-07-29 16:40:04 +08:00
def create_text_message(self, text: str, save_as: str = "") -> ToolInvokeMessage:
"""
2024-07-29 16:40:04 +08:00
create a text message
2024-07-29 16:40:04 +08:00
:param text: the text
:return: the text message
"""
return ToolInvokeMessage(
2024-09-14 02:47:01 +08:00
type=ToolInvokeMessage.MessageType.TEXT, message=ToolInvokeMessage.TextMessage(text=text), save_as=save_as
)
2024-07-29 16:40:04 +08:00
2024-09-14 02:47:01 +08:00
def create_blob_message(self, blob: bytes, meta: Optional[dict] = None, save_as: str = "") -> ToolInvokeMessage:
"""
2024-07-29 16:40:04 +08:00
create a blob message
2024-07-29 16:40:04 +08:00
:param blob: the blob
:return: the blob message
"""
2024-08-29 14:06:10 +08:00
return ToolInvokeMessage(
2024-09-14 02:47:01 +08:00
type=ToolInvokeMessage.MessageType.BLOB,
message=ToolInvokeMessage.BlobMessage(blob=blob),
meta=meta,
save_as=save_as,
2024-08-29 14:06:10 +08:00
)
def create_json_message(self, object: dict) -> ToolInvokeMessage:
"""
2024-07-29 16:40:04 +08:00
create a json message
"""
2024-08-29 14:06:10 +08:00
return ToolInvokeMessage(
2024-09-14 02:47:01 +08:00
type=ToolInvokeMessage.MessageType.JSON, message=ToolInvokeMessage.JsonMessage(json_object=object)
2024-08-29 14:06:10 +08:00
)