mirror of
https://github.com/langgenius/dify.git
synced 2026-07-27 23:18:33 +08:00
feat: support tool multi-select input (#39346)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com> Co-authored-by: WH-2099 <wh2099@pm.me>
This commit is contained in:
parent
4da4fa72cd
commit
d5788cc019
@ -58,7 +58,6 @@ class Tool(ABC):
|
||||
if self.runtime and self.runtime.runtime_parameters:
|
||||
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(
|
||||
@ -87,14 +86,14 @@ class Tool(ABC):
|
||||
return result
|
||||
|
||||
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
|
||||
"""Transform declared tool parameter values without resolving runtime schemas."""
|
||||
result = deepcopy(tool_parameters)
|
||||
for parameter in self.entity.parameters or []:
|
||||
if parameter.name in tool_parameters:
|
||||
result[parameter.name] = parameter.type.cast_value(tool_parameters[parameter.name])
|
||||
if parameter.multiple:
|
||||
result[parameter.name] = parameter.init_frontend_parameter(result.get(parameter.name))
|
||||
else:
|
||||
result[parameter.name] = parameter.type.cast_value(tool_parameters[parameter.name])
|
||||
|
||||
return result
|
||||
|
||||
@ -196,17 +195,31 @@ class Tool(ABC):
|
||||
}:
|
||||
continue
|
||||
|
||||
parameter_schema: dict[str, Any] = (
|
||||
{
|
||||
"type": parameter.type.as_normal_type(),
|
||||
"description": parameter.llm_description or "",
|
||||
}
|
||||
if parameter.input_schema is None
|
||||
else deepcopy(parameter.input_schema)
|
||||
)
|
||||
is_multiple_select = parameter.multiple and parameter.type in {
|
||||
ToolParameter.ToolParameterType.SELECT,
|
||||
ToolParameter.ToolParameterType.DYNAMIC_SELECT,
|
||||
}
|
||||
if is_multiple_select:
|
||||
item_schema: dict[str, Any] = {"type": "string"}
|
||||
if parameter.type == ToolParameter.ToolParameterType.SELECT and parameter.options:
|
||||
item_schema["enum"] = [option.value for option in parameter.options]
|
||||
parameter_schema: dict[str, Any] = {"type": "array", "items": item_schema}
|
||||
else:
|
||||
parameter_schema = (
|
||||
{
|
||||
"type": parameter.type.as_normal_type(),
|
||||
"description": parameter.llm_description or "",
|
||||
}
|
||||
if parameter.input_schema is None
|
||||
else deepcopy(parameter.input_schema)
|
||||
)
|
||||
parameter_schema.setdefault("description", parameter.llm_description or "")
|
||||
|
||||
if parameter.type == ToolParameter.ToolParameterType.SELECT and parameter.options:
|
||||
if (
|
||||
not is_multiple_select
|
||||
and parameter.type == ToolParameter.ToolParameterType.SELECT
|
||||
and parameter.options
|
||||
):
|
||||
parameter_schema["enum"] = [option.value for option in parameter.options]
|
||||
|
||||
schema["properties"][parameter.name] = parameter_schema
|
||||
|
||||
@ -292,9 +292,7 @@ class ToolInvokeMessageBinary(BaseModel):
|
||||
|
||||
|
||||
class ToolParameter(PluginParameter):
|
||||
"""
|
||||
Overrides type
|
||||
"""
|
||||
"""Tool-specific parameter declaration and invocation-value normalization."""
|
||||
|
||||
class ToolParameterType(StrEnum):
|
||||
"""
|
||||
@ -333,12 +331,28 @@ class ToolParameter(PluginParameter):
|
||||
LLM = auto() # will be set by LLM
|
||||
|
||||
type: ToolParameterType = Field(..., description="The type of the parameter")
|
||||
multiple: bool = Field(
|
||||
default=False,
|
||||
description="Whether the parameter is multiple select, only valid for select or dynamic-select type",
|
||||
)
|
||||
human_description: I18nObject | None = Field(default=None, description="The description presented to the user")
|
||||
form: ToolParameterForm = Field(..., description="The form of the parameter, schema/form/llm")
|
||||
llm_description: str | None = None
|
||||
# MCP object and array type parameters use this field to store the schema
|
||||
input_schema: dict[str, Any] | None = None
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_multiple(self) -> ToolParameter:
|
||||
supports_multiple = self.type in {
|
||||
self.ToolParameterType.SELECT,
|
||||
self.ToolParameterType.DYNAMIC_SELECT,
|
||||
}
|
||||
if self.multiple and not supports_multiple:
|
||||
raise ValueError("multiple is only valid for select and dynamic-select parameters")
|
||||
if supports_multiple and self.default is not None and (isinstance(self.default, list) != self.multiple):
|
||||
raise ValueError("default must be a list exactly when multiple is true")
|
||||
return self
|
||||
|
||||
@classmethod
|
||||
def get_simple_instance(
|
||||
cls,
|
||||
@ -378,8 +392,25 @@ class ToolParameter(PluginParameter):
|
||||
options=option_objs,
|
||||
)
|
||||
|
||||
def init_frontend_parameter(self, value: Any):
|
||||
return init_frontend_parameter(self, self.type, value)
|
||||
def init_frontend_parameter(self, value: Any) -> Any:
|
||||
"""Normalize a value against this tool parameter's full declaration."""
|
||||
if not self.multiple:
|
||||
return init_frontend_parameter(self, self.type, value)
|
||||
|
||||
parameter_value = self.default if value is None else value
|
||||
if parameter_value is None:
|
||||
parameter_value = []
|
||||
if not isinstance(parameter_value, list):
|
||||
raise ValueError(f"tool parameter {self.name} must be a list when multiple is true")
|
||||
if not all(isinstance(item, str) for item in parameter_value):
|
||||
raise ValueError(f"tool parameter {self.name} must contain only strings")
|
||||
if self.required and not parameter_value:
|
||||
raise ValueError(f"tool parameter {self.name} not found in tool config")
|
||||
if self.type == self.ToolParameterType.SELECT:
|
||||
options = [option.value for option in self.options]
|
||||
if any(item not in options for item in parameter_value):
|
||||
raise ValueError(f"tool parameter {self.name} value {parameter_value} not in options {options}")
|
||||
return parameter_value
|
||||
|
||||
|
||||
class ToolProviderIdentity(BaseModel):
|
||||
|
||||
@ -22330,7 +22330,7 @@ Tool label
|
||||
|
||||
#### ToolParameter
|
||||
|
||||
Overrides type
|
||||
Tool-specific parameter declaration and invocation-value normalization.
|
||||
|
||||
| Name | Type | Description | Required |
|
||||
| ---- | ---- | ----------- | -------- |
|
||||
@ -22343,6 +22343,7 @@ Overrides type
|
||||
| llm_description | string | | No |
|
||||
| max | number<br>integer | | No |
|
||||
| min | number<br>integer | | No |
|
||||
| multiple | boolean | Whether the parameter is multiple select, only valid for select or dynamic-select type | No |
|
||||
| name | string | The name of the parameter | Yes |
|
||||
| options | [ [PluginParameterOption](#pluginparameteroption) ] | | No |
|
||||
| placeholder | [I18nObject](#i18nobject) | The placeholder presented to the user | No |
|
||||
|
||||
@ -125,6 +125,9 @@ class TestPluginParameterEntities:
|
||||
parameter = PluginParameter(name="p", label=self._label(), options="invalid") # type: ignore[arg-type]
|
||||
assert parameter.options == []
|
||||
|
||||
def test_plugin_parameter_excludes_tool_specific_multiple_declaration(self):
|
||||
assert "multiple" not in PluginParameter.model_fields
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("parameter_type", "expected"),
|
||||
[
|
||||
|
||||
@ -5,6 +5,8 @@ from dataclasses import dataclass
|
||||
from typing import Any, cast
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from core.app.entities.app_invoke_entities import InvokeFrom
|
||||
from core.tools.__base.tool import Tool
|
||||
from core.tools.__base.tool_runtime import ToolRuntime
|
||||
@ -33,6 +35,7 @@ class DummyParameter:
|
||||
options: list[Any] | None = None
|
||||
llm_description: str | None = None
|
||||
input_schema: dict[str, Any] | None = None
|
||||
multiple: bool = False
|
||||
|
||||
|
||||
class DummyTool(Tool):
|
||||
@ -129,6 +132,26 @@ def test_invoke_supports_single_message_and_parameter_casting():
|
||||
}
|
||||
|
||||
|
||||
def test_invoke_preserves_multiple_select_values():
|
||||
tool = _build_tool()
|
||||
parameter = ToolParameter.get_simple_instance(
|
||||
name="choice",
|
||||
llm_description="Choice",
|
||||
typ=ToolParameter.ToolParameterType.SELECT,
|
||||
required=True,
|
||||
options=["a", "b"],
|
||||
)
|
||||
parameter.multiple = True
|
||||
tool.entity.parameters = [parameter]
|
||||
|
||||
list(tool.invoke(session=MagicMock(), user_id="user-1", tool_parameters={"choice": ["a", "b"]}))
|
||||
|
||||
assert tool.last_invocation is not None
|
||||
assert tool.last_invocation["tool_parameters"] == {"choice": ["a", "b"]}
|
||||
with pytest.raises(ValueError, match="must be a list"):
|
||||
tool.invoke(session=MagicMock(), user_id="user-1", tool_parameters={"choice": "a"})
|
||||
|
||||
|
||||
def test_invoke_supports_list_and_generator_results():
|
||||
tool = _build_tool()
|
||||
tool.result = [tool.create_text_message("a"), tool.create_text_message("b")]
|
||||
@ -214,6 +237,21 @@ def test_get_llm_parameters_json_schema_uses_effective_runtime_parameters():
|
||||
required=False,
|
||||
options=["global", "cn"],
|
||||
)
|
||||
regions_parameter = ToolParameter.get_simple_instance(
|
||||
name="regions",
|
||||
llm_description="Search regions",
|
||||
typ=ToolParameter.ToolParameterType.SELECT,
|
||||
required=False,
|
||||
options=["global", "cn"],
|
||||
)
|
||||
regions_parameter.multiple = True
|
||||
tags_parameter = ToolParameter.get_simple_instance(
|
||||
name="tags",
|
||||
llm_description="Search tags",
|
||||
typ=ToolParameter.ToolParameterType.DYNAMIC_SELECT,
|
||||
required=False,
|
||||
)
|
||||
tags_parameter.multiple = True
|
||||
hidden_parameter = ToolParameter.get_simple_instance(
|
||||
name="api_key",
|
||||
llm_description="Hidden api key",
|
||||
@ -241,7 +279,15 @@ def test_get_llm_parameters_json_schema_uses_effective_runtime_parameters():
|
||||
"properties": {"nested": {"type": "string"}},
|
||||
},
|
||||
)
|
||||
tool.entity.parameters = [query_parameter, region_parameter, hidden_parameter, file_parameter, payload_parameter]
|
||||
tool.entity.parameters = [
|
||||
query_parameter,
|
||||
region_parameter,
|
||||
regions_parameter,
|
||||
tags_parameter,
|
||||
hidden_parameter,
|
||||
file_parameter,
|
||||
payload_parameter,
|
||||
]
|
||||
|
||||
query_override = ToolParameter.get_simple_instance(
|
||||
name="query",
|
||||
@ -262,6 +308,16 @@ def test_get_llm_parameters_json_schema_uses_effective_runtime_parameters():
|
||||
"description": "Search region",
|
||||
"enum": ["global", "cn"],
|
||||
},
|
||||
"regions": {
|
||||
"type": "array",
|
||||
"items": {"type": "string", "enum": ["global", "cn"]},
|
||||
"description": "Search regions",
|
||||
},
|
||||
"tags": {
|
||||
"type": "array",
|
||||
"items": {"type": "string"},
|
||||
"description": "Search tags",
|
||||
},
|
||||
"payload": {
|
||||
"type": "object",
|
||||
"properties": {"nested": {"type": "string"}},
|
||||
|
||||
@ -1,5 +1,8 @@
|
||||
import pytest
|
||||
from pydantic import ValidationError
|
||||
|
||||
from core.tools.entities.common_entities import I18nObject
|
||||
from core.tools.entities.tool_entities import ToolEntity, ToolIdentity, ToolInvokeMessage
|
||||
from core.tools.entities.tool_entities import ToolEntity, ToolIdentity, ToolInvokeMessage, ToolParameter
|
||||
|
||||
|
||||
def _make_identity() -> ToolIdentity:
|
||||
@ -11,6 +14,65 @@ def _make_identity() -> ToolIdentity:
|
||||
)
|
||||
|
||||
|
||||
def _make_select_parameter(**updates: object) -> ToolParameter:
|
||||
data = ToolParameter.get_simple_instance(
|
||||
name="choice",
|
||||
llm_description="Choice",
|
||||
typ=ToolParameter.ToolParameterType.SELECT,
|
||||
required=False,
|
||||
options=["a", "b"],
|
||||
).model_dump()
|
||||
data.update(updates)
|
||||
return ToolParameter.model_validate(data)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("updates", "message"),
|
||||
[
|
||||
({"type": ToolParameter.ToolParameterType.STRING, "multiple": True}, "multiple is only valid"),
|
||||
({"multiple": True, "default": "a"}, "default must be a list"),
|
||||
({"default": ["a"]}, "default must be a list"),
|
||||
],
|
||||
)
|
||||
def test_tool_parameter_rejects_invalid_multiple_declarations(updates: dict[str, object], message: str):
|
||||
with pytest.raises(ValidationError, match=message):
|
||||
_make_select_parameter(**updates)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"parameter_type",
|
||||
[ToolParameter.ToolParameterType.SELECT, ToolParameter.ToolParameterType.DYNAMIC_SELECT],
|
||||
)
|
||||
def test_tool_parameter_accepts_multiple_select_declarations(parameter_type: ToolParameter.ToolParameterType):
|
||||
parameter = _make_select_parameter(type=parameter_type, multiple=True, default=["a"])
|
||||
|
||||
assert parameter.multiple is True
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("value", "message"),
|
||||
[
|
||||
("a", "must be a list"),
|
||||
(["a", 1], "only strings"),
|
||||
(["missing"], "not in options"),
|
||||
([], "not found in tool config"),
|
||||
],
|
||||
)
|
||||
def test_multiple_select_normalization_rejects_invalid_values(value: object, message: str):
|
||||
parameter = _make_select_parameter(multiple=True, required=True)
|
||||
|
||||
with pytest.raises(ValueError, match=message):
|
||||
parameter.init_frontend_parameter(value)
|
||||
|
||||
|
||||
def test_multiple_select_normalization_preserves_explicit_empty_list():
|
||||
parameter = _make_select_parameter(multiple=True, default=["a"])
|
||||
|
||||
assert parameter.init_frontend_parameter(None) == ["a"]
|
||||
assert parameter.init_frontend_parameter([]) == []
|
||||
assert parameter.init_frontend_parameter(["a", "b"]) == ["a", "b"]
|
||||
|
||||
|
||||
def test_log_message_metadata_none_defaults_to_empty_dict():
|
||||
log_message = ToolInvokeMessage.LogMessage(
|
||||
id="log-1",
|
||||
|
||||
@ -9,30 +9,29 @@ from services.tools.tools_transform_service import ToolTransformService
|
||||
MODULE = "services.tools.tools_transform_service"
|
||||
|
||||
|
||||
def _parameter(
|
||||
name: str,
|
||||
label: str,
|
||||
form: ToolParameter.ToolParameterForm = ToolParameter.ToolParameterForm.FORM,
|
||||
) -> ToolParameter:
|
||||
return ToolParameter(
|
||||
name=name,
|
||||
label=I18nObject(en_US=label),
|
||||
human_description=I18nObject(en_US=label),
|
||||
type=ToolParameter.ToolParameterType.STRING,
|
||||
form=form,
|
||||
)
|
||||
|
||||
|
||||
class TestToolTransformService:
|
||||
"""Test cases for ToolTransformService.convert_tool_entity_to_api_entity method"""
|
||||
|
||||
def test_convert_tool_with_parameter_override(self):
|
||||
"""Test that runtime parameters correctly override base parameters"""
|
||||
# Create mock base parameters
|
||||
base_param1 = Mock(spec=ToolParameter)
|
||||
base_param1.name = "param1"
|
||||
base_param1.form = ToolParameter.ToolParameterForm.FORM
|
||||
base_param1.type = "string"
|
||||
base_param1.label = "Base Param 1"
|
||||
base_param1 = _parameter("param1", "Base Param 1")
|
||||
base_param2 = _parameter("param2", "Base Param 2")
|
||||
|
||||
base_param2 = Mock(spec=ToolParameter)
|
||||
base_param2.name = "param2"
|
||||
base_param2.form = ToolParameter.ToolParameterForm.FORM
|
||||
base_param2.type = "string"
|
||||
base_param2.label = "Base Param 2"
|
||||
|
||||
# Create mock runtime parameters that override base parameters
|
||||
runtime_param1 = Mock(spec=ToolParameter)
|
||||
runtime_param1.name = "param1"
|
||||
runtime_param1.form = ToolParameter.ToolParameterForm.FORM
|
||||
runtime_param1.type = "string"
|
||||
runtime_param1.label = "Runtime Param 1" # Different label to verify override
|
||||
runtime_param1 = _parameter("param1", "Runtime Param 1")
|
||||
|
||||
# Create mock tool
|
||||
mock_tool = Mock(spec=Tool)
|
||||
@ -63,34 +62,19 @@ class TestToolTransformService:
|
||||
# Find the overridden parameter
|
||||
overridden_param = next((p for p in result.parameters if p.name == "param1"), None)
|
||||
assert overridden_param is not None
|
||||
assert overridden_param.label == "Runtime Param 1" # Should be runtime version
|
||||
assert overridden_param.label.en_US == "Runtime Param 1" # Should be runtime version
|
||||
|
||||
# Find the non-overridden parameter
|
||||
original_param = next((p for p in result.parameters if p.name == "param2"), None)
|
||||
assert original_param is not None
|
||||
assert original_param.label == "Base Param 2" # Should be base version
|
||||
assert original_param.label.en_US == "Base Param 2" # Should be base version
|
||||
|
||||
def test_convert_tool_with_additional_runtime_parameters(self):
|
||||
"""Test that additional runtime parameters are added to the final list"""
|
||||
# Create mock base parameters
|
||||
base_param1 = Mock(spec=ToolParameter)
|
||||
base_param1.name = "param1"
|
||||
base_param1.form = ToolParameter.ToolParameterForm.FORM
|
||||
base_param1.type = "string"
|
||||
base_param1.label = "Base Param 1"
|
||||
base_param1 = _parameter("param1", "Base Param 1")
|
||||
|
||||
# Create mock runtime parameters - one that overrides and one that's new
|
||||
runtime_param1 = Mock(spec=ToolParameter)
|
||||
runtime_param1.name = "param1"
|
||||
runtime_param1.form = ToolParameter.ToolParameterForm.FORM
|
||||
runtime_param1.type = "string"
|
||||
runtime_param1.label = "Runtime Param 1"
|
||||
|
||||
runtime_param2 = Mock(spec=ToolParameter)
|
||||
runtime_param2.name = "runtime_only"
|
||||
runtime_param2.form = ToolParameter.ToolParameterForm.FORM
|
||||
runtime_param2.type = "string"
|
||||
runtime_param2.label = "Runtime Only Param"
|
||||
runtime_param1 = _parameter("param1", "Runtime Param 1")
|
||||
runtime_param2 = _parameter("runtime_only", "Runtime Only Param")
|
||||
|
||||
# Create mock tool
|
||||
mock_tool = Mock(spec=Tool)
|
||||
@ -124,34 +108,19 @@ class TestToolTransformService:
|
||||
# Verify the overridden parameter has runtime version
|
||||
overridden_param = next((p for p in result.parameters if p.name == "param1"), None)
|
||||
assert overridden_param is not None
|
||||
assert overridden_param.label == "Runtime Param 1"
|
||||
assert overridden_param.label.en_US == "Runtime Param 1"
|
||||
|
||||
# Verify the new runtime parameter is included
|
||||
new_param = next((p for p in result.parameters if p.name == "runtime_only"), None)
|
||||
assert new_param is not None
|
||||
assert new_param.label == "Runtime Only Param"
|
||||
assert new_param.label.en_US == "Runtime Only Param"
|
||||
|
||||
def test_convert_tool_with_non_form_runtime_parameters(self):
|
||||
"""Test that non-FORM runtime parameters are not added as new parameters"""
|
||||
# Create mock base parameters
|
||||
base_param1 = Mock(spec=ToolParameter)
|
||||
base_param1.name = "param1"
|
||||
base_param1.form = ToolParameter.ToolParameterForm.FORM
|
||||
base_param1.type = "string"
|
||||
base_param1.label = "Base Param 1"
|
||||
base_param1 = _parameter("param1", "Base Param 1")
|
||||
|
||||
# Create mock runtime parameters with different forms
|
||||
runtime_param1 = Mock(spec=ToolParameter)
|
||||
runtime_param1.name = "param1"
|
||||
runtime_param1.form = ToolParameter.ToolParameterForm.FORM
|
||||
runtime_param1.type = "string"
|
||||
runtime_param1.label = "Runtime Param 1"
|
||||
|
||||
runtime_param2 = Mock(spec=ToolParameter)
|
||||
runtime_param2.name = "llm_param"
|
||||
runtime_param2.form = ToolParameter.ToolParameterForm.LLM
|
||||
runtime_param2.type = "string"
|
||||
runtime_param2.label = "LLM Param"
|
||||
runtime_param1 = _parameter("param1", "Runtime Param 1")
|
||||
runtime_param2 = _parameter("llm_param", "LLM Param", ToolParameter.ToolParameterForm.LLM)
|
||||
|
||||
# Create mock tool
|
||||
mock_tool = Mock(spec=Tool)
|
||||
@ -236,38 +205,15 @@ class TestToolTransformService:
|
||||
|
||||
def test_convert_tool_parameter_order_preserved(self):
|
||||
"""Test that parameter order is preserved correctly"""
|
||||
# Create mock base parameters in specific order
|
||||
base_param1 = Mock(spec=ToolParameter)
|
||||
base_param1.name = "param1"
|
||||
base_param1.form = ToolParameter.ToolParameterForm.FORM
|
||||
base_param1.type = "string"
|
||||
base_param1.label = "Base Param 1"
|
||||
|
||||
base_param2 = Mock(spec=ToolParameter)
|
||||
base_param2.name = "param2"
|
||||
base_param2.form = ToolParameter.ToolParameterForm.FORM
|
||||
base_param2.type = "string"
|
||||
base_param2.label = "Base Param 2"
|
||||
|
||||
base_param3 = Mock(spec=ToolParameter)
|
||||
base_param3.name = "param3"
|
||||
base_param3.form = ToolParameter.ToolParameterForm.FORM
|
||||
base_param3.type = "string"
|
||||
base_param3.label = "Base Param 3"
|
||||
base_param1 = _parameter("param1", "Base Param 1")
|
||||
base_param2 = _parameter("param2", "Base Param 2")
|
||||
base_param3 = _parameter("param3", "Base Param 3")
|
||||
|
||||
# Create runtime parameter that overrides middle parameter
|
||||
runtime_param2 = Mock(spec=ToolParameter)
|
||||
runtime_param2.name = "param2"
|
||||
runtime_param2.form = ToolParameter.ToolParameterForm.FORM
|
||||
runtime_param2.type = "string"
|
||||
runtime_param2.label = "Runtime Param 2"
|
||||
runtime_param2 = _parameter("param2", "Runtime Param 2")
|
||||
|
||||
# Create new runtime parameter
|
||||
runtime_param4 = Mock(spec=ToolParameter)
|
||||
runtime_param4.name = "param4"
|
||||
runtime_param4.form = ToolParameter.ToolParameterForm.FORM
|
||||
runtime_param4.type = "string"
|
||||
runtime_param4.label = "Runtime Param 4"
|
||||
runtime_param4 = _parameter("param4", "Runtime Param 4")
|
||||
|
||||
# Create mock tool
|
||||
mock_tool = Mock(spec=Tool)
|
||||
@ -300,7 +246,7 @@ class TestToolTransformService:
|
||||
# Verify that param2 was overridden with runtime version
|
||||
param2 = result.parameters[1]
|
||||
assert param2.name == "param2"
|
||||
assert param2.label == "Runtime Param 2"
|
||||
assert param2.label.en_US == "Runtime Param 2"
|
||||
|
||||
|
||||
class TestWorkflowProviderToUserProvider:
|
||||
|
||||
@ -1999,6 +1999,7 @@ export type ToolParameter = {
|
||||
llm_description?: string | null
|
||||
max?: number | number | null
|
||||
min?: number | number | null
|
||||
multiple?: boolean
|
||||
name: string
|
||||
options?: Array<PluginParameterOption>
|
||||
placeholder?: I18nObject | null
|
||||
|
||||
@ -3057,7 +3057,7 @@ export const zPluginParameterAutoGenerate = z.object({
|
||||
/**
|
||||
* ToolParameter
|
||||
*
|
||||
* Overrides type
|
||||
* Tool-specific parameter declaration and invocation-value normalization.
|
||||
*/
|
||||
export const zToolParameter = z.object({
|
||||
auto_generate: zPluginParameterAutoGenerate.nullish(),
|
||||
@ -3078,6 +3078,7 @@ export const zToolParameter = z.object({
|
||||
llm_description: z.string().nullish(),
|
||||
max: z.union([z.number(), z.int()]).nullish(),
|
||||
min: z.union([z.number(), z.int()]).nullish(),
|
||||
multiple: z.boolean().optional().default(false),
|
||||
name: z.string(),
|
||||
options: z.array(zPluginParameterOption).optional(),
|
||||
placeholder: zI18nObject.nullish(),
|
||||
|
||||
@ -67,6 +67,17 @@ describe('form-input-item helpers', () => {
|
||||
expect(filesState.isFile).toBe(true)
|
||||
expect(filesState.showVariableSelector).toBe(true)
|
||||
expect(getTargetVarType(filesState)).toBe(VarType.arrayFile)
|
||||
|
||||
const multipleSelectState = getFormInputState(
|
||||
createSchema({ multiple: true, type: FormTypeEnum.select }),
|
||||
{ type: VarKindType.variable, value: ['node', 'formats'] },
|
||||
)
|
||||
expect(getTargetVarType(multipleSelectState)).toBe(VarType.arrayString)
|
||||
expect(getFilterVar(multipleSelectState)?.({ type: VarType.arrayString } as Var)).toBe(true)
|
||||
expect(getFilterVar(multipleSelectState)?.({ type: VarType.array } as Var)).toBe(true)
|
||||
expect(getFilterVar(multipleSelectState)?.({ type: VarType.arrayNumber } as Var)).toBe(false)
|
||||
expect(getFilterVar(multipleSelectState)?.({ type: VarType.arrayObject } as Var)).toBe(false)
|
||||
expect(getFilterVar(multipleSelectState)?.({ type: VarType.string } as Var)).toBe(false)
|
||||
})
|
||||
|
||||
it('should return filter functions and var kind types by schema mode', () => {
|
||||
|
||||
@ -140,6 +140,7 @@ export const getTargetVarType = (state: FormInputState) => {
|
||||
if (state.isString) return VarType.string
|
||||
if (state.isNumber) return VarType.number
|
||||
if (state.isFile) return state.isFiles ? VarType.arrayFile : VarType.file
|
||||
if (state.isSelect && state.isMultipleSelect) return VarType.arrayString
|
||||
if (state.isSelect) return VarType.string
|
||||
if (state.isBoolean) return VarType.boolean
|
||||
if (state.isObject) return VarType.object
|
||||
@ -149,6 +150,8 @@ export const getTargetVarType = (state: FormInputState) => {
|
||||
|
||||
export const getFilterVar = (state: FormInputState) => {
|
||||
if (state.isNumber) return (varPayload: Var) => varPayload.type === VarType.number
|
||||
if (state.isSelect && state.isMultipleSelect)
|
||||
return (varPayload: Var) => [VarType.array, VarType.arrayString].includes(varPayload.type)
|
||||
if (state.isString)
|
||||
return (varPayload: Var) =>
|
||||
[VarType.string, VarType.number, VarType.secret].includes(varPayload.type)
|
||||
|
||||
@ -0,0 +1,64 @@
|
||||
import type { ToolNodeType } from '../types'
|
||||
import { BlockEnum } from '@/app/components/workflow/types'
|
||||
import { withSelectorKey } from '@/test/i18n-mock'
|
||||
import nodeDefault from '../default'
|
||||
|
||||
const t = withSelectorKey((key: string) => key, 'workflow')
|
||||
|
||||
describe('tool default node validation', () => {
|
||||
it('should reject an empty required multi-select input', () => {
|
||||
const payload = {
|
||||
...nodeDefault.defaultValue,
|
||||
title: 'Tool',
|
||||
desc: '',
|
||||
type: BlockEnum.Tool,
|
||||
tool_parameters: {
|
||||
formats: {
|
||||
type: 'constant',
|
||||
value: [],
|
||||
},
|
||||
},
|
||||
} as ToolNodeType
|
||||
|
||||
const result = nodeDefault.checkValid(payload, t, {
|
||||
toolInputsSchema: [
|
||||
{ variable: 'formats', required: true, label: 'Formats', type: 'select', multiple: true },
|
||||
],
|
||||
toolSettingSchema: [],
|
||||
language: 'en_US',
|
||||
notAuthed: false,
|
||||
})
|
||||
|
||||
expect(result).toEqual({
|
||||
isValid: false,
|
||||
errorMessage: 'errorMsg.fieldRequired',
|
||||
})
|
||||
})
|
||||
|
||||
it('should accept an explicit empty array for a required array input', () => {
|
||||
const payload = {
|
||||
...nodeDefault.defaultValue,
|
||||
title: 'Tool',
|
||||
desc: '',
|
||||
type: BlockEnum.Tool,
|
||||
tool_parameters: {
|
||||
items: {
|
||||
type: 'constant',
|
||||
value: [],
|
||||
},
|
||||
},
|
||||
} as ToolNodeType
|
||||
|
||||
const result = nodeDefault.checkValid(payload, t, {
|
||||
toolInputsSchema: [{ variable: 'items', required: true, label: 'Items', type: 'array' }],
|
||||
toolSettingSchema: [],
|
||||
language: 'en_US',
|
||||
notAuthed: false,
|
||||
})
|
||||
|
||||
expect(result).toEqual({
|
||||
isValid: true,
|
||||
errorMessage: '',
|
||||
})
|
||||
})
|
||||
})
|
||||
@ -96,4 +96,19 @@ describe('ToolNode', () => {
|
||||
expect(container).toBeEmptyDOMElement()
|
||||
})
|
||||
})
|
||||
|
||||
it('should render multi-select configuration values', () => {
|
||||
render(
|
||||
<Node
|
||||
id="tool-node-1"
|
||||
data={createNodeData({
|
||||
tool_configurations: {
|
||||
formats: { type: 'constant', value: ['png', 'svg'] },
|
||||
},
|
||||
})}
|
||||
/>,
|
||||
)
|
||||
|
||||
expect(screen.getByTitle('png, svg')).toHaveTextContent('png, svg')
|
||||
})
|
||||
})
|
||||
|
||||
@ -51,7 +51,15 @@ const nodeDefault: NodeDefault<ToolNodeType> = {
|
||||
field: field.label,
|
||||
})
|
||||
} else {
|
||||
if (!errorMessages && (value === undefined || value === null || value === ''))
|
||||
const isEmptyMultiSelect =
|
||||
field.type === 'select' &&
|
||||
field.multiple &&
|
||||
Array.isArray(value) &&
|
||||
value.length === 0
|
||||
if (
|
||||
!errorMessages &&
|
||||
(value === undefined || value === null || value === '' || isEmptyMultiSelect)
|
||||
)
|
||||
errorMessages = t(($) => $[`${i18nPrefix}.fieldRequired`], {
|
||||
ns: 'workflow',
|
||||
field: field.label,
|
||||
@ -67,7 +75,12 @@ const nodeDefault: NodeDefault<ToolNodeType> = {
|
||||
})
|
||||
.forEach((field: any) => {
|
||||
const value = payload.tool_configurations[field.variable]
|
||||
if (!errorMessages && (value === undefined || value === null || value === ''))
|
||||
const isEmptyMultiSelect =
|
||||
field.type === 'select' && field.multiple && Array.isArray(value) && value.length === 0
|
||||
if (
|
||||
!errorMessages &&
|
||||
(value === undefined || value === null || value === '' || isEmptyMultiSelect)
|
||||
)
|
||||
errorMessages = t(($) => $[`${i18nPrefix}.fieldRequired`], {
|
||||
ns: 'workflow',
|
||||
field: field.label[language],
|
||||
|
||||
@ -76,6 +76,14 @@ const Node: FC<NodeProps<ToolNodeType>> = ({ data }) => {
|
||||
: tool_configurations[key].value}
|
||||
</div>
|
||||
)}
|
||||
{Array.isArray(tool_configurations[key].value) && (
|
||||
<div
|
||||
title={tool_configurations[key].value.join(', ')}
|
||||
className="w-0 shrink-0 grow truncate text-right text-xs font-normal text-text-secondary"
|
||||
>
|
||||
{tool_configurations[key].value.join(', ')}
|
||||
</div>
|
||||
)}
|
||||
{typeof tool_configurations[key] !== 'string' &&
|
||||
tool_configurations[key]?.type === FormTypeEnum.modelSelector && (
|
||||
<div
|
||||
|
||||
Loading…
Reference in New Issue
Block a user