diff --git a/team/apps.py b/team/apps.py index ee3a09f92..31808a1d0 100644 --- a/team/apps.py +++ b/team/apps.py @@ -4,3 +4,6 @@ class TeamConfig(AppConfig): default_auto_field = "django.db.models.BigAutoField" name = "team" + + def ready(self): + import team.signals # noqa F401 diff --git a/team/management/__init__.py b/team/management/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/team/management/commands/__init__.py b/team/management/commands/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/team/management/commands/create_user_groups.py b/team/management/commands/create_user_groups.py new file mode 100644 index 000000000..2fe6b8d68 --- /dev/null +++ b/team/management/commands/create_user_groups.py @@ -0,0 +1,53 @@ +from django.contrib.auth.models import Group, Permission +from django.contrib.contenttypes.models import ContentType +from django.core.management.base import BaseCommand + +from team.models import ( + GROUP_NAMES, + COLLECTION_TEAM_ADMIN, + JOURNAL_TEAM_ADMIN, + CollectionTeamMember, + Company, + CompanyTeamMember, + JournalCompanyContract, + JournalTeamMember, +) + + +class Command(BaseCommand): + help = "Create default user groups and assign permissions for team management" + + def handle(self, *args, **options): + for name in GROUP_NAMES: + Group.objects.get_or_create(name=name) + self.stdout.write(f"Group '{name}' ensured.") + + self._assign_permissions() + self.stdout.write(self.style.SUCCESS("User groups created/updated successfully.")) + + def _assign_permissions(self): + # COLLECTION_TEAM_ADMIN: can manage all team members and Company CRUD + collection_admin_group, _ = Group.objects.get_or_create(name=COLLECTION_TEAM_ADMIN) + collection_admin_permissions = self._get_model_permissions( + [CollectionTeamMember, Company, JournalTeamMember, CompanyTeamMember, JournalCompanyContract] + ) + collection_admin_group.permissions.set(collection_admin_permissions) + + # JOURNAL_TEAM_ADMIN: can manage journal team members and Company Contracts CRUD + journal_admin_group, _ = Group.objects.get_or_create(name=JOURNAL_TEAM_ADMIN) + journal_admin_permissions = self._get_model_permissions( + [JournalTeamMember, JournalCompanyContract] + ) + journal_admin_group.permissions.set(journal_admin_permissions) + + # COMPANY_TEAM_ADMIN: can manage company team members + company_admin_group, _ = Group.objects.get_or_create(name=COMPANY_TEAM_ADMIN) + company_admin_permissions = self._get_model_permissions([CompanyTeamMember]) + company_admin_group.permissions.set(company_admin_permissions) + + def _get_model_permissions(self, models): + permissions = [] + for model in models: + ct = ContentType.objects.get_for_model(model) + permissions.extend(Permission.objects.filter(content_type=ct)) + return permissions diff --git a/team/models.py b/team/models.py index 29fb007a3..0272bc174 100644 --- a/team/models.py +++ b/team/models.py @@ -16,6 +16,31 @@ ALLOWED_COLLECTIONS = ["dom", "spa", "scl", "pan"] +# Django group names for team-based access control. +# COLLECTION_TEAM_ADMIN: collection managers — can CRUD Company, JournalTeamMember, +# CompanyTeamMember, and members of their own collections. +COLLECTION_TEAM_ADMIN = "COLLECTION_TEAM_ADMIN" +# COLLECTION_TEAM_MEMBER: regular collection members — read-only access to own record. +COLLECTION_TEAM_MEMBER = "COLLECTION_TEAM_MEMBER" +# JOURNAL_TEAM_ADMIN: journal managers — can CRUD JournalTeamMember and JournalCompanyContract +# for their managed journals. +JOURNAL_TEAM_ADMIN = "JOURNAL_TEAM_ADMIN" +# JOURNAL_TEAM_MEMBER: regular journal members — read-only access to own record. +JOURNAL_TEAM_MEMBER = "JOURNAL_TEAM_MEMBER" +# COMPANY_TEAM_ADMIN: company managers — can CRUD CompanyTeamMember for their companies. +COMPANY_TEAM_ADMIN = "COMPANY_TEAM_ADMIN" +# COMPANY_MEMBER: regular company members — read-only access to own record. +COMPANY_MEMBER = "COMPANY_MEMBER" + +GROUP_NAMES = [ + COLLECTION_TEAM_ADMIN, + COLLECTION_TEAM_MEMBER, + JOURNAL_TEAM_ADMIN, + JOURNAL_TEAM_MEMBER, + COMPANY_TEAM_ADMIN, + COMPANY_MEMBER, +] + class TeamRole(models.TextChoices): """Role types for team members.""" @@ -206,6 +231,20 @@ def members(user, is_active_member=None): def has_upload_permission(cls, user): return cls.objects.filter(user=user, collection__acron__in=ALLOWED_COLLECTIONS).exists() + @classmethod + def get_queryset_for_user(cls, user, qs): + """Return the queryset of CollectionTeamMember records visible to the user. + + - Managers see all members of their own collection(s). + - Regular members see only their own record. + """ + managed_collection_ids = cls.objects.filter( + user=user, role=TeamRole.MANAGER, is_active_member=True + ).values_list("collection", flat=True) + if managed_collection_ids: + return qs.filter(collection__in=managed_collection_ids) + return qs.filter(user=user) + class Company(VisualIdentityMixin, CommonControlField): """ @@ -265,6 +304,23 @@ def get_members(cls, company_id): is_active_member=True ) + @classmethod + def get_queryset_for_user(cls, user, qs): + """Return the queryset of Company records visible to the user. + + - COLLECTION_TEAM_ADMIN (collection managers) can see all companies. + - Company members see only the companies they belong to. + """ + is_collection_manager = CollectionTeamMember.objects.filter( + user=user, role=TeamRole.MANAGER, is_active_member=True + ).exists() + if is_collection_manager: + return qs + company_ids = CompanyTeamMember.objects.filter( + user=user, is_active_member=True + ).values_list("company", flat=True) + return qs.filter(id__in=company_ids) + class JournalTeamMember(TeamMember): """ @@ -339,6 +395,26 @@ def get_user_journals(cls, user, role=None, is_active=True): filters["is_active_member"] = is_active return cls.objects.filter(**filters).select_related("journal") + @classmethod + def get_queryset_for_user(cls, user, qs): + """Return the queryset of JournalTeamMember records visible to the user. + + - COLLECTION_TEAM_ADMIN sees all journal team members. + - JOURNAL_TEAM_ADMIN sees members of their managed journals. + - Regular members see only their own record. + """ + is_collection_manager = CollectionTeamMember.objects.filter( + user=user, role=TeamRole.MANAGER, is_active_member=True + ).exists() + if is_collection_manager: + return qs + managed_journal_ids = cls.objects.filter( + user=user, role=TeamRole.MANAGER, is_active_member=True + ).values_list("journal", flat=True) + if managed_journal_ids: + return qs.filter(journal__in=managed_journal_ids) + return qs.filter(user=user) + class CompanyTeamMember(TeamMember): """ @@ -413,6 +489,26 @@ def get_user_companies(cls, user, role=None, is_active=True): filters["is_active_member"] = is_active return cls.objects.filter(**filters).select_related("company") + @classmethod + def get_queryset_for_user(cls, user, qs): + """Return the queryset of CompanyTeamMember records visible to the user. + + - COLLECTION_TEAM_ADMIN sees all company team members. + - COMPANY_TEAM_ADMIN sees members of their managed companies. + - Regular members see only their own record. + """ + is_collection_manager = CollectionTeamMember.objects.filter( + user=user, role=TeamRole.MANAGER, is_active_member=True + ).exists() + if is_collection_manager: + return qs + managed_company_ids = cls.objects.filter( + user=user, role=TeamRole.MANAGER, is_active_member=True + ).values_list("company", flat=True) + if managed_company_ids: + return qs.filter(company__in=managed_company_ids) + return qs.filter(user=user) + class JournalCompanyContract(CommonControlField): """ @@ -478,3 +574,23 @@ def get_company_journals(cls, company, is_active=True): def can_manage_contract(cls, user, journal): """Check if a user can manage contracts for a journal (must be a journal manager).""" return JournalTeamMember.user_is_manager(user, journal) + + @classmethod + def get_queryset_for_user(cls, user, qs): + """Return the queryset of JournalCompanyContract records visible to the user. + + - COLLECTION_TEAM_ADMIN sees all contracts. + - JOURNAL_TEAM_ADMIN sees contracts for their managed journals. + - All others see no contracts. + """ + is_collection_manager = CollectionTeamMember.objects.filter( + user=user, role=TeamRole.MANAGER, is_active_member=True + ).exists() + if is_collection_manager: + return qs + managed_journal_ids = JournalTeamMember.objects.filter( + user=user, role=TeamRole.MANAGER, is_active_member=True + ).values_list("journal", flat=True) + if managed_journal_ids: + return qs.filter(journal__in=managed_journal_ids) + return qs.none() diff --git a/team/signals.py b/team/signals.py new file mode 100644 index 000000000..9f14c9ecb --- /dev/null +++ b/team/signals.py @@ -0,0 +1,74 @@ +from django.contrib.auth.models import Group +from django.db.models.signals import post_delete, post_save + +from .models import ( + COLLECTION_TEAM_ADMIN, + COLLECTION_TEAM_MEMBER, + COMPANY_MEMBER, + COMPANY_TEAM_ADMIN, + JOURNAL_TEAM_ADMIN, + JOURNAL_TEAM_MEMBER, + CollectionTeamMember, + CompanyTeamMember, + JournalTeamMember, + TeamRole, +) + + +def _roles_for_user(model_class, user): + """Return the set of active roles the user holds in a team model.""" + return set( + model_class.objects.filter(user=user, is_active_member=True) + .values_list("role", flat=True) + ) + + +def update_user_groups(user): + """ + Synchronise a user's auth.Group memberships to reflect their current + active team-member roles. Called after any team member is saved or deleted. + """ + if user is None: + return + + collection_roles = _roles_for_user(CollectionTeamMember, user) + journal_roles = _roles_for_user(JournalTeamMember, user) + company_roles = _roles_for_user(CompanyTeamMember, user) + + _sync_group(user, COLLECTION_TEAM_ADMIN, TeamRole.MANAGER in collection_roles) + _sync_group(user, COLLECTION_TEAM_MEMBER, TeamRole.MEMBER in collection_roles) + _sync_group(user, JOURNAL_TEAM_ADMIN, TeamRole.MANAGER in journal_roles) + _sync_group(user, JOURNAL_TEAM_MEMBER, TeamRole.MEMBER in journal_roles) + _sync_group(user, COMPANY_TEAM_ADMIN, TeamRole.MANAGER in company_roles) + _sync_group(user, COMPANY_MEMBER, TeamRole.MEMBER in company_roles) + + +def _sync_group(user, group_name, should_belong): + """Add or remove a user from a group, creating the group if needed.""" + group, _ = Group.objects.get_or_create(name=group_name) + if should_belong: + user.groups.add(group) + else: + user.groups.remove(group) + + +def _make_signal_handler(description): + def handler(sender, instance, **kwargs): + update_user_groups(instance.user) + handler.__name__ = description + return handler + + +_TEAM_MODELS = [CollectionTeamMember, JournalTeamMember, CompanyTeamMember] + +for _model in _TEAM_MODELS: + post_save.connect( + _make_signal_handler(f"sync_{_model.__name__.lower()}_groups_on_save"), + sender=_model, + weak=False, + ) + post_delete.connect( + _make_signal_handler(f"sync_{_model.__name__.lower()}_groups_on_delete"), + sender=_model, + weak=False, + ) diff --git a/team/tests.py b/team/tests.py index 0d46b39ef..9e248d129 100644 --- a/team/tests.py +++ b/team/tests.py @@ -1,10 +1,17 @@ from django.contrib.auth import get_user_model from django.db import IntegrityError -from django.test import TestCase +from django.test import RequestFactory, TestCase from collection.models import Collection from journal.models import Journal from team.models import ( + COLLECTION_TEAM_ADMIN, + COLLECTION_TEAM_MEMBER, + COMPANY_MEMBER, + COMPANY_TEAM_ADMIN, + GROUP_NAMES, + JOURNAL_TEAM_ADMIN, + JOURNAL_TEAM_MEMBER, CollectionTeamMember, Company, CompanyTeamMember, @@ -540,3 +547,444 @@ def test_can_manage_contract_as_non_manager(self): self.assertFalse( JournalCompanyContract.can_manage_contract(self.user, self.journal) ) + + +class GroupNamesTest(TestCase): + """Test that group name constants are defined correctly.""" + + def test_group_names_constants(self): + """Test that all group name constants are defined.""" + self.assertEqual(COLLECTION_TEAM_ADMIN, "COLLECTION_TEAM_ADMIN") + self.assertEqual(COLLECTION_TEAM_MEMBER, "COLLECTION_TEAM_MEMBER") + self.assertEqual(JOURNAL_TEAM_ADMIN, "JOURNAL_TEAM_ADMIN") + self.assertEqual(JOURNAL_TEAM_MEMBER, "JOURNAL_TEAM_MEMBER") + self.assertEqual(COMPANY_TEAM_ADMIN, "COMPANY_TEAM_ADMIN") + self.assertEqual(COMPANY_MEMBER, "COMPANY_MEMBER") + + def test_group_names_list(self): + """Test that GROUP_NAMES contains all expected group names.""" + self.assertIn(COLLECTION_TEAM_ADMIN, GROUP_NAMES) + self.assertIn(COLLECTION_TEAM_MEMBER, GROUP_NAMES) + self.assertIn(JOURNAL_TEAM_ADMIN, GROUP_NAMES) + self.assertIn(JOURNAL_TEAM_MEMBER, GROUP_NAMES) + self.assertIn(COMPANY_TEAM_ADMIN, GROUP_NAMES) + self.assertIn(COMPANY_MEMBER, GROUP_NAMES) + self.assertEqual(len(GROUP_NAMES), 6) + + +class GetQuerysetFilteringTest(TestCase): + """Test the queryset filtering logic used in wagtail_hooks get_queryset methods.""" + + def setUp(self): + self.superuser = User.objects.create_superuser( + username="superuser", email="super@example.com", password="pass" + ) + self.collection_manager = User.objects.create_user( + username="col_manager", email="col_manager@example.com", password="pass" + ) + self.collection_member = User.objects.create_user( + username="col_member", email="col_member@example.com", password="pass" + ) + self.journal_manager = User.objects.create_user( + username="jour_manager", email="jour_manager@example.com", password="pass" + ) + self.journal_member = User.objects.create_user( + username="jour_member", email="jour_member@example.com", password="pass" + ) + self.company_manager = User.objects.create_user( + username="comp_manager", email="comp_manager@example.com", password="pass" + ) + self.company_member_user = User.objects.create_user( + username="comp_member", email="comp_member@example.com", password="pass" + ) + + self.collection = Collection.objects.create( + acron="TST", name="Test Collection", creator=self.superuser + ) + self.other_collection = Collection.objects.create( + acron="OTH", name="Other Collection", creator=self.superuser + ) + self.journal = Journal.objects.create(title="Test Journal", creator=self.superuser) + self.other_journal = Journal.objects.create(title="Other Journal", creator=self.superuser) + self.company = Company.objects.create(name="Test Company", creator=self.superuser) + self.other_company = Company.objects.create(name="Other Company", creator=self.superuser) + + # Set up collection team members + CollectionTeamMember.objects.create( + user=self.collection_manager, + collection=self.collection, + role=TeamRole.MANAGER, + is_active_member=True, + creator=self.superuser, + ) + CollectionTeamMember.objects.create( + user=self.collection_member, + collection=self.collection, + role=TeamRole.MEMBER, + is_active_member=True, + creator=self.superuser, + ) + + # Set up journal team members + JournalTeamMember.objects.create( + user=self.journal_manager, + journal=self.journal, + role=TeamRole.MANAGER, + is_active_member=True, + creator=self.superuser, + ) + JournalTeamMember.objects.create( + user=self.journal_member, + journal=self.journal, + role=TeamRole.MEMBER, + is_active_member=True, + creator=self.superuser, + ) + # journal_manager also member of other_journal (as member) + JournalTeamMember.objects.create( + user=self.journal_manager, + journal=self.other_journal, + role=TeamRole.MEMBER, + is_active_member=True, + creator=self.superuser, + ) + + # Set up company team members + CompanyTeamMember.objects.create( + user=self.company_manager, + company=self.company, + role=TeamRole.MANAGER, + is_active_member=True, + creator=self.superuser, + ) + CompanyTeamMember.objects.create( + user=self.company_member_user, + company=self.company, + role=TeamRole.MEMBER, + is_active_member=True, + creator=self.superuser, + ) + + # Contracts + self.contract = JournalCompanyContract.objects.create( + journal=self.journal, + company=self.company, + is_active=True, + creator=self.superuser, + ) + self.other_contract = JournalCompanyContract.objects.create( + journal=self.other_journal, + company=self.other_company, + is_active=True, + creator=self.superuser, + ) + + # --- CollectionTeamMember queryset filtering --- + + def test_collection_team_qs_superuser_sees_all(self): + """Superuser should see all CollectionTeamMember records.""" + qs = CollectionTeamMember.objects.all() + self.assertEqual(qs.count(), 2) + + def test_collection_team_qs_manager_sees_own_collection_members(self): + """COLLECTION_TEAM_ADMIN sees members of their collection(s).""" + managed_ids = CollectionTeamMember.objects.filter( + user=self.collection_manager, role=TeamRole.MANAGER, is_active_member=True + ).values_list("collection", flat=True) + filtered = CollectionTeamMember.objects.filter(collection__in=managed_ids) + # Both manager and member of the collection should be visible + self.assertEqual(filtered.count(), 2) + + def test_collection_team_qs_member_sees_only_self(self): + """COLLECTION_TEAM_MEMBER sees only their own record.""" + managed_ids = CollectionTeamMember.objects.filter( + user=self.collection_member, role=TeamRole.MANAGER, is_active_member=True + ).values_list("collection", flat=True) + self.assertFalse(managed_ids.exists()) + # Falls back to filter(user=self.collection_member) + filtered = CollectionTeamMember.objects.filter(user=self.collection_member) + self.assertEqual(filtered.count(), 1) + + # --- Company queryset filtering --- + + def test_company_qs_collection_manager_sees_all(self): + """COLLECTION_TEAM_ADMIN can see all companies.""" + is_collection_manager = CollectionTeamMember.objects.filter( + user=self.collection_manager, role=TeamRole.MANAGER, is_active_member=True + ).exists() + self.assertTrue(is_collection_manager) + # Collection manager should see all companies + qs = Company.objects.all() + self.assertEqual(qs.count(), 2) + + def test_company_qs_company_member_sees_own(self): + """Company member sees only their own companies.""" + is_collection_manager = CollectionTeamMember.objects.filter( + user=self.company_member_user, role=TeamRole.MANAGER, is_active_member=True + ).exists() + self.assertFalse(is_collection_manager) + company_ids = CompanyTeamMember.objects.filter( + user=self.company_member_user, is_active_member=True + ).values_list("company", flat=True) + filtered = Company.objects.filter(id__in=company_ids) + self.assertEqual(filtered.count(), 1) + self.assertEqual(filtered.first(), self.company) + + # --- JournalTeamMember queryset filtering --- + + def test_journal_team_qs_collection_manager_sees_all(self): + """COLLECTION_TEAM_ADMIN sees all journal team members.""" + is_collection_manager = CollectionTeamMember.objects.filter( + user=self.collection_manager, role=TeamRole.MANAGER, is_active_member=True + ).exists() + self.assertTrue(is_collection_manager) + # Should see all journal team members + self.assertEqual(JournalTeamMember.objects.count(), 3) + + def test_journal_team_qs_journal_manager_sees_own_journal_members(self): + """JOURNAL_TEAM_ADMIN sees members of their managed journals.""" + is_collection_manager = CollectionTeamMember.objects.filter( + user=self.journal_manager, role=TeamRole.MANAGER, is_active_member=True + ).exists() + self.assertFalse(is_collection_manager) + managed_journal_ids = JournalTeamMember.objects.filter( + user=self.journal_manager, role=TeamRole.MANAGER, is_active_member=True + ).values_list("journal", flat=True) + self.assertTrue(managed_journal_ids.exists()) + filtered = JournalTeamMember.objects.filter(journal__in=managed_journal_ids) + # journal_manager and journal_member are in the managed journal + self.assertEqual(filtered.count(), 2) + + def test_journal_team_qs_journal_member_sees_only_self(self): + """JOURNAL_TEAM_MEMBER sees only their own record.""" + managed_journal_ids = JournalTeamMember.objects.filter( + user=self.journal_member, role=TeamRole.MANAGER, is_active_member=True + ).values_list("journal", flat=True) + self.assertFalse(managed_journal_ids.exists()) + filtered = JournalTeamMember.objects.filter(user=self.journal_member) + self.assertEqual(filtered.count(), 1) + + # --- CompanyTeamMember queryset filtering --- + + def test_company_team_qs_collection_manager_sees_all(self): + """COLLECTION_TEAM_ADMIN sees all company team members.""" + is_collection_manager = CollectionTeamMember.objects.filter( + user=self.collection_manager, role=TeamRole.MANAGER, is_active_member=True + ).exists() + self.assertTrue(is_collection_manager) + self.assertEqual(CompanyTeamMember.objects.count(), 2) + + def test_company_team_qs_manager_sees_own_company_members(self): + """COMPANY_TEAM_ADMIN sees members of their managed companies.""" + managed_company_ids = CompanyTeamMember.objects.filter( + user=self.company_manager, role=TeamRole.MANAGER, is_active_member=True + ).values_list("company", flat=True) + self.assertTrue(managed_company_ids.exists()) + filtered = CompanyTeamMember.objects.filter(company__in=managed_company_ids) + self.assertEqual(filtered.count(), 2) + + def test_company_team_qs_member_sees_only_self(self): + """COMPANY_MEMBER sees only their own record.""" + managed_company_ids = CompanyTeamMember.objects.filter( + user=self.company_member_user, role=TeamRole.MANAGER, is_active_member=True + ).values_list("company", flat=True) + self.assertFalse(managed_company_ids.exists()) + filtered = CompanyTeamMember.objects.filter(user=self.company_member_user) + self.assertEqual(filtered.count(), 1) + + # --- JournalCompanyContract queryset filtering --- + + def test_contract_qs_collection_manager_sees_all(self): + """COLLECTION_TEAM_ADMIN sees all contracts.""" + is_collection_manager = CollectionTeamMember.objects.filter( + user=self.collection_manager, role=TeamRole.MANAGER, is_active_member=True + ).exists() + self.assertTrue(is_collection_manager) + self.assertEqual(JournalCompanyContract.objects.count(), 2) + + def test_contract_qs_journal_manager_sees_own_journal_contracts(self): + """JOURNAL_TEAM_ADMIN sees contracts for their managed journals.""" + is_collection_manager = CollectionTeamMember.objects.filter( + user=self.journal_manager, role=TeamRole.MANAGER, is_active_member=True + ).exists() + self.assertFalse(is_collection_manager) + managed_journal_ids = JournalTeamMember.objects.filter( + user=self.journal_manager, role=TeamRole.MANAGER, is_active_member=True + ).values_list("journal", flat=True) + filtered = JournalCompanyContract.objects.filter(journal__in=managed_journal_ids) + self.assertEqual(filtered.count(), 1) + self.assertEqual(filtered.first(), self.contract) + + def test_contract_qs_non_manager_sees_none(self): + """Users with no manager role see no contracts.""" + is_collection_manager = CollectionTeamMember.objects.filter( + user=self.journal_member, role=TeamRole.MANAGER, is_active_member=True + ).exists() + self.assertFalse(is_collection_manager) + managed_journal_ids = JournalTeamMember.objects.filter( + user=self.journal_member, role=TeamRole.MANAGER, is_active_member=True + ).values_list("journal", flat=True) + filtered = JournalCompanyContract.objects.filter(journal__in=managed_journal_ids) + self.assertEqual(filtered.count(), 0) + + +class UserGroupSyncTest(TestCase): + """Test that team member save/delete signals keep auth.Group memberships in sync.""" + + def setUp(self): + self.user = User.objects.create_user( + username="syncuser", email="syncuser@example.com", password="testpass123" + ) + self.collection = Collection.objects.create( + acron="SYN", name="Sync Collection", creator=self.user + ) + self.journal = Journal.objects.create(title="Sync Journal", creator=self.user) + self.company = Company.objects.create(name="Sync Company", creator=self.user) + + def _group_names(self): + return set(self.user.groups.values_list("name", flat=True)) + + # --- CollectionTeamMember --- + + def test_collection_manager_gets_collection_team_admin_group(self): + """Creating a MANAGER CollectionTeamMember adds COLLECTION_TEAM_ADMIN to the user.""" + CollectionTeamMember.objects.create( + user=self.user, + collection=self.collection, + role=TeamRole.MANAGER, + is_active_member=True, + creator=self.user, + ) + self.assertIn(COLLECTION_TEAM_ADMIN, self._group_names()) + self.assertNotIn(COLLECTION_TEAM_MEMBER, self._group_names()) + + def test_collection_member_gets_collection_team_member_group(self): + """Creating a MEMBER CollectionTeamMember adds COLLECTION_TEAM_MEMBER to the user.""" + CollectionTeamMember.objects.create( + user=self.user, + collection=self.collection, + role=TeamRole.MEMBER, + is_active_member=True, + creator=self.user, + ) + self.assertIn(COLLECTION_TEAM_MEMBER, self._group_names()) + self.assertNotIn(COLLECTION_TEAM_ADMIN, self._group_names()) + + def test_deactivating_collection_member_removes_group(self): + """Setting is_active_member=False removes the user from the group.""" + member = CollectionTeamMember.objects.create( + user=self.user, + collection=self.collection, + role=TeamRole.MANAGER, + is_active_member=True, + creator=self.user, + ) + self.assertIn(COLLECTION_TEAM_ADMIN, self._group_names()) + member.is_active_member = False + member.save() + self.assertNotIn(COLLECTION_TEAM_ADMIN, self._group_names()) + + def test_deleting_collection_member_removes_group(self): + """Deleting a CollectionTeamMember removes the user from the group.""" + member = CollectionTeamMember.objects.create( + user=self.user, + collection=self.collection, + role=TeamRole.MANAGER, + is_active_member=True, + creator=self.user, + ) + self.assertIn(COLLECTION_TEAM_ADMIN, self._group_names()) + member.delete() + self.assertNotIn(COLLECTION_TEAM_ADMIN, self._group_names()) + + # --- JournalTeamMember --- + + def test_journal_manager_gets_journal_team_admin_group(self): + """Creating a MANAGER JournalTeamMember adds JOURNAL_TEAM_ADMIN to the user.""" + JournalTeamMember.objects.create( + user=self.user, + journal=self.journal, + role=TeamRole.MANAGER, + is_active_member=True, + creator=self.user, + ) + self.assertIn(JOURNAL_TEAM_ADMIN, self._group_names()) + self.assertNotIn(JOURNAL_TEAM_MEMBER, self._group_names()) + + def test_journal_member_gets_journal_team_member_group(self): + """Creating a MEMBER JournalTeamMember adds JOURNAL_TEAM_MEMBER to the user.""" + JournalTeamMember.objects.create( + user=self.user, + journal=self.journal, + role=TeamRole.MEMBER, + is_active_member=True, + creator=self.user, + ) + self.assertIn(JOURNAL_TEAM_MEMBER, self._group_names()) + self.assertNotIn(JOURNAL_TEAM_ADMIN, self._group_names()) + + # --- CompanyTeamMember --- + + def test_company_manager_gets_company_team_admin_group(self): + """Creating a MANAGER CompanyTeamMember adds COMPANY_TEAM_ADMIN to the user.""" + CompanyTeamMember.objects.create( + user=self.user, + company=self.company, + role=TeamRole.MANAGER, + is_active_member=True, + creator=self.user, + ) + self.assertIn(COMPANY_TEAM_ADMIN, self._group_names()) + self.assertNotIn(COMPANY_MEMBER, self._group_names()) + + def test_company_member_gets_company_member_group(self): + """Creating a MEMBER CompanyTeamMember adds COMPANY_MEMBER to the user.""" + CompanyTeamMember.objects.create( + user=self.user, + company=self.company, + role=TeamRole.MEMBER, + is_active_member=True, + creator=self.user, + ) + self.assertIn(COMPANY_MEMBER, self._group_names()) + self.assertNotIn(COMPANY_TEAM_ADMIN, self._group_names()) + + def test_role_change_updates_groups(self): + """Changing role from MANAGER to MEMBER updates the groups correctly.""" + member = CollectionTeamMember.objects.create( + user=self.user, + collection=self.collection, + role=TeamRole.MANAGER, + is_active_member=True, + creator=self.user, + ) + self.assertIn(COLLECTION_TEAM_ADMIN, self._group_names()) + member.role = TeamRole.MEMBER + member.save() + # Refresh from DB + self.user.refresh_from_db() + self.assertNotIn(COLLECTION_TEAM_ADMIN, self._group_names()) + self.assertIn(COLLECTION_TEAM_MEMBER, self._group_names()) + + def test_user_manager_of_one_collection_and_member_of_another_gets_both_groups(self): + """User who is MANAGER of one collection and MEMBER of another gets both groups.""" + other_collection = Collection.objects.create( + acron="OTH", name="Other Collection", creator=self.user + ) + CollectionTeamMember.objects.create( + user=self.user, + collection=self.collection, + role=TeamRole.MANAGER, + is_active_member=True, + creator=self.user, + ) + CollectionTeamMember.objects.create( + user=self.user, + collection=other_collection, + role=TeamRole.MEMBER, + is_active_member=True, + creator=self.user, + ) + group_names = self._group_names() + self.assertIn(COLLECTION_TEAM_ADMIN, group_names) + self.assertIn(COLLECTION_TEAM_MEMBER, group_names) diff --git a/team/wagtail_hooks.py b/team/wagtail_hooks.py index b8f5ef9ec..9e4f7ed09 100644 --- a/team/wagtail_hooks.py +++ b/team/wagtail_hooks.py @@ -40,7 +40,7 @@ def get_queryset(self, request): qs = super().get_queryset(request) if request.user.is_superuser: return qs - return CollectionTeamMember.members(request.user) + return CollectionTeamMember.get_queryset_for_user(request.user, qs) class CompanyViewSet(SnippetViewSet): @@ -66,6 +66,12 @@ class CompanyViewSet(SnippetViewSet): "url", ) + def get_queryset(self, request): + qs = super().get_queryset(request) + if request.user.is_superuser: + return qs + return Company.get_queryset_for_user(request.user, qs) + class JournalTeamMemberViewSet(SnippetViewSet): model = JournalTeamMember @@ -89,6 +95,12 @@ class JournalTeamMemberViewSet(SnippetViewSet): "journal__title", ) + def get_queryset(self, request): + qs = super().get_queryset(request) + if request.user.is_superuser: + return qs + return JournalTeamMember.get_queryset_for_user(request.user, qs) + class CompanyTeamMemberViewSet(SnippetViewSet): model = CompanyTeamMember @@ -112,6 +124,12 @@ class CompanyTeamMemberViewSet(SnippetViewSet): "company__name", ) + def get_queryset(self, request): + qs = super().get_queryset(request) + if request.user.is_superuser: + return qs + return CompanyTeamMember.get_queryset_for_user(request.user, qs) + class JournalCompanyContractViewSet(SnippetViewSet): model = JournalCompanyContract @@ -133,6 +151,12 @@ class JournalCompanyContractViewSet(SnippetViewSet): "company__name", ) + def get_queryset(self, request): + qs = super().get_queryset(request) + if request.user.is_superuser: + return qs + return JournalCompanyContract.get_queryset_for_user(request.user, qs) + class TeamViewSetGroup(SnippetViewSetGroup): """