Commit d41fdb6b authored by Mickaël Desfrênes's avatar Mickaël Desfrênes
Browse files

optimize acl queries

parent 997f451c
Loading
Loading
Loading
Loading
+87 −16
Original line number Diff line number Diff line
@@ -57,7 +57,7 @@ def _climb_tree(
def _class_name_from_instance(object_instance: Union[str, Resource, Collection]) -> str:
    # if str, considered to be already a class name
    if type(object_instance) is str:
        return object_instance
        return str(object_instance)
    if isinstance(object_instance, Resource):
        return PERM_CLASS_RESOURCE
    if isinstance(object_instance, Collection):
@@ -231,6 +231,83 @@ class UserAccess:
        self.user = user
        self.project = project
        self._permissions = None
        self._class_permissions_by_class = {}
        self._object_permissions_by_key = {}
        self._has_object_permissions = False
        self._tree_permission_keys_cache = {}

    def _reset_permissions_cache_and_indexes(self):
        self._permissions = None
        self._class_permissions_by_class = {}
        self._object_permissions_by_key = {}
        self._has_object_permissions = False
        self._tree_permission_keys_cache = {}

    def _build_permission_indexes(self):
        self._class_permissions_by_class = {}
        self._object_permissions_by_key = {}
        self._has_object_permissions = False
        for permission in self._permissions:
            if permission.object_pk is None:
                self._class_permissions_by_class.setdefault(
                    permission.object_class, permission
                )
                continue

            self._has_object_permissions = True
            permission_key = (permission.object_class, permission.object_pk)
            self._object_permissions_by_key.setdefault(permission_key, []).append(
                permission
            )

    def _collection_tree_permission_keys(self, collection: Collection):
        collection_key = (PERM_CLASS_COLLECTION, collection.pk)
        if collection_key not in self._tree_permission_keys_cache:
            self._tree_permission_keys_cache[collection_key] = tuple(
                (_class_name_from_instance(node), node.pk)
                for node in _climb_tree(collection)
            )
        return self._tree_permission_keys_cache[collection_key]

    def _tree_permission_keys(self, object_instance: Union[Resource, Collection]):
        if isinstance(object_instance, Collection):
            return self._collection_tree_permission_keys(object_instance)

        permission_keys = [
            (_class_name_from_instance(object_instance), object_instance.pk)
        ]
        for collection in object_instance.available_collections():
            permission_keys.extend(self._collection_tree_permission_keys(collection))
        return tuple(permission_keys)

    def _has_crud_access(
        self,
        object_instance: Union[str, Resource, Collection],
        crud_access: str,
    ) -> bool:
        crud_access = str(crud_access).lower().strip()
        if crud_access not in CRUD_VERBS:
            raise ValueError(
                f"not a crud access right: {crud_access}. Must be one of {', '.join(CRUD_VERBS)}."
            )

        # Load permissions and build the derived indexes before lookup.
        _ = self.permissions
        object_class_name = _class_name_from_instance(object_instance)
        if self._has_object_permissions and isinstance(
            object_instance, Resource | Collection
        ):
            for permission_key in self._tree_permission_keys(object_instance):
                permissions = self._object_permissions_by_key.get(permission_key)
                if permissions:
                    return False not in [
                        _crud_val(permission, crud_access) for permission in permissions
                    ]

        class_permission = self._class_permissions_by_class.get(object_class_name)
        if class_permission is not None:
            return _crud_val(class_permission, crud_access)
        return False

    @property
    def permissions(self) -> List[ObjectPermission]:
@@ -246,38 +323,32 @@ class UserAccess:
                role__projectaccess__project=self.project,
            ):
                self._permissions.append(perm)
            self._build_permission_indexes()
        return self._permissions

    def can_crud(self, object_or_class, crud_access) -> bool:
        return (
            has_crud_access(self.permissions, object_or_class, crud_access)
            and self.user.is_active
            self._has_crud_access(object_or_class, crud_access) and self.user.is_active
        )

    def can_create(self, object_or_class) -> bool:
        return (
            has_crud_access(self.permissions, object_or_class, CRUD_CREATE)
            and self.user.is_active
            self._has_crud_access(object_or_class, CRUD_CREATE) and self.user.is_active
        )

    def can_read(self, object_or_class) -> bool:
        # if self.user.is_superuser and self.user.is_active:  # superuser gets a pass
        #    return True
        return (
            has_crud_access(self.permissions, object_or_class, CRUD_READ)
            and self.user.is_active
        )
        return self._has_crud_access(object_or_class, CRUD_READ) and self.user.is_active

    def can_update(self, object_or_class) -> bool:
        return (
            has_crud_access(self.permissions, object_or_class, CRUD_UPDATE)
            and self.user.is_active
            self._has_crud_access(object_or_class, CRUD_UPDATE) and self.user.is_active
        )

    def can_delete(self, object_or_class) -> bool:
        return (
            has_crud_access(self.permissions, object_or_class, CRUD_DELETE)
            and self.user.is_active
            self._has_crud_access(object_or_class, CRUD_DELETE) and self.user.is_active
        )

    def check_create(self, object_or_class):
@@ -331,7 +402,7 @@ class UserAccess:
    def user_roles(self) -> Iterator[Role]:
        for access in ProjectAccess.objects.filter(
            project=self.project, user=self.user
        ):
        ).select_related("role"):
            yield access.role

    def project_roles(self) -> Iterator[Role]:
@@ -344,7 +415,7 @@ class UserAccess:
        access, _ = ProjectAccess.objects.get_or_create(
            project=self.project, user=self.user, role=role
        )
        self._permissions = None
        self._reset_permissions_cache_and_indexes()
        return access

    def remove_role_from_user(self, role: Role):
@@ -353,4 +424,4 @@ class UserAccess:
        ProjectAccess.objects.filter(
            project=self.project, user=self.user, role=role
        ).delete()
        self._permissions = None
        self._reset_permissions_cache_and_indexes()
+83 −1
Original line number Diff line number Diff line
@@ -15,7 +15,12 @@ import json
import os
import stat

from resources.acl import make_or_create_global_readonly_role_for_project, UserAccess
from resources.acl import (
    PERM_CLASS_COLLECTION,
    PERM_CLASS_RESOURCE,
    make_or_create_global_readonly_role_for_project,
    UserAccess,
)
from resources import helpers
from resources import snapshots
from resources import tasks
@@ -50,6 +55,83 @@ class ServiceTestCase(TestCase):
        self.assertFalse(access.can_update("resource"))
        self.assertFalse(access.can_delete("resource"))

    def test_object_permission_on_ancestor_applies_to_resource(self):
        user = User.objects.create(username="ancestor_reader")
        project = Project.objects.create(label="ancestor project", description="")
        root = models.Collection.objects.create(
            title="Root Collection", project=project, parent=None
        )
        child = models.Collection.objects.create(
            title="Child Collection", project=project, parent=root
        )
        file_type = FileType.objects.create(title="Text", mime="text/plain")
        file_instance = File.objects.create(
            title="Document",
            original_name="document.txt",
            project=project,
            hash="c" * 64,
            file_type=file_type,
            size=123,
        )
        other_file = File.objects.create(
            title="Other Document",
            original_name="other-document.txt",
            project=project,
            hash="d" * 64,
            file_type=file_type,
            size=123,
        )
        child.resources.add(file_instance)
        child.resources.add(other_file)
        role = Role.objects.create(label="ancestor reader", project=project)
        ObjectPermission.objects.create(
            role=role,
            object_class="resource",
            object_pk=None,
            object_read=False,
        )
        ObjectPermission.objects.create(
            role=role,
            object_class="collection",
            object_pk=root.pk,
            object_read=True,
        )
        ProjectAccess.objects.create(project=project, user=user, role=role)

        access = UserAccess(user, project)

        self.assertTrue(access.can_read(file_instance))
        self.assertTrue(access.can_read(other_file))
        self.assertFalse(access.can_update(file_instance))
        self.assertEqual(
            access._tree_permission_keys_cache,
            {
                (PERM_CLASS_COLLECTION, child.pk): (
                    (PERM_CLASS_COLLECTION, child.pk),
                    (PERM_CLASS_COLLECTION, root.pk),
                )
            },
        )
        self.assertNotIn(
            (PERM_CLASS_RESOURCE, file_instance.pk), access._tree_permission_keys_cache
        )
        self.assertNotIn(
            (PERM_CLASS_RESOURCE, other_file.pk), access._tree_permission_keys_cache
        )

    def test_user_roles_loads_roles_with_accesses(self):
        user = User.objects.create(username="role_user")
        project = Project.objects.create(label="role project", description="")
        role = Role.objects.create(label="reader", project=project)
        ProjectAccess.objects.create(project=project, user=user, role=role)
        access = UserAccess(user, project)

        with CaptureQueriesContext(connection) as queries:
            labels = [role.label for role in access.user_roles()]

        self.assertEqual(labels, ["reader"])
        self.assertEqual(len(queries), 1)

    def test_media_command_uses_configured_timeout(self):
        with patch.object(
            helpers.settings, "JAMA_MEDIA_SUBPROCESS_TIMEOUT_SECONDS", 42