diff --git a/api/core/tools/__base/tool.py b/api/core/tools/__base/tool.py index b16f80169fc..e023ce117b1 100644 --- a/api/core/tools/__base/tool.py +++ b/api/core/tools/__base/tool.py @@ -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 diff --git a/api/core/tools/entities/tool_entities.py b/api/core/tools/entities/tool_entities.py index 0c77693dde4..786910d91d4 100644 --- a/api/core/tools/entities/tool_entities.py +++ b/api/core/tools/entities/tool_entities.py @@ -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): diff --git a/api/openapi/markdown/console-openapi.md b/api/openapi/markdown/console-openapi.md index 02a20e95504..22496868c8c 100644 --- a/api/openapi/markdown/console-openapi.md +++ b/api/openapi/markdown/console-openapi.md @@ -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
integer | | No | | min | number
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 | diff --git a/api/tests/unit_tests/core/plugin/test_plugin_entities.py b/api/tests/unit_tests/core/plugin/test_plugin_entities.py index deac0ba1da5..0a532646abb 100644 --- a/api/tests/unit_tests/core/plugin/test_plugin_entities.py +++ b/api/tests/unit_tests/core/plugin/test_plugin_entities.py @@ -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"), [ diff --git a/api/tests/unit_tests/core/tools/test_base_tool.py b/api/tests/unit_tests/core/tools/test_base_tool.py index 9e80e086472..f164e3fddea 100644 --- a/api/tests/unit_tests/core/tools/test_base_tool.py +++ b/api/tests/unit_tests/core/tools/test_base_tool.py @@ -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"}}, diff --git a/api/tests/unit_tests/core/tools/test_tool_entities.py b/api/tests/unit_tests/core/tools/test_tool_entities.py index a5b7e8a9a34..82ab59752c2 100644 --- a/api/tests/unit_tests/core/tools/test_tool_entities.py +++ b/api/tests/unit_tests/core/tools/test_tool_entities.py @@ -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", diff --git a/api/tests/unit_tests/services/tools/test_tools_transform_service.py b/api/tests/unit_tests/services/tools/test_tools_transform_service.py index 32c1a00d301..dcc02d0cf00 100644 --- a/api/tests/unit_tests/services/tools/test_tools_transform_service.py +++ b/api/tests/unit_tests/services/tools/test_tools_transform_service.py @@ -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: diff --git a/packages/contracts/generated/api/console/workspaces/types.gen.ts b/packages/contracts/generated/api/console/workspaces/types.gen.ts index 3a3b4ad620c..350b46dd1a2 100644 --- a/packages/contracts/generated/api/console/workspaces/types.gen.ts +++ b/packages/contracts/generated/api/console/workspaces/types.gen.ts @@ -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 placeholder?: I18nObject | null diff --git a/packages/contracts/generated/api/console/workspaces/zod.gen.ts b/packages/contracts/generated/api/console/workspaces/zod.gen.ts index 1f1530897fa..77bfc7dfe19 100644 --- a/packages/contracts/generated/api/console/workspaces/zod.gen.ts +++ b/packages/contracts/generated/api/console/workspaces/zod.gen.ts @@ -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(), diff --git a/web/app/components/workflow/nodes/_base/components/__tests__/form-input-item.helpers.spec.ts b/web/app/components/workflow/nodes/_base/components/__tests__/form-input-item.helpers.spec.ts index e9353369108..19f498c6c95 100644 --- a/web/app/components/workflow/nodes/_base/components/__tests__/form-input-item.helpers.spec.ts +++ b/web/app/components/workflow/nodes/_base/components/__tests__/form-input-item.helpers.spec.ts @@ -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', () => { diff --git a/web/app/components/workflow/nodes/_base/components/form-input-item.helpers.ts b/web/app/components/workflow/nodes/_base/components/form-input-item.helpers.ts index 12f70742041..a0e045aa130 100644 --- a/web/app/components/workflow/nodes/_base/components/form-input-item.helpers.ts +++ b/web/app/components/workflow/nodes/_base/components/form-input-item.helpers.ts @@ -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) diff --git a/web/app/components/workflow/nodes/tool/__tests__/default.spec.ts b/web/app/components/workflow/nodes/tool/__tests__/default.spec.ts new file mode 100644 index 00000000000..e8b464a4be7 --- /dev/null +++ b/web/app/components/workflow/nodes/tool/__tests__/default.spec.ts @@ -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: '', + }) + }) +}) diff --git a/web/app/components/workflow/nodes/tool/__tests__/node.spec.tsx b/web/app/components/workflow/nodes/tool/__tests__/node.spec.tsx index cc7c6b4e0d1..6b411922922 100644 --- a/web/app/components/workflow/nodes/tool/__tests__/node.spec.tsx +++ b/web/app/components/workflow/nodes/tool/__tests__/node.spec.tsx @@ -96,4 +96,19 @@ describe('ToolNode', () => { expect(container).toBeEmptyDOMElement() }) }) + + it('should render multi-select configuration values', () => { + render( + , + ) + + expect(screen.getByTitle('png, svg')).toHaveTextContent('png, svg') + }) }) diff --git a/web/app/components/workflow/nodes/tool/default.ts b/web/app/components/workflow/nodes/tool/default.ts index 12830cb9136..3287ee2c144 100644 --- a/web/app/components/workflow/nodes/tool/default.ts +++ b/web/app/components/workflow/nodes/tool/default.ts @@ -51,7 +51,15 @@ const nodeDefault: NodeDefault = { 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 = { }) .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], diff --git a/web/app/components/workflow/nodes/tool/node.tsx b/web/app/components/workflow/nodes/tool/node.tsx index 80bb7cde0b0..b41b1a0f940 100644 --- a/web/app/components/workflow/nodes/tool/node.tsx +++ b/web/app/components/workflow/nodes/tool/node.tsx @@ -76,6 +76,14 @@ const Node: FC> = ({ data }) => { : tool_configurations[key].value} )} + {Array.isArray(tool_configurations[key].value) && ( +
+ {tool_configurations[key].value.join(', ')} +
+ )} {typeof tool_configurations[key] !== 'string' && tool_configurations[key]?.type === FormTypeEnum.modelSelector && (