From 3935466e63dec0727f6b4effc1f858f01d2c62d0 Mon Sep 17 00:00:00 2001 From: Prowler Bot Date: Mon, 6 Jul 2026 16:20:13 +0200 Subject: [PATCH] fix: handle invitations in social and SAML auth (#11852) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-authored-by: Adrián Peña Co-authored-by: alejandrobailo Co-authored-by: Pedro Martín --- api/CHANGELOG.md | 1 + api/src/backend/api/adapters.py | 68 +++++-- api/src/backend/api/tests/test_adapters.py | 40 +++- api/src/backend/api/tests/test_views.py | 132 ++++++++++--- api/src/backend/api/utils.py | 41 +++- api/src/backend/api/v1/serializers.py | 3 + api/src/backend/api/v1/views.py | 91 +++++++-- ui/CHANGELOG.md | 8 + ui/actions/integrations/saml.ts | 11 +- ui/app/(auth)/(guest-only)/sign-up/page.tsx | 2 + ui/app/api/auth/callback/github/route.ts | 20 +- ui/app/api/auth/callback/google/route.ts | 20 +- ui/app/api/auth/callback/saml/route.ts | 6 +- ui/components/auth/oss/auth-form.tsx | 3 + ui/components/auth/oss/sign-in-form.tsx | 6 +- ui/components/auth/oss/sign-up-form.tsx | 20 +- ui/components/auth/oss/social-buttons.tsx | 208 +++++++++++++------- ui/lib/auth-callback-url.test.ts | 108 ++++++++++ ui/lib/auth-callback-url.ts | 65 ++++++ 19 files changed, 695 insertions(+), 158 deletions(-) create mode 100644 ui/lib/auth-callback-url.test.ts create mode 100644 ui/lib/auth-callback-url.ts diff --git a/api/CHANGELOG.md b/api/CHANGELOG.md index a781e57043..07f2203f2c 100644 --- a/api/CHANGELOG.md +++ b/api/CHANGELOG.md @@ -7,6 +7,7 @@ All notable changes to the **Prowler API** are documented in this file. ### 🐞 Fixed - Attack Paths: Scan rows now have database defaults for `is_migrated` and `sink_backend` so `scan-perform-scheduled` inserts survive deploy skew [(#11826)](https://github.com/prowler-cloud/prowler/pull/11826) +- Invited users now keep their invitation context when completing authentication with Google, GitHub, or SAML, so the invitation is accepted during login [(#11752)](https://github.com/prowler-cloud/prowler/pull/11752) ### 🔐 Security diff --git a/api/src/backend/api/adapters.py b/api/src/backend/api/adapters.py index cbd6795731..fa0abacb20 100644 --- a/api/src/backend/api/adapters.py +++ b/api/src/backend/api/adapters.py @@ -9,6 +9,7 @@ from api.models import ( User, UserRoleRelationship, ) +from api.utils import accept_invitation_for_user from django.db import transaction @@ -20,6 +21,22 @@ class ProwlerSocialAccountAdapter(DefaultSocialAccountAdapter): except User.DoesNotExist: return None + @staticmethod + def _get_invitation_token(request): + for source_name in ("data", "POST"): + data = getattr(request, source_name, None) or {} + if not hasattr(data, "get"): + continue + invitation_token = data.get("invitation_token") + if invitation_token: + return invitation_token + + wrapped_request = getattr(request, "_request", None) + if wrapped_request and wrapped_request is not request: + return ProwlerSocialAccountAdapter._get_invitation_token(wrapped_request) + + return None + def pre_social_login(self, request, sociallogin): # Link existing accounts with the same email address email = sociallogin.account.extra_data.get("email") @@ -83,29 +100,38 @@ class ProwlerSocialAccountAdapter(DefaultSocialAccountAdapter): user.name = social_account_name user.save(using=MainRouter.admin_db) - tenant = Tenant.objects.using(MainRouter.admin_db).create( - name=f"{user.email.split('@')[0]} default tenant" - ) - with rls_transaction(str(tenant.id)): - Membership.objects.using(MainRouter.admin_db).create( - user=user, tenant=tenant, role=Membership.RoleChoices.OWNER - ) - role = Role.objects.using(MainRouter.admin_db).create( - name="admin", - tenant_id=tenant.id, - manage_users=True, - manage_account=True, - manage_billing=True, - manage_providers=True, - manage_integrations=True, - manage_scans=True, - unlimited_visibility=True, - ) - UserRoleRelationship.objects.using(MainRouter.admin_db).create( + invitation_token = self._get_invitation_token(request) + if invitation_token: + invitation, _ = accept_invitation_for_user( user=user, - role=role, - tenant_id=tenant.id, + invitation_token=invitation_token, ) + request.prowler_invitation_token = invitation_token + request.prowler_invitation_tenant_id = str(invitation.tenant_id) + else: + tenant = Tenant.objects.using(MainRouter.admin_db).create( + name=f"{user.email.split('@')[0]} default tenant" + ) + with rls_transaction(str(tenant.id)): + Membership.objects.using(MainRouter.admin_db).create( + user=user, tenant=tenant, role=Membership.RoleChoices.OWNER + ) + role = Role.objects.using(MainRouter.admin_db).create( + name="admin", + tenant_id=tenant.id, + manage_users=True, + manage_account=True, + manage_billing=True, + manage_providers=True, + manage_integrations=True, + manage_scans=True, + unlimited_visibility=True, + ) + UserRoleRelationship.objects.using(MainRouter.admin_db).create( + user=user, + role=role, + tenant_id=tenant.id, + ) else: request.session["saml_user_created"] = str(user.id) diff --git a/api/src/backend/api/tests/test_adapters.py b/api/src/backend/api/tests/test_adapters.py index 91d3bb054a..1df9908ee5 100644 --- a/api/src/backend/api/tests/test_adapters.py +++ b/api/src/backend/api/tests/test_adapters.py @@ -5,7 +5,7 @@ import pytest from allauth.socialaccount.models import SocialLogin from api.adapters import ProwlerSocialAccountAdapter from api.db_router import MainRouter -from api.models import SAMLConfiguration +from api.models import Invitation, Membership, SAMLConfiguration, Tenant from django.contrib.auth import get_user_model User = get_user_model() @@ -188,6 +188,44 @@ class TestProwlerSocialAccountAdapter: _, called_user = call_args[0] assert called_user.email == create_test_user.email + def test_save_user_social_with_invitation_joins_invited_tenant( + self, rf, create_test_user, tenants_fixture + ): + adapter = ProwlerSocialAccountAdapter() + invited_tenant = tenants_fixture[2] + invited_email = "frank-invited@example.com" + invitation = Invitation.objects.create( + tenant=invited_tenant, + email=invited_email, + inviter=create_test_user, + ) + request = rf.post("/", data={"invitation_token": invitation.token}) + request.session = {} + + sociallogin = MagicMock(spec=SocialLogin) + sociallogin.provider = MagicMock() + sociallogin.provider.id = "google" + sociallogin.account = MagicMock() + sociallogin.account.extra_data = {"name": "Frank"} + + real_user = User.objects.create_user( + name="Frank", email=invited_email, password="Secret123!" + ) + tenants_before = Tenant.objects.count() + + with patch("api.adapters.super") as mock_super: + mock_super.return_value.save_user.return_value = real_user + adapter.save_user(request, sociallogin) + + invitation.refresh_from_db() + assert invitation.state == Invitation.State.ACCEPTED + assert Tenant.objects.count() == tenants_before + assert Membership.objects.filter( + user=real_user, + tenant=invited_tenant, + role=Membership.RoleChoices.MEMBER, + ).exists() + def test_save_user_saml_sets_session_flag(self, rf): adapter = ProwlerSocialAccountAdapter() request = rf.get("/") diff --git a/api/src/backend/api/tests/test_views.py b/api/src/backend/api/tests/test_views.py index b159a284d3..a4951d936e 100644 --- a/api/src/backend/api/tests/test_views.py +++ b/api/src/backend/api/tests/test_views.py @@ -59,7 +59,11 @@ from api.models import ( from api.rls import Tenant from api.uuid_utils import datetime_to_uuid7 from api.v1.serializers import TokenSerializer -from api.v1.views import ComplianceOverviewViewSet, TenantFinishACSView +from api.v1.views import ( + ComplianceOverviewViewSet, + CustomSAMLLoginView, + TenantFinishACSView, +) from botocore.exceptions import ClientError, NoCredentialsError from conftest import ( API_JSON_CONTENT_TYPE, @@ -8598,16 +8602,14 @@ class TestInvitationViewSet: expires_at=self.TOMORROW, ) - data = { - "invitation_token": invitation.token, - } + data = {"invitation_token": invitation.token} assert not Membership.objects.filter( user__email__iexact=user.email, tenant=tenant ).exists() response = authenticated_client.post( - reverse("invitation-accept"), data=data, format="json" + reverse("invitation-accept"), data=data, format="vnd.api+json" ) assert response.status_code == status.HTTP_201_CREATED invitation.refresh_from_db() @@ -8616,13 +8618,46 @@ class TestInvitationViewSet: ).exists() assert invitation.state == Invitation.State.ACCEPTED.value - def test_invitations_accept_invitation_invalid_token(self, authenticated_client): - data = { - "invitation_token": "invalid_token", - } + def test_invitations_accept_invitation_existing_membership( + self, + authenticated_client, + create_test_user, + tenants_fixture, + ): + *_, tenant = tenants_fixture + user = create_test_user + + invitation = Invitation.objects.create( + tenant=tenant, + email=TEST_USER, + inviter=user, + expires_at=self.TOMORROW, + ) + Membership.objects.create(user=user, tenant=tenant) + + data = {"invitation_token": invitation.token} response = authenticated_client.post( - reverse("invitation-accept"), data=data, format="json" + reverse("invitation-accept"), + data=data, + format="vnd.api+json", + ) + + assert response.status_code == status.HTTP_201_CREATED + invitation.refresh_from_db() + assert invitation.state == Invitation.State.ACCEPTED.value + assert ( + Membership.objects.filter( + user__email__iexact=user.email, tenant=tenant + ).count() + == 1 + ) + + def test_invitations_accept_invitation_invalid_token(self, authenticated_client): + data = {"invitation_token": "invalid_token"} + + response = authenticated_client.post( + reverse("invitation-accept"), data=data, format="vnd.api+json" ) assert response.status_code == status.HTTP_404_NOT_FOUND @@ -8636,12 +8671,10 @@ class TestInvitationViewSet: invitation.email = TEST_USER invitation.save() - data = { - "invitation_token": invitation.token, - } + data = {"invitation_token": invitation.token} response = authenticated_client.post( - reverse("invitation-accept"), data=data, format="json" + reverse("invitation-accept"), data=data, format="vnd.api+json" ) assert response.status_code == status.HTTP_410_GONE @@ -8677,12 +8710,10 @@ class TestInvitationViewSet: invitation.email = TEST_USER invitation.save() - data = { - "invitation_token": invitation.token, - } + data = {"invitation_token": invitation.token} response = authenticated_client.post( - reverse("invitation-accept"), data=data, format="json" + reverse("invitation-accept"), data=data, format="vnd.api+json" ) assert response.status_code == status.HTTP_400_BAD_REQUEST @@ -8700,12 +8731,10 @@ class TestInvitationViewSet: invitation.email = TEST_USER invitation.save() - data = { - "invitation_token": invitation.token, - } + data = {"invitation_token": invitation.token} response = authenticated_client.post( - reverse("invitation-accept"), data=data, format="json" + reverse("invitation-accept"), data=data, format="vnd.api+json" ) assert response.status_code == status.HTTP_400_BAD_REQUEST @@ -13697,6 +13726,26 @@ class TestSAMLTokenValidation: assert response2.status_code == status.HTTP_404_NOT_FOUND +@pytest.mark.django_db +class TestCustomSAMLLoginView: + def test_dispatch_clears_stale_callback_url_when_request_has_none(self): + request = RequestFactory().get("/api/v1/saml/login/testtenant/") + request.session = { + "saml_callback_url": "/invitation/accept?invitation_token=old-token" + } + + with patch( + "allauth.socialaccount.providers.saml.views.LoginView.dispatch", + return_value=JsonResponse({}), + ): + response = CustomSAMLLoginView.as_view()( + request, organization_slug="testtenant" + ) + + assert response.status_code == status.HTTP_200_OK + assert "saml_callback_url" not in request.session + + @pytest.mark.django_db class TestSAMLInitiateAPIView: def test_valid_email_domain_and_certificates( @@ -13708,7 +13757,7 @@ class TestSAMLInitiateAPIView: url = reverse("api_saml_initiate") payload = {"email_domain": saml_setup["email"]} - response = authenticated_client.post(url, data=payload, format="json") + response = authenticated_client.post(url, data=payload, format="vnd.api+json") assert response.status_code == status.HTTP_302_FOUND assert ( @@ -13717,11 +13766,42 @@ class TestSAMLInitiateAPIView: ) assert "SAMLRequest" not in response.url + def test_valid_email_domain_preserves_safe_callback_url( + self, authenticated_client, saml_setup + ): + url = reverse("api_saml_initiate") + callback_url = "/invitation/accept?invitation_token=test-token" + payload = { + "email_domain": saml_setup["email"], + "callback_url": callback_url, + } + + response = authenticated_client.post(url, data=payload, format="vnd.api+json") + + assert response.status_code == status.HTTP_302_FOUND + query_params = parse_qs(urlparse(response.url).query) + assert query_params["callback_url"] == [callback_url] + + def test_valid_email_domain_rejects_external_callback_url( + self, authenticated_client, saml_setup + ): + url = reverse("api_saml_initiate") + payload = { + "email_domain": saml_setup["email"], + "callback_url": "https://attacker.example/invitation", + } + + response = authenticated_client.post(url, data=payload, format="vnd.api+json") + + assert response.status_code == status.HTTP_302_FOUND + query_params = parse_qs(urlparse(response.url).query) + assert "callback_url" not in query_params + def test_invalid_email_domain(self, authenticated_client): url = reverse("api_saml_initiate") payload = {"email_domain": "user@unauthorized.com"} - response = authenticated_client.post(url, data=payload, format="json") + response = authenticated_client.post(url, data=payload, format="vnd.api+json") assert response.status_code == status.HTTP_403_FORBIDDEN assert response.json()["errors"]["detail"] == "Unauthorized domain." @@ -13904,7 +13984,8 @@ class TestTenantFinishACSView: ) ) request.user = user - request.session = {} + callback_url = "/invitation/accept?invitation_token=test-token" + request.session = {"saml_callback_url": callback_url} with ( patch( @@ -13946,6 +14027,7 @@ class TestTenantFinishACSView: assert parsed_url.netloc == expected_callback_host query_params = parse_qs(parsed_url.query) assert "id" in query_params + assert query_params["callbackUrl"] == [callback_url] token_id = query_params["id"][0] token_obj = SAMLToken.objects.get(id=token_id) diff --git a/api/src/backend/api/utils.py b/api/src/backend/api/utils.py index ce1dc0f10d..bb636b1bfa 100644 --- a/api/src/backend/api/utils.py +++ b/api/src/backend/api/utils.py @@ -7,9 +7,19 @@ from allauth.socialaccount.providers.oauth2.client import OAuth2Client from api.db_router import MainRouter from api.db_utils import rls_transaction from api.exceptions import InvitationTokenExpiredException -from api.models import Integration, Invitation, Processor, Provider, Resource +from api.models import ( + Integration, + Invitation, + Membership, + Processor, + Provider, + Resource, + Role, + UserRoleRelationship, +) from api.v1.serializers import FindingMetadataSerializer from django.contrib.postgres.aggregates import ArrayAgg +from django.db import transaction from django.db.models import Subquery from prowler.lib.outputs.jira.jira import Jira, JiraBasicAuthError from prowler.providers.aws.lib.s3.s3 import S3 @@ -538,6 +548,35 @@ def validate_invitation( return invitation +def accept_invitation_for_user( + *, user, invitation_token: str, raise_not_found: bool = False +): + with transaction.atomic(using=MainRouter.admin_db): + invitation = validate_invitation( + invitation_token, user.email, raise_not_found=raise_not_found + ) + with rls_transaction(str(invitation.tenant_id), using=MainRouter.admin_db): + membership, _ = Membership.objects.using(MainRouter.admin_db).get_or_create( + user=user, + tenant=invitation.tenant, + defaults={"role": Membership.RoleChoices.MEMBER}, + ) + invitation_roles = Role.objects.using(MainRouter.admin_db).filter( + invitations=invitation + ) + for role in invitation_roles: + UserRoleRelationship.objects.using(MainRouter.admin_db).get_or_create( + user=user, + role=role, + defaults={"tenant": invitation.tenant}, + ) + + invitation.state = Invitation.State.ACCEPTED + invitation.save(using=MainRouter.admin_db) + + return invitation, membership + + # ToRemove after removing the fallback mechanism in /findings/metadata def get_findings_metadata_no_aggregations(tenant_id: str, filtered_queryset): filtered_ids = filtered_queryset.order_by().values("id") diff --git a/api/src/backend/api/v1/serializers.py b/api/src/backend/api/v1/serializers.py index 1d160b4048..a88f53fc7d 100644 --- a/api/src/backend/api/v1/serializers.py +++ b/api/src/backend/api/v1/serializers.py @@ -3147,6 +3147,9 @@ class ProcessorUpdateSerializer(BaseWriteSerializer): class SamlInitiateSerializer(BaseSerializerV1): email_domain = serializers.CharField() + callback_url = serializers.CharField( + required=False, allow_blank=True, max_length=2048 + ) class JSONAPIMeta: resource_name = "saml-initiate" diff --git a/api/src/backend/api/v1/views.py b/api/src/backend/api/v1/views.py index 588c65c30d..e7907cf6f8 100644 --- a/api/src/backend/api/v1/views.py +++ b/api/src/backend/api/v1/views.py @@ -9,7 +9,7 @@ from collections import defaultdict from copy import deepcopy from datetime import UTC, datetime, timedelta from decimal import ROUND_HALF_UP, Decimal, InvalidOperation -from urllib.parse import urljoin +from urllib.parse import urlencode, urljoin import sentry_sdk from allauth.socialaccount.models import SocialAccount, SocialApp @@ -129,6 +129,7 @@ from api.renderers import APIJSONRenderer, PlainTextRenderer from api.rls import Tenant from api.utils import ( CustomOAuth2Client, + accept_invitation_for_user, get_findings_metadata_no_aggregations, initialize_prowler_integration, initialize_prowler_provider, @@ -542,6 +543,46 @@ class SchemaView(SpectacularAPIView): return super().get(request, *args, **kwargs) +SAML_CALLBACK_SESSION_KEY = "saml_callback_url" + + +def _safe_callback_path(value): + if not value or not isinstance(value, str): + return None + if not value.startswith("/") or value.startswith("//"): + return None + return value + + +def _get_request_invitation_token(request): + for source_name in ("data", "POST"): + data = getattr(request, source_name, None) or {} + if not hasattr(data, "get"): + continue + invitation_token = data.get("invitation_token") + if invitation_token: + return invitation_token + + wrapped_request = getattr(request, "_request", None) + if wrapped_request and wrapped_request is not request: + return _get_request_invitation_token(wrapped_request) + + return None + + +def _accept_social_invitation(request, user): + invitation_token = _get_request_invitation_token(request) + tenant_id = getattr(request, "prowler_invitation_tenant_id", None) + if invitation_token and not tenant_id: + invitation, _ = accept_invitation_for_user( + user=user, + invitation_token=invitation_token, + raise_not_found=True, + ) + tenant_id = str(invitation.tenant_id) + return tenant_id + + @extend_schema(exclude=True) class GoogleSocialLoginView(SocialLoginView): adapter_class = GoogleOAuth2Adapter @@ -552,7 +593,11 @@ class GoogleSocialLoginView(SocialLoginView): original_response = super().get_response() if self.user and self.user.is_authenticated: - serializer = TokenSocialLoginSerializer(data={"email": self.user.email}) + tenant_id = _accept_social_invitation(self.request, self.user) + serializer_data = {"email": self.user.email} + if tenant_id: + serializer_data["tenant_id"] = tenant_id + serializer = TokenSocialLoginSerializer(data=serializer_data) try: serializer.is_valid(raise_exception=True) except TokenError as e: @@ -577,7 +622,11 @@ class GithubSocialLoginView(SocialLoginView): original_response = super().get_response() if self.user and self.user.is_authenticated: - serializer = TokenSocialLoginSerializer(data={"email": self.user.email}) + tenant_id = _accept_social_invitation(self.request, self.user) + serializer_data = {"email": self.user.email} + if tenant_id: + serializer_data["tenant_id"] = tenant_id + serializer = TokenSocialLoginSerializer(data=serializer_data) try: serializer.is_valid(raise_exception=True) @@ -637,6 +686,10 @@ class CustomSAMLLoginView(LoginView): This approach maintains security while providing better UX. """ + callback_url = _safe_callback_path(request.GET.get("callback_url")) + request.session.pop(SAML_CALLBACK_SESSION_KEY, None) + if callback_url: + request.session[SAML_CALLBACK_SESSION_KEY] = callback_url if request.method == "GET": # Convert GET to POST while preserving parameters request.method = "POST" @@ -681,6 +734,11 @@ class SAMLInitiateAPIView(GenericAPIView): "saml_login", kwargs={"organization_slug": config.email_domain} ) login_url = urljoin(api_host, login_path) + callback_url = _safe_callback_path( + serializer.validated_data.get("callback_url") + ) + if callback_url: + login_url = f"{login_url}?{urlencode({'callback_url': callback_url})}" return redirect(login_url) @@ -896,7 +954,13 @@ class TenantFinishACSView(FinishACSView): token=token_data, user=user ) callback_url = env.str("SAML_SSO_CALLBACK_URL") - redirect_url = f"{callback_url}?id={saml_token.id}" + redirect_params = {"id": str(saml_token.id)} + saml_callback_url = _safe_callback_path( + request.session.pop(SAML_CALLBACK_SESSION_KEY, None) + ) + if saml_callback_url: + redirect_params["callbackUrl"] = saml_callback_url + redirect_url = f"{callback_url}?{urlencode(redirect_params)}" request.session.pop("saml_user_created", None) return redirect(redirect_url) @@ -4389,25 +4453,12 @@ class InvitationAcceptViewSet(BaseRLSViewSet): invitation_token = serializer.validated_data["invitation_token"] user_email = request.user.email - invitation = validate_invitation( - invitation_token, user_email, raise_not_found=True - ) - - # Proceed with accepting the invitation user = User.objects.using(MainRouter.admin_db).get(email=user_email) - membership = Membership.objects.using(MainRouter.admin_db).create( + invitation, membership = accept_invitation_for_user( user=user, - tenant=invitation.tenant, + invitation_token=invitation_token, + raise_not_found=True, ) - user_role = [] - for role in invitation.roles.all(): - user_role.append( - UserRoleRelationship.objects.using(MainRouter.admin_db).create( - user=user, role=role, tenant=invitation.tenant - ) - ) - invitation.state = Invitation.State.ACCEPTED - invitation.save(using=MainRouter.admin_db) self.response_serializer_class = MembershipSerializer membership_serializer = self.get_serializer(membership) diff --git a/ui/CHANGELOG.md b/ui/CHANGELOG.md index 7a0a7de47a..d74a3111c1 100644 --- a/ui/CHANGELOG.md +++ b/ui/CHANGELOG.md @@ -2,6 +2,14 @@ All notable changes to the **Prowler UI** are documented in this file. +## [1.32.1] (Prowler UNRELEASED) + +### 🐞 Fixed + +- Invitation callback paths are now preserved when invited users continue with Google, GitHub, or SAML authentication [(#11752)](https://github.com/prowler-cloud/prowler/pull/11752) + +--- + ## [1.32.0] (Prowler v5.32.0) ### 🚀 Added diff --git a/ui/actions/integrations/saml.ts b/ui/actions/integrations/saml.ts index 05a58e9e5e..41e153a46a 100644 --- a/ui/actions/integrations/saml.ts +++ b/ui/actions/integrations/saml.ts @@ -166,8 +166,13 @@ export const deleteSamlConfig = async (id: string) => { } }; -export const initiateSamlAuth = async (email: string) => { +export const initiateSamlAuth = async (email: string, callbackUrl = "/") => { try { + const attributes = { + email_domain: email, + ...(callbackUrl !== "/" && { callback_url: callbackUrl }), + }; + const response = await fetch(`${apiBaseUrl}/auth/saml/initiate/`, { method: "POST", headers: { @@ -176,9 +181,7 @@ export const initiateSamlAuth = async (email: string) => { body: JSON.stringify({ data: { type: "saml-initiate", - attributes: { - email_domain: email, - }, + attributes, }, }), redirect: "manual", diff --git a/ui/app/(auth)/(guest-only)/sign-up/page.tsx b/ui/app/(auth)/(guest-only)/sign-up/page.tsx index b846c59a93..4854de012b 100644 --- a/ui/app/(auth)/(guest-only)/sign-up/page.tsx +++ b/ui/app/(auth)/(guest-only)/sign-up/page.tsx @@ -13,6 +13,7 @@ const SignUp = async ({ typeof resolvedSearchParams?.invitation_token === "string" ? resolvedSearchParams.invitation_token : null; + const isCloudEnv = process.env.NEXT_PUBLIC_IS_CLOUD_ENV === "true"; const GOOGLE_AUTH_URL = getAuthUrl("google"); const GITHUB_AUTH_URL = getAuthUrl("github"); @@ -21,6 +22,7 @@ const SignUp = async ({ { const samlError = searchParams.get("sso_saml_failed"); @@ -102,7 +103,7 @@ export const SignInForm = ({ form.setValue("password", ""); } - const result = await initiateSamlAuth(email); + const result = await initiateSamlAuth(email, callbackUrl); if (result.success && result.redirectUrl) { window.location.href = result.redirectUrl; @@ -181,6 +182,7 @@ export const SignInForm = ({ diff --git a/ui/components/auth/oss/sign-up-form.tsx b/ui/components/auth/oss/sign-up-form.tsx index e18ded5845..df5cb0ba12 100644 --- a/ui/components/auth/oss/sign-up-form.tsx +++ b/ui/components/auth/oss/sign-up-form.tsx @@ -41,12 +41,14 @@ const FORM_ERROR_TYPE = { export const SignUpForm = ({ invitationToken, + isCloudEnv, googleAuthUrl, githubAuthUrl, isGoogleOAuthEnabled, isGithubOAuthEnabled, }: { invitationToken?: string | null; + isCloudEnv?: boolean; googleAuthUrl?: string; githubAuthUrl?: string; isGoogleOAuthEnabled?: boolean; @@ -54,6 +56,9 @@ export const SignUpForm = ({ }) => { const router = useRouter(); const { toast } = useToast(); + const callbackUrl = invitationToken + ? `/invitation/accept?invitation_token=${encodeURIComponent(invitationToken)}` + : "/"; const form = useForm({ resolver: zodResolver(signUpSchema), @@ -75,8 +80,14 @@ export const SignUpForm = ({ name: "password", defaultValue: "", }); + const termsAccepted = useWatch({ + control: form.control, + name: "termsAndConditions", + defaultValue: false, + }); const isLoading = form.formState.isSubmitting; + const isSocialAuthDisabled = Boolean(isCloudEnv && !termsAccepted); const onSubmit = async (data: SignUpFormData) => { const newUser = await createNewUser(data); @@ -88,7 +99,7 @@ export const SignUpForm = ({ }); form.reset(); - if (process.env.NEXT_PUBLIC_IS_CLOUD_ENV === "true") { + if (isCloudEnv) { router.push("/email-verification"); } else { router.push("/sign-in"); @@ -200,7 +211,7 @@ export const SignUpForm = ({ /> )} - {process.env.NEXT_PUBLIC_IS_CLOUD_ENV === "true" && ( + {isCloudEnv && ( - {!invitationToken && ( + {(!invitationToken || isCloudEnv) && ( <>
diff --git a/ui/components/auth/oss/social-buttons.tsx b/ui/components/auth/oss/social-buttons.tsx index 0dd92c48c0..65a891a743 100644 --- a/ui/components/auth/oss/social-buttons.tsx +++ b/ui/components/auth/oss/social-buttons.tsx @@ -1,83 +1,153 @@ -import { Tooltip } from "@heroui/tooltip"; import { Icon } from "@iconify/react"; +import type { ReactNode } from "react"; -import { Button } from "@/components/shadcn"; +import { + Button, + Tooltip, + TooltipContent, + TooltipTrigger, +} from "@/components/shadcn"; import { CustomLink } from "@/components/ui/custom/custom-link"; +import { appendCallbackState } from "@/lib/auth-callback-url"; + +type SocialProvider = { + key: string; + label: string; + url?: string; + isOAuthEnabled?: boolean; + enabledIcon: string; + disabledIcon: string; + disabledDocs: { + message: string; + href: string; + }; +}; + +const SocialButton = ({ + provider, + isDisabled, + disabledTooltipContent, +}: { + provider: SocialProvider; + isDisabled: boolean; + disabledTooltipContent: ReactNode; +}) => { + const button = ( + + ); + + if (!isDisabled) { + return button; + } + + return ( + + + {button} + + + {provider.isOAuthEnabled ? ( + disabledTooltipContent + ) : ( +
+ {provider.disabledDocs.message}{" "} + + Read the docs + +
+ )} +
+
+ ); +}; export const SocialButtons = ({ googleAuthUrl, githubAuthUrl, + callbackUrl = "/", isGoogleOAuthEnabled, isGithubOAuthEnabled, + isDisabled = false, + disabledTooltipContent, }: { googleAuthUrl?: string; githubAuthUrl?: string; + callbackUrl?: string; isGoogleOAuthEnabled?: boolean; isGithubOAuthEnabled?: boolean; -}) => ( - <> - - Social Login with Google is not enabled.{" "} - - Read the docs - - - } - placement="top" - shadow="sm" - isDisabled={isGoogleOAuthEnabled} - className="w-96" - > - - - - - - Social Login with Github is not enabled.{" "} - - Read the docs - - - } - placement="top" - shadow="sm" - isDisabled={isGithubOAuthEnabled} - className="w-96" - > - - - - - -); + isDisabled?: boolean; + disabledTooltipContent?: ReactNode; +}) => { + const googleUrl = googleAuthUrl + ? appendCallbackState(googleAuthUrl, callbackUrl) + : undefined; + const githubUrl = githubAuthUrl + ? appendCallbackState(githubAuthUrl, callbackUrl) + : undefined; + const socialDisabledTooltip = + disabledTooltipContent || "Social login is currently unavailable."; + + const providers: SocialProvider[] = [ + { + key: "google", + label: "Continue with Google", + url: googleUrl, + isOAuthEnabled: isGoogleOAuthEnabled, + enabledIcon: "flat-color-icons:google", + disabledIcon: "simple-icons:google", + disabledDocs: { + message: "Social Login with Google is not enabled.", + href: "https://docs.prowler.com/projects/prowler-open-source/en/latest/tutorials/prowler-app-social-login/#google-oauth-configuration", + }, + }, + { + key: "github", + label: "Continue with Github", + url: githubUrl, + isOAuthEnabled: isGithubOAuthEnabled, + enabledIcon: "simple-icons:github", + disabledIcon: "simple-icons:github", + disabledDocs: { + message: "Social Login with Github is not enabled.", + href: "https://docs.prowler.com/projects/prowler-open-source/en/latest/tutorials/prowler-app-social-login/#github-oauth-configuration", + }, + }, + ]; + + return ( + <> + {providers.map((provider) => ( + + ))} + + ); +}; diff --git a/ui/lib/auth-callback-url.test.ts b/ui/lib/auth-callback-url.test.ts new file mode 100644 index 0000000000..71161c7c98 --- /dev/null +++ b/ui/lib/auth-callback-url.test.ts @@ -0,0 +1,108 @@ +import { describe, expect, it } from "vitest"; + +import { + appendCallbackState, + getInvitationTokenFromCallbackPath, + getSafeCallbackPath, +} from "@/lib/auth-callback-url"; + +describe("auth callback URL helpers", () => { + describe("when appending OAuth state", () => { + it("should add a relative callback path as provider state", () => { + const authUrl = "https://accounts.example.com/oauth?client_id=client"; + const callbackPath = "/invitation/accept?invitation_token=test-token"; + + const result = appendCallbackState(authUrl, callbackPath); + + expect(new URL(result).searchParams.get("state")).toBe(callbackPath); + }); + + it("should not add state for the default callback path", () => { + const authUrl = "https://accounts.example.com/oauth?client_id=client"; + + const result = appendCallbackState(authUrl, "/"); + + expect(new URL(result).searchParams.has("state")).toBe(false); + }); + }); + + describe("when reading callback paths", () => { + it("should return relative callback paths", () => { + const params = new URLSearchParams({ + state: "/invitation/accept?invitation_token=test-token", + }); + + const result = getSafeCallbackPath(params); + + expect(result).toBe("/invitation/accept?invitation_token=test-token"); + }); + + it("should reject external callback URLs", () => { + const params = new URLSearchParams({ + state: "https://attacker.example/phishing", + }); + + const result = getSafeCallbackPath(params); + + expect(result).toBe("/"); + }); + + it("should reject protocol-relative callback URLs", () => { + const params = new URLSearchParams({ + state: "//attacker.example/phishing", + }); + + const result = getSafeCallbackPath(params); + + expect(result).toBe("/"); + }); + + it("should reject backslash-normalized callback URLs", () => { + const params = new URLSearchParams({ state: "/\\attacker.example" }); + + const result = getSafeCallbackPath(params); + + expect(result).toBe("/"); + }); + + it("should reject callback URLs with control characters before the host", () => { + const params = new URLSearchParams({ state: "/\t/attacker.example" }); + + const result = getSafeCallbackPath(params); + + expect(result).toBe("/"); + }); + + it("should preserve the query string of relative callback paths", () => { + const params = new URLSearchParams({ + state: "/invitation/accept?invitation_token=test-token&foo=bar", + }); + + const result = getSafeCallbackPath(params); + + expect(result).toBe( + "/invitation/accept?invitation_token=test-token&foo=bar", + ); + }); + }); + + describe("when appending OAuth state for unsafe paths", () => { + it("should not add a backslash-normalized path as provider state", () => { + const authUrl = "https://accounts.example.com/oauth?client_id=client"; + + const result = appendCallbackState(authUrl, "/\\attacker.example"); + + expect(new URL(result).searchParams.has("state")).toBe(false); + }); + }); + + describe("when reading invitation tokens", () => { + it("should return invitation tokens from safe callback paths", () => { + const callbackPath = "/invitation/accept?invitation_token=test-token"; + + const result = getInvitationTokenFromCallbackPath(callbackPath); + + expect(result).toBe("test-token"); + }); + }); +}); diff --git a/ui/lib/auth-callback-url.ts b/ui/lib/auth-callback-url.ts new file mode 100644 index 0000000000..c2bdd6be9c --- /dev/null +++ b/ui/lib/auth-callback-url.ts @@ -0,0 +1,65 @@ +const DEFAULT_CALLBACK_PATH = "/"; +const INVITATION_TOKEN_PARAM = "invitation_token"; +// Origin used only to resolve relative paths; never part of the returned value. +const INTERNAL_ORIGIN = "http://localhost"; + +type CallbackSearchParams = { + get(name: string): string | null; +}; + +export const getSafeCallbackPathFromValue = ( + value: string | null | undefined, +) => { + if (!value || !value.startsWith("/") || value.startsWith("//")) { + return DEFAULT_CALLBACK_PATH; + } + + // A prefix check is not enough: the URL parser normalizes backslashes and + // control characters, so "/\evil.com" or "/\t/evil.com" pass the check above + // yet resolve to an external origin. Resolve against a fixed origin and + // confirm it stayed internal before trusting the path. + try { + const url = new URL(value, INTERNAL_ORIGIN); + if (url.origin !== INTERNAL_ORIGIN) { + return DEFAULT_CALLBACK_PATH; + } + + return `${url.pathname}${url.search}${url.hash}`; + } catch (_error) { + return DEFAULT_CALLBACK_PATH; + } +}; + +export const getSafeCallbackPath = ( + searchParams: CallbackSearchParams, + key = "state", +) => getSafeCallbackPathFromValue(searchParams.get(key)); + +export const appendCallbackState = (authUrl: string, callbackPath: string) => { + const safeCallbackPath = getSafeCallbackPathFromValue(callbackPath); + if (safeCallbackPath === DEFAULT_CALLBACK_PATH) { + return authUrl; + } + + try { + const url = new URL(authUrl); + url.searchParams.set("state", safeCallbackPath); + return url.toString(); + } catch (_error) { + return authUrl; + } +}; + +export const getInvitationTokenFromCallbackPath = (callbackPath: string) => { + const safeCallbackPath = getSafeCallbackPathFromValue(callbackPath); + if (safeCallbackPath === DEFAULT_CALLBACK_PATH) { + return null; + } + + try { + const url = new URL(safeCallbackPath, "http://localhost"); + return url.searchParams.get(INVITATION_TOKEN_PARAM); + } catch (_error) { + return null; + } +};