Fix async tool requests for user defined tools

This commit is contained in:
John Chilton
2025-10-29 16:34:13 -04:00
parent ed5ff25b39
commit 0e1940e6bb
5 changed files with 42 additions and 8 deletions
+21 -5
View File
@@ -59,6 +59,7 @@ from galaxy.schema.tasks import (
PurgeDatasetsTaskRequest,
QueueJobs,
SetupHistoryExportJob,
TOOL_SOURCE_CLASS,
WriteHistoryContentTo,
WriteHistoryTo,
WriteInvocationTo,
@@ -74,9 +75,14 @@ log = get_logger(__name__)
@lru_cache
def cached_create_tool_from_representation(app: MinimalManagerApp, raw_tool_source: str, tool_dir: str = ""):
def cached_create_tool_from_representation(
app: MinimalManagerApp,
raw_tool_source: str,
tool_source_class: TOOL_SOURCE_CLASS,
tool_dir: str = "",
):
return create_tool_from_representation(
app=app, raw_tool_source=raw_tool_source, tool_dir=tool_dir, tool_source_class="XmlToolSource"
app=app, raw_tool_source=raw_tool_source, tool_dir=tool_dir, tool_source_class=tool_source_class
)
@@ -225,11 +231,14 @@ def setup_fetch_data(
self,
job_id: int,
raw_tool_source: str,
tool_source_class: TOOL_SOURCE_CLASS,
app: MinimalManagerApp,
sa_session: galaxy_scoped_session,
task_user_id: Optional[int] = None,
):
tool = cached_create_tool_from_representation(app=app, raw_tool_source=raw_tool_source)
tool = cached_create_tool_from_representation(
app=app, raw_tool_source=raw_tool_source, tool_source_class=tool_source_class
)
job = sa_session.get(Job, job_id)
assert job
# self.request.hostname is the actual worker name given by the `-n` argument, not the hostname as you might think.
@@ -258,11 +267,14 @@ def setup_fetch_data(
def finish_job(
job_id: int,
raw_tool_source: str,
tool_source_class: TOOL_SOURCE_CLASS,
app: MinimalManagerApp,
sa_session: galaxy_scoped_session,
task_user_id: Optional[int] = None,
):
tool = cached_create_tool_from_representation(app=app, raw_tool_source=raw_tool_source)
tool = cached_create_tool_from_representation(
app=app, raw_tool_source=raw_tool_source, tool_source_class=tool_source_class
)
job = sa_session.get(Job, job_id)
assert job
# TODO: assert state ?
@@ -335,8 +347,12 @@ def fetch_data(
@galaxy_task(action="queuing up submitted jobs")
def queue_jobs(request: QueueJobs, app: MinimalManagerApp, job_submitter: JobSubmitter):
tool = cached_create_tool_from_representation(
app, request.tool_source.raw_tool_source, tool_dir=request.tool_source.tool_dir
app=app,
raw_tool_source=request.tool_source.raw_tool_source,
tool_dir=request.tool_source.tool_dir,
tool_source_class=request.tool_source.tool_source_class,
)
job_submitter.queue_jobs(
tool,
request,
+5
View File
@@ -1,5 +1,6 @@
from enum import Enum
from typing import (
Literal,
Optional,
)
from uuid import UUID
@@ -153,9 +154,13 @@ class TaskResult(Model):
)
TOOL_SOURCE_CLASS = Literal["XmlToolSource", "YamlToolSource", "CwlToolSource"]
class ToolSource(Model):
raw_tool_source: str
tool_dir: str
tool_source_class: TOOL_SOURCE_CLASS = "XmlToolSource"
class QueueJobs(Model):
+2
View File
@@ -332,6 +332,8 @@ def _parse_test(i, test_dict) -> ToolSourceTest:
def to_test_assert_list(assertions) -> AssertionList:
assertions = assertions or []
def expand_dict_form(item):
key, value = item
new_value = value.copy()
+13 -3
View File
@@ -314,16 +314,26 @@ def _execute(
# task_user_id parameter is used to do task user rate limiting. It is only passed
# to first task in chain because it is only necessary to rate limit the first
# task in a chain.
tool_source_class = type(tool.tool_source).__name__
async_result = (
setup_fetch_data.s(
job_id, raw_tool_source=raw_tool_source, task_user_id=getattr(trans.user, "id", None)
job_id,
raw_tool_source=raw_tool_source,
task_user_id=getattr(trans.user, "id", None),
tool_source_class=tool_source_class,
)
| fetch_data.s(job_id=job_id)
| set_job_metadata.s(
extended_metadata_collection="extended" in tool.app.config.metadata_strategy,
job_id=job_id,
).set(link_error=finish_job.si(job_id=job_id, raw_tool_source=raw_tool_source))
| finish_job.si(job_id=job_id, raw_tool_source=raw_tool_source)
).set(
link_error=finish_job.si(
job_id=job_id,
raw_tool_source=raw_tool_source,
tool_source_class=tool_source_class,
)
)
| finish_job.si(job_id=job_id, raw_tool_source=raw_tool_source, tool_source_class=tool_source_class)
)()
job2.set_runner_external_id(async_result.task_id)
continue
@@ -273,6 +273,7 @@ class JobsService(ServiceBase):
tool_source = ToolSource(
raw_tool_source=tool.tool_source.to_string(),
tool_dir=tool.tool_dir,
tool_source_class=type(tool.tool_source).__name__,
)
task_request = QueueJobs(
user=trans.async_request_user,