Skip to content
Merged
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
Original file line number Diff line number Diff line change
Expand Up @@ -1495,15 +1495,21 @@ def update_user(self, user: User) -> bool:
new_group_ids = {grp.id for grp in user.groups}
if existing_role_ids != new_role_ids or existing_group_ids != new_group_ids:
user.changed_on = datetime.datetime.now(tz=datetime.timezone.utc)
self.session.merge(user)
merged_user = self.session.merge(user)
self.session.commit()
self._reset_user_permissions_cache(merged_user)
log.info(const.LOGMSG_INF_SEC_UPD_USER, user)
except Exception as e:
log.error(const.LOGMSG_ERR_SEC_UPD_USER, e)
self.session.rollback()
return False
return True

@staticmethod
def _reset_user_permissions_cache(user: User) -> None:
"""Invalidate cached permissions to avoid stale auth checks after role updates."""
user._perms = None

def del_register_user(self, register_user) -> bool:
"""
Delete registration object from database.
Expand Down Expand Up @@ -1986,6 +1992,7 @@ def auth_user_ldap(self, username, password, rotate_session_id=True) -> User | N
# Sync the user's roles
if user and user_attributes and self.auth_roles_sync_at_login:
user.roles = self._ldap_calculate_user_roles(user_attributes)
self._reset_user_permissions_cache(user)
log.debug("Calculated new roles for user=%r as: %s", user_dn, user.roles)

# If the user is new, register them
Expand Down Expand Up @@ -2013,6 +2020,8 @@ def auth_user_ldap(self, username, password, rotate_session_id=True) -> User | N
if rotate_session_id:
self._rotate_session_id()
self.update_user_auth_stat(user)
self.session.expire(user, ["roles", "groups"])
self._reset_user_permissions_cache(user)
return user
return None

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -26,9 +26,11 @@

from airflow.providers.fab.auth_manager.models import (
Action,
Group,
Permission,
Resource,
Role,
User,
)
from airflow.providers.fab.auth_manager.security_manager.override import FabAirflowSecurityManagerOverride

Expand Down Expand Up @@ -193,6 +195,28 @@ def test_check_password_not_match(self, check_password):
check_password.return_value = False
assert not sm.check_password("test_user", "test_password")

def test_update_user_clears_cached_permissions(self):
sm = EmptySecurityManager()
user = Mock(
spec=User,
id=1,
roles=[Mock(spec=Role, id=2)],
groups=[Mock(spec=Group, id=3)],
_perms={("can_read", "DAG")},
)
existing_user = Mock(spec=User, roles=[Mock(spec=Role, id=4)], groups=[Mock(spec=Group, id=5)])
mock_merged_user = Mock(spec=User, _perms={("can_edit", "DAG")})
mock_session = Mock(spec=Session)
Comment thread
Pranaykarvi marked this conversation as resolved.
mock_session.get.return_value = existing_user
mock_session.merge.return_value = mock_merged_user

with mock.patch.object(EmptySecurityManager, "session", mock_session):
assert sm.update_user(user)

assert user._perms == {("can_read", "DAG")}
assert mock_merged_user._perms is None
mock_session.commit.assert_called_once_with()

@pytest.mark.parametrize(
("provider", "resp", "user_info"),
[
Expand Down
Loading