diff --git a/aqueduct/management/auth.py b/aqueduct/management/auth.py index 96b7b5f9..58847f5a 100644 --- a/aqueduct/management/auth.py +++ b/aqueduct/management/auth.py @@ -110,6 +110,89 @@ def _update_user_role( user.is_superuser = is_superuser user.save() + @staticmethod + def _add_team_membership(profile: UserProfile, teams): + enable_creation = getattr(settings, "ENABLE_OAUTH_GROUP_CREATION", True) + org = profile.org + for team_name, original_group_name in teams: + # Look up by oauth_group_name first, so renaming the mapping + # function renames existing teams instead of creating duplicates. + team = Team.objects.filter(oauth_group_name=original_group_name, org=org).first() + created = False + + if team is not None: + if team.name != team_name: + collision = ( + Team.objects.filter(name=team_name, org=org).exclude(pk=team.pk).exists() + ) + + if collision: + log.warning( + "Cannot rename team '%s' -> '%s': name collision (org: %s, " + "oauth_group: '%s'). Reusing existing team as-is.", + team.name, + team_name, + org.name, + original_group_name, + ) + else: + log.info( + "Renaming team '%s' -> '%s' (org: %s, oauth_group: '%s')", + team.name, + team_name, + org.name, + original_group_name, + ) + team.name = team_name + team.save(update_fields=["name"]) + elif not enable_creation and not Team.objects.filter(name=team_name, org=org).exists(): + log.info("Skipping team '%s' (ENABLE_OAUTH_GROUP_CREATION=False)", team_name) + continue + else: + team, created = Team.objects.get_or_create( + name=team_name, org=org, defaults={"oauth_group_name": original_group_name} + ) + + if created: + log.info("Created team '%s' for org '%s'", team_name, org.name) + else: + log.info("Reused existing team '%s' for org '%s'", team_name, org.name) + + TeamMembership.objects.get_or_create(user_profile=profile, team=team) + log.info("Added user '%s' to team '%s' (%s)", profile.user.email, team_name, org.name) + + @staticmethod + def _remove_team_membership(profile: UserProfile, teams): + enable_removal = getattr(settings, "ENABLE_OAUTH_GROUP_REMOVAL", True) + org = profile.org + + for team_name in teams: + try: + team = Team.objects.get(name=team_name, org=org) + is_oauth_managed = bool(team.oauth_group_name) + + if is_oauth_managed or enable_removal: + TeamMembership.objects.filter(user_profile=profile, team=team).delete() + log.info( + "Removed user '%s' from team '%s' (%s)", + profile.user.email, + team_name, + org.name, + ) + else: + log.info( + "Skipping removal from non-OAuth team '%s' for user '%s'", + team_name, + profile.user.email, + ) + except Team.DoesNotExist: + log.warning( + "Team '%s' not found for removal (org: %s, user: %s)", + team_name, + org.name, + profile.user.email, + ) + def _sync_team_membership(self, user: User, profile: UserProfile, team_names: list[str]): """ Synchronize team membership based on OAuth claims. @@ -122,6 +205,10 @@ def _sync_team_membership(self, user: User, profile: UserProfile, team_names: li - Respects org boundaries (teams must belong to user's org) """ if not getattr(settings, "ENABLE_OAUTH_GROUP_MANAGEMENT", False): + log.info( + "Skipping synchronization of teams %s (ENABLE_OAUTH_GROUP_MANAGEMENT=False)", + sorted(team_names), + ) return if not team_names: @@ -129,12 +216,11 @@ def _sync_team_membership(self, user: User, profile: UserProfile, team_names: li team_mappings = self._get_teams(team_names=team_names) - org = profile.org with transaction.atomic(): existing_memberships = set( - TeamMembership.objects.filter(user_profile=profile).values_list( - "team__name", flat=True - ) + TeamMembership.objects.filter(user_profile=profile) + .exclude(team__oauth_group_name__in=["", None]) + .values_list("team__name", flat=True) ) target_team_names = {team_name for team_name, _ in team_mappings} @@ -144,88 +230,10 @@ def _sync_team_membership(self, user: User, profile: UserProfile, team_names: li (name, team_name_to_original[name]) for name in target_team_names - existing_memberships ] - teams_to_remove = existing_memberships - target_team_names - - enable_creation = getattr(settings, "ENABLE_OAUTH_GROUP_CREATION", True) - enable_removal = getattr(settings, "ENABLE_OAUTH_GROUP_REMOVAL", True) - - for team_name, original_group_name in teams_to_add: - # Look up by oauth_group_name first, so renaming the mapping - # function renames existing teams instead of creating duplicates. - existing = Team.objects.filter( - oauth_group_name=original_group_name, org=org - ).first() - - if existing is not None: - if existing.name != team_name: - # Check for name collision before renaming - collision = ( - Team.objects.filter(name=team_name, org=org) - .exclude(pk=existing.pk) - .exists() - ) - if collision: - log.warning( - "Cannot rename team '%s' -> '%s': name collision (org: %s, " - "oauth_group: '%s'). Reusing existing team as-is.", - existing.name, - team_name, - org.name, - original_group_name, - ) - else: - log.info( - "Renaming team '%s' -> '%s' (org: %s, oauth_group: '%s')", - existing.name, - team_name, - org.name, - original_group_name, - ) - existing.name = team_name - existing.save(update_fields=["name"]) - team = existing - created = False - elif ( - not enable_creation - and not Team.objects.filter(name=team_name, org=org).exists() - ): - log.info("Skipping team '%s' (ENABLE_OAUTH_GROUP_CREATION=False)", team_name) - continue - else: - team, created = Team.objects.get_or_create( - name=team_name, org=org, defaults={"oauth_group_name": original_group_name} - ) - - if created: - log.info("Created team '%s' for org '%s'", team_name, org.name) - else: - log.info("Reused existing team '%s' for org '%s'", team_name, org.name) + self._add_team_membership(profile=profile, teams=teams_to_add) - TeamMembership.objects.get_or_create(user_profile=profile, team=team) - log.info("Added user '%s' to team '%s' (%s)", user.email, team_name, org.name) - - for team_name in teams_to_remove: - try: - team = Team.objects.get(name=team_name, org=org) - is_oauth_managed = bool(team.oauth_group_name) - if is_oauth_managed or enable_removal: - TeamMembership.objects.filter(user_profile=profile, team=team).delete() - log.info( - "Removed user '%s' from team '%s' (%s)", user.email, team_name, org.name - ) - else: - log.info( - "Skipping removal from non-OAuth team '%s' for user '%s'", - team_name, - user.email, - ) - except Team.DoesNotExist: - log.warning( - "Team '%s' not found for removal (org: %s, user: %s)", - team_name, - org.name, - user.email, - ) + teams_to_remove = existing_memberships - target_team_names + self._remove_team_membership(profile=profile, teams=teams_to_remove) def create_user(self, claims: dict[str, Any]) -> User | None: org = self._org(claims) diff --git a/aqueduct/management/tests/test_oauth_team_creation.py b/aqueduct/management/tests/test_oauth_team_creation.py index a654a520..b2965b54 100644 --- a/aqueduct/management/tests/test_oauth_team_creation.py +++ b/aqueduct/management/tests/test_oauth_team_creation.py @@ -268,6 +268,24 @@ def test_user_removed_from_oauth_managed_teams_even_when_removal_disabled(self): self.assertEqual(memberships.count(), 1) self.assertEqual(memberships.first().team.name, "E123") + def test_user_removed_from_oauth_managed_teams(self): + """Test that user is removed from OAuth-managed teams when removal is enabled.""" + user = User.objects.create_user(username="testuser", email="test@example.com") + user.groups.add(self.user_group) + profile = UserProfile.objects.create(user=user, org=self.org) + + initial_groups = {"email": "test@example.com", "groups": ["E123-Students", "E456-Staff"]} + sync_teams(self.backend, user, profile, initial_groups) + + self.assertEqual(TeamMembership.objects.filter(user_profile=profile).count(), 2) + + updated_groups = {"email": "test@example.com", "groups": ["E123-Students"]} + sync_teams(self.backend, user, profile, updated_groups) + + memberships = TeamMembership.objects.filter(user_profile=profile) + self.assertEqual(memberships.count(), 1) + self.assertEqual(memberships.first().team.name, "E123") + @override_settings(ENABLE_OAUTH_GROUP_REMOVAL=False) def test_user_not_removed_from_non_oauth_teams_when_removal_disabled(self): """Test that user stays in non-OAuth teams when ENABLE_OAUTH_GROUP_REMOVAL=False.""" @@ -294,6 +312,44 @@ def test_user_not_removed_from_non_oauth_teams_when_removal_disabled(self): self.assertIn("E123", team_names) self.assertIn("ManualTeam", team_names) + def test_user_not_removed_from_non_oauth_teams_when_removal_enabled(self): + """Regression: non-OAuth teams are never removed, even with removal enabled.""" + user = User.objects.create_user(username="testuser", email="test@example.com") + user.groups.add(self.user_group) + profile = UserProfile.objects.create(user=user, org=self.org) + + manual_team = Team.objects.create(name="ManualTeam", org=self.org, oauth_group_name="") + TeamMembership.objects.create(user_profile=profile, team=manual_team) + + # Default ENABLE_OAUTH_GROUP_REMOVAL=True; group no longer maps to the manual team. + updated_groups = {"email": "test@example.com", "groups": ["OtherGroup"]} + sync_teams(self.backend, user, profile, updated_groups) + + memberships = TeamMembership.objects.filter(user_profile=profile) + self.assertEqual(memberships.count(), 1) + self.assertEqual(memberships.first().team.name, "ManualTeam") + + def test_only_oauth_teams_managed_and_non_oauth_preserved(self): + """Only OAuth-managed teams are added/removed; non-OAuth memberships are preserved.""" + user = User.objects.create_user(username="testuser", email="test@example.com") + user.groups.add(self.user_group) + profile = UserProfile.objects.create(user=user, org=self.org) + + initial_groups = {"email": "test@example.com", "groups": ["E123-Students", "E456-Staff"]} + sync_teams(self.backend, user, profile, initial_groups) + self.assertEqual(TeamMembership.objects.filter(user_profile=profile).count(), 2) + + # Manually attach to a non-OAuth (manual) team; sync must never touch it. + manual_team = Team.objects.create(name="ManualTeam", org=self.org, oauth_group_name="") + TeamMembership.objects.create(user_profile=profile, team=manual_team) + + # E456-Staff is dropped, E123-Students kept; OtherGroup maps to no OAuth team. + updated_groups = {"email": "test@example.com", "groups": ["E123-Students", "OtherGroup"]} + sync_teams(self.backend, user, profile, updated_groups) + + team_names = {m.team.name for m in TeamMembership.objects.filter(user_profile=profile)} + self.assertEqual(team_names, {"E123", "ManualTeam"}) + def test_membership_sync_on_update(self): """Test membership sync on update_user().""" claims_initial = {"email": "test@example.com", "groups": ["E123-Students"]} @@ -484,6 +540,46 @@ def test_creation_disabled_flag(self): self.assertEqual(Team.objects.filter(org=self.org).count(), 0) self.assertEqual(TeamMembership.objects.filter(user_profile=profile).count(), 0) + def test_creation_disabled_adds_membership_to_existing_team(self): + """ENABLE_OAUTH_GROUP_CREATION=False still adds memberships to existing teams.""" + seed_active_config(SNIPPET_TEAM_NAMES_AND_MAP) + Team.objects.create(name="E123", org=self.org, oauth_group_name="E123-Students") + + with override_settings(ENABLE_OAUTH_GROUP_CREATION=False): + claims = {"email": "test@example.com", "groups": ["E123-Students", "E456-Staff"]} + + user = User.objects.create_user(username="testuser", email="test@example.com") + user.groups.add(self.user_group) + profile = UserProfile.objects.create(user=user, org=self.org) + + sync_teams(self.backend, user, profile, claims) + + # E456 is not created; E123 already exists and the user is added to it. + self.assertEqual(Team.objects.filter(org=self.org).count(), 1) + self.assertEqual(TeamMembership.objects.filter(user_profile=profile).count(), 1) + self.assertEqual(TeamMembership.objects.get(user_profile=profile).team.name, "E123") + + def test_creation_disabled_reuses_manual_team_by_name(self): + """Creation disabled reuses an existing manual team matching the display name.""" + seed_active_config(SNIPPET_TEAM_NAMES_AND_MAP) + manual_team = Team.objects.create(name="E123", org=self.org, oauth_group_name="") + + with override_settings(ENABLE_OAUTH_GROUP_CREATION=False): + claims = {"email": "test@example.com", "groups": ["E123-Students"]} + + user = User.objects.create_user(username="testuser", email="test@example.com") + user.groups.add(self.user_group) + profile = UserProfile.objects.create(user=user, org=self.org) + + sync_teams(self.backend, user, profile, claims) + + # No new team created; the manual team is reused and the user is added to it. + self.assertEqual(Team.objects.filter(org=self.org).count(), 1) + memberships = TeamMembership.objects.filter(user_profile=profile) + self.assertEqual(memberships.count(), 1) + self.assertEqual(memberships.first().team, manual_team) + self.assertEqual(manual_team.oauth_group_name, "") + def test_manual_team_not_affected_by_oauth_sync(self): """Test that manually created teams without oauth_group_name are not affected.""" manual_team = Team.objects.create(name="ManualTeam", org=self.org, oauth_group_name="")