diff --git a/lib/galaxy/managers/users.py b/lib/galaxy/managers/users.py index 5ad3a7952b0..450aeff7620 100644 --- a/lib/galaxy/managers/users.py +++ b/lib/galaxy/managers/users.py @@ -17,7 +17,6 @@ from markupsafe import escape from sqlalchemy import ( and_, exc, - func, select, true, ) @@ -316,24 +315,20 @@ class UserManager(base.ModelManager, deletable.PurgableManagerMixin): :raises exceptions.Conflict: if any are found """ - # TODO: remove this check when unique=True is added to the email column - if self.by_email(email) is not None: + if self.by_email(email, case_sensitive=False) is not None: raise exceptions.Conflict("Email must be unique", email=email) def by_id(self, user_id: int) -> Optional[model.User]: return self.app.model.session.get(self.model_class, user_id) # ---- filters - def by_email(self, email: str, filters=None, **kwargs) -> Optional[model.User]: + def by_email(self, email: str, case_sensitive: bool = True, deleted: bool | None = None) -> model.User | None: """ Find a user by their email. """ - filters = combine_lists(self.model_class.email == email, filters) - try: - # TODO: use one_or_none - return super().one(filters=filters, **kwargs) - except exceptions.ObjectNotFound: - return None + return get_user_by_email( + self.session(), email, self.model_class, case_sensitive=case_sensitive, deleted=deleted + ) def by_api_key(self, api_key: str, sa_session=None): """ @@ -425,10 +420,10 @@ class UserManager(base.ModelManager, deletable.PurgableManagerMixin): user = None if VALID_EMAIL_RE.match(identity): # VALID_PUBLICNAME and VALID_EMAIL do not overlap, so 'identity' here is an email address - user = get_user_by_email(self.session(), identity, self.model_class) + user = self.by_email(identity) if not user: # Try a case-insensitive match on the email - user = self._get_user_by_email_case_insensitive(self.session(), identity) + user = self.by_email(identity, case_sensitive=False) else: user = get_user_by_username(self.session(), identity, self.model_class) return user @@ -622,9 +617,9 @@ class UserManager(base.ModelManager, deletable.PurgableManagerMixin): return None def get_reset_token(self, trans, email): - reset_user = get_user_by_email(trans.sa_session, email, self.app.model.User) + reset_user = self.by_email(email) if not reset_user: - reset_user = self._get_user_by_email_case_insensitive(trans.sa_session, email) + reset_user = self.by_email(email, case_sensitive=False) if reset_user and not reset_user.deleted: prt = self.app.model.PasswordResetToken(reset_user) trans.sa_session.add(prt) @@ -660,7 +655,10 @@ class UserManager(base.ModelManager, deletable.PurgableManagerMixin): return None if getattr(self.app.config, "normalize_remote_user_email", False): remote_user_email = remote_user_email.lower() - user = get_user_by_email(self.session(), remote_user_email, self.app.model.User) + user = self.by_email(remote_user_email) + if not user: + # Try a case-insensitive match on the email + user = self.by_email(remote_user_email, case_sensitive=False) if user: # Ensure a private role and default permissions are set for remote users (remote user creation bug existed prior to 2009) self.app.security_agent.get_private_user_role(user, auto_create=True) @@ -669,13 +667,14 @@ class UserManager(base.ModelManager, deletable.PurgableManagerMixin): self.app.security_agent.user_set_default_permissions(user) self.app.security_agent.user_set_default_permissions(user, history=True, dataset=True) elif user is None: + session = self.session() random.seed() - user = self.app.model.User(email=remote_user_email) + username = username_from_email(session, remote_user_email, self.model_class) + user = self.model_class(email=remote_user_email, username=username) user.set_random_password(length=12) user.external = True - user.username = username_from_email(self.session(), remote_user_email, self.app.model.User) - self.session().add(user) - self.session().commit() + session.add(user) + session.commit() self.app.security_agent.create_private_user_role(user) # We set default user permissions, before we log in and set the default history permissions if self.app_type == "galaxy": @@ -683,10 +682,6 @@ class UserManager(base.ModelManager, deletable.PurgableManagerMixin): # self.log_event( "Automatically created account '%s'", user.email ) return user - def _get_user_by_email_case_insensitive(self, session, email): - stmt = select(self.app.model.User).where(func.lower(self.app.model.User.email) == email.lower()).limit(1) - return session.scalars(stmt).first() - class UserSerializer(base.ModelSerializer, deletable.PurgableSerializerMixin): model_manager_class = UserManager diff --git a/lib/galaxy/model/db/user.py b/lib/galaxy/model/db/user.py index 00f824f5e1e..16ce1dff810 100644 --- a/lib/galaxy/model/db/user.py +++ b/lib/galaxy/model/db/user.py @@ -1,6 +1,11 @@ -from typing import Optional +from collections.abc import ( + Iterable, + Sequence, +) +from typing import TYPE_CHECKING from sqlalchemy import ( + and_, false, func, or_, @@ -12,10 +17,12 @@ from galaxy.model import ( Role, User, ) -from galaxy.model.scoped_session import galaxy_scoped_session + +if TYPE_CHECKING: + from sqlalchemy.orm import scoped_session -def get_users_by_ids(session: galaxy_scoped_session, user_ids): +def get_users_by_ids(session: "scoped_session", user_ids: Iterable[int]) -> Sequence[User]: stmt = select(User).where(User.id.in_(user_ids)) return session.scalars(stmt).all() @@ -24,29 +31,33 @@ def get_users_by_ids(session: galaxy_scoped_session, user_ids): # the tool_shed app, which has its own User model, which is different from # galaxy.model.User. In that case, the tool_shed user model should be passed as # the model_class argument. -def get_user_by_email(session, email: str, model_class=User, case_sensitive=True): +def get_user_by_email( + session: "scoped_session", email: str, model_class=User, case_sensitive: bool = True, deleted: bool | None = None +) -> User | None: filter_clause = model_class.email == email if not case_sensitive: - filter_clause = func.lower(model_class.email) == func.lower(email) + filter_clause = func.lower(model_class.email) == email.lower() + if deleted is not None: + filter_clause = and_(filter_clause, model_class.deleted == deleted) stmt = select(model_class).where(filter_clause).limit(1) return session.scalars(stmt).first() -def get_user_by_username(session, username: str, model_class=User): +def get_user_by_username(session: "scoped_session", username: str, model_class=User) -> User | None: stmt = select(model_class).filter(model_class.username == username).limit(1) return session.scalars(stmt).first() def get_users_for_index( - session, + session: "scoped_session", deleted: bool, - f_email: Optional[str] = None, - f_name: Optional[str] = None, - f_any: Optional[str] = None, + f_email: str | None = None, + f_name: str | None = None, + f_any: str | None = None, is_admin: bool = False, expose_user_email: bool = False, expose_user_name: bool = False, -): +) -> Sequence[User]: stmt = select(User) if f_email and (is_admin or expose_user_email): stmt = stmt.where(User.email.like(f"%{f_email}%")) @@ -69,7 +80,7 @@ def get_users_for_index( return session.scalars(stmt).all() -def _cleanup_nonprivate_user_roles(session, user, private_role_id): +def _cleanup_nonprivate_user_roles(session: "scoped_session", user: User, private_role_id: int) -> None: """ Delete UserRoleAssociations EXCEPT FOR THE PRIVATE ROLE; Delete sharing roles that are associated with this user only; diff --git a/lib/galaxy/visualization/genomes.py b/lib/galaxy/visualization/genomes.py index f7ae75be582..2c7195a0523 100644 --- a/lib/galaxy/visualization/genomes.py +++ b/lib/galaxy/visualization/genomes.py @@ -376,6 +376,7 @@ class Genomes: dbkey_owner, dbkey = decode_dbkey(dbkey) if dbkey_owner: dbkey_user = get_user_by_username(trans.sa_session, dbkey_owner) + assert dbkey_user is not None else: dbkey_user = trans.user diff --git a/lib/galaxy/webapps/base/webapp.py b/lib/galaxy/webapps/base/webapp.py index a894c83fa75..495affb05e8 100644 --- a/lib/galaxy/webapps/base/webapp.py +++ b/lib/galaxy/webapps/base/webapp.py @@ -626,7 +626,7 @@ class GalaxyWebTransaction(base.DefaultWebTransaction, context.ProvidesHistoryCo galaxy_session_requires_flush = True elif ( remote_user_email - and galaxy_session.user.email != remote_user_email + and galaxy_session.user.email.lower() != remote_user_email.lower() and ( not self.app.config.allow_user_impersonation or remote_user_email not in self.app.config.admin_users_list diff --git a/lib/galaxy/webapps/galaxy/services/quotas.py b/lib/galaxy/webapps/galaxy/services/quotas.py index 564f864d599..c45f7742c7c 100644 --- a/lib/galaxy/webapps/galaxy/services/quotas.py +++ b/lib/galaxy/webapps/galaxy/services/quotas.py @@ -131,7 +131,13 @@ class QuotasService(ServiceBase): try: return trans.security.decode_id(item) except Exception: - return get_user_by_email(trans.sa_session, item).id + user = get_user_by_email(trans.sa_session, item) + if not user: + # Try a case-insensitive match on the email + user = get_user_by_email(trans.sa_session, item, case_sensitive=False) + if not user: + raise ValueError(f"User with email address '{item}' not found.") + return user.id def get_group_id(item): try: diff --git a/lib/galaxy/webapps/galaxy/services/sharable.py b/lib/galaxy/webapps/galaxy/services/sharable.py index 57cce95213a..cdc92fa74cf 100644 --- a/lib/galaxy/webapps/galaxy/services/sharable.py +++ b/lib/galaxy/webapps/galaxy/services/sharable.py @@ -4,8 +4,6 @@ from typing import ( Union, ) -from sqlalchemy import false - from galaxy.managers import base from galaxy.managers.sharable import ( SharableModelManager, @@ -156,9 +154,12 @@ class ShareableService: email_address = email_or_id.strip() if not email_address: continue - send_to_user = self.manager.user_manager.by_email( - email_address, filters=[User.table.c.deleted == false()] - ) + send_to_user = self.manager.user_manager.by_email(email_address, deleted=False) + if not send_to_user: + # Try a case-insensitive match on the email + send_to_user = self.manager.user_manager.by_email( + email_address, case_sensitive=False, deleted=False + ) if not send_to_user: send_to_err.add(f"{email_or_id} is not a valid Galaxy user.") diff --git a/test/integration/test_user_preferences.py b/test/integration/test_user_preferences.py index 62e76aa108a..4bc715a55ef 100644 --- a/test/integration/test_user_preferences.py +++ b/test/integration/test_user_preferences.py @@ -21,6 +21,7 @@ class TestUserPreferences(integration_util.IntegrationTestCase): app = cast(Any, self._test_driver.app if self._test_driver else None) db_user = get_user_by_email(app.model.session, user["email"]) + assert db_user is not None # create some initial data put(url) diff --git a/test/unit/app/managers/test_UserManager.py b/test/unit/app/managers/test_UserManager.py index 04a5662e6e4..90844e233c5 100644 --- a/test/unit/app/managers/test_UserManager.py +++ b/test/unit/app/managers/test_UserManager.py @@ -31,7 +31,6 @@ user2_data = dict(email="user2@user2.user2", username="user2", password=default_ user3_data = dict(email="user3@user3.user3", username="user3", password=default_password) user4_data = dict(email="user4@user4.user4", username="user4", password=default_password) uppercase_email_user = dict(email="USER5@USER5.USER5", username="USER5", password=default_password) -lowercase_email_user = dict(email="user5@user5.user5", username="user5", password=default_password) # ============================================================================= @@ -74,12 +73,23 @@ class TestUserManager(BaseTestCase): self.log("emails must be unique") with self.assertRaises(exceptions.Conflict): self.user_manager.create( - **dict(email="user2@user2.user2", username="user2a", password=default_password), + email=user2_data["email"], + username="user2a", + password=default_password, + ) + self.log("emails must be case-insensitive unique") + with self.assertRaises(exceptions.Conflict): + self.user_manager.create( + email=user2_data["email"].capitalize(), + username="user2a", + password=default_password, ) self.log("usernames must be unique") with self.assertRaises(exceptions.Conflict): self.user_manager.create( - **dict(email="user2a@user2.user2", username="user2", password=default_password), + email="user2a@user2.user2", + username=user2_data["username"], + password=default_password, ) def test_trimming(self): @@ -250,27 +260,10 @@ class TestUserManager(BaseTestCase): assert uppercase_user.username == uppercase_email_user["username"] assert self.user_manager.get_user_by_identity(uppercase_user.email) == uppercase_user assert self.user_manager.get_user_by_identity(uppercase_user.username) == uppercase_user - # Create another user with the same email just differently capitalized. - # This is not normally allowed now, since registration goes through user_manager.register(), - # which checks for that, but was possible in earlier releases of Galaxy - lowercase_user = self.user_manager.create(**lowercase_email_user) - assert lowercase_user.email == lowercase_email_user["email"] - assert lowercase_user.username == lowercase_email_user["username"] - assert self.user_manager.get_user_by_identity(lowercase_user.email) == lowercase_user - assert self.user_manager.get_user_by_identity(lowercase_user.username) == lowercase_user - # assert uppercase user can still be retrieved - assert self.user_manager.get_user_by_identity(uppercase_user.email) == uppercase_user - assert self.user_manager.get_user_by_identity(uppercase_user.username) == uppercase_user # username matches need to be exact - assert self.user_manager.get_user_by_identity(uppercase_user.username.capitalize()) is None - # email matches can ignore capitalization - ignore_email_capitalization_user = self.user_manager.create( - email="user123@nopassword.com", username="someusername123" - ) - assert ( - self.user_manager.get_user_by_identity(ignore_email_capitalization_user.email.capitalize()) - == ignore_email_capitalization_user - ) + assert self.user_manager.get_user_by_identity(uppercase_email_user["username"].capitalize()) is None + # Email lookups should be case-insensitive + assert self.user_manager.get_user_by_identity(uppercase_email_user["email"].capitalize()) == uppercase_user # ============================================================================= diff --git a/test/unit/data/model/db/conftest.py b/test/unit/data/model/db/conftest.py index 182b667daf4..357933a9f2a 100644 --- a/test/unit/data/model/db/conftest.py +++ b/test/unit/data/model/db/conftest.py @@ -37,7 +37,7 @@ def engine(db_url: str) -> "Engine": @pytest.fixture def session(engine: "Engine") -> Session: session = Session(engine) - # For sqlite, we need to explicitly enale foreign key constraints. + # For sqlite, we need to explicitly enable foreign key constraints. if engine.name == "sqlite": session.execute(text("PRAGMA foreign_keys = ON;")) return session