diff --git a/pyproject.toml b/pyproject.toml index eb41480f1..e0292bcc0 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -68,6 +68,7 @@ version = "5.2.0" [project.optional-dependencies] all = [ "social-auth-core[azuread]", + "social-auth-core[cas]", "social-auth-core[google-onetap]", "social-auth-core[saml]", "social-auth-core[shopify]" @@ -75,6 +76,9 @@ all = [ allpy3 = [ "social-auth-core[all]" ] +cas = [ + "python-cas>=1.7.2" +] # This is present until pip implements supports for PEP 735 # see https://github.com/pypa/pip/issues/12963 dev = [ @@ -129,6 +133,7 @@ check_untyped_defs = true disallow_untyped_defs = true ignore_missing_imports = true module = [ + "cas", "google.appengine.*", "onelogin.*", "openid.*", diff --git a/social_core/backends/cas_generic.py b/social_core/backends/cas_generic.py new file mode 100644 index 000000000..303bb1d48 --- /dev/null +++ b/social_core/backends/cas_generic.py @@ -0,0 +1,67 @@ +""" +Generic CAS backend + +Backend to authenticat with a generic CAS server. +""" + +from __future__ import annotations + +from cas import CASClient, CASClientV1, CASClientV2, CASClientV3, CASClientWithSAMLV1 + +from social_core.exceptions import AuthTokenError, SocialAuthImproperlyConfiguredError + +from .base import BaseAuth + + +class CasAuth(BaseAuth): + name = "cas" + title = "Cas" + + ID_KEY = "cas_user" + SERVER_URL: str | None = None + + def auth_url(self) -> str: + client = self.get_cas_client() + + url = client.get_login_url() + self.log_debug(f"Redirecting to CAS login: {url}") + return url + + def get_cas_client( + self, + ) -> CASClientV1 | CASClientV2 | CASClientV3 | CASClientWithSAMLV1: + """ + initializes the CASClient according to + the CAS_* settings + """ + server_url = self.setting("SERVER_URL", self.SERVER_URL) + service_url = self.redirect_uri + + if not server_url: + raise SocialAuthImproperlyConfiguredError + + version = self.setting("VERSION", 3) + kwargs = { + "service_url": service_url, + "server_url": server_url, + "extra_login_params": self.setting("EXTRA_LOGIN_PARAMS", []), + } + + # ty checks don't like __new__ to return a different type + return CASClient.__new__(CASClient, version=version, **kwargs) + + def auth_complete(self, *args, **kwargs): + client = self.get_cas_client() + ticket = self.strategy.request_data().get("ticket") + user, attributes, _pgtiou = client.verify_ticket(ticket) + if user is None or attributes is None: + raise AuthTokenError(self, "Token verification did not succeed") + kwargs.update({"response": attributes | {"cas_user": user}, "backend": self}) + return self.strategy.authenticate(*args, **kwargs) + + def get_user_details(self, response): + return { + "username": response.get(self.id_key()), + "email": response.get(self.setting("EMAIL_FIELD", "email")), + "fullname": response.get(self.setting("NAME_FIELD", "name")), + } diff --git a/social_core/tests/backends/test_cas_generic.py b/social_core/tests/backends/test_cas_generic.py new file mode 100644 index 000000000..8f98fcec7 --- /dev/null +++ b/social_core/tests/backends/test_cas_generic.py @@ -0,0 +1,73 @@ +from __future__ import annotations + +import urllib.parse + +import responses + +from social_core.tests.strategy import TEST_URI + +from .base import BaseBackendTest + +VALIDATION_PAYLOAD = b""" + + + test + + 1970-01-01T00:00:00+00:00 + false + true + test@example.com + Firstname Lastname + + + + + + + + +""" + + +class CasAuthTest(BaseBackendTest): + backend_path = "social_core.backends.cas_generic.CasAuth" + expected_username = "test" + ticket = "ST-12345678" + server_url = "http://cas-server.com/" + service_url = TEST_URI + validate_endpoint = "p3/serviceValidate" + + def setUp(self) -> None: + super().setUp() + params = [("ticket", self.ticket), ("service", self.service_url)] + + self.strategy.set_settings( + { + "SOCIAL_AUTH_CAS_SERVER_URL": self.server_url, + } + ) + + url = ( + urllib.parse.urljoin(self.server_url, self.validate_endpoint) + + "?" + + urllib.parse.urlencode(params) + ) + responses.add(responses.GET, url, VALIDATION_PAYLOAD) + + def do_start(self): + self.strategy.set_request_data({"ticket": self.ticket}, self.backend) + return self.backend.complete() + + def test_login(self) -> None: + self.do_login() + + def test_partial_pipeline(self) -> None: + self.do_partial_pipeline() + + def test_auth_url(self) -> None: + self.assertEqual( + self.backend.auth_url(), + urllib.parse.urljoin(self.server_url, "login") + + "?" + + urllib.parse.urlencode({"service": self.service_url}), + )