diff --git a/server/auth/auth.py b/server/auth/auth.py index 145616cc..03184e8d 100644 --- a/server/auth/auth.py +++ b/server/auth/auth.py @@ -43,6 +43,10 @@ class AuthTypeBase(ABC): """Return the name of the user (string)""" pass + def get_user_picture(self): + """Return the location to the user's picture""" + return None + class AuthTypeClientBase(AuthTypeBase): """Base type for all authentication types that require the client to login""" diff --git a/server/auth/auth_oauth.py b/server/auth/auth_oauth.py index d7162318..b2724c57 100644 --- a/server/auth/auth_oauth.py +++ b/server/auth/auth_oauth.py @@ -146,21 +146,19 @@ class AuthTypeOAuth(AuthTypeClientBase): def get_user_id(self): payload = self.get_userinfo() - if payload and payload.get("sub"): - return payload.get("sub") - return None + return payload.get("sub") if payload else None def get_user_name(self): payload = self.get_userinfo() - if payload and payload.get("name"): - return payload.get("name") - return None + return payload.get("name") if payload else None def get_user_email(self): payload = self.get_userinfo() - if payload and payload.get("email"): - return payload.get("email") - return None + return payload.get("email") if payload else None + + def get_user_picture(self): + payload = self.get_userinfo() + return payload.get("picture") if payload else None def update_response(self, response): response.cache_control.update(dict(public=True, max_age=0, no_store=True, no_cache=True, must_revalidate=True)) diff --git a/server/auth/auth_test.py b/server/auth/auth_test.py index 4594dae0..06b92cc7 100644 --- a/server/auth/auth_test.py +++ b/server/auth/auth_test.py @@ -10,12 +10,14 @@ class AuthTypeTest(AuthTypeClientBase): CXGUID = "cxguid_test" CXGUNAME = "cxguname_test" CXGUEMAIL = "cxguemail_test" + CXGUPICTURE = "cxgupicture_test" def __init__(self, app_config): super().__init__() self.user_name = "test_account" self.user_id = "id0001" self.user_email = "test_account@test.com" + self.user_picture = None def is_valid_authentication_type(self): return True @@ -42,12 +44,16 @@ class AuthTypeTest(AuthTypeClientBase): def get_user_email(self): return session.get(self.CXGUEMAIL) + def get_user_picture(self): + return session.get(self.CXGUPICTURE) + def login(self): args = request.args return_to = args.get("dataset", "/") session[self.CXGUID] = args.get("userid", self.user_id) session[self.CXGUNAME] = args.get("username", self.user_name) session[self.CXGUEMAIL] = args.get("email", self.user_email) + session[self.CXGUPICTURE] = args.get("picture", self.user_picture) return redirect(return_to) def logout(self): diff --git a/server/common/config/client_config.py b/server/common/config/client_config.py index ccddf2c1..8a70c0b6 100644 --- a/server/common/config/client_config.py +++ b/server/common/config/client_config.py @@ -117,5 +117,6 @@ def get_client_userinfo(app_config, data_adaptor): "username": auth.get_user_name(), "user_id": auth.get_user_id(), "email": auth.get_user_email(), + "picture": auth.get_user_picture(), } return userinfo diff --git a/server/test/unit/auth/test_auth.py b/server/test/unit/auth/test_auth.py index f7f6f358..0b8c0bf1 100644 --- a/server/test/unit/auth/test_auth.py +++ b/server/test/unit/auth/test_auth.py @@ -82,6 +82,7 @@ class AuthTest(unittest.TestCase): userinfo = session.get(f"{server}/auth/pbmc3k.cxg/api/v0.2/userinfo").json() self.assertTrue(userinfo["userinfo"]["is_authenticated"]) self.assertEqual(userinfo["userinfo"]["username"], "test_account") + self.assertEqual(userinfo["userinfo"]["picture"], None) self.assertTrue(config["config"]["parameters"]["annotations"]) r = session.get(f"{server}/{logout_uri}") @@ -100,6 +101,12 @@ class AuthTest(unittest.TestCase): self.assertIsNone(userinfo) self.assertFalse(config["config"]["parameters"]["annotations"]) + # login with a picture + r = session.get(f"{server}/{login_uri}&picture=myimage.png") + userinfo = session.get(f"{server}/auth/pbmc3k.cxg/api/v0.2/userinfo").json() + self.assertTrue(userinfo["userinfo"]["is_authenticated"]) + self.assertEqual(userinfo["userinfo"]["picture"], "myimage.png") + def test_auth_test_single(self): c = AppConfig() c.update_server_config(