Refactor error handling logic

This commit is contained in:
davelopez
2021-02-18 14:50:40 +01:00
parent 71214ae567
commit 559ba46302
2 changed files with 46 additions and 59 deletions
+43 -57
View File
@@ -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
+3 -2
View File
@@ -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