diff --git a/client/galaxy/scripts/components/ToolRecommendation.vue b/client/galaxy/scripts/components/ToolRecommendation.vue new file mode 100644 index 00000000000..880f3aa93cb --- /dev/null +++ b/client/galaxy/scripts/components/ToolRecommendation.vue @@ -0,0 +1,241 @@ + + + diff --git a/client/galaxy/scripts/mvc/tool/tool-form.js b/client/galaxy/scripts/mvc/tool/tool-form.js index 0682b2f5840..59ced2add5e 100644 --- a/client/galaxy/scripts/mvc/tool/tool-form.js +++ b/client/galaxy/scripts/mvc/tool/tool-form.js @@ -13,6 +13,7 @@ import ToolFormBase from "mvc/tool/tool-form-base"; import Webhooks from "mvc/webhooks"; import Vue from "vue"; import ToolEntryPoints from "components/ToolEntryPoints/ToolEntryPoints"; +import ToolRecommendation from "components/ToolRecommendation"; const View = Backbone.View.extend({ initialize: function(options) { @@ -268,6 +269,19 @@ const View = Backbone.View.extend({ } } this.$el.append(this._templateSuccess(response, job_def)); + const enable_tool_recommendations = Galaxy.config.enable_tool_recommendations; + if (enable_tool_recommendations === true || enable_tool_recommendations === "true") { + // show tool recommendations + const ToolRecommendationInstance = Vue.extend(ToolRecommendation); + const vm = document.createElement("div"); + this.$el.append(vm); + const instance = new ToolRecommendationInstance({ + propsData: { + toolId: job_def.tool_id + } + }); + instance.$mount(vm); + } this.$el.parent().scrollTop(0); // Show Webhook if job is running if (response.jobs && response.jobs.length > 0) { diff --git a/client/galaxy/style/scss/base.scss b/client/galaxy/style/scss/base.scss index ed488b4f892..ac567172780 100644 --- a/client/galaxy/style/scss/base.scss +++ b/client/galaxy/style/scss/base.scss @@ -1650,3 +1650,38 @@ body.reports { bottom: 0; background: url(../../images/largespinner.gif) no-repeat center center fixed; } + +/* Used for tree in tool recommendations */ +.tool-recommendation-view { + .node { + cursor: pointer; + circle { + fill: lighten($brand-primary, 20%); + stroke: lighten($brand-primary, 20%); + stroke-width: 0.3rem; + } + text { + font: 0.75rem sans-serif; + } + } + + .node-enter { + fill-opacity: 1e-6; + } + + .node-update { + fill-opacity: 1; + } + + .tree-size { + width: 96%; + height: 50%; + position: absolute; + } + + .link { + fill: none; + stroke: lighten($brand-primary, 20%); + stroke-width: 0.3rem; + } +} diff --git a/config/tool_recommendations_overwrite.yml.sample b/config/tool_recommendations_overwrite.yml.sample new file mode 100644 index 00000000000..b1fc424a766 --- /dev/null +++ b/config/tool_recommendations_overwrite.yml.sample @@ -0,0 +1,48 @@ +# Provide a list of tools which are deprecated. These tools would be removed the list of +# recommended tools by the deep learning model. These tools will be removed from the recommendations and a warning message, set as the 'text_message', +# is shown instead of their recommendations and when they are executed. +# Format: +# <>: +# - is_deprecate: True +# - text_message: <> +# +# For the following example, the tool 'cufflinks' is deprecated. It will be removed from the recommendations and a warning message, set as the 'text_message', +# is shown instead of its recommendations and when it is executed: +# +# cufflinks: +# - is_deprecated: True +# text_message: It is deprecated +# +# +# +# Provide list of tools to be recommended. These tools will either be appended to the tools recommended by deep learning model or +# completely overwrite them with these tools. +# Format: +# <>: +# - tool_id: <> +# name: <> +# i_extensions: +# - <> +# - <> +# - tool_id: <> +# name: <> +# i_extensions: +# - <> +# - <> +# <>: +# - tool_id ... +# +# For the following example, the tools 'cat1' and 'sort1' are shown as the recommendations for tool 'Filter1': +# +# Filter1: +# - tool_id: 'cat1' +# name: 'Concatenate datasets tail-to-head' +# i_extensions: +# - tabular +# - txt +# - tool_id: 'sort1' +# name: 'Sort data in ascending or descending order' +# i_extensions: +# - tabular +# - txt + diff --git a/doc/source/admin/galaxy_options.rst b/doc/source/admin/galaxy_options.rst index 185f18428c9..6d387d8ee36 100644 --- a/doc/source/admin/galaxy_options.rst +++ b/doc/source/admin/galaxy_options.rst @@ -4042,4 +4042,67 @@ :Type: int +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ +``enable_tool_recommendations`` +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +:Description: + Allow the display of tool recommendations in workflow editor and + after tool execution. If it is enabled and set to true, please + enable 'tool_recommendation_model_path' as well +:Default: ``false`` +:Type: bool + + +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ +``tool_recommendation_model_path`` +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +:Description: + Set remote path of the trained model (HDF5 file) for tool + recommendation. +:Default: ``https://github.com/galaxyproject/galaxy-test-data/raw/master/tool_recommendation_model.hdf5`` +:Type: str + + +~~~~~~~~~~~~~~~~~~~~~~~~ +``topk_recommendations`` +~~~~~~~~~~~~~~~~~~~~~~~~ + +:Description: + Set the number of predictions/recommendations to be made by the + model +:Default: ``20`` +:Type: int + + +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ +``admin_tool_recommendations_path`` +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +:Description: + Set path to the additional tool preferences from Galaxy admins. It + has two blocks. One for listing deprecated tools which will be + removed from the recommendations and another is for adding + additional tools to be recommended along side those from the deep + learning model. +:Default: ``tool_recommendations_overwrite.yml`` +:Type: str + + +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ +``overwrite_model_recommendations`` +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +:Description: + Overwrite or append to the tool recommendations by the deep + learning model. When set to true, all the recommendations by the + deep learning model are overwritten by the recommendations set by + an admin in a config file 'tool_recommendations_overwrite.yml'. + When set to false, the recommended tools by admins and predicted + by the deep learning model are shown. +:Default: ``false`` +:Type: bool + + diff --git a/lib/galaxy/config/sample/galaxy.yml.sample b/lib/galaxy/config/sample/galaxy.yml.sample index 03a187c83be..5c6eefc5d46 100644 --- a/lib/galaxy/config/sample/galaxy.yml.sample +++ b/lib/galaxy/config/sample/galaxy.yml.sample @@ -1977,3 +1977,31 @@ galaxy: # as threshold (above threshold: regular select fields will be used) #select_type_workflow_threshold: -1 + # Allow the display of tool recommendations in workflow editor and + # after tool execution. If it is enabled and set to true, please + # enable 'tool_recommendation_model_path' as well + #enable_tool_recommendations: false + + # Set remote path of the trained model (HDF5 file) for tool + # recommendation. + #tool_recommendation_model_path: https://github.com/galaxyproject/galaxy-test-data/raw/master/tool_recommendation_model.hdf5 + + # Set the number of predictions/recommendations to be made by the + # model + #topk_recommendations: 20 + + # Set path to the additional tool preferences from Galaxy admins. It + # has two blocks. One for listing deprecated tools which will be + # removed from the recommendations and another is for adding + # additional tools to be recommended along side those from the deep + # learning model. + #admin_tool_recommendations_path: tool_recommendations_overwrite.yml + + # Overwrite or append to the tool recommendations by the deep learning + # model. When set to true, all the recommendations by the deep + # learning model are overwritten by the recommendations set by an + # admin in a config file 'tool_recommendations_overwrite.yml'. When + # set to false, the recommended tools by admins and predicted by the + # deep learning model are shown. + #overwrite_model_recommendations: false + diff --git a/lib/galaxy/dependencies/__init__.py b/lib/galaxy/dependencies/__init__.py index 4d398f0f230..328e553cb73 100644 --- a/lib/galaxy/dependencies/__init__.py +++ b/lib/galaxy/dependencies/__init__.py @@ -194,6 +194,12 @@ class ConditionalDependencies(object): def check_influxdb(self): return 'influxdb' in self.error_report_modules + def check_keras(self): + return asbool(self.config["enable_tool_recommendations"]) + + def check_tensorflow(self): + return asbool(self.config["enable_tool_recommendations"]) + def optional(config_file=None): if not config_file: diff --git a/lib/galaxy/dependencies/conditional-requirements.txt b/lib/galaxy/dependencies/conditional-requirements.txt index b57b56bfcbb..b1f96df9525 100644 --- a/lib/galaxy/dependencies/conditional-requirements.txt +++ b/lib/galaxy/dependencies/conditional-requirements.txt @@ -28,3 +28,7 @@ watchdog python-gitlab pygithub influxdb + +# Deep learning packages for tool recommendation +keras==2.2.4 +tensorflow==1.12.2 diff --git a/lib/galaxy/managers/configuration.py b/lib/galaxy/managers/configuration.py index 6df4ace6bbe..4812e33091f 100644 --- a/lib/galaxy/managers/configuration.py +++ b/lib/galaxy/managers/configuration.py @@ -62,6 +62,11 @@ class ConfigSerializer(base.ModelSerializer): 'communication_server_port' : _use_config, 'communication_server_host' : _use_config, 'persistent_communication_rooms' : _use_config, + 'enable_tool_recommendations' : _use_config, + 'tool_recommendation_model_path' : _use_config, + 'admin_tool_recommendations_path' : _use_config, + 'overwrite_model_recommendations' : _use_config, + 'topk_recommendations' : _use_config, 'allow_user_impersonation' : _use_config, 'allow_user_creation' : _defaults_to(False), # schema default is True 'use_remote_user' : _defaults_to(None), # schema default is False; or config.single_user diff --git a/lib/galaxy/tools/recommendations.py b/lib/galaxy/tools/recommendations.py new file mode 100644 index 00000000000..c8df15f36a7 --- /dev/null +++ b/lib/galaxy/tools/recommendations.py @@ -0,0 +1,219 @@ +""" Compute tool recommendations """ + +import json +import logging +import os + +import h5py +import numpy as np +import requests +import yaml + +from galaxy.tools.parameters import populate_state +from galaxy.tools.parameters.basic import workflow_building_modes +from galaxy.workflow.modules import module_factory + +log = logging.getLogger(__name__) + + +class ToolRecommendations(): + + def __init__(self): + self.tool_recommendation_model_path = None + self.admin_tool_recommendations_path = None + self.deprecated_tools = dict() + self.admin_recommendations = dict() + + def set_model(self, trans, remote_model_url): + """ + Create model and associated dictionaries for recommendations + """ + if not self.tool_recommendation_model_path: + # import moves from the top of file: in case the tool recommendation feature is disabled, + # keras is not downloaded because of conditional requirement and Galaxy does not build + try: + from keras.models import model_from_json + except Exception: + trans.response.status = 400 + return False + self.tool_recommendation_model_path = self.download_model(remote_model_url) + self.all_tools = dict() + model_weights = list() + counter_layer_weights = 0 + # collect ids and names of all the installed tools + for tool_id, tool in trans.app.toolbox.tools(): + t_id_renamed = tool_id + if t_id_renamed.find("/") > -1: + t_id_renamed = t_id_renamed.split("/")[-2] + self.all_tools[t_id_renamed] = (tool_id, tool.name) + # read the hdf5 attributes + trained_model = h5py.File(self.tool_recommendation_model_path, 'r') + model_config = json.loads(trained_model['model_config'][()]) + self.loaded_model = model_from_json(model_config) + # set the dictionary of tools + self.model_data_dictionary = json.loads(trained_model['data_dictionary'][()]) + self.reverse_dictionary = dict((v, k) for k, v in self.model_data_dictionary.items()) + # set the list of compatible tools + self.compatible_tools = json.loads(trained_model['compatible_tools'][()]) + self.tool_weights = json.loads(trained_model['class_weights'][()]) + self.tool_weights_sorted = dict() + # sort the tools' usage dictionary + tool_pos_sorted = [int(key) for key in self.tool_weights.keys()] + for k in tool_pos_sorted: + self.tool_weights_sorted[k] = self.tool_weights[str(k)] + # iterate through all the attributes of the model to find weights of neural network layers + for item in trained_model.keys(): + if "weight_" in item: + d_key = "weight_" + str(counter_layer_weights) + weights = trained_model[d_key][()] + model_weights.append(weights) + counter_layer_weights += 1 + # set the model weights + self.loaded_model.set_weights(model_weights) + return True + + def collect_admin_preferences(self, admin_path): + """ + Collect preferences for recommendations of tools + set by admins as dictionaries of deprecated tools and + additional recommendations + """ + if not self.admin_tool_recommendations_path and admin_path is not None: + self.admin_tool_recommendations_path = os.path.join(os.getcwd(), admin_path) + if os.path.exists(self.admin_tool_recommendations_path): + with open(self.admin_tool_recommendations_path) as admin_recommendations: + admin_recommendation_preferences = yaml.safe_load(admin_recommendations) + if admin_recommendation_preferences: + for tool_id in admin_recommendation_preferences: + tool_info = admin_recommendation_preferences[tool_id] + if 'is_deprecated' in tool_info[0]: + self.deprecated_tools[tool_id] = tool_info[0]["text_message"] + else: + if tool_id not in self.admin_recommendations: + self.admin_recommendations[tool_id] = tool_info + + def download_model(self, model_url, download_local='database/'): + """ + Download the model from remote server + """ + local_dir = os.path.join(os.getcwd(), download_local, 'tool_recommendation_model.hdf5') + # read model from remote + model_binary = requests.get(model_url) + # save model to a local directory + with open(local_dir, 'wb') as model_file: + model_file.write(model_binary.content) + return local_dir + + def get_tool_extensions(self, trans, tool_id): + """ + Get the input and output extensions of a tool + """ + payload = {'type': 'tool', 'tool_id': tool_id, '_': 'true'} + inputs = payload.get('inputs', {}) + trans.workflow_building_mode = workflow_building_modes.ENABLED + module = module_factory.from_dict(trans, payload) + if 'tool_state' not in payload: + module_state = {} + populate_state(trans, module.get_inputs(), inputs, module_state, check=False) + module.recover_state(module_state) + inputs = module.get_all_inputs(connectable_only=True) + outputs = module.get_all_outputs() + input_extensions = list() + output_extensions = list() + for i_ext in inputs: + input_extensions.extend(i_ext['extensions']) + for o_ext in outputs: + output_extensions.extend(o_ext['extensions']) + return input_extensions, output_extensions + + def filter_tool_predictions(self, trans, prediction_data, tool_ids, tool_scores, last_tool_name): + """ + Filter tool predictions based on datatype compatibility and tool connections. + Add admin preferences to recommendations. + """ + last_compatible_tools = list() + if last_tool_name in self.compatible_tools: + last_compatible_tools = self.compatible_tools[last_tool_name].split(",") + prediction_data["is_deprecated"] = False + t_ids_scores = zip(tool_ids, tool_scores) + # form the payload of the predicted tools to be shown + for child, score in t_ids_scores: + c_dict = dict() + for t_id in self.all_tools: + # select the name and tool id if it is installed in Galaxy + if t_id == child and score > 0.0 and child in last_compatible_tools and child not in self.deprecated_tools: + full_tool_id = self.all_tools[t_id][0] + pred_input_extensions, _ = self.get_tool_extensions(trans, full_tool_id) + c_dict["name"] = self.all_tools[t_id][1] + " (" + str(score) + "%)" + c_dict["tool_id"] = full_tool_id + c_dict["i_extensions"] = list(set(pred_input_extensions)) + prediction_data["children"].append(c_dict) + break + # incorporate preferences set by admins + if self.admin_tool_recommendations_path: + # filter out deprecated tools + t_ids_scores = [(tid, score) for tid, score in zip(tool_ids, tool_scores) if tid not in self.deprecated_tools] + # set the property if the last tool of the sequence is deprecated + if last_tool_name in self.deprecated_tools: + prediction_data["is_deprecated"] = True + prediction_data["message"] = self.deprecated_tools[last_tool_name] + # add the recommendations given by admins + for tool_id in self.admin_recommendations: + if last_tool_name == tool_id: + admin_recommendations = self.admin_recommendations[tool_id] + if trans.app.config.overwrite_model_recommendations is True: + prediction_data["children"] = admin_recommendations + else: + prediction_data["children"].extend(admin_recommendations) + break + # get the root name for displaying after tool run + for t_id in self.all_tools: + if t_id == last_tool_name: + prediction_data["name"] = self.all_tools[t_id][1] + break + return prediction_data + + def compute_tool_prediction(self, trans, tool_sequence): + """ + Compute the predicted tools for a tool sequences + Return a payload with the tool sequences and recommended tools + Return an empty payload with just the tool sequence if anything goes wrong within the try block + """ + max_seq_len = 25 + topk = trans.app.config.topk_recommendations + prediction_data = dict() + tool_sequence = tool_sequence.split(",")[::-1] + prediction_data["name"] = ",".join(tool_sequence) + prediction_data["children"] = list() + last_tool_name = tool_sequence[-1] + # do prediction only if the last is present in the collections of tools + if last_tool_name in self.model_data_dictionary: + sample = np.zeros(max_seq_len) + # get the list of datatype extensions of the last tool of the tool sequence + _, last_output_extensions = self.get_tool_extensions(trans, self.all_tools[last_tool_name][0]) + prediction_data["o_extensions"] = list(set(last_output_extensions)) + # get tool names without slashes and create a sequence vector + for idx, tool_name in enumerate(tool_sequence): + if tool_name.find("/") > -1: + tool_name = tool_name.split("/")[-2] + sample[idx] = int(self.model_data_dictionary[tool_name]) + sample = np.reshape(sample, (1, max_seq_len)) + # predict next tools for a test path + prediction = self.loaded_model.predict(sample, verbose=0) + prediction = np.reshape(prediction, (prediction.shape[1],)) + # boost the predicted scores using tools' usage + weight_values = list(self.tool_weights_sorted.values()) + prediction = prediction * weight_values + # normalize the predicted scores with max and sort the predictions + max_prediction = float(np.max(prediction)) + if max_prediction == 0.0: + max_prediction = 1.0 + prediction = prediction / max_prediction + prediction_pos = np.argsort(prediction, axis=-1) + # get topk prediction + topk_prediction_pos = prediction_pos[-topk:] + # read tool names using reverse dictionary + pred_tool_ids = [self.reverse_dictionary[int(tool_pos)] for tool_pos in topk_prediction_pos] + predicted_scores = [int(prediction[pos] * 100) for pos in topk_prediction_pos] + prediction_data = self.filter_tool_predictions(trans, prediction_data, pred_tool_ids[::-1], predicted_scores[::-1], last_tool_name) + return prediction_data diff --git a/lib/galaxy/webapps/galaxy/api/workflows.py b/lib/galaxy/webapps/galaxy/api/workflows.py index e015c234709..9cfe322b53f 100644 --- a/lib/galaxy/webapps/galaxy/api/workflows.py +++ b/lib/galaxy/webapps/galaxy/api/workflows.py @@ -26,6 +26,7 @@ from galaxy.managers import ( from galaxy.managers.jobs import fetch_job_states, invocation_job_source_iter from galaxy.model.item_attrs import UsesAnnotations from galaxy.tool_shed.galaxy_install.install_manager import InstallRepositoryManager +from galaxy.tools import recommendations from galaxy.tools.parameters import populate_state from galaxy.tools.parameters.basic import workflow_building_modes from galaxy.util.sanitize_html import sanitize_html @@ -57,6 +58,7 @@ class WorkflowsAPIController(BaseAPIController, UsesStoredWorkflowMixin, UsesAnn self.history_manager = histories.HistoryManager(app) self.workflow_manager = workflows.WorkflowsManager(app) self.workflow_contents_manager = workflows.WorkflowContentsManager(app) + self.tool_recommendations = recommendations.ToolRecommendations() def __get_full_shed_url(self, url): for name, shed_url in self.app.tool_shed_registry.tool_sheds.items(): @@ -631,6 +633,37 @@ class WorkflowsAPIController(BaseAPIController, UsesStoredWorkflowMixin, UsesAnn 'post_job_actions' : module.get_post_job_actions(inputs) } + @expose_api + def get_tool_predictions(self, trans, payload, **kwd): + """ + POST /api/workflows/get_tool_predictions + Fetch predicted tools for a workflow + :type payload: dict + :param payload: a dictionary containing two parameters: + 'tool_sequence' - comma separated sequence of tool ids + 'remote_model_url' - (optional) path to the deep learning model + """ + remote_model_url = payload.get('remote_model_url', None) + if remote_model_url is None: + remote_model_url = trans.app.config.tool_recommendation_model_path + if 'tool_sequence' not in payload or remote_model_url is None: + return + is_model_set = True + recommended_tools = dict() + tool_sequence = "" + # collect tool recommendation preferences if set by admin + self.tool_recommendations.collect_admin_preferences(trans.app.config.admin_tool_recommendations_path) + # recreate the neural network model to be used for prediction + is_model_set = self.tool_recommendations.set_model(trans, remote_model_url) + if is_model_set is True: + # get the recommended tools for a tool sequence + tool_sequence = payload.get('tool_sequence', "") + recommended_tools = self.tool_recommendations.compute_tool_prediction(trans, tool_sequence) + return { + "current_tool": tool_sequence, + "predicted_data": recommended_tools + } + # # -- Helper methods -- # diff --git a/lib/galaxy/webapps/galaxy/buildapp.py b/lib/galaxy/webapps/galaxy/buildapp.py index f998ca73765..5438ee12b66 100644 --- a/lib/galaxy/webapps/galaxy/buildapp.py +++ b/lib/galaxy/webapps/galaxy/buildapp.py @@ -404,6 +404,7 @@ def populate_api_routes(webapp, app): webapp.mapper.connect('/api/container_resolvers/{index}/toolbox', action="resolve_toolbox", controller="container_resolution", conditions=dict(method=["GET"])) webapp.mapper.connect('/api/container_resolvers/{index}/resolve/install', action="resolve_with_install", controller="container_resolution", conditions=dict(method=["POST"])) webapp.mapper.connect('/api/container_resolvers/{index}/toolbox/install', action="resolve_toolbox_with_install", controller="container_resolution", conditions=dict(method=["POST"])) + webapp.mapper.connect('/api/workflows/get_tool_predictions', action='get_tool_predictions', controller="workflows", conditions=dict(method=["POST"])) webapp.mapper.resource_with_deleted('user', 'users', path_prefix='/api') webapp.mapper.resource('genome', 'genomes', path_prefix='/api') diff --git a/lib/galaxy/webapps/galaxy/config_schema.yml b/lib/galaxy/webapps/galaxy/config_schema.yml index 3d5f3bdb657..84c37488ec8 100644 --- a/lib/galaxy/webapps/galaxy/config_schema.yml +++ b/lib/galaxy/webapps/galaxy/config_schema.yml @@ -2982,3 +2982,44 @@ mapping: use 0 in order to always use select2 fields, use -1 (default) in order to always use the regular select fields, use any other positive number as threshold (above threshold: regular select fields will be used) + + enable_tool_recommendations: + type: bool + default: false + required: false + desc: | + Allow the display of tool recommendations in workflow editor and after tool execution. + If it is enabled and set to true, please enable 'tool_recommendation_model_path' as well + + tool_recommendation_model_path: + type: str + default: 'https://github.com/galaxyproject/galaxy-test-data/raw/master/tool_recommendation_model.hdf5' + required: false + desc: | + Set remote path of the trained model (HDF5 file) for tool recommendation. + + topk_recommendations: + type: int + default: 20 + required: false + desc: | + Set the number of predictions/recommendations to be made by the model + + admin_tool_recommendations_path: + type: str + required: false + default: 'tool_recommendations_overwrite.yml' + path_resolves_to: config_dir + desc: | + Set path to the additional tool preferences from Galaxy admins. + It has two blocks. One for listing deprecated tools which will be removed from the recommendations and + another is for adding additional tools to be recommended along side those from the deep learning model. + + overwrite_model_recommendations: + type: bool + default: false + required: false + desc: | + Overwrite or append to the tool recommendations by the deep learning model. When set to true, all the recommendations by the deep learning model + are overwritten by the recommendations set by an admin in a config file 'tool_recommendations_overwrite.yml'. When set to false, the recommended tools + by admins and predicted by the deep learning model are shown. diff --git a/lib/galaxy_test/api/test_workflows.py b/lib/galaxy_test/api/test_workflows.py index a40ab737f8a..bb0a662896a 100644 --- a/lib/galaxy_test/api/test_workflows.py +++ b/lib/galaxy_test/api/test_workflows.py @@ -331,6 +331,29 @@ class WorkflowsApiTestCase(BaseWorkflowsApiTestCase): self._assert_user_has_workflow_with_name(name) return upload_response + def test_get_tool_predictions(self): + request = {"tool_sequence": "Cut1", "remote_model_url": "https://github.com/galaxyproject/galaxy-test-data/raw/master/tool_recommendation_model.hdf5"} + actual_recommendations = ['Filter1', 'cat1', 'addValue', 'comp1', 'Grep1'] + route = "workflows/get_tool_predictions" + response = self._post(route, data=request) + recommendation_response = response.json() + is_empty = bool(recommendation_response["current_tool"]) + if is_empty is False: + self._assert_status_code_is(response, 400) + else: + # check Ok response from the API + self._assert_status_code_is(response, 200) + recommendation_response = response.json() + # check the input tool sequence + assert recommendation_response["current_tool"] == request["tool_sequence"] + # check non-empty predictions list + predicted_tools = recommendation_response["predicted_data"]["children"] + assert len(predicted_tools) > 0 + # check for the correct predictions + for tool in predicted_tools: + assert tool["tool_id"] in actual_recommendations + break + def test_update(self): original_workflow = self.workflow_populator.load_workflow(name="test_import") uuids = {}