Cleanup tool request hacks for legacy using new relaxed request tool state.

A strict vs relaxed state for incoming requests lets us preserve some older behavior if it is needed for tests and such while also allowing us to default to "the correct thing".
This commit is contained in:
John Chilton
2025-10-29 16:33:27 -04:00
parent e929bc648b
commit a685889a32
11 changed files with 225 additions and 38 deletions
@@ -10,6 +10,12 @@ input_state: Dict[str, Any]
+ {abstract} _to_base_model(parameters: ToolParameterBundle): Optional[Type[BaseModel]]
}
class RelaxedRequestToolState {
state_representation = "request"
+ _to_base_model(parameters: ToolParameterBundle): Type[BaseModel]
}
note bottom: Object references of the form \n{src: "hda", id: <encoded_id>}.\n Allow mapping/reduce constructs. Relaxed syntax to allow some of the odd stuff allowed in legacy tool input format.
class RequestToolState {
state_representation = "request"
+ _to_base_model(parameters: ToolParameterBundle): Type[BaseModel]
@@ -53,6 +59,7 @@ state_representation = "workflow_step_linked"
}
note bottom: Expect pre-process ``in`` dictionaries and bring in representation\n of links and defaults and validate them in model.\n
ToolState <|-- RelaxedRequestToolState
ToolState <|-- RequestToolState
ToolState <|-- RequestInternalToolState
ToolState <|-- RequestInternalDereferencedToolState
@@ -61,6 +68,8 @@ ToolState <|-- TestCaseToolState
ToolState <|-- WorkflowStepToolState
ToolState <|-- WorkflowStepLinkedToolState
RelaxedRequestToolState - RequestToolState : strictify >
RequestToolState - RequestInternalToolState : decode >
RequestInternalToolState - RequestInternalDereferencedToolState : dereference >
@@ -46,6 +46,7 @@ from .convert import (
fill_static_defaults,
landing_decode,
landing_encode,
strictify,
)
from .factory import (
from_input_source,
@@ -63,6 +64,7 @@ from .model_validation import (
validate_internal_request,
validate_internal_request_dereferenced,
validate_landing_request,
validate_relaxed_request,
validate_request,
validate_test_case,
validate_workflow_step,
@@ -74,6 +76,7 @@ from .state import (
JobInternalToolState,
LandingRequestInternalToolState,
LandingRequestToolState,
RelaxedRequestToolState,
RequestInternalDereferencedToolState,
RequestInternalToolState,
RequestToolState,
@@ -138,6 +141,7 @@ __all__ = (
"validate_internal_request",
"validate_internal_request_dereferenced",
"validate_landing_request",
"validate_relaxed_request",
"validate_request",
"validate_test_case",
"validate_workflow_step",
@@ -151,6 +155,7 @@ __all__ = (
"to_json_schema_string",
"test_case_state",
"validate_test_cases_for_tool_source",
"RelaxedRequestToolState",
"RequestToolState",
"RequestInternalToolState",
"RequestInternalDereferencedToolState",
@@ -168,6 +173,7 @@ __all__ = (
"landing_decode",
"landing_encode",
"dereference",
"strictify",
"WorkflowStepToolState",
"WorkflowStepLinkedToolState",
)
+95 -18
View File
@@ -1,6 +1,7 @@
"""Utilities for converting between request states."""
import logging
from copy import deepcopy
from typing import (
Any,
Callable,
@@ -35,6 +36,7 @@ from .state import (
JobInternalToolState,
LandingRequestInternalToolState,
LandingRequestToolState,
RelaxedRequestToolState,
RequestInternalDereferencedToolState,
RequestInternalToolState,
RequestToolState,
@@ -80,7 +82,10 @@ def cwl_runtime_model(input_models: ToolParameterBundle):
def decode(
external_state: RequestToolState, input_models: ToolParameterBundle, decode_id: Callable[[str], int], name_base: Optional[str] = None
external_state: RequestToolState,
input_models: ToolParameterBundle,
decode_id: Callable[[str], int],
name_base: Optional[str] = None,
) -> RequestInternalToolState:
"""Prepare an internal representation of tool state (request_internal) for storing in the database."""
@@ -147,6 +152,59 @@ def landing_encode(
return request_state
def strictify(relaxed_state: RelaxedRequestToolState, input_models: ToolParameterBundle) -> RequestToolState:
"""Convert a relaxed request state into a strict request state by applying legacy behavior."""
tool_state = deepcopy(relaxed_state.input_state)
def _strictify_parameter(tool_state: Dict[str, Any], parameter: ToolParameterT) -> None:
if parameter.parameter_type == "gx_conditional":
conditional_state = _initialize_conditional_state(parameter, tool_state)
test_parameter = parameter.test_parameter
test_parameter_name = test_parameter.name
explicit_test_value: Optional[DiscriminatorType] = (
conditional_state[test_parameter_name] if test_parameter_name in conditional_state else None
)
test_value = validate_explicit_conditional_test_value(test_parameter_name, explicit_test_value)
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":
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":
section_state = _initialize_section_state(parameter, tool_state)
_fill_defaults(section_state, parameter)
elif parameter.parameter_type == "gx_text":
parameter_name = parameter.name
text_parameter = parameter
if parameter_name not in tool_state:
if not text_parameter.optional:
# restore legacy behavior of allowing empty string implicit default
# for these non-optional inputs.
tool_state[parameter_name] = ""
else:
tool_state[parameter_name] = 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] = ""
def _strictify_parameters(tool_state: Dict[str, Any], input_models: ToolParameterBundle) -> None:
for parameter in input_models.parameters:
_strictify_parameter(tool_state, parameter)
_strictify_parameters(tool_state, input_models)
request_state = RequestToolState(tool_state)
request_state.validate(input_models)
return request_state
def dereference(
internal_state: RequestInternalToolState, input_models: ToolParameterBundle, dereference: DereferenceCallable
) -> RequestInternalDereferencedToolState:
@@ -296,12 +354,7 @@ def _fill_default_for(tool_state: Dict[str, Any], parameter: ToolParameterT) ->
if option is not None:
tool_state[parameter_name] = option
elif parameter.parameter_type == "gx_conditional":
if parameter_name not in tool_state:
tool_state[parameter_name] = {}
raw_conditional_state = tool_state[parameter_name]
assert isinstance(raw_conditional_state, dict)
conditional_state = cast(Dict[str, Any], raw_conditional_state)
conditional_state = _initialize_conditional_state(parameter, tool_state)
test_parameter = parameter.test_parameter
test_parameter_name = test_parameter.name
@@ -314,19 +367,11 @@ def _fill_default_for(tool_state: Dict[str, Any], parameter: ToolParameterT) ->
_fill_default_for(conditional_state, test_parameter)
_fill_defaults(conditional_state, when)
elif parameter.parameter_type == "gx_repeat":
if parameter_name not in tool_state:
tool_state[parameter_name] = []
repeat_instances = cast(List[Dict[str, Any]], tool_state[parameter_name])
if parameter.min:
while len(repeat_instances) < parameter.min:
repeat_instances.append({})
for instance_state in tool_state[parameter_name]:
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":
if parameter_name not in tool_state:
tool_state[parameter_name] = {}
section_state = cast(Dict[str, Any], tool_state[parameter_name])
section_state = _initialize_section_state(parameter, tool_state)
_fill_defaults(section_state, parameter)
elif parameter.parameter_type == "gx_data_collection":
collection_parameter = parameter
@@ -348,6 +393,38 @@ def _fill_default_for(tool_state: Dict[str, Any], parameter: ToolParameterT) ->
tool_state[parameter_name] = ""
def _initialize_section_state(parameter: ToolParameterT, tool_state: Dict[str, Any]) -> Dict[str, Any]:
assert parameter.parameter_type == "gx_section"
parameter_name = parameter.name
if parameter_name not in tool_state:
tool_state[parameter_name] = {}
section_state = cast(Dict[str, Any], tool_state[parameter_name])
return section_state
def _initialize_conditional_state(parameter: ToolParameterT, tool_state: Dict[str, Any]) -> Dict[str, Any]:
assert parameter.parameter_type == "gx_conditional"
parameter_name = parameter.name
if parameter_name not in tool_state:
tool_state[parameter_name] = {}
raw_conditional_state = tool_state[parameter_name]
assert isinstance(raw_conditional_state, dict)
conditional_state = cast(Dict[str, Any], raw_conditional_state)
return conditional_state
def _initialize_repeat_state(parameter: ToolParameterT, tool_state: Dict[str, Any]) -> List[Dict[str, Any]]:
assert parameter.parameter_type == "gx_repeat"
parameter_name = parameter.name
if parameter_name not in tool_state:
tool_state[parameter_name] = []
repeat_instances = cast(List[Dict[str, Any]], tool_state[parameter_name])
if parameter.min:
while len(repeat_instances) < parameter.min:
repeat_instances.append({})
return repeat_instances
def _select_which_when(
conditional: ConditionalParameterModel, test_value: Optional[DiscriminatorType], conditional_state: Dict[str, Any]
@@ -47,6 +47,7 @@ def validate_model_type_factory(state_representation: StateRepresentationT) -> V
return validate_request
validate_relaxed_request = validate_model_type_factory("relaxed_request")
validate_request = validate_model_type_factory("request")
validate_internal_request = validate_model_type_factory("request_internal")
validate_internal_request_dereferenced = validate_model_type_factory("request_internal_dereferenced")
+9
View File
@@ -18,6 +18,7 @@ from galaxy.tool_util_models.parameters import (
create_job_internal_model,
create_landing_request_internal_model,
create_landing_request_model,
create_relaxed_request_model,
create_request_internal_dereferenced_model,
create_request_internal_model,
create_request_model,
@@ -73,6 +74,14 @@ class ToolState(ABC):
"""Return a model type for this tool state kind."""
class RelaxedRequestToolState(ToolState):
state_representation: Literal["relaxed_request"] = "relaxed_request"
@classmethod
def _parameter_model_for(cls, parameters: ToolParameterBundle, name: Optional[str] = None) -> Type[BaseModel]:
return create_relaxed_request_model(parameters, name)
class RequestToolState(ToolState):
state_representation: Literal["request"] = "request"
+12 -13
View File
@@ -76,6 +76,7 @@ from .tool_source import (
# + request_internal: This is a pydantic model to validate what Galaxy expects to find in the database,
# in particular dataset and collection references should be decoded integers.
StateRepresentationT = Literal[
"relaxed_request",
"request",
"request_internal",
"request_internal_dereferenced",
@@ -275,15 +276,15 @@ class TextParameterModel(BaseGalaxyToolParameterModelDefinition):
return optional_if_needed(StrictStr, self.optional)
@property
def py_type_request(self) -> Type:
def py_type_relaxed_request(self) -> Type:
# such a hack but explicit nulls are always allowed in the API even for non-optional
# parameters - it becomes "" in the internal state.
return optional(StrictStr)
def pydantic_template(self, state_representation: StateRepresentationT) -> DynamicModelInformation:
py_type = self.py_type
if state_representation in ["request", "request_internal", "request_internal_dereferenced"]:
py_type = self.py_type_request
if state_representation == "relaxed_request":
py_type = self.py_type_relaxed_request
py_type = decorate_type_with_validators_if_needed(py_type, self.validators)
if state_representation == "workflow_step_linked":
py_type = allow_connected_value(py_type)
@@ -313,15 +314,8 @@ class IntegerParameterModel(BaseGalaxyToolParameterModelDefinition):
def py_type(self) -> Type:
return optional_if_needed(StrictInt, self.optional)
#@property
#def py_type_request(self) -> Type:
# # ugh... we allow explicit nulls in the API even if the input is not optional
# return optional_if_needed(StrictInt, True)
def pydantic_template(self, state_representation: StateRepresentationT) -> DynamicModelInformation:
py_type = self.py_type
# if state_representation == "request":
# py_type = self.py_type_request
validators = self.validators[:]
if self.min is not None or self.max is not None:
validators.append(InRangeParameterValidatorModel(min=self.min, max=self.max, implicit=True))
@@ -685,9 +679,9 @@ class DataParameterModel(BaseGalaxyToolParameterModelDefinition):
return optional_if_needed(base_model, self.optional)
def pydantic_template(self, state_representation: StateRepresentationT) -> DynamicModelInformation:
if state_representation == "request":
if state_representation in ["request", "relaxed_request"]:
return allow_batching(dynamic_model_information_from_py_type(self, self.py_type), BatchDataInstance)
if state_representation == "landing_request":
elif state_representation == "landing_request":
return allow_batching(
dynamic_model_information_from_py_type(self, self.py_type, requires_value=False), BatchDataInstance
)
@@ -715,6 +709,10 @@ class DataParameterModel(BaseGalaxyToolParameterModelDefinition):
return dynamic_model_information_from_py_type(self, type(None), requires_value=False)
elif state_representation == "workflow_step_linked":
return dynamic_model_information_from_py_type(self, ConnectedValue)
else:
raise NotImplementedError(
f"Have not implemented data collection parameter models for state representation {state_representation}"
)
@property
def request_requires_value(self) -> bool:
@@ -844,7 +842,7 @@ class DataCollectionParameterModel(BaseGalaxyToolParameterModelDefinition):
return optional_if_needed(base_type, self.optional)
def pydantic_template(self, state_representation: StateRepresentationT) -> DynamicModelInformation:
if state_representation == "request":
if state_representation in ["request", "relaxed_request"]:
return allow_batching(dynamic_model_information_from_py_type(self, self.py_type))
elif state_representation == "landing_request":
return allow_batching(dynamic_model_information_from_py_type(self, self.py_type, requires_value=False))
@@ -1821,6 +1819,7 @@ def create_model_factory(state_representation: StateRepresentationT):
return create_method
create_relaxed_request_model = create_model_factory("relaxed_request")
create_request_model = create_model_factory("request")
create_request_internal_model = create_model_factory("request_internal")
create_request_internal_dereferenced_model = create_model_factory("request_internal_dereferenced")
+14 -1
View File
@@ -54,7 +54,9 @@ from galaxy.schema.tasks import (
from galaxy.security.idencoding import IdEncodingHelper
from galaxy.tool_util.parameters import (
decode,
RelaxedRequestToolState,
RequestToolState,
strictify,
)
from galaxy.webapps.galaxy.services.base import (
async_task_summary,
@@ -77,6 +79,11 @@ class JobRequest(BaseModel):
tool_version: Optional[str] = Field(default=None, title="tool_version", description="TODO")
history_id: Optional[DecodedDatabaseIdField] = Field(default=None, title="history_id", description="TODO")
inputs: Optional[dict[str, Any]] = Field(default_factory=lambda: {}, title="Inputs", description="TODO")
strict: bool = Field(
default=True,
title="Strict",
description="Turn on strict validation of the inputs that drops support for some inconsistent legacy behavior.",
)
use_cached_jobs: Optional[bool] = Field(default=None, title="use_cached_jobs")
rerun_remap_job_id: Optional[DecodedDatabaseIdField] = Field(
default=None, title="rerun_remap_job_id", description="TODO"
@@ -239,7 +246,13 @@ class JobsService(ServiceBase):
if history_id is not None:
target_history = self.history_manager.get_owned(history_id, trans.user, current_history=trans.history)
inputs = job_request.inputs
request_state = RequestToolState(inputs or {})
strict = job_request.strict
if not strict:
relaxed_request_state = RelaxedRequestToolState(inputs or {})
relaxed_request_state.validate(tool, f"{tool.id} (relaxed request model)")
request_state = strictify(relaxed_request_state, tool)
else:
request_state = RequestToolState(inputs or {})
request_state.validate(tool, f"{tool.id} (request model)")
request_internal_state = decode(request_state, tool, trans.security.decode_id)
tool_request = ToolRequest()
+5 -2
View File
@@ -1211,11 +1211,12 @@ class BaseDatasetPopulator(BasePopulator):
payload = self.run_tool_payload(tool_id, inputs, history_id, **kwds)
return self.tools_post(payload)
def tool_request_raw(self, tool_id: str, inputs: dict[str, Any], history_id: str) -> Response:
def tool_request_raw(self, tool_id: str, inputs: dict[str, Any], history_id: str, strict: bool = True) -> Response:
payload = {
"tool_id": tool_id,
"history_id": history_id,
"inputs": inputs,
"strict": strict,
}
response = self._post("jobs", data=payload, json=True)
return response
@@ -4128,7 +4129,9 @@ class DescribeToolExecution:
kwds["input_format"] = self._input_format
history_id = self._ensure_history_id
if self._input_format == "request":
execute_response = self._dataset_populator.tool_request_raw(self._tool_id, self._inputs, history_id)
execute_response = self._dataset_populator.tool_request_raw(
self._tool_id, self._inputs, history_id, strict=False
)
if execute_response.status_code == 200:
response_json = execute_response.json()
tool_request_id = response_json.get("tool_request_id")
@@ -16,17 +16,32 @@
# - non optional, multi-selects require a selection (see TODO below...)
# - https://github.com/galaxyproject/galaxy/issues/18541
gx_int:
request_valid:
request_valid: &gx_int_request_valid
- parameter: 5
- parameter: 6
# galaxy parameters created with a value - so doesn't need to appear in request even though non-optional
- {}
request_invalid:
request_invalid: &gx_int_request_invalid
- parameter: "5"
- parameter: null
- parameter: "null"
- parameter: "None"
- parameter: { 5 }
- parameter: {__class__: 'ConnectedValue'}
# int parameters have no differences between relaxed and strict, internal or external
# requests, or dereferenced or not.
relaxed_request_valid:
*gx_int_request_valid
relaxed_request_invalid:
*gx_int_request_invalid
request_internal_valid:
*gx_int_request_valid
request_internal_invalid:
*gx_int_request_invalid
request_internal_dereferenced_valid:
*gx_int_request_valid
request_internal_dereferenced_invalid:
*gx_int_request_invalid
job_internal_valid:
- parameter: 5
job_internal_invalid:
@@ -114,16 +129,28 @@ gx_boolean:
- parameter: {__class__: 'ConnectedValue3'}
gx_int_optional:
request_valid:
request_valid: &gx_int_optional_request_valid
- parameter: 5
- parameter: null
- {}
request_invalid:
request_invalid: &gx_int_optional_request_invalid
- parameter: "5"
- parameter: "None"
- parameter: "null"
- parameter: [5]
- parameter: {__class__: 'ConnectedValue'}
relaxed_request_valid:
*gx_int_optional_request_valid
relaxed_request_invalid:
*gx_int_optional_request_invalid
request_internal_valid:
*gx_int_optional_request_valid
request_internal_invalid:
*gx_int_optional_request_invalid
request_internal_dereferenced_valid:
*gx_int_optional_request_valid
request_internal_dereferenced_invalid:
*gx_int_optional_request_invalid
job_internal_valid:
- parameter: 5
- parameter: null
@@ -363,6 +390,20 @@ gx_text_optional:
- parameter: {}
- parameter: { "moo": "cow" }
gx_text_optional_false:
request_valid:
- parameter: "mytext"
# Should this be invalid? -John
- {}
request_invalid:
- parameter: null
relaxed_request_valid:
- parameter: "mytext"
- parameter: null
- {}
relaxed_request_invalid:
- parameter: 5
gx_text_length_validation:
request_valid:
- parameter: "mytext"
@@ -15,9 +15,11 @@ from galaxy.tool_util.parameters import (
landing_decode,
landing_encode,
LandingRequestToolState,
RelaxedRequestToolState,
RequestInternalDereferencedToolState,
RequestInternalToolState,
RequestToolState,
strictify,
)
from galaxy.tool_util.parser.util import parse_profile_version
from .test_parameter_test_cases import tool_source_for
@@ -259,6 +261,25 @@ def test_fill_defaults():
assert with_defaults["conditional_parameter"]["boolean_parameter"] is False
def test_strictify():
strict_state = strictify_for({"parameter": 1}, "parameters/gx_int")
assert strict_state["parameter"] == 1
strict_state = strictify_for({}, "parameters/gx_text_optional_false")
assert strict_state["parameter"] == ""
strict_state = strictify_for({"parameter": None}, "parameters/gx_text_optional_false")
assert strict_state["parameter"] == ""
def strictify_for(tool_state: Dict[str, Any], tool_path: str) -> Dict[str, Any]:
tool_source = tool_source_for(tool_path)
bundle = input_models_for_tool_source(tool_source)
relaxed_state = RelaxedRequestToolState(tool_state)
relaxed_state.validate(bundle)
return strictify(relaxed_state, bundle).input_state
def _fake_dereference(input: DataRequestUri) -> DataRequestInternalHda:
return DataRequestInternalHda(id=EXAMPLE_ID_1, src="hda")
@@ -20,6 +20,7 @@ from galaxy.tool_util.parameters import (
validate_internal_request,
validate_internal_request_dereferenced,
validate_landing_request,
validate_relaxed_request,
validate_request,
validate_test_case,
validate_workflow_step,
@@ -91,6 +92,8 @@ def _test_file(file: str, specification=None, parameter_bundle: Optional[ToolPar
assert parameter_bundle
assertion_functions = {
"relaxed_request_valid": _assert_relaxed_requests_validate,
"relaxed_request_invalid": _assert_relaxed_requests_invalid,
"request_valid": _assert_requests_validate,
"request_invalid": _assert_requests_invalid,
"request_internal_valid": _assert_internal_requests_validate,
@@ -150,6 +153,9 @@ def model_assertion_function_factory(validate_function: ValidationFunctionT, wha
return _assert_validates, _assert_invalid
_assert_relaxed_request_validates, _assert_relaxed_request_invalid = model_assertion_function_factory(
validate_relaxed_request, "relaxed request"
)
_assert_request_validates, _assert_request_invalid = model_assertion_function_factory(validate_request, "request")
_assert_internal_request_validates, _assert_internal_request_invalid = model_assertion_function_factory(
validate_internal_request, "internal request"
@@ -176,6 +182,8 @@ _assert_internal_landing_request_validates, _assert_internal_landing_request_inv
validate_internal_landing_request, " internallanding request"
)
_assert_relaxed_requests_validate = partial(_for_each, _assert_relaxed_request_validates)
_assert_relaxed_requests_invalid = partial(_for_each, _assert_relaxed_request_invalid)
_assert_requests_validate = partial(_for_each, _assert_request_validates)
_assert_requests_invalid = partial(_for_each, _assert_request_invalid)
_assert_internal_requests_validate = partial(_for_each, _assert_internal_request_validates)