From 9a520d5d288d37b4b7f5d80941d94d7b6e8891cd Mon Sep 17 00:00:00 2001 From: Nicola Soranzo Date: Mon, 29 Sep 2025 09:12:49 +0100 Subject: [PATCH] Remove/replace unnecessary casts --- lib/galaxy/managers/jobs.py | 2 +- lib/galaxy/tool_util/parameters/convert.py | 93 ++++++------ lib/galaxy/tools/parameters/__init__.py | 153 +++++++++----------- lib/galaxy/webapps/galaxy/api/tools.py | 5 +- lib/galaxy/webapps/galaxy/services/tools.py | 5 +- test/functional/test_toolbox_pytest.py | 5 +- 6 files changed, 124 insertions(+), 139 deletions(-) diff --git a/lib/galaxy/managers/jobs.py b/lib/galaxy/managers/jobs.py index 1e3ad580c3d..5e58076c730 100644 --- a/lib/galaxy/managers/jobs.py +++ b/lib/galaxy/managers/jobs.py @@ -2042,7 +2042,7 @@ class JobSubmitter: def _tool_request(self, tool_request_id: int) -> ToolRequest: sa_session = self.app.model.context - tool_request: ToolRequest = cast(ToolRequest, sa_session.query(ToolRequest).get(tool_request_id)) + tool_request = sa_session.get(ToolRequest, tool_request_id) if tool_request is None: raise Exception(f"Problem fetching request with ID {tool_request_id}") return tool_request diff --git a/lib/galaxy/tool_util/parameters/convert.py b/lib/galaxy/tool_util/parameters/convert.py index c940c80ef4b..df6fe701a80 100644 --- a/lib/galaxy/tool_util/parameters/convert.py +++ b/lib/galaxy/tool_util/parameters/convert.py @@ -9,21 +9,29 @@ from typing import ( Dict, List, Optional, - Union, ) from galaxy.tool_util_models.parameters import ( + BooleanParameterModel, ConditionalParameterModel, ConditionalWhen, create_job_runtime_model, + DataCollectionParameterModel, DataCollectionRequest, + DataColumnParameterModel, + DataParameterModel, DataRequestHda, DataRequestInternalHda, DataRequestUri, DiscriminatorType, + DrillDownParameterModel, FloatParameterModel, + GenomeBuildParameterModel, HiddenParameterModel, IntegerParameterModel, + RepeatParameterModel, + SectionParameterModel, + SelectParameterModel, TextParameterModel, ToolParameterBundle, ToolParameterT, @@ -158,7 +166,7 @@ def strictify(relaxed_state: RelaxedRequestToolState, input_models: ToolParamete tool_state = deepcopy(relaxed_state.input_state) def _strictify_parameter(tool_state: Dict[str, Any], parameter: ToolParameterT) -> None: - if parameter.parameter_type == "gx_conditional": + if isinstance(parameter, ConditionalParameterModel): conditional_state = _initialize_conditional_state(parameter, tool_state) test_parameter = parameter.test_parameter @@ -171,18 +179,17 @@ def strictify(relaxed_state: RelaxedRequestToolState, input_models: ToolParamete when = _select_which_when(parameter, test_value, conditional_state) _strictify_parameter(conditional_state, test_parameter) _strictify_parameters(conditional_state, when) - elif parameter.parameter_type == "gx_repeat": + elif isinstance(parameter, RepeatParameterModel): repeat_instances = _initialize_repeat_state(parameter, tool_state) for instance_state in repeat_instances: _strictify_parameters(instance_state, parameter) - elif parameter.parameter_type == "gx_section": + elif isinstance(parameter, SectionParameterModel): section_state = _initialize_section_state(parameter, tool_state) _fill_defaults(section_state, parameter) - elif parameter.parameter_type == "gx_text": + elif isinstance(parameter, TextParameterModel): parameter_name = parameter.name - text_parameter = parameter if parameter_name not in tool_state: - if not text_parameter.optional: + if not parameter.optional: # restore legacy behavior of allowing empty string implicit default # for these non-optional inputs. tool_state[parameter_name] = "" @@ -191,7 +198,7 @@ def strictify(relaxed_state: RelaxedRequestToolState, input_models: ToolParamete else: # legacy behavior of converting explicit None into implicit null. We should introduce # a layer somewhere to deal with this behavior further up the stack and clean up these models. - if not text_parameter.optional and tool_state[parameter_name] is None: + if not parameter.optional and tool_state[parameter_name] is None: tool_state[parameter_name] = "" def _strictify_parameters(tool_state: Dict[str, Any], input_models: ToolParameterBundle) -> None: @@ -219,7 +226,7 @@ def dereference( return src_dict def dereference_callback(parameter: ToolParameterT, value: Any): - if parameter.parameter_type == "gx_data": + if isinstance(parameter, DataParameterModel): if value is None: return VISITOR_NO_REPLACEMENT if parameter.multiple and isinstance(value, list): @@ -248,7 +255,7 @@ def encode_test( ): def encode_callback(parameter: ToolParameterT, value: Any): - if parameter.parameter_type == "gx_data": + if isinstance(parameter, DataParameterModel): if value is not None: if parameter.multiple: assert isinstance(value, list), str(value) @@ -258,22 +265,22 @@ def encode_test( assert isinstance(value, dict), str(value) test_dataset = cast(JsonTestDatasetDefDict, value) return adapt_datasets(test_dataset).model_dump() - elif parameter.parameter_type == "gx_data_collection": + elif isinstance(parameter, DataCollectionParameterModel): if value is not None: assert isinstance(value, dict), str(value) test_collection = cast(JsonTestCollectionDefDict, value) return adapt_collections(test_collection).model_dump() - elif parameter.parameter_type == "gx_select": + elif isinstance(parameter, SelectParameterModel): if parameter.multiple and value is not None: return [v.strip() for v in value.split(",")] else: return VISITOR_NO_REPLACEMENT - elif parameter.parameter_type == "gx_drill_down": + elif isinstance(parameter, DrillDownParameterModel): if parameter.multiple and value is not None: return [v.strip() for v in value.split(",")] else: return VISITOR_NO_REPLACEMENT - elif parameter.parameter_type == "gx_data_column": + elif isinstance(parameter, DataColumnParameterModel): if parameter.multiple and value is not None and isinstance(value, (str,)): return [int(v.strip()) for v in value.split(",")] else: @@ -313,27 +320,19 @@ def _fill_defaults(tool_state: Dict[str, Any], input_models: ToolParameterBundle def _fill_default_for(tool_state: Dict[str, Any], parameter: ToolParameterT) -> None: parameter_name = parameter.name - if parameter.parameter_type == "gx_boolean": + if isinstance(parameter, BooleanParameterModel): if parameter_name not in tool_state: # even optional parameters default to false if not in the body of the request :_( # see test_tools.py -> expression_null_handling_boolean or test cases for gx_boolean_optional.xml tool_state[parameter_name] = parameter.value or False - if parameter.parameter_type in ["gx_integer", "gx_float", "gx_hidden"]: - has_value_parameter = cast( - Union[ - IntegerParameterModel, - FloatParameterModel, - HiddenParameterModel, - ], - parameter, - ) + if isinstance(parameter, (IntegerParameterModel, FloatParameterModel, HiddenParameterModel)): if parameter_name not in tool_state: - tool_state[parameter_name] = has_value_parameter.value - elif parameter.parameter_type == "gx_genomebuild": + tool_state[parameter_name] = parameter.value + elif isinstance(parameter, GenomeBuildParameterModel): if parameter_name not in tool_state and parameter.optional: tool_state[parameter_name] = None - elif parameter.parameter_type == "gx_select": + elif isinstance(parameter, SelectParameterModel): # don't fill in dynamic parameters - wait for runtime to specify the default if parameter.dynamic_options: return @@ -343,7 +342,7 @@ def _fill_default_for(tool_state: Dict[str, Any], parameter: ToolParameterT) -> tool_state[parameter_name] = parameter.default_value else: tool_state[parameter_name] = None - elif parameter.parameter_type == "gx_drill_down": + elif isinstance(parameter, DrillDownParameterModel): if parameter_name not in tool_state: if parameter.multiple: options = parameter.default_options @@ -353,7 +352,7 @@ def _fill_default_for(tool_state: Dict[str, Any], parameter: ToolParameterT) -> option = parameter.default_option if option is not None: tool_state[parameter_name] = option - elif parameter.parameter_type == "gx_conditional": + elif isinstance(parameter, ConditionalParameterModel): conditional_state = _initialize_conditional_state(parameter, tool_state) test_parameter = parameter.test_parameter @@ -366,35 +365,33 @@ def _fill_default_for(tool_state: Dict[str, Any], parameter: ToolParameterT) -> when = _select_which_when(parameter, test_value, conditional_state) _fill_default_for(conditional_state, test_parameter) _fill_defaults(conditional_state, when) - elif parameter.parameter_type == "gx_repeat": + elif isinstance(parameter, RepeatParameterModel): repeat_instances = _initialize_repeat_state(parameter, tool_state) for instance_state in repeat_instances: _fill_defaults(instance_state, parameter) - elif parameter.parameter_type == "gx_section": + elif isinstance(parameter, SectionParameterModel): section_state = _initialize_section_state(parameter, tool_state) _fill_defaults(section_state, parameter) - elif parameter.parameter_type == "gx_data_collection": + elif isinstance(parameter, DataCollectionParameterModel): collection_parameter = parameter if parameter_name not in tool_state and collection_parameter.optional: tool_state[parameter_name] = None - elif parameter.parameter_type in ["gx_text"]: - text_parameter = cast(TextParameterModel, parameter) + elif isinstance(parameter, TextParameterModel): if parameter_name not in tool_state: - if not text_parameter.optional: + if not parameter.optional: # restore legacy behavior of allowing empty string implicit default # for these non-optional inputs. - tool_state[parameter_name] = text_parameter.default_value or "" + tool_state[parameter_name] = parameter.default_value or "" else: - tool_state[parameter_name] = text_parameter.default_value or None + tool_state[parameter_name] = parameter.default_value or None else: # legacy behavior of converting explicit None into implicit null. We should introduce # a layer somewhere to deal with this behavior further up the stack and clean up these models. - if not text_parameter.optional and tool_state[parameter_name] is None: - tool_state[parameter_name] = text_parameter.default_value or "" + if not parameter.optional and tool_state[parameter_name] is None: + tool_state[parameter_name] = parameter.default_value or "" -def _initialize_section_state(parameter: ToolParameterT, tool_state: Dict[str, Any]) -> Dict[str, Any]: - assert parameter.parameter_type == "gx_section" +def _initialize_section_state(parameter: SectionParameterModel, tool_state: Dict[str, Any]) -> Dict[str, Any]: parameter_name = parameter.name if parameter_name not in tool_state: tool_state[parameter_name] = {} @@ -402,8 +399,7 @@ def _initialize_section_state(parameter: ToolParameterT, tool_state: Dict[str, A return section_state -def _initialize_conditional_state(parameter: ToolParameterT, tool_state: Dict[str, Any]) -> Dict[str, Any]: - assert parameter.parameter_type == "gx_conditional" +def _initialize_conditional_state(parameter: ConditionalParameterModel, tool_state: Dict[str, Any]) -> Dict[str, Any]: parameter_name = parameter.name if parameter_name not in tool_state: tool_state[parameter_name] = {} @@ -414,8 +410,7 @@ def _initialize_conditional_state(parameter: ToolParameterT, tool_state: Dict[st return conditional_state -def _initialize_repeat_state(parameter: ToolParameterT, tool_state: Dict[str, Any]) -> List[Dict[str, Any]]: - assert parameter.parameter_type == "gx_repeat" +def _initialize_repeat_state(parameter: RepeatParameterModel, tool_state: Dict[str, Any]) -> List[Dict[str, Any]]: parameter_name = parameter.name if parameter_name not in tool_state: tool_state[parameter_name] = [] @@ -460,13 +455,13 @@ def _encode_callback_for(encode_id: EncodeFunctionT) -> Callback: return encode_src_dict(element) def encode_callback(parameter: ToolParameterT, value: Any): - if parameter.parameter_type == "gx_data": + if isinstance(parameter, DataParameterModel): if parameter.multiple and isinstance(value, list): return list(map(encode_element, value)) else: assert isinstance(value, dict), str(value) return encode_element(value) - elif parameter.parameter_type == "gx_data_collection": + elif isinstance(parameter, DataCollectionParameterModel): assert isinstance(value, dict), str(value) return encode_element(value) else: @@ -495,7 +490,7 @@ def _decode_callback_for(decode_id: DecodeFunctionT) -> Callback: return decode_src_dict(element) def decode_callback(parameter: ToolParameterT, value: Any): - if parameter.parameter_type == "gx_data": + if isinstance(parameter, DataParameterModel): if value is None: return VISITOR_NO_REPLACEMENT if parameter.multiple and isinstance(value, list): @@ -503,7 +498,7 @@ def _decode_callback_for(decode_id: DecodeFunctionT) -> Callback: else: assert isinstance(value, dict), str(value) return decode_element(value) - elif parameter.parameter_type == "gx_data_collection": + elif isinstance(parameter, DataCollectionParameterModel): if value is None: return VISITOR_NO_REPLACEMENT assert isinstance(value, dict), str(value) diff --git a/lib/galaxy/tools/parameters/__init__.py b/lib/galaxy/tools/parameters/__init__.py index 79abac27f78..b43eb8b0282 100644 --- a/lib/galaxy/tools/parameters/__init__.py +++ b/lib/galaxy/tools/parameters/__init__.py @@ -22,6 +22,7 @@ from galaxy.util import unicodify from galaxy.util.expressions import ExpressionContext from galaxy.util.json import safe_loads from .basic import ( + ColumnListParameter, DataCollectionToolParameter, DataToolParameter, ParameterValueError, @@ -455,13 +456,12 @@ def populate_state( elif input_format == "21.01": context = ExpressionContext(state, context) for input in inputs.values(): - state[input.name] = input.get_initial_value(request_context, context) - group_state = state[input.name] + input_name = input.name + group_state = state[input_name] = input.get_initial_value(request_context, context) if isinstance(input, Repeat): - repeat_name = input.name - repeat_incoming = incoming.get(repeat_name) or [] + repeat_incoming = incoming.get(input_name) or [] if repeat_incoming and (len(repeat_incoming) > input.max or len(repeat_incoming) < input.min): - errors[repeat_name] = "The number of repeat elements is outside the range specified by the tool." + errors[input_name] = "The number of repeat elements is outside the range specified by the tool." else: del group_state[:] for rep in repeat_incoming: @@ -480,12 +480,12 @@ def populate_state( input_format=input_format, ) if repeat_errors: - errors[input.name] = repeat_errors + errors[input_name] = repeat_errors elif isinstance(input, Conditional): test_param = input.test_param assert test_param is not None - incoming_group_state = incoming.get(input.name, {}) + incoming_group_state = incoming.get(input_name, {}) if test_param.name in incoming_group_state: test_param_value = incoming_group_state.get(test_param.name) else: @@ -500,9 +500,9 @@ def populate_state( else: try: current_case = input.get_current_case(value) - group_state = state[input.name] = {} + group_state = state[input_name] = {} cast_errors: ParameterValidationErrorsT = {} - incoming_for_conditional = cast(ToolStateJobInstanceT, incoming.get(input.name) or {}) + incoming_for_conditional = cast(ToolStateJobInstanceT, incoming.get(input_name) or {}) populate_state( request_context, input.cases[current_case].inputs, @@ -515,7 +515,7 @@ def populate_state( input_format=input_format, ) if cast_errors: - errors[input.name] = cast_errors + errors[input_name] = cast_errors group_state["__current_case__"] = current_case except Exception: errors[test_param.name] = "The selected case is unavailable/invalid." @@ -523,7 +523,7 @@ def populate_state( elif isinstance(input, Section): section_errors: ParameterValidationErrorsT = {} - incoming_for_state = cast(ToolStateJobInstanceT, incoming.get(input.name) or {}) + incoming_for_state = cast(ToolStateJobInstanceT, incoming.get(input_name) or {}) populate_state( request_context, input.inputs, @@ -536,22 +536,22 @@ def populate_state( input_format=input_format, ) if section_errors: - errors[input.name] = section_errors + errors[input_name] = section_errors - elif input.type == "upload_dataset": + elif isinstance(input, UploadDataset): raise NotImplementedError else: assert isinstance(input, ToolParameter) - param_value = _get_incoming_value(incoming, input.name, state.get(input.name)) + param_value = _get_incoming_value(incoming, input_name, state.get(input_name)) value, error = ( check_param(request_context, input, param_value, context, simple_errors=simple_errors) if check else [param_value, None] ) if error: - errors[input.name] = error - state[input.name] = value + errors[input_name] = error + state[input_name] = value else: raise RequestParameterInvalidException( f"Input format {input_format} not recognized; input_format must be either legacy or 21.01." @@ -701,85 +701,78 @@ def populate_state_async( ): context = ExpressionContext(state, context) for input in inputs.values(): - initial_value = input.get_initial_value(request_context, context) input_name = input.name - state[input_name] = initial_value - group_state = state[input_name] - if input.type == "repeat": - repeat_input = cast(Repeat, input) - if ( - len(incoming[repeat_input.name]) > repeat_input.max - or len(incoming[repeat_input.name]) < repeat_input.min - ): - errors[repeat_input.name] = "The number of repeat elements is outside the range specified by the tool." + group_state = state[input_name] = input.get_initial_value(request_context, context) + if isinstance(input, Repeat): + repeat_incoming = incoming[input_name] + if len(repeat_incoming) > input.max or len(repeat_incoming) < input.min: + errors[input_name] = "The number of repeat elements is outside the range specified by the tool." else: del group_state[:] - for rep in incoming[repeat_input.name]: + for rep in repeat_incoming: new_state: ToolStateJobInstancePopulatedT = {} group_state.append(new_state) repeat_errors: ParameterValidationErrorsT = {} populate_state_async( request_context, - repeat_input.inputs, + input.inputs, rep, new_state, repeat_errors, context=context, ) if repeat_errors: - errors[repeat_input.name] = repeat_errors + errors[input_name] = repeat_errors - elif input.type == "conditional": - conditional_input = cast(Conditional, input) - test_param = cast(ToolParameter, conditional_input.test_param) - test_param_value = incoming.get(conditional_input.name, {}).get(test_param.name) + elif isinstance(input, Conditional): + test_param = cast(ToolParameter, input.test_param) + test_param_value = incoming.get(input_name, {}).get(test_param.name) value, error = check_param(request_context, test_param, test_param_value, context) if error: errors[test_param.name] = error else: try: - current_case = conditional_input.get_current_case(value) - group_state = state[conditional_input.name] = {} + current_case = input.get_current_case(value) + group_state = state[input_name] = {} cast_errors: ParameterValidationErrorsT = {} populate_state_async( request_context, - conditional_input.cases[current_case].inputs, - cast(ToolStateJobInstanceT, incoming.get(conditional_input.name)), + input.cases[current_case].inputs, + cast(ToolStateJobInstanceT, incoming.get(input_name)), group_state, cast_errors, context=context, ) if cast_errors: - errors[conditional_input.name] = cast_errors + errors[input_name] = cast_errors group_state["__current_case__"] = current_case except Exception: errors[test_param.name] = "The selected case is unavailable/invalid." group_state[test_param.name] = value - elif input.type == "section": - section_input = cast(Section, input) + elif isinstance(input, Section): section_errors: ParameterValidationErrorsT = {} populate_state_async( request_context, - section_input.inputs, - cast(ToolStateJobInstanceT, incoming.get(section_input.name)), + input.inputs, + cast(ToolStateJobInstanceT, incoming.get(input_name)), group_state, section_errors, context=context, ) if section_errors: - errors[section_input.name] = section_errors + errors[input_name] = section_errors - elif input.type == "upload_dataset": + elif isinstance(input, UploadDataset): raise NotImplementedError else: assert isinstance(input, ToolParameter) - param_value = _get_incoming_value(incoming, input.name, state.get(input.name)) + param_value = _get_incoming_value(incoming, input_name, state.get(input_name)) value, error = check_param(request_context, input, param_value, context, simple_errors=False) if error: - errors[input.name] = error - state[input.name] = value + errors[input_name] = error + state[input_name] = value def to_internal_single(value): if isinstance(value, HistoryDatasetCollectionAssociation): @@ -797,18 +790,17 @@ def populate_state_async( return to_internal_single(value) if input_name not in incoming: - if input.type == "data_column": + if isinstance(input, ColumnListParameter): if isinstance(value, str): incoming[input_name] = int(value) elif isinstance(value, list): incoming[input_name] = [int(v) for v in value] else: incoming[input_name] = value - elif input.type == "text": - text_input = cast(TextToolParameter, input) + elif isinstance(input, TextToolParameter): # see behavior of tools in test_tools.py::test_null_to_text_tools # these parameters act as empty string in this context - if value is None and not text_input.optional: + if value is None and not input.optional: incoming[input_name] = "" else: incoming[input_name] = value @@ -828,65 +820,62 @@ def fill_dynamic_defaults( """ context = ExpressionContext(job_tool_state, job_tool_state) for input in inputs.values(): - if input.type == "repeat": - repeat_input = cast(Repeat, input) - repeat_name = repeat_input.name - for rep, rep_params in enumerate(job_tool_state[repeat_name]): + input_name = input.name + if isinstance(input, Repeat): + for rep, rep_params in enumerate(job_tool_state[input_name]): fill_dynamic_defaults( request_context, - repeat_input.inputs, + input.inputs, rep_params, - params[repeat_name][rep], + params[input_name][rep], context=context, ) - elif input.type == "conditional": - conditional_input = cast(Conditional, input) - test_param = cast(ToolParameter, conditional_input.test_param) - test_param_value = job_tool_state.get(conditional_input.name, {}).get(test_param.name) + elif isinstance(input, Conditional): + test_param = cast(ToolParameter, input.test_param) + test_param_value = job_tool_state.get(input_name, {}).get(test_param.name) try: - current_case = conditional_input.get_current_case(test_param_value) + current_case = input.get_current_case(test_param_value) fill_dynamic_defaults( request_context, - conditional_input.cases[current_case].inputs, - cast(ToolStateJobInstanceT, job_tool_state.get(conditional_input.name)), - cast(ToolStateJobInstancePopulatedT, params.get(conditional_input.name)), + input.cases[current_case].inputs, + cast(ToolStateJobInstanceT, job_tool_state.get(input_name)), + cast(ToolStateJobInstancePopulatedT, params.get(input_name)), context=context, ) except Exception: raise Exception("The selected case is unavailable/invalid.") - elif input.type == "section": - section_input = cast(Section, input) + elif isinstance(input, Section): fill_dynamic_defaults( request_context, - section_input.inputs, - cast(ToolStateJobInstanceT, job_tool_state.get(section_input.name)), - cast(ToolStateJobInstancePopulatedT, params.get(section_input.name)), + input.inputs, + cast(ToolStateJobInstanceT, job_tool_state.get(input_name)), + cast(ToolStateJobInstancePopulatedT, params.get(input_name)), context=context, ) - elif input.type == "upload_dataset": + elif isinstance(input, UploadDataset): raise NotImplementedError else: - if input.name not in job_tool_state and input.name in params: - if input.type == "data_column": - if isinstance(params[input.name], str): - job_tool_state[input.name] = int(params[input.name]) - elif isinstance(params[input.name], list): - job_tool_state[input.name] = [int(v) for v in params[input.name]] + if input_name not in job_tool_state and input_name in params: + if isinstance(input, ColumnListParameter): + if isinstance(params[input_name], str): + job_tool_state[input_name] = int(params[input_name]) + elif isinstance(params[input_name], list): + job_tool_state[input_name] = [int(v) for v in params[input_name]] else: - job_tool_state[input.name] = params[input.name] - elif input.type == "data_collection": - data_collection = params[input.name] + job_tool_state[input_name] = params[input_name] + elif isinstance(input, DataCollectionToolParameter): + data_collection = params[input_name] if data_collection: - job_tool_state[input.name] = { + job_tool_state[input_name] = { "src": "hdca", "id": data_collection.id, } else: - job_tool_state[input.name] = params[input.name] + job_tool_state[input_name] = params[input_name] def _get_incoming_value(incoming, key, default): diff --git a/lib/galaxy/webapps/galaxy/api/tools.py b/lib/galaxy/webapps/galaxy/api/tools.py index 531b3cf8b6a..219f538e814 100644 --- a/lib/galaxy/webapps/galaxy/api/tools.py +++ b/lib/galaxy/webapps/galaxy/api/tools.py @@ -332,12 +332,9 @@ class FetchTools: def _get_tool_request_or_raise_not_found( self, trans: ProvidesHistoryContext, id: DecodedDatabaseIdField ) -> ToolRequest: - tool_request: Optional[ToolRequest] = cast( - Optional[ToolRequest], trans.app.model.context.query(ToolRequest).get(id) - ) + tool_request = trans.app.model.context.get(ToolRequest, id) if tool_request is None: raise exceptions.ObjectNotFound() - assert tool_request return tool_request @router.post("/api/tool_landings", public=True, allow_cors=True) diff --git a/lib/galaxy/webapps/galaxy/services/tools.py b/lib/galaxy/webapps/galaxy/services/tools.py index a2c8ad7a2d8..8396ced129c 100644 --- a/lib/galaxy/webapps/galaxy/services/tools.py +++ b/lib/galaxy/webapps/galaxy/services/tools.py @@ -6,6 +6,7 @@ from json import dumps from typing import ( Any, cast, + get_args, Optional, Union, ) @@ -335,9 +336,9 @@ class ToolsService(ServiceBase): preferred_object_store_id = payload.get("preferred_object_store_id") credentials_context = payload.get("credentials_context") input_format = str(payload.get("input_format", "legacy")) - if input_format not in ("legacy", "21.01"): + if input_format not in get_args(InputFormatT): raise exceptions.RequestParameterInvalidException(f"input_format invalid {input_format}") - input_format = cast(InputFormatT, input_format) + input_format = cast(InputFormatT, input_format) # https://github.com/python/mypy/issues/15106 if "data_manager_mode" in payload: incoming["__data_manager_mode"] = payload["data_manager_mode"] vars = tool.handle_input( diff --git a/test/functional/test_toolbox_pytest.py b/test/functional/test_toolbox_pytest.py index 82915fba890..b6b999de4ed 100644 --- a/test/functional/test_toolbox_pytest.py +++ b/test/functional/test_toolbox_pytest.py @@ -1,6 +1,7 @@ import os from typing import ( cast, + get_args, NamedTuple, ) @@ -65,7 +66,9 @@ class TestFrameworkTools(ApiTestCase): @pytest.mark.parametrize("testcase", cases(), ids=idfn) def test_tool(self, testcase: ToolTest): - use_legacy_api = cast(UseLegacyApiT, os.environ.get("GALAXY_TEST_USE_LEGACY_TOOL_API", DEFAULT_USE_LEGACY_API)) + use_legacy_api = os.environ.get("GALAXY_TEST_USE_LEGACY_TOOL_API", DEFAULT_USE_LEGACY_API) + assert use_legacy_api in get_args(UseLegacyApiT) + cast(UseLegacyApiT, use_legacy_api) # https://github.com/python/mypy/issues/15106 self._test_driver.run_tool_test( testcase.tool_id, testcase.test_index, tool_version=testcase.tool_version, use_legacy_api=use_legacy_api )