Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
178 changes: 93 additions & 85 deletions aqueduct/management/auth.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand All @@ -122,19 +205,22 @@ 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:
return

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}
Expand All @@ -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)
Expand Down
96 changes: 96 additions & 0 deletions aqueduct/management/tests/test_oauth_team_creation.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""
Expand All @@ -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"]}
Expand Down Expand Up @@ -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="")
Expand Down
Loading