diff --git a/.env b/.env index 136531ec2d..3f80a2460b 100644 --- a/.env +++ b/.env @@ -110,11 +110,20 @@ DJANGO_OUTPUT_S3_AWS_SECRET_ACCESS_KEY="" DJANGO_OUTPUT_S3_AWS_SESSION_TOKEN="" # The AWS region where your S3 bucket is located (e.g., "us-east-1") +# Required if the bucket uses SSE-KMS: download URLs are then signed with SigV4, which +# is scoped to this region, so it must match the bucket's DJANGO_OUTPUT_S3_AWS_DEFAULT_REGION="" # The name of the S3 bucket where scan output should be stored DJANGO_OUTPUT_S3_AWS_OUTPUT_BUCKET="" +# The storage endpoint the API and Celery workers use to upload and list scan output +# (e.g. "http://minio:9000"). Leave empty on AWS S3. Set it when scan output is stored on +# S3-compatible object storage such as MinIO instead of real S3. +# If set without DJANGO_OUTPUT_S3_AWS_PUBLIC_ENDPOINT_URL below, report download URLs are +# signed against this internal host, and a browser outside the container network cannot open them. +DJANGO_OUTPUT_S3_AWS_ENDPOINT_URL="" + # The storage address the browser can reach, used only to sign report download URLs # (e.g. "https://storage.example.com"). Leave empty on AWS S3. Set it when storage is # only reachable inside the container network, such as MinIO on "http://minio:9000". @@ -165,7 +174,7 @@ SENTRY_RELEASE=local # REO_DEV_CLIENT_ID= #### Prowler release version #### -NEXT_PUBLIC_PROWLER_RELEASE_VERSION=v5.43.0 +NEXT_PUBLIC_PROWLER_RELEASE_VERSION=v5.44.0 # Social login credentials SOCIAL_GOOGLE_OAUTH_CALLBACK_URL="${AUTH_URL}/api/auth/callback/google" diff --git a/.github/workflows/ui-e2e-tests-v2.yml b/.github/workflows/ui-e2e-tests-v2.yml index 2bf6db8a5a..8cf43f9d66 100644 --- a/.github/workflows/ui-e2e-tests-v2.yml +++ b/.github/workflows/ui-e2e-tests-v2.yml @@ -11,6 +11,7 @@ on: - "v5.*" paths: - ".github/workflows/ui-e2e-tests-v2.yml" + - ".github/workflows/test-impact-analysis.yml" - ".github/test-impact.yml" - "ui/**" - "api/**" # API changes can affect UI E2E @@ -237,8 +238,8 @@ jobs: - name: Add AWS credentials for testing run: | - echo "AWS_ACCESS_KEY_ID=${{ secrets.E2E_AWS_PROVIDER_ACCESS_KEY }}" >> .env - echo "AWS_SECRET_ACCESS_KEY=${{ secrets.E2E_AWS_PROVIDER_SECRET_KEY }}" >> .env + echo "AWS_ACCESS_KEY_ID=${E2E_AWS_PROVIDER_ACCESS_KEY}" >> .env + echo "AWS_SECRET_ACCESS_KEY=${E2E_AWS_PROVIDER_SECRET_KEY}" >> .env - name: Build API image from current code # docker-compose.yml references prowlercloud/prowler-api:latest from the registry, diff --git a/.grype.yaml b/.grype.yaml index 26fec99bdd..5028944529 100644 --- a/.grype.yaml +++ b/.grype.yaml @@ -115,3 +115,14 @@ ignore: - vulnerability: CVE-2026-9669 package: name: python + # CVE-2026-82049 (tarfile data/tar filter bypass via a hard link to a symlink) has no + # fixed CPython release on any branch: the fix is merged on main and 3.13 only, and the + # 3.12 backport is still open. Grype records 3.14.0b1 as the fix, so only-fixed does not + # drop it, yet python:3.12.14-slim-trixie reports it too. Prowler never extracts tar + # archives to disk: the ECR image inspection reads members in memory with extractfile(). + # Remove once the base image ships a 3.12 release that includes the backport. + # https://github.com/python/cpython/issues/157190 + # https://github.com/python/cpython/pull/157454 + - vulnerability: CVE-2026-82049 + package: + name: python diff --git a/api/CHANGELOG.md b/api/CHANGELOG.md index 318cd30d4f..1a927baa9c 100644 --- a/api/CHANGELOG.md +++ b/api/CHANGELOG.md @@ -4,6 +4,22 @@ All notable changes to the **Prowler API** are documented in this file. +## [1.44.0] (Prowler v5.43.0) + +### 🐞 Fixed + +- Report download URLs can be signed against a browser-reachable storage host via `DJANGO_OUTPUT_S3_AWS_PUBLIC_ENDPOINT_URL`, so downloads complete on deployments where storage is only reachable inside the container network [(#12552)](https://github.com/prowler-cloud/prowler/pull/12552) +- A scan report download no longer fails with a server error when `DJANGO_OUTPUT_S3_AWS_DEFAULT_REGION` is unset, which is common on storage with no meaningful region [(#12552)](https://github.com/prowler-cloud/prowler/pull/12552) +- Lapsed pending invitations are reported as expired and no longer block a new invitation for the same email [(#12831)](https://github.com/prowler-cloud/prowler/pull/12831) + +### 🔐 Security + +- `libsqlite3-0`, `gzip`, `perl-base` and `libpcre2-8-0` upgraded in the API container image, patching high Debian CVEs [(#12804)](https://github.com/prowler-cloud/prowler/pull/12804) +- PowerShell from 7.5.9 to 7.5.11 in the API container image, bundling .NET runtime 9.0.20 and patching CVE-2026-62901 [(#12811)](https://github.com/prowler-cloud/prowler/pull/12811) +- Bumped `anyio` to 4.14.2 to resolve CVE-2026-63374 [(#12848)](https://github.com/prowler-cloud/prowler/pull/12848) + +--- + ## [1.43.0] (Prowler v5.42.0) ### 🔄 Changed diff --git a/api/changelog.d/api-image-debian-cves.security.md b/api/changelog.d/api-image-debian-cves.security.md deleted file mode 100644 index 2d8b0403ad..0000000000 --- a/api/changelog.d/api-image-debian-cves.security.md +++ /dev/null @@ -1 +0,0 @@ -`libsqlite3-0`, `gzip`, `perl-base` and `libpcre2-8-0` upgraded in the API container image, patching high Debian CVEs diff --git a/api/changelog.d/api-image-powershell-dotnet-cve.security.md b/api/changelog.d/api-image-powershell-dotnet-cve.security.md deleted file mode 100644 index 28b91a32a0..0000000000 --- a/api/changelog.d/api-image-powershell-dotnet-cve.security.md +++ /dev/null @@ -1 +0,0 @@ -PowerShell from 7.5.9 to 7.5.11 in the API container image, bundling .NET runtime 9.0.20 and patching CVE-2026-62901 diff --git a/api/changelog.d/attack-paths-tmp-db-reaper.fixed.md b/api/changelog.d/attack-paths-tmp-db-reaper.fixed.md new file mode 100644 index 0000000000..b165fc6d40 --- /dev/null +++ b/api/changelog.d/attack-paths-tmp-db-reaper.fixed.md @@ -0,0 +1 @@ +Adds a periodic sweep that drops orphaned Attack Paths temp Neo4j scan databases left behind when a worker or Neo4j crashes mid-scan, before they accumulate unbounded diff --git a/api/changelog.d/ephemeral-resource-count-reset.fixed.md b/api/changelog.d/ephemeral-resource-count-reset.fixed.md new file mode 100644 index 0000000000..0c56ae2e5d --- /dev/null +++ b/api/changelog.d/ephemeral-resource-count-reset.fixed.md @@ -0,0 +1 @@ +Resources no longer keep a stale failed findings count forever when a scoped or imported scan for the same provider completes after a full scan, which used to make the full scan skip its own cleanup diff --git a/api/changelog.d/lapsed-invitations-block-re-invites.fixed.md b/api/changelog.d/lapsed-invitations-block-re-invites.fixed.md deleted file mode 100644 index 29a8e99235..0000000000 --- a/api/changelog.d/lapsed-invitations-block-re-invites.fixed.md +++ /dev/null @@ -1 +0,0 @@ -Lapsed pending invitations are reported as expired and no longer block a new invitation for the same email diff --git a/api/changelog.d/latest-scan-null-completed-at.fixed.md b/api/changelog.d/latest-scan-null-completed-at.fixed.md new file mode 100644 index 0000000000..9932210055 --- /dev/null +++ b/api/changelog.d/latest-scan-null-completed-at.fixed.md @@ -0,0 +1 @@ +Providers whose most recent completed scan has no `completed_at` timestamp are no longer missing from every endpoint that reports a provider's latest scan, which now falls back to scan creation order instead of skipping the provider diff --git a/api/changelog.d/latest-scan-selector.changed.md b/api/changelog.d/latest-scan-selector.changed.md new file mode 100644 index 0000000000..417cf9e997 --- /dev/null +++ b/api/changelog.d/latest-scan-selector.changed.md @@ -0,0 +1 @@ +Unify how every endpoint resolves a provider latest completed scan, so overlapping scans no longer make findings, compliance and mute rules read from different scans diff --git a/api/changelog.d/report-download-public-storage-endpoint.fixed.md b/api/changelog.d/report-download-public-storage-endpoint.fixed.md deleted file mode 100644 index 87004d8a56..0000000000 --- a/api/changelog.d/report-download-public-storage-endpoint.fixed.md +++ /dev/null @@ -1 +0,0 @@ -Report download URLs can be signed against a browser-reachable storage host via `DJANGO_OUTPUT_S3_AWS_PUBLIC_ENDPOINT_URL`, so downloads complete on deployments where storage is only reachable inside the container network diff --git a/api/changelog.d/s3-client-default-region.fixed.md b/api/changelog.d/s3-client-default-region.fixed.md deleted file mode 100644 index 62efc579f2..0000000000 --- a/api/changelog.d/s3-client-default-region.fixed.md +++ /dev/null @@ -1 +0,0 @@ -A scan report download no longer fails with a server error when `DJANGO_OUTPUT_S3_AWS_DEFAULT_REGION` is unset, which is common on storage with no meaningful region diff --git a/api/changelog.d/s3-output-internal-endpoint.added.md b/api/changelog.d/s3-output-internal-endpoint.added.md new file mode 100644 index 0000000000..88159f1ef8 --- /dev/null +++ b/api/changelog.d/s3-output-internal-endpoint.added.md @@ -0,0 +1 @@ +Scan output uploads and downloads can now target S3-compatible object storage such as MinIO directly via `DJANGO_OUTPUT_S3_AWS_ENDPOINT_URL`, instead of relying on process-wide AWS environment variables that also hijacked unrelated AWS API calls diff --git a/api/changelog.d/s3-report-download-sigv4.fixed.md b/api/changelog.d/s3-report-download-sigv4.fixed.md new file mode 100644 index 0000000000..7c11870c11 --- /dev/null +++ b/api/changelog.d/s3-report-download-sigv4.fixed.md @@ -0,0 +1 @@ +Scan report downloads from an S3 bucket with default SSE-KMS encryption no longer fail with an `InvalidArgument` error: when `DJANGO_OUTPUT_S3_AWS_DEFAULT_REGION` is set, presigned download URLs are signed with AWS Signature Version 4 for that region diff --git a/api/changelog.d/scan-create-task-args.fixed.md b/api/changelog.d/scan-create-task-args.fixed.md new file mode 100644 index 0000000000..26d42cf0c4 --- /dev/null +++ b/api/changelog.d/scan-create-task-args.fixed.md @@ -0,0 +1 @@ +`POST /api/v1/scans` again returns the new scan id in the response `task_args`, which had been empty since the scan broker publish moved to transaction commit diff --git a/api/changelog.d/worker-logging-restart-policy.fixed.md b/api/changelog.d/worker-logging-restart-policy.fixed.md new file mode 100644 index 0000000000..dad995fdd9 --- /dev/null +++ b/api/changelog.d/worker-logging-restart-policy.fixed.md @@ -0,0 +1 @@ +Celery loggers are now declared explicitly in `custom_logging.py` so fatal worker errors are no longer silenced by `disable_existing_loggers=True`. All long-running services in `docker-compose.yml` now have `restart: unless-stopped` so containers recover automatically after unexpected crashes. diff --git a/api/pyproject.toml b/api/pyproject.toml index 00b50e8300..a349bd72ea 100644 --- a/api/pyproject.toml +++ b/api/pyproject.toml @@ -71,7 +71,7 @@ name = "prowler-api" package-mode = false # Needed for the SDK compatibility requires-python = ">=3.11,<3.13" -version = "1.44.0" +version = "1.45.0" # Shared ruff baseline (kept in sync with mcp_server/pyproject.toml). # target-version tracks this project's lowest supported Python. @@ -137,7 +137,7 @@ constraint-dependencies = [ "aliyun-log-fastpb==0.2.0", "amqp==5.3.1", "annotated-types==0.7.0", - "anyio==4.12.1", + "anyio==4.14.2", "applicationinsights==0.11.10", "apscheduler==3.11.2", "argcomplete==3.5.3", diff --git a/api/src/backend/api/attack_paths/database.py b/api/src/backend/api/attack_paths/database.py index 3ef55b7eca..1db76c4307 100644 --- a/api/src/backend/api/attack_paths/database.py +++ b/api/src/backend/api/attack_paths/database.py @@ -207,6 +207,11 @@ def drop_database(database: str) -> None: sink_module.get_backend().drop_database(database) +def list_databases() -> list[str]: + """List database names on the ingest cluster. Temp scan DBs always live here.""" + return ingest.list_databases() + + def drop_subgraph(database: str, provider_id: str) -> int: return sink_module.get_backend().drop_subgraph(database, provider_id) diff --git a/api/src/backend/api/attack_paths/ingest/__init__.py b/api/src/backend/api/attack_paths/ingest/__init__.py index 5833b8b373..d95482ed85 100644 --- a/api/src/backend/api/attack_paths/ingest/__init__.py +++ b/api/src/backend/api/attack_paths/ingest/__init__.py @@ -13,6 +13,7 @@ from api.attack_paths.ingest.driver import ( get_session, get_uri, init_driver, + list_databases, run_cypher, ) @@ -25,5 +26,6 @@ __all__ = [ "get_session", "get_uri", "init_driver", + "list_databases", "run_cypher", ] diff --git a/api/src/backend/api/attack_paths/ingest/driver.py b/api/src/backend/api/attack_paths/ingest/driver.py index 1b05c721e7..5c8b573ad7 100644 --- a/api/src/backend/api/attack_paths/ingest/driver.py +++ b/api/src/backend/api/attack_paths/ingest/driver.py @@ -165,6 +165,14 @@ def drop_database(database: str) -> None: session.run(f"DROP DATABASE `{database}` IF EXISTS DESTROY DATA") +def list_databases() -> list[str]: + """List every database name on the Neo4j temp-database cluster.""" + # A cluster returns one row per hosting server, so dedupe on name + with get_session() as session: + result = session.run("SHOW DATABASES YIELD name RETURN DISTINCT name") + return [record["name"] for record in result] + + def clear_cache(database: str) -> None: """Best-effort cache clear for a Neo4j database.""" from api.attack_paths.database import GraphDatabaseQueryException diff --git a/api/src/backend/api/db_utils.py b/api/src/backend/api/db_utils.py index b6d3fdada1..b73fe127fb 100644 --- a/api/src/backend/api/db_utils.py +++ b/api/src/backend/api/db_utils.py @@ -409,7 +409,7 @@ def batch_delete(tenant_id, queryset, batch_size=settings.DJANGO_DELETION_BATCH_ Args: tenant_id (str): Tenant ID the queryset belongs to. - queryset (QuerySet): The queryset of objects to delete. + queryset: The queryset of objects to delete. batch_size (int): The number of objects to delete in each batch. Returns: diff --git a/api/src/backend/api/migrations/0100_attack_paths_tmp_db_reap_periodic_task.py b/api/src/backend/api/migrations/0100_attack_paths_tmp_db_reap_periodic_task.py new file mode 100644 index 0000000000..ff27aff249 --- /dev/null +++ b/api/src/backend/api/migrations/0100_attack_paths_tmp_db_reap_periodic_task.py @@ -0,0 +1,48 @@ +from django.db import migrations + +TASK_NAME = "attack-paths-reap-orphaned-tmp-databases" +INTERVAL_HOURS = 6 + + +def create_periodic_task(apps, schema_editor): + IntervalSchedule = apps.get_model("django_celery_beat", "IntervalSchedule") + PeriodicTask = apps.get_model("django_celery_beat", "PeriodicTask") + + schedule, _ = IntervalSchedule.objects.get_or_create( + every=INTERVAL_HOURS, + period="hours", + ) + + PeriodicTask.objects.update_or_create( + name=TASK_NAME, + defaults={ + "task": TASK_NAME, + "interval": schedule, + "enabled": True, + }, + ) + + +def delete_periodic_task(apps, schema_editor): + IntervalSchedule = apps.get_model("django_celery_beat", "IntervalSchedule") + PeriodicTask = apps.get_model("django_celery_beat", "PeriodicTask") + + PeriodicTask.objects.filter(name=TASK_NAME).delete() + + # Clean up the schedule if no other task references it + IntervalSchedule.objects.filter( + every=INTERVAL_HOURS, + period="hours", + periodictask__isnull=True, + ).delete() + + +class Migration(migrations.Migration): + dependencies = [ + ("api", "0099_delete_tenant_onboarding_profile"), + ("django_celery_beat", "0019_alter_periodictasks_options"), + ] + + operations = [ + migrations.RunPython(create_periodic_task, delete_periodic_task), + ] diff --git a/api/src/backend/api/models.py b/api/src/backend/api/models.py index 6a18cd1994..18454654ec 100644 --- a/api/src/backend/api/models.py +++ b/api/src/backend/api/models.py @@ -617,9 +617,66 @@ class Task(RowLevelSecurityProtectedModel): resource_name = "tasks" +class ScanQuerySet(models.QuerySet): + """Shared selectors for "the latest scan of a provider". + + The queryset must already be scoped by the caller: manager, tenant, RBAC, + providers and database alias. + """ + + # How "which completed scan is the provider's current one" is ordered. + LATEST_ORDER_BY = ( + models.F("completed_at").desc(nulls_last=True), + models.F("inserted_at").desc(), + models.F("id").desc(), + ) + + def _eligible_for_latest(self) -> "ScanQuerySet": + """Restrict to the scans that may be a provider's latest. + + Returns: + ScanQuerySet: The completed scans. + """ + return self.filter(state=StateChoices.COMPLETED) + + def latest_per_provider(self) -> "ScanQuerySet": + """Pick each provider's latest scan with `DISTINCT ON (provider_id)`. + + Returns: + ScanQuerySet: One scan per provider, the latest one. + """ + return ( + self._eligible_for_latest() + .order_by("provider_id", *self.LATEST_ORDER_BY) + .distinct("provider_id") + ) + + def latest_ids_per_provider(self) -> list[UUID]: + """Evaluate `latest_per_provider` and return the scan ids. + + The ids are materialised so callers can pass them as a literal `IN` + list; as a subquery Postgres misestimates the row count and picks a + slow nested loop. + + Returns: + list[UUID]: The id of each provider's latest scan. + """ + return list(self.latest_per_provider().values_list("id", flat=True)) + + def latest_first(self) -> "ScanQuerySet": + """Order eligible scans newest first, without deduplicating per provider. + + Expects the queryset to be already filtered to a single provider. + + Returns: + ScanQuerySet: The eligible scans, latest first. + """ + return self._eligible_for_latest().order_by(*self.LATEST_ORDER_BY) + + class Scan(RowLevelSecurityProtectedModel): - objects = ActiveProviderManager() - all_objects = models.Manager() + objects = ActiveProviderManager.from_queryset(ScanQuerySet)() + all_objects = ScanQuerySet.as_manager() _SCOPING_SCANNER_ARG_KEYS_CACHE: tuple[str, ...] | None = None @@ -726,6 +783,12 @@ class Scan(RowLevelSecurityProtectedModel): name="scans_prov_state_ins_desc_idx", ), # TODO This might replace `scans_prov_state_ins_desc_idx` completely. Review usage + # Since `ScanQuerySet`, no code path reads a provider's + # completed scans by `-inserted_at`. The only query left that + # matches this index (and `scans_prov_state_ins_desc_idx` above) + # is `GET /scans?filter[provider]=…&filter[state]=completed` with + # the default sort. Both are candidates to drop in a follow-up + # once production `pg_stat_user_indexes.idx_scan` confirms it. models.Index( fields=["tenant_id", "provider_id", "-inserted_at"], condition=Q(state=StateChoices.COMPLETED), diff --git a/api/src/backend/api/specs/v1.yaml b/api/src/backend/api/specs/v1.yaml index ab57f94702..78b46029bf 100644 --- a/api/src/backend/api/specs/v1.yaml +++ b/api/src/backend/api/specs/v1.yaml @@ -1,7 +1,7 @@ openapi: 3.0.3 info: title: Prowler API - version: 1.44.0 + version: 1.45.0 description: |- Prowler API specification. diff --git a/api/src/backend/api/tests/test_attack_paths_database.py b/api/src/backend/api/tests/test_attack_paths_database.py index c4aca45928..a563667852 100644 --- a/api/src/backend/api/tests/test_attack_paths_database.py +++ b/api/src/backend/api/tests/test_attack_paths_database.py @@ -187,6 +187,27 @@ class TestRoutingByDatabasePrefix: sink_backend_stub.drop_database.assert_called_once_with("db-tenant-abc") mock_ingest.drop_database.assert_not_called() + def test_list_databases_always_routes_to_ingest(self, sink_backend_stub): + with patch("api.attack_paths.database.ingest") as mock_ingest: + mock_ingest.list_databases.return_value = ["db-tmp-scan-uuid-1"] + + assert db_module.list_databases() == ["db-tmp-scan-uuid-1"] + + mock_ingest.list_databases.assert_called_once_with() + + def test_ingest_list_databases_dedupes_cluster_rows(self): + from api.attack_paths.ingest import driver as ingest_driver + + with patch.object(ingest_driver, "get_session") as mock_get_session: + session = mock_get_session.return_value.__enter__.return_value + session.run.return_value = [{"name": "db-tmp-scan-uuid-1"}] + + assert ingest_driver.list_databases() == ["db-tmp-scan-uuid-1"] + + session.run.assert_called_once_with( + "SHOW DATABASES YIELD name RETURN DISTINCT name" + ) + def test_clear_cache_routes_temp_to_ingest(self, sink_backend_stub): with patch("api.attack_paths.database.ingest") as mock_ingest: db_module.clear_cache("db-tmp-scan-uuid-1") diff --git a/api/src/backend/api/tests/test_models.py b/api/src/backend/api/tests/test_models.py index 5095da3a0e..1619fc6980 100644 --- a/api/src/backend/api/tests/test_models.py +++ b/api/src/backend/api/tests/test_models.py @@ -1,14 +1,16 @@ -from datetime import UTC, datetime +from datetime import UTC, datetime, timedelta import pytest from allauth.socialaccount.models import SocialApp from api.db_router import MainRouter from api.models import ( + Provider, ProviderComplianceScore, Resource, ResourceTag, SAMLConfiguration, SAMLDomainIndex, + Scan, StateChoices, StatusChoices, TenantComplianceSummary, @@ -524,3 +526,226 @@ class TestTenantComplianceSummaryModel: assert summary1.id != summary2.id assert summary1.requirements_passed != summary2.requirements_passed + + +def _latest_scan_fixture(tenant, provider, *, completed_at, inserted_at=None, **kwargs): + scan = Scan.objects.create( + tenant_id=tenant.id, + provider=provider, + trigger=kwargs.pop("trigger", Scan.TriggerChoices.MANUAL), + state=kwargs.pop("state", StateChoices.COMPLETED), + completed_at=completed_at, + **kwargs, + ) + if inserted_at is not None: + # `inserted_at` is auto_now_add, so it has to be forced after the fact. + Scan.all_objects.filter(pk=scan.pk).update(inserted_at=inserted_at) + scan.refresh_from_db() + return scan + + +@pytest.mark.django_db +class TestScanQuerySetOrdering: + def test_completed_later_wins_over_inserted_later( + self, tenants_fixture, aws_provider + ): + """The scan that FINISHED last is current, not the one that started last.""" + tenant, *_ = tenants_fixture + now = datetime.now(UTC) + + finished_last = _latest_scan_fixture( + tenant, + aws_provider, + inserted_at=now - timedelta(hours=3), + completed_at=now, + ) + _latest_scan_fixture( + tenant, + aws_provider, + inserted_at=now - timedelta(hours=1), + completed_at=now - timedelta(hours=1), + ) + + assert Scan.all_objects.filter( + tenant_id=tenant.id + ).latest_ids_per_provider() == [finished_last.id] + + def test_null_completed_at_provider_is_still_returned( + self, tenants_fixture, aws_provider + ): + """NULLS LAST, not `completed_at__isnull=False`. + + Excluding NULL `completed_at` would drop the provider from every + "latest" endpoint instead of falling back to `inserted_at`. + """ + tenant, *_ = tenants_fixture + only_scan = _latest_scan_fixture(tenant, aws_provider, completed_at=None) + + assert Scan.all_objects.filter( + tenant_id=tenant.id + ).latest_ids_per_provider() == [only_scan.id] + + def test_null_completed_at_never_outranks_a_finished_scan( + self, tenants_fixture, aws_provider + ): + """Postgres sorts NULLs first under DESC; NULLS LAST is what fixes it.""" + tenant, *_ = tenants_fixture + now = datetime.now(UTC) + + finished = _latest_scan_fixture( + tenant, + aws_provider, + inserted_at=now - timedelta(hours=2), + completed_at=now - timedelta(hours=2), + ) + _latest_scan_fixture( + tenant, + aws_provider, + inserted_at=now, + completed_at=None, + ) + + assert Scan.all_objects.filter( + tenant_id=tenant.id + ).latest_ids_per_provider() == [finished.id] + + def test_id_breaks_an_exact_timestamp_tie_deterministically( + self, tenants_fixture, aws_provider + ): + tenant, *_ = tenants_fixture + now = datetime.now(UTC) + + scans = [ + _latest_scan_fixture( + tenant, aws_provider, inserted_at=now, completed_at=now + ) + for _ in range(3) + ] + expected = max(scan.id for scan in scans) + + picks = { + Scan.all_objects.filter(tenant_id=tenant.id).latest_ids_per_provider()[0] + for _ in range(5) + } + assert picks == {expected} + + +@pytest.mark.django_db +class TestScanQuerySetEligibility: + def test_unfinished_scans_are_excluded(self, tenants_fixture, aws_provider): + tenant, *_ = tenants_fixture + _latest_scan_fixture( + tenant, + aws_provider, + completed_at=None, + state=StateChoices.EXECUTING, + ) + assert ( + Scan.all_objects.filter(tenant_id=tenant.id).latest_ids_per_provider() == [] + ) + + +@pytest.mark.django_db +class TestScanQuerySetManagerChoice: + def test_active_manager_hides_soft_deleted_providers( + self, tenants_fixture, aws_provider + ): + """`Scan.objects` drops soft-deleted providers, `all_objects` keeps them. + + The queryset must not decide this for the caller. + """ + tenant, *_ = tenants_fixture + scan = _latest_scan_fixture( + tenant, aws_provider, completed_at=datetime.now(UTC) + ) + + Provider.all_objects.filter(pk=aws_provider.pk).update(is_deleted=True) + + assert Scan.all_objects.filter( + tenant_id=tenant.id + ).latest_ids_per_provider() == [scan.id] + assert Scan.objects.filter(tenant_id=tenant.id).latest_ids_per_provider() == [] + + +@pytest.mark.django_db +class TestScanQuerySetPerProviderScoping: + def test_one_scan_per_provider(self, tenants_fixture, aws_provider_pair): + tenant, *_ = tenants_fixture + provider_one, provider_two = aws_provider_pair + now = datetime.now(UTC) + + newest_one = _latest_scan_fixture(tenant, provider_one, completed_at=now) + _latest_scan_fixture(tenant, provider_one, completed_at=now - timedelta(days=1)) + newest_two = _latest_scan_fixture(tenant, provider_two, completed_at=now) + + assert set( + Scan.all_objects.filter(tenant_id=tenant.id).latest_ids_per_provider() + ) == {newest_one.id, newest_two.id} + + def test_caller_filters_are_preserved(self, tenants_fixture, aws_provider_pair): + tenant, *_ = tenants_fixture + provider_one, provider_two = aws_provider_pair + now = datetime.now(UTC) + + scan_one = _latest_scan_fixture(tenant, provider_one, completed_at=now) + _latest_scan_fixture(tenant, provider_two, completed_at=now) + + assert Scan.all_objects.filter( + tenant_id=tenant.id, provider__in=[provider_one] + ).latest_ids_per_provider() == [scan_one.id] + + def test_latest_first_is_ordered_not_deduplicated( + self, tenants_fixture, aws_provider + ): + tenant, *_ = tenants_fixture + now = datetime.now(UTC) + + newest = _latest_scan_fixture(tenant, aws_provider, completed_at=now) + older = _latest_scan_fixture( + tenant, aws_provider, completed_at=now - timedelta(days=1) + ) + + ordered = list( + Scan.all_objects.filter( + tenant_id=tenant.id, provider_id=aws_provider.id + ).latest_first() + ) + assert [scan.id for scan in ordered] == [newest.id, older.id] + + def test_tenant_isolation(self, tenants_fixture, aws_provider): + tenant, other_tenant, *_ = tenants_fixture + _latest_scan_fixture(tenant, aws_provider, completed_at=datetime.now(UTC)) + + assert ( + Scan.all_objects.filter(tenant_id=other_tenant.id).latest_ids_per_provider() + == [] + ) + + def test_empty_queryset_returns_empty_list(self, tenants_fixture): + tenant, *_ = tenants_fixture + assert ( + Scan.all_objects.filter(tenant_id=tenant.id).latest_ids_per_provider() == [] + ) + + +@pytest.mark.django_db +class TestScanQuerySetPerProviderQuerysetShape: + def test_returns_a_queryset_not_a_list(self, tenants_fixture, aws_provider): + tenant, *_ = tenants_fixture + _latest_scan_fixture(tenant, aws_provider, completed_at=datetime.now(UTC)) + + qs = Scan.all_objects.filter(tenant_id=tenant.id).latest_per_provider() + # Callers chain .values(...) / .values_list(...) onto this. + assert qs.values_list("provider_id", flat=True).count() == 1 + + +@pytest.mark.django_db +class TestScanQuerySetRelatedManager: + def test_reverse_relation_exposes_the_methods(self, tenants_fixture, aws_provider): + tenant, *_ = tenants_fixture + now = datetime.now(UTC) + + newest = _latest_scan_fixture(tenant, aws_provider, completed_at=now) + _latest_scan_fixture(tenant, aws_provider, completed_at=now - timedelta(days=1)) + + assert aws_provider.scans.latest_first().first().id == newest.id diff --git a/api/src/backend/api/tests/test_views.py b/api/src/backend/api/tests/test_views.py index 43a13e4d4c..64a2250082 100644 --- a/api/src/backend/api/tests/test_views.py +++ b/api/src/backend/api/tests/test_views.py @@ -3951,6 +3951,43 @@ class TestScanViewSet: mock_enqueue_scan_execution.assert_called_once() # assert scan.scanner_args == expected_scanner_args + @patch("api.v1.views.enqueue_scan_execution_on_commit") + def test_scans_create_returns_the_scan_id_in_task_args( + self, + mock_enqueue_scan_execution, + authenticated_client, + okta_provider, + ): + """The 202 is a task, so `task_args` is the only place the scan id is. + + It is serialized before the on_commit publish that would otherwise fill + the kwargs, so the record has to carry them from the start. + """ + payload = { + "data": { + "type": "scans", + "attributes": {"name": "New Scan"}, + "relationships": { + "provider": { + "data": {"type": "providers", "id": str(okta_provider.id)} + } + }, + } + } + + response = authenticated_client.post( + reverse("scan-list"), + data=payload, + content_type=API_JSON_CONTENT_TYPE, + ) + + assert response.status_code == status.HTTP_202_ACCEPTED + scan = Scan.objects.get() + assert response.json()["data"]["attributes"]["task_args"] == { + "scan_id": str(scan.id), + "provider_id": str(okta_provider.id), + } + @patch("tasks.tasks.perform_scan_task.apply_async") def test_scans_create_queues_scan_when_provider_has_active_scan( self, diff --git a/api/src/backend/api/v1/views.py b/api/src/backend/api/v1/views.py index 1bb40f4fd3..9642b36d9d 100644 --- a/api/src/backend/api/v1/views.py +++ b/api/src/backend/api/v1/views.py @@ -2822,6 +2822,15 @@ class ScanViewSet(ProviderVisibilityMixin, BaseRLSViewSet): tenant_id=self.request.tenant_id, task_id=pre_task_id, task_status=(QUEUED_SCAN_TASK_STATE if active_scan else None), + # This response is serialized before the on_commit publish, + # so without these the caller gets a task id and no scan id. + # Kept in step with what `enqueue_scan_execution_on_commit` + # publishes below. + task_kwargs={ + "tenant_id": str(self.request.tenant_id), + "scan_id": str(scan.id), + "provider_id": str(scan.provider_id), + }, ) if not active_scan: @@ -3478,12 +3487,8 @@ class ResourceViewSet(PaginateByPkMixin, BaseRLSViewSet): filtered_queryset = self.filter_queryset(self.get_queryset()) latest_scans = ( - Scan.all_objects.filter( - tenant_id=tenant_id, - state=StateChoices.COMPLETED, - ) - .order_by("provider_id", "-inserted_at") - .distinct("provider_id") + Scan.all_objects.filter(tenant_id=tenant_id) + .latest_per_provider() .values("provider_id") ) @@ -3615,11 +3620,9 @@ class ResourceViewSet(PaginateByPkMixin, BaseRLSViewSet): tenant_id = request.tenant_id query_params = request.query_params - latest_scans_queryset = ( - Scan.all_objects.filter(tenant_id=tenant_id, state=StateChoices.COMPLETED) - .order_by("provider_id", "-inserted_at") - .distinct("provider_id") - ) + latest_scans_queryset = Scan.all_objects.filter( + tenant_id=tenant_id + ).latest_per_provider() queryset = ResourceScanSummary.objects.filter( tenant_id=tenant_id, @@ -4206,12 +4209,9 @@ class FindingViewSet(PaginateByPkMixin, BaseRLSViewSet): tenant_id = request.tenant_id filtered_queryset = self.filter_queryset(self.get_queryset()) - latest_scan_ids = list( - Scan.all_objects.filter(tenant_id=tenant_id, state=StateChoices.COMPLETED) - .order_by("provider_id", "-inserted_at") - .distinct("provider_id") - .values_list("id", flat=True) - ) + latest_scan_ids = Scan.all_objects.filter( + tenant_id=tenant_id + ).latest_ids_per_provider() filtered_queryset = filtered_queryset.filter( tenant_id=tenant_id, scan_id__in=latest_scan_ids ) @@ -4234,11 +4234,9 @@ class FindingViewSet(PaginateByPkMixin, BaseRLSViewSet): tenant_id = request.tenant_id query_params = request.query_params - latest_scans_queryset = ( - Scan.all_objects.filter(tenant_id=tenant_id, state=StateChoices.COMPLETED) - .order_by("provider_id", "-inserted_at") - .distinct("provider_id") - ) + latest_scans_queryset = Scan.all_objects.filter( + tenant_id=tenant_id + ).latest_per_provider() raw_latest_scans_ids = list( latest_scans_queryset.values_list("id", "unique_resource_count") ) @@ -4979,11 +4977,7 @@ class ComplianceOverviewViewSet( if provider_filters: scans = scans.filter(**provider_filters) - return list( - scans.order_by("provider_id", "-inserted_at") - .distinct("provider_id") - .values_list("id", flat=True) - ) + return scans.latest_ids_per_provider() def _filtered_queryset_for_latest_provider_scans(self, latest_scan_ids=None): if latest_scan_ids is None: @@ -5725,14 +5719,9 @@ class OverviewViewSet(ProviderFilterParamsMixin, BaseRLSViewSet): else {} ) - latest_scan_ids = ( - Scan.all_objects.filter( - tenant_id=tenant_id, state=StateChoices.COMPLETED, **provider_filter - ) - .order_by("provider_id", "-inserted_at") - .distinct("provider_id") - .values_list("id", flat=True) - ) + latest_scan_ids = Scan.all_objects.filter( + tenant_id=tenant_id, **provider_filter + ).latest_ids_per_provider() return filtered_queryset.filter( tenant_id=tenant_id, scan_id__in=latest_scan_ids @@ -5759,16 +5748,10 @@ class OverviewViewSet(ProviderFilterParamsMixin, BaseRLSViewSet): def _latest_scan_ids_for_allowed_providers(self, tenant_id, provider_filters=None): provider_filter = self._get_provider_filter() - queryset = Scan.all_objects.filter( - tenant_id=tenant_id, state=StateChoices.COMPLETED, **provider_filter - ) + queryset = Scan.all_objects.filter(tenant_id=tenant_id, **provider_filter) if provider_filters: queryset = queryset.filter(**provider_filters) - return ( - queryset.order_by("provider_id", "-inserted_at") - .distinct("provider_id") - .values_list("id", flat=True) - ) + return queryset.latest_ids_per_provider() @action(detail=False, methods=["get"], url_name="providers") def providers(self, request): @@ -5780,14 +5763,9 @@ class OverviewViewSet(ProviderFilterParamsMixin, BaseRLSViewSet): else {} ) - latest_scan_ids = ( - Scan.all_objects.filter( - tenant_id=tenant_id, state=StateChoices.COMPLETED, **provider_filter - ) - .order_by("provider_id", "-inserted_at") - .distinct("provider_id") - .values_list("id", flat=True) - ) + latest_scan_ids = Scan.all_objects.filter( + tenant_id=tenant_id, **provider_filter + ).latest_ids_per_provider() findings_aggregated = ( queryset.filter(scan_id__in=latest_scan_ids) @@ -7934,18 +7912,13 @@ class FindingGroupViewSet(JsonApiFilterMixin, BaseRLSViewSet): def _get_latest_findings_per_provider(self, filtered_queryset): """Keep only findings from each provider's most recent completed scan.""" - # Materialize to a literal IN list. Left as a subquery, Postgres can't - # estimate the match count and picks a serial nested loop on - # resource_finding_mappings when one scan dominates findings - latest_scan_ids = list( - Scan.objects.filter( - tenant_id=self.request.tenant_id, - state=StateChoices.COMPLETED, - ) - .order_by("provider_id", "-completed_at", "-inserted_at") - .distinct("provider_id") - .values_list("id", flat=True) - ) + # `latest_ids_per_provider` materializes to a literal IN list. + # Left as a subquery, Postgres can't estimate the match count and picks + # a serial nested loop on resource_finding_mappings when one scan + # dominates findings + latest_scan_ids = Scan.objects.filter( + tenant_id=self.request.tenant_id + ).latest_ids_per_provider() return filtered_queryset.filter(scan_id__in=latest_scan_ids) def _post_process_aggregation(self, aggregated_data): @@ -8891,16 +8864,14 @@ class FindingGroupViewSet(JsonApiFilterMixin, BaseRLSViewSet): tenant_id = request.tenant_id queryset = self._get_finding_queryset() - # Order by -completed_at (matching the /latest summary path and the - # daily summary upsert keyed on midnight(completed_at)) so that - # overlapping scans do not make /resources and /latest read from - # different scans and report diverging counts. - latest_scan_ids = ( - Scan.objects.filter(tenant_id=tenant_id, state=StateChoices.COMPLETED) - .order_by("provider_id", "-completed_at", "-inserted_at") - .distinct("provider_id") - .values_list("id", flat=True) - ) + # The shared selector orders by -completed_at (matching the /latest + # summary path and the daily summary upsert keyed on + # midnight(completed_at)) so that overlapping scans do not make + # /resources and /latest read from different scans and report + # diverging counts. + latest_scan_ids = Scan.objects.filter( + tenant_id=tenant_id + ).latest_ids_per_provider() normalized_params = self._normalize_jsonapi_params(request.query_params) # Remove date filters since we're using latest diff --git a/api/src/backend/config/custom_logging.py b/api/src/backend/config/custom_logging.py index a601c8e2cd..f60651f8d4 100644 --- a/api/src/backend/config/custom_logging.py +++ b/api/src/backend/config/custom_logging.py @@ -231,6 +231,62 @@ LOGGING = { "level": LEVEL, "propagate": False, }, + # Celery loggers must be declared explicitly because + # disable_existing_loggers=True silences any logger that exists at + # dictConfig time but is not named here. Without these, fatal worker + # errors (e.g. celery.worker CRITICAL) produce no output. + # "celery" must keep propagating: get_task_logger() parents task + # loggers under celery.task, so blocking here hides them from root. + "celery": { + "level": LEVEL, + "propagate": True, + }, + "celery.worker": { + "handlers": ["tasks_console"], + "level": LEVEL, + "propagate": False, + }, + "celery.worker.consumer": { + "handlers": ["tasks_console"], + "level": LEVEL, + "propagate": False, + }, + "celery.worker.consumer.consumer": { + "handlers": ["tasks_console"], + "level": LEVEL, + "propagate": False, + }, + "kombu": { + "handlers": ["tasks_console"], + "level": LEVEL, + "propagate": False, + }, + "kombu.transport.redis": { + "handlers": ["tasks_console"], + "level": LEVEL, + "propagate": False, + }, + "billiard": { + "handlers": ["tasks_console"], + "level": LEVEL, + "propagate": False, + }, + "amqp": { + "handlers": ["tasks_console"], + "level": LEVEL, + "propagate": False, + }, + # WARNING keeps task failures but skips one "succeeded" line per task. + "celery.app.trace": { + "handlers": ["tasks_console"], + "level": "WARNING", + "propagate": False, + }, + "celery.beat": { + "handlers": ["tasks_console"], + "level": LEVEL, + "propagate": False, + }, }, # Gunicorn required configuration "root": { diff --git a/api/src/backend/config/django/base.py b/api/src/backend/config/django/base.py index a208dda915..664fcd1c4a 100644 --- a/api/src/backend/config/django/base.py +++ b/api/src/backend/config/django/base.py @@ -7,6 +7,7 @@ from config.settings.eventstream import * # noqa from config.settings.partitions import * # noqa from config.settings.sentry import * # noqa from config.settings.social_login import * # noqa +from django.core.exceptions import ImproperlyConfigured SECRET_KEY = env("SECRET_KEY", default="secret") DEBUG = env.bool("DJANGO_DEBUG", default=False) @@ -295,6 +296,9 @@ DJANGO_OUTPUT_S3_AWS_SECRET_ACCESS_KEY = env.str( ) DJANGO_OUTPUT_S3_AWS_SESSION_TOKEN = env.str("DJANGO_OUTPUT_S3_AWS_SESSION_TOKEN", "") DJANGO_OUTPUT_S3_AWS_DEFAULT_REGION = env.str("DJANGO_OUTPUT_S3_AWS_DEFAULT_REGION", "") +# Storage endpoint the API and Celery workers use to talk to S3-compatible object storage +# such as MinIO. Empty means the real AWS S3 endpoint, which is unaffected. +DJANGO_OUTPUT_S3_AWS_ENDPOINT_URL = env.str("DJANGO_OUTPUT_S3_AWS_ENDPOINT_URL", "") # Browser-reachable storage host used to sign download URLs. Empty means sign against the # same endpoint the API talks to, which is what Prowler Cloud on S3 does. DJANGO_OUTPUT_S3_AWS_PUBLIC_ENDPOINT_URL = env.str( @@ -325,6 +329,17 @@ ATTACK_PATHS_SCAN_STALE_THRESHOLD_MINUTES = env.int( "ATTACK_PATHS_SCAN_STALE_THRESHOLD_MINUTES", 960 ) # 16h +# Minimum age (of the scan row, or of the scan id itself when the row is gone) before +# the periodic reaper will drop an orphaned temp Neo4j database. Keeps a scan that is +# still legitimately in flight from ever losing its staging database mid-run. +ATTACK_PATHS_TMP_DB_REAP_SAFETY_MARGIN_HOURS = env.int( + "ATTACK_PATHS_TMP_DB_REAP_SAFETY_MARGIN_HOURS", 6 +) +if ATTACK_PATHS_TMP_DB_REAP_SAFETY_MARGIN_HOURS <= 0: + raise ImproperlyConfigured( + "ATTACK_PATHS_TMP_DB_REAP_SAFETY_MARGIN_HOURS must be a positive number of hours" + ) + # Selects where the persistent attack-paths graph is stored. The scan # temporary database is always Neo4j; only the sink is configurable. # Valid values: "neo4j" (default, OSS and local dev), "neptune" (hosted). diff --git a/api/src/backend/tasks/jobs/attack_paths/tmp_db_reaper.py b/api/src/backend/tasks/jobs/attack_paths/tmp_db_reaper.py new file mode 100644 index 0000000000..3d575d41d4 --- /dev/null +++ b/api/src/backend/tasks/jobs/attack_paths/tmp_db_reaper.py @@ -0,0 +1,107 @@ +"""Periodic reaper for orphaned temp Neo4j scan databases. + +`scan.py` creates a throw-away `db-tmp-scan-` database per +scan and drops it once the scan finishes, success or failure. When the worker +or Neo4j itself dies mid-scan, that drop never runs and nothing else ever +revisits the database - it sits there forever. This sweep lists every temp +database on the ingest cluster and drops the ones whose scan is gone or has +been finished for longer than the configured safety margin. +""" + +from datetime import UTC, datetime, timedelta + +from api.attack_paths import database as graph_database +from api.db_router import MainRouter +from api.models import AttackPathsScan, StateChoices +from api.uuid_utils import datetime_from_uuid7 +from celery.utils.log import get_task_logger +from config.django.base import ATTACK_PATHS_TMP_DB_REAP_SAFETY_MARGIN_HOURS +from uuid6 import UUID as UUID7 + +logger = get_task_logger(__name__) + +TERMINAL_STATES = ( + StateChoices.COMPLETED, + StateChoices.FAILED, + StateChoices.CANCELLED, +) + + +def reap_orphaned_tmp_databases() -> dict: + """Drop temp Neo4j scan databases whose scan is gone or long finished. + + A failure listing databases aborts the whole sweep (nothing to iterate). + A failure reaping one database is logged and skipped so the rest of the + sweep still runs. + """ + now = datetime.now(tz=UTC) + safety_margin = timedelta(hours=ATTACK_PATHS_TMP_DB_REAP_SAFETY_MARGIN_HOURS) + + try: + databases = graph_database.list_databases() + except Exception: + logger.exception("Failed to list ingest Neo4j databases for temp-db reap") + return {"dropped_count": 0, "databases": []} + + tmp_databases = [ + name for name in databases if name.startswith(graph_database.TEMP_DB_PREFIX) + ] + + dropped: list[str] = [] + for database in tmp_databases: + try: + if _is_orphaned(database, now, safety_margin): + graph_database.drop_database(database) + dropped.append(database) + logger.info(f"Dropped orphaned temp Neo4j database `{database}`") + except Exception: + logger.exception(f"Failed to reap temp Neo4j database `{database}`") + + logger.info(f"Temp Neo4j database reap: {len(dropped)} dropped") + return {"dropped_count": len(dropped), "databases": dropped} + + +def _is_orphaned(database: str, now: datetime, safety_margin: timedelta) -> bool: + """Decide whether a temp database is safe to drop. + + No scan row: the row was hard-deleted (tenant/provider cleanup) or was + never created. Falls back to the scan id's own UUIDv7 timestamp so a + database created moments ago is never touched even without a row to check. + + Scan row present: only reapable once it reached a terminal state and has + been finished for longer than the safety margin, so a scan still + legitimately executing is never touched. + """ + scan_id = database[len(graph_database.TEMP_DB_PREFIX) :] + + try: + scan_uuid = UUID7(scan_id) + except ValueError: + logger.warning( + f"Temp database `{database}` has an unparseable scan id, skipping" + ) + return False + + # Global sweep with no tenant context: admin_db bypasses RLS on purpose, the same + # way cleanup_stale_attack_paths_scans finds stale scans across every tenant. + scan = ( + AttackPathsScan.all_objects.using(MainRouter.admin_db) + .filter(id=scan_uuid) + .first() + ) + + if scan is None: + if scan_uuid.version != 7: + logger.warning( + f"Temp database `{database}` has no scan row and a non-UUIDv7 id, " + "skipping" + ) + return False + return now - datetime_from_uuid7(scan_uuid) >= safety_margin + + if scan.state not in TERMINAL_STATES: + return False + + # `mark_scan_finished` does not touch `updated_at`, so prefer `completed_at` + finished_at = scan.completed_at or scan.updated_at + return now - finished_at >= safety_margin diff --git a/api/src/backend/tasks/jobs/backfill.py b/api/src/backend/tasks/jobs/backfill.py index 56cb626786..a72da95779 100644 --- a/api/src/backend/tasks/jobs/backfill.py +++ b/api/src/backend/tasks/jobs/backfill.py @@ -499,10 +499,13 @@ def backfill_provider_compliance_scores(tenant_id: str) -> dict: provider_id__in=existing_providers ) + # `completed_scans` keeps its own `completed_at__isnull=False`: this + # task writes a *dated* ProviderComplianceScore row, so unlike the read + # paths it genuinely cannot use a scan without a `completed_at`. scan_info = list( - completed_scans.order_by("provider_id", "-completed_at") - .distinct("provider_id") - .values("id", "provider_id", "completed_at") + completed_scans.latest_per_provider().values( + "id", "provider_id", "completed_at" + ) ) if not scan_info: diff --git a/api/src/backend/tasks/jobs/export.py b/api/src/backend/tasks/jobs/export.py index aacb22578a..49b91c219f 100644 --- a/api/src/backend/tasks/jobs/export.py +++ b/api/src/backend/tasks/jobs/export.py @@ -207,15 +207,21 @@ def get_s3_client(): This function attempts to initialize an S3 client by reading the AWS access key, secret key, session token, and region from environment variables. It then validates the client by listing available S3 buckets. If an error occurs during this process (for example, due to missing or - invalid credentials), it falls back to creating an S3 client without explicitly provided credentials, - which may rely on other configuration sources (e.g., IAM roles). + invalid credentials), it falls back to creating an S3 client without explicitly provided + credentials, which may rely on other configuration sources (e.g., IAM roles). + + That fallback is only safe when no explicit endpoint is configured: with an endpoint set, the + explicit client already targets the intended S3-compatible storage, and the fallback client + would go to the AWS default provider chain instead, an unrelated real-AWS account reachable + from the host. So when an endpoint is configured, the original error propagates instead. Returns: boto3.client: A configured S3 client instance. Raises: - ClientError, NoCredentialsError, or ParamValidationError if both attempts to create a client fail. + ClientError, NoCredentialsError, or ParamValidationError if the client cannot be created. """ + endpoint = settings.DJANGO_OUTPUT_S3_AWS_ENDPOINT_URL s3_client = None try: s3_client = boto3.client( @@ -226,9 +232,12 @@ def get_s3_client(): # Storage that has no meaningful region, MinIO among it, is usually configured # without one, and botocore rejects an empty region before any request is made. region_name=settings.DJANGO_OUTPUT_S3_AWS_DEFAULT_REGION or "us-east-1", + endpoint_url=endpoint or None, ) s3_client.list_buckets() except (ClientError, NoCredentialsError, ParamValidationError, ValueError): + if endpoint: + raise s3_client = boto3.client("s3") s3_client.list_buckets() @@ -236,13 +245,21 @@ def get_s3_client(): def get_s3_presign_client(): - """Return a client that signs URLs against the public storage host. + """Return a client that signs download URLs with SigV4. - None means no public host is configured and the caller should presign with its own - client, which leaves deployments on real S3 with the URL they get today. + It is used when a public or internal storage host is configured, or when the bucket's + region is: boto3 otherwise presigns S3 URLs with SigV2, which S3 rejects for SSE-KMS + objects. None means none of those is set and the caller should presign with its own + client, which leaves those deployments with the URL they get today. + + The public endpoint wins when both are set: the internal endpoint may only be reachable + from inside the cluster, and a URL signed against it would not open in a browser. """ - public_endpoint = settings.DJANGO_OUTPUT_S3_AWS_PUBLIC_ENDPOINT_URL - if not public_endpoint: + endpoint = ( + settings.DJANGO_OUTPUT_S3_AWS_PUBLIC_ENDPOINT_URL + or settings.DJANGO_OUTPUT_S3_AWS_ENDPOINT_URL + ) + if not endpoint and not settings.DJANGO_OUTPUT_S3_AWS_DEFAULT_REGION: return None # Blank keys are signed as-is (empty credential scope) instead of deferring to the @@ -266,9 +283,10 @@ def get_s3_presign_client(): # SigV4 puts the region in the credential scope, and MinIO answers to us-east-1 # unless it was told otherwise, so an empty region would sign an unusable URL. region_name=settings.DJANGO_OUTPUT_S3_AWS_DEFAULT_REGION or "us-east-1", - endpoint_url=public_endpoint, + endpoint_url=endpoint or None, # The signature covers the host, so the addressing style has to be pinned rather - # than guessed from the endpoint: MinIO serves path-style. + # than guessed from the endpoint: MinIO serves path-style, and on AWS it keeps the + # regional host instead of the global one, which redirects for new buckets. config=Config(signature_version="s3v4", s3={"addressing_style": "path"}), ) diff --git a/api/src/backend/tasks/jobs/muting.py b/api/src/backend/tasks/jobs/muting.py index 839889a8ef..bc1d7f1c43 100644 --- a/api/src/backend/tasks/jobs/muting.py +++ b/api/src/backend/tasks/jobs/muting.py @@ -1,7 +1,7 @@ from collections.abc import Iterable from api.db_utils import rls_transaction -from api.models import Finding, MuteRule, Scan, StateChoices +from api.models import Finding, MuteRule, Scan from celery.utils.log import get_task_logger logger = get_task_logger(__name__) @@ -39,17 +39,9 @@ def mute_findings_in_latest_scans( with rls_transaction(tenant_id): mute_rule = MuteRule.objects.get(id=mute_rule_id, tenant_id=tenant_id) - latest_scans = list( - Scan.objects.filter( - tenant_id=tenant_id, - provider_id__in=provider_ids, - state=StateChoices.COMPLETED, - completed_at__isnull=False, - ) - .order_by("provider_id", "-completed_at", "-inserted_at", "-id") - .distinct("provider_id") - .values_list("id", flat=True) - ) + latest_scans = Scan.objects.filter( + tenant_id=tenant_id, provider_id__in=provider_ids + ).latest_ids_per_provider() changed_scan_ids = [] findings_muted = 0 diff --git a/api/src/backend/tasks/jobs/orphan_recovery.py b/api/src/backend/tasks/jobs/orphan_recovery.py index c8cda54cd2..7290fb313f 100644 --- a/api/src/backend/tasks/jobs/orphan_recovery.py +++ b/api/src/backend/tasks/jobs/orphan_recovery.py @@ -84,6 +84,7 @@ _SKIP_RECOVERY = { "scan-perform-scheduled", "attack-paths-scan-perform", "attack-paths-cleanup-stale-scans", + "attack-paths-reap-orphaned-tmp-databases", "reconcile-orphan-tasks", } diff --git a/api/src/backend/tasks/jobs/scan.py b/api/src/backend/tasks/jobs/scan.py index 8b57ebeecd..26b64993bd 100644 --- a/api/src/backend/tasks/jobs/scan.py +++ b/api/src/backend/tasks/jobs/scan.py @@ -2700,20 +2700,37 @@ def reset_ephemeral_resource_findings_count(tenant_id: str, scan_id: str) -> dic # refreshed). Wiping based on the older scan would zero counts the newer # scan just set. Skip and let the newer scan's reset task do the work; if # this task was delayed in the queue, that's the correct outcome. - # `completed_at__isnull=False` is required: Postgres orders NULL first in - # DESC, so a sibling COMPLETED scan with a missing completed_at would sort - # as "newest" and incorrectly cause us to skip. + # + # The comparison must be against the newest *full-scope* scan, which is + # what this variable has always been named after but did not use to be: + # the query filtered nothing about scope, so any newer scan that is not + # full-scope (an imported one, for instance) made the full-scope scan + # skip its own cleanup and leave ephemeral resources with a stale + # failed_findings_count permanently. + # + # `is_full_scope()` reads `trigger` plus the scoping keys inside + # `scanner_args`, which is not expressible as a WHERE clause, so the + # candidates are walked newest-first in Python until the first full-scope + # one. The walk needs no cap: `scan` is itself a full-scope candidate, so + # it stops at `scan` at the latest, after reading only the scans newer than + # it. A fixed window would return None once more newer scoped scans had + # landed than it inspected, and skip the cleanup exactly like the bug above. + # + # NULL `completed_at` no longer needs an explicit filter here: the shared + # ordering in `ScanQuerySet.LATEST_ORDER_BY` sorts NULLs last + # rather than excluding them, which also fixes the case where a provider + # whose completed scans all have a NULL `completed_at` resolved to None and + # therefore never ran the reset at all. with rls_transaction(tenant_id): - latest_full_scope_scan_id = ( - Scan.objects.filter( - tenant_id=tenant_id, - provider_id=scan.provider_id, - state=StateChoices.COMPLETED, - completed_at__isnull=False, - ) - .order_by("-completed_at", "-inserted_at") - .values_list("id", flat=True) - .first() + candidates = ( + Scan.objects.filter(tenant_id=tenant_id, provider_id=scan.provider_id) + .latest_first() + .only("id", "trigger", "scanner_args") + .iterator(chunk_size=100) + ) + latest_full_scope_scan_id = next( + (candidate.id for candidate in candidates if candidate.is_full_scope()), + None, ) if latest_full_scope_scan_id != scan.id: logger.info( diff --git a/api/src/backend/tasks/tasks.py b/api/src/backend/tasks/tasks.py index d160a2b6c7..9ec6526963 100644 --- a/api/src/backend/tasks/tasks.py +++ b/api/src/backend/tasks/tasks.py @@ -1,3 +1,4 @@ +import json import os from datetime import UTC, datetime, timedelta from pathlib import Path @@ -42,6 +43,7 @@ from tasks.jobs.attack_paths import ( ) from tasks.jobs.attack_paths import db_utils as attack_paths_db_utils from tasks.jobs.attack_paths.cleanup import cleanup_stale_attack_paths_scans +from tasks.jobs.attack_paths.tmp_db_reaper import reap_orphaned_tmp_databases from tasks.jobs.backfill import ( aggregate_scan_category_summaries, aggregate_scan_resource_group_summaries, @@ -163,13 +165,28 @@ def create_scan_task_record( task_id: str, task_name: str = "scan-perform", task_status: str | None = states.PENDING, + task_kwargs: dict | None = None, ) -> Task: + """Pre-create the TaskResult + Task rows for a pre-generated task id. + + Pass ``task_kwargs`` when the response built from this record is serialized + before the broker publish. ``task_kwargs`` is otherwise only written by the + ``before_task_publish`` signal (``api/signals.py``), and the scan publish is + deferred to ``on_commit``, so the 202 would carry an empty ``task_args`` and + the caller would have no way to learn the scan id it was just handed a task + for. The publish later overwrites the field with the same kwargs as a Python + repr; both forms decode to the same dict (``decode_celery_field``). + """ if task_status is None: task_status = states.PENDING + defaults = {"status": task_status, "task_name": task_name} + if task_kwargs is not None: + defaults["task_kwargs"] = json.dumps(task_kwargs) + task_result, _ = TaskResult.objects.update_or_create( task_id=str(task_id), - defaults={"status": task_status, "task_name": task_name}, + defaults=defaults, ) prowler_task, _ = Task.objects.update_or_create( id=str(task_id), @@ -711,6 +728,13 @@ def cleanup_stale_attack_paths_scans_task(): return cleanup_stale_attack_paths_scans() +@shared_task( + name="attack-paths-reap-orphaned-tmp-databases", queue="attack-paths-scans" +) +def reap_orphaned_attack_paths_tmp_databases_task(): + return reap_orphaned_tmp_databases() + + @shared_task(name="reconcile-orphan-tasks", queue="celery") def reconcile_orphan_tasks_task(): """Periodic watchdog: recover tasks whose worker is gone (deploys, crashes).""" diff --git a/api/src/backend/tasks/tests/test_attack_paths_tmp_db_reaper.py b/api/src/backend/tasks/tests/test_attack_paths_tmp_db_reaper.py new file mode 100644 index 0000000000..ab8d478383 --- /dev/null +++ b/api/src/backend/tasks/tests/test_attack_paths_tmp_db_reaper.py @@ -0,0 +1,286 @@ +from datetime import UTC, datetime, timedelta +from unittest.mock import patch +from uuid import uuid4 + +import pytest +from api.attack_paths.database import TEMP_DB_PREFIX +from api.models import AttackPathsScan, StateChoices +from api.uuid_utils import datetime_to_uuid7 +from config.django.base import ATTACK_PATHS_TMP_DB_REAP_SAFETY_MARGIN_HOURS + +MARGIN = timedelta(hours=ATTACK_PATHS_TMP_DB_REAP_SAFETY_MARGIN_HOURS) + + +def _tmp_db_name(scan_uuid) -> str: + return f"{TEMP_DB_PREFIX}{scan_uuid}" + + +@pytest.mark.django_db +class TestReapOrphanedTmpDatabases: + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.drop_database") + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.list_databases") + def test_ignores_databases_without_the_temp_prefix(self, mock_list, mock_drop): + from tasks.jobs.attack_paths.tmp_db_reaper import reap_orphaned_tmp_databases + + mock_list.return_value = ["db-tenant-abc123", "system", "neo4j"] + + result = reap_orphaned_tmp_databases() + + assert result == {"dropped_count": 0, "databases": []} + mock_drop.assert_not_called() + + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.drop_database") + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.list_databases") + def test_drops_temp_db_with_no_scan_row_past_safety_margin( + self, mock_list, mock_drop + ): + from tasks.jobs.attack_paths.tmp_db_reaper import reap_orphaned_tmp_databases + + old_scan_id = datetime_to_uuid7( + datetime.now(tz=UTC) - MARGIN - timedelta(hours=1) + ) + database = _tmp_db_name(old_scan_id) + mock_list.return_value = [database] + + result = reap_orphaned_tmp_databases() + + assert result == {"dropped_count": 1, "databases": [database]} + mock_drop.assert_called_once_with(database) + + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.drop_database") + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.list_databases") + def test_preserves_temp_db_with_no_scan_row_inside_safety_margin( + self, mock_list, mock_drop + ): + from tasks.jobs.attack_paths.tmp_db_reaper import reap_orphaned_tmp_databases + + recent_scan_id = datetime_to_uuid7(datetime.now(tz=UTC) - timedelta(minutes=5)) + database = _tmp_db_name(recent_scan_id) + mock_list.return_value = [database] + + result = reap_orphaned_tmp_databases() + + assert result == {"dropped_count": 0, "databases": []} + mock_drop.assert_not_called() + + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.drop_database") + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.list_databases") + def test_preserves_temp_db_with_unparseable_scan_id(self, mock_list, mock_drop): + from tasks.jobs.attack_paths.tmp_db_reaper import reap_orphaned_tmp_databases + + database = f"{TEMP_DB_PREFIX}not-a-uuid" + mock_list.return_value = [database] + + result = reap_orphaned_tmp_databases() + + assert result == {"dropped_count": 0, "databases": []} + mock_drop.assert_not_called() + + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.drop_database") + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.list_databases") + def test_drops_terminal_scan_past_safety_margin( + self, mock_list, mock_drop, tenants_fixture, aws_provider + ): + from tasks.jobs.attack_paths.tmp_db_reaper import reap_orphaned_tmp_databases + + tenant = tenants_fixture[0] + old_updated_at = datetime.now(tz=UTC) - MARGIN - timedelta(hours=1) + scan = AttackPathsScan.objects.create( + tenant_id=tenant.id, + provider=aws_provider, + state=StateChoices.COMPLETED, + ) + AttackPathsScan.objects.filter(id=scan.id).update(updated_at=old_updated_at) + + database = _tmp_db_name(scan.id) + mock_list.return_value = [database] + + result = reap_orphaned_tmp_databases() + + assert result == {"dropped_count": 1, "databases": [database]} + mock_drop.assert_called_once_with(database) + + @pytest.mark.parametrize( + "state", + [StateChoices.FAILED, StateChoices.CANCELLED], + ) + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.drop_database") + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.list_databases") + def test_drops_other_terminal_states_past_safety_margin( + self, mock_list, mock_drop, tenants_fixture, aws_provider, state + ): + from tasks.jobs.attack_paths.tmp_db_reaper import reap_orphaned_tmp_databases + + tenant = tenants_fixture[0] + old_updated_at = datetime.now(tz=UTC) - MARGIN - timedelta(hours=1) + scan = AttackPathsScan.objects.create( + tenant_id=tenant.id, + provider=aws_provider, + state=state, + ) + AttackPathsScan.objects.filter(id=scan.id).update(updated_at=old_updated_at) + + database = _tmp_db_name(scan.id) + mock_list.return_value = [database] + + result = reap_orphaned_tmp_databases() + + assert result["dropped_count"] == 1 + mock_drop.assert_called_once_with(database) + + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.drop_database") + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.list_databases") + def test_preserves_terminal_scan_inside_safety_margin( + self, mock_list, mock_drop, tenants_fixture, aws_provider + ): + from tasks.jobs.attack_paths.tmp_db_reaper import reap_orphaned_tmp_databases + + tenant = tenants_fixture[0] + scan = AttackPathsScan.objects.create( + tenant_id=tenant.id, + provider=aws_provider, + state=StateChoices.COMPLETED, + ) + + database = _tmp_db_name(scan.id) + mock_list.return_value = [database] + + result = reap_orphaned_tmp_databases() + + assert result == {"dropped_count": 0, "databases": []} + mock_drop.assert_not_called() + + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.drop_database") + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.list_databases") + def test_margin_counts_from_completed_at_not_updated_at( + self, mock_list, mock_drop, tenants_fixture, aws_provider + ): + from tasks.jobs.attack_paths.tmp_db_reaper import reap_orphaned_tmp_databases + + tenant = tenants_fixture[0] + old = datetime.now(tz=UTC) - MARGIN - timedelta(hours=1) + scan = AttackPathsScan.objects.create( + tenant_id=tenant.id, + provider=aws_provider, + state=StateChoices.COMPLETED, + ) + AttackPathsScan.objects.filter(id=scan.id).update( + updated_at=old, completed_at=datetime.now(tz=UTC) + ) + + database = _tmp_db_name(scan.id) + mock_list.return_value = [database] + + result = reap_orphaned_tmp_databases() + + assert result == {"dropped_count": 0, "databases": []} + mock_drop.assert_not_called() + + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.drop_database") + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.list_databases") + def test_drops_scan_completed_past_safety_margin( + self, mock_list, mock_drop, tenants_fixture, aws_provider + ): + from tasks.jobs.attack_paths.tmp_db_reaper import reap_orphaned_tmp_databases + + tenant = tenants_fixture[0] + old = datetime.now(tz=UTC) - MARGIN - timedelta(hours=1) + scan = AttackPathsScan.objects.create( + tenant_id=tenant.id, + provider=aws_provider, + state=StateChoices.FAILED, + ) + AttackPathsScan.objects.filter(id=scan.id).update( + updated_at=old, completed_at=old + ) + + database = _tmp_db_name(scan.id) + mock_list.return_value = [database] + + result = reap_orphaned_tmp_databases() + + assert result == {"dropped_count": 1, "databases": [database]} + mock_drop.assert_called_once_with(database) + + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.drop_database") + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.list_databases") + def test_never_drops_an_executing_scan_regardless_of_age( + self, mock_list, mock_drop, tenants_fixture, aws_provider + ): + from tasks.jobs.attack_paths.tmp_db_reaper import reap_orphaned_tmp_databases + + tenant = tenants_fixture[0] + very_old = datetime.now(tz=UTC) - timedelta(days=30) + scan = AttackPathsScan.objects.create( + tenant_id=tenant.id, + provider=aws_provider, + state=StateChoices.EXECUTING, + ) + AttackPathsScan.objects.filter(id=scan.id).update(updated_at=very_old) + + database = _tmp_db_name(scan.id) + mock_list.return_value = [database] + + result = reap_orphaned_tmp_databases() + + assert result == {"dropped_count": 0, "databases": []} + mock_drop.assert_not_called() + + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.drop_database") + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.list_databases") + def test_one_failed_drop_does_not_stop_the_rest_of_the_sweep( + self, mock_list, mock_drop + ): + from tasks.jobs.attack_paths.tmp_db_reaper import reap_orphaned_tmp_databases + + old_time = datetime.now(tz=UTC) - MARGIN - timedelta(hours=1) + failing_scan_id = datetime_to_uuid7(old_time) + succeeding_scan_id = datetime_to_uuid7(old_time) + failing_db = _tmp_db_name(failing_scan_id) + succeeding_db = _tmp_db_name(succeeding_scan_id) + mock_list.return_value = [failing_db, succeeding_db] + mock_drop.side_effect = [Exception("boom"), None] + + result = reap_orphaned_tmp_databases() + + assert result == {"dropped_count": 1, "databases": [succeeding_db]} + assert mock_drop.call_count == 2 + + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.list_databases") + def test_returns_empty_result_when_listing_databases_fails(self, mock_list): + from tasks.jobs.attack_paths.tmp_db_reaper import reap_orphaned_tmp_databases + + mock_list.side_effect = Exception("neo4j unreachable") + + result = reap_orphaned_tmp_databases() + + assert result == {"dropped_count": 0, "databases": []} + + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.drop_database") + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.list_databases") + def test_preserves_temp_db_with_random_uuid_and_no_row(self, mock_list, mock_drop): + """A non-UUIDv7 id with no matching row has no reliable timestamp, so it + must be left alone rather than guessed at.""" + from tasks.jobs.attack_paths.tmp_db_reaper import reap_orphaned_tmp_databases + + database = _tmp_db_name(uuid4()) + mock_list.return_value = [database] + + result = reap_orphaned_tmp_databases() + + assert result == {"dropped_count": 0, "databases": []} + mock_drop.assert_not_called() + + +class TestReapOrphanedTmpDatabasesTask: + @patch( + "tasks.tasks.reap_orphaned_tmp_databases", + return_value={"dropped_count": 2, "databases": ["db-tmp-scan-a"]}, + ) + def test_task_invokes_the_reaper(self, mock_reap): + from tasks.tasks import reap_orphaned_attack_paths_tmp_databases_task + + result = reap_orphaned_attack_paths_tmp_databases_task.run() + + assert result == {"dropped_count": 2, "databases": ["db-tmp-scan-a"]} + mock_reap.assert_called_once_with() diff --git a/api/src/backend/tasks/tests/test_export.py b/api/src/backend/tasks/tests/test_export.py index bbec0b742e..3fca90e065 100644 --- a/api/src/backend/tasks/tests/test_export.py +++ b/api/src/backend/tasks/tests/test_export.py @@ -3,9 +3,10 @@ import uuid import zipfile from datetime import datetime from pathlib import Path -from unittest.mock import MagicMock, patch +from unittest.mock import MagicMock, call, patch from urllib.parse import parse_qs, urlparse +import boto3 import pytest from botocore.exceptions import ClientError from django.test import override_settings @@ -63,15 +64,46 @@ class TestOutputs: assert mock_boto_client.call_args.kwargs["region_name"] == "us-east-1" + @patch("tasks.jobs.export.boto3.client") + @override_settings(DJANGO_OUTPUT_S3_AWS_ENDPOINT_URL="http://minio:9000") + def test_get_s3_client_passes_the_endpoint_when_set(self, mock_boto_client): + get_s3_client() + + assert mock_boto_client.call_args.kwargs["endpoint_url"] == "http://minio:9000" + + @patch("tasks.jobs.export.boto3.client") + @override_settings(DJANGO_OUTPUT_S3_AWS_ENDPOINT_URL="") + def test_get_s3_client_endpoint_empty_by_default(self, mock_boto_client): + """Empty keeps today's behavior: no endpoint override, real S3 is used.""" + get_s3_client() + + assert mock_boto_client.call_args.kwargs["endpoint_url"] is None + + @patch("tasks.jobs.export.boto3.client") + @override_settings(DJANGO_OUTPUT_S3_AWS_ENDPOINT_URL="http://minio:9000") + def test_get_s3_client_does_not_fall_back_when_endpoint_set(self, mock_boto_client): + """A configured endpoint means the explicit client failed talking to it. The fallback + goes to the default provider chain (e.g. an EC2 instance role) against real AWS, so it + must not be used: the original error propagates instead.""" + error = ClientError({"Error": {"Code": "403"}}, "ListBuckets") + mock_boto_client.side_effect = error + + with pytest.raises(ClientError): + get_s3_client() + + mock_boto_client.assert_called_once() + @patch("tasks.jobs.export.boto3.client") @patch("tasks.jobs.export.settings") def test_get_s3_client_fallback(self, mock_settings, mock_boto_client): + mock_settings.DJANGO_OUTPUT_S3_AWS_ENDPOINT_URL = "" mock_boto_client.side_effect = [ ClientError({"Error": {"Code": "403"}}, "ListBuckets"), MagicMock(), ] client = get_s3_client() assert client is not None + assert mock_boto_client.call_args_list[1] == call("s3") @patch("tasks.jobs.export.get_s3_client") @patch("tasks.jobs.export.base") @@ -278,10 +310,70 @@ def _presign(client): class TestS3PresignClient: - @override_settings(**PRESIGN_SETTINGS, DJANGO_OUTPUT_S3_AWS_PUBLIC_ENDPOINT_URL="") - def test_no_public_endpoint_returns_none(self): + @override_settings( + **{**PRESIGN_SETTINGS, "DJANGO_OUTPUT_S3_AWS_DEFAULT_REGION": ""}, + DJANGO_OUTPUT_S3_AWS_PUBLIC_ENDPOINT_URL="", + ) + def test_no_public_endpoint_and_no_region_returns_none(self): + # Without a region, SigV4 would have to guess one and break other regions. assert get_s3_presign_client() is None + @override_settings(**PRESIGN_SETTINGS, DJANGO_OUTPUT_S3_AWS_PUBLIC_ENDPOINT_URL="") + def test_region_without_public_endpoint_signs_sigv4_on_the_regional_host(self): + # SSE-KMS objects reject the SigV2 URLs boto3 presigns by default, and the + # global host redirects for new buckets, which breaks a SigV4 signature. + url = urlparse(_presign(get_s3_presign_client())) + query = parse_qs(url.query) + + assert url.netloc == "s3.eu-west-1.amazonaws.com" + assert url.path == "/output-bucket/tenant/scan/report.zip" + assert query["X-Amz-Algorithm"] == ["AWS4-HMAC-SHA256"] + assert "/eu-west-1/s3/aws4_request" in query["X-Amz-Credential"][0] + + @override_settings( + **{ + **PRESIGN_SETTINGS, + "DJANGO_OUTPUT_S3_AWS_ACCESS_KEY_ID": "", + "DJANGO_OUTPUT_S3_AWS_SECRET_ACCESS_KEY": "", + }, + DJANGO_OUTPUT_S3_AWS_PUBLIC_ENDPOINT_URL="", + ) + def test_region_without_static_keys_signs_with_the_default_chain(self, monkeypatch): + # An ECS task role reaches boto3 through the default chain, like the env here. + # A fresh default session keeps these keys from being cached for later tests. + monkeypatch.setattr(boto3, "DEFAULT_SESSION", None) + monkeypatch.setenv("AWS_ACCESS_KEY_ID", "role-access-key") + monkeypatch.setenv("AWS_SECRET_ACCESS_KEY", "role-secret-key") + monkeypatch.setenv("AWS_DEFAULT_REGION", "us-east-1") + + query = parse_qs(urlparse(_presign(get_s3_presign_client())).query) + + assert query["X-Amz-Credential"][0].startswith("role-access-key/") + assert "/eu-west-1/s3/aws4_request" in query["X-Amz-Credential"][0] + + @override_settings( + **{**PRESIGN_SETTINGS, "DJANGO_OUTPUT_S3_AWS_DEFAULT_REGION": ""}, + DJANGO_OUTPUT_S3_AWS_PUBLIC_ENDPOINT_URL="", + DJANGO_OUTPUT_S3_AWS_ENDPOINT_URL="http://minio:9000", + ) + def test_internal_endpoint_without_public_endpoint_signs_against_it(self): + # No browser-reachable host was configured, so the internal one is the best + # available target instead of falling through to the real AWS host. + url = urlparse(_presign(get_s3_presign_client())) + + assert url.netloc == "minio:9000" + assert url.path == "/output-bucket/tenant/scan/report.zip" + + @override_settings( + **PRESIGN_SETTINGS, + DJANGO_OUTPUT_S3_AWS_PUBLIC_ENDPOINT_URL="https://storage.example.com", + DJANGO_OUTPUT_S3_AWS_ENDPOINT_URL="http://minio:9000", + ) + def test_public_endpoint_wins_over_the_internal_endpoint(self): + url = urlparse(_presign(get_s3_presign_client())) + + assert url.netloc == "storage.example.com" + @override_settings( **PRESIGN_SETTINGS, DJANGO_OUTPUT_S3_AWS_PUBLIC_ENDPOINT_URL="https://storage.example.com", diff --git a/api/src/backend/tasks/tests/test_scan.py b/api/src/backend/tasks/tests/test_scan.py index 0cee24351e..efe36341fc 100644 --- a/api/src/backend/tasks/tests/test_scan.py +++ b/api/src/backend/tasks/tests/test_scan.py @@ -5967,6 +5967,112 @@ class TestResetEphemeralResourceFindingsCount: resource2.refresh_from_db() assert resource2.failed_findings_count == 5 + def test_runs_when_newer_scan_is_not_full_scope( + self, tenants_fixture, scans_fixture, aws_provider, resources_fixture + ): + """A newer scoped scan must not block the full-scope scan's cleanup. + + The race guard used to pick the newest COMPLETED scan of any scope + despite being named after full-scope ones, so a single scoped scan + landing after a complete one made the complete scan skip its cleanup + and leave ephemeral resources with a stale count permanently. + """ + from datetime import timedelta + + tenant, *_ = tenants_fixture + scan1, *_ = scans_fixture + resource1, resource2, _ = resources_fixture + + Resource.objects.filter(id=resource2.id).update(failed_findings_count=5) + self._make_scan_summary(tenant.id, scan1.id, resource1) + + newer_completed_at = scan1.completed_at + timedelta(minutes=5) + Scan.objects.create( + name="Newer scoped scan", + provider=aws_provider, + trigger=Scan.TriggerChoices.MANUAL, + state=StateChoices.COMPLETED, + tenant_id=tenant.id, + started_at=newer_completed_at, + completed_at=newer_completed_at, + scanner_args={"checks": ["check1"]}, + ) + + result = reset_ephemeral_resource_findings_count( + tenant_id=str(tenant.id), scan_id=str(scan1.id) + ) + + assert result["status"] == "completed" + + resource2.refresh_from_db() + assert resource2.failed_findings_count == 0 + + def test_runs_when_many_newer_scans_are_not_full_scope( + self, tenants_fixture, scans_fixture, aws_provider, resources_fixture + ): + """The walk must not give up before it reaches the full-scope scan. + + A fixed look-back window returned None once more newer scoped scans + had landed than it inspected, and skipped the cleanup exactly like the + original bug did with one. + """ + from datetime import timedelta + + tenant, *_ = tenants_fixture + scan1, *_ = scans_fixture + resource1, resource2, _ = resources_fixture + + Resource.objects.filter(id=resource2.id).update(failed_findings_count=5) + self._make_scan_summary(tenant.id, scan1.id, resource1) + + for minutes in range(1, 41): + newer_completed_at = scan1.completed_at + timedelta(minutes=minutes) + Scan.objects.create( + name=f"Newer scoped scan {minutes}", + provider=aws_provider, + trigger=Scan.TriggerChoices.MANUAL, + state=StateChoices.COMPLETED, + tenant_id=tenant.id, + started_at=newer_completed_at, + completed_at=newer_completed_at, + scanner_args={"checks": ["check1"]}, + ) + + result = reset_ephemeral_resource_findings_count( + tenant_id=str(tenant.id), scan_id=str(scan1.id) + ) + + assert result["status"] == "completed" + + resource2.refresh_from_db() + assert resource2.failed_findings_count == 0 + + def test_runs_when_completed_at_is_null( + self, tenants_fixture, scans_fixture, aws_provider, resources_fixture + ): + """NULL `completed_at` used to make the reset never run at all. + + The old guard filtered `completed_at__isnull=False`, so a provider + whose completed scans all had a NULL `completed_at` resolved the + "latest" scan to None, which never equals `scan.id`. + """ + tenant, *_ = tenants_fixture + scan1, *_ = scans_fixture + resource1, resource2, _ = resources_fixture + + Scan.all_objects.filter(id=scan1.id).update(completed_at=None) + Resource.objects.filter(id=resource2.id).update(failed_findings_count=5) + self._make_scan_summary(tenant.id, scan1.id, resource1) + + result = reset_ephemeral_resource_findings_count( + tenant_id=str(tenant.id), scan_id=str(scan1.id) + ) + + assert result["status"] == "completed" + + resource2.refresh_from_db() + assert resource2.failed_findings_count == 0 + def test_does_not_touch_other_providers_resources( self, tenants_fixture, scans_fixture, aws_provider, resources_fixture ): diff --git a/api/src/backend/tasks/tests/test_tasks.py b/api/src/backend/tasks/tests/test_tasks.py index b4854fe196..5ab1cdcf54 100644 --- a/api/src/backend/tasks/tests/test_tasks.py +++ b/api/src/backend/tasks/tests/test_tasks.py @@ -1,3 +1,4 @@ +import json import uuid from contextlib import contextmanager from datetime import UTC, datetime @@ -14,6 +15,7 @@ from api.models import ( StateChoices, Task, ) +from api.v1.serializers import TaskSerializer from botocore.exceptions import ClientError from celery import states from django_celery_beat.models import IntervalSchedule, PeriodicTask @@ -32,6 +34,7 @@ from tasks.tasks import ( _scan_tmp_output_directory, check_integrations_task, check_lighthouse_provider_connection_task, + create_scan_task_record, generate_outputs_task, mute_findings_in_latest_scans_task, perform_attack_paths_scan_task, @@ -3384,3 +3387,79 @@ class TestTaskTimeLimits: "lighthouse-provider-connection-check", ): assert celery_app.tasks[name].time_limit < default + + +@pytest.mark.django_db +class TestCreateScanTaskRecord: + """`task_kwargs` is what a response built before the publish can report.""" + + def _scan(self, tenant, provider): + """A manual scan, like the one `POST /api/v1/scans` creates.""" + return Scan.objects.create( + tenant_id=tenant.id, + provider=provider, + name="Manual scan", + trigger=Scan.TriggerChoices.MANUAL, + state=StateChoices.AVAILABLE, + ) + + def _publish_kwargs(self, tenant, scan): + """What `enqueue_scan_execution_on_commit` publishes for this scan.""" + return { + "tenant_id": str(tenant.id), + "scan_id": str(scan.id), + "provider_id": str(scan.provider_id), + } + + def _task_args(self, task): + """Read the record back the way `TaskSerializer` does.""" + return TaskSerializer(task).data["task_args"] + + def test_the_stored_kwargs_are_the_ones_the_publish_would_send( + self, tenants_fixture, aws_provider + ): + """The 202 reports what is stored here, so it has to be the dispatch kwargs.""" + tenant = tenants_fixture[0] + scan = self._scan(tenant, aws_provider) + + task = create_scan_task_record( + tenant_id=str(tenant.id), + task_id=str(uuid.uuid4()), + task_kwargs=self._publish_kwargs(tenant, scan), + ) + + assert self._task_args(task) == { + "scan_id": str(scan.id), + "provider_id": str(aws_provider.id), + } + + def test_a_record_created_without_kwargs_reports_none(self, tenants_fixture): + """The argument is optional, so the other callers keep their behaviour.""" + task = create_scan_task_record( + tenant_id=str(tenants_fixture[0].id), + task_id=str(uuid.uuid4()), + ) + + assert self._task_args(task) == {} + + def test_the_publish_can_overwrite_the_stored_kwargs( + self, tenants_fixture, aws_provider + ): + """django-celery-results stores a Python repr; both must decode alike.""" + tenant = tenants_fixture[0] + scan = self._scan(tenant, aws_provider) + task_id = str(uuid.uuid4()) + kwargs = self._publish_kwargs(tenant, scan) + + task = create_scan_task_record( + tenant_id=str(tenant.id), task_id=task_id, task_kwargs=kwargs + ) + before = self._task_args(task) + + # What `before_task_publish` writes once the task reaches the broker. + task_result = TaskResult.objects.get(task_id=task_id) + task_result.task_kwargs = json.dumps(repr(kwargs)) + task_result.save(update_fields=["task_kwargs"]) + task.refresh_from_db() + + assert self._task_args(task) == before diff --git a/api/uv.lock b/api/uv.lock index a835e2f223..7f6151d6d1 100644 --- a/api/uv.lock +++ b/api/uv.lock @@ -53,7 +53,7 @@ constraints = [ { name = "aliyun-log-fastpb", specifier = "==0.2.0" }, { name = "amqp", specifier = "==5.3.1" }, { name = "annotated-types", specifier = "==0.7.0" }, - { name = "anyio", specifier = "==4.12.1" }, + { name = "anyio", specifier = "==4.14.2" }, { name = "applicationinsights", specifier = "==0.11.10" }, { name = "apscheduler", specifier = "==3.11.2" }, { name = "argcomplete", specifier = "==3.5.3" }, @@ -969,15 +969,15 @@ wheels = [ [[package]] name = "anyio" -version = "4.12.1" +version = "4.14.2" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "idna" }, { name = "typing-extensions" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/96/f0/5eb65b2bb0d09ac6776f2eb54adee6abe8228ea05b20a5ad0e4945de8aac/anyio-4.12.1.tar.gz", hash = "sha256:41cfcc3a4c85d3f05c932da7c26d0201ac36f72abd4435ba90d0464a3ffed703", size = 228685, upload-time = "2026-01-06T11:45:21.246Z" } +sdist = { url = "https://files.pythonhosted.org/packages/61/cc/a381afa6efea9f496eff839d4a6a1aed3bfafc7b3ab4b0d1b243a12573dd/anyio-4.14.2.tar.gz", hash = "sha256:cfa139f3ed1a23ee8f88a145ddb5ac7605b8bbfd8592baacd7ce3d8bb4313c7f", size = 260176, upload-time = "2026-07-12T20:29:07.082Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/38/0e/27be9fdef66e72d64c0cdc3cc2823101b80585f8119b5c112c2e8f5f7dab/anyio-4.12.1-py3-none-any.whl", hash = "sha256:d405828884fc140aa80a3c667b8beed277f1dfedec42ba031bd6ac3db606ab6c", size = 113592, upload-time = "2026-01-06T11:45:19.497Z" }, + { url = "https://files.pythonhosted.org/packages/da/35/f2287558c17e29fafc8ef3daf819bb9834061cfa43bff8014f7df7f63bdc/anyio-4.14.2-py3-none-any.whl", hash = "sha256:9f505dda5ac9f0c8309b5e8bd445a8c2bf7246f3ce950121e45ea15bc41d1494", size = 125813, upload-time = "2026-07-12T20:29:05.763Z" }, ] [[package]] @@ -4938,7 +4938,7 @@ dependencies = [ [[package]] name = "prowler-api" -version = "1.44.0" +version = "1.45.0" source = { virtual = "." } dependencies = [ { name = "cartography" }, diff --git a/docker-compose.yml b/docker-compose.yml index 5ed0e97a28..4b836aba6c 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -15,6 +15,7 @@ services: api: hostname: "prowler-api" image: prowlercloud/prowler-api:${PROWLER_API_VERSION:-stable} + restart: unless-stopped env_file: - path: .env required: false @@ -44,6 +45,7 @@ services: ui: image: prowlercloud/prowler-ui:${PROWLER_UI_VERSION:-stable} + restart: unless-stopped env_file: - path: .env required: false @@ -61,6 +63,7 @@ services: postgres: image: postgres:16-alpine@sha256:57c72fd2a128e416c7fcc499958864df5301e940bca0a56f58fddf30ffc07777 + restart: unless-stopped hostname: "postgres-db" volumes: - ./_data/postgres:/var/lib/postgresql/data @@ -81,6 +84,7 @@ services: valkey: image: valkey/valkey:8-alpine@sha256:a038175878d66b9d274fbf8be73c0305e93798b83917647f167e18cef3c71eec + restart: unless-stopped hostname: "valkey" volumes: - ./_data/valkey:/data @@ -97,6 +101,7 @@ services: neo4j: image: graphstack/dozerdb:5.26.27.0@sha256:9b54d6b3a98a76c00bd23e8e78d8c82081ff168162aebd47b25c234e092cb0a0 + restart: unless-stopped hostname: "neo4j" volumes: - ./_data/neo4j:/data @@ -129,6 +134,7 @@ services: worker: image: prowlercloud/prowler-api:${PROWLER_API_VERSION:-stable} + restart: unless-stopped # Give Celery soft shutdown time to drain/re-queue in-flight tasks on stop. stop_grace_period: 120s env_file: @@ -149,6 +155,7 @@ services: worker-beat: image: prowlercloud/prowler-api:${PROWLER_API_VERSION:-stable} + restart: unless-stopped env_file: - path: ./.env required: false @@ -165,6 +172,7 @@ services: mcp-server: image: prowlercloud/prowler-mcp:${PROWLER_MCP_VERSION:-stable} + restart: unless-stopped environment: - PROWLER_MCP_TRANSPORT_MODE=http env_file: diff --git a/docs/changelog.mdx b/docs/changelog.mdx index 7c01d35a17..cfbd393fd8 100644 --- a/docs/changelog.mdx +++ b/docs/changelog.mdx @@ -4,6 +4,73 @@ description: "New features and improvements in each Prowler release" rss: true --- + + ### 🏛️ Compliance — FedRAMP 20x Consolidated Rules 2026 + + The FedRAMP 20x Phase One pilot frameworks (`fedramp_20x_ksi_low_aws`, `fedramp_20x_ksi_low_azure` and `fedramp_20x_ksi_low_gcp`) are replaced by two universal frameworks built from the FedRAMP Consolidated Rules for 2026, each covering AWS, Azure, GCP, Kubernetes and Microsoft 365 from a single definition: + + - **FedRAMP 20x KSI** (`fedramp_20x_ksi_2026`): the 46 Key Security Indicators across 10 themes. There is one indicator catalog for every class instead of a separate Low fork; each indicator carries its class applicability and NIST SP 800-53 controls. + - **FedRAMP 20x Class C FRR** (`fedramp_20x_frr_class_c_2026`): the 158 FedRAMP Rules of the Class C ruleset that bind cloud service providers. Most are program obligations (reports, notifications, certification package) and stay manual; checks are mapped only where they evidence part of the rule text. Configurable checks carry configuration requirements, so a relaxed `audit_config` cannot turn a requirement green. + + Automation or stored results that reference the pilot framework IDs need to move to `fedramp_20x_ksi_2026`. In Prowler App, both frameworks show the per-provider breakdown in the cross-provider compliance view and can be downloaded as OCSF. + + Read more in the [Compliance documentation](https://docs.prowler.com/user-guide/compliance/tutorials/compliance). + + ### 🔎 AWS — Inspector Coverage, CISA KEV and FIPS Checks + + Seven new AWS checks back the vulnerability detection and cryptography rules of FedRAMP 20x Class C: + + - `inspector2_coverage_scan_status_active` and `inspector2_coverage_recently_scanned` report resources Amazon Inspector is not scanning, or last scanned more than `inspector2_max_days_since_last_scan` days ago (default 3). + - `inspector2_active_findings_no_known_exploited_vulnerabilities` and `inspector2_active_findings_kev_within_due_date` report active findings whose CVE is in the CISA Known Exploited Vulnerabilities catalog, and those still open past the CISA due date. The KEV data comes from Inspector itself through `inspector2:BatchGetFindingDetails`, so no external feed is needed. + - `inspector2_active_findings_within_max_age` reports active findings first observed more than `inspector2_active_finding_max_age_days` days ago (default 192). + - `elbv2_listener_fips_tls_enabled` and `transfer_server_fips_security_policy_enabled` report HTTPS/TLS load balancer listeners and Transfer Family servers without a FIPS security policy. + + `inspector2:BatchGetFindingDetails` is not part of `SecurityAudit`, so it is now included in the Prowler additions policy and the CloudFormation scan role. Without it, the KEV checks report `MANUAL` naming the missing permission instead of a false `FAIL`. + + Explore all AWS checks at [Prowler Hub](https://hub.prowler.com/check?provider=aws). + + ### ☁️ AWS — Partition Bootstrap Failover + + When `PROWLER_AWS_PARTITION` is set, the bootstrap STS calls (validating credentials, assuming a role and getting an MFA session token) now try up to two more regions of the partition if the first one cannot be reached. A GovCloud host whose configured region belongs to another partition was still sent to `us-gov-east-1`, and on a network that routes only to `us-gov-west-1` the connection check and the scan failed on perfectly valid credentials. Only connection errors and timeouts move on to the next region; credential errors are reported from the first one as before. Later STS calls reuse the region that answered, and nothing changes when `PROWLER_AWS_PARTITION` is unset. + + Read more in the [AWS Regions and Partitions documentation](https://docs.prowler.com/user-guide/providers/aws/regions-and-partitions). + + ### 🌐 Azure — Sovereign Cloud Endpoints for Defender and Key Vault + + Defender security contacts and Key Vault key rotation policies now use the endpoints of the cloud selected with `--azure-region` instead of the hardcoded `management.azure.com` and `vault.azure.net` hosts, so both work on `AzureUSGovernment` and `AzureChinaCloud`. Key Vault clients are built from the vault URI that Azure returns for each vault. + + Read more in the [Azure non-default cloud documentation](https://docs.prowler.com/user-guide/providers/azure/use-non-default-cloud). + + ### ✉️ Invitations — Expired Invitations No Longer Block Re-Invites + + A pending invitation past its expiry date is now reported as expired, and inviting the same email again marks it as expired and creates the new invitation instead of returning a generic error. The Invitations table disables Edit and Revoke on expired and revoked invitations, and `filter[state__in]` on the invitations endpoint no longer returns a server error. + + In Prowler Cloud, new organizations are offered an **Invite your team** step once the first provider is connected, and Prowler Private Cloud deployments can set `UI_SELF_REGISTRATION_ENABLED=false` to make sign-up invitation-only. + + Read more in the [Invitations documentation](https://docs.prowler.com/user-guide/tutorials/prowler-app-rbac#invitations). + + ### 📄 Reports — Downloads on Self-Hosted Storage + + Report downloads no longer depend on the storage host being the same inside and outside the container network. `DJANGO_OUTPUT_S3_AWS_PUBLIC_ENDPOINT_URL` signs the download URL against a browser-reachable host, so a deployment whose object storage answers only on an internal address serves the file instead of a link the browser cannot open. Leaving `DJANGO_OUTPUT_S3_AWS_DEFAULT_REGION` unset, common on S3-compatible storage with no meaningful region, no longer makes the download fail with a server error. + + ### 🔍 Checks + + - **Huawei Cloud:** new `smn_topic_subscriptions` check reports SMN topics without any subscription. + - **Cloudflare:** the API token links in the provider wizard request the SSL and Certificates, Bot Management and Zone WAF read permissions the checks need, so a token created from the wizard no longer produces failures on permissions it was never granted. + - **Microsoft 365:** five Defender malware, anti-phishing and inbound anti-spam checks no longer fail with `KeyError` on tenants that use the Standard or Strict preset security policies, which dropped every finding of those checks. Preset policies are covered by `defender_strict_preset_security_policy_enabled`. + - **Google Workspace:** `security_2sv_enforced` reports domain-wide 2-Step Verification failures as `FAIL` even when every failing setting is overridden for a group or organizational unit. + + Explore all checks at [Prowler Hub](https://hub.prowler.com/check). + + ### 🔐 Security Updates + + - `libsqlite3-0`, `gzip`, `perl-base`, `libssh2-1t64` and `libpcre2-8-0` upgraded in the SDK and API container images, patching high-severity Debian CVEs. + - PowerShell upgraded to 7.5.11 in the SDK and API container images, bundling .NET runtime 9.0.20 and patching CVE-2026-62901. + - `anyio` upgraded to 4.14.2 in the SDK, the API and the MCP Server, patching CVE-2026-63374. + + See the [full release notes on GitHub](https://github.com/prowler-cloud/prowler/releases/tag/5.43.0) for the complete list of changes. + + ### ☁️ AWS — ISO Partitions diff --git a/docs/getting-started/installation/prowler-app.mdx b/docs/getting-started/installation/prowler-app.mdx index b27b7142d2..c1a00312aa 100644 --- a/docs/getting-started/installation/prowler-app.mdx +++ b/docs/getting-started/installation/prowler-app.mdx @@ -128,8 +128,8 @@ To update the environment file: Edit the `.env` file and change version values: ```env -PROWLER_UI_VERSION="5.42.0" -PROWLER_API_VERSION="5.42.0" +PROWLER_UI_VERSION="5.43.0" +PROWLER_API_VERSION="5.43.0" ``` diff --git a/docs/introduction.mdx b/docs/introduction.mdx index 01ba5ccb63..56bfe4a169 100644 --- a/docs/introduction.mdx +++ b/docs/introduction.mdx @@ -71,7 +71,7 @@ Prowler supports a wide range of providers organized by category: | [LLM](/user-guide/providers/llm/getting-started-llm) | Official | Models | CLI | | [M365](/user-guide/providers/microsoft365/getting-started-m365) | Official | Tenants | UI, API, CLI | | [MongoDB Atlas](/user-guide/providers/mongodbatlas/getting-started-mongodbatlas) | Official | Organizations | UI, API, CLI | -| [Okta](/user-guide/providers/okta/getting-started-okta) | Official | Organizations | CLI | +| [Okta](/user-guide/providers/okta/getting-started-okta) | Official | Organizations | UI, API, CLI | | [Vercel](/user-guide/providers/vercel/getting-started-vercel) | Official | Teams / Projects | UI, API, CLI | ### Kubernetes diff --git a/docs/user-guide/providers/aws/boto3-configuration.mdx b/docs/user-guide/providers/aws/boto3-configuration.mdx index 80fd4a30fb..070f6f6557 100644 --- a/docs/user-guide/providers/aws/boto3-configuration.mdx +++ b/docs/user-guide/providers/aws/boto3-configuration.mdx @@ -29,6 +29,22 @@ Boto3 defaults both timeouts to 60 seconds. In networks with restricted egress ( +## Retries Configuration + + + +The number of retries is set with `--aws-retries-max-attempts`, where `0` disables retries. It can also be set through an environment variable, which is the way to tune it in Prowler Cloud and other deployments without a CLI: + +```console +export PROWLER_AWS_BOTO3_RETRIES_MAX_ATTEMPTS=0 +``` + +The CLI flag takes precedence over the environment variable. The value must be a non-negative integer; when neither is set, Prowler uses 3 retries. + + +The environment variable is process-wide: it applies to every AWS provider built in the process where it is set, not only to a connection check. A scan started in that same process picks it up too. Boto3's Standard retry mode, which Prowler uses, also retries service-side throttling responses (see the errors listed below), so `0` disables retries for those as well. On a large account a scan can hit throttling under normal load, and with retries disabled that throttling becomes a hard failure instead of a retried call. Set the variable only on the processes that run connection checks; leave scan workers on the default, or raise their retry count instead of lowering it. + + ## Retry Behavior Overview Boto3's Standard retry mode includes the following mechanisms: diff --git a/docs/user-guide/tutorials/prowler-app.mdx b/docs/user-guide/tutorials/prowler-app.mdx index ddbd489cdc..db0916e4a4 100644 --- a/docs/user-guide/tutorials/prowler-app.mdx +++ b/docs/user-guide/tutorials/prowler-app.mdx @@ -89,6 +89,10 @@ After adding your cloud account credentials, click the `Check connection` button Test Connection + +For a single AWS account, Prowler tests the connection as part of the `Connect account` step, so the wizard moves straight to launching the scan. + + ## Step 6: Scan Started After the connection check succeeds, save the provider and start your first scan with the `Launch Scan` button. The `Scans` section shows the scan in progress: diff --git a/mcp_server/CHANGELOG.md b/mcp_server/CHANGELOG.md index 80bdabd28a..07e92992fc 100644 --- a/mcp_server/CHANGELOG.md +++ b/mcp_server/CHANGELOG.md @@ -4,6 +4,14 @@ All notable changes to the **Prowler MCP Server** are documented in this file. +## [0.12.2] (Prowler v5.43.0) + +### 🔐 Security + +- Bumped `anyio` to 4.14.2 to resolve CVE-2026-63374 [(#12848)](https://github.com/prowler-cloud/prowler/pull/12848) + +--- + ## [0.12.1] (Prowler v5.42.0) ### 🔐 Security diff --git a/mcp_server/uv.lock b/mcp_server/uv.lock index 849847e3c8..125087fb20 100644 --- a/mcp_server/uv.lock +++ b/mcp_server/uv.lock @@ -39,15 +39,15 @@ wheels = [ [[package]] name = "anyio" -version = "4.13.0" +version = "4.14.2" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "idna" }, { name = "typing-extensions", marker = "python_full_version < '3.13'" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/19/14/2c5dd9f512b66549ae92767a9c7b330ae88e1932ca57876909410251fe13/anyio-4.13.0.tar.gz", hash = "sha256:334b70e641fd2221c1505b3890c69882fe4a2df910cba14d97019b90b24439dc", size = 231622, upload-time = "2026-03-24T12:59:09.671Z" } +sdist = { url = "https://files.pythonhosted.org/packages/61/cc/a381afa6efea9f496eff839d4a6a1aed3bfafc7b3ab4b0d1b243a12573dd/anyio-4.14.2.tar.gz", hash = "sha256:cfa139f3ed1a23ee8f88a145ddb5ac7605b8bbfd8592baacd7ce3d8bb4313c7f", size = 260176, upload-time = "2026-07-12T20:29:07.082Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/da/42/e921fccf5015463e32a3cf6ee7f980a6ed0f395ceeaa45060b61d86486c2/anyio-4.13.0-py3-none-any.whl", hash = "sha256:08b310f9e24a9594186fd75b4f73f4a4152069e3853f1ed8bfbf58369f4ad708", size = 114353, upload-time = "2026-03-24T12:59:08.246Z" }, + { url = "https://files.pythonhosted.org/packages/da/35/f2287558c17e29fafc8ef3daf819bb9834061cfa43bff8014f7df7f63bdc/anyio-4.14.2-py3-none-any.whl", hash = "sha256:9f505dda5ac9f0c8309b5e8bd445a8c2bf7246f3ce950121e45ea15bc41d1494", size = 125813, upload-time = "2026-07-12T20:29:05.763Z" }, ] [[package]] diff --git a/prowler/CHANGELOG.md b/prowler/CHANGELOG.md index 7232bd41a0..3a7bc157bf 100644 --- a/prowler/CHANGELOG.md +++ b/prowler/CHANGELOG.md @@ -4,6 +4,34 @@ All notable changes to the **Prowler SDK** are documented in this file. +## [5.43.0] (Prowler v5.43.0) + +### 🚀 Added + +- `FedRAMP-20x-KSI` universal compliance framework (`fedramp_20x_ksi_2026`) with the 46 Key Security Indicators from the FedRAMP Consolidated Rules 2026 mapped for AWS, Azure, GCP, Kubernetes and M365 [(#11701)](https://github.com/prowler-cloud/prowler/pull/11701) +- `smn_topic_subscriptions` check for Huawei Cloud provider: SMN topics have at least one subscription configured [(#12186)](https://github.com/prowler-cloud/prowler/pull/12186) +- `inspector2_coverage_scan_status_active`, `inspector2_coverage_recently_scanned`, `inspector2_active_findings_no_known_exploited_vulnerabilities`, `inspector2_active_findings_kev_within_due_date`, `inspector2_active_findings_within_max_age`, `elbv2_listener_fips_tls_enabled` and `transfer_server_fips_security_policy_enabled` checks for AWS provider, covering FedRAMP 20x Class C vulnerability detection, CISA KEV remediation and FIPS cryptography rules; the KEV checks require `inspector2:BatchGetFindingDetails`, now in the Prowler additions policy [(#12808)](https://github.com/prowler-cloud/prowler/pull/12808) +- `FedRAMP-20x-FRR-Class-C` universal compliance framework (`fedramp_20x_frr_class_c_2026`) with the 158 provider rules of the FedRAMP 20x Class C ruleset from the FedRAMP Consolidated Rules 2026 for AWS, Azure, GCP, Kubernetes and M365 [(#12808)](https://github.com/prowler-cloud/prowler/pull/12808) + +### 🔄 Changed + +- FedRAMP 20x Phase One pilot frameworks `fedramp_20x_ksi_low_aws`, `fedramp_20x_ksi_low_azure` and `fedramp_20x_ksi_low_gcp` replaced by `fedramp_20x_ksi_2026` [(#12855)](https://github.com/prowler-cloud/prowler/pull/12855) + +### 🐞 Fixed + +- `security_2sv_enforced` reports domain-wide 2-Step Verification failures as FAIL even when every failing setting is overridden for a group or organizational unit [(#12700)](https://github.com/prowler-cloud/prowler/pull/12700) +- Bootstrap STS calls now try up to two more regions of the partition declared in `PROWLER_AWS_PARTITION` when the first one cannot be reached, so a deployment that routes to only one region of its partition no longer fails on an endpoint it has no path to. This covers validating credentials, assuming a role and getting an MFA session token [(#12799)](https://github.com/prowler-cloud/prowler/pull/12799) +- `KeyError` in M365 Defender malware, anti-phishing and inbound anti-spam checks when the tenant has Standard or Strict preset security policies [(#12809)](https://github.com/prowler-cloud/prowler/pull/12809) +- Azure Defender security contacts and Key Vault key rotation policies now use the endpoints of the selected cloud (`--azure-region`) instead of the hardcoded `management.azure.com` and `vault.azure.net` hosts, so both work on `AzureUSGovernment` and `AzureChinaCloud` [(#12813)](https://github.com/prowler-cloud/prowler/pull/12813) + +### 🔐 Security + +- `libsqlite3-0`, `gzip`, `perl-base`, `libssh2-1t64` and `libpcre2-8-0` upgraded in the SDK container image, patching nine high Debian CVEs [(#12804)](https://github.com/prowler-cloud/prowler/pull/12804) +- PowerShell from 7.5.9 to 7.5.11 in the SDK container image, bundling .NET runtime 9.0.20 and patching CVE-2026-62901 [(#12811)](https://github.com/prowler-cloud/prowler/pull/12811) +- Bumped `anyio` to 4.14.2 to resolve CVE-2026-63374 [(#12848)](https://github.com/prowler-cloud/prowler/pull/12848) + +--- + ## [5.42.0] (Prowler v5.42.0) ### 🚀 Added diff --git a/prowler/changelog.d/aws-inspector2-fips-checks.added.md b/prowler/changelog.d/aws-inspector2-fips-checks.added.md deleted file mode 100644 index 40c36f8062..0000000000 --- a/prowler/changelog.d/aws-inspector2-fips-checks.added.md +++ /dev/null @@ -1 +0,0 @@ -`inspector2_coverage_scan_status_active`, `inspector2_coverage_recently_scanned`, `inspector2_active_findings_no_known_exploited_vulnerabilities`, `inspector2_active_findings_kev_within_due_date`, `inspector2_active_findings_within_max_age`, `elbv2_listener_fips_tls_enabled` and `transfer_server_fips_security_policy_enabled` checks for AWS provider, covering FedRAMP 20x Class C vulnerability detection, CISA KEV remediation and FIPS cryptography rules; the KEV checks require `inspector2:BatchGetFindingDetails`, now in the Prowler additions policy diff --git a/prowler/changelog.d/aws-partition-bootstrap-falls-back-to-the-next-region.fixed.md b/prowler/changelog.d/aws-partition-bootstrap-falls-back-to-the-next-region.fixed.md deleted file mode 100644 index 4bd9c855d1..0000000000 --- a/prowler/changelog.d/aws-partition-bootstrap-falls-back-to-the-next-region.fixed.md +++ /dev/null @@ -1 +0,0 @@ -Bootstrap STS calls now try up to two more regions of the partition declared in `PROWLER_AWS_PARTITION` when the first one cannot be reached, so a deployment that routes to only one region of its partition no longer fails on an endpoint it has no path to. This covers validating credentials, assuming a role and getting an MFA session token diff --git a/prowler/changelog.d/aws-retries-max-attempts-env.added.md b/prowler/changelog.d/aws-retries-max-attempts-env.added.md new file mode 100644 index 0000000000..d88cc59507 --- /dev/null +++ b/prowler/changelog.d/aws-retries-max-attempts-env.added.md @@ -0,0 +1 @@ +`PROWLER_AWS_BOTO3_RETRIES_MAX_ATTEMPTS` environment variable to set the Boto3 retries for deployments without CLI flags diff --git a/prowler/changelog.d/aws-sts-reuse-answering-region.fixed.md b/prowler/changelog.d/aws-sts-reuse-answering-region.fixed.md new file mode 100644 index 0000000000..431426522a --- /dev/null +++ b/prowler/changelog.d/aws-sts-reuse-answering-region.fixed.md @@ -0,0 +1 @@ +STS calls after role assumption use the answering region, avoiding a second wait for an unreachable partition region diff --git a/prowler/changelog.d/azure-sovereign-cloud-defender-keyvault-hosts.fixed.md b/prowler/changelog.d/azure-sovereign-cloud-defender-keyvault-hosts.fixed.md deleted file mode 100644 index 99ef3354e0..0000000000 --- a/prowler/changelog.d/azure-sovereign-cloud-defender-keyvault-hosts.fixed.md +++ /dev/null @@ -1 +0,0 @@ -Azure Defender security contacts and Key Vault key rotation policies now use the endpoints of the selected cloud (`--azure-region`) instead of the hardcoded `management.azure.com` and `vault.azure.net` hosts, so both work on `AzureUSGovernment` and `AzureChinaCloud` diff --git a/prowler/changelog.d/fedramp-20x-frr-class-c-2026.added.md b/prowler/changelog.d/fedramp-20x-frr-class-c-2026.added.md deleted file mode 100644 index 147ad052f3..0000000000 --- a/prowler/changelog.d/fedramp-20x-frr-class-c-2026.added.md +++ /dev/null @@ -1 +0,0 @@ -`FedRAMP-20x-FRR-Class-C` universal compliance framework (`fedramp_20x_frr_class_c_2026`) with the 158 provider rules of the FedRAMP 20x Class C ruleset from the FedRAMP Consolidated Rules 2026 for AWS, Azure, GCP, Kubernetes and M365 diff --git a/prowler/changelog.d/fedramp-20x-ksi-2026.added.md b/prowler/changelog.d/fedramp-20x-ksi-2026.added.md deleted file mode 100644 index 5fae5bb4a2..0000000000 --- a/prowler/changelog.d/fedramp-20x-ksi-2026.added.md +++ /dev/null @@ -1 +0,0 @@ -`FedRAMP-20x-KSI` universal compliance framework (`fedramp_20x_ksi_2026`) with the 46 Key Security Indicators from the FedRAMP Consolidated Rules 2026 mapped for AWS, Azure, GCP, Kubernetes and M365 diff --git a/prowler/changelog.d/fedramp-20x-ksi-low-pilot.removed.md b/prowler/changelog.d/fedramp-20x-ksi-low-pilot.removed.md deleted file mode 100644 index fdca7d54e5..0000000000 --- a/prowler/changelog.d/fedramp-20x-ksi-low-pilot.removed.md +++ /dev/null @@ -1 +0,0 @@ -`fedramp_20x_ksi_low_aws`, `fedramp_20x_ksi_low_azure` and `fedramp_20x_ksi_low_gcp` FedRAMP 20x Phase One pilot frameworks, superseded by `fedramp_20x_ksi_2026` diff --git a/prowler/changelog.d/googleworkspace-2sv-all-users-overrides.fixed.md b/prowler/changelog.d/googleworkspace-2sv-all-users-overrides.fixed.md deleted file mode 100644 index 1283ac91ad..0000000000 --- a/prowler/changelog.d/googleworkspace-2sv-all-users-overrides.fixed.md +++ /dev/null @@ -1 +0,0 @@ -`security_2sv_enforced` reports domain-wide 2-Step Verification failures as FAIL even when every failing setting is overridden for a group or organizational unit diff --git a/prowler/changelog.d/m365-defender-preset-policy-keyerror.fixed.md b/prowler/changelog.d/m365-defender-preset-policy-keyerror.fixed.md deleted file mode 100644 index 2a50037e84..0000000000 --- a/prowler/changelog.d/m365-defender-preset-policy-keyerror.fixed.md +++ /dev/null @@ -1 +0,0 @@ -`KeyError` in M365 Defender malware, anti-phishing and inbound anti-spam checks when the tenant has Standard or Strict preset security policies diff --git a/prowler/changelog.d/sdk-image-debian-cves.security.md b/prowler/changelog.d/sdk-image-debian-cves.security.md deleted file mode 100644 index 50bd4186f3..0000000000 --- a/prowler/changelog.d/sdk-image-debian-cves.security.md +++ /dev/null @@ -1 +0,0 @@ -`libsqlite3-0`, `gzip`, `perl-base`, `libssh2-1t64` and `libpcre2-8-0` upgraded in the SDK container image, patching nine high Debian CVEs diff --git a/prowler/changelog.d/sdk-image-powershell-dotnet-cve.security.md b/prowler/changelog.d/sdk-image-powershell-dotnet-cve.security.md deleted file mode 100644 index 370aa37289..0000000000 --- a/prowler/changelog.d/sdk-image-powershell-dotnet-cve.security.md +++ /dev/null @@ -1 +0,0 @@ -PowerShell from 7.5.9 to 7.5.11 in the SDK container image, bundling .NET runtime 9.0.20 and patching CVE-2026-62901 diff --git a/prowler/changelog.d/smn-topic-subscriptions.added.md b/prowler/changelog.d/smn-topic-subscriptions.added.md deleted file mode 100644 index 87ca0c8c28..0000000000 --- a/prowler/changelog.d/smn-topic-subscriptions.added.md +++ /dev/null @@ -1 +0,0 @@ -`smn_topic_subscriptions` check for Huawei Cloud provider: SMN topics have at least one subscription configured diff --git a/prowler/changelog.d/ui-e2e-secrets-via-env.security.md b/prowler/changelog.d/ui-e2e-secrets-via-env.security.md new file mode 100644 index 0000000000..476f690e3c --- /dev/null +++ b/prowler/changelog.d/ui-e2e-secrets-via-env.security.md @@ -0,0 +1 @@ +Pass the E2E AWS credentials to the UI E2E workflow through environment variables instead of template expansion diff --git a/prowler/config/config.py b/prowler/config/config.py index 21929b5099..928f7fd520 100644 --- a/prowler/config/config.py +++ b/prowler/config/config.py @@ -52,7 +52,7 @@ class _MutableTimestamp: timestamp = _MutableTimestamp(datetime.today()) timestamp_utc = _MutableTimestamp(datetime.now(timezone.utc)) -prowler_version = "5.43.0" +prowler_version = "5.44.0" html_logo_url = "https://github.com/prowler-cloud/prowler/" square_logo_img = "https://raw.githubusercontent.com/prowler-cloud/prowler/dc7d2d5aeb92fdf12e8604f42ef6472cd3e8e889/docs/img/prowler-logo-black.png" aws_logo = "https://user-images.githubusercontent.com/38561120/235953920-3e3fba08-0795-41dc-b480-9bea57db9f2e.png" diff --git a/prowler/providers/aws/aws_provider.py b/prowler/providers/aws/aws_provider.py index 784daf0caa..37067bb516 100644 --- a/prowler/providers/aws/aws_provider.py +++ b/prowler/providers/aws/aws_provider.py @@ -112,7 +112,7 @@ class AwsProvider(Provider): def __init__( self, - retries_max_attempts: int = 3, + retries_max_attempts: Optional[int] = None, role_arn: str = None, session_duration: int = 3600, external_id: str = None, @@ -141,6 +141,7 @@ class AwsProvider(Provider): Args: - retries_max_attempts: The maximum number of retries for the AWS client. + Defaults to the PROWLER_AWS_BOTO3_RETRIES_MAX_ATTEMPTS environment variable or, if unset, to 3. - role_arn: The ARN of the IAM role to assume. - session_duration: The duration of the session in seconds, between 900 and 43200. - external_id: The external ID to use when assuming the IAM role. @@ -1230,7 +1231,8 @@ class AwsProvider(Provider): Args: - session: The AWS session object - - assumed_role_info: The AWSAssumeRoleInfo object + - assumed_role_info: The AWSAssumeRoleInfo object. Its sts_region is + updated to the region that answered, so later calls go straight there Returns: - AWSCredentials: The AWS credentials for the assumed role @@ -1256,11 +1258,14 @@ class AwsProvider(Provider): mfa_info = AwsProvider.input_role_mfa_token_and_code() assume_role_arguments["SerialNumber"] = mfa_info.arn assume_role_arguments["TokenCode"] = mfa_info.totp - _, assumed_credentials = AwsProvider.sts_call_with_partition_failover( - session, - assumed_role_info.sts_region, - lambda sts_client: sts_client.assume_role(**assume_role_arguments), + sts_region, assumed_credentials = ( + AwsProvider.sts_call_with_partition_failover( + session, + assumed_role_info.sts_region, + lambda sts_client: sts_client.assume_role(**assume_role_arguments), + ) ) + assumed_role_info.sts_region = sts_region # Convert the UTC datetime object to your local timezone credentials_expiration_local_time = ( assumed_credentials["Credentials"]["Expiration"] @@ -1558,6 +1563,8 @@ class AwsProvider(Provider): session, assumed_role_information, ) + # Validate where the role was assumed, not where it timed out + aws_region = assumed_role_information.sts_region session = Session( aws_access_key_id=assumed_role_credentials.aws_access_key_id, aws_secret_access_key=assumed_role_credentials.aws_secret_access_key, diff --git a/prowler/providers/aws/aws_regions_by_service.json b/prowler/providers/aws/aws_regions_by_service.json index bb51556b78..d7620009bf 100644 --- a/prowler/providers/aws/aws_regions_by_service.json +++ b/prowler/providers/aws/aws_regions_by_service.json @@ -122,6 +122,21 @@ "aws-us-gov": [] } }, + "account-access": { + "regions": { + "aws": [ + "ap-southeast-1", + "us-west-1" + ], + "aws-cn": [], + "aws-eusc": [], + "aws-iso": [], + "aws-iso-b": [], + "aws-iso-e": [], + "aws-iso-f": [], + "aws-us-gov": [] + } + }, "acm": { "regions": { "aws": [ @@ -242,6 +257,24 @@ ] } }, + "agent-registry": { + "regions": { + "aws": [ + "ap-northeast-1", + "ap-southeast-2", + "eu-west-1", + "us-east-1", + "us-west-2" + ], + "aws-cn": [], + "aws-eusc": [], + "aws-iso": [], + "aws-iso-b": [], + "aws-iso-e": [], + "aws-iso-f": [], + "aws-us-gov": [] + } + }, "ahl": { "regions": { "aws": [ @@ -266,14 +299,20 @@ "aiops": { "regions": { "aws": [ + "ap-east-1", "ap-northeast-1", "ap-south-1", "ap-southeast-1", "ap-southeast-2", "ap-southeast-5", "ap-southeast-7", + "eu-central-1", "eu-north-1", - "eu-south-2" + "eu-south-2", + "eu-west-1", + "us-east-1", + "us-east-2", + "us-west-2" ], "aws-cn": [], "aws-eusc": [], @@ -973,12 +1012,15 @@ "aws": [ "ap-northeast-1", "ap-northeast-2", + "ap-northeast-3", "ap-south-1", "ap-southeast-1", "ap-southeast-2", "ap-southeast-5", "ca-central-1", + "ca-west-1", "eu-central-1", + "eu-central-2", "eu-south-1", "eu-south-2", "eu-west-1", @@ -1545,37 +1587,11 @@ ] } }, - "awstransform": { - "regions": { - "aws": [ - "ap-northeast-1", - "ap-northeast-2", - "ap-south-1", - "ap-southeast-2", - "ca-central-1", - "eu-central-1", - "eu-west-2", - "sa-east-1", - "us-east-1" - ], - "aws-cn": [], - "aws-eusc": [], - "aws-iso": [], - "aws-iso-b": [], - "aws-iso-e": [], - "aws-iso-f": [], - "aws-us-gov": [] - } - }, "b2bi": { "regions": { "aws": [ - "ap-south-2", - "ap-southeast-2", "ca-central-1", - "eu-central-1", "eu-west-1", - "eu-west-3", "us-east-1", "us-east-2", "us-west-2" @@ -1889,6 +1905,7 @@ "ap-northeast-1", "ap-northeast-2", "ap-south-1", + "ap-south-2", "ap-southeast-1", "ap-southeast-2", "ap-southeast-5", @@ -1904,6 +1921,7 @@ "sa-east-1", "us-east-1", "us-east-2", + "us-west-1", "us-west-2" ], "aws-cn": [], @@ -4202,16 +4220,33 @@ "dax": { "regions": { "aws": [ + "af-south-1", + "ap-east-1", + "ap-east-2", "ap-northeast-1", + "ap-northeast-2", + "ap-northeast-3", "ap-south-1", + "ap-south-2", "ap-southeast-1", "ap-southeast-2", + "ap-southeast-3", + "ap-southeast-4", + "ap-southeast-5", + "ap-southeast-6", + "ap-southeast-7", + "ca-central-1", + "ca-west-1", "eu-central-1", + "eu-central-2", "eu-north-1", + "eu-south-1", "eu-south-2", "eu-west-1", "eu-west-2", "eu-west-3", + "il-central-1", + "mx-central-1", "sa-east-1", "us-east-1", "us-east-2", @@ -6006,8 +6041,10 @@ "evs": { "regions": { "aws": [ + "ap-east-2", "ap-northeast-1", "ap-northeast-2", + "ap-northeast-3", "ap-south-1", "ap-south-2", "ap-southeast-1", @@ -6019,9 +6056,11 @@ "eu-central-2", "eu-north-1", "eu-south-1", + "eu-south-2", "eu-west-1", "eu-west-2", "eu-west-3", + "il-central-1", "mx-central-1", "sa-east-1", "us-east-1", @@ -9218,9 +9257,12 @@ "aws": [ "af-south-1", "ap-east-1", + "ap-east-2", "ap-northeast-1", "ap-northeast-2", + "ap-northeast-3", "ap-south-1", + "ap-south-2", "ap-southeast-1", "ap-southeast-2", "ap-southeast-5", @@ -9228,12 +9270,17 @@ "ca-central-1", "ca-west-1", "eu-central-1", + "eu-central-2", "eu-north-1", + "eu-south-1", + "eu-south-2", "eu-west-1", "eu-west-2", "eu-west-3", + "il-central-1", "me-central-1", "me-south-1", + "mx-central-1", "sa-east-1", "us-east-1", "us-east-2", @@ -9965,7 +10012,9 @@ "aws-iso-b": [], "aws-iso-e": [], "aws-iso-f": [], - "aws-us-gov": [] + "aws-us-gov": [ + "us-gov-west-1" + ] } }, "neptune": { @@ -10427,7 +10476,9 @@ "us-west-2" ], "aws-cn": [], - "aws-eusc": [], + "aws-eusc": [ + "eusc-de-east-1" + ], "aws-iso": [], "aws-iso-b": [], "aws-iso-e": [], @@ -10902,7 +10953,10 @@ "aws-iso-b": [], "aws-iso-e": [], "aws-iso-f": [], - "aws-us-gov": [] + "aws-us-gov": [ + "us-gov-east-1", + "us-gov-west-1" + ] } }, "pca-connector-scep": { @@ -10964,6 +11018,8 @@ "ap-southeast-1", "ap-southeast-2", "ap-southeast-3", + "ap-southeast-4", + "ap-southeast-6", "eu-central-1", "eu-north-1", "eu-south-1", @@ -13414,6 +13470,21 @@ "aws-us-gov": [] } }, + "securityagent": { + "regions": { + "aws": [ + "us-east-1", + "us-west-2" + ], + "aws-cn": [], + "aws-eusc": [], + "aws-iso": [], + "aws-iso-b": [], + "aws-iso-e": [], + "aws-iso-f": [], + "aws-us-gov": [] + } + }, "securityhub": { "regions": { "aws": [ @@ -13877,6 +13948,7 @@ "ap-southeast-3", "ap-southeast-4", "ap-southeast-5", + "ap-southeast-6", "ap-southeast-7", "ca-central-1", "ca-west-1", @@ -14839,6 +14911,18 @@ "aws-us-gov": [] } }, + "supportauthz": { + "regions": { + "aws": [], + "aws-cn": [], + "aws-eusc": [], + "aws-iso": [], + "aws-iso-b": [], + "aws-iso-e": [], + "aws-iso-f": [], + "aws-us-gov": [] + } + }, "sustainability": { "regions": { "aws": [ @@ -15062,8 +15146,10 @@ "aws": [ "ap-northeast-1", "ap-south-1", + "ap-southeast-1", "ap-southeast-2", "eu-central-1", + "eu-north-1", "eu-west-1", "us-east-1", "us-east-2", @@ -15083,20 +15169,28 @@ "timestream-influxdb": { "regions": { "aws": [ + "af-south-1", + "ap-east-1", "ap-northeast-1", + "ap-northeast-2", "ap-northeast-3", "ap-south-1", + "ap-south-2", "ap-southeast-1", "ap-southeast-2", "ap-southeast-3", + "ap-southeast-4", + "ap-southeast-7", "ca-central-1", "eu-central-1", + "eu-central-2", "eu-north-1", "eu-south-1", "eu-south-2", "eu-west-1", "eu-west-2", "eu-west-3", + "il-central-1", "me-central-1", "mx-central-1", "sa-east-1", @@ -15162,29 +15256,6 @@ ] } }, - "tnb": { - "regions": { - "aws": [ - "ap-northeast-2", - "ap-southeast-2", - "ca-central-1", - "eu-central-1", - "eu-north-1", - "eu-south-2", - "eu-west-3", - "sa-east-1", - "us-east-1", - "us-west-2" - ], - "aws-cn": [], - "aws-eusc": [], - "aws-iso": [], - "aws-iso-b": [], - "aws-iso-e": [], - "aws-iso-f": [], - "aws-us-gov": [] - } - }, "transcribe": { "regions": { "aws": [ @@ -15437,7 +15508,38 @@ "uxc": { "regions": { "aws": [ - "us-east-1" + "af-south-1", + "ap-east-1", + "ap-east-2", + "ap-northeast-1", + "ap-northeast-2", + "ap-northeast-3", + "ap-south-1", + "ap-south-2", + "ap-southeast-1", + "ap-southeast-2", + "ap-southeast-3", + "ap-southeast-4", + "ap-southeast-5", + "ap-southeast-6", + "ap-southeast-7", + "ca-central-1", + "ca-west-1", + "eu-central-1", + "eu-central-2", + "eu-north-1", + "eu-south-1", + "eu-south-2", + "eu-west-1", + "eu-west-2", + "eu-west-3", + "il-central-1", + "mx-central-1", + "sa-east-1", + "us-east-1", + "us-east-2", + "us-west-1", + "us-west-2" ], "aws-cn": [], "aws-eusc": [], @@ -16123,6 +16225,7 @@ "aws": [ "ap-east-1", "ap-northeast-2", + "ap-southeast-1", "ap-southeast-5", "eu-south-2", "me-central-1", diff --git a/prowler/providers/aws/config.py b/prowler/providers/aws/config.py index ed2ca503d0..dd365bb8cf 100644 --- a/prowler/providers/aws/config.py +++ b/prowler/providers/aws/config.py @@ -2,7 +2,10 @@ import os from botocore.config import Config -from prowler.providers.aws.exceptions.exceptions import AWSInvalidBoto3TimeoutError +from prowler.providers.aws.exceptions.exceptions import ( + AWSInvalidBoto3RetriesError, + AWSInvalidBoto3TimeoutError, +) AWS_STS_GLOBAL_ENDPOINT_REGION = "us-east-1" AWS_REGION_US_EAST_1 = "us-east-1" @@ -27,10 +30,28 @@ def get_boto3_timeout_from_env(name: str, default: int) -> int: return int(raw) +def get_boto3_retries_from_env(name: str, default: int) -> int: + """Non-negative integer retries read from the environment, or default when unset.""" + raw = os.getenv(name, "").strip() + if not raw: + return default + if not raw.isdecimal(): + raise AWSInvalidBoto3RetriesError( + file=os.path.basename(__file__), + message=f"{name} must be a non-negative integer number of retries, got {raw!r}", + ) + return int(raw) + + def get_default_session_config() -> Config: return Config( user_agent_extra=BOTO3_USER_AGENT_EXTRA, - retries={"max_attempts": BOTO3_RETRIES_MAX_ATTEMPTS, "mode": "standard"}, + retries={ + "max_attempts": get_boto3_retries_from_env( + "PROWLER_AWS_BOTO3_RETRIES_MAX_ATTEMPTS", BOTO3_RETRIES_MAX_ATTEMPTS + ), + "mode": "standard", + }, connect_timeout=get_boto3_timeout_from_env( "PROWLER_AWS_BOTO3_CONNECT_TIMEOUT", BOTO3_CONNECT_TIMEOUT ), diff --git a/prowler/providers/aws/exceptions/exceptions.py b/prowler/providers/aws/exceptions/exceptions.py index 089e0d99c7..dae23408e9 100644 --- a/prowler/providers/aws/exceptions/exceptions.py +++ b/prowler/providers/aws/exceptions/exceptions.py @@ -82,6 +82,10 @@ class AWSBaseException(ProwlerException): "message": "The Boto3 timeout configured through the environment is invalid", "remediation": "Set PROWLER_AWS_BOTO3_CONNECT_TIMEOUT and PROWLER_AWS_BOTO3_READ_TIMEOUT to a positive integer number of seconds.", }, + (1919, "AWSInvalidBoto3RetriesError"): { + "message": "The Boto3 retries configured through the environment are invalid", + "remediation": "Set PROWLER_AWS_BOTO3_RETRIES_MAX_ATTEMPTS to a non-negative integer, 0 disables retries.", + }, } def __init__(self, code, file=None, original_exception=None, message=None): @@ -244,3 +248,12 @@ class AWSInvalidBoto3TimeoutError(AWSBaseException): super().__init__( 1918, file=file, original_exception=original_exception, message=message ) + + +class AWSInvalidBoto3RetriesError(AWSBaseException): + """Boto3 retries configured through the environment are not a non-negative integer.""" + + def __init__(self, file=None, original_exception=None, message=None): + super().__init__( + 1919, file=file, original_exception=original_exception, message=message + ) diff --git a/prowler/providers/aws/lib/s3/s3.py b/prowler/providers/aws/lib/s3/s3.py index a4bbb42cc7..d2f511fd09 100644 --- a/prowler/providers/aws/lib/s3/s3.py +++ b/prowler/providers/aws/lib/s3/s3.py @@ -85,7 +85,7 @@ class S3: aws_access_key_id: str = None, aws_secret_access_key: str = None, aws_session_token: Optional[str] = None, - retries_max_attempts: int = 3, + retries_max_attempts: Optional[int] = None, regions: set = set(), ) -> None: """ diff --git a/prowler/providers/aws/lib/security_hub/security_hub.py b/prowler/providers/aws/lib/security_hub/security_hub.py index bc372d1ddd..ae4c2350e4 100644 --- a/prowler/providers/aws/lib/security_hub/security_hub.py +++ b/prowler/providers/aws/lib/security_hub/security_hub.py @@ -106,7 +106,7 @@ class SecurityHub: aws_access_key_id: str = None, aws_secret_access_key: str = None, aws_session_token: Optional[str] = None, - retries_max_attempts: int = 3, + retries_max_attempts: Optional[int] = None, regions: set = set(), ) -> "SecurityHub": """ diff --git a/prowler/providers/aws/lib/session/aws_set_up_session.py b/prowler/providers/aws/lib/session/aws_set_up_session.py index 8f0b4130ca..6b00bb437b 100644 --- a/prowler/providers/aws/lib/session/aws_set_up_session.py +++ b/prowler/providers/aws/lib/session/aws_set_up_session.py @@ -40,7 +40,7 @@ class AwsSetUpSession: aws_access_key_id: str = None, aws_secret_access_key: str = None, aws_session_token: Optional[str] = None, - retries_max_attempts: int = 3, + retries_max_attempts: Optional[int] = None, regions: set = set(), connect_timeout: Optional[int] = None, read_timeout: Optional[int] = None, @@ -106,6 +106,8 @@ class AwsSetUpSession: session=self._session.current_session, aws_region=sts_region, ) + # Later STS calls go where validation got an answer, not where it timed out + sts_region = caller_identity.region logger.info("Credentials validated") ######## diff --git a/pyproject.toml b/pyproject.toml index a05fa6a6ee..ac90a11d60 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -144,7 +144,7 @@ maintainers = [{name = "Prowler Engineering", email = "engineering@prowler.com"} name = "prowler" readme = "README.md" requires-python = ">=3.10,<3.14" -version = "5.43.0" +version = "5.44.0" [project.scripts] prowler = "prowler.__main__:prowler" @@ -214,7 +214,7 @@ constraint-dependencies = [ "aliyun-log-fastpb==0.3.0", "annotated-types==0.7.0", "antlr4-python3-runtime==4.13.2", - "anyio==4.13.0", + "anyio==4.14.2", "apscheduler==3.11.2", "astroid==3.3.11", "async-timeout==5.0.1", diff --git a/tests/providers/aws/aws_provider_test.py b/tests/providers/aws/aws_provider_test.py index 2d759f362a..2cf9d828b5 100644 --- a/tests/providers/aws/aws_provider_test.py +++ b/tests/providers/aws/aws_provider_test.py @@ -28,15 +28,19 @@ from prowler.providers.aws.config import ( AWS_STS_GLOBAL_ENDPOINT_REGION, BOTO3_CONNECT_TIMEOUT, BOTO3_READ_TIMEOUT, + BOTO3_RETRIES_MAX_ATTEMPTS, BOTO3_USER_AGENT_EXTRA, ROLE_SESSION_NAME, + get_boto3_retries_from_env, get_boto3_timeout_from_env, get_default_session_config, ) from prowler.providers.aws.exceptions.exceptions import ( AWSAccessKeyIDInvalidError, AWSArgumentTypeValidationError, + AWSAssumeRoleError, AWSIAMRoleARNInvalidResourceTypeError, + AWSInvalidBoto3RetriesError, AWSInvalidBoto3TimeoutError, AWSInvalidPartitionError, AWSInvalidProviderIdError, @@ -1845,6 +1849,143 @@ aws: ] assert isinstance(credentials, AWSCredentials) assert credentials.aws_access_key_id == "AKIAIOSFODNN7EXAMPLE" + # Refreshing the credentials later goes straight to the region that answered + assert assumed_role_info.sts_region == AWS_REGION_GOV_CLOUD_US_WEST_1 + + def test_assume_role_does_not_retry_a_credential_error(self, monkeypatch): + monkeypatch.setenv("PROWLER_AWS_PARTITION", AWS_GOV_CLOUD_PARTITION) + current_session = session.Session(region_name=AWS_REGION_US_EAST_1) + attempted_regions = [] + + def create_sts_session(session, aws_region): + attempted_regions.append(aws_region) + sts_client = mock.MagicMock() + sts_client.assume_role.side_effect = botocore.exceptions.ClientError( + {"Error": {"Code": "AccessDenied", "Message": "denied"}}, + "AssumeRole", + ) + return sts_client + + assumed_role_info = AWSAssumeRoleInfo( + role_arn=ARN( + arn=f"arn:{AWS_GOV_CLOUD_PARTITION}:iam::{AWS_ACCOUNT_NUMBER}:role/test-role" + ), + session_duration=3600, + external_id=None, + mfa_enabled=False, + role_session_name=ROLE_SESSION_NAME, + sts_region=AWS_REGION_GOV_CLOUD_US_EAST_1, + ) + + with patch( + "prowler.providers.aws.aws_provider.AwsProvider.create_sts_session", + new=create_sts_session, + ): + with raises(AWSAssumeRoleError): + AwsProvider.assume_role(current_session, assumed_role_info) + + assert attempted_regions == [AWS_REGION_GOV_CLOUD_US_EAST_1] + assert assumed_role_info.sts_region == AWS_REGION_GOV_CLOUD_US_EAST_1 + + def test_test_connection_role_validates_where_the_role_was_assumed( + self, monkeypatch + ): + monkeypatch.setenv("PROWLER_AWS_PARTITION", AWS_GOV_CLOUD_PARTITION) + monkeypatch.delenv("AWS_DEFAULT_REGION", raising=False) + attempted_calls = [] + + def create_sts_session(session, aws_region): + if aws_region == AWS_REGION_GOV_CLOUD_US_EAST_1: + attempted_calls.append(aws_region) + raise botocore.exceptions.ConnectTimeoutError( + endpoint_url=f"https://sts.{aws_region}.amazonaws.com" + ) + sts_client = mock.MagicMock() + + def assume_role(**_): + attempted_calls.append(("AssumeRole", aws_region)) + return { + "Credentials": { + "AccessKeyId": "AKIAIOSFODNN7EXAMPLE", + "SecretAccessKey": "secret", + "SessionToken": "token", + "Expiration": datetime.now() + timedelta(seconds=3600), + } + } + + def get_caller_identity(): + attempted_calls.append(("GetCallerIdentity", aws_region)) + return { + "UserId": "test-user-id", + "Account": AWS_ACCOUNT_NUMBER, + "Arn": AWS_GOV_CLOUD_ACCOUNT_ARN, + } + + sts_client.assume_role.side_effect = assume_role + sts_client.get_caller_identity.side_effect = get_caller_identity + return sts_client + + with patch( + "prowler.providers.aws.aws_provider.AwsProvider.create_sts_session", + new=create_sts_session, + ): + connection = AwsProvider.test_connection( + role_arn=f"arn:{AWS_GOV_CLOUD_PARTITION}:iam::{AWS_ACCOUNT_NUMBER}:role/test-role", + aws_access_key_id="test-access-key", + aws_secret_access_key="test-secret-key", + raise_on_exception=False, + ) + + assert connection.is_connected + # The unreachable region is paid for once, not again for the validation + assert attempted_calls == [ + AWS_REGION_GOV_CLOUD_US_EAST_1, + ("AssumeRole", AWS_REGION_GOV_CLOUD_US_WEST_1), + ("GetCallerIdentity", AWS_REGION_GOV_CLOUD_US_WEST_1), + ] + + @mock_aws + def test_aws_set_up_session_assumes_the_role_where_validation_got_an_answer( + self, monkeypatch + ): + monkeypatch.setenv("PROWLER_AWS_PARTITION", AWS_GOV_CLOUD_PARTITION) + monkeypatch.setenv("AWS_DEFAULT_REGION", AWS_REGION_US_EAST_1) + answered = AWSCallerIdentity( + user_id="test-user-id", + account=AWS_ACCOUNT_NUMBER, + arn=ARN(AWS_GOV_CLOUD_ACCOUNT_ARN), + region=AWS_REGION_GOV_CLOUD_US_WEST_1, + ) + sts_regions = [] + + class RoleAssumed(Exception): + pass + + def assume_role(session, assumed_role_info): + sts_regions.append(assumed_role_info.sts_region) + raise RoleAssumed + + with ( + patch( + "prowler.providers.aws.aws_provider.AwsProvider.validate_credentials", + return_value=answered, + ), + patch( + "prowler.providers.aws.aws_provider.AwsProvider.assume_role", + side_effect=assume_role, + ), + ): + with raises(RoleAssumed): + AwsSetUpSession( + role_arn=f"arn:{AWS_GOV_CLOUD_PARTITION}:iam::{AWS_ACCOUNT_NUMBER}:role/test-role", + session_duration=900, + external_id="test-external-id", + role_session_name=ROLE_SESSION_NAME, + aws_access_key_id="testing", + aws_secret_access_key="testing", + ) + + assert sts_regions == [AWS_REGION_GOV_CLOUD_US_WEST_1] def test_setup_session_mfa_falls_back_to_the_next_partition_region( self, monkeypatch @@ -3352,6 +3493,123 @@ aws: ): get_boto3_timeout_from_env("PROWLER_AWS_BOTO3_CONNECT_TIMEOUT", 10) + def test_get_default_session_config_retries_from_env(self): + with mock.patch.dict( + os.environ, {"PROWLER_AWS_BOTO3_RETRIES_MAX_ATTEMPTS": "1"} + ): + config = get_default_session_config() + + assert config.retries == {"max_attempts": 1, "mode": "standard"} + + def test_get_default_session_config_retries_from_env_0_disables_retries(self): + with mock.patch.dict( + os.environ, {"PROWLER_AWS_BOTO3_RETRIES_MAX_ATTEMPTS": "0"} + ): + config = get_default_session_config() + + assert config.retries == {"max_attempts": 0, "mode": "standard"} + + def test_set_session_config_argument_overrides_env_retries(self): + with mock.patch.dict( + os.environ, {"PROWLER_AWS_BOTO3_RETRIES_MAX_ATTEMPTS": "1"} + ): + config = AwsProvider.set_session_config(5) + + assert config.retries == {"max_attempts": 5, "mode": "standard"} + + @mock_aws + def test_aws_provider_without_retries_argument_uses_env_retries(self): + with mock.patch.dict( + os.environ, {"PROWLER_AWS_BOTO3_RETRIES_MAX_ATTEMPTS": "0"} + ): + aws_provider = AwsProvider() + client = aws_provider.session.current_session.client( + "ec2", region_name=AWS_REGION_US_EAST_1 + ) + + # botocore rewrites max_attempts into total_max_attempts (retries + 1) + assert client.meta.config.retries["total_max_attempts"] == 1 + + @mock_aws + def test_aws_provider_retries_argument_overrides_env_retries(self): + with mock.patch.dict( + os.environ, {"PROWLER_AWS_BOTO3_RETRIES_MAX_ATTEMPTS": "0"} + ): + aws_provider = AwsProvider(retries_max_attempts=7) + client = aws_provider.session.current_session.client( + "ec2", region_name=AWS_REGION_US_EAST_1 + ) + + assert client.meta.config.retries["total_max_attempts"] == 8 + + @mock_aws + def test_aws_set_up_session_without_retries_argument_uses_env_retries(self): + with mock.patch.dict( + os.environ, {"PROWLER_AWS_BOTO3_RETRIES_MAX_ATTEMPTS": "1"} + ): + aws_session = AwsSetUpSession( + aws_access_key_id="testing", + aws_secret_access_key="testing", + ) + client = aws_session._session.current_session.client( + "ec2", region_name=AWS_REGION_US_EAST_1 + ) + + assert client.meta.config.retries["total_max_attempts"] == 2 + + def test_test_connection_session_uses_env_retries(self): + with ( + mock.patch.dict( + os.environ, {"PROWLER_AWS_BOTO3_RETRIES_MAX_ATTEMPTS": "0"} + ), + mock.patch.object( + AwsProvider, + "validate_credentials", + return_value=AWSCallerIdentity( + user_id="test-user-id", + account=AWS_ACCOUNT_NUMBER, + arn=ARN(AWS_ACCOUNT_ARN), + region=AWS_REGION_US_EAST_1, + ), + ) as mock_validate_credentials, + ): + connection = AwsProvider.test_connection( + aws_access_key_id="test-access-key", + aws_secret_access_key="test-secret-key", + raise_on_exception=False, + ) + + assert connection.is_connected + validated_session = mock_validate_credentials.call_args.args[0] + assert validated_session._session.get_default_client_config().retries == { + "max_attempts": 0, + "mode": "standard", + } + + @pytest.mark.parametrize("raw", ["-1", "three", "1.5"]) + def test_get_boto3_retries_from_env_rejects_anything_but_non_negative_integers( + self, raw + ): + with mock.patch.dict( + os.environ, {"PROWLER_AWS_BOTO3_RETRIES_MAX_ATTEMPTS": raw} + ): + with raises( + AWSInvalidBoto3RetriesError, + match="PROWLER_AWS_BOTO3_RETRIES_MAX_ATTEMPTS", + ): + get_boto3_retries_from_env("PROWLER_AWS_BOTO3_RETRIES_MAX_ATTEMPTS", 3) + + def test_get_boto3_retries_from_env_blank_falls_back_to_default(self): + with mock.patch.dict( + os.environ, {"PROWLER_AWS_BOTO3_RETRIES_MAX_ATTEMPTS": " "} + ): + assert ( + get_boto3_retries_from_env( + "PROWLER_AWS_BOTO3_RETRIES_MAX_ATTEMPTS", BOTO3_RETRIES_MAX_ATTEMPTS + ) + == BOTO3_RETRIES_MAX_ATTEMPTS + ) + def test_get_boto3_timeout_from_env_blank_falls_back_to_default(self): with mock.patch.dict(os.environ, {"PROWLER_AWS_BOTO3_CONNECT_TIMEOUT": " "}): assert ( diff --git a/ui/CHANGELOG.md b/ui/CHANGELOG.md index 8687ddde54..9ea71d69e9 100644 --- a/ui/CHANGELOG.md +++ b/ui/CHANGELOG.md @@ -4,6 +4,24 @@ All notable changes to the **Prowler UI** are documented in this file. +## [1.43.0] (Prowler v5.43.0) + +### 🚀 Added + +- Registry marketplace and external provider onboarding for Private Cloud, with permission-based access independent of billing, confirmed artifact installation, schema-driven credentials, connection checks, and scan launch [(#12494)](https://github.com/prowler-cloud/prowler/pull/12494) +- AWS Marketplace button variant with outlined styling for light and dark themes [(#12803)](https://github.com/prowler-cloud/prowler/pull/12803) +- `UI_SELF_REGISTRATION_ENABLED` flag for Prowler Private Cloud deployments; when `"false"`, `/sign-up` only opens with an invitation, the sign-in page drops the "Sign up" link and the profile hides **Create organization** [(#12815)](https://github.com/prowler-cloud/prowler/pull/12815) +- "Invite your team" step offered once after the first provider is connected, before the onboarding checkpoint, reusing the invitation form tagged with `source=onboarding` [(#12819)](https://github.com/prowler-cloud/prowler/pull/12819) + +### 🐞 Fixed + +- Automatic onboarding stays hidden on billing pages and remains available after leaving billing [(#12803)](https://github.com/prowler-cloud/prowler/pull/12803) +- Per-provider breakdown and OCSF download for FedRAMP 20x KSI and Class C FRR in the cross-provider compliance view [(#12810)](https://github.com/prowler-cloud/prowler/pull/12810) +- Edit and Revoke actions are disabled for expired and revoked invitations [(#12831)](https://github.com/prowler-cloud/prowler/pull/12831) +- Cloudflare API token links in the provider wizard request the SSL and Certificates, Bot Management and Zone WAF read permissions the scan needs [(#12842)](https://github.com/prowler-cloud/prowler/pull/12842) + +--- + ## [1.42.0] (Prowler v5.42.0) ### 🚀 Added diff --git a/ui/__tests__/msw/handlers/organizations.ts b/ui/__tests__/msw/handlers/organizations.ts index 2755e9281c..d1750a9fd5 100644 --- a/ui/__tests__/msw/handlers/organizations.ts +++ b/ui/__tests__/msw/handlers/organizations.ts @@ -302,6 +302,7 @@ export const handlersForOrganizations = ( organizations.map((o) => o.secretId).filter((id): id is string => !!id), ); let orgSeq = 0; + let providerSeq = 0; let secretSeq = 0; /** Reads per connection task, so `executingPolls` can hold one task running. */ const connectionTaskReads = new Map(); @@ -565,6 +566,37 @@ export const handlersForOrganizations = ( HttpResponse.json({ data: [], meta: collectionMeta(0) }), ), + // --- single-account connect (AWS one-step form) ----------------------- + http.post(`${API}/providers`, async ({ request }) => { + const body = (await request.json()) as { + data: { attributes: { provider: string; uid: string; alias?: string } }; + }; + providerSeq += 1; + return HttpResponse.json( + { + data: { + id: `provider-created-${providerSeq}`, + type: "providers", + attributes: { + ...body.data.attributes, + connection: { connected: false, last_checked_at: null }, + }, + }, + }, + { status: 201 }, + ); + }), + + http.post(`${API}/providers/secrets`, () => { + secretSeq += 1; + return HttpResponse.json( + { + data: { id: `secret-created-${secretSeq}`, type: "provider-secrets" }, + }, + { status: 201 }, + ); + }), + // --- providers (uid resolution) + connection testing ----------------- http.get<{ id: string }>(`${API}/providers/:id`, ({ params }) => { const provider = fx.providers.find((p) => p.id === params.id); diff --git a/ui/actions/findings/findings-by-resource.adapter.test.ts b/ui/actions/findings/findings-by-resource.adapter.test.ts index 24be4cd6ed..2c456a78a1 100644 --- a/ui/actions/findings/findings-by-resource.adapter.test.ts +++ b/ui/actions/findings/findings-by-resource.adapter.test.ts @@ -215,3 +215,72 @@ describe("adaptFindingsByResourceResponse — malformed input", () => { ); }); }); + +describe("adaptFindingsByResourceResponse — provider id", () => { + beforeEach(() => { + vi.clearAllMocks(); + }); + + it("should carry the provider id resolved through the scan include", () => { + // Given — scan.provider include path, as the drawer requests it + createDictMock.mockImplementation((type: string) => { + if (type === "scans") { + return { + "scan-1": { + id: "scan-1", + attributes: {}, + relationships: { provider: { data: { id: "provider-1" } } }, + }, + }; + } + if (type === "providers") { + return { + "provider-1": { + id: "provider-1", + attributes: { provider: "aws", alias: "prod", uid: "123" }, + }, + }; + } + return {}; + }); + + const input = { + data: { + id: "finding-1", + attributes: { + uid: "uid-1", + check_id: "s3_check", + status: "FAIL", + severity: "high", + check_metadata: {}, + }, + relationships: { + resources: { data: [] }, + scan: { data: { id: "scan-1" } }, + }, + }, + included: [], + }; + + // When + const [finding] = adaptFindingsByResourceResponse(input); + + // Then — the partial-scan request needs the id, not only the uid + expect(finding.providerId).toBe("provider-1"); + expect(finding.providerUid).toBe("123"); + }); + + it("should leave the provider id empty when the scan is not included", () => { + createDictMock.mockReturnValue({}); + + const [finding] = adaptFindingsByResourceResponse({ + data: { + id: "finding-1", + attributes: { uid: "uid-1", check_id: "s3_check", status: "FAIL" }, + relationships: { resources: { data: [] }, scan: { data: null } }, + }, + }); + + expect(finding.providerId).toBe(""); + }); +}); diff --git a/ui/actions/findings/findings-by-resource.adapter.ts b/ui/actions/findings/findings-by-resource.adapter.ts index 3ebae3c6c5..e4cd21abb1 100644 --- a/ui/actions/findings/findings-by-resource.adapter.ts +++ b/ui/actions/findings/findings-by-resource.adapter.ts @@ -64,6 +64,7 @@ export interface ResourceDrawerFinding { resourceDetails: string | null; resourceMetadata: Record | string | null; // Provider + providerId: string; providerType: ProviderType; providerAlias: string; providerUid: string; @@ -280,6 +281,7 @@ export function adaptFindingsByResourceResponse( | null | undefined) ?? null, // Provider + providerId: providerRelId ?? "", providerType: ((providerAttrs.provider as string | undefined) || "aws") as ProviderType, providerAlias: (providerAttrs.alias as string | undefined) || "", diff --git a/ui/actions/mute-rules/mute-rules.test.ts b/ui/actions/mute-rules/mute-rules.test.ts new file mode 100644 index 0000000000..cfb7f5daf9 --- /dev/null +++ b/ui/actions/mute-rules/mute-rules.test.ts @@ -0,0 +1,87 @@ +import { beforeEach, describe, expect, it, vi } from "vitest"; + +const { fetchMock, getAuthHeadersMock, revalidatePathMock } = vi.hoisted( + () => ({ + fetchMock: vi.fn(), + getAuthHeadersMock: vi.fn(), + revalidatePathMock: vi.fn(), + }), +); + +vi.mock("@/lib/helper", () => ({ + apiBaseUrl: "https://api.test/api/v1", + getAuthHeaders: getAuthHeadersMock, +})); + +vi.mock("next/cache", () => ({ + revalidatePath: revalidatePathMock, +})); + +import { createMuteRule } from "./mute-rules"; + +const NAME_CONFLICT_DETAIL = "A mute rule with this name already exists."; + +const errorResponse = (contentType: string, body: string, status = 400) => + new Response(body, { + status, + headers: { "Content-Type": contentType }, + }); + +const nameConflictBody = JSON.stringify({ + errors: [ + { + detail: NAME_CONFLICT_DETAIL, + status: "400", + source: { pointer: "/data/attributes/name" }, + code: "invalid", + }, + ], +}); + +const muteRuleFormData = () => { + const formData = new FormData(); + formData.set("name", "Root account has a hardware MFA device enabled"); + formData.set("reason", "Not our approach here with SSO"); + formData.set("finding_ids", JSON.stringify(["finding-1"])); + return formData; +}; + +beforeEach(() => { + vi.clearAllMocks(); + vi.stubGlobal("fetch", fetchMock); + vi.spyOn(console, "error").mockImplementation(() => {}); + getAuthHeadersMock.mockResolvedValue({ Authorization: "Bearer token" }); +}); + +describe("createMuteRule", () => { + it("should return only the error detail for a JSON:API error response", async () => { + fetchMock.mockResolvedValue( + errorResponse("application/vnd.api+json", nameConflictBody), + ); + + const result = await createMuteRule(null, muteRuleFormData()); + + expect(result?.errors?.general).toBe(NAME_CONFLICT_DETAIL); + expect(revalidatePathMock).not.toHaveBeenCalled(); + }); + + it("should return only the error detail for a plain JSON error response", async () => { + fetchMock.mockResolvedValue( + errorResponse("application/json", nameConflictBody), + ); + + const result = await createMuteRule(null, muteRuleFormData()); + + expect(result?.errors?.general).toBe(NAME_CONFLICT_DETAIL); + }); + + it("should return the response text for a non-JSON error response", async () => { + fetchMock.mockResolvedValue( + errorResponse("text/plain", "Bad gateway", 502), + ); + + const result = await createMuteRule(null, muteRuleFormData()); + + expect(result?.errors?.general).toBe("Bad gateway"); + }); +}); diff --git a/ui/actions/mute-rules/mute-rules.ts b/ui/actions/mute-rules/mute-rules.ts index c74c4a7e2a..1ce9a8b3bc 100644 --- a/ui/actions/mute-rules/mute-rules.ts +++ b/ui/actions/mute-rules/mute-rules.ts @@ -156,7 +156,8 @@ export const createMuteRule = async ( let errorMessage = `Failed to create mute rule: ${response.statusText}`; const responseContentType = response.headers.get("content-type"); try { - if (responseContentType?.includes("application/json")) { + // The API answers with application/vnd.api+json + if (responseContentType?.includes("json")) { const errorData = await response.json(); const jsonApiError = ( errorData as { diff --git a/ui/actions/organizations/organizations.ts b/ui/actions/organizations/organizations.ts index 69509f0376..761c08acfb 100644 --- a/ui/actions/organizations/organizations.ts +++ b/ui/actions/organizations/organizations.ts @@ -546,7 +546,7 @@ export const applyDiscovery = async ( ); // No `include`: the apply view rejects the parameter outright and fails the // whole request. The created providers' uids are read afterwards instead, with - // `getProviderUidsByIds`. + // `getProviderUidsAndConnectionBaselines`. const attributes = buildApplyAttributes(payload); diff --git a/ui/actions/providers/providers.ts b/ui/actions/providers/providers.ts index a862c488af..76e57b30ba 100644 --- a/ui/actions/providers/providers.ts +++ b/ui/actions/providers/providers.ts @@ -151,28 +151,28 @@ export const getProvider = async (formData: FormData) => { const PROVIDERS_PAGE_MAX = 100; /** - * Uids of the given providers, keyed by provider id. A provider's `uid` is the - * candidate it was created for (AWS account id / GCP project id), so this is what - * matches an apply's created providers back to the selection. Batched with - * `filter[id__in]` rather than one `GET /providers/{id}` per id. + * Providers matching the given ids, batched with `filter[id__in]` (page size + * `PROVIDERS_PAGE_MAX`, the server max, which also bounds the id batch size) + * rather than one `GET /providers/{id}` per id. Shared by every action below + * that resolves providers by id; a batch that fails to fetch leaves its + * providers out of the result rather than failing the rest, and it is on each + * caller to say what "missing" means for its own map. Not exported: an + * exported function in this `"use server"` module becomes a callable server + * action. */ -export const getProviderUidsByIds = async ( +const fetchProvidersByIds = async ( providerIds: string[], -): Promise> => { +): Promise => { const uniqueIds = Array.from(new Set(providerIds.filter(Boolean))); if (uniqueIds.length === 0) { - return {}; + return []; } const headers = await getAuthHeaders({ contentType: false }); - const batches: string[][] = []; + const providers: ProvidersApiResponse["data"] = []; + for (let start = 0; start < uniqueIds.length; start += PROVIDERS_PAGE_MAX) { - batches.push(uniqueIds.slice(start, start + PROVIDERS_PAGE_MAX)); - } - - const uidById: Record = {}; - - for (const batch of batches) { + const batch = uniqueIds.slice(start, start + PROVIDERS_PAGE_MAX); const url = new URL(`${apiBaseUrl}/providers`); url.searchParams.set("filter[id__in]", batch.join(",")); url.searchParams.set("page[size]", String(PROVIDERS_PAGE_MAX)); @@ -183,20 +183,93 @@ export const getProviderUidsByIds = async ( | ProvidersApiResponse | undefined; - for (const provider of result?.data ?? []) { - const uid = provider?.attributes?.uid; - if (typeof provider?.id === "string" && typeof uid === "string") { - uidById[provider.id] = uid; - } - } + providers.push(...(result?.data ?? [])); } catch { - // A failed batch leaves its providers unmapped rather than failing the rest. + // A failed batch leaves its providers out of the result rather than + // failing the rest. + } + } + + return providers; +}; + +/** + * Uids of the given providers, keyed by provider id. A provider's `uid` is the + * candidate it was created for (AWS account id / GCP project id), so this is what + * matches an apply's created providers back to the selection. + */ +export const getProviderUidsByIds = async ( + providerIds: string[], +): Promise> => { + const uidById: Record = {}; + + for (const provider of await fetchProvidersByIds(providerIds)) { + const uid = provider?.attributes?.uid; + if (typeof provider?.id === "string" && typeof uid === "string") { + uidById[provider.id] = uid; } } return uidById; }; +/** + * `connection.last_checked_at` for each given provider, keyed by id. Read before + * a batch of connection checks is dispatched, so `resolveProviderConnectionState` + * can tell a check's own result apart from an older one already on record by + * comparing values, never by comparing the browser's clock against the server's + * (see that function for why). A provider missing from the response -- the batch + * read failed, or it was deleted mid-flight -- is left out of the map rather than + * defaulted, so callers can tell "no prior check" (`null`) from "unknown". + */ +export const getProviderConnectionBaselines = async ( + providerIds: string[], +): Promise> => { + const baselineById: Record = {}; + + for (const provider of await fetchProvidersByIds(providerIds)) { + if (typeof provider?.id === "string") { + baselineById[provider.id] = + provider.attributes?.connection?.last_checked_at ?? null; + } + } + + return baselineById; +}; + +/** + * Uid and `connection.last_checked_at` for each given provider, keyed by id, read + * with a single batched `filter[id__in]` request. The organization onboarding + * apply step needs both right after creating providers: the uid to match each one + * back to the candidate it was created for (see `getProviderUidsByIds`), and the + * baseline to compare a dispatched check's result against (see + * `getProviderConnectionBaselines`). Reading them together avoids fetching the + * same set of providers twice. + */ +export const getProviderUidsAndConnectionBaselines = async ( + providerIds: string[], +): Promise<{ + uidById: Record; + baselineById: Record; +}> => { + const uidById: Record = {}; + const baselineById: Record = {}; + + for (const provider of await fetchProvidersByIds(providerIds)) { + if (typeof provider?.id !== "string") { + continue; + } + const uid = provider.attributes?.uid; + if (typeof uid === "string") { + uidById[provider.id] = uid; + } + baselineById[provider.id] = + provider.attributes?.connection?.last_checked_at ?? null; + } + + return { uidById, baselineById }; +}; + export const updateProvider = async (formData: FormData) => { const headers = await getAuthHeaders({ contentType: true }); const providerId = formData.get(ProviderCredentialFields.PROVIDER_ID); diff --git a/ui/actions/registry/registry.adapter.ts b/ui/actions/registry/registry.adapter.ts index 4ae9de0c7d..e84fab1218 100644 --- a/ui/actions/registry/registry.adapter.ts +++ b/ui/actions/registry/registry.adapter.ts @@ -24,6 +24,7 @@ const REGISTRY_TASK_PATH_PREFIX = "/api/v1/tasks/"; const REGISTRY_ERROR_CODE = { KEY_REJECTED: "registry_key_rejected", UNAVAILABLE: "registry_unavailable", + PAGE_NOT_FOUND: "registry_page_not_found", } as const; // Opposite remedies, so a 409 is never read without its code. const REGISTRY_REMOVAL_CONFLICT_CODE = { @@ -216,18 +217,29 @@ export async function classifyRegistryFailure( return { status: REGISTRY_FAILURE.ERROR }; } +// An enabled backend also answers 404 (missing page): tell them apart by code. +export async function isRegistryDisabledResponse(response: Response) { + if (response.status !== 404) return false; + const codes = await getRegistryErrorCodes(response); + return !codes.includes(REGISTRY_ERROR_CODE.PAGE_NOT_FOUND); +} + function isRegistryDiscoveryEndpoint(endpoint: RegistryEndpoint) { return registryDiscoveryEndpoints.has(endpoint); } -async function getRegistryErrorCode(response: Response) { +async function getRegistryErrorCodes(response: Response) { const parsed = errorDocumentSchema.safeParse( await response .clone() .json() .catch(() => undefined), ); - return parsed.success ? parsed.data.errors[0]?.code : undefined; + return parsed.success ? parsed.data.errors.map(({ code }) => code) : []; +} + +async function getRegistryErrorCode(response: Response) { + return (await getRegistryErrorCodes(response))[0]; } const REGISTRY_CATALOG_PAGE_SIZE = 100; diff --git a/ui/actions/registry/registry.test.ts b/ui/actions/registry/registry.test.ts index 7ae458cc3f..a3a44ffcf9 100644 --- a/ui/actions/registry/registry.test.ts +++ b/ui/actions/registry/registry.test.ts @@ -119,14 +119,17 @@ describe("installed Registry provider discovery", () => { emptyMetadata = false, failedEndpoint, failureStatus = 500, + failureBody = {}, }: { emptyMetadata?: boolean; failedEndpoint?: string; failureStatus?: number; + failureBody?: unknown; } = {}) { fetchMock.mockImplementation((url: string) => { const endpoint = new URL(url).pathname.split("/").pop(); - if (endpoint === failedEndpoint) return jsonResponse({}, failureStatus); + if (endpoint === failedEndpoint) + return jsonResponse(failureBody, failureStatus); if (endpoint === "available-artifacts") return jsonResponse({ data: [ @@ -242,6 +245,45 @@ describe("installed Registry provider discovery", () => { }); }, ); + + it.each(["available-artifacts", "artifacts", "providers"])( + "hides Registry when the backend has it disabled and %s answers 404", + async (failedEndpoint) => { + mockDiscovery({ failedEndpoint, failureStatus: 404 }); + expect(await getInstalledRegistryProviderOptions()).toEqual({ + status: "access_denied", + }); + }, + ); + + it.each(["available-artifacts", "providers"])( + "keeps Registry visible when an enabled backend answers 404 for a missing %s page", + async (failedEndpoint) => { + mockDiscovery({ + failedEndpoint, + failureStatus: 404, + failureBody: { + errors: [{ status: "404", code: "registry_page_not_found" }], + }, + }); + expect(await getInstalledRegistryProviderOptions()).toEqual({ + status: "error", + }); + }, + ); + + it("keeps Registry visible when the missing-page code is not the first error", async () => { + mockDiscovery({ + failedEndpoint: "available-artifacts", + failureStatus: 404, + failureBody: { + errors: [{ code: "not_found" }, { code: "registry_page_not_found" }], + }, + }); + expect(await getInstalledRegistryProviderOptions()).toEqual({ + status: "error", + }); + }); }); describe("Registry guarded reads", () => { diff --git a/ui/actions/registry/registry.ts b/ui/actions/registry/registry.ts index 976373a965..4d470bfc6b 100644 --- a/ui/actions/registry/registry.ts +++ b/ui/actions/registry/registry.ts @@ -47,6 +47,7 @@ import { classifyRegistryRemovalConflict, collectCompleteRegistryCatalog, isRegistryCollection, + isRegistryDisabledResponse, parseRegistryArtifactSubmission, parseRegistryCredentialSubmission, RegistryCatalogPageError, @@ -86,6 +87,9 @@ async function readRegistryResponse( return { status: REGISTRY_FAILURE.ERROR }; } if (response.ok) return response; + // A backend with Registry disabled answers 404: hide it like a denial. + if (await isRegistryDisabledResponse(response)) + return { status: REGISTRY_FAILURE.ACCESS_DENIED }; return endpoint === REGISTRY_ENDPOINT.PROVIDERS || endpoint === REGISTRY_ENDPOINT.AVAILABLE_ARTIFACTS diff --git a/ui/actions/scans/scans.test.ts b/ui/actions/scans/scans.test.ts index bf91484ee7..6e08a6f2a9 100644 --- a/ui/actions/scans/scans.test.ts +++ b/ui/actions/scans/scans.test.ts @@ -6,12 +6,14 @@ const { getAuthHeadersMock, handleApiErrorMock, handleApiResponseMock, + isReportDownloadLockedMock, } = vi.hoisted(() => ({ addScanOperationMock: vi.fn(), fetchMock: vi.fn(), getAuthHeadersMock: vi.fn(), handleApiErrorMock: vi.fn(), handleApiResponseMock: vi.fn(), + isReportDownloadLockedMock: vi.fn(), })); vi.mock("@/lib", () => ({ @@ -23,6 +25,10 @@ vi.mock("@/lib", () => ({ error instanceof Error ? error.message : String(error), })); +vi.mock("next/cache", () => ({ + revalidatePath: vi.fn(), +})); + vi.mock("@/lib/server-actions-helper", () => ({ handleApiError: handleApiErrorMock, handleApiResponse: handleApiResponseMock, @@ -32,7 +38,17 @@ vi.mock("@/lib/sentry-breadcrumbs", () => ({ addScanOperation: addScanOperationMock, })); +vi.mock("@/lib/report-download-access", () => ({ + REPORT_DOWNLOAD_LOCKED_ERROR: + "Report downloads require an active subscription.", + isReportDownloadLocked: isReportDownloadLockedMock, +})); + import { + createPartialScan, + getComplianceCsv, + getComplianceOcsf, + getCompliancePdfReport, getExportsZip, launchOrganizationScans, scheduleOrganizationDailyScans, @@ -156,6 +172,7 @@ describe("getExportsZip", () => { vi.clearAllMocks(); vi.stubGlobal("fetch", fetchMock); getAuthHeadersMock.mockResolvedValue({ Authorization: "Bearer token" }); + isReportDownloadLockedMock.mockResolvedValue(false); }); it("returns a generic server error when the report endpoint returns HTML", async () => { @@ -181,3 +198,98 @@ describe("getExportsZip", () => { }); }); }); + +describe("report downloads for subscription-only tenants", () => { + beforeEach(() => { + vi.clearAllMocks(); + vi.stubGlobal("fetch", fetchMock); + getAuthHeadersMock.mockResolvedValue({ Authorization: "Bearer token" }); + isReportDownloadLockedMock.mockResolvedValue(true); + }); + + it.each([ + { name: "scan ZIP", download: () => getExportsZip("scan-123") }, + { + name: "compliance CSV", + download: () => getComplianceCsv("scan-123", "cis_2.0_aws"), + }, + { + name: "compliance OCSF", + download: () => getComplianceOcsf("scan-123", "dora_aws"), + }, + { + name: "compliance PDF", + download: () => getCompliancePdfReport("scan-123", "threatscore"), + }, + ])("rejects the $name without calling the API", async ({ download }) => { + // When + const result = await download(); + + // Then + expect(fetchMock).not.toHaveBeenCalled(); + expect(result).toEqual({ + error: "Report downloads require an active subscription.", + }); + }); +}); + +describe("createPartialScan", () => { + beforeEach(() => { + vi.clearAllMocks(); + vi.stubGlobal("fetch", fetchMock); + getAuthHeadersMock.mockResolvedValue({ Authorization: "Bearer token" }); + fetchMock.mockResolvedValue(new Response(null, { status: 202 })); + handleApiResponseMock.mockResolvedValue({ data: { id: "scan-1" } }); + }); + + it("posts the resource uids as a scan of one provider", async () => { + // When + const result = await createPartialScan({ + providerId: "provider-1", + resourceUids: ["arn:aws:s3:::bucket"], + }); + + // Then + expect(fetchMock).toHaveBeenCalledWith( + "https://api.example.com/api/v1/scans", + expect.objectContaining({ + method: "POST", + body: JSON.stringify({ + data: { + type: "scans", + attributes: { resource_uids: ["arn:aws:s3:::bucket"] }, + relationships: { + provider: { data: { type: "providers", id: "provider-1" } }, + }, + }, + }), + }), + ); + expect(handleApiResponseMock).toHaveBeenCalledWith( + expect.any(Response), + "/scans", + ); + expect(result).toEqual({ data: { id: "scan-1" } }); + expect(addScanOperationMock).toHaveBeenCalledWith("start", "scan-1"); + }); + + it("refuses more resources than the API accepts without calling it", async () => { + // Given — the API caps a partial scan at 10 resources. + const resourceUids = Array.from( + { length: 11 }, + (_, index) => `arn:aws:s3:::bucket-${index}`, + ); + + // When + const result = await createPartialScan({ + providerId: "provider-1", + resourceUids, + }); + + // Then + expect(fetchMock).not.toHaveBeenCalled(); + expect(result).toEqual({ + error: "Select between 1 and 10 resources to re-check", + }); + }); +}); diff --git a/ui/actions/scans/scans.ts b/ui/actions/scans/scans.ts index 89f832552c..cd08f5ce5a 100644 --- a/ui/actions/scans/scans.ts +++ b/ui/actions/scans/scans.ts @@ -15,9 +15,14 @@ import { } from "@/lib/compliance/compliance-report-types"; import { runWithConcurrencyLimit } from "@/lib/concurrency"; import { appendSanitizedProviderTypeFilters } from "@/lib/provider-filters"; +import { + isReportDownloadLocked, + REPORT_DOWNLOAD_LOCKED_ERROR, +} from "@/lib/report-download-access"; import { addScanOperation } from "@/lib/sentry-breadcrumbs"; import { handleApiError, handleApiResponse } from "@/lib/server-actions-helper"; import { SCAN_STATES } from "@/types/attack-paths"; +import { PARTIAL_SCAN_MAX_RESOURCES } from "@/types/partial-scans"; const ORGANIZATION_SCAN_CONCURRENCY_LIMIT = 5; @@ -182,6 +187,70 @@ export const scanOnDemand = async (formData: FormData) => { } }; +/** Prowler Cloud only: re-check up to PARTIAL_SCAN_MAX_RESOURCES resources of one provider. */ +export const createPartialScan = async ({ + providerId, + resourceUids, +}: { + providerId: string; + resourceUids: string[]; +}) => { + if (!providerId) { + return { error: "Provider ID is required" }; + } + if ( + resourceUids.length === 0 || + resourceUids.length > PARTIAL_SCAN_MAX_RESOURCES + ) { + return { + error: `Select between 1 and ${PARTIAL_SCAN_MAX_RESOURCES} resources to re-check`, + }; + } + + const headers = await getAuthHeaders({ contentType: true }); + + addScanOperation("create", undefined, { + provider_id: providerId, + partial: true, + resource_count: resourceUids.length, + }); + + const url = new URL(`${apiBaseUrl}/scans`); + + try { + const requestBody = { + data: { + type: "scans", + attributes: { resource_uids: resourceUids }, + relationships: { + provider: { + data: { + type: "providers", + id: providerId, + }, + }, + }, + }, + }; + + const response = await fetch(url.toString(), { + method: "POST", + headers, + body: JSON.stringify(requestBody), + }); + + const result = await handleApiResponse(response, "/scans"); + if (result?.data?.id) { + addScanOperation("start", result.data.id); + revalidatePath("/scans"); + } + return result; + } catch (error) { + addScanOperation("create"); + return handleApiError(error); + } +}; + export const scheduleDaily = async (formData: FormData) => { const headers = await getAuthHeaders({ contentType: true }); @@ -377,6 +446,10 @@ export const updateScan = async (formData: FormData) => { }; export const getExportsZip = async (scanId: string) => { + if (await isReportDownloadLocked()) { + return { error: REPORT_DOWNLOAD_LOCKED_ERROR }; + } + const headers = await getAuthHeaders({ contentType: false }); const url = new URL(`${apiBaseUrl}/scans/${scanId}/report`); @@ -457,6 +530,10 @@ const _fetchScanBinary = async ( filename: string, errorLabel: string, ): Promise => { + if (await isReportDownloadLocked()) { + return { error: REPORT_DOWNLOAD_LOCKED_ERROR }; + } + const headers = await getAuthHeaders({ contentType: false }); const url = new URL(`${apiBaseUrl}/scans/${scanId}/${urlPath}`); diff --git a/ui/app/(prowler)/compliance/[compliancetitle]/page.tsx b/ui/app/(prowler)/compliance/[compliancetitle]/page.tsx index b153c203f9..6ba928cb21 100644 --- a/ui/app/(prowler)/compliance/[compliancetitle]/page.tsx +++ b/ui/app/(prowler)/compliance/[compliancetitle]/page.tsx @@ -38,6 +38,7 @@ import { } from "@/lib/compliance/compliance-report-types"; import { LIGHTHOUSE_COMPLIANCE_CONTEXT_MODE } from "@/lib/lighthouse/context/constants"; import { buildComplianceContext } from "@/lib/lighthouse/context/contributions"; +import { isReportDownloadLocked } from "@/lib/report-download-access"; import { isCloud } from "@/lib/shared/env"; import { cn } from "@/lib/utils"; import type { SearchParamsProps } from "@/types"; @@ -77,6 +78,8 @@ export default async function ComplianceDetail({ notFound(); } + const subscriptionOnlyPromise = isReportDownloadLocked(); + // Cross-provider mode replaces the per-scan pipeline with the universal // roll-up view. Prowler Cloud-only: the OSS API has no such endpoint, so // the route is blocked in OSS the same way the compliance tab is. @@ -85,6 +88,7 @@ export default async function ComplianceDetail({ redirect("/compliance"); } + const subscriptionOnly = await subscriptionOnlyPromise; return ( ); @@ -124,6 +129,7 @@ export default async function ComplianceDetail({ } const crossAccountTitle = compliancetitle.split("-").join(" "); + const subscriptionOnly = await subscriptionOnlyPromise; return ( @@ -172,18 +179,23 @@ export default async function ComplianceDetail({ let selectedScan: ScanEntity | null = null; const selectedScanId = scanId || null; - const [metadataInfoData, attributesData, selectedScanResponse] = - await Promise.all([ - getComplianceOverviewMetadataInfo({ - filters: { - "filter[scan_id]": selectedScanId ?? undefined, - }, - }), - getComplianceAttributes(complianceId, selectedScanId ?? undefined), - selectedScanId - ? getScan(selectedScanId, { include: "provider" }) - : Promise.resolve(null), - ]); + const [ + metadataInfoData, + attributesData, + selectedScanResponse, + subscriptionOnly, + ] = await Promise.all([ + getComplianceOverviewMetadataInfo({ + filters: { + "filter[scan_id]": selectedScanId ?? undefined, + }, + }), + getComplianceAttributes(complianceId, selectedScanId ?? undefined), + selectedScanId + ? getScan(selectedScanId, { include: "provider" }) + : Promise.resolve(null), + subscriptionOnlyPromise, + ]); // The compliance catalog is still warming after a deploy/restart. Show the // "still loading" state with a Try Again instead of rendering an empty page. @@ -309,6 +321,7 @@ export default async function ComplianceDetail({ complianceId, latestCisIds.has(complianceId), )} + subscriptionOnly={subscriptionOnly} /> )} diff --git a/ui/app/(prowler)/compliance/_actions/cross-provider.test.ts b/ui/app/(prowler)/compliance/_actions/cross-provider.test.ts index edd1481aac..9f0cfa5d87 100644 --- a/ui/app/(prowler)/compliance/_actions/cross-provider.test.ts +++ b/ui/app/(prowler)/compliance/_actions/cross-provider.test.ts @@ -5,11 +5,13 @@ const { getAuthHeadersMock, handleApiResponseMock, captureExceptionMock, + isReportDownloadLockedMock, } = vi.hoisted(() => ({ fetchMock: vi.fn(), getAuthHeadersMock: vi.fn(), handleApiResponseMock: vi.fn(), captureExceptionMock: vi.fn(), + isReportDownloadLockedMock: vi.fn(), })); vi.mock("@/lib", () => ({ @@ -28,6 +30,12 @@ vi.mock("@sentry/nextjs", () => ({ captureException: captureExceptionMock, })); +vi.mock("@/lib/report-download-access", () => ({ + REPORT_DOWNLOAD_LOCKED_ERROR: + "Report downloads require an active subscription.", + isReportDownloadLocked: isReportDownloadLockedMock, +})); + import { generateCrossProviderPdf, getCrossProviderComplianceOverview, @@ -56,6 +64,29 @@ beforeEach(() => { Authorization: "Bearer test-token", }); handleApiResponseMock.mockResolvedValue({ data: null }); + isReportDownloadLockedMock.mockResolvedValue(false); +}); + +describe("cross-provider PDF reports for subscription-only tenants", () => { + it.each([ + { + name: "generation", + run: () => generateCrossProviderPdf({ complianceId: "csa_ccm_4.0" }), + }, + { name: "download", run: () => getCrossProviderPdfBinary("task-1") }, + ])("rejects the $name without calling the API", async ({ run }) => { + // Given + isReportDownloadLockedMock.mockResolvedValue(true); + + // When + const result = await run(); + + // Then + expect(fetchMock).not.toHaveBeenCalled(); + expect(result).toEqual({ + error: "Report downloads require an active subscription.", + }); + }); }); describe("getCrossProviderComplianceOverview", () => { diff --git a/ui/app/(prowler)/compliance/_components/cross-account-detail.tsx b/ui/app/(prowler)/compliance/_components/cross-account-detail.tsx index 9c662558f5..53e502fcb9 100644 --- a/ui/app/(prowler)/compliance/_components/cross-account-detail.tsx +++ b/ui/app/(prowler)/compliance/_components/cross-account-detail.tsx @@ -48,6 +48,7 @@ interface CrossAccountDetailProps { providerType: KnownProviderType; searchParams: Record; targetSection?: string; + subscriptionOnly?: boolean; } /** @@ -63,6 +64,7 @@ export const CrossAccountDetail = async ({ providerType, searchParams, targetSection, + subscriptionOnly = false, }: CrossAccountDetailProps) => { const filters = parseCrossAccountFilters(searchParams); @@ -205,6 +207,7 @@ export const CrossAccountDetail = async ({ providerType={providerType} filters={{ ...filters, scanIds: attrs.scan_ids }} latestPdf={latestPdf} + subscriptionOnly={subscriptionOnly} /> } filters={ diff --git a/ui/app/(prowler)/compliance/_components/cross-provider-detail.tsx b/ui/app/(prowler)/compliance/_components/cross-provider-detail.tsx index 9bdfc91688..8b91362db0 100644 --- a/ui/app/(prowler)/compliance/_components/cross-provider-detail.tsx +++ b/ui/app/(prowler)/compliance/_components/cross-provider-detail.tsx @@ -44,6 +44,7 @@ interface CrossProviderDetailProps { complianceId: string; searchParams: Record; targetSection?: string; + subscriptionOnly?: boolean; } /** @@ -57,6 +58,7 @@ export const CrossProviderDetail = async ({ complianceId, searchParams, targetSection, + subscriptionOnly = false, }: CrossProviderDetailProps) => { const filters = parseCrossProviderFilters(searchParams); @@ -206,6 +208,7 @@ export const CrossProviderDetail = async ({ complianceId={complianceId} filters={{ ...filters, scanIds: attrs.scan_ids }} latestPdf={latestPdf} + subscriptionOnly={subscriptionOnly} /> } filters={ diff --git a/ui/app/(prowler)/compliance/_components/cross-provider-pdf-button.test.tsx b/ui/app/(prowler)/compliance/_components/cross-provider-pdf-button.test.tsx index 12b74a41bf..923f6cdcbd 100644 --- a/ui/app/(prowler)/compliance/_components/cross-provider-pdf-button.test.tsx +++ b/ui/app/(prowler)/compliance/_components/cross-provider-pdf-button.test.tsx @@ -2,6 +2,9 @@ import { render, screen, waitFor } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; import { beforeAll, beforeEach, describe, expect, it, vi } from "vitest"; +import { useCloudUpgradeStore } from "@/store/cloud-upgrade/store"; +import { PAID_PLAN_UPGRADE_FEATURE } from "@/types/cloud-upgrade"; + import { CrossProviderPdfButton } from "./cross-provider-pdf-button"; // Radix dialogs/dropdowns rely on pointer-capture and scrollIntoView, which @@ -99,6 +102,7 @@ describe("CrossProviderPdfButton", () => { vi.clearAllMocks(); storeState.tasks = {}; generatePdfMock.mockResolvedValue({ taskId: "task-1" }); + useCloudUpgradeStore.getState().closeCloudUpgrade(); }); const openGenerateModal = async ( @@ -192,6 +196,34 @@ describe("CrossProviderPdfButton", () => { await waitFor(() => expect(downloadPdfMock).toHaveBeenCalledWith("task-7")); }); + it.each([/download latest/i, /generate new report/i])( + "opens the paid plan upgrade instead of %s for subscription-only tenants", + async (label) => { + // Given + const user = userEvent.setup(); + render( + , + ); + + // When + await user.click(screen.getByRole("button", { name: /report/i })); + await user.click(await screen.findByRole("menuitem", { name: label })); + + // Then + expect(downloadPdfMock).not.toHaveBeenCalled(); + expect( + screen.queryByRole("dialog", { name: /generate/i }), + ).not.toBeInTheDocument(); + expect(useCloudUpgradeStore.getState().activeFeature).toBe( + PAID_PLAN_UPGRADE_FEATURE.REPORT_DOWNLOAD, + ); + }, + ); + it("keeps a completed report downloadable after the ready toast closes", async () => { // Given storeState.tasks = { diff --git a/ui/app/(prowler)/compliance/_components/cross-provider-pdf-button.tsx b/ui/app/(prowler)/compliance/_components/cross-provider-pdf-button.tsx index 3e19460feb..1056fbf9d1 100644 --- a/ui/app/(prowler)/compliance/_components/cross-provider-pdf-button.tsx +++ b/ui/app/(prowler)/compliance/_components/cross-provider-pdf-button.tsx @@ -12,6 +12,7 @@ import { FormButtons } from "@/components/shadcn/form"; import { Input } from "@/components/shadcn/input/input"; import { Modal } from "@/components/shadcn/modal"; import { toast } from "@/components/shadcn/toast"; +import { useReportDownload } from "@/hooks/use-report-download"; import { TASK_WATCHER_STATUS, trackAndPollTask, @@ -47,6 +48,8 @@ interface CrossProviderPdfButtonProps { /** Already-generated report matching these filters, if any — offered as an * instant download instead of forcing a re-generate. */ latestPdf: LatestCrossProviderPdf | null; + /** Prowler Cloud tenants without a paid plan cannot download reports. */ + subscriptionOnly?: boolean; } export const CrossProviderPdfButton = ({ @@ -54,7 +57,9 @@ export const CrossProviderPdfButton = ({ providerType, filters, latestPdf, + subscriptionOnly = false, }: CrossProviderPdfButtonProps) => { + const runReportDownload = useReportDownload(subscriptionOnly); const [dialogOpen, setDialogOpen] = useState(false); const [reportName, setReportName] = useState(""); const [submitting, setSubmitting] = useState(false); @@ -185,13 +190,15 @@ export const CrossProviderPdfButton = ({ icon={} label={`Download latest${formatGeneratedAt(availablePdf.completedAt)}`} description={availablePdf.filename} - onSelect={() => downloadPdf(availablePdf.taskId)} + onSelect={() => + runReportDownload(() => downloadPdf(availablePdf.taskId)) + } /> )} } label="Generate new report…" - onSelect={() => setDialogOpen(true)} + onSelect={() => runReportDownload(() => setDialogOpen(true))} /> )} diff --git a/ui/app/(prowler)/compliance/_lib/aggregated-compliance-actions.ts b/ui/app/(prowler)/compliance/_lib/aggregated-compliance-actions.ts index 6c8303dc83..6d6004e5e5 100644 --- a/ui/app/(prowler)/compliance/_lib/aggregated-compliance-actions.ts +++ b/ui/app/(prowler)/compliance/_lib/aggregated-compliance-actions.ts @@ -7,6 +7,10 @@ import { getErrorMessage, } from "@/lib"; import { hasActionError, type ActionErrorResult } from "@/lib/action-errors"; +import { + isReportDownloadLocked, + REPORT_DOWNLOAD_LOCKED_ERROR, +} from "@/lib/report-download-access"; import { handleApiResponse } from "@/lib/server-actions-helper"; import { SentryErrorSource, SentryErrorType } from "@/sentry"; @@ -171,6 +175,10 @@ export const generateAggregatedCompliancePdf = async ( url: URL, operation: string, ): Promise<{ taskId: string } | { error: string }> => { + if (await isReportDownloadLocked()) { + return { error: REPORT_DOWNLOAD_LOCKED_ERROR }; + } + const headers = await getAuthHeaders({ contentType: false }); try { @@ -211,6 +219,10 @@ export const getAggregatedCompliancePdfBinary = async ({ operation: string; defaultFilename: string; }): Promise => { + if (await isReportDownloadLocked()) { + return { error: REPORT_DOWNLOAD_LOCKED_ERROR }; + } + const headers = await getAuthHeaders({ contentType: false }); try { diff --git a/ui/app/(prowler)/compliance/page.test.tsx b/ui/app/(prowler)/compliance/page.test.tsx index 759ac4f32a..8f5e89a023 100644 --- a/ui/app/(prowler)/compliance/page.test.tsx +++ b/ui/app/(prowler)/compliance/page.test.tsx @@ -15,18 +15,24 @@ import { beforeEach, describe, expect, it, vi } from "vitest"; import Compliance from "./page"; const { + complianceFiltersSpy, complianceOverviewGridSpy, getComplianceOverviewMetadataInfoMock, getCompliancesOverviewMock, + getScanMock, getScansMock, getThreatScoreMock, + isCloudMock, loadComplianceWatchlistContextMock, } = vi.hoisted(() => ({ + complianceFiltersSpy: vi.fn(), complianceOverviewGridSpy: vi.fn(), getComplianceOverviewMetadataInfoMock: vi.fn(), getCompliancesOverviewMock: vi.fn(), + getScanMock: vi.fn(), getScansMock: vi.fn(), getThreatScoreMock: vi.fn(), + isCloudMock: vi.fn(() => false), loadComplianceWatchlistContextMock: vi.fn(), })); @@ -41,12 +47,13 @@ vi.mock("@/actions/overview", () => ({ })); vi.mock("@/actions/scans", () => ({ + getScan: getScanMock, getScans: getScansMock, getScansByState: vi.fn(), })); vi.mock("@/lib/shared/env", () => ({ - isCloud: () => false, + isCloud: isCloudMock, })); vi.mock("./_lib/watchlist-context", () => ({ @@ -83,7 +90,10 @@ vi.mock("@/components/compliance", () => ({ })); vi.mock("@/components/compliance/compliance-header/compliance-filters", () => ({ - ComplianceFilters: () =>
Compliance filters
, + ComplianceFilters: (props: { scans: Array<{ id: string }> }) => { + complianceFiltersSpy(props); + return
Compliance filters
; + }, })); vi.mock("@/components/compliance/compliance-overview-grid", () => ({ @@ -148,6 +158,8 @@ describe("Compliance overview page", () => { describe("Compliance overview task response", () => { beforeEach(() => { vi.clearAllMocks(); + isCloudMock.mockReturnValue(false); + getScanMock.mockResolvedValue(undefined); getScansMock.mockResolvedValue({ data: [ { @@ -233,6 +245,138 @@ describe("Compliance overview task response", () => { ); }); + it("keeps partial scans out of the per-scan selector", async () => { + // Given - a Cloud partial scan among the completed scans; it computes no + // compliance, so selecting it would show nothing and its downloads fail + isCloudMock.mockReturnValue(true); + getScansMock.mockResolvedValue({ + data: [ + { + id: "scan-1", + attributes: { + name: "Production scan", + completed_at: "2026-08-05T17:00:00Z", + }, + relationships: { provider: { data: { id: "provider-1" } } }, + }, + { + id: "scan-partial", + attributes: { + name: "Re-check", + completed_at: "2026-08-06T09:00:00Z", + is_partial: true, + }, + relationships: { provider: { data: { id: "provider-1" } } }, + }, + ], + included: [ + { + type: "providers", + id: "provider-1", + attributes: { provider: "aws", uid: "123456789012", alias: "prod" }, + }, + ], + }); + getCompliancesOverviewMock.mockResolvedValue({ data: [] }); + + // When + const page = await Compliance({ + searchParams: Promise.resolve({ scanId: "scan-1" }), + }); + render(page as ReactElement); + + // Then - only the full scan reaches the selector, and the API is asked + // for the flag that tells them apart + expect(complianceFiltersSpy).toHaveBeenCalledWith( + expect.objectContaining({ + scans: [expect.objectContaining({ id: "scan-1" })], + }), + ); + expect(getScansMock).toHaveBeenCalledWith( + expect.objectContaining({ + // Filtered at the API too, so a page full of re-checks cannot hide + // the full scans behind it. + filters: expect.objectContaining({ "filter[is_partial]": "false" }), + fields: { scans: "name,completed_at,provider,is_partial" }, + }), + ); + }); + + it("does not send the Cloud-only partial filter outside Prowler Cloud", async () => { + getCompliancesOverviewMock.mockResolvedValue({ data: [] }); + + await Compliance({ searchParams: Promise.resolve({ scanId: "scan-1" }) }); + + expect(getScansMock).toHaveBeenCalledWith( + expect.objectContaining({ + filters: { "filter[state]": "completed" }, + }), + ); + }); + + it("falls back to the first full scan when the URL names a partial scan", async () => { + // Given - a stale link to a partial scan, absent from the eligible list + getScanMock.mockResolvedValue({ + data: { id: "scan-partial", attributes: { is_partial: true } }, + }); + getCompliancesOverviewMock.mockResolvedValue({ data: [] }); + + // When + const page = await Compliance({ + searchParams: Promise.resolve({ scanId: "scan-partial" }), + }); + render(page as ReactElement); + + // Then - the page selects a scan that has compliance instead + expect(getScanMock).toHaveBeenCalledWith("scan-partial"); + expect(complianceFiltersSpy).toHaveBeenCalledWith( + expect.objectContaining({ selectedScanId: "scan-1" }), + ); + expect(getCompliancesOverviewMock).toHaveBeenCalledWith( + expect.objectContaining({ scanId: "scan-1" }), + ); + }); + + it("falls back to the first full scan when the URL scan cannot be found", async () => { + // Given - a deleted or mistyped id: the lookup returns an error, no data + getScanMock.mockResolvedValue({ error: "Not found", status: 404 }); + getCompliancesOverviewMock.mockResolvedValue({ data: [] }); + + // When + const page = await Compliance({ + searchParams: Promise.resolve({ scanId: "scan-missing" }), + }); + render(page as ReactElement); + + // Then - no compliance request goes out for an id that does not exist + expect(complianceFiltersSpy).toHaveBeenCalledWith( + expect.objectContaining({ selectedScanId: "scan-1" }), + ); + expect(getCompliancesOverviewMock).not.toHaveBeenCalledWith( + expect.objectContaining({ scanId: "scan-missing" }), + ); + }); + + it("keeps trusting a URL scan id that is older than the listed page", async () => { + // Given - a full scan beyond the first page, shaped like an OSS response + // that carries no is_partial field at all + getScanMock.mockResolvedValue({ + data: { id: "scan-old", attributes: { name: "Old scan" } }, + }); + getCompliancesOverviewMock.mockResolvedValue({ data: [] }); + + // When + const page = await Compliance({ + searchParams: Promise.resolve({ scanId: "scan-old" }), + }); + render(page as ReactElement); + + // Then + expect(complianceFiltersSpy).toHaveBeenCalledWith( + expect.objectContaining({ selectedScanId: "scan-old" }), + ); + }); + it("shows the invalid scan message for a JSON:API error response", async () => { // Given - handleApiResponse converted a client error to its error result getCompliancesOverviewMock.mockResolvedValue({ diff --git a/ui/app/(prowler)/compliance/page.tsx b/ui/app/(prowler)/compliance/page.tsx index 015521f6fc..6a31432d99 100644 --- a/ui/app/(prowler)/compliance/page.tsx +++ b/ui/app/(prowler)/compliance/page.tsx @@ -7,7 +7,7 @@ import { getCompliancesOverview, } from "@/actions/compliances"; import { getThreatScore } from "@/actions/overview"; -import { getScans, getScansByState } from "@/actions/scans"; +import { getScan, getScans, getScansByState } from "@/actions/scans"; import { ComplianceSkeletonGrid, NoScansAvailable, @@ -20,6 +20,7 @@ import { Alert, AlertDescription } from "@/components/shadcn/alert"; import { Card, CardContent } from "@/components/shadcn/card/card"; import { ContentLayout } from "@/components/shadcn/content-layout"; import { pickLatestCisPerProvider } from "@/lib/compliance/compliance-report-types"; +import { isReportDownloadLocked } from "@/lib/report-download-access"; import { isCloud } from "@/lib/shared/env"; import { ExpandedScanData, @@ -40,6 +41,31 @@ import { import type { ComplianceWatchlistContext } from "./_lib/watchlist-context"; import { loadComplianceWatchlistContext } from "./_lib/watchlist-context"; +/** + * A scan id from the URL is trusted when it is listed, or when a single lookup + * returns a scan that is not partial (older full scans keep working). A partial + * scan, reached through a stale link, or an id the API cannot find, has no + * compliance to show, so the caller falls back to the first eligible scan. + */ +async function resolveUrlScanId( + scanIdFromUrl: string | undefined, + eligibleScans: ExpandedScanData[], +): Promise { + if (!scanIdFromUrl) return undefined; + if (eligibleScans.some((scan) => scan.id === scanIdFromUrl)) { + return scanIdFromUrl; + } + + const urlScan = (await getScan(scanIdFromUrl)) as + | { data?: { attributes?: { is_partial?: boolean } } } + | undefined; + const scan = urlScan?.data; + if (!scan) return undefined; + + // OSS scans carry no is_partial at all, so only an explicit true excludes. + return scan.attributes?.is_partial === true ? undefined : scanIdFromUrl; +} + export default async function Compliance({ searchParams, }: { @@ -124,16 +150,23 @@ export default async function Compliance({ ); } - const scansData = await getScans({ - filters: { - "filter[state]": "completed", - }, - pageSize: 50, - fields: { - scans: "name,completed_at,provider", - }, - include: "provider", - }); + const [subscriptionOnly, scansData] = await Promise.all([ + isReportDownloadLocked(), + getScans({ + filters: { + "filter[state]": "completed", + // Partial scans compute no compliance. Exclude them at the API so the + // page below never fills up with them; the filter is Cloud-only. + ...(isCloud() ? { "filter[is_partial]": "false" } : {}), + }, + pageSize: 50, + fields: { + // is_partial is Cloud-only; the OSS API ignores unknown sparse fields. + scans: "name,completed_at,provider,is_partial", + }, + include: "provider", + }), + ]); if (!scansData?.data) { return ( @@ -157,7 +190,10 @@ export default async function Compliance({ ); } + // Belt and braces for an API without the filter: partial scans never + // compute compliance, so they have nothing to show or download here. const expandedScansData: ExpandedScanData[] = scansData.data + .filter((scan: ScanProps) => !scan.attributes?.is_partial) .filter((scan: ScanProps) => scan.relationships?.provider?.data?.id) .map((scan: ScanProps) => { const providerId = scan.relationships!.provider!.data!.id; @@ -187,7 +223,9 @@ export default async function Compliance({ ? scanIdParam[0] : scanIdParam; const selectedScanId: string | null = - scanIdFromUrl || expandedScansData[0]?.id || null; + (await resolveUrlScanId(scanIdFromUrl, expandedScansData)) || + expandedScansData[0]?.id || + null; const onboardingAction = selectedScanId ? { flowId: "view-compliance" } : { @@ -256,6 +294,7 @@ export default async function Compliance({ provider={selectedScan.providerInfo.provider} selectedScan={selectedScanData} sectionScores={threatScoreData.sectionScores} + subscriptionOnly={subscriptionOnly} /> )} @@ -273,6 +312,7 @@ export default async function Compliance({ scanId={selectedScanId} selectedScan={selectedScanData} watchlistPromise={watchlistPromise} + subscriptionOnly={subscriptionOnly} /> @@ -302,11 +342,13 @@ const SSRComplianceGrid = async ({ scanId, selectedScan, watchlistPromise, + subscriptionOnly, }: { searchParams: SearchParamsProps; scanId: string | null; selectedScan?: ScanEntity; watchlistPromise: Promise; + subscriptionOnly: boolean; }) => { const regionFilter = searchParams["filter[region__in]"]?.toString() || ""; @@ -388,6 +430,7 @@ const SSRComplianceGrid = async ({ catalogEntries={watchlist.entries} providerType={providerType} canManageWatchlist={watchlist.canManage} + subscriptionOnly={subscriptionOnly} /> ); diff --git a/ui/app/(prowler)/layout.tsx b/ui/app/(prowler)/layout.tsx index 4a6244ffec..ebb8a9c53c 100644 --- a/ui/app/(prowler)/layout.tsx +++ b/ui/app/(prowler)/layout.tsx @@ -59,6 +59,10 @@ export default async function RootLayout({ // Skip Cloud-only onboarding fetches and orchestrators in OSS. const cloudEnabled = isCloud(); + // Every deployment needs the provider count: it drives the first-run redirect + // and the sidebar's Add Provider action. + const providersPromise = getProviders({ page: 1, pageSize: 1 }); + // One-time server-side Registry gate per request: only an ELIGIBLE answer // shows the sidebar entry; UNKNOWN and INELIGIBLE both hide it. Started // here so it resolves in parallel with the Cloud onboarding fetches. @@ -69,29 +73,27 @@ export default async function RootLayout({ // Fail-open: unknown scan state is treated as "has data" so the banner never blocks // progression on a fetch error. let hasCompletedScan = true; - // Tri-state: true = has providers, false = zero providers, undefined = fetch failed (gate fails open). - let hasProviders: boolean | undefined = false; // Scopes the onboarding steps' local markers, so resolving them for one // tenant does not silence them for another. let tenantId: string | null = null; if (cloudEnabled) { - const [providersData, scansByState] = await Promise.all([ - getProviders({ page: 1, pageSize: 1 }), - getScansByState(), - ]); + const scansByState = await getScansByState(); hasCompletedScan = Array.isArray(scansByState?.data) ? scansByState.data.some( (scan: { attributes?: { state?: string } }) => scan.attributes?.state === SCAN_STATES.COMPLETED, ) : true; - hasProviders = Array.isArray(providersData?.data) - ? providersData.data.length > 0 - : undefined; tenantId = (await auth())?.tenantId ?? null; } + const providersData = await providersPromise; + // Tri-state: true = has providers, false = zero providers, undefined = fetch failed (gate fails open). + const hasProviders: boolean | undefined = Array.isArray(providersData?.data) + ? providersData.data.length > 0 + : undefined; + const registryEligible = (await registryAccessPromise).status === REGISTRY_ACCESS.ELIGIBLE; @@ -114,13 +116,12 @@ export default async function RootLayout({ - {/* Store uses boolean; gate receives tri-state to fail open on fetch errors. */} - + {/* Tri-state for both: an unknown count leaves the store unresolved and the gate closed. */} + + {/* Every deployment: an empty tenant lands on the add-provider wizard once. */} + {cloudEnabled && ( <> - {/* Single mount point so the watcher survives post-connect navigation. */} {/* Persistent banner shown only while a guided sequence is active. */} diff --git a/ui/app/(prowler)/providers/providers-page.harness.tsx b/ui/app/(prowler)/providers/providers-page.harness.tsx index 4c10b5b07b..61a96f64c2 100644 --- a/ui/app/(prowler)/providers/providers-page.harness.tsx +++ b/ui/app/(prowler)/providers/providers-page.harness.tsx @@ -54,6 +54,19 @@ export class ProvidersPageHarness extends BrowserHarness { return this.countRequests("POST", "/apply"); } + /** `POST /providers` alone: secrets and connection checks nest under it. */ + get providerCreateCallCount(): number { + return this.requestLog.filter( + (request) => + request.method === "POST" && + new URL(request.url).pathname.replace(/\/$/, "").endsWith("/providers"), + ).length; + } + + get secretCreateCallCount(): number { + return this.countRequests("POST", "/providers/secrets"); + } + /** * Whether any apply asked the endpoint to include related resources, which it * rejects outright — a tripwire, not a preference. @@ -83,7 +96,7 @@ export class ProvidersPageHarness extends BrowserHarness { ).length; } - private get connectionCallCount(): number { + get connectionCallCount(): number { return this.countRequests("POST", "/connection"); } @@ -198,10 +211,75 @@ export class ProvidersPageHarness extends BrowserHarness { /** Enter the AWS Organizations onboarding flow from a fresh wizard. */ async chooseAwsOrganizations(): Promise { await this.selectProviderType(/Amazon Web Services/); - await this.chooseMethod(/Add Multiple Accounts With AWS Organizations/); + // AWS hosts its single-account/organization switch as tabs on its connect step. + const tab = await this.waitFor(() => + this.byRoleName("tab", /Full AWS Organization/), + ); + await this.user.click(tab); await this.waitForText(/Organization Details/); } + /** Labels of the wizard's progress stepper, top to bottom. */ + stepperLabels(): string[] { + return Array.from( + document.querySelectorAll('nav[aria-label="Wizard progress"] span'), + (node) => node.textContent?.trim() ?? "", + ); + } + + /** Wait until the AWS single-account connect step (with its method tabs) is showing. */ + async waitForAwsConnectStep(): Promise { + await this.waitFor( + () => this.byRoleName("tab", /Single AWS Account/), + undefined, + "AWS connect step", + ); + } + + /** Type the account id and static keys on the AWS one-step connect form. */ + async fillAwsAccountKeys({ + accountId, + accessKeyId, + secretAccessKey, + }: { + accountId: string; + accessKeyId: string; + secretAccessKey: string; + }): Promise { + const accountInput = await this.waitFor(() => + this.inputByName("providerUid"), + ); + await this.user.fill(accountInput, accountId); + const keyInput = await this.waitFor(() => + this.inputByName("aws_access_key_id"), + ); + await this.user.fill(keyInput, accessKeyId); + const secretInput = await this.waitFor(() => + this.inputByName("aws_secret_access_key"), + ); + await this.user.fill(secretInput, secretAccessKey); + } + + /** Submit the AWS one-step form; waits for its action to become enabled. */ + async connectAccount(): Promise { + await this.clickPrimary(/Connect account/); + } + + /** Wait until the provider wizard reached its launch step. */ + async waitForProviderLaunchStep(timeoutMs = 20000): Promise { + await this.waitForText(/Scan Schedule/, timeoutMs); + } + + /** Switch back to a single account from the organization flow's tabs. */ + async switchToAwsSingleAccount(): Promise { + const tab = await this.waitFor( + () => this.byRoleName("tab", /Single AWS Account/), + undefined, + "Single AWS Account tab", + ); + await this.user.click(tab); + } + /** Select GCP and open the GCP Organization method card (no advance wait). */ async chooseGcpOrganizationsMethod(): Promise { await this.selectProviderType(/Google Cloud Platform/); diff --git a/ui/app/(prowler)/providers/providers-page.integration.test.tsx b/ui/app/(prowler)/providers/providers-page.integration.test.tsx index 33b6cb91d0..5f3fd66bfc 100644 --- a/ui/app/(prowler)/providers/providers-page.integration.test.tsx +++ b/ui/app/(prowler)/providers/providers-page.integration.test.tsx @@ -212,6 +212,83 @@ describe("Organization onboarding wizard", () => { }, 60000); }); + describe("Wizard progress", () => { + it("drops the credentials and test rows once AWS is picked, since one step covers them", async () => { + const harness = new ProvidersPageHarness(awsOnboardingFixture()); + await harness.mount(); + expect(harness.stepperLabels()).toEqual([ + "Link a Provider", + "Authenticate Credentials", + "Validate Connection", + "Launch Scan", + ]); + + await harness.selectProviderType(/Amazon Web Services/); + await harness.waitForAwsConnectStep(); + + expect(harness.stepperLabels()).toEqual([ + "Link a Provider", + "Launch Scan", + ]); + }, 40000); + }); + + describe("Single account with access keys", () => { + // Runs compiled by the React Compiler, unlike the unit suite: it guards the + // form's validity being read as a reactive value, not frozen in a memo. + it("enables Connect account once the form is filled and jumps to the launch step", async () => { + const harness = new ProvidersPageHarness(awsOnboardingFixture()); + await harness.mount(); + await harness.selectProviderType(/Amazon Web Services/); + await harness.waitForAwsConnectStep(); + await harness.chooseMethod(/Static access keys/); + await harness.fillAwsAccountKeys({ + accountId: "210987654321", + accessKeyId: "AKIAEXAMPLE", + secretAccessKey: "secret-value", + }); + + await harness.connectAccount(); + + await harness.waitForProviderLaunchStep(); + expect(harness.providerCreateCallCount).toBe(1); + expect(harness.secretCreateCallCount).toBe(1); + // Reaching launch above is what proves no separate test step ran; this + // pins that the one step really did check the connection. + expect(harness.connectionCallCount).toBe(1); + const secret = await harness.lastRequestBody<{ + data: { relationships: { provider: { data: { id: string } } } }; + }>("POST", "/providers/secrets"); + expect(secret?.data.relationships.provider.data.id).toBe( + "provider-created-1", + ); + }, 40000); + }); + + describe("Leaving the organization flow", () => { + it("keeps the method tabs on Organization Details and switches back to a single account", async () => { + const harness = new ProvidersPageHarness(awsOnboardingFixture()); + await harness.mount(); + await harness.chooseAwsOrganizations(); + + await harness.switchToAwsSingleAccount(); + + await harness.waitForAwsConnectStep(); + expect(harness.hasOrganizationSetupStep()).toBe(false); + }, 40000); + + it("returns to the AWS single-account step when going back from Organization Details", async () => { + const harness = new ProvidersPageHarness(awsOnboardingFixture()); + await harness.mount(); + await harness.chooseAwsOrganizations(); + + await harness.goBack(); + + await harness.waitForAwsConnectStep(); + expect(harness.hasOrganizationSetupStep()).toBe(false); + }, 40000); + }); + describe("Account selection", () => { it("disables blocked accounts and excludes them from the selectable count", async () => { const harness = new ProvidersPageHarness(awsOnboardingFixture()); diff --git a/ui/app/(prowler)/scans/page.tsx b/ui/app/(prowler)/scans/page.tsx index 4b308921e4..fe4e93c98e 100644 --- a/ui/app/(prowler)/scans/page.tsx +++ b/ui/app/(prowler)/scans/page.tsx @@ -22,6 +22,7 @@ import { import { SkeletonTableScans } from "@/components/scans/table"; import { ScanJobsTable } from "@/components/scans/table/scan-jobs-table"; import { ContentLayout } from "@/components/shadcn/content-layout"; +import { isReportDownloadLocked } from "@/lib/report-download-access"; import { buildProviderScheduleSummary, buildSchedulesByProviderId, @@ -196,7 +197,10 @@ export default async function Scans({ const hasManageIngestionsPermission = Boolean( session?.user?.permissions?.manage_ingestions, ); - const activeScanCount = await getActiveScanCount(resolvedSearchParams); + const [activeScanCount, reportDownloadLocked] = await Promise.all([ + getActiveScanCount(resolvedSearchParams), + isReportDownloadLocked(), + ]); // Mirrors ScansPageShell's launch gate: it only mounts the view-first-scan trigger // when Launch Scan is usable (manage_scans + a connected provider). Without the // permission nothing can consume the navbar action, so offer none rather than an @@ -234,6 +238,7 @@ export default async function Scans({ @@ -245,10 +250,12 @@ const SSRDataTableScans = async ({ searchParams, providers, scanScheduleCapability, + subscriptionOnly, }: { searchParams: SearchParamsProps; providers: ProviderProps[]; scanScheduleCapability?: ScanScheduleCapability; + subscriptionOnly: boolean; }) => { const tab = getScanJobsTab(searchParams.tab); @@ -294,6 +301,7 @@ const SSRDataTableScans = async ({ tab={tab} hasFilters={hasUserFilters} scanScheduleCapability={capability} + subscriptionOnly={subscriptionOnly} /> ); } @@ -389,6 +397,7 @@ const SSRDataTableScans = async ({ tab={tab} hasFilters={hasUserFilters} scanScheduleCapability={scanScheduleCapability} + subscriptionOnly={subscriptionOnly} /> ); }; diff --git a/ui/app/api/scans/[scanId]/report/route.test.ts b/ui/app/api/scans/[scanId]/report/route.test.ts index ad5e9194ea..468482a54c 100644 --- a/ui/app/api/scans/[scanId]/report/route.test.ts +++ b/ui/app/api/scans/[scanId]/report/route.test.ts @@ -1,9 +1,10 @@ -import { afterEach, describe, expect, it, vi } from "vitest"; +import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; import { GET } from "./route"; -const { getAuthHeadersMock } = vi.hoisted(() => ({ +const { getAuthHeadersMock, isReportDownloadLockedMock } = vi.hoisted(() => ({ getAuthHeadersMock: vi.fn(), + isReportDownloadLockedMock: vi.fn(), })); vi.mock("@/lib", () => ({ @@ -11,12 +12,48 @@ vi.mock("@/lib", () => ({ getAuthHeaders: getAuthHeadersMock, })); +vi.mock("@/lib/report-download-access", () => ({ + REPORT_DOWNLOAD_LOCKED_ERROR: + "Report downloads require an active subscription.", + isReportDownloadLocked: isReportDownloadLockedMock, +})); + describe("GET /api/scans/[scanId]/report", () => { + beforeEach(() => { + isReportDownloadLockedMock.mockResolvedValue(false); + }); + afterEach(() => { vi.unstubAllGlobals(); vi.clearAllMocks(); }); + it.each([ + { label: "download", url: "http://localhost/api" }, + { label: "preflight", url: "http://localhost/api?preflight=1" }, + ])( + "rejects the $label without reaching the API when downloads are locked", + async ({ url }) => { + // Given + const fetchMock = vi.fn(); + vi.stubGlobal("fetch", fetchMock); + isReportDownloadLockedMock.mockResolvedValue(true); + + // When + const response = await GET(new Request(url), { + params: Promise.resolve({ scanId: "scan-123" }), + }); + + // Then + expect(fetchMock).not.toHaveBeenCalled(); + expect(response.status).toBe(403); + expect(response.headers.get("cache-control")).toBe("no-store"); + await expect(response.text()).resolves.toBe( + "Report downloads require an active subscription.", + ); + }, + ); + it("streams the upstream report body without buffering it", async () => { const upstreamBody = new ReadableStream({ start(controller) { diff --git a/ui/app/api/scans/[scanId]/report/route.ts b/ui/app/api/scans/[scanId]/report/route.ts index 5b891385cb..7c0c7d1083 100644 --- a/ui/app/api/scans/[scanId]/report/route.ts +++ b/ui/app/api/scans/[scanId]/report/route.ts @@ -1,6 +1,10 @@ import { NextResponse } from "next/server"; import { apiBaseUrl, getAuthHeaders } from "@/lib"; +import { + isReportDownloadLocked, + REPORT_DOWNLOAD_LOCKED_ERROR, +} from "@/lib/report-download-access"; export const dynamic = "force-dynamic"; export const runtime = "nodejs"; @@ -63,6 +67,14 @@ export async function GET( { params }: ScanReportRouteContext, ) { const { scanId } = await params; + + if (await isReportDownloadLocked()) { + return new Response(REPORT_DOWNLOAD_LOCKED_ERROR, { + status: 403, + headers: { "Cache-Control": "no-store", "Content-Type": "text/plain" }, + }); + } + const headers = await getAuthHeaders({ contentType: false }); const upstreamUrl = `${apiBaseUrl}/scans/${encodeURIComponent(scanId)}/report`; const isPreflight = diff --git a/ui/changelog.d/aws-marketplace-button.added.md b/ui/changelog.d/aws-marketplace-button.added.md deleted file mode 100644 index 235112967f..0000000000 --- a/ui/changelog.d/aws-marketplace-button.added.md +++ /dev/null @@ -1 +0,0 @@ -AWS Marketplace button variant with outlined styling for light and dark themes diff --git a/ui/changelog.d/aws-one-step-connect.changed.md b/ui/changelog.d/aws-one-step-connect.changed.md new file mode 100644 index 0000000000..924f32b1bb --- /dev/null +++ b/ui/changelog.d/aws-one-step-connect.changed.md @@ -0,0 +1 @@ +AWS accounts are connected in a single wizard step: the account is read from the role ARN, or typed for access keys, the role is assumed with Prowler's own credentials, and the credentials are stored and tested with the account diff --git a/ui/changelog.d/cloudflare-token-permissions.fixed.md b/ui/changelog.d/cloudflare-token-permissions.fixed.md deleted file mode 100644 index 9a862ee88c..0000000000 --- a/ui/changelog.d/cloudflare-token-permissions.fixed.md +++ /dev/null @@ -1 +0,0 @@ -Cloudflare API token links in the provider wizard request the SSL and Certificates, Bot Management and Zone WAF read permissions the scan needs diff --git a/ui/changelog.d/defer-onboarding-on-billing.fixed.md b/ui/changelog.d/defer-onboarding-on-billing.fixed.md deleted file mode 100644 index 5afcc36bcc..0000000000 --- a/ui/changelog.d/defer-onboarding-on-billing.fixed.md +++ /dev/null @@ -1 +0,0 @@ -Automatic onboarding stays hidden on billing pages and remains available after leaving billing diff --git a/ui/changelog.d/disable-self-registration.added.md b/ui/changelog.d/disable-self-registration.added.md deleted file mode 100644 index 46e62eb1b3..0000000000 --- a/ui/changelog.d/disable-self-registration.added.md +++ /dev/null @@ -1 +0,0 @@ -`UI_SELF_REGISTRATION_ENABLED` flag for Prowler Private Cloud deployments; when `"false"`, `/sign-up` only opens with an invitation, the sign-in page drops the "Sign up" link and the profile hides **Create organization** diff --git a/ui/changelog.d/fedramp-20x-cross-provider-breakdown.fixed.md b/ui/changelog.d/fedramp-20x-cross-provider-breakdown.fixed.md deleted file mode 100644 index dc7c71ce1b..0000000000 --- a/ui/changelog.d/fedramp-20x-cross-provider-breakdown.fixed.md +++ /dev/null @@ -1 +0,0 @@ -Per-provider breakdown and OCSF download for FedRAMP 20x KSI and Class C FRR in the cross-provider compliance view diff --git a/ui/changelog.d/first-run-add-provider.changed.md b/ui/changelog.d/first-run-add-provider.changed.md new file mode 100644 index 0000000000..a40ae414e3 --- /dev/null +++ b/ui/changelog.d/first-run-add-provider.changed.md @@ -0,0 +1 @@ +New tenants without providers land on the Add Provider wizard on first sign-in instead of a welcome modal diff --git a/ui/changelog.d/invitation-row-actions-non-pending.fixed.md b/ui/changelog.d/invitation-row-actions-non-pending.fixed.md deleted file mode 100644 index dda0527b0d..0000000000 --- a/ui/changelog.d/invitation-row-actions-non-pending.fixed.md +++ /dev/null @@ -1 +0,0 @@ -Edit and Revoke actions are disabled for expired and revoked invitations diff --git a/ui/changelog.d/mute-rule-error-toast-raw-json.fixed.md b/ui/changelog.d/mute-rule-error-toast-raw-json.fixed.md new file mode 100644 index 0000000000..d59b218727 --- /dev/null +++ b/ui/changelog.d/mute-rule-error-toast-raw-json.fixed.md @@ -0,0 +1 @@ +Mute rule creation errors show the API error message instead of the raw JSON:API response body diff --git a/ui/changelog.d/onboarding-invite-step.added.md b/ui/changelog.d/onboarding-invite-step.added.md deleted file mode 100644 index c2220888b5..0000000000 --- a/ui/changelog.d/onboarding-invite-step.added.md +++ /dev/null @@ -1 +0,0 @@ -"Invite your team" step offered once after the first provider is connected, before the onboarding checkpoint, reusing the invitation form tagged with `source=onboarding` diff --git a/ui/changelog.d/provider-connection-check-wait.fixed.md b/ui/changelog.d/provider-connection-check-wait.fixed.md new file mode 100644 index 0000000000..779efa5e74 --- /dev/null +++ b/ui/changelog.d/provider-connection-check-wait.fixed.md @@ -0,0 +1 @@ +Provider connection test no longer reports `Max retries exceeded` for checks that take longer than 30 seconds, such as networks where some AWS endpoints are unreachable; the wait now covers the backend task's full time limit and falls back to the provider's current connection state if it is still exhausted diff --git a/ui/changelog.d/registry-private-cloud.added.md b/ui/changelog.d/registry-private-cloud.added.md deleted file mode 100644 index 0f5b9fb6d3..0000000000 --- a/ui/changelog.d/registry-private-cloud.added.md +++ /dev/null @@ -1 +0,0 @@ -Registry marketplace and external provider onboarding for Private Cloud, with permission-based access independent of billing, confirmed artifact installation, schema-driven credentials, connection checks, and scan launch diff --git a/ui/changelog.d/sidebar-add-provider-action.added.md b/ui/changelog.d/sidebar-add-provider-action.added.md new file mode 100644 index 0000000000..29cc8aec45 --- /dev/null +++ b/ui/changelog.d/sidebar-add-provider-action.added.md @@ -0,0 +1 @@ +Sidebar action reads Add Provider while the tenant has no providers diff --git a/ui/changelog.d/sidebar-mode-hydration.fixed.md b/ui/changelog.d/sidebar-mode-hydration.fixed.md new file mode 100644 index 0000000000..c495c3071d --- /dev/null +++ b/ui/changelog.d/sidebar-mode-hydration.fixed.md @@ -0,0 +1 @@ +Sidebar no longer throws a React hydration error on full page loads for users who last used the chat mode diff --git a/ui/components/compliance/compliance-card.tsx b/ui/components/compliance/compliance-card.tsx index e5bc4f57c0..a3fd54fa69 100644 --- a/ui/components/compliance/compliance-card.tsx +++ b/ui/components/compliance/compliance-card.tsx @@ -50,6 +50,8 @@ interface ComplianceCardProps { * viewer cannot curate the organization's watchlist. */ watchlistAction?: ReactNode; + /** Prowler Cloud tenants without a paid plan cannot download reports. */ + subscriptionOnly?: boolean; } export const ComplianceCard: React.FC = ({ @@ -62,6 +64,7 @@ export const ComplianceCard: React.FC = ({ id, isLatestCisForProvider = false, watchlistAction, + subscriptionOnly = false, }) => { const searchParams = useSearchParams(); const router = useRouter(); @@ -174,6 +177,7 @@ export const ComplianceCard: React.FC = ({ isLatestCisForProvider, )} disabled={hasRegionFilter} + subscriptionOnly={subscriptionOnly} /> {watchlistAction} diff --git a/ui/components/compliance/compliance-download-container.test.tsx b/ui/components/compliance/compliance-download-container.test.tsx index 34f877a789..814a32b2b3 100644 --- a/ui/components/compliance/compliance-download-container.test.tsx +++ b/ui/components/compliance/compliance-download-container.test.tsx @@ -27,6 +27,9 @@ vi.mock("@/components/shadcn", async (importOriginal) => ({ toast: {}, })); +import { useCloudUpgradeStore } from "@/store/cloud-upgrade/store"; +import { PAID_PLAN_UPGRADE_FEATURE } from "@/types/cloud-upgrade"; + import { ComplianceDownloadContainer } from "./compliance-download-container"; describe("ComplianceDownloadContainer", () => { @@ -36,6 +39,7 @@ describe("ComplianceDownloadContainer", () => { beforeEach(() => { vi.clearAllMocks(); + useCloudUpgradeStore.getState().closeCloudUpgrade(); }); it("uses the shared action dropdown for the card actions mode", () => { @@ -136,6 +140,42 @@ describe("ComplianceDownloadContainer", () => { ); }); + it.each([ + { label: /Download CSV report/i, complianceId: "compliance-1" }, + { label: /Download OCSF report/i, complianceId: "dora_2022_2554" }, + { label: /Download PDF report/i, complianceId: "compliance-1" }, + ])( + "should open the paid plan upgrade instead of $label for subscription-only tenants", + async ({ label, complianceId }) => { + // Given + const user = userEvent.setup(); + render( + , + ); + + // When + await user.click( + screen.getByRole("button", { name: "Open compliance export actions" }), + ); + await user.click(screen.getByRole("menuitem", { name: label })); + + // Then + expect(downloadComplianceCsvMock).not.toHaveBeenCalled(); + expect(downloadComplianceOcsfMock).not.toHaveBeenCalled(); + expect(downloadCompliancePdfMock).not.toHaveBeenCalled(); + expect(useCloudUpgradeStore.getState().activeFeature).toBe( + PAID_PLAN_UPGRADE_FEATURE.REPORT_DOWNLOAD, + ); + }, + ); + it("should hide the OCSF action for frameworks without OCSF support", async () => { const user = userEvent.setup(); diff --git a/ui/components/compliance/compliance-download-container.tsx b/ui/components/compliance/compliance-download-container.tsx index 9c49b1df80..520694ceae 100644 --- a/ui/components/compliance/compliance-download-container.tsx +++ b/ui/components/compliance/compliance-download-container.tsx @@ -14,6 +14,7 @@ import { TooltipContent, TooltipTrigger, } from "@/components/shadcn/tooltip"; +import { useReportDownload } from "@/hooks/use-report-download"; import { type ComplianceReportType, isOcsfSupported, @@ -37,6 +38,8 @@ interface ComplianceDownloadContainerProps { /** Custom dropdown trigger (e.g. an outline "Report" button); only used * when presentation is "dropdown". Defaults to the dots icon. */ dropdownTrigger?: React.ReactNode; + /** Prowler Cloud tenants without a paid plan cannot download reports. */ + subscriptionOnly?: boolean; } export const ComplianceDownloadContainer = ({ @@ -49,7 +52,9 @@ export const ComplianceDownloadContainer = ({ buttonWidth = "auto", presentation = "buttons", dropdownTrigger, + subscriptionOnly = false, }: ComplianceDownloadContainerProps) => { + const runReportDownload = useReportDownload(subscriptionOnly); const [isDownloadingCsv, setIsDownloadingCsv] = useState(false); const [isDownloadingOcsf, setIsDownloadingOcsf] = useState(false); const [isDownloadingPdf, setIsDownloadingPdf] = useState(false); @@ -60,35 +65,38 @@ export const ComplianceDownloadContainer = ({ // action everywhere else so the user never hits a guaranteed 404. const ocsfAvailable = isOcsfSupported(complianceId); - const handleDownloadCsv = async () => { - if (isDownloadingCsv) return; - setIsDownloadingCsv(true); - try { - await downloadComplianceCsv(scanId, complianceId, toast); - } finally { - setIsDownloadingCsv(false); - } - }; + const handleDownloadCsv = () => + runReportDownload(async () => { + if (isDownloadingCsv) return; + setIsDownloadingCsv(true); + try { + await downloadComplianceCsv(scanId, complianceId, toast); + } finally { + setIsDownloadingCsv(false); + } + }); - const handleDownloadOcsf = async () => { - if (!ocsfAvailable || isDownloadingOcsf) return; - setIsDownloadingOcsf(true); - try { - await downloadComplianceOcsf(scanId, complianceId, toast); - } finally { - setIsDownloadingOcsf(false); - } - }; + const handleDownloadOcsf = () => + runReportDownload(async () => { + if (!ocsfAvailable || isDownloadingOcsf) return; + setIsDownloadingOcsf(true); + try { + await downloadComplianceOcsf(scanId, complianceId, toast); + } finally { + setIsDownloadingOcsf(false); + } + }); - const handleDownloadPdf = async () => { - if (!reportType || isDownloadingPdf) return; - setIsDownloadingPdf(true); - try { - await downloadCompliancePdf(scanId, reportType, toast); - } finally { - setIsDownloadingPdf(false); - } - }; + const handleDownloadPdf = () => + runReportDownload(async () => { + if (!reportType || isDownloadingPdf) return; + setIsDownloadingPdf(true); + try { + await downloadCompliancePdf(scanId, reportType, toast); + } finally { + setIsDownloadingPdf(false); + } + }); const buttonClassName = cn( "border-button-primary text-button-primary hover:bg-button-primary/10", diff --git a/ui/components/compliance/compliance-overview-grid.tsx b/ui/components/compliance/compliance-overview-grid.tsx index 16500dd5e8..e28e7566f3 100644 --- a/ui/components/compliance/compliance-overview-grid.tsx +++ b/ui/components/compliance/compliance-overview-grid.tsx @@ -47,6 +47,7 @@ interface ComplianceOverviewGridProps { catalogEntries?: ComplianceCatalogEntry[]; providerType?: string; canManageWatchlist?: boolean; + subscriptionOnly?: boolean; } export const ComplianceOverviewGrid = ({ @@ -57,6 +58,7 @@ export const ComplianceOverviewGrid = ({ catalogEntries, providerType, canManageWatchlist = false, + subscriptionOnly = false, }: ComplianceOverviewGridProps) => { const router = useRouter(); const searchParams = useSearchParams(); @@ -146,6 +148,7 @@ export const ComplianceOverviewGrid = ({ id={id} selectedScan={selectedScan} isLatestCisForProvider={latestCisIds?.has(id) ?? false} + subscriptionOnly={subscriptionOnly} watchlistAction={ watchlistEnabled && canManageWatchlist ? ( ({ + downloadComplianceCsvMock: vi.fn(), + downloadComplianceReportPdfMock: vi.fn(), + })); + +vi.mock("next/navigation", () => ({ + useRouter: () => ({ push: vi.fn() }), + useSearchParams: () => new URLSearchParams(), +})); + +vi.mock("@/lib/helper", () => ({ + downloadComplianceCsv: downloadComplianceCsvMock, + downloadComplianceReportPdf: downloadComplianceReportPdfMock, +})); describe("ThreatScoreBadge", () => { const currentDir = path.dirname(fileURLToPath(import.meta.url)); @@ -17,6 +40,44 @@ describe("ThreatScoreBadge", () => { expect(source).not.toContain("ComplianceDownloadContainer"); }); + describe("for subscription-only tenants", () => { + beforeEach(() => { + vi.clearAllMocks(); + useCloudUpgradeStore.getState().closeCloudUpgrade(); + }); + + it.each([/Download CSV report/i, /Download PDF report/i])( + "opens the paid plan upgrade instead of %s", + async (label) => { + // Given + const user = userEvent.setup(); + render( + , + ); + + // When + await user.click( + screen.getByRole("button", { + name: "Open compliance export actions", + }), + ); + await user.click(screen.getByRole("menuitem", { name: label })); + + // Then + expect(downloadComplianceCsvMock).not.toHaveBeenCalled(); + expect(downloadComplianceReportPdfMock).not.toHaveBeenCalled(); + expect(useCloudUpgradeStore.getState().activeFeature).toBe( + PAID_PLAN_UPGRADE_FEATURE.REPORT_DOWNLOAD, + ); + }, + ); + }); + it("does not use Collapsible components", () => { expect(source).not.toContain("Collapsible"); expect(source).not.toContain("CollapsibleTrigger"); diff --git a/ui/components/compliance/threatscore-badge.tsx b/ui/components/compliance/threatscore-badge.tsx index f3e526fc3b..7d4308e504 100644 --- a/ui/components/compliance/threatscore-badge.tsx +++ b/ui/components/compliance/threatscore-badge.tsx @@ -13,6 +13,7 @@ import { ActionDropdownItem, } from "@/components/shadcn/dropdown"; import { Progress } from "@/components/shadcn/progress"; +import { useReportDownload } from "@/hooks/use-report-download"; import { COMPLIANCE_REPORT_TYPES } from "@/lib/compliance/compliance-report-types"; import { getScoreColor, @@ -35,6 +36,8 @@ interface ThreatScoreBadgeProps { provider: string; selectedScan?: ScanEntity; sectionScores?: SectionScores; + /** Prowler Cloud tenants without a paid plan cannot download reports. */ + subscriptionOnly?: boolean; } export const ThreatScoreBadge = ({ @@ -42,8 +45,10 @@ export const ThreatScoreBadge = ({ scanId, provider, sectionScores, + subscriptionOnly = false, }: ThreatScoreBadgeProps) => { const router = useRouter(); + const runReportDownload = useReportDownload(subscriptionOnly); const searchParams = useSearchParams(); const [isDownloadingCsv, setIsDownloadingCsv] = useState(false); const [isDownloadingPdf, setIsDownloadingPdf] = useState(false); @@ -83,29 +88,31 @@ export const ThreatScoreBadge = ({ const pillars = getOrderedPillars(sectionScores); - const handleDownloadCsv = async () => { - if (isDownloadingCsv) return; - setIsDownloadingCsv(true); - try { - await downloadComplianceCsv(scanId, complianceId, toast); - } finally { - setIsDownloadingCsv(false); - } - }; + const handleDownloadCsv = () => + runReportDownload(async () => { + if (isDownloadingCsv) return; + setIsDownloadingCsv(true); + try { + await downloadComplianceCsv(scanId, complianceId, toast); + } finally { + setIsDownloadingCsv(false); + } + }); - const handleDownloadPdf = async () => { - if (isDownloadingPdf) return; - setIsDownloadingPdf(true); - try { - await downloadComplianceReportPdf( - scanId, - COMPLIANCE_REPORT_TYPES.THREATSCORE, - toast, - ); - } finally { - setIsDownloadingPdf(false); - } - }; + const handleDownloadPdf = () => + runReportDownload(async () => { + if (isDownloadingPdf) return; + setIsDownloadingPdf(true); + try { + await downloadComplianceReportPdf( + scanId, + COMPLIANCE_REPORT_TYPES.THREATSCORE, + toast, + ); + } finally { + setIsDownloadingPdf(false); + } + }); return ( diff --git a/ui/components/findings/recheck-resource-action-item.test.tsx b/ui/components/findings/recheck-resource-action-item.test.tsx new file mode 100644 index 0000000000..723e456018 --- /dev/null +++ b/ui/components/findings/recheck-resource-action-item.test.tsx @@ -0,0 +1,89 @@ +import { render, screen } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { beforeEach, describe, expect, it, vi } from "vitest"; + +const { isCloudMock, hasPermissionMock } = vi.hoisted(() => ({ + isCloudMock: vi.fn(() => true), + hasPermissionMock: vi.fn(() => true), +})); + +vi.mock("@/lib/shared/env", () => ({ + isCloud: isCloudMock, +})); + +vi.mock("@/hooks/use-auth", () => ({ + useAuth: () => ({ hasPermission: hasPermissionMock }), +})); + +vi.mock("@/components/shadcn/dropdown", () => ({ + ActionDropdownItem: ({ + label, + onSelect, + }: { + label: string; + onSelect?: () => void; + }) => ( + + ), +})); + +import { usePartialScanStore } from "@/store/partial-scan/store"; + +import { + RECHECK_RESOURCE_LABEL, + RecheckResourceActionItem, +} from "./recheck-resource-action-item"; + +const target = { + providerId: "provider-1", + providerUid: "123456789012", + providerType: "aws", + providerAlias: "prod", + resourceUid: "arn:aws:s3:::bucket", + resourceName: "bucket", +}; + +describe("RecheckResourceActionItem", () => { + beforeEach(() => { + vi.clearAllMocks(); + isCloudMock.mockReturnValue(true); + hasPermissionMock.mockReturnValue(true); + usePartialScanStore.getState().closePartialScan(); + }); + + it("opens the confirmation with the resource as target", async () => { + const user = userEvent.setup(); + render(); + + await user.click( + screen.getByRole("button", { name: RECHECK_RESOURCE_LABEL }), + ); + + expect(usePartialScanStore.getState().activeTarget).toEqual(target); + }); + + it("is hidden outside Prowler Cloud", () => { + isCloudMock.mockReturnValue(false); + render(); + + expect(screen.queryByRole("button")).not.toBeInTheDocument(); + }); + + it("is hidden without the manage_scans permission", () => { + hasPermissionMock.mockReturnValue(false); + render(); + + expect(hasPermissionMock).toHaveBeenCalledWith("manage_scans"); + expect(screen.queryByRole("button")).not.toBeInTheDocument(); + }); + + it("is hidden when the row cannot name a resource", () => { + render( + , + ); + + expect(screen.queryByRole("button")).not.toBeInTheDocument(); + }); +}); diff --git a/ui/components/findings/recheck-resource-action-item.tsx b/ui/components/findings/recheck-resource-action-item.tsx new file mode 100644 index 0000000000..9489a255fb --- /dev/null +++ b/ui/components/findings/recheck-resource-action-item.tsx @@ -0,0 +1,33 @@ +"use client"; + +import { RefreshCw } from "lucide-react"; + +import { ActionDropdownItem } from "@/components/shadcn/dropdown"; +import { usePartialScanTarget } from "@/hooks/use-partial-scan-target"; +import { usePartialScanStore } from "@/store"; +import type { PartialScanTarget } from "@/types/partial-scans"; + +export const RECHECK_RESOURCE_LABEL = "Re-check resource"; + +interface RecheckResourceActionItemProps { + target: Partial | null | undefined; +} + +/** Prowler Cloud only: opens the partial-scan confirmation for one resource. */ +export const RecheckResourceActionItem = ({ + target, +}: RecheckResourceActionItemProps) => { + const resolvedTarget = usePartialScanTarget(target); + const openPartialScan = usePartialScanStore((state) => state.openPartialScan); + + if (!resolvedTarget) return null; + + return ( + } + label={RECHECK_RESOURCE_LABEL} + aria-label={RECHECK_RESOURCE_LABEL} + onSelect={() => openPartialScan(resolvedTarget)} + /> + ); +}; diff --git a/ui/components/findings/recheck-resource-icon-button.test.tsx b/ui/components/findings/recheck-resource-icon-button.test.tsx new file mode 100644 index 0000000000..ae19aa2161 --- /dev/null +++ b/ui/components/findings/recheck-resource-icon-button.test.tsx @@ -0,0 +1,78 @@ +import { render, screen } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { beforeEach, describe, expect, it, vi } from "vitest"; + +const { isCloudMock, hasPermissionMock } = vi.hoisted(() => ({ + isCloudMock: vi.fn(() => true), + hasPermissionMock: vi.fn(() => true), +})); + +vi.mock("@/lib/shared/env", () => ({ + isCloud: isCloudMock, +})); + +vi.mock("@/hooks/use-auth", () => ({ + useAuth: () => ({ hasPermission: hasPermissionMock }), +})); + +import { usePartialScanStore } from "@/store/partial-scan/store"; + +import { RECHECK_RESOURCE_LABEL } from "./recheck-resource-action-item"; +import { RecheckResourceIconButton } from "./recheck-resource-icon-button"; + +const target = { + providerUid: "123456789012", + providerType: "aws", + providerAlias: "prod", + resourceUid: "arn:aws:s3:::bucket", + resourceName: "bucket", +}; + +describe("RecheckResourceIconButton", () => { + beforeEach(() => { + vi.clearAllMocks(); + isCloudMock.mockReturnValue(true); + hasPermissionMock.mockReturnValue(true); + usePartialScanStore.getState().closePartialScan(); + }); + + it("opens the confirmation for the resource without triggering the row", async () => { + // The button sits inside a clickable row that opens the detail drawer. + const user = userEvent.setup(); + const onRowClick = vi.fn(); + render( +
+ +
, + ); + + await user.click( + screen.getByRole("button", { name: RECHECK_RESOURCE_LABEL }), + ); + + expect(usePartialScanStore.getState().activeTarget).toEqual(target); + expect(onRowClick).not.toHaveBeenCalled(); + }); + + it("is always green and pulsing", () => { + render(); + + const button = screen.getByRole("button", { name: RECHECK_RESOURCE_LABEL }); + expect(button.className).toContain("text-button-primary"); + expect(button.className).toContain("animate-pulse"); + }); + + it("is hidden outside Prowler Cloud", () => { + isCloudMock.mockReturnValue(false); + render(); + + expect(screen.queryByRole("button")).not.toBeInTheDocument(); + }); + + it("is hidden without the manage_scans permission", () => { + hasPermissionMock.mockReturnValue(false); + render(); + + expect(screen.queryByRole("button")).not.toBeInTheDocument(); + }); +}); diff --git a/ui/components/findings/recheck-resource-icon-button.tsx b/ui/components/findings/recheck-resource-icon-button.tsx new file mode 100644 index 0000000000..dc57bcce81 --- /dev/null +++ b/ui/components/findings/recheck-resource-icon-button.tsx @@ -0,0 +1,62 @@ +"use client"; + +import { RefreshCw } from "lucide-react"; +import type { MouseEvent } from "react"; + +import { Button } from "@/components/shadcn/button/button"; +import { + Tooltip, + TooltipContent, + TooltipTrigger, +} from "@/components/shadcn/tooltip"; +import { usePartialScanTarget } from "@/hooks/use-partial-scan-target"; +import { cn } from "@/lib/utils"; +import { usePartialScanStore } from "@/store"; +import type { PartialScanTarget } from "@/types/partial-scans"; + +import { RECHECK_RESOURCE_LABEL } from "./recheck-resource-action-item"; + +interface RecheckResourceIconButtonProps { + target: Partial | null | undefined; + className?: string; +} + +// Beside "last seen": always green and pulsing, so the re-check is noticed. +// It is an action, not a notification, so it never settles into a seen state. +export function RecheckResourceIconButton({ + target, + className, +}: RecheckResourceIconButtonProps) { + const resolvedTarget = usePartialScanTarget(target); + const openPartialScan = usePartialScanStore((state) => state.openPartialScan); + + if (!resolvedTarget) return null; + + const handleClick = (event: MouseEvent) => { + // The row itself opens the detail drawer on click. + event.stopPropagation(); + openPartialScan(resolvedTarget); + }; + + return ( + + + + + {RECHECK_RESOURCE_LABEL} + + ); +} diff --git a/ui/components/findings/recheck-resource-modal-host.test.tsx b/ui/components/findings/recheck-resource-modal-host.test.tsx new file mode 100644 index 0000000000..3a1a68d3c0 --- /dev/null +++ b/ui/components/findings/recheck-resource-modal-host.test.tsx @@ -0,0 +1,56 @@ +import { render } from "@testing-library/react"; +import { beforeEach, describe, expect, it, vi } from "vitest"; + +const { RecheckResourceModalMock } = vi.hoisted(() => ({ + RecheckResourceModalMock: vi.fn( + (_props: { + isOpen: boolean; + onOpenChange: (open: boolean) => void; + target: unknown; + }) => null, + ), +})); + +vi.mock("./recheck-resource-modal", () => ({ + RecheckResourceModal: RecheckResourceModalMock, +})); + +import { usePartialScanStore } from "@/store/partial-scan/store"; + +import { RecheckResourceModalHost } from "./recheck-resource-modal-host"; + +const target = { + providerId: "provider-1", + providerUid: "123456789012", + providerType: "aws", + resourceUid: "arn:aws:s3:::bucket", + resourceName: "bucket", +}; + +describe("RecheckResourceModalHost", () => { + beforeEach(() => { + vi.clearAllMocks(); + usePartialScanStore.getState().closePartialScan(); + }); + + it("renders nothing without an active target", () => { + render(); + + expect(RecheckResourceModalMock).not.toHaveBeenCalled(); + }); + + it("mounts the modal for the active target and closes through the store", () => { + usePartialScanStore.getState().openPartialScan(target); + + render(); + + expect(RecheckResourceModalMock).toHaveBeenCalledWith( + expect.objectContaining({ isOpen: true, target }), + undefined, + ); + + RecheckResourceModalMock.mock.calls[0][0].onOpenChange(false); + + expect(usePartialScanStore.getState().activeTarget).toBeNull(); + }); +}); diff --git a/ui/components/findings/recheck-resource-modal-host.tsx b/ui/components/findings/recheck-resource-modal-host.tsx new file mode 100644 index 0000000000..2cfcaaeb69 --- /dev/null +++ b/ui/components/findings/recheck-resource-modal-host.tsx @@ -0,0 +1,25 @@ +"use client"; + +import { usePartialScanStore } from "@/store"; + +import { RecheckResourceModal } from "./recheck-resource-modal"; + +// One global modal, like Jira dispatch: it remounts per target so its state +// (pending, error) never leaks between resources. +export const RecheckResourceModalHost = () => { + const activeTarget = usePartialScanStore((state) => state.activeTarget); + const closePartialScan = usePartialScanStore( + (state) => state.closePartialScan, + ); + + if (!activeTarget) return null; + + return ( + !open && closePartialScan()} + target={activeTarget} + /> + ); +}; diff --git a/ui/components/findings/recheck-resource-modal.test.tsx b/ui/components/findings/recheck-resource-modal.test.tsx new file mode 100644 index 0000000000..52b236dfdc --- /dev/null +++ b/ui/components/findings/recheck-resource-modal.test.tsx @@ -0,0 +1,223 @@ +import { render, screen, waitFor } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { beforeEach, describe, expect, it, vi } from "vitest"; + +const { createPartialScanMock, getProvidersMock, refreshMock, toastMock } = + vi.hoisted(() => ({ + createPartialScanMock: vi.fn(), + getProvidersMock: vi.fn(), + refreshMock: vi.fn(), + toastMock: vi.fn(), + })); + +vi.mock("@/actions/scans", () => ({ + createPartialScan: createPartialScanMock, +})); + +vi.mock("@/actions/providers", () => ({ + getProviders: getProvidersMock, +})); + +vi.mock("next/navigation", () => ({ + useRouter: () => ({ refresh: refreshMock }), +})); + +vi.mock("@/components/shadcn/toast", () => ({ + toast: toastMock, + ToastAction: ({ children }: { children: React.ReactNode }) => <>{children}, +})); + +vi.mock("@/components/shadcn/modal", () => ({ + // The close button stands in for Escape and backdrop clicks. + Modal: ({ + open, + title, + children, + onOpenChange, + }: { + open: boolean; + title: string; + children: React.ReactNode; + onOpenChange: (open: boolean) => void; + }) => + open ? ( +
+ + {children} +
+ ) : null, +})); + +import { PARTIAL_SCAN_LAUNCH_ERROR } from "@/lib/partial-scans"; + +import { + PROVIDER_NOT_FOUND_ERROR, + RECHECK_RESOURCE_SUBMIT_LABEL, + RecheckResourceModal, +} from "./recheck-resource-modal"; + +const target = { + providerId: "provider-1", + providerUid: "123456789012", + providerType: "aws", + providerAlias: "prod", + resourceUid: "arn:aws:s3:::bucket", + resourceName: "bucket", +}; + +const submit = async () => { + const user = userEvent.setup(); + await user.click( + screen.getByRole("button", { name: RECHECK_RESOURCE_SUBMIT_LABEL }), + ); +}; + +describe("RecheckResourceModal", () => { + beforeEach(() => { + vi.clearAllMocks(); + createPartialScanMock.mockResolvedValue({ data: { id: "scan-1" } }); + }); + + it("launches a partial scan for the one resource and closes", async () => { + const onOpenChange = vi.fn(); + render( + , + ); + + await submit(); + + expect(createPartialScanMock).toHaveBeenCalledWith({ + providerId: "provider-1", + resourceUids: ["arn:aws:s3:::bucket"], + }); + expect(getProvidersMock).not.toHaveBeenCalled(); + expect(toastMock).toHaveBeenCalledWith( + expect.objectContaining({ title: "Re-check launched" }), + ); + expect(onOpenChange).toHaveBeenCalledWith(false); + expect(refreshMock).toHaveBeenCalled(); + }); + + it("resolves the provider id from its uid and type when the row lacks it", async () => { + getProvidersMock.mockResolvedValue({ + data: [ + { id: "aws-1", attributes: { uid: "123456789012", provider: "aws" } }, + ], + }); + const { providerId: _omitted, ...rowTarget } = target; + render( + , + ); + + await submit(); + + expect(getProvidersMock).toHaveBeenCalledWith({ + filters: { "filter[uid]": "123456789012", "filter[provider]": "aws" }, + }); + expect(createPartialScanMock).toHaveBeenCalledWith({ + providerId: "aws-1", + resourceUids: ["arn:aws:s3:::bucket"], + }); + }); + + it("explains when the provider cannot be found and launches nothing", async () => { + getProvidersMock.mockResolvedValue({ data: [] }); + const { providerId: _omitted, ...rowTarget } = target; + render( + , + ); + + await submit(); + + expect(await screen.findByRole("alert")).toHaveTextContent( + PROVIDER_NOT_FOUND_ERROR, + ); + expect(createPartialScanMock).not.toHaveBeenCalled(); + }); + + it("keeps the modal open and shows the API reason when the re-check is refused", async () => { + // The 409 detail already tells the user what to do, so it is shown verbatim. + createPartialScanMock.mockResolvedValue({ + error: + "A scan is already running for this provider. Re-check these resources once it finishes.", + status: 409, + }); + const onOpenChange = vi.fn(); + render( + , + ); + + await submit(); + + expect(await screen.findByRole("alert")).toHaveTextContent( + "A scan is already running for this provider.", + ); + expect(onOpenChange).not.toHaveBeenCalled(); + expect(toastMock).not.toHaveBeenCalled(); + expect(refreshMock).not.toHaveBeenCalled(); + await waitFor(() => + expect( + screen.getByRole("button", { name: RECHECK_RESOURCE_SUBMIT_LABEL }), + ).toBeEnabled(), + ); + }); + + it("treats a response without a scan id as a failed launch", async () => { + // An empty 2xx becomes { success: true } and there is no scan to follow. + createPartialScanMock.mockResolvedValue({ success: true, status: 202 }); + const onOpenChange = vi.fn(); + render( + , + ); + + await submit(); + + expect(await screen.findByRole("alert")).toHaveTextContent( + PARTIAL_SCAN_LAUNCH_ERROR, + ); + expect(toastMock).not.toHaveBeenCalled(); + expect(onOpenChange).not.toHaveBeenCalled(); + }); + + it("ignores a dismiss while the request is in flight", async () => { + let resolveLaunch!: (value: unknown) => void; + createPartialScanMock.mockReturnValue( + new Promise((resolve) => { + resolveLaunch = resolve; + }), + ); + const onOpenChange = vi.fn(); + const user = userEvent.setup(); + render( + , + ); + await user.click( + screen.getByRole("button", { name: RECHECK_RESOURCE_SUBMIT_LABEL }), + ); + + await user.click(screen.getByRole("button", { name: "Dismiss" })); + + expect(onOpenChange).not.toHaveBeenCalled(); + + resolveLaunch({ data: { id: "scan-1" } }); + await waitFor(() => expect(onOpenChange).toHaveBeenCalledWith(false)); + }); +}); diff --git a/ui/components/findings/recheck-resource-modal.tsx b/ui/components/findings/recheck-resource-modal.tsx new file mode 100644 index 0000000000..29177d8bfa --- /dev/null +++ b/ui/components/findings/recheck-resource-modal.tsx @@ -0,0 +1,176 @@ +"use client"; + +import Link from "next/link"; +import { useRouter } from "next/navigation"; +import { type FormEvent, useState } from "react"; + +import { getProviders } from "@/actions/providers"; +import { createPartialScan } from "@/actions/scans"; +import { Button } from "@/components/shadcn"; +import { Modal } from "@/components/shadcn/modal"; +import { toast, ToastAction } from "@/components/shadcn/toast"; +import { + findProviderIdForTarget, + getPartialScanErrorMessage, + PARTIAL_SCAN_LAUNCH_ERROR, +} from "@/lib/partial-scans"; +import type { PartialScanTarget } from "@/types/partial-scans"; + +export const RECHECK_RESOURCE_SUBMIT_LABEL = "Re-check resource"; +export const PROVIDER_NOT_FOUND_ERROR = + "We couldn't find the provider of this resource. Refresh the page and try again."; + +interface RecheckResourceModalProps { + isOpen: boolean; + onOpenChange: (open: boolean) => void; + target: PartialScanTarget; +} + +const hasDisplayName = (name: string) => name.trim() !== "" && name !== "-"; + +export function RecheckResourceModal({ + isOpen, + onOpenChange, + target, +}: RecheckResourceModalProps) { + const router = useRouter(); + const [isPending, setIsPending] = useState(false); + const [error, setError] = useState(null); + + const resourceLabel = hasDisplayName(target.resourceName) + ? target.resourceName + : target.resourceUid; + + // Drill-down rows only know the provider by uid + type; the API needs its id. + const resolveProviderId = async () => { + if (target.providerId) return target.providerId; + + const response = await getProviders({ + filters: { + "filter[uid]": target.providerUid, + "filter[provider]": target.providerType, + }, + }); + + return findProviderIdForTarget(response?.data ?? [], target); + }; + + const handleSubmit = async (event: FormEvent) => { + event.preventDefault(); + if (isPending) return; + + setIsPending(true); + setError(null); + + try { + const providerId = await resolveProviderId(); + if (!providerId) { + setError(PROVIDER_NOT_FOUND_ERROR); + return; + } + + const result = await createPartialScan({ + providerId, + resourceUids: [target.resourceUid], + }); + const errorMessage = getPartialScanErrorMessage(result); + if (errorMessage) { + setError(errorMessage); + return; + } + // An empty 2xx has no scan to follow, so it is not a launch. + if (!result?.data?.id) { + setError(PARTIAL_SCAN_LAUNCH_ERROR); + return; + } + + toast({ + title: "Re-check launched", + description: `Only ${resourceLabel} is being re-checked. Its findings update once the scan completes.`, + action: ( + + View scan + + ), + }); + onOpenChange(false); + router.refresh(); + } catch { + setError(PARTIAL_SCAN_LAUNCH_ERROR); + } finally { + setIsPending(false); + } + }; + + return ( + { + if (!open && isPending) return; + onOpenChange(open); + }} + title="Re-check this resource" + description="Run a partial scan on a single resource instead of the whole provider." + size="lg" + > +
+
+

+ Resource +

+

+ {resourceLabel} +

+ {resourceLabel !== target.resourceUid && ( +

+ {target.resourceUid} +

+ )} + {target.providerAlias && ( +

+ Provider:{" "} + + {target.providerAlias} + +

+ )} +
+ +
+

+ Only the checks that last reported on this resource run again. Its + findings update when the scan completes; every other resource keeps + the results of the latest full scan. +

+

+ Overviews and compliance do not change until the next full scan. A + re-check is refused while the provider has a scan running or queued. +

+
+ + {error && ( +

+ {error} +

+ )} + +
+ + +
+
+
+ ); +} diff --git a/ui/components/findings/table/column-finding-resources.test.tsx b/ui/components/findings/table/column-finding-resources.test.tsx index ab2f0d6070..8c31599cbe 100644 --- a/ui/components/findings/table/column-finding-resources.test.tsx +++ b/ui/components/findings/table/column-finding-resources.test.tsx @@ -189,6 +189,14 @@ vi.mock("@/lib/shared/env", () => ({ isCloud: isCloudMock, })); +// The re-check menu item reads the session for manage_scans; grant it here. +vi.mock("@/hooks/use-auth", () => ({ + useAuth: () => ({ + permissions: { manage_scans: true }, + hasPermission: () => true, + }), +})); + vi.mock("./lighthouse-skills-launch", async (importOriginal) => { const actual = await importOriginal(); @@ -208,6 +216,7 @@ vi.mock("./notification-indicator", () => ({ })); import { useJiraDispatchStore } from "@/store/jira-dispatch/store"; +import { usePartialScanStore } from "@/store/partial-scan/store"; import type { FindingResourceRow } from "@/types"; import { FINDING_TRIAGE_DISABLED_REASON, @@ -312,6 +321,24 @@ function renderResourceActionsCell({ render(
{CellComponent({ row: { original: resource, index: 0 } })}
); } +function renderLastSeenCell(resource: FindingResourceRow = makeResource()) { + const columns = getColumnFindingResources({ + rowSelection: {}, + selectableRowCount: 1, + }); + const column = columns.find( + (col) => (col as { id?: string }).id === "lastSeen", + ); + if (!column?.cell) { + throw new Error("lastSeen column not found"); + } + const CellComponent = column.cell as (props: { + row: { original: FindingResourceRow; index: number }; + }) => ReactNode; + + render(
{CellComponent({ row: { original: resource, index: 0 } })}
); +} + describe("column-finding-resources", () => { beforeEach(() => { vi.clearAllMocks(); @@ -320,6 +347,55 @@ describe("column-finding-resources", () => { useJiraDispatchStore.getState().closeJiraDispatch(); }); + it("offers a Cloud re-check identified by provider uid and type", async () => { + // Given — drill-down rows carry no provider id, only its uid and type + const user = userEvent.setup(); + isCloudMock.mockReturnValue(true); + usePartialScanStore.getState().closePartialScan(); + renderResourceActionsCell(); + + // When + await user.click(screen.getByRole("button", { name: "Re-check resource" })); + + // Then + expect(usePartialScanStore.getState().activeTarget).toEqual({ + providerUid: "123456789", + providerType: "aws", + providerAlias: "production", + resourceUid: "arn:aws:s3:::my-bucket", + resourceName: "my-bucket", + }); + }); + + it("offers the re-check beside Last seen on Cloud rows", async () => { + // The quiet icon next to the timestamp opens the same confirmation as ⋮. + const user = userEvent.setup(); + isCloudMock.mockReturnValue(true); + usePartialScanStore.getState().closePartialScan(); + renderLastSeenCell(); + + await user.click(screen.getByRole("button", { name: "Re-check resource" })); + + expect(usePartialScanStore.getState().activeTarget).toEqual( + expect.objectContaining({ resourceUid: "arn:aws:s3:::my-bucket" }), + ); + }); + + it("keeps Last seen plain outside Prowler Cloud", () => { + renderLastSeenCell(); + + expect(screen.getByText("2024-01-01T00:00:00Z")).toBeInTheDocument(); + expect(screen.queryByRole("button")).not.toBeInTheDocument(); + }); + + it("hides the re-check outside Prowler Cloud", () => { + renderResourceActionsCell(); + + expect( + screen.queryByRole("button", { name: "Re-check resource" }), + ).not.toBeInTheDocument(); + }); + it("opens the finding drawer and launches a row skill with full context", async () => { // Given const user = userEvent.setup(); diff --git a/ui/components/findings/table/column-finding-resources.tsx b/ui/components/findings/table/column-finding-resources.tsx index c5cc656ebb..31a5e25e74 100644 --- a/ui/components/findings/table/column-finding-resources.tsx +++ b/ui/components/findings/table/column-finding-resources.tsx @@ -6,6 +6,8 @@ import { useContext, useState } from "react"; import { JiraDispatchActionItem } from "@/components/findings/jira-dispatch-action-item"; import { MuteFindingsModal } from "@/components/findings/mute-findings-modal"; +import { RecheckResourceActionItem } from "@/components/findings/recheck-resource-action-item"; +import { RecheckResourceIconButton } from "@/components/findings/recheck-resource-icon-button"; import { Checkbox } from "@/components/shadcn"; import { ActionDropdown, @@ -71,6 +73,16 @@ const buildResourceFindingItem = (resource: FindingResourceRow) => region: resource.region, }); +// Shared by the ⋮ menu item and the "Last seen" icon so both open the same +// confirmation for the same resource. +const buildResourceRecheckTarget = (resource: FindingResourceRow) => ({ + providerUid: resource.providerUid, + providerType: resource.providerType, + providerAlias: resource.providerAlias, + resourceUid: resource.resourceUid, + resourceName: resource.resourceName, +}); + const ResourceRowActions = ({ row, findingTitle, @@ -199,6 +211,9 @@ const ResourceRowActions = ({ })} payload={jiraPayload} /> + {isCloud() && ( { @@ -384,7 +399,12 @@ export function getColumnFindingResources({ ), cell: ({ row }) => ( - + + + + ), enableSorting: false, diff --git a/ui/components/findings/table/data-table-row-actions.test.tsx b/ui/components/findings/table/data-table-row-actions.test.tsx index daad634da2..7483c040e7 100644 --- a/ui/components/findings/table/data-table-row-actions.test.tsx +++ b/ui/components/findings/table/data-table-row-actions.test.tsx @@ -3,6 +3,7 @@ import userEvent from "@testing-library/user-event"; import { beforeEach, describe, expect, it, vi } from "vitest"; import { useJiraDispatchStore } from "@/store/jira-dispatch/store"; +import { usePartialScanStore } from "@/store/partial-scan/store"; import { FINDING_TRIAGE_DISABLED_REASON, FINDING_TRIAGE_STATUS, @@ -44,6 +45,14 @@ vi.mock("@/lib/shared/env", () => ({ isCloud: isCloudMock, })); +// The re-check menu item reads the session for manage_scans; grant it here. +vi.mock("@/hooks/use-auth", () => ({ + useAuth: () => ({ + permissions: { manage_scans: true }, + hasPermission: () => true, + }), +})); + vi.mock("./lighthouse-skills-launch", async (importOriginal) => { const actual = await importOriginal(); @@ -170,6 +179,58 @@ describe("DataTableRowActions", () => { useJiraDispatchStore.getState().closeJiraDispatch(); }); + it("offers a Cloud re-check of the row's resource with its provider id", async () => { + // Given — a flat finding row expanded with its resource and provider + const user = userEvent.setup(); + isCloudMock.mockReturnValue(true); + usePartialScanStore.getState().closePartialScan(); + render( + , + ); + + // When + await user.click(screen.getByRole("button", { name: "Re-check resource" })); + + // Then + expect(usePartialScanStore.getState().activeTarget).toEqual({ + providerId: "provider-1", + providerUid: "123", + providerType: "aws", + providerAlias: "prod", + resourceUid: "arn:aws:s3:::my-bucket", + resourceName: "my-bucket", + }); + }); + + it("does not offer a re-check outside Prowler Cloud", () => { + render( + , + ); + + expect( + screen.queryByRole("button", { name: "Re-check resource" }), + ).not.toBeInTheDocument(); + }); + it("launches a Lighthouse skill from the row submenu with finding context", async () => { // Given const user = userEvent.setup(); @@ -212,6 +273,10 @@ describe("DataTableRowActions", () => { ); expect(screen.queryByText("Lighthouse Skills")).not.toBeInTheDocument(); + // A group spans many resources, so a single-resource re-check is not offered. + expect( + screen.queryByRole("button", { name: "Re-check resource" }), + ).not.toBeInTheDocument(); expect( screen.queryByRole("button", { name: "Triage Decision" }), ).not.toBeInTheDocument(); diff --git a/ui/components/findings/table/data-table-row-actions.tsx b/ui/components/findings/table/data-table-row-actions.tsx index 3033003e57..ecef7f1da3 100644 --- a/ui/components/findings/table/data-table-row-actions.tsx +++ b/ui/components/findings/table/data-table-row-actions.tsx @@ -7,6 +7,7 @@ import { useContext, useState } from "react"; import { JiraDispatchActionItem } from "@/components/findings/jira-dispatch-action-item"; import { MuteFindingsModal } from "@/components/findings/mute-findings-modal"; +import { RecheckResourceActionItem } from "@/components/findings/recheck-resource-action-item"; import { ActionDropdown, ActionDropdownItem, @@ -53,12 +54,17 @@ export interface FindingRowData { resource?: { attributes?: { name?: string; + uid?: string; }; }; provider?: { + // Expanded included providers carry their id at the top level. + id?: string; + data?: { id?: string }; attributes?: { alias?: string; provider?: string; + uid?: string; }; }; }; @@ -246,6 +252,20 @@ export function DataTableRowActions({ router.refresh(); }; + // Partial scans re-check one resource, so group rows never offer them. + const recheckTarget = isGroup + ? null + : { + providerId: + finding.relationships?.provider?.id ?? + finding.relationships?.provider?.data?.id, + providerUid: finding.relationships?.provider?.attributes?.uid, + providerType: finding.relationships?.provider?.attributes?.provider, + providerAlias: finding.relationships?.provider?.attributes?.alias, + resourceUid: finding.relationships?.resource?.attributes?.uid, + resourceName: finding.relationships?.resource?.attributes?.name, + }; + const launchSkill = useLighthouseSkillLaunch(); const launchPrompt = useLighthousePromptLaunch(); // Skills are finding-level only: group rows carry check ids, not finding @@ -300,6 +320,7 @@ export function DataTableRowActions({ onSelect={handleMuteClick} /> + {isCloud() && !isGroup && ( ({ isCloud: mockIsCloud, })); +// The re-check menu item reads the session for manage_scans; grant it here. +vi.mock("@/hooks/use-auth", () => ({ + useAuth: () => ({ + permissions: { manage_scans: true }, + hasPermission: () => true, + }), +})); + vi.mock("@/app/(prowler)/lighthouse/_lib/panel-chat-store", () => ({ requestPanelChatMessage: mockRequestPanelChatMessage, requestPanelSkillLaunch: mockRequestPanelSkillLaunch, @@ -543,6 +551,7 @@ vi.mock("../../muted", () => ({ // --------------------------------------------------------------------------- import type { ResourceDrawerFinding } from "@/actions/findings"; +import { usePartialScanStore } from "@/store/partial-scan/store"; import { SIDE_PANEL_TAB, useSidePanelStore } from "@/store/side-panel"; import type { FindingResourceRow } from "@/types"; import type { FindingComplianceFramework } from "@/types/compliance-watchlist"; @@ -655,6 +664,7 @@ const mockFinding: ResourceDrawerFinding = { resourceGroup: "default", resourceDetails: null, resourceMetadata: null, + providerId: "provider-1", providerType: "aws", providerAlias: "prod", providerUid: "123456789", @@ -2241,3 +2251,80 @@ describe("ResourceDetailDrawerContent — Metadata tab", () => { ).not.toBeInTheDocument(); }); }); + +describe("ResourceDetailDrawerContent — re-check resource", () => { + beforeEach(() => { + usePartialScanStore.getState().closePartialScan(); + }); + + it("should offer a re-check of the current resource with its provider id", async () => { + // Given — the drawer finding was fetched with scan.provider, so it knows the id + const user = userEvent.setup(); + mockIsCloud.mockReturnValue(true); + render( + , + ); + + // When + await user.click( + within(screen.getByRole("menu", { name: "Resource actions" })).getByRole( + "button", + { name: "Re-check resource" }, + ), + ); + + // Then + expect(usePartialScanStore.getState().activeTarget).toEqual({ + providerId: "provider-1", + providerUid: "123456789", + providerType: "aws", + providerAlias: "prod", + resourceUid: "arn:aws:s3:::bucket", + resourceName: "my-bucket", + }); + }); + + it("should offer the re-check beside Last detected with the provider id", async () => { + const user = userEvent.setup(); + mockIsCloud.mockReturnValue(true); + render( + , + ); + + const metadataRow = screen.getByTestId( + "resource-detail-secondary-metadata-row", + ); + await user.click( + within(metadataRow).getByRole("button", { name: "Re-check resource" }), + ); + + expect(usePartialScanStore.getState().activeTarget).toEqual( + expect.objectContaining({ + providerId: "provider-1", + resourceUid: "arn:aws:s3:::bucket", + }), + ); + }); +}); diff --git a/ui/components/findings/table/resource-detail-drawer/resource-detail-drawer-content.tsx b/ui/components/findings/table/resource-detail-drawer/resource-detail-drawer-content.tsx index 1f8d3cb75b..a02617e3c1 100644 --- a/ui/components/findings/table/resource-detail-drawer/resource-detail-drawer-content.tsx +++ b/ui/components/findings/table/resource-detail-drawer/resource-detail-drawer-content.tsx @@ -27,6 +27,8 @@ import { import { JiraDispatchActionItem } from "@/components/findings/jira-dispatch-action-item"; import { MarkdownContainer } from "@/components/findings/markdown-container"; import { MuteFindingsModal } from "@/components/findings/mute-findings-modal"; +import { RecheckResourceActionItem } from "@/components/findings/recheck-resource-action-item"; +import { RecheckResourceIconButton } from "@/components/findings/recheck-resource-icon-button"; import { getComplianceIcon } from "@/components/icons"; import { Badge, @@ -408,6 +410,14 @@ export function ResourceDetailDrawerContent({ const resourceRegionLabel = resourceRegion || "-"; const firstSeenAt = currentResource?.firstSeenAt ?? f?.firstSeenAt ?? null; const lastSeenAt = currentResource?.lastSeenAt ?? f?.updatedAt ?? null; + const recheckTarget = { + providerId: f?.providerId, + providerUid, + providerType, + providerAlias, + resourceUid, + resourceName, + }; const hasPrev = currentIndex > 0; const hasNext = currentIndex < totalResources - 1; const selectedScanIds = parseSelectedScanIds( @@ -726,7 +736,12 @@ export function ResourceDetailDrawerContent({ variant="compact" className="min-w-0" > - + + + {f && ( + + )} + + {externalResourceTarget && ( } diff --git a/ui/components/findings/table/resource-detail-drawer/resource-detail-drawer.test.tsx b/ui/components/findings/table/resource-detail-drawer/resource-detail-drawer.test.tsx index 20761aeb36..12e95fd8fe 100644 --- a/ui/components/findings/table/resource-detail-drawer/resource-detail-drawer.test.tsx +++ b/ui/components/findings/table/resource-detail-drawer/resource-detail-drawer.test.tsx @@ -178,6 +178,7 @@ function drawerFinding( resourceGroup: "storage", resourceDetails: null, resourceMetadata: null, + providerId: "provider-1", providerType: "aws", providerAlias: "Production", providerUid: "123456789012", diff --git a/ui/components/findings/table/resource-detail-drawer/use-resource-detail-drawer.test.ts b/ui/components/findings/table/resource-detail-drawer/use-resource-detail-drawer.test.ts index 49c0118ce9..88495ce62c 100644 --- a/ui/components/findings/table/resource-detail-drawer/use-resource-detail-drawer.test.ts +++ b/ui/components/findings/table/resource-detail-drawer/use-resource-detail-drawer.test.ts @@ -110,6 +110,7 @@ function makeDrawerFinding( resourceGroup: "default", resourceDetails: null, resourceMetadata: null, + providerId: "provider-1", providerType: "aws", providerAlias: "prod", providerUid: "123", diff --git a/ui/components/layout/app-sidebar/app-sidebar-content.test.tsx b/ui/components/layout/app-sidebar/app-sidebar-content.test.tsx index 8e33a31b86..055d173ae3 100644 --- a/ui/components/layout/app-sidebar/app-sidebar-content.test.tsx +++ b/ui/components/layout/app-sidebar/app-sidebar-content.test.tsx @@ -13,11 +13,13 @@ const { openCloudUpgradeMock, openLaunchScanModalMock, pathnameValue, + permissionsValue, pushMock, } = vi.hoisted(() => ({ openCloudUpgradeMock: vi.fn(), openLaunchScanModalMock: vi.fn(), pathnameValue: { current: "/findings" }, + permissionsValue: { current: {} as Record }, pushMock: vi.fn(), })); @@ -27,7 +29,7 @@ vi.mock("next/navigation", () => ({ })); vi.mock("@/hooks", () => ({ - useAuth: () => ({ permissions: {} }), + useAuth: () => ({ permissions: permissionsValue.current }), })); vi.mock("@/hooks/use-runtime-config", () => ({ @@ -54,11 +56,16 @@ vi.mock("@/app/(prowler)/lighthouse/_components/navigation", () => ({ describe("AppSidebarContent", () => { beforeEach(() => { pathnameValue.current = "/findings"; + permissionsValue.current = { manage_providers: true }; pushMock.mockClear(); openCloudUpgradeMock.mockClear(); openLaunchScanModalMock.mockClear(); useAppSidebarMode.setState({ mode: APP_SIDEBAR_MODE.BROWSE }); - useUIStore.setState({ registryEligible: false }); + useUIStore.setState({ + registryEligible: false, + hasProviders: false, + hasProvidersResolved: false, + }); }); afterEach(() => { @@ -163,6 +170,67 @@ describe("AppSidebarContent", () => { expect(pushMock).toHaveBeenCalledWith("/lighthouse"); }); + it("offers Add Provider instead of Launch Scan once the tenant is known to have no providers", () => { + // Given + vi.stubEnv("UI_CLOUD_ENABLED", "false"); + useUIStore.setState({ hasProviders: false, hasProvidersResolved: true }); + + // When + render(); + + // Then + expect(screen.getByRole("link", { name: "Add Provider" })).toHaveAttribute( + "href", + "/providers?addProvider=true&addProviderSource=sidebar_cta", + ); + expect( + screen.queryByRole("link", { name: "Launch Scan" }), + ).not.toBeInTheDocument(); + }); + + it("keeps Launch Scan for a user who cannot add providers", () => { + // Given: an empty list may only mean limited visibility. + vi.stubEnv("UI_CLOUD_ENABLED", "false"); + permissionsValue.current = { manage_providers: false }; + useUIStore.setState({ hasProviders: false, hasProvidersResolved: true }); + + // When + render(); + + // Then + expect(screen.getByRole("link", { name: "Launch Scan" })).toBeVisible(); + expect( + screen.queryByRole("link", { name: "Add Provider" }), + ).not.toBeInTheDocument(); + }); + + it("keeps Launch Scan while the provider count is still unresolved", () => { + // Given + vi.stubEnv("UI_CLOUD_ENABLED", "false"); + useUIStore.setState({ hasProviders: false, hasProvidersResolved: false }); + + // When + render(); + + // Then + expect(screen.getByRole("link", { name: "Launch Scan" })).toBeVisible(); + expect( + screen.queryByRole("link", { name: "Add Provider" }), + ).not.toBeInTheDocument(); + }); + + it("keeps Launch Scan for a tenant that already has providers", () => { + // Given + vi.stubEnv("UI_CLOUD_ENABLED", "false"); + useUIStore.setState({ hasProviders: true, hasProvidersResolved: true }); + + // When + render(); + + // Then + expect(screen.getByRole("link", { name: "Launch Scan" })).toBeVisible(); + }); + it("opens the current scan modal instead of navigating from the scans route", async () => { // Given vi.stubEnv("UI_CLOUD_ENABLED", "true"); diff --git a/ui/components/layout/app-sidebar/app-sidebar-content.tsx b/ui/components/layout/app-sidebar/app-sidebar-content.tsx index f157e2f7ae..e98e0cb9b3 100644 --- a/ui/components/layout/app-sidebar/app-sidebar-content.tsx +++ b/ui/components/layout/app-sidebar/app-sidebar-content.tsx @@ -10,7 +10,7 @@ import { useRuntimeConfig } from "@/hooks/use-runtime-config"; import { isCloud } from "@/lib/shared/env"; import { useUIStore } from "@/store/ui/store"; -import { useAppSidebarMode } from "./app-sidebar-mode-store"; +import { useHydratedAppSidebarMode } from "./app-sidebar-mode-store"; import { AppSidebarModeToggle } from "./app-sidebar-mode-toggle"; import { LaunchScanAction } from "./launch-scan-action"; import { getNavigationConfig } from "./navigation-config"; @@ -28,7 +28,7 @@ export function AppSidebarContent({ onSelect }: AppSidebarContentProps) { // One-time server decision per request, seeded by the root layout. const registryEligible = useUIStore((state) => state.registryEligible); const { apiDocsUrl, cloudBillingEnabled } = useRuntimeConfig(); - const mode = useAppSidebarMode((state) => state.mode); + const mode = useHydratedAppSidebarMode(); const isCloudEnvironment = isCloud(); const sections = getNavigationConfig({ pathname, diff --git a/ui/components/layout/app-sidebar/app-sidebar-mode-store.test.ts b/ui/components/layout/app-sidebar/app-sidebar-mode-store.test.ts index ba85902736..d019319ef2 100644 --- a/ui/components/layout/app-sidebar/app-sidebar-mode-store.test.ts +++ b/ui/components/layout/app-sidebar/app-sidebar-mode-store.test.ts @@ -1,8 +1,10 @@ +import { renderHook } from "@testing-library/react"; import { beforeEach, describe, expect, it } from "vitest"; import { migrateAppSidebarState, useAppSidebarMode, + useHydratedAppSidebarMode, } from "./app-sidebar-mode-store"; import { APP_SIDEBAR_MODE } from "./types"; @@ -70,4 +72,21 @@ describe("app sidebar mode store", () => { // Then expect(useAppSidebarMode.getState().mode).toBe(APP_SIDEBAR_MODE.CHAT); }); + + it("answers browse on the first render and the persisted mode once mounted", () => { + // Given — a tenant that last used the chat; the server rendered browse. + useAppSidebarMode.setState({ mode: APP_SIDEBAR_MODE.CHAT }); + const renders: string[] = []; + + // When + renderHook(() => { + const mode = useHydratedAppSidebarMode(); + renders.push(mode); + return mode; + }); + + // Then — the hydrating render matches the HTML, the next one the store. + expect(renders[0]).toBe(APP_SIDEBAR_MODE.BROWSE); + expect(renders[renders.length - 1]).toBe(APP_SIDEBAR_MODE.CHAT); + }); }); diff --git a/ui/components/layout/app-sidebar/app-sidebar-mode-store.ts b/ui/components/layout/app-sidebar/app-sidebar-mode-store.ts index fc076c0296..b9cf055319 100644 --- a/ui/components/layout/app-sidebar/app-sidebar-mode-store.ts +++ b/ui/components/layout/app-sidebar/app-sidebar-mode-store.ts @@ -1,6 +1,8 @@ import { create } from "zustand"; import { createJSONStorage, persist } from "zustand/middleware"; +import { useStore } from "@/hooks/use-store"; + import { APP_SIDEBAR_MODE, type AppSidebarMode } from "./types"; interface PersistedAppSidebarState { @@ -52,3 +54,15 @@ export const useAppSidebarMode = create()( }, ), ); + +// The persisted mode is only known in the browser: `persist` rehydrates from +// localStorage before the first client render, while the server always +// rendered `browse`. Reading the store directly makes a tenant that last used +// the chat hydrate a different sidebar than the one in the HTML (React #418). +// This read answers `browse` until mounted, then the persisted value. +export function useHydratedAppSidebarMode(): AppSidebarMode { + return ( + useStore(useAppSidebarMode, (state) => state.mode) ?? + APP_SIDEBAR_MODE.BROWSE + ); +} diff --git a/ui/components/layout/app-sidebar/app-sidebar-mode-toggle.tsx b/ui/components/layout/app-sidebar/app-sidebar-mode-toggle.tsx index 2f62956f51..6f1cc10bcf 100644 --- a/ui/components/layout/app-sidebar/app-sidebar-mode-toggle.tsx +++ b/ui/components/layout/app-sidebar/app-sidebar-mode-toggle.tsx @@ -18,7 +18,10 @@ import { import { useCloudUpgradeStore } from "@/store"; import { CLOUD_UPGRADE_FEATURE } from "@/types/cloud-upgrade"; -import { useAppSidebarMode } from "./app-sidebar-mode-store"; +import { + useAppSidebarMode, + useHydratedAppSidebarMode, +} from "./app-sidebar-mode-store"; import { APP_SIDEBAR_MODE, type AppSidebarMode, @@ -49,7 +52,7 @@ export function AppSidebarModeToggle({ }: AppSidebarModeToggleProps) { const router = useRouter(); const pathname = usePathname(); - const mode = useAppSidebarMode((state) => state.mode); + const mode = useHydratedAppSidebarMode(); const setMode = useAppSidebarMode((state) => state.setMode); const openCloudUpgrade = useCloudUpgradeStore( (state) => state.openCloudUpgrade, diff --git a/ui/components/layout/app-sidebar/launch-scan-action.tsx b/ui/components/layout/app-sidebar/launch-scan-action.tsx index 1da429811f..26f0e55002 100644 --- a/ui/components/layout/app-sidebar/launch-scan-action.tsx +++ b/ui/components/layout/app-sidebar/launch-scan-action.tsx @@ -1,15 +1,28 @@ "use client"; -import { ScanLine } from "lucide-react"; +import { CloudCog, ScanLine } from "lucide-react"; import Link from "next/link"; import { usePathname } from "next/navigation"; import { Button } from "@/components/shadcn/button/button"; +import { useAuth } from "@/hooks"; +import { + dispatchProviderFunnel, + PROVIDER_FUNNEL_STEP, + SIDEBAR_CTA_VARIANT, + WIZARD_OPEN_SOURCE, +} from "@/lib/provider-funnel/provider-funnel-events"; +import { buildAddProviderHref } from "@/lib/providers-navigation"; import { LAUNCH_SCAN_HREF } from "@/lib/scans-navigation"; import { useScansStore } from "@/store"; +import { useUIStore } from "@/store/ui/store"; import type { AppSidebarSelectionHandler } from "./types"; +const ADD_PROVIDER_FROM_SIDEBAR_HREF = buildAddProviderHref( + WIZARD_OPEN_SOURCE.SIDEBAR_CTA, +); + interface LaunchScanActionProps { onSelect?: AppSidebarSelectionHandler; } @@ -28,8 +41,43 @@ export function LaunchScanAction({ onSelect }: LaunchScanActionProps) { const openLaunchScanModal = useScansStore( (state) => state.openLaunchScanModal, ); + const { permissions } = useAuth(); + // Only a confirmed empty tenant swaps the action; an unresolved count keeps Launch Scan. + const hasNoProviders = useUIStore( + (state) => state.hasProvidersResolved && !state.hasProviders, + ); + // Without the permission an empty list may just be limited visibility. + const needsFirstProvider = + hasNoProviders && permissions.manage_providers === true; const isScansPage = pathname.startsWith("/scans"); + if (needsFirstProvider) { + return ( + + ); + } + + const trackLaunchScan = () => + dispatchProviderFunnel({ + step: PROVIDER_FUNNEL_STEP.SIDEBAR_CTA_CLICKED, + variant: SIDEBAR_CTA_VARIANT.LAUNCH_SCAN, + }); + if (isScansPage) { return ( diff --git a/ui/components/layout/main-layout/main-layout.test.tsx b/ui/components/layout/main-layout/main-layout.test.tsx index 6198f0daf9..27e4dd6e8b 100644 --- a/ui/components/layout/main-layout/main-layout.test.tsx +++ b/ui/components/layout/main-layout/main-layout.test.tsx @@ -15,6 +15,12 @@ vi.mock("@/components/findings/jira-dispatch-modal-host", () => ({ JiraDispatchModalHost: () =>
, })); +vi.mock("@/components/findings/recheck-resource-modal-host", () => ({ + RecheckResourceModalHost: () => ( +
+ ), +})); + describe("MainLayout", () => { it("mounts the shared Cloud upgrade modal with page content", () => { render( diff --git a/ui/components/layout/main-layout/main-layout.tsx b/ui/components/layout/main-layout/main-layout.tsx index 68193326ed..e863389e41 100644 --- a/ui/components/layout/main-layout/main-layout.tsx +++ b/ui/components/layout/main-layout/main-layout.tsx @@ -4,6 +4,7 @@ import { usePathname } from "next/navigation"; import { type ReactNode, Suspense } from "react"; import { JiraDispatchModalHost } from "@/components/findings/jira-dispatch-modal-host"; +import { RecheckResourceModalHost } from "@/components/findings/recheck-resource-modal-host"; import { AppSidebar } from "@/components/layout/app-sidebar"; import { CloudUpgradeModal } from "@/components/shared/cloud-upgrade-modal"; import { useMediaQuery } from "@/hooks/use-media-query"; @@ -40,6 +41,7 @@ export default function MainLayout({ children }: { children: ReactNode }) { +
is the reference for the app's (container-query) // breakpoints, so pushing it with the side panel re-evaluates them. diff --git a/ui/components/onboarding/__tests__/onboarding-gate.test.tsx b/ui/components/onboarding/__tests__/onboarding-gate.test.tsx index 9a5f7cf88e..f83c8a0385 100644 --- a/ui/components/onboarding/__tests__/onboarding-gate.test.tsx +++ b/ui/components/onboarding/__tests__/onboarding-gate.test.tsx @@ -1,18 +1,23 @@ -import { render, screen, waitFor } from "@testing-library/react"; -import userEvent from "@testing-library/user-event"; +import { render, waitFor } from "@testing-library/react"; import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; +import { isFirstRunHandled } from "@/lib/onboarding/first-run-marker"; import { addProviderTour } from "@/lib/tours/add-provider.tour"; import { localStorageAdapter } from "@/lib/tours/store/local-storage-adapter"; import { OnboardingGate } from "../onboarding-gate"; -const pushMock = vi.fn(); +const replaceMock = vi.fn(); const armMock = vi.fn(); const pathnameMock = vi.fn(); +const permissionsMock = vi.fn(); + +vi.mock("@/hooks/use-auth", () => ({ + useAuth: () => ({ permissions: permissionsMock() }), +})); vi.mock("next/navigation", () => ({ - useRouter: () => ({ push: pushMock, replace: vi.fn() }), + useRouter: () => ({ push: vi.fn(), replace: replaceMock }), usePathname: () => pathnameMock(), })); @@ -27,20 +32,31 @@ const addProviderTourId = { version: addProviderTour.version, }; +const TENANT_A = "11111111-1111-4111-8111-111111111111"; +const TENANT_B = "22222222-2222-4222-8222-222222222222"; + +const CLOUD_FIRST_RUN_HREF = + "/providers?addProvider=true&addProviderSource=first_run&onboarding=add-provider"; +const OSS_FIRST_RUN_HREF = + "/providers?addProvider=true&addProviderSource=first_run"; + describe("OnboardingGate", () => { beforeEach(() => { window.localStorage.clear(); - pushMock.mockClear(); + replaceMock.mockClear(); armMock.mockClear(); pathnameMock.mockReturnValue("/"); + permissionsMock.mockReturnValue({ manage_providers: true }); + vi.stubEnv("UI_CLOUD_ENABLED", "true"); }); afterEach(() => { + vi.unstubAllEnvs(); vi.restoreAllMocks(); }); it.each(["/billing", "/billing/", "/billing/checkout"])( - "defers onboarding on %s without resolving it", + "defers the first run on %s without resolving it", (pathname) => { // Given pathnameMock.mockReturnValue(pathname); @@ -49,16 +65,13 @@ describe("OnboardingGate", () => { render(); // Then - expect( - screen.queryByRole("button", { name: /get started/i }), - ).not.toBeInTheDocument(); - expect(localStorageAdapter.get(addProviderTourId)).toBeNull(); + expect(replaceMock).not.toHaveBeenCalled(); expect(armMock).not.toHaveBeenCalled(); - expect(pushMock).not.toHaveBeenCalled(); + expect(isFirstRunHandled()).toBe(false); }, ); - it("offers onboarding after leaving billing without remounting the gate", async () => { + it("sends the user to add a provider after leaving billing, without remounting the gate", async () => { // Given pathnameMock.mockReturnValue("/billing"); const { rerender } = render(); @@ -68,170 +81,158 @@ describe("OnboardingGate", () => { rerender(); // Then - expect( - await screen.findByRole("button", { name: /get started/i }), - ).toBeInTheDocument(); - expect(localStorageAdapter.get(addProviderTourId)).toBeNull(); - expect(armMock).not.toHaveBeenCalled(); + await waitFor(() => + expect(replaceMock).toHaveBeenCalledExactlyOnceWith(CLOUD_FIRST_RUN_HREF), + ); }); - it("does not suppress onboarding on a route that only shares the billing prefix", async () => { + it("does not defer on a route that only shares the billing prefix", async () => { // Given - pathnameMock.mockReturnValue("/billing-settings"); + pathnameMock.mockReturnValue("/billing-history"); // When render(); // Then - expect( - await screen.findByRole("button", { name: /get started/i }), - ).toBeInTheDocument(); + await waitFor(() => expect(replaceMock).toHaveBeenCalledOnce()); }); - describe("when the user has no providers and no completion record", () => { - it("shows the Welcome modal", async () => { + describe("when a Cloud tenant has no providers and never went through the first run", () => { + it("opens the add-provider wizard with its tour and arms the checkpoint", async () => { + // Given / When render(); - expect( - await screen.findByRole("button", { name: /get started/i }), - ).toBeInTheDocument(); - }); - }); - - describe("when the user already has providers", () => { - it("does not show the Welcome modal", async () => { - render(); - - await waitFor(() => { - expect( - screen.queryByRole("button", { name: /get started/i }), - ).not.toBeInTheDocument(); - }); - }); - }); - - describe("when a completion record already exists in this browser", () => { - it("does not show the Welcome modal", async () => { - localStorageAdapter.set(addProviderTourId, { - tourId: addProviderTour.id, - version: addProviderTour.version, - state: "dismissed", - completedAt: new Date().toISOString(), - }); - - render(); - - await waitFor(() => { - expect( - screen.queryByRole("button", { name: /get started/i }), - ).not.toBeInTheDocument(); - }); - }); - }); - - describe("when the gate flow is dismissed but later sequence flows are incomplete", () => { - it("does not show the Welcome modal for a later flow", async () => { - // Later flows are only reachable via the checkpoint/sequence, never the gate. - localStorageAdapter.set(addProviderTourId, { - tourId: addProviderTour.id, - version: addProviderTour.version, - state: "dismissed", - completedAt: new Date().toISOString(), - }); - - render(); - - await waitFor(() => { - expect( - screen.queryByRole("button", { name: /get started/i }), - ).not.toBeInTheDocument(); - }); - }); - }); - - describe("when hasProviders is undefined (fail-open)", () => { - it("does not show the Welcome modal", async () => { - // `undefined` mirrors the tri-state layout forwards on a failed provider fetch. - render(); - - await waitFor(() => { - expect( - screen.queryByRole("button", { name: /get started/i }), - ).not.toBeInTheDocument(); - }); - }); - - it("can be mounted with the prop omitted entirely (fail-open)", async () => { - render(); - - await waitFor(() => { - expect( - screen.queryByRole("button", { name: /get started/i }), - ).not.toBeInTheDocument(); - }); - }); - }); - - describe("when the user accepts the Welcome modal", () => { - it("navigates to the flow route with the onboarding query param and writes no record", async () => { - const user = userEvent.setup(); - render(); - const getStarted = await screen.findByRole("button", { - name: /get started/i, - }); - - await user.click(getStarted); - - expect(pushMock).toHaveBeenCalledWith( - "/providers?onboarding=add-provider", + // Then + await waitFor(() => + expect(replaceMock).toHaveBeenCalledExactlyOnceWith( + CLOUD_FIRST_RUN_HREF, + ), ); - expect(localStorageAdapter.get(addProviderTourId)).toBeNull(); + expect(armMock).toHaveBeenCalledOnce(); }); - it("arms the onboarding checkpoint", async () => { - const user = userEvent.setup(); - render(); - const getStarted = await screen.findByRole("button", { - name: /get started/i, - }); + it("happens only once per tenant on this browser", async () => { + // Given + const { unmount } = render( + , + ); + await waitFor(() => expect(replaceMock).toHaveBeenCalledOnce()); + unmount(); + replaceMock.mockClear(); - await user.click(getStarted); + // When + render(); - expect(armMock).toHaveBeenCalledTimes(1); + // Then + expect(isFirstRunHandled(TENANT_A)).toBe(true); + expect(replaceMock).not.toHaveBeenCalled(); + }); + + it("honours a browser-wide marker written before markers were tenant-scoped", () => { + // Given: e2e storage state and pre-existing browsers set the bare key. + window.localStorage.setItem("prowler.onboarding.first-run", "true"); + + // When + render(); + + // Then + expect(replaceMock).not.toHaveBeenCalled(); + expect(isFirstRunHandled(TENANT_A)).toBe(true); + }); + + it("still runs for a different empty tenant on the same browser", async () => { + // Given + const { unmount } = render( + , + ); + await waitFor(() => expect(replaceMock).toHaveBeenCalledOnce()); + unmount(); + replaceMock.mockClear(); + + // When + render(); + + // Then + await waitFor(() => expect(replaceMock).toHaveBeenCalledOnce()); + expect(isFirstRunHandled(TENANT_B)).toBe(true); }); }); - describe("when the user dismisses the Welcome modal", () => { - it("writes a dismissed record and stops showing the modal", async () => { - const user = userEvent.setup(); + describe("when a self-hosted deployment has no providers", () => { + it("opens the add-provider wizard without the Cloud-only tour or checkpoint", async () => { + // Given + vi.stubEnv("UI_CLOUD_ENABLED", "false"); + + // When render(); - const skip = await screen.findByRole("button", { - name: /skip for now/i, - }); - await user.click(skip); - - await waitFor(() => { - expect( - screen.queryByRole("button", { name: /skip for now/i }), - ).not.toBeInTheDocument(); - }); - const record = localStorageAdapter.get(addProviderTourId); - expect(record).not.toBeNull(); - expect(record?.state).toBe("dismissed"); - }); - - it("does NOT arm the onboarding checkpoint", async () => { - const user = userEvent.setup(); - render(); - const skip = await screen.findByRole("button", { - name: /skip for now/i, - }); - - await user.click(skip); - - // Skipping must never arm the checkpoint (user opted out). + // Then + await waitFor(() => + expect(replaceMock).toHaveBeenCalledExactlyOnceWith(OSS_FIRST_RUN_HREF), + ); expect(armMock).not.toHaveBeenCalled(); }); }); + + describe("when the user cannot add providers", () => { + it("leaves the user where they are, since an empty list may just be limited visibility", () => { + // Given + permissionsMock.mockReturnValue({ manage_providers: false }); + + // When + render(); + + // Then + expect(replaceMock).not.toHaveBeenCalled(); + expect(isFirstRunHandled()).toBe(false); + }); + }); + + describe("when the tenant already has providers", () => { + it("leaves the user where they are", () => { + // Given / When + render(); + + // Then + expect(replaceMock).not.toHaveBeenCalled(); + expect(isFirstRunHandled()).toBe(false); + }); + }); + + describe("when the add-provider tour was already resolved in this browser", () => { + it("leaves the user where they are", () => { + // Given + localStorageAdapter.set(addProviderTourId, { + tourId: addProviderTour.id, + version: addProviderTour.version, + state: "dismissed", + completedAt: new Date().toISOString(), + }); + + // When + render(); + + // Then + expect(replaceMock).not.toHaveBeenCalled(); + expect(armMock).not.toHaveBeenCalled(); + }); + }); + + describe("when the provider count is unknown (fail-open)", () => { + it("leaves the user where they are when the fetch failed", () => { + // Given / When + render(); + + // Then + expect(replaceMock).not.toHaveBeenCalled(); + }); + + it("can be mounted with the prop omitted entirely", () => { + // Given / When + render(); + + // Then + expect(replaceMock).not.toHaveBeenCalled(); + }); + }); }); diff --git a/ui/components/onboarding/__tests__/onboarding-trigger.test.tsx b/ui/components/onboarding/__tests__/onboarding-trigger.test.tsx index e0f040f832..fdf884d257 100644 --- a/ui/components/onboarding/__tests__/onboarding-trigger.test.tsx +++ b/ui/components/onboarding/__tests__/onboarding-trigger.test.tsx @@ -118,6 +118,24 @@ describe("OnboardingTrigger", () => { ); }); + it("starts at the step the page asks for, skipping the ones before it", async () => { + // Given + searchParamsValue = new URLSearchParams("onboarding=add-provider"); + + // When + render( + , + ); + + // Then + await waitFor(() => + expect(startMock).toHaveBeenCalledExactlyOnceWith("provider-type"), + ); + }); + it("strips only the onboarding param and preserves other query params", async () => { searchParamsValue = new URLSearchParams( "scanId=scan-1&onboarding=add-provider&tab=completed", diff --git a/ui/components/onboarding/__tests__/onboarding-welcome-modal.test.tsx b/ui/components/onboarding/__tests__/onboarding-welcome-modal.test.tsx deleted file mode 100644 index a3665031c2..0000000000 --- a/ui/components/onboarding/__tests__/onboarding-welcome-modal.test.tsx +++ /dev/null @@ -1,83 +0,0 @@ -import { render, screen } from "@testing-library/react"; -import userEvent from "@testing-library/user-event"; -import { describe, expect, it, vi } from "vitest"; - -import { OnboardingWelcomeModal } from "../onboarding-welcome-modal"; - -describe("OnboardingWelcomeModal", () => { - describe("when open is true", () => { - it("renders the flow title and description", () => { - render( - , - ); - - expect(screen.getByText("Add your first provider")).toBeInTheDocument(); - expect( - screen.getByText( - "Connect a cloud account so Prowler has something to scan.", - ), - ).toBeInTheDocument(); - }); - - it("calls onAccept when the primary action is clicked", async () => { - const user = userEvent.setup(); - const onAccept = vi.fn(); - const onDismiss = vi.fn(); - render( - , - ); - - await user.click(screen.getByRole("button", { name: /get started/i })); - - expect(onAccept).toHaveBeenCalledTimes(1); - expect(onDismiss).not.toHaveBeenCalled(); - }); - - it("calls onDismiss when the skip action is clicked", async () => { - const user = userEvent.setup(); - const onAccept = vi.fn(); - const onDismiss = vi.fn(); - render( - , - ); - - await user.click(screen.getByRole("button", { name: /skip for now/i })); - - expect(onDismiss).toHaveBeenCalledTimes(1); - expect(onAccept).not.toHaveBeenCalled(); - }); - }); - - describe("when open is false", () => { - it("does not render the modal content", () => { - render( - , - ); - - expect( - screen.queryByRole("button", { name: /get started/i }), - ).not.toBeInTheDocument(); - }); - }); -}); diff --git a/ui/components/onboarding/onboarding-gate.tsx b/ui/components/onboarding/onboarding-gate.tsx index f32d411938..6f76318415 100644 --- a/ui/components/onboarding/onboarding-gate.tsx +++ b/ui/components/onboarding/onboarding-gate.tsx @@ -1,26 +1,40 @@ "use client"; import { usePathname, useRouter } from "next/navigation"; -import { useState } from "react"; -import { getOrderedFlows, shouldStartOnboarding } from "@/lib/onboarding"; +import { useAuth } from "@/hooks/use-auth"; +import { useMountEffect } from "@/hooks/use-mount-effect"; +import { + getOrderedFlows, + type OnboardingFlow, + shouldStartOnboarding, +} from "@/lib/onboarding"; +import { + isFirstRunHandled, + markFirstRunHandled, +} from "@/lib/onboarding/first-run-marker"; +import { WIZARD_OPEN_SOURCE } from "@/lib/provider-funnel/provider-funnel-events"; +import { buildAddProviderHref } from "@/lib/providers-navigation"; +import { isCloud } from "@/lib/shared/env"; import { localStorageAdapter } from "@/lib/tours/store/local-storage-adapter"; -import { TOUR_COMPLETION_STATES } from "@/lib/tours/tour-types"; import { useTourCompletion } from "@/lib/tours/use-tour-completion"; import { useOnboardingCheckpointStore } from "@/store/onboarding-checkpoint"; -import { OnboardingWelcomeModal } from "./onboarding-welcome-modal"; - interface OnboardingGateProps { - // `undefined` = fetch failed/ambiguous; fail-open (never force the modal). + // `undefined` = fetch failed/ambiguous; fail-open (never force the first run). hasProviders?: boolean; + // Scopes the first-run marker so one tenant's first run never silences another's. + tenantId?: string | null; } -// Mandatory new-user gate. Mounted once in the layout; decision derived during render -// via useSyncExternalStore — server renders nothing, no hydration mismatch. -export function OnboardingGate({ hasProviders }: OnboardingGateProps) { - const router = useRouter(); +// New-tenant gate. Mounted once in the layout: an empty tenant is sent straight to +// the add-provider wizard, once per tenant and browser. Renders nothing. +export function OnboardingGate({ + hasProviders, + tenantId = null, +}: OnboardingGateProps) { const pathname = usePathname(); + const { permissions } = useAuth(); // Billing must stay usable before onboarding; leaving it keeps the gate eligible. const isBillingRoute = pathname === "/billing" || pathname?.startsWith("/billing/"); @@ -28,52 +42,53 @@ export function OnboardingGate({ hasProviders }: OnboardingGateProps) { // Gate forces only the first flow (`add-provider`); remaining flows come via checkpoint/replay. const flow = getOrderedFlows()[0] ?? null; - // Returns null on server/first render — gate stays closed until resolved client-side. + // Returns null on server/first render; the redirect re-reads storage before acting. const completionRecord = useTourCompletion(flow?.tour ?? null); - // Session flag prevents the gate re-opening after accept/dismiss within this mount. - const [resolvedThisSession, setResolvedThisSession] = useState(false); - - const activeFlow = - flow && + const shouldRedirect = + flow !== null && !isBillingRoute && - !resolvedThisSession && - shouldStartOnboarding({ hasProviders, completionRecord }) - ? flow - : null; + shouldStartOnboarding({ + hasProviders, + canManageProviders: permissions.manage_providers === true, + completionRecord, + }); - if (!activeFlow) return null; + if (!shouldRedirect) return null; - const handleAccept = () => { - // Arm checkpoint only on explicit accept — skip must never arm it. + return ; +} + +interface FirstRunRedirectProps { + flow: OnboardingFlow; + tenantId: string | null; +} + +function FirstRunRedirect({ flow, tenantId }: FirstRunRedirectProps) { + const router = useRouter(); + + useMountEffect(() => { + // Hydration renders with an empty completion snapshot, so decide from storage here. + const tourId = { id: flow.tour.id, version: flow.tour.version }; + if ( + isFirstRunHandled(tenantId) || + localStorageAdapter.get(tourId) !== null + ) { + return; + } + + markFirstRunHandled(tenantId); + + const addProviderHref = buildAddProviderHref(WIZARD_OPEN_SOURCE.FIRST_RUN); + if (!isCloud()) { + router.replace(addProviderHref); + return; + } + + // Tours and the post-connect checkpoint are Cloud-only. useOnboardingCheckpointStore.getState().arm(); - setResolvedThisSession(true); - // Routes may already carry a query string, so pick the right separator. - const separator = activeFlow.route.includes("?") ? "&" : "?"; - router.push(`${activeFlow.route}${separator}onboarding=${activeFlow.id}`); - }; - - const handleDismiss = () => { - // Persist dismissal so the gate silently skips on future visits. - localStorageAdapter.set( - { id: activeFlow.tour.id, version: activeFlow.tour.version }, - { - tourId: activeFlow.tour.id, - version: activeFlow.tour.version, - state: TOUR_COMPLETION_STATES.DISMISSED, - completedAt: new Date().toISOString(), - }, - ); - setResolvedThisSession(true); - }; + router.replace(`${addProviderHref}&onboarding=${flow.id}`); + }); - return ( - - ); + return null; } diff --git a/ui/components/onboarding/onboarding-trigger.tsx b/ui/components/onboarding/onboarding-trigger.tsx index 2a1854f3c8..270afbd721 100644 --- a/ui/components/onboarding/onboarding-trigger.tsx +++ b/ui/components/onboarding/onboarding-trigger.tsx @@ -34,6 +34,8 @@ interface OnboardingTriggerProps { flow: OnboardingFlow; // force-started when the sequence names it or `?onboarding=` matches stepHandlers?: { [K in TTarget]?: TourStepHandlers }; configOverrides?: Partial; + // Step to begin from when the page already did what the earlier steps ask for. + startAtTarget?: TTarget; } // Latched per-trigger: `key` mounts a fresh runner on each re-trigger; `mode` drives param-strip logic. @@ -50,6 +52,7 @@ export function OnboardingTrigger({ flow, stepHandlers, configOverrides, + startAtTarget, }: OnboardingTriggerProps) { const searchParams = useSearchParams(); const param = searchParams?.get(ONBOARDING_PARAM) ?? null; // null outside Suspense context @@ -108,6 +111,7 @@ export function OnboardingTrigger({ queryString={request.queryString} stepHandlers={stepHandlers} configOverrides={configOverrides} + startAtTarget={startAtTarget} /> ); } @@ -118,6 +122,7 @@ interface OnboardingTourRunnerProps { queryString: string; stepHandlers?: { [K in TTarget]?: TourStepHandlers }; configOverrides?: Partial; + startAtTarget?: TTarget; } function OnboardingTourRunner({ @@ -126,6 +131,7 @@ function OnboardingTourRunner({ queryString, stepHandlers, configOverrides, + startAtTarget, }: OnboardingTourRunnerProps) { // onClosed is intentionally inert — the banner owns advance/exit for both modes. const { start } = useDriverTour(flow.tour, { @@ -144,7 +150,7 @@ function OnboardingTourRunner({ queueMicrotask(() => { if (cancelled) return; - start(); + start(startAtTarget); if (mode === "replay") { // Only strip when the param actually started this replay; a same-route // in-memory request leaves the URL untouched (no replaceState needed). diff --git a/ui/components/onboarding/onboarding-welcome-modal.tsx b/ui/components/onboarding/onboarding-welcome-modal.tsx deleted file mode 100644 index 566db34cba..0000000000 --- a/ui/components/onboarding/onboarding-welcome-modal.tsx +++ /dev/null @@ -1,42 +0,0 @@ -"use client"; - -import { Button } from "@/components/shadcn"; -import { DialogFooter } from "@/components/shadcn/dialog"; -import { Modal } from "@/components/shadcn/modal/modal"; - -interface OnboardingWelcomeModalProps { - open: boolean; - flowTitle?: string; - flowDescription?: string; - onAccept: () => void; - onDismiss: () => void; -} - -export function OnboardingWelcomeModal({ - open, - flowTitle, - flowDescription, - onAccept, - onDismiss, -}: OnboardingWelcomeModalProps) { - return ( - { - if (!next) onDismiss(); - }} - > - - {/* Outline matches the app's modal secondary action (e.g. Launch Scan's Cancel). */} - - - - - ); -} diff --git a/ui/components/providers/forms/delete-organization-form.tsx b/ui/components/providers/forms/delete-organization-form.tsx index 4a12c18587..c777575066 100644 --- a/ui/components/providers/forms/delete-organization-form.tsx +++ b/ui/components/providers/forms/delete-organization-form.tsx @@ -12,6 +12,7 @@ import { pollTaskCompletion } from "@/components/providers/organizations/org-acc import { Button, useToast } from "@/components/shadcn"; import { getNodeLabel } from "@/lib/organizations"; import { NodeKind, OrganizationType } from "@/types/organizations"; +import { CONNECTION_CHECK_STATUS } from "@/types/providers"; import { PROVIDERS_GROUP_KIND, ProvidersGroupKind, @@ -85,11 +86,11 @@ export function DeleteOrganizationForm({ const taskId = extractTaskId(result); const taskResult = taskId ? await pollTaskCompletion(taskId) - : { success: true as const }; + : { status: CONNECTION_CHECK_STATUS.SUCCESS }; setIsLoading(false); - if (!taskResult.success) { + if (taskResult.status !== CONNECTION_CHECK_STATUS.SUCCESS) { toast({ variant: "destructive", title: "Deletion did not complete", diff --git a/ui/components/providers/organizations/aws-method-selector.test.tsx b/ui/components/providers/organizations/aws-method-selector.test.tsx deleted file mode 100644 index 7ea1d9a95e..0000000000 --- a/ui/components/providers/organizations/aws-method-selector.test.tsx +++ /dev/null @@ -1,43 +0,0 @@ -import { render, screen } from "@testing-library/react"; -import userEvent from "@testing-library/user-event"; -import { afterEach, describe, expect, it, vi } from "vitest"; - -import { useCloudUpgradeStore } from "@/store"; -import { CLOUD_UPGRADE_FEATURE } from "@/types/cloud-upgrade"; - -import { AwsMethodSelector } from "./aws-method-selector"; - -describe("AwsMethodSelector", () => { - afterEach(() => { - vi.unstubAllEnvs(); - useCloudUpgradeStore.getState().closeCloudUpgrade(); - }); - - it("opens the AWS Organizations upgrade in Local Server", async () => { - // Given - vi.stubEnv("UI_CLOUD_ENABLED", "false"); - const user = userEvent.setup(); - const onSelectOrganizations = vi.fn(); - - // When - render( - , - ); - - // Then - await user.click( - screen.getByRole("radio", { - name: /add multiple accounts with aws organizations/i, - }), - ); - - expect(onSelectOrganizations).not.toHaveBeenCalled(); - expect(screen.getByText("Cloud")).toBeVisible(); - expect(useCloudUpgradeStore.getState().activeFeature).toBe( - CLOUD_UPGRADE_FEATURE.AWS_ORGANIZATIONS, - ); - }); -}); diff --git a/ui/components/providers/organizations/aws-method-selector.tsx b/ui/components/providers/organizations/aws-method-selector.tsx deleted file mode 100644 index 882adef1b2..0000000000 --- a/ui/components/providers/organizations/aws-method-selector.tsx +++ /dev/null @@ -1,50 +0,0 @@ -"use client"; - -import { Box, Boxes } from "lucide-react"; - -import { RadioCard } from "@/components/providers/radio-card"; -import { Badge } from "@/components/shadcn/badge/badge"; -import { isCloud } from "@/lib/shared/env"; -import { useCloudUpgradeStore } from "@/store"; -import { CLOUD_UPGRADE_FEATURE } from "@/types/cloud-upgrade"; - -interface AwsMethodSelectorProps { - onSelectSingle: () => void; - onSelectOrganizations: () => void; -} - -export function AwsMethodSelector({ - onSelectSingle, - onSelectOrganizations, -}: AwsMethodSelectorProps) { - const isCloudEnv = isCloud(); - const openCloudUpgrade = useCloudUpgradeStore( - (state) => state.openCloudUpgrade, - ); - - return ( -
-

- Select a method to add your accounts to Prowler. -

- - - - - isCloudEnv - ? onSelectOrganizations() - : openCloudUpgrade(CLOUD_UPGRADE_FEATURE.AWS_ORGANIZATIONS) - } - > - {!isCloudEnv && Cloud} - -
- ); -} diff --git a/ui/components/providers/organizations/hooks/use-org-account-selection-flow.test.ts b/ui/components/providers/organizations/hooks/use-org-account-selection-flow.test.ts index 46833ae4f0..78b7e4430d 100644 --- a/ui/components/providers/organizations/hooks/use-org-account-selection-flow.test.ts +++ b/ui/components/providers/organizations/hooks/use-org-account-selection-flow.test.ts @@ -8,6 +8,7 @@ import { type GcpOrgHierarchy, ORGANIZATION_TYPE, } from "@/types/organizations"; +import { CONNECTION_CHECK_STATUS } from "@/types/providers"; import { useOrgAccountSelectionFlow } from "./use-org-account-selection-flow"; @@ -15,13 +16,25 @@ const organizationsActionsMock = vi.hoisted(() => ({ applyDiscovery: vi.fn(), })); const providersActionsMock = vi.hoisted(() => ({ - getProviderUidsByIds: vi.fn(), + getProviderConnectionBaselines: vi.fn(), + getProviderUidsAndConnectionBaselines: vi.fn(), revalidateProviders: vi.fn(), startProviderConnectionChecks: vi.fn(), })); const tasksActionsMock = vi.hoisted(() => ({ getTasksByIds: vi.fn(), })); +const providerHelpersMock = vi.hoisted(() => ({ + resolveProviderConnectionState: vi.fn(), +})); +const pollConnectionTasksMock = vi.hoisted(() => vi.fn()); +// Mutable holder for the real `pollConnectionTasks`, captured once the module +// mock factory below runs, and re-applied in `beforeEach` since +// `mockReset: true` clears `pollConnectionTasksMock`'s implementation before +// every test. +const realPollConnectionTasksHolder = vi.hoisted( + () => ({}) as { current?: (...args: unknown[]) => unknown }, +); vi.mock( "@/actions/organizations/organizations", @@ -29,6 +42,15 @@ vi.mock( ); vi.mock("@/actions/providers/providers", () => providersActionsMock); vi.mock("@/actions/task/tasks", () => tasksActionsMock); +vi.mock("@/lib/provider-helpers", () => providerHelpersMock); +vi.mock("../org-account-selection.utils", async (importOriginal) => { + const actual = + await importOriginal(); + realPollConnectionTasksHolder.current = actual.pollConnectionTasks as ( + ...args: unknown[] + ) => unknown; + return { ...actual, pollConnectionTasks: pollConnectionTasksMock }; +}); const ORGANIZATION_UID = "organizations/123456789012"; const PROJECT_UID = "projects/acme-prod"; @@ -55,6 +77,7 @@ function seedAppliedSelection() { interface RenderedFlow { onNext: ReturnType; startTesting: () => Promise; + getFooterConfig: () => WizardFooterConfig | null; } function renderFlow(): RenderedFlow { @@ -79,6 +102,7 @@ function renderFlow(): RenderedFlow { footerConfig?.onAction?.(); }); }, + getFooterConfig: () => footerConfig, }; } @@ -91,18 +115,27 @@ describe("useOrgAccountSelectionFlow", () => { ...Object.values(organizationsActionsMock), ...Object.values(providersActionsMock), ...Object.values(tasksActionsMock), + ...Object.values(providerHelpersMock), ]) { mockFn.mockReset(); } + pollConnectionTasksMock.mockReset(); + pollConnectionTasksMock.mockImplementation((...args: unknown[]) => + realPollConnectionTasksHolder.current?.(...args), + ); organizationsActionsMock.applyDiscovery.mockResolvedValue({ data: { relationships: { providers: { data: [{ id: PROVIDER_ID }] } }, }, }); - providersActionsMock.getProviderUidsByIds.mockResolvedValue({ - [PROVIDER_ID]: PROJECT_UID, - }); + providersActionsMock.getProviderUidsAndConnectionBaselines.mockResolvedValue( + { + uidById: { [PROVIDER_ID]: PROJECT_UID }, + baselineById: {}, + }, + ); + providersActionsMock.getProviderConnectionBaselines.mockResolvedValue({}); providersActionsMock.revalidateProviders.mockResolvedValue(undefined); }); @@ -156,5 +189,250 @@ describe("useOrgAccountSelectionFlow", () => { }); expect(onNext).toHaveBeenCalledTimes(1); }); + + it("resolves a still-pending task from the provider's persisted state once the wait is exhausted", async () => { + // Given the batch poll never settles the task before retries run out. + seedAppliedSelection(); + providersActionsMock.startProviderConnectionChecks.mockResolvedValue({ + [PROVIDER_ID]: { taskId: "task-1" }, + }); + providersActionsMock.getProviderUidsAndConnectionBaselines.mockResolvedValue( + { + uidById: { [PROVIDER_ID]: PROJECT_UID }, + baselineById: { [PROVIDER_ID]: "2025-01-01T00:00:00Z" }, + }, + ); + providerHelpersMock.resolveProviderConnectionState.mockResolvedValue({ + status: CONNECTION_CHECK_STATUS.SUCCESS, + error: null, + }); + pollConnectionTasksMock.mockImplementation( + async (taskIds: string[], { onSettled, resolveExhausted }) => { + for (const taskId of taskIds) { + const resolved = resolveExhausted + ? await resolveExhausted(taskId) + : null; + onSettled( + taskId, + resolved ?? { + status: CONNECTION_CHECK_STATUS.FAILED, + error: "Connection test timed out.", + }, + ); + } + }, + ); + const { onNext, startTesting } = renderFlow(); + + // When + await startTesting(); + + // Then: read from the provider's own record, not reported as a timeout, + // using the baseline captured for this provider before dispatch. + await waitFor(() => { + expect(useOrgSetupStore.getState().connectionResults[PROVIDER_ID]).toBe( + CONNECTION_TEST_STATUS.SUCCESS, + ); + }); + expect( + providersActionsMock.getProviderUidsAndConnectionBaselines, + ).toHaveBeenCalledWith([PROVIDER_ID]); + expect( + providerHelpersMock.resolveProviderConnectionState, + ).toHaveBeenCalledWith(PROVIDER_ID, "2025-01-01T00:00:00Z"); + expect(onNext).toHaveBeenCalledTimes(1); + }); + + it("does not report a still-running fallback as a connection failure", async () => { + // Given: the wait exhausts and the provider's own record cannot confirm + // an outcome either (the backend check is genuinely still running). + seedAppliedSelection(); + providersActionsMock.startProviderConnectionChecks.mockResolvedValue({ + [PROVIDER_ID]: { taskId: "task-1" }, + }); + providerHelpersMock.resolveProviderConnectionState.mockResolvedValue({ + status: CONNECTION_CHECK_STATUS.PENDING, + error: "The connection test is still running.", + }); + pollConnectionTasksMock.mockImplementation( + async (taskIds: string[], { onSettled, resolveExhausted }) => { + for (const taskId of taskIds) { + const resolved = resolveExhausted + ? await resolveExhausted(taskId) + : null; + onSettled( + taskId, + resolved ?? { + status: CONNECTION_CHECK_STATUS.FAILED, + error: "Connection test timed out.", + }, + ); + } + }, + ); + const { onNext, startTesting } = renderFlow(); + + // When + await startTesting(); + + // Then: neither a success (does not auto-advance) nor an error. + await waitFor(() => { + expect( + providerHelpersMock.resolveProviderConnectionState, + ).toHaveBeenCalled(); + }); + expect(useOrgSetupStore.getState().connectionResults[PROVIDER_ID]).toBe( + CONNECTION_TEST_STATUS.PENDING, + ); + expect(onNext).not.toHaveBeenCalled(); + }); + + it("keeps the retry control available when every unresolved account is pending, not failed", async () => { + // Given: no confirmed error, only a wait exhausted with no verdict -- + // `hasConnectionErrors` alone would hide "Test Connections" here. + seedAppliedSelection(); + providersActionsMock.startProviderConnectionChecks.mockResolvedValue({ + [PROVIDER_ID]: { taskId: "task-1" }, + }); + providerHelpersMock.resolveProviderConnectionState.mockResolvedValue({ + status: CONNECTION_CHECK_STATUS.PENDING, + error: "The connection test is still running.", + }); + pollConnectionTasksMock.mockImplementation( + async (taskIds: string[], { onSettled, resolveExhausted }) => { + for (const taskId of taskIds) { + const resolved = resolveExhausted + ? await resolveExhausted(taskId) + : null; + onSettled( + taskId, + resolved ?? { + status: CONNECTION_CHECK_STATUS.FAILED, + error: "Connection test timed out.", + }, + ); + } + }, + ); + const { startTesting, getFooterConfig } = renderFlow(); + + // When + await startTesting(); + + // Then: the action stays visible and enabled for a retry. + await waitFor(() => { + expect(useOrgSetupStore.getState().connectionResults[PROVIDER_ID]).toBe( + CONNECTION_TEST_STATUS.PENDING, + ); + }); + const footerConfig = getFooterConfig(); + expect(footerConfig?.showAction).toBe(true); + expect(footerConfig?.actionDisabled).toBe(false); + }); + + it("retries only the still-pending account, not one that already succeeded", async () => { + // Given: two accounts selected, one project and one folder-scoped project + // under the same GCP org so both resolve from a single apply. + const OTHER_UID = "projects/acme-staging"; + const hierarchyWithTwoProjects: GcpOrgHierarchy = { + ...GCP_HIERARCHY, + candidates: [ + ...GCP_HIERARCHY.candidates, + { uid: OTHER_UID, label: "Acme Staging", parentId: ORGANIZATION_UID }, + ], + }; + const OTHER_PROVIDER_ID = "provider-2"; + const store = useOrgSetupStore.getState(); + store.setOrganizationType(ORGANIZATION_TYPE.GCP); + store.setOrganization("org-1", "Acme", ORGANIZATION_UID); + store.setDiscovery("discovery-1", hierarchyWithTwoProjects); + store.setSelectedCandidateIds([PROJECT_UID, OTHER_UID]); + + organizationsActionsMock.applyDiscovery.mockResolvedValue({ + data: { + relationships: { + providers: { + data: [{ id: PROVIDER_ID }, { id: OTHER_PROVIDER_ID }], + }, + }, + }, + }); + providersActionsMock.getProviderUidsAndConnectionBaselines.mockResolvedValue( + { + uidById: { + [PROVIDER_ID]: PROJECT_UID, + [OTHER_PROVIDER_ID]: OTHER_UID, + }, + baselineById: {}, + }, + ); + providersActionsMock.startProviderConnectionChecks.mockResolvedValue({ + [PROVIDER_ID]: { taskId: "task-1" }, + [OTHER_PROVIDER_ID]: { taskId: "task-2" }, + }); + tasksActionsMock.getTasksByIds.mockResolvedValue({ + "task-1": { + data: { + attributes: { state: "completed", result: { connected: true } }, + }, + }, + "task-2": { data: { attributes: { state: "executing" } } }, + }); + providerHelpersMock.resolveProviderConnectionState.mockImplementation( + async (providerId: string) => + providerId === OTHER_PROVIDER_ID + ? { + status: CONNECTION_CHECK_STATUS.PENDING, + error: "The connection test is still running.", + } + : { status: CONNECTION_CHECK_STATUS.SUCCESS, error: null }, + ); + pollConnectionTasksMock.mockImplementation( + async (taskIds: string[], { onSettled, resolveExhausted }) => { + for (const taskId of taskIds) { + if (taskId === "task-1") { + onSettled(taskId, { status: CONNECTION_CHECK_STATUS.SUCCESS }); + continue; + } + const resolved = resolveExhausted + ? await resolveExhausted(taskId) + : null; + onSettled( + taskId, + resolved ?? { + status: CONNECTION_CHECK_STATUS.FAILED, + error: "Connection test timed out.", + }, + ); + } + }, + ); + const { startTesting } = renderFlow(); + + // First pass: one account succeeds, the other is left pending. + await startTesting(); + await waitFor(() => { + expect( + useOrgSetupStore.getState().connectionResults[OTHER_PROVIDER_ID], + ).toBe(CONNECTION_TEST_STATUS.PENDING); + }); + expect(useOrgSetupStore.getState().connectionResults[PROVIDER_ID]).toBe( + CONNECTION_TEST_STATUS.SUCCESS, + ); + providersActionsMock.startProviderConnectionChecks.mockClear(); + + // When: pressing "Test Connections" again to retry. + await startTesting(); + + // Then: only the still-pending account is re-dispatched. + await waitFor(() => { + expect( + providersActionsMock.startProviderConnectionChecks, + ).toHaveBeenCalled(); + }); + expect( + providersActionsMock.startProviderConnectionChecks, + ).toHaveBeenCalledWith([OTHER_PROVIDER_ID]); + }); }); }); diff --git a/ui/components/providers/organizations/hooks/use-org-account-selection-flow.ts b/ui/components/providers/organizations/hooks/use-org-account-selection-flow.ts index b5130760ba..65d4b4599c 100644 --- a/ui/components/providers/organizations/hooks/use-org-account-selection-flow.ts +++ b/ui/components/providers/organizations/hooks/use-org-account-selection-flow.ts @@ -5,7 +5,8 @@ import { useEffect, useRef, useState } from "react"; import { applyDiscovery } from "@/actions/organizations/organizations"; import { buildApplyPayload } from "@/actions/organizations/organizations.adapter"; import { - getProviderUidsByIds, + getProviderConnectionBaselines, + getProviderUidsAndConnectionBaselines, revalidateProviders, startProviderConnectionChecks, } from "@/actions/providers/providers"; @@ -13,12 +14,14 @@ import { WIZARD_FOOTER_ACTION_TYPE, WizardFooterConfig, } from "@/components/providers/wizard/steps/footer-controls"; +import { resolveProviderConnectionState } from "@/lib/provider-helpers"; import { useOrgSetupStore } from "@/store/organizations/store"; import { CONNECTION_TEST_STATUS, ConnectionTestStatus, PROVIDER_SECRET_STATE, } from "@/types/organizations"; +import { CONNECTION_CHECK_STATUS } from "@/types/providers"; import { TREE_ITEM_STATUS, TreeDataItem } from "@/types/tree"; import { @@ -26,6 +29,7 @@ import { canAdvanceToLaunchStep, getLaunchableProviderIds, pollConnectionTasks, + type PollConnectionTaskResult, } from "../org-account-selection.utils"; import { extractErrorMessage } from "./error-utils"; @@ -141,13 +145,22 @@ function buildTreeWithConnectionState( status = TREE_ITEM_STATUS.ERROR; errorMessage = (providerId && connectionErrors[providerId]) || "Connection failed."; - } else if ( - showPendingState || - connectionStatus === CONNECTION_TEST_STATUS.PENDING - ) { + } else if (showPendingState) { + // A batch test is actively in flight -- genuinely waiting on a response, + // so the spinner is accurate. isLoading = true; status = undefined; errorMessage = undefined; + } else if (connectionStatus === CONNECTION_TEST_STATUS.PENDING) { + // The wait was exhausted with no confirmed outcome, and nothing is + // polling this account any more -- a spinner here would be misleading. + // A static icon marks it as unresolved instead; "Test Connections" + // retries it (see `hasUnresolvedConnections`). + isLoading = false; + status = TREE_ITEM_STATUS.PENDING; + errorMessage = + (providerId && connectionErrors[providerId]) || + "The connection test is still running. Refresh in a moment to see the result."; } else if (hasAppliedProviders) { // Applied, but no outcome ever arrived for this account — typically an // unresolved provider uid. Without this the row falls back to a plain @@ -241,6 +254,13 @@ export function useOrgAccountSelectionFlow({ const hasConnectionErrors = Object.values(connectionResults).some( (status) => status === CONNECTION_TEST_STATUS.ERROR, ); + // A wait exhausted with no verdict, distinct from a confirmed error: it does + // not earn the error banner (see `org-account-selection.tsx`), but it still + // needs a way back to a resolved state, so it counts toward `canRetry` below. + const hasPendingConnections = Object.values(connectionResults).some( + (status) => status === CONNECTION_TEST_STATUS.PENDING, + ); + const hasUnresolvedConnections = hasConnectionErrors || hasPendingConnections; const willReplaceSelectedNames = sanitizedSelectedCandidateIds .map((id) => candidateLookup.get(id)) .filter( @@ -281,7 +301,10 @@ export function useOrgAccountSelectionFlow({ }; }, []); - const testAllConnections = async (providerIds: string[]) => { + const testAllConnections = async ( + providerIds: string[], + precomputedBaselines?: Record, + ) => { connectionTestAbortControllerRef.current?.abort(); const abortController = new AbortController(); connectionTestAbortControllerRef.current = abortController; @@ -296,26 +319,49 @@ export function useOrgAccountSelectionFlow({ const settleProvider = ( providerId: string, - result: { success: boolean; error?: string }, + result: PollConnectionTaskResult, ) => { if (!isMountedRef.current || signal.aborted) { return; } + + // Still running past the wait -- neither a pass nor a fail. Leaves the + // account pending rather than reporting a failure the backend never gave; + // the message is kept (not nulled) so the tree can explain the static + // pending icon it now shows once `isTesting` stops. + if (result.status === CONNECTION_CHECK_STATUS.PENDING) { + setConnectionResult(providerId, CONNECTION_TEST_STATUS.PENDING); + setConnectionError(providerId, result.error ?? null); + return; + } + + const succeeded = result.status === CONNECTION_CHECK_STATUS.SUCCESS; setConnectionResult( providerId, - result.success + succeeded ? CONNECTION_TEST_STATUS.SUCCESS : CONNECTION_TEST_STATUS.ERROR, ); setConnectionError( providerId, - result.success + succeeded ? null : result.error || "Connection failed for this account.", ); }; try { + // Read before dispatch, so the fallback below can tell each provider's own + // check result apart from whatever (possibly stale) result was already on + // record -- by comparing values, not by comparing timestamps against the + // browser's clock. See `resolveProviderConnectionState`. The initial apply + // already reads this alongside the created providers' uids (see + // `handleApplyAndTest`) and passes it in, so a retry is the only path that + // fetches it here. + const connectionBaselines = + precomputedBaselines ?? + (await getProviderConnectionBaselines(providerIds)); + // One action dispatches every check and one reads every pending task per // round: Next runs client-invoked server actions one at a time, so a loop // here would serialize the batch whatever concurrency it asked for. @@ -341,7 +387,7 @@ export function useOrgAccountSelectionFlow({ // No task id means no check ever ran, so it cannot count as passing. if (!outcome.taskId) { settleProvider(providerId, { - success: false, + status: CONNECTION_CHECK_STATUS.FAILED, error: "Connection test did not start.", }); continue; @@ -358,6 +404,17 @@ export function useOrgAccountSelectionFlow({ settleProvider(providerId, result); } }, + resolveExhausted: async (taskId) => { + const providerId = providerIdByTaskId.get(taskId); + if (!providerId) { + return null; + } + const state = await resolveProviderConnectionState( + providerId, + connectionBaselines[providerId], + ); + return { status: state.status, error: state.error ?? undefined }; + }, }); } catch { if (isMountedRef.current && !signal.aborted) { @@ -443,10 +500,21 @@ export function useOrgAccountSelectionFlow({ ) ?? []; setCreatedProviderIds(providerIds); + + // One filtered `/providers` read for both: the apply view rejects `include`, + // so the created providers' uids are read back separately, and the flow needs + // their connection baselines before dispatch anyway (see `testAllConnections`). + // Reading them together avoids fetching the same provider ids twice. + const { uidById, baselineById } = + await getProviderUidsAndConnectionBaselines(providerIds); + if (!isMountedRef.current) { + return; + } + const mapping = await buildCandidateToProviderMap({ selectedCandidateIds: currentSelectedCandidateIds, providerIds, - resolveProviderUids: getProviderUidsByIds, + resolveProviderUids: async () => uidById, }); if (!isMountedRef.current) { return; @@ -456,7 +524,7 @@ export function useOrgAccountSelectionFlow({ setIsApplying(false); lastAppliedSelectionKeyRef.current = currentSelectionKey; - await testAllConnections(providerIds); + await testAllConnections(providerIds, baselineById); }; const handleStartTesting = () => { @@ -482,12 +550,18 @@ export function useOrgAccountSelectionFlow({ return; } - const failedProviderIds = createdProviderIds.filter( + // Retries both confirmed failures and accounts a previous wait exhausted + // without a verdict -- otherwise a still-pending account has no way back to + // a resolved state once the batch that produced it has stopped polling. + const unresolvedProviderIds = createdProviderIds.filter( (providerId) => - connectionResults[providerId] === CONNECTION_TEST_STATUS.ERROR, + connectionResults[providerId] === CONNECTION_TEST_STATUS.ERROR || + connectionResults[providerId] === CONNECTION_TEST_STATUS.PENDING, ); const providerIdsToTest = - failedProviderIds.length > 0 ? failedProviderIds : createdProviderIds; + unresolvedProviderIds.length > 0 + ? unresolvedProviderIds + : createdProviderIds; void testAllConnections(providerIdsToTest); }; startTestingActionRef.current = handleStartTesting; @@ -513,7 +587,7 @@ export function useOrgAccountSelectionFlow({ return; } - const canRetry = hasConnectionErrors || Boolean(applyError); + const canRetry = hasUnresolvedConnections || Boolean(applyError); const hasSelectedAccounts = selectedCount > 0; onFooterChange({ @@ -542,7 +616,7 @@ export function useOrgAccountSelectionFlow({ }); }, [ applyError, - hasConnectionErrors, + hasUnresolvedConnections, isApplying, isTesting, isTestingView, diff --git a/ui/components/providers/organizations/org-account-selection.utils.test.ts b/ui/components/providers/organizations/org-account-selection.utils.test.ts index 366ec715b9..11c01e69bb 100644 --- a/ui/components/providers/organizations/org-account-selection.utils.test.ts +++ b/ui/components/providers/organizations/org-account-selection.utils.test.ts @@ -1,10 +1,13 @@ import { describe, expect, it, vi } from "vitest"; import { CONNECTION_TEST_STATUS } from "@/types/organizations"; +import { CONNECTION_CHECK_STATUS } from "@/types/providers"; import { buildCandidateToProviderMap, canAdvanceToLaunchStep, + CONNECTION_CHECK_DEFAULT_DELAYS_MS, + CONNECTION_CHECK_MAX_RETRIES, getLaunchableProviderIds, pollConnectionTasks, } from "./org-account-selection.utils"; @@ -89,8 +92,14 @@ describe("pollConnectionTasks", () => { // Then — the fast one is reported after round 1 and dropped from later reads, // while the slow one is still pending. expect(settled).toEqual([ - ["task-fast", { success: true }], - ["task-slow", { success: false, error: "Role trust policy mismatch." }], + ["task-fast", { status: CONNECTION_CHECK_STATUS.SUCCESS }], + [ + "task-slow", + { + status: CONNECTION_CHECK_STATUS.FAILED, + error: "Role trust policy mismatch.", + }, + ], ]); expect(rounds).toEqual([ ["task-fast", "task-slow"], @@ -150,8 +159,99 @@ describe("pollConnectionTasks", () => { // Then — the settled result stands; the pending one is reported cancelled. expect(getTasksByIds).toHaveBeenCalledTimes(1); expect(settled).toEqual([ - ["task-a", { success: true }], - ["task-b", { success: false, error: "Connection test cancelled." }], + ["task-a", { status: CONNECTION_CHECK_STATUS.SUCCESS }], + [ + "task-b", + { + status: CONNECTION_CHECK_STATUS.FAILED, + error: "Connection test cancelled.", + }, + ], + ]); + }); + + it("reports cancellation instead of accepting the resolver's result when abort lands mid-await", async () => { + // Given: the wait exhausts with one task still pending, and the caller's + // `resolveExhausted` aborts the flow while its own lookup is in flight + // (e.g. the wizard unmounted). The abort must win even though the + // resolver still returns a result. + const abortController = new AbortController(); + const getTasksByIds = vi.fn(async () => ({ + "task-a": executing, + })); + const settled: Array<[string, unknown]> = []; + const resolveExhausted = vi.fn(async () => { + abortController.abort(); + return { status: CONNECTION_CHECK_STATUS.SUCCESS }; + }); + + // When + await pollConnectionTasks(["task-a"], { + onSettled: (taskId, result) => settled.push([taskId, result]), + getTasksByIds, + sleep: async () => {}, + maxRetries: 1, + signal: abortController.signal, + resolveExhausted, + }); + + // Then: cancelled, not the resolver's (stale) success. + expect(resolveExhausted).toHaveBeenCalledWith("task-a"); + expect(settled).toEqual([ + [ + "task-a", + { + status: CONNECTION_CHECK_STATUS.FAILED, + error: "Connection test cancelled.", + }, + ], + ]); + }); + + it("stops resolving further tasks once abort lands between resolveExhausted calls", async () => { + // Given: two tasks are still pending at exhaustion; abort fires while the + // first is being resolved, so the second must never be looked up. + const abortController = new AbortController(); + const getTasksByIds = vi.fn(async () => ({ + "task-a": executing, + "task-b": executing, + })); + const settled: Array<[string, unknown]> = []; + const resolveExhausted = vi.fn(async (taskId: string) => { + if (taskId === "task-a") { + abortController.abort(); + } + return { status: CONNECTION_CHECK_STATUS.SUCCESS }; + }); + + // When + await pollConnectionTasks(["task-a", "task-b"], { + onSettled: (taskId, result) => settled.push([taskId, result]), + getTasksByIds, + sleep: async () => {}, + maxRetries: 1, + signal: abortController.signal, + resolveExhausted, + }); + + // Then + expect(resolveExhausted).toHaveBeenCalledTimes(1); + expect(resolveExhausted).toHaveBeenCalledWith("task-a"); + expect(settled).toEqual([ + [ + "task-a", + { + status: CONNECTION_CHECK_STATUS.FAILED, + error: "Connection test cancelled.", + }, + ], + [ + "task-b", + { + status: CONNECTION_CHECK_STATUS.FAILED, + error: "Connection test cancelled.", + }, + ], ]); }); @@ -173,8 +273,124 @@ describe("pollConnectionTasks", () => { // Then expect(settled).toEqual([ - ["task-a", { success: true }], - ["task-b", { success: false, error: "Connection test timed out." }], + ["task-a", { status: CONNECTION_CHECK_STATUS.SUCCESS }], + [ + "task-b", + { + status: CONNECTION_CHECK_STATUS.FAILED, + error: "Connection test timed out.", + }, + ], + ]); + }); + + it("sizes the default wait past the backend's 120s provider-connection-check time limit", () => { + // The last delay in the ladder repeats for every retry beyond it, so the + // worst-case total wait is (maxRetries - 1) * lastDelay. + const lastDelay = + CONNECTION_CHECK_DEFAULT_DELAYS_MS[ + CONNECTION_CHECK_DEFAULT_DELAYS_MS.length - 1 + ]; + const worstCaseWaitMs = (CONNECTION_CHECK_MAX_RETRIES - 1) * lastDelay; + + expect(worstCaseWaitMs).toBeGreaterThan(120_000); + }); + + it("resolves a still-pending task from the caller once the wait is exhausted", async () => { + // Given: the batch read never settles "task-b" before retries run out. + const getTasksByIds = vi.fn(async () => ({ + "task-a": completed(true), + "task-b": executing, + })); + const settled: Array<[string, unknown]> = []; + const resolveExhausted = vi.fn(async (taskId: string) => + taskId === "task-b" ? { status: CONNECTION_CHECK_STATUS.SUCCESS } : null, + ); + + // When + await pollConnectionTasks(["task-a", "task-b"], { + onSettled: (taskId, result) => settled.push([taskId, result]), + getTasksByIds, + sleep: async () => {}, + maxRetries: 2, + resolveExhausted, + }); + + // Then: the exhausted task is settled from the fallback, not a timeout. + expect(resolveExhausted).toHaveBeenCalledWith("task-b"); + expect(settled).toEqual([ + ["task-a", { status: CONNECTION_CHECK_STATUS.SUCCESS }], + ["task-b", { status: CONNECTION_CHECK_STATUS.SUCCESS }], + ]); + }); + + it("reports a still-running fallback as pending, not as a failure", async () => { + // Given: the batch read never settles "task-b", and the caller's fallback + // cannot confirm an outcome either (the backend task is still running). + const getTasksByIds = vi.fn(async () => ({ + "task-a": completed(true), + "task-b": executing, + })); + const settled: Array<[string, unknown]> = []; + const resolveExhausted = vi.fn(async (taskId: string) => + taskId === "task-b" + ? { + status: CONNECTION_CHECK_STATUS.PENDING, + error: "The connection test is still running.", + } + : null, + ); + + // When + await pollConnectionTasks(["task-a", "task-b"], { + onSettled: (taskId, result) => settled.push([taskId, result]), + getTasksByIds, + sleep: async () => {}, + maxRetries: 2, + resolveExhausted, + }); + + // Then: pending, distinct from both success and failure. + expect(settled).toEqual([ + ["task-a", { status: CONNECTION_CHECK_STATUS.SUCCESS }], + [ + "task-b", + { + status: CONNECTION_CHECK_STATUS.PENDING, + error: "The connection test is still running.", + }, + ], + ]); + }); + + it("falls back to the timeout message when the fallback cannot resolve a task", async () => { + // Given + const getTasksByIds = vi.fn(async () => ({ + "task-a": completed(true), + "task-b": executing, + })); + const settled: Array<[string, unknown]> = []; + const resolveExhausted = vi.fn(async () => null); + + // When + await pollConnectionTasks(["task-a", "task-b"], { + onSettled: (taskId, result) => settled.push([taskId, result]), + getTasksByIds, + sleep: async () => {}, + maxRetries: 2, + resolveExhausted, + }); + + // Then + expect(settled).toEqual([ + ["task-a", { status: CONNECTION_CHECK_STATUS.SUCCESS }], + [ + "task-b", + { + status: CONNECTION_CHECK_STATUS.FAILED, + error: "Connection test timed out.", + }, + ], ]); }); @@ -197,8 +413,11 @@ describe("pollConnectionTasks", () => { // Then expect(getTasksByIds).toHaveBeenCalledTimes(1); expect(settled).toEqual([ - ["task-a", { success: true }], - ["task-b", { success: false, error: "Task not found." }], + ["task-a", { status: CONNECTION_CHECK_STATUS.SUCCESS }], + [ + "task-b", + { status: CONNECTION_CHECK_STATUS.FAILED, error: "Task not found." }, + ], ]); }); }); diff --git a/ui/components/providers/organizations/org-account-selection.utils.ts b/ui/components/providers/organizations/org-account-selection.utils.ts index 59d74b90d0..23929dde94 100644 --- a/ui/components/providers/organizations/org-account-selection.utils.ts +++ b/ui/components/providers/organizations/org-account-selection.utils.ts @@ -2,8 +2,21 @@ import { CONNECTION_TEST_STATUS, ConnectionTestStatus, } from "@/types/organizations"; +import { + CONNECTION_CHECK_STATUS, + type ConnectionCheckStatus, +} from "@/types/providers"; const DEFAULT_POLL_DELAYS_MS = [2000, 3000, 5000] as const; +export const CONNECTION_CHECK_DEFAULT_DELAYS_MS = DEFAULT_POLL_DELAYS_MS; + +/** + * `provider-connection-check` has a 120s hard time limit in Celery + * (api/src/backend/config/celery.py `task_annotations`). With the delay ladder + * above -- 2s, 3s, then 5s repeating -- 32 retries cover roughly 155s, + * comfortably past the task's hard limit plus queueing/network slack. + */ +export const CONNECTION_CHECK_MAX_RETRIES = 32; interface BuildCandidateToProviderMapParams { selectedCandidateIds: string[]; @@ -27,10 +40,20 @@ interface PollConnectionTasksOptions /** Called once per task, the round it reaches a terminal state. */ onSettled: (taskId: string, result: PollConnectionTaskResult) => void; getTasksByIds?: (taskIds: string[]) => Promise>; + /** + * Called once per task still pending after `maxRetries` is exhausted, so the + * caller can re-read the provider's persisted connection state instead of + * reporting a flat timeout -- the backend task may still be running past the + * wait, or may have already finished with the UI no longer polling it. + * Returning `null` falls back to the timeout message. + */ + resolveExhausted?: ( + taskId: string, + ) => Promise; } export interface PollConnectionTaskResult { - success: boolean; + status: ConnectionCheckStatus; error?: string; } @@ -123,7 +146,10 @@ function readConnectionOutcome( taskResponse: unknown, ): PollConnectionTaskResult | null { if (isRecord(taskResponse) && typeof taskResponse.error === "string") { - return { success: false, error: taskResponse.error }; + return { + status: CONNECTION_CHECK_STATUS.FAILED, + error: taskResponse.error, + }; } const data = @@ -138,10 +164,10 @@ function readConnectionOutcome( const connected = typeof result?.connected === "boolean" ? result.connected : true; if (connected) { - return { success: true }; + return { status: CONNECTION_CHECK_STATUS.SUCCESS }; } return { - success: false, + status: CONNECTION_CHECK_STATUS.FAILED, error: (typeof result?.error === "string" && result.error) || "Connection failed for this account.", @@ -150,7 +176,7 @@ function readConnectionOutcome( if (state === "failed") { return { - success: false, + status: CONNECTION_CHECK_STATUS.FAILED, error: (typeof result?.error === "string" && result.error) || "Connection test task failed.", @@ -158,7 +184,10 @@ function readConnectionOutcome( } if (!state || !IN_PROGRESS_TASK_STATES.has(state)) { - return { success: false, error: "Unexpected task state." }; + return { + status: CONNECTION_CHECK_STATUS.FAILED, + error: "Unexpected task state.", + }; } return null; @@ -180,9 +209,10 @@ export async function pollConnectionTasks( getTasksByIds, sleep = async (ms: number) => new Promise((resolve) => setTimeout(resolve, ms)), - maxRetries = 20, + maxRetries = CONNECTION_CHECK_MAX_RETRIES, delaysMs = [...DEFAULT_POLL_DELAYS_MS], signal, + resolveExhausted, }: PollConnectionTasksOptions, ): Promise { const pending = new Set(taskIds.filter(Boolean)); @@ -199,7 +229,7 @@ export async function pollConnectionTasks( const settleRemaining = (error: string) => { for (const taskId of Array.from(pending)) { - onSettled(taskId, { success: false, error }); + onSettled(taskId, { status: CONNECTION_CHECK_STATUS.FAILED, error }); } pending.clear(); }; @@ -239,13 +269,46 @@ export async function pollConnectionTasks( await sleepWithAbort(getPollingDelay(attempt, delaysMs), sleep, signal); } + if (resolveExhausted) { + // Sequential, not `Promise.all`: each call goes through its own + // `getProvider` server action, and client-invoked server actions run one + // at a time through Next's action queue (see `pollConnectionTasks`'s own + // batched read above) -- running them "concurrently" from here would not + // shorten the wait, only reorder it. + for (const taskId of Array.from(pending)) { + if (signal?.aborted) { + settleRemaining("Connection test cancelled."); + return; + } + + const resolved = await resolveExhausted(taskId); + + // The signal can abort while `resolveExhausted` itself is in flight; its + // result must not be accepted after that, or a check the caller has + // already moved on from could still report success. + if (signal?.aborted) { + settleRemaining("Connection test cancelled."); + return; + } + + if (resolved) { + pending.delete(taskId); + onSettled(taskId, resolved); + } + } + } + settleRemaining("Connection test timed out."); } /** * Polls a generic async task until it settles. Unlike {@link pollConnectionTasks} * it does not interpret a connection result; it is used for organization/node - * deletion, which the API answers with a `202` + task. + * deletion, which the API answers with a `202` + task. Its result is typed with + * `ConnectionCheckStatus` only because that is the connection-specific alias of + * the generic `TASK_OUTCOME` (`types/tasks.ts`) already in scope here -- the + * three outcomes (succeeded / failed / still running) apply to any polled task, + * not just a connection check. */ export async function pollTaskCompletion( taskId: string, @@ -267,16 +330,25 @@ export async function pollTaskCompletion( for (let attempt = 0; attempt < maxRetries; attempt += 1) { if (signal?.aborted) { - return { success: false, error: "Deletion cancelled." }; + return { + status: CONNECTION_CHECK_STATUS.FAILED, + error: "Deletion cancelled.", + }; } const taskResponse = await taskFetcher(taskId); if (signal?.aborted) { - return { success: false, error: "Deletion cancelled." }; + return { + status: CONNECTION_CHECK_STATUS.FAILED, + error: "Deletion cancelled.", + }; } if (isRecord(taskResponse) && typeof taskResponse.error === "string") { - return { success: false, error: taskResponse.error }; + return { + status: CONNECTION_CHECK_STATUS.FAILED, + error: taskResponse.error, + }; } const data = @@ -289,12 +361,12 @@ export async function pollTaskCompletion( const result = isRecord(attributes?.result) ? attributes.result : null; if (state === "completed") { - return { success: true }; + return { status: CONNECTION_CHECK_STATUS.SUCCESS }; } if (state === "failed") { return { - success: false, + status: CONNECTION_CHECK_STATUS.FAILED, error: (typeof result?.error === "string" && result.error) || "The deletion task failed.", @@ -303,17 +375,26 @@ export async function pollTaskCompletion( // A cancelled task is a real terminal state, not an unreadable one. if (state === "cancelled") { - return { success: false, error: "The deletion was cancelled." }; + return { + status: CONNECTION_CHECK_STATUS.FAILED, + error: "The deletion was cancelled.", + }; } if (!state || !IN_PROGRESS_TASK_STATES.has(state)) { - return { success: false, error: "Unexpected task state." }; + return { + status: CONNECTION_CHECK_STATUS.FAILED, + error: "Unexpected task state.", + }; } await sleepWithAbort(getPollingDelay(attempt, delaysMs), sleep, signal); } - return { success: false, error: "Deletion timed out." }; + return { + status: CONNECTION_CHECK_STATUS.FAILED, + error: "Deletion timed out.", + }; } export function getLaunchableProviderIds( diff --git a/ui/components/providers/organizations/org-setup-form.tsx b/ui/components/providers/organizations/org-setup-form.tsx index 5d2640dd4b..839cd3a365 100644 --- a/ui/components/providers/organizations/org-setup-form.tsx +++ b/ui/components/providers/organizations/org-setup-form.tsx @@ -10,6 +10,10 @@ import { z } from "zod"; import { updateOrganizationName } from "@/actions/organizations/organizations"; import { AWSProviderBadge } from "@/components/icons/providers-badge"; +import { + AWS_ONBOARDING_METHOD, + AwsOnboardingMethodTabs, +} from "@/components/providers/wizard/steps/aws/aws-onboarding-method-tabs"; import type { WizardFooterConfig } from "@/components/providers/wizard/steps/footer-controls"; import { WIZARD_FOOTER_ACTION_TYPE } from "@/components/providers/wizard/steps/footer-controls"; import type { OrgWizardIntent } from "@/components/providers/wizard/types"; @@ -71,6 +75,8 @@ interface OrgSetupFormProps { onBack: () => void; onClose?: () => void; onNext: () => void; + /** Keeps the single/organization tabs on screen; absent when the flow was entered directly. */ + onSelectSingleAccount?: () => void; onFooterChange: (config: WizardFooterConfig) => void; onPhaseChange: (phase: OrgSetupPhase) => void; initialPhase?: OrgSetupPhase; @@ -82,6 +88,7 @@ export function OrgSetupForm({ onBack, onClose, onNext, + onSelectSingleAccount, onFooterChange, onPhaseChange, initialPhase = ORG_SETUP_PHASE.DETAILS, @@ -326,6 +333,13 @@ export function OrgSetupForm({
+ {onSelectSingleAccount && ( + + )} +

Enter the Organization ID for the accounts you want to add to Prowler. diff --git a/ui/components/providers/providers-accounts-view.test.tsx b/ui/components/providers/providers-accounts-view.test.tsx index da2d3c10ef..abd1487e0d 100644 --- a/ui/components/providers/providers-accounts-view.test.tsx +++ b/ui/components/providers/providers-accounts-view.test.tsx @@ -1,18 +1,24 @@ import { render, screen } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; import type { ReactNode } from "react"; -import { afterEach, describe, expect, it, vi } from "vitest"; +import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; +import { + PROVIDER_FUNNEL_EVENT, + type ProviderFunnelDetail, +} from "@/lib/provider-funnel/provider-funnel-events"; import type { FilterOption, MetaDataProps, ProviderProps } from "@/types"; import type { ProvidersTableRow } from "@/types/providers-table"; import { SCAN_SCHEDULE_CAPABILITY } from "@/types/schedules"; const { + onboardingTriggerSpy, providersAccountsTableSpy, refreshMock, replaceMock, searchParamsValue, } = vi.hoisted(() => ({ + onboardingTriggerSpy: vi.fn(), providersAccountsTableSpy: vi.fn(), refreshMock: vi.fn(), replaceMock: vi.fn(), @@ -28,6 +34,14 @@ vi.mock("next/navigation", () => ({ useSearchParams: () => new URLSearchParams(searchParamsValue.current), })); +vi.mock("@/components/onboarding", () => ({ + OnboardingTrigger: (props: { startAtTarget?: string }) => { + onboardingTriggerSpy(props); + return null; + }, + PageReady: () => null, +})); + vi.mock("@/components/providers/table", () => ({ SkeletonTableProviders: () =>

, })); @@ -132,9 +146,21 @@ const disconnectedProviders: ProviderProps[] = [ ]; describe("ProvidersAccountsView", () => { + const funnelSignals: ProviderFunnelDetail[] = []; + const recordFunnelSignal: EventListener = (event) => { + funnelSignals.push((event as CustomEvent).detail); + }; + + beforeEach(() => { + funnelSignals.length = 0; + window.addEventListener(PROVIDER_FUNNEL_EVENT, recordFunnelSignal); + }); + afterEach(() => { + window.removeEventListener(PROVIDER_FUNNEL_EVENT, recordFunnelSignal); vi.restoreAllMocks(); providersAccountsTableSpy.mockClear(); + onboardingTriggerSpy.mockClear(); searchParamsValue.current = ""; window.history.replaceState({}, "", "/"); }); @@ -273,6 +299,143 @@ describe("ProvidersAccountsView", () => { expect(replaceMock).not.toHaveBeenCalled(); }); + it("signals which control opened the wizard", async () => { + // Given + const user = userEvent.setup(); + render( + , + ); + + // When + await user.click(screen.getByRole("button", { name: "Add Provider" })); + + // Then + expect(funnelSignals).toEqual([ + { step: "wizard_opened", source: "page_button" }, + ]); + }); + + it("signals the empty-state CTA as the wizard entry point", async () => { + // Given + const user = userEvent.setup(); + render( + , + ); + + // When + await user.click( + screen.getByRole("button", { name: /open add provider modal/i }), + ); + + // Then + expect(funnelSignals).toEqual([ + { step: "wizard_opened", source: "empty_state" }, + ]); + }); + + it("signals the entry point carried in the URL and cleans it on close", async () => { + // Given + searchParamsValue.current = + "tab=connected&addProvider=true&addProviderSource=sidebar_cta"; + const replaceStateSpy = vi.spyOn(window.history, "replaceState"); + const user = userEvent.setup(); + + render( + , + ); + + // Then + expect(funnelSignals).toEqual([ + { step: "wizard_opened", source: "sidebar_cta" }, + ]); + + // When + await user.click(screen.getByRole("button", { name: /close/i })); + + // Then + expect(replaceStateSpy).toHaveBeenCalledWith( + null, + "", + "/providers?tab=connected", + ); + }); + + it("treats an unknown URL entry point as a plain URL open", () => { + // Given + searchParamsValue.current = "addProvider=true&addProviderSource=made_up"; + + // When + render( + , + ); + + // Then + expect(funnelSignals).toEqual([{ step: "wizard_opened", source: "url" }]); + }); + + it("starts the tour at the provider-type step when the wizard is already open", () => { + // Given + searchParamsValue.current = "addProvider=true&onboarding=add-provider"; + + // When + render( + , + ); + + // Then: the welcome and "open the wizard" steps have nothing left to ask for. + expect(onboardingTriggerSpy).toHaveBeenLastCalledWith( + expect.objectContaining({ startAtTarget: "provider-type" }), + ); + }); + + it("lets the tour start from its first step while the wizard is closed", () => { + // Given / When + render( + , + ); + + // Then + expect(onboardingTriggerSpy).toHaveBeenLastCalledWith( + expect.objectContaining({ startAtTarget: undefined }), + ); + }); + it("keeps filters and table visible when providers are disconnected", () => { // Given/When render( diff --git a/ui/components/providers/providers-accounts-view.tsx b/ui/components/providers/providers-accounts-view.tsx index 2bbd24c0a1..b94202af38 100644 --- a/ui/components/providers/providers-accounts-view.tsx +++ b/ui/components/providers/providers-accounts-view.tsx @@ -16,10 +16,19 @@ import type { ProviderWizardInitialData, } from "@/components/providers/wizard/types"; import { Alert, AlertDescription } from "@/components/shadcn/alert"; +import { useMountEffect } from "@/hooks/use-mount-effect"; import { getFlowById } from "@/lib/onboarding"; +import { + dispatchProviderFunnel, + PROVIDER_FUNNEL_STEP, + WIZARD_OPEN_SOURCE, + type WizardOpenSource, +} from "@/lib/provider-funnel/provider-funnel-events"; import { ADD_PROVIDER_SEARCH_PARAM, ADD_PROVIDER_SEARCH_VALUE, + ADD_PROVIDER_SOURCE_PARAM, + resolveAddProviderSource, } from "@/lib/providers-navigation"; import { ADD_PROVIDER_TOUR_TARGETS, @@ -102,7 +111,25 @@ export function ProvidersAccountsView({ OrgWizardInitialData | undefined >(undefined); - const openProviderWizard = (initialData?: ProviderWizardInitialData) => { + const signalWizardOpened = (source: WizardOpenSource) => + dispatchProviderFunnel({ + step: PROVIDER_FUNNEL_STEP.WIZARD_OPENED, + source, + }); + + // A URL-opened wizard never goes through openProviderWizard, so signal it on mount. + useMountEffect(() => { + if (!shouldOpenProviderWizardFromUrl) return; + signalWizardOpened( + resolveAddProviderSource(searchParams.get(ADD_PROVIDER_SOURCE_PARAM)), + ); + }); + + const openProviderWizard = ( + source: WizardOpenSource, + initialData?: ProviderWizardInitialData, + ) => { + signalWizardOpened(source); setOrgWizardInitialData(undefined); setProviderWizardInitialData(initialData); setIsProviderWizardOpen(true); @@ -130,6 +157,7 @@ export function ProvidersAccountsView({ if (searchParams.has(ADD_PROVIDER_SEARCH_PARAM)) { const params = new URLSearchParams(searchParams.toString()); params.delete(ADD_PROVIDER_SEARCH_PARAM); + params.delete(ADD_PROVIDER_SOURCE_PARAM); const query = params.toString(); window.history.replaceState( null, @@ -146,6 +174,12 @@ export function ProvidersAccountsView({ {/* Signals the navbar that this route's data has loaded (enables the replay icon). */} @@ -154,7 +188,9 @@ export function ProvidersAccountsView({ openProviderWizard()} + onOpenWizard={() => + openProviderWizard(WIZARD_OPEN_SOURCE.EMPTY_STATE) + } ctaTourId="add-provider-trigger" /> ) : ( @@ -175,7 +211,11 @@ export function ProvidersAccountsView({ actions={ <> - openProviderWizard()} /> + + openProviderWizard(WIZARD_OPEN_SOURCE.PAGE_BUTTON) + } + /> } /> @@ -186,7 +226,9 @@ export function ProvidersAccountsView({ scanScheduleCapability={scanScheduleCapability} scanConfigs={scanConfigs} scanConfigStatus={scanConfigStatus} - onOpenProviderWizard={openProviderWizard} + onOpenProviderWizard={(initialData) => + openProviderWizard(WIZARD_OPEN_SOURCE.ROW_ACTION, initialData) + } onOpenOrganizationWizard={openOrganizationWizard} />
diff --git a/ui/components/providers/table/data-table-row-actions.test.tsx b/ui/components/providers/table/data-table-row-actions.test.tsx index c65863abd2..ae0fe7c244 100644 --- a/ui/components/providers/table/data-table-row-actions.test.tsx +++ b/ui/components/providers/table/data-table-row-actions.test.tsx @@ -24,6 +24,7 @@ import { ORGANIZATION_TYPE, type OrganizationType, } from "@/types/organizations"; +import { CONNECTION_CHECK_STATUS } from "@/types/providers"; import { PROVIDERS_GROUP_KIND, PROVIDERS_ROW_TYPE, @@ -35,18 +36,36 @@ import { SCAN_SCHEDULE_CAPABILITY } from "@/types/schedules"; const { checkConnectionProviderMock, + getProviderConnectionBaselinesMock, getScheduleMock, getTasksByIdsMock, + pollConnectionTasksMock, pushMock, + realPollConnectionTasksHolder, + resolveProviderConnectionStateMock, revalidateProvidersMock, startProviderConnectionChecksMock, + testProviderConnectionMock, + toastMock, } = vi.hoisted(() => ({ checkConnectionProviderMock: vi.fn(), + getProviderConnectionBaselinesMock: vi.fn(), getScheduleMock: vi.fn(), getTasksByIdsMock: vi.fn(), + pollConnectionTasksMock: vi.fn(), pushMock: vi.fn(), + // Mutable holder for the real `pollConnectionTasks`, captured once the + // module mock factory below runs, and read fresh in `beforeEach` since + // `mockReset: true` clears `pollConnectionTasksMock`'s implementation + // before every test. + realPollConnectionTasksHolder: {} as { + current?: (...args: unknown[]) => unknown; + }, + resolveProviderConnectionStateMock: vi.fn(), revalidateProvidersMock: vi.fn(), startProviderConnectionChecksMock: vi.fn(), + testProviderConnectionMock: vi.fn(), + toastMock: vi.fn(), })); vi.mock("next/navigation", () => ({ @@ -59,6 +78,7 @@ vi.mock("@/actions/organizations/organizations", () => ({ vi.mock("@/actions/providers/providers", () => ({ checkConnectionProvider: checkConnectionProviderMock, + getProviderConnectionBaselines: getProviderConnectionBaselinesMock, revalidateProviders: revalidateProvidersMock, startProviderConnectionChecks: startProviderConnectionChecksMock, })); @@ -124,13 +144,28 @@ vi.mock("@/components/scans/schedule/edit-scan-schedule-modal", () => ({ vi.mock("@/components/shadcn", async (importOriginal) => ({ ...(await importOriginal>()), - useToast: () => ({ toast: vi.fn() }), + useToast: () => ({ toast: toastMock }), })); vi.mock("@/lib/provider-helpers", () => ({ - testProviderConnection: vi.fn(), + resolveProviderConnectionState: resolveProviderConnectionStateMock, + testProviderConnection: testProviderConnectionMock, })); +vi.mock( + "@/components/providers/organizations/org-account-selection.utils", + async (importOriginal) => { + const actual = + await importOriginal< + typeof import("@/components/providers/organizations/org-account-selection.utils") + >(); + realPollConnectionTasksHolder.current = actual.pollConnectionTasks as ( + ...args: unknown[] + ) => unknown; + return { ...actual, pollConnectionTasks: pollConnectionTasksMock }; + }, +); + import { DataTableRowActions } from "./data-table-row-actions"; const createRow = (hasSecret = false) => @@ -301,6 +336,12 @@ describe("DataTableRowActions", () => { }); beforeEach(() => { + // `mockReset: true` (vitest.config.ts) clears this before every test, so + // the real implementation is the default and tests only override it when + // they need to simulate an exhausted poll. + pollConnectionTasksMock.mockImplementation((...args: unknown[]) => + realPollConnectionTasksHolder.current?.(...args), + ); getScheduleMock.mockResolvedValue({ data: { type: "schedules", @@ -311,6 +352,7 @@ describe("DataTableRowActions", () => { }, }, }); + getProviderConnectionBaselinesMock.mockResolvedValue({}); }); it("renders Add Credentials for provider rows without credentials", async () => { @@ -780,6 +822,173 @@ describe("DataTableRowActions", () => { expect(checkConnectionProviderMock).not.toHaveBeenCalled(); }); + it("falls back to the provider's persisted state for a task still pending once the bulk wait is exhausted", async () => { + // Given: the batch poll exhausts its retries for both tasks; the component + // must re-read each provider's connection state instead of reporting a + // flat timeout. + const user = userEvent.setup(); + const testableProviderIds = ["provider-child-1", "provider-standalone"]; + startProviderConnectionChecksMock.mockResolvedValue({ + "provider-child-1": { taskId: "task-1" }, + "provider-standalone": { taskId: "task-2" }, + }); + // The baseline read before dispatch: one provider has a prior stored check, + // the other has never been checked. + getProviderConnectionBaselinesMock.mockResolvedValue({ + "provider-child-1": "2025-01-01T00:00:00Z", + "provider-standalone": null, + }); + pollConnectionTasksMock.mockImplementation( + async (taskIds: string[], { onSettled, resolveExhausted }) => { + for (const taskId of taskIds) { + const resolved = resolveExhausted + ? await resolveExhausted(taskId) + : null; + onSettled( + taskId, + resolved ?? { + status: CONNECTION_CHECK_STATUS.FAILED, + error: "Connection test timed out.", + }, + ); + } + }, + ); + resolveProviderConnectionStateMock.mockImplementation( + async (providerId: string) => + providerId === "provider-standalone" + ? { status: CONNECTION_CHECK_STATUS.SUCCESS, error: null } + : { + status: CONNECTION_CHECK_STATUS.FAILED, + error: "Connection was not confirmed. Test the connection again.", + }, + ); + + render( + , + ); + + // When + await user.click(screen.getByRole("button")); + await user.click(screen.getByText("Test Connections (2)")); + + // Then: each pending task is resolved from the provider's own record, using + // the baseline captured for that specific provider before dispatch. + await vi.waitFor(() => + expect(revalidateProvidersMock).toHaveBeenCalledTimes(1), + ); + expect(getProviderConnectionBaselinesMock).toHaveBeenCalledWith( + testableProviderIds, + ); + expect(resolveProviderConnectionStateMock).toHaveBeenCalledWith( + "provider-child-1", + "2025-01-01T00:00:00Z", + ); + expect(resolveProviderConnectionStateMock).toHaveBeenCalledWith( + "provider-standalone", + null, + ); + }); + + it("shows a neutral toast, not a failure, when a single test is still running past the wait", async () => { + // Given: the exhausted single-provider test cannot confirm an outcome yet. + const user = userEvent.setup(); + testProviderConnectionMock.mockResolvedValue({ + status: CONNECTION_CHECK_STATUS.PENDING, + error: + "The connection test is still running. Refresh in a moment to see the result.", + }); + + render( + , + ); + + // When + await user.click(screen.getByRole("button")); + await user.click(screen.getByText("Test Connection")); + + // Then: no destructive toast for a check that is merely still running. + await vi.waitFor(() => expect(toastMock).toHaveBeenCalledTimes(1)); + expect(toastMock).toHaveBeenCalledWith( + expect.not.objectContaining({ variant: "destructive" }), + ); + expect(toastMock.mock.calls[0][0].title).not.toMatch(/failed/i); + }); + + it("does not count a still-running bulk result as failed", async () => { + // Given: one provider settles successfully, the other is still running + // once the bulk wait is exhausted. + const user = userEvent.setup(); + const testableProviderIds = ["provider-child-1", "provider-standalone"]; + startProviderConnectionChecksMock.mockResolvedValue({ + "provider-child-1": { taskId: "task-1" }, + "provider-standalone": { taskId: "task-2" }, + }); + pollConnectionTasksMock.mockImplementation( + async (taskIds: string[], { onSettled, resolveExhausted }) => { + for (const taskId of taskIds) { + const resolved = resolveExhausted + ? await resolveExhausted(taskId) + : null; + onSettled( + taskId, + resolved ?? { + status: CONNECTION_CHECK_STATUS.FAILED, + error: "Connection test timed out.", + }, + ); + } + }, + ); + resolveProviderConnectionStateMock.mockImplementation( + async (providerId: string) => + providerId === "provider-standalone" + ? { status: CONNECTION_CHECK_STATUS.SUCCESS, error: null } + : { + status: CONNECTION_CHECK_STATUS.PENDING, + error: "The connection test is still running.", + }, + ); + + render( + , + ); + + // When + await user.click(screen.getByRole("button")); + await user.click(screen.getByText("Test Connections (2)")); + + // Then: not styled as a failure — no destructive toast. + await vi.waitFor(() => expect(toastMock).toHaveBeenCalledTimes(1)); + expect(toastMock).toHaveBeenCalledWith( + expect.not.objectContaining({ variant: "destructive" }), + ); + }); + it("shows selected provider count in Test Connections when OU row has active selection", async () => { const user = userEvent.setup(); render( diff --git a/ui/components/providers/table/data-table-row-actions.tsx b/ui/components/providers/table/data-table-row-actions.tsx index bcbbd8e583..95692ced80 100644 --- a/ui/components/providers/table/data-table-row-actions.tsx +++ b/ui/components/providers/table/data-table-row-actions.tsx @@ -16,6 +16,7 @@ import { useState } from "react"; import { updateOrganizationName } from "@/actions/organizations/organizations"; import { updateProvider } from "@/actions/providers"; import { + getProviderConnectionBaselines, revalidateProviders, startProviderConnectionChecks, } from "@/actions/providers/providers"; @@ -42,7 +43,10 @@ import { getNodeLabel, organizationNameFallbackHint, } from "@/lib/organizations"; -import { testProviderConnection } from "@/lib/provider-helpers"; +import { + resolveProviderConnectionState, + testProviderConnection, +} from "@/lib/provider-helpers"; import { getScanScheduleCapability } from "@/lib/schedules"; import { isCloud } from "@/lib/shared/env"; import { @@ -52,6 +56,7 @@ import { OrgFlowType, } from "@/types/organizations"; import { PROVIDER_WIZARD_MODE } from "@/types/provider-wizard"; +import { CONNECTION_CHECK_STATUS } from "@/types/providers"; import { isProvidersOrganizationRow, PROVIDERS_GROUP_KIND, @@ -394,9 +399,15 @@ export function DataTableRowActions({ // asks for. let succeeded = 0; let failed = 0; - const pendingTaskIds: string[] = []; + let pending = 0; + const providerIdByTaskId = new Map(); try { + // Read before dispatch, so the fallback below can tell each provider's own + // check result apart from whatever (possibly stale) result was already on + // record -- by comparing values, not by comparing timestamps against the + // browser's clock. See `resolveProviderConnectionState`. + const connectionBaselines = await getProviderConnectionBaselines(ids); const outcomes = await startProviderConnectionChecks(ids); for (const id of ids) { @@ -408,34 +419,54 @@ export function DataTableRowActions({ continue; } - pendingTaskIds.push(outcome.taskId); + providerIdByTaskId.set(outcome.taskId, id); } - await pollConnectionTasks(pendingTaskIds, { + await pollConnectionTasks(Array.from(providerIdByTaskId.keys()), { onSettled: (_taskId, result) => { - if (result.success) { + if (result.status === CONNECTION_CHECK_STATUS.SUCCESS) { succeeded += 1; + } else if (result.status === CONNECTION_CHECK_STATUS.PENDING) { + pending += 1; } else { failed += 1; } }, + resolveExhausted: async (taskId) => { + const id = providerIdByTaskId.get(taskId); + if (!id) { + return null; + } + const state = await resolveProviderConnectionState( + id, + connectionBaselines[id], + ); + return { status: state.status, error: state.error ?? undefined }; + }, }); } catch { - failed = ids.length - succeeded; + failed = ids.length - succeeded - pending; } await revalidateProviders(); - if (failed === 0) { + if (failed === 0 && pending === 0) { toast({ title: "Connection test completed", description: `${succeeded} ${succeeded === 1 ? "provider" : "providers"} tested successfully.`, }); + } else if (failed === 0) { + toast({ + title: "Connection test still running", + description: `${succeeded} succeeded, ${pending} still running. Refresh in a moment to see the rest.`, + }); } else { toast({ variant: "destructive", title: "Connection test completed", - description: `${succeeded} succeeded, ${failed} failed out of ${ids.length} providers.`, + description: `${succeeded} succeeded, ${failed} failed${ + pending ? `, ${pending} still running` : "" + } out of ${ids.length} providers.`, }); } @@ -454,17 +485,22 @@ export function DataTableRowActions({ const result = await testProviderConnection(providerId); setLoading(false); - if (!result.connected) { + if (result.status === CONNECTION_CHECK_STATUS.SUCCESS) { + toast({ + title: "Connection test completed", + description: "Provider tested successfully.", + }); + } else if (result.status === CONNECTION_CHECK_STATUS.PENDING) { + toast({ + title: "Connection test still running", + description: result.error ?? "Refresh in a moment to see the result.", + }); + } else { toast({ variant: "destructive", title: "Connection test failed", description: result.error ?? "Unknown error", }); - } else { - toast({ - title: "Connection test completed", - description: "Provider tested successfully.", - }); } } }; diff --git a/ui/components/providers/wizard/hooks/use-provider-wizard-controller.test.tsx b/ui/components/providers/wizard/hooks/use-provider-wizard-controller.test.tsx index ea8892ea5b..459d25106a 100644 --- a/ui/components/providers/wizard/hooks/use-provider-wizard-controller.test.tsx +++ b/ui/components/providers/wizard/hooks/use-provider-wizard-controller.test.tsx @@ -1,6 +1,10 @@ import { act, renderHook, waitFor } from "@testing-library/react"; -import { beforeEach, describe, expect, it, vi } from "vitest"; +import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; +import { + PROVIDER_FUNNEL_EVENT, + type ProviderFunnelDetail, +} from "@/lib/provider-funnel/provider-funnel-events"; import { useOrgSetupStore } from "@/store/organizations/store"; import { useProviderWizardStore } from "@/store/provider-wizard/store"; import { ORG_WIZARD_STEP, ORGANIZATION_TYPE } from "@/types/organizations"; @@ -40,7 +44,18 @@ vi.mock("next-auth/react", () => ({ })); describe("useProviderWizardController", () => { + const funnelSignals: ProviderFunnelDetail[] = []; + const recordFunnelSignal: EventListener = (event) => { + funnelSignals.push((event as CustomEvent).detail); + }; + + afterEach(() => { + window.removeEventListener(PROVIDER_FUNNEL_EVENT, recordFunnelSignal); + }); + beforeEach(() => { + funnelSignals.length = 0; + window.addEventListener(PROVIDER_FUNNEL_EVENT, recordFunnelSignal); vi.useRealTimers(); vi.clearAllMocks(); requestOpenOnWizardCloseMock.mockClear(); @@ -144,6 +159,70 @@ describe("useProviderWizardController", () => { expect(refreshMock).toHaveBeenCalledTimes(1); }); + it("signals the step where the wizard was left and that no provider was created", () => { + // Given + const { result } = renderHook(() => + useProviderWizardController({ open: true, onOpenChange: vi.fn() }), + ); + + // When + act(() => { + result.current.handleClose(); + }); + + // Then + expect(funnelSignals).toEqual([ + { step: "wizard_closed", lastStep: "connect", providerCreated: false }, + ]); + }); + + it("signals a close after the provider was created, from the step reached", () => { + // Given + const { result } = renderHook(() => + useProviderWizardController({ open: true, onOpenChange: vi.fn() }), + ); + act(() => { + useProviderWizardStore.getState().setProvider({ + id: "provider-1", + type: "aws", + uid: "123456789012", + alias: null, + }); + result.current.setCurrentStep(PROVIDER_WIZARD_STEP.TEST); + }); + + // When + act(() => { + result.current.handleClose(); + }); + + // Then + expect(funnelSignals).toEqual([ + { step: "wizard_closed", lastStep: "test", providerCreated: true }, + ]); + }); + + it("signals the organization method when the organizations flow opens", () => { + // Given + const { result } = renderHook(() => + useProviderWizardController({ open: true, onOpenChange: vi.fn() }), + ); + + // When + act(() => { + result.current.openOrganizationsFlow(ORGANIZATION_TYPE.AZURE); + }); + + // Then + expect(funnelSignals).toEqual([ + { + step: "method_selected", + providerType: "azure", + method: "organization", + }, + ]); + }); + it("hydrates update mode when initial data is provided", async () => { // Given const onOpenChange = vi.fn(); @@ -251,6 +330,8 @@ describe("useProviderWizardController", () => { expect(result.current.wizardVariant).toBe("provider"); expect(result.current.isProviderFlow).toBe(true); expect(result.current.currentStep).toBe(PROVIDER_WIZARD_STEP.CONNECT); + // Back lands on the AWS connect step the tabs live on, not the provider picker. + expect(result.current.providerTypeHint).toBe("aws"); }); it("moves to launch step after a successful connection test in add mode", () => { diff --git a/ui/components/providers/wizard/hooks/use-provider-wizard-controller.ts b/ui/components/providers/wizard/hooks/use-provider-wizard-controller.ts index 0059c43a77..ee78b91ed9 100644 --- a/ui/components/providers/wizard/hooks/use-provider-wizard-controller.ts +++ b/ui/components/providers/wizard/hooks/use-provider-wizard-controller.ts @@ -4,6 +4,11 @@ import { useRouter } from "next/navigation"; import { useEffect, useRef, useState } from "react"; import { DOCS_URLS, getProviderHelpText } from "@/lib/external-urls"; +import { + dispatchProviderFunnel, + PROVIDER_FUNNEL_METHOD, + PROVIDER_FUNNEL_STEP, +} from "@/lib/provider-funnel/provider-funnel-events"; import { isCloud } from "@/lib/shared/env"; import { endActiveTour } from "@/lib/tours/use-driver-tour"; import { useOnboardingCheckpointStore } from "@/store/onboarding-checkpoint"; @@ -44,6 +49,20 @@ const ORG_DOCS_URL = { [ORGANIZATION_TYPE.GCP]: DOCS_URLS.GCP_ORGANIZATIONS, } as const satisfies Record; +// Stable names for the abandonment signal; the numeric step ids are not a contract. +const PROVIDER_STEP_NAME = { + [PROVIDER_WIZARD_STEP.CONNECT]: "connect", + [PROVIDER_WIZARD_STEP.CREDENTIALS]: "credentials", + [PROVIDER_WIZARD_STEP.TEST]: "test", + [PROVIDER_WIZARD_STEP.LAUNCH]: "launch", +} as const satisfies Record; + +const ORG_STEP_NAME = { + [ORG_WIZARD_STEP.SETUP]: "organizations_setup", + [ORG_WIZARD_STEP.VALIDATE]: "organizations_validate", + [ORG_WIZARD_STEP.LAUNCH]: "organizations_launch", +} as const satisfies Record; + const EMPTY_FOOTER_CONFIG: WizardFooterConfig = { showBack: false, backLabel: "Back", @@ -196,6 +215,11 @@ export function useProviderWizardController({ ]); const isOrgDirectEntry = Boolean(orgInitialData); + // Opened on an existing account's credentials, so the one-step AWS flow is not + // in play. Same three fields the hydration above requires to start on CREDENTIALS. + const isDirectCredentialsEntry = Boolean( + initialProviderId && initialProviderType && initialProviderUid, + ); const handleClose = () => { // Closing the wizard at any point ends the add-provider tour; the checkpoint @@ -205,6 +229,15 @@ export function useProviderWizardController({ // Read providerId before reset clears it — non-null means a provider was connected. const connectedProviderId = useProviderWizardStore.getState().providerId; + dispatchProviderFunnel({ + step: PROVIDER_FUNNEL_STEP.WIZARD_CLOSED, + lastStep: + wizardVariant === WIZARD_VARIANT.PROVIDER + ? PROVIDER_STEP_NAME[currentStep] + : ORG_STEP_NAME[orgCurrentStep], + providerCreated: connectedProviderId !== null, + }); + resetProviderWizard(); resetOrgWizard(); setWizardVariant(WIZARD_VARIANT.PROVIDER); @@ -251,6 +284,11 @@ export function useProviderWizardController({ // Organizations diverges from the credentials path the tour guides toward; end // it so it doesn't dangle on a step that no longer fits. No-op off-onboarding. endActiveTour(); + dispatchProviderFunnel({ + step: PROVIDER_FUNNEL_STEP.METHOD_SELECTED, + providerType: orgType, + method: PROVIDER_FUNNEL_METHOD.ORGANIZATION, + }); resetOrgWizard(); setOrganizationType(orgType); setWizardVariant(WIZARD_VARIANT.ORGANIZATIONS); @@ -261,11 +299,14 @@ export function useProviderWizardController({ }; const backToProviderFlow = () => { + // The AWS organization flow is entered from the AWS connect step's tabs, so + // going back lands on that step again instead of the provider picker. + const cameFromAwsConnect = organizationType === ORGANIZATION_TYPE.AWS; resetOrgWizard(); setWizardVariant(WIZARD_VARIANT.PROVIDER); setCurrentStep(PROVIDER_WIZARD_STEP.CONNECT); setFooterConfig(EMPTY_FOOTER_CONFIG); - setProviderTypeHint(null); + setProviderTypeHint(cameFromAwsConnect ? "aws" : null); setOrgSetupPhase(ORG_SETUP_PHASE.DETAILS); }; @@ -287,6 +328,7 @@ export function useProviderWizardController({ handleClose, handleDialogOpenChange, handleTestSuccess, + isDirectCredentialsEntry, isOrgDirectEntry, isProviderFlow, mode, diff --git a/ui/components/providers/wizard/provider-wizard-modal.test.tsx b/ui/components/providers/wizard/provider-wizard-modal.test.tsx index fc913924bc..5a1115dbd9 100644 --- a/ui/components/providers/wizard/provider-wizard-modal.test.tsx +++ b/ui/components/providers/wizard/provider-wizard-modal.test.tsx @@ -1,24 +1,55 @@ import { act, render, screen, waitFor } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; -import { beforeEach, describe, expect, it, vi } from "vitest"; +import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; import { Toaster } from "@/components/shadcn/toast/Toaster"; import { resetToasts } from "@/components/shadcn/toast/use-toast"; +import { + PROVIDER_FUNNEL_EVENT, + type ProviderFunnelDetail, +} from "@/lib/provider-funnel/provider-funnel-events"; +import { endActiveTour } from "@/lib/tours/use-driver-tour"; import { useProviderWizardStore } from "@/store/provider-wizard/store"; +import { useUIStore } from "@/store/ui/store"; +import { CONNECTION_CHECK_STATUS } from "@/types/providers"; import { ProviderWizardModal } from "./provider-wizard-modal"; -const { addRegistryProvider, getInstalledRegistryProviderOptions } = vi.hoisted( - () => ({ - addRegistryProvider: vi.fn(), - getInstalledRegistryProviderOptions: vi.fn(), - }), -); +const { + addCredentialsProvider, + addProvider, + addRegistryProvider, + getInstalledRegistryProviderOptions, + testProviderConnection, + updateCredentialsProvider, + updateProvider, +} = vi.hoisted(() => ({ + addCredentialsProvider: vi.fn(), + addProvider: vi.fn(), + addRegistryProvider: vi.fn(), + getInstalledRegistryProviderOptions: vi.fn(), + testProviderConnection: vi.fn(), + updateCredentialsProvider: vi.fn(), + updateProvider: vi.fn(), +})); vi.mock("next/navigation", () => ({ useRouter: () => ({ refresh: vi.fn(), push: vi.fn() }), })); -vi.mock("@/actions/providers/providers", () => ({ addProvider: vi.fn() })); +vi.mock("next-auth/react", () => ({ + useSession: () => ({ + data: { tenantId: "tenant-abc" }, + status: "authenticated", + }), +})); +vi.mock("@/actions/providers/providers", () => ({ + addCredentialsProvider, + addProvider, + updateCredentialsProvider, + updateProvider, +})); +// The real module reaches next-auth through lib/helper -> auth.config. +vi.mock("@/lib/provider-helpers", () => ({ testProviderConnection })); vi.mock("@/actions/providers/registry-provider", () => ({ addRegistryProvider, })); @@ -37,12 +68,39 @@ vi.mock("@/lib/tours/use-driver-tour", () => ({ endActiveTour: vi.fn(), })); vi.mock("./steps/credentials-step", () => ({ - CredentialsStep: () =>

Credential details

, + CredentialsStep: ({ onBack }: { onBack: () => void }) => ( + <> +

Credential details

+ + + ), })); vi.mock("./steps/test-connection-step", () => ({ - TestConnectionStep: () => null, + TestConnectionStep: ({ + onResetCredentials, + }: { + onResetCredentials: () => void; + }) => ( + <> +

Connection test

+ + + ), +})); +vi.mock("./steps/launch-step", () => ({ + LaunchStep: ({ onBack }: { onBack: () => void }) => ( + <> +

Launch scan

+ + + ), })); -vi.mock("./steps/launch-step", () => ({ LaunchStep: () => null })); vi.mock("../organizations/azure-org-setup-form", () => ({ AzureOrgSetupForm: () => null, })); @@ -89,12 +147,24 @@ async function enterAccountDetails() { describe("provider wizard account creation", () => { beforeEach(() => { + // Registry discovery only runs in Cloud. + vi.stubEnv("UI_CLOUD_ENABLED", "true"); useProviderWizardStore.getState().reset(); resetToasts(); getInstalledRegistryProviderOptions.mockResolvedValue({ status: "ready", options: [{ type: "acme", label: "Acme Cloud" }], }); + testProviderConnection.mockResolvedValue({ + status: CONNECTION_CHECK_STATUS.SUCCESS, + error: null, + }); + updateCredentialsProvider.mockResolvedValue({ data: { id: "secret-1" } }); + updateProvider.mockResolvedValue({ data: { id: "provider-1" } }); + }); + + afterEach(() => { + vi.unstubAllEnvs(); }); it("shows progress, blocks repeat clicks, and advances after creation", async () => { @@ -126,6 +196,20 @@ describe("provider wizard account creation", () => { expect(await screen.findByText("Credential details")).toBeVisible(); }); + it("tells the rest of the app the tenant now has a provider", async () => { + // Given + useUIStore.setState({ hasProviders: false, hasProvidersResolved: true }); + addRegistryProvider.mockResolvedValueOnce(createdAccount); + const user = await enterAccountDetails(); + + // When + await user.click(screen.getByRole("button", { name: "Next" })); + await screen.findByText("Credential details"); + + // Then: the sidebar stops offering Add Provider without waiting for a reload. + expect(useUIStore.getState().hasProviders).toBe(true); + }); + it("restores Next after a failed creation and retries the same account", async () => { // Given const failure = { errors: [{ detail: "Creation failed. Try again." }] }; @@ -200,6 +284,198 @@ describe("provider wizard account creation", () => { expect(addRegistryProvider).toHaveBeenCalledTimes(2); }); + it("signals the provider type the user picked, once", async () => { + // Given + const funnelSignals: ProviderFunnelDetail[] = []; + const recordFunnelSignal: EventListener = (event) => { + funnelSignals.push((event as CustomEvent).detail); + }; + window.addEventListener(PROVIDER_FUNNEL_EVENT, recordFunnelSignal); + const user = userEvent.setup(); + render(); + + await screen.findByRole("option", { name: "Acme Cloud Registry" }); + + // When + await user.click( + screen.getByRole("option", { name: /Amazon Web Services/ }), + ); + await screen.findByRole("radio", { name: /IAM Role/ }); + window.removeEventListener(PROVIDER_FUNNEL_EVENT, recordFunnelSignal); + + // Then + expect(funnelSignals).toEqual([ + { step: "provider_type_selected", providerType: "aws" }, + ]); + }); + + describe("when the user picks AWS", () => { + const ROLE_ARN = "arn:aws:iam::123456789012:role/ProwlerScan"; + + async function pickAws() { + const user = userEvent.setup(); + render(); + await screen.findByRole("option", { name: "Acme Cloud Registry" }); + await user.click( + screen.getByRole("option", { name: /Amazon Web Services/ }), + ); + await screen.findByRole("textbox", { name: /Role ARN/ }); + return user; + } + + it("connects the account and its credentials in one step, then tests the connection", async () => { + // Given + addProvider.mockResolvedValue({ data: { id: "provider-1" } }); + addCredentialsProvider.mockResolvedValue({ data: { id: "secret-1" } }); + const user = await pickAws(); + + // When + await user.type( + screen.getByRole("textbox", { name: /Role ARN/ }), + ROLE_ARN, + ); + const connect = screen.getByRole("button", { name: "Connect account" }); + await waitFor(() => expect(connect).toBeEnabled()); + await user.click(connect); + + // Then: neither the credentials step nor the connection test shows up. + expect(await screen.findByText("Launch scan")).toBeVisible(); + expect(screen.queryByText("Credential details")).not.toBeInTheDocument(); + expect(screen.queryByText("Connection test")).not.toBeInTheDocument(); + expect(testProviderConnection).toHaveBeenCalledWith("provider-1"); + expect(useProviderWizardStore.getState()).toMatchObject({ + providerId: "provider-1", + secretId: "secret-1", + via: "role", + }); + }); + + it("keeps the account on the one-step form when the connection is refused", async () => { + // Given + addProvider.mockResolvedValue({ data: { id: "provider-1" } }); + addCredentialsProvider.mockResolvedValue({ data: { id: "secret-1" } }); + testProviderConnection.mockResolvedValue({ + status: CONNECTION_CHECK_STATUS.FAILED, + error: "The role could not be assumed.", + }); + const user = await pickAws(); + + // When + await user.type( + screen.getByRole("textbox", { name: /Role ARN/ }), + ROLE_ARN, + ); + const connect = screen.getByRole("button", { name: "Connect account" }); + await waitFor(() => expect(connect).toBeEnabled()); + await user.click(connect); + + // Then + expect( + await screen.findByText("The role could not be assumed."), + ).toBeVisible(); + expect(screen.queryByText("Launch scan")).not.toBeInTheDocument(); + expect(screen.getByRole("textbox", { name: /Role ARN/ })).toBeVisible(); + }); + + it("closes instead of launching a scan when AWS credentials are updated", async () => { + // Given: the row action opens an existing AWS provider's credentials. + addProvider.mockResolvedValue({ data: { id: "provider-1" } }); + addCredentialsProvider.mockResolvedValue({ data: { id: "secret-1" } }); + const onOpenChange = vi.fn(); + const user = userEvent.setup(); + render( + , + ); + + // When: Back reaches the AWS one-step form, still in update mode. + await user.click( + await screen.findByRole("button", { name: "Back to provider" }), + ); + await user.type( + await screen.findByRole("textbox", { name: /Role ARN/ }), + ROLE_ARN, + ); + const connect = screen.getByRole("button", { name: "Connect account" }); + await waitFor(() => expect(connect).toBeEnabled()); + await user.click(connect); + + // Then: an update never offers a scan. + await waitFor(() => expect(onOpenChange).toHaveBeenCalledWith(false)); + expect(screen.queryByText("Launch scan")).not.toBeInTheDocument(); + }); + + it("returns to the one-step form when the launch step is stepped back from", async () => { + // Given + addProvider.mockResolvedValue({ data: { id: "provider-1" } }); + addCredentialsProvider.mockResolvedValue({ data: { id: "secret-1" } }); + const user = await pickAws(); + await user.type( + screen.getByRole("textbox", { name: /Role ARN/ }), + ROLE_ARN, + ); + const connect = screen.getByRole("button", { name: "Connect account" }); + await waitFor(() => expect(connect).toBeEnabled()); + await user.click(connect); + await screen.findByText("Launch scan"); + + // When + await user.click(screen.getByRole("button", { name: "Back to form" })); + + // Then: AWS has no separate credentials step, so it lands on its own form. + expect( + await screen.findByRole("textbox", { name: /Role ARN/ }), + ).toBeVisible(); + expect(screen.queryByText("Credential details")).not.toBeInTheDocument(); + }); + + it("steps the tour aside once the account can be connected", async () => { + // Given + vi.mocked(endActiveTour).mockClear(); + const user = await pickAws(); + + // When + await user.type( + screen.getByRole("textbox", { name: /Role ARN/ }), + ROLE_ARN, + ); + + // Then: the footer sits outside the tour's spotlight, so the tour ends + // right when the user is ready to press Connect account. + await waitFor(() => + expect( + screen.getByRole("button", { name: "Connect account" }), + ).toBeEnabled(), + ); + expect(endActiveTour).toHaveBeenCalled(); + }); + + it("goes back to the provider list", async () => { + // Given + const user = await pickAws(); + + // When + await user.click(screen.getByRole("button", { name: "Back" })); + + // Then + expect( + await screen.findByRole("option", { name: /Microsoft Azure/ }), + ).toBeVisible(); + expect( + screen.queryByRole("textbox", { name: /Role ARN/ }), + ).not.toBeInTheDocument(); + }); + }); + it("keeps native providers available during a Registry discovery error and retries", async () => { // Given getInstalledRegistryProviderOptions.mockRejectedValueOnce( diff --git a/ui/components/providers/wizard/provider-wizard-modal.tsx b/ui/components/providers/wizard/provider-wizard-modal.tsx index fab1a4440f..f5bf93f789 100644 --- a/ui/components/providers/wizard/provider-wizard-modal.tsx +++ b/ui/components/providers/wizard/provider-wizard-modal.tsx @@ -12,22 +12,26 @@ import { DialogHeader, DialogTitle } from "@/components/shadcn/dialog"; import { Modal } from "@/components/shadcn/modal"; import { useScanScheduleCapability } from "@/hooks/use-scan-schedule-capability"; import { useScrollHint } from "@/hooks/use-scroll-hint"; +import { + dispatchProviderFunnel, + PROVIDER_FUNNEL_STEP, +} from "@/lib/provider-funnel/provider-funnel-events"; import { advanceActiveTour, endActiveTour } from "@/lib/tours/use-driver-tour"; import { ORG_SETUP_PHASE, ORG_WIZARD_STEP, ORGANIZATION_TYPE, } from "@/types/organizations"; -import { - PROVIDER_WIZARD_MODE, - PROVIDER_WIZARD_STEP, -} from "@/types/provider-wizard"; +import { PROVIDER_WIZARD_STEP } from "@/types/provider-wizard"; import type { ScanScheduleCapability } from "@/types/schedules"; import { useProviderWizardController } from "./hooks/use-provider-wizard-controller"; import { + getCredentialsRetryStep, + getLaunchBackStep, getOrganizationsStepperOffset, getProviderWizardDocsDestination, + getProviderWizardStepper, } from "./provider-wizard-modal.utils"; import { ConnectStep } from "./steps/connect-step"; import { CredentialsStep } from "./steps/credentials-step"; @@ -35,12 +39,7 @@ import { WIZARD_FOOTER_ACTION_TYPE } from "./steps/footer-controls"; import { LaunchStep } from "./steps/launch-step"; import { TestConnectionStep } from "./steps/test-connection-step"; import type { OrgWizardInitialData, ProviderWizardInitialData } from "./types"; -import { PROVIDER_WIZARD_STEPS, WizardStepper } from "./wizard-stepper"; - -const UPDATE_MODE_WIZARD_STEPS = PROVIDER_WIZARD_STEPS.slice( - 0, - PROVIDER_WIZARD_STEP.LAUNCH, -); +import { WizardStepper } from "./wizard-stepper"; interface ProviderWizardModalProps { open: boolean; @@ -70,6 +69,7 @@ export function ProviderWizardModal({ handleClose, handleDialogOpenChange, handleTestSuccess, + isDirectCredentialsEntry, isOrgDirectEntry, isProviderFlow, mode, @@ -78,6 +78,7 @@ export function ProviderWizardModal({ organizationType, orgCurrentStep, orgSetupPhase, + providerTypeHint, resolvedFooterConfig, setCurrentStep, setFooterConfig, @@ -102,6 +103,12 @@ export function ProviderWizardModal({ isScheduleCapabilityLoading, } = useScanScheduleCapability(scanScheduleCapability); const docsDestination = getProviderWizardDocsDestination(docsLink); + const providerStepper = getProviderWizardStepper({ + mode, + providerType: providerTypeHint, + currentStep, + isDirectCredentialsEntry, + }); return ( - {/* Anchors the add-provider tour's final step to the wizard content and - footer, keeping the real form controls clickable under the overlay. */} -
+
{isProviderFlow ? ( ) : ( -
+ {/* Anchors the add-provider tour's final step to the form column only, so + its popover has room on the left, under the stepper. */} +
{ setCurrentStep(PROVIDER_WIZARD_STEP.CREDENTIALS); // Reaching credentials is the tour's handoff point: end it so the // user continues on their own. No-op off-onboarding. endActiveTour(); }} + onCredentialsSaved={() => { + // AWS stored its credentials and tested the connection in this + // step, so it takes the same exit the test step took: an update + // closes the wizard, an add moves on to the launch step. + handleTestSuccess(); + endActiveTour(); + }} onSelectOrganizations={openOrganizationsFlow} onFooterChange={setFooterConfig} onProviderTypeChange={(providerType) => { // Picking a type reveals the account-detail inputs. Advance the tour // to its wizard-body step, pinned beside the form. No-op off-onboarding. if (providerType) advanceActiveTour(); + // The form re-reports the same type on re-render; signal a pick once. + if (providerType && providerType !== providerTypeHint) { + dispatchProviderFunnel({ + step: PROVIDER_FUNNEL_STEP.PROVIDER_TYPE_SELECTED, + providerType, + }); + } setProviderTypeHint(providerType); }} /> @@ -200,7 +219,13 @@ export function ProviderWizardModal({ - setCurrentStep(PROVIDER_WIZARD_STEP.CREDENTIALS) + setCurrentStep( + getCredentialsRetryStep({ + mode, + providerType: providerTypeHint, + isDirectCredentialsEntry, + }), + ) } onFooterChange={setFooterConfig} /> @@ -209,7 +234,14 @@ export function ProviderWizardModal({ {isProviderFlow && currentStep === PROVIDER_WIZARD_STEP.LAUNCH && ( setCurrentStep(PROVIDER_WIZARD_STEP.TEST)} + onBack={() => + setCurrentStep( + getLaunchBackStep({ + providerType: providerTypeHint, + isDirectCredentialsEntry, + }), + ) + } onClose={handleClose} onFooterChange={setFooterConfig} capability={resolvedScanScheduleCapability} @@ -225,6 +257,9 @@ export function ProviderWizardModal({ onBack={ isOrgDirectEntry ? handleClose : backToProviderFlow } + onSelectSingleAccount={ + isOrgDirectEntry ? undefined : backToProviderFlow + } onClose={handleClose} onNext={() => { setOrgCurrentStep(ORG_WIZARD_STEP.VALIDATE); @@ -347,7 +382,8 @@ export function ProviderWizardModal({ {(resolvedFooterConfig.showBack || resolvedFooterConfig.showSecondaryAction || resolvedFooterConfig.showAction) && ( -
+ // Outside the tour's spotlight, yet the way forward: keep it clickable. +
{resolvedFooterConfig.showBack && ( diff --git a/ui/components/providers/wizard/provider-wizard-modal.utils.test.ts b/ui/components/providers/wizard/provider-wizard-modal.utils.test.ts index 6e2b090fc9..1664f2a4cb 100644 --- a/ui/components/providers/wizard/provider-wizard-modal.utils.test.ts +++ b/ui/components/providers/wizard/provider-wizard-modal.utils.test.ts @@ -9,11 +9,162 @@ import { import { type KnownProviderType, PROVIDER_TYPES } from "@/types/providers"; import { + getCredentialsRetryStep, + getLaunchBackStep, getOrganizationsStepperOffset, getProviderWizardDocsDestination, getProviderWizardModalTitle, + getProviderWizardStepper, } from "./provider-wizard-modal.utils"; +describe("getProviderWizardStepper", () => { + const labels = (steps: { label: string }[]) => steps.map((s) => s.label); + + it("lists the four generic steps until a provider is picked", () => { + const stepper = getProviderWizardStepper({ + mode: PROVIDER_WIZARD_MODE.ADD, + providerType: null, + currentStep: PROVIDER_WIZARD_STEP.CONNECT, + }); + + expect(labels(stepper.steps)).toEqual([ + "Link a Provider", + "Authenticate Credentials", + "Validate Connection", + "Launch Scan", + ]); + expect(stepper.stepOffset).toBe(0); + }); + + it("leaves only two rows when adding an AWS account", () => { + const stepper = getProviderWizardStepper({ + mode: PROVIDER_WIZARD_MODE.ADD, + providerType: "aws", + currentStep: PROVIDER_WIZARD_STEP.CONNECT, + }); + + expect(labels(stepper.steps)).toEqual(["Link a Provider", "Launch Scan"]); + expect(stepper.stepOffset).toBe(0); + }); + + it("keeps the first AWS row active if the wizard ever lands on a folded step", () => { + const stepper = getProviderWizardStepper({ + mode: PROVIDER_WIZARD_MODE.ADD, + providerType: "aws", + currentStep: PROVIDER_WIZARD_STEP.CREDENTIALS, + }); + + expect(stepper.stepOffset).toBe(-PROVIDER_WIZARD_STEP.CREDENTIALS); + }); + + it("puts the AWS launch step on the second row", () => { + const stepper = getProviderWizardStepper({ + mode: PROVIDER_WIZARD_MODE.ADD, + providerType: "aws", + currentStep: PROVIDER_WIZARD_STEP.LAUNCH, + }); + + // LAUNCH is index 3 in the wizard but the second row of the AWS stepper. + expect(stepper.stepOffset).toBe(-2); + }); + + it("keeps the generic rows when adding credentials to a registered AWS account", () => { + const stepper = getProviderWizardStepper({ + mode: PROVIDER_WIZARD_MODE.ADD, + providerType: "aws", + currentStep: PROVIDER_WIZARD_STEP.CREDENTIALS, + isDirectCredentialsEntry: true, + }); + + expect(labels(stepper.steps)).toEqual([ + "Link a Provider", + "Authenticate Credentials", + "Validate Connection", + "Launch Scan", + ]); + expect(stepper.stepOffset).toBe(0); + }); + + it("keeps the generic rows for a provider that is not AWS", () => { + const stepper = getProviderWizardStepper({ + mode: PROVIDER_WIZARD_MODE.ADD, + providerType: "azure", + currentStep: PROVIDER_WIZARD_STEP.TEST, + }); + + expect(stepper.steps).toHaveLength(4); + expect(stepper.stepOffset).toBe(0); + }); + + it("still shows the credentials step when updating AWS credentials", () => { + const stepper = getProviderWizardStepper({ + mode: PROVIDER_WIZARD_MODE.UPDATE, + providerType: "aws", + currentStep: PROVIDER_WIZARD_STEP.CREDENTIALS, + }); + + expect(labels(stepper.steps)).toEqual([ + "Link a Provider", + "Authenticate Credentials", + "Validate Connection", + ]); + expect(stepper.stepOffset).toBe(0); + }); +}); + +describe("getCredentialsRetryStep", () => { + it("returns an AWS account being added to its one-step form", () => { + expect( + getCredentialsRetryStep({ + mode: PROVIDER_WIZARD_MODE.ADD, + providerType: "aws", + }), + ).toBe(PROVIDER_WIZARD_STEP.CONNECT); + }); + + it("returns to the credentials step when AWS credentials were added from the list", () => { + expect( + getCredentialsRetryStep({ + mode: PROVIDER_WIZARD_MODE.ADD, + providerType: "aws", + isDirectCredentialsEntry: true, + }), + ).toBe(PROVIDER_WIZARD_STEP.CREDENTIALS); + }); + + it("returns every other provider to the credentials step", () => { + expect( + getCredentialsRetryStep({ + mode: PROVIDER_WIZARD_MODE.ADD, + providerType: "azure", + }), + ).toBe(PROVIDER_WIZARD_STEP.CREDENTIALS); + }); +}); + +describe("getLaunchBackStep", () => { + it("returns an AWS account to its one-step form", () => { + expect(getLaunchBackStep({ providerType: "aws" })).toBe( + PROVIDER_WIZARD_STEP.CONNECT, + ); + }); + + it("returns every other provider to the connection test", () => { + expect(getLaunchBackStep({ providerType: "azure" })).toBe( + PROVIDER_WIZARD_STEP.TEST, + ); + }); + + it("returns to the connection test when AWS credentials were added from the list", () => { + expect( + getLaunchBackStep({ + providerType: "aws", + isDirectCredentialsEntry: true, + }), + ).toBe(PROVIDER_WIZARD_STEP.TEST); + }); +}); + describe("getOrganizationsStepperOffset", () => { it("keeps step 1 active during organization details", () => { const offset = getOrganizationsStepperOffset( diff --git a/ui/components/providers/wizard/provider-wizard-modal.utils.ts b/ui/components/providers/wizard/provider-wizard-modal.utils.ts index 1e8f6f9248..1a3932ff47 100644 --- a/ui/components/providers/wizard/provider-wizard-modal.utils.ts +++ b/ui/components/providers/wizard/provider-wizard-modal.utils.ts @@ -6,8 +6,93 @@ import { } from "@/types/organizations"; import { PROVIDER_WIZARD_MODE, + PROVIDER_WIZARD_STEP, ProviderWizardMode, + ProviderWizardStep, } from "@/types/provider-wizard"; +import type { ProviderType } from "@/types/providers"; + +import { + AWS_PROVIDER_WIZARD_STEPS, + PROVIDER_WIZARD_STEPS, +} from "./wizard-stepper"; + +const UPDATE_MODE_WIZARD_STEPS = PROVIDER_WIZARD_STEPS.slice( + 0, + PROVIDER_WIZARD_STEP.LAUNCH, +); + +const AWS_CONNECT_STEPPER_ROW = 0; +const AWS_LAUNCH_STEPPER_ROW = 1; + +interface ProviderWizardStepperInput { + mode: ProviderWizardMode; + providerType: ProviderType | null; + currentStep: ProviderWizardStep; + // "Add credentials" on a registered account opens on CREDENTIALS and still walks + // the separate steps, so it keeps the generic rows. + isDirectCredentialsEntry?: boolean; +} + +/** Rows for the provider-flow stepper plus the offset that maps `currentStep` onto them. */ +export function getProviderWizardStepper({ + mode, + providerType, + currentStep, + isDirectCredentialsEntry = false, +}: ProviderWizardStepperInput) { + if (mode === PROVIDER_WIZARD_MODE.UPDATE) { + return { steps: UPDATE_MODE_WIZARD_STEPS, stepOffset: 0 }; + } + if (providerType === "aws" && !isDirectCredentialsEntry) { + // Only CONNECT and LAUNCH are reachable here; CREDENTIALS and TEST have no + // row of their own, so anything short of LAUNCH folds onto the first row. + const stepOffset = + currentStep === PROVIDER_WIZARD_STEP.LAUNCH + ? AWS_LAUNCH_STEPPER_ROW - PROVIDER_WIZARD_STEP.LAUNCH + : AWS_CONNECT_STEPPER_ROW - currentStep; + return { steps: AWS_PROVIDER_WIZARD_STEPS, stepOffset }; + } + return { steps: PROVIDER_WIZARD_STEPS, stepOffset: 0 }; +} + +interface CredentialsRetryStepInput { + mode: ProviderWizardMode; + providerType: ProviderType | null; + isDirectCredentialsEntry?: boolean; +} + +/** Where "Back" from the connection test lands: AWS re-enters its one-step form. */ +export function getCredentialsRetryStep({ + mode, + providerType, + isDirectCredentialsEntry = false, +}: CredentialsRetryStepInput): ProviderWizardStep { + if ( + mode === PROVIDER_WIZARD_MODE.ADD && + providerType === "aws" && + !isDirectCredentialsEntry + ) { + return PROVIDER_WIZARD_STEP.CONNECT; + } + return PROVIDER_WIZARD_STEP.CREDENTIALS; +} + +interface LaunchBackStepInput { + providerType: ProviderType | null; + isDirectCredentialsEntry?: boolean; +} + +/** Where "Back" from the launch step lands: AWS returns to its one-step form. */ +export function getLaunchBackStep({ + providerType, + isDirectCredentialsEntry = false, +}: LaunchBackStepInput): ProviderWizardStep { + if (providerType === "aws" && !isDirectCredentialsEntry) { + return PROVIDER_WIZARD_STEP.CONNECT; + } + return PROVIDER_WIZARD_STEP.TEST; +} export function getOrganizationsStepperOffset( currentStep: OrgWizardStep, diff --git a/ui/components/providers/wizard/steps/aws/aws-connect-step.test.tsx b/ui/components/providers/wizard/steps/aws/aws-connect-step.test.tsx new file mode 100644 index 0000000000..73cf42fea8 --- /dev/null +++ b/ui/components/providers/wizard/steps/aws/aws-connect-step.test.tsx @@ -0,0 +1,757 @@ +import { act, render, screen, waitFor } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { useState } from "react"; +import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; + +import { + PROVIDER_FUNNEL_EVENT, + type ProviderFunnelDetail, +} from "@/lib/provider-funnel/provider-funnel-events"; +import { useProviderWizardStore } from "@/store/provider-wizard/store"; +import { + CONNECTION_CHECK_STATUS, + type ConnectionCheckStatus, +} from "@/types/providers"; + +import { AwsConnectStep } from "./aws-connect-step"; +import type { AwsConnectUiState } from "./types"; + +const { + addProvider, + addCredentialsProvider, + updateProvider, + updateCredentialsProvider, + testProviderConnection, + openCloudUpgradeMock, +} = vi.hoisted(() => ({ + addProvider: vi.fn(), + addCredentialsProvider: vi.fn(), + updateProvider: vi.fn(), + updateCredentialsProvider: vi.fn(), + testProviderConnection: vi.fn(), + openCloudUpgradeMock: vi.fn(), +})); + +vi.mock("next-auth/react", () => ({ + useSession: () => ({ + data: { tenantId: "tenant-abc" }, + status: "authenticated", + }), +})); +vi.mock("@/actions/providers/providers", () => ({ + addProvider, + addCredentialsProvider, + updateProvider, + updateCredentialsProvider, +})); +// The real module reaches next-auth through lib/helper -> auth.config. +vi.mock("@/lib/provider-helpers", () => ({ testProviderConnection })); +vi.mock("@/store", () => ({ + useCloudUpgradeStore: ( + selector: (state: { + openCloudUpgrade: typeof openCloudUpgradeMock; + }) => unknown, + ) => selector({ openCloudUpgrade: openCloudUpgradeMock }), +})); + +const FORM_ID = "aws-connect-test-form"; +const ROLE_ARN = "arn:aws:iam::123456789012:role/ProwlerScan"; + +// Stands in for the wizard footer: the step only publishes its UI state. +function Harness({ + onConnected, + onSelectOrganizations, +}: { + onConnected: () => void; + onSelectOrganizations: () => void; +}) { + const [uiState, setUiState] = useState(null); + return ( + <> + + + + ); +} + +function renderStep() { + const onConnected = vi.fn(); + const onSelectOrganizations = vi.fn(); + const { unmount } = render( + , + ); + return { + onConnected, + onSelectOrganizations, + unmount, + user: userEvent.setup(), + }; +} + +const connectButton = () => + screen.getByRole("button", { name: "Connect account" }); + +describe("AwsConnectStep", () => { + const funnelSignals: ProviderFunnelDetail[] = []; + const recordFunnelSignal: EventListener = (event) => { + funnelSignals.push((event as CustomEvent).detail); + }; + + beforeEach(() => { + funnelSignals.length = 0; + window.addEventListener(PROVIDER_FUNNEL_EVENT, recordFunnelSignal); + vi.clearAllMocks(); + sessionStorage.clear(); + useProviderWizardStore.getState().reset(); + addProvider.mockResolvedValue({ data: { id: "provider-1" } }); + addCredentialsProvider.mockResolvedValue({ data: { id: "secret-1" } }); + updateProvider.mockResolvedValue({ data: { id: "provider-1" } }); + updateCredentialsProvider.mockResolvedValue({ data: { id: "secret-1" } }); + testProviderConnection.mockResolvedValue({ + status: CONNECTION_CHECK_STATUS.SUCCESS, + error: null, + }); + }); + + afterEach(() => { + window.removeEventListener(PROVIDER_FUNNEL_EVENT, recordFunnelSignal); + vi.unstubAllEnvs(); + }); + + describe("in Prowler Cloud", () => { + beforeEach(() => { + vi.stubEnv("UI_CLOUD_ENABLED", "true"); + }); + + it("creates the role from the shared stack and connects with just its ARN", async () => { + // Given + const { onConnected, user } = renderStep(); + + // Then: the button opens the shared template with the External ID filled in; + // the AccountId parameter defaults to Prowler Cloud's account there. + const quickCreate = screen.getByRole("link", { + name: /Create the IAM role in AWS/i, + }); + expect(quickCreate).toHaveAttribute( + "href", + expect.stringContaining("prowler-scan-role.yml"), + ); + expect(quickCreate).toHaveAttribute( + "href", + expect.stringContaining("param_ExternalId=tenant-abc"), + ); + expect(connectButton()).toBeDisabled(); + + // When + await user.type( + screen.getByRole("textbox", { name: /Role ARN/ }), + ROLE_ARN, + ); + + // Then + expect( + await screen.findByText(/Account 123456789012 will be added/), + ).toBeVisible(); + await waitFor(() => expect(connectButton()).toBeEnabled()); + + // When + await user.click(connectButton()); + + // Then + await waitFor(() => expect(onConnected).toHaveBeenCalledOnce()); + const secret = Object.fromEntries( + (addCredentialsProvider.mock.calls[0][0] as FormData).entries(), + ); + expect(secret).toMatchObject({ + providerId: "provider-1", + role_arn: ROLE_ARN, + external_id: "tenant-abc", + credentials_type: "aws-sdk-default", + }); + expect(funnelSignals).toContainEqual({ + step: "account_submitted", + providerType: "aws", + via: "role", + outcome: "success", + }); + }); + + it("shows an account the API already knows on the ARN field and stays on the step", async () => { + // Given + addProvider.mockResolvedValueOnce({ + errors: [ + { + detail: "Provider with this uid already exists.", + source: { pointer: "/data/attributes/uid" }, + }, + ], + }); + const { onConnected, user } = renderStep(); + await user.type( + screen.getByRole("textbox", { name: /Role ARN/ }), + ROLE_ARN, + ); + await waitFor(() => expect(connectButton()).toBeEnabled()); + + // When + await user.click(connectButton()); + + // Then + expect( + await screen.findByText("Provider with this uid already exists."), + ).toBeVisible(); + expect(onConnected).not.toHaveBeenCalled(); + expect(funnelSignals).toContainEqual({ + step: "account_submitted", + providerType: "aws", + via: "role", + outcome: "error", + }); + }); + + it("hands the whole-organization choice to the organizations flow", async () => { + // Given + const { onSelectOrganizations, user } = renderStep(); + + // When + await user.click( + screen.getByRole("tab", { name: /Full AWS Organization/ }), + ); + + // Then + expect(onSelectOrganizations).toHaveBeenCalledOnce(); + }); + + it("connects with access keys and the typed account id", async () => { + // Given + const { onConnected, user } = renderStep(); + + // When + await user.click( + screen.getByRole("radio", { name: /Static access keys/ }), + ); + await user.type( + screen.getByRole("textbox", { name: /Account ID/ }), + "210987654321", + ); + await user.type( + screen.getByPlaceholderText("Enter the AWS Access Key ID"), + "AKIAEXAMPLE", + ); + await user.type( + screen.getByPlaceholderText("Enter the AWS Secret Access Key"), + "secret-value", + ); + await waitFor(() => expect(connectButton()).toBeEnabled()); + await user.click(connectButton()); + + // Then + await waitFor(() => expect(onConnected).toHaveBeenCalledOnce()); + const provider = Object.fromEntries( + (addProvider.mock.calls[0][0] as FormData).entries(), + ); + expect(provider).toEqual({ + providerType: "aws", + providerUid: "210987654321", + }); + }); + }); + + describe("with access keys, when the API refuses the account", () => { + beforeEach(() => { + vi.stubEnv("UI_CLOUD_ENABLED", "true"); + }); + + it("shows the refusal on the Account ID field and stays on the step", async () => { + // Given + addProvider.mockResolvedValueOnce({ + errors: [ + { + detail: "Provider with this uid already exists.", + source: { pointer: "/data/attributes/uid" }, + }, + ], + }); + const { onConnected, user } = renderStep(); + await user.click( + screen.getByRole("radio", { name: /Static access keys/ }), + ); + await user.type( + screen.getByRole("textbox", { name: /Account ID/ }), + "210987654321", + ); + await user.type( + screen.getByPlaceholderText("Enter the AWS Access Key ID"), + "AKIAEXAMPLE", + ); + await user.type( + screen.getByPlaceholderText("Enter the AWS Secret Access Key"), + "secret-value", + ); + await waitFor(() => expect(connectButton()).toBeEnabled()); + + // When + await user.click(connectButton()); + + // Then + expect( + await screen.findByText("Provider with this uid already exists."), + ).toBeVisible(); + // The field wrapper carries the invalid state for the Account ID input. + expect( + screen + .getByRole("textbox", { name: /Account ID/ }) + .closest("[aria-invalid]"), + ).toHaveAttribute("aria-invalid", "true"); + expect(onConnected).not.toHaveBeenCalled(); + }); + }); + + describe("when the connection is tested", () => { + beforeEach(() => { + vi.stubEnv("UI_CLOUD_ENABLED", "true"); + }); + + const submitRole = async () => { + const step = renderStep(); + await step.user.type( + screen.getByRole("textbox", { name: /Role ARN/ }), + ROLE_ARN, + ); + await waitFor(() => expect(connectButton()).toBeEnabled()); + await step.user.click(connectButton()); + return step; + }; + + const submitKeys = async () => { + const step = renderStep(); + await step.user.click( + screen.getByRole("radio", { name: /Static access keys/ }), + ); + await step.user.type( + screen.getByRole("textbox", { name: /Account ID/ }), + "210987654321", + ); + await step.user.type( + screen.getByPlaceholderText("Enter the AWS Access Key ID"), + "AKIAEXAMPLE", + ); + await step.user.type( + screen.getByPlaceholderText("Enter the AWS Secret Access Key"), + "secret-value", + ); + await waitFor(() => expect(connectButton()).toBeEnabled()); + await step.user.click(connectButton()); + return step; + }; + + it("reports the test in progress and blocks the action while it runs", async () => { + // Given: a test that has not answered yet. + let settle!: (result: { + status: ConnectionCheckStatus; + error: string | null; + }) => void; + testProviderConnection.mockImplementation( + () => + new Promise((resolve) => { + settle = resolve; + }), + ); + + // When + const { onConnected } = await submitRole(); + + // Then + expect(await screen.findByRole("status")).toHaveTextContent( + /testing the connection/i, + ); + expect( + screen.getByRole("button", { name: "Testing connection..." }), + ).toBeDisabled(); + expect(onConnected).not.toHaveBeenCalled(); + + // When / Then + await act(async () => + settle({ status: CONNECTION_CHECK_STATUS.SUCCESS, error: null }), + ); + await waitFor(() => expect(onConnected).toHaveBeenCalledOnce()); + }); + + it("ignores a result that lands after the step was closed", async () => { + // Given: the wizard is closed (or switched to organizations) mid-test. + let settle!: (result: { + status: ConnectionCheckStatus; + error: string | null; + }) => void; + testProviderConnection.mockImplementation( + () => + new Promise((resolve) => { + settle = resolve; + }), + ); + const { onConnected, unmount } = await submitRole(); + await screen.findByRole("status"); + + // When + unmount(); + await act(async () => + settle({ status: CONNECTION_CHECK_STATUS.SUCCESS, error: null }), + ); + + // Then: a reset wizard must not be pushed to the launch step. + expect(onConnected).not.toHaveBeenCalled(); + }); + + it("tests the account that was connected with static keys too", async () => { + // When + const { onConnected } = await submitKeys(); + + // Then + await waitFor(() => expect(onConnected).toHaveBeenCalledOnce()); + expect(testProviderConnection).toHaveBeenCalledWith("provider-1"); + }); + + it("stays on the keys form when the connection is refused", async () => { + // Given + testProviderConnection.mockResolvedValue({ + status: CONNECTION_CHECK_STATUS.FAILED, + error: "The access keys were rejected.", + }); + + // When + const { onConnected } = await submitKeys(); + + // Then + expect(await screen.findByRole("alert")).toHaveTextContent( + "The access keys were rejected.", + ); + expect(onConnected).not.toHaveBeenCalled(); + expect(screen.getByRole("textbox", { name: /Account ID/ })).toBeVisible(); + }); + + it("tests the registered account before leaving the step", async () => { + // When + const { onConnected } = await submitRole(); + + // Then + await waitFor(() => expect(onConnected).toHaveBeenCalledOnce()); + expect(testProviderConnection).toHaveBeenCalledWith("provider-1"); + }); + + it("stays on the form and offers a retry when the connection is refused", async () => { + // Given + testProviderConnection.mockResolvedValue({ + status: CONNECTION_CHECK_STATUS.FAILED, + error: "The role could not be assumed.", + }); + + // When + const { onConnected } = await submitRole(); + + // Then + expect(await screen.findByRole("alert")).toHaveTextContent( + "The role could not be assumed.", + ); + expect(onConnected).not.toHaveBeenCalled(); + expect( + screen.getByRole("button", { name: "Retry connection" }), + ).toBeEnabled(); + }); + + it("shows a neutral message, not a failure, when the check is still pending", async () => { + // Given: the wait was exhausted with no confirmed outcome (the backend + // check is genuinely still running past the wait). + testProviderConnection.mockResolvedValue({ + status: CONNECTION_CHECK_STATUS.PENDING, + error: + "The connection test is still running. Refresh in a moment to see the result.", + }); + + // When + const { onConnected } = await submitRole(); + + // Then: announced neutrally, not as an alert, and the account stays + // registered rather than reporting a failure the backend never gave. + expect(await screen.findByRole("status")).toHaveTextContent( + /still running/i, + ); + expect(screen.queryByRole("alert")).not.toBeInTheDocument(); + // Nor does it advance: the outcome is still unknown. + expect(onConnected).not.toHaveBeenCalled(); + expect(screen.getByRole("button", { name: "Check again" })).toBeEnabled(); + }); + + // The helper always supplies a reason today; this guards the alert against a + // future contract that does not. + it("falls back to a generic reason when the API gives none", async () => { + // Given + testProviderConnection.mockResolvedValue({ + status: CONNECTION_CHECK_STATUS.FAILED, + error: null, + }); + + // When + await submitRole(); + + // Then + expect(await screen.findByRole("alert")).toHaveTextContent( + /could not connect/i, + ); + }); + + it("recovers when the connection test itself fails", async () => { + // Given: task polling rejects on a 5xx instead of reporting a failure. + testProviderConnection.mockRejectedValue(new Error("Server error (500)")); + + // When + const { onConnected } = await submitRole(); + + // Then + expect(await screen.findByRole("alert")).toHaveTextContent( + /account is saved/i, + ); + expect(onConnected).not.toHaveBeenCalled(); + await waitFor(() => + expect( + screen.getByRole("button", { name: "Retry connection" }), + ).toBeEnabled(), + ); + }); + + it("drops the failure as soon as the form is edited again", async () => { + // Given + testProviderConnection.mockResolvedValue({ + status: CONNECTION_CHECK_STATUS.FAILED, + error: "The role could not be assumed.", + }); + const { user } = await submitRole(); + await screen.findByRole("alert"); + + // When + await user.type( + screen.getByRole("textbox", { name: /Role ARN/ }), + "-extra", + ); + + // Then + await waitFor(() => + expect(screen.queryByRole("alert")).not.toBeInTheDocument(), + ); + }); + + it("moves on once a retry connects", async () => { + // Given + testProviderConnection + .mockResolvedValueOnce({ + status: CONNECTION_CHECK_STATUS.FAILED, + error: "Denied.", + }) + .mockResolvedValueOnce({ + status: CONNECTION_CHECK_STATUS.SUCCESS, + error: null, + }); + const { onConnected, user } = await submitRole(); + + // When + await user.click( + await screen.findByRole("button", { name: "Retry connection" }), + ); + + // Then + await waitFor(() => expect(onConnected).toHaveBeenCalledOnce()); + // The account is registered once; the retry only rewrites its secret. + expect(addProvider).toHaveBeenCalledOnce(); + }); + }); + + describe("when the step is left and reopened within the same wizard", () => { + beforeEach(() => { + vi.stubEnv("UI_CLOUD_ENABLED", "true"); + }); + + it("keeps what was typed, including the chosen access method", async () => { + // Given + const onConnected = vi.fn(); + const onSelectOrganizations = vi.fn(); + const user = userEvent.setup(); + const { unmount } = render( + , + ); + await user.click( + screen.getByRole("radio", { name: /Static access keys/ }), + ); + await user.type( + screen.getByRole("textbox", { name: /Account ID/ }), + "210987654321", + ); + await user.type( + screen.getByRole("textbox", { name: /Provider alias/ }), + "Staging", + ); + + // When: the organizations tab or the connection test unmounts the step. + unmount(); + render( + , + ); + + // Then + expect( + screen.getByRole("radio", { name: /Static access keys/ }), + ).toHaveAttribute("aria-checked", "true"); + expect(screen.getByRole("textbox", { name: /Account ID/ })).toHaveValue( + "210987654321", + ); + expect( + screen.getByRole("textbox", { name: /Provider alias/ }), + ).toHaveValue("Staging"); + }); + + it("starts blank again once the wizard is reset", async () => { + // Given + const user = userEvent.setup(); + const { unmount } = render( + , + ); + await user.type( + screen.getByRole("textbox", { name: /Role ARN/ }), + ROLE_ARN, + ); + unmount(); + + // When + useProviderWizardStore.getState().reset(); + render(); + + // Then + expect(screen.getByRole("textbox", { name: /Role ARN/ })).toHaveValue(""); + }); + }); + + describe("in Prowler Cloud, role creation", () => { + beforeEach(() => { + vi.stubEnv("UI_CLOUD_ENABLED", "true"); + }); + + it("leads with the one-click stack and keeps the other templates behind a toggle", async () => { + // Given + const { user } = renderStep(); + + // Then + expect( + screen.queryByRole("link", { name: /CloudFormation Template/i }), + ).not.toBeInTheDocument(); + expect( + screen.queryByRole("link", { name: /Terraform Code/i }), + ).not.toBeInTheDocument(); + + // When + await user.click( + screen.getByRole("button", { name: /Other ways to create the role/i }), + ); + + // Then + expect( + screen.getByRole("link", { name: /CloudFormation Template/i }), + ).toHaveAttribute("href", expect.stringContaining("prowler-scan-role")); + expect( + screen.getByRole("link", { name: /Terraform Code/i }), + ).toBeVisible(); + }); + + it("never asks which credentials assume the role: Prowler Cloud does", async () => { + // Given + const { user } = renderStep(); + await user.click( + screen.getByRole("button", { name: /Advanced options/i }), + ); + + // Then + expect(screen.queryByRole("combobox")).not.toBeInTheDocument(); + expect( + screen.queryByPlaceholderText("Enter the AWS Access Key ID"), + ).not.toBeInTheDocument(); + expect( + screen.getByPlaceholderText("Enter the role session name"), + ).toBeVisible(); + }); + }); + + describe("in a self-hosted deployment", () => { + beforeEach(() => { + vi.stubEnv("UI_CLOUD_ENABLED", "false"); + }); + + it("offers the same one-click role setup, on the shared template", async () => { + // Given + const { user } = renderStep(); + + // Then: the template keeps the AccountId parameter self-hosted users must edit. + expect( + screen.getByRole("link", { name: /Create the IAM role in AWS/i }), + ).toHaveAttribute( + "href", + expect.stringContaining("prowler-scan-role.yml"), + ); + + // When + await user.click( + screen.getByRole("button", { name: /Advanced options/i }), + ); + + // Then: keys belong to the "Static access keys" method, never to the role one. + expect(screen.queryByRole("combobox")).not.toBeInTheDocument(); + expect( + screen.queryByPlaceholderText("Enter the AWS Access Key ID"), + ).not.toBeInTheDocument(); + expect( + screen.getByPlaceholderText("Enter the role session name"), + ).toBeVisible(); + }); + + it("assumes the role with the credentials of the host running Prowler", async () => { + // Given + const { onConnected, user } = renderStep(); + + // When + await user.type( + screen.getByRole("textbox", { name: /Role ARN/ }), + ROLE_ARN, + ); + await screen.findByText(/Account 123456789012 will be added/); + await waitFor(() => expect(connectButton()).toBeEnabled()); + await user.click(connectButton()); + + // Then + await waitFor(() => expect(onConnected).toHaveBeenCalledOnce()); + const secret = Object.fromEntries( + (addCredentialsProvider.mock.calls[0][0] as FormData).entries(), + ); + expect(secret).toMatchObject({ + role_arn: ROLE_ARN, + credentials_type: "aws-sdk-default", + }); + expect(secret).not.toHaveProperty("aws_access_key_id"); + }); + }); +}); diff --git a/ui/components/providers/wizard/steps/aws/aws-connect-step.tsx b/ui/components/providers/wizard/steps/aws/aws-connect-step.tsx new file mode 100644 index 0000000000..9cca93a5dd --- /dev/null +++ b/ui/components/providers/wizard/steps/aws/aws-connect-step.tsx @@ -0,0 +1,628 @@ +"use client"; + +import { zodResolver } from "@hookform/resolvers/zod"; +import { + ChevronDownIcon, + CircleAlert, + KeyRound, + Loader2, + ShieldCheck, +} from "lucide-react"; +import { useSession } from "next-auth/react"; +import { useEffect, useRef, useState } from "react"; +import { + Control, + FieldValues, + Resolver, + UseFormReturn, + useForm, + useFormState, + useWatch, +} from "react-hook-form"; + +import { ConnectionPending } from "@/components/icons"; +import { RadioCard } from "@/components/providers/radio-card"; +import { CredentialsRoleHelper } from "@/components/providers/workflow/credentials-role-helper"; +import { WizardInputField } from "@/components/providers/workflow/forms/fields"; +import { AwsRoleOptionalFields } from "@/components/providers/workflow/forms/select-credentials-type/aws/credentials-type/aws-role-optional-fields"; +import { AWSStaticCredentialsForm } from "@/components/providers/workflow/forms/select-credentials-type/aws/credentials-type/aws-static-credentials-form"; +import { ProviderTitleDocs } from "@/components/providers/workflow/provider-title-docs"; +import { Badge } from "@/components/shadcn/badge/badge"; +import { Button } from "@/components/shadcn/button/button"; +import { + Collapsible, + CollapsibleContent, + CollapsibleTrigger, +} from "@/components/shadcn/collapsible"; +import { Form } from "@/components/shadcn/form"; +import { useFormServerErrors } from "@/hooks/use-form-server-errors"; +import { useMountEffect } from "@/hooks/use-mount-effect"; +import { PROVIDER_CREDENTIALS_ERROR_MAPPING } from "@/lib/error-mappings"; +import { getAWSCredentialsTemplateLinks } from "@/lib/external-urls"; +import { ProviderCredentialFields } from "@/lib/provider-credentials/provider-credential-fields"; +import { + ACCOUNT_SUBMIT_OUTCOME, + dispatchProviderFunnel, + PROVIDER_FUNNEL_STEP, +} from "@/lib/provider-funnel/provider-funnel-events"; +import { testProviderConnection } from "@/lib/provider-helpers"; +import { useProviderWizardStore } from "@/store/provider-wizard/store"; +import type { AWSCredentials, AWSCredentialsRole } from "@/types"; +import type { AwsConnectDraft } from "@/types/provider-wizard"; +import { CONNECTION_CHECK_STATUS } from "@/types/providers"; + +import { + awsKeysConnectSchema, + type AwsKeysConnectValues, + awsRoleConnectSchema, + type AwsRoleConnectValues, +} from "./aws-connect.schema"; +import { + AWS_ONBOARDING_METHOD, + AwsOnboardingMethodTabs, +} from "./aws-onboarding-method-tabs"; +import { parseAwsAccountIdFromRoleArn } from "./aws-role-arn"; +import { + AWS_UID_ERROR_POINTER, + connectAwsAccount, +} from "./connect-aws-account"; +import { + AWS_ACCESS_METHOD, + type AwsAccessMethod, + type AwsConnectUiState, +} from "./types"; + +const ALIAS_ERROR_POINTER = "/data/attributes/alias"; +const UNIQUE_TOGETHER_ERROR_POINTER = "/data/attributes/__all__"; + +// What the user typed survives the step unmounting (organizations tab, a step +// back from the launch step) until the wizard closes. +const readDraft = () => useProviderWizardStore.getState().awsConnectDraft; + +const initialMethod = (): AwsAccessMethod => + readDraft()?.method === AWS_ACCESS_METHOD.CREDENTIALS + ? AWS_ACCESS_METHOD.CREDENTIALS + : AWS_ACCESS_METHOD.ROLE; + +function useDraftValues( + form: UseFormReturn, + key: keyof Pick, +) { + const values = useWatch({ control: form.control }); + useEffect(() => { + useProviderWizardStore + .getState() + .setAwsConnectDraft({ [key]: values as AwsConnectDraft[typeof key] }); + }, [key, values]); +} + +interface AwsConnectStepProps { + formId: string; + onConnected: () => void; + onSelectOrganizations: () => void; + onUiStateChange: (state: AwsConnectUiState) => void; +} + +/** One form to register an AWS account, store its credentials and test the connection. */ +export function AwsConnectStep({ + formId, + onConnected, + onSelectOrganizations, + onUiStateChange, +}: AwsConnectStepProps) { + // Local state needed: the access method only matters until the account is connected. + const [method, setMethod] = useState(initialMethod); + // Local state needed: the active form reports it so the method cannot change mid-submit. + const [isBusy, setIsBusy] = useState(false); + + const isRole = method === AWS_ACCESS_METHOD.ROLE; + + const chooseMethod = (next: AwsAccessMethod) => { + setMethod(next); + useProviderWizardStore.getState().setAwsConnectDraft({ method: next }); + }; + + return ( +
+ + + + +
+

+ Choose how Prowler should access your account. +

+ chooseMethod(AWS_ACCESS_METHOD.ROLE)} + > + + Recommended + + + chooseMethod(AWS_ACCESS_METHOD.CREDENTIALS)} + /> +
+ + {isRole ? ( + + ) : ( + + )} +
+ ); +} + +interface ConnectFormProps + extends Pick< + AwsConnectStepProps, + "formId" | "onConnected" | "onUiStateChange" + > { + onBusyChange: (isBusy: boolean) => void; +} + +interface UseAwsConnectSubmitOptions { + form: UseFormReturn; + method: AwsAccessMethod; + // The field an account-level API error belongs to for this method. + accountField: string; + // Beyond form validity: the role form also needs an account read from the ARN. + accountResolved?: boolean; + extraValues?: Record; + onConnected: () => void; + onBusyChange: (isBusy: boolean) => void; + onUiStateChange: (state: AwsConnectUiState) => void; +} + +const CONNECTION_FAILED_MESSAGE = + "Prowler could not connect with these credentials. Review them and try again."; + +const CONNECTION_UNREACHABLE_MESSAGE = + "The connection test could not be completed. The account is saved, so you can try again."; + +// Fallback only: `testProviderConnection` already carries this same message on +// `error` for a pending result (see `resolveProviderConnectionState`). +const CONNECTION_PENDING_MESSAGE = + "The connection test is still running. Refresh in a moment to see the result."; + +/** Footer label for the one-step form: the test and the retry share the submit. */ +const resolveActionLabel = ({ + isTesting, + isSubmitting, + hasFailed, + hasPending, +}: { + isTesting: boolean; + isSubmitting: boolean; + hasFailed: boolean; + hasPending: boolean; +}) => { + if (isTesting) return "Testing connection..."; + if (isSubmitting) return "Connecting account..."; + if (hasFailed) return "Retry connection"; + return hasPending ? "Check again" : "Connect account"; +}; + +/** Registers the account, stores its credentials and tests the connection in one submit. */ +function useAwsConnectSubmit({ + form, + method, + accountField, + accountResolved = true, + extraValues, + onConnected, + onBusyChange, + onUiStateChange, +}: UseAwsConnectSubmitOptions) { + const { handleServerResponse } = useFormServerErrors(form, { + ...PROVIDER_CREDENTIALS_ERROR_MAPPING, + [AWS_UID_ERROR_POINTER]: accountField, + [UNIQUE_TOGETHER_ERROR_POINTER]: accountField, + [ALIAS_ERROR_POINTER]: ProviderCredentialFields.PROVIDER_ALIAS, + }); + // Local state needed: the connection test runs inside the submit, and its + // outcome belongs to this step rather than to any form field. + const [isTesting, setIsTesting] = useState(false); + const [connectionError, setConnectionError] = useState(null); + // Still running past the wait -- neither a pass nor a fail. Kept separate + // from `connectionError` so it never renders with the destructive styling a + // confirmed failure gets, and never counts as one. + const [connectionPending, setConnectionPending] = useState( + null, + ); + // A hook, not `form.formState.isValid` read inline: the React Compiler keys + // its memo on the stable `form` object and would freeze a proxy read at false. + const { isSubmitting, isValid } = useFormState({ control: form.control }); + const canSubmit = isValid && accountResolved; + const isBusy = isSubmitting || isTesting; + // Closing the wizard (or switching to organizations) unmounts the step while a + // test may still be running; its result must not advance a wizard already reset. + const isActiveRef = useRef(true); + useMountEffect(() => { + isActiveRef.current = true; + return () => { + isActiveRef.current = false; + }; + }); + + // Same contract ConnectAccountForm uses: the wizard footer lives outside the step. + // Both callbacks must be stable setters, or this effect would loop. + useEffect(() => { + onBusyChange(isBusy); + onUiStateChange({ + showBack: true, + showAction: true, + actionLabel: resolveActionLabel({ + isTesting, + isSubmitting, + hasFailed: connectionError !== null, + hasPending: connectionPending !== null, + }), + actionDisabled: !canSubmit || isBusy, + isLoading: isBusy, + }); + }, [ + canSubmit, + connectionError, + connectionPending, + isBusy, + isSubmitting, + isTesting, + onBusyChange, + onUiStateChange, + ]); + + // A past failure or a still-pending result must not sit above the field the + // user is already correcting. + useEffect(() => { + if (connectionError === null && connectionPending === null) return; + const subscription = form.watch(() => { + setConnectionError(null); + setConnectionPending(null); + }); + return () => subscription.unsubscribe(); + }, [connectionError, connectionPending, form]); + + const onSubmit = form.handleSubmit(async (values) => { + setConnectionError(null); + setConnectionPending(null); + const result = await connectAwsAccount({ + method, + values: { ...values, ...extraValues }, + }); + dispatchProviderFunnel({ + step: PROVIDER_FUNNEL_STEP.ACCOUNT_SUBMITTED, + providerType: "aws", + via: method, + outcome: result.ok + ? ACCOUNT_SUBMIT_OUTCOME.SUCCESS + : ACCOUNT_SUBMIT_OUTCOME.ERROR, + }); + if (!result.ok) { + // Maps API pointers onto the form's fields; anything unmapped becomes a toast. + handleServerResponse({ errors: result.errors }); + return; + } + + // The account stays registered whatever the test says; resubmitting edits it + // in place. Task polling rejects on a 5xx, so the flag has to be cleared in a + // finally or the step would stay stuck on "Testing connection...". + let connected = false; + setIsTesting(true); + try { + const connection = await testProviderConnection(result.providerId); + connected = connection.status === CONNECTION_CHECK_STATUS.SUCCESS; + if (connection.status === CONNECTION_CHECK_STATUS.FAILED) { + setConnectionError(connection.error || CONNECTION_FAILED_MESSAGE); + } else if (connection.status === CONNECTION_CHECK_STATUS.PENDING) { + setConnectionPending(connection.error || CONNECTION_PENDING_MESSAGE); + } + } catch { + setConnectionError(CONNECTION_UNREACHABLE_MESSAGE); + } finally { + setIsTesting(false); + } + + if (connected && isActiveRef.current) onConnected(); + }); + + return { onSubmit, isTesting, connectionError, connectionPending }; +} + +/** + * Progress line while the test runs, the neutral message once the wait is + * exhausted with no verdict, or the API's reason once it is refused. + */ +function ConnectionFeedback({ + isTesting, + error, + pending, +}: { + isTesting: boolean; + error: string | null; + pending: string | null; +}) { + const alertRef = useRef(null); + + // The form scrolls inside the modal and the action button sits outside it, so + // an error or a pending result raised from the footer can land above the fold. + useEffect(() => { + if (!error && !pending) return; + // Guarded: jsdom has no scrollIntoView, and a throw here would unmount the step. + alertRef.current?.scrollIntoView?.({ block: "start", behavior: "smooth" }); + }, [error, pending]); + + if (isTesting) { + return ( +

+ + Testing the connection. This usually takes a few seconds. +

+ ); + } + + if (pending) { + return ( +
+
+ ); + } + + if (!error) return null; + + return ( +
+ +

+ {error} +

+
+ ); +} + +function AwsRoleConnectForm({ + formId, + onConnected, + onBusyChange, + onUiStateChange, +}: ConnectFormProps) { + const { data: session } = useSession(); + const externalId = session?.tenantId ?? ""; + + const form = useForm({ + resolver: zodResolver( + awsRoleConnectSchema, + ) as unknown as Resolver, + mode: "onChange", + defaultValues: { + [ProviderCredentialFields.PROVIDER_ID]: "", + [ProviderCredentialFields.PROVIDER_TYPE]: "aws", + [ProviderCredentialFields.PROVIDER_ALIAS]: "", + // The role is assumed with Prowler's own credentials (Cloud's identity or the + // host's AWS SDK chain); static keys are a method of their own, never mixed in. + [ProviderCredentialFields.CREDENTIALS_TYPE]: + ProviderCredentialFields.CREDENTIALS_TYPE_AWS, + [ProviderCredentialFields.ROLE_ARN]: "", + [ProviderCredentialFields.AWS_ACCESS_KEY_ID]: "", + [ProviderCredentialFields.AWS_SECRET_ACCESS_KEY]: "", + [ProviderCredentialFields.AWS_SESSION_TOKEN]: "", + [ProviderCredentialFields.ROLE_SESSION_NAME]: "", + [ProviderCredentialFields.SESSION_DURATION]: "3600", + ...readDraft()?.roleValues, + }, + }); + useDraftValues(form, "roleValues"); + + const roleArn = useWatch({ + control: form.control, + name: ProviderCredentialFields.ROLE_ARN, + }); + const detectedAccountId = parseAwsAccountIdFromRoleArn(roleArn ?? ""); + + const { onSubmit, isTesting, connectionError, connectionPending } = + useAwsConnectSubmit({ + form, + method: AWS_ACCESS_METHOD.ROLE, + accountField: ProviderCredentialFields.ROLE_ARN, + accountResolved: detectedAccountId !== null, + // The external id is the tenant's, never user input, so it joins at submit time. + extraValues: { [ProviderCredentialFields.EXTERNAL_ID]: externalId }, + onConnected, + onBusyChange, + onUiStateChange, + }); + + // One template for every build: self-hosted users set the account that assumes + // the role, so the AccountId parameter must stay editable in the console. + const templateLinks = getAWSCredentialsTemplateLinks(externalId); + const roleControl = form.control as unknown as Control; + + return ( +
+ + + +
+

1. Create the IAM role

+ +
+ +
+

2. Paste the role ARN

+
+ + {detectedAccountId && ( +

+ Account {detectedAccountId} will be added to Prowler. +

+ )} +
+ } + /> +
+ + + + + + + + + + + + ); +} + +function AwsKeysConnectForm({ + formId, + onConnected, + onBusyChange, + onUiStateChange, +}: ConnectFormProps) { + const form = useForm({ + resolver: zodResolver( + awsKeysConnectSchema, + ) as unknown as Resolver, + mode: "onChange", + defaultValues: { + [ProviderCredentialFields.PROVIDER_ID]: "", + [ProviderCredentialFields.PROVIDER_TYPE]: "aws", + [ProviderCredentialFields.PROVIDER_UID]: "", + [ProviderCredentialFields.PROVIDER_ALIAS]: "", + [ProviderCredentialFields.AWS_ACCESS_KEY_ID]: "", + [ProviderCredentialFields.AWS_SECRET_ACCESS_KEY]: "", + [ProviderCredentialFields.AWS_SESSION_TOKEN]: "", + ...readDraft()?.keysValues, + }, + }); + useDraftValues(form, "keysValues"); + + const { onSubmit, isTesting, connectionError, connectionPending } = + useAwsConnectSubmit({ + form, + method: AWS_ACCESS_METHOD.CREDENTIALS, + accountField: ProviderCredentialFields.PROVIDER_UID, + onConnected, + onBusyChange, + onUiStateChange, + }); + + return ( +
+ + + + value.replace(/\D/g, "").slice(0, 12)} + /> + } + /> + } /> + + + ); +} + +function AliasField({ control }: { control: Control }) { + return ( + + ); +} diff --git a/ui/components/providers/wizard/steps/aws/aws-connect.schema.ts b/ui/components/providers/wizard/steps/aws/aws-connect.schema.ts new file mode 100644 index 0000000000..2d17d23858 --- /dev/null +++ b/ui/components/providers/wizard/steps/aws/aws-connect.schema.ts @@ -0,0 +1,52 @@ +import { z } from "zod"; + +import { ProviderCredentialFields } from "@/lib/provider-credentials/provider-credential-fields"; +import { + addCredentialsFormSchema, + addCredentialsRoleFormSchema, +} from "@/types/formSchemas"; + +const aliasField = { + [ProviderCredentialFields.PROVIDER_ALIAS]: z.string().trim().optional(), +}; + +// The shared credential schemas stay the source of truth; the step only adds the +// account fields it collects in the same form. +export const awsRoleConnectSchema = addCredentialsRoleFormSchema("aws").and( + z.object(aliasField), +); + +export const awsKeysConnectSchema = addCredentialsFormSchema("aws").and( + z.object({ + ...aliasField, + [ProviderCredentialFields.PROVIDER_UID]: z + .string() + .trim() + .regex(/^\d{12}$/, "AWS Account ID must be exactly 12 digits"), + }), +); + +// The shared schema factories take a plain string, so their inferred type is the +// union of every provider; the step's forms declare the AWS shape explicitly. +interface AwsConnectAccountValues { + [ProviderCredentialFields.PROVIDER_ID]: string; + [ProviderCredentialFields.PROVIDER_TYPE]: string; + [ProviderCredentialFields.PROVIDER_ALIAS]?: string; +} + +export interface AwsRoleConnectValues extends AwsConnectAccountValues { + [ProviderCredentialFields.ROLE_ARN]: string; + [ProviderCredentialFields.CREDENTIALS_TYPE]?: string; + [ProviderCredentialFields.AWS_ACCESS_KEY_ID]?: string; + [ProviderCredentialFields.AWS_SECRET_ACCESS_KEY]?: string; + [ProviderCredentialFields.AWS_SESSION_TOKEN]?: string; + [ProviderCredentialFields.ROLE_SESSION_NAME]?: string; + [ProviderCredentialFields.SESSION_DURATION]?: string; +} + +export interface AwsKeysConnectValues extends AwsConnectAccountValues { + [ProviderCredentialFields.PROVIDER_UID]: string; + [ProviderCredentialFields.AWS_ACCESS_KEY_ID]: string; + [ProviderCredentialFields.AWS_SECRET_ACCESS_KEY]: string; + [ProviderCredentialFields.AWS_SESSION_TOKEN]?: string; +} diff --git a/ui/components/providers/wizard/steps/aws/aws-onboarding-method-tabs.test.tsx b/ui/components/providers/wizard/steps/aws/aws-onboarding-method-tabs.test.tsx new file mode 100644 index 0000000000..6a4b7a279f --- /dev/null +++ b/ui/components/providers/wizard/steps/aws/aws-onboarding-method-tabs.test.tsx @@ -0,0 +1,77 @@ +import { render, screen } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { afterEach, describe, expect, it, vi } from "vitest"; + +import { useCloudUpgradeStore } from "@/store"; +import { CLOUD_UPGRADE_FEATURE } from "@/types/cloud-upgrade"; + +import { + AWS_ONBOARDING_METHOD, + AwsOnboardingMethodTabs, +} from "./aws-onboarding-method-tabs"; + +describe("AwsOnboardingMethodTabs", () => { + afterEach(() => { + vi.unstubAllEnvs(); + useCloudUpgradeStore.getState().closeCloudUpgrade(); + }); + + it("switches to the organization flow in Cloud", async () => { + vi.stubEnv("UI_CLOUD_ENABLED", "true"); + const user = userEvent.setup(); + const onSelectOrganizations = vi.fn(); + render( + , + ); + + await user.click( + screen.getByRole("tab", { name: /Full AWS Organization/ }), + ); + + expect(onSelectOrganizations).toHaveBeenCalledOnce(); + expect(screen.queryByText("Cloud")).not.toBeInTheDocument(); + }); + + it("opens the AWS Organizations upgrade in Local Server", async () => { + vi.stubEnv("UI_CLOUD_ENABLED", "false"); + const user = userEvent.setup(); + const onSelectOrganizations = vi.fn(); + render( + , + ); + + await user.click( + screen.getByRole("tab", { name: /Full AWS Organization/ }), + ); + + expect(onSelectOrganizations).not.toHaveBeenCalled(); + expect(screen.getByText("Cloud")).toBeVisible(); + expect(useCloudUpgradeStore.getState().activeFeature).toBe( + CLOUD_UPGRADE_FEATURE.AWS_ORGANIZATIONS, + ); + }); + + it("returns to the single account flow from the organization tab", async () => { + const user = userEvent.setup(); + const onSelectSingle = vi.fn(); + render( + , + ); + + await user.click(screen.getByRole("tab", { name: "Single AWS Account" })); + + expect(onSelectSingle).toHaveBeenCalledOnce(); + }); +}); diff --git a/ui/components/providers/wizard/steps/aws/aws-onboarding-method-tabs.tsx b/ui/components/providers/wizard/steps/aws/aws-onboarding-method-tabs.tsx new file mode 100644 index 0000000000..2f24f177d4 --- /dev/null +++ b/ui/components/providers/wizard/steps/aws/aws-onboarding-method-tabs.tsx @@ -0,0 +1,72 @@ +"use client"; + +import { Badge } from "@/components/shadcn/badge/badge"; +import { Tabs, TabsList, TabsTrigger } from "@/components/shadcn/tabs/tabs"; +import { isCloud } from "@/lib/shared/env"; +import { useCloudUpgradeStore } from "@/store"; +import { CLOUD_UPGRADE_FEATURE } from "@/types/cloud-upgrade"; + +export const AWS_ONBOARDING_METHOD = { + SINGLE: "single", + ORGANIZATION: "organization", +} as const; + +export type AwsOnboardingMethod = + (typeof AWS_ONBOARDING_METHOD)[keyof typeof AWS_ONBOARDING_METHOD]; + +interface AwsOnboardingMethodTabsProps { + value: AwsOnboardingMethod; + onSelectSingle?: () => void; + onSelectOrganizations?: () => void; +} + +/** Single account vs. whole organization switch at the top of the AWS connect step. */ +export function AwsOnboardingMethodTabs({ + value, + onSelectSingle, + onSelectOrganizations, +}: AwsOnboardingMethodTabsProps) { + const isCloudEnv = isCloud(); + const openCloudUpgrade = useCloudUpgradeStore( + (state) => state.openCloudUpgrade, + ); + + const handleValueChange = (next: string) => { + if (next === value) return; + if (next === AWS_ONBOARDING_METHOD.SINGLE) { + onSelectSingle?.(); + return; + } + if (isCloudEnv) { + onSelectOrganizations?.(); + return; + } + openCloudUpgrade(CLOUD_UPGRADE_FEATURE.AWS_ORGANIZATIONS); + }; + + return ( + + + + Single AWS Account + + + Cloud + + ) + } + > + Full AWS Organization + + + + ); +} diff --git a/ui/components/providers/wizard/steps/aws/aws-role-arn.test.ts b/ui/components/providers/wizard/steps/aws/aws-role-arn.test.ts new file mode 100644 index 0000000000..11929a20b7 --- /dev/null +++ b/ui/components/providers/wizard/steps/aws/aws-role-arn.test.ts @@ -0,0 +1,25 @@ +import { describe, expect, it } from "vitest"; + +import { parseAwsAccountIdFromRoleArn } from "./aws-role-arn"; + +describe("parseAwsAccountIdFromRoleArn", () => { + it.each([ + ["arn:aws:iam::123456789012:role/ProwlerScan", "123456789012"], + [" arn:aws:iam::123456789012:role/path/ProwlerScan ", "123456789012"], + ["arn:aws-cn:iam::123456789012:role/ProwlerScan", "123456789012"], + ["arn:aws-us-gov:iam::123456789012:role/Prowler@Scan", "123456789012"], + ])("extracts the account id from %s", (arn, expected) => { + expect(parseAwsAccountIdFromRoleArn(arn)).toBe(expected); + }); + + it.each([ + "", + "123456789012", + "arn:aws:iam::12345678901:role/ProwlerScan", + "arn:aws:iam::123456789012:user/prowler", + "arn:aws:s3:::bucket", + "arn:aws:iam::123456789012:role/", + ])("returns null for %s", (arn) => { + expect(parseAwsAccountIdFromRoleArn(arn)).toBeNull(); + }); +}); diff --git a/ui/components/providers/wizard/steps/aws/aws-role-arn.ts b/ui/components/providers/wizard/steps/aws/aws-role-arn.ts new file mode 100644 index 0000000000..47b7b79153 --- /dev/null +++ b/ui/components/providers/wizard/steps/aws/aws-role-arn.ts @@ -0,0 +1,9 @@ +const AWS_ROLE_ARN_PATTERN = + /^arn:aws(?:-[a-z]+)*:iam::(\d{12}):role\/[\w+=,.@/-]+$/; + +export const AWS_ROLE_ARN_MESSAGE = + "Must be a valid IAM Role ARN (e.g. arn:aws:iam::123456789012:role/ProwlerScan)"; + +/** The 12-digit account id embedded in an IAM role ARN, or null when malformed. */ +export const parseAwsAccountIdFromRoleArn = (roleArn: string) => + AWS_ROLE_ARN_PATTERN.exec(roleArn.trim())?.[1] ?? null; diff --git a/ui/components/providers/wizard/steps/aws/connect-aws-account.test.ts b/ui/components/providers/wizard/steps/aws/connect-aws-account.test.ts new file mode 100644 index 0000000000..3bbc0f9f26 --- /dev/null +++ b/ui/components/providers/wizard/steps/aws/connect-aws-account.test.ts @@ -0,0 +1,340 @@ +import { beforeEach, describe, expect, it, vi } from "vitest"; + +import { useProviderWizardStore } from "@/store/provider-wizard/store"; +import { useUIStore } from "@/store/ui/store"; + +import { connectAwsAccount } from "./connect-aws-account"; +import { AWS_ACCESS_METHOD } from "./types"; + +const { + addProvider, + addCredentialsProvider, + updateProvider, + updateCredentialsProvider, +} = vi.hoisted(() => ({ + addProvider: vi.fn(), + addCredentialsProvider: vi.fn(), + updateProvider: vi.fn(), + updateCredentialsProvider: vi.fn(), +})); + +vi.mock("@/actions/providers/providers", () => ({ + addProvider, + addCredentialsProvider, + updateProvider, + updateCredentialsProvider, +})); + +const ROLE_ARN = "arn:aws:iam::123456789012:role/ProwlerScan"; + +const roleValues = { + providerId: "", + providerType: "aws", + providerAlias: "Production", + role_arn: ROLE_ARN, + external_id: "tenant-1", + credentials_type: "aws-sdk-default", + aws_access_key_id: "", + aws_secret_access_key: "", + aws_session_token: "", + role_session_name: "", + session_duration: "3600", +}; + +const formEntries = (call: number, mock: typeof addProvider) => + Object.fromEntries((mock.mock.calls[call][0] as FormData).entries()); + +describe("connectAwsAccount", () => { + beforeEach(() => { + vi.clearAllMocks(); + sessionStorage.clear(); + useProviderWizardStore.getState().reset(); + useUIStore.setState({ hasProviders: false, hasProvidersResolved: true }); + addProvider.mockResolvedValue({ data: { id: "provider-1" } }); + addCredentialsProvider.mockResolvedValue({ data: { id: "secret-1" } }); + updateProvider.mockResolvedValue({ data: { id: "provider-1" } }); + updateCredentialsProvider.mockResolvedValue({ data: { id: "secret-1" } }); + }); + + describe("when connecting through an IAM role", () => { + it("registers the account read from the ARN and stores its credentials in one go", async () => { + // When + const result = await connectAwsAccount({ + method: AWS_ACCESS_METHOD.ROLE, + values: roleValues, + }); + + // Then + expect(result).toEqual({ ok: true, providerId: "provider-1" }); + expect(formEntries(0, addProvider)).toEqual({ + providerType: "aws", + providerUid: "123456789012", + providerAlias: "Production", + }); + expect(formEntries(0, addCredentialsProvider)).toEqual({ + providerId: "provider-1", + providerType: "aws", + role_arn: ROLE_ARN, + external_id: "tenant-1", + credentials_type: "aws-sdk-default", + session_duration: "3600", + }); + }); + + it("leaves the wizard ready for the connection test", async () => { + // When + await connectAwsAccount({ + method: AWS_ACCESS_METHOD.ROLE, + values: roleValues, + }); + + // Then + expect(useProviderWizardStore.getState()).toMatchObject({ + providerId: "provider-1", + providerType: "aws", + providerUid: "123456789012", + providerAlias: "Production", + via: "role", + secretId: "secret-1", + mode: "add", + }); + expect(useUIStore.getState().hasProviders).toBe(true); + }); + + it("rejects a malformed ARN on its field without calling the API", async () => { + // When + const result = await connectAwsAccount({ + method: AWS_ACCESS_METHOD.ROLE, + values: { ...roleValues, role_arn: "arn:aws:s3:::bucket" }, + }); + + // Then + expect(result).toEqual({ + ok: false, + errors: [ + expect.objectContaining({ + source: { pointer: "/data/attributes/uid" }, + }), + ], + }); + expect(addProvider).not.toHaveBeenCalled(); + }); + }); + + describe("when connecting with access keys", () => { + it("registers the typed account id and sends only the keys as the secret", async () => { + // When + const result = await connectAwsAccount({ + method: AWS_ACCESS_METHOD.CREDENTIALS, + values: { + providerId: "", + providerType: "aws", + providerUid: "210987654321", + providerAlias: "", + aws_access_key_id: "AKIAEXAMPLE", + aws_secret_access_key: "secret", + aws_session_token: "", + }, + }); + + // Then + expect(result).toEqual({ ok: true, providerId: "provider-1" }); + expect(formEntries(0, addProvider)).toEqual({ + providerType: "aws", + providerUid: "210987654321", + }); + expect(formEntries(0, addCredentialsProvider)).toEqual({ + providerId: "provider-1", + providerType: "aws", + aws_access_key_id: "AKIAEXAMPLE", + aws_secret_access_key: "secret", + }); + expect(useProviderWizardStore.getState().via).toBe("credentials"); + }); + }); + + describe("when the API refuses the account", () => { + it("returns the provider errors and stores nothing", async () => { + // Given + const errors = [ + { + detail: "Provider with this uid already exists.", + source: { pointer: "/data/attributes/uid" }, + }, + ]; + addProvider.mockResolvedValueOnce({ errors }); + + // When + const result = await connectAwsAccount({ + method: AWS_ACCESS_METHOD.ROLE, + values: roleValues, + }); + + // Then + expect(result).toEqual({ ok: false, errors }); + expect(addCredentialsProvider).not.toHaveBeenCalled(); + expect(useProviderWizardStore.getState().providerId).toBeNull(); + }); + }); + + describe("when the API fails without field errors", () => { + it("reports the account failure instead of throwing", async () => { + // Given + addProvider.mockResolvedValueOnce({ error: "Server is unavailable." }); + + // When + const result = await connectAwsAccount({ + method: AWS_ACCESS_METHOD.ROLE, + values: roleValues, + }); + + // Then + expect(result).toEqual({ + ok: false, + errors: [{ detail: "Server is unavailable." }], + }); + expect(addCredentialsProvider).not.toHaveBeenCalled(); + expect(useProviderWizardStore.getState().providerId).toBeNull(); + }); + + it("reports an account response without an id instead of stalling", async () => { + // Given + addProvider.mockResolvedValueOnce({ data: {} }); + + // When + const result = await connectAwsAccount({ + method: AWS_ACCESS_METHOD.ROLE, + values: roleValues, + }); + + // Then + expect(result).toEqual({ + ok: false, + errors: [{ detail: expect.stringMatching(/try again/i) }], + }); + expect(addCredentialsProvider).not.toHaveBeenCalled(); + }); + + it("reports the credentials failure and keeps the account for a retry", async () => { + // Given + addCredentialsProvider.mockResolvedValueOnce({ + error: "Server is unavailable.", + }); + + // When + const result = await connectAwsAccount({ + method: AWS_ACCESS_METHOD.ROLE, + values: roleValues, + }); + + // Then + expect(result).toEqual({ + ok: false, + errors: [{ detail: "Server is unavailable." }], + }); + expect(useProviderWizardStore.getState()).toMatchObject({ + providerId: "provider-1", + secretId: null, + }); + }); + }); + + describe("when the credentials are refused after the account was registered", () => { + it("reuses the registered account on the next attempt instead of creating it twice", async () => { + // Given + const errors = [ + { + detail: "Invalid role ARN.", + source: { pointer: "/data/attributes/secret/role_arn" }, + }, + ]; + addCredentialsProvider.mockResolvedValueOnce({ errors }); + + // When + const first = await connectAwsAccount({ + method: AWS_ACCESS_METHOD.ROLE, + values: roleValues, + }); + const second = await connectAwsAccount({ + method: AWS_ACCESS_METHOD.ROLE, + values: roleValues, + }); + + // Then + expect(first).toEqual({ ok: false, errors }); + expect(second).toEqual({ ok: true, providerId: "provider-1" }); + expect(addProvider).toHaveBeenCalledOnce(); + expect(addCredentialsProvider).toHaveBeenCalledTimes(2); + expect(updateProvider).not.toHaveBeenCalled(); + }); + + it("renames the registered account when the alias changed before the retry", async () => { + // Given + addCredentialsProvider.mockResolvedValueOnce({ + errors: [{ detail: "Invalid role ARN." }], + }); + await connectAwsAccount({ + method: AWS_ACCESS_METHOD.ROLE, + values: roleValues, + }); + + // When + const second = await connectAwsAccount({ + method: AWS_ACCESS_METHOD.ROLE, + values: { ...roleValues, providerAlias: "Production EU" }, + }); + + // Then + expect(second).toEqual({ ok: true, providerId: "provider-1" }); + expect(addProvider).toHaveBeenCalledOnce(); + expect(formEntries(0, updateProvider)).toEqual({ + providerId: "provider-1", + providerAlias: "Production EU", + }); + expect(useProviderWizardStore.getState().providerAlias).toBe( + "Production EU", + ); + }); + }); + + describe("when the account was already connected in this wizard session", () => { + it("updates the stored credentials instead of creating a second secret", async () => { + // Given + await connectAwsAccount({ + method: AWS_ACCESS_METHOD.ROLE, + values: roleValues, + }); + updateCredentialsProvider.mockResolvedValueOnce({ + data: { id: "secret-1" }, + }); + + // When + const result = await connectAwsAccount({ + method: AWS_ACCESS_METHOD.ROLE, + values: { + ...roleValues, + role_arn: "arn:aws:iam::123456789012:role/ProwlerScanV2", + }, + }); + + // Then + expect(result).toEqual({ ok: true, providerId: "provider-1" }); + expect(addProvider).toHaveBeenCalledOnce(); + expect(addCredentialsProvider).toHaveBeenCalledOnce(); + expect(updateCredentialsProvider).toHaveBeenCalledExactlyOnceWith( + "secret-1", + expect.any(FormData), + ); + expect( + Object.fromEntries( + (updateCredentialsProvider.mock.calls[0][1] as FormData).entries(), + ), + ).toMatchObject({ + providerId: "provider-1", + providerType: "aws", + role_arn: "arn:aws:iam::123456789012:role/ProwlerScanV2", + }); + expect(useProviderWizardStore.getState().secretId).toBe("secret-1"); + }); + }); +}); diff --git a/ui/components/providers/wizard/steps/aws/connect-aws-account.ts b/ui/components/providers/wizard/steps/aws/connect-aws-account.ts new file mode 100644 index 0000000000..009a28e117 --- /dev/null +++ b/ui/components/providers/wizard/steps/aws/connect-aws-account.ts @@ -0,0 +1,192 @@ +import { + addCredentialsProvider, + addProvider, + updateCredentialsProvider, + updateProvider, +} from "@/actions/providers/providers"; +import { ProviderCredentialFields } from "@/lib/provider-credentials/provider-credential-fields"; +import { useProviderWizardStore } from "@/store/provider-wizard/store"; +import { useUIStore } from "@/store/ui/store"; +import type { ApiError } from "@/types"; +import { PROVIDER_WIZARD_MODE } from "@/types/provider-wizard"; + +import { + AWS_ROLE_ARN_MESSAGE, + parseAwsAccountIdFromRoleArn, +} from "./aws-role-arn"; +import { AWS_ACCESS_METHOD, type AwsAccessMethod } from "./types"; + +export const AWS_UID_ERROR_POINTER = "/data/attributes/uid"; + +// Account fields travel with the provider, never with its secret. +const ACCOUNT_FIELDS: readonly string[] = [ + ProviderCredentialFields.PROVIDER_UID, + ProviderCredentialFields.PROVIDER_ALIAS, +]; + +export interface AwsConnectInput { + method: AwsAccessMethod; + values: Record; +} + +interface AwsConnectSuccess { + ok: true; + providerId: string; +} + +interface AwsConnectFailure { + ok: false; + errors: ApiError[]; +} + +export type AwsConnectResult = AwsConnectSuccess | AwsConnectFailure; + +const asText = (value: unknown) => (typeof value === "string" ? value : ""); + +// Blank fields are left out, so optional credentials never reach the API as "". +const toFormData = (values: Record) => { + const formData = new FormData(); + Object.entries(values).forEach(([key, value]) => { + const text = asText(value).trim(); + if (text) formData.append(key, text); + }); + return formData; +}; + +interface CreatedResource { + id?: unknown; +} + +interface CreateActionResponse { + data?: CreatedResource; + error?: string; + errors?: ApiError[]; +} + +const UNCONFIRMED_RESPONSE_MESSAGE = + "The API did not confirm the request. Please try again."; + +// Actions resolve { errors } on a refusal and { error } on a crash, never throwing. +// A body with no id is reported too, or the step would stall without feedback. +const readCreatedId = (response: unknown) => { + const body = response as CreateActionResponse | undefined; + if (body?.errors?.length) return { id: null, errors: body.errors }; + if (body?.error) return { id: null, errors: [{ detail: body.error }] }; + const id = body?.data?.id; + if (typeof id !== "string" || !id) { + return { id: null, errors: [{ detail: UNCONFIRMED_RESPONSE_MESSAGE }] }; + } + return { id, errors: null }; +}; + +const resolveAccountId = ({ method, values }: AwsConnectInput) => + method === AWS_ACCESS_METHOD.ROLE + ? parseAwsAccountIdFromRoleArn( + asText(values[ProviderCredentialFields.ROLE_ARN]), + ) + : asText(values[ProviderCredentialFields.PROVIDER_UID]).trim() || null; + +// A retry may carry a new alias; the account registered earlier has to follow it. +const renameProvider = async (providerId: string, alias: string) => { + const store = useProviderWizardStore.getState(); + if ((store.providerAlias ?? "") === alias) + return { providerId, errors: null }; + + const updated = readCreatedId( + await updateProvider( + toFormData({ + [ProviderCredentialFields.PROVIDER_ID]: providerId, + [ProviderCredentialFields.PROVIDER_ALIAS]: alias, + }), + ), + ); + if (!updated.id) return { providerId: null, errors: updated.errors }; + + store.setProvider({ + id: providerId, + type: "aws", + uid: store.providerUid ?? "", + alias: alias || null, + }); + return { providerId, errors: null }; +}; + +// A retry after a refused secret must not register the same account twice. +const ensureProvider = async (uid: string, alias: string) => { + const store = useProviderWizardStore.getState(); + if (store.providerId && store.providerUid === uid) { + return renameProvider(store.providerId, alias); + } + + const created = readCreatedId( + await addProvider( + toFormData({ + [ProviderCredentialFields.PROVIDER_TYPE]: "aws", + [ProviderCredentialFields.PROVIDER_UID]: uid, + [ProviderCredentialFields.PROVIDER_ALIAS]: alias, + }), + ), + ); + if (!created.id) return { providerId: null, errors: created.errors }; + + const providerId = created.id; + store.setProvider({ + id: providerId, + type: "aws", + uid, + alias: alias || null, + }); + store.setSecretId(null); + store.setMode(PROVIDER_WIZARD_MODE.ADD); + // The layout only re-counts providers on a server render; flip the shared flag now. + useUIStore.getState().setHasProviders(true); + return { providerId, errors: null }; +}; + +/** Registers the AWS account and stores its credentials in a single submit. */ +export async function connectAwsAccount( + input: AwsConnectInput, +): Promise { + const uid = resolveAccountId(input); + if (!uid) { + return { + ok: false, + errors: [ + { + detail: AWS_ROLE_ARN_MESSAGE, + source: { pointer: AWS_UID_ERROR_POINTER }, + } as ApiError, + ], + }; + } + + const alias = asText( + input.values[ProviderCredentialFields.PROVIDER_ALIAS], + ).trim(); + const provider = await ensureProvider(uid, alias); + if (!provider.providerId) return { ok: false, errors: provider.errors ?? [] }; + + const secretValues = Object.fromEntries( + Object.entries(input.values).filter( + ([key]) => !ACCOUNT_FIELDS.includes(key), + ), + ); + const secretFormData = toFormData({ + ...secretValues, + [ProviderCredentialFields.PROVIDER_ID]: provider.providerId, + [ProviderCredentialFields.PROVIDER_TYPE]: "aws", + }); + // A provider holds one secret: resubmitting a connected account edits it in place. + const storedSecretId = useProviderWizardStore.getState().secretId; + const secret = readCreatedId( + storedSecretId + ? await updateCredentialsProvider(storedSecretId, secretFormData) + : await addCredentialsProvider(secretFormData), + ); + if (!secret.id) return { ok: false, errors: secret.errors ?? [] }; + + const store = useProviderWizardStore.getState(); + store.setSecretId(secret.id); + store.setVia(input.method); + return { ok: true, providerId: provider.providerId }; +} diff --git a/ui/components/providers/wizard/steps/aws/types.ts b/ui/components/providers/wizard/steps/aws/types.ts new file mode 100644 index 0000000000..6273d05e6d --- /dev/null +++ b/ui/components/providers/wizard/steps/aws/types.ts @@ -0,0 +1,16 @@ +export const AWS_ACCESS_METHOD = { + ROLE: "role", + CREDENTIALS: "credentials", +} as const; + +export type AwsAccessMethod = + (typeof AWS_ACCESS_METHOD)[keyof typeof AWS_ACCESS_METHOD]; + +/** What the step publishes so the wizard can draw its footer. */ +export interface AwsConnectUiState { + showBack: boolean; + showAction: boolean; + actionLabel: string; + actionDisabled: boolean; + isLoading: boolean; +} diff --git a/ui/components/providers/wizard/steps/connect-step.tsx b/ui/components/providers/wizard/steps/connect-step.tsx index 649f0d1562..23b45b46c0 100644 --- a/ui/components/providers/wizard/steps/connect-step.tsx +++ b/ui/components/providers/wizard/steps/connect-step.tsx @@ -6,11 +6,14 @@ import { ConnectAccountForm, ConnectAccountSuccessData, } from "@/components/providers/workflow/forms"; +import { endActiveTour } from "@/lib/tours/use-driver-tour"; import { useProviderWizardStore } from "@/store/provider-wizard/store"; -import { OrgFlowType } from "@/types/organizations"; +import { useUIStore } from "@/store/ui/store"; +import { ORGANIZATION_TYPE, OrgFlowType } from "@/types/organizations"; import { PROVIDER_WIZARD_MODE } from "@/types/provider-wizard"; import { ProviderType } from "@/types/providers"; +import { AwsConnectStep } from "./aws/aws-connect-step"; import { WIZARD_FOOTER_ACTION_TYPE, WizardFooterConfig, @@ -18,20 +21,28 @@ import { interface ConnectStepProps { onNext: () => void; + /** AWS registers, stores and tests the account in this step, so it skips ahead. */ + onCredentialsSaved: () => void; onSelectOrganizations: (orgType: OrgFlowType) => void; onFooterChange: (config: WizardFooterConfig) => void; onProviderTypeChange: (providerType: ProviderType | null) => void; + /** Provider the user was already working with, e.g. when returning from the AWS organization flow. */ + initialProviderType?: ProviderType | null; } export function ConnectStep({ onNext, + onCredentialsSaved, onSelectOrganizations, onFooterChange, onProviderTypeChange, + initialProviderType = null, }: ConnectStepProps) { const { setProvider, setVia, setSecretId, setMode } = useProviderWizardStore(); const backHandlerRef = useRef<(() => void) | null>(null); + // Local state needed: AWS swaps the generic account form for its one-step form. + const [isAwsFlow, setIsAwsFlow] = useState(initialProviderType === "aws"); const [uiState, setUiState] = useState({ showBack: false, showAction: false, @@ -52,15 +63,25 @@ export function ConnectStep({ setVia(null); setSecretId(null); setMode(PROVIDER_WIZARD_MODE.ADD); + // The layout only re-counts providers on a server render; flip the shared flag now. + useUIStore.getState().setHasProviders(true); onNext(); }; useEffect(() => { + // The footer sits outside the tour's spotlight, so once the user can continue + // the tour has done its job and gets out of the way. No-op off-onboarding. + if (uiState.showAction && !uiState.actionDisabled && !uiState.isLoading) { + endActiveTour(); + } onFooterChange({ showBack: uiState.showBack, backLabel: "Back", backDisabled: uiState.isLoading, - onBack: () => backHandlerRef.current?.(), + // Leaving AWS remounts the generic form on a fresh provider list. + onBack: isAwsFlow + ? () => setIsAwsFlow(false) + : () => backHandlerRef.current?.(), showAction: uiState.showAction, actionLabel: uiState.actionLabel, actionLoading: uiState.isLoading, @@ -68,7 +89,25 @@ export function ConnectStep({ actionType: WIZARD_FOOTER_ACTION_TYPE.SUBMIT, actionFormId: formId, }); - }, [onFooterChange, uiState]); + }, [isAwsFlow, onFooterChange, uiState]); + + const handleProviderTypeChange = (providerType: ProviderType | null) => { + onProviderTypeChange(providerType); + if (providerType === "aws") setIsAwsFlow(true); + }; + + if (isAwsFlow) { + return ( + + onSelectOrganizations(ORGANIZATION_TYPE.AWS) + } + onUiStateChange={setUiState} + /> + ); + } return ( { backHandlerRef.current = handler; diff --git a/ui/components/providers/wizard/wizard-stepper.tsx b/ui/components/providers/wizard/wizard-stepper.tsx index 4d86680977..5ef115eff5 100644 --- a/ui/components/providers/wizard/wizard-stepper.tsx +++ b/ui/components/providers/wizard/wizard-stepper.tsx @@ -46,6 +46,18 @@ const STEPS: StepConfig[] = [ export const PROVIDER_WIZARD_STEPS = STEPS; +// AWS registers the account, stores its credentials and tests the connection in +// one step, so the wizard goes straight from CONNECT to LAUNCH. +export const AWS_PROVIDER_WIZARD_STEPS: StepConfig[] = [ + { + label: "Link a Provider", + description: + "Enter the account details and the credentials Prowler will use, then test the connection.", + icon: FolderGit2, + }, + STEPS[3], +]; + export function WizardStepper({ currentStep, stepOffset = 0, diff --git a/ui/components/providers/workflow/credentials-role-helper.test.tsx b/ui/components/providers/workflow/credentials-role-helper.test.tsx new file mode 100644 index 0000000000..e1908a4f34 --- /dev/null +++ b/ui/components/providers/workflow/credentials-role-helper.test.tsx @@ -0,0 +1,89 @@ +import { render, screen } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { afterEach, beforeEach, describe, expect, it } from "vitest"; + +import { + PROVIDER_FUNNEL_EVENT, + type ProviderFunnelDetail, +} from "@/lib/provider-funnel/provider-funnel-events"; + +import { CredentialsRoleHelper } from "./credentials-role-helper"; + +const templateLinks = { + cloudformation: "https://example.com/template.yml", + cloudformationQuickLink: "https://example.com/quick-create", + terraform: "https://example.com/terraform", +}; + +describe("CredentialsRoleHelper", () => { + const funnelSignals: ProviderFunnelDetail[] = []; + const recordFunnelSignal: EventListener = (event) => { + funnelSignals.push((event as CustomEvent).detail); + }; + + beforeEach(() => { + funnelSignals.length = 0; + window.addEventListener(PROVIDER_FUNNEL_EVENT, recordFunnelSignal); + }); + + afterEach(() => { + window.removeEventListener(PROVIDER_FUNNEL_EVENT, recordFunnelSignal); + }); + + describe("when connecting a provider", () => { + it("signals which role template the user opened", async () => { + // Given + const user = userEvent.setup(); + render( + , + ); + + // When + await user.click( + screen.getByRole("link", { name: /Create the IAM role in AWS/i }), + ); + await user.click( + screen.getByRole("button", { name: /Other ways to create the role/i }), + ); + await user.click( + screen.getByRole("link", { name: "CloudFormation Template" }), + ); + await user.click(screen.getByRole("link", { name: "Terraform Code" })); + + // Then + expect(funnelSignals).toEqual([ + { + step: "role_template_opened", + template: "cloudformation_quick_create", + }, + { step: "role_template_opened", template: "cloudformation_template" }, + { step: "role_template_opened", template: "terraform" }, + ]); + }); + }); + + describe("when configuring an integration", () => { + it("stays out of the provider funnel", async () => { + // Given + const user = userEvent.setup(); + render( + , + ); + + // When + await user.click( + screen.getByRole("link", { name: /Create the IAM role in AWS/i }), + ); + + // Then + expect(funnelSignals).toEqual([]); + }); + }); +}); diff --git a/ui/components/providers/workflow/credentials-role-helper.tsx b/ui/components/providers/workflow/credentials-role-helper.tsx index 12e0bd3502..76bee3ec44 100644 --- a/ui/components/providers/workflow/credentials-role-helper.tsx +++ b/ui/components/providers/workflow/credentials-role-helper.tsx @@ -1,103 +1,132 @@ "use client"; +import { ChevronDownIcon, ExternalLink } from "lucide-react"; + import { IdIcon } from "@/components/icons"; -import { Button } from "@/components/shadcn"; +import { Button } from "@/components/shadcn/button/button"; import { CodeSnippet } from "@/components/shadcn/code-snippet/code-snippet"; +import { + Collapsible, + CollapsibleContent, + CollapsibleTrigger, +} from "@/components/shadcn/collapsible"; +import { + dispatchProviderFunnel, + PROVIDER_FUNNEL_STEP, + ROLE_TEMPLATE_KIND, + type RoleTemplateKind, +} from "@/lib/provider-funnel/provider-funnel-events"; +import { isCloud } from "@/lib/shared/env"; import { IntegrationType } from "@/types/integrations"; +interface CredentialsRoleTemplateLinks { + cloudformation: string; + cloudformationQuickLink: string; + terraform: string; +} + interface CredentialsRoleHelperProps { externalId: string; - templateLinks: { - cloudformation: string; - cloudformationQuickLink: string; - terraform: string; - }; + templateLinks: CredentialsRoleTemplateLinks; integrationType?: IntegrationType; } +const describeRole = (integrationType?: IntegrationType) => { + if (integrationType === "amazon_s3") { + return "A read-only IAM role must be manually created or updated. Open the AWS console to do it from the stack; the External ID comes filled in."; + } + if (integrationType) { + return "A read-only IAM role must be manually created. Open the AWS console to create it from the stack; the External ID comes filled in."; + } + return isCloud() + ? "Open the AWS console to create a read-only IAM role that Prowler Cloud can assume. The stack comes with your External ID filled in." + : "Open the AWS console to create a read-only IAM role that Prowler can assume. Fill in the AWS account Prowler runs from; the External ID comes filled in."; +}; + +/** One button creates the IAM role; the raw templates stay tucked away. */ export const CredentialsRoleHelper = ({ externalId, templateLinks, integrationType, }: CredentialsRoleHelperProps) => { - const isAmazonS3 = integrationType === "amazon_s3"; + // Integrations reuse this helper; only the add-provider journey is signalled. + const signalTemplateOpened = (template: RoleTemplateKind) => { + if (integrationType) return; + dispatchProviderFunnel({ + step: PROVIDER_FUNNEL_STEP.ROLE_TEMPLATE_OPENED, + template, + }); + }; return ( -
-
-

- A read-only IAM role must be manually created - {isAmazonS3 ? " or updated" : ""} -

+
+

+ {describeRole(integrationType)} +

- + Create the IAM role in AWS + + + -
-
- - or - -
-
+
+ + External ID: + + } /> +
-

- {isAmazonS3 - ? "Refer to the documentation" - : "Use one of the following templates to create the IAM role"} -

- - - -
- - External ID: - - } /> -
-
+ +
); }; diff --git a/ui/components/providers/workflow/forms/connect-account-form.test.tsx b/ui/components/providers/workflow/forms/connect-account-form.test.tsx index 2a906d431c..763cbcb4b3 100644 --- a/ui/components/providers/workflow/forms/connect-account-form.test.tsx +++ b/ui/components/providers/workflow/forms/connect-account-form.test.tsx @@ -1,6 +1,6 @@ import { render, screen, waitFor } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; -import { beforeEach, describe, expect, it, vi } from "vitest"; +import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; const { addProvider, updateProvider, getInstalledRegistryProviderOptions } = vi.hoisted(() => ({ @@ -86,6 +86,32 @@ describe("provider account aliases", () => { describe("Registry provider source tabs", () => { beforeEach(() => { vi.clearAllMocks(); + vi.stubEnv("UI_CLOUD_ENABLED", "true"); + }); + + afterEach(() => { + vi.unstubAllEnvs(); + }); + + it("never mentions Registry outside Cloud, even when discovery would fail", async () => { + // Given + vi.stubEnv("UI_CLOUD_ENABLED", "false"); + getInstalledRegistryProviderOptions.mockRejectedValue(new Error("network")); + + // When + render(); + + // Then + expect( + await screen.findByRole("option", { name: /Amazon Web Services/ }), + ).toBeVisible(); + expect(getInstalledRegistryProviderOptions).not.toHaveBeenCalled(); + expect( + screen.queryByText("Registry providers could not be loaded"), + ).not.toBeInTheDocument(); + expect( + screen.queryByRole("tab", { name: "Registry" }), + ).not.toBeInTheDocument(); }); it("hides the Registry tab when discovery denies access (Local or flag off)", async () => { diff --git a/ui/components/providers/workflow/forms/connect-account-form.tsx b/ui/components/providers/workflow/forms/connect-account-form.tsx index ddc880e3e7..f0199a0c80 100644 --- a/ui/components/providers/workflow/forms/connect-account-form.tsx +++ b/ui/components/providers/workflow/forms/connect-account-form.tsx @@ -9,7 +9,6 @@ import { useForm, UseFormReturn } from "react-hook-form"; import { addProvider, updateProvider } from "@/actions/providers/providers"; import { addRegistryProvider } from "@/actions/providers/registry-provider"; import { getInstalledRegistryProviderOptions } from "@/actions/registry/registry"; -import { AwsMethodSelector } from "@/components/providers/organizations/aws-method-selector"; import { AzureMethodSelector } from "@/components/providers/organizations/azure-method-selector"; import { GcpMethodSelector } from "@/components/providers/organizations/gcp-method-selector"; import { WizardInputField } from "@/components/providers/workflow/forms/fields"; @@ -22,6 +21,7 @@ import { REGISTRY_PROVIDER_DISCOVERY, type RegistryProviderOption, } from "@/lib/registry/provider-options"; +import { isCloud } from "@/lib/shared/env"; import { createAddProviderFormSchema, AddProviderFormValues, @@ -49,12 +49,13 @@ export interface ConnectAccountSuccessData { /** * Provider types that offer an organization-onboarding method choice: exactly the - * ones with an onboarding flow, so a new flow type cannot miss the fork. + * ones with an onboarding flow, so a new flow type cannot miss the fork. AWS is the + * exception: the wizard's own AWS step hosts its single-account/organization switch. */ function providerHasOrgMethod( providerType: ProviderType | undefined, ): providerType is OrgFlowType { - return toOrgFlowType(providerType) !== undefined; + return providerType !== "aws" && toOrgFlowType(providerType) !== undefined; } interface ConnectAccountFormProps { @@ -233,6 +234,8 @@ export const ConnectAccountForm = ({ const createdAccount = useRef(null); useEffect(() => { + // Registry is Cloud-only: elsewhere never ask, so a failure cannot surface it. + if (!isCloud()) return; let active = true; const load = async () => { try { @@ -504,18 +507,6 @@ export const ConnectAccountForm = ({ />
)} - {/* Step 2: AWS method selector (before choosing a method) */} - {prevStep === 2 && providerType === "aws" && method === null && ( - <> - - setMethod("single")} - onSelectOrganizations={() => - onSelectOrganizations?.(ORGANIZATION_TYPE.AWS) - } - /> - - )} {/* Step 2: Azure method selector (before choosing a method) */} {prevStep === 2 && providerType === "azure" && method === null && ( <> diff --git a/ui/components/providers/workflow/forms/select-credentials-type/aws/credentials-type/aws-role-credentials-form.tsx b/ui/components/providers/workflow/forms/select-credentials-type/aws/credentials-type/aws-role-credentials-form.tsx index e344b7e92b..082cd1443b 100644 --- a/ui/components/providers/workflow/forms/select-credentials-type/aws/credentials-type/aws-role-credentials-form.tsx +++ b/ui/components/providers/workflow/forms/select-credentials-type/aws/credentials-type/aws-role-credentials-form.tsx @@ -3,21 +3,16 @@ import { Control, UseFormSetValue, useWatch } from "react-hook-form"; import { CredentialsRoleHelper } from "@/components/providers/workflow"; import { WizardInputField } from "@/components/providers/workflow/forms/fields"; -import { Badge } from "@/components/shadcn/badge/badge"; import { Checkbox } from "@/components/shadcn/checkbox/checkbox"; -import { - Select, - SelectContent, - SelectItem, - SelectTrigger, - SelectValue, -} from "@/components/shadcn/select/select"; import { Separator } from "@/components/shadcn/separator/separator"; import { ProviderCredentialFields } from "@/lib/provider-credentials/provider-credential-fields"; import { isCloud } from "@/lib/shared/env"; import { AWSCredentialsRole } from "@/types"; import { IntegrationType } from "@/types/integrations"; +import { AwsRoleCredentialsSource } from "./aws-role-credentials-source"; +import { AwsRoleOptionalFields } from "./aws-role-optional-fields"; + export const AWSRoleCredentialsForm = ({ control, setValue, @@ -80,81 +75,12 @@ export const AWSRoleCredentialsForm = ({ )}
- - Specify which AWS credentials to use - - -
- -
- - {credentialsType === "access-secret-key" && ( - <> - - - - - )} + {type === "providers" ? ( @@ -210,31 +136,7 @@ export const AWSRoleCredentialsForm = ({ isRequired /> - - Optional fields - -
- - -
+ )} diff --git a/ui/components/providers/workflow/forms/select-credentials-type/aws/credentials-type/aws-role-credentials-source.tsx b/ui/components/providers/workflow/forms/select-credentials-type/aws/credentials-type/aws-role-credentials-source.tsx new file mode 100644 index 0000000000..491f24df81 --- /dev/null +++ b/ui/components/providers/workflow/forms/select-credentials-type/aws/credentials-type/aws-role-credentials-source.tsx @@ -0,0 +1,105 @@ +import { Control, UseFormSetValue } from "react-hook-form"; + +import { WizardInputField } from "@/components/providers/workflow/forms/fields"; +import { Badge } from "@/components/shadcn/badge/badge"; +import { + Select, + SelectContent, + SelectItem, + SelectTrigger, + SelectValue, +} from "@/components/shadcn/select/select"; +import { ProviderCredentialFields } from "@/lib/provider-credentials/provider-credential-fields"; +import { AWSCredentialsRole } from "@/types"; + +interface AwsRoleCredentialsSourceProps { + control: Control; + setValue: UseFormSetValue; + credentialsType: string; + isCloudEnv: boolean; +} + +/** Which credentials Prowler uses to assume the role, plus the keys when they are static. */ +export const AwsRoleCredentialsSource = ({ + control, + setValue, + credentialsType, + isCloudEnv, +}: AwsRoleCredentialsSourceProps) => ( + <> +
+ + Specify which AWS credentials to use + + +
+ + {credentialsType === "access-secret-key" && ( + <> + + + + + )} + +); diff --git a/ui/components/providers/workflow/forms/select-credentials-type/aws/credentials-type/aws-role-optional-fields.tsx b/ui/components/providers/workflow/forms/select-credentials-type/aws/credentials-type/aws-role-optional-fields.tsx new file mode 100644 index 0000000000..9522da6ed2 --- /dev/null +++ b/ui/components/providers/workflow/forms/select-credentials-type/aws/credentials-type/aws-role-optional-fields.tsx @@ -0,0 +1,37 @@ +import { Control } from "react-hook-form"; + +import { WizardInputField } from "@/components/providers/workflow/forms/fields"; +import { ProviderCredentialFields } from "@/lib/provider-credentials/provider-credential-fields"; +import { AWSCredentialsRole } from "@/types"; + +interface AwsRoleOptionalFieldsProps { + control: Control; +} + +/** Session name and duration of the assumed role; both optional. */ +export const AwsRoleOptionalFields = ({ + control, +}: AwsRoleOptionalFieldsProps) => ( +
+ + +
+); diff --git a/ui/components/providers/workflow/forms/test-connection-form.test.tsx b/ui/components/providers/workflow/forms/test-connection-form.test.tsx new file mode 100644 index 0000000000..1740f1dde5 --- /dev/null +++ b/ui/components/providers/workflow/forms/test-connection-form.test.tsx @@ -0,0 +1,185 @@ +import { render, screen } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { beforeEach, describe, expect, it, vi } from "vitest"; + +const { pushMock, testProviderConnectionMock } = vi.hoisted(() => ({ + pushMock: vi.fn(), + testProviderConnectionMock: vi.fn(), +})); + +vi.mock("next/navigation", () => ({ + useRouter: () => ({ push: pushMock, back: vi.fn() }), +})); + +vi.mock("@/actions/providers", () => ({ + deleteCredentials: vi.fn(), +})); + +vi.mock("@/lib/provider-helpers", () => ({ + testProviderConnection: testProviderConnectionMock, +})); + +import { CONNECTION_CHECK_STATUS } from "@/types/providers"; + +import { + TestConnectionForm, + type TestConnectionProviderData, +} from "./test-connection-form"; + +const providerData: TestConnectionProviderData = { + data: { + id: "provider-1", + type: "providers", + attributes: { + uid: "111111111111", + connection: { connected: false, last_checked_at: null }, + provider: "aws", + alias: "Production", + scanner_args: {}, + }, + relationships: { + secret: { data: { type: "provider-secrets", id: "secret-1" } }, + }, + }, +}; + +describe("TestConnectionForm", () => { + beforeEach(() => { + pushMock.mockReset(); + testProviderConnectionMock.mockReset(); + }); + + it("advances on a confirmed successful connection", async () => { + // Given + testProviderConnectionMock.mockResolvedValue({ + status: CONNECTION_CHECK_STATUS.SUCCESS, + error: null, + }); + const onSuccess = vi.fn(); + const user = userEvent.setup(); + + render( + , + ); + + // When + await user.click(screen.getByRole("button", { name: /continue/i })); + + // Then + expect(onSuccess).toHaveBeenCalledTimes(1); + expect( + screen.queryByText(/issue with your credentials/i), + ).not.toBeInTheDocument(); + }); + + it("shows a destructive failure message for a confirmed failed connection", async () => { + // Given + testProviderConnectionMock.mockResolvedValue({ + status: CONNECTION_CHECK_STATUS.FAILED, + error: "Role trust policy mismatch.", + }); + const user = userEvent.setup(); + + render( + , + ); + + // When + await user.click(screen.getByRole("button", { name: /continue/i })); + + // Then + expect(screen.getByText("Role trust policy mismatch.")).toBeInTheDocument(); + expect( + screen.getByText(/issue with your credentials/i), + ).toBeInTheDocument(); + expect( + screen.getByRole("button", { name: /reset credentials/i }), + ).toBeInTheDocument(); + // Announced to screen readers as soon as it appears, not only on focus. + expect(screen.getByRole("status")).toHaveTextContent( + "Role trust policy mismatch.", + ); + }); + + it("shows a neutral still-running message, not a failure, when the check is still pending", async () => { + // Given: the wait was exhausted and the provider's stored state could not + // confirm an outcome yet (the backend check is genuinely still running). + testProviderConnectionMock.mockResolvedValue({ + status: CONNECTION_CHECK_STATUS.PENDING, + error: + "The connection test is still running. Refresh in a moment to see the result.", + }); + const onSuccess = vi.fn(); + const user = userEvent.setup(); + + render( + , + ); + + // When + await user.click(screen.getByRole("button", { name: /continue/i })); + + // Then: the neutral message shows, but nothing reads as a credentials failure. + expect(screen.getByText(/still running/i)).toBeInTheDocument(); + // Announced to screen readers, same as the failure banner. + expect(screen.getByRole("status")).toHaveTextContent(/still running/i); + expect( + screen.queryByText(/issue with your credentials/i), + ).not.toBeInTheDocument(); + expect( + screen.queryByRole("button", { name: /reset credentials/i }), + ).not.toBeInTheDocument(); + // Does not advance either -- the outcome is still unknown. + expect(onSuccess).not.toHaveBeenCalled(); + expect(pushMock).not.toHaveBeenCalled(); + // The retry control is explicit about what pressing it does now: it is no + // longer the first check, so "Continue" would be misleading. + expect( + screen.getByRole("button", { name: /check again/i }), + ).toBeInTheDocument(); + }); + + it("re-runs the check when 'Check again' is pressed on a still-pending result", async () => { + // Given: the first attempt came back pending. + testProviderConnectionMock.mockResolvedValueOnce({ + status: CONNECTION_CHECK_STATUS.PENDING, + error: "The connection test is still running.", + }); + const user = userEvent.setup(); + + render( + , + ); + await user.click(screen.getByRole("button", { name: /continue/i })); + expect( + screen.getByRole("button", { name: /check again/i }), + ).toBeInTheDocument(); + + // When: pressing "Check again" resolves this time. + testProviderConnectionMock.mockResolvedValueOnce({ + status: CONNECTION_CHECK_STATUS.SUCCESS, + error: null, + }); + await user.click(screen.getByRole("button", { name: /check again/i })); + + // Then + expect(testProviderConnectionMock).toHaveBeenCalledTimes(2); + expect(screen.queryByText(/still running/i)).not.toBeInTheDocument(); + }); +}); diff --git a/ui/components/providers/workflow/forms/test-connection-form.tsx b/ui/components/providers/workflow/forms/test-connection-form.tsx index 687b036cce..a0beb2cb4c 100644 --- a/ui/components/providers/workflow/forms/test-connection-form.tsx +++ b/ui/components/providers/workflow/forms/test-connection-form.tsx @@ -10,11 +10,15 @@ import { useForm } from "react-hook-form"; import { z } from "zod"; import { deleteCredentials } from "@/actions/providers"; -import { CheckIcon } from "@/components/icons"; +import { CheckIcon, ConnectionPending } from "@/components/icons"; import { Button } from "@/components/shadcn"; import { Form } from "@/components/shadcn/form"; -import { testProviderConnection } from "@/lib/provider-helpers"; +import { + testProviderConnection, + type TestConnectionResult, +} from "@/lib/provider-helpers"; import { ProviderType, testConnectionFormSchema } from "@/types"; +import { CONNECTION_CHECK_STATUS } from "@/types/providers"; import { ProviderConnectionInfo } from "./provider-connection-info"; @@ -69,10 +73,8 @@ export const TestConnectionForm = ({ const providerId = searchParams.id; const [apiErrorMessage, setApiErrorMessage] = useState(null); - const [connectionStatus, setConnectionStatus] = useState<{ - connected: boolean; - error: string | null; - } | null>(null); + const [connectionStatus, setConnectionStatus] = + useState(null); const [isResettingCredentials, setIsResettingCredentials] = useState(false); const formSchema = testConnectionFormSchema; @@ -103,7 +105,7 @@ export const TestConnectionForm = ({ setConnectionStatus(result); - if (result.connected) { + if (result.status === CONNECTION_CHECK_STATUS.SUCCESS) { if (onSuccess) { onSuccess(); return; @@ -167,13 +169,17 @@ export const TestConnectionForm = ({
)} - {connectionStatus && !connectionStatus.connected && ( + {connectionStatus?.status === CONNECTION_CHECK_STATUS.FAILED && ( <> -
+
@@ -189,6 +195,30 @@ export const TestConnectionForm = ({ )} + {connectionStatus?.status === CONNECTION_CHECK_STATUS.PENDING && ( +
+
+ {/* Static, not spinning: nothing is polling any more once the wait + is exhausted, so an animated spinner would misrepresent this as + still in progress. */} +
+
+

+ {connectionStatus.error || + "The connection test is still running. Refresh in a moment to see the result."} +

+
+
+ )} + Back to providers - ) : connectionStatus?.error ? ( + ) : connectionStatus?.status === CONNECTION_CHECK_STATUS.FAILED ? (
diff --git a/ui/components/scans/table/cells/scan-info-cell.tsx b/ui/components/scans/table/cells/scan-info-cell.tsx index 226732d8fa..c0a524b4a8 100644 --- a/ui/components/scans/table/cells/scan-info-cell.tsx +++ b/ui/components/scans/table/cells/scan-info-cell.tsx @@ -27,12 +27,25 @@ export function ScanInfoCell({ scan }: { scan: ScanProps }) { } return ( -
+
+ {scan.attributes.is_partial && ( + + + + Partial + + + + Re-checked a few resources. Overviews still reflect the latest full + scan. + + + )}
); } diff --git a/ui/components/scans/table/scan-jobs-columns.test.tsx b/ui/components/scans/table/scan-jobs-columns.test.tsx index 496fd11e14..8d0e15143c 100644 --- a/ui/components/scans/table/scan-jobs-columns.test.tsx +++ b/ui/components/scans/table/scan-jobs-columns.test.tsx @@ -11,8 +11,17 @@ import { vi.mock("@/components/shadcn", async (importOriginal) => ({ ...(await importOriginal>()), - Badge: ({ children }: { children: ReactNode }) => {children}, + Badge: ({ + children, + tabIndex, + }: { + children: ReactNode; + tabIndex?: number; + }) => {children}, Progress: () =>
, + Tooltip: ({ children }: { children: ReactNode }) => <>{children}, + TooltipTrigger: ({ children }: { children: ReactNode }) => <>{children}, + TooltipContent: () => null, StackedCell: ({ primary, secondary, @@ -188,6 +197,25 @@ describe("getScanJobsColumns", () => { expect(screen.getByText("ID: scan-1")).toBeInTheDocument(); }); + it("labels a partial scan next to its alias", () => { + // Prowler Cloud exposes is_partial for re-checks of a few resources. + const scan = makeCompletedScan(); + renderCell("scanInfo", { + ...scan, + attributes: { ...scan.attributes, is_partial: true }, + }); + + expect(screen.getByText("Production scan")).toBeInTheDocument(); + // Focusable so keyboard users can reach the tooltip. + expect(screen.getByText("Partial")).toHaveAttribute("tabindex", "0"); + }); + + it("shows no partial label on a full scan", () => { + renderCell("scanInfo", makeCompletedScan()); + + expect(screen.queryByText("Partial")).not.toBeInTheDocument(); + }); + it("renders the completed duration column", () => { renderCell("duration", makeCompletedScan()); diff --git a/ui/components/scans/table/scan-jobs-columns.tsx b/ui/components/scans/table/scan-jobs-columns.tsx index dc040368fa..0a9e4ce311 100644 --- a/ui/components/scans/table/scan-jobs-columns.tsx +++ b/ui/components/scans/table/scan-jobs-columns.tsx @@ -20,9 +20,13 @@ import { } from "./cells"; import { ScanJobsRowActions } from "./scan-jobs-row-actions"; -interface GetScanJobsColumnsOptions { - tab: ScanJobsTab; +interface ScanJobsRowActionOptions { capability?: ScanScheduleCapability; + subscriptionOnly?: boolean; +} + +interface GetScanJobsColumnsOptions extends ScanJobsRowActionOptions { + tab: ScanJobsTab; } const accountColumn: ColumnDef = { @@ -121,12 +125,17 @@ const resourcesColumn: ColumnDef = { const actionsColumn = ( tab: ScanJobsTab, - capability?: ScanScheduleCapability, + { capability, subscriptionOnly }: ScanJobsRowActionOptions, ): ColumnDef => ({ id: "actions", header: ({ column }) => , cell: ({ row }) => ( - + ), enableSorting: false, }); @@ -141,7 +150,7 @@ const durationColumn: ColumnDef = { }; const activeColumns = ( - capability?: ScanScheduleCapability, + rowActionOptions: ScanJobsRowActionOptions, ): ColumnDef[] => [ accountColumn, scanInfoColumn, @@ -166,11 +175,11 @@ const activeColumns = ( ), enableSorting: false, }, - actionsColumn(SCAN_JOBS_TAB.ACTIVE, capability), + actionsColumn(SCAN_JOBS_TAB.ACTIVE, rowActionOptions), ]; const completedColumns = ( - capability?: ScanScheduleCapability, + rowActionOptions: ScanJobsRowActionOptions, ): ColumnDef[] => [ accountColumn, scanInfoColumn, @@ -197,28 +206,30 @@ const completedColumns = ( ), cell: ({ row }) => renderDateCell(row.original.attributes.completed_at), }, - actionsColumn(SCAN_JOBS_TAB.COMPLETED, capability), + actionsColumn(SCAN_JOBS_TAB.COMPLETED, rowActionOptions), ]; const scheduledColumns = ( - capability?: ScanScheduleCapability, + rowActionOptions: ScanJobsRowActionOptions, ): ColumnDef[] => [ accountColumn, scanInfoColumn, scheduledScanScheduleColumn, nextScanColumn, lastScanColumn, - actionsColumn(SCAN_JOBS_TAB.SCHEDULED, capability), + actionsColumn(SCAN_JOBS_TAB.SCHEDULED, rowActionOptions), ]; export function getScanJobsColumns( options: GetScanJobsColumnsOptions, ): ColumnDef[] { - if (options.tab === SCAN_JOBS_TAB.SCHEDULED) { - return scheduledColumns(options.capability); + const { tab, ...rowActionOptions } = options; + + if (tab === SCAN_JOBS_TAB.SCHEDULED) { + return scheduledColumns(rowActionOptions); } - if (options.tab === SCAN_JOBS_TAB.ACTIVE) { - return activeColumns(options.capability); + if (tab === SCAN_JOBS_TAB.ACTIVE) { + return activeColumns(rowActionOptions); } - return completedColumns(options.capability); + return completedColumns(rowActionOptions); } diff --git a/ui/components/scans/table/scan-jobs-row-actions.test.tsx b/ui/components/scans/table/scan-jobs-row-actions.test.tsx index 07ea83ba94..d9b8fba781 100644 --- a/ui/components/scans/table/scan-jobs-row-actions.test.tsx +++ b/ui/components/scans/table/scan-jobs-row-actions.test.tsx @@ -2,7 +2,9 @@ import { render, screen } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; +import { useCloudUpgradeStore } from "@/store/cloud-upgrade/store"; import type { ScanProps } from "@/types"; +import { PAID_PLAN_UPGRADE_FEATURE } from "@/types/cloud-upgrade"; import { SCAN_SCHEDULE_CAPABILITY } from "@/types/schedules"; import { ScanJobsRowActions } from "./scan-jobs-row-actions"; @@ -122,6 +124,7 @@ describe("ScanJobsRowActions", () => { afterEach(() => { vi.unstubAllEnvs(); vi.clearAllMocks(); + useCloudUpgradeStore.getState().closeCloudUpgrade(); }); it("opens the Edit modal seeded with the current scan name", async () => { @@ -384,6 +387,65 @@ describe("ScanJobsRowActions", () => { expect(downloadScanZipMock).toHaveBeenCalledWith("scan-1", toastMock); }); + it("offers neither report download nor compliance for a partial scan", async () => { + // A partial scan re-checks a few resources: no report files, no compliance. + const user = userEvent.setup(); + render( + , + ); + + await user.click( + screen.getByRole("button", { name: /open actions menu/i }), + ); + + expect( + screen.queryByRole("menuitem", { name: /download scan reports/i }), + ).not.toBeInTheDocument(); + expect( + screen.queryByRole("menuitem", { name: /view compliance/i }), + ).not.toBeInTheDocument(); + // The rest of the completed-scan actions stay available. + expect( + screen.getByRole("menuitem", { name: /view findings/i }), + ).toBeInTheDocument(); + }); + + it("opens the paid plan upgrade instead of downloading subscription-only reports", async () => { + // Given + const user = userEvent.setup(); + render( + , + ); + + // When + await user.click( + screen.getByRole("button", { name: /open actions menu/i }), + ); + await user.click( + screen.getByRole("menuitem", { name: /download scan reports/i }), + ); + + // Then + expect(downloadScanZipMock).not.toHaveBeenCalled(); + expect(useCloudUpgradeStore.getState().activeFeature).toBe( + PAID_PLAN_UPGRADE_FEATURE.REPORT_DOWNLOAD, + ); + }); + it("opens failed scan error details from the actions menu", async () => { // Given const user = userEvent.setup(); diff --git a/ui/components/scans/table/scan-jobs-row-actions.tsx b/ui/components/scans/table/scan-jobs-row-actions.tsx index fc8b67ff52..1de2fb6ef3 100644 --- a/ui/components/scans/table/scan-jobs-row-actions.tsx +++ b/ui/components/scans/table/scan-jobs-row-actions.tsx @@ -29,6 +29,7 @@ import { ActionDropdown, ActionDropdownItem, } from "@/components/shadcn/dropdown"; +import { useReportDownload } from "@/hooks/use-report-download"; import { buildPerScanComplianceHref } from "@/lib/compliance/compliance-tab-url"; import { downloadScanZip } from "@/lib/helper"; import { getScanScheduleCapability } from "@/lib/schedules"; @@ -54,14 +55,18 @@ interface ScanJobsRowActionsProps { * Schedule capability override. Only for Prowler Cloud. */ capability?: ScanScheduleCapability; + /** Prowler Cloud tenants without a paid plan cannot download reports. */ + subscriptionOnly?: boolean; } export function ScanJobsRowActions({ scan, tab, capability, + subscriptionOnly = false, }: ScanJobsRowActionsProps) { const router = useRouter(); + const runReportDownload = useReportDownload(subscriptionOnly); const canEditSchedule = (capability ?? getScanScheduleCapability(isCloud())) === SCAN_SCHEDULE_CAPABILITY.ADVANCED; @@ -78,6 +83,9 @@ export function ScanJobsRowActions({ const scanState = scan.attributes.state; const isCompleted = scanState === "completed"; const isFailed = scanState === "failed"; + // Prowler Cloud partial scans re-check a few resources: they compute no + // compliance and write no report files, so neither entry applies. + const isPartial = scan.attributes.is_partial === true; const taskId = scan.relationships.task.data?.id; // The findings page bounds the UTC day range with completed_at; without it the // range collapses to the start day and can miss later findings. @@ -205,16 +213,22 @@ export function ScanJobsRowActions({ onSelect={openFindings} disabled={!isCompleted || !hasCompletedAt} /> - } - label="View Compliance" - onSelect={openCompliance} - /> - } - label="Download Scan Reports" - onSelect={() => downloadScanZip(scan.id, toast)} - /> + {!isPartial && ( + } + label="View Compliance" + onSelect={openCompliance} + /> + )} + {!isPartial && ( + } + label="Download Scan Reports" + onSelect={() => + runReportDownload(() => downloadScanZip(scan.id, toast)) + } + /> + )} )} {isFailed && ( diff --git a/ui/components/scans/table/scan-jobs-table.tsx b/ui/components/scans/table/scan-jobs-table.tsx index 294ae6865c..4d0c5ccf8a 100644 --- a/ui/components/scans/table/scan-jobs-table.tsx +++ b/ui/components/scans/table/scan-jobs-table.tsx @@ -23,6 +23,7 @@ interface ScanJobsTableProps { tab: ScanJobsTab; hasFilters?: boolean; scanScheduleCapability?: ScanScheduleCapability; + subscriptionOnly?: boolean; } const REFRESHING_STATES = ["available", "executing"] as const; @@ -33,6 +34,7 @@ export function ScanJobsTable({ tab, hasFilters = false, scanScheduleCapability, + subscriptionOnly, }: ScanJobsTableProps) { const searchParams = useSearchParams(); const hasRefreshingScan = data.some((scan) => @@ -43,6 +45,7 @@ export function ScanJobsTable({ const columns = getScanJobsColumns({ tab, capability: scanScheduleCapability, + subscriptionOnly, }); const showEmptyState = data.length === 0 && !hasFilters; const selectedScanId = searchParams?.get("scanId"); diff --git a/ui/components/shadcn/discovery-callout/discovery-callout.test.tsx b/ui/components/shadcn/discovery-callout/discovery-callout.test.tsx index e8870c4f0b..0f6ab44169 100644 --- a/ui/components/shadcn/discovery-callout/discovery-callout.test.tsx +++ b/ui/components/shadcn/discovery-callout/discovery-callout.test.tsx @@ -55,6 +55,22 @@ describe("DiscoveryCallout", () => { expect(onDismiss).toHaveBeenCalledTimes(1); }); + it("stays open when focus moves elsewhere, e.g. to a product tour", () => { + // Given + const onDismiss = vi.fn(); + renderCallout(true, onDismiss); + const elsewhere = document.createElement("button"); + document.body.appendChild(elsewhere); + + // When: driver.js focuses its own popover as a tour starts. + fireEvent.focusIn(elsewhere); + + // Then: a passive hint is not dismissed by focus it never held. + expect(onDismiss).not.toHaveBeenCalled(); + expect(screen.getByRole("button", { name: "Got it" })).toBeInTheDocument(); + elsewhere.remove(); + }); + it("keeps focus free when it opens", () => { // Given / When: the callout opens on its own (not user-invoked) renderCallout(true, vi.fn()); diff --git a/ui/components/shadcn/discovery-callout/discovery-callout.tsx b/ui/components/shadcn/discovery-callout/discovery-callout.tsx index f2c2aad8b4..70a9d1db27 100644 --- a/ui/components/shadcn/discovery-callout/discovery-callout.tsx +++ b/ui/components/shadcn/discovery-callout/discovery-callout.tsx @@ -62,8 +62,14 @@ export function DiscoveryCalloutContent({ side={side} align={align} sideOffset={8} - // A discovery hint must never steal focus from what the user is doing. + // A hint is worth showing even mid-tour: above driver.js's overlay (z 10000) + // and still dismissible while the tour locks the rest of the page. + className="z-[10001]" + data-tour-interactive + // A discovery hint must never steal focus from what the user is doing, + // nor vanish because something else took it (a tour popover, a form). onOpenAutoFocus={(event) => event.preventDefault()} + onFocusOutside={(event) => event.preventDefault()} data-testid={testId} >
diff --git a/ui/components/shadcn/tree-view/tree-status-icon.tsx b/ui/components/shadcn/tree-view/tree-status-icon.tsx index 1ef34024b1..c79d6e9838 100644 --- a/ui/components/shadcn/tree-view/tree-status-icon.tsx +++ b/ui/components/shadcn/tree-view/tree-status-icon.tsx @@ -1,6 +1,6 @@ "use client"; -import { CircleCheckIcon, CircleXIcon } from "lucide-react"; +import { CircleCheckIcon, CircleXIcon, InfoIcon } from "lucide-react"; import { cn } from "@/lib/utils"; import { TREE_ITEM_STATUS, TreeItemStatus } from "@/types/tree"; @@ -11,11 +11,14 @@ interface TreeStatusIconProps { } /** - * TreeStatusIcon component - displays success or error status for tree nodes. + * TreeStatusIcon component - displays success, error, or pending status for + * tree nodes. * * Features: * - CircleCheck icon for success (green) * - CircleX icon for error (red) + * - Static Info icon for pending: an item nothing is polling any more but with + * no confirmed outcome, so a spinner would misrepresent it as in progress * - Same size as TreeSpinner for consistent layout */ export function TreeStatusIcon({ status, className }: TreeStatusIconProps) { @@ -37,5 +40,14 @@ export function TreeStatusIcon({ status, className }: TreeStatusIconProps) { ); } + if (status === TREE_ITEM_STATUS.PENDING) { + return ( + + ); + } + return null; } diff --git a/ui/components/shadcn/tree-view/tree-status-indicator.tsx b/ui/components/shadcn/tree-view/tree-status-indicator.tsx index 4a7adb06e4..3dee23b30f 100644 --- a/ui/components/shadcn/tree-view/tree-status-indicator.tsx +++ b/ui/components/shadcn/tree-view/tree-status-indicator.tsx @@ -5,7 +5,7 @@ import { TooltipContent, TooltipTrigger, } from "@/components/shadcn/tooltip"; -import { TreeItemStatus } from "@/types/tree"; +import { TREE_ITEM_STATUS, TreeItemStatus } from "@/types/tree"; import { TreeStatusIcon } from "./tree-status-icon"; @@ -22,7 +22,11 @@ export function TreeStatusIndicator({ return null; } - if (status === "error" && errorMessage) { + if ( + (status === TREE_ITEM_STATUS.ERROR || + status === TREE_ITEM_STATUS.PENDING) && + errorMessage + ) { return ( diff --git a/ui/components/shared/cloud-upgrade-modal.test.tsx b/ui/components/shared/cloud-upgrade-modal.test.tsx index 2ac1ec1c68..0687d9c564 100644 --- a/ui/components/shared/cloud-upgrade-modal.test.tsx +++ b/ui/components/shared/cloud-upgrade-modal.test.tsx @@ -3,7 +3,10 @@ import userEvent from "@testing-library/user-event"; import { afterEach, describe, expect, it, vi } from "vitest"; import { useCloudUpgradeStore } from "@/store/cloud-upgrade/store"; -import { CLOUD_UPGRADE_FEATURE } from "@/types/cloud-upgrade"; +import { + CLOUD_UPGRADE_FEATURE, + PAID_PLAN_UPGRADE_FEATURE, +} from "@/types/cloud-upgrade"; import { CloudUpgradeModal } from "./cloud-upgrade-modal"; @@ -11,6 +14,16 @@ const modalTestState = vi.hoisted(() => ({ keepContentMounted: false, })); +const authState = vi.hoisted(() => ({ + canManageBilling: true, +})); + +vi.mock("@/hooks/use-auth", () => ({ + useAuth: () => ({ + permissions: { manage_billing: authState.canManageBilling }, + }), +})); + vi.mock("@/components/shadcn/modal", async (importOriginal) => { const actual = await importOriginal(); @@ -41,6 +54,7 @@ describe("CloudUpgradeModal", () => { afterEach(() => { cleanup(); modalTestState.keepContentMounted = false; + authState.canManageBilling = true; vi.unstubAllEnvs(); useCloudUpgradeStore.getState().closeCloudUpgrade(); }); @@ -219,7 +233,7 @@ describe("CloudUpgradeModal", () => { }, ); - it("does not render upgrade UI in Prowler Cloud", () => { + it("does not render Local Server upgrades in Prowler Cloud", () => { // Given vi.stubEnv("UI_CLOUD_ENABLED", "true"); useCloudUpgradeStore @@ -232,4 +246,83 @@ describe("CloudUpgradeModal", () => { // Then expect(screen.queryByRole("dialog")).not.toBeInTheDocument(); }); + + it("does not render paid plan upgrades in Local Server", () => { + // Given + vi.stubEnv("UI_CLOUD_ENABLED", "false"); + useCloudUpgradeStore + .getState() + .openCloudUpgrade(PAID_PLAN_UPGRADE_FEATURE.REPORT_DOWNLOAD); + + // When + render(); + + // Then + expect(screen.queryByRole("dialog")).not.toBeInTheDocument(); + }); + + it("renders the report download upgrade in Prowler Cloud", async () => { + // Given + vi.stubEnv("UI_CLOUD_ENABLED", "true"); + useCloudUpgradeStore + .getState() + .openCloudUpgrade(PAID_PLAN_UPGRADE_FEATURE.REPORT_DOWNLOAD); + + // When + render(); + + // Then + expect( + await screen.findByRole("dialog", { name: "Download Your Scan Reports" }), + ).toBeVisible(); + expect(screen.getByText("Available on paid plans")).toBeVisible(); + expect( + screen.queryByText("Available in Prowler Cloud"), + ).not.toBeInTheDocument(); + const upgradeLink = screen.getByRole("link", { + name: "Upgrade to Download", + }); + expect(upgradeLink).toHaveAttribute( + "href", + "/billing?feature=report_download", + ); + expect(upgradeLink).not.toHaveAttribute("target"); + const pricingLink = screen.getByRole("link", { + name: "View Plans & Pricing", + }); + expect(pricingLink).toHaveAttribute( + "href", + "https://prowler.com/pricing?utm_source=prowler-cloud&utm_content=report-download", + ); + expect(pricingLink).toHaveAttribute("target", "_blank"); + expect( + screen.queryByText(/Your Prowler Local Server remains unchanged/), + ).not.toBeInTheDocument(); + }); + + it("asks users without billing access to contact an admin", async () => { + // Given + vi.stubEnv("UI_CLOUD_ENABLED", "true"); + authState.canManageBilling = false; + useCloudUpgradeStore + .getState() + .openCloudUpgrade(PAID_PLAN_UPGRADE_FEATURE.REPORT_DOWNLOAD); + + // When + render(); + + // Then + expect( + await screen.findByRole("dialog", { name: "Download Your Scan Reports" }), + ).toBeVisible(); + expect( + screen.queryByRole("link", { name: "Upgrade to Download" }), + ).not.toBeInTheDocument(); + expect( + screen.getByText("Ask an account admin to upgrade your plan."), + ).toBeVisible(); + expect( + screen.getByRole("link", { name: "View Plans & Pricing" }), + ).toBeVisible(); + }); }); diff --git a/ui/components/shared/cloud-upgrade-modal.tsx b/ui/components/shared/cloud-upgrade-modal.tsx index 3c7080427b..df288d7d6c 100644 --- a/ui/components/shared/cloud-upgrade-modal.tsx +++ b/ui/components/shared/cloud-upgrade-modal.tsx @@ -5,18 +5,211 @@ import { Check, Cloud } from "lucide-react"; import { Badge } from "@/components/shadcn/badge/badge"; import { Button } from "@/components/shadcn/button/button"; import { Modal } from "@/components/shadcn/modal"; +import { useAuth } from "@/hooks/use-auth"; import { CLOUD_UPGRADE_CONTENT, CLOUD_UPGRADE_FOOTER_NOTE, CLOUD_UPGRADE_SECONDARY_CTA, getCloudUpgradeCompareUrl, getCloudUpgradePrimaryUrl, + getPaidPlanUpgradeBillingHref, + getPaidPlanUpgradeCompareUrl, + isCloudUpgradeFeature, + isPaidPlanUpgradeFeature, + PAID_PLAN_UPGRADE_ADMIN_NOTE, + PAID_PLAN_UPGRADE_BADGE, + PAID_PLAN_UPGRADE_CONTENT, } from "@/lib/cloud-upgrade"; import { isCloud } from "@/lib/shared/env"; import { useCloudUpgradeStore } from "@/store"; +import type { + CloudUpgradeFeature, + PaidPlanUpgradeFeature, +} from "@/types/cloud-upgrade"; const allowInitialAutoFocus = () => {}; +const CTA_CLASS_NAME = + "h-auto min-h-9 w-full min-w-0 shrink whitespace-normal md:flex-1"; + +interface UpgradeModalCta { + label: string; + href: string; + opensInNewTab: boolean; +} + +interface UpgradeModalVariantProps { + open: boolean; + onClose: () => void; + returnFocusElement: HTMLElement | null; +} + +interface UpgradeModalLayoutProps extends UpgradeModalVariantProps { + title: string; + description: string; + badge: string; + benefits: readonly string[]; + primaryCta?: UpgradeModalCta; + secondaryCta: UpgradeModalCta; + footerNote?: string; +} + +interface UpgradeModalLinkProps { + cta: UpgradeModalCta; + isSecondary?: boolean; +} + +const UpgradeModalLink = ({ cta, isSecondary }: UpgradeModalLinkProps) => ( + +); + +const UpgradeModalLayout = ({ + open, + onClose, + returnFocusElement, + title, + description, + badge, + benefits, + primaryCta, + secondaryCta, + footerNote, +}: UpgradeModalLayoutProps) => ( + !nextOpen && onClose()} + onOpenAutoFocus={allowInitialAutoFocus} + onCloseAutoFocus={(event) => { + event.preventDefault(); + returnFocusElement?.focus(); + }} + title={title} + description={description} + size="2xl" + > +
+
+
+
+ {badge} +
+ +
    + {benefits.map((benefit) => ( +
  • +
  • + ))} +
+ +
+ {primaryCta && } + +
+ + {footerNote && ( +

+ {footerNote} +

+ )} +
+
+); + +interface LocalServerUpgradeModalProps extends UpgradeModalVariantProps { + feature: CloudUpgradeFeature; +} + +const LocalServerUpgradeModal = ({ + feature, + ...modalProps +}: LocalServerUpgradeModalProps) => { + const content = CLOUD_UPGRADE_CONTENT[feature]; + + return ( + + ); +}; + +interface PaidPlanUpgradeModalProps extends UpgradeModalVariantProps { + feature: PaidPlanUpgradeFeature; +} + +const PaidPlanUpgradeModal = ({ + feature, + ...modalProps +}: PaidPlanUpgradeModalProps) => { + const { permissions } = useAuth(); + const content = PAID_PLAN_UPGRADE_CONTENT[feature]; + // The /billing route redirects users without billing access to /profile. + const canManageBilling = permissions.manage_billing === true; + + return ( + + ); +}; + export const CloudUpgradeModal = () => { const activeFeature = useCloudUpgradeStore((state) => state.activeFeature); const retainedFeature = useCloudUpgradeStore( @@ -29,81 +222,21 @@ export const CloudUpgradeModal = () => { (state) => state.returnFocusElement, ); - if (isCloud()) return null; - const feature = activeFeature ?? retainedFeature; - const content = CLOUD_UPGRADE_CONTENT[feature]; + const modalProps = { + open: activeFeature !== null, + onClose: closeCloudUpgrade, + returnFocusElement, + }; - return ( - !open && closeCloudUpgrade()} - onOpenAutoFocus={allowInitialAutoFocus} - onCloseAutoFocus={(event) => { - event.preventDefault(); - returnFocusElement?.focus(); - }} - title={content.title} - description={content.description} - size="2xl" - > -
-
-
-
- Available in Prowler Cloud -
+ // Cloud only upsells paid plans; Local Server only upsells Prowler Cloud. + if (isCloud()) { + return isPaidPlanUpgradeFeature(feature) ? ( + + ) : null; + } -
    - {content.benefits.map((benefit) => ( -
  • -
  • - ))} -
- - - -

- {CLOUD_UPGRADE_FOOTER_NOTE} -

-
-
- ); + return isCloudUpgradeFeature(feature) ? ( + + ) : null; }; diff --git a/ui/components/side-panel/side-panel-trigger.test.tsx b/ui/components/side-panel/side-panel-trigger.test.tsx index e5435a46e3..309e0cfc52 100644 --- a/ui/components/side-panel/side-panel-trigger.test.tsx +++ b/ui/components/side-panel/side-panel-trigger.test.tsx @@ -29,6 +29,7 @@ describe("SidePanelTrigger discovery callout", () => { }); afterEach(() => { + document.body.classList.remove("driver-active"); vi.useRealTimers(); }); @@ -49,6 +50,20 @@ describe("SidePanelTrigger discovery callout", () => { ).toBeInTheDocument(); }); + it("stays usable above a running product tour", () => { + // Given: driver.js dims and disables everything outside its spotlight. + document.body.classList.add("driver-active"); + render(); + + // When + act(() => vi.advanceTimersByTime(HINT_DELAY_MS)); + + // Then: the callout opts out of that, so it is visible and dismissible. + expect(screen.getByTestId("side-panel-ai-hint")).toHaveAttribute( + "data-tour-interactive", + ); + }); + it("never surfaces the callout again once seen", () => { // Given: a returning user useSidePanelStore.setState({ hasSeenAiTriggerHint: true }); diff --git a/ui/hooks/use-partial-scan-target.ts b/ui/hooks/use-partial-scan-target.ts new file mode 100644 index 0000000000..ba36f4659b --- /dev/null +++ b/ui/hooks/use-partial-scan-target.ts @@ -0,0 +1,23 @@ +"use client"; + +import { useAuth } from "@/hooks/use-auth"; +import { + isPartialScanAvailable, + isPartialScanTarget, +} from "@/lib/partial-scans"; +import { isCloud } from "@/lib/shared/env"; +import type { PartialScanTarget } from "@/types/partial-scans"; + +/** The target to re-check, or null when the feature or the row cannot offer it. */ +export function usePartialScanTarget( + target: Partial | null | undefined, +): PartialScanTarget | null { + const { hasPermission } = useAuth(); + + const isAvailable = isPartialScanAvailable({ + cloudEnabled: isCloud(), + canManageScans: hasPermission("manage_scans"), + }); + + return isAvailable && isPartialScanTarget(target) ? target : null; +} diff --git a/ui/hooks/use-report-download.ts b/ui/hooks/use-report-download.ts new file mode 100644 index 0000000000..ce7bd43b1b --- /dev/null +++ b/ui/hooks/use-report-download.ts @@ -0,0 +1,18 @@ +import { useCloudUpgradeStore } from "@/store"; +import { PAID_PLAN_UPGRADE_FEATURE } from "@/types/cloud-upgrade"; + +/** Wraps report downloads so subscription-only tenants get the paid plan upgrade instead. */ +export const useReportDownload = (subscriptionOnly = false) => { + const openCloudUpgrade = useCloudUpgradeStore( + (state) => state.openCloudUpgrade, + ); + + return (download: () => void | Promise) => { + if (subscriptionOnly) { + openCloudUpgrade(PAID_PLAN_UPGRADE_FEATURE.REPORT_DOWNLOAD); + return; + } + + return download(); + }; +}; diff --git a/ui/lib/cloud-upgrade.ts b/ui/lib/cloud-upgrade.ts index 2dfc79f0d7..a3a6dc0e8c 100644 --- a/ui/lib/cloud-upgrade.ts +++ b/ui/lib/cloud-upgrade.ts @@ -1,6 +1,9 @@ import { CLOUD_UPGRADE_FEATURE, type CloudUpgradeFeature, + PAID_PLAN_UPGRADE_FEATURE, + type PaidPlanUpgradeFeature, + type UpgradeFeature, } from "@/types/cloud-upgrade"; import { MAX_SAML_ADDITIONAL_EMAIL_DOMAINS } from "@/types/saml"; @@ -200,3 +203,48 @@ export const getCloudUpgradePrimaryUrl = (feature: CloudUpgradeFeature) => export const getCloudUpgradeCompareUrl = (feature: CloudUpgradeFeature) => buildCloudUpgradeUrl(PRICING_URL, feature); + +export const isCloudUpgradeFeature = ( + feature: UpgradeFeature, +): feature is CloudUpgradeFeature => feature in CLOUD_UPGRADE_CONTENT; + +export const PAID_PLAN_UPGRADE_BADGE = "Available on paid plans"; +export const PAID_PLAN_UPGRADE_ADMIN_NOTE = + "Ask an account admin to upgrade your plan."; + +const CLOUD_UTM_SOURCE = "prowler-cloud"; + +const PAID_PLAN_UPGRADE_UTM_CONTENT = { + [PAID_PLAN_UPGRADE_FEATURE.REPORT_DOWNLOAD]: "report-download", +} as const satisfies Record; + +export const PAID_PLAN_UPGRADE_CONTENT = { + [PAID_PLAN_UPGRADE_FEATURE.REPORT_DOWNLOAD]: { + title: "Download Your Scan Reports", + description: "Report downloads are included in Prowler Cloud paid plans.", + benefits: [ + "Download the full scan output in CSV, JSON-OCSF, and HTML", + "Export compliance reports as CSV, OCSF, and PDF", + "Share evidence with auditors and your team", + ], + primaryCta: "Upgrade to Download", + }, +} as const satisfies Record; + +export const isPaidPlanUpgradeFeature = ( + feature: UpgradeFeature, +): feature is PaidPlanUpgradeFeature => feature in PAID_PLAN_UPGRADE_CONTENT; + +export const getPaidPlanUpgradeBillingHref = ( + feature: PaidPlanUpgradeFeature, +) => `/billing?${new URLSearchParams({ feature })}`; + +export const getPaidPlanUpgradeCompareUrl = ( + feature: PaidPlanUpgradeFeature, +) => { + const url = new URL(PRICING_URL); + url.searchParams.set("utm_source", CLOUD_UTM_SOURCE); + url.searchParams.set("utm_content", PAID_PLAN_UPGRADE_UTM_CONTENT[feature]); + + return url.toString(); +}; diff --git a/ui/lib/helper.test.ts b/ui/lib/helper.test.ts index f59dda487b..5bc5682f66 100644 --- a/ui/lib/helper.test.ts +++ b/ui/lib/helper.test.ts @@ -1,9 +1,11 @@ import { afterEach, describe, expect, it, vi } from "vitest"; import { + checkTaskStatus, downloadScanZip, getErrorMessage, permissionFormFields, + TASK_STATUS_MAX_RETRIES_ERROR, } from "./helper"; vi.mock("@/actions/scans", () => ({ @@ -11,8 +13,9 @@ vi.mock("@/actions/scans", () => ({ getCompliancePdfReport: vi.fn(), })); +const { getTask } = vi.hoisted(() => ({ getTask: vi.fn() })); vi.mock("@/actions/task", () => ({ - getTask: vi.fn(), + getTask, })); vi.mock("@/auth.config", () => ({ @@ -140,6 +143,42 @@ describe("getErrorMessage", () => { }); }); +describe("checkTaskStatus", () => { + afterEach(() => { + vi.restoreAllMocks(); + }); + + it("keeps polling past a caller's default retry window and reports success once the task completes", async () => { + let calls = 0; + getTask.mockImplementation(async () => { + calls += 1; + if (calls < 25) { + return { data: { attributes: { state: "executing" } } }; + } + return { data: { attributes: { state: "completed" } } }; + }); + + // 25 retries exceeds the generic 20-retry default, simulating a task that + // outlives the caller's usual wait. + const result = await checkTaskStatus("task-id", 40, 1); + + expect(result.completed).toBe(true); + expect(calls).toBe(25); + }); + + it("reports the exhausted-retries error once maxRetries is used up", async () => { + getTask.mockResolvedValue({ data: { attributes: { state: "executing" } } }); + + const result = await checkTaskStatus("task-id", 3, 1); + + expect(result).toEqual({ + completed: false, + error: TASK_STATUS_MAX_RETRIES_ERROR, + }); + expect(getTask).toHaveBeenCalledTimes(3); + }); +}); + describe("permissionFormFields", () => { it("describes Unlimited Visibility as organization-wide", () => { // Given diff --git a/ui/lib/helper.ts b/ui/lib/helper.ts index d62ddaf65a..5612cbdd99 100644 --- a/ui/lib/helper.ts +++ b/ui/lib/helper.ts @@ -351,6 +351,10 @@ export const isGithubOAuthEnabled = !!process.env.SOCIAL_GITHUB_OAUTH_CLIENT_ID && !!process.env.SOCIAL_GITHUB_OAUTH_CLIENT_SECRET; +/** Returned by {@link checkTaskStatus} when `maxRetries` is exhausted, so callers + * can tell an exhausted wait apart from a real task failure. */ +export const TASK_STATUS_MAX_RETRIES_ERROR = "Max retries exceeded"; + /** * Polls a task until it settles. The settled task comes back with the verdict so * callers can read its result without fetching the same task again. @@ -390,7 +394,7 @@ export const checkTaskStatus = async ( } } - return { completed: false, error: "Max retries exceeded" }; + return { completed: false, error: TASK_STATUS_MAX_RETRIES_ERROR }; }; export const wait = (ms: number) => diff --git a/ui/lib/onboarding/README.md b/ui/lib/onboarding/README.md index 0b76ec6ced..dd5f7f51a9 100644 --- a/ui/lib/onboarding/README.md +++ b/ui/lib/onboarding/README.md @@ -19,11 +19,33 @@ posts to the API. | Per-route trigger | `ui/components/onboarding/onboarding-trigger.tsx` | | Ephemeral sequence slice | `ui/store/onboarding-sequence.ts` | | Checkpoint watcher + dialog | `ui/components/onboarding/onboarding-checkpoint-{watcher,dialog}.tsx` | -| Mandatory new-user gate | `ui/components/onboarding/onboarding-gate.tsx` | +| New-tenant gate (first-run redirect) | `ui/components/onboarding/onboarding-gate.tsx` | +| First-run marker (once per tenant) | `ui/lib/onboarding/first-run-marker.ts` | | Step outcome events (window) | `ui/lib/onboarding/onboarding-events.ts` | | Invite step before the checkpoint | `ui/components/onboarding/onboarding-invite-{step,dialog}.tsx` | | Manual replay list | `ui/components/ui/user-nav/user-nav.tsx` | +## First run + +The gate is mounted in every deployment. When the tenant provably has no +providers (`hasProviders === false`), the user holds `manage_providers` and +neither the first-run marker (`prowler.onboarding.first-run.`, so a +first run in one tenant never silences it for another on the same browser; the +bare `prowler.onboarding.first-run` key is a browser-wide opt-out, which is what +the e2e storage state sets) nor an add-provider completion record exists, it +replaces the route once with +`/providers?addProvider=true&addProviderSource=first_run`, so the add-provider +wizard is already open. Billing routes defer it; an unknown provider count or a +user without the permission (an empty list may only mean limited visibility) +never triggers it. + +In Cloud the URL also carries `&onboarding=add-provider` and the checkpoint is +armed. Because the wizard is already open, the providers page passes +`startAtTarget="provider-type"` to its ``, which skips the +tour's welcome and "open the wizard" steps. A navbar replay with the wizard +closed still starts from the first step. Self-hosted deployments get the +redirect only: tours and the checkpoint stay Cloud-only. + ## How the guided sequence works 1. The `(prowler)/layout.tsx` derives a tri-state `hasProviders` on every diff --git a/ui/lib/onboarding/__tests__/gate-decision.test.ts b/ui/lib/onboarding/__tests__/gate-decision.test.ts index 833987be7e..dbe551c2dc 100644 --- a/ui/lib/onboarding/__tests__/gate-decision.test.ts +++ b/ui/lib/onboarding/__tests__/gate-decision.test.ts @@ -18,14 +18,26 @@ describe("shouldStartOnboarding", () => { it("returns true for a zero-provider user with no completion record", () => { const result = shouldStartOnboarding({ hasProviders: false, + canManageProviders: true, completionRecord: null, }); expect(result).toBe(true); }); + it("returns false when the user cannot add providers, even in an empty tenant", () => { + // Limited-visibility users see zero providers without the tenant being empty. + const result = shouldStartOnboarding({ + hasProviders: false, + canManageProviders: false, + completionRecord: null, + }); + expect(result).toBe(false); + }); + it("returns false when the user already has providers", () => { const result = shouldStartOnboarding({ hasProviders: true, + canManageProviders: true, completionRecord: null, }); expect(result).toBe(false); @@ -34,6 +46,7 @@ describe("shouldStartOnboarding", () => { it("returns false when a dismissed record exists", () => { const result = shouldStartOnboarding({ hasProviders: false, + canManageProviders: true, completionRecord: recordWithState(TOUR_COMPLETION_STATES.DISMISSED), }); expect(result).toBe(false); @@ -42,6 +55,7 @@ describe("shouldStartOnboarding", () => { it("returns false when a completed record exists", () => { const result = shouldStartOnboarding({ hasProviders: false, + canManageProviders: true, completionRecord: recordWithState(TOUR_COMPLETION_STATES.COMPLETED), }); expect(result).toBe(false); @@ -50,6 +64,7 @@ describe("shouldStartOnboarding", () => { it("returns false when a skipped record exists", () => { const result = shouldStartOnboarding({ hasProviders: false, + canManageProviders: true, completionRecord: recordWithState(TOUR_COMPLETION_STATES.SKIPPED), }); expect(result).toBe(false); @@ -59,6 +74,7 @@ describe("shouldStartOnboarding", () => { // strict === false check rejects non-false values; don't force onboarding on unknown state const result = shouldStartOnboarding({ hasProviders: undefined, + canManageProviders: true, completionRecord: null, }); expect(result).toBe(false); @@ -67,6 +83,7 @@ describe("shouldStartOnboarding", () => { it("fails open when hasProviders is null", () => { const result = shouldStartOnboarding({ hasProviders: null as unknown as boolean, + canManageProviders: true, completionRecord: null, }); expect(result).toBe(false); diff --git a/ui/lib/onboarding/first-run-marker.ts b/ui/lib/onboarding/first-run-marker.ts new file mode 100644 index 0000000000..12e4a52ec9 --- /dev/null +++ b/ui/lib/onboarding/first-run-marker.ts @@ -0,0 +1,44 @@ +// Durable "this browser already went through the first-run redirect" memory. +// Self-hosted deployments run no tour, so no completion record would ever be +// written there; without this marker an empty tenant would be redirected on +// every page load. +// +// Scoped per tenant, like the other onboarding markers: going through the first +// run in one tenant must not silence it for another one on the same browser. +// The bare key is a browser-wide opt-out: written before markers were scoped, +// by e2e storage state, or when no usable tenant id exists. +const FIRST_RUN_MARKER_KEY = "prowler.onboarding.first-run"; + +// Tenant ids are UUIDs; anything else is refused rather than concatenated +// into a storage key. +const TENANT_ID_PATTERN = + /^[0-9a-f]{8}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{12}$/i; + +export function firstRunMarkerKey(tenantId?: string | null): string { + if (!tenantId || !TENANT_ID_PATTERN.test(tenantId)) { + return FIRST_RUN_MARKER_KEY; + } + return `${FIRST_RUN_MARKER_KEY}.${tenantId.toLowerCase()}`; +} + +export function isFirstRunHandled(tenantId?: string | null): boolean { + if (typeof window === "undefined") return true; + try { + return ( + window.localStorage.getItem(FIRST_RUN_MARKER_KEY) !== null || + window.localStorage.getItem(firstRunMarkerKey(tenantId)) !== null + ); + } catch { + // Unreadable storage must not redirect forever: treat as handled. + return true; + } +} + +export function markFirstRunHandled(tenantId?: string | null): void { + if (typeof window === "undefined") return; + try { + window.localStorage.setItem(firstRunMarkerKey(tenantId), "true"); + } catch { + // Non-fatal: a repeated redirect beats a thrown render. + } +} diff --git a/ui/lib/onboarding/gate-decision.ts b/ui/lib/onboarding/gate-decision.ts index 656a3b0cb4..39f7437a69 100644 --- a/ui/lib/onboarding/gate-decision.ts +++ b/ui/lib/onboarding/gate-decision.ts @@ -3,15 +3,19 @@ import type { TourCompletionRecord } from "@/lib/tours/tour-types"; export interface GateDecisionInput { // `undefined` allowed; strict `=== false` check below fails open on ambiguous signals. hasProviders: boolean | undefined; + // Limited-visibility users list zero providers in a tenant that is not empty. + canManageProviders: boolean; completionRecord: TourCompletionRecord | null; } -// Only forces onboarding when providers are provably absent and no record exists. +// Only forces onboarding when providers are provably absent, the user can add one +// and no record exists. export function shouldStartOnboarding({ hasProviders, + canManageProviders, completionRecord, }: GateDecisionInput): boolean { const hasNoRecord = completionRecord === null || completionRecord === undefined; - return hasProviders === false && hasNoRecord; + return hasProviders === false && canManageProviders && hasNoRecord; } diff --git a/ui/lib/partial-scans.test.ts b/ui/lib/partial-scans.test.ts new file mode 100644 index 0000000000..aa0b8961eb --- /dev/null +++ b/ui/lib/partial-scans.test.ts @@ -0,0 +1,111 @@ +import { describe, expect, it } from "vitest"; + +import { + findProviderIdForTarget, + getPartialScanErrorMessage, + isPartialScanAvailable, + isPartialScanTarget, + PARTIAL_SCAN_LAUNCH_ERROR, +} from "./partial-scans"; + +describe("isPartialScanAvailable", () => { + it("needs both Prowler Cloud and the manage_scans permission", () => { + expect( + isPartialScanAvailable({ cloudEnabled: true, canManageScans: true }), + ).toBe(true); + expect( + isPartialScanAvailable({ cloudEnabled: false, canManageScans: true }), + ).toBe(false); + expect( + isPartialScanAvailable({ cloudEnabled: true, canManageScans: false }), + ).toBe(false); + }); +}); + +describe("isPartialScanTarget", () => { + it("accepts a resource uid with a provider id", () => { + expect( + isPartialScanTarget({ + providerId: "provider-1", + resourceUid: "arn:aws:s3:::bucket", + resourceName: "bucket", + }), + ).toBe(true); + }); + + it("accepts a resource uid with a provider uid and type", () => { + expect( + isPartialScanTarget({ + providerUid: "123456789012", + providerType: "aws", + resourceUid: "arn:aws:s3:::bucket", + resourceName: "bucket", + }), + ).toBe(true); + }); + + it("rejects placeholder or missing identifiers", () => { + // Adapters fill unknown values with "-", which is not a real uid. + expect( + isPartialScanTarget({ + providerId: "provider-1", + resourceUid: "-", + resourceName: "bucket", + }), + ).toBe(false); + expect( + isPartialScanTarget({ + providerUid: "", + providerType: "aws", + resourceUid: "arn:aws:s3:::bucket", + }), + ).toBe(false); + expect(isPartialScanTarget(null)).toBe(false); + }); +}); + +describe("findProviderIdForTarget", () => { + const providers = [ + { id: "aws-1", attributes: { uid: "123456789012", provider: "aws" } }, + { id: "gcp-1", attributes: { uid: "123456789012", provider: "gcp" } }, + ] as never[]; + + it("matches on uid and provider type together", () => { + expect( + findProviderIdForTarget(providers, { + providerUid: "123456789012", + providerType: "gcp", + }), + ).toBe("gcp-1"); + }); + + it("returns undefined when nothing matches", () => { + expect( + findProviderIdForTarget(providers, { + providerUid: "000000000000", + providerType: "aws", + }), + ).toBeUndefined(); + }); +}); + +describe("getPartialScanErrorMessage", () => { + it("returns null for a created scan", () => { + expect(getPartialScanErrorMessage({ data: { id: "scan-1" } })).toBeNull(); + }); + + it("surfaces the API detail for a refused re-check", () => { + expect( + getPartialScanErrorMessage({ + error: "A scan is already running for this provider.", + status: 409, + }), + ).toBe("A scan is already running for this provider."); + }); + + it("falls back to a generic message when the error carries no detail", () => { + expect(getPartialScanErrorMessage({ status: 500 })).toBe( + PARTIAL_SCAN_LAUNCH_ERROR, + ); + }); +}); diff --git a/ui/lib/partial-scans.ts b/ui/lib/partial-scans.ts new file mode 100644 index 0000000000..b1c5f67f20 --- /dev/null +++ b/ui/lib/partial-scans.ts @@ -0,0 +1,51 @@ +import { + type ActionErrorResult, + getActionErrorMessage, + hasActionError, +} from "@/lib/action-errors"; +import type { PartialScanTarget } from "@/types/partial-scans"; +import type { ProviderProps } from "@/types/providers"; + +export const PARTIAL_SCAN_LAUNCH_ERROR = + "The re-check could not be launched. Please try again."; + +interface PartialScanAvailability { + cloudEnabled: boolean; + canManageScans: boolean; +} + +/** Partial scans are a Cloud feature that needs the manage_scans permission. */ +export const isPartialScanAvailable = ({ + cloudEnabled, + canManageScans, +}: PartialScanAvailability): boolean => cloudEnabled && canManageScans; + +const isMeaningful = (value: string | undefined): value is string => + typeof value === "string" && value.trim() !== "" && value !== "-"; + +/** A target needs a real resource uid and a way to identify its provider. */ +export const isPartialScanTarget = ( + target: Partial | null | undefined, +): target is PartialScanTarget => { + if (!target || !isMeaningful(target.resourceUid)) return false; + if (isMeaningful(target.providerId)) return true; + return isMeaningful(target.providerUid) && isMeaningful(target.providerType); +}; + +/** Provider uid is unique per provider type, so both together pick one row. */ +export const findProviderIdForTarget = ( + providers: Pick[], + target: Pick, +): string | undefined => + providers.find( + (provider) => + provider.attributes.uid === target.providerUid && + provider.attributes.provider === target.providerType, + )?.id; + +export const getPartialScanErrorMessage = ( + result: (ActionErrorResult & { data?: unknown }) | null | undefined, +): string | null => + hasActionError(result) + ? getActionErrorMessage(result, { fallback: PARTIAL_SCAN_LAUNCH_ERROR }) + : null; diff --git a/ui/lib/provider-funnel/provider-funnel-events.test.ts b/ui/lib/provider-funnel/provider-funnel-events.test.ts new file mode 100644 index 0000000000..bd85ba01c6 --- /dev/null +++ b/ui/lib/provider-funnel/provider-funnel-events.test.ts @@ -0,0 +1,56 @@ +import { afterEach, describe, expect, it, vi } from "vitest"; + +import { + dispatchProviderFunnel, + PROVIDER_FUNNEL_EVENT, + PROVIDER_FUNNEL_STEP, + type ProviderFunnelDetail, + WIZARD_OPEN_SOURCE, +} from "./provider-funnel-events"; + +describe("dispatchProviderFunnel", () => { + const listeners: EventListener[] = []; + + const listen = (listener: (detail: ProviderFunnelDetail) => void) => { + const handler: EventListener = (event) => + listener((event as CustomEvent).detail); + listeners.push(handler); + window.addEventListener(PROVIDER_FUNNEL_EVENT, handler); + }; + + afterEach(() => { + listeners + .splice(0) + .forEach((handler) => + window.removeEventListener(PROVIDER_FUNNEL_EVENT, handler), + ); + }); + + it("delivers the step detail to an outside window listener", () => { + const received = vi.fn(); + listen(received); + + dispatchProviderFunnel({ + step: PROVIDER_FUNNEL_STEP.WIZARD_OPENED, + source: WIZARD_OPEN_SOURCE.SIDEBAR_CTA, + }); + + expect(received).toHaveBeenCalledExactlyOnceWith({ + step: "wizard_opened", + source: "sidebar_cta", + }); + }); + + it("is a no-op during server rendering, where there is no window", () => { + vi.stubGlobal("window", undefined); + + expect(() => + dispatchProviderFunnel({ + step: PROVIDER_FUNNEL_STEP.PROVIDER_TYPE_SELECTED, + providerType: "aws", + }), + ).not.toThrow(); + + vi.unstubAllGlobals(); + }); +}); diff --git a/ui/lib/provider-funnel/provider-funnel-events.ts b/ui/lib/provider-funnel/provider-funnel-events.ts new file mode 100644 index 0000000000..2d485b9536 --- /dev/null +++ b/ui/lib/provider-funnel/provider-funnel-events.ts @@ -0,0 +1,125 @@ +// Window events the add-provider journey dispatches as the user moves through +// it. They carry no listener of their own: a deployment that wants to observe +// the funnel (product analytics, for instance) subscribes from outside, so the +// UI stays free of any tracking dependency. Details are low-cardinality only: +// never a provider uid, alias, ARN or anything typed into a credentials form. +export const PROVIDER_FUNNEL_EVENT = "prowler:provider-funnel"; + +export const PROVIDER_FUNNEL_STEP = { + SIDEBAR_CTA_CLICKED: "sidebar_cta_clicked", + WIZARD_OPENED: "wizard_opened", + PROVIDER_TYPE_SELECTED: "provider_type_selected", + METHOD_SELECTED: "method_selected", + ROLE_TEMPLATE_OPENED: "role_template_opened", + ACCOUNT_SUBMITTED: "account_submitted", + WIZARD_CLOSED: "wizard_closed", +} as const; + +export type ProviderFunnelStep = + (typeof PROVIDER_FUNNEL_STEP)[keyof typeof PROVIDER_FUNNEL_STEP]; + +export const SIDEBAR_CTA_VARIANT = { + ADD_PROVIDER: "add_provider", + LAUNCH_SCAN: "launch_scan", +} as const; + +export type SidebarCtaVariant = + (typeof SIDEBAR_CTA_VARIANT)[keyof typeof SIDEBAR_CTA_VARIANT]; + +export const WIZARD_OPEN_SOURCE = { + FIRST_RUN: "first_run", + SIDEBAR_CTA: "sidebar_cta", + URL: "url", + PAGE_BUTTON: "page_button", + EMPTY_STATE: "empty_state", + ROW_ACTION: "row_action", +} as const; + +export type WizardOpenSource = + (typeof WIZARD_OPEN_SOURCE)[keyof typeof WIZARD_OPEN_SOURCE]; + +export const PROVIDER_FUNNEL_METHOD = { + SINGLE: "single", + ORGANIZATION: "organization", +} as const; + +export type ProviderFunnelMethod = + (typeof PROVIDER_FUNNEL_METHOD)[keyof typeof PROVIDER_FUNNEL_METHOD]; + +export const ROLE_TEMPLATE_KIND = { + CLOUDFORMATION_QUICK_CREATE: "cloudformation_quick_create", + CLOUDFORMATION_TEMPLATE: "cloudformation_template", + TERRAFORM: "terraform", +} as const; + +export type RoleTemplateKind = + (typeof ROLE_TEMPLATE_KIND)[keyof typeof ROLE_TEMPLATE_KIND]; + +export interface SidebarCtaClickedDetail { + step: typeof PROVIDER_FUNNEL_STEP.SIDEBAR_CTA_CLICKED; + variant: SidebarCtaVariant; +} + +export interface WizardOpenedDetail { + step: typeof PROVIDER_FUNNEL_STEP.WIZARD_OPENED; + source: WizardOpenSource; +} + +export interface ProviderTypeSelectedDetail { + step: typeof PROVIDER_FUNNEL_STEP.PROVIDER_TYPE_SELECTED; + providerType: string; +} + +export interface MethodSelectedDetail { + step: typeof PROVIDER_FUNNEL_STEP.METHOD_SELECTED; + providerType: string; + method: ProviderFunnelMethod; +} + +export interface RoleTemplateOpenedDetail { + step: typeof PROVIDER_FUNNEL_STEP.ROLE_TEMPLATE_OPENED; + template: RoleTemplateKind; +} + +export const ACCOUNT_SUBMIT_OUTCOME = { + SUCCESS: "success", + ERROR: "error", +} as const; + +export type AccountSubmitOutcome = + (typeof ACCOUNT_SUBMIT_OUTCOME)[keyof typeof ACCOUNT_SUBMIT_OUTCOME]; + +// Account and credentials sent together (the one-step AWS form). +export interface AccountSubmittedDetail { + step: typeof PROVIDER_FUNNEL_STEP.ACCOUNT_SUBMITTED; + providerType: string; + via: string; + outcome: AccountSubmitOutcome; +} + +export interface WizardClosedDetail { + step: typeof PROVIDER_FUNNEL_STEP.WIZARD_CLOSED; + lastStep: string; + // A provider record exists; it may still lack credentials or a connection. + providerCreated: boolean; +} + +export type ProviderFunnelDetail = + | SidebarCtaClickedDetail + | WizardOpenedDetail + | ProviderTypeSelectedDetail + | MethodSelectedDetail + | RoleTemplateOpenedDetail + | AccountSubmittedDetail + | WizardClosedDetail; + +export function dispatchProviderFunnel(detail: ProviderFunnelDetail): void { + if (typeof window === "undefined") return; + try { + window.dispatchEvent( + new CustomEvent(PROVIDER_FUNNEL_EVENT, { detail }), + ); + } catch { + // A listener that throws must never break the journey it observes. + } +} diff --git a/ui/lib/provider-helpers.test.ts b/ui/lib/provider-helpers.test.ts index 8efba9ac80..0a3a4ccc95 100644 --- a/ui/lib/provider-helpers.test.ts +++ b/ui/lib/provider-helpers.test.ts @@ -1,16 +1,155 @@ import { beforeEach, describe, expect, it, vi } from "vitest"; -const { checkConnectionProvider, checkTaskStatus } = vi.hoisted(() => ({ - checkConnectionProvider: vi.fn(), - checkTaskStatus: vi.fn(), +const { checkConnectionProvider, checkTaskStatus, getProvider } = vi.hoisted( + () => ({ + checkConnectionProvider: vi.fn(), + checkTaskStatus: vi.fn(), + getProvider: vi.fn(), + }), +); +vi.mock("@/actions/providers/providers", () => ({ + checkConnectionProvider, + getProvider, +})); +vi.mock("./helper", () => ({ + checkTaskStatus, + TASK_STATUS_MAX_RETRIES_ERROR: "Max retries exceeded", })); -vi.mock("@/actions/providers/providers", () => ({ checkConnectionProvider })); -vi.mock("./helper", () => ({ checkTaskStatus })); -import { testProviderConnection } from "./provider-helpers"; +import { + CONNECTION_CHECK_STATUS, + PROVIDER_CONNECTION_CHECK_MAX_RETRIES, + PROVIDER_CONNECTION_CHECK_POLL_DELAY_MS, + resolveProviderConnectionState, + testProviderConnection, +} from "./provider-helpers"; + +describe("resolveProviderConnectionState", () => { + it("trusts a stored connected=true once last_checked_at differs from the baseline", async () => { + getProvider.mockResolvedValue({ + data: { + attributes: { + connection: { + connected: true, + last_checked_at: "2026-01-01T00:00:10.000Z", + }, + }, + }, + }); + + expect(await resolveProviderConnectionState("account", null)).toEqual({ + status: CONNECTION_CHECK_STATUS.SUCCESS, + error: null, + }); + }); + + it("treats a stored connected=true as pending while last_checked_at still matches the baseline", async () => { + // Given: the provider was connected from a previous check, and this check -- + // dispatched after that baseline was captured -- is still running. + const baseline = "2025-12-31T23:59:59.000Z"; + getProvider.mockResolvedValue({ + data: { + attributes: { + connection: { connected: true, last_checked_at: baseline }, + }, + }, + }); + + const result = await resolveProviderConnectionState("account", baseline); + + expect(result.status).toBe(CONNECTION_CHECK_STATUS.PENDING); + expect(result.error).toMatch(/still running/i); + }); + + it("reports success once last_checked_at changes, even when the new value sorts earlier than the baseline", async () => { + // Given: a deployment whose clocks are not synchronised -- the server writes + // a `last_checked_at` that, read as a plain string/date, sorts *before* the + // baseline captured moments earlier from the same server. A timestamp + // comparison would misread this as stale; a value comparison does not. + const baseline = "2026-06-01T00:00:00.000Z"; + getProvider.mockResolvedValue({ + data: { + attributes: { + connection: { + connected: true, + last_checked_at: "2020-01-01T00:00:00.000Z", + }, + }, + }, + }); + + const result = await resolveProviderConnectionState("account", baseline); + + expect(result).toEqual({ + status: CONNECTION_CHECK_STATUS.SUCCESS, + error: null, + }); + }); + + it("treats a missing last_checked_at as pending even when connected is true", async () => { + getProvider.mockResolvedValue({ + data: { + attributes: { + connection: { connected: true, last_checked_at: null }, + }, + }, + }); + + const result = await resolveProviderConnectionState("account", null); + + expect(result.status).toBe(CONNECTION_CHECK_STATUS.PENDING); + }); + + it("treats an unknown baseline as unresolved, never trusting a pre-existing stored result", async () => { + // Given: the pre-dispatch baseline read failed, so there is nothing to + // compare this result against. + getProvider.mockResolvedValue({ + data: { + attributes: { + connection: { + connected: true, + last_checked_at: "2026-01-01T00:00:10.000Z", + }, + }, + }, + }); + + const result = await resolveProviderConnectionState("account", undefined); + + expect(result.status).toBe(CONNECTION_CHECK_STATUS.PENDING); + }); + + it("reports a changed connected=false as a confirmed failure", async () => { + getProvider.mockResolvedValue({ + data: { + attributes: { + connection: { + connected: false, + last_checked_at: "2026-01-01T00:00:10.000Z", + }, + }, + }, + }); + + const result = await resolveProviderConnectionState("account", null); + + expect(result).toEqual({ + status: CONNECTION_CHECK_STATUS.FAILED, + error: expect.stringMatching(/test the connection again/i), + }); + }); +}); describe("provider connection confirmation", () => { beforeEach(() => { + // The baseline read (before dispatch) and the exhausted-wait fallback read + // both go through `getProvider`; default to "no prior check" unless a test + // overrides it. + getProvider.mockResolvedValue({ + data: { + attributes: { connection: { connected: null, last_checked_at: null } }, + }, + }); checkConnectionProvider.mockResolvedValue({ data: { id: "task" } }); }); it.each([undefined, {}, { connected: "true" }, { connected: false }])( @@ -20,7 +159,9 @@ describe("provider connection confirmation", () => { completed: true, task: { data: { attributes: { result } } }, }); - expect((await testProviderConnection("account")).connected).toBe(false); + expect((await testProviderConnection("account")).status).toBe( + CONNECTION_CHECK_STATUS.FAILED, + ); }, ); it("advances on an explicitly successful connection", async () => { @@ -29,8 +170,199 @@ describe("provider connection confirmation", () => { task: { data: { attributes: { result: { connected: true } } } }, }); expect(await testProviderConnection("account")).toEqual({ - connected: true, + status: CONNECTION_CHECK_STATUS.SUCCESS, error: null, }); }); + + it("reads the baseline before dispatching the check", async () => { + checkTaskStatus.mockResolvedValue({ + completed: true, + task: { data: { attributes: { result: { connected: true } } } }, + }); + + const callOrder: string[] = []; + getProvider.mockImplementation(async () => { + callOrder.push("getProvider"); + return { + data: { + attributes: { + connection: { connected: null, last_checked_at: null }, + }, + }, + }; + }); + checkConnectionProvider.mockImplementation(async () => { + callOrder.push("checkConnectionProvider"); + return { data: { id: "task" } }; + }); + + await testProviderConnection("account"); + + expect(callOrder).toEqual(["getProvider", "checkConnectionProvider"]); + }); + + it("sizes the wait past the backend's 120s provider-connection-check time limit", async () => { + checkTaskStatus.mockResolvedValue({ + completed: true, + task: { data: { attributes: { result: { connected: true } } } }, + }); + + await testProviderConnection("account"); + + expect(checkTaskStatus).toHaveBeenCalledWith( + "task", + PROVIDER_CONNECTION_CHECK_MAX_RETRIES, + PROVIDER_CONNECTION_CHECK_POLL_DELAY_MS, + ); + expect( + PROVIDER_CONNECTION_CHECK_MAX_RETRIES * + PROVIDER_CONNECTION_CHECK_POLL_DELAY_MS, + ).toBeGreaterThan(120_000); + }); + + describe("when the wait is exhausted", () => { + beforeEach(() => { + checkTaskStatus.mockResolvedValue({ + completed: false, + error: "Max retries exceeded", + }); + }); + + it("reports success when the stored state changed after the baseline was captured", async () => { + // Baseline (read before dispatch): no prior check. Fallback read: a fresh + // result landed while the UI was waiting. + getProvider + .mockResolvedValueOnce({ + data: { + attributes: { + connection: { connected: null, last_checked_at: null }, + }, + }, + }) + .mockResolvedValueOnce({ + data: { + attributes: { + connection: { + connected: true, + last_checked_at: "2999-01-01T00:00:00Z", + }, + }, + }, + }); + + expect(await testProviderConnection("account")).toEqual({ + status: CONNECTION_CHECK_STATUS.SUCCESS, + error: null, + }); + }); + + it("does not report success from a stale stored result predating this check", async () => { + // Given: the provider was already connected from a previous check -- + // captured as the baseline before this check was dispatched -- and the + // fallback read comes back with that exact same (unchanged) value, + // meaning the new check has not written a result yet. + const priorLastCheckedAt = "2025-06-01T00:00:00Z"; + getProvider.mockResolvedValue({ + data: { + attributes: { + connection: { + connected: true, + last_checked_at: priorLastCheckedAt, + }, + }, + }, + }); + + const result = await testProviderConnection("account"); + + expect(result.status).not.toBe(CONNECTION_CHECK_STATUS.SUCCESS); + expect(result.status).toBe(CONNECTION_CHECK_STATUS.PENDING); + }); + + it("reports pending, not connected, when the browser clock is behind the server and the stored result predates the check", async () => { + // Given: a deployment whose clocks are not synchronised -- the browser's + // clock is behind the server's. Under a clock-based comparison this could + // make an older stored result look newer than the check's start and be + // reported as connected. The baseline here is captured from the server's + // own prior value, so the comparison never depends on the browser's clock + // at all: an unchanged value is still unchanged. + const staleServerTimestamp = "2026-03-01T12:00:00Z"; + getProvider.mockResolvedValue({ + data: { + attributes: { + connection: { + connected: true, + last_checked_at: staleServerTimestamp, + }, + }, + }, + }); + + const result = await testProviderConnection("account"); + + expect(result.status).toBe(CONNECTION_CHECK_STATUS.PENDING); + expect(result.status).not.toBe(CONNECTION_CHECK_STATUS.SUCCESS); + }); + + it("reports the failure when the provider is confirmed not connected by this check", async () => { + getProvider + .mockResolvedValueOnce({ + data: { + attributes: { + connection: { connected: null, last_checked_at: null }, + }, + }, + }) + .mockResolvedValueOnce({ + data: { + attributes: { + connection: { + connected: false, + last_checked_at: "2999-01-01T00:00:00Z", + }, + }, + }, + }); + + const result = await testProviderConnection("account"); + + expect(result.status).toBe(CONNECTION_CHECK_STATUS.FAILED); + expect(result.error).toMatch(/test the connection again/i); + }); + + it("shows a neutral still-checking message when the provider state is undetermined", async () => { + getProvider.mockResolvedValue({ + data: { + attributes: { + connection: { connected: null, last_checked_at: null }, + }, + }, + }); + + const result = await testProviderConnection("account"); + + expect(result.status).toBe(CONNECTION_CHECK_STATUS.PENDING); + expect(result.error).toMatch(/still running|refresh/i); + expect(result.error?.toLowerCase()).not.toContain("error"); + expect(result.error?.toLowerCase()).not.toContain("failed"); + }); + }); + + it("does not re-read provider state for a real task failure", async () => { + checkTaskStatus.mockResolvedValue({ + completed: false, + error: "Unexpected task state", + }); + + const result = await testProviderConnection("account"); + + // The baseline read (before dispatch) still happens, but nothing reads + // provider state again to resolve this failure. + expect(getProvider).toHaveBeenCalledTimes(1); + expect(result).toEqual({ + status: CONNECTION_CHECK_STATUS.FAILED, + error: "Unexpected task state", + }); + }); }); diff --git a/ui/lib/provider-helpers.ts b/ui/lib/provider-helpers.ts index 329833f384..22e26edd29 100644 --- a/ui/lib/provider-helpers.ts +++ b/ui/lib/provider-helpers.ts @@ -1,12 +1,24 @@ -import { checkConnectionProvider } from "@/actions/providers/providers"; import { + checkConnectionProvider, + getProvider, +} from "@/actions/providers/providers"; +import { + CONNECTION_CHECK_STATUS, + type ConnectionCheckStatus, ProviderEntity, ProviderProps, ProvidersApiResponse, ProviderType, } from "@/types/providers"; -import { checkTaskStatus } from "./helper"; +import { checkTaskStatus, TASK_STATUS_MAX_RETRIES_ERROR } from "./helper"; + +// Re-exported so callers that only need the status enum (e.g. +// `org-account-selection.utils.ts` and its tests) can import it from +// `@/types/providers` directly, without pulling in this module's server-action +// dependencies. +export { CONNECTION_CHECK_STATUS }; +export type { ConnectionCheckStatus }; export const extractProviderUIDs = ( providersData: ProvidersApiResponse, @@ -172,10 +184,119 @@ export const requiresBackButton = (via?: string | null): boolean => { }; export interface TestConnectionResult { - connected: boolean; + status: ConnectionCheckStatus; error: string | null; } +/** + * The `provider-connection-check` Celery task has a 120s hard time limit + * (api/src/backend/config/celery.py `task_annotations`). Poll long enough to + * cover a full run plus queueing/network slack, instead of the generic 30s + * default, which cuts the wait off well before the backend gives up. + */ +export const PROVIDER_CONNECTION_CHECK_TASK_TIME_LIMIT_MS = 120_000; +const PROVIDER_CONNECTION_CHECK_POLL_BUFFER_MS = 30_000; +export const PROVIDER_CONNECTION_CHECK_POLL_DELAY_MS = 1_500; +export const PROVIDER_CONNECTION_CHECK_MAX_RETRIES = Math.ceil( + (PROVIDER_CONNECTION_CHECK_TASK_TIME_LIMIT_MS + + PROVIDER_CONNECTION_CHECK_POLL_BUFFER_MS) / + PROVIDER_CONNECTION_CHECK_POLL_DELAY_MS, +); + +const CONNECTION_NOT_CONFIRMED_MESSAGE = + "Connection was not confirmed. Test the connection again."; +const CONNECTION_STILL_RUNNING_MESSAGE = + "The connection test is still running. Refresh in a moment to see the result."; + +/** + * Reads a provider's current `connection.last_checked_at`, to be captured + * *before* a connection check is dispatched for it. Passed on to + * `resolveProviderConnectionState` as the value a later read is compared against + * -- see that function for why the comparison is by value, not by clock. + * + * `undefined` means the read failed (network error, provider not found) and is + * distinct from `null` ("no prior check exists"): `resolveProviderConnectionState` + * treats an unknown baseline as impossible to clear, never as an implicit change. + */ +export async function captureConnectionBaseline( + providerId: string, +): Promise { + const formData = new FormData(); + formData.append("id", providerId); + + const providerResponse = await getProvider(formData); + if (!providerResponse?.data) { + return undefined; + } + + return providerResponse.data.attributes?.connection?.last_checked_at ?? null; +} + +/** + * Re-reads a provider's persisted connection state from the API. Used when a + * connection-check wait is exhausted: the backend task may still be running (or + * may already have finished after the UI stopped waiting on it), so this reports + * whatever the provider record currently says instead of a flat error. + * + * The stored `connection` is the result of the *last check that finished*, not + * necessarily the one this call is following up on -- e.g. a provider was + * connected, the user changed its credentials, and the new check is still + * running past the wait. `baseline` -- the provider's `last_checked_at` captured + * (via {@link captureConnectionBaseline}) before this check was dispatched -- + * guards against reporting that stale result as current: the stored state is + * only trusted once `last_checked_at` has changed from it. + * + * This compares the two values directly rather than comparing timestamps, so it + * holds even in deployments whose clocks are not synchronised: `last_checked_at` + * is written by the server, but the previous check used a wait deadline taken + * from the browser's clock, and a browser clock that drifts from the server's + * could make an older result look newer than the check, or a finished check look + * like it is still pending. + */ +export async function resolveProviderConnectionState( + providerId: string, + baseline: string | null | undefined, +): Promise { + const formData = new FormData(); + formData.append("id", providerId); + + const providerResponse = await getProvider(formData); + const connection = providerResponse?.data?.attributes?.connection; + const lastCheckedAt = connection?.last_checked_at; + + // An unknown baseline (the pre-dispatch read failed) can never be cleared -- + // there is nothing to compare against, so the result cannot be trusted yet. + const isCurrent = + baseline !== undefined && !!lastCheckedAt && lastCheckedAt !== baseline; + + if (!isCurrent) { + // No persisted result yet, or it predates this check -- the backend task + // may still be running. + return { + status: CONNECTION_CHECK_STATUS.PENDING, + error: CONNECTION_STILL_RUNNING_MESSAGE, + }; + } + + if (connection?.connected === true) { + return { status: CONNECTION_CHECK_STATUS.SUCCESS, error: null }; + } + + if (connection?.connected === false) { + return { + status: CONNECTION_CHECK_STATUS.FAILED, + error: CONNECTION_NOT_CONFIRMED_MESSAGE, + }; + } + + // Current per its timestamp, but `connected` is null -- treat the same as + // still pending rather than guessing a pass or fail. + return { + status: CONNECTION_CHECK_STATUS.PENDING, + error: CONNECTION_STILL_RUNNING_MESSAGE, + }; +} + /** * Tests a provider's connection end-to-end: submits the task, polls until * completion, and returns the real connection result. @@ -186,6 +307,10 @@ export interface TestConnectionResult { export async function testProviderConnection( providerId: string, ): Promise { + // Captured before the check is dispatched, so any `last_checked_at` this check + // eventually writes is guaranteed to differ from it. + const baseline = await captureConnectionBaseline(providerId); + const formData = new FormData(); formData.append("providerId", providerId); @@ -193,21 +318,31 @@ export async function testProviderConnection( if (data?.errors && data.errors.length > 0) { return { - connected: false, + status: CONNECTION_CHECK_STATUS.FAILED, error: data.errors[0]?.detail ?? "Unknown error", }; } const taskId = data?.data?.id; if (!taskId) { - return { connected: false, error: "No task ID returned" }; + return { + status: CONNECTION_CHECK_STATUS.FAILED, + error: "No task ID returned", + }; } - const taskResult = await checkTaskStatus(taskId); + const taskResult = await checkTaskStatus( + taskId, + PROVIDER_CONNECTION_CHECK_MAX_RETRIES, + PROVIDER_CONNECTION_CHECK_POLL_DELAY_MS, + ); if (!taskResult.completed) { + if (taskResult.error === TASK_STATUS_MAX_RETRIES_ERROR) { + return resolveProviderConnectionState(providerId, baseline); + } return { - connected: false, + status: CONNECTION_CHECK_STATUS.FAILED, error: taskResult.error ?? "Connection test timed out", }; } @@ -217,10 +352,9 @@ export async function testProviderConnection( const connected = result?.connected === true; return { - connected, - error: connected - ? null - : result?.error || - "Connection was not confirmed. Test the connection again.", + status: connected + ? CONNECTION_CHECK_STATUS.SUCCESS + : CONNECTION_CHECK_STATUS.FAILED, + error: connected ? null : result?.error || CONNECTION_NOT_CONFIRMED_MESSAGE, }; } diff --git a/ui/lib/providers-navigation.ts b/ui/lib/providers-navigation.ts index 7428eed540..53b0b03ad3 100644 --- a/ui/lib/providers-navigation.ts +++ b/ui/lib/providers-navigation.ts @@ -1,3 +1,25 @@ +import { + WIZARD_OPEN_SOURCE, + type WizardOpenSource, +} from "@/lib/provider-funnel/provider-funnel-events"; + export const ADD_PROVIDER_SEARCH_PARAM = "addProvider"; export const ADD_PROVIDER_SEARCH_VALUE = "true"; export const ADD_PROVIDER_HREF = `/providers?${ADD_PROVIDER_SEARCH_PARAM}=${ADD_PROVIDER_SEARCH_VALUE}`; + +// Optional hint telling the providers page which entry point opened the wizard. +export const ADD_PROVIDER_SOURCE_PARAM = "addProviderSource"; + +export const buildAddProviderHref = (source: WizardOpenSource): string => + `${ADD_PROVIDER_HREF}&${ADD_PROVIDER_SOURCE_PARAM}=${source}`; + +const WIZARD_OPEN_SOURCES: readonly string[] = + Object.values(WIZARD_OPEN_SOURCE); + +// Unknown or missing hints fall back to a plain URL-driven open. +export const resolveAddProviderSource = ( + value: string | null | undefined, +): WizardOpenSource => + value && WIZARD_OPEN_SOURCES.includes(value) + ? (value as WizardOpenSource) + : WIZARD_OPEN_SOURCE.URL; diff --git a/ui/lib/report-download-access.ts b/ui/lib/report-download-access.ts new file mode 100644 index 0000000000..afc565b89a --- /dev/null +++ b/ui/lib/report-download-access.ts @@ -0,0 +1,10 @@ +export const REPORT_DOWNLOAD_LOCKED_ERROR = + "Report downloads require an active subscription."; + +/** + * Whether the current tenant must upgrade before downloading reports. + * Self-hosted deployments never lock downloads; the Prowler Cloud overlay + * replaces this body with its billing lookup. + */ +export const isReportDownloadLocked = (): Promise => + Promise.resolve(false); diff --git a/ui/lib/tours/__tests__/use-driver-tour.lifecycle.test.tsx b/ui/lib/tours/__tests__/use-driver-tour.lifecycle.test.tsx index 1c300787d4..55cc5f6cb0 100644 --- a/ui/lib/tours/__tests__/use-driver-tour.lifecycle.test.tsx +++ b/ui/lib/tours/__tests__/use-driver-tour.lifecycle.test.tsx @@ -11,6 +11,7 @@ const driverHarness = vi.hoisted(() => { drive: ReturnType; isActive: ReturnType; isLastStep: ReturnType; + setSteps: ReturnType; }> = []; const driverMock = vi.fn((config: { onDestroyed?: () => void }) => { @@ -27,6 +28,7 @@ const driverHarness = vi.hoisted(() => { isLastStep: vi.fn(() => false), moveNext: vi.fn(), movePrevious: vi.fn(), + setSteps: vi.fn(), }; instances.push(instance); return instance; @@ -195,4 +197,200 @@ describe("useDriverTour lifecycle", () => { // ...but no completion record was persisted, so the tour can reappear later. expect(store.get({ id: tour.id, version: tour.version })).toBeNull(); }); + + describe("when asked to start at an anchored step", () => { + const anchoredTour = { + id: "anchored-tour", + version: 1, + coversFiles: [], + steps: [ + { title: "Welcome", description: "Intro" }, + { target: "late", title: "Late anchor", description: "Inside a modal" }, + ], + } satisfies TourDefinition; + + function AnchoredProbe({ + onResult, + }: { + onResult: (result: UseDriverTourResult) => void; + }) { + onResult( + useDriverTour(anchoredTour, { autoOpen: false, store: createStore() }), + ); + return null; + } + + afterEach(() => { + document.body.innerHTML = ""; + }); + + it("renumbers the tour from the anchor once it is in the DOM", async () => { + // Given + let latestResult: UseDriverTourResult | undefined; + render( (latestResult = result)} />); + const anchor = document.createElement("div"); + anchor.setAttribute("data-tour-id", "anchored-tour-late"); + document.body.appendChild(anchor); + + // When + await act(async () => { + latestResult?.start("late"); + }); + + // Then: the skipped steps describe UI the user already went through, + // so the tour reads "Step 1 of 1", not "Step 2 of 2". + const [instance] = driverHarness.instances; + expect(instance.setSteps).toHaveBeenCalledExactlyOnceWith([ + expect.objectContaining({ + popover: expect.objectContaining({ title: "Late anchor" }), + }), + ]); + expect(instance.drive).toHaveBeenCalledExactlyOnceWith(); + }); + + it("takes over from a pending auto-open so the tour starts at the anchor", async () => { + // Given: auto-open is armed but the caller asks for the anchored start first. + vi.useFakeTimers(); + let latestResult: UseDriverTourResult | undefined; + function AutoOpenAnchoredProbe() { + latestResult = useDriverTour(anchoredTour, { + autoOpen: true, + store: createStore(), + }); + return null; + } + render(); + act(() => { + latestResult?.start("late"); + }); + + // When: the auto-open delay elapses before the anchor exists. + act(() => { + vi.advanceTimersByTime(50); + }); + + // Then: nothing opens from the top. + const [instance] = driverHarness.instances; + expect(instance.drive).not.toHaveBeenCalled(); + + // When: the anchor mounts. + const anchor = document.createElement("div"); + anchor.setAttribute("data-tour-id", "anchored-tour-late"); + document.body.appendChild(anchor); + await act(async () => { + await vi.runAllTimersAsync(); + }); + + // Then: the tour is driven once, renumbered from the anchor. + expect(instance.setSteps).toHaveBeenCalledExactlyOnceWith([ + expect.objectContaining({ + popover: expect.objectContaining({ title: "Late anchor" }), + }), + ]); + expect(instance.drive).toHaveBeenCalledExactlyOnceWith(); + }); + + it("plays the whole tour again when started from the top afterwards", async () => { + // Given + let latestResult: UseDriverTourResult | undefined; + render( (latestResult = result)} />); + const anchor = document.createElement("div"); + anchor.setAttribute("data-tour-id", "anchored-tour-late"); + document.body.appendChild(anchor); + await act(async () => { + latestResult?.start("late"); + }); + + // When + await act(async () => { + latestResult?.stop(); + latestResult?.start(); + }); + + // Then + const [instance] = driverHarness.instances; + expect(instance.setSteps).toHaveBeenLastCalledWith([ + expect.objectContaining({ + popover: expect.objectContaining({ title: "Welcome" }), + }), + expect.objectContaining({ + popover: expect.objectContaining({ title: "Late anchor" }), + }), + ]); + }); + + it("starts at the step when only its fallback anchor is in the DOM", async () => { + // Given + const tourWithFallback = { + ...anchoredTour, + id: "fallback-tour", + steps: [ + anchoredTour.steps[0], + { ...anchoredTour.steps[1], fallbackTarget: "stable" }, + ], + } satisfies TourDefinition; + let latestResult: UseDriverTourResult | undefined; + function FallbackProbe() { + latestResult = useDriverTour(tourWithFallback, { + autoOpen: false, + store: createStore(), + }); + return null; + } + render(); + const anchor = document.createElement("div"); + anchor.setAttribute("data-tour-id", "fallback-tour-stable"); + document.body.appendChild(anchor); + + // When + await act(async () => { + latestResult?.start("late"); + }); + + // Then + const [instance] = driverHarness.instances; + expect(instance.setSteps).toHaveBeenCalledExactlyOnceWith([ + expect.objectContaining({ + popover: expect.objectContaining({ title: "Late anchor" }), + }), + ]); + expect(instance.drive).toHaveBeenCalledExactlyOnceWith(); + }); + + it("stays closed when stopped before the anchor mounts", async () => { + // Given + let latestResult: UseDriverTourResult | undefined; + render( (latestResult = result)} />); + await act(async () => { + latestResult?.start("late"); + }); + + // When + await act(async () => { + latestResult?.stop(); + const anchor = document.createElement("div"); + anchor.setAttribute("data-tour-id", "anchored-tour-late"); + document.body.appendChild(anchor); + }); + + // Then + expect(driverHarness.instances[0].drive).not.toHaveBeenCalled(); + }); + + it("starts from the first step when the target is not part of the tour", async () => { + // Given + let latestResult: UseDriverTourResult | undefined; + render( (latestResult = result)} />); + + // When + await act(async () => { + latestResult?.start("unknown"); + }); + + // Then + expect( + driverHarness.instances[0].drive, + ).toHaveBeenCalledExactlyOnceWith(); + }); + }); }); diff --git a/ui/lib/tours/add-provider.tour.ts b/ui/lib/tours/add-provider.tour.ts index 614735d79b..6c4831babd 100644 --- a/ui/lib/tours/add-provider.tour.ts +++ b/ui/lib/tours/add-provider.tour.ts @@ -8,10 +8,11 @@ import { export const ADD_PROVIDER_TOUR_TARGETS = { TRIGGER: "trigger", PROVIDER_TYPE: "provider-type", - // Wraps the whole wizard modal so the final step's spotlight covers every input - // (UID, alias) and the footer — driver.js only keeps the highlighted element and - // its descendants interactive, so anchoring here stops the overlay from freezing - // those inputs. + // Wraps the wizard's form column so the final step's spotlight covers every + // input (UID, alias) — driver.js only keeps the highlighted element and its + // descendants interactive, so anchoring here stops the overlay from freezing + // those inputs. The footer sits outside the anchor and stays clickable through + // `data-tour-interactive` (see styles/tours.css). WIZARD_BODY: "wizard-body", } as const; @@ -53,15 +54,16 @@ export const addProviderTour = defineTour({ }, { target: "wizard-body", - // Pinned to the left of the form column, mirroring the provider-type step. + // Left of the form column, in the gap under the stepper and level with the + // footer the user continues from, so it never covers the form itself. side: TOUR_STEP_SIDES.LEFT, - align: TOUR_STEP_ALIGNMENTS.START, + align: TOUR_STEP_ALIGNMENTS.END, // Final step: stays until the user closes it or advances to credentials, which // the wizard ends the tour from. No Next button. autoAdvance: true, title: "Add your account details", description: - "Enter your account ID and an optional alias, then continue. From here you'll add credentials, test the connection, and launch your first scan — at your own pace.", + "Fill in the connection details for this provider, then continue. Prowler checks the connection, then you launch your first scan — at your own pace.", }, ], }); diff --git a/ui/lib/tours/use-driver-tour.ts b/ui/lib/tours/use-driver-tour.ts index d7228cd303..8fa4ac94a8 100644 --- a/ui/lib/tours/use-driver-tour.ts +++ b/ui/lib/tours/use-driver-tour.ts @@ -102,7 +102,8 @@ export interface UseDriverTourOptions { } export interface UseDriverTourResult { - start: () => void; + /** Optional step `target` to begin from, skipping the steps before it. */ + start: (startAtTarget?: string) => void; stop: () => void; /** True if a completion record exists for `(tour.id, tour.version)`. */ hasCompleted: boolean; @@ -259,6 +260,11 @@ export function useDriverTour( // tour would be marked resolved forever after a simple theme toggle. const teardownRef = useRef(false); + // Bumped by start() and stop() so a pending anchored start knows it went stale. + const startGenerationRef = useRef(0); + // Every adapted step, so start() can hand driver.js a trimmed or full list. + const stepsRef = useRef([]); + const tourId = tour.id; const tourVersion = tour.version; const existing = store.get({ id: tourId, version: tourVersion }); @@ -361,6 +367,7 @@ export function useDriverTour( }); driverRef.current = driver(config); + stepsRef.current = steps; return () => { const instance = driverRef.current; @@ -381,7 +388,11 @@ export function useDriverTour( const instance = driverRef.current; if (!instance || instance.isActive()) return; + // A start()/stop() issued meanwhile takes over: an anchored start must not + // be pre-empted by the full tour opening from the top. + const generation = startGenerationRef.current; const timer = window.setTimeout(() => { + if (startGenerationRef.current !== generation) return; if (!instance.isActive()) { activeTourInstance = instance; instance.drive(); @@ -394,13 +405,46 @@ export function useDriverTour( }, [autoOpen, enabled, hasCompleted, tourId, tourVersion]); return { - start: () => { + start: (startAtTarget) => { const instance = driverRef.current; if (!instance) return; - activeTourInstance = instance; - instance.drive(); + const generation = ++startGenerationRef.current; + + const startIndex = startAtTarget + ? tour.steps.findIndex((step) => step.target === startAtTarget) + : -1; + if (!startAtTarget || startIndex <= 0) { + instance.setSteps(stepsRef.current); + activeTourInstance = instance; + instance.drive(); + return; + } + + // The anchor may mount right after the caller (e.g. a modal opening), so wait for it. + // Either anchor will do, mirroring how adaptStep resolves the step's element. + const fallbackTarget = tour.steps[startIndex].fallbackTarget; + const anchorSelector = [startAtTarget, fallbackTarget] + .filter((target): target is string => target !== undefined) + .map((target) => getTourTargetSelector(tourId, target)) + .join(", "); + waitForElement(anchorSelector) + .then(() => { + if (startGenerationRef.current !== generation) return; + if (driverRef.current !== instance || instance.isActive()) return; + // The skipped steps describe UI the caller already went through, so + // the tour is renumbered from the anchor ("Step 1 of 2", not "3 of 4"). + instance.setSteps(stepsRef.current.slice(startIndex)); + activeTourInstance = instance; + instance.drive(); + }) + .catch(() => { + // Anchor never appeared (e.g. the modal was dismissed); skip the tour. + }); + }, + stop: () => { + startGenerationRef.current += 1; + driverRef.current?.destroy(); }, - stop: () => driverRef.current?.destroy(), hasCompleted, }; } diff --git a/ui/store/cloud-upgrade/store.ts b/ui/store/cloud-upgrade/store.ts index 21e421ccb0..ee806c4ccb 100644 --- a/ui/store/cloud-upgrade/store.ts +++ b/ui/store/cloud-upgrade/store.ts @@ -2,15 +2,15 @@ import { create } from "zustand"; import { CLOUD_UPGRADE_FEATURE, - type CloudUpgradeFeature, + type UpgradeFeature, } from "@/types/cloud-upgrade"; interface CloudUpgradeStoreState { - activeFeature: CloudUpgradeFeature | null; - retainedFeature: CloudUpgradeFeature; + activeFeature: UpgradeFeature | null; + retainedFeature: UpgradeFeature; returnFocusElement: HTMLElement | null; openCloudUpgrade: ( - feature: CloudUpgradeFeature, + feature: UpgradeFeature, returnFocusElement?: HTMLElement, ) => void; closeCloudUpgrade: () => void; diff --git a/ui/store/index.ts b/ui/store/index.ts index 57fc46f86d..99a002f407 100644 --- a/ui/store/index.ts +++ b/ui/store/index.ts @@ -2,6 +2,7 @@ export * from "./cloud-upgrade/store"; export * from "./compliance/store"; export * from "./jira-dispatch/store"; export * from "./organizations/store"; +export * from "./partial-scan/store"; export * from "./provider-wizard/store"; export * from "./scans/store"; export * from "./ui/store"; diff --git a/ui/store/partial-scan/store.ts b/ui/store/partial-scan/store.ts new file mode 100644 index 0000000000..f804a7938e --- /dev/null +++ b/ui/store/partial-scan/store.ts @@ -0,0 +1,17 @@ +import { create } from "zustand"; + +import type { PartialScanTarget } from "@/types/partial-scans"; + +interface PartialScanStoreState { + activeTarget: PartialScanTarget | null; + openPartialScan: (target: PartialScanTarget) => void; + closePartialScan: () => void; +} + +// Menu items live inside dropdowns that unmount on select, so the confirmation +// modal is hosted once globally and driven through this store. +export const usePartialScanStore = create((set) => ({ + activeTarget: null, + openPartialScan: (activeTarget) => set({ activeTarget }), + closePartialScan: () => set({ activeTarget: null }), +})); diff --git a/ui/store/provider-wizard/store.ts b/ui/store/provider-wizard/store.ts index 90182dc88b..2d94eff7db 100644 --- a/ui/store/provider-wizard/store.ts +++ b/ui/store/provider-wizard/store.ts @@ -2,6 +2,7 @@ import { create } from "zustand"; import { createJSONStorage, persist } from "zustand/middleware"; import { + AwsConnectDraft, PROVIDER_WIZARD_MODE, ProviderWizardIdentity, ProviderWizardMode, @@ -16,10 +17,12 @@ interface ProviderWizardState { via: string | null; secretId: string | null; mode: ProviderWizardMode; + awsConnectDraft: AwsConnectDraft | null; setProvider: (provider: ProviderWizardIdentity) => void; setVia: (via: string | null) => void; setSecretId: (secretId: string | null) => void; setMode: (mode: ProviderWizardMode) => void; + setAwsConnectDraft: (patch: Partial) => void; reset: () => void; } @@ -31,6 +34,13 @@ const initialState = { via: null, secretId: null, mode: PROVIDER_WIZARD_MODE.ADD, + awsConnectDraft: null, +}; + +const EMPTY_AWS_CONNECT_DRAFT: AwsConnectDraft = { + method: "role", + roleValues: {}, + keysValues: {}, }; export const useProviderWizardStore = create()( @@ -47,11 +57,21 @@ export const useProviderWizardStore = create()( setVia: (via) => set({ via }), setSecretId: (secretId) => set({ secretId }), setMode: (mode) => set({ mode }), + setAwsConnectDraft: (patch) => + set((state) => ({ + awsConnectDraft: { + ...EMPTY_AWS_CONNECT_DRAFT, + ...state.awsConnectDraft, + ...patch, + }, + })), reset: () => set(initialState), }), { name: "provider-wizard-store", storage: createJSONStorage(() => sessionStorage), + // The draft may hold access keys: it never leaves memory. + partialize: ({ awsConnectDraft: _draft, ...persisted }) => persisted, }, ), ); diff --git a/ui/store/ui/store-initializer.test.tsx b/ui/store/ui/store-initializer.test.tsx index 01a3545bb9..77428af1c5 100644 --- a/ui/store/ui/store-initializer.test.tsx +++ b/ui/store/ui/store-initializer.test.tsx @@ -7,7 +7,11 @@ import { StoreInitializer } from "./store-initializer"; describe("StoreInitializer", () => { beforeEach(() => { localStorage.clear(); - useUIStore.setState({ hasProviders: false, registryEligible: false }); + useUIStore.setState({ + hasProviders: false, + hasProvidersResolved: false, + registryEligible: false, + }); }); it("keeps Registry hidden when the server sends no eligibility decision", () => { @@ -33,4 +37,23 @@ describe("StoreInitializer", () => { expect(persisted.state?.hasProviders).toBe(true); expect(persisted.state).not.toHaveProperty("registryEligible"); }); + + it("leaves the provider count unresolved when the server could not determine it", () => { + // Given / When + render(); + + // Then + expect(useUIStore.getState().hasProvidersResolved).toBe(false); + }); + + it("resolves a confirmed empty tenant without persisting the resolution", () => { + // Given / When + render(); + + // Then + expect(useUIStore.getState().hasProviders).toBe(false); + expect(useUIStore.getState().hasProvidersResolved).toBe(true); + const persisted = JSON.parse(localStorage.getItem("ui-store") ?? "{}"); + expect(persisted.state).not.toHaveProperty("hasProvidersResolved"); + }); }); diff --git a/ui/store/ui/store.ts b/ui/store/ui/store.ts index 5aa321e0b1..522658995a 100644 --- a/ui/store/ui/store.ts +++ b/ui/store/ui/store.ts @@ -4,6 +4,8 @@ import { persist } from "zustand/middleware"; interface UIStoreState { isSideMenuOpen: boolean; hasProviders: boolean; + // True once the server reported a definitive provider count for this session. + hasProvidersResolved: boolean; registryEligible: boolean; openSideMenu: () => void; @@ -17,17 +19,19 @@ export const useUIStore = create()( (set) => ({ isSideMenuOpen: false, hasProviders: false, + hasProvidersResolved: false, registryEligible: false, openSideMenu: () => set({ isSideMenuOpen: true }), closeSideMenu: () => set({ isSideMenuOpen: false }), - setHasProviders: (value: boolean) => set({ hasProviders: value }), + setHasProviders: (value: boolean) => + set({ hasProviders: value, hasProvidersResolved: true }), setRegistryEligible: (value: boolean) => set({ registryEligible: value }), }), { name: "ui-store", - // Registry eligibility is a per-request server decision; persisting it - // would resurface a stale entry on the next session before the server - // seed corrects it. + // Registry eligibility and the provider-count resolution are per-request + // server decisions; persisting them would resurface a stale entry on the + // next session before the server seed corrects it. partialize: ({ isSideMenuOpen, hasProviders }) => ({ isSideMenuOpen, hasProviders, diff --git a/ui/styles/tours.css b/ui/styles/tours.css index 3764e606b8..36828183ed 100644 --- a/ui/styles/tours.css +++ b/ui/styles/tours.css @@ -10,6 +10,14 @@ pointer-events: auto; } +/* Surfaces that must stay usable while a tour drives (the wizard footer, a + * discovery callout). Clickable regardless of the spotlight; a surface that lives + * in its own top-level layer (a portal) also rises above the overlay (z 10000). */ +.driver-active [data-tour-interactive], +.driver-active [data-tour-interactive] * { + pointer-events: auto; +} + @keyframes driver-fade-in { 0% { opacity: 0; diff --git a/ui/tests/providers/providers-page.ts b/ui/tests/providers/providers-page.ts index e9bd96cf08..ecb8755ccd 100644 --- a/ui/tests/providers/providers-page.ts +++ b/ui/tests/providers/providers-page.ts @@ -397,19 +397,21 @@ export class ProvidersPage extends BasePage { // "Add Provider" control; with zero providers the page renders the empty // state whose CTA is labelled "Open Add Provider modal" (button on // /providers, link on /scans). Only one of these is ever in the DOM at once. - this.addProviderButton = page + // Scoped to
: an empty tenant also gets an "Add Provider" CTA in the sidebar. + const main = page.getByRole("main"); + this.addProviderButton = main .getByRole("button", { name: "Add Provider", exact: true, }) .or( - page.getByRole("link", { + main.getByRole("link", { name: "Add Provider", exact: true, }), ) - .or(page.getByRole("button", { name: "Open Add Provider modal" })) - .or(page.getByRole("link", { name: "Open Add Provider modal" })); + .or(main.getByRole("button", { name: "Open Add Provider modal" })) + .or(main.getByRole("link", { name: "Open Add Provider modal" })); // Table displaying existing providers this.providersTable = page.getByRole("table"); @@ -701,13 +703,15 @@ export class ProvidersPage extends BasePage { await this.selectProviderRadio(this.githubProviderRadio); } - async selectAWSSingleAccountMethod(): Promise { - const singleAccountOption = this.page.getByRole("radio", { - name: "Add A Single AWS Cloud Account", - exact: true, - }); - await expect(singleAccountOption).toBeVisible({ timeout: 10000 }); - await singleAccountOption.click(); + // AWS picks its access method on the same step that registers the account. + async selectAwsAccessMethod(type: AWSCredentialType): Promise { + const name = + type === AWS_CREDENTIAL_OPTIONS.AWS_CREDENTIALS + ? "Static access keys" + : /IAM Role/; + const accessMethod = this.wizardModal.getByRole("radio", { name }); + await expect(accessMethod).toBeVisible({ timeout: 10000 }); + await accessMethod.click(); } async selectAzureSingleSubscriptionMethod(): Promise { @@ -729,12 +733,7 @@ export class ProvidersPage extends BasePage { } async selectAWSOrganizationsMethod(): Promise { - await this.page - .getByRole("radio", { - name: "Add Multiple Accounts With AWS Organizations", - exact: true, - }) - .click(); + await this.page.getByRole("tab", { name: /Full AWS Organization/ }).click(); } async verifyOrganizationsAuthenticationStepLoaded(): Promise { @@ -774,10 +773,12 @@ export class ProvidersPage extends BasePage { await this.page.getByRole("option", { name: optionName }).click(); } + // The account id is only typed for access keys; with a role it is read from the ARN. async fillAWSProviderDetails(data: AWSProviderData): Promise { - await this.selectAWSSingleAccountMethod(); - await expect(this.accountIdInput).toBeVisible({ timeout: 10000 }); - await this.accountIdInput.fill(data.accountId); + await expect(this.aliasInput).toBeVisible({ timeout: 10000 }); + if (await this.accountIdInput.isVisible().catch(() => false)) { + await this.accountIdInput.fill(data.accountId); + } if (data.alias) { await this.aliasInput.fill(data.alias); @@ -881,6 +882,7 @@ export class ProvidersPage extends BasePage { const actionNames = [ "Go to scans", "Authenticate", + "Connect account", "Next", "Save", "Check connection", @@ -1123,41 +1125,6 @@ export class ProvidersPage extends BasePage { async fillRoleCredentials(credentials: AWSProviderCredential): Promise { await expect(this.roleArnInput).toBeVisible({ timeout: 10000 }); - const accessKeyInputInWizard = this.wizardModal.getByPlaceholder( - "Enter the AWS Access Key ID", - ); - const secretKeyInputInWizard = this.wizardModal.getByPlaceholder( - "Enter the AWS Secret Access Key", - ); - const accessKeyId = - credentials.accessKeyId || process.env.E2E_AWS_PROVIDER_ACCESS_KEY; - const secretAccessKey = - credentials.secretAccessKey || process.env.E2E_AWS_PROVIDER_SECRET_KEY; - - const shouldFillStaticKeys = Boolean(accessKeyId || secretAccessKey); - if (shouldFillStaticKeys) { - const accessKeyIsVisible = await accessKeyInputInWizard - .isVisible() - .catch(() => false); - - // In cloud env the default can be SDK mode, so expose Access/Secret explicitly. - if (!accessKeyIsVisible) { - await this.selectAuthenticationMethod( - AWS_CREDENTIAL_OPTIONS.AWS_ROLE_ARN, - ); - } - } - - if (accessKeyId) { - await expect(accessKeyInputInWizard).toBeVisible({ timeout: 10000 }); - await accessKeyInputInWizard.fill(accessKeyId); - await expect(accessKeyInputInWizard).toHaveValue(accessKeyId); - } - if (secretAccessKey) { - await expect(secretKeyInputInWizard).toBeVisible({ timeout: 10000 }); - await secretKeyInputInWizard.fill(secretAccessKey); - await expect(secretKeyInputInWizard).toHaveValue(secretAccessKey); - } if (credentials.roleArn) { await this.roleArnInput.fill(credentials.roleArn); } @@ -1673,33 +1640,6 @@ export class ProvidersPage extends BasePage { } } - async selectAuthenticationMethod(method: AWSCredentialType): Promise { - // Select the authentication method (shadcn Select renders as combobox + listbox) - - const trigger = this.page.locator('[role="combobox"]').filter({ - hasText: /AWS SDK Default|Prowler Cloud will assume|Access & Secret Key/i, - }); - - await trigger.click(); - - const listbox = this.page.getByRole("listbox"); - await expect(listbox).toBeVisible({ timeout: 10000 }); - - if (method === AWS_CREDENTIAL_OPTIONS.AWS_ROLE_ARN) { - await this.page - .getByRole("option", { name: "Access & Secret Key" }) - .click({ force: true }); - } else if (method === AWS_CREDENTIAL_OPTIONS.AWS_SDK_DEFAULT) { - await this.page - .getByRole("option", { - name: /AWS SDK Default|Prowler Cloud will assume your IAM role/i, - }) - .click({ force: true }); - } else { - throw new Error(`Invalid authentication method: ${method}`); - } - } - async clickProviderRowActions(providerUid: string): Promise { // Click the actions dropdown for a specific provider row const row = this.providersTable.locator("tbody tr", { diff --git a/ui/tests/providers/providers.md b/ui/tests/providers/providers.md index 6aacca06da..0bc1b9d47d 100644 --- a/ui/tests/providers/providers.md +++ b/ui/tests/providers/providers.md @@ -29,9 +29,9 @@ 1. Navigate to providers page 2. Click "Add Provider" button 3. Select AWS provider type -4. Fill provider details (account ID and alias) -5. Select "credentials" authentication type -6. Fill static credentials (access key and secret key) +4. On the single AWS step, select the "Access keys" access method +5. Fill provider details (account ID and alias) +6. Fill static credentials (access key and secret key) and click "Connect account" 7. Confirm provider connection without launching a scan 8. Verify return to Providers page 9. Verify provider exists in Providers table @@ -61,63 +61,6 @@ --- -## Test Case: `PROVIDER-E2E-002` - Add AWS Provider with Assume Role Credentials Access Key and Secret Key - -**Priority:** `critical` - -**Tags:** - -- type → @e2e, @serial -- feature → @providers -- provider → @aws - -**Description/Objective:** Validates the complete flow of adding a new AWS provider using role-based authentication with Access Key and Secret Key - -**Preconditions:** - -- Admin user authentication required (admin.auth.setup setup) -- Environment variables configured: E2E_AWS_PROVIDER_ACCOUNT_ID, E2E_AWS_PROVIDER_ACCESS_KEY, E2E_AWS_PROVIDER_SECRET_KEY, E2E_AWS_PROVIDER_ROLE_ARN -- Remove any existing provider with the same Account ID before starting the test -- This test must be run serially and never in parallel with other tests, as it requires the Account ID not to be already registered beforehand. - -### Flow Steps - -1. Navigate to providers page -2. Click "Add Provider" button -3. Select AWS provider type -4. Fill provider details (account ID and alias) -5. Select "role" authentication type -6. Fill role credentials (access key, secret key, and role ARN) -7. Confirm provider connection without launching a scan -8. Verify return to Providers page -9. Verify provider exists in Providers table - -### Expected Result - -- AWS provider successfully added with role credentials -- Provider connection validated without launching a scan -- User returned to Providers page -- Provider appears in Providers table with the expected UID - -### Key verification points - -- Provider page loads correctly -- Connect account page displays AWS option -- Role credentials form accepts all required fields -- Launch step appears -- Successful return to Providers page after closing the launch step -- Provider exists in Providers table (verified by account ID) -- Provider UID matches the expected value - -### Notes - -- Test uses environment variables for AWS credentials and role ARN -- Provider cleanup performed before each test to ensure clean state -- Requires valid AWS account with role assumption permissions -- Role ARN must be properly configured - ---- - ## Test Case: `PROVIDER-E2E-003` - Add Azure Provider with Static Credentials **Priority:** `critical` @@ -618,10 +561,10 @@ 1. Navigate to providers page 2. Click "Add Provider" button 3. Select AWS provider type -4. Fill provider details (account ID and alias) -5. Select "role" authentication type -6. Switch authentication method to "Use AWS SDK default credentials" -7. Fill role ARN using AWS SDK credential inputs +4. On the single AWS step, keep the "IAM Role" access method +5. Fill the alias (the account ID is read from the role ARN) +6. Nothing to choose: the role is assumed with the host's AWS SDK default credentials +7. Fill the role ARN and click "Connect account" 8. Confirm provider connection without launching a scan 9. Verify return to Providers page 10. Verify provider exists in Providers table @@ -637,8 +580,8 @@ - Provider page loads correctly - Connect account page displays AWS option -- Credentials form exposes AWS SDK default authentication method -- Role ARN field accepts provided value when SDK method is selected +- The role form asks for no credentials of its own: the AWS SDK default chain assumes the role +- Role ARN field accepts the provided value - Launch step appears - Successful return to Providers page after closing the launch step - Provider exists in Providers table (verified by account ID) diff --git a/ui/tests/providers/providers.spec.ts b/ui/tests/providers/providers.spec.ts index 724428df5e..4052a1d4c0 100644 --- a/ui/tests/providers/providers.spec.ts +++ b/ui/tests/providers/providers.spec.ts @@ -108,16 +108,11 @@ test.describe("Add Provider", () => { // Select AWS provider await providersPage.selectAWSProvider(); - // Fill provider details - await providersPage.fillAWSProviderDetails(awsProviderData); - await providersPage.clickNext(); - - await providersPage.verifyCredentialsPageLoaded(); - - // Select static credentials type - await providersPage.selectCredentialsType( + // AWS registers the account and its credentials in a single step + await providersPage.selectAwsAccessMethod( AWS_CREDENTIAL_OPTIONS.AWS_CREDENTIALS, ); + await providersPage.fillAWSProviderDetails(awsProviderData); // Fill static credentials await providersPage.fillStaticCredentials(staticCredentials); @@ -130,73 +125,6 @@ test.describe("Add Provider", () => { }, ); - test( - "should add a new AWS provider with assume role credentials with Access Key and Secret Key", - { - tag: [ - "@critical", - "@e2e", - "@providers", - "@aws", - "@serial", - "@PROVIDER-E2E-002", - ], - }, - async ({ page }) => { - // Validate required environment variables - if (!roleArn) { - throw new Error( - "E2E_AWS_PROVIDER_ROLE_ARN environment variable is not set", - ); - } - - // Prepare test data for AWS provider - const awsProviderData: AWSProviderData = { - accountId: accountId, - alias: "Test E2E AWS Account - Credentials", - }; - - // Prepare role-based credentials - const roleCredentials: AWSProviderCredential = { - type: AWS_CREDENTIAL_OPTIONS.AWS_ROLE_ARN, - accessKeyId: accessKey, - secretAccessKey: secretKey, - roleArn: roleArn, - }; - - // Navigate to providers page - await providersPage.goto(); - await providersPage.verifyPageLoaded(); - - // Start adding new provider - await providersPage.clickAddProvider(); - await providersPage.verifyConnectAccountPageLoaded(); - - // Select AWS provider - await providersPage.selectAWSProvider(); - - // Fill provider details - await providersPage.fillAWSProviderDetails(awsProviderData); - await providersPage.clickNext(); - - await providersPage.verifyCredentialsPageLoaded(); - - // Select role credentials type - await providersPage.selectCredentialsType( - AWS_CREDENTIAL_OPTIONS.AWS_ROLE_ARN, - ); - - // Fill role credentials - await providersPage.fillRoleCredentials(roleCredentials); - await providersPage.clickNext(); - - // Confirm the provider connection without launching a scan - await providersPage.completeProviderConnectionWithoutLaunchingScan( - accountId, - ); - }, - ); - test( "should add a new AWS provider with assume role credentials using AWS SDK", { @@ -240,21 +168,15 @@ test.describe("Add Provider", () => { // Select AWS provider await providersPage.selectAWSProvider(); - // Fill provider details - await providersPage.fillAWSProviderDetails(awsProviderData); - await providersPage.clickNext(); - - // Select role credentials type - await providersPage.selectCredentialsType( + // AWS registers the account (read from the role ARN) and its + // credentials in a single step + await providersPage.selectAwsAccessMethod( AWS_CREDENTIAL_OPTIONS.AWS_ROLE_ARN, ); - await providersPage.verifyCredentialsPageLoaded(); - - // Select Authentication Method - await providersPage.selectAuthenticationMethod( - AWS_CREDENTIAL_OPTIONS.AWS_SDK_DEFAULT, - ); + await providersPage.fillAWSProviderDetails(awsProviderData); + // The role is assumed with the credentials of the host running Prowler + // (AWS SDK default); the wizard asks for nothing else. // Fill role credentials await providersPage.fillRoleCredentials(roleCredentials); await providersPage.clickNext(); diff --git a/ui/tests/scans/scans-page.ts b/ui/tests/scans/scans-page.ts index e75c2d3508..9a18d341c0 100644 --- a/ui/tests/scans/scans-page.ts +++ b/ui/tests/scans/scans-page.ts @@ -24,9 +24,9 @@ export class ScansPage extends BasePage { super(page); // Scan provider selection elements - // The sidebar exposes its own icon-button labeled "Launch Scan" - // (aria-label, wrapped in a Tooltip), so scoping by accessible name - // alone hits a strict-mode duplicate. Scope to the page-shell's + // The sidebar exposes its own action labeled "Launch Scan" (it reads + // "Add Provider" only while the tenant has no providers), so scoping by + // accessible name alone hits a strict-mode duplicate. Scope to the page-shell's // tabs-and-actions group, which only contains the visible-text // Launch Scan button. this.launchScanButton = page diff --git a/ui/tests/setups/manage-registry.auth.setup.ts b/ui/tests/setups/manage-registry.auth.setup.ts index bf3daac561..8fd5195af1 100644 --- a/ui/tests/setups/manage-registry.auth.setup.ts +++ b/ui/tests/setups/manage-registry.auth.setup.ts @@ -21,6 +21,8 @@ authManageRegistrySetup( const signInPage = new SignInPage(page); await signInPage.goto(); + // The fixture tenant has no providers: keep the first-run redirect out of the way. + await signInPage.skipFirstRunRedirect(); await signInPage.login(fixtureCredentials); await page.waitForURL("/"); await new RegistryPage(page).dismissWelcomeDialog(); diff --git a/ui/tests/sign-in-base/sign-in-base-page.ts b/ui/tests/sign-in-base/sign-in-base-page.ts index 8966e26831..799841f68a 100644 --- a/ui/tests/sign-in-base/sign-in-base-page.ts +++ b/ui/tests/sign-in-base/sign-in-base-page.ts @@ -395,7 +395,23 @@ export class SignInPage extends BasePage { ); } - await this.loginAndVerify(credentials); + await this.goto(); + await this.skipFirstRunRedirect(); + await this.login(credentials); + await this.verifySuccessfulLogin(); await this.page.context().storageState({ path: storagePath }); } + + /** + * An empty tenant redirects each fresh browser context to the add-provider + * wizard once. Suites expect a plain landing, so mark that first run as done + * with the browser-wide key (the per-tenant one is only written by the app) + * before signing in; sign-up.spec covers the redirect itself with a brand-new + * tenant. Call it on a page already on the app origin. + */ + async skipFirstRunRedirect(): Promise { + await this.page.evaluate(() => { + window.localStorage.setItem("prowler.onboarding.first-run", "true"); + }); + } } diff --git a/ui/tests/sign-up/sign-up.md b/ui/tests/sign-up/sign-up.md index 542d1c0778..00b4639dbf 100644 --- a/ui/tests/sign-up/sign-up.md +++ b/ui/tests/sign-up/sign-up.md @@ -33,13 +33,14 @@ ### Expected Result - Sign-up succeeds and redirects to Login. -- User can log in successfully using the created credentials and reach the home page. +- User can log in successfully using the created credentials. +- Because the new tenant has no providers, the first run lands on `/providers` with the add-provider wizard already open (instead of the home page). ### Key verification points - After submitting sign-up, the URL changes to `/sign-in`. - The newly created credentials can be used to sign in successfully. -- After login, the user lands on the home (`/`) and main content is visible. +- After login, the user lands on `/providers` and the "Adding A Provider" wizard is visible. ### Notes diff --git a/ui/tests/sign-up/sign-up.spec.ts b/ui/tests/sign-up/sign-up.spec.ts index 3c732c6497..62336db7c6 100644 --- a/ui/tests/sign-up/sign-up.spec.ts +++ b/ui/tests/sign-up/sign-up.spec.ts @@ -1,6 +1,7 @@ -import { test } from "@playwright/test"; +import { expect, test } from "@playwright/test"; import { SignUpPage } from "./sign-up-page"; import { SignInPage } from "../sign-in-base/sign-in-base-page"; +import { ProvidersPage } from "../providers/providers-page"; import { makeSuffix } from "../helpers"; test.describe("Sign Up Flow", () => { @@ -45,7 +46,12 @@ test.describe("Sign Up Flow", () => { email: uniqueEmail, password: password, }); - await signInPage.verifySuccessfulLogin(); + + // A brand-new tenant has no providers, so the first run lands on the + // add-provider wizard instead of the Overview. + const providersPage = new ProvidersPage(page); + await expect(page).toHaveURL(/\/providers/); + await providersPage.verifyWizardModalOpen(); }, ); }); diff --git a/ui/types/cloud-upgrade.ts b/ui/types/cloud-upgrade.ts index 4c821390f5..39d7fd113e 100644 --- a/ui/types/cloud-upgrade.ts +++ b/ui/types/cloud-upgrade.ts @@ -16,3 +16,13 @@ export const CLOUD_UPGRADE_FEATURE = { export type CloudUpgradeFeature = (typeof CLOUD_UPGRADE_FEATURE)[keyof typeof CLOUD_UPGRADE_FEATURE]; + +// Prowler Cloud features gated behind a paid plan. +export const PAID_PLAN_UPGRADE_FEATURE = { + REPORT_DOWNLOAD: "report_download", +} as const; + +export type PaidPlanUpgradeFeature = + (typeof PAID_PLAN_UPGRADE_FEATURE)[keyof typeof PAID_PLAN_UPGRADE_FEATURE]; + +export type UpgradeFeature = CloudUpgradeFeature | PaidPlanUpgradeFeature; diff --git a/ui/types/partial-scans.ts b/ui/types/partial-scans.ts new file mode 100644 index 0000000000..d5eedcc89f --- /dev/null +++ b/ui/types/partial-scans.ts @@ -0,0 +1,15 @@ +import type { ProviderType } from "./providers"; + +/** Mirrors `PARTIAL_SCAN_MAX_RESOURCES` in the Cloud API. */ +export const PARTIAL_SCAN_MAX_RESOURCES = 10; + +/** One resource to re-check. Prowler Cloud only. */ +export interface PartialScanTarget { + /** Provider UUID. Resolved from `providerUid` + `providerType` when absent. */ + providerId?: string; + providerUid: string; + providerType: ProviderType; + providerAlias?: string; + resourceUid: string; + resourceName: string; +} diff --git a/ui/types/provider-wizard.ts b/ui/types/provider-wizard.ts index c4ad805edd..8f8dd5c17a 100644 --- a/ui/types/provider-wizard.ts +++ b/ui/types/provider-wizard.ts @@ -24,3 +24,12 @@ export interface ProviderWizardIdentity { uid: string | null; alias: string | null; } + +export type AwsConnectDraftValues = Record; + +/** What the AWS connect step typed so far; in memory only, gone with the wizard. */ +export interface AwsConnectDraft { + method: string; + roleValues: AwsConnectDraftValues; + keysValues: AwsConnectDraftValues; +} diff --git a/ui/types/providers.ts b/ui/types/providers.ts index f6f898ce48..272578bf32 100644 --- a/ui/types/providers.ts +++ b/ui/types/providers.ts @@ -1,4 +1,5 @@ import type { ScheduleFrequency } from "./schedules"; +import { TASK_OUTCOME, type TaskOutcome } from "./tasks"; export const PROVIDER_TYPES = [ "aws", @@ -22,6 +23,17 @@ export const PROVIDER_TYPES = [ /** The closed set of provider types this UI build ships bespoke assets for. */ export type KnownProviderType = (typeof PROVIDER_TYPES)[number]; +/** + * Outcome of a provider connection check (or a poll of one): confirmed + * connected, confirmed failed, or still running past the wait. An alias of the + * generic `TASK_OUTCOME` (see `types/tasks.ts`), which `pollTaskCompletion` also + * returns for the unrelated task it polls (organization/node deletion). Kept + * import-free (types only) so it can be used by test doubles and UI-only code + * without pulling in `lib/provider-helpers.ts`'s server-action dependencies. + */ +export const CONNECTION_CHECK_STATUS = TASK_OUTCOME; +export type ConnectionCheckStatus = TaskOutcome; + // Autocomplete for predefined + open for dynamic providers export type ProviderType = KnownProviderType | (string & {}); diff --git a/ui/types/scans.ts b/ui/types/scans.ts index 14b1c45a41..9d691aed7b 100644 --- a/ui/types/scans.ts +++ b/ui/types/scans.ts @@ -62,6 +62,8 @@ export interface ScanAttributes { completed_at: string | null; scheduled_at: string | null; next_scan_at: string | null; + /** Prowler Cloud only: true when the scan re-checked a few resources. */ + is_partial?: boolean; } export interface ScanRelationships { diff --git a/ui/types/tasks.ts b/ui/types/tasks.ts index a3ccb40b3f..838e1dc7a4 100644 --- a/ui/types/tasks.ts +++ b/ui/types/tasks.ts @@ -1,3 +1,18 @@ +/** + * Generic settle outcome for a polled async task: it succeeded, it failed, or the + * wait was exhausted with the task still running. `CONNECTION_CHECK_STATUS` in + * `types/providers.ts` is a connection-specific alias of this same shape, kept as + * its own export so a caller that only cares about a connection result does not + * have to name a generic task type to use it. + */ +export const TASK_OUTCOME = { + SUCCESS: "success", + FAILED: "failed", + PENDING: "pending", +} as const; + +export type TaskOutcome = (typeof TASK_OUTCOME)[keyof typeof TASK_OUTCOME]; + export type TaskState = | "available" | "scheduled" diff --git a/ui/types/tree.ts b/ui/types/tree.ts index d9b2f88c9f..a0ab5005f0 100644 --- a/ui/types/tree.ts +++ b/ui/types/tree.ts @@ -6,11 +6,14 @@ */ /** - * Status indicator for tree items after loading completes + * Status indicator for tree items after loading completes. `PENDING` is for an + * item whose outcome never arrived even though nothing is polling it any more -- + * distinct from `isLoading`, which is for an item actively being polled. */ export const TREE_ITEM_STATUS = { SUCCESS: "success", ERROR: "error", + PENDING: "pending", } as const; export type TreeItemStatus = @@ -33,7 +36,7 @@ export interface TreeDataItem { disabled?: boolean; /** Whether the item is in a loading state (shows spinner) */ isLoading?: boolean; - /** Status indicator shown after loading (success/error) */ + /** Status indicator shown after loading (success/error/pending) */ status?: TreeItemStatus; /** Optional error detail used by status icon tooltip */ errorMessage?: string; diff --git a/uv.lock b/uv.lock index fbaa60fb40..57e4816cd1 100644 --- a/uv.lock +++ b/uv.lock @@ -46,7 +46,7 @@ constraints = [ { name = "aliyun-log-fastpb", specifier = "==0.3.0" }, { name = "annotated-types", specifier = "==0.7.0" }, { name = "antlr4-python3-runtime", specifier = "==4.13.2" }, - { name = "anyio", specifier = "==4.13.0" }, + { name = "anyio", specifier = "==4.14.2" }, { name = "apscheduler", specifier = "==3.11.2" }, { name = "astroid", specifier = "==3.3.11" }, { name = "async-timeout", specifier = "==5.0.1" }, @@ -790,16 +790,16 @@ wheels = [ [[package]] name = "anyio" -version = "4.13.0" +version = "4.14.2" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "exceptiongroup", marker = "python_full_version < '3.11'" }, { name = "idna" }, { name = "typing-extensions", marker = "python_full_version < '3.13'" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/19/14/2c5dd9f512b66549ae92767a9c7b330ae88e1932ca57876909410251fe13/anyio-4.13.0.tar.gz", hash = "sha256:334b70e641fd2221c1505b3890c69882fe4a2df910cba14d97019b90b24439dc", size = 231622, upload-time = "2026-03-24T12:59:09.671Z" } +sdist = { url = "https://files.pythonhosted.org/packages/61/cc/a381afa6efea9f496eff839d4a6a1aed3bfafc7b3ab4b0d1b243a12573dd/anyio-4.14.2.tar.gz", hash = "sha256:cfa139f3ed1a23ee8f88a145ddb5ac7605b8bbfd8592baacd7ce3d8bb4313c7f", size = 260176, upload-time = "2026-07-12T20:29:07.082Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/da/42/e921fccf5015463e32a3cf6ee7f980a6ed0f395ceeaa45060b61d86486c2/anyio-4.13.0-py3-none-any.whl", hash = "sha256:08b310f9e24a9594186fd75b4f73f4a4152069e3853f1ed8bfbf58369f4ad708", size = 114353, upload-time = "2026-03-24T12:59:08.246Z" }, + { url = "https://files.pythonhosted.org/packages/da/35/f2287558c17e29fafc8ef3daf819bb9834061cfa43bff8014f7df7f63bdc/anyio-4.14.2-py3-none-any.whl", hash = "sha256:9f505dda5ac9f0c8309b5e8bd445a8c2bf7246f3ce950121e45ea15bc41d1494", size = 125813, upload-time = "2026-07-12T20:29:05.763Z" }, ] [[package]] @@ -3764,7 +3764,7 @@ wheels = [ [[package]] name = "prowler" -version = "5.43.0" +version = "5.44.0" source = { editable = "." } dependencies = [ { name = "alibabacloud-actiontrail20200706" },