Merge pull request #22637 from nsoranzo/improve_wf_management_type_annot

Improve type annotation of workflow management methods
This commit is contained in:
Marius van den Beek
2026-05-06 16:13:13 +02:00
committed by GitHub
9 changed files with 128 additions and 91 deletions
+95 -74
View File
@@ -113,6 +113,7 @@ from galaxy.work.context import WorkRequestContext
from galaxy.workflow.modules import (
module_factory,
PickValueModule,
SubWorkflowModule,
ToolModule,
WorkflowModule,
WorkflowModuleInjector,
@@ -615,6 +616,30 @@ def _step_pja_dict(step):
return pja_dict
class WorkflowStateResolutionOptions(BaseModel):
# fill in default tool state when updating, may change tool_state
fill_defaults: bool = False
# If True, assume all tool state coming from generated form instead of potentially simpler json stored in DB/exported
from_tool_form: bool = False
# If False, allow running with less exact tool versions
exact_tools: bool = True
class WorkflowUpdateOptions(WorkflowStateResolutionOptions):
# Only used internally, don't set. If using the API assume updating the workflows
# representation with name or annotation for instance, updates the corresponding
# stored workflow
update_stored_workflow_attributes: bool = True
allow_missing_tools: bool = False
dry_run: bool = False
class RawWorkflowDescription:
def __init__(self, as_dict, workflow_path: str | None = None):
self.as_dict = as_dict
self.workflow_path = workflow_path
class WorkflowContentsManager(UsesAnnotations):
def __init__(self, app: MinimalManagerApp, trs_proxy: TrsProxy):
@@ -740,8 +765,12 @@ class WorkflowContentsManager(UsesAnnotations):
return CreatedWorkflow(stored_workflow=stored, workflow=workflow, missing_tools=missing_tool_tups)
def update_workflow_from_raw_description(
self, trans, stored_workflow, raw_workflow_description, workflow_update_options
):
self,
trans,
stored_workflow: StoredWorkflow,
raw_workflow_description: RawWorkflowDescription,
workflow_update_options: WorkflowUpdateOptions,
) -> tuple[Workflow, list[str]]:
raw_workflow_description = self.ensure_raw_description(raw_workflow_description)
# Put parameters in workflow mode
@@ -951,13 +980,13 @@ class WorkflowContentsManager(UsesAnnotations):
def workflow_to_dict(
self,
trans,
stored,
style="export",
version=None,
history=None,
instance_id=None,
preserve_external_subworkflow_links=False,
):
stored: StoredWorkflow,
style: str = "export",
version: int | None = None,
history: History | None = None,
instance_id: int | None = None,
preserve_external_subworkflow_links: bool = False,
) -> dict[str, Any]:
"""Export the workflow contents to a dictionary ready for JSON-ification and to be
sent out via API for instance. There are three styles of export allowed 'export', 'instance', and
'editor'. The Galaxy team will do its best to preserve the backward compatibility of the
@@ -974,11 +1003,7 @@ class WorkflowContentsManager(UsesAnnotations):
def to_format_2(wf_dict, **kwds):
return from_galaxy_native(wf_dict, None, **kwds)
if version == "":
version = None
if version is not None:
version = int(version)
elif instance_id:
if version is None and instance_id:
# If the instance_id is provided, we need to extract the workflow instance via the version.
for i, workflow in enumerate(reversed(stored.workflows)):
if workflow.id == instance_id:
@@ -1000,24 +1025,24 @@ class WorkflowContentsManager(UsesAnnotations):
elif style == "format2":
wf_dict = self._workflow_to_dict_export(
trans,
stored,
workflow=workflow,
stored=stored,
preserve_external_subworkflow_links=preserve_external_subworkflow_links,
)
wf_dict = to_format_2(wf_dict)
elif style == "format2_wrapped_yaml":
wf_dict = self._workflow_to_dict_export(
trans,
stored,
workflow=workflow,
stored=stored,
preserve_external_subworkflow_links=preserve_external_subworkflow_links,
)
wf_dict = to_format_2(wf_dict, json_wrapper=True)
elif style == "ga":
wf_dict = self._workflow_to_dict_export(
trans,
stored,
workflow=workflow,
stored=stored,
preserve_external_subworkflow_links=preserve_external_subworkflow_links,
)
else:
@@ -1031,17 +1056,24 @@ class WorkflowContentsManager(UsesAnnotations):
wf_dict["version"] = len(stored.workflows) - 1
return wf_dict
def _sync_stored_workflow(self, trans, stored_workflow):
def _sync_stored_workflow(self, trans, stored_workflow: StoredWorkflow) -> None:
if trans.user_is_admin:
workflow_path = stored_workflow.from_path
self.store_workflow_to_path(workflow_path, stored_workflow, stored_workflow.latest_workflow, trans=trans)
assert workflow_path is not None
self.store_workflow_to_path(
workflow_path, stored_workflow=stored_workflow, workflow=stored_workflow.latest_workflow, trans=trans
)
def store_workflow_artifacts(self, directory, filename_base, workflow, **kwd):
def store_workflow_artifacts(
self, directory: str, filename_base: str, workflow: Workflow, history: History, user: User | None
) -> None:
modern_workflow_path = os.path.join(directory, f"{filename_base}.gxwf.yml")
legacy_workflow_path = os.path.join(directory, f"{filename_base}.ga")
abstract_cwl_workflow_path = os.path.join(directory, f"{filename_base}.abstract.cwl")
for path in [legacy_workflow_path, modern_workflow_path, abstract_cwl_workflow_path]:
self.app.workflow_contents_manager.store_workflow_to_path(path, workflow.stored_workflow, workflow, **kwd)
self.store_workflow_to_path(
path, stored_workflow=workflow.stored_workflow, workflow=workflow, history=history, user=user
)
try:
cytoscape_path = os.path.join(directory, f"{filename_base}.html")
to_cytoscape(modern_workflow_path, cytoscape_path)
@@ -1049,26 +1081,32 @@ class WorkflowContentsManager(UsesAnnotations):
# completely optional and currently broken so ignore...
pass
def store_workflow_to_path(self, workflow_path, stored_workflow, workflow, **kwd):
trans = kwd.get("trans")
def store_workflow_to_path(
self,
workflow_path: str,
stored_workflow: StoredWorkflow,
workflow: Workflow,
trans=None,
history: History | None = None,
user: User | None = None,
) -> None:
if trans is None:
trans = WorkRequestContext(app=self.app, user=kwd.get("user"), history=kwd.get("history"))
trans = WorkRequestContext(app=self.app, user=user, history=history)
workflow = stored_workflow.latest_workflow
with open(workflow_path, "w") as f:
wf_dict = self._workflow_to_dict_export(trans, workflow=workflow, stored=stored_workflow)
if workflow_path.endswith(".ga"):
wf_dict = self._workflow_to_dict_export(trans, stored_workflow, workflow=workflow)
json.dump(wf_dict, f, indent=4)
elif workflow_path.endswith(".abstract.cwl"):
wf_dict = self._workflow_to_dict_export(trans, stored_workflow, workflow=workflow)
abstract_dict = from_dict(wf_dict)
ordered_dump(abstract_dict, f)
else:
wf_dict = self._workflow_to_dict_export(trans, stored_workflow, workflow=workflow)
wf_dict = from_galaxy_native(wf_dict, None, json_wrapper=True)
f.write(wf_dict["yaml_content"])
def _workflow_to_dict_run(self, trans: ProvidesUserContext, stored, workflow, history=None):
def _workflow_to_dict_run(
self, trans: ProvidesUserContext, stored: StoredWorkflow, workflow: Workflow, history: History | None = None
) -> dict[str, Any]:
"""
Builds workflow dictionary used by run workflow form
"""
@@ -1079,7 +1117,7 @@ class WorkflowContentsManager(UsesAnnotations):
trans.workflow_building_mode = workflow_building_modes.USE_HISTORY
module_injector = WorkflowModuleInjector(trans)
has_upgrade_messages = False
step_version_changes = []
step_version_changes: list[str] = []
missing_tools = []
errors = {}
module_injector.inject_all(workflow, exact_tools=False, ignore_tool_missing_exception=True)
@@ -1094,20 +1132,20 @@ class WorkflowContentsManager(UsesAnnotations):
if step.upgrade_messages:
has_upgrade_messages = True
if step.type in ("tool", "subworkflow", None):
assert isinstance(step.module, (ToolModule, SubWorkflowModule))
if step.module.version_changes:
step_version_changes.extend(step.module.version_changes)
step_errors = step.module.get_errors()
if step_errors:
errors[step.id] = step_errors
if missing_tools:
workflow.annotation = self.get_item_annotation_str(trans.sa_session, trans.user, workflow)
raise exceptions.MessageException(f"Following tools missing: {', '.join(missing_tools)}")
workflow.annotation = self.get_item_annotation_str(trans.sa_session, trans.user, workflow)
step_order_indices = {}
for step in workflow.steps:
step_order_indices[step.id] = step.order_index
step_models = []
for step in workflow.steps:
assert step.module is not None
step_model = None
if step.type == "tool":
incoming: dict[str, Any] = {}
@@ -1118,6 +1156,7 @@ class WorkflowContentsManager(UsesAnnotations):
raise exceptions.MessageException(
f"Following tool missing or inaccessible: '{step.tool_id}/{step.tool_uuid}'"
)
assert step.state is not None
params_to_incoming(incoming, tool.inputs, step.state.inputs, trans.app)
step_model = tool.to_json(
trans, incoming, workflow_building_mode=workflow_building_modes.USE_HISTORY, history=history
@@ -1141,15 +1180,17 @@ class WorkflowContentsManager(UsesAnnotations):
step_model["step_name"] = step.module.get_name()
step_model["step_version"] = step.module.get_version()
step_model["step_index"] = step.order_index
step_model["output_connections"] = [
{
"input_step_index": step_order_indices.get(oc.input_step_id),
"output_step_index": step_order_indices.get(oc.output_step_id),
"input_name": oc.input_name,
"output_name": oc.output_name,
}
for oc in step.output_connections
]
step_model["output_connections"] = []
for oc in step.output_connections:
assert oc.output_step_id is not None
step_model["output_connections"].append(
{
"input_step_index": step_order_indices.get(oc.input_step_id),
"output_step_index": step_order_indices.get(oc.output_step_id),
"input_name": oc.input_name,
"output_name": oc.output_name,
}
)
if step.annotations:
step_model["annotation"] = step.annotations[0].annotation
if step.upgrade_messages:
@@ -1537,12 +1578,12 @@ class WorkflowContentsManager(UsesAnnotations):
def _workflow_to_dict_export(
self,
trans,
stored=None,
workflow=None,
internal=False,
allow_upgrade=False,
preserve_external_subworkflow_links=False,
):
workflow: Workflow,
stored: StoredWorkflow | None = None,
internal: bool = False,
allow_upgrade: bool = False,
preserve_external_subworkflow_links: bool = False,
) -> dict[str, Any]:
"""Export the workflow contents to a dictionary ready for JSON-ification and export.
If internal, use content_ids instead subworkflow definitions.
@@ -1664,6 +1705,7 @@ class WorkflowContentsManager(UsesAnnotations):
del step_dict["tool_version"]
del step_dict["tool_state"]
subworkflow = step.subworkflow
assert subworkflow is not None
if preserve_external_subworkflow_links and subworkflow.source_metadata:
del step_dict["content_id"]
source_metadata = subworkflow.source_metadata
@@ -1677,8 +1719,8 @@ class WorkflowContentsManager(UsesAnnotations):
del step_dict["content_id"]
subworkflow_as_dict = self._workflow_to_dict_export(
trans,
stored=None,
workflow=subworkflow,
stored=None,
preserve_external_subworkflow_links=preserve_external_subworkflow_links,
)
step_dict["subworkflow"] = subworkflow_as_dict
@@ -1782,7 +1824,9 @@ class WorkflowContentsManager(UsesAnnotations):
steps[step.order_index] = step_dict
return data
def _workflow_to_dict_instance(self, trans, stored, workflow, legacy=True):
def _workflow_to_dict_instance(
self, trans, stored: StoredWorkflow, workflow: Workflow, legacy: bool = True
) -> dict[str, Any]:
encode = self.app.security.encode_id
sa_session = self.app.model.context
item = stored.to_dict(view="element")
@@ -1844,6 +1888,7 @@ class WorkflowContentsManager(UsesAnnotations):
del step_dict["tool_id"]
del step_dict["tool_version"]
del step_dict["tool_inputs"]
assert step.subworkflow is not None
step_dict["workflow_id"] = step.subworkflow.id
for conn in step.input_connections:
@@ -2244,7 +2289,7 @@ class WorkflowContentsManager(UsesAnnotations):
workflow = stored_workflow.get_internal_version(refactor_request.version)
as_dict = self._workflow_to_dict_export(
trans, stored_workflow, workflow=workflow, internal=True, allow_upgrade=True
trans, workflow=workflow, stored=stored_workflow, internal=True, allow_upgrade=True
)
raw_workflow_description = self.normalize_workflow_format(trans, as_dict)
workflow_update_options = WorkflowUpdateOptions(
@@ -2406,24 +2451,6 @@ class RefactorResponse(BaseModel):
dry_run: bool
class WorkflowStateResolutionOptions(BaseModel):
# fill in default tool state when updating, may change tool_state
fill_defaults: bool = False
# If True, assume all tool state coming from generated form instead of potentially simpler json stored in DB/exported
from_tool_form: bool = False
# If False, allow running with less exact tool versions
exact_tools: bool = True
class WorkflowUpdateOptions(WorkflowStateResolutionOptions):
# Only used internally, don't set. If using the API assume updating the workflows
# representation with name or annotation for instance, updates the corresponding
# stored workflow
update_stored_workflow_attributes: bool = True
allow_missing_tools: bool = False
dry_run: bool = False
# Workflow update options but with some different defaults - we allow creating
# workflows with missing tools by default but not updating.
class WorkflowCreateOptions(WorkflowStateResolutionOptions):
@@ -2476,12 +2503,6 @@ class MissingToolsException(exceptions.MessageException):
self.errors = errors
class RawWorkflowDescription:
def __init__(self, as_dict, workflow_path=None):
self.as_dict = as_dict
self.workflow_path = workflow_path
def _get_stored_workflow(session, workflow_uuid, workflow_id, by_stored_id):
stmt = select(StoredWorkflow)
if workflow_uuid is not None:
+9 -10
View File
@@ -417,11 +417,11 @@ def set_datatypes_registry(d_registry):
_datatypes_registry = d_registry
class HasTags:
class HasTags(Dictifiable):
dict_collection_visible_keys = ["tags"]
dict_element_visible_keys = ["tags"]
def to_dict(self, *args, **kwargs):
def to_dict(self, *args, **kwargs) -> dict[str, Any]:
rval = super().to_dict(*args, **kwargs)
rval["tags"] = self.make_tag_string_list()
return rval
@@ -3508,7 +3508,7 @@ class HistoryAudit(Base):
session.execute(q)
class History(Base, HasTags, Dictifiable, UsesAnnotations, HasName, Serializable, UsesCreateAndUpdateTime):
class History(Base, HasTags, UsesAnnotations, HasName, Serializable, UsesCreateAndUpdateTime):
__tablename__ = "history"
__table_args__ = (Index("ix_history_slug", "slug", mysql_length=200),)
@@ -5875,7 +5875,7 @@ class DatasetInstance(RepresentById, UsesCreateAndUpdateTime, _HasTable):
rval["file_metadata"] = file_metadata
class HistoryDatasetAssociation(DatasetInstance, HasTags, Dictifiable, UsesAnnotations, HasName, Serializable):
class HistoryDatasetAssociation(DatasetInstance, HasTags, UsesAnnotations, HasName, Serializable):
"""
Resource class that creates a relation between a dataset and a user history.
"""
@@ -7830,7 +7830,6 @@ class HistoryDatasetCollectionAssociation(
Base,
DatasetCollectionInstance,
HasTags,
Dictifiable,
UsesAnnotations,
Serializable,
):
@@ -8582,7 +8581,7 @@ class GalaxySessionToHistoryAssociation(Base, RepresentById):
self.history = history
class StoredWorkflow(Base, HasTags, Dictifiable, RepresentById, UsesCreateAndUpdateTime):
class StoredWorkflow(Base, HasTags, RepresentById, UsesCreateAndUpdateTime):
"""
StoredWorkflow represents the root node of a tree of objects that compose a workflow, including workflow revisions, steps, and subworkflows.
It is responsible for the metadata associated with a workflow including owner, name, published, and create/update time.
@@ -8703,7 +8702,7 @@ class StoredWorkflow(Base, HasTags, Dictifiable, RepresentById, UsesCreateAndUpd
self.workflows = listify(workflow)
self.hidden = hidden
def get_internal_version(self, version: Optional[int] = None):
def get_internal_version(self, version: Optional[int] = None) -> "Workflow":
if version is None:
return self.latest_workflow
if len(self.workflows) <= version:
@@ -8756,7 +8755,7 @@ class StoredWorkflow(Base, HasTags, Dictifiable, RepresentById, UsesCreateAndUpd
rows_as_dict = dict(r for r in rows if r[0] is not None) # type: ignore[arg-type, var-annotated]
return InvocationsStateCounts(rows_as_dict)
def to_dict(self, view="collection", value_mapper=None):
def to_dict(self, view="collection", value_mapper=None) -> dict[str, Any]:
rval = super().to_dict(view=view, value_mapper=value_mapper)
rval["latest_workflow_uuid"] = (lambda uuid: str(uuid) if self.latest_workflow.uuid else None)(
self.latest_workflow.uuid
@@ -11520,7 +11519,7 @@ class UserAuthnzToken(Base, UserMixin, RepresentById):
return instance
class Page(Base, HasTags, Dictifiable, RepresentById, UsesCreateAndUpdateTime):
class Page(Base, HasTags, RepresentById, UsesCreateAndUpdateTime):
__tablename__ = "page"
__table_args__ = (Index("ix_page_slug", "slug", mysql_length=200),)
@@ -11641,7 +11640,7 @@ class PageUserShareAssociation(Base, UserShareAssociation):
page: Mapped["Page"] = relationship(back_populates="users_shared_with")
class Visualization(Base, HasTags, Dictifiable, RepresentById, UsesCreateAndUpdateTime):
class Visualization(Base, HasTags, RepresentById, UsesCreateAndUpdateTime):
__tablename__ = "visualization"
__table_args__ = (
Index("ix_visualization_dbkey", "dbkey", mysql_length=200),
+1 -1
View File
@@ -2579,7 +2579,7 @@ class DirectoryModelExportStore(ModelExportStore):
if not self.app:
raise Exception(f"Missing self.app in {self}.")
self.app.workflow_contents_manager.store_workflow_artifacts(
workflows_directory, workflow_key, workflow, user=history.user, history=history
workflows_directory, workflow_key, workflow, history=history, user=history.user
)
invocations_attrs.append(invocation_attrs)
+6 -1
View File
@@ -6,6 +6,7 @@ import logging
from collections.abc import Callable
from typing import (
Any,
TYPE_CHECKING,
)
from webob.exc import (
@@ -59,6 +60,9 @@ from galaxy.web.form_builder import (
)
from galaxy.workflow.modules import WorkflowModuleInjector
if TYPE_CHECKING:
from galaxy.structured_app import StructuredApp
log = logging.getLogger(__name__)
# States for passing messages
@@ -582,6 +586,7 @@ class UsesVisualizationMixin(UsesLibraryMixinItems):
class UsesStoredWorkflowMixin(SharableItemSecurityMixin, UsesAnnotations):
"""Mixin for controllers that use StoredWorkflow objects."""
app: "StructuredApp"
slug_builder = SlugBuilder()
def get_stored_workflow(self, trans, id, check_ownership=True, check_accessible=False):
@@ -637,7 +642,7 @@ class UsesStoredWorkflowMixin(SharableItemSecurityMixin, UsesAnnotations):
session.commit()
return imported_stored
def _workflow_to_dict(self, trans, stored):
def _workflow_to_dict(self, trans, stored: StoredWorkflow) -> dict[str, Any]:
"""
Converts a workflow to a dict of attributes suitable for exporting.
"""
@@ -343,6 +343,11 @@ class WorkflowsAPIController(
style = kwd.get("style", "export")
download_format = kwd.get("format")
version = kwd.get("version")
if version is not None:
try:
version = int(version)
except ValueError:
raise exceptions.RequestParameterInvalidException("Invalid version specified.")
history = None
if history_id := kwd.get("history_id"):
history = self.history_manager.get_accessible(
@@ -24,6 +24,7 @@ log = logging.getLogger(__name__)
class WorkflowController(BaseUIController, SharableMixin, UsesStoredWorkflowMixin, UsesItemRatings):
app: StructuredApp
history_manager: HistoryManager = depends(HistoryManager)
slug_builder = SlugBuilder()
+1 -1
View File
@@ -777,7 +777,7 @@ class SubWorkflowModule(WorkflowModule):
return self._modules
@property
def version_changes(self):
def version_changes(self) -> list[str]:
version_changes = []
for m in self.get_modules():
if hasattr(m, "version_changes"):
+9 -3
View File
@@ -1544,7 +1544,9 @@ steps:
"trs_tool_id": "#workflow/github.com/jmchilton/galaxy-workflow-dockstore-example-1/mycoolworkflow",
"trs_version_id": "master",
}
workflow_id = self._post("workflows", data=trs_payload).json()["id"]
response = self._post("workflows", data=trs_payload)
response.raise_for_status()
workflow_id = response.json()["id"]
original_workflow = self._download_workflow(workflow_id)
assert "Test Workflow" in original_workflow["name"]
assert original_workflow.get("source_metadata").get("trs_tool_id") == trs_payload["trs_tool_id"]
@@ -1571,7 +1573,9 @@ steps:
"%23workflow%2Fgithub.com%2Fjmchilton%2Fgalaxy-workflow-dockstore-example-1%2Fmycoolworkflow/"
"versions/master",
}
workflow_id = self._post("workflows", data=trs_payload).json()["id"]
response = self._post("workflows", data=trs_payload)
response.raise_for_status()
workflow_id = response.json()["id"]
original_workflow = self._download_workflow(workflow_id)
assert "Test Workflow" in original_workflow["name"]
assert (
@@ -1604,7 +1608,9 @@ steps:
"archive_source": "trs_tool",
"trs_url": "https://workflowhub.eu/ga4gh/trs/v2/tools/109/versions/5",
}
workflow_id = self._post("workflows", data=trs_payload).json()["id"]
response = self._post("workflows", data=trs_payload)
response.raise_for_status()
workflow_id = response.json()["id"]
original_workflow = self._download_workflow(workflow_id)
assert "COVID-19: variation analysis reporting" in original_workflow["name"]
assert original_workflow.get("source_metadata").get("trs_tool_id") == "109"
+1 -1
View File
@@ -6,7 +6,7 @@ ignore_missing_imports = True
check_untyped_defs = True
exclude = (?x)(
^build/
| ^test/functional/
| .*test/functional/tools/cwl_tools/
| .*tool_shed/test/test_data/repos/
)
pretty = True