From 559ba46302a57b8a6621bd14ea2656228bb4e886 Mon Sep 17 00:00:00 2001 From: davelopez <46503462+davelopez@users.noreply.github.com> Date: Thu, 18 Feb 2021 14:33:09 +0100 Subject: [PATCH] Refactor error handling logic --- lib/galaxy/managers/groups.py | 100 ++++++++++-------------- lib/galaxy/webapps/galaxy/api/groups.py | 5 +- 2 files changed, 46 insertions(+), 59 deletions(-) diff --git a/lib/galaxy/managers/groups.py b/lib/galaxy/managers/groups.py index 6150b539a3f..1c85418b2a0 100644 --- a/lib/galaxy/managers/groups.py +++ b/lib/galaxy/managers/groups.py @@ -6,6 +6,13 @@ from typing import ( from sqlalchemy import false from galaxy import model +from galaxy.app import StructuredApp +from galaxy.exceptions import ( + Conflict, + ObjectAttributeMissingException, + ObjectNotFound, +) +from galaxy.managers.base import decode_id from galaxy.managers.context import ProvidesAppContext from galaxy.schema.fields import EncodedDatabaseIdField from galaxy.web import url_for @@ -14,102 +21,81 @@ from galaxy.web import url_for class GroupsManager: """Interface/service object shared by controllers for interacting with groups.""" + def __init__(self, app: StructuredApp) -> None: + self._app = app + def index(self, trans: ProvidesAppContext): """ Displays a collection (list) of groups. """ rval = [] for group in trans.sa_session.query(model.Group).filter(model.Group.table.c.deleted == false()): - if trans.user_is_admin: - item = group.to_dict(value_mapper={'id': trans.security.encode_id}) - encoded_id = trans.security.encode_id(group.id) - item['url'] = url_for('group', id=encoded_id) - rval.append(item) + item = group.to_dict(value_mapper={'id': trans.security.encode_id}) + encoded_id = trans.security.encode_id(group.id) + item['url'] = url_for('group', id=encoded_id) + rval.append(item) return rval def create(self, trans: ProvidesAppContext, payload: Dict[str, Any]): """ Creates a new group. """ - if not trans.user_is_admin: - trans.response.status = 403 - return "You are not authorized to create a new group." name = payload.get('name', None) - if not name: - trans.response.status = 400 - return "Enter a valid name" - if trans.sa_session.query(model.Group).filter(model.Group.table.c.name == name).first(): - trans.response.status = 400 - return "A group with that name already exists" + if name is None: + raise ObjectAttributeMissingException("Missing required name") + self._check_duplicated_group_name(trans, name) group = model.Group(name=name) trans.sa_session.add(group) user_ids = payload.get('user_ids', []) - users = [trans.sa_session.query(model.User).get(trans.security.decode_id(i)) for i in user_ids] + users = [trans.sa_session.query(model.User).get(self._decode_id(i)) for i in user_ids] role_ids = payload.get('role_ids', []) - roles = [trans.sa_session.query(model.Role).get(trans.security.decode_id(i)) for i in role_ids] + roles = [trans.sa_session.query(model.Role).get(self._decode_id(i)) for i in role_ids] trans.app.security_agent.set_entity_group_associations(groups=[group], roles=roles, users=users) - """ - # Create the UserGroupAssociations - for user in users: - trans.app.security_agent.associate_user_group( user, group ) - # Create the GroupRoleAssociations - for role in roles: - trans.app.security_agent.associate_group_role( group, role ) - """ trans.sa_session.flush() + encoded_id = trans.security.encode_id(group.id) item = group.to_dict(view='element', value_mapper={'id': trans.security.encode_id}) item['url'] = url_for('group', id=encoded_id) return [item] - def show(self, trans: ProvidesAppContext, id: EncodedDatabaseIdField): + def show(self, trans: ProvidesAppContext, encoded_id: EncodedDatabaseIdField): """ Displays information about a group. """ - group_id = id - try: - decoded_group_id = trans.security.decode_id(group_id) - except TypeError: - trans.response.status = 400 - return "Malformed group id ( %s ) specified, unable to decode." % str(group_id) - try: - group = trans.sa_session.query(model.Group).get(decoded_group_id) - except Exception: - group = None - if not group: - trans.response.status = 400 - return "Invalid group id ( %s ) specified." % str(group_id) + group = self._get_group(trans, encoded_id) item = group.to_dict(view='element', value_mapper={'id': trans.security.encode_id}) - item['url'] = url_for('group', id=group_id) - item['users_url'] = url_for('group_users', group_id=group_id) - item['roles_url'] = url_for('group_roles', group_id=group_id) + item['url'] = url_for('group', id=encoded_id) + item['users_url'] = url_for('group_users', group_id=encoded_id) + item['roles_url'] = url_for('group_roles', group_id=encoded_id) return item - def update(self, trans: ProvidesAppContext, id: EncodedDatabaseIdField, payload: Dict[str, Any]): + def update(self, trans: ProvidesAppContext, encoded_id: EncodedDatabaseIdField, payload: Dict[str, Any]): """ Modifies a group. """ - group_id = id - try: - decoded_group_id = trans.security.decode_id(group_id) - except TypeError: - trans.response.status = 400 - return "Malformed group id ( %s ) specified, unable to decode." % str(group_id) - try: - group = trans.sa_session.query(model.Group).get(decoded_group_id) - except Exception: - group = None - if not group: - trans.response.status = 400 - return "Invalid group id ( %s ) specified." % str(group_id) + group = self._get_group(trans, encoded_id) name = payload.get('name', None) if name: group.name = name trans.sa_session.add(group) user_ids = payload.get('user_ids', []) - users = [trans.sa_session.query(model.User).get(trans.security.decode_id(i)) for i in user_ids] + users = [trans.sa_session.query(model.User).get(self._decode_id(i)) for i in user_ids] role_ids = payload.get('role_ids', []) - roles = [trans.sa_session.query(model.Role).get(trans.security.decode_id(i)) for i in role_ids] + roles = [trans.sa_session.query(model.Role).get(self._decode_id(i)) for i in role_ids] trans.app.security_agent.set_entity_group_associations(groups=[group], roles=roles, users=users, delete_existing_assocs=False) trans.sa_session.flush() + + def _decode_id(self, encoded_id: EncodedDatabaseIdField) -> int: + return decode_id(self._app, encoded_id) + + def _check_duplicated_group_name(self, trans: ProvidesAppContext, group_name: str) -> None: + if trans.sa_session.query(model.Group).filter(model.Group.table.c.name == group_name).first(): + raise Conflict(f"A group with name '{group_name}' already exists") + + def _get_group(self, trans: ProvidesAppContext, encoded_id: EncodedDatabaseIdField) -> model.Group: + decoded_group_id = self._decode_id(encoded_id) + group = trans.sa_session.query(model.Group).get(decoded_group_id) + if group is None: + raise ObjectNotFound(f"Group with id {encoded_id} was not found.") + return group diff --git a/lib/galaxy/webapps/galaxy/api/groups.py b/lib/galaxy/webapps/galaxy/api/groups.py index 43386c22b93..9433dab6ff6 100644 --- a/lib/galaxy/webapps/galaxy/api/groups.py +++ b/lib/galaxy/webapps/galaxy/api/groups.py @@ -7,6 +7,7 @@ from typing import ( Dict, ) +from galaxy.app import StructuredApp from galaxy.managers.context import ProvidesAppContext from galaxy.managers.groups import GroupsManager from galaxy.schema.fields import EncodedDatabaseIdField @@ -21,9 +22,9 @@ log = logging.getLogger(__name__) class GroupAPIController(BaseAPIController): - def __init__(self, app): + def __init__(self, app: StructuredApp): super().__init__(app) - self.manager = GroupsManager() + self.manager = GroupsManager(app) @expose_api @require_admin