Merge pull request #23358 from mvdbeek/lazy-file-source-oauth-tokens

Serialize only file sources required by a job
This commit is contained in:
Marius van den Beek
2026-08-25 16:41:33 +02:00
committed by GitHub
22 changed files with 621 additions and 33 deletions
+37 -6
View File
@@ -51,6 +51,8 @@ class ProvidesFileSourcesTransaction(Protocol):
log = logging.getLogger(__name__)
USER_FILE_SOURCES_SCHEME = "gxuserfiles"
class FileSourcePath(NamedTuple):
file_source: BaseFilesSource
@@ -70,8 +72,7 @@ class UserDefinedFileSources(Protocol):
"""Entry-point for Galaxy to inject user-defined file sources.
Supplied object of this class is used to write out concrete
description of file sources when serializing all file sources
available to a user.
descriptions of user file sources selected for serialization.
"""
def validate_uri_root(self, uri: str, user_context: "FileSourcesUserContext") -> None:
@@ -87,6 +88,7 @@ class UserDefinedFileSources(Protocol):
browsable_only: bool | None = False,
include_kind: set[PluginKind] | None = None,
exclude_kind: set[PluginKind] | None = None,
referenced_uris: set[str] | None = None,
) -> list[dict[str, Any]]:
"""Write out user file sources as list of config dictionaries."""
# config_dicts: List[FilesSourceProperties] = []
@@ -111,6 +113,7 @@ class NullUserDefinedFileSources(UserDefinedFileSources):
browsable_only: bool | None = False,
include_kind: set[PluginKind] | None = None,
exclude_kind: set[PluginKind] | None = None,
referenced_uris: set[str] | None = None,
) -> list[dict[str, Any]]:
return []
@@ -209,6 +212,15 @@ class ConfiguredFileSources:
def _parse_plugin_source(self, plugin_source: PluginConfigSource):
return self._plugin_loader.load_plugins(plugin_source, self._file_sources_config)
@staticmethod
def _best_score(scores: list[FileSourceScore]) -> FileSourceScore | None:
best = max(scores, key=lambda candidate: candidate.score, default=None)
return best if best is not None and best.score > 0 else None
def _best_configured_match(self, url: str) -> FileSourceScore | None:
scores = [FileSourceScore(file_source, file_source.score_url_match(url)) for file_source in self._file_sources]
return self._best_score(scores)
def find_best_match(self, url: str) -> BaseFilesSource | None:
"""Returns the best matching file source for handling a particular url. Each filesource scores its own
ability to match a particular url, and the highest scorer with a score > 0 is selected."""
@@ -216,8 +228,8 @@ class ConfiguredFileSources:
user_best_score = self._user_defined_file_sources.find_best_match(url)
if user_best_score is not None:
scores.append(user_best_score)
scores.sort(key=lambda f: f.score, reverse=True)
return next((fsscore.file_source for fsscore in scores if fsscore.score > 0), None)
best = self._best_score(scores)
return best.file_source if best is not None else None
def get_file_source_path(self, uri):
"""Parse uri into a FileSource object and a path relative to its base."""
@@ -280,8 +292,14 @@ class ConfiguredFileSources:
browsable_only: bool | None = False,
include_kind: set[PluginKind] | None = None,
exclude_kind: set[PluginKind] | None = None,
referenced_uris: set[str] | None = None,
) -> list[dict[str, Any]]:
rval: list[dict[str, Any]] = []
referenced_file_sources = None
if referenced_uris is not None:
referenced_file_sources = [
match.file_source for uri in referenced_uris if (match := self._best_configured_match(uri)) is not None
]
for file_source in self._file_sources:
if not file_source.user_has_access(user_context):
continue
@@ -291,6 +309,9 @@ class ConfiguredFileSources:
continue
if browsable_only and not file_source.get_browsable():
continue
# Skip sources that are not the best match for any URI required by this job.
if referenced_file_sources is not None and file_source not in referenced_file_sources:
continue
el = file_source.to_dict(for_serialization=for_serialization, user_context=user_context)
rval.append(el)
if user_context:
@@ -301,13 +322,23 @@ class ConfiguredFileSources:
browsable_only=browsable_only,
include_kind=include_kind,
exclude_kind=exclude_kind,
referenced_uris=referenced_uris,
)
)
return rval
def to_dict(self, for_serialization: bool = False, user_context: "OptionalUserContext" = None) -> dict[str, Any]:
def to_dict(
self,
for_serialization: bool = False,
user_context: "OptionalUserContext" = None,
referenced_uris: set[str] | None = None,
) -> dict[str, Any]:
return {
"file_sources": self.plugins_to_dict(for_serialization=for_serialization, user_context=user_context),
"file_sources": self.plugins_to_dict(
for_serialization=for_serialization,
user_context=user_context,
referenced_uris=referenced_uris,
),
"config": self._file_sources_config.to_dict(),
}
-7
View File
@@ -178,16 +178,9 @@ class eLabFTWFilesSource(BaseFilesSource[eLabFTWFileSourceTemplateConfiguration,
def get_prefix(self) -> str | None:
endpoint: ParseResult = self._get_endpoint()
return self.id if self.scheme not in {"elabftw", DEFAULT_SCHEME} else (endpoint.netloc or None)
# it would make better sense to return
# `self.id if self.scheme == USER_FILE_SOURCES_SCHEME else (endpoint.netloc or None)`, where
# `USER_FILE_SOURCES_SCHEME` comes from `galaxy.managers.file_source_instances`; however, that would lead to a
# circular import (maybe `USER_FILE_SOURCES_SCHEME` should be moved to a module in a layer deeper than
# `galaxy.managers`)
def get_scheme(self) -> str:
return self.scheme if self.scheme and self.scheme != DEFAULT_SCHEME else "elabftw"
# it would make better sense to return `self.scheme if self.scheme == USER_FILE_SOURCES_SCHEME else "elabftw"`,
# but the same circular import issue as above arises
def score_url_match(self, url: str) -> int:
parsed_url = urlparse(url)
+40 -11
View File
@@ -106,7 +106,10 @@ from galaxy.tools.evaluation import (
ToolEvaluator,
UserToolEvaluator,
)
from galaxy.tools.parameters import params_to_json_internal
from galaxy.tools.parameters import (
collect_directory_uris,
params_to_json_internal,
)
from galaxy.util import (
parse_xml_string,
RWXRWXRWX,
@@ -1081,6 +1084,22 @@ class MinimalJobWrapper(HasResourceParameters):
if authnz_manager and trans.user:
authnz_manager.refresh_expiring_oidc_tokens(trans, trans.user)
def _referenced_file_source_uris(self, job: Job) -> set[str] | None:
"""Return required URIs, or ``None`` when action discovery is incomplete."""
uris: set[str] = set()
if self.tool is not None:
tool_action = self.tool.tool_action
if not tool_action.has_complete_file_source_uri_discovery():
return None
param_dict = self.get_param_dict(job)
uris.update(collect_directory_uris(self.tool.inputs, param_dict))
uris.update(tool_action.iter_referenced_file_source_uris(param_dict))
for input_association in job.input_datasets + job.input_library_datasets:
dataset = input_association.dataset
if dataset is not None and dataset.has_deferred_data and dataset.dataset is not None:
uris.update(dataset.dataset.source_uris)
return uris
@property
def job_io(self) -> JobIO:
if self._job_io is None:
@@ -1088,6 +1107,7 @@ class MinimalJobWrapper(HasResourceParameters):
work_request = WorkRequestContext(self.app, user=job.user, galaxy_session=job.galaxy_session)
user_context = ProvidesFileSourcesUserContext(work_request)
self._refresh_oidc_tokens_for_job(work_request)
referenced_uris = self._referenced_file_source_uris(job)
tool_source = self.tool.tool_source.to_string() if self.tool else None
tool_dir = self.tool.tool_dir if self.tool else None
self._job_io = JobIO(
@@ -1106,7 +1126,11 @@ class MinimalJobWrapper(HasResourceParameters):
new_file_path=self.app.config.new_file_path,
builds_file_path=self.app.config.builds_file_path,
len_file_path=self.app.config.len_file_path,
file_sources_dict=self.app.file_sources.to_dict(for_serialization=True, user_context=user_context),
file_sources_dict=self.app.file_sources.to_dict(
for_serialization=True,
user_context=user_context,
referenced_uris=referenced_uris,
),
user_context=user_context,
check_job_script_integrity=self.app.config.check_job_script_integrity,
check_job_script_integrity_count=self.app.config.check_job_script_integrity_count,
@@ -1240,15 +1264,18 @@ class MinimalJobWrapper(HasResourceParameters):
return os.path.abspath(os.path.join(self.working_directory, "outputs", COMMAND_VERSION_FILENAME))
def __prepare_upload_paramfile(self, job):
"""Special case paramfile handling for the upload tool. Copies the paramfile to the working directory"""
"""Copy the upload paramfile into the working directory and use the stable path."""
new = os.path.join(self.working_directory, "upload_params.json")
param_file_path = json.loads(next(iter(param.value for param in job.parameters if param.name == "paramfile")))
try:
shutil.copy2(param_file_path, new)
except OSError as exc:
# It won't exist at the old path if setup was interrupted and tried again later
if exc.errno != errno.ENOENT or not os.path.exists(new):
raise
paramfile_parameter = next(iter(param for param in job.parameters if param.name == "paramfile"))
param_file_path = json.loads(paramfile_parameter.value)
if param_file_path != new:
try:
shutil.copy2(param_file_path, new)
except OSError as exc:
# It won't exist at the old path if setup was interrupted and tried again later
if exc.errno != errno.ENOENT or not os.path.exists(new):
raise
paramfile_parameter.value = json.dumps(new)
def prepare(self, compute_environment=None):
"""
@@ -1413,7 +1440,9 @@ class MinimalJobWrapper(HasResourceParameters):
return tool_evaluator
def _fix_output_permissions(self):
for path in [dp.real_path for dp in self.job_io.get_mutable_output_fnames()]:
if self._job_io is None:
return
for path in [dp.real_path for dp in self._job_io.get_mutable_output_fnames()]:
if os.path.exists(path):
util.umask_fix_perms(path, self.app.config.umask, 0o666, self.app.config.gid)
+32 -3
View File
@@ -4,6 +4,7 @@ from typing import (
cast,
Literal,
)
from urllib.parse import urlsplit
from uuid import uuid4
from pydantic import (
@@ -26,6 +27,7 @@ from galaxy.files import (
FileSourceScore,
FileSourcesUserContext,
ProvidesFileSourcesUserContext,
USER_FILE_SOURCES_SCHEME,
UserDefinedFileSources,
)
from galaxy.files.plugins import (
@@ -53,6 +55,7 @@ from galaxy.files.templates.capabilities import (
)
from galaxy.managers.context import ProvidesUserContext
from galaxy.model import (
get_uuid,
User,
UserFileSource,
)
@@ -111,7 +114,21 @@ from ._config_templates import (
log = logging.getLogger(__name__)
USER_FILE_SOURCES_SCHEME = "gxuserfiles"
def referenced_user_file_source_ids(referenced_uris: set[str]) -> set[str]:
"""Return canonical UUID hex strings addressed by ``gxuserfiles`` URIs."""
ids: set[str] = set()
for uri in referenced_uris:
if not uri.startswith(f"{USER_FILE_SOURCES_SCHEME}://"):
continue
try:
split = urlsplit(uri)
if not split.netloc:
raise ValueError("URI has no authority")
ids.add(get_uuid(split.netloc).hex)
except ValueError as exc:
raise RequestParameterInvalidException(f"Invalid user file source URI [{uri}]") from exc
return ids
class UserFileSourceModel(BaseModel):
@@ -675,15 +692,24 @@ class UserDefinedFileSourcesImpl(UserDefinedFileSources):
)[0]
return file_source
def _all_user_file_source_properties(self, user_context: FileSourcesUserContext) -> list[dict[str, Any]]:
def _all_user_file_source_properties(
self,
user_context: FileSourcesUserContext,
referenced_uris: set[str] | None = None,
) -> list[dict[str, Any]]:
username_filter = User.__table__.c.username == user_context.username
user: User | None = self._sa_session.query(User).filter(username_filter).one_or_none()
if user is None:
return []
referenced_ids = None if referenced_uris is None else referenced_user_file_source_ids(referenced_uris)
all_file_source_properties: list[dict[str, Any]] = []
for user_file_source in user.file_sources:
if user_file_source.hidden:
continue
# Filter before resolving properties because resolution can access the vault or mint
# an OAuth access token.
if referenced_ids is not None and get_uuid(user_file_source.uuid).hex not in referenced_ids:
continue
try:
files_source_properties = self._file_source_properties(user_file_source)
except ValidationError:
@@ -736,13 +762,16 @@ class UserDefinedFileSourcesImpl(UserDefinedFileSources):
browsable_only: bool | None = False,
include_kind: set[PluginKind] | None = None,
exclude_kind: set[PluginKind] | None = None,
referenced_uris: set[str] | None = None,
) -> list[dict[str, Any]]:
"""Write out user file sources as list of config dictionaries."""
if user_context.anonymous:
return []
as_dicts = []
for files_source_properties in self._all_user_file_source_properties(user_context):
for files_source_properties in self._all_user_file_source_properties(
user_context, referenced_uris=referenced_uris
):
files_source_type = files_source_properties["type"]
plugin_type_class = self._plugin_loader.get_plugin_type_class(files_source_type)
plugin_kind = plugin_type_class.plugin_kind
+5
View File
@@ -4858,6 +4858,11 @@ class Dataset(Base, StorableObject, Serializable):
def is_new(self):
return self.state == self.states.NEW
@property
def source_uris(self) -> list[str]:
"""The URIs this dataset was populated from (e.g. remote/deferred sources)."""
return [source.source_uri for source in self.sources if source.source_uri]
def in_ready_state(self):
return self.state in self.ready_states
+11
View File
@@ -4,6 +4,7 @@ import os
import re
from abc import abstractmethod
from collections.abc import (
Iterable,
Mapping,
MutableMapping,
)
@@ -98,6 +99,15 @@ class ToolAction:
"""
produces_real_jobs: bool
file_source_uri_discovery_complete = False
def has_complete_file_source_uri_discovery(self) -> bool:
"""Return whether this concrete action has audited URI discovery."""
return type(self).__dict__.get("file_source_uri_discovery_complete", False)
def iter_referenced_file_source_uris(self, param_dict: ToolStateJobInstancePopulatedT) -> Iterable[str]:
"""Yield file source URIs embedded in action-specific parameters."""
return ()
@abstractmethod
def execute(
@@ -141,6 +151,7 @@ class DefaultToolAction(ToolAction):
"""Default tool action is to run an external command"""
produces_real_jobs: bool = True
file_source_uri_discovery_complete = True
def _collect_input_datasets(
self,
+2
View File
@@ -29,6 +29,8 @@ log = logging.getLogger(__name__)
class DataManagerToolAction(DefaultToolAction):
"""Tool action used for Data Manager Tools"""
file_source_uri_discovery_complete = True
def execute(
self,
tool,
+2
View File
@@ -8,6 +8,8 @@ log = logging.getLogger(__name__)
class DataSourceToolAction(DefaultToolAction):
"""Tool action used for Data Source Tools"""
file_source_uri_discovery_complete = True
def _get_default_data_name(
self, dataset, tool, on_text=None, trans=None, incoming=None, history=None, params=None, job_params=None, **kwd
):
@@ -2,6 +2,7 @@ import datetime
import logging
import os
import tempfile
from collections.abc import Iterable
from galaxy.job_execution.setup import JobWorkingDirectory
from galaxy.model import (
@@ -40,6 +41,13 @@ class ImportHistoryToolAction(ToolAction):
"""Tool action used for importing a history to an archive."""
produces_real_jobs: bool = True
file_source_uri_discovery_complete = True
def iter_referenced_file_source_uris(self, param_dict: ToolStateJobInstancePopulatedT) -> Iterable[str]:
if param_dict.get("__ARCHIVE_TYPE__") == "url":
archive_source = param_dict.get("__ARCHIVE_SOURCE__")
if isinstance(archive_source, str) and archive_source:
yield archive_source
def execute(
self,
@@ -117,6 +125,7 @@ class ExportHistoryToolAction(ToolAction):
"""Tool action used for exporting a history to an archive."""
produces_real_jobs: bool = True
file_source_uri_discovery_complete = True
def execute(
self,
+1
View File
@@ -40,6 +40,7 @@ class SetMetadataToolAction(ToolAction):
produces_real_jobs: bool = False
set_output_hid: bool = False
file_source_uri_discovery_complete = True
def execute(
self,
@@ -40,6 +40,7 @@ log = logging.getLogger(__name__)
class ModelOperationToolAction(DefaultToolAction):
produces_real_jobs: bool = False
file_source_uri_discovery_complete = True
def check_inputs_ready(self, tool, trans, incoming, history, execution_cache=None, collection_info=None):
if execution_cache is None:
+31 -1
View File
@@ -1,8 +1,12 @@
import json
import logging
import os
from collections.abc import Iterable
from galaxy.exceptions import RequestParameterMissingException
from galaxy.exceptions import (
RequestParameterInvalidException,
RequestParameterMissingException,
)
from galaxy.job_execution.output_collect import copy_collection_metadata_from_target_dict
from galaxy.managers.context import ProvidesHistoryContext
from galaxy.model import (
@@ -16,6 +20,7 @@ from galaxy.model.dataset_collections.structure import UninitializedTree
from galaxy.schema.credentials import CredentialsContext
from galaxy.tools._types import ToolStateJobInstancePopulatedT
from galaxy.tools.actions import upload_common
from galaxy.tools.data_fetch_utils import iter_fetch_request_urls
from galaxy.tools.execute import (
DatasetCollectionElementsSliceT,
DEFAULT_DATASET_COLLECTION_ELEMENTS,
@@ -38,6 +43,7 @@ log = logging.getLogger(__name__)
class BaseUploadToolAction(ToolAction):
produces_real_jobs = True
file_source_uri_discovery_complete = False
def execute(
self,
@@ -86,6 +92,25 @@ class BaseUploadToolAction(ToolAction):
class UploadToolAction(BaseUploadToolAction):
file_source_uri_discovery_complete = True
def iter_referenced_file_source_uris(self, param_dict: ToolStateJobInstancePopulatedT) -> Iterable[str]:
paramfile = param_dict.get("paramfile")
if not isinstance(paramfile, str):
raise RequestParameterInvalidException("Legacy upload job is missing its paramfile")
with open(paramfile) as f:
upload_params = json.load(f)
if not isinstance(upload_params, list):
raise RequestParameterInvalidException("Legacy upload paramfile must contain a list")
for upload_param in upload_params:
if not isinstance(upload_param, dict):
raise RequestParameterInvalidException("Legacy upload paramfile entries must be objects")
if upload_param.get("type") == "url":
path = upload_param.get("path")
if not isinstance(path, str) or not path:
raise RequestParameterInvalidException("Legacy URL upload entry is missing its path")
yield path
def _setup_job(
self, tool, trans: ProvidesHistoryContext, incoming, dataset_upload_inputs, history, preferred_object_store_id
):
@@ -112,6 +137,11 @@ class UploadToolAction(BaseUploadToolAction):
class FetchUploadToolAction(BaseUploadToolAction):
file_source_uri_discovery_complete = True
def iter_referenced_file_source_uris(self, param_dict: ToolStateJobInstancePopulatedT) -> Iterable[str]:
return iter_fetch_request_urls(param_dict)
def _setup_job(
self, tool, trans: ProvidesHistoryContext, incoming, dataset_upload_inputs, history, preferred_object_store_id
):
+12
View File
@@ -1,3 +1,5 @@
import json
from collections.abc import Iterator
from datetime import (
datetime,
timezone,
@@ -6,6 +8,7 @@ from typing import Any
from galaxy.authnz.psa_authnz import locate_token_expiration
from galaxy.model import User
from galaxy.tools._types import ToolStateJobInstancePopulatedT
def iter_fetch_urls(value: Any):
@@ -19,6 +22,15 @@ def iter_fetch_urls(value: Any):
yield from iter_fetch_urls(child)
def iter_fetch_request_urls(param_dict: ToolStateJobInstancePopulatedT) -> Iterator[str]:
"""Yield URLs from a data-fetch tool's serialized request."""
request_json = param_dict.get("request_json")
if request_json:
for url in iter_fetch_urls(json.loads(request_json)):
if isinstance(url, str) and url:
yield url
def fetch_uses_authorization_header(request: dict[str, Any], file_sources, user_context) -> bool:
for url in iter_fetch_urls(request):
file_source_path = file_sources.get_file_source_path(url)
+16
View File
@@ -22,6 +22,7 @@ from .basic import (
ColumnListParameter,
DataCollectionToolParameter,
DataToolParameter,
DirectoryUriToolParameter,
ParameterValueError,
SelectToolParameter,
TextToolParameter,
@@ -286,6 +287,21 @@ def visit_input_values(
)
def collect_directory_uris(
inputs: ToolInputsT,
input_values: ToolStateJobInstancePopulatedT,
) -> set[str]:
"""Collect the values of every ``directory_uri`` parameter (file source write destinations)."""
uris: set[str] = set()
def _collect(input, value, **kwargs):
if isinstance(input, DirectoryUriToolParameter) and isinstance(value, str) and value:
uris.add(value)
visit_input_values(inputs, input_values, _collect)
return uris
def check_param(
trans, param: ToolParameter, incoming_value, param_values, simple_errors: bool = True
) -> tuple[Any, str | ValueError | None]:
+87
View File
@@ -1,10 +1,12 @@
import abc
import os
from contextlib import contextmanager
from types import SimpleNamespace
from typing import (
cast,
TYPE_CHECKING,
)
from uuid import uuid4
from galaxy.app_unittest_utils.tools_support import (
MockContext,
@@ -12,6 +14,7 @@ from galaxy.app_unittest_utils.tools_support import (
)
from galaxy.jobs import (
JobWrapper,
MinimalJobWrapper,
TaskWrapper,
)
from galaxy.jobs.handler import BaseJobHandlerQueue
@@ -23,6 +26,8 @@ from galaxy.model import (
)
from galaxy.objectstore import BaseObjectStore
from galaxy.tools import ToolBox
from galaxy.tools.parameters.basic import DirectoryUriToolParameter
from galaxy.util import XML
from galaxy.util.bunch import Bunch
from galaxy.util.unittest import TestCase
@@ -150,6 +155,14 @@ class MockTool:
self.home_target = None
self.tmp_target = None
self.tool_source = Bunch(to_string=lambda: "")
self.inputs = {}
self.tool_action = SimpleNamespace(
has_complete_file_source_uri_discovery=lambda: True,
iter_referenced_file_source_uris=lambda param_dict: (),
)
def params_from_strings(self, param_dict):
return param_dict
def get_job_destination(self, params):
return Bunch(runner="local", id="local", params={})
@@ -197,3 +210,77 @@ class MockObjectStore:
if kwds.get("base_dir", "") == "job_work":
return self.working_directory
return None
def _minimal_wrapper(param_dict=None, inputs=None, action_uris=(), action_discovery_complete=True):
tool_action = SimpleNamespace(
has_complete_file_source_uri_discovery=lambda: action_discovery_complete,
iter_referenced_file_source_uris=lambda param_dict: action_uris,
)
return SimpleNamespace(
tool=SimpleNamespace(inputs=inputs or {}, tool_action=tool_action),
get_param_dict=lambda job: param_dict or {},
)
def _job_with_file_source_inputs(input_datasets=None, input_library_datasets=None):
return SimpleNamespace(
id=1,
input_datasets=input_datasets or [],
input_library_datasets=input_library_datasets or [],
)
def test_referenced_file_source_uris_reads_tool_parameters_and_action():
destination = "gxfiles://good/out"
fetched = f"gxuserfiles://{uuid4().hex}/input"
destination_param = DirectoryUriToolParameter(None, XML('<param name="destination" type="directory_uri"/>'))
wrapper = _minimal_wrapper(
param_dict={"destination": destination},
inputs={"destination": destination_param},
action_uris=(fetched,),
)
assert MinimalJobWrapper._referenced_file_source_uris(wrapper, _job_with_file_source_inputs()) == {
destination,
fetched,
}
def test_referenced_file_source_uris_empty_for_job_without_sources():
assert MinimalJobWrapper._referenced_file_source_uris(_minimal_wrapper(), _job_with_file_source_inputs()) == set()
def test_referenced_file_source_uris_unknown_for_unaudited_action():
wrapper = _minimal_wrapper(action_discovery_complete=False)
assert MinimalJobWrapper._referenced_file_source_uris(wrapper, _job_with_file_source_inputs()) is None
def test_referenced_file_source_uris_adds_regular_and_library_input_sources():
regular_src = f"gxuserfiles://{uuid4().hex}/regular.txt"
library_src = f"gxuserfiles://{uuid4().hex}/library.txt"
hda = SimpleNamespace(has_deferred_data=True, dataset=SimpleNamespace(source_uris=[regular_src]))
ldda = SimpleNamespace(has_deferred_data=True, dataset=SimpleNamespace(source_uris=[library_src]))
job = _job_with_file_source_inputs(
input_datasets=[SimpleNamespace(dataset=hda)],
input_library_datasets=[SimpleNamespace(dataset=ldda)],
)
assert MinimalJobWrapper._referenced_file_source_uris(_minimal_wrapper(), job) == {regular_src, library_src}
def test_referenced_file_source_uris_ignores_materialized_input_sources():
source = f"gxuserfiles://{uuid4().hex}/materialized.txt"
hda = SimpleNamespace(has_deferred_data=False, dataset=SimpleNamespace(source_uris=[source]))
job = _job_with_file_source_inputs(input_datasets=[SimpleNamespace(dataset=hda)])
assert MinimalJobWrapper._referenced_file_source_uris(_minimal_wrapper(), job) == set()
def test_fix_output_permissions_does_not_initialize_job_io():
class WrapperWithoutJobIO:
_job_io = None
@property
def job_io(self):
raise AssertionError("job_io should not be initialized during failure cleanup")
wrapper = cast(MinimalJobWrapper, WrapperWithoutJobIO())
MinimalJobWrapper._fix_output_permissions(wrapper)
@@ -2,7 +2,10 @@ import os
from typing import (
cast,
)
from uuid import uuid4
from uuid import (
UUID,
uuid4,
)
import pytest
import responses
@@ -15,7 +18,12 @@ from galaxy.exceptions import (
RequestParameterInvalidException,
RequestParameterMissingException,
)
from galaxy.files import FileSourcesUserContext
from galaxy.files import (
ConfiguredFileSources,
FileSourcesUserContext,
USER_FILE_SOURCES_SCHEME,
)
from galaxy.files.plugins import FileSourcePluginsConfig
from galaxy.files.sources import dropbox
from galaxy.files.templates import ConfiguredFileSourceTemplates
from galaxy.files.templates.examples import get_example
@@ -24,13 +32,13 @@ from galaxy.managers.file_source_instances import (
CreateInstancePayload,
FileSourceInstancesManager,
ModifyInstancePayload,
referenced_user_file_source_ids,
TemplateFormDataRequest,
TestUpdateInstancePayload,
TestUpgradeInstancePayload,
UpdateInstancePayload,
UpdateInstanceSecretPayload,
UpgradeInstancePayload,
USER_FILE_SOURCES_SCHEME,
UserDefinedFileSourcesConfig,
UserDefinedFileSourcesImpl,
UserFileSourceModel,
@@ -326,6 +334,97 @@ class TestFileSourcesTestCase(BaseTestCase):
assert not status.connection.is_not_ok
assert fsspec_fs_init_kwd["token"] == "my_test_access_token"
def test_serialization_mints_no_tokens_when_no_sources_referenced(self, tmp_path, monkeypatch):
self._init_dropbox_env(tmp_path, monkeypatch)
self._create_dropbox_oauth_source(uuid4().hex)
self._create_dropbox_oauth_source(uuid4().hex)
calls = self._count_refresh_token_mints(monkeypatch)
as_dicts = self.file_sources.user_file_sources_to_dicts(
True,
cast(FileSourcesUserContext, self.trans),
referenced_uris=set(),
)
assert calls == []
assert as_dicts == []
def test_serialization_mints_only_referenced_source(self, tmp_path, monkeypatch):
self._init_dropbox_env(tmp_path, monkeypatch)
referenced_uuid = uuid4().hex
self._create_dropbox_oauth_source(referenced_uuid)
self._create_dropbox_oauth_source(uuid4().hex)
calls = self._count_refresh_token_mints(monkeypatch)
as_dicts = self.file_sources.user_file_sources_to_dicts(
True,
cast(FileSourcesUserContext, self.trans),
referenced_uris={f"gxuserfiles://{referenced_uuid}/some/path"},
)
assert calls == [f"refresh_token_{referenced_uuid}"]
assert len(as_dicts) == 1
assert UUID(as_dicts[0]["id"]).hex == referenced_uuid
assert as_dicts[0]["access_token"] == "my_test_access_token"
def test_serialization_without_reference_filter_includes_all_sources(self, tmp_path, monkeypatch):
self._init_dropbox_env(tmp_path, monkeypatch)
self._create_dropbox_oauth_source(uuid4().hex)
self._create_dropbox_oauth_source(uuid4().hex)
calls = self._count_refresh_token_mints(monkeypatch)
as_dicts = self.file_sources.user_file_sources_to_dicts(
True,
cast(FileSourcesUserContext, self.trans),
)
assert len(calls) == 2
assert len(as_dicts) == 2
def test_configured_file_sources_to_dict_threads_referenced_uris(self, tmp_path, monkeypatch):
self._init_dropbox_env(tmp_path, monkeypatch)
referenced_uuid = uuid4().hex
self._create_dropbox_oauth_source(referenced_uuid)
self._create_dropbox_oauth_source(uuid4().hex)
calls = self._count_refresh_token_mints(monkeypatch)
configured = ConfiguredFileSources(
FileSourcePluginsConfig(),
user_defined_file_sources=self.file_sources,
)
as_dict = configured.to_dict(
for_serialization=True,
user_context=cast(FileSourcesUserContext, self.trans),
referenced_uris={f"gxuserfiles://{referenced_uuid}/x"},
)
ids = [s["id"] for s in as_dict["file_sources"]]
assert calls == [f"refresh_token_{referenced_uuid}"]
assert len(ids) == 1
assert UUID(ids[0]).hex == referenced_uuid
def _create_dropbox_oauth_source(self, uuid: str) -> None:
config_secret_key = UserFileSource.vault_key_from_uuid(uuid, "_oauth2_refresh_token", None)
# Seed a per-source refresh token so a captured mint identifies which source minted.
self.trans.user_vault.write_secret(config_secret_key, f"refresh_token_{uuid}")
self._create_instance(
CreateInstancePayload(
name=SIMPLE_FILE_SOURCE_NAME,
description=SIMPLE_FILE_SOURCE_DESCRIPTION,
template_id="dropbox",
template_version=0,
variables={},
secrets={},
uuid=uuid,
)
)
def _count_refresh_token_mints(self, monkeypatch) -> list:
calls: list = []
class MockDropboxDriveFileSystem:
pass
def mock_get_token_from_refresh_raw(refresh_token, client_pair, config):
calls.append(refresh_token)
return MockResponse({"access_token": "my_test_access_token"})
monkeypatch.setattr(config_templates, "get_token_from_refresh_raw", mock_get_token_from_refresh_raw)
monkeypatch.setattr(dropbox.DropboxFilesSource, "required_module", MockDropboxDriveFileSystem)
return calls
def test_onedrive_oauth2_flow(self, tmp_path, monkeypatch):
json = {
"refresh_token": "my_test_refresh_token",
@@ -1197,3 +1296,33 @@ class OneDriveMockResponse:
def json(self):
return self._json_data
def test_referenced_user_file_source_ids_selects_only_user_sources():
user_uuid = uuid4().hex
uris = {
f"gxuserfiles://{user_uuid}/some/path",
"gxfiles://dropbox/other",
"https://example.com/file.txt",
"drs://example.org/abc",
}
assert referenced_user_file_source_ids(uris) == {user_uuid}
def test_referenced_user_file_source_ids_normalizes_dashed_uuid():
dashed = str(uuid4())
assert referenced_user_file_source_ids({f"gxuserfiles://{dashed}/x"}) == {dashed.replace("-", "")}
def test_referenced_user_file_source_ids_handles_no_match():
uris = {"gxfiles://dropbox/x", "not-a-uri", ""}
assert referenced_user_file_source_ids(uris) == set()
@pytest.mark.parametrize(
"uri",
["gxuserfiles://[invalid-authority/x", "gxuserfiles://not-a-uuid/x", "gxuserfiles:///x"],
)
def test_referenced_user_file_source_ids_rejects_invalid_user_source_uri(uri):
with pytest.raises(RequestParameterInvalidException, match="Invalid user file source URI"):
referenced_user_file_source_ids({uri})
@@ -0,0 +1,30 @@
from galaxy.tools.parameters import collect_directory_uris
from galaxy.tools.parameters.basic import (
DirectoryUriToolParameter,
TextToolParameter,
)
from galaxy.util import XML
def _directory_uri_param(name):
return DirectoryUriToolParameter(None, XML(f'<param name="{name}" type="directory_uri"/>'))
def _text_param(name):
return TextToolParameter(None, XML(f'<param name="{name}" type="text"/>'))
def test_collect_directory_uris_selects_only_directory_uri_params():
inputs = {"dest": _directory_uri_param("dest"), "other": _text_param("other")}
values = {"dest": "gxfiles://target/out", "other": "gxfiles://not-a-destination/x"}
assert collect_directory_uris(inputs, values) == {"gxfiles://target/out"}
def test_collect_directory_uris_empty_when_no_directory_uri_params():
inputs = {"other": _text_param("other")}
assert collect_directory_uris(inputs, {"other": "text"}) == set()
def test_collect_directory_uris_skips_empty_values():
inputs = {"dest": _directory_uri_param("dest")}
assert collect_directory_uris(inputs, {"dest": ""}) == set()
+45 -2
View File
@@ -1,12 +1,26 @@
import json
from datetime import (
datetime,
timedelta,
timezone,
)
from typing import cast
from typing import (
Any,
cast,
)
from galaxy.model import User
from galaxy.tools.data_fetch_utils import compute_token_expiry_for_provider
from galaxy.tools.actions.upload import FetchUploadToolAction
from galaxy.tools.data_fetch_utils import (
compute_token_expiry_for_provider,
iter_fetch_request_urls,
iter_fetch_urls,
)
class ConcreteFetchUploadToolAction(FetchUploadToolAction):
def get_output_name(self, *args: Any, **kwargs: Any) -> str:
raise NotImplementedError
class DummyToken:
@@ -62,3 +76,32 @@ def test_compute_token_expiry_for_provider_returns_none_when_token_missing_auth_
token.extra_data = {}
user = DummyUser([token])
assert compute_token_expiry_for_provider(cast(User, user), "oidc") is None
def test_iter_fetch_request_urls_extracts_urls():
request = {
"targets": [
{"elements": [{"src": "url", "url": "gxuserfiles://abc/x"}, {"src": "url", "url": "https://e.com/y"}]}
]
}
param_dict = {"request_json": json.dumps(request)}
assert set(iter_fetch_request_urls(param_dict)) == {"gxuserfiles://abc/x", "https://e.com/y"}
def test_iter_fetch_request_urls_empty_without_request_json():
assert list(iter_fetch_request_urls({})) == []
assert list(iter_fetch_request_urls({"request_json": ""})) == []
def test_iter_fetch_request_urls_ignores_non_string_and_empty_urls():
request = {"targets": [{"elements": [{"src": "url", "url": None}, {"src": "url", "url": ""}]}]}
assert list(iter_fetch_urls(request)) == [None, ""]
assert list(iter_fetch_request_urls({"request_json": json.dumps(request)})) == []
def test_fetch_upload_action_reports_referenced_file_source_uris():
request = {"targets": [{"elements": [{"src": "url", "url": "gxfiles://source/input"}]}]}
action = ConcreteFetchUploadToolAction()
assert list(action.iter_referenced_file_source_uris({"request_json": json.dumps(request)})) == [
"gxfiles://source/input"
]
@@ -3,6 +3,7 @@ import os
import tarfile
import tempfile
from shutil import rmtree
from typing import Any
from unittest.mock import Mock
from sqlalchemy import select
@@ -12,6 +13,7 @@ from galaxy.app_unittest_utils.galaxy_mock import MockApp
from galaxy.exceptions import MalformedContents
from galaxy.model.orm.util import add_object_to_object_session
from galaxy.objectstore.unittest_utils import Config as TestConfig
from galaxy.tools.actions.history_imp_exp import ImportHistoryToolAction
from galaxy.tools.imp_exp import (
JobExportHistoryArchiveWrapper,
JobImportHistoryArchiveWrapper,
@@ -28,6 +30,24 @@ HISTORY_ATTRS = """{"hid_counter": 2, "update_time": "2016-02-08 18:38:38.705058
JOBS_ATTRS = """[{"info": null, "tool_id": "upload1", "update_time": "2016-02-08T18:39:23.356482", "stdout": "", "input_mapping": {}, "tool_version": "1.1.4", "traceback": null, "command_line": "python /galaxy/tools/data_source/upload.py /galaxy /scratch/tmppwU9rD /scratch/tmpP4_45Y 1:/scratch/jobs/000/dataset_1_files:/data/000/dataset_1.dat", "exit_code": 0, "output_datasets": [1], "state": "ok", "create_time": "2016-02-08T18:38:39.153873", "params": {"files": [{"to_posix_lines": "Yes", "NAME": "None", "file_data": null, "space_to_tab": null, "url_paste": "/scratch/strio_url_paste_o6nrv8", "__index__": 0, "ftp_files": "", "uuid": "None"}], "paramfile": "/scratch/tmpP4_45Y", "file_type": "auto", "files_metadata": {"file_type": "auto", "__current_case__": 41}, "async_datasets": "None", "dbkey": "?"}, "stderr": ""}]"""
class ConcreteImportHistoryToolAction(ImportHistoryToolAction):
def get_output_name(self, *args: Any, **kwargs: Any) -> str:
raise NotImplementedError
def test_import_history_action_reports_url_archive_source():
action = ConcreteImportHistoryToolAction()
archive_source = "https://example.org/history.tar.gz"
assert list(
action.iter_referenced_file_source_uris({"__ARCHIVE_TYPE__": "url", "__ARCHIVE_SOURCE__": archive_source})
) == [archive_source]
assert not list(
action.iter_referenced_file_source_uris(
{"__ARCHIVE_TYPE__": "file", "__ARCHIVE_SOURCE__": "/tmp/history.tar.gz"}
)
)
def t_data_path(name):
return os.path.join(galaxy_directory(), "test-data", name)
@@ -0,0 +1,65 @@
import json
import tempfile
from typing import Any
import pytest
from galaxy.exceptions import RequestParameterInvalidException
from galaxy.tools.actions.upload import UploadToolAction
class ConcreteUploadToolAction(UploadToolAction):
def get_output_name(self, *args: Any, **kwargs: Any) -> str:
raise NotImplementedError
def test_upload_action_reports_urls_from_paramfile():
upload_params = [
{"type": "url", "path": "https://example.org/input.txt"},
{"type": "file", "path": "/tmp/pasted-input.txt"},
]
with tempfile.NamedTemporaryFile(mode="w") as paramfile:
json.dump(upload_params, paramfile)
paramfile.flush()
action = ConcreteUploadToolAction()
assert list(action.iter_referenced_file_source_uris({"paramfile": paramfile.name})) == [
"https://example.org/input.txt"
]
@pytest.mark.parametrize(
("upload_params", "message"),
[
({"type": "url", "path": "https://example.org/input.txt"}, "must contain a list"),
(["not an object"], "entries must be objects"),
([{"type": "url", "path": None}], "URL upload entry is missing its path"),
],
)
def test_upload_action_rejects_invalid_paramfile_shape(upload_params, message):
with tempfile.NamedTemporaryFile(mode="w") as paramfile:
json.dump(upload_params, paramfile)
paramfile.flush()
action = ConcreteUploadToolAction()
with pytest.raises(RequestParameterInvalidException, match=message):
list(action.iter_referenced_file_source_uris({"paramfile": paramfile.name}))
def test_upload_action_rejects_missing_paramfile_parameter():
action = ConcreteUploadToolAction()
with pytest.raises(RequestParameterInvalidException, match="missing its paramfile"):
list(action.iter_referenced_file_source_uris({}))
def test_upload_action_does_not_hide_missing_paramfile(tmp_path):
action = ConcreteUploadToolAction()
with pytest.raises(FileNotFoundError):
list(action.iter_referenced_file_source_uris({"paramfile": str(tmp_path / "missing.json")}))
def test_upload_action_does_not_hide_malformed_paramfile():
with tempfile.NamedTemporaryFile(mode="w") as paramfile:
paramfile.write("{")
paramfile.flush()
action = ConcreteUploadToolAction()
with pytest.raises(json.JSONDecodeError):
list(action.iter_referenced_file_source_uris({"paramfile": paramfile.name}))
+12
View File
@@ -40,6 +40,18 @@ def test_file_source_http_specific():
assert_realizes_as(file_sources, test_url, "hello specific world", user_context=user_context)
def test_plugins_to_dict_serializes_only_best_matching_http_source():
test_url = "https://www.usegalaxy.org/myfile.txt"
user_context = user_context_fixture()
file_sources = configured_file_sources(FILE_SOURCES_CONF)
plugins = file_sources.plugins_to_dict(
for_serialization=True,
user_context=user_context,
referenced_uris={test_url},
)
assert [plugin["id"] for plugin in plugins] == ["test1"]
def test_file_source_another_http_specific():
test_url = "http://www.galaxyproject.org/anotherfile.txt"
+31
View File
@@ -21,6 +21,7 @@ from galaxy.files.unittest_utils import (
from ._util import (
assert_realizes_as,
assert_realizes_throws_exception,
configured_file_sources,
find,
find_file_a,
list_dir,
@@ -556,3 +557,33 @@ def test_get_file_source_path_strips_whitespace():
resolved = file_sources.get_file_source_path("\ngxfiles://test1/a\n")
assert resolved.file_source is not None
assert resolved.path == "/a"
def _two_posix_file_sources(tmp_path):
root_good = tmp_path / "good"
root_other = tmp_path / "other"
root_good.mkdir()
root_other.mkdir()
return configured_file_sources(
[
{"type": "posix", "id": "good", "root": str(root_good)},
{"type": "posix", "id": "other", "root": str(root_other)},
]
)
def test_plugins_to_dict_serializes_only_referenced_sources(tmp_path):
file_sources = _two_posix_file_sources(tmp_path)
plugins = file_sources.plugins_to_dict(for_serialization=True, referenced_uris={"gxfiles://good/some/file"})
assert [p["id"] for p in plugins] == ["good"]
def test_plugins_to_dict_serializes_nothing_when_no_uris_referenced(tmp_path):
file_sources = _two_posix_file_sources(tmp_path)
assert file_sources.plugins_to_dict(for_serialization=True, referenced_uris=set()) == []
def test_plugins_to_dict_serializes_all_when_referenced_uris_none(tmp_path):
file_sources = _two_posix_file_sources(tmp_path)
plugins = file_sources.plugins_to_dict(for_serialization=True)
assert {p["id"] for p in plugins} == {"good", "other"}