Compare commits

...
4 Commits
23 changed files with 1297 additions and 867 deletions
+1 -1
View File
@@ -158,7 +158,7 @@ SENTRY_RELEASE=local
# REO_DEV_CLIENT_ID=
#### Prowler release version ####
NEXT_PUBLIC_PROWLER_RELEASE_VERSION=v5.40.0
NEXT_PUBLIC_PROWLER_RELEASE_VERSION=v5.40.1
# Social login credentials
SOCIAL_GOOGLE_OAUTH_CALLBACK_URL="${AUTH_URL}/api/auth/callback/google"
@@ -0,0 +1 @@
`POST /api/v1/mute-rules` now updates only each affected provider's latest completed scan and future scans, preventing historical reaggregation from flooding Celery queues
+1 -1
View File
@@ -71,7 +71,7 @@ name = "prowler-api"
package-mode = false
# Needed for the SDK compatibility
requires-python = ">=3.11,<3.13"
version = "1.41.0"
version = "1.41.1"
# Shared ruff baseline (kept in sync with mcp_server/pyproject.toml).
# target-version tracks this project's lowest supported Python.
+1 -1
View File
@@ -1,7 +1,7 @@
openapi: 3.0.3
info:
title: Prowler API
version: 1.41.0
version: 1.41.1
description: |-
Prowler API specification.
+18 -27
View File
@@ -18333,19 +18333,14 @@ class TestMuteRuleViewSet:
assert len(data) == 2
assert data[0]["id"] == str(mute_rules_fixture[first_index].id)
@patch("api.v1.views.chain")
@patch("api.v1.views.reaggregate_all_finding_group_summaries_task.si")
@patch("api.v1.views.mute_historical_findings_task.si")
@patch("api.v1.views.mute_findings_in_latest_scans_task.apply_async")
@patch("api.v1.views.transaction.on_commit", side_effect=lambda fn: fn())
def test_mute_rules_create_valid(
self,
_mock_on_commit,
mock_mute_signature,
mock_reaggregate_signature,
mock_chain,
mock_mute_task,
authenticated_client,
findings_fixture,
create_test_user,
):
"""Test creating a valid mute rule."""
finding_ids = [str(findings_fixture[0].id)]
@@ -18372,24 +18367,20 @@ class TestMuteRuleViewSet:
assert response_data["attributes"]["name"] == "New Mute Rule"
assert response_data["attributes"]["reason"] == "Security exception approved"
# Verify the finding was immediately muted
from api.models import Finding
finding = Finding.objects.get(id=findings_fixture[0].id)
assert finding.muted is True
assert finding.muted_at is not None
assert finding.muted_reason == "Security exception approved"
assert finding.muted is False
assert finding.muted_at is None
assert finding.muted_reason is None
# Verify background task chain was called: mute → reaggregate all
mock_mute_signature.assert_called_once()
mock_reaggregate_signature.assert_called_once()
mock_chain.assert_called_once_with(
mock_mute_signature.return_value,
mock_reaggregate_signature.return_value,
mock_mute_task.assert_called_once_with(
kwargs={
"tenant_id": str(finding.tenant_id),
"mute_rule_id": response_data["id"],
"provider_ids": [str(finding.scan.provider_id)],
}
)
mock_chain.return_value.apply_async.assert_called_once()
@patch("tasks.tasks.mute_historical_findings_task.apply_async")
@patch("api.v1.views.mute_findings_in_latest_scans_task.apply_async")
def test_mute_rules_create_converts_finding_ids_to_uids(
self,
mock_task,
@@ -18425,7 +18416,7 @@ class TestMuteRuleViewSet:
]
assert set(mute_rule.finding_uids) == set(expected_uids)
@patch("tasks.tasks.mute_historical_findings_task.apply_async")
@patch("api.v1.views.mute_findings_in_latest_scans_task.apply_async")
def test_mute_rules_deduplicates_uids(
self,
mock_task,
@@ -18492,10 +18483,10 @@ class TestMuteRuleViewSet:
finding1.refresh_from_db()
finding2.refresh_from_db()
assert finding1.muted is True
assert finding2.muted is True
assert finding1.muted is False
assert finding2.muted is False
@patch("tasks.tasks.mute_historical_findings_task.apply_async")
@patch("api.v1.views.mute_findings_in_latest_scans_task.apply_async")
def test_mute_rules_create_overlap_detection_active(
self,
mock_task,
@@ -18528,7 +18519,7 @@ class TestMuteRuleViewSet:
"already muted" in error_detail.lower() or "overlap" in error_detail.lower()
)
@patch("tasks.tasks.mute_historical_findings_task.apply_async")
@patch("api.v1.views.mute_findings_in_latest_scans_task.apply_async")
def test_mute_rules_create_no_overlap_with_inactive(
self,
mock_task,
@@ -18584,7 +18575,7 @@ class TestMuteRuleViewSet:
== "/data/attributes/finding_ids"
)
@patch("tasks.tasks.mute_historical_findings_task.apply_async")
@patch("api.v1.views.mute_findings_in_latest_scans_task.apply_async")
def test_mute_rules_create_invalid_finding_ids(
self, mock_task, authenticated_client
):
+18 -27
View File
@@ -244,7 +244,6 @@ from api.v1.serializers import (
UserUpdateSerializer,
)
from botocore.exceptions import ClientError, NoCredentialsError, ParamValidationError
from celery import chain
from celery.result import AsyncResult
from config.custom_logging import BackendLogger
from config.env import env
@@ -342,8 +341,7 @@ from tasks.tasks import (
enqueue_scan_execution_on_commit,
get_active_provider_scan,
jira_integration_task,
mute_historical_findings_task,
reaggregate_all_finding_group_summaries_task,
mute_findings_in_latest_scans_task,
refresh_lighthouse_provider_models_task,
)
@@ -7551,35 +7549,28 @@ class MuteRuleViewSet(BaseRLSViewSet):
serializer = self.get_serializer(data=request.data)
serializer.is_valid(raise_exception=True)
# Create the mute rule
tenant_id = str(request.tenant_id)
finding_ids = serializer.validated_data["finding_ids"]
provider_ids = list(
dict.fromkeys(
Finding.all_objects.filter(
id__in=finding_ids, tenant_id=tenant_id
).values_list("scan__provider_id", flat=True)
)
)
mute_rule = serializer.save()
tenant_id = str(request.tenant_id)
finding_ids = request.data.get("finding_ids", [])
# Immediately mute the selected findings
Finding.all_objects.filter(
id__in=finding_ids, tenant_id=tenant_id, muted=False
).update(
muted=True,
muted_at=mute_rule.inserted_at,
muted_reason=mute_rule.reason,
)
# Launch background task for historical muting + reaggregation
transaction.on_commit(
lambda: chain(
mute_historical_findings_task.si(
tenant_id=tenant_id,
mute_rule_id=str(mute_rule.id),
),
reaggregate_all_finding_group_summaries_task.si(
tenant_id=tenant_id,
),
).apply_async()
lambda: mute_findings_in_latest_scans_task.apply_async(
kwargs={
"tenant_id": tenant_id,
"mute_rule_id": str(mute_rule.id),
"provider_ids": [str(provider_id) for provider_id in provider_ids],
}
)
)
# Return the created mute rule
serializer = self.get_serializer(mute_rule)
return Response(
data=serializer.data,
+86 -45
View File
@@ -1,63 +1,104 @@
from collections.abc import Iterable
from api.db_utils import rls_transaction
from api.models import Finding, MuteRule
from api.models import Finding, MuteRule, Scan, StateChoices
from celery.utils.log import get_task_logger
from config.django.base import DJANGO_FINDINGS_BATCH_SIZE
from tasks.utils import batched
logger = get_task_logger(__name__)
def mute_historical_findings(tenant_id: str, mute_rule_id: str):
"""
Mute historical findings that match the given mute rule.
def _mute_findings_for_rule(
*,
tenant_id: str,
scan_id: str,
finding_uids: Iterable[str],
muted_at,
muted_reason: str,
) -> int:
finding_uids = list(finding_uids)
if not finding_uids:
return 0
This function processes findings in batches, updating their muted status
and adding the mute reason.
return Finding.all_objects.filter(
tenant_id=tenant_id,
scan_id=scan_id,
uid__in=finding_uids,
muted=False,
).update(
muted=True,
muted_at=muted_at,
muted_reason=muted_reason,
)
Args:
tenant_id (str): The tenant ID for RLS context
mute_rule_id (str): The ID of the mute rule to apply
Returns:
dict: Summary of the muting operation with findings_muted count
"""
findings_muted_count = 0
def mute_findings_in_latest_scans(
tenant_id: str, mute_rule_id: str, provider_ids: list[str]
) -> dict:
"""Apply a mute rule to the latest completed scan of each provider."""
provider_ids = list(dict.fromkeys(provider_ids))
# Get the list of UIDs to mute and the reason
with rls_transaction(tenant_id):
mute_rule = MuteRule.objects.get(id=mute_rule_id, tenant_id=tenant_id)
finding_uids = mute_rule.finding_uids
mute_reason = mute_rule.reason
muted_at = mute_rule.inserted_at
# Query findings that match the UIDs and are not already muted
with rls_transaction(tenant_id):
findings_to_mute = Finding.objects.filter(
tenant_id=tenant_id, uid__in=finding_uids, muted=False
)
total_findings = findings_to_mute.count()
logger.info(
f"Processing {total_findings} findings for mute rule {mute_rule_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)
)
if total_findings > 0:
for batch, is_last in batched(
findings_to_mute.iterator(), DJANGO_FINDINGS_BATCH_SIZE
):
batch_ids = [f.id for f in batch]
updated_count = Finding.all_objects.filter(
id__in=batch_ids, tenant_id=tenant_id
).update(
muted=True,
muted_at=muted_at,
muted_reason=mute_reason,
)
findings_muted_count += updated_count
logger.info(f"Muted {findings_muted_count} findings for rule {mute_rule_id}")
changed_scan_ids = []
findings_muted = 0
for scan_id in latest_scans:
updated = _mute_findings_for_rule(
tenant_id=tenant_id,
scan_id=str(scan_id),
finding_uids=mute_rule.finding_uids,
muted_at=mute_rule.inserted_at,
muted_reason=mute_rule.reason,
)
if updated:
findings_muted += updated
changed_scan_ids.append(str(scan_id))
logger.info(
"Muted %d findings in %d latest scans for rule %s",
findings_muted,
len(changed_scan_ids),
mute_rule_id,
)
return {
"findings_muted": findings_muted_count,
"findings_muted": findings_muted,
"rule_id": mute_rule_id,
"scan_ids": changed_scan_ids,
}
def reconcile_scan_mute_rules(tenant_id: str, scan_id: str) -> dict:
"""Apply the current enabled mute rules to one completed scan."""
findings_muted = 0
with rls_transaction(tenant_id):
mute_rules = MuteRule.objects.filter(tenant_id=tenant_id, enabled=True).values(
"finding_uids", "reason", "inserted_at"
)
for mute_rule in mute_rules:
findings_muted += _mute_findings_for_rule(
tenant_id=tenant_id,
scan_id=scan_id,
finding_uids=mute_rule["finding_uids"],
muted_at=mute_rule["inserted_at"],
muted_reason=mute_rule["reason"],
)
logger.info(
"Reconciled mute rules for scan %s; muted %d findings",
scan_id,
findings_muted,
)
return {"findings_muted": findings_muted, "scan_id": str(scan_id)}
+45 -99
View File
@@ -73,7 +73,10 @@ from tasks.jobs.lighthouse_providers import (
check_lighthouse_provider_connection,
refresh_lighthouse_provider_models,
)
from tasks.jobs.muting import mute_historical_findings
from tasks.jobs.muting import (
mute_findings_in_latest_scans,
reconcile_scan_mute_rules,
)
from tasks.jobs.orphan_recovery import reconcile_orphans
from tasks.jobs.report import (
STALE_TMP_OUTPUT_MAX_AGE_HOURS,
@@ -526,6 +529,7 @@ def perform_scan_task(
provider_id=provider_id,
checks_to_execute=checks_to_execute,
)
reconcile_scan_mute_rules(tenant_id, scan_id)
_perform_scan_complete_tasks(tenant_id, scan_id, provider_id)
return result
finally:
@@ -635,6 +639,7 @@ def perform_scheduled_scan_task(self, tenant_id: str, provider_id: str):
scan_id=str(scan_instance.id),
provider_id=provider_id,
)
reconcile_scan_mute_rules(tenant_id, str(scan_instance.id))
_perform_scan_complete_tasks(tenant_id, str(scan_instance.id), provider_id)
return result
finally:
@@ -1188,85 +1193,48 @@ def aggregate_finding_group_summaries_task(tenant_id: str, scan_id: str):
return aggregate_finding_group_summaries(tenant_id=tenant_id, scan_id=scan_id)
@shared_task(
base=RLSTask, name="reaggregate-all-finding-group-summaries", queue="overview"
)
@set_tenant(keep_tenant=True)
def reaggregate_all_finding_group_summaries_task(tenant_id: str):
"""Reaggregate every pre-aggregated summary table for this tenant.
def _dispatch_scan_summary_reaggregation(tenant_id: str, scan_ids: list[str]) -> None:
if not scan_ids:
return
Mirrors the unbounded scope of `mute_historical_findings_task`: that task
rewrites every Finding row whose UID matches a mute rule, with no time
limit. To keep the pre-aggregated tables consistent with that update,
this task re-runs the same per-scan aggregation pipeline that scan
completion runs on the latest completed scan of every (provider, day)
pair, rebuilding the tables that power the read endpoints:
- `ScanSummary` and `DailySeveritySummary` -> `/overviews/findings`,
`/overviews/findings-severity`, `/overviews/services`.
- `FindingGroupDailySummary` -> `/finding-groups` and
`/finding-groups/latest`.
- `ScanGroupSummary` -> `/overviews/resource-groups` (resource
inventory).
- `ScanCategorySummary` -> `/overviews/categories`.
- `AttackSurfaceOverview` -> `/overviews/attack-surfaces`.
Per-scan pipelines are dispatched in parallel via a Celery group so
wallclock scales with the worker pool.
"""
completed_scans = list(
Scan.objects.filter(
tenant_id=tenant_id,
state=StateChoices.COMPLETED,
completed_at__isnull=False,
)
.order_by("-completed_at")
.values("id", "completed_at", "provider_id")
logger.info(
"Reaggregating overview/finding summaries for %d latest scans",
len(scan_ids),
)
# Keep the latest scan per (provider, day) pair so the daily summary row
# the aggregator writes is the most recent snapshot of that day for that
# provider. Iterating from most recent to oldest means the first scan we
# see for a given key wins.
latest_scans: dict[tuple, str] = {}
for scan in completed_scans:
key = (scan["provider_id"], scan["completed_at"].date())
if key not in latest_scans:
latest_scans[key] = str(scan["id"])
scan_ids = list(latest_scans.values())
if scan_ids:
logger.info(
"Reaggregating overview/finding summaries for %d scans (provider x day)",
len(scan_ids),
)
# DailySeveritySummary reads from ScanSummary, so ScanSummary must be
# recomputed first; the other aggregators read Finding directly and
# can run in parallel with the severity step.
group(
chain(
perform_scan_summary_task.si(tenant_id=tenant_id, scan_id=scan_id),
group(
aggregate_daily_severity_task.si(
tenant_id=tenant_id, scan_id=scan_id
),
aggregate_finding_group_summaries_task.si(
tenant_id=tenant_id, scan_id=scan_id
),
aggregate_scan_resource_group_summaries_task.si(
tenant_id=tenant_id, scan_id=scan_id
),
aggregate_scan_category_summaries_task.si(
tenant_id=tenant_id, scan_id=scan_id
),
aggregate_attack_surface_task.si(
tenant_id=tenant_id, scan_id=scan_id
),
group(
chain(
perform_scan_summary_task.si(tenant_id=tenant_id, scan_id=scan_id),
group(
aggregate_daily_severity_task.si(tenant_id=tenant_id, scan_id=scan_id),
aggregate_finding_group_summaries_task.si(
tenant_id=tenant_id, scan_id=scan_id
),
)
for scan_id in scan_ids
).apply_async()
return {"scans_reaggregated": len(scan_ids)}
aggregate_scan_resource_group_summaries_task.si(
tenant_id=tenant_id, scan_id=scan_id
),
aggregate_scan_category_summaries_task.si(
tenant_id=tenant_id, scan_id=scan_id
),
aggregate_attack_surface_task.si(tenant_id=tenant_id, scan_id=scan_id),
),
)
for scan_id in scan_ids
).apply_async()
@shared_task(base=RLSTask, name="findings-mute-latest-scans", queue="overview")
@set_tenant(keep_tenant=True)
def mute_findings_in_latest_scans_task(
tenant_id: str, mute_rule_id: str, provider_ids: list[str]
):
"""Apply a mute rule to current scans and rebuild only changed summaries."""
result = mute_findings_in_latest_scans(
tenant_id=tenant_id,
mute_rule_id=mute_rule_id,
provider_ids=provider_ids,
)
_dispatch_scan_summary_reaggregation(tenant_id, result["scan_ids"])
return result
@shared_task(base=RLSTask, name="lighthouse-connection-check")
@@ -1467,25 +1435,3 @@ def generate_compliance_reports_task(tenant_id: str, scan_id: str, provider_id:
generate_csa=True,
generate_cis=True,
)
@shared_task(name="findings-mute-historical")
def mute_historical_findings_task(tenant_id: str, mute_rule_id: str):
"""
Background task to mute all historical findings matching a mute rule.
This task processes findings in batches to avoid memory issues with large datasets.
It updates the Finding.muted, Finding.muted_at, and Finding.muted_reason fields
for all findings whose UID is in the mute rule's finding_uids list.
Args:
tenant_id (str): The tenant ID for RLS context.
mute_rule_id (str): The primary key of the MuteRule to apply.
Returns:
dict: A dictionary containing:
- 'findings_muted' (int): Total number of findings muted.
- 'rule_id' (str): The mute rule ID.
- 'status' (str): Final status ('completed').
"""
return mute_historical_findings(tenant_id, mute_rule_id)
+176 -502
View File
@@ -1,531 +1,205 @@
from datetime import UTC, datetime
from datetime import UTC, datetime, timedelta
from uuid import uuid4
import pytest
from api.models import Finding, MuteRule
from django.core.exceptions import ObjectDoesNotExist
from api.models import Finding, MuteRule, Scan, StateChoices
from prowler.lib.check.models import Severity
from prowler.lib.outputs.finding import Status
from tasks.jobs.muting import mute_historical_findings
from tasks.jobs.muting import (
mute_findings_in_latest_scans,
reconcile_scan_mute_rules,
)
def _create_finding(scan: Scan, uid: str) -> Finding:
return Finding.objects.create(
tenant_id=scan.tenant_id,
uid=uid,
scan=scan,
status=Status.FAIL,
status_extended="Test finding",
impact=Severity.high,
severity=Severity.high,
raw_result={},
check_id="test_check",
check_metadata={"CheckId": "test_check"},
muted=False,
)
def _create_mute_rule(tenant_id, user, finding_uids, *, enabled=True) -> MuteRule:
return MuteRule.objects.create(
tenant_id=tenant_id,
name=f"Mute rule {uuid4()}",
reason="Approved exception",
enabled=enabled,
created_by=user,
finding_uids=finding_uids,
)
@pytest.mark.django_db
class TestMuteHistoricalFindings:
"""
Test suite for the mute_historical_findings function.
class TestMuteFindingsInLatestScans:
def test_mutes_latest_scan_and_leaves_older_scan_unchanged(
self, scans_fixture, create_test_user
):
latest_scan = scans_fixture[0]
older_scan = Scan.objects.create(
tenant_id=latest_scan.tenant_id,
provider=latest_scan.provider,
name="Older scan",
trigger=Scan.TriggerChoices.MANUAL,
state=StateChoices.COMPLETED,
started_at=datetime.now(UTC) - timedelta(days=1),
completed_at=datetime.now(UTC) - timedelta(days=1),
)
uid = "latest-scan-only"
older_finding = _create_finding(older_scan, uid)
latest_finding = _create_finding(latest_scan, uid)
mute_rule = _create_mute_rule(latest_scan.tenant_id, create_test_user, [uid])
This class tests the batch processing of findings to update their muted status
based on MuteRule criteria.
"""
result = mute_findings_in_latest_scans(
str(latest_scan.tenant_id),
str(mute_rule.id),
[str(latest_scan.provider_id)],
)
@pytest.fixture(scope="function")
def test_user(self, create_test_user):
"""Create a test user for mute rule creation."""
return create_test_user
older_finding.refresh_from_db()
latest_finding.refresh_from_db()
assert older_finding.muted is False
assert latest_finding.muted is True
assert latest_finding.muted_at == mute_rule.inserted_at
assert latest_finding.muted_reason == mute_rule.reason
assert result == {
"findings_muted": 1,
"rule_id": str(mute_rule.id),
"scan_ids": [str(latest_scan.id)],
}
@pytest.fixture(scope="function")
def mute_rule_with_findings(self, tenants_fixture, findings_fixture, test_user):
"""
Create a mute rule that targets the first finding in the fixture.
"""
def test_mutes_one_latest_scan_per_provider(self, scans_fixture, create_test_user):
first_scan, second_scan, _ = scans_fixture
uid = "shared-selected-uid"
first_finding = _create_finding(first_scan, uid)
second_finding = _create_finding(second_scan, uid)
mute_rule = _create_mute_rule(first_scan.tenant_id, create_test_user, [uid])
result = mute_findings_in_latest_scans(
str(first_scan.tenant_id),
str(mute_rule.id),
[str(first_scan.provider_id), str(second_scan.provider_id)],
)
first_finding.refresh_from_db()
second_finding.refresh_from_db()
assert first_finding.muted is True
assert second_finding.muted is True
assert result["findings_muted"] == 2
assert set(result["scan_ids"]) == {str(first_scan.id), str(second_scan.id)}
def test_provider_without_completed_scan_does_nothing(
self, tenants_fixture, provider_factory, create_test_user
):
tenant = tenants_fixture[0]
finding = findings_fixture[0]
mute_rule = MuteRule.objects.create(
tenant_id=tenant.id,
name="Test Mute Rule",
reason="Testing mute functionality",
enabled=True,
created_by=test_user,
finding_uids=[finding.uid],
provider = provider_factory()
mute_rule = _create_mute_rule(
tenant.id, create_test_user, ["future-scan-finding"]
)
return mute_rule
result = mute_findings_in_latest_scans(
str(tenant.id), str(mute_rule.id), [str(provider.id)]
)
@pytest.fixture(scope="function")
def mute_rule_multiple_findings(self, scans_fixture, test_user):
"""
Create multiple unmuted findings and a mute rule targeting all of them.
"""
assert result == {
"findings_muted": 0,
"rule_id": str(mute_rule.id),
"scan_ids": [],
}
def test_retry_does_not_report_changed_scans_twice(
self, scans_fixture, create_test_user
):
scan = scans_fixture[0]
tenant_id = scan.tenant_id
# Create 5 unmuted findings
finding_uids = []
for i in range(5):
finding = Finding.objects.create(
tenant_id=tenant_id,
uid=f"test_finding_uid_mute_{i}",
scan=scan,
status=Status.FAIL,
status_extended=f"Test status {i}",
impact=Severity.high,
severity=Severity.high,
raw_result={
"status": Status.FAIL,
"impact": Severity.high,
"severity": Severity.high,
},
check_id=f"test_check_id_{i}",
check_metadata={
"CheckId": f"test_check_id_{i}",
"Description": f"Test description {i}",
},
muted=False,
)
finding_uids.append(finding.uid)
# Create mute rule targeting all findings
mute_rule = MuteRule.objects.create(
tenant_id=tenant_id,
name="Test Multiple Findings Mute Rule",
reason="Testing batch muting",
enabled=True,
created_by=test_user,
finding_uids=finding_uids,
finding = _create_finding(scan, "idempotent-mute")
mute_rule = _create_mute_rule(scan.tenant_id, create_test_user, [finding.uid])
args = (
str(scan.tenant_id),
str(mute_rule.id),
[str(scan.provider_id)],
)
return mute_rule, finding_uids
first_result = mute_findings_in_latest_scans(*args)
second_result = mute_findings_in_latest_scans(*args)
@pytest.fixture(scope="function")
def mute_rule_already_muted(self, findings_fixture, test_user):
"""
Create a mute rule that targets an already-muted finding.
"""
tenant_id = findings_fixture[1].tenant_id
already_muted_finding = findings_fixture[1]
assert first_result["scan_ids"] == [str(scan.id)]
assert second_result["findings_muted"] == 0
assert second_result["scan_ids"] == []
mute_rule = MuteRule.objects.create(
tenant_id=tenant_id,
name="Test Already Muted Rule",
reason="Testing already muted findings",
enabled=True,
created_by=test_user,
finding_uids=[already_muted_finding.uid],
def test_does_not_cross_tenant_boundary(
self, tenants_fixture, provider_factory, create_test_user
):
tenant = tenants_fixture[0]
other_tenant = tenants_fixture[2]
other_provider = provider_factory(tenant=other_tenant)
other_scan = Scan.objects.create(
tenant_id=other_tenant.id,
provider=other_provider,
name="Other tenant scan",
trigger=Scan.TriggerChoices.MANUAL,
state=StateChoices.COMPLETED,
started_at=datetime.now(UTC),
completed_at=datetime.now(UTC),
)
other_finding = _create_finding(other_scan, "tenant-isolated-uid")
mute_rule = _create_mute_rule(tenant.id, create_test_user, [other_finding.uid])
result = mute_findings_in_latest_scans(
str(tenant.id), str(mute_rule.id), [str(other_provider.id)]
)
return mute_rule
other_finding.refresh_from_db()
assert other_finding.muted is False
assert result["scan_ids"] == []
@pytest.fixture(scope="function")
def mute_rule_mixed_findings(self, scans_fixture, test_user):
"""
Create a mute rule with a mix of muted and unmuted findings.
"""
def test_nonexistent_rule_raises(self, tenants_fixture):
with pytest.raises(MuteRule.DoesNotExist):
mute_findings_in_latest_scans(str(tenants_fixture[0].id), str(uuid4()), [])
@pytest.mark.django_db
class TestReconcileScanMuteRules:
def test_applies_only_enabled_rules_to_requested_scan(
self, scans_fixture, create_test_user
):
scan = scans_fixture[0]
tenant_id = scan.tenant_id
# Create 3 unmuted findings
unmuted_uids = []
for i in range(3):
finding = Finding.objects.create(
tenant_id=tenant_id,
uid=f"unmuted_finding_{i}",
scan=scan,
status=Status.FAIL,
status_extended=f"Unmuted status {i}",
impact=Severity.medium,
severity=Severity.medium,
raw_result={
"status": Status.FAIL,
"impact": Severity.medium,
"severity": Severity.medium,
},
check_id=f"unmuted_check_{i}",
check_metadata={
"CheckId": f"unmuted_check_{i}",
"Description": f"Unmuted description {i}",
},
muted=False,
)
unmuted_uids.append(finding.uid)
# Create 2 already muted findings
muted_uids = []
for i in range(2):
finding = Finding.objects.create(
tenant_id=tenant_id,
uid=f"muted_finding_{i}",
scan=scan,
status=Status.FAIL,
status_extended=f"Muted status {i}",
impact=Severity.low,
severity=Severity.low,
raw_result={
"status": Status.FAIL,
"impact": Severity.low,
"severity": Severity.low,
},
check_id=f"muted_check_{i}",
check_metadata={
"CheckId": f"muted_check_{i}",
"Description": f"Muted description {i}",
},
muted=True,
muted_at=datetime.now(UTC),
muted_reason="Already muted",
)
muted_uids.append(finding.uid)
# Create mute rule targeting all findings
all_uids = unmuted_uids + muted_uids
mute_rule = MuteRule.objects.create(
tenant_id=tenant_id,
name="Test Mixed Findings Rule",
reason="Testing mixed muted/unmuted findings",
enabled=True,
created_by=test_user,
finding_uids=all_uids,
active_finding = _create_finding(scan, "active-rule-uid")
disabled_finding = _create_finding(scan, "disabled-rule-uid")
active_rule = _create_mute_rule(
scan.tenant_id, create_test_user, [active_finding.uid]
)
return mute_rule, unmuted_uids, muted_uids
@pytest.fixture(scope="function")
def mute_rule_batch_test(self, scans_fixture, test_user):
"""
Create enough findings to test batch processing (>1000 for default batch size).
"""
scan = scans_fixture[0]
tenant_id = scan.tenant_id
# Create 1500 findings to exceed default batch size of 1000
finding_uids = []
for i in range(1500):
finding = Finding.objects.create(
tenant_id=tenant_id,
uid=f"batch_test_finding_{i}",
scan=scan,
status=Status.FAIL,
status_extended=f"Batch test status {i}",
impact=Severity.critical,
severity=Severity.critical,
raw_result={
"status": Status.FAIL,
"impact": Severity.critical,
"severity": Severity.critical,
},
check_id=f"batch_test_check_{i}",
check_metadata={
"CheckId": f"batch_test_check_{i}",
"Description": f"Batch test description {i}",
},
muted=False,
)
finding_uids.append(finding.uid)
# Create mute rule targeting all findings
mute_rule = MuteRule.objects.create(
tenant_id=tenant_id,
name="Test Batch Processing Rule",
reason="Testing batch processing functionality",
enabled=True,
created_by=test_user,
finding_uids=finding_uids,
_create_mute_rule(
scan.tenant_id,
create_test_user,
[disabled_finding.uid],
enabled=False,
)
return mute_rule, finding_uids
def test_mute_historical_findings_single_finding(
self, mute_rule_with_findings, findings_fixture
):
"""
Test muting a single historical finding.
"""
mute_rule = mute_rule_with_findings
tenant_id = str(mute_rule.tenant_id)
finding = findings_fixture[0]
# Ensure the finding is not muted before execution
finding.refresh_from_db()
assert finding.muted is False
assert finding.muted_at is None
assert finding.muted_reason is None
# Execute the muting function
result = mute_historical_findings(tenant_id, str(mute_rule.id))
# Verify return value
assert result["findings_muted"] == 1
assert result["rule_id"] == str(mute_rule.id)
# Verify the finding was muted
finding.refresh_from_db()
assert finding.muted is True
assert finding.muted_at == mute_rule.inserted_at
assert finding.muted_reason == mute_rule.reason
def test_mute_historical_findings_multiple_findings(
self, mute_rule_multiple_findings
):
"""
Test muting multiple historical findings.
"""
mute_rule, finding_uids = mute_rule_multiple_findings
tenant_id = str(mute_rule.tenant_id)
# Verify all findings are unmuted
findings = Finding.objects.filter(tenant_id=tenant_id, uid__in=finding_uids)
assert findings.count() == 5
for finding in findings:
assert finding.muted is False
# Execute the muting function
result = mute_historical_findings(tenant_id, str(mute_rule.id))
# Verify return value
assert result["findings_muted"] == 5
assert result["rule_id"] == str(mute_rule.id)
# Verify all findings were muted
findings = Finding.objects.filter(tenant_id=tenant_id, uid__in=finding_uids)
for finding in findings:
assert finding.muted is True
assert finding.muted_at == mute_rule.inserted_at
assert finding.muted_reason == mute_rule.reason
def test_mute_historical_findings_already_muted(
self, mute_rule_already_muted, findings_fixture
):
"""
Test that already-muted findings are not counted or updated.
"""
mute_rule = mute_rule_already_muted
tenant_id = str(mute_rule.tenant_id)
finding = findings_fixture[1]
# Verify the finding is already muted
finding.refresh_from_db()
assert finding.muted is True
original_muted_at = finding.muted_at
original_muted_reason = finding.muted_reason
# Execute the muting function
result = mute_historical_findings(tenant_id, str(mute_rule.id))
# Verify no findings were muted
assert result["findings_muted"] == 0
assert result["rule_id"] == str(mute_rule.id)
# Verify the finding's mute status did not change
finding.refresh_from_db()
assert finding.muted is True
assert finding.muted_at == original_muted_at
assert finding.muted_reason == original_muted_reason
def test_mute_historical_findings_mixed_status(self, mute_rule_mixed_findings):
"""
Test muting when some findings are already muted and others are not.
"""
mute_rule, unmuted_uids, muted_uids = mute_rule_mixed_findings
tenant_id = str(mute_rule.tenant_id)
# Execute the muting function
result = mute_historical_findings(tenant_id, str(mute_rule.id))
# Verify only unmuted findings were counted
assert result["findings_muted"] == 3
assert result["rule_id"] == str(mute_rule.id)
# Verify unmuted findings are now muted
unmuted_findings = Finding.objects.filter(
tenant_id=tenant_id, uid__in=unmuted_uids
older_scan = Scan.objects.create(
tenant_id=scan.tenant_id,
provider=scan.provider,
name="Older matching scan",
trigger=Scan.TriggerChoices.MANUAL,
state=StateChoices.COMPLETED,
started_at=datetime.now(UTC) - timedelta(days=1),
completed_at=datetime.now(UTC) - timedelta(days=1),
)
for finding in unmuted_findings:
assert finding.muted is True
assert finding.muted_at == mute_rule.inserted_at
assert finding.muted_reason == mute_rule.reason
older_finding = _create_finding(older_scan, active_finding.uid)
# Verify already-muted findings remained unchanged
already_muted_findings = Finding.objects.filter(
tenant_id=tenant_id, uid__in=muted_uids
)
for finding in already_muted_findings:
assert finding.muted is True
assert finding.muted_reason == "Already muted"
result = reconcile_scan_mute_rules(str(scan.tenant_id), str(scan.id))
def test_mute_historical_findings_nonexistent_rule(self, tenants_fixture):
"""
Test that a nonexistent mute rule raises ObjectDoesNotExist.
"""
tenant_id = str(tenants_fixture[0].id)
nonexistent_rule_id = str(uuid4())
with pytest.raises(ObjectDoesNotExist):
mute_historical_findings(tenant_id, nonexistent_rule_id)
def test_mute_historical_findings_no_matching_findings(
self, tenants_fixture, test_user
):
"""
Test muting when no findings match the rule's UIDs.
"""
tenant_id = str(tenants_fixture[0].id)
# Create a mute rule with non-existent finding UIDs
mute_rule = MuteRule.objects.create(
tenant_id=tenant_id,
name="Test No Match Rule",
reason="Testing no matching findings",
enabled=True,
created_by=test_user,
finding_uids=[
"nonexistent_uid_1",
"nonexistent_uid_2",
"nonexistent_uid_3",
],
)
# Execute the muting function
result = mute_historical_findings(tenant_id, str(mute_rule.id))
# Verify no findings were muted
assert result["findings_muted"] == 0
assert result["rule_id"] == str(mute_rule.id)
def test_mute_historical_findings_batch_processing(self, mute_rule_batch_test):
"""
Test that large numbers of findings are processed in batches correctly.
"""
mute_rule, finding_uids = mute_rule_batch_test
tenant_id = str(mute_rule.tenant_id)
# Verify all findings exist and are unmuted
findings = Finding.objects.filter(tenant_id=tenant_id, uid__in=finding_uids)
assert findings.count() == 1500
for finding in findings:
assert finding.muted is False
# Execute the muting function
result = mute_historical_findings(tenant_id, str(mute_rule.id))
# Verify return value
assert result["findings_muted"] == 1500
assert result["rule_id"] == str(mute_rule.id)
# Verify all findings were muted
findings = Finding.objects.filter(tenant_id=tenant_id, uid__in=finding_uids)
for finding in findings:
assert finding.muted is True
assert finding.muted_at == mute_rule.inserted_at
assert finding.muted_reason == mute_rule.reason
def test_mute_historical_findings_preserves_muted_at_timestamp(
self, mute_rule_with_findings, findings_fixture
):
"""
Test that muted_at is set to the rule's inserted_at, not the current time.
"""
mute_rule = mute_rule_with_findings
tenant_id = str(mute_rule.tenant_id)
finding = findings_fixture[0]
# Execute the muting function
result = mute_historical_findings(tenant_id, str(mute_rule.id))
# Verify the finding was muted
assert result["findings_muted"] == 1
# Verify muted_at matches the rule's inserted_at timestamp
finding.refresh_from_db()
assert finding.muted_at == mute_rule.inserted_at
assert finding.muted_at is not None
def test_mute_historical_findings_partial_match(self, scans_fixture, test_user):
"""
Test muting when only some of the rule's UIDs exist as findings.
"""
scan = scans_fixture[0]
tenant_id = str(scan.tenant_id)
# Create 3 findings
existing_uids = []
for i in range(3):
finding = Finding.objects.create(
tenant_id=tenant_id,
uid=f"partial_match_finding_{i}",
scan=scan,
status=Status.FAIL,
status_extended=f"Partial match status {i}",
impact=Severity.high,
severity=Severity.high,
raw_result={
"status": Status.FAIL,
"impact": Severity.high,
"severity": Severity.high,
},
check_id=f"partial_match_check_{i}",
check_metadata={
"CheckId": f"partial_match_check_{i}",
"Description": f"Partial match description {i}",
},
muted=False,
)
existing_uids.append(finding.uid)
# Create a mute rule with both existing and non-existing UIDs
all_uids = existing_uids + [
"nonexistent_uid_1",
"nonexistent_uid_2",
]
mute_rule = MuteRule.objects.create(
tenant_id=tenant_id,
name="Test Partial Match Rule",
reason="Testing partial matching",
enabled=True,
created_by=test_user,
finding_uids=all_uids,
)
# Execute the muting function
result = mute_historical_findings(tenant_id, str(mute_rule.id))
# Verify only existing findings were muted
assert result["findings_muted"] == 3
assert result["rule_id"] == str(mute_rule.id)
# Verify the existing findings were muted
findings = Finding.objects.filter(tenant_id=tenant_id, uid__in=existing_uids)
assert findings.count() == 3
for finding in findings:
assert finding.muted is True
assert finding.muted_at == mute_rule.inserted_at
assert finding.muted_reason == mute_rule.reason
def test_mute_historical_findings_empty_uids(self, tenants_fixture, test_user):
"""
Test muting when the rule has an empty finding_uids array.
"""
tenant_id = str(tenants_fixture[0].id)
# Create a mute rule with empty finding_uids
mute_rule = MuteRule.objects.create(
tenant_id=tenant_id,
name="Test Empty UIDs Rule",
reason="Testing empty UIDs",
enabled=True,
created_by=test_user,
finding_uids=[],
)
# Execute the muting function
result = mute_historical_findings(tenant_id, str(mute_rule.id))
# Verify no findings were muted
assert result["findings_muted"] == 0
assert result["rule_id"] == str(mute_rule.id)
def test_mute_historical_findings_return_format(self, mute_rule_with_findings):
"""
Test that the return value has the correct format and fields.
"""
mute_rule = mute_rule_with_findings
tenant_id = str(mute_rule.tenant_id)
result = mute_historical_findings(tenant_id, str(mute_rule.id))
# Verify return value structure
assert isinstance(result, dict)
assert "findings_muted" in result
assert "rule_id" in result
assert isinstance(result["findings_muted"], int)
assert isinstance(result["rule_id"], str)
assert result["rule_id"] == str(mute_rule.id)
active_finding.refresh_from_db()
disabled_finding.refresh_from_db()
older_finding.refresh_from_db()
assert active_finding.muted is True
assert active_finding.muted_at == active_rule.inserted_at
assert disabled_finding.muted is False
assert older_finding.muted is False
assert result == {"findings_muted": 1, "scan_id": str(scan.id)}
+69 -124
View File
@@ -1,6 +1,6 @@
import uuid
from contextlib import contextmanager
from datetime import UTC, datetime, timedelta
from datetime import UTC, datetime
from unittest.mock import MagicMock, patch
import httpx
@@ -33,10 +33,10 @@ from tasks.tasks import (
check_integrations_task,
check_lighthouse_provider_connection_task,
generate_outputs_task,
mute_findings_in_latest_scans_task,
perform_attack_paths_scan_task,
perform_scan_task,
perform_scheduled_scan_task,
reaggregate_all_finding_group_summaries_task,
refresh_lighthouse_provider_models_task,
s3_integration_task,
security_hub_integration_task,
@@ -2959,6 +2959,7 @@ class TestPerformScheduledScanTask:
with (
patch("tasks.tasks.perform_prowler_scan", side_effect=_complete_scan),
patch("tasks.tasks._perform_scan_complete_tasks"),
patch("tasks.tasks.reconcile_scan_mute_rules") as mock_reconcile,
self._override_task_request(perform_scheduled_scan_task, id=task_id),
):
perform_scheduled_scan_task.run(
@@ -2982,6 +2983,13 @@ class TestPerformScheduledScanTask:
).count()
== 1
)
completed_scan = Scan.objects.get(
tenant_id=tenant.id,
provider=provider,
trigger=Scan.TriggerChoices.SCHEDULED,
state=StateChoices.COMPLETED,
)
mock_reconcile.assert_called_once_with(str(tenant.id), str(completed_scan.id))
assert (
Scan.objects.filter(
tenant_id=tenant.id,
@@ -3176,7 +3184,10 @@ class TestPerformScanTask:
task=queued_task,
)
events = []
def _complete_scan(tenant_id, scan_id, provider_id, checks_to_execute=None):
events.append("scan")
scan_instance = Scan.objects.get(id=scan_id)
scan_instance.state = StateChoices.COMPLETED
scan_instance.save()
@@ -3184,7 +3195,14 @@ class TestPerformScanTask:
with (
patch("tasks.tasks.perform_prowler_scan", side_effect=_complete_scan),
patch("tasks.tasks._perform_scan_complete_tasks"),
patch(
"tasks.tasks.reconcile_scan_mute_rules",
side_effect=lambda *_args: events.append("reconcile"),
),
patch(
"tasks.tasks._perform_scan_complete_tasks",
side_effect=lambda *_args: events.append("summaries"),
),
patch("tasks.tasks.perform_scan_task.apply_async") as mock_apply_async,
):
with django_capture_on_commit_callbacks(execute=True):
@@ -3196,6 +3214,7 @@ class TestPerformScanTask:
queued_task_result.refresh_from_db()
assert result == {"status": "ok"}
assert events == ["scan", "reconcile", "summaries"]
assert queued_task_result.status == states.PENDING
mock_apply_async.assert_called_once_with(
kwargs={
@@ -3241,10 +3260,7 @@ class TestPerformScanTask:
@pytest.mark.django_db
class TestReaggregateAllFindingGroupSummaries:
def setup_method(self):
self.tenant_id = str(uuid.uuid4())
class TestMuteFindingsInLatestScansTask:
@patch("tasks.tasks.chain")
@patch("tasks.tasks.group")
@patch("tasks.tasks.aggregate_attack_surface_task")
@@ -3253,10 +3269,10 @@ class TestReaggregateAllFindingGroupSummaries:
@patch("tasks.tasks.aggregate_finding_group_summaries_task")
@patch("tasks.tasks.aggregate_daily_severity_task")
@patch("tasks.tasks.perform_scan_summary_task")
@patch("tasks.tasks.Scan.objects.filter")
def test_dispatches_subtasks_for_each_provider_per_day(
@patch("tasks.tasks.mute_findings_in_latest_scans")
def test_reaggregates_only_changed_scans(
self,
mock_scan_filter,
mock_mute_findings,
mock_scan_summary_task,
mock_daily_severity_task,
mock_finding_group_task,
@@ -3265,119 +3281,36 @@ class TestReaggregateAllFindingGroupSummaries:
mock_attack_surface_task,
mock_group,
mock_chain,
tenants_fixture,
):
provider_id_1 = uuid.uuid4()
provider_id_2 = uuid.uuid4()
scan_id_today_p1 = uuid.uuid4()
scan_id_yesterday_p1 = uuid.uuid4()
scan_id_today_p2 = uuid.uuid4()
today = datetime.now(tz=UTC)
yesterday = today - timedelta(days=1)
mock_outer_group_result = MagicMock()
# The first `group()` call wraps the inner parallel step; subsequent
# calls wrap the outer per-scan generator.
mock_group.side_effect = lambda *args, **kwargs: (
list(args[0]) if args and hasattr(args[0], "__iter__") else None,
mock_outer_group_result,
)[1]
mock_scan_filter.return_value.order_by.return_value.values.return_value = [
{
"id": scan_id_today_p1,
"completed_at": today,
"provider_id": provider_id_1,
},
{
"id": scan_id_today_p2,
"completed_at": today,
"provider_id": provider_id_2,
},
{
"id": scan_id_yesterday_p1,
"completed_at": yesterday,
"provider_id": provider_id_1,
},
]
result = reaggregate_all_finding_group_summaries_task(tenant_id=self.tenant_id)
assert result == {"scans_reaggregated": 3}
expected_scan_ids = {
str(scan_id_today_p1),
str(scan_id_today_p2),
str(scan_id_yesterday_p1),
tenant_id = str(tenants_fixture[0].id)
mute_rule_id = str(uuid.uuid4())
provider_ids = [str(uuid.uuid4()), str(uuid.uuid4())]
scan_ids = [str(uuid.uuid4()), str(uuid.uuid4())]
result = {
"findings_muted": 2,
"rule_id": mute_rule_id,
"scan_ids": scan_ids,
}
for task_mock in (
mock_scan_summary_task,
mock_daily_severity_task,
mock_finding_group_task,
mock_resource_group_task,
mock_category_task,
mock_attack_surface_task,
):
assert task_mock.si.call_count == 3
dispatched = {
call.kwargs["scan_id"] for call in task_mock.si.call_args_list
}
assert dispatched == expected_scan_ids
for call in task_mock.si.call_args_list:
assert call.kwargs["tenant_id"] == self.tenant_id
assert mock_chain.call_count == 3
mock_outer_group_result.apply_async.assert_called_once()
@patch("tasks.tasks.chain")
@patch("tasks.tasks.group")
@patch("tasks.tasks.aggregate_attack_surface_task")
@patch("tasks.tasks.aggregate_scan_category_summaries_task")
@patch("tasks.tasks.aggregate_scan_resource_group_summaries_task")
@patch("tasks.tasks.aggregate_finding_group_summaries_task")
@patch("tasks.tasks.aggregate_daily_severity_task")
@patch("tasks.tasks.perform_scan_summary_task")
@patch("tasks.tasks.Scan.objects.filter")
def test_dedupes_scans_to_latest_per_provider_per_day(
self,
mock_scan_filter,
mock_scan_summary_task,
mock_daily_severity_task,
mock_finding_group_task,
mock_resource_group_task,
mock_category_task,
mock_attack_surface_task,
mock_group,
mock_chain,
):
"""When several scans run on the same day for the same provider, only
the latest one is dispatched (matching the daily summary unique key)."""
provider_id = uuid.uuid4()
latest_scan_today = uuid.uuid4()
earlier_scan_today = uuid.uuid4()
today_late = datetime.now(tz=UTC)
today_early = today_late - timedelta(hours=4)
mock_mute_findings.return_value = result
mock_outer_group_result = MagicMock()
mock_group.side_effect = lambda *args, **kwargs: (
list(args[0]) if args and hasattr(args[0], "__iter__") else None,
mock_outer_group_result,
)[1]
# Returned ordered by `-completed_at`, so the most recent comes first.
mock_scan_filter.return_value.order_by.return_value.values.return_value = [
{
"id": latest_scan_today,
"completed_at": today_late,
"provider_id": provider_id,
},
{
"id": earlier_scan_today,
"completed_at": today_early,
"provider_id": provider_id,
},
]
task_result = mute_findings_in_latest_scans_task(
tenant_id=tenant_id,
mute_rule_id=mute_rule_id,
provider_ids=provider_ids,
)
result = reaggregate_all_finding_group_summaries_task(tenant_id=self.tenant_id)
assert result == {"scans_reaggregated": 1}
assert task_result == result
mock_mute_findings.assert_called_once_with(
tenant_id=tenant_id,
mute_rule_id=mute_rule_id,
provider_ids=provider_ids,
)
for task_mock in (
mock_scan_summary_task,
mock_daily_severity_task,
@@ -3386,23 +3319,35 @@ class TestReaggregateAllFindingGroupSummaries:
mock_category_task,
mock_attack_surface_task,
):
task_mock.si.assert_called_once_with(
tenant_id=self.tenant_id, scan_id=str(latest_scan_today)
)
mock_chain.assert_called_once()
assert task_mock.si.call_count == 2
assert {
call.kwargs["scan_id"] for call in task_mock.si.call_args_list
} == set(scan_ids)
assert mock_chain.call_count == 2
mock_outer_group_result.apply_async.assert_called_once()
@patch("tasks.tasks.chain")
@patch("tasks.tasks.group")
@patch("tasks.tasks.Scan.objects.filter")
def test_no_completed_scans_skips_dispatch(
self, mock_scan_filter, mock_group, mock_chain
@patch("tasks.tasks.mute_findings_in_latest_scans")
def test_skips_reaggregation_when_no_scan_changed(
self, mock_mute_findings, mock_group, mock_chain, tenants_fixture
):
mock_scan_filter.return_value.order_by.return_value.values.return_value = []
tenant_id = str(tenants_fixture[0].id)
mute_rule_id = str(uuid.uuid4())
result = {
"findings_muted": 0,
"rule_id": mute_rule_id,
"scan_ids": [],
}
mock_mute_findings.return_value = result
result = reaggregate_all_finding_group_summaries_task(tenant_id=self.tenant_id)
task_result = mute_findings_in_latest_scans_task(
tenant_id=tenant_id,
mute_rule_id=mute_rule_id,
provider_ids=[],
)
assert result == {"scans_reaggregated": 0}
assert task_result == result
mock_group.assert_not_called()
mock_chain.assert_not_called()
Generated
+1 -1
View File
@@ -4938,7 +4938,7 @@ dependencies = [
[[package]]
name = "prowler-api"
version = "1.41.0"
version = "1.41.1"
source = { virtual = "." }
dependencies = [
{ name = "cartography" },
@@ -0,0 +1 @@
GitHub repository discovery for unscoped scans now paginates beyond the first 100 accessible repositories instead of silently scanning only the first page
+1 -1
View File
@@ -52,7 +52,7 @@ class _MutableTimestamp:
timestamp = _MutableTimestamp(datetime.today())
timestamp_utc = _MutableTimestamp(datetime.now(timezone.utc))
prowler_version = "5.40.0"
prowler_version = "5.40.1"
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"
@@ -10,6 +10,8 @@ from prowler.lib.logger import logger
from prowler.providers.github.lib.service.service import GithubService
from prowler.providers.github.models import GithubAppIdentityInfo
GITHUB_GRAPHQL_TIMEOUT = (10, 60)
class Repository(GithubService):
def __init__(self, provider):
@@ -65,41 +67,105 @@ class Repository(GithubService):
"Content-Type": "application/json",
}
query = """
{
query ($cursor: String) {
viewer {
repositories(first: 100, affiliations: [OWNER, ORGANIZATION_MEMBER]) {
repositories(
first: 100
after: $cursor
affiliations: [OWNER, ORGANIZATION_MEMBER]
) {
nodes {
nameWithOwner
}
pageInfo {
hasNextPage
endCursor
}
}
}
}
"""
repositories = []
cursor = None
seen_cursors = set()
try:
response = requests.post(
graphql_url, json={"query": query}, headers=headers
)
response.raise_for_status()
data = response.json()
while True:
response = requests.post(
graphql_url,
json={"query": query, "variables": {"cursor": cursor}},
headers=headers,
timeout=GITHUB_GRAPHQL_TIMEOUT,
)
response.raise_for_status()
data = response.json()
if "errors" in data:
logger.error(f"Error in GraphQL query: {data['errors']}")
return []
errors = data.get("errors") if isinstance(data, dict) else None
repository_connection = (
((data.get("data") or {}).get("viewer") or {}).get("repositories")
if isinstance(data, dict)
else None
)
if not isinstance(repository_connection, dict):
logger.error(
f"Error in GraphQL query: {errors or 'invalid response'}"
)
return repositories
if errors:
# GitHub returns partial responses: repositories the token
# cannot access (e.g. behind organization SAML enforcement)
# come back as null nodes together with an "errors" entry,
# while the rest of the page is valid.
logger.warning(
f"GitHub GraphQL returned errors while discovering repositories, "
f"some repositories may be skipped: {errors}"
)
repo_nodes = (
data.get("data", {})
.get("viewer", {})
.get("repositories", {})
.get("nodes", [])
)
return [repo["nameWithOwner"] for repo in repo_nodes]
repo_nodes = repository_connection.get("nodes")
page_info = repository_connection.get("pageInfo")
if (
not isinstance(repo_nodes, list)
or not isinstance(page_info, dict)
or not isinstance(page_info.get("hasNextPage"), bool)
):
logger.error(
"GitHub GraphQL returned an invalid repositories page; "
"repository discovery may be incomplete."
)
return repositories
except requests.exceptions.RequestException as error:
for repo_node in repo_nodes:
if not repo_node:
logger.warning(
"Skipping a repository the token cannot access during discovery."
)
continue
repositories.append(repo_node["nameWithOwner"])
if not page_info["hasNextPage"]:
return repositories
cursor = page_info.get("endCursor")
if not cursor or cursor in seen_cursors:
logger.error(
"GitHub GraphQL pagination returned an invalid cursor; "
"repository discovery may be incomplete."
)
return repositories
seen_cursors.add(cursor)
except (
requests.exceptions.RequestException,
ValueError,
KeyError,
TypeError,
AttributeError,
) as error:
logger.error(
f"{error.__class__.__name__}[{error.__traceback__.tb_lineno}]: {error}"
)
return []
return repositories
def _default_branch_matches_rule_pattern(
self, pattern: str, default_branch: str
+1 -1
View File
@@ -143,7 +143,7 @@ maintainers = [{name = "Prowler Engineering", email = "engineering@prowler.com"}
name = "prowler"
readme = "README.md"
requires-python = ">=3.10,<3.14"
version = "5.40.0"
version = "5.40.1"
[project.scripts]
prowler = "prowler.__main__:prowler"
@@ -1165,3 +1165,194 @@ class Test_Repository_List_Rate_Limit_Propagation:
):
with raises(RateLimitExceededException):
repository_service._list_repositories()
class Test_Repository_GraphQL_Pagination:
"""Accessible repository discovery must follow GraphQL pagination."""
def _repository_service(self):
provider = set_mocked_github_provider()
provider.repositories = []
provider.organizations = []
with patch.object(Repository, "__init__", lambda *_: None):
repository_service = Repository(provider)
repository_service.clients = [MagicMock()]
repository_service.provider = provider
return repository_service
@staticmethod
def _graphql_response(nodes, has_next_page=False, end_cursor=None, errors=None):
response = MagicMock()
payload = {
"data": {
"viewer": {
"repositories": {
"nodes": nodes,
"pageInfo": {
"hasNextPage": has_next_page,
"endCursor": end_cursor,
},
}
}
}
}
if errors is not None:
payload["errors"] = errors
response.json.return_value = payload
return response
def test_graphql_paginates_accessible_repositories(self):
repository_service = self._repository_service()
first_page_names = [f"owner/repo-{index}" for index in range(100)]
second_page_names = ["owner/repo-100"]
pages = [
self._graphql_response(
[{"nameWithOwner": name} for name in first_page_names],
has_next_page=True,
end_cursor="page-2",
),
self._graphql_response(
[{"nameWithOwner": name} for name in second_page_names]
),
]
with patch("requests.post", side_effect=pages) as mock_post:
repositories = repository_service._get_accessible_repos_graphql()
assert repositories == first_page_names + second_page_names
assert mock_post.call_count == 2
first_request, second_request = mock_post.call_args_list
assert "after: $cursor" in first_request.kwargs["json"]["query"]
assert first_request.kwargs["json"]["variables"] == {"cursor": None}
assert second_request.kwargs["json"]["variables"] == {"cursor": "page-2"}
assert first_request.kwargs["timeout"] == (10, 60)
def test_graphql_keeps_accessible_repositories_on_partial_errors(self):
"""Per-node errors (e.g. SAML-protected repositories) must not drop the page."""
repository_service = self._repository_service()
response = self._graphql_response(
[{"nameWithOwner": "owner/visible"}, None],
errors=[
{
"type": "FORBIDDEN",
"path": ["viewer", "repositories", "nodes", 1],
"message": "Resource protected by organization SAML enforcement.",
}
],
)
with (
patch("requests.post", return_value=response),
patch(
"prowler.providers.github.services.repository.repository_service.logger"
) as mock_logger,
):
repositories = repository_service._get_accessible_repos_graphql()
assert repositories == ["owner/visible"]
assert mock_logger.warning.call_count == 2
assert "SAML" in str(mock_logger.warning.call_args_list[0])
def test_graphql_later_page_failure_keeps_collected_repositories(self):
repository_service = self._repository_service()
first_page = self._graphql_response(
[{"nameWithOwner": "owner/first"}], has_next_page=True, end_cursor="page-2"
)
with (
patch(
"requests.post",
side_effect=[
first_page,
requests.exceptions.Timeout("second page timed out"),
],
) as mock_post,
patch(
"prowler.providers.github.services.repository.repository_service.logger"
) as mock_logger,
):
repositories = repository_service._get_accessible_repos_graphql()
assert repositories == ["owner/first"]
assert mock_post.call_count == 2
mock_logger.error.assert_called_once()
assert "Timeout" in str(mock_logger.error.call_args)
def test_graphql_errors_without_data_return_empty_list(self):
repository_service = self._repository_service()
response = MagicMock()
response.json.return_value = {
"data": None,
"errors": [{"type": "RATE_LIMITED", "message": "API rate limit exceeded"}],
}
with (
patch("requests.post", return_value=response),
patch(
"prowler.providers.github.services.repository.repository_service.logger"
) as mock_logger,
):
repositories = repository_service._get_accessible_repos_graphql()
assert repositories == []
assert "RATE_LIMITED" in str(mock_logger.error.call_args)
def test_graphql_stops_on_repeated_cursor(self):
repository_service = self._repository_service()
pages = [
self._graphql_response(
[{"nameWithOwner": "owner/a"}], has_next_page=True, end_cursor="same"
),
self._graphql_response(
[{"nameWithOwner": "owner/b"}], has_next_page=True, end_cursor="same"
),
]
with (
patch("requests.post", side_effect=pages) as mock_post,
patch(
"prowler.providers.github.services.repository.repository_service.logger"
) as mock_logger,
):
repositories = repository_service._get_accessible_repos_graphql()
assert repositories == ["owner/a", "owner/b"]
assert mock_post.call_count == 2
mock_logger.error.assert_called_once()
@pytest.mark.parametrize(
"nodes, page_info",
[
("not-a-list", {"hasNextPage": False}),
({}, {"hasNextPage": False}),
(None, {"hasNextPage": False}),
([{"nameWithOwner": "owner/b"}], ["invalid"]),
([{"nameWithOwner": "owner/b"}], "invalid"),
([{"nameWithOwner": "owner/b"}], []),
([{"nameWithOwner": "owner/b"}], None),
([{"nameWithOwner": "owner/b"}], {}),
([{"nameWithOwner": "owner/b"}], {"hasNextPage": "false"}),
],
)
def test_graphql_invalid_page_keeps_collected_repositories(self, nodes, page_info):
repository_service = self._repository_service()
first_page = self._graphql_response(
[{"nameWithOwner": "owner/a"}], has_next_page=True, end_cursor="page-2"
)
invalid_page = MagicMock()
invalid_page.json.return_value = {
"data": {
"viewer": {"repositories": {"nodes": nodes, "pageInfo": page_info}}
}
}
with (
patch("requests.post", side_effect=[first_page, invalid_page]),
patch(
"prowler.providers.github.services.repository.repository_service.logger"
) as mock_logger,
):
repositories = repository_service._get_accessible_repos_graphql()
assert repositories == ["owner/a"]
mock_logger.error.assert_called_once()
+92
View File
@@ -154,4 +154,96 @@ describe("auth actions", () => {
expect(result.permissions.manage_lighthouse_ai_configuration).toBe(false);
expect(result.permissions.manage_users).toBe(true);
});
it("should forward an abort signal when loading the current user", async () => {
// Given
mockUserMe({ manage_users: true });
const abortController = new AbortController();
// When
await getUserByMe("access-token", abortController.signal);
// Then
expect(fetchMock).toHaveBeenCalledWith(
"https://api.example.com/api/v1/users/me?include=roles",
expect.objectContaining({ signal: abortController.signal }),
);
});
it.each([
{
status: 401,
detail: "Rejected by API",
message: "Invalid or expired token",
},
{
status: 403,
detail: "Database password: super-secret",
message: "Access denied",
},
{ status: 404, detail: "Rejected by API", message: "User not found" },
])(
"should preserve a $status status when loading the current user fails",
async ({ status, detail, message }) => {
// Given
fetchMock.mockResolvedValue(
new Response(JSON.stringify({ errors: [{ detail }] }), { status }),
);
// When
const result = getUserByMe("access-token");
// Then
await expect(result).rejects.toMatchObject({ message, status });
},
);
it("should preserve a 401 status when the error body is not JSON", async () => {
// Given
fetchMock.mockResolvedValue(new Response("Unauthorized", { status: 401 }));
// When
const result = getUserByMe("access-token");
// Then
await expect(result).rejects.toMatchObject({
message: "Invalid or expired token",
status: 401,
});
});
it("should preserve a 403 status when the error body is not JSON", async () => {
// Given
fetchMock.mockResolvedValue(new Response("Forbidden", { status: 403 }));
// When
const result = getUserByMe("access-token");
// Then
await expect(result).rejects.toMatchObject({
message: "Access denied",
status: 403,
});
});
it("should not expose upstream details for unexpected errors", async () => {
// Given
fetchMock.mockResolvedValue(
new Response(
JSON.stringify({
errors: [{ detail: "Database password: super-secret" }],
}),
{ status: 500 },
),
);
// When
const result = getUserByMe("access-token");
// Then
await expect(result).rejects.toMatchObject({
message: "Unable to load user",
status: 500,
});
});
});
+25 -17
View File
@@ -4,6 +4,7 @@ import { AuthError } from "next-auth";
import { signIn, signOut } from "@/auth.config";
import { apiBaseUrl } from "@/lib";
import { UserMeError } from "@/lib/auth-errors";
import { addAuthEvent } from "@/lib/sentry-breadcrumbs";
import type { UtmParams } from "@/lib/utm";
import type { SignInFormData, SignUpFormData } from "@/types";
@@ -140,7 +141,10 @@ export const getToken = async (formData: SignInFormData) => {
}
};
export const getUserByMe = async (accessToken: string) => {
export const getUserByMe = async (
accessToken: string,
signal?: AbortSignal,
) => {
const url = new URL(`${apiBaseUrl}/users/me?include=roles`);
try {
@@ -150,25 +154,23 @@ export const getUserByMe = async (accessToken: string) => {
Accept: "application/vnd.api+json",
Authorization: `Bearer ${accessToken}`,
},
signal,
});
const parsedResponse = await response.json();
if (!response.ok) {
// Handle different HTTP error codes
switch (response.status) {
case 401:
throw new Error("Invalid or expired token");
case 403:
throw new Error(parsedResponse.errors?.[0]?.detail);
case 404:
throw new Error("User not found");
default:
throw new Error(
parsedResponse.errors?.[0]?.detail || "Unknown error",
);
}
const errorMessage =
response.status === 401
? "Invalid or expired token"
: response.status === 403
? "Access denied"
: response.status === 404
? "User not found"
: "Unable to load user";
throw new UserMeError(errorMessage, response.status);
}
const parsedResponse = await response.json();
const userRole = parsedResponse.included?.find(
(item: any) => item.type === "roles",
);
@@ -193,8 +195,14 @@ export const getUserByMe = async (accessToken: string) => {
dateJoined: parsedResponse.data.attributes.date_joined,
permissions,
};
} catch (error: any) {
throw new Error(error.message || "Network error or server unreachable");
} catch (error: unknown) {
if (error instanceof UserMeError) throw error;
throw new UserMeError(
error instanceof Error
? error.message
: "Network error or server unreachable",
);
}
};
+444
View File
@@ -1,6 +1,7 @@
import { beforeEach, describe, expect, it, vi } from "vitest";
import { authConfig } from "./auth.config";
import { UserMeError } from "./lib/auth-errors";
import type { RolePermissionAttributes } from "./types/users";
const { getUserByMeMock } = vi.hoisted(() => ({
@@ -45,6 +46,35 @@ const ELEVATED_PERMISSIONS: RolePermissionAttributes = {
manage_scans: true,
};
const accessTokenFor = (tenantId: string, expiration: number) =>
`header.${Buffer.from(
JSON.stringify({ sub: "user-1", tenant_id: tenantId, exp: expiration }),
).toString("base64url")}.signature`;
const successfulRefreshResponse = (accessToken: string, refreshToken: string) =>
new Response(
JSON.stringify({
data: {
attributes: {
access: accessToken,
refresh: refreshToken,
},
},
}),
{ status: 200 },
);
const mockSuccessfulRefresh = (
accessToken: string,
refreshToken = "new-refresh-token",
) => {
const fetchMock = vi
.fn()
.mockResolvedValue(successfulRefreshResponse(accessToken, refreshToken));
vi.stubGlobal("fetch", fetchMock);
return fetchMock;
};
describe("authConfig JWT callback", () => {
beforeEach(() => {
vi.clearAllMocks();
@@ -194,4 +224,418 @@ describe("authConfig JWT callback", () => {
});
expect(result.error).toBeUndefined();
});
it("should replace restricted permissions after access token refresh", async () => {
// Given
const currentAccessToken = accessTokenFor("stale-tenant", 1);
const newAccessToken = accessTokenFor("tenant-1", 4_102_444_800);
mockSuccessfulRefresh(newAccessToken);
getUserByMeMock.mockResolvedValue({
name: "Tenant User",
email: "tenant@example.com",
company: "Tenant Company",
dateJoined: "2026-01-01",
permissions: ELEVATED_PERMISSIONS,
});
const jwtCallback = authConfig.callbacks?.jwt;
if (!jwtCallback) throw new Error("JWT callback is not configured");
// When
const result = await jwtCallback({
token: {
accessToken: currentAccessToken,
refreshToken: "current-refresh-token",
user: {
name: "Tenant User",
email: "tenant@example.com",
dateJoined: "2026-01-01",
permissions: RESTRICTED_PERMISSIONS,
},
},
user: {} as Parameters<typeof jwtCallback>[0]["user"],
});
// Then
expect(getUserByMeMock).toHaveBeenCalledWith(
newAccessToken,
expect.any(AbortSignal),
);
expect(result).toMatchObject({
accessToken: newAccessToken,
refreshToken: "new-refresh-token",
tenant_id: "tenant-1",
user: {
permissions: ELEVATED_PERMISSIONS,
},
});
expect(result.error).toBeUndefined();
});
it("should replace elevated permissions when access is revoked", async () => {
// Given
const currentAccessToken = accessTokenFor("tenant-1", 1);
const newAccessToken = accessTokenFor("tenant-1", 4_102_444_800);
mockSuccessfulRefresh(newAccessToken);
getUserByMeMock.mockResolvedValue({
name: "Tenant User",
email: "tenant@example.com",
company: "Tenant Company",
dateJoined: "2026-01-01",
permissions: RESTRICTED_PERMISSIONS,
});
const jwtCallback = authConfig.callbacks?.jwt;
if (!jwtCallback) throw new Error("JWT callback is not configured");
// When
const result = await jwtCallback({
token: {
accessToken: currentAccessToken,
refreshToken: "current-refresh-token",
user: {
name: "Tenant User",
email: "tenant@example.com",
dateJoined: "2026-01-01",
permissions: ELEVATED_PERMISSIONS,
},
},
user: {} as Parameters<typeof jwtCallback>[0]["user"],
});
// Then
expect(getUserByMeMock).toHaveBeenCalledWith(
newAccessToken,
expect.any(AbortSignal),
);
expect(result.user?.permissions).toEqual(RESTRICTED_PERMISSIONS);
expect(result.error).toBeUndefined();
});
it("should keep the refreshed tokens and cached user when reloading the user fails", async () => {
// Given
const currentAccessToken = accessTokenFor("tenant-1", 1);
const newAccessToken = accessTokenFor("tenant-1", 4_102_444_800);
const warnSpy = vi
.spyOn(console, "warn")
.mockImplementation(() => undefined);
mockSuccessfulRefresh(newAccessToken);
getUserByMeMock.mockRejectedValue(
new UserMeError("Sensitive backend detail", 500),
);
const jwtCallback = authConfig.callbacks?.jwt;
if (!jwtCallback) throw new Error("JWT callback is not configured");
const sessionCallback = authConfig.callbacks?.session;
if (!sessionCallback) throw new Error("Session callback is not configured");
const cachedUser = {
name: "Tenant User",
email: "tenant@example.com",
dateJoined: "2026-01-01",
permissions: ELEVATED_PERMISSIONS,
};
// When
const result = await jwtCallback({
token: {
accessToken: currentAccessToken,
refreshToken: "current-refresh-token",
user: cachedUser,
},
user: {} as Parameters<typeof jwtCallback>[0]["user"],
});
const session = await sessionCallback({
session: {
expires: "2026-12-31T23:59:59.999Z",
user: { name: "Tenant User" },
},
token: result,
} as Parameters<typeof sessionCallback>[0]);
// Then
expect(session).toMatchObject({
accessToken: newAccessToken,
refreshToken: "new-refresh-token",
tenantId: "tenant-1",
user: cachedUser,
});
expect(result.error).toBeUndefined();
expect(warnSpy).toHaveBeenCalledWith(
"Unable to refresh user after access token refresh",
);
});
it("should bound a pending user reload and keep the refreshed session", async () => {
// Given
const currentAccessToken = accessTokenFor("tenant-1", 1);
const newAccessToken = accessTokenFor("tenant-1", 4_102_444_800);
const abortController = new AbortController();
const timeoutSpy = vi
.spyOn(AbortSignal, "timeout")
.mockReturnValue(abortController.signal);
vi.spyOn(console, "warn").mockImplementation(() => undefined);
mockSuccessfulRefresh(newAccessToken);
getUserByMeMock.mockImplementation(
(_accessToken: string, signal?: AbortSignal) => {
if (!signal) return Promise.reject(new Error("Missing abort signal"));
return new Promise((_resolve, reject) => {
signal.addEventListener("abort", () => reject(signal.reason), {
once: true,
});
});
},
);
const jwtCallback = authConfig.callbacks?.jwt;
if (!jwtCallback) throw new Error("JWT callback is not configured");
const cachedUser = {
name: "Tenant User",
email: "tenant@example.com",
dateJoined: "2026-01-01",
permissions: ELEVATED_PERMISSIONS,
};
// When
const resultPromise = jwtCallback({
token: {
accessToken: currentAccessToken,
refreshToken: "current-refresh-token",
user: cachedUser,
},
user: {} as Parameters<typeof jwtCallback>[0]["user"],
});
await vi.waitFor(() => expect(getUserByMeMock).toHaveBeenCalled());
// Then
expect(timeoutSpy).toHaveBeenCalledWith(5_000);
expect(getUserByMeMock).toHaveBeenCalledWith(
newAccessToken,
abortController.signal,
);
abortController.abort(
new DOMException("Request timed out", "TimeoutError"),
);
await expect(resultPromise).resolves.toMatchObject({
accessToken: newAccessToken,
refreshToken: "new-refresh-token",
user: cachedUser,
error: undefined,
});
});
it.each([401, 403, 404])(
"should invalidate the session when reloading the user returns %i",
async (status) => {
// Given
const currentAccessToken = accessTokenFor("tenant-1", 1);
const newAccessToken = accessTokenFor("tenant-1", 4_102_444_800);
vi.spyOn(console, "warn").mockImplementation(() => undefined);
mockSuccessfulRefresh(newAccessToken);
getUserByMeMock.mockRejectedValue(
new UserMeError("Unable to load user", status),
);
const jwtCallback = authConfig.callbacks?.jwt;
if (!jwtCallback) throw new Error("JWT callback is not configured");
const sessionCallback = authConfig.callbacks?.session;
if (!sessionCallback)
throw new Error("Session callback is not configured");
// When
const result = await jwtCallback({
token: {
accessToken: currentAccessToken,
refreshToken: "current-refresh-token",
user: {
name: "Tenant User",
email: "tenant@example.com",
dateJoined: "2026-01-01",
permissions: ELEVATED_PERMISSIONS,
},
},
user: {} as Parameters<typeof jwtCallback>[0]["user"],
});
const session = await sessionCallback({
session: {
expires: "2026-12-31T23:59:59.999Z",
user: { name: "Tenant User" },
},
token: result,
} as Parameters<typeof sessionCallback>[0]);
// Then
expect(result.user).toBeUndefined();
expect(result.accessToken).toBeUndefined();
expect(result.refreshToken).toBeUndefined();
expect(result.error).toBe("RefreshAccessTokenError");
expect(session.user).toBeUndefined();
expect(session.accessToken).toBeUndefined();
expect(session.refreshToken).toBeUndefined();
},
);
it("should invalidate the session when access token refresh fails", async () => {
// Given
const currentAccessToken = accessTokenFor("tenant-1", 1);
vi.spyOn(console, "warn").mockImplementation(() => undefined);
vi.stubGlobal(
"fetch",
vi
.fn()
.mockResolvedValue(
new Response(
JSON.stringify({ errors: [{ detail: "Refresh token expired" }] }),
{ status: 401 },
),
),
);
const jwtCallback = authConfig.callbacks?.jwt;
if (!jwtCallback) throw new Error("JWT callback is not configured");
const sessionCallback = authConfig.callbacks?.session;
if (!sessionCallback) throw new Error("Session callback is not configured");
// When
const result = await jwtCallback({
token: {
accessToken: currentAccessToken,
refreshToken: "expired-refresh-token",
user: {
name: "Tenant User",
email: "tenant@example.com",
dateJoined: "2026-01-01",
permissions: ELEVATED_PERMISSIONS,
},
},
user: {} as Parameters<typeof jwtCallback>[0]["user"],
});
const session = await sessionCallback({
session: {
expires: "2026-12-31T23:59:59.999Z",
user: { name: "Tenant User" },
},
token: result,
} as Parameters<typeof sessionCallback>[0]);
// Then
expect(getUserByMeMock).not.toHaveBeenCalled();
expect(result.error).toBe("RefreshAccessTokenError");
expect(session.error).toBe("RefreshAccessTokenError");
expect(session.user).toBeUndefined();
expect(session.accessToken).toBeUndefined();
expect(session.refreshToken).toBeUndefined();
expect(session.tenantId).toBeUndefined();
});
it("should deduplicate concurrent token and user refreshes", async () => {
// Given
const currentAccessToken = accessTokenFor("tenant-1", 1);
const newAccessToken = accessTokenFor("tenant-1", 4_102_444_800);
const fetchMock = mockSuccessfulRefresh(newAccessToken);
getUserByMeMock.mockResolvedValue({
name: "Tenant User",
email: "tenant@example.com",
company: "Tenant Company",
dateJoined: "2026-01-01",
permissions: ELEVATED_PERMISSIONS,
});
const jwtCallback = authConfig.callbacks?.jwt;
if (!jwtCallback) throw new Error("JWT callback is not configured");
const currentToken = {
accessToken: currentAccessToken,
refreshToken: "shared-refresh-token",
user: {
name: "Tenant User",
email: "tenant@example.com",
dateJoined: "2026-01-01",
permissions: RESTRICTED_PERMISSIONS,
},
};
// When
const [firstResult, secondResult] = await Promise.all([
jwtCallback({
token: { ...currentToken },
user: {} as Parameters<typeof jwtCallback>[0]["user"],
}),
jwtCallback({
token: { ...currentToken },
user: {} as Parameters<typeof jwtCallback>[0]["user"],
}),
]);
// Then
expect(fetchMock).toHaveBeenCalledTimes(1);
expect(getUserByMeMock).toHaveBeenCalledTimes(1);
expect(firstResult).toEqual(secondResult);
expect(firstResult.user?.permissions).toEqual(ELEVATED_PERMISSIONS);
});
it("should retry reloading the user on the next token rotation", async () => {
// Given
const currentAccessToken = accessTokenFor("tenant-1", 1);
const firstAccessToken = accessTokenFor("tenant-1", 1);
const secondAccessToken = accessTokenFor("tenant-1", 4_102_444_800);
vi.spyOn(console, "warn").mockImplementation(() => undefined);
const fetchMock = vi
.fn()
.mockResolvedValueOnce(
successfulRefreshResponse(
firstAccessToken,
"first-rotated-refresh-token",
),
)
.mockResolvedValueOnce(
successfulRefreshResponse(
secondAccessToken,
"second-rotated-refresh-token",
),
);
vi.stubGlobal("fetch", fetchMock);
getUserByMeMock
.mockRejectedValueOnce(new Error("Temporary API failure"))
.mockResolvedValueOnce({
name: "Tenant User",
email: "tenant@example.com",
company: "Tenant Company",
dateJoined: "2026-01-01",
permissions: RESTRICTED_PERMISSIONS,
});
const jwtCallback = authConfig.callbacks?.jwt;
if (!jwtCallback) throw new Error("JWT callback is not configured");
// When
const firstResult = await jwtCallback({
token: {
accessToken: currentAccessToken,
refreshToken: "current-refresh-token",
user: {
name: "Tenant User",
email: "tenant@example.com",
dateJoined: "2026-01-01",
permissions: ELEVATED_PERMISSIONS,
},
},
user: {} as Parameters<typeof jwtCallback>[0]["user"],
});
const secondResult = await jwtCallback({
token: firstResult,
user: {} as Parameters<typeof jwtCallback>[0]["user"],
});
// Then
expect(fetchMock).toHaveBeenCalledTimes(2);
expect(getUserByMeMock).toHaveBeenNthCalledWith(
1,
firstAccessToken,
expect.any(AbortSignal),
);
expect(getUserByMeMock).toHaveBeenNthCalledWith(
2,
secondAccessToken,
expect.any(AbortSignal),
);
expect(firstResult.user?.permissions).toEqual(ELEVATED_PERMISSIONS);
expect(secondResult).toMatchObject({
accessToken: secondAccessToken,
refreshToken: "second-rotated-refresh-token",
user: { permissions: RESTRICTED_PERMISSIONS },
});
});
});
+29
View File
@@ -11,6 +11,7 @@ import { z } from "zod";
import { getToken, getUserByMe } from "./actions/auth";
import { apiBaseUrl } from "./lib";
import { UserMeError } from "./lib/auth-errors";
import {
SLACK_CALLBACK_PATH,
SLACK_EXPIRED_CALLBACK_URL,
@@ -64,6 +65,9 @@ const DEFAULT_PERMISSIONS: RolePermissionAttributes = {
const TENANT_SWITCH_ERROR = "TenantSwitchError";
const NON_RETRYABLE_USER_ME_STATUSES = new Set([401, 403, 404]);
const USER_REFRESH_TIMEOUT_MS = 5_000;
type TokenUserInput = Partial<TokenUser> & { company?: string };
type JwtCallback = NonNullable<NonNullable<NextAuthConfig["callbacks"]>["jwt"]>;
@@ -206,6 +210,31 @@ const refreshAccessToken = async (token: AuthToken): Promise<AuthToken> => {
applyDecodedClaims(nextToken, newAccessToken, "refreshed access token");
try {
const userMeResponse = await getUserByMe(
newAccessToken,
AbortSignal.timeout(USER_REFRESH_TIMEOUT_MS),
);
nextToken.user = tokenUserFromApi(userMeResponse);
} catch (error) {
if (
error instanceof UserMeError &&
error.status !== undefined &&
NON_RETRYABLE_USER_ME_STATUSES.has(error.status)
) {
return {
...nextToken,
accessToken: undefined,
refreshToken: undefined,
user: undefined,
error: "RefreshAccessTokenError",
};
}
// eslint-disable-next-line no-console
console.warn("Unable to refresh user after access token refresh");
}
return nextToken;
} catch (error) {
// eslint-disable-next-line no-console
@@ -0,0 +1 @@
Cached permissions now refresh from `/users/me?include=roles` after access token rotation
+9
View File
@@ -0,0 +1,9 @@
export class UserMeError extends Error {
constructor(
message: string,
public readonly status?: number,
) {
super(message);
this.name = "UserMeError";
}
}
Generated
+1 -1
View File
@@ -3752,7 +3752,7 @@ wheels = [
[[package]]
name = "prowler"
version = "5.40.0"
version = "5.40.1"
source = { editable = "." }
dependencies = [
{ name = "alibabacloud-actiontrail20200706" },