Merge pull request #10539 from mvdbeek/apply_rules_performance

Improve database operation job submission performance
This commit is contained in:
Marius van den Beek
2020-10-26 15:14:46 +01:00
committed by GitHub
9 changed files with 126 additions and 85 deletions
+58 -8
View File
@@ -40,6 +40,7 @@ from sqlalchemy.orm import (
joinedload,
object_session,
Query,
reconstructor,
)
from sqlalchemy.schema import UniqueConstraint
@@ -1757,11 +1758,33 @@ class History(HasTags, Dictifiable, UsesAnnotations, HasName, RepresentById):
self.datasets = []
self.galaxy_sessions = []
self.tags = []
# Objects to eventually add to history
self._pending_additions = []
@reconstructor
def init_on_load(self):
# Restores properties that are not tracked in the database
self._pending_additions = []
def stage_addition(self, items):
history_id = self.id
for item in listify(items):
if history_id:
item.history_id = history_id
else:
item.history = self
self._pending_additions.append(item)
@property
def empty(self):
return self.hid_counter == 1
def add_pending_datasets(self, set_output_hid=True):
# These are assumed to be either copies of existing datasets or new, empty datasets,
# so we don't need to set the quota.
self.add_datasets(object_session(self), self._pending_additions, set_hid=set_output_hid, quota=False, flush=False)
self._pending_additions = []
def _next_hid(self, n=1):
# this is overriden in mapping.py db_next_hid() method
if len(self.datasets) == 0:
@@ -2284,6 +2307,11 @@ class StorableObject:
else:
self.uuid = UUID(str(uuid))
def flush(self):
sa_session = object_session(self)
if sa_session:
sa_session.flush()
class Dataset(StorableObject, RepresentById):
states = Bunch(NEW='new',
@@ -4109,14 +4137,36 @@ class DatasetCollection(Dictifiable, UsesAnnotations, RepresentById):
@property
def dataset_instances(self):
instances = []
for element in self.elements:
if element.is_collection:
instances.extend(element.child_collection.dataset_instances)
else:
instance = element.dataset_instance
instances.append(instance)
return instances
db_session = object_session(self)
if db_session and self.id:
dc = alias(DatasetCollection.table)
de = alias(DatasetCollectionElement.table)
hda = alias(HistoryDatasetAssociation.table)
depth_collection_type = self.collection_type
select_from = dc.outerjoin(de, de.c.dataset_collection_id == dc.c.id)
while ":" in depth_collection_type:
child_collection = alias(DatasetCollection.table)
child_collection_element = alias(DatasetCollectionElement.table)
select_from = select_from.outerjoin(child_collection, child_collection.c.id == de.c.child_collection_id)
select_from = select_from.outerjoin(child_collection_element, child_collection_element.c.dataset_collection_id == child_collection.c.id)
de = child_collection_element
depth_collection_type = depth_collection_type.split(":", 1)[1]
select_from = select_from.outerjoin(hda, hda.c.id == de.c.hda_id)
select_stmt = select([hda]).select_from(select_from).where(dc.c.id == self.id).distinct()
return db_session.query(HistoryDatasetAssociation).select_entity_from(select_stmt).all()
else:
# Sessionless context
instances = []
for element in self.elements:
if element.is_collection:
instances.extend(element.child_collection.dataset_instances)
else:
instance = element.dataset_instance
instances.append(instance)
return instances
@property
def dataset_elements(self):
-1
View File
@@ -283,7 +283,6 @@ class ModelPersistenceContext(metaclass=abc.ABCMeta):
association_name = f'__new_primary_file_{name}|{element_identifier_str}__'
self.add_output_dataset_association(association_name, dataset)
self.flush()
self.update_object_store_with_datasets(datasets=element_datasets['datasets'], paths=element_datasets['paths'], extra_files=element_datasets['extra_files'])
add_datasets_timer = ExecutionTimer()
self.add_datasets_to_history(element_datasets['datasets'])
+5 -1
View File
@@ -256,7 +256,11 @@ class BaseObjectStore(ObjectStore):
def _get_object_id(self, obj):
if hasattr(obj, self.store_by):
return getattr(obj, self.store_by)
obj_id = getattr(obj, self.store_by)
if obj_id is None:
obj.flush()
return obj.id
return obj_id
else:
# job's don't have uuids, so always use ID in this case when creating
# job working directories.
+26 -31
View File
@@ -27,7 +27,6 @@ from galaxy import (
from galaxy.job_execution import output_collect
from galaxy.managers.jobs import JobSearch
from galaxy.metadata import get_metadata_compute_strategy
from galaxy.model.tags import GalaxyTagHandler
from galaxy.tool_shed.util.repository_util import get_installed_repository
from galaxy.tool_shed.util.shed_util_common import set_image_paths
from galaxy.tool_util.deps import (
@@ -1647,9 +1646,8 @@ class Tool(Dictifiable):
)
job = rval[0]
out_data = rval[1]
if len(rval) == 4:
execution_slice.datasets_to_persist = rval[2]
execution_slice.history = rval[3]
if len(rval) > 2:
execution_slice.history = rval[2]
except (webob.exc.HTTPFound, exceptions.MessageException) as e:
# if it's a webob redirect exception, pass it up the stack
raise e
@@ -2757,36 +2755,33 @@ class DatabaseOperationTool(Tool):
return not self.require_dataset_ok
def check_inputs_ready(self, input_datasets, input_dataset_collections):
def check_dataset_instance(input_dataset):
if input_dataset.is_pending:
def check_dataset_state(state):
if state in model.Dataset.non_ready_states:
raise ToolInputsNotReadyException("An input dataset is pending.")
if self.require_dataset_ok:
if input_dataset.state != input_dataset.dataset.states.OK:
if state != model.Dataset.states.OK:
raise ValueError("Tool requires inputs to be in valid state.")
for input_dataset in input_datasets.values():
check_dataset_instance(input_dataset)
check_dataset_state(input_dataset.state)
for input_dataset_collection_pairs in input_dataset_collections.values():
for input_dataset_collection, _ in input_dataset_collection_pairs:
if not input_dataset_collection.collection.populated:
if not input_dataset_collection.collection.populated_optimized:
raise ToolInputsNotReadyException("An input collection is not populated.")
for dataset_instance in input_dataset_collection.dataset_instances:
check_dataset_instance(dataset_instance)
states, _ = input_dataset_collection.collection.dataset_states_and_extensions_summary
for state in states:
check_dataset_state(state)
def _add_datasets_to_history(self, history, elements, datasets_visible=False):
datasets = []
for element_object in elements:
if getattr(element_object, "history_content_type", None) == "dataset":
element_object.visible = datasets_visible
datasets.append(element_object)
history.stage_addition(element_object)
if datasets:
history.add_datasets(self.sa_session, datasets, set_hid=True)
def produce_outputs(self, trans, out_data, output_collections, incoming, history):
def produce_outputs(self, trans, out_data, output_collections, incoming, history, **kwds):
return self._outputs_dict()
def _outputs_dict(self):
@@ -2796,7 +2791,7 @@ class DatabaseOperationTool(Tool):
class UnzipCollectionTool(DatabaseOperationTool):
tool_type = 'unzip_collection'
def produce_outputs(self, trans, out_data, output_collections, incoming, history, tags=None):
def produce_outputs(self, trans, out_data, output_collections, incoming, history, tags=None, **kwds):
has_collection = incoming["input"]
if hasattr(has_collection, "element_type"):
# It is a DCE
@@ -2834,7 +2829,7 @@ class ZipCollectionTool(DatabaseOperationTool):
class BuildListCollectionTool(DatabaseOperationTool):
tool_type = 'build_list'
def produce_outputs(self, trans, out_data, output_collections, incoming, history, tags=None):
def produce_outputs(self, trans, out_data, output_collections, incoming, history, tags=None, **kwds):
new_elements = OrderedDict()
for i, incoming_repeat in enumerate(incoming["datasets"]):
@@ -2850,7 +2845,7 @@ class BuildListCollectionTool(DatabaseOperationTool):
class ExtractDatasetCollectionTool(DatabaseOperationTool):
tool_type = 'extract_dataset'
def produce_outputs(self, trans, out_data, output_collections, incoming, history, tags=None):
def produce_outputs(self, trans, out_data, output_collections, incoming, history, tags=None, **kwds):
has_collection = incoming["input"]
if hasattr(has_collection, "element_type"):
# It is a DCE
@@ -3153,16 +3148,17 @@ class RelabelFromFileTool(DatabaseOperationTool):
class ApplyRulesTool(DatabaseOperationTool):
tool_type = 'apply_rules'
def produce_outputs(self, trans, out_data, output_collections, incoming, history, **kwds):
def produce_outputs(self, trans, out_data, output_collections, incoming, history, tag_handler, **kwds):
hdca = incoming["input"]
rule_set = RuleSet(incoming["rules"])
copied_datasets = []
def copy_dataset(dataset, tags):
copied_dataset = dataset.copy(flush=False)
copied_datasets.append(copied_dataset)
if tags is not None:
self.app.tag_handler.set_tags_from_list(trans.get_user(), copied_dataset, tags)
tag_handler.set_tags_from_list(trans.get_user(), copied_dataset, tags, flush=False)
copied_dataset.history_id = history.id
copied_datasets.append(copied_dataset)
return copied_dataset
new_elements = self.app.dataset_collections_service.apply_rules(
@@ -3170,19 +3166,18 @@ class ApplyRulesTool(DatabaseOperationTool):
)
self._add_datasets_to_history(history, copied_datasets)
output_collections.create_collection(
next(iter(self.outputs.values())), "output", collection_type=rule_set.collection_type, elements=new_elements
next(iter(self.outputs.values())), "output", collection_type=rule_set.collection_type, elements=new_elements,
)
class TagFromFileTool(DatabaseOperationTool):
tool_type = 'tag_from_file'
def produce_outputs(self, trans, out_data, output_collections, incoming, history, **kwds):
def produce_outputs(self, trans, out_data, output_collections, incoming, history, tag_handler, **kwds):
hdca = incoming["input"]
how = incoming['how']
new_tags_dataset_assoc = incoming["tags"]
new_elements = OrderedDict()
tags_manager = GalaxyTagHandler(trans.app.model.context)
new_datasets = []
def add_copied_value_to_new_elements(new_tags_dict, dce):
@@ -3195,13 +3190,13 @@ class TagFromFileTool(DatabaseOperationTool):
if new_tags:
if how in ('add', 'remove') and dce.element_object.tags:
# We need get the original tags and update them with the new tags
old_tags = {tag for tag in tags_manager.get_tags_str(dce.element_object.tags).split(',') if tag}
old_tags = {tag for tag in tag_handler.get_tags_str(dce.element_object.tags).split(',') if tag}
if how == 'add':
old_tags.update(set(new_tags))
elif how == 'remove':
old_tags = old_tags - set(new_tags)
new_tags = old_tags
tags_manager.add_tags_from_list(user=history.user, item=copied_value, new_tags_list=new_tags)
tag_handler.add_tags_from_list(user=history.user, item=copied_value, new_tags_list=new_tags, flush=False)
else:
# We have a collection, and we copy the elements so that we don't manipulate the original tags
copied_value = dce.element_object.copy(element_destination=history)
@@ -3211,14 +3206,14 @@ class TagFromFileTool(DatabaseOperationTool):
new_element.element_object.visible = False
new_tags = new_tags_dict.get(new_element.element_identifier)
if how in ('add', 'remove'):
old_tags = {tag for tag in tags_manager.get_tags_str(old_element.element_object.tags).split(',') if tag}
old_tags = {tag for tag in tag_handler.get_tags_str(old_element.element_object.tags).split(',') if tag}
if new_tags:
if how == 'add':
old_tags.update(set(new_tags))
elif how == 'remove':
old_tags = old_tags - set(new_tags)
new_tags = old_tags
tags_manager.add_tags_from_list(user=history.user, item=new_element.element_object, new_tags_list=new_tags)
tag_handler.add_tags_from_list(user=history.user, item=new_element.element_object, new_tags_list=new_tags, flush=False)
new_elements[dce.element_identifier] = copied_value
new_tags_path = new_tags_dataset_assoc.file_name
@@ -3268,10 +3263,10 @@ class FilterFromFileTool(DatabaseOperationTool):
discarded_elements[element_identifier] = copied_value
self._add_datasets_to_history(history, filtered_elements.values())
self._add_datasets_to_history(history, discarded_elements.values())
output_collections.create_collection(
self.outputs["output_filtered"], "output_filtered", elements=filtered_elements
)
self._add_datasets_to_history(history, discarded_elements.values())
output_collections.create_collection(
self.outputs["output_discarded"], "output_discarded", elements=discarded_elements
)
+27 -31
View File
@@ -177,16 +177,23 @@ class DefaultToolAction:
for action, role_id in action_tuples:
record_permission(action, role_id)
replace_collection = False
_, extensions = collection.dataset_states_and_extensions_summary
conversion_required = False
for ext in extensions:
if ext:
datatype = trans.app.datatypes_registry.get_datatype_by_extension(ext)
if not datatype.matches_any(input.formats):
conversion_required = True
break
processed_dataset_dict = {}
for i, v in enumerate(collection.dataset_instances):
processed_dataset = process_dataset(v)
if processed_dataset is not v:
replace_collection = True
processed_dataset_dict[v] = processed_dataset
input_datasets[prefix + input.name + str(i + 1)] = processed_dataset
if replace_collection:
processed_dataset = None
if conversion_required:
processed_dataset = process_dataset(v)
if processed_dataset is not v:
processed_dataset_dict[v] = processed_dataset
input_datasets[prefix + input.name + str(i + 1)] = processed_dataset or v
if conversion_required:
collection_type_description = trans.app.dataset_collections_service.collection_type_descriptions.for_collection_type(collection.collection_type)
collection_builder = CollectionBuilder(collection_type_description)
collection_builder.replace_elements_in_collection(
@@ -245,12 +252,9 @@ class DefaultToolAction:
def _collect_inputs(self, tool, trans, incoming, history, current_user_roles, collection_info):
""" Collect history as well as input datasets and collections. """
app = trans.app
# Set history.
if not history:
history = tool.get_default_history_by_trans(trans, create=True)
if history not in trans.sa_session:
history = trans.sa_session.query(app.model.History).get(history.id)
# Track input dataset collections - but replace with simply lists so collect
# input datasets can process these normally.
@@ -327,8 +331,7 @@ class DefaultToolAction:
if not completed_job:
# Determine output dataset permission/roles list
existing_datasets = [inp for inp in inp_data.values() if inp]
if existing_datasets:
if all_permissions:
output_permissions = app.security_agent.guess_derived_permissions(all_permissions)
else:
# No valid inputs, we will use history defaults
@@ -455,7 +458,6 @@ class DefaultToolAction:
# Flush all datasets at once.
return data
datasets_to_persist = []
for name, output in tool.outputs.items():
if not filter_output(output, incoming):
handle_output_timer = ExecutionTimer()
@@ -492,7 +494,7 @@ class DefaultToolAction:
effective_output_name = output_part_def.effective_output_name
element = handle_output(effective_output_name, output_part_def.output_def, hidden=True)
datasets_to_persist.append(element)
history.stage_addition(element)
# TODO: this shouldn't exist in the top-level of the history at all
# but for now we are still working around that by hiding the contents
# there.
@@ -509,17 +511,13 @@ class DefaultToolAction:
element_kwds = dict(elements=collections_manager.ELEMENTS_UNINITIALIZED)
else:
element_kwds = dict(element_identifiers=element_identifiers)
hdca = output_collections.create_collection(
output_collections.create_collection(
output=output,
name=name,
set_hid=True if flush_job else False,
flush=flush_job,
completed_job=completed_job,
**element_kwds
)
if hdca:
datasets_to_persist.append(hdca)
log.info(f"Handled collection output named {name} for tool {tool.id} {handle_output_timer}")
log.info("Handled collection output named {} for tool {} {}".format(name, tool.id, handle_output_timer))
else:
handle_output(name, output)
log.info(f"Handled output named {name} for tool {tool.id} {handle_output_timer}")
@@ -531,7 +529,7 @@ class DefaultToolAction:
# Add all the top-level (non-child) datasets to the history unless otherwise specified
for name, data in out_data.items():
if name not in child_dataset_names and name not in incoming: # don't add children; or already existing datasets, i.e. async created
datasets_to_persist.append(data)
history.stage_addition(data)
# Add all the children to their parents
for parent_name, child_name in parent_to_child_pairs:
@@ -556,14 +554,13 @@ class DefaultToolAction:
if app.config.track_jobs_in_database and rerun_remap_job_id is not None:
# We need a flush here and get hids in order to rewrite jobs parameter,
# but remapping jobs should only affect single jobs anyway, so this is not too costly.
history.add_pending_datasets(set_output_hid=set_output_hid)
trans.sa_session.flush()
history.add_datasets(trans.sa_session, datasets_to_persist, set_hid=set_output_hid, quota=False, flush=False)
self._remap_job_on_rerun(trans=trans,
galaxy_session=galaxy_session,
rerun_remap_job_id=rerun_remap_job_id,
current_job=job,
out_data=out_data)
datasets_to_persist = []
log.info(f"Setup for job {job.log_str()} complete, ready to be enqueued {job_setup_timer}")
# Some tools are not really executable, but jobs are still created for them ( for record keeping ).
@@ -589,8 +586,7 @@ class DefaultToolAction:
else:
if flush_job:
# Set HID and add to history.
# This is brand new and certainly empty so don't worry about quota.
history.add_datasets(trans.sa_session, datasets_to_persist, set_hid=set_output_hid, quota=False, flush=False)
history.add_pending_datasets(set_output_hid=set_output_hid)
job_flush_timer = ExecutionTimer()
trans.sa_session.flush()
log.info(f"Flushed transaction for job {job.log_str()} {job_flush_timer}")
@@ -598,7 +594,7 @@ class DefaultToolAction:
# Dispatch to a job handler. enqueue() is responsible for flushing the job
app.job_manager.enqueue(job, tool=tool)
trans.log_event("Added job to the job queue, id: %s" % str(job.id), tool_id=job.tool_id)
return job, out_data, datasets_to_persist, history
return job, out_data, history
def _remap_job_on_rerun(self, trans, galaxy_session, rerun_remap_job_id, current_job, out_data):
"""
@@ -821,7 +817,7 @@ class OutputCollections:
self.out_collection_instances = {}
self.tags = tags
def create_collection(self, output, name, collection_type=None, set_hid=True, flush=True, completed_job=None, **element_kwds):
def create_collection(self, output, name, collection_type=None, completed_job=None, **element_kwds):
input_collections = self.input_collections
collections_manager = self.trans.app.dataset_collections_service
collection_type = collection_type or output.structure.collection_type
@@ -900,16 +896,16 @@ class OutputCollections:
collection_type=collection_type,
trusted_identifiers=True,
tags=self.tags,
set_hid=set_hid,
flush=flush,
set_hid=False,
flush=False,
completed_job=completed_job,
output_name=name,
**element_kwds
)
# name here is name of the output element - not name
# of the hdca.
self.history.stage_addition(hdca)
self.out_collection_instances[name] = hdca
return hdca
def on_text_for_names(input_names):
+3 -4
View File
@@ -62,16 +62,16 @@ class ModelOperationToolAction(DefaultToolAction):
self._record_outputs(job, out_data, output_collections)
job.state = job.states.OK
trans.sa_session.add(job)
trans.sa_session.flush() # ensure job.id are available
# Queue the job for execution
# trans.app.job_manager.job_queue.put( job.id, tool.id )
# trans.log_event( "Added database job action to the job queue, id: %s" % str(job.id), tool_id=job.tool_id )
log.info("Calling produce_outputs, tool is %s" % tool)
return job, out_data
return job, out_data, history
def _produce_outputs(self, trans, tool, out_data, output_collections, incoming, history, tags):
tool.produce_outputs(trans, out_data, output_collections, incoming, history=history, tags=tags)
tag_handler = trans.app.tag_handler.create_tag_handler_session()
tool.produce_outputs(trans, out_data, output_collections, incoming, history=history, tags=tags, tag_handler=tag_handler)
mapped_over_elements = output_collections.dataset_collection_elements
if mapped_over_elements:
for name, value in out_data.items():
@@ -80,4 +80,3 @@ class ModelOperationToolAction(DefaultToolAction):
mapped_over_elements[name].hda = value
trans.sa_session.add_all(out_data.values())
trans.sa_session.flush()
+4 -7
View File
@@ -90,7 +90,7 @@ def execute(trans, tool, mapping_params, history, rerun_remap_job_id=None, colle
jobs_executed = 0
has_remaining_jobs = False
datasets_to_persist = []
execution_slice = None
for i, execution_slice in enumerate(execution_tracker.new_execution_slices()):
if max_num_jobs and jobs_executed >= max_num_jobs:
@@ -100,12 +100,10 @@ def execute(trans, tool, mapping_params, history, rerun_remap_job_id=None, colle
execute_single_job(execution_slice, completed_jobs[i])
history = execution_slice.history or history
jobs_executed += 1
if execution_slice.datasets_to_persist:
datasets_to_persist.extend(execution_slice.datasets_to_persist)
if datasets_to_persist:
history.add_datasets(trans.sa_session, datasets_to_persist, set_hid=True, quota=False, flush=False)
# a side effect of history.add_datasets is a commit within db_next_hid (even with flush=False).
if execution_slice:
# a side effect of adding datasets to a history is a commit within db_next_hid (even with flush=False).
history.add_pending_datasets()
else:
# Make sure collections, implicit jobs etc are flushed even if there are no precreated output datasets
trans.sa_session.flush()
@@ -130,7 +128,6 @@ class ExecutionSlice:
self.job_index = job_index
self.param_combination = param_combination
self.dataset_collection_elements = dataset_collection_elements
self.datasets_to_persist = None
self.history = None
+2 -1
View File
@@ -81,6 +81,7 @@ def test_job_context_discover_outputs_flushes_once(mocker):
final_job_state=job_context.final_job_state,
)
collection_builder.populate()
assert spy.call_count == 1
assert spy.call_count == 0
sa_session.flush()
assert len(collection.dataset_instances) == 10
assert collection.dataset_instances[0].dataset.file_size == 1
+1 -1
View File
@@ -134,7 +134,7 @@ class DefaultToolActionTestCase(unittest.TestCase, tools_support.UsesApp, tools_
if incoming is None:
incoming = dict(param1="moo")
self._init_tool(contents)
job, out_data, _, _ = self.action.execute(
job, out_data, _, = self.action.execute(
tool=self.tool,
trans=self.trans,
history=self.history,