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

224 lines
6.8 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-10-21 20:01:49 +08:00
if TYPE_CHECKING:
from models.model import File
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,
)
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
def invoke(
self,
user_id: str,
tool_parameters: dict[str, Any],
conversation_id: Optional[str] = None,
app_id: Optional[str] = None,
message_id: Optional[str] = None,
) -> Generator[ToolInvokeMessage]:
2024-08-29 14:06:10 +08:00
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,
conversation_id=conversation_id,
app_id=app_id,
message_id=message_id,
)
2024-08-30 18:11:38 +08:00
if isinstance(result, ToolInvokeMessage):
2025-01-09 16:53:30 +08:00
def single_generator() -> Generator[ToolInvokeMessage, None, None]:
2024-08-30 18:11:38 +08:00
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
2025-01-09 16:53:30 +08:00
def generator() -> Generator[ToolInvokeMessage, None, None]:
2024-08-30 18:11:38 +08:00
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-10-21 20:01:49 +08:00
for parameter in self.entity.parameters or []:
if parameter.name in tool_parameters:
2024-10-21 20:01:49 +08:00
result[parameter.name] = parameter.type.cast_value(tool_parameters[parameter.name])
return result
@abstractmethod
def _invoke(
self,
user_id: str,
tool_parameters: dict[str, Any],
conversation_id: Optional[str] = None,
app_id: Optional[str] = None,
message_id: Optional[str] = None,
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,
conversation_id: Optional[str] = None,
app_id: Optional[str] = None,
message_id: Optional[str] = None,
) -> 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
def get_merged_runtime_parameters(
self,
conversation_id: Optional[str] = None,
app_id: Optional[str] = None,
message_id: Optional[str] = None,
) -> list[ToolParameter]:
2024-03-08 20:31:13 +08:00
"""
2024-09-23 18:06:16 +08:00
get merged runtime parameters
2024-03-08 20:31:13 +08:00
2024-09-23 18:06:16 +08:00
:return: merged 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,
) -> 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)
2024-09-14 02:47:01 +08:00
)
2024-07-29 16:40:04 +08:00
2024-10-21 20:01:49 +08:00
def create_file_message(self, file: "File") -> ToolInvokeMessage:
return ToolInvokeMessage(
2024-10-21 20:01:49 +08:00
type=ToolInvokeMessage.MessageType.FILE,
message=ToolInvokeMessage.FileMessage(),
meta={"file": file},
)
2024-07-29 16:40:04 +08:00
def create_link_message(self, link: 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)
2024-09-14 02:47:01 +08:00
)
2024-07-29 16:40:04 +08:00
def create_text_message(self, text: 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(
type=ToolInvokeMessage.MessageType.TEXT,
message=ToolInvokeMessage.TextMessage(text=text),
)
2024-07-29 16:40:04 +08:00
def create_blob_message(self, blob: bytes, meta: Optional[dict] = None) -> 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,
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
)