diff --git a/.env b/.env index 8294fac91c..49a03a7c69 100644 --- a/.env +++ b/.env @@ -72,8 +72,8 @@ NEO4J_APOC_IMPORT_FILE_ENABLED=false NEO4J_APOC_IMPORT_FILE_USE_NEO4J_CONFIG=true NEO4J_APOC_TRIGGER_ENABLED=false NEO4J_DBMS_CONNECTOR_BOLT_LISTEN_ADDRESS=0.0.0.0:7687 -# Neo4j Prowler settings -ATTACK_PATHS_BATCH_SIZE=1000 +# Attack Paths graph settings +ATTACK_PATHS_GRAPH_MUTATION_BATCH_SIZE=1000 ATTACK_PATHS_SERVICE_UNAVAILABLE_MAX_RETRIES=3 ATTACK_PATHS_READ_QUERY_TIMEOUT_SECONDS=30 ATTACK_PATHS_MAX_CUSTOM_QUERY_NODES=250 diff --git a/api/changelog.d/attack-paths-scan-reliability.fixed.md b/api/changelog.d/attack-paths-scan-reliability.fixed.md new file mode 100644 index 0000000000..1532943dc8 --- /dev/null +++ b/api/changelog.d/attack-paths-scan-reliability.fixed.md @@ -0,0 +1 @@ +Attack Paths scans handle provider deletion races cleanly, detect stale tasks after 16 hours, use backend-specific graph synchronization batches, and report exhausted Neptune write retries with the original database error diff --git a/api/src/backend/api/attack_paths/database.py b/api/src/backend/api/attack_paths/database.py index 3a33b964b7..3ef55b7eca 100644 --- a/api/src/backend/api/attack_paths/database.py +++ b/api/src/backend/api/attack_paths/database.py @@ -27,6 +27,7 @@ from django.conf import ( MAX_CUSTOM_QUERY_NODES = env.int("ATTACK_PATHS_MAX_CUSTOM_QUERY_NODES", default=250) TEMP_DB_PREFIX = "db-tmp-scan-" +DATABASE_NOT_FOUND_CODE = "Neo.ClientError.Database.DatabaseNotFound" # Exceptions @@ -44,6 +45,10 @@ class GraphDatabaseQueryException(Exception): return self.message +class NeptuneWriteRetryExhaustedException(GraphDatabaseQueryException): + pass + + class WriteQueryNotAllowedException(GraphDatabaseQueryException): pass diff --git a/api/src/backend/api/attack_paths/retryable_session.py b/api/src/backend/api/attack_paths/retryable_session.py index 2cd32f54ee..e7e78e9798 100644 --- a/api/src/backend/api/attack_paths/retryable_session.py +++ b/api/src/backend/api/attack_paths/retryable_session.py @@ -10,6 +10,28 @@ import neo4j.exceptions logger = logging.getLogger(__name__) +class RetryExhaustedError(Exception): + def __init__( + self, + *, + retry_context: str, + method_name: str, + attempts: int, + elapsed_seconds: float, + last_error: Exception, + ) -> None: + self.retry_context = retry_context + self.method_name = method_name + self.attempts = attempts + self.elapsed_seconds = elapsed_seconds + self.last_error = last_error + last_message = getattr(last_error, "message", None) or str(last_error) + super().__init__( + f"{retry_context} {method_name} failed after {attempts} attempts over " + f"{elapsed_seconds:.3f}s. Last error: {last_message}" + ) + + class RetryableSession: """Wrapper around ``neo4j.Session`` with a refreshable retry policy.""" @@ -19,11 +41,13 @@ class RetryableSession: max_retries: int, retry_if: Callable[[Exception], bool] | None = None, initial_retry_delay_seconds: float = 0, + retry_context: str | None = None, ) -> None: self._session_factory = session_factory self._max_retries = max(0, max_retries) self._retry_if = retry_if self._initial_retry_delay_seconds = max(0.0, initial_retry_delay_seconds) + self._retry_context = retry_context self._session = self._session_factory() def close(self) -> None: @@ -54,6 +78,7 @@ class RetryableSession: def _call_with_retry(self, method_name: str, *args: Any, **kwargs: Any) -> Any: attempt = 0 last_exc: Exception | None = None + started_at = time.monotonic() while attempt <= self._max_retries: try: @@ -68,17 +93,38 @@ class RetryableSession: attempt += 1 if attempt > self._max_retries: + if self._retry_context is not None: + raise RetryExhaustedError( + retry_context=self._retry_context, + method_name=method_name, + attempts=attempt, + elapsed_seconds=time.monotonic() - started_at, + last_error=exc, + ) from exc raise delay = self._retry_delay(attempt) - logger.warning( - "Graph session %s failed with %s; retry %s/%s in %.3fs", - method_name, - type(exc).__name__, - attempt, - self._max_retries, - delay, - ) + if self._retry_context is not None: + error_message = getattr(exc, "message", None) or str(exc) + logger.warning( + "%s %s failed with %s: %s; retry %s/%s in %.3fs", + self._retry_context, + method_name, + type(exc).__name__, + error_message, + attempt, + self._max_retries, + delay, + ) + else: + logger.warning( + "Graph session %s failed with %s; retry %s/%s in %.3fs", + method_name, + type(exc).__name__, + attempt, + self._max_retries, + delay, + ) self._refresh_session() if delay: time.sleep(delay) diff --git a/api/src/backend/api/attack_paths/sink/base.py b/api/src/backend/api/attack_paths/sink/base.py index 0ba4737f5e..0134e0a41f 100644 --- a/api/src/backend/api/attack_paths/sink/base.py +++ b/api/src/backend/api/attack_paths/sink/base.py @@ -15,6 +15,8 @@ class SinkDatabase(Protocol): has a single graph, and isolation is label-based). """ + sync_batch_size: int + def init(self) -> None: ... def close(self) -> None: ... diff --git a/api/src/backend/api/attack_paths/sink/neo4j.py b/api/src/backend/api/attack_paths/sink/neo4j.py index cb7d4889b0..c820ed9d93 100644 --- a/api/src/backend/api/attack_paths/sink/neo4j.py +++ b/api/src/backend/api/attack_paths/sink/neo4j.py @@ -54,6 +54,8 @@ DATABASE_NOT_FOUND_CODE = "Neo.ClientError.Database.DatabaseNotFound" class Neo4jSink(SinkDatabase): """Neo4j-backed sink. Multi-database cluster; tenant isolation is physical.""" + sync_batch_size = env.int("ATTACK_PATHS_NEO4J_SYNC_BATCH_SIZE", default=1000) + def __init__(self) -> None: self._driver: neo4j.Driver | None = None self._lock = threading.Lock() @@ -203,7 +205,7 @@ class Neo4jSink(SinkDatabase): """ from api.attack_paths.database import GraphDatabaseQueryException from tasks.jobs.attack_paths.config import ( - BATCH_SIZE, + GRAPH_MUTATION_BATCH_SIZE, PROVIDER_RESOURCE_LABEL, get_provider_label, ) @@ -251,7 +253,7 @@ class Neo4jSink(SinkDatabase): total_key="rels", deleted_key="deleted_rels", initial_total=deleted_relationships, - batch_size=BATCH_SIZE, + batch_size=GRAPH_MUTATION_BATCH_SIZE, drop_t0=drop_t0, ) relationship_batches += phase_batches @@ -270,7 +272,7 @@ class Neo4jSink(SinkDatabase): total_key="nodes", deleted_key="deleted_nodes", initial_total=0, - batch_size=BATCH_SIZE, + batch_size=GRAPH_MUTATION_BATCH_SIZE, drop_t0=drop_t0, ) diff --git a/api/src/backend/api/attack_paths/sink/neptune.py b/api/src/backend/api/attack_paths/sink/neptune.py index c68b6f8277..022e3c8669 100644 --- a/api/src/backend/api/attack_paths/sink/neptune.py +++ b/api/src/backend/api/attack_paths/sink/neptune.py @@ -25,7 +25,7 @@ from urllib.parse import urlsplit import neo4j import neo4j.exceptions -from api.attack_paths.retryable_session import RetryableSession +from api.attack_paths.retryable_session import RetryableSession, RetryExhaustedError from api.attack_paths.sink.base import SinkDatabase from api.attack_paths.sink.drop import ( NODE_DELETE_QUERY_TEMPLATE, @@ -85,6 +85,8 @@ def _is_retryable_write_error(exc: Exception) -> bool: class NeptuneSink(SinkDatabase): """Neptune-backed sink. Single database; isolation is label-based.""" + sync_batch_size = env.int("ATTACK_PATHS_NEPTUNE_SYNC_BATCH_SIZE", default=500) + def __init__(self) -> None: self._writer: neo4j.Driver | None = None self._reader: neo4j.Driver | None = None @@ -206,6 +208,7 @@ class NeptuneSink(SinkDatabase): from api.attack_paths.database import ( ClientStatementException, GraphDatabaseQueryException, + NeptuneWriteRetryExhaustedException, WriteQueryNotAllowedException, ) @@ -227,9 +230,17 @@ class NeptuneSink(SinkDatabase): initial_retry_delay_seconds=( NEPTUNE_WRITE_RETRY_DELAY_SECONDS if is_write_session else 0 ), + retry_context="Neptune write" if is_write_session else None, ) yield session_wrapper + except RetryExhaustedError as exc: + last_error = exc.last_error + raise NeptuneWriteRetryExhaustedException( + message=str(exc), + code=getattr(last_error, "code", None), + ) from last_error + except neo4j.exceptions.Neo4jError as exc: if ( default_access_mode == neo4j.READ_ACCESS @@ -291,7 +302,7 @@ class NeptuneSink(SinkDatabase): graph's branching factor. """ from tasks.jobs.attack_paths.config import ( - BATCH_SIZE, + GRAPH_MUTATION_BATCH_SIZE, PROVIDER_RESOURCE_LABEL, get_provider_label, ) @@ -330,7 +341,7 @@ class NeptuneSink(SinkDatabase): total_key="rels", deleted_key="deleted_rels", initial_total=deleted_relationships, - batch_size=BATCH_SIZE, + batch_size=GRAPH_MUTATION_BATCH_SIZE, drop_t0=drop_t0, ) relationship_batches += phase_batches @@ -349,7 +360,7 @@ class NeptuneSink(SinkDatabase): total_key="nodes", deleted_key="deleted_nodes", initial_total=0, - batch_size=BATCH_SIZE, + batch_size=GRAPH_MUTATION_BATCH_SIZE, drop_t0=drop_t0, ) diff --git a/api/src/backend/api/decorators.py b/api/src/backend/api/decorators.py index a055b2252f..2dd2d5fea4 100644 --- a/api/src/backend/api/decorators.py +++ b/api/src/backend/api/decorators.py @@ -1,12 +1,13 @@ import uuid from functools import wraps +from api.attack_paths.database import GraphDatabaseQueryException from api.db_router import READ_REPLICA_ALIAS from api.db_utils import POSTGRES_TENANT_VAR, SET_CONFIG_QUERY, rls_transaction from api.exceptions import ProviderDeletedException -from api.models import Provider, Scan +from api.models import Membership, Provider, Scan, Tenant from django.core.exceptions import ObjectDoesNotExist -from django.db import DatabaseError, connection, transaction +from django.db import DEFAULT_DB_ALIAS, DatabaseError, connection, transaction from rest_framework_json_api.serializers import ValidationError @@ -75,9 +76,11 @@ def handle_provider_deletion(func): """ Decorator that raises `ProviderDeletedException` if provider was deleted during execution. - Catches `ObjectDoesNotExist` and `DatabaseError` (including `IntegrityError`), checks if - provider still exists, and raises `ProviderDeletedException` if not. Otherwise, - re-raises original exception. + Catches `ObjectDoesNotExist`, `DatabaseError` (including `IntegrityError`), and + `GraphDatabaseQueryException`, checks if provider still exists, and raises + `ProviderDeletedException` if not. Graph database errors also check whether the + tenant still exists and has memberships. Otherwise, re-raises the original + exception. Requires `tenant_id` and `provider_id` in kwargs. @@ -92,11 +95,16 @@ def handle_provider_deletion(func): def wrapper(*args, **kwargs): try: return func(*args, **kwargs) - except (ObjectDoesNotExist, DatabaseError): + except (ObjectDoesNotExist, DatabaseError, GraphDatabaseQueryException) as exc: tenant_id = kwargs.get("tenant_id") provider_id = kwargs.get("provider_id") + database_alias = ( + DEFAULT_DB_ALIAS + if isinstance(exc, GraphDatabaseQueryException) + else READ_REPLICA_ALIAS + ) - with rls_transaction(tenant_id, using=READ_REPLICA_ALIAS): + with rls_transaction(tenant_id, using=database_alias): if provider_id is None: scan_id = kwargs.get("scan_id") if scan_id is None: @@ -113,6 +121,13 @@ def handle_provider_deletion(func): raise ProviderDeletedException( f"Provider '{provider_id}' was deleted during the scan" ) from None + if isinstance(exc, GraphDatabaseQueryException) and ( + not Tenant.objects.filter(pk=tenant_id).exists() + or not Membership.objects.filter(tenant_id=tenant_id).exists() + ): + raise ProviderDeletedException( + f"Tenant '{tenant_id}' was deleted during the scan" + ) from None raise return wrapper diff --git a/api/src/backend/api/tests/test_decorators.py b/api/src/backend/api/tests/test_decorators.py index 5c2897730f..0bf58340a1 100644 --- a/api/src/backend/api/tests/test_decorators.py +++ b/api/src/backend/api/tests/test_decorators.py @@ -2,11 +2,12 @@ import uuid from unittest.mock import call, patch import pytest +from api.attack_paths.database import GraphDatabaseQueryException from api.db_utils import POSTGRES_TENANT_VAR, SET_CONFIG_QUERY from api.decorators import handle_provider_deletion, set_tenant from api.exceptions import ProviderDeletedException from django.core.exceptions import ObjectDoesNotExist -from django.db import DatabaseError, IntegrityError +from django.db import DEFAULT_DB_ALIAS, DatabaseError, IntegrityError @pytest.mark.django_db @@ -204,6 +205,106 @@ class TestHandleProviderDeletionDecorator: with pytest.raises(DatabaseError): task_func(tenant_id=str(tenant.id), provider_id=str(provider.id)) + @patch("api.decorators.rls_transaction") + @patch("api.decorators.Provider.objects.filter") + def test_graph_database_error_provider_missing_or_soft_deleted( + self, mock_provider_filter, mock_rls, tenants_fixture + ): + tenant = tenants_fixture[0] + provider_id = str(uuid.uuid4()) + + mock_rls.return_value.__enter__ = lambda s: None + mock_rls.return_value.__exit__ = lambda s, *args: None + mock_provider_filter.return_value.exists.return_value = False + + @handle_provider_deletion + def task_func(**kwargs): + raise GraphDatabaseQueryException("Temporary database not found") + + with pytest.raises(ProviderDeletedException): + task_func(tenant_id=str(tenant.id), provider_id=provider_id) + + @patch("api.decorators.rls_transaction") + @patch("api.decorators.Tenant.objects.filter") + @patch("api.decorators.Provider.objects.filter") + def test_graph_database_error_tenant_missing( + self, mock_provider_filter, mock_tenant_filter, mock_rls, tenants_fixture + ): + tenant = tenants_fixture[0] + provider_id = str(uuid.uuid4()) + + mock_rls.return_value.__enter__ = lambda s: None + mock_rls.return_value.__exit__ = lambda s, *args: None + mock_provider_filter.return_value.exists.return_value = True + mock_tenant_filter.return_value.exists.return_value = False + + @handle_provider_deletion + def task_func(**kwargs): + raise GraphDatabaseQueryException("Temporary database not found") + + with pytest.raises(ProviderDeletedException): + task_func(tenant_id=str(tenant.id), provider_id=provider_id) + + @patch("api.decorators.rls_transaction") + @patch("api.decorators.Membership.objects.filter") + @patch("api.decorators.Tenant.objects.filter") + @patch("api.decorators.Provider.objects.filter") + def test_graph_database_error_tenant_without_memberships( + self, + mock_provider_filter, + mock_tenant_filter, + mock_membership_filter, + mock_rls, + tenants_fixture, + ): + tenant = tenants_fixture[0] + provider_id = str(uuid.uuid4()) + + mock_rls.return_value.__enter__ = lambda s: None + mock_rls.return_value.__exit__ = lambda s, *args: None + mock_provider_filter.return_value.exists.return_value = True + mock_tenant_filter.return_value.exists.return_value = True + mock_membership_filter.return_value.exists.return_value = False + + @handle_provider_deletion + def task_func(**kwargs): + raise GraphDatabaseQueryException("Temporary database not found") + + with pytest.raises(ProviderDeletedException): + task_func(tenant_id=str(tenant.id), provider_id=provider_id) + + @patch("api.decorators.rls_transaction") + @patch("api.decorators.Membership.objects.filter") + @patch("api.decorators.Tenant.objects.filter") + @patch("api.decorators.Provider.objects.filter") + def test_graph_database_error_active_provider_and_tenant_reraises( + self, + mock_provider_filter, + mock_tenant_filter, + mock_membership_filter, + mock_rls, + tenants_fixture, + ): + tenant = tenants_fixture[0] + provider_id = str(uuid.uuid4()) + graph_error = GraphDatabaseQueryException("Temporary database not found") + + mock_rls.return_value.__enter__ = lambda s: None + mock_rls.return_value.__exit__ = lambda s, *args: None + mock_provider_filter.return_value.exists.return_value = True + mock_tenant_filter.return_value.exists.return_value = True + mock_membership_filter.return_value.exists.return_value = True + + @handle_provider_deletion + def task_func(**kwargs): + raise graph_error + + with pytest.raises(GraphDatabaseQueryException) as exc_info: + task_func(tenant_id=str(tenant.id), provider_id=provider_id) + + assert exc_info.value is graph_error + mock_rls.assert_called_once_with(str(tenant.id), using=DEFAULT_DB_ALIAS) + def test_missing_provider_and_scan_raises_assertion(self, tenants_fixture): """Raises AssertionError when neither provider_id nor scan_id in kwargs.""" diff --git a/api/src/backend/api/tests/test_retryable_session.py b/api/src/backend/api/tests/test_retryable_session.py index 8bb457b8b6..09b184c223 100644 --- a/api/src/backend/api/tests/test_retryable_session.py +++ b/api/src/backend/api/tests/test_retryable_session.py @@ -1,7 +1,7 @@ from unittest.mock import MagicMock, patch import pytest -from api.attack_paths.retryable_session import RetryableSession +from api.attack_paths.retryable_session import RetryableSession, RetryExhaustedError from neo4j.exceptions import ServiceUnavailable @@ -24,6 +24,7 @@ class TestRetryableSession: max_retries=3, retry_if=lambda exc: exc is retryable_error, initial_retry_delay_seconds=2, + retry_context="Neptune write", ) assert session.execute_write(work) == "success" @@ -54,6 +55,7 @@ class TestRetryableSession: max_retries=3, retry_if=lambda _: False, initial_retry_delay_seconds=2, + retry_context="Neptune write", ) with pytest.raises(RuntimeError) as exc_info: @@ -83,3 +85,81 @@ class TestRetryableSession: driver_sessions[0].close.assert_called_once_with() driver_sessions[1].close.assert_called_once_with() driver_sessions[2].close.assert_not_called() + + def test_retry_exhaustion_with_context_reports_attempts_and_elapsed_time(self): + error = RuntimeError("still retryable") + driver_sessions = [MagicMock() for _ in range(3)] + for driver_session in driver_sessions: + driver_session.execute_write.side_effect = error + session = RetryableSession( + session_factory=MagicMock(side_effect=driver_sessions), + max_retries=2, + retry_if=lambda _: True, + retry_context="Neptune write", + ) + + with ( + patch( + "api.attack_paths.retryable_session.time.monotonic", + side_effect=[100.0, 127.1234], + ), + pytest.raises(RetryExhaustedError) as exc_info, + ): + session.execute_write(MagicMock()) + + assert exc_info.value.method_name == "execute_write" + assert exc_info.value.attempts == 3 + assert exc_info.value.elapsed_seconds == pytest.approx(27.1234) + assert exc_info.value.last_error is error + assert exc_info.value.__cause__ is error + assert str(exc_info.value) == ( + "Neptune write execute_write failed after 3 attempts over 27.123s. " + "Last error: still retryable" + ) + + def test_retry_exhaustion_with_zero_retries_reports_one_attempt(self): + error = ServiceUnavailable("still unavailable") + driver_session = MagicMock() + driver_session.execute_write.side_effect = error + session = RetryableSession( + session_factory=MagicMock(return_value=driver_session), + max_retries=0, + retry_context="Neptune write", + ) + + with pytest.raises(RetryExhaustedError) as exc_info: + session.execute_write(MagicMock()) + + assert exc_info.value.attempts == 1 + + @patch("api.attack_paths.retryable_session.time.sleep") + @patch("api.attack_paths.retryable_session.random.uniform", return_value=3.0) + def test_contextual_retry_warning_includes_original_error( + self, _mock_uniform, _mock_sleep + ): + error = RuntimeError("retryable detail") + first_session = MagicMock() + first_session.execute_write.side_effect = error + second_session = MagicMock() + second_session.execute_write.return_value = "success" + session = RetryableSession( + session_factory=MagicMock(side_effect=[first_session, second_session]), + max_retries=1, + retry_if=lambda _: True, + initial_retry_delay_seconds=2, + retry_context="Neptune write", + ) + + with patch("api.attack_paths.retryable_session.logger.warning") as mock_warning: + assert session.execute_write(MagicMock()) == "success" + + mock_warning.assert_called_once_with( + "%s %s failed with %s: %s; retry %s/%s in %.3fs", + "Neptune write", + "execute_write", + "RuntimeError", + "retryable detail", + 1, + 1, + 3.0, + ) diff --git a/api/src/backend/api/tests/test_sentry.py b/api/src/backend/api/tests/test_sentry.py index 082f563808..67ba1ee01e 100644 --- a/api/src/backend/api/tests/test_sentry.py +++ b/api/src/backend/api/tests/test_sentry.py @@ -1,6 +1,7 @@ import logging from unittest.mock import MagicMock, patch +import pytest from config.settings import sentry as sentry_settings from config.settings.sentry import before_send @@ -82,6 +83,45 @@ def test_before_send_passes_through_non_ignored_log(): assert result == event +def test_before_send_ignores_cartography_missing_temporary_database_log(): + log_record = _make_log_record( + msg="Cartography job failed with %s for database %s", + name="cartography.graph.job", + args=( + "Neo.ClientError.Database.DatabaseNotFound", + "db-tmp-scan-12345678", + ), + ) + + event = MagicMock() + + assert before_send(event, {"log_record": log_record}) is None + + +@pytest.mark.parametrize( + ("logger_name", "message"), + [ + ( + "cartography.graph.job.worker", + "Neo.ClientError.Database.DatabaseNotFound for db-tmp-scan-12345678", + ), + ( + "cartography.graph.job", + "DatabaseNotFound for db-tmp-scan-12345678", + ), + ( + "cartography.graph.job", + "Neo.ClientError.Database.DatabaseNotFound for db-tenant-12345678", + ), + ], +) +def test_before_send_passes_through_similar_cartography_logs(logger_name, message): + log_record = _make_log_record(msg=message, name=logger_name) + event = MagicMock() + + assert before_send(event, {"log_record": log_record}) is event + + def test_before_send_passes_through_non_ignored_exception(): """Test that before_send passes through exceptions that don't contain ignored exceptions.""" exc_info = (Exception, Exception("Some other error message"), None) diff --git a/api/src/backend/api/tests/test_sink.py b/api/src/backend/api/tests/test_sink.py index 487519923e..c778626b62 100644 --- a/api/src/backend/api/tests/test_sink.py +++ b/api/src/backend/api/tests/test_sink.py @@ -11,7 +11,11 @@ from unittest.mock import MagicMock, patch import neo4j import pytest from api.attack_paths import sink as sink_module -from api.attack_paths.database import GraphDatabaseQueryException +from api.attack_paths.database import ( + GraphDatabaseQueryException, + NeptuneWriteRetryExhaustedException, +) +from api.attack_paths.retryable_session import RetryExhaustedError from api.attack_paths.sink import factory from api.attack_paths.sink.neo4j import DATABASE_NOT_FOUND_CODE, Neo4jSink from api.attack_paths.sink.neptune import ( @@ -123,6 +127,14 @@ class TestSinkFactory: assert mock_driver.call_count == 1 +def test_neo4j_sync_batch_size_defaults_to_1000(): + assert Neo4jSink.sync_batch_size == 1000 + + +def test_neptune_sync_batch_size_defaults_to_500(): + assert NeptuneSink.sync_batch_size == 500 + + class TestGetBackendForScan: """``get_backend_for_scan`` routes by the row's recorded sink backend.""" @@ -372,6 +384,7 @@ class TestNeptuneRetryPolicy: assert ( kwargs["initial_retry_delay_seconds"] == NEPTUNE_WRITE_RETRY_DELAY_SECONDS ) + assert kwargs["retry_context"] == "Neptune write" @patch("api.attack_paths.sink.neptune.RetryableSession") def test_reader_session_does_not_enable_write_retry_policy(self, retryable_session): @@ -384,6 +397,48 @@ class TestNeptuneRetryPolicy: kwargs = retryable_session.call_args.kwargs assert kwargs["retry_if"] is None assert kwargs["initial_retry_delay_seconds"] == 0 + assert kwargs["retry_context"] is None + + def test_writer_retry_exhaustion_preserves_neptune_error_details(self): + message = ( + "Unexpected server exception 'Operation failed due to conflicting " + "concurrent operations (please retry), 0 transactions are currently " + "rolling back.'" + ) + error = neo4j.exceptions.Neo4jError._hydrate_neo4j( + code="BoltProtocol.unexpectedException", + message=message, + ) + retry_error = RetryExhaustedError( + retry_context="Neptune write", + method_name="execute_write", + attempts=4, + elapsed_seconds=27.1234, + last_error=error, + ) + sink = NeptuneSink() + driver = MagicMock() + retryable_session = MagicMock() + retryable_session.execute_write.side_effect = retry_error + + with ( + patch.object(sink, "_get_writer", return_value=driver), + patch( + "api.attack_paths.sink.neptune.RetryableSession", + return_value=retryable_session, + ), + pytest.raises(NeptuneWriteRetryExhaustedException) as exc_info, + ): + with sink.get_session() as session: + session.execute_write(MagicMock()) + + assert exc_info.value.code == "BoltProtocol.unexpectedException" + assert str(exc_info.value) == ( + "BoltProtocol.unexpectedException: Neptune write execute_write failed " + "after 4 attempts over 27.123s. Last error: " + f"{message}" + ) + assert exc_info.value.__cause__ is error class TestNeptuneSinkDropSubgraph: diff --git a/api/src/backend/config/django/base.py b/api/src/backend/config/django/base.py index 04251b4479..a079942600 100644 --- a/api/src/backend/config/django/base.py +++ b/api/src/backend/config/django/base.py @@ -312,8 +312,8 @@ ATTACK_PATHS_SCAN_INACTIVITY_THRESHOLD_MINUTES = env.int( "ATTACK_PATHS_SCAN_INACTIVITY_THRESHOLD_MINUTES", 30 ) ATTACK_PATHS_SCAN_STALE_THRESHOLD_MINUTES = env.int( - "ATTACK_PATHS_SCAN_STALE_THRESHOLD_MINUTES", 2880 -) # 48h + "ATTACK_PATHS_SCAN_STALE_THRESHOLD_MINUTES", 960 +) # 16h # Selects where the persistent attack-paths graph is stored. The scan # temporary database is always Neo4j; only the sink is configurable. diff --git a/api/src/backend/config/settings/sentry.py b/api/src/backend/config/settings/sentry.py index d75f9360e5..a3acbe37bb 100644 --- a/api/src/backend/config/settings/sentry.py +++ b/api/src/backend/config/settings/sentry.py @@ -91,6 +91,13 @@ def before_send(event, hint): log_msg = log_record.getMessage() log_lvl = log_record.levelno + if ( + getattr(log_record, "name", "") == "cartography.graph.job" + and "Neo.ClientError.Database.DatabaseNotFound" in log_msg + and "db-tmp-scan-" in log_msg + ): + return None + # The Neo4j driver logs transient connection errors (defunct # connections, resets) at ERROR level via the `neo4j.io` logger. # `RetryableSession` handles these with retries. If all retries diff --git a/api/src/backend/tasks/jobs/attack_paths/aws.py b/api/src/backend/tasks/jobs/attack_paths/aws.py index 15ecd86a19..4b2b96dd1b 100644 --- a/api/src/backend/tasks/jobs/attack_paths/aws.py +++ b/api/src/backend/tasks/jobs/attack_paths/aws.py @@ -8,6 +8,8 @@ import aioboto3 import boto3 import botocore import neo4j +import neo4j.exceptions +from api.attack_paths.database import DATABASE_NOT_FOUND_CODE from api.models import ( AttackPathsScan as ProwlerAPIAttackPathsScan, ) @@ -347,6 +349,12 @@ def sync_aws_account( ) except Exception as e: + if ( + isinstance(e, neo4j.exceptions.Neo4jError) + and e.code == DATABASE_NOT_FOUND_CODE + ): + raise + logger.info( f"Synced function {func_name} for AWS account {prowler_api_provider.uid} in {time.perf_counter() - func_t0:.3f}s (FAILED)" ) diff --git a/api/src/backend/tasks/jobs/attack_paths/config.py b/api/src/backend/tasks/jobs/attack_paths/config.py index d8ed63a8fc..122fc8a9e0 100644 --- a/api/src/backend/tasks/jobs/attack_paths/config.py +++ b/api/src/backend/tasks/jobs/attack_paths/config.py @@ -10,13 +10,10 @@ NormalizedList = _provider_config.NormalizedList PROVIDER_CONFIGS = _provider_config.PROVIDER_CONFIGS ProviderConfig = _provider_config.ProviderConfig -# Batch size for Neo4j write operations (resource labeling, cleanup) -BATCH_SIZE = env.int("ATTACK_PATHS_BATCH_SIZE", 1000) +# Batch size for graph mutation operations (resource labeling and subgraph deletion) +GRAPH_MUTATION_BATCH_SIZE = env.int("ATTACK_PATHS_GRAPH_MUTATION_BATCH_SIZE", 1000) # Batch size for Postgres findings fetch (keyset pagination page size) FINDINGS_BATCH_SIZE = env.int("ATTACK_PATHS_FINDINGS_BATCH_SIZE", 1000) -# Batch size for temp-to-tenant graph sync (nodes and relationships per cursor page) -SYNC_BATCH_SIZE = env.int("ATTACK_PATHS_SYNC_BATCH_SIZE", 1000) - # Neo4j internal labels (Prowler-specific, not provider-specific) # - `Internet`: Singleton node representing external internet access for exposed-resource queries # - `ProwlerFinding`: Label for finding nodes created by Prowler and linked to cloud resources diff --git a/api/src/backend/tasks/jobs/attack_paths/findings.py b/api/src/backend/tasks/jobs/attack_paths/findings.py index 6cc7ddb2e0..c47c9f1149 100644 --- a/api/src/backend/tasks/jobs/attack_paths/findings.py +++ b/api/src/backend/tasks/jobs/attack_paths/findings.py @@ -21,8 +21,8 @@ from cartography.config import Config as CartographyConfig from celery.utils.log import get_task_logger from prowler.config import config as ProwlerConfig from tasks.jobs.attack_paths.config import ( - BATCH_SIZE, FINDINGS_BATCH_SIZE, + GRAPH_MUTATION_BATCH_SIZE, get_node_uid_field, get_provider_resource_label, get_root_node_label, @@ -135,7 +135,7 @@ def add_resource_label( while labeled_count > 0: result = neo4j_session.run( query, - {"provider_uid": provider_uid, "batch_size": BATCH_SIZE}, + {"provider_uid": provider_uid, "batch_size": GRAPH_MUTATION_BATCH_SIZE}, ) labeled_count = result.single().get("labeled_count", 0) total_labeled += labeled_count diff --git a/api/src/backend/tasks/jobs/attack_paths/scan.py b/api/src/backend/tasks/jobs/attack_paths/scan.py index 32337c6832..e161c9eb8d 100644 --- a/api/src/backend/tasks/jobs/attack_paths/scan.py +++ b/api/src/backend/tasks/jobs/attack_paths/scan.py @@ -372,7 +372,19 @@ def run(tenant_id: str, scan_id: str, task_id: str) -> dict[str, Any]: except Exception as e: exception_message = utils.stringify_exception(e, "Attack Paths scan failed") - logger.exception(exception_message) + temporary_database_missing = ( + isinstance(e, graph_database.GraphDatabaseQueryException) + and e.code == graph_database.DATABASE_NOT_FOUND_CODE + and tmp_database_name in str(e) + ) + if temporary_database_missing: + logger.warning(exception_message) + else: + logger.exception(exception_message) + cleanup_log_level = ( + logging.WARNING if temporary_database_missing else logging.ERROR + ) + cleanup_exc_info = not temporary_database_missing ingestion_exceptions["global_error"] = exception_message # Recover `graph_data_ready` based on how far the swap got @@ -387,19 +399,24 @@ def run(tenant_id: str, scan_id: str, task_id: str) -> dict[str, Any]: ) except Exception: - logger.error( - f"Failed to recover `graph_data_ready` for provider {attack_paths_scan.provider_id}", - exc_info=True, + logger.log( + cleanup_log_level, + "Failed to recover `graph_data_ready` for provider " + f"{attack_paths_scan.provider_id}", + exc_info=cleanup_exc_info, ) # Dropping the temporary database if it still exists try: graph_database.drop_database(tmp_cartography_config.neo4j_database) - except Exception as e: - logger.error( - f"Failed to drop temporary Neo4j database `{tmp_cartography_config.neo4j_database}` during cleanup: {e}", - exc_info=True, + except Exception as cleanup_error: + logger.log( + cleanup_log_level, + "Failed to drop temporary Neo4j database " + f"`{tmp_cartography_config.neo4j_database}` during cleanup: " + f"{cleanup_error}", + exc_info=cleanup_exc_info, ) # Set Attack Paths scan state to FAILED @@ -407,10 +424,12 @@ def run(tenant_id: str, scan_id: str, task_id: str) -> dict[str, Any]: db_utils.finish_attack_paths_scan( attack_paths_scan, StateChoices.FAILED, ingestion_exceptions ) - except Exception as e: - logger.error( - f"Could not mark Attack Paths scan {attack_paths_scan.id} as `FAILED` (row may have been deleted): {e}", - exc_info=True, + except Exception as cleanup_error: + logger.log( + cleanup_log_level, + f"Could not mark Attack Paths scan {attack_paths_scan.id} as `FAILED` " + f"(row may have been deleted): {cleanup_error}", + exc_info=cleanup_exc_info, ) raise diff --git a/api/src/backend/tasks/jobs/attack_paths/sync.py b/api/src/backend/tasks/jobs/attack_paths/sync.py index 00c2c585c7..98a52bd48b 100644 --- a/api/src/backend/tasks/jobs/attack_paths/sync.py +++ b/api/src/backend/tasks/jobs/attack_paths/sync.py @@ -30,7 +30,6 @@ from tasks.jobs.attack_paths.config import ( PROVIDER_CONFIGS, PROVIDER_ISOLATION_PROPERTIES, PROVIDER_RESOURCE_LABEL, - SYNC_BATCH_SIZE, NormalizedList, get_provider_label, get_tenant_label, @@ -116,6 +115,7 @@ def sync_nodes( Source and target sessions are opened sequentially per batch to avoid holding two Bolt connections simultaneously for the entire sync duration. """ + batch_size = sink.sync_batch_size t0 = time.perf_counter() last_id = -1 parents_synced = 0 @@ -137,7 +137,7 @@ def sync_nodes( with graph_database.get_session(source_database) as source_session: result = source_session.run( NODE_FETCH_QUERY, - {"last_id": last_id, "batch_size": SYNC_BATCH_SIZE}, + {"last_id": last_id, "batch_size": batch_size}, ) for record in result: batch_count += 1 @@ -156,17 +156,17 @@ def sync_nodes( for labels, batch in parent_groups.items(): rendered_labels = _render_labels(labels, extra_labels) - for sink_batch in _iter_sink_batches(batch): + for sink_batch in _iter_sink_batches(batch, batch_size): sink.write_nodes(target_database, rendered_labels, sink_batch) for child_label, batch in child_groups.items(): rendered_labels = _render_labels((child_label,), extra_labels) - for sink_batch in _iter_sink_batches(batch): + for sink_batch in _iter_sink_batches(batch, batch_size): sink.write_nodes(target_database, rendered_labels, sink_batch) children_synced += len(batch) for rel_type, batch in rel_groups.items(): - for sink_batch in _iter_sink_batches(batch): + for sink_batch in _iter_sink_batches(batch, batch_size): sink.write_relationships( target_database, rel_type, provider_id, sink_batch ) @@ -205,6 +205,7 @@ def sync_relationships( Source and target sessions are opened sequentially per batch to avoid holding two Bolt connections simultaneously for the entire sync duration. """ + batch_size = sink.sync_batch_size t0 = time.perf_counter() last_id = -1 total_synced = 0 @@ -217,7 +218,7 @@ def sync_relationships( with graph_database.get_session(source_database) as source_session: result = source_session.run( RELATIONSHIPS_FETCH_QUERY, - {"last_id": last_id, "batch_size": SYNC_BATCH_SIZE}, + {"last_id": last_id, "batch_size": batch_size}, ) for record in result: batch_count += 1 @@ -229,7 +230,7 @@ def sync_relationships( break for rel_type, batch in grouped.items(): - for sink_batch in _iter_sink_batches(batch): + for sink_batch in _iter_sink_batches(batch, batch_size): sink.write_relationships( target_database, rel_type, provider_id, sink_batch ) @@ -247,10 +248,9 @@ def sync_relationships( def _iter_sink_batches( rows: list[dict[str, Any]], - batch_size: int | None = None, + batch_size: int, ) -> Iterator[list[dict[str, Any]]]: """Yield final sink write batches after source rows have been transformed.""" - batch_size = SYNC_BATCH_SIZE if batch_size is None else batch_size if batch_size <= 0: raise ValueError("Sink batch size must be greater than zero") diff --git a/api/src/backend/tasks/tasks.py b/api/src/backend/tasks/tasks.py index 7a1ff54131..a91ff85c01 100644 --- a/api/src/backend/tasks/tasks.py +++ b/api/src/backend/tasks/tasks.py @@ -11,6 +11,7 @@ from api.compliance import ( from api.db_router import READ_REPLICA_ALIAS from api.db_utils import delete_related_daily_task, rls_transaction from api.decorators import handle_provider_deletion, set_tenant +from api.exceptions import ProviderDeletedException from api.models import ( Finding, Integration, @@ -666,7 +667,13 @@ class AttackPathsScanRLSTask(RLSTask): scan_id = kwargs.get("scan_id") if tenant_id and scan_id: - logger.error(f"Attack paths scan task {task_id} failed: {exc}") + if isinstance(exc, ProviderDeletedException): + logger.warning( + f"Attack paths scan task {task_id} stopped because its provider " + f"or tenant was deleted: {exc}" + ) + else: + logger.error(f"Attack paths scan task {task_id} failed: {exc}") attack_paths_db_utils.fail_attack_paths_scan(tenant_id, scan_id, str(exc)) diff --git a/api/src/backend/tasks/tests/test_attack_paths_aws.py b/api/src/backend/tasks/tests/test_attack_paths_aws.py new file mode 100644 index 0000000000..dc2c59d614 --- /dev/null +++ b/api/src/backend/tasks/tests/test_attack_paths_aws.py @@ -0,0 +1,102 @@ +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + +import neo4j.exceptions +import pytest +from tasks.jobs.attack_paths import aws + +DATABASE_NOT_FOUND_CODE = "Neo.ClientError.Database.DatabaseNotFound" + + +def _make_neo4j_error(code: str) -> neo4j.exceptions.Neo4jError: + return neo4j.exceptions.Neo4jError._hydrate_neo4j( + code=code, + message="graph query failed", + ) + + +def _resource_functions(failing_sync, following_sync): + return { + "failing_sync": failing_sync, + "following_sync": following_sync, + "permission_relationships": MagicMock(), + "resourcegroupstaggingapi": MagicMock(), + } + + +def test_sync_aws_account_reraises_database_not_found_immediately(): + error = _make_neo4j_error(DATABASE_NOT_FOUND_CODE) + failing_sync = MagicMock(side_effect=error) + following_sync = MagicMock() + + with ( + patch.object( + aws.cartography_aws, + "RESOURCE_FUNCTIONS", + _resource_functions(failing_sync, following_sync), + ), + patch.object(aws.db_utils, "update_attack_paths_scan_progress"), + patch.object(aws.utils, "stringify_exception") as stringify_exception, + patch.object(aws.logger, "warning") as warning, + pytest.raises(neo4j.exceptions.Neo4jError) as exc_info, + ): + aws.sync_aws_account( + SimpleNamespace(uid="123456789012"), + [ + "failing_sync", + "following_sync", + "permission_relationships", + "resourcegroupstaggingapi", + ], + {}, + MagicMock(), + ) + + assert exc_info.value is error + following_sync.assert_not_called() + stringify_exception.assert_not_called() + warning.assert_not_called() + + +@pytest.mark.parametrize( + "error", + [ + _make_neo4j_error("Neo.ClientError.Statement.SyntaxError"), + RuntimeError("resource sync failed"), + ], + ids=["different-neo4j-error", "non-neo4j-error"], +) +def test_sync_aws_account_warns_and_continues_for_other_exceptions(error): + failing_sync = MagicMock(side_effect=error) + following_sync = MagicMock() + + with ( + patch.object( + aws.cartography_aws, + "RESOURCE_FUNCTIONS", + _resource_functions(failing_sync, following_sync), + ), + patch.object(aws.db_utils, "update_attack_paths_scan_progress"), + patch.object( + aws.utils, + "stringify_exception", + return_value="formatted failure", + ), + patch.object(aws.logger, "warning") as warning, + ): + failed_syncs = aws.sync_aws_account( + SimpleNamespace(uid="123456789012"), + [ + "failing_sync", + "following_sync", + "permission_relationships", + "resourcegroupstaggingapi", + ], + {}, + MagicMock(), + ) + + assert failed_syncs == {"failing_sync": "formatted failure"} + following_sync.assert_called_once_with() + warning.assert_called_once() + assert "Continuing to the next AWS sync function" in warning.call_args.args[0] diff --git a/api/src/backend/tasks/tests/test_attack_paths_scan.py b/api/src/backend/tasks/tests/test_attack_paths_scan.py index 3be885a7fc..4409e5f19d 100644 --- a/api/src/backend/tasks/tests/test_attack_paths_scan.py +++ b/api/src/backend/tasks/tests/test_attack_paths_scan.py @@ -1,3 +1,4 @@ +import logging from contextlib import nullcontext from datetime import UTC, datetime, timedelta from types import SimpleNamespace @@ -5,7 +6,9 @@ from unittest.mock import MagicMock, call, patch from uuid import uuid4 import pytest +from api.attack_paths.database import GraphDatabaseQueryException from api.db_utils import rls_transaction +from api.exceptions import ProviderDeletedException from api.models import ( AttackPathsScan, Finding, @@ -250,6 +253,32 @@ class TestAttackPathsRun: mock_starting.assert_not_called() mock_create_db.assert_not_called() + @pytest.mark.parametrize( + ("ingestion_error", "temporary_database_missing"), + [ + (RuntimeError("ingestion boom"), False), + ( + GraphDatabaseQueryException( + message="Graph not found: db-scan-id", + code="Neo.ClientError.Database.DatabaseNotFound", + ), + True, + ), + ( + GraphDatabaseQueryException( + message="Graph not found: db-tenant-id", + code="Neo.ClientError.Database.DatabaseNotFound", + ), + False, + ), + ], + ids=[ + "regular-error", + "temporary-database-missing", + "sink-database-missing", + ], + ) + @patch("tasks.jobs.attack_paths.scan.logger") @patch( "tasks.jobs.attack_paths.scan.utils.stringify_exception", return_value="Cartography failed: ingestion boom", @@ -302,6 +331,9 @@ class TestAttackPathsRun: mock_drop_db, mock_event_loop, mock_stringify, + mock_logger, + ingestion_error, + temporary_database_missing, tenants_fixture, aws_provider, scans_fixture, @@ -321,7 +353,11 @@ class TestAttackPathsRun: session_ctx = MagicMock() session_ctx.__enter__.return_value = mock_session session_ctx.__exit__.return_value = False - ingestion_fn = MagicMock(side_effect=RuntimeError("ingestion boom")) + ingestion_fn = MagicMock(side_effect=ingestion_error) + if temporary_database_missing: + mock_finish.side_effect = DatabaseError( + "Save with update_fields did not affect any rows" + ) with ( patch( @@ -337,13 +373,28 @@ class TestAttackPathsRun: return_value=ingestion_fn, ), ): - with pytest.raises(RuntimeError, match="ingestion boom"): + with pytest.raises(type(ingestion_error)): attack_paths_run(str(tenant.id), str(scan.id), "task-456") failure_args = mock_finish.call_args[0] assert failure_args[0] is attack_paths_scan assert failure_args[1] == StateChoices.FAILED assert failure_args[2] == {"global_error": "Cartography failed: ingestion boom"} + mock_drop_db.assert_called_once_with("db-scan-id") + if temporary_database_missing: + mock_logger.warning.assert_any_call("Cartography failed: ingestion boom") + mock_logger.exception.assert_not_called() + mock_logger.log.assert_called_once_with( + logging.WARNING, + f"Could not mark Attack Paths scan {attack_paths_scan.id} as `FAILED` " + "(row may have been deleted): Save with update_fields did not affect " + "any rows", + exc_info=False, + ) + else: + mock_logger.exception.assert_called_once_with( + "Cartography failed: ingestion boom" + ) @patch( "tasks.jobs.attack_paths.scan.utils.stringify_exception", @@ -1265,6 +1316,33 @@ class TestAttackPathsScanRLSTaskOnFailure: mock_fail.assert_called_once_with("t-1", "s-1", "boom") + def test_on_failure_logs_provider_deletion_as_warning(self): + from tasks.tasks import AttackPathsScanRLSTask + + task = AttackPathsScanRLSTask() + error = ProviderDeletedException("provider deleted") + + with ( + patch("tasks.tasks.logger") as mock_logger, + patch( + "tasks.tasks.attack_paths_db_utils.fail_attack_paths_scan" + ) as mock_fail, + ): + task.on_failure( + exc=error, + task_id="task-abc", + args=(), + kwargs={"tenant_id": "t-1", "scan_id": "s-1"}, + _einfo=None, + ) + + mock_logger.warning.assert_called_once_with( + "Attack paths scan task task-abc stopped because its provider or tenant " + "was deleted: provider deleted" + ) + mock_logger.error.assert_not_called() + mock_fail.assert_called_once_with("t-1", "s-1", "provider deleted") + def test_on_failure_skips_when_missing_kwargs(self): from tasks.tasks import AttackPathsScanRLSTask @@ -1896,7 +1974,7 @@ class TestSyncNodes: mock_source_1.run.return_value = [row] mock_source_2 = MagicMock() mock_source_2.run.return_value = [] - sink = MagicMock() + sink = MagicMock(sync_batch_size=1000) with patch( "tasks.jobs.attack_paths.sync.graph_database.get_session", @@ -1933,7 +2011,7 @@ class TestSyncNodes: src_1.run.return_value = [row] src_2 = MagicMock() src_2.run.return_value = [] - sink = MagicMock() + sink = MagicMock(sync_batch_size=1000) sink.write_nodes.side_effect = lambda *_a, **_kw: call_order.append( "sink:write" ) @@ -1969,18 +2047,15 @@ class TestSyncNodes: src_2.run.return_value = [row_b] src_3 = MagicMock() src_3.run.return_value = [] - sink = MagicMock() + sink = MagicMock(sync_batch_size=1) - with ( - patch( - "tasks.jobs.attack_paths.sync.graph_database.get_session", - side_effect=[ - _make_session_ctx(src_1), - _make_session_ctx(src_2), - _make_session_ctx(src_3), - ], - ), - patch("tasks.jobs.attack_paths.sync.SYNC_BATCH_SIZE", 1), + with patch( + "tasks.jobs.attack_paths.sync.graph_database.get_session", + side_effect=[ + _make_session_ctx(src_1), + _make_session_ctx(src_2), + _make_session_ctx(src_3), + ], ): result = sync_module.sync_nodes("src", "tgt", "t-1", "p-1", sink, []) @@ -2009,17 +2084,14 @@ class TestSyncNodes: src_1.run.return_value = [row] src_2 = MagicMock() src_2.run.return_value = [] - sink = MagicMock() + sink = MagicMock(sync_batch_size=2) - with ( - patch( - "tasks.jobs.attack_paths.sync.graph_database.get_session", - side_effect=[ - _make_session_ctx(src_1), - _make_session_ctx(src_2), - ], - ), - patch("tasks.jobs.attack_paths.sync.SYNC_BATCH_SIZE", 2), + with patch( + "tasks.jobs.attack_paths.sync.graph_database.get_session", + side_effect=[ + _make_session_ctx(src_1), + _make_session_ctx(src_2), + ], ): result = sync_module.sync_nodes( "src", "tgt", "t-1", "p-1", sink, normalized_lists @@ -2037,7 +2109,7 @@ class TestSyncNodes: def test_sync_nodes_empty_source_returns_zero(self): src = MagicMock() src.run.return_value = [] - sink = MagicMock() + sink = MagicMock(sync_batch_size=1000) with patch( "tasks.jobs.attack_paths.sync.graph_database.get_session", @@ -2066,7 +2138,7 @@ class TestSyncRelationships: src_1.run.return_value = [row] src_2 = MagicMock() src_2.run.return_value = [] - sink = MagicMock() + sink = MagicMock(sync_batch_size=1000) sink.write_relationships.side_effect = lambda *_a, **_kw: call_order.append( "sink:write" ) @@ -2104,18 +2176,15 @@ class TestSyncRelationships: src_2.run.return_value = [row_b] src_3 = MagicMock() src_3.run.return_value = [] - sink = MagicMock() + sink = MagicMock(sync_batch_size=1) - with ( - patch( - "tasks.jobs.attack_paths.sync.graph_database.get_session", - side_effect=[ - _make_session_ctx(src_1), - _make_session_ctx(src_2), - _make_session_ctx(src_3), - ], - ), - patch("tasks.jobs.attack_paths.sync.SYNC_BATCH_SIZE", 1), + with patch( + "tasks.jobs.attack_paths.sync.graph_database.get_session", + side_effect=[ + _make_session_ctx(src_1), + _make_session_ctx(src_2), + _make_session_ctx(src_3), + ], ): total = sync_module.sync_relationships("src", "tgt", "p-1", sink) @@ -2140,17 +2209,14 @@ class TestSyncRelationships: src_1.run.return_value = rows src_2 = MagicMock() src_2.run.return_value = [] - sink = MagicMock() + sink = MagicMock(sync_batch_size=2) - with ( - patch( - "tasks.jobs.attack_paths.sync.graph_database.get_session", - side_effect=[ - _make_session_ctx(src_1), - _make_session_ctx(src_2), - ], - ), - patch("tasks.jobs.attack_paths.sync.SYNC_BATCH_SIZE", 2), + with patch( + "tasks.jobs.attack_paths.sync.graph_database.get_session", + side_effect=[ + _make_session_ctx(src_1), + _make_session_ctx(src_2), + ], ): total = sync_module.sync_relationships("src", "tgt", "p-1", sink) @@ -2163,7 +2229,7 @@ class TestSyncRelationships: def test_sync_relationships_empty_source_returns_zero(self): src = MagicMock() src.run.return_value = [] - sink = MagicMock() + sink = MagicMock(sync_batch_size=1000) with patch( "tasks.jobs.attack_paths.sync.graph_database.get_session", @@ -3056,6 +3122,61 @@ class TestCleanupStaleAttackPathsScans: ap_scan.refresh_from_db() assert ap_scan.state == StateChoices.FAILED + @pytest.mark.parametrize( + ("age_seconds", "should_clean"), + [ + (960 * 60 - 1, False), + (960 * 60, False), + (960 * 60 + 1, True), + ], + ) + @patch("tasks.jobs.attack_paths.cleanup.recover_graph_data_ready") + @patch("tasks.jobs.attack_paths.cleanup.graph_database.drop_database") + @patch( + "tasks.jobs.attack_paths.cleanup.rls_transaction", + new=lambda *args, **kwargs: nullcontext(), + ) + @patch("tasks.jobs.attack_paths.cleanup._revoke_task") + @patch("tasks.jobs.attack_paths.cleanup._ping_workers") + def test_stale_threshold_boundary_is_strict( + self, + mock_ping, + mock_revoke, + mock_drop_db, + mock_recover, + age_seconds, + should_clean, + tenants_fixture, + aws_provider, + ): + from tasks.jobs.attack_paths.cleanup import cleanup_stale_attack_paths_scans + + now = datetime.now(tz=UTC) + ap_scan, task_result = self._create_executing_scan( + tenants_fixture[0], + aws_provider, + started_at=now - timedelta(seconds=age_seconds), + worker="live-worker@host", + ) + mock_ping.return_value = ({"live-worker@host"}, set()) + + with patch("tasks.jobs.attack_paths.cleanup.datetime") as mock_datetime: + mock_datetime.now.return_value = now + result = cleanup_stale_attack_paths_scans() + + assert result["cleaned_up_count"] == int(should_clean) + ap_scan.refresh_from_db() + expected_state = StateChoices.FAILED if should_clean else StateChoices.EXECUTING + assert ap_scan.state == expected_state + if should_clean: + mock_revoke.assert_called_once_with(task_result, terminate=True) + mock_drop_db.assert_called_once() + mock_recover.assert_called_once() + else: + mock_revoke.assert_not_called() + mock_drop_db.assert_not_called() + mock_recover.assert_not_called() + @patch("tasks.jobs.attack_paths.cleanup.recover_graph_data_ready") @patch("tasks.jobs.attack_paths.cleanup.graph_database.drop_database") @patch( diff --git a/claude_plugins/prowler/.claude-plugin/plugin.json b/claude_plugins/prowler/.claude-plugin/plugin.json index 7bf822e2e7..c77187c3b0 100644 --- a/claude_plugins/prowler/.claude-plugin/plugin.json +++ b/claude_plugins/prowler/.claude-plugin/plugin.json @@ -22,7 +22,7 @@ "api_key": { "type": "string", "title": "Prowler API key", - "description": "API key token used to authenticate with Prowler Cloud / Prowler App via the Prowler MCP server. Create one at https://cloud.prowler.com.", + "description": "API key token used to authenticate with Prowler (Prowler Cloud, Prowler Private Cloud, or Prowler Local Server) via the Prowler MCP server. Create one at https://cloud.prowler.com.", "sensitive": true, "required": true } diff --git a/claude_plugins/prowler/skills/framework-compliance-triage/SKILL.md b/claude_plugins/prowler/skills/framework-compliance-triage/SKILL.md index 1af29f82b9..f8a9ded0c6 100644 --- a/claude_plugins/prowler/skills/framework-compliance-triage/SKILL.md +++ b/claude_plugins/prowler/skills/framework-compliance-triage/SKILL.md @@ -38,12 +38,12 @@ If the framework is not supported, tell the user, suggest they request it or con ### 1.1 Connect to Prowler Cloud -Verify the Prowler MCP connection by calling `prowler_app_search_providers` — a successful response returns the list of providers. If the call fails, walk the user through troubleshooting: internet connectivity, Prowler Cloud credentials, and permissions on the Prowler Cloud account. +Verify the Prowler MCP connection by calling `prowler_search_providers` — a successful response returns the list of providers. If the call fails, walk the user through troubleshooting: internet connectivity, Prowler Cloud credentials, and permissions on the Prowler Cloud account. For getting accurate information about configurations use `prowler_docs_search` to pull relevant instructions from the Prowler documentation. ### 1.2 Verify the provider is configured (or configure it) -Call `prowler_app_search_providers` to check whether the target provider (AWS account, Azure Subscription, GitHub Account...) exists in the user's Prowler Cloud account. Handle the result based on what's found: +Call `prowler_search_providers` to check whether the target provider (AWS account, Azure Subscription, GitHub Account...) exists in the user's Prowler Cloud account. Handle the result based on what's found: - **Provider not present.** Guide the user through adding and configuring it. Retrieve the relevant connection, credential, and permission instructions with `prowler_docs_search`. - **Provider present but misconfigured** (missing credentials, insufficient permissions, etc.). Walk the user through fixing the configuration, pulling the relevant guidance with `prowler_docs_search`. @@ -57,15 +57,15 @@ Call `prowler_app_search_providers` to check whether the target provider (AWS ac The flow needs at least one completed scan with a compliance report available. -Look for a completed scan first: call `prowler_app_list_scans` with the selected `provider_id` and `state: ["completed"]`, then call `prowler_app_get_compliance_overview` with each `scan_id` to find one whose compliance report is available. If one is found, continue to the next section. +Look for a completed scan first: call `prowler_list_scans` with the selected `provider_id` and `state: ["completed"]`, then call `prowler_get_compliance_overview` with each `scan_id` to find one whose compliance report is available. If one is found, continue to the next section. -If no completed scan has a report, call `prowler_app_list_scans` again with `state: ["available", "executing"]` to detect a scan in progress. +If no completed scan has a report, call `prowler_list_scans` again with `state: ["available", "executing"]` to detect a scan in progress. > **Checkpoint — Scan-in-progress decision** *(conditional: an in-progress scan was detected)* > > Tell the user a scan is already running and ask whether to wait for it to complete or start a fresh one. Wait for the answer. -If no scan is running (or the user chose to start a fresh one), trigger a new scan with `prowler_app_trigger_scan` and the `provider_id`. The link `https://cloud.prowler.com/scans?filter%5Bprovider_uid__in%5D={provider_id}` lets the user monitor progress. +If no scan is running (or the user chose to start a fresh one), trigger a new scan with `prowler_trigger_scan` and the `provider_id`. The link `https://cloud.prowler.com/scans?filter%5Bprovider_uid__in%5D={provider_id}` lets the user monitor progress. When a scan is in progress (either pre-existing and elected to wait, or just triggered), stop the flow and ask the user to return when it's completed — restart this section to re-check the results. @@ -85,7 +85,7 @@ Status taxonomy for failed requirements and their findings: ### Report template -A fresh report is rendered like this (substituting values from the `prowler_app_get_compliance_framework_state_details` Prowler MCP tool response): +A fresh report is rendered like this (substituting values from the `prowler_get_compliance_framework_state_details` Prowler MCP tool response): ````markdown # Compliance report: @@ -120,7 +120,7 @@ A fresh report is rendered like this (substituting values from the `prowler_app_ Resolve the report path for the current `compliance_id` and provider account. -If the file does not exist, call `prowler_app_get_compliance_framework_state_details` for the target scan, render the template above, and write the file with one initialization entry in the activity log. +If the file does not exist, call `prowler_get_compliance_framework_state_details` for the target scan, render the template above, and write the file with one initialization entry in the activity log. If the file exists, read it and compare its `Scan ID` to the target scan from section 1.3. When the scan matches, reuse the file and summarize remaining `[FAIL]` and `[IN PROGRESS]` items in chat. @@ -128,7 +128,7 @@ If the file exists, read it and compare its `Scan ID` to the target scan from se > > Tell the user the report on disk was generated from a different scan and ask whether to refresh it from the new scan. Wait for the answer. -On confirmation, regenerate the failed-requirements section from the new `prowler_app_get_compliance_framework_state_details` response, carry forward the **Global remediation approach** block and the full activity log, and append an activity-log entry noting the scan change. +On confirmation, regenerate the failed-requirements section from the new `prowler_get_compliance_framework_state_details` response, carry forward the **Global remediation approach** block and the full activity log, and append an activity-log entry noting the scan change. Once the file is current, surface the top failing requirements in chat: sort by finding count descending, show the top 5 with their codes and counts, and point to the file path for the full list. @@ -174,7 +174,7 @@ Once approved, the loop proceeds through the batch without further prompts unles Pick the first `[FAIL]` requirement at the top of the failed-requirements section. Move its status and every finding under it to `[IN PROGRESS]`, and add a `**Fix plan**:` sub-bullet describing what will be done. -Call `prowler_app_get_finding_details` for each `finding_id` to retrieve the failing resource and the Prowler Hub's remediation guidance for that check using the tool `prowler_hub_get_check_details` with the `check_id` from the finding details. Summarize the guidance in chat, and append it to the `**Fix plan**` note for each finding. +Call `prowler_get_finding_details` for each `finding_id` to retrieve the failing resource and the Prowler Hub's remediation guidance for that check using the tool `prowler_hub_get_check_details` with the `check_id` from the finding details. Summarize the guidance in chat, and append it to the `**Fix plan**` note for each finding. If a finding does not apply to the target resource (Organization-only check on a User account, paid-tier feature, missing resource type, etc.), set the requirement status to `[SKIPPED]` with the reason, log it in the activity log, and move on without attempting the fix — even if it was missed during §3.2. @@ -194,6 +194,6 @@ Move to the next `[FAIL]` requirement and repeat from section 3.3. > **Checkpoint — Rescan trigger** *(conditional: no `[FAIL]` requirements remain; all are `[FIXED-UNVERIFIED]` or `[SKIPPED]`)* > -> Summarize what was applied, list any `[SKIPPED]` items with reasons, and ask whether to trigger a fresh scan with `prowler_app_trigger_scan` to verify the fixes end-to-end. Wait for the answer. +> Summarize what was applied, list any `[SKIPPED]` items with reasons, and ask whether to trigger a fresh scan with `prowler_trigger_scan` to verify the fixes end-to-end. Wait for the answer. On confirmation, trigger the rescan. When it completes, restart section 2.1 with the carry-forward path — requirements no longer in the new FAIL list move to `[PASS]`, anything still failing reverts to `[FAIL]` with the previous fix attempt visible in the activity log. diff --git a/docs/developer-guide/lighthouse-architecture.mdx b/docs/developer-guide/lighthouse-architecture.mdx index 631e0e7ae4..e12eff9549 100644 --- a/docs/developer-guide/lighthouse-architecture.mdx +++ b/docs/developer-guide/lighthouse-architecture.mdx @@ -132,7 +132,7 @@ The MCP client manages connections to the Prowler MCP Server using a singleton p - **Connection Management**: Retry logic with configurable attempts and delays - **Tool Discovery**: Fetches available tools from MCP server on initialization -- **Authentication Injection**: Automatically adds JWT tokens to `prowler_app_*` tool calls +- **Authentication Injection**: Automatically adds JWT tokens to `prowler_*` tool calls - **Reconnection**: Supports forced reconnection after server restarts Key constants: @@ -141,10 +141,14 @@ Key constants: - `RECONNECT_INTERVAL_MS`: 5 minutes before retry after failure ```typescript -// Authentication injection for prowler_app tools +// Authentication injection for core prowler_ tools (Hub/Docs excluded) private handleBeforeToolCall = ({ name, args }) => { - // Only inject auth for prowler_app_* tools (user-specific data) - if (!name.startsWith("prowler_app_")) { + // Only inject auth for prowler_* tools (user-specific data). + // The legacy prowler_app_ prefix is also accepted for a resilient rollout. + if ( + !name.startsWith("prowler_") && + !name.startsWith("prowler_app_") + ) { return { args }; } @@ -307,7 +311,7 @@ MCP tools are organized into three namespaces based on authentication requiremen | Namespace | Auth Required | Description | |-----------|---------------|-------------| -| `prowler_app_*` | Yes (JWT) | Prowler Cloud and Prowler Local Server tools for findings, providers, scans, resources | +| `prowler_*` | Yes (JWT) | Prowler Cloud, Prowler Private Cloud, and Prowler Local Server tools for findings, providers, scans, resources | | `prowler_hub_*` | No | Security checks catalog, compliance frameworks | | `prowler_docs_*` | No | Documentation search and retrieval | @@ -315,7 +319,7 @@ MCP tools are organized into three namespaces based on authentication requiremen 1. User authenticates with Prowler Local Server, receiving a JWT token 2. Token is stored in session and propagated via `authContextStorage` -3. MCP client injects `Authorization: Bearer ` header for `prowler_app_*` calls +3. MCP client injects `Authorization: Bearer ` header for `prowler_*` calls 4. MCP Server validates token and applies RLS filtering ### Tool Execution Pattern @@ -323,7 +327,7 @@ MCP tools are organized into three namespaces based on authentication requiremen The agent uses meta-tools rather than direct tool registration: ``` -Agent needs data → describe_tool("prowler_app_search_findings") +Agent needs data → describe_tool("prowler_search_findings") → Returns parameter schema → execute_tool with parameters → MCP client adds auth header → MCP Server executes → Results returned to agent → Agent continues reasoning diff --git a/docs/developer-guide/mcp-server.mdx b/docs/developer-guide/mcp-server.mdx index 1b9c49a18b..6d409ccc02 100644 --- a/docs/developer-guide/mcp-server.mdx +++ b/docs/developer-guide/mcp-server.mdx @@ -18,11 +18,15 @@ The Prowler MCP Server brings the entire Prowler ecosystem to AI assistants thro The server follows a modular architecture with three independent sub-servers: -| Sub-Server | Auth Required | Description | -|------------|---------------|-------------| -| `prowler_app` | Yes | Full access to Prowler Cloud and Prowler Local Server features | -| Prowler Hub | No | Security checks catalog with **over 2,000 checks**, fixers, and **70+ compliance frameworks** | -| Prowler Documentation | No | Full-text search and retrieval of official documentation | +| Sub-Server | Tool Prefix | Auth Required | Description | +|------------|-------------|---------------|-------------| +| Prowler | `prowler_` | Yes | Full access to Prowler Cloud, Prowler Private Cloud, and Prowler Local Server features | +| Prowler Hub | `prowler_hub_` | No | Security checks catalog with **over 2,000 checks**, fixers, and **70+ compliance frameworks** | +| Prowler Documentation | `prowler_docs_` | No | Full-text search and retrieval of official documentation | + + +The core Prowler sub-server is served under the `prowler_` tool prefix, while its source lives in the `prowler_app/` module for historical reasons. Tool names use the prefix; import paths use the module. + For a complete list of tools and their descriptions, see the [Tools Reference](/getting-started/basic-usage/prowler-mcp-tools). @@ -413,7 +417,7 @@ uv run prowler-mcp uv run prowler-mcp --transport http --host 0.0.0.0 --port 8000 # Run with environment variables -PROWLER_APP_API_KEY="pk_xxx" uv run prowler-mcp +PROWLER_API_KEY="pk_xxx" uv run prowler-mcp ``` For complete installation and deployment options, see: diff --git a/docs/docs.json b/docs/docs.json index d9ac4c2bd8..b9bc195182 100644 --- a/docs/docs.json +++ b/docs/docs.json @@ -72,9 +72,9 @@ "group": "Prowler MCP", "pages": [ "getting-started/products/prowler-mcp", - "getting-started/installation/prowler-mcp", "getting-started/basic-usage/prowler-mcp", - "getting-started/basic-usage/prowler-mcp-tools" + "getting-started/basic-usage/prowler-mcp-tools", + "getting-started/installation/prowler-mcp" ] }, { diff --git a/docs/getting-started/basic-usage/prowler-mcp-tools.mdx b/docs/getting-started/basic-usage/prowler-mcp-tools.mdx index d225617f0e..acb50c9e81 100644 --- a/docs/getting-started/basic-usage/prowler-mcp-tools.mdx +++ b/docs/getting-started/basic-usage/prowler-mcp-tools.mdx @@ -10,7 +10,7 @@ Complete reference guide for all tools available in the Prowler MCP Server. Tool |----------|------------|------------------------| | Prowler Hub | 10 tools | No | | Prowler Documentation | 2 tools | No | -| Prowler Cloud & Prowler Local Server | 32 tools | Yes | +| Prowler Cloud, Private Cloud & Local Server | 32 tools | Yes | ## Tool Naming Convention @@ -18,11 +18,11 @@ All tools follow a consistent naming pattern with prefixes: - `prowler_hub_*` - Prowler Hub catalog and compliance tools - `prowler_docs_*` - Prowler documentation search and retrieval -- `prowler_app_*` - Prowler Cloud and App (Self-Managed) management tools +- `prowler_*` - Prowler Cloud, Prowler Private Cloud & Prowler Local Server management tools -## Prowler Cloud and Prowler Local Server Tools +## Prowler Tools -Manage Prowler Cloud or Prowler Local Server features. **Requires authentication.** +Manage your Prowler deployment — Prowler Cloud, Prowler Private Cloud, or Prowler Local Server. **Requires authentication.** These tools require a valid API key. See the [Configuration Guide](/getting-started/basic-usage/prowler-mcp) for authentication setup. @@ -32,44 +32,44 @@ These tools require a valid API key. See the [Configuration Guide](/getting-star Tools for searching, viewing, and analyzing security findings across all cloud providers. -- **`prowler_app_search_security_findings`** - Search and filter security findings with advanced filtering options (severity, status, provider, region, service, check ID, date range, muted status) -- **`prowler_app_get_finding_details`** - Get comprehensive details about a specific finding including remediation guidance, check metadata, and resource relationships -- **`prowler_app_get_findings_overview`** - Get aggregate statistics and trends about security findings as a markdown report +- **`prowler_search_security_findings`** - Search and filter security findings with advanced filtering options (severity, status, provider, region, service, check ID, date range, muted status) +- **`prowler_get_finding_details`** - Get comprehensive details about a specific finding including remediation guidance, check metadata, and resource relationships +- **`prowler_get_findings_overview`** - Get aggregate statistics and trends about security findings as a markdown report ### Finding Groups Management Tools for listing finding groups aggregated by check ID, viewing complete group counters, and drilling down into affected resources. -- **`prowler_app_list_finding_groups`** - List latest or historical finding groups with filters for provider, region, service, resource, category, check, severity, status, muted state, delta, date range, and sorting -- **`prowler_app_get_finding_group_details`** - Get complete details for a specific finding group including counters, description, timestamps, and impacted providers -- **`prowler_app_list_finding_group_resources`** - List actionable unmuted resources affected by a finding group by default, including nested resource and provider data plus the `finding_id` for remediation details. Set `include_muted` to include suppressed resources +- **`prowler_list_finding_groups`** - List latest or historical finding groups with filters for provider, region, service, resource, category, check, severity, status, muted state, delta, date range, and sorting +- **`prowler_get_finding_group_details`** - Get complete details for a specific finding group including counters, description, timestamps, and impacted providers +- **`prowler_list_finding_group_resources`** - List actionable unmuted resources affected by a finding group by default, including nested resource and provider data plus the `finding_id` for remediation details. Set `include_muted` to include suppressed resources ### Provider Management Tools for managing cloud provider connections in Prowler. -- **`prowler_app_search_providers`** - Search and view configured providers with their connection status -- **`prowler_app_connect_provider`** - Register and connect a provider with credentials for security scanning -- **`prowler_app_delete_provider`** - Permanently remove a provider from Prowler +- **`prowler_search_providers`** - Search and view configured providers with their connection status +- **`prowler_connect_provider`** - Register and connect a provider with credentials for security scanning +- **`prowler_delete_provider`** - Permanently remove a provider from Prowler ### Scan Management Tools for managing and monitoring security scans. -- **`prowler_app_list_scans`** - List and filter security scans across all providers -- **`prowler_app_get_scan`** - Get comprehensive details about a specific scan (progress, duration, resource counts) -- **`prowler_app_trigger_scan`** - Trigger a manual security scan for a provider -- **`prowler_app_schedule_daily_scan`** - Schedule automated daily scans for continuous monitoring -- **`prowler_app_update_scan`** - Update scan name for better organization +- **`prowler_list_scans`** - List and filter security scans across all providers +- **`prowler_get_scan`** - Get comprehensive details about a specific scan (progress, duration, resource counts) +- **`prowler_trigger_scan`** - Trigger a manual security scan for a provider +- **`prowler_schedule_daily_scan`** - Schedule automated daily scans for continuous monitoring +- **`prowler_update_scan`** - Update scan name for better organization ### Resources Management Tools for searching, viewing, and analyzing cloud resources discovered by Prowler. -- **`prowler_app_list_resources`** - List and filter cloud resources with advanced filtering options (provider, region, service, resource type, tags) -- **`prowler_app_get_resource`** - Get comprehensive details about a specific resource including configuration, metadata, and finding relationships -- **`prowler_app_get_resource_events`** - Get the timeline of cloud API actions performed on a resource (AWS CloudTrail). Shows who did what and when, with full request/response payloads -- **`prowler_app_get_resources_overview`** - Get aggregate statistics about cloud resources as a markdown report +- **`prowler_list_resources`** - List and filter cloud resources with advanced filtering options (provider, region, service, resource type, tags) +- **`prowler_get_resource`** - Get comprehensive details about a specific resource including configuration, metadata, and finding relationships +- **`prowler_get_resource_events`** - Get the timeline of cloud API actions performed on a resource (AWS CloudTrail). Shows who did what and when, with full request/response payloads +- **`prowler_get_resources_overview`** - Get aggregate statistics about cloud resources as a markdown report ### Muting Management @@ -77,33 +77,33 @@ Tools for managing finding muting, including pattern-based bulk muting (mutelist #### Mutelist (Pattern-Based Muting) -- **`prowler_app_get_mutelist`** - Retrieve the current mutelist configuration for the tenant -- **`prowler_app_set_mutelist`** - Create or update the mutelist configuration for pattern-based bulk muting -- **`prowler_app_delete_mutelist`** - Remove the mutelist configuration from the tenant +- **`prowler_get_mutelist`** - Retrieve the current mutelist configuration for the tenant +- **`prowler_set_mutelist`** - Create or update the mutelist configuration for pattern-based bulk muting +- **`prowler_delete_mutelist`** - Remove the mutelist configuration from the tenant #### Mute Rules (Finding-Specific Muting) -- **`prowler_app_list_mute_rules`** - Search and filter mute rules with pagination support -- **`prowler_app_get_mute_rule`** - Retrieve comprehensive details about a specific mute rule -- **`prowler_app_create_mute_rule`** - Create a new mute rule to mute specific findings with documentation and audit trail -- **`prowler_app_update_mute_rule`** - Update a mute rule's name, reason, or enabled status -- **`prowler_app_delete_mute_rule`** - Delete a mute rule from the system +- **`prowler_list_mute_rules`** - Search and filter mute rules with pagination support +- **`prowler_get_mute_rule`** - Retrieve comprehensive details about a specific mute rule +- **`prowler_create_mute_rule`** - Create a new mute rule to mute specific findings with documentation and audit trail +- **`prowler_update_mute_rule`** - Update a mute rule's name, reason, or enabled status +- **`prowler_delete_mute_rule`** - Delete a mute rule from the system ### Attack Paths Analysis Tools for analyzing privilege escalation chains and security misconfigurations using graph-based analysis. Attack Paths maps relationships between cloud resources, permissions, and security findings to detect how privileges can be escalated and how misconfigurations can be exploited. -- **`prowler_app_list_attack_paths_scans`** - List Attack Paths scans with filtering by provider, provider type, and scan state (available, scheduled, executing, completed, failed, cancelled) -- **`prowler_app_list_attack_paths_queries`** - Discover available Attack Paths queries for a completed scan, including query names, descriptions, and required parameters -- **`prowler_app_run_attack_paths_query`** - Execute an Attack Paths query against a completed scan and retrieve graph results with nodes (cloud resources, findings, virtual nodes) and relationships (access paths, role assumptions, security group memberships) -- **`prowler_app_get_attack_paths_cartography_schema`** - Retrieve the Cartography graph schema (node labels, relationships, properties) for writing accurate custom openCypher queries +- **`prowler_list_attack_paths_scans`** - List Attack Paths scans with filtering by provider, provider type, and scan state (available, scheduled, executing, completed, failed, cancelled) +- **`prowler_list_attack_paths_queries`** - Discover available Attack Paths queries for a completed scan, including query names, descriptions, and required parameters +- **`prowler_run_attack_paths_query`** - Execute an Attack Paths query against a completed scan and retrieve graph results with nodes (cloud resources, findings, virtual nodes) and relationships (access paths, role assumptions, security group memberships) +- **`prowler_get_attack_paths_cartography_schema`** - Retrieve the Cartography graph schema (node labels, relationships, properties) for writing accurate custom openCypher queries ### Compliance Management Tools for viewing compliance status and framework details across all cloud providers. -- **`prowler_app_get_compliance_overview`** - Get high-level compliance status across all frameworks for a specific scan or provider, including pass/fail statistics per framework -- **`prowler_app_get_compliance_framework_state_details`** - Get detailed requirement-level breakdown for a specific compliance framework, including failed requirements and associated finding IDs +- **`prowler_get_compliance_overview`** - Get high-level compliance status across all frameworks for a specific scan or provider, including pass/fail statistics per framework +- **`prowler_get_compliance_framework_state_details`** - Get detailed requirement-level breakdown for a specific compliance framework, including failed requirements and associated finding IDs ## Prowler Hub Tools @@ -145,7 +145,7 @@ Search and access official Prowler documentation. **No authentication required.* - Use natural language to interact with the tools through your AI assistant - Tools can be combined for complex workflows - Filter options are available on most list tools -- Authentication is only required for Prowler Cloud and Prowler Local Server tools +- Authentication is only required for Prowler tools (Prowler Cloud, Prowler Private Cloud, or Prowler Local Server) ## Additional Resources diff --git a/docs/getting-started/basic-usage/prowler-mcp.mdx b/docs/getting-started/basic-usage/prowler-mcp.mdx index a06f54375e..120e5c038b 100644 --- a/docs/getting-started/basic-usage/prowler-mcp.mdx +++ b/docs/getting-started/basic-usage/prowler-mcp.mdx @@ -7,10 +7,10 @@ Configure your MCP client to connect to Prowler MCP Server. ## Step 1: Get Your API Key -**Authentication is optional**: Prowler Hub and Prowler Documentation features work without authentication. An API key is only required for Prowler Cloud and Prowler Local Server features. +**Authentication is optional**: Prowler Hub and Prowler Documentation features work without authentication. An API key is only required for Prowler tools (Prowler Cloud, Prowler Private Cloud, or Prowler Local Server). -To use Prowler Cloud or Prowler Local Server features. To get the API key, please refer to the [API Keys](/user-guide/tutorials/prowler-app-api-keys) guide. +An API key authenticates the Prowler tools (Prowler Cloud, Prowler Private Cloud, or Prowler Local Server). To get the API key, please refer to the [API Keys](/user-guide/tutorials/prowler-app-api-keys) guide. Keep the API key secure. Never share it publicly or commit it to version control. @@ -18,12 +18,14 @@ Keep the API key secure. Never share it publicly or commit it to version control ## Step 2: Configure Your MCP Host/Client -Choose the configuration based on your deployment: +Most users should use the **Cloud MCP Server** — it needs no installation and is maintained by Prowler. The [Local MCP Server](#local-mcp-server-configuration) configuration is provided afterwards for users who run the server themselves. -- **HTTP Mode**: Prowler Cloud MCP Server or self-hosted Prowler MCP Server. -- **STDIO Mode**: Local installation only (runs as subprocess of your MCP client). +- **Cloud MCP Server (HTTP)**: the managed server at `https://mcp.prowler.com/mcp` (or your own self-hosted HTTP server). +- **Local MCP Server (STDIO)**: local installation only (runs as a subprocess of your MCP client). -### HTTP Mode +## Cloud MCP Server Configuration (Recommended) + +Connect to the **Cloud MCP Server** at `https://mcp.prowler.com/mcp` over HTTP. This is the recommended path — no installation, always up to date. The same configuration works for a self-hosted HTTP server: just swap the URL. @@ -61,10 +63,10 @@ Choose the configuration based on your deployment: "args": [ "https://mcp.prowler.com/mcp", // or your self-hosted Prowler MCP Server URL "--header", - "Authorization: Bearer ${PROWLER_APP_API_KEY}" + "Authorization: Bearer ${PROWLER_API_KEY}" ], "env": { - "PROWLER_APP_API_KEY": "" + "PROWLER_API_KEY": "" } } } @@ -96,10 +98,10 @@ Choose the configuration based on your deployment: "args": [ "https://mcp.prowler.com/mcp", "--header", - "Authorization: Bearer ${PROWLER_APP_API_KEY}" + "Authorization: Bearer ${PROWLER_API_KEY}" ], "env": { - "PROWLER_APP_API_KEY": "" + "PROWLER_API_KEY": "" } } } @@ -110,8 +112,8 @@ Choose the configuration based on your deployment: Run the following command: ```bash - export PROWLER_APP_API_KEY="" - claude mcp add --transport http prowler https://mcp.prowler.com/mcp --header "Authorization: Bearer $PROWLER_APP_API_KEY" --scope user + export PROWLER_API_KEY="" + claude mcp add --transport http prowler https://mcp.prowler.com/mcp --header "Authorization: Bearer $PROWLER_API_KEY" --scope user ``` @@ -137,9 +139,9 @@ Choose the configuration based on your deployment: -### STDIO Mode +## Local MCP Server Configuration -STDIO mode is only available when running the MCP server locally. +STDIO mode is only available when running the **Local MCP Server** on your own machine. See the [Installation guide](/getting-started/installation/prowler-mcp) to set it up first. @@ -152,7 +154,7 @@ STDIO mode is only available when running the MCP server locally. "command": "uvx", "args": ["/absolute/path/to/prowler/mcp_server/"], "env": { - "PROWLER_APP_API_KEY": "", + "PROWLER_API_KEY": "", "API_BASE_URL": "https://api.prowler.com/api/v1" } } @@ -179,7 +181,7 @@ STDIO mode is only available when running the MCP server locally. "--rm", "-i", "--env", - "PROWLER_APP_API_KEY=", + "PROWLER_API_KEY=", "--env", "API_BASE_URL=https://api.prowler.com/api/v1", "prowlercloud/prowler-mcp" @@ -205,7 +207,7 @@ Restart your MCP client and start asking questions: ## Authentication Methods -Prowler MCP Server supports two authentication methods to connect to Prowler Cloud or Prowler Local Server: +Prowler MCP Server supports two authentication methods to connect to Prowler (Prowler Cloud, Prowler Private Cloud, or Prowler Local Server): ### API Key (Recommended) diff --git a/docs/getting-started/installation/prowler-mcp.mdx b/docs/getting-started/installation/prowler-mcp.mdx index 06f204da6a..5be469b5b2 100644 --- a/docs/getting-started/installation/prowler-mcp.mdx +++ b/docs/getting-started/installation/prowler-mcp.mdx @@ -5,12 +5,12 @@ title: "Installation" There are **two ways** to use Prowler MCP Server: - + **No installation required** - Just configuration Use `https://mcp.prowler.com/mcp` - + **Local installation** - Full control Install via Docker, PyPI, or source code @@ -18,8 +18,8 @@ There are **two ways** to use Prowler MCP Server: -For "Option 1: Managed by Prowler", go directly to the [Configuration Guide](/getting-started/basic-usage/prowler-mcp#hosted-server-configuration-recommended) to set up your Claude Desktop, Cursor, or other MCP client. -**This guide is focused on local installation, "Option 2: Run Locally"**. +For the Cloud MCP Server, go directly to the [Configuration Guide](/getting-started/basic-usage/prowler-mcp#cloud-mcp-server-configuration-recommended) to set up your Claude Desktop, Cursor, or other MCP client. +**This guide is focused on local installation, the Local MCP Server**. ## Installation Methods @@ -51,7 +51,7 @@ Choose one of the following installation methods: ```bash docker run --rm -i \ - -e PROWLER_APP_API_KEY="pk_your_api_key" \ + -e PROWLER_API_KEY="pk_your_api_key" \ -e API_BASE_URL="https://api.prowler.com/api/v1" \ prowlercloud/prowler-mcp ``` @@ -143,7 +143,7 @@ Choose one of the following installation methods: ## Updating Prowler MCP Server -When running Prowler MCP Server locally ("Option 2: Run Locally"), upgrade to the latest version using the same method chosen for installation. The hosted server (`https://mcp.prowler.com/mcp`) is always kept up to date by Prowler and requires no action. +When running the Local MCP Server, upgrade to the latest version using the same method chosen for installation. The Cloud MCP Server (`https://mcp.prowler.com/mcp`) is always kept up to date by Prowler and requires no action. @@ -219,19 +219,19 @@ Configure the server using environment variables: | Variable | Description | Required | Default | |----------|-------------|----------|---------| -| `PROWLER_APP_API_KEY` | Prowler API key | Only for STDIO mode | - | +| `PROWLER_API_KEY` | Prowler API key | Only for STDIO mode | - | | `API_BASE_URL` | Custom Prowler API endpoint | No | `https://api.prowler.com/api/v1` | | `PROWLER_MCP_TRANSPORT_MODE` | Default transport mode (overwritten by `--transport` argument) | No | `stdio` | ```bash macOS/Linux -export PROWLER_APP_API_KEY="pk_your_api_key_here" +export PROWLER_API_KEY="pk_your_api_key_here" export API_BASE_URL="https://api.prowler.com/api/v1" export PROWLER_MCP_TRANSPORT_MODE="http" ``` ```bash Windows PowerShell -$env:PROWLER_APP_API_KEY="pk_your_api_key_here" +$env:PROWLER_API_KEY="pk_your_api_key_here" $env:API_BASE_URL="https://api.prowler.com/api/v1" $env:PROWLER_MCP_TRANSPORT_MODE="http" ``` @@ -246,7 +246,7 @@ Never commit your API key to version control. Use environment variables or secur For convenience, create a `.env` file in the `mcp_server` directory: ```bash .env -PROWLER_APP_API_KEY=pk_your_api_key_here +PROWLER_API_KEY=pk_your_api_key_here API_BASE_URL=https://api.prowler.com/api/v1 PROWLER_MCP_TRANSPORT_MODE=stdio ``` diff --git a/docs/getting-started/products/prowler-mcp.mdx b/docs/getting-started/products/prowler-mcp.mdx index b35ac3d611..93159b2e7d 100644 --- a/docs/getting-started/products/prowler-mcp.mdx +++ b/docs/getting-started/products/prowler-mcp.mdx @@ -8,8 +8,29 @@ title: "Overview" **Preview Feature**: This MCP server is currently under active development. Features and functionality may change. We welcome your feedback—please report any issues on [GitHub](https://github.com/prowler-cloud/prowler/issues) or join our [Slack community](https://goto.prowler.com/slack) to discuss and share your thoughts. +## Quickest Way to Connect: Cloud MCP Server + +The fastest way to get started is the **Cloud MCP Server** at `https://mcp.prowler.com/mcp` — no installation, always up to date, and maintained by Prowler. Just point your MCP client at the URL and authenticate with a [Prowler API key](/user-guide/tutorials/prowler-app-api-keys) as a Bearer token: + +```json +{ + "mcpServers": { + "prowler": { + "url": "https://mcp.prowler.com/mcp", + "headers": { + "Authorization": "Bearer " + } + } + } +} +``` + + + Step-by-step setup for Claude Desktop, Claude Code, Cursor, and other clients. + + -Prowler MCP Server can run as a local instance, or as the hosted **Prowler MCP** at `https://mcp.prowler.com/mcp`. The hosted server also provides tools for Prowler Cloud-specific features such as [Alerts](/user-guide/tutorials/prowler-alerts), [Scan Scheduling](/user-guide/tutorials/prowler-scan-scheduling), and [Findings Triage](/user-guide/tutorials/prowler-app-findings-triage). See [Deployment Options](#deployment-options). +Prefer to run it yourself? The **Local MCP Server** runs on your own machine or infrastructure. The Cloud MCP Server additionally provides tools for Prowler Cloud-specific features such as [Alerts](/user-guide/tutorials/prowler-alerts), [Scan Scheduling](/user-guide/tutorials/prowler-scan-scheduling), and [Findings Triage](/user-guide/tutorials/prowler-app-findings-triage). See [Cloud vs Local MCP Server](#cloud-vs-local-mcp-server). ## What is the Model Context Protocol? @@ -20,9 +41,9 @@ The [Model Context Protocol (MCP)](https://modelcontextprotocol.io) is an open s The Prowler MCP Server provides three main integration points: -### 1. Prowler Cloud and Prowler Local Server +### 1. Prowler Cloud, Private Cloud & Local Server -Full access to Prowler Cloud and Prowler Local Server for: +Full access to your Prowler deployment — Prowler Cloud, Prowler Private Cloud, or Prowler Local Server — for: - **Findings Analysis**: Query, filter, and analyze security findings across all your cloud environments - **Provider Management**: Create, configure, and manage your configured Prowler providers (AWS, Azure, GCP, etc.) - **Scan Orchestration**: Trigger on-demand scans and schedule recurring security assessments @@ -48,12 +69,53 @@ Search and retrieve official Prowler documentation: ## MCP Server Architecture -The following diagram illustrates the Prowler MCP Server architecture and its integration points: +The following diagram illustrates the Prowler MCP Server architecture and its integration points. MCP clients connect to either the **Cloud MCP Server** (recommended) or a **Local MCP Server**; both expose the same tools and reach the same Prowler backends: -![Prowler MCP Server Schema](/images/prowler_mcp_schema.png) +```mermaid +flowchart LR + subgraph HOSTS["MCP Clients"] + chat["Chat Interfaces
(Claude Desktop, LobeChat)"] + ide["IDEs and Code Editors
(Claude Code, Cursor)"] + apps["Other AI Applications
(5ire, custom agents)"] + end + + subgraph SERVERS["Prowler MCP Server"] + direction TB + cloud["☁️ Cloud MCP Server (Recommended)
mcp.prowler.com/mcp · HTTP
Managed by Prowler · always up to date
Adds Cloud-only tools (Alerts,
Scan Scheduling, Findings Triage)"] + local["💻 Local MCP Server
Self-run · STDIO or HTTP
Python 3.12+ or Docker
You manage updates"] + end + + subgraph TOOLS["Prowler MCP Tools"] + prowler_tools["prowler_* tools
(API key or JWT auth)
Findings · Providers · Scans
Resources · Muting · Compliance
Attack Paths"] + hub_tools["prowler_hub_* tools
(no auth)
Checks Catalog · Check Code
Fixers · Compliance Frameworks"] + docs_tools["prowler_docs_* tools
(no auth)
Search · Document Retrieval"] + end + + api["Prowler API (REST)
Cloud · Private Cloud · Local Server"] + hub["hub.prowler.com
(REST)"] + docs["docs.prowler.com
(Mintlify)"] + + chat -->|HTTP| cloud + ide -->|HTTP| cloud + apps -->|HTTP| cloud + chat -->|STDIO or HTTP| local + ide -->|STDIO or HTTP| local + apps -->|STDIO or HTTP| local + + cloud --> prowler_tools + cloud --> hub_tools + cloud --> docs_tools + local --> prowler_tools + local --> hub_tools + local --> docs_tools + + prowler_tools -->|REST| api + hub_tools -->|REST| hub + docs_tools -->|REST| docs +``` The architecture shows how AI assistants connect through the MCP protocol to access Prowler's three main components: -- Prowler Cloud and Prowler Local Server for security operations +- Prowler Cloud, Prowler Private Cloud, or Prowler Local Server for security operations - Prowler Hub for security knowledge - Prowler Documentation for guidance and reference. @@ -92,8 +154,8 @@ REQUIREMENTS: DATA TO FETCH: Use these MCP tools in this order: -1. Prowler app list providers - To get all available configured provider in the account -2. Prowler app get latest findings - To get findings information, if there are so many you can use the filter_fields to get less information, or pagination to get in different batches +1. Prowler list providers - To get all available configured provider in the account +2. Prowler get latest findings - To get findings information, if there are so many you can use the filter_fields to get less information, or pagination to get in different batches 3. For most critical findings you can get more context and remediation with Prowler Hub to get remediations for example DESIGN REQUIREMENTS: @@ -130,66 +192,48 @@ Generate the complete HTML file and display it > -## Deployment Options +## Cloud vs Local MCP Server -Prowler MCP Server can be used in three ways: +There are two ways to run the Prowler MCP Server. For almost everyone, the **Cloud MCP Server** is the right choice — it needs no installation and is maintained by Prowler. The **Local MCP Server** exists for users who need to run it on their own machine or infrastructure. -### 1. Prowler Cloud MCP Server +| | ☁️ **Cloud MCP Server** (Recommended) | 💻 **Local MCP Server** | +|---|---|---| +| **Endpoint** | `https://mcp.prowler.com/mcp` | Runs on your machine or infrastructure | +| **Setup** | Just configure your MCP client | Install via Docker, or source | +| **Transport** | HTTP | STDIO (subprocess) or self-hosted HTTP | +| **Maintenance** | Managed by Prowler, always up to date | You manage updates | +| **Requirements** | None (just an MCP client) | Python 3.12+ or Docker | +| **Cloud-only tools** | ✅ Alerts, Scan Scheduling, Findings Triage | ❌ Not available | +| **Authentication** | API key or JWT token | API key/JWT (HTTP) or env vars (STDIO) | -**Use Prowler's managed MCP server at `https://mcp.prowler.com/mcp`** +### ☁️ Cloud MCP Server (Recommended) -- No installation required. -- Managed and maintained by Prowler team. -- Authentication to Prowler Cloud or Prowler Local Server via API key or JWT token. -- Includes tools for Prowler Cloud-specific features such as Alerts, Scan Scheduling, and Findings Triage. +Prowler's managed MCP server at `https://mcp.prowler.com/mcp`. No installation, always up to date, and it includes tools for Prowler Cloud-specific features such as Alerts, Scan Scheduling, and Findings Triage. This is the path we recommend for nearly all users — go straight to the [Configuration guide](/getting-started/basic-usage/prowler-mcp#cloud-mcp-server-configuration-recommended). -### 2. Local STDIO Mode +### 💻 Local MCP Server -**Run the server locally on your machine** +Run the server yourself when you need full control over the deployment. It connects to Prowler Cloud, Prowler Private Cloud, or Prowler Local Server and can run in two modes: -- Runs as a subprocess of the MCP client. -- Possibility to connect to Prowler Local Server. -- Authentication to Prowler Cloud or Prowler Local Server via environment variables. -- Requires Python 3.12+ or Docker. +- **STDIO mode** — the server runs as a subprocess of your MCP client. Authentication via environment variables. +- **Self-hosted HTTP mode** — deploy your own remote HTTP server. Authentication via API key or JWT token. -### 3. Self-Hosted HTTP Mode - -**Deploy your own remote MCP server** - -- Full control over deployment. -- Possibility to connect to Prowler Local Server. -- Authentication to Prowler Local Server via API key or JWT token. -- Requires Python 3.12+ or Docker. - -## Requirements - -Requirements vary based on deployment option: - -**For Prowler Cloud MCP Server:** -- Prowler Cloud account and API key (only for Prowler Cloud and Prowler Local Server features) - -**For self-hosted STDIO/HTTP Mode:** -- Python 3.12+ or Docker -- Network access to: - - `https://hub.prowler.com` (for Prowler Hub) - - `https://docs.prowler.com` (for Prowler Documentation) - - Prowler Cloud API or Prowler Local Server API (for Prowler Cloud and Prowler Local Server features) +Both require Python 3.12+ or Docker, plus network access to `https://hub.prowler.com` (Prowler Hub), `https://docs.prowler.com` (Prowler Documentation), and the Prowler API or Prowler Local Server API (Prowler features). See the [Installation guide](/getting-started/installation/prowler-mcp) to get started. -**No Authentication Required**: Prowler Hub and Prowler Documentation features work without authentication in both deployment options. A Prowler API key is only required to access Prowler Cloud or Prowler Local Server features. +**No Authentication Required**: Prowler Hub and Prowler Documentation features work without authentication on both the Cloud and Local MCP Server. A Prowler API key is only required to access Prowler features (Prowler Cloud, Prowler Private Cloud, or Prowler Local Server). ## Next Steps - - Install the Prowler MCP Server using uv or Docker - - Configure your MCP client to connect to the server + Connect your MCP client to the Cloud MCP Server + + + Explore all available tools and capabilities - - Explore all available tools and capabilities + + Run the Local MCP Server yourself using Docker, source, or uvx diff --git a/docs/images/lighthouse-architecture.mmd b/docs/images/lighthouse-architecture.mmd index 47407544e9..6798801fb8 100644 --- a/docs/images/lighthouse-architecture.mmd +++ b/docs/images/lighthouse-architecture.mmd @@ -15,7 +15,7 @@ flowchart TB llm["LLM Provider
(OpenAI / Bedrock / OpenAI-compatible)"] subgraph MCP["Prowler MCP Server"] - app_tools["prowler_app_* tools
(auth required)"] + app_tools["prowler_* tools
(auth required)"] hub_tools["prowler_hub_* tools
(no auth)"] docs_tools["prowler_docs_* tools
(no auth)"] end @@ -29,7 +29,7 @@ flowchart TB agent <-->|LLM API| llm agent --> metatools metatools --> mcpclient - mcpclient -->|MCP HTTP · Bearer token
for prowler_app_* only| app_tools + mcpclient -->|MCP HTTP · Bearer token
for prowler_* only| app_tools mcpclient -->|MCP HTTP| hub_tools mcpclient -->|MCP HTTP| docs_tools app_tools -->|REST| api diff --git a/docs/images/organizations/authentication-details.png b/docs/images/organizations/authentication-details.png index 2b4ae782cf..aec5060afd 100644 Binary files a/docs/images/organizations/authentication-details.png and b/docs/images/organizations/authentication-details.png differ diff --git a/docs/images/organizations/onboarding-flow.svg b/docs/images/organizations/onboarding-flow.svg index f6e11fc0a3..b5ba7858a0 100644 --- a/docs/images/organizations/onboarding-flow.svg +++ b/docs/images/organizations/onboarding-flow.svg @@ -3,41 +3,37 @@ + + + Onboarding Flow - + 1 - Create Management - Account Role - - Quick Create or Manual - Allows Prowler to - discover your org - structure + Start the Wizard + + In Prowler Cloud + Enter your Org ID + and OU/root target - - - - - 2 - Deploy StackSet - - In AWS Console - Creates ProwlerScan - role in every - member account + Deploy the Roles + + Single CF Stack + Management role + + StackSet to members + in one CF stack @@ -46,11 +42,11 @@ 3 - Run the Wizard - - In Prowler Cloud - Discovers accounts, - tests connections + Discover & Connect + + In Prowler Cloud + Discovers accounts, + tests connections @@ -59,13 +55,13 @@ 4 - Launch Scans - - Automatic - Scans run on all - connected accounts - on your schedule + Launch Scans + + Automatic + Scans run on all + connected accounts + on your schedule - Steps 1 and 2 are done once in AWS | Steps 3 and 4 are done in Prowler Cloud + Step 2 runs once in AWS | Steps 1, 3 and 4 are in Prowler Cloud diff --git a/docs/images/organizations/two-roles-architecture.svg b/docs/images/organizations/two-roles-architecture.svg index c67588b049..f8c40d5b21 100644 --- a/docs/images/organizations/two-roles-architecture.svg +++ b/docs/images/organizations/two-roles-architecture.svg @@ -47,7 +47,7 @@ - Deploy: Quick Create link or Manual + Deploy: single stack or standalone @@ -86,7 +86,7 @@ - Deploy: via CloudFormation StackSet + Deploy: StackSet (single stack) Prowler discovers diff --git a/docs/images/prowler_mcp_schema.mmd b/docs/images/prowler_mcp_schema.mmd index 96973546f6..2b8508cd18 100644 --- a/docs/images/prowler_mcp_schema.mmd +++ b/docs/images/prowler_mcp_schema.mmd @@ -1,29 +1,40 @@ flowchart LR - subgraph HOSTS["MCP Hosts"] + subgraph HOSTS["MCP Clients"] chat["Chat Interfaces
(Claude Desktop, LobeChat)"] ide["IDEs and Code Editors
(Claude Code, Cursor)"] apps["Other AI Applications
(5ire, custom agents)"] end - subgraph MCP["Prowler MCP Server"] - app_tools["prowler_app_* tools
(JWT or API key auth)
Findings · Providers · Scans
Resources · Muting · Compliance
Attack Paths"] + subgraph SERVERS["Prowler MCP Server"] + direction TB + cloud["☁️ Cloud MCP Server (Recommended)
mcp.prowler.com/mcp · HTTP
Managed by Prowler · always up to date
Adds Cloud-only tools (Alerts,
Scan Scheduling, Findings Triage)"] + local["💻 Local MCP Server
Self-run · STDIO or HTTP
Python 3.12+ or Docker
You manage updates"] + end + + subgraph TOOLS["Prowler MCP Tools"] + prowler_tools["prowler_* tools
(API key or JWT auth)
Findings · Providers · Scans
Resources · Muting · Compliance
Attack Paths"] hub_tools["prowler_hub_* tools
(no auth)
Checks Catalog · Check Code
Fixers · Compliance Frameworks"] docs_tools["prowler_docs_* tools
(no auth)
Search · Document Retrieval"] end - api["Prowler API
(REST)"] + api["Prowler API (REST)
Cloud · Private Cloud · Local Server"] hub["hub.prowler.com
(REST)"] docs["docs.prowler.com
(Mintlify)"] - chat -->|STDIO or HTTP| app_tools - chat -->|STDIO or HTTP| hub_tools - chat -->|STDIO or HTTP| docs_tools - ide -->|STDIO or HTTP| app_tools - ide -->|STDIO or HTTP| hub_tools - ide -->|STDIO or HTTP| docs_tools - apps -->|STDIO or HTTP| app_tools - apps -->|STDIO or HTTP| hub_tools - apps -->|STDIO or HTTP| docs_tools - app_tools -->|REST| api + chat -->|HTTP| cloud + ide -->|HTTP| cloud + apps -->|HTTP| cloud + chat -->|STDIO or HTTP| local + ide -->|STDIO or HTTP| local + apps -->|STDIO or HTTP| local + + cloud --> prowler_tools + cloud --> hub_tools + cloud --> docs_tools + local --> prowler_tools + local --> hub_tools + local --> docs_tools + + prowler_tools -->|REST| api hub_tools -->|REST| hub docs_tools -->|REST| docs diff --git a/docs/images/prowler_mcp_schema.png b/docs/images/prowler_mcp_schema.png deleted file mode 100644 index 8a8884fa5e..0000000000 Binary files a/docs/images/prowler_mcp_schema.png and /dev/null differ diff --git a/docs/user-guide/providers/aws/organizations.mdx b/docs/user-guide/providers/aws/organizations.mdx index c0e9ae0e6b..b48df1c442 100644 --- a/docs/user-guide/providers/aws/organizations.mdx +++ b/docs/user-guide/providers/aws/organizations.mdx @@ -2,6 +2,8 @@ title: 'AWS Organizations in Prowler' --- +import { VersionBadge } from "/snippets/version-badge.mdx" + **Using Prowler Cloud?** You can onboard your entire AWS Organization through the UI with automatic account discovery, OU-aware tree selection, and bulk connection testing — no scripts or YAML files required. @@ -71,11 +73,43 @@ The additional fields in CSV header output are as follows: ## Deploying Prowler IAM Roles Across AWS Organizations + + When onboarding multiple AWS accounts into Prowler Cloud, it is important to deploy the Prowler Scan IAM Role in each account. The most efficient way to do this across an AWS Organization is by leveraging AWS CloudFormation StackSets, which rolls out infrastructure—like IAM roles—to all accounts centrally from the Management or Delegated Admin account. -When using Infrastructure as Code (IaC), Terraform is recommended to manage this deployment systematically. +### Native CloudFormation StackSet Deployment (Recommended) -### Recommended Approach +The [Prowler Scan IAM Role CloudFormation template](https://github.com/prowler-cloud/prowler/blob/master/permissions/templates/cloudformation/prowler-scan-role.yml) can deploy the role across your entire AWS Organization on its own—no third-party modules required. When launched in the **Management Account** (or a **Delegated Administrator** account) with `DeployStackSet=true` and `EnableOrganizations=true`, it creates a service-managed CloudFormation StackSet that rolls the ProwlerScan role out to every account under the target Organizational Unit (or the organization root), and keeps new accounts covered automatically through auto-deployment. + +To deploy from the CloudFormation console: open **CloudFormation → Create stack → With new resources**, choose **Upload a template file** and select `prowler-scan-role.yml` (or paste its S3 URL), then set the parameters below on the **Specify stack details** step. Leave the **Configure stack options** step at its defaults. + +Deploy a single CloudFormation Stack in the Management Account with the following parameters: + +| Parameter | Description | Default | +| --- | --- | --- | +| `ExternalId` | External ID provided by Prowler Cloud to secure role assumption. | — | +| `DeployLocalRole` | Create the ProwlerScan role in this (Management) account. | `true` | +| `DeployStackSet` | Create a service-managed StackSet that deploys the role to member accounts. | `false` | +| `AWSOrganizationalUnitId` | Target OU (`ou-xxxx-yyyyyyyy`) or organization root (`r-xxxx`) for the StackSet. Required when `DeployStackSet=true`. | `""` | +| `DeployFromDelegatedAdmin` | Set to `true` when deploying from a Delegated Administrator account instead of the Management Account (uses `CallAs: DELEGATED_ADMIN`). | `false` | +| `EnableOrganizations` | Add AWS Organizations permissions to the Management Account role: read-only account discovery plus the StackSet-management permissions the deployment needs. Set to `true` when deploying in the Management Account. | `false` | +| `FailureTolerancePercentage` | Percentage of accounts in which the StackSet operation can fail before CloudFormation stops the operation. | `10` | +| `RetainStacksOnAccountRemoval` | Keep the role in an account after it leaves the Organization or OU. | `false` | + + +On the review step, select **"I acknowledge that AWS CloudFormation might create IAM resources with custom names"** — the template provisions the named `ProwlerScan` IAM role, so the stack requires the `CAPABILITY_NAMED_IAM` capability and fails without this acknowledgment. (The quick-create link handles this for you.) + + + +The service-managed StackSet does **not** deploy to the Management Account itself. Keeping `DeployLocalRole=true` ensures the role also exists there, so a single stack covers both the Management and member accounts. + +Trusted access for CloudFormation StackSets must be enabled in the Organization (see the note at the top of this page) before `DeployStackSet` will work. + +Deploying for the CLI or a self-hosted Prowler (not Prowler Cloud)? Also set `AccountId` to the account you assume the role from and `IAMPrincipal` to your identity — the defaults target Prowler Cloud. See [Aligning the trust policy with your identity](/user-guide/providers/aws/authentication#trust-policy-align-iamprincipal-with-your-identity). + + + +### Alternative: Deploy with Terraform - **Use StackSets** from the **Management Account** (or a Delegated Admin/Security Account). - **Use Terraform** to orchestrate the deployment. diff --git a/docs/user-guide/providers/aws/role-assumption.mdx b/docs/user-guide/providers/aws/role-assumption.mdx index b714beafbd..54dd0e214a 100644 --- a/docs/user-guide/providers/aws/role-assumption.mdx +++ b/docs/user-guide/providers/aws/role-assumption.mdx @@ -77,6 +77,15 @@ The template requires the following parameters: - **AccountId:** *(Optional)* AWS Account ID that will assume the role (default: Prowler Cloud account) - **IAMPrincipal:** *(Optional)* The IAM principal allowed to assume the role (default: `role/prowler*`) + +From the CLI you assume the role with **your own** identity, not from Prowler Cloud. The `AccountId` and `IAMPrincipal` defaults target Prowler Cloud, so set **`AccountId`** to the account you run Prowler from and **`IAMPrincipal`** to your identity (for example `role/` or `user/`). Otherwise `sts:AssumeRole` fails with `AccessDenied`. See [Aligning the trust policy with your identity](/user-guide/providers/aws/authentication#trust-policy-align-iamprincipal-with-your-identity). + + + +To deploy the role across an entire AWS Organization from a single stack (Management Account role plus a service-managed StackSet for the member accounts), the template also accepts `DeployLocalRole`, `DeployStackSet`, `AWSOrganizationalUnitId`, `DeployFromDelegatedAdmin`, `EnableOrganizations`, `FailureTolerancePercentage`, and `RetainStacksOnAccountRemoval`. See [AWS Organizations in Prowler](/user-guide/providers/aws/organizations#native-cloudformation-stackset-deployment-recommended) for the full parameter reference. + + + When running Prowler CLI, include the External ID using the `-I/--external-id` flag: ```sh diff --git a/docs/user-guide/tutorials/aws-organizations-bulk-provisioning.mdx b/docs/user-guide/tutorials/aws-organizations-bulk-provisioning.mdx index 2e8b5eb62e..188ac3540d 100644 --- a/docs/user-guide/tutorials/aws-organizations-bulk-provisioning.mdx +++ b/docs/user-guide/tutorials/aws-organizations-bulk-provisioning.mdx @@ -272,6 +272,8 @@ python aws_org_generator.py \ 4. Deploy to all organizational units 5. Use a unique external ID (e.g., `prowler-org-2024-abc123`) + Alternatively, deploy the same template as a **single stack** with `DeployStackSet=true` and `AWSOrganizationalUnitId` set to your root/OU ID — it creates the StackSet for you. See [Native CloudFormation StackSet Deployment](../providers/aws/organizations#native-cloudformation-stackset-deployment-recommended). + {/* TODO: Add screenshot of CloudFormation StackSets deployment */} diff --git a/docs/user-guide/tutorials/prowler-app-attack-paths.mdx b/docs/user-guide/tutorials/prowler-app-attack-paths.mdx index 32df3c50ba..22e56e5634 100644 --- a/docs/user-guide/tutorials/prowler-app-attack-paths.mdx +++ b/docs/user-guide/tutorials/prowler-app-attack-paths.mdx @@ -283,7 +283,7 @@ In addition to the upstream schema, Prowler enriches the graph with: AI assistants connected through Prowler MCP Server can fetch the exact Cartography schema for the active scan via the - `prowler_app_get_attack_paths_cartography_schema` tool. This guarantees that + `prowler_get_attack_paths_cartography_schema` tool. This guarantees that generated queries match the schema version pinned by the running Prowler release. @@ -427,10 +427,10 @@ Attack Paths capabilities are also available through the [Prowler MCP Server](/g The following MCP tools are available for Attack Paths: -- **`prowler_app_list_attack_paths_scans`** - List and filter Attack Paths scans. -- **`prowler_app_list_attack_paths_queries`** - Discover available queries for a completed scan. -- **`prowler_app_run_attack_paths_query`** - Execute a query and retrieve graph results with nodes and relationships. -- **`prowler_app_get_attack_paths_cartography_schema`** - Retrieve the Cartography graph schema for custom openCypher queries. +- **`prowler_list_attack_paths_scans`** - List and filter Attack Paths scans. +- **`prowler_list_attack_paths_queries`** - Discover available queries for a completed scan. +- **`prowler_run_attack_paths_query`** - Execute a query and retrieve graph results with nodes and relationships. +- **`prowler_get_attack_paths_cartography_schema`** - Retrieve the Cartography graph schema for custom openCypher queries. ### Example Questions diff --git a/docs/user-guide/tutorials/prowler-app-scan-configuration.mdx b/docs/user-guide/tutorials/prowler-app-scan-configuration.mdx index 8d50f0072c..60e0db2f4c 100644 --- a/docs/user-guide/tutorials/prowler-app-scan-configuration.mdx +++ b/docs/user-guide/tutorials/prowler-app-scan-configuration.mdx @@ -8,7 +8,7 @@ import { SubscriptionBanner } from "/snippets/subscription-banner.mdx" -Scan Configuration lets you override, per provider, specific values in the default configuration Prowler's checks use during a scan. Each configuration modifies how specific checks behave, e.g.: thresholds, allowed values, retention windows, and you attach it to the providers that you want to use it on their next scan. +Scan Configuration lets you override, per provider, specific values in the default configuration Prowler's checks use during a scan. Each configuration can modify how specific checks behave, such as thresholds, allowed values, and retention windows, or exclude checks and services from the scan scope. Attach it to the providers that should use it on their next scan. @@ -54,6 +54,24 @@ gcp: storage_min_retention_days: 30 ``` +### Limiting the Scan Scope + + + +Use `excluded_checks` to skip individual checks and `excluded_services` to skip every check in a service for the matching provider type: + +```yaml +aws: + excluded_checks: + - s3_bucket_public_access + excluded_services: + - ec2 +``` + + +When a Scan Configuration excludes checks or services, Prowler calculates overviews, aggregations, and other result-based information from the reduced scan scope. The displayed information reflects only the checks and services that ran, not a complete assessment of the provider. Consider the applied Scan Configuration when interpreting totals and security posture. + + ## Creating a Scan Configuration diff --git a/docs/user-guide/tutorials/prowler-cloud-aws-organizations.mdx b/docs/user-guide/tutorials/prowler-cloud-aws-organizations.mdx index 8bb6e320ae..8b2124e004 100644 --- a/docs/user-guide/tutorials/prowler-cloud-aws-organizations.mdx +++ b/docs/user-guide/tutorials/prowler-cloud-aws-organizations.mdx @@ -8,12 +8,14 @@ import { SubscriptionBanner } from "/snippets/subscription-banner.mdx" -Prowler Cloud enables you to onboard all AWS accounts in your Organization through a single guided wizard. Instead of connecting accounts one by one, you can discover every account in your AWS Organization, select the ones you want to monitor, test connectivity, and launch scans — all from the Prowler Cloud UI. +Prowler Cloud onboards every AWS account in your Organization through a single guided wizard. Instead of connecting accounts one by one, you can discover every account in your AWS Organization, select the ones you want to monitor, test connectivity, and launch scans — all from the Prowler Cloud UI. For CLI-based multi-account scanning, see [AWS Organizations in Prowler CLI](/user-guide/providers/aws/organizations). +To follow this guide you need an active [Prowler Cloud](https://cloud.prowler.com) account and access to your AWS Organization [management account](https://docs.aws.amazon.com/organizations/latest/userguide/orgs_introduction.html) (or a registered delegated administrator account). + ## Overview ### Individual Accounts vs Organizations @@ -25,225 +27,17 @@ For CLI-based multi-account scanning, see [AWS Organizations in Prowler CLI](/us ### How It Works -Before using the AWS Organizations wizard, you need to deploy **two Identity and Access Management (IAM) roles** in your AWS environment. The onboarding follows this sequence: + + +Onboarding deploys the **ProwlerScan Identity and Access Management (IAM) role** in your management account and in every member account. A **single CloudFormation stack** — launched from the wizard's **Create Stack in Management Account** button ([Step 2](#step-2-authenticate-with-your-management-account)) — creates the management account role **and** a service-managed StackSet that rolls the role out to your member accounts in one operation. Prefer to deploy the roles yourself? See [Deploy the Roles Manually](#deploy-the-roles-manually). - Onboarding flow: 1. Create Management Account Role (Quick Create or Manual), 2. Deploy StackSet, 3. Run the Wizard, 4. Launch Scans + Onboarding flow: 1. Start the Wizard, 2. Deploy the Roles (single CloudFormation stack), 3. Discover and Connect, 4. Launch Scans -## Key Concepts +## Step 1: Start the Organization Wizard -### What Is an External ID? - -An **External ID** is a security token that Prowler generates unique to your tenant. When Prowler assumes the IAM role in your AWS account, it presents this External ID to prove its identity. - -This prevents the [confused deputy problem](https://docs.aws.amazon.com/IAM/latest/UserGuide/confused-deputy.html) — a scenario where an unauthorized party could trick AWS into granting access to your account. By requiring the External ID, only your specific Prowler tenant can assume the role. - -You don't need to create the External ID yourself — Prowler generates it automatically and displays it in the wizard for you to copy. - -### Two Roles Architecture - -Prowler requires **two separate IAM roles** deployed in different places, each with a distinct purpose: - -| Role | Where it lives | What it does | How to deploy it | -|------|---------------|--------------|------------------| -| **ProwlerScan** (management account) | Your management (root) account only | Discovers the Organization structure **and** scans the management account. Has additional Organizations discovery permissions. | Via **Quick Create** link or **manually** in the IAM Console ([Step 1](#step-1-create-the-management-account-role)). Cannot be deployed via StackSet. | -| **ProwlerScan** (member accounts) | Every member account | Scans the account for security findings. | Via **CloudFormation StackSet** ([Step 2](#step-2-deploy-the-cloudformation-stackset)). Automated across all accounts. | - - - Two Roles Architecture: ProwlerScan in management account (Quick Create or Manual, discovery + scanning) and ProwlerScan in member accounts (via StackSet, scanning only) - - - -**Same name, different permissions.** Both roles are named `ProwlerScan` — Prowler expects a consistent role name across all accounts. The management account role has the same scanning permissions as member accounts, plus additional Organizations discovery permissions (see [Step 1](#step-1-create-the-management-account-role) for the full list). - - -### What Is a CloudFormation StackSet? - -A [CloudFormation StackSet](https://docs.aws.amazon.com/AWSCloudFormation/latest/UserGuide/what-is-cfnstacksets.html) lets you deploy the same CloudFormation template across multiple AWS accounts in a single operation. Prowler uses a StackSet to deploy the **ProwlerScan** IAM role into every member account of your organization, so you don't have to create the role manually in each account. - -## Prerequisites - -### Prowler Cloud Account - -You need an active [Prowler Cloud](https://cloud.prowler.com) account. Each AWS account you connect will count as a provider in your subscription. See [Billing Impact](#billing-impact) for details. - -### AWS Organization Enabled - -Your AWS environment must have [AWS Organizations](https://docs.aws.amazon.com/organizations/latest/userguide/orgs_introduction.html) enabled. You will need access to the **management account** (or a delegated administrator account) to provide the Organization ID and IAM Role ARN. - -## Step 1: Create the Management Account Role - -The first role you need to create is the **management account role**. This role allows Prowler to discover your Organization structure — listing accounts, OUs, and hierarchy. - - -**StackSets do not deploy to the management account.** Organizational CloudFormation StackSets with service-managed permissions only target member accounts — this is an AWS limitation, not a Prowler one. You must create the management account role separately, either via the Quick Create link ([Option A](#option-a-quick-create-link-fastest)) or manually ([Option B](#option-b-create-the-role-manually)). - - - -**The role must be named `ProwlerScan`** — the same name as the role deployed to member accounts via StackSet. Prowler expects a consistent role name across all accounts in the Organization. If you use a different name, connection tests and scans will fail for the management account. - - -### Option A: Quick Create Link (Fastest) - -The Prowler wizard provides a one-click link that opens the AWS Console with the CloudFormation template pre-configured. This creates a **CloudFormation Stack** (not a StackSet) that deploys the ProwlerScan role with Organizations permissions enabled in your management account. - - -**[Open Quick Create Stack in AWS Console →](https://us-east-1.console.aws.amazon.com/cloudformation/home?region=us-east-1#/stacks/quickcreate?templateURL=https%3A%2F%2Fprowler-cloud-public.s3.eu-west-1.amazonaws.com%2Fpermissions%2Ftemplates%2Faws%2Fcloudformation%2Fprowler-scan-role.yml&stackName=Prowler¶m_EnableOrganizations=true)** - -Opens the CloudFormation Console with the Prowler scan role template and `EnableOrganizations=true` pre-filled. You will need to enter the **ExternalId** parameter manually — copy it from the Prowler wizard ([Step 4](#step-4-authenticate-with-your-management-account)). - - -1. Click **[Open Quick Create Stack in AWS Console →](https://us-east-1.console.aws.amazon.com/cloudformation/home?region=us-east-1#/stacks/quickcreate?templateURL=https%3A%2F%2Fprowler-cloud-public.s3.eu-west-1.amazonaws.com%2Fpermissions%2Ftemplates%2Faws%2Fcloudformation%2Fprowler-scan-role.yml&stackName=Prowler¶m_EnableOrganizations=true)** or use the **Create Stack in Management Account** button in the Prowler wizard (which also pre-fills the ExternalId). -2. Enter the **ExternalId** parameter if not pre-filled. -3. Check **"I acknowledge that AWS CloudFormation might create IAM resources with custom names"** and click **Create stack**. -4. Wait for the stack to reach **CREATE_COMPLETE** status. - -Take note of the **Role ARN** from the stack's **Outputs** tab — you will need it in the wizard. - -### Option B: Create the Role Manually - -1. Sign in to the [AWS IAM Console](https://console.aws.amazon.com/iam/) in your **management account**. - -2. Go to **Roles > Create role** and select **Custom trust policy**. - -3. Paste the following trust policy. This allows Prowler Cloud to assume the role using your tenant's External ID (you will get this from the Prowler wizard in [Step 3](#step-3-start-the-organization-wizard)): - -```json -{ - "Version": "2012-10-17", - "Statement": [ - { - "Effect": "Allow", - "Principal": { - "AWS": "arn:aws:iam::232136659152:root" - }, - "Action": "sts:AssumeRole", - "Condition": { - "StringEquals": { - "sts:ExternalId": "" - }, - "StringLike": { - "aws:PrincipalArn": "arn:aws:iam::232136659152:role/prowler*" - } - } - } - ] -} -``` - -Replace `` with the External ID shown in the Prowler wizard. - -4. Attach the following AWS managed policies: - - **SecurityAudit** - - **ViewOnlyAccess** - - This allows Prowler to also scan the management account for security findings, just like any other account. - -5. Create an additional inline policy with the following permissions. These are specific to the management account and allow Prowler to discover your Organization structure: - -```json -{ - "Version": "2012-10-17", - "Statement": [ - { - "Sid": "ProwlerOrganizationDiscovery", - "Effect": "Allow", - "Action": [ - "organizations:DescribeAccount", - "organizations:DescribeOrganization", - "organizations:ListAccounts", - "organizations:ListAccountsForParent", - "organizations:ListOrganizationalUnitsForParent", - "organizations:ListRoots", - "organizations:ListTagsForResource" - ], - "Resource": "*" - }, - { - "Sid": "ProwlerStackSetManagement", - "Effect": "Allow", - "Action": [ - "organizations:RegisterDelegatedAdministrator", - "iam:CreateServiceLinkedRole" - ], - "Resource": "*" - } - ] -} -``` - - -You can optionally restrict the `Resource` field to your specific Organization ARN (e.g., `arn:aws:organizations::123456789012:organization/o-abc123def4`) instead of `"*"` to minimize the blast radius. - - -6. Name the role **`ProwlerScan`** and click **Create role**. Take note of the **Role ARN** — you will need it in the Prowler wizard. - -The ARN follows this format: `arn:aws:iam:::role/ProwlerScan` - - -The role **must** be named `ProwlerScan`. Do not use a different name. - - - -If you just created the role, it may take up to **60 seconds** for AWS to propagate it. If you get an error in the Prowler wizard, wait a moment and try again. - - -## Step 2: Deploy the CloudFormation StackSet - -After creating the management account role, the next step is to deploy the **ProwlerScan** role to your member accounts using a CloudFormation StackSet. This is the recommended method for consistent, scalable deployment across your entire organization. - -The StackSet uses **service-managed permissions**, which means AWS Organizations handles the cross-account deployment automatically — you don't need to create execution roles manually in each account. The StackSet deploys the ProwlerScan IAM role in every target member account, enabling Prowler to assume that role for cross-account scanning. - - -**Trusted access required:** CloudFormation StackSets must have trusted access enabled in your management account. Verify this in the AWS Console under **AWS Organizations > Settings > Trusted access for AWS CloudFormation StackSets**. - - - -**The Quick Create link creates a Stack, not a StackSet.** The link in the Prowler wizard creates a CloudFormation **Stack** that deploys the ProwlerScan role in your management account only ([Step 1](#step-1-create-the-management-account-role)). To deploy the role across **member accounts**, you must create a StackSet manually as described below. AWS does not support Quick Create links for StackSets. - - - -**[Open StackSets Console →](https://us-east-1.console.aws.amazon.com/cloudformation/home?region=us-east-1#/stacksets/create)** - -Opens the CloudFormation StackSets creation page directly. You will need to paste the template URL and ExternalId manually. - - -1. Click the link above or navigate to **CloudFormation > StackSets > Create StackSet** in your management account. -2. Choose **Service-managed permissions**. -3. Select **Amazon S3 URL** as the template source and paste the following URL: - ``` - https://prowler-cloud-public.s3.eu-west-1.amazonaws.com/permissions/templates/aws/cloudformation/prowler-scan-role.yml - ``` -4. Set the **ExternalId** parameter to the External ID shown in the Prowler wizard. -5. Choose your deployment targets (entire organization or specific OUs). -6. Select the AWS regions where you want the role deployed. -7. Click **Create StackSet**. - -### Verify StackSet Deployment - -After deploying, verify that all stack instances completed successfully: - -1. In the CloudFormation Console, go to **StackSets** and select your Prowler StackSet. -2. Click the **Stack instances** tab. -3. Confirm that all instances show **Status: CURRENT** and **Stack status: CREATE_COMPLETE**. - -Deployment typically takes **2–5 minutes** for medium-sized organizations. Large organizations (500+ accounts) may take longer. - - -**Prefer Terraform?** You can deploy the ProwlerScan role using Terraform instead. See the [StackSets deployment guide](/user-guide/providers/aws/organizations#deploying-prowler-iam-roles-across-aws-organizations) for the Terraform module. - - -### Key Considerations - -- **Service-managed permissions**: Always select **Service-managed permissions** when creating the StackSet. This lets AWS Organizations manage the deployment automatically across current and future member accounts. -- **Least privilege**: The ProwlerScan role deployed by the StackSet uses `SecurityAudit` and `ViewOnlyAccess` — AWS managed policies that grant read-only access — plus a small set of additional read-only permissions for services not covered by those policies. See the [CloudFormation template](https://prowler-cloud-public.s3.eu-west-1.amazonaws.com/permissions/templates/aws/cloudformation/prowler-scan-role.yml) for the full list. Prowler does not make any changes to your accounts. -- **New accounts**: When you add new accounts to your AWS Organization, the StackSet automatically deploys the ProwlerScan role to them if you targeted the organization root or the relevant OU. Combined with Prowler's 6-hour automatic sync, new accounts are onboarded end-to-end without manual intervention. -- **Management account**: Organizational StackSets **do not deploy to the management account itself**. If you want to scan the management account, you need to create the ProwlerScan role there separately using a regular CloudFormation Stack. - -## Step 3: Start the Organization Wizard - -Now that both roles are deployed — the management account role (Step 1) and the ProwlerScan role in member accounts (Step 2) — you can start the Prowler wizard. +The Prowler wizard walks you through the entire flow: deploying both roles from a single CloudFormation stack, discovering your accounts, testing connectivity, and launching scans. ### Open the Wizard @@ -280,29 +74,50 @@ Now that both roles are deployed — the management account role (Step 1) and th Click **Next** to proceed to the authentication phase. -## Step 4: Authenticate with Your Management Account +## Step 2: Authenticate with Your Management Account -The wizard's **Authentication Details** page guides you through three actions: deploying the roles in AWS, entering the management account Role ARN, and confirming the deployment. +The **Authentication Details** page guides you through three actions: deploying the roles in AWS, entering the deployment account Role ARN, and confirming the deployment. The deployment account is either the management account or, when delegated administrator mode is selected, the delegated administrator account. ### External ID -The wizard displays a **Prowler External ID** at the top — auto-generated and unique to your tenant. Click the copy icon to copy it. You will need this External ID for both the management account Stack and the member accounts StackSet. +The wizard displays a **Prowler External ID** at the top — auto-generated and unique to your tenant. Click the copy icon to copy it. The External ID is pre-filled into the deployment link, and the single stack applies it to both the management account role and the member-account StackSet. Learn more in [What Is an External ID?](#what-is-an-external-id). ### Deploy the Roles -The wizard provides two deployment actions: + -1. **Create Stack in Management Account** — opens a Quick Create link that deploys the ProwlerScan role with `EnableOrganizations=true` in your management account ([Step 1](#step-1-create-the-management-account-role)). The External ID is pre-filled. +The wizard deploys the deployment account role and the member-account StackSet in a **single** CloudFormation Stack: -2. **Open StackSets Console** — links to the CloudFormation StackSets console where you create a StackSet for member accounts ([Step 2](#step-2-deploy-the-cloudformation-stackset)). Copy the template URL shown in the wizard and paste the External ID manually. + +**Prefer to use your own role?** You do not have to use the Quick Create template. Create the ProwlerScan role yourself — through the IAM Console, Terraform, or your own CloudFormation [(Following this guide)](#deploy-the-roles-manually) — and paste its ARN into the Role ARN field below. The role must use the external ID from the earlier step and include the trust policy and permissions described in [Deploy the Roles Manually](#deploy-the-roles-manually). + + +1. **Organizational Unit or Root ID** — enter the AWS OU (`ou-xxxx-yyyyyyyy`) or organization root (`r-xxxx`) you want to onboard. Prowler rolls the ProwlerScan role out to every member account under this target. Find it in the [AWS Organizations Console](https://console.aws.amazon.com/organizations/); use the **root ID** (`r-`) to cover the entire organization or an **OU ID** (`ou-`) to target a specific unit. + +2. *(Optional)* Check **"I'm deploying from a delegated administrator account"** if you launch the stack from a delegated administrator account instead of the management account. + +3. **Create Stack in Management Account** — or **Create Stack in Delegated Administrator Account** when delegated administrator mode is selected — opens a Quick Create link that deploys, in a single stack: the ProwlerScan role in the account where you launch the stack (`DeployLocalRole`, with `EnableOrganizations=true`) **and** a service-managed StackSet (`DeployStackSet`) that rolls the role out to your member accounts. The External ID, OU/Root ID, and deployment options are pre-filled. - Authentication Details form showing External ID, two deployment buttons (Create Stack in Management Account and Open StackSets Console), Management Account Role ARN field, and deployment confirmation checkbox + Authentication Details form showing External ID, Organizational Unit or Root ID field, delegated administrator checkbox, deployment account stack button, deployment account Role ARN field, and deployment confirmation checkbox -### Enter the Management Account Role ARN + +**Finding your Organizational Unit or Root ID.** In the [AWS Organizations Console](https://console.aws.amazon.com/organizations/) the root (`r-…`) and OU (`ou-…`) IDs appear in the account tree, or run these from your management account: -Paste the **Role ARN** of the management account role you created in [Step 1](#step-1-create-the-management-account-role) into the **Management Account Role ARN** field. +```bash +# Root ID — deploys the role to the entire organization +aws organizations list-roots --query 'Roots[0].Id' --output text + +# OU IDs under the root — to target a specific unit instead +aws organizations list-organizational-units-for-parent --parent-id r-xxxx \ + --query 'OrganizationalUnits[].{Name:Name,Id:Id}' --output table +``` + + +### Enter the Deployment Account Role ARN + +Paste the **Role ARN** created by the stack above into the **Management Account Role ARN** field or, when delegated administrator mode is selected, the **Delegated Administrator Account Role ARN** field. The ARN follows this format: ``` @@ -312,12 +127,16 @@ arn:aws:iam:::role/ProwlerScan For example: `arn:aws:iam::123456789012:role/ProwlerScan` - Management Account Role ARN field in the Authentication Details form + Deployment account Role ARN field in the Authentication Details form + +It may take up to **60 seconds** for AWS to generate the IAM Role ARN after the stack completes. If the wizard reports an error, wait a moment and try again. + + ### Confirm and Discover -1. Check the box: **"The Stack and StackSet have been successfully deployed in AWS"**. +1. Check the box: **"The Stack has been successfully deployed in AWS"**. 2. Click **Authenticate**. Here's what happens behind the scenes: @@ -325,7 +144,7 @@ Here's what happens behind the scenes: - An asynchronous discovery is triggered to query your AWS Organization structure. - You will see a **"Gathering AWS Accounts..."** spinner — this typically takes **30 seconds to 2 minutes** depending on your organization size. -## Step 5: Select Accounts to Scan +## Step 3: Select Accounts to Scan ### Understanding the Tree View @@ -336,6 +155,7 @@ Once discovery completes, the wizard displays a **hierarchical tree view** of yo - The tree supports up to **5 levels of nesting** (Root > OUs > Sub-OUs > Accounts). +- If you deployed the stack for just one OU, that OU will be preselected in the tree. - **Selecting an OU** automatically selects all accounts within it. - **Individual overrides**: deselect specific accounts even if the parent OU is selected. - The header shows **"X of Y accounts selected"** to track your selection. @@ -352,14 +172,12 @@ Only **ACTIVE** accounts can be selected for scanning: | **CLOSED** | No | Account has been closed. | -**Your existing data is safe.** If an AWS account is already connected to Prowler as an individual provider, it will appear in the tree with a checkmark indicator. +**Your existing data is safe.** If an AWS account is already connected to Prowler as an individual provider, it appears in the tree with a checkmark indicator. When you proceed: - The existing provider is **linked** to the organization — it is **not** duplicated. - All your **historical scan data and findings are preserved** — nothing is overwritten. - There is **no additional billing** — the existing provider is reused. - -This is completely safe. You are simply associating the account with the organization for easier management. ### Custom Aliases @@ -368,14 +186,9 @@ You can edit the display name for each account before connecting. This alias is ### Blocked Accounts -Some accounts may appear as **blocked** (grayed out, not selectable). This happens when: -- The account is **already linked to a different organization** in Prowler (`linked_to_other_organization`). +Some accounts may appear as **blocked** (grayed out, not selectable) when the account is **already linked to a different organization** in Prowler (`linked_to_other_organization`). Hover over the blocked account to see the specific reason. -Hover over the blocked account to see the specific reason. - -## Step 6: Test Connections - -### How Connection Testing Works +## Step 4: Test Connections Click **Test Connections** to verify that Prowler can assume the **ProwlerScan** role in each selected member account. @@ -383,154 +196,225 @@ Click **Test Connections** to verify that Prowler can assume the **ProwlerScan** Connection testing in progress with spinners on each account -- Each account shows a real-time status indicator: - - **Spinner** — test in progress - - **Green checkmark (✓)** — connection successful - - **Red icon (✗)** — connection failed (hover to see the error) - -### All Tests Pass +Each account shows a real-time status indicator: +- **Spinner** — test in progress +- **Green checkmark (✓)** — connection successful +- **Red icon (✗)** — connection failed (hover to see the error) If every account connects successfully, you automatically advance to the next step. -### Some Tests Fail +### When Some Tests Fail -An error banner appears: **"There was a problem connecting to some accounts."** - -You have two options: +An error banner appears: **"There was a problem connecting to some accounts."** You have two options: **a) Fix and retry:** 1. Go to the AWS Console and verify the StackSet deployed to the failing accounts. 2. Check that the External ID in the StackSet matches the one shown in Prowler. 3. Return to Prowler and click **Test Connections** — only the **failed accounts are re-tested** (smart retry). Accounts that already passed are not tested again. - - Test Connections button - - **b) Skip and continue:** -Click **Skip Connection Validation** to proceed with only the accounts that connected successfully. The failed accounts will not be scanned. +Click **Skip Connection Validation** to proceed with only the accounts that connected successfully. The failed accounts will not be scanned. This option is only available when at least one account connected successfully. Connection test results showing failed accounts with error banner and Skip Connection Validation button - -**Skip Connection Validation** is only available when at least one account connected successfully. - +If **no accounts** connected successfully, you cannot proceed. Fix the underlying connection issues — see [Troubleshooting](#troubleshooting) — and retry before launching scans. -### All Tests Fail - -If **no accounts** connected successfully, you cannot proceed: - -> *"No accounts connected successfully. Fix the connection errors and retry before launching scans."* - -You must fix the underlying connection issues before continuing. See [Updating Credentials](#updating-credentials) below. - -### Updating Credentials - -If connection tests fail, here's how to fix common issues: - -1. Open the [CloudFormation Console](https://console.aws.amazon.com/cloudformation/) and check that your StackSet instances show **CREATE_COMPLETE** for the failing accounts. If not, update the StackSet to include the missing OUs. -2. Compare the **ExternalId** parameter in your StackSet with the External ID displayed in the Prowler wizard. They must match exactly. -3. After fixing the issue in AWS, return to Prowler and click **Test Connections**. Only the previously failed accounts will be re-tested. - -## Step 7: Launch Scans - -### Choose Scan Schedule +## Step 5: Launch Scans The Organizations wizard uses the same schedule controls described in [Scan Scheduling](/user-guide/tutorials/prowler-scan-scheduling#schedule-options). -### Launch - -Click **Save**, **Save and launch scan**, or **Launch scan**, depending on the selected schedule option. A toast notification confirms whether the schedule was saved, scans were launched, or both. The toast includes a link to the **Scans** page. Prowler redirects to the **Providers** page. - -Scans are only launched for accounts that are accessible (passed connection testing) and were selected. +Click **Save**, **Save and launch scan**, or **Launch scan**, depending on the selected schedule option. A toast notification confirms whether the schedule was saved, scans were launched, or both, and includes a link to the **Scans** page. Prowler then redirects to the **Providers** page. Scans launch only for accounts that passed connection testing and were selected. Launch Scan step showing Accounts Connected confirmation, scan schedule selector, and Launch scan button -### What Happens Next - +After launching: - Scans appear in the **Scans** page as they start and complete. - Results populate the **Overview** and **Findings** pages. -- Prowler runs an **automatic sync every 6 hours** to detect new accounts added to your Organization or accounts that have been removed. New accounts are onboarded automatically based on the parent OU configuration. +- Prowler runs an **automatic sync every 6 hours** to detect accounts added to or removed from your Organization. New accounts under the targeted OU or root are onboarded automatically. ## Billing Impact Each AWS account you connect through the Organizations wizard counts as one **provider** in your Prowler Cloud subscription. - **Already-connected accounts**: if an account was already linked as a provider, adding it to the organization does **not** incur additional billing. The existing provider is reused. -- **Large organizations**: connecting a 500-account organization will result in up to 500 providers on your subscription. Review your plan limits before proceeding. +- **Large organizations**: connecting a 500-account organization results in up to 500 providers on your subscription. Review your plan limits before proceeding. - **Deleted providers**: if you later remove an account, the deleted provider no longer counts toward your subscription. For pricing details, see [Prowler Cloud Pricing](https://prowler.com/pricing). ## Troubleshooting -### Invalid AWS Organization ID +### Only Some Accounts Connect -*"Must be a valid AWS Organization ID"* +Discovery succeeds and the tree view appears, but only one account — or a handful — passes the connection test. This almost always means the ProwlerScan role reached the deployment account but not every member account. -- Verify the Organization ID format: `o-` followed by 10–32 lowercase alphanumeric characters (e.g., `o-abc123def4`) -- Copy it directly from the [AWS Organizations Console](https://console.aws.amazon.com/organizations/) to avoid typos +- **Confirm the StackSet deployed.** Open the [CloudFormation Console](https://console.aws.amazon.com/cloudformation/) in the deployment account, select your Prowler StackSet, open the **Stack instances** tab, and confirm every instance shows **Status: CURRENT** and **Stack status: CREATE_COMPLETE**. Instances still in progress or in a failed state explain the missing accounts. +- **Check the targeted OU or root.** The single stack only rolls the role out to accounts under the **Organizational Unit or Root ID** you entered in [Step 2](#step-2-authenticate-with-your-management-account). Accounts in other OUs are not covered — redeploy targeting the organization root (`r-`) or add the missing OUs. +- **Verify the deployment account.** The role is created only in the account where you launched the stack. If you deployed from a **delegated administrator account**, confirm that account is a **registered delegated administrator** for CloudFormation StackSets (registered through AWS Organizations), not just a regular member account. A regular member account cannot create a service-managed StackSet, so only its own role is created — leaving every other account without the role. +- **Suspended accounts** cannot be scanned. Deselect them and proceed. -### Invalid IAM Role ARN +### No Accounts Connect -*"Must be a valid IAM Role ARN"* +No account passes the connection test. -- Verify the ARN format: `arn:aws:iam::<12-digit-account-id>:role/` -- Copy the ARN directly from the [IAM Console](https://console.aws.amazon.com/iam/) in your management account +- **External ID mismatch.** Compare the **ExternalId** parameter in your StackSet with the External ID shown in the Prowler wizard. They must match exactly. +- **StackSet not deployed.** Confirm the StackSet exists and its instances reached **CREATE_COMPLETE**. If you deployed the roles manually, verify [trusted access for CloudFormation StackSets](#member-account-role-stackset) is enabled. +- **IP-based policies.** If your accounts restrict access by IP, allow the [Prowler Cloud egress IPs](/security/networking). -### Authentication Failed +### Authentication Fails or Times Out -*"Authentication failed. Please verify the StackSet deployment and Role ARN"* +*"Authentication failed. Please verify the StackSet deployment and Role ARN"* or *"Authentication timed out"* -- Verify the management account role exists and was created in [Step 1](#step-1-create-the-management-account-role) -- Confirm the trust policy includes the correct External ID from the wizard -- Check the role has all Organizations discovery permissions listed in [Step 1](#step-1-create-the-management-account-role) -- Double-check the Role ARN format and account ID for typos +- Verify the deployment account role exists and is named exactly `ProwlerScan`. +- Confirm the trust policy includes the correct External ID from the wizard. +- Check the role has the Organizations discovery permissions listed in [Deploy the Roles Manually](#management-account-role). +- Double-check the Role ARN format and account ID for typos. +- Retry — the role can take up to **60 seconds** to propagate, and a second attempt often succeeds. For very large organizations (500+ accounts), allow extra time for discovery. -### Authentication Timed Out +### Invalid Organization ID or Role ARN -*"Authentication timed out"* +*"Must be a valid AWS Organization ID"* or *"Must be a valid IAM Role ARN"* -- Retry the authentication step — the second attempt often succeeds -- Check for AWS API rate limiting on the Organizations service -- For very large organizations (500+ accounts), allow extra time for discovery - -### Connection Test Fails for All Accounts - -No accounts pass the connection test. - -- Verify the CloudFormation StackSet was deployed — complete [Step 2](#step-2-deploy-the-cloudformation-stackset) and wait for stack instances to reach **CREATE_COMPLETE** -- Check that the **ExternalId** parameter in the StackSet matches the External ID shown in the Prowler wizard -- If your accounts use IP-based IAM policies, allow [Prowler Cloud egress IPs](/security/networking) - -### Connection Test Fails for Some Accounts - -Some accounts show a red icon while others pass. - -- Expand the StackSet deployment to include the OUs containing the failing accounts -- Suspended accounts cannot be scanned — deselect them and proceed -- Ensure the [STS regional endpoint](https://docs.aws.amazon.com/IAM/latest/UserGuide/id_credentials_temp_enable-regions.html) is enabled in the account's region -- After fixing, click **Test Connections** — only the failed accounts will be re-tested - -### No Accounts Connected Successfully - -*"No accounts connected successfully. Fix the connection errors and retry before launching scans."* - -- Hover over the red icon on each account to see the specific error -- Fix the underlying issues using the guidance above -- Click **Test Connections** to retry +- Organization ID format: `o-` followed by 10–32 lowercase alphanumeric characters (e.g., `o-abc123def4`). +- Role ARN format: `arn:aws:iam::<12-digit-account-id>:role/ProwlerScan`. +- Copy both directly from the AWS Console to avoid typos. ### Failed to Apply Discovery *"Failed to apply discovery"* -- Check the `blocked_reasons` field for any blocked accounts -- Retry the operation -- If the error persists, contact [Prowler Support](mailto:support@prowler.com) +- Check the `blocked_reasons` field for any blocked accounts and retry the operation. +- If the error persists, contact [Prowler Support](mailto:support@prowler.com). + +## Deploy the Roles Manually + +The wizard's **Create Stack** button is the fastest path, but you can create both roles yourself — for example with Terraform or your own CloudFormation — and paste the management account Role ARN into [Step 2](#step-2-authenticate-with-your-management-account). Both roles must be named `ProwlerScan`, since Prowler expects a consistent role name across all accounts. + + +**Prefer Terraform?** You can deploy the ProwlerScan role across the organization with Terraform instead of CloudFormation. See the [StackSets deployment guide](/user-guide/providers/aws/organizations#deploying-prowler-iam-roles-across-aws-organizations) for the module. + + +### Management Account Role + +The management account role lets Prowler discover your Organization structure — listing accounts, OUs, and hierarchy — and scan the management account itself. StackSets with service-managed permissions do not deploy to the management account, so this role is always created separately from the member-account StackSet. + +1. Sign in to the [AWS IAM Console](https://console.aws.amazon.com/iam/) in your **management account** (or delegated administrator account). +2. Go to **Roles > Create role** and select **Custom trust policy**. +3. Paste the following trust policy, replacing `` with the External ID shown in the Prowler wizard: + +```json +{ + "Version": "2012-10-17", + "Statement": [ + { + "Effect": "Allow", + "Principal": { + "AWS": "arn:aws:iam::232136659152:root" + }, + "Action": "sts:AssumeRole", + "Condition": { + "StringEquals": { + "sts:ExternalId": "" + }, + "StringLike": { + "aws:PrincipalArn": "arn:aws:iam::232136659152:role/prowler*" + } + } + } + ] +} +``` + +4. Attach the AWS managed policies **SecurityAudit** and **ViewOnlyAccess** so Prowler can scan the management account for security findings. +5. Add an inline policy with the Organizations discovery permissions: + +```json +{ + "Version": "2012-10-17", + "Statement": [ + { + "Sid": "ProwlerOrganizationDiscovery", + "Effect": "Allow", + "Action": [ + "organizations:DescribeAccount", + "organizations:DescribeOrganization", + "organizations:ListAccounts", + "organizations:ListAccountsForParent", + "organizations:ListOrganizationalUnitsForParent", + "organizations:ListRoots", + "organizations:ListTagsForResource" + ], + "Resource": "*" + }, + { + "Sid": "ProwlerStackSetManagement", + "Effect": "Allow", + "Action": [ + "organizations:RegisterDelegatedAdministrator", + "iam:CreateServiceLinkedRole" + ], + "Resource": "*" + } + ] +} +``` + + +You can restrict the `Resource` field to your specific Organization ARN (e.g., `arn:aws:organizations::123456789012:organization/o-abc123def4`) instead of `"*"` to minimize the blast radius. + + +6. Name the role **`ProwlerScan`** and click **Create role**. The ARN follows the format `arn:aws:iam:::role/ProwlerScan` — paste it into the wizard. + +### Member Account Role (StackSet) + +Deploy the ProwlerScan role to every member account with a [CloudFormation StackSet](https://docs.aws.amazon.com/AWSCloudFormation/latest/UserGuide/what-is-cfnstacksets.html), so you don't create the role manually in each account. + + +**Trusted access required.** CloudFormation StackSets must have trusted access enabled in your management account. Verify this under **AWS Organizations > Settings > Trusted access for AWS CloudFormation StackSets**. + + +1. In your management account, navigate to **CloudFormation > StackSets > Create StackSet** ([open directly](https://us-east-1.console.aws.amazon.com/cloudformation/home?region=us-east-1#/stacksets/create)). +2. Choose **Service-managed permissions** so AWS Organizations deploys the role automatically across current and future member accounts. +3. Select **Amazon S3 URL** as the template source and paste: + ``` + https://prowler-cloud-public.s3.eu-west-1.amazonaws.com/permissions/templates/aws/cloudformation/prowler-scan-role.yml + ``` +4. Set the **ExternalId** parameter to the External ID shown in the Prowler wizard. +5. Choose your deployment targets (entire organization or specific OUs) and regions, then click **Create StackSet**. +6. Open the **Stack instances** tab and confirm every instance shows **Status: CURRENT** and **Stack status: CREATE_COMPLETE**. Deployment typically takes **2–5 minutes**; large organizations (500+ accounts) may take longer. + +The StackSet role uses read-only access only (`SecurityAudit`, `ViewOnlyAccess`, plus a small set of additional read-only permissions). Prowler makes no changes to your accounts. See the [CloudFormation template](https://prowler-cloud-public.s3.eu-west-1.amazonaws.com/permissions/templates/aws/cloudformation/prowler-scan-role.yml) for the full list. When you add new accounts under the targeted OU or root, the StackSet deploys the role automatically, and Prowler's 6-hour sync onboards them end-to-end. + +## Key Concepts + +### What Is an External ID? + +An **External ID** is a security token that Prowler generates unique to your tenant. When Prowler assumes the IAM role in your AWS account, it presents this External ID to prove its identity. + +This prevents the [confused deputy problem](https://docs.aws.amazon.com/IAM/latest/UserGuide/confused-deputy.html) — a scenario where an unauthorized party could trick AWS into granting access to your account. By requiring the External ID, only your specific Prowler tenant can assume the role. Prowler generates it automatically and displays it in the wizard for you to copy. + +### Two Roles Architecture + +Prowler uses **two IAM roles**, both named `ProwlerScan` but deployed in different places: + +| Role | Where it lives | What it does | +|------|---------------|--------------| +| **ProwlerScan** (management account) | Your management (or delegated administrator) account | Discovers the Organization structure **and** scans that account. Includes additional Organizations discovery permissions. | +| **ProwlerScan** (member accounts) | Every member account | Scans the account for security findings. | + +Both roles share the name `ProwlerScan` because Prowler expects a consistent role name across all accounts. The single CloudFormation stack in [Step 2](#step-2-authenticate-with-your-management-account) deploys both at once. + + + Two Roles Architecture: ProwlerScan in management account (discovery + scanning) and ProwlerScan in member accounts (via StackSet, scanning only) + + +### What Is a CloudFormation StackSet? + +A [CloudFormation StackSet](https://docs.aws.amazon.com/AWSCloudFormation/latest/UserGuide/what-is-cfnstacksets.html) deploys the same CloudFormation template across multiple AWS accounts in a single operation. Prowler uses a service-managed StackSet to deploy the **ProwlerScan** IAM role into every member account of your organization, so you don't create the role manually in each account. StackSets do not deploy to the management account, which is why that role is created separately. ## What's Next diff --git a/mcp_server/.env.template b/mcp_server/.env.template index 11b8caa724..10aabd84cf 100644 --- a/mcp_server/.env.template +++ b/mcp_server/.env.template @@ -1,3 +1,3 @@ -PROWLER_APP_API_KEY="pk_your_api_key_here" +PROWLER_API_KEY="pk_your_api_key_here" API_BASE_URL="https://api.prowler.com/api/v1" PROWLER_MCP_TRANSPORT_MODE="stdio" diff --git a/mcp_server/AGENTS.md b/mcp_server/AGENTS.md index a82cc42e33..1d3d808e6a 100644 --- a/mcp_server/AGENTS.md +++ b/mcp_server/AGENTS.md @@ -25,7 +25,7 @@ The Prowler MCP Server provides AI agents access to the Prowler ecosystem throug ## CRITICAL RULES ### Tool Implementation -- ALWAYS: Extend `BaseTool` ABC for Prowler App tools (auto-registration) +- ALWAYS: Extend `BaseTool` ABC for Prowler tools (auto-registration) - ALWAYS: Use `@mcp.tool()` decorator for Hub/Docs tools - NEVER: Manually register BaseTool subclasses - NEVER: Import tools directly in server.py @@ -56,7 +56,7 @@ await prowler_mcp_server.import_server(docs_mcp_server, prefix="prowler_docs") ### Tool Naming - `prowler_hub_*` - Catalog and compliance (no auth) - `prowler_docs_*` - Documentation search (no auth) -- `prowler_app_*` - Cloud/App management (auth required) +- `prowler_*` - Prowler Cloud, Private Cloud & Local Server management (auth required) --- diff --git a/mcp_server/README.md b/mcp_server/README.md index e990f0f363..c0e4fdb0fe 100644 --- a/mcp_server/README.md +++ b/mcp_server/README.md @@ -6,9 +6,9 @@ ## Key Capabilities -### Prowler Cloud and Prowler App (Self-Managed) +### Prowler Cloud, Prowler Private Cloud & Prowler Local Server -Full access to Prowler Cloud platform and self-managed Prowler App for: +Full access to your Prowler data (Prowler Cloud, Prowler Private Cloud, or Prowler Local Server) for: - **Findings Analysis**: Query, filter, and analyze security findings across all your cloud environments - **Finding Groups Analysis**: Triage findings grouped by check ID and drill down into affected resources - **Provider Management**: Create, configure, and manage your configured Prowler providers (AWS, Azure, GCP, etc.) @@ -49,7 +49,7 @@ For comprehensive guides and tutorials, see the official documentation: Prowler MCP Server can be used in three ways: -### 1. Prowler Cloud MCP Server (Recommended) +### 1. Hosted Prowler MCP (Recommended) **Use Prowler's managed MCP server at `https://mcp.prowler.com/mcp`** @@ -126,7 +126,7 @@ For complete tool descriptions and parameters, see the [Tools Reference](https:/ ### Tool Naming Convention All tools follow a consistent naming pattern with prefixes: -- `prowler_app_*` - Prowler Cloud and App (Self-Managed) management tools +- `prowler_*` - Prowler Cloud, Prowler Private Cloud & Prowler Local Server management tools - `prowler_hub_*` - Prowler Hub catalog and compliance tools - `prowler_docs_*` - Prowler documentation search and retrieval @@ -146,7 +146,7 @@ prowler_mcp_server/ **Key Features:** - **Modular Design**: Three independent sub-servers with prefixed namespacing -- **Auto-Discovery**: Prowler App tools are automatically discovered and registered +- **Auto-Discovery**: Prowler tools are automatically discovered and registered - **LLM Optimization**: Response models minimize token usage by excluding empty values - **Dual Transport**: Supports both STDIO (local) and HTTP (remote) modes @@ -174,17 +174,17 @@ The Prowler MCP Server enables powerful workflows through AI assistants: ## Requirements -**For Prowler Cloud MCP Server:** -- Prowler Cloud account and API key (only for Prowler Cloud/App features) +**For the hosted Prowler MCP:** +- Prowler Cloud account and API key (only for Prowler features) **For self-hosted STDIO/HTTP Mode:** - Python 3.12+ or Docker - Network access to: - `https://hub.prowler.com` (for Prowler Hub) - `https://docs.prowler.com` (for Prowler Documentation) - - Prowler Cloud API or self-hosted Prowler App API (for Prowler Cloud/App features) + - Prowler Cloud API or Prowler Local Server API (for Prowler features) -> **No Authentication Required**: Prowler Hub and Prowler Documentation features work without authentication. A Prowler API key is only required to access Prowler Cloud or Prowler App (Self-Managed) features. +> **No Authentication Required**: Prowler Hub and Prowler Documentation features work without authentication. A Prowler API key is only required to access Prowler features (Prowler Cloud, Prowler Private Cloud, or Prowler Local Server). ## Configuring MCP Hosts @@ -200,7 +200,7 @@ For developers looking to extend the MCP server with new tools or features: ## Related Products - **[Prowler Hub](https://hub.prowler.com)**: Browse security checks and compliance frameworks -- **[Prowler Cloud](https://cloud.prowler.com)**: Managed Prowler platform +- **[Prowler Cloud](https://cloud.prowler.com)**: Fully managed Prowler in the cloud - **[Lighthouse AI](https://docs.prowler.com/getting-started/products/prowler-lighthouse-ai)**: AI security analyst ## License diff --git a/mcp_server/changelog.d/prowler-tools-namespace.changed.md b/mcp_server/changelog.d/prowler-tools-namespace.changed.md new file mode 100644 index 0000000000..0ca03d15ca --- /dev/null +++ b/mcp_server/changelog.d/prowler-tools-namespace.changed.md @@ -0,0 +1 @@ +Core Prowler tool namespace from the `prowler_app_*` prefix to `prowler_*` diff --git a/mcp_server/prowler_mcp_server/__init__.py b/mcp_server/prowler_mcp_server/__init__.py index fe7af2dcea..1956a6ec41 100644 --- a/mcp_server/prowler_mcp_server/__init__.py +++ b/mcp_server/prowler_mcp_server/__init__.py @@ -5,7 +5,7 @@ This package provides MCP tools for accessing: - Prowler Hub: All security artifacts (detections, remediations and frameworks) supported by Prowler """ -__version__ = "0.5.0" +__version__ = "0.8.0" __author__ = "Prowler Team" __email__ = "engineering@prowler.com" diff --git a/mcp_server/prowler_mcp_server/prowler_app/models/__init__.py b/mcp_server/prowler_mcp_server/prowler_app/models/__init__.py index 899ab4e866..1b15e35ac9 100644 --- a/mcp_server/prowler_mcp_server/prowler_app/models/__init__.py +++ b/mcp_server/prowler_mcp_server/prowler_app/models/__init__.py @@ -1,4 +1,4 @@ -"""Pydantic models for Prowler App MCP Server.""" +"""Pydantic models for Prowler MCP Server.""" from prowler_mcp_server.prowler_app.models.base import MinimalSerializerMixin from prowler_mcp_server.prowler_app.models.findings import ( diff --git a/mcp_server/prowler_mcp_server/prowler_app/models/finding_groups.py b/mcp_server/prowler_mcp_server/prowler_app/models/finding_groups.py index ae8431ba63..c2429012c3 100644 --- a/mcp_server/prowler_mcp_server/prowler_app/models/finding_groups.py +++ b/mcp_server/prowler_mcp_server/prowler_app/models/finding_groups.py @@ -228,7 +228,7 @@ class FindingGroupResource(MinimalSerializerMixin): resource: FindingGroupResourceInfo = Field(description="Affected resource") provider: FindingGroupProviderInfo = Field(description="Affected provider") finding_id: str = Field( - description="Finding UUID to use with prowler_app_get_finding_details" + description="Finding UUID to use with prowler_get_finding_details" ) status: FindingStatus = Field(description="Finding status for this resource") severity: FindingSeverity = Field(description="Finding severity") diff --git a/mcp_server/prowler_mcp_server/prowler_app/tools/__init__.py b/mcp_server/prowler_mcp_server/prowler_app/tools/__init__.py index 4d740b6efe..8b5be76076 100644 --- a/mcp_server/prowler_mcp_server/prowler_app/tools/__init__.py +++ b/mcp_server/prowler_mcp_server/prowler_app/tools/__init__.py @@ -1,4 +1,4 @@ -"""Domain-specific tools for Prowler App MCP Server. +"""Domain-specific tools for Prowler MCP Server. Each module in this package contains a BaseTool subclass that registers and implements tools for a specific domain (findings, providers, scans, etc.). diff --git a/mcp_server/prowler_mcp_server/prowler_app/tools/attack_paths.py b/mcp_server/prowler_mcp_server/prowler_app/tools/attack_paths.py index b08bbfe01f..5bd66760fa 100644 --- a/mcp_server/prowler_mcp_server/prowler_app/tools/attack_paths.py +++ b/mcp_server/prowler_mcp_server/prowler_app/tools/attack_paths.py @@ -1,4 +1,4 @@ -"""Attack Paths tools for Prowler App MCP Server. +"""Attack Paths tools for Prowler MCP Server. This module provides tools for analyzing Attack Paths data from Neo4j graph database. Attack Paths help identify security risks by tracing potential attack vectors @@ -22,16 +22,16 @@ class AttackPathsTools(BaseTool): """Tools for Attack Paths analysis. Provides tools for: - - prowler_app_list_attack_paths_scans: Find completed scans ready for analysis - - prowler_app_list_attack_paths_queries: Discover available queries for a scan - - prowler_app_run_attack_paths_query: Execute query and analyze attack paths + - prowler_list_attack_paths_scans: Find completed scans ready for analysis + - prowler_list_attack_paths_queries: Discover available queries for a scan + - prowler_run_attack_paths_query: Execute query and analyze attack paths """ async def list_attack_paths_scans( self, provider_id: list[str] = Field( default=[], - description="Filter by Prowler's internal UUID(s) (v4) for specific provider(s). Use `prowler_app_search_providers` tool to find provider IDs", + description="Filter by Prowler's internal UUID(s) (v4) for specific provider(s). Use `prowler_search_providers` tool to find provider IDs", ), provider_type: list[str] = Field( default=[], @@ -73,8 +73,8 @@ class AttackPathsTools(BaseTool): Workflow: 1. Use this tool to find completed attack paths scans - 2. Use prowler_app_list_attack_paths_queries to see available queries for a scan - 3. Use prowler_app_run_attack_paths_query to execute analysis + 2. Use prowler_list_attack_paths_queries to see available queries for a scan + 3. Use prowler_run_attack_paths_query to execute analysis """ try: # Validate pagination @@ -113,7 +113,7 @@ class AttackPathsTools(BaseTool): async def list_attack_paths_queries( self, scan_id: str = Field( - description="UUID of a COMPLETED attack paths scan. Use `prowler_app_list_attack_paths_scans` with state=['completed'] to find scan IDs" + description="UUID of a COMPLETED attack paths scan. Use `prowler_list_attack_paths_scans` with state=['completed'] to find scan IDs" ), ) -> list[dict[str, Any]]: """Discover available Attack Paths queries for a completed scan. @@ -133,9 +133,9 @@ class AttackPathsTools(BaseTool): - aws-ec2-instances-internet-exposed: Find internet-exposed EC2 instances Workflow: - 1. Use prowler_app_list_attack_paths_scans to find a completed scan + 1. Use prowler_list_attack_paths_scans to find a completed scan 2. Use this tool to discover available queries - 3. Use prowler_app_run_attack_paths_query with query_id and any required parameters + 3. Use prowler_run_attack_paths_query with query_id and any required parameters """ try: api_response = await self.api_client.get( @@ -158,7 +158,7 @@ class AttackPathsTools(BaseTool): description="UUID of a COMPLETED attack paths scan. The scan must be in 'completed' state" ), query_id: str = Field( - description="Query ID to execute (e.g., 'aws-internet-exposed-ec2-sensitive-s3-access'). Use `prowler_app_list_attack_paths_queries` to discover available queries" + description="Query ID to execute (e.g., 'aws-internet-exposed-ec2-sensitive-s3-access'). Use `prowler_list_attack_paths_queries` to discover available queries" ), parameters: dict[str, str] = Field( default_factory=dict, @@ -194,7 +194,7 @@ class AttackPathsTools(BaseTool): Workflow: 1. Ensure scan is completed - 2. List available queries (use prowler_app_list_attack_paths_queries) + 2. List available queries (use prowler_list_attack_paths_queries) 3. Execute this tool with appropriate parameters 4. Analyze the returned graph for security insights """ @@ -231,7 +231,7 @@ class AttackPathsTools(BaseTool): async def get_attack_paths_cartography_schema( self, scan_id: str = Field( - description="UUID of a COMPLETED attack paths scan. Use `prowler_app_list_attack_paths_scans` with state=['completed'] to find scan IDs" + description="UUID of a COMPLETED attack paths scan. Use `prowler_list_attack_paths_scans` with state=['completed'] to find scan IDs" ), ) -> dict[str, Any]: """Retrieve the Cartography graph schema for a completed attack paths scan. @@ -253,10 +253,10 @@ class AttackPathsTools(BaseTool): - schema_content: Full Cartography schema markdown with node/relationship definitions Workflow: - 1. Use prowler_app_list_attack_paths_scans to find a completed scan + 1. Use prowler_list_attack_paths_scans to find a completed scan 2. Use this tool to get the schema for the scan's provider 3. Use the schema to craft custom openCypher queries - 4. Execute queries with prowler_app_run_attack_paths_query + 4. Execute queries with prowler_run_attack_paths_query """ try: api_response = await self.api_client.get( diff --git a/mcp_server/prowler_mcp_server/prowler_app/tools/compliance.py b/mcp_server/prowler_mcp_server/prowler_app/tools/compliance.py index 360dd5510d..33cdd22a69 100644 --- a/mcp_server/prowler_mcp_server/prowler_app/tools/compliance.py +++ b/mcp_server/prowler_mcp_server/prowler_app/tools/compliance.py @@ -1,4 +1,4 @@ -"""Compliance framework tools for Prowler App MCP Server. +"""Compliance framework tools for Prowler MCP Server. This module provides tools for viewing compliance status and requirement details across all cloud providers. @@ -50,7 +50,7 @@ class ComplianceTools(BaseTool): if not scans_data: raise ValueError( f"No completed scans found for provider {provider_id}. " - "Run a scan first using prowler_app_trigger_scan." + "Run a scan first using prowler_trigger_scan." ) scan_id = scans_data[0]["id"] @@ -60,11 +60,11 @@ class ComplianceTools(BaseTool): self, scan_id: str | None = Field( default=None, - description="UUID of a specific scan to get compliance data for. Required if provider_id is not specified. Use `prowler_app_list_scans` to find scan IDs.", + description="UUID of a specific scan to get compliance data for. Required if provider_id is not specified. Use `prowler_list_scans` to find scan IDs.", ), provider_id: str | None = Field( default=None, - description="Prowler's internal UUID (v4) for a specific provider. If provided without scan_id, the tool will automatically find the latest completed scan for this provider. Use `prowler_app_search_providers` tool to find provider IDs.", + description="Prowler's internal UUID (v4) for a specific provider. If provided without scan_id, the tool will automatically find the latest completed scan for this provider. Use `prowler_search_providers` tool to find provider IDs.", ), ) -> dict[str, Any]: """Get high-level compliance overview across all frameworks for a specific scan. @@ -90,11 +90,11 @@ class ComplianceTools(BaseTool): Workflow: 1. Use this tool to get an overview of all compliance frameworks - 2. Use prowler_app_get_compliance_framework_state_details with a specific compliance_id to see which requirements failed + 2. Use prowler_get_compliance_framework_state_details with a specific compliance_id to see which requirements failed """ if not scan_id and not provider_id: return { - "error": "Either scan_id or provider_id must be provided. Use prowler_app_search_providers to find provider IDs or prowler_app_list_scans to find scan IDs." + "error": "Either scan_id or provider_id must be provided. Use prowler_search_providers to find provider IDs or prowler_list_scans to find scan IDs." } elif scan_id and provider_id: return { @@ -254,7 +254,7 @@ class ComplianceTools(BaseTool): async def get_compliance_framework_state_details( self, compliance_id: str = Field( - description="Compliance framework ID to get details for (e.g., 'cis_1.5_aws', 'pci_dss_v4.0_aws'). You can get compliance IDs from prowler_app_get_compliance_overview or consulting Prowler Hub/Prowler Documentation that you can also find in form of tools in this MCP Server", + description="Compliance framework ID to get details for (e.g., 'cis_1.5_aws', 'pci_dss_v4.0_aws'). You can get compliance IDs from prowler_get_compliance_overview or consulting Prowler Hub/Prowler Documentation that you can also find in form of tools in this MCP Server", ), scan_id: str | None = Field( default=None, @@ -262,14 +262,14 @@ class ComplianceTools(BaseTool): ), provider_id: str | None = Field( default=None, - description="Prowler's internal UUID (v4) for a specific provider. If provided without scan_id, the tool will automatically find the latest completed scan for this provider. Use `prowler_app_search_providers` tool to find provider IDs.", + description="Prowler's internal UUID (v4) for a specific provider. If provided without scan_id, the tool will automatically find the latest completed scan for this provider. Use `prowler_search_providers` tool to find provider IDs.", ), ) -> dict[str, Any]: """Get detailed requirement-level breakdown for a specific compliance framework. IMPORTANT: This tool returns DETAILED requirement information for a single compliance framework, focusing on FAILED requirements and their associated FAILED finding IDs. - Use this after prowler_app_get_compliance_overview to drill down into specific frameworks. + Use this after prowler_get_compliance_overview to drill down into specific frameworks. The markdown report includes: @@ -280,7 +280,7 @@ class ComplianceTools(BaseTool): 2. Failed Requirements Breakdown: - Each failed requirement's ID and description - Associated failed finding IDs for each failed requirement - - Use prowler_app_get_finding_details with these finding IDs for more details and remediation guidance + - Use prowler_get_finding_details with these finding IDs for more details and remediation guidance Default behavior: - Requires either scan_id OR provider_id @@ -289,14 +289,14 @@ class ComplianceTools(BaseTool): - Only shows failed requirements with their associated failed finding IDs Workflow: - 1. Use prowler_app_get_compliance_overview to identify frameworks with failures + 1. Use prowler_get_compliance_overview to identify frameworks with failures 2. Use this tool with the compliance_id to see failed requirements and their finding IDs - 3. Use prowler_app_get_finding_details with the finding IDs to get remediation guidance + 3. Use prowler_get_finding_details with the finding IDs to get remediation guidance """ # Validate that either scan_id or provider_id is provided if not scan_id and not provider_id: return { - "error": "Either scan_id or provider_id must be provided. Use prowler_app_search_providers to find provider IDs or prowler_app_list_scans to find scan IDs." + "error": "Either scan_id or provider_id must be provided. Use prowler_search_providers to find provider IDs or prowler_list_scans to find scan IDs." } # Resolve provider_id to latest scan_id if needed @@ -395,7 +395,7 @@ class ComplianceTools(BaseTool): report_lines.append("**Failed Finding IDs**: None found") report_lines.append("") report_lines.append( - "*Use `prowler_app_get_finding_details` with these finding IDs to get remediation guidance.*" + "*Use `prowler_get_finding_details` with these finding IDs to get remediation guidance.*" ) report_lines.append("") diff --git a/mcp_server/prowler_mcp_server/prowler_app/tools/finding_groups.py b/mcp_server/prowler_mcp_server/prowler_app/tools/finding_groups.py index 905a352740..05adf8db2b 100644 --- a/mcp_server/prowler_mcp_server/prowler_app/tools/finding_groups.py +++ b/mcp_server/prowler_mcp_server/prowler_app/tools/finding_groups.py @@ -1,4 +1,4 @@ -"""Finding Groups tools for Prowler App MCP Server. +"""Finding Groups tools for Prowler MCP Server. This module provides read-only tools for finding group triage and drill-downs. """ @@ -233,8 +233,8 @@ class FindingGroupsTools(BaseTool): `date_to`, this uses `/finding-groups` with a maximum 2-day date window. Use this tool to find noisy or high-impact checks, then call - prowler_app_get_finding_group_details for complete counters or - prowler_app_list_finding_group_resources to drill into affected resources. + prowler_get_finding_group_details for complete counters or + prowler_list_finding_group_resources to drill into affected resources. """ try: self.api_client.validate_page_size(page_size) @@ -423,7 +423,7 @@ class FindingGroupsTools(BaseTool): Default behavior returns FAIL, unmuted resources so the result is actionable. Set `include_muted=True` to include accepted/suppressed resources too. Each row includes nested resource and provider data plus - `finding_id`. Use `prowler_app_get_finding_details(finding_id)` to + `finding_id`. Use `prowler_get_finding_details(finding_id)` to retrieve complete remediation guidance for a specific resource finding. """ try: diff --git a/mcp_server/prowler_mcp_server/prowler_app/tools/findings.py b/mcp_server/prowler_mcp_server/prowler_app/tools/findings.py index ec492c6a43..b556101cab 100644 --- a/mcp_server/prowler_mcp_server/prowler_app/tools/findings.py +++ b/mcp_server/prowler_mcp_server/prowler_app/tools/findings.py @@ -1,4 +1,4 @@ -"""Security Findings tools for Prowler App MCP Server. +"""Security Findings tools for Prowler MCP Server. This module provides tools for searching, viewing, and analyzing security findings across all cloud providers. @@ -92,7 +92,7 @@ class FindingsTools(BaseTool): """Search and filter security findings across all cloud providers with rich filtering capabilities. IMPORTANT: This tool returns LIGHTWEIGHT findings. Use this for fast searching and filtering across many findings. - For complete details use prowler_app_get_finding_details on specific findings. + For complete details use prowler_get_finding_details on specific findings. Default behavior: - Returns latest findings from most recent scans (no date parameters needed) @@ -111,7 +111,7 @@ class FindingsTools(BaseTool): Workflow: 1. Use this tool to search and filter findings by severity, status, provider, service, region, etc. - 2. Use prowler_app_get_finding_details with the finding 'id' to get complete information about the finding + 2. Use prowler_get_finding_details with the finding 'id' to get complete information about the finding """ # Validate page_size parameter self.api_client.validate_page_size(page_size) @@ -187,9 +187,9 @@ class FindingsTools(BaseTool): """Retrieve comprehensive details about a specific security finding by its ID. IMPORTANT: This tool returns COMPLETE finding details. - Use this after finding a specific finding via prowler_app_search_security_findings + Use this after finding a specific finding via prowler_search_security_findings - This tool provides ALL information that prowler_app_search_security_findings returns PLUS: + This tool provides ALL information that prowler_search_security_findings returns PLUS: 1. Check Metadata (information about the check script that generated the finding): - title: Human-readable phrase used to summarize the check @@ -217,7 +217,7 @@ class FindingsTools(BaseTool): - resource_ids: List of UUIDs for cloud resources associated with this finding Workflow: - 1. Use prowler_app_search_security_findings to browse and filter findings + 1. Use prowler_search_security_findings to browse and filter findings 2. Use this tool with the finding 'id' to get remediation guidance and complete context """ params = { diff --git a/mcp_server/prowler_mcp_server/prowler_app/tools/muting.py b/mcp_server/prowler_mcp_server/prowler_app/tools/muting.py index 639f1ec3b7..37e1504165 100644 --- a/mcp_server/prowler_mcp_server/prowler_app/tools/muting.py +++ b/mcp_server/prowler_mcp_server/prowler_app/tools/muting.py @@ -1,4 +1,4 @@ -"""Muting tools for Prowler App MCP Server. +"""Muting tools for Prowler MCP Server. This module provides tools for managing finding muting in Prowler, including: - Mutelist management (pattern-based bulk muting) @@ -43,7 +43,7 @@ class MutingTools(BaseTool): Workflow: 1. Use this tool to check if a mutelist is configured 2. Examine current muting patterns before making updates - 3. Use prowler_app_set_mutelist to create or update the configuration + 3. Use prowler_set_mutelist to create or update the configuration """ self.logger.info("Retrieving mutelist configuration...") @@ -61,7 +61,7 @@ class MutingTools(BaseTool): if len(data) == 0: return { "error": "No mutelist found", - "message": "No mutelist configuration exists for this tenant. Use prowler_app_set_mutelist to create one.", + "message": "No mutelist configuration exists for this tenant. Use prowler_set_mutelist to create one.", } # Return the first (and only) mutelist @@ -116,10 +116,10 @@ Structure: - Exceptions: Accounts, Regions, Resources to exclude from muting Workflow: - 1. Use prowler_app_get_mutelist to check existing configuration + 1. Use prowler_get_mutelist to check existing configuration 2. Build configuration object following Prowler mutelist format 3. Use this tool to create or update the mutelist - 4. Verify with prowler_app_get_mutelist + 4. Verify with prowler_get_mutelist """ self.logger.info("Setting mutelist configuration...") @@ -171,12 +171,12 @@ Structure: """Remove the mutelist configuration from the tenant. WARNING: This is a destructive operation that cannot be undone. - - The mutelist will need to be re-created with prowler_app_set_mutelist + - The mutelist will need to be re-created with prowler_set_mutelist - New findings from future scans will NOT be muted by the deleted mutelist - Previously muted findings remain muted (deletion doesn't un-mute them) Workflow: - 1. Use prowler_app_get_mutelist to confirm what will be deleted + 1. Use prowler_get_mutelist to confirm what will be deleted 2. Use this tool to permanently remove the mutelist 3. New scans will no longer apply mutelist-based muting """ @@ -229,7 +229,7 @@ Structure: """Search and filter mute rules with pagination support. IMPORTANT: This tool returns LIGHTWEIGHT mute rules without the full list of finding UIDs. - Use prowler_app_get_mute_rule to get complete details including all finding UIDs and creator information. + Use prowler_get_mute_rule to get complete details including all finding UIDs and creator information. Default behavior: - Returns all mute rules (both enabled and disabled) @@ -237,15 +237,15 @@ Structure: - Includes basic rule information without full finding UID lists Each mute rule includes: - - Core identification: id (UUID for prowler_app_get_mute_rule), name + - Core identification: id (UUID for prowler_get_mute_rule), name - Contextual information: reason, enabled status - State tracking: finding_count (number of findings currently muted) - Temporal data: inserted_at, updated_at timestamps Workflow: 1. Use this tool to search and filter mute rules by name, enabled status, or keywords - 2. Use prowler_app_get_mute_rule with the mute rule 'id' to get complete details including all finding UIDs - 3. Use prowler_app_update_mute_rule or prowler_app_delete_mute_rule to modify rules + 2. Use prowler_get_mute_rule with the mute rule 'id' to get complete details including all finding UIDs + 3. Use prowler_update_mute_rule or prowler_delete_mute_rule to modify rules """ self.logger.info("Listing mute rules...") self.api_client.validate_page_size(page_size) @@ -289,17 +289,17 @@ Structure: """Retrieve comprehensive details about a specific mute rule by its ID. IMPORTANT: This tool returns COMPLETE mute rule details including the full list of finding UIDs. - Use this after finding a rule via prowler_app_list_mute_rules. + Use this after finding a rule via prowler_list_mute_rules. - This tool provides ALL information that prowler_app_list_mute_rules returns PLUS: + This tool provides ALL information that prowler_list_mute_rules returns PLUS: - finding_uids: Complete list of finding UIDs that are muted by this rule - user_creator_id: UUID of the user who created the rule (audit trail) Workflow: - 1. Use prowler_app_list_mute_rules to find rules by name or filter criteria + 1. Use prowler_list_mute_rules to find rules by name or filter criteria 2. Use this tool with the rule 'id' to get complete details 3. Examine finding_uids list to understand which findings are muted - 4. Use prowler_app_update_mute_rule or prowler_app_delete_mute_rule to modify if needed + 4. Use prowler_update_mute_rule or prowler_delete_mute_rule to modify if needed """ self.logger.info(f"Retrieving mute rule {rule_id}...") @@ -323,7 +323,7 @@ Structure: description="Reason for muting these findings. Document why this security issue is acceptable or intentional (e.g., 'Development environment with controlled access', 'Legacy application requires IMDSv1')." ), finding_ids: list[str] = Field( - description="List of finding IDs (UUIDs) to mute. Get these from the prowler_app_search_security_findings tool. Must provide at least 1 finding ID." + description="List of finding IDs (UUIDs) to mute. Get these from the prowler_search_security_findings tool. Must provide at least 1 finding ID." ), ) -> dict[str, Any]: """Create a new mute rule to mute specific findings with documentation and audit trail. @@ -337,15 +337,15 @@ Structure: - Records creator for audit trail The mute rule includes: - - Core identification: id (UUID for prowler_app_get_mute_rule), name, reason + - Core identification: id (UUID for prowler_get_mute_rule), name, reason - Configuration: enabled status, finding_uids list - Audit trail: user_creator_id (UUID of the Prowler user from the tenant that created the rule), timestamps when the rule was created and last modified Workflow: - 1. Use prowler_app_search_security_findings to identify findings to mute + 1. Use prowler_search_security_findings to identify findings to mute 2. Use this tool with finding IDs, descriptive name, and documented reason - 3. Verify with prowler_app_get_mute_rule to confirm rule creation - 4. Check findings are muted with prowler_app_search_security_findings (filter by muted=true) + 3. Verify with prowler_get_mute_rule to confirm rule creation + 4. Check findings are muted with prowler_search_security_findings (filter by muted=true) """ self.logger.info(f"Creating mute rule '{name}'...") @@ -399,9 +399,9 @@ Structure: - enabled: Toggle rule active status (doesn't affect already-muted findings) Workflow: - 1. Use prowler_app_get_mute_rule to see current rule state + 1. Use prowler_get_mute_rule to see current rule state 2. Use this tool to update name, reason, or enabled status - 3. Verify changes with prowler_app_get_mute_rule + 3. Verify changes with prowler_get_mute_rule """ self.logger.info(f"Updating mute rule {rule_id}...") @@ -451,9 +451,9 @@ Structure: - Cannot be undone - rule must be recreated to restore Workflow: - 1. Use prowler_app_get_mute_rule to review what will be deleted + 1. Use prowler_get_mute_rule to review what will be deleted 2. Use this tool to permanently remove the rule - 3. Verify deletion with prowler_app_list_mute_rules (rule should no longer appear) + 3. Verify deletion with prowler_list_mute_rules (rule should no longer appear) """ self.logger.info(f"Deleting mute rule {rule_id}...") diff --git a/mcp_server/prowler_mcp_server/prowler_app/tools/providers.py b/mcp_server/prowler_mcp_server/prowler_app/tools/providers.py index b22d57d7b9..3ba417d677 100644 --- a/mcp_server/prowler_mcp_server/prowler_app/tools/providers.py +++ b/mcp_server/prowler_mcp_server/prowler_app/tools/providers.py @@ -1,4 +1,4 @@ -"""Provider Management tools for Prowler App MCP Server. +"""Provider Management tools for Prowler MCP Server. This module provides tools for managing provider connections, including searching, connecting, and deleting providers. @@ -19,9 +19,9 @@ class ProvidersTools(BaseTool): """Tools for provider management operations Provides tools for: - - prowler_app_search_providers: Search and view configured providers with their connection status - - prowler_app_connect_provider: Connect or register a provider for security scanning in Prowler - - prowler_app_delete_provider: Permanently remove a provider from Prowler + - prowler_search_providers: Search and view configured providers with their connection status + - prowler_connect_provider: Connect or register a provider for security scanning in Prowler + - prowler_delete_provider: Permanently remove a provider from Prowler """ async def search_providers( @@ -145,7 +145,7 @@ class ProvidersTools(BaseTool): ) -> dict[str, Any]: """Register a provider to be scanned with Prowler. - This tool will register a provider in Prowler App, even if the UID is wrong. + This tool will register a provider in Prowler, even if the UID is wrong. If the provider is already registered, it will be updated with the new provided alias or credentials if provided. If credentials are provided, they will be added to the indicated provider, if the provider does not exist, it will be created and the credentials will be added to it. If the connection test is successful, the provider will be connected. @@ -292,13 +292,13 @@ class ProvidersTools(BaseTool): async def delete_provider( self, provider_id: str = Field( - description="Prowler's internal UUID (v4) for the provider to permanently remove, generated when the provider was registered in the system. Use `prowler_app_search_providers` tool to find the provider_id if you only know the alias or the provider's own identifier (provider_uid)" + description="Prowler's internal UUID (v4) for the provider to permanently remove, generated when the provider was registered in the system. Use `prowler_search_providers` tool to find the provider_id if you only know the alias or the provider's own identifier (provider_uid)" ), ) -> dict[str, Any]: """Permanently remove a registered provider from Prowler. WARNING: This is a destructive operation that cannot be undone. The provider will need to be - re-added with prowler_app_connect_provider if you want to scan it again. + re-added with prowler_connect_provider if you want to scan it again. The tool always returns the deletion status and message. """ diff --git a/mcp_server/prowler_mcp_server/prowler_app/tools/resources.py b/mcp_server/prowler_mcp_server/prowler_app/tools/resources.py index 011e013b91..88fcca25ae 100644 --- a/mcp_server/prowler_mcp_server/prowler_app/tools/resources.py +++ b/mcp_server/prowler_mcp_server/prowler_app/tools/resources.py @@ -1,4 +1,4 @@ -"""Cloud Resources tools for Prowler App MCP Server. +"""Cloud Resources tools for Prowler MCP Server. This module provides tools for searching, viewing, and analyzing cloud resources across all providers. @@ -86,7 +86,7 @@ class ResourcesTools(BaseTool): IMPORTANT: This tool returns LIGHTWEIGHT resource information. Use this for fast searching and filtering across many resources. For complete configuration details, metadata, and finding - relationships, use prowler_app_get_resource on specific resources of interest. + relationships, use prowler_get_resource on specific resources of interest. This is the primary tool for browsing resources with rich filtering capabilities. Returns current state by default (latest scan per provider). Specify dates to query @@ -102,16 +102,16 @@ class ResourcesTools(BaseTool): - With dates: queries historical resource state (2-day maximum range between date_from and date_to) Each resource includes: - - Core identification: id (UUID for prowler_app_get_resource), uid, name + - Core identification: id (UUID for prowler_get_resource), uid, name - Location context: region, service, type - Security context: failed_findings_count (number of active security issues) - Tags: tags associated with the resource Useful Workflow: 1. Use this tool to search and filter resources by provider, region, service, tags, etc. - 2. Use prowler_app_get_resource with the resource 'id' to get complete configuration and metadata - 3. Use prowler_app_search_security_findings to find security issues for specific resources - 4. Use prowler_app_get_finding_details to get details about the security issues for specific resources + 2. Use prowler_get_resource with the resource 'id' to get complete configuration and metadata + 3. Use prowler_search_security_findings to find security issues for specific resources + 4. Use prowler_get_finding_details to get details about the security issues for specific resources """ # Validate page_size parameter self.api_client.validate_page_size(page_size) @@ -177,15 +177,15 @@ class ResourcesTools(BaseTool): async def get_resource( self, resource_id: str = Field( - description="Prowler's internal UUID (v4) for the resource to retrieve, generated when the resource was discovered in the system. Use `prowler_app_list_resources` tool to find the right ID" + description="Prowler's internal UUID (v4) for the resource to retrieve, generated when the resource was discovered in the system. Use `prowler_list_resources` tool to find the right ID" ), ) -> dict[str, Any]: """Retrieve comprehensive details about a specific resource by its ID. IMPORTANT: This tool provides COMPLETE resource details with all available information. - Use this after finding a specific resource via prowler_app_list_resources. + Use this after finding a specific resource via prowler_list_resources. - This tool provides ALL information that prowler_app_list_resources returns PLUS: + This tool provides ALL information that prowler_list_resources returns PLUS: 1. Configuration Details: - metadata: Provider-specific configuration (tags, policies, encryption settings, network rules) @@ -197,12 +197,12 @@ class ResourcesTools(BaseTool): 3. Security Relationships: - finding_ids: Prowler's internal UUIDs (v4) of all security findings associated with this resource - - Use prowler_app_get_finding_details on these IDs to get remediation guidance + - Use prowler_get_finding_details on these IDs to get remediation guidance Useful Workflow: - 1. Use prowler_app_list_resources to browse and filter across many resources + 1. Use prowler_list_resources to browse and filter across many resources 2. Use this tool to drill down into specific resources of interest - 3. Use prowler_app_get_finding_details to get details about the security issues for specific resources + 3. Use prowler_get_finding_details to get details about the security issues for specific resources """ params = {} @@ -348,7 +348,7 @@ class ResourcesTools(BaseTool): async def get_resource_events( self, resource_id: str = Field( - description="Prowler's internal UUID (v4) for the resource. Use `prowler_app_list_resources` to find the right ID, or get it from a finding's resource relationship via `prowler_app_get_finding_details`." + description="Prowler's internal UUID (v4) for the resource. Use `prowler_list_resources` to find the right ID, or get it from a finding's resource relationship via `prowler_get_finding_details`." ), lookback_days: int = Field( default=90, @@ -386,8 +386,8 @@ class ResourcesTools(BaseTool): - Identifying unauthorized or unexpected modifications Workflows: - 1. Resource browsing: prowler_app_list_resources → find resource → this tool for event history - 2. Incident investigation: prowler_app_get_finding_details → get resource ID from finding → this tool to identify who caused the issue, what they changed, and when + 1. Resource browsing: prowler_list_resources → find resource → this tool for event history + 2. Incident investigation: prowler_get_finding_details → get resource ID from finding → this tool to identify who caused the issue, what they changed, and when """ params = { "lookback_days": lookback_days, diff --git a/mcp_server/prowler_mcp_server/prowler_app/tools/scans.py b/mcp_server/prowler_mcp_server/prowler_app/tools/scans.py index 1df636ffc0..21d1431b71 100644 --- a/mcp_server/prowler_mcp_server/prowler_app/tools/scans.py +++ b/mcp_server/prowler_mcp_server/prowler_app/tools/scans.py @@ -1,4 +1,4 @@ -"""Security Scans tools for Prowler App MCP Server. +"""Security Scans tools for Prowler MCP Server. This module provides tools for managing and monitoring Prowler security scans. """ @@ -20,18 +20,18 @@ class ScansTools(BaseTool): """Tools for security scan operations. Provides tools for: - - prowler_app_list_scans: Search and filter scans with rich filtering capabilities - - prowler_app_get_scan: Get comprehensive details about a specific scan - - prowler_app_trigger_scan: Trigger manual security scans for providers - - prowler_app_schedule_daily_scan: Schedule automated daily scans for continuous monitoring - - prowler_app_update_scan: Update scan names for better organization + - prowler_list_scans: Search and filter scans with rich filtering capabilities + - prowler_get_scan: Get comprehensive details about a specific scan + - prowler_trigger_scan: Trigger manual security scans for providers + - prowler_schedule_daily_scan: Schedule automated daily scans for continuous monitoring + - prowler_update_scan: Update scan names for better organization """ async def list_scans( self, provider_id: list[str] = Field( default=[], - description="Filter by Prowler's internal UUID(s) (v4) for specific provider(s), generated when the provider was registered. Use `prowler_app_search_providers` tool to find provider IDs", + description="Filter by Prowler's internal UUID(s) (v4) for specific provider(s), generated when the provider was registered. Use `prowler_search_providers` tool to find provider IDs", ), provider_type: list[str] = Field( default=[], @@ -56,7 +56,7 @@ class ScansTools(BaseTool): ), trigger: Literal["manual", "scheduled"] | None = Field( default=None, - description="Filter by how the scan was initiated. Options: 'manual' (user-initiated via prowler_app_trigger_scan), 'scheduled' (automated via prowler_app_schedule_daily_scan)", + description="Filter by how the scan was initiated. Options: 'manual' (user-initiated via prowler_trigger_scan), 'scheduled' (automated via prowler_schedule_daily_scan)", ), name: str | None = Field( default=None, @@ -75,7 +75,7 @@ class ScansTools(BaseTool): IMPORTANT: This tool returns LIGHTWEIGHT scan information. Use this for fast searching and filtering across many scans. For complete scan details including progress, duration, and resource counts, - use prowler_app_get_scan on specific scans of interest. + use prowler_get_scan on specific scans of interest. Default behavior: - Returns all scans @@ -83,15 +83,15 @@ class ScansTools(BaseTool): - Includes all scan states (available, scheduled, executing, completed, failed, cancelled) Each scan includes: - - Core identification: id (UUID for prowler_app_get_scan), name + - Core identification: id (UUID for prowler_get_scan), name - Execution context: state, trigger (manual/scheduled) - Temporal data: started_at, completed_at - Provider relationship: provider_id Workflow: 1. Use this tool to search and filter scans by provider, state, or date range - 2. Use prowler_app_get_scan with the scan 'id' to get progress, duration, and resource counts - 3. Use prowler_app_search_security_findings filtered by scan dates to analyze scan results + 2. Use prowler_get_scan with the scan 'id' to get progress, duration, and resource counts + 3. Use prowler_search_security_findings filtered by scan dates to analyze scan results """ # Validate pagination self.api_client.validate_page_size(page_size) @@ -128,15 +128,15 @@ class ScansTools(BaseTool): async def get_scan( self, scan_id: str = Field( - description="Prowler's internal UUID (v4) for the scan to retrieve, generated when the scan was created (e.g., '123e4567-e89b-12d3-a456-426614174000'). Use `prowler_app_list_scans` tool to find scan IDs" + description="Prowler's internal UUID (v4) for the scan to retrieve, generated when the scan was created (e.g., '123e4567-e89b-12d3-a456-426614174000'). Use `prowler_list_scans` tool to find scan IDs" ), ) -> dict[str, Any]: """Retrieve comprehensive details about a specific scan by its ID. IMPORTANT: This tool returns COMPLETE scan details. - Use this after finding a specific scan via prowler_app_list_scans. + Use this after finding a specific scan via prowler_list_scans. - This tool provides ALL information that prowler_app_list_scans returns PLUS: + This tool provides ALL information that prowler_list_scans returns PLUS: 1. Execution Details: - progress: Scan completion progress as percentage (0-100%) @@ -155,9 +155,9 @@ class ScansTools(BaseTool): - Understanding scan scheduling patterns Workflow: - 1. Use prowler_app_list_scans to browse and filter scans + 1. Use prowler_list_scans to browse and filter scans 2. Use this tool with the scan 'id' to monitor progress or view detailed results - 3. For completed scans, use prowler_app_search_security_findings filtered by date to analyze findings + 3. For completed scans, use prowler_search_security_findings filtered by date to analyze findings """ # Fetch scan with all fields params = { @@ -172,7 +172,7 @@ class ScansTools(BaseTool): async def trigger_scan( self, provider_id: str = Field( - description="Prowler's internal UUID (v4) for the provider to scan, generated when the provider was registered in the system (e.g., '4d0e2614-6385-4fa7-bf0b-c2e2f75c6877'). Use `prowler_app_search_providers` tool to find the provider ID" + description="Prowler's internal UUID (v4) for the provider to scan, generated when the provider was registered in the system (e.g., '4d0e2614-6385-4fa7-bf0b-c2e2f75c6877'). Use `prowler_search_providers` tool to find the provider ID" ), name: str | None = Field( default=None, @@ -182,14 +182,14 @@ class ScansTools(BaseTool): """Trigger a manual security scan for a provider. IMPORTANT: This tool returns immediately once the scan is created. - The scan will continue running in the background. Use `prowler_app_get_scan` + The scan will continue running in the background. Use `prowler_get_scan` with the returned scan ID to monitor progress and check when it completes. Example Useful Workflow: - 1. Use `prowler_app_search_providers` to find the provider_id you want to scan + 1. Use `prowler_search_providers` to find the provider_id you want to scan 2. Use this tool to trigger the scan - 3. Use `prowler_app_get_scan` with the returned scan 'id' to monitor progress - 4. Once completed, use `prowler_app_search_security_findings` to analyze results + 3. Use `prowler_get_scan` with the returned scan 'id' to monitor progress + 4. Once completed, use `prowler_search_security_findings` to analyze results """ try: # Build request data @@ -231,7 +231,7 @@ class ScansTools(BaseTool): return ScanCreationResult( scan=scan_info, status="success", - message=f"Scan {scan_id} created successfully. The scan may take some time to complete. Use prowler_app_get_scan tool with this ID to monitor progress.", + message=f"Scan {scan_id} created successfully. The scan may take some time to complete. Use prowler_get_scan tool with this ID to monitor progress.", ).model_dump() except Exception as e: @@ -245,7 +245,7 @@ class ScansTools(BaseTool): async def schedule_daily_scan( self, provider_id: str = Field( - description="Prowler's internal UUID (v4) for the provider to scan, generated when the provider was registered in the system (e.g., '4d0e2614-6385-4fa7-bf0b-c2e2f75c6877'). Use `prowler_app_search_providers` tool to find the provider ID" + description="Prowler's internal UUID (v4) for the provider to scan, generated when the provider was registered in the system (e.g., '4d0e2614-6385-4fa7-bf0b-c2e2f75c6877'). Use `prowler_search_providers` tool to find the provider ID" ), ) -> dict[str, Any]: """Schedule automated daily scans for a provider for continuous security monitoring. @@ -256,17 +256,17 @@ class ScansTools(BaseTool): you're not actively using the system. IMPORTANT: This tool returns immediately once the daily schedule is created. - The schedule will be set up in the background. Use `prowler_app_list_scans` + The schedule will be set up in the background. Use `prowler_list_scans` filtered by provider_id and trigger='scheduled' to view scheduled scans. IMPORTANT: This creates a PERSISTENT schedule. The provider will be scanned automatically every 24 hours until the provider is deleted. Example Useful Workflow: - 1. Use `prowler_app_search_providers` to find the provider_id you want to monitor + 1. Use `prowler_search_providers` to find the provider_id you want to monitor 2. Use this tool to create the daily schedule - 3. Use `prowler_app_list_scans` filtered by provider_id to view scheduled and completed scans - 4. Monitor findings over time with `prowler_app_search_security_findings` + 3. Use `prowler_list_scans` filtered by provider_id to view scheduled and completed scans + 4. Monitor findings over time with `prowler_search_security_findings` """ self.logger.info(f"Creating daily schedule for provider {provider_id}") task_response = await self.api_client.post( @@ -285,7 +285,7 @@ class ScansTools(BaseTool): ) if task_state == "available": - return_message = "Daily schedule created successfully. The schedule is being set up in the background. Use prowler_app_list_scans with provider_id filter to view scheduled scans." + return_message = "Daily schedule created successfully. The schedule is being set up in the background. Use prowler_list_scans with provider_id filter to view scheduled scans." else: return_message = "Daily schedule creation failed. Please try again later." @@ -297,7 +297,7 @@ class ScansTools(BaseTool): async def update_scan( self, scan_id: str = Field( - description="Prowler's internal UUID (v4) for the scan to update, generated when the scan was created (e.g., '123e4567-e89b-12d3-a456-426614174000'). Use `prowler_app_list_scans` tool to find the scan ID if you only know the provider or scan name. Returns an error if the scan ID is invalid or not found." + description="Prowler's internal UUID (v4) for the scan to update, generated when the scan was created (e.g., '123e4567-e89b-12d3-a456-426614174000'). Use `prowler_list_scans` tool to find the scan ID if you only know the provider or scan name. Returns an error if the scan ID is invalid or not found." ), name: str = Field( description="New human-friendly name for the scan (3-100 characters). Use descriptive names to improve organization and tracking, e.g., 'Production Security Audit - Q4 2025', 'Post-Deployment Compliance Check'. IMPORTANT: Only the scan name can be updated - other attributes (state, progress, duration) are read-only and managed by the system." @@ -309,7 +309,7 @@ class ScansTools(BaseTool): (state, progress, duration, etc.) are read-only and managed by the system. Example Useful Workflow: - 1. Use `prowler_app_list_scans` to find the scan you want to rename + 1. Use `prowler_list_scans` to find the scan you want to rename 2. Use this tool with the scan 'id' and new name """ api_response = await self.api_client.patch( diff --git a/mcp_server/prowler_mcp_server/prowler_app/utils/api_client.py b/mcp_server/prowler_mcp_server/prowler_app/utils/api_client.py index a6aacc3ce1..187364bee1 100644 --- a/mcp_server/prowler_mcp_server/prowler_app/utils/api_client.py +++ b/mcp_server/prowler_mcp_server/prowler_app/utils/api_client.py @@ -1,4 +1,4 @@ -"""Shared API client utilities for Prowler App tools.""" +"""Shared API client utilities for Prowler tools.""" import asyncio from datetime import datetime, timedelta diff --git a/mcp_server/prowler_mcp_server/prowler_app/utils/auth.py b/mcp_server/prowler_mcp_server/prowler_app/utils/auth.py index 72c06000df..eff5d3a117 100644 --- a/mcp_server/prowler_mcp_server/prowler_app/utils/auth.py +++ b/mcp_server/prowler_mcp_server/prowler_app/utils/auth.py @@ -10,7 +10,7 @@ from prowler_mcp_server.lib.logger import logger class ProwlerAppAuth: - """Handles authentication for Prowler App API using API keys or JWT tokens.""" + """Handles authentication for Prowler API using API keys or JWT tokens.""" def __init__( self, @@ -18,19 +18,23 @@ class ProwlerAppAuth: base_url: str = os.getenv("API_BASE_URL", "https://api.prowler.com/api/v1"), ): self.base_url = base_url.rstrip("/") - logger.info(f"Using Prowler App API base URL: {self.base_url}") + logger.info(f"Using Prowler API base URL: {self.base_url}") self.mode = mode self.access_token: str | None = None self.api_key: str | None = None if mode == "stdio": # STDIO mode - self.api_key = os.getenv("PROWLER_APP_API_KEY") + # PROWLER_API_KEY is the current variable; PROWLER_APP_API_KEY is kept + # as a backward-compatible fallback so existing setups keep working. + self.api_key = os.getenv("PROWLER_API_KEY") or os.getenv( + "PROWLER_APP_API_KEY" + ) if not self.api_key: - raise ValueError("PROWLER_APP_API_KEY environment variable is required") + raise ValueError("PROWLER_API_KEY environment variable is required") if not self.api_key.startswith("pk_"): - raise ValueError("Prowler App API key format is incorrect") + raise ValueError("Prowler API key format is incorrect") def _parse_jwt(self, token: str) -> dict | None: """Parse JWT token and return payload diff --git a/mcp_server/prowler_mcp_server/prowler_app/utils/tool_loader.py b/mcp_server/prowler_mcp_server/prowler_app/utils/tool_loader.py index b85c13af35..3a00474b5f 100644 --- a/mcp_server/prowler_mcp_server/prowler_app/utils/tool_loader.py +++ b/mcp_server/prowler_mcp_server/prowler_app/utils/tool_loader.py @@ -13,18 +13,27 @@ from prowler_mcp_server.lib.logger import logger from prowler_mcp_server.prowler_app.tools.base import BaseTool -def load_all_tools(mcp: FastMCP) -> None: - """Auto-discover and load all BaseTool subclasses from the tools package. +def load_all_tools( + mcp: FastMCP, + tools_package: str = "prowler_mcp_server.prowler_app.tools", +) -> None: + """Auto-discover and load all BaseTool subclasses from a tools package. This function: - 1. Dynamically imports all Python modules in the tools package - 2. Discovers all concrete BaseTool subclasses + 1. Dynamically imports all Python modules in the given tools package + 2. Discovers all concrete BaseTool subclasses defined in that package 3. Instantiates each tool class 4. Registers all tools with the provided FastMCP instance + ``BaseTool.__subclasses__()`` returns every subclass in the process, so the + discovered classes are filtered by ``__module__`` prefix. This keeps sibling + sub-servers (e.g. ``prowler_app`` and ``prowler_cloud``) from cross-registering + each other's tools, regardless of import order. + Args: mcp: The FastMCP instance to register tools with - TOOLS_PACKAGE: The package path containing tool modules (default: prowler_mcp_server.prowler_app.tools) + tools_package: The package path containing tool modules + (default: prowler_mcp_server.prowler_app.tools) Example: from fastmcp import FastMCP @@ -33,7 +42,7 @@ def load_all_tools(mcp: FastMCP) -> None: app = FastMCP("prowler-app") load_all_tools(app) """ - TOOLS_PACKAGE = "prowler_mcp_server.prowler_app.tools" + TOOLS_PACKAGE = tools_package logger.info(f"Auto-discovering tools from package: {TOOLS_PACKAGE}") # Import the tools package @@ -59,11 +68,14 @@ def load_all_tools(mcp: FastMCP) -> None: except Exception as e: logger.error(f"Failed to import module {module_name}: {e}") - # Discover all concrete BaseTool subclasses + # Discover all concrete BaseTool subclasses defined in this package only. + # __subclasses__() is process-wide, so filter by module to avoid sibling + # sub-servers cross-registering each other's tools. concrete_tools = [ tool_class for tool_class in BaseTool.__subclasses__() if not getattr(tool_class, "__abstractmethods__", None) + and tool_class.__module__.startswith(TOOLS_PACKAGE) ] logger.info(f"Discovered {len(concrete_tools)} tool classes") diff --git a/mcp_server/prowler_mcp_server/server.py b/mcp_server/prowler_mcp_server/server.py index 7c85641dee..a46ca12672 100644 --- a/mcp_server/prowler_mcp_server/server.py +++ b/mcp_server/prowler_mcp_server/server.py @@ -19,15 +19,15 @@ def setup_main_server(): except Exception as e: logger.error(f"Failed to mount Prowler Hub server: {e}") - # Mount Prowler App tools with prowler_app_ namespace + # Mount core Prowler tools with prowler_ namespace try: - logger.info("Mounting Prowler App server...") + logger.info("Mounting Prowler tools server...") from prowler_mcp_server.prowler_app.server import app_mcp_server - prowler_mcp_server.mount(app_mcp_server, namespace="prowler_app") - logger.info("Successfully mounted Prowler App server") + prowler_mcp_server.mount(app_mcp_server, namespace="prowler") + logger.info("Successfully mounted Prowler tools server") except Exception as e: - logger.error(f"Failed to mount Prowler App server: {e}") + logger.error(f"Failed to mount Prowler tools server: {e}") # Mount Prowler Documentation tools with prowler_docs_ namespace try: diff --git a/mcp_server/pyproject.toml b/mcp_server/pyproject.toml index 63aefdf931..f27e08b743 100644 --- a/mcp_server/pyproject.toml +++ b/mcp_server/pyproject.toml @@ -19,11 +19,13 @@ description = "MCP server for Prowler ecosystem" name = "prowler-mcp" readme = "README.md" requires-python = ">=3.12" -version = "0.5.0" +version = "0.8.0" [project.scripts] prowler-mcp = "prowler_mcp_server.main:main" +[tool.pytest] + [tool.pytest.ini_options] testpaths = ["tests"] diff --git a/mcp_server/uv.lock b/mcp_server/uv.lock index 3258767442..a9afbeafd7 100644 --- a/mcp_server/uv.lock +++ b/mcp_server/uv.lock @@ -676,7 +676,7 @@ wheels = [ [[package]] name = "prowler-mcp" -version = "0.5.0" +version = "0.8.0" source = { editable = "." } dependencies = [ { name = "fastmcp" }, diff --git a/prowler/changelog.d/sdk-scan-config-exclusions.added.md b/prowler/changelog.d/sdk-scan-config-exclusions.added.md new file mode 100644 index 0000000000..6955af0e70 --- /dev/null +++ b/prowler/changelog.d/sdk-scan-config-exclusions.added.md @@ -0,0 +1 @@ +`excluded_checks` and `excluded_services` in scan configurations to narrow the execution scope diff --git a/prowler/config/scan_config_schema.py b/prowler/config/scan_config_schema.py index ac00250c78..431c4cec60 100644 --- a/prowler/config/scan_config_schema.py +++ b/prowler/config/scan_config_schema.py @@ -9,21 +9,49 @@ The Prowler App, however, needs to surface those errors to the user when they save a Scan Config from the UI, and to expose the schema as JSON so the UI can validate live with `ajv`. This module provides: -- `validate_scan_config(payload)` — STRICT: returns a list of - `{path, message}` errors without silently dropping anything. The DRF - serializer (`api/.../v1/serializers.py:validate_scan_config_payload`) - turns each entry into a `ValidationError`. +- `validate_and_normalize_scan_config(payload)` — STRICT: returns + ``(normalized, errors)``. When ``errors`` is non-empty the normalized + dictionary is empty so callers never persist a partially validated + configuration. On success the normalized payload is JSON-serializable + (`model_dump(mode="json", exclude_unset=True)`), so the API can store + it directly in a Django ``JSONField`` and consume it at scan time + without re-running schema validation. + +- `validate_scan_config(payload)` — thin backward-compatible wrapper that + returns only the validation errors, preserved for callers that don't + need the normalized payload. - `SCAN_CONFIG_SCHEMA` — aggregated JSON Schema derived from the Pydantic models via `model_json_schema()`. Served by the `/scan-configs/schema` endpoint and consumed by the UI editor for in-editor live validation. """ +import json +from functools import lru_cache from typing import Any from pydantic import ValidationError from prowler.config.schema.registry import SCHEMAS +from prowler.lib.check.check import list_services +from prowler.lib.check.models import CheckMetadata + +# Pydantic v2 prefixes messages emitted from a ``field_validator`` that +# raises ``ValueError`` with this string. Strip it so the message that +# reaches the UI is the one the validator actually wrote. +_PYDANTIC_VALUE_ERROR_PREFIX = "Value error, " + + +@lru_cache(maxsize=None) +def _get_provider_check_ids(provider: str) -> frozenset[str]: + """Return cached check identifiers for a provider.""" + return frozenset(CheckMetadata.get_bulk(provider)) + + +@lru_cache(maxsize=None) +def _get_provider_services(provider: str) -> frozenset[str]: + """Return cached service identifiers for a provider.""" + return frozenset(list_services(provider)) def _format_loc(loc: tuple) -> str: @@ -50,48 +78,145 @@ def _format_loc(loc: tuple) -> str: return ".".join(parts) if parts else "" -def validate_scan_config(payload: Any) -> list[dict]: - """Validate a scan config payload against the registered provider schemas. +def validate_and_normalize_scan_config( + payload: Any, +) -> tuple[dict, list[dict[str, str]]]: + """Strict validation and normalization of a scan configuration payload. - Strict by design: every Pydantic violation surfaces as a `{path, message}` - entry so the caller can decide how to present it. Unknown provider - sections are accepted (consistent with `additionalProperties: True` at - the top level — the SDK simply has no opinion on them). + Returns ``(normalized, errors)``: + + - ``normalized`` is a JSON-serializable dict that mirrors the layout of + ``prowler/config/config.yaml`` (keyed by provider type). Registered + provider sections are dumped from their Pydantic models with + ``mode="json"`` (so the API can persist the result in a Django + ``JSONField``) and ``exclude_unset=True`` (so omitted defaults are + not injected into pre-existing configurations). Unknown provider + sections and unknown keys inside registered sections are preserved + untouched for forward compatibility with plugin-provided keys. + - ``errors`` is a list of ``{"path": , "message": }`` + entries, one per schema or exclusion-catalog violation. When any error + is present the normalized dictionary is returned empty so the caller + never persists a partially validated configuration. + + The input payload is never mutated. """ if not isinstance(payload, dict): - return [ + return {}, [ { "path": "", "message": "Scan config must be a mapping with provider sections.", } ] - errors: list[dict] = [] + errors: list[dict[str, str]] = [] + normalized: dict[str, Any] = {} + for provider, section in payload.items(): - schema_cls = SCHEMAS.get(provider) + # Reject non-string provider keys so distinct entries like ``123`` + # and ``"123"`` don't collide after ``str()`` in the normalized dict. + # YAML always produces string keys at this level; anything else + # comes from a hand-built payload and is a caller bug. + if not isinstance(provider, str): + errors.append( + { + "path": repr(provider), + "message": "provider keys must be strings.", + } + ) + continue + + provider_key = provider + schema_cls = SCHEMAS.get(provider_key) if schema_cls is None: - # Unknown provider type: tolerated. The SDK will simply ignore it. + # Unknown provider type: tolerated, but only when its contents + # are already JSON-serializable. The API persists the returned + # payload in a Django ``JSONField`` and would blow up at write + # time if we let a ``set()`` or similar through here. + try: + json.dumps(section) + except (TypeError, ValueError) as exc: + errors.append( + { + "path": provider_key, + "message": ( + "unknown provider section is not JSON-serializable: " + f"{exc}" + ), + } + ) + continue + normalized[provider_key] = section continue if not isinstance(section, dict): errors.append( { - "path": str(provider), + "path": provider_key, "message": "section must be a mapping.", } ) continue try: - schema_cls.model_validate(section) + model = schema_cls.model_validate(section) except ValidationError as exc: for err in exc.errors(): loc = err.get("loc") or () - path = _format_loc((str(provider), *loc)) - errors.append( - { - "path": path, - "message": err.get("msg", "validation error"), - } - ) + path = _format_loc((provider_key, *loc)) + message = err.get("msg", "validation error") + # Only strip on the specific error type that pydantic + # prefixes — a legitimate future message that happens to + # start with "Value error, " keeps its text intact. + if err.get("type") == "value_error" and message.startswith( + _PYDANTIC_VALUE_ERROR_PREFIX + ): + message = message[len(_PYDANTIC_VALUE_ERROR_PREFIX) :] + errors.append({"path": path, "message": message}) + continue + + if model.excluded_checks: + available_checks = _get_provider_check_ids(provider_key) + for index, check in enumerate(model.excluded_checks): + if check not in available_checks: + errors.append( + { + "path": f"{provider_key}.excluded_checks[{index}]", + "message": ( + f"Unknown check '{check}' for provider " + f"'{provider_key}'." + ), + } + ) + + if model.excluded_services: + available_services = _get_provider_services(provider_key) + for index, service in enumerate(model.excluded_services): + if service not in available_services: + errors.append( + { + "path": f"{provider_key}.excluded_services[{index}]", + "message": ( + f"Unknown service '{service}' for provider " + f"'{provider_key}'." + ), + } + ) + + normalized[provider_key] = model.model_dump(mode="json", exclude_unset=True) + + if errors: + return {}, errors + return normalized, [] + + +def validate_scan_config(payload: Any) -> list[dict]: + """Backward-compatible wrapper returning only validation errors. + + Preserved for callers that only need the strict-validation error list + (e.g. the DRF serializer that turns each entry into a + ``ValidationError``). New callers should prefer + :func:`validate_and_normalize_scan_config` to also receive the + normalized payload. + """ + _, errors = validate_and_normalize_scan_config(payload) return errors diff --git a/prowler/config/schema/base.py b/prowler/config/schema/base.py index cc473a4545..fc5a76af43 100644 --- a/prowler/config/schema/base.py +++ b/prowler/config/schema/base.py @@ -1,4 +1,12 @@ -from pydantic import BaseModel, ConfigDict +from typing import Annotated + +from pydantic import BaseModel, ConfigDict, Field, StringConstraints, field_validator + +# Item type for excluded_checks / excluded_services list entries. Item +# whitespace is stripped via ``str_strip_whitespace`` on the base +# ``model_config`` (no second stripping implementation added here), so +# ``min_length=1`` catches "", " ", and any all-whitespace input uniformly. +NonEmptyScopeIdentifier = Annotated[str, StringConstraints(min_length=1)] class ProviderConfigBase(BaseModel): @@ -15,3 +23,28 @@ class ProviderConfigBase(BaseModel): str_strip_whitespace=True, validate_assignment=False, ) + + excluded_checks: list[NonEmptyScopeIdentifier] = Field( + default_factory=list, + description="Check identifiers to exclude from the scan scope.", + json_schema_extra={"default": [], "uniqueItems": True}, + ) + excluded_services: list[NonEmptyScopeIdentifier] = Field( + default_factory=list, + description="Service identifiers to exclude from the scan scope.", + json_schema_extra={"default": [], "uniqueItems": True}, + ) + + @field_validator("excluded_checks", "excluded_services") + @classmethod + def _reject_duplicates(cls, value: list[str]) -> list[str]: + seen: set[str] = set() + duplicates: set[str] = set() + for item in value: + if item in seen: + duplicates.add(item) + else: + seen.add(item) + if duplicates: + raise ValueError(f"duplicate values are not allowed: {sorted(duplicates)}") + return value diff --git a/prowler/lib/scan/scan.py b/prowler/lib/scan/scan.py index 4bef660d33..b87cfcdf1d 100644 --- a/prowler/lib/scan/scan.py +++ b/prowler/lib/scan/scan.py @@ -178,25 +178,58 @@ class Scan: ) ) - # Exclude checks + # Validate excluded checks against the FULL provider catalog — not + # just the selected scope — so a global config can exclude a valid + # check even when that check is not part of a particular scoped run. + excluded_check_set: set[str] = set() if excluded_checks: - for check in excluded_checks: - if check in self._checks_to_execute: - self._checks_to_execute.remove(check) - else: - raise ScanInvalidCheckError( - f"Invalid check provided: {check}. Check does not exist in the provider." - ) + excluded_check_set = set(excluded_checks) + if len(excluded_check_set) != len(excluded_checks): + raise ScanInvalidCheckError( + "Duplicate excluded checks are not allowed." + ) + unknown_checks = excluded_check_set.difference(self._bulk_checks_metadata) + if unknown_checks: + raise ScanInvalidCheckError( + f"Invalid excluded check(s) provided: {sorted(unknown_checks)}." + ) - # Exclude services + # Validate excluded services against the provider service catalog. + # Only resolve the catalog when there is something to check to avoid + # walking the provider package tree unnecessarily. + excluded_service_set: set[str] = set() if excluded_services: - for check in self._checks_to_execute: - if get_service_name_from_check_name(check) in excluded_services: - self._checks_to_execute.remove(check) - else: - raise ScanInvalidServiceError( - f"Invalid service provided: {check}. Service does not exist in the provider." - ) + excluded_service_set = set(excluded_services) + if len(excluded_service_set) != len(excluded_services): + raise ScanInvalidServiceError( + "Duplicate excluded services are not allowed." + ) + unknown_services = excluded_service_set.difference( + list_services(provider.type) + ) + if unknown_services: + raise ScanInvalidServiceError( + f"Invalid excluded service(s) provided: {sorted(unknown_services)}." + ) + + if excluded_check_set or excluded_service_set: + previous_scope = self._checks_to_execute + selected_checks = { + check + for check in previous_scope + if check not in excluded_check_set + and get_service_name_from_check_name(check) not in excluded_service_set + } + # Only complain when exclusions actually emptied a non-empty + # scope. If the scope was already empty (e.g. a severity or + # category filter matched nothing) the exclusions did not + # cause the emptiness and the misleading error would obscure + # the real reason. + if previous_scope and not selected_checks: + raise ScanInvalidCheckError( + "The scan configuration excludes every selected check." + ) + self._checks_to_execute = sorted(selected_checks) self._number_of_checks_to_execute = len(self._checks_to_execute) diff --git a/skills/prowler-mcp/SKILL.md b/skills/prowler-mcp/SKILL.md index af3c597771..704acf8553 100644 --- a/skills/prowler-mcp/SKILL.md +++ b/skills/prowler-mcp/SKILL.md @@ -19,7 +19,7 @@ The Prowler MCP Server uses three sub-servers with prefixed namespacing: | Sub-Server | Prefix | Auth | Purpose | |------------|--------|------|---------| -| Prowler App | `prowler_app_*` | Required | Cloud management tools | +| Prowler | `prowler_*` | Required | Prowler Cloud, Private Cloud & Local Server management tools | | Prowler Hub | `prowler_hub_*` | No | Security checks catalog | | Prowler Docs | `prowler_docs_*` | No | Documentation search | @@ -27,7 +27,7 @@ For complete architecture, patterns, and examples, see [docs/developer-guide/mcp --- -## Critical Rules (Prowler App Only) +## Critical Rules (Prowler Tools Only) ### Tool Implementation @@ -56,7 +56,7 @@ Use `@mcp.tool()` decorator directly—no BaseTool or models required. --- -## Quick Reference: New Prowler App Tool +## Quick Reference: New Prowler Tool 1. Create tool class in `prowler_app/tools/` extending `BaseTool` 2. Create models in `prowler_app/models/` using `MinimalSerializerMixin` @@ -64,7 +64,7 @@ Use `@mcp.tool()` decorator directly—no BaseTool or models required. --- -## QA Checklist (Prowler App) +## QA Checklist (Prowler Tools) - [ ] Tool docstrings describe LLM-relevant behavior - [ ] Models use `MinimalSerializerMixin` diff --git a/tests/config/schema/exclusions_test.py b/tests/config/schema/exclusions_test.py new file mode 100644 index 0000000000..a13ed2f786 --- /dev/null +++ b/tests/config/schema/exclusions_test.py @@ -0,0 +1,113 @@ +"""Coverage for the ``excluded_checks`` / ``excluded_services`` fields +added to :class:`prowler.config.schema.base.ProviderConfigBase`. + +Because the fields live on the base class, every registered provider +schema exposes them and every provider must therefore share the same +whitespace / uniqueness / non-empty guarantees. These tests lock in that +contract at the base level and at the JSON-Schema level (which the UI +editor consumes via ``ajv``). +""" + +import pytest +from pydantic import ValidationError + +from prowler.config.scan_config_schema import SCAN_CONFIG_SCHEMA +from prowler.config.schema.aws import AWSProviderConfig +from prowler.config.schema.registry import SCHEMAS +from prowler.config.schema.validator import validate_provider_config + +EXCLUSION_FIELDS = ("excluded_checks", "excluded_services") + + +class Test_JSON_Schema_Exposes_Exclusion_Fields: + @pytest.mark.parametrize("provider", sorted(SCHEMAS)) + @pytest.mark.parametrize("field", EXCLUSION_FIELDS) + def test_field_shape(self, provider, field): + field_schema = SCAN_CONFIG_SCHEMA["properties"][provider]["properties"][field] + assert field_schema["type"] == "array" + assert field_schema["items"] == {"type": "string", "minLength": 1} + assert field_schema["uniqueItems"] is True + assert field_schema["default"] == [] + + +class Test_Exclusion_Field_Validation: + def _model(self, **kwargs): + return AWSProviderConfig.model_validate(kwargs) + + @pytest.mark.parametrize("field", EXCLUSION_FIELDS) + def test_empty_string_is_rejected(self, field): + with pytest.raises(ValidationError): + self._model(**{field: [""]}) + + @pytest.mark.parametrize("field", EXCLUSION_FIELDS) + def test_whitespace_only_string_is_rejected(self, field): + with pytest.raises(ValidationError): + self._model(**{field: [" "]}) + + @pytest.mark.parametrize("field", EXCLUSION_FIELDS) + def test_raw_duplicates_are_rejected(self, field): + with pytest.raises(ValidationError): + self._model(**{field: ["s3", "s3"]}) + + @pytest.mark.parametrize("field", EXCLUSION_FIELDS) + def test_normalized_duplicates_are_rejected(self, field): + # After whitespace normalization ``" s3 "`` collapses to ``"s3"`` + # and must be caught by the duplicate check. + with pytest.raises(ValidationError): + self._model(**{field: ["s3", " s3 "]}) + + @pytest.mark.parametrize("field", EXCLUSION_FIELDS) + def test_whitespace_is_stripped(self, field): + model = self._model(**{field: [" identifier "]}) + assert getattr(model, field) == ["identifier"] + + @pytest.mark.parametrize("field", EXCLUSION_FIELDS) + def test_non_string_item_is_rejected(self, field): + with pytest.raises(ValidationError): + self._model(**{field: [123]}) + + +class Test_Exclusion_Defaults_Are_Not_Injected: + """The strict normalization path uses ``model_dump(exclude_unset=True)`` + so pre-existing configs that never set an exclusion field must round-trip + without the default empty list being materialized.""" + + def test_absent_fields_are_not_injected_by_validator(self): + # The lenient SDK runtime path also uses ``exclude_unset=True`` under + # the hood; asserting on the validator output guards the promise + # against future refactors. + assert validate_provider_config("aws", {}, SCHEMAS["aws"]) == {} + + def test_absent_fields_are_not_injected_when_other_keys_are_present(self): + assert validate_provider_config( + "aws", + {"max_ec2_instance_age_in_days": 180}, + SCHEMAS["aws"], + ) == {"max_ec2_instance_age_in_days": 180} + + def test_explicit_empty_list_round_trips(self): + # Explicitly setting ``excluded_checks: []`` is different from + # omitting it — the empty list is user-provided and must be + # preserved by the strict-normalization contract. + assert validate_provider_config( + "aws", + {"excluded_checks": []}, + SCHEMAS["aws"], + ) == {"excluded_checks": []} + + +class Test_Extra_Fields_Are_Preserved: + """``extra="allow"`` must keep plugin-provided keys around so the + ecosystem contract in ``validator_test.py`` still holds after adding + the exclusion fields.""" + + def test_unknown_keys_are_preserved_alongside_exclusions(self): + out = validate_provider_config( + "aws", + {"excluded_checks": ["s3_bucket_public_access"], "plugin_option": "kept"}, + SCHEMAS["aws"], + ) + assert out == { + "excluded_checks": ["s3_bucket_public_access"], + "plugin_option": "kept", + } diff --git a/tests/config/schema/scan_config_schema_test.py b/tests/config/schema/scan_config_schema_test.py new file mode 100644 index 0000000000..037fcd9501 --- /dev/null +++ b/tests/config/schema/scan_config_schema_test.py @@ -0,0 +1,393 @@ +"""Coverage for the strict scan-config validation and normalization +contract exposed to the Prowler App backend. + +Split from :mod:`tests.config.schema.validator_test` because the strict +API (``validate_and_normalize_scan_config``) has different guarantees: +it never silently drops keys, and it returns a JSON-serializable payload +the backend can persist verbatim in a Django ``JSONField``. +""" + +import json +from unittest.mock import call, patch + +import pytest + +from prowler.config.scan_config_schema import ( + _get_provider_check_ids, + _get_provider_services, + validate_and_normalize_scan_config, + validate_scan_config, +) + + +@pytest.fixture(autouse=True) +def clear_provider_catalog_caches(): + """Keep provider catalog cache state isolated between tests.""" + _get_provider_check_ids.cache_clear() + _get_provider_services.cache_clear() + yield + _get_provider_check_ids.cache_clear() + _get_provider_services.cache_clear() + + +class Test_Non_Dict_Root: + @pytest.mark.parametrize("payload", [None, "string", 42, [], (1, 2)]) + def test_non_mapping_root_is_rejected(self, payload): + normalized, errors = validate_and_normalize_scan_config(payload) + assert normalized == {} + assert len(errors) == 1 + assert errors[0]["path"] == "" + + +class Test_Registered_Provider_Section_Must_Be_Mapping: + @pytest.mark.parametrize("section", ["a-string", 42, ["s3"], None]) + def test_non_mapping_section_reports_provider_path(self, section): + normalized, errors = validate_and_normalize_scan_config({"aws": section}) + assert normalized == {} + assert errors == [{"path": "aws", "message": "section must be a mapping."}] + + +class Test_Success_Path: + def test_whitespace_is_normalized_in_exclusions(self): + normalized, errors = validate_and_normalize_scan_config( + { + "aws": { + "excluded_checks": [" s3_bucket_default_encryption "], + "excluded_services": [" s3 "], + } + } + ) + assert errors == [] + assert normalized == { + "aws": { + "excluded_checks": ["s3_bucket_default_encryption"], + "excluded_services": ["s3"], + } + } + + def test_plugin_options_are_preserved(self): + # Third-party plugins inject arbitrary keys inside a provider + # section; ``extra="allow"`` on the schema keeps them alive + # through the dump/normalize round-trip. + normalized, errors = validate_and_normalize_scan_config( + {"aws": {"plugin_option": "preserved", "another": 42}} + ) + assert errors == [] + assert normalized == {"aws": {"plugin_option": "preserved", "another": 42}} + + def test_plugin_catalog_identifiers_are_accepted_and_catalogs_are_cached(self): + payload = { + "aws": { + "excluded_checks": ["plugin_check"], + "excluded_services": ["plugin_service"], + } + } + with ( + patch( + "prowler.config.scan_config_schema.CheckMetadata.get_bulk", + return_value={"plugin_check": object()}, + ) as check_catalog, + patch( + "prowler.config.scan_config_schema.list_services", + return_value=["plugin_service"], + ) as service_catalog, + ): + first_result = validate_and_normalize_scan_config(payload) + second_result = validate_and_normalize_scan_config(payload) + + normalized, errors = first_result + assert errors == [] + assert normalized == { + "aws": { + "excluded_checks": ["plugin_check"], + "excluded_services": ["plugin_service"], + } + } + assert second_result == first_result + check_catalog.assert_called_once_with("aws") + service_catalog.assert_called_once_with("aws") + + def test_catalog_caches_are_keyed_by_provider(self): + with ( + patch( + "prowler.config.scan_config_schema.CheckMetadata.get_bulk", + side_effect=lambda provider: {f"{provider}_plugin_check": object()}, + ) as check_catalog, + patch( + "prowler.config.scan_config_schema.list_services", + side_effect=lambda provider: [f"{provider}_plugin_service"], + ) as service_catalog, + ): + payload = { + "aws": { + "excluded_checks": ["aws_plugin_check"], + "excluded_services": ["aws_plugin_service"], + }, + "azure": { + "excluded_checks": ["azure_plugin_check"], + "excluded_services": ["azure_plugin_service"], + }, + } + first_result = validate_and_normalize_scan_config(payload) + second_result = validate_and_normalize_scan_config(payload) + + assert first_result[1] == [] + assert second_result == first_result + assert check_catalog.call_args_list == [call("aws"), call("azure")] + assert service_catalog.call_args_list == [call("aws"), call("azure")] + + def test_omitted_defaults_are_not_injected(self): + normalized, errors = validate_and_normalize_scan_config( + {"aws": {"max_ec2_instance_age_in_days": 90}} + ) + assert errors == [] + assert normalized == {"aws": {"max_ec2_instance_age_in_days": 90}} + assert "excluded_checks" not in normalized["aws"] + assert "excluded_services" not in normalized["aws"] + + def test_unknown_provider_sections_are_preserved_verbatim(self): + payload = {"future_provider": {"custom_option": True, "nested": {"k": 1}}} + normalized, errors = validate_and_normalize_scan_config(payload) + assert errors == [] + assert normalized == payload + + def test_normalized_payload_is_json_serializable(self): + normalized, _ = validate_and_normalize_scan_config( + { + "aws": { + "excluded_checks": ["s3_bucket_public_access"], + "excluded_services": ["s3"], + } + } + ) + # If ``model_dump(mode="json", ...)`` is ever dropped this + # ``json.dumps`` call is what will notice. + json.dumps(normalized) + + def test_input_payload_is_not_mutated(self): + payload = { + "aws": { + "excluded_checks": [" s3_bucket_public_access "], + "excluded_services": [" s3 "], + } + } + snapshot = json.loads(json.dumps(payload)) + validate_and_normalize_scan_config(payload) + assert payload == snapshot + + +class Test_Error_Path: + def test_unknown_excluded_check_is_rejected(self): + normalized, errors = validate_and_normalize_scan_config( + {"aws": {"excluded_checks": ["aws_check_that_does_not_exist"]}} + ) + assert normalized == {} + assert errors == [ + { + "path": "aws.excluded_checks[0]", + "message": ( + "Unknown check 'aws_check_that_does_not_exist' for provider " + "'aws'." + ), + } + ] + + def test_unknown_excluded_service_is_rejected(self): + normalized, errors = validate_and_normalize_scan_config( + {"aws": {"excluded_services": ["not_a_real_aws_service"]}} + ) + assert normalized == {} + assert errors == [ + { + "path": "aws.excluded_services[0]", + "message": ( + "Unknown service 'not_a_real_aws_service' for provider 'aws'." + ), + } + ] + + def test_multiple_unknown_exclusions_return_deterministic_errors(self): + normalized, errors = validate_and_normalize_scan_config( + { + "aws": { + "excluded_checks": [ + "unknown_check_one", + "s3_bucket_default_encryption", + "unknown_check_two", + ], + "excluded_services": [ + "unknown_service_one", + "s3", + "unknown_service_two", + ], + } + } + ) + assert normalized == {} + assert errors == [ + { + "path": "aws.excluded_checks[0]", + "message": "Unknown check 'unknown_check_one' for provider 'aws'.", + }, + { + "path": "aws.excluded_checks[2]", + "message": "Unknown check 'unknown_check_two' for provider 'aws'.", + }, + { + "path": "aws.excluded_services[0]", + "message": ( + "Unknown service 'unknown_service_one' for provider 'aws'." + ), + }, + { + "path": "aws.excluded_services[2]", + "message": ( + "Unknown service 'unknown_service_two' for provider 'aws'." + ), + }, + ] + + def test_check_from_another_provider_is_rejected(self): + azure_check = "postgresql_flexible_server_allow_access_services_disabled" + normalized, errors = validate_and_normalize_scan_config( + {"aws": {"excluded_checks": [azure_check]}} + ) + assert normalized == {} + assert errors == [ + { + "path": "aws.excluded_checks[0]", + "message": f"Unknown check '{azure_check}' for provider 'aws'.", + } + ] + + def test_invalid_input_returns_empty_normalized_and_errors(self): + normalized, errors = validate_and_normalize_scan_config( + {"aws": {"excluded_services": ["s3", " s3 "]}} + ) + assert normalized == {} + assert errors + assert any(err["path"].startswith("aws.excluded_services") for err in errors) + + def test_partial_error_zeros_the_normalized_payload(self): + # One valid provider + one invalid provider must not leak the + # valid section into a partially normalized result. + normalized, errors = validate_and_normalize_scan_config( + { + "aws": {"excluded_services": ["s3", "s3"]}, + "azure": {"vm_backup_min_daily_retention_days": 7}, + } + ) + assert normalized == {} + assert errors + assert any(err["path"].startswith("aws.") for err in errors) + + def test_value_error_prefix_is_stripped_from_user_facing_messages(self): + # Pydantic prefixes messages emitted from ``field_validator`` + # ValueError with ``"Value error, "``. If this test starts to fail + # because the prefix reappears, either pydantic changed the format + # or the strip in ``validate_and_normalize_scan_config`` was + # dropped — either way the UI would render the noisy prefix, so + # we lock the cleaned message in explicitly. + _, errors = validate_and_normalize_scan_config( + {"aws": {"excluded_services": ["s3", "s3"]}} + ) + assert errors + message = errors[0]["message"] + assert not message.startswith("Value error, ") + assert "duplicate values are not allowed" in message + + def test_all_errors_are_reported_not_only_the_first(self): + normalized, errors = validate_and_normalize_scan_config( + { + "aws": { + "excluded_checks": ["", ""], + "excluded_services": ["", ""], + } + } + ) + assert normalized == {} + # ``excluded_checks`` yields per-item empty-string errors AND a + # duplicate error; ``excluded_services`` yields the same set. + paths = {err["path"] for err in errors} + assert any(p.startswith("aws.excluded_checks") for p in paths) + assert any(p.startswith("aws.excluded_services") for p in paths) + + +class Test_Non_String_Provider_Keys: + """The normalized payload is later persisted in a Django JSONField + keyed by provider. Two entries whose ``str()`` collide (e.g. ``123`` + and ``"123"``) would silently overwrite each other, so non-string + keys are rejected up front instead of silently coerced.""" + + def test_non_string_key_is_rejected(self): + normalized, errors = validate_and_normalize_scan_config({123: {}}) + assert normalized == {} + assert errors == [{"path": "123", "message": "provider keys must be strings."}] + + def test_string_and_int_collision_does_not_silently_overwrite(self): + # If only ``str()`` coercion happened both keys would collapse to + # ``"aws"`` in the output — this test guards against that regression. + normalized, errors = validate_and_normalize_scan_config( + {"aws": {}, 123: {"a": 1}} + ) + assert normalized == {} + assert any(err["path"] == "123" for err in errors) + + +class Test_Unknown_Sections_Must_Be_JSON_Serializable: + """``normalized`` is persisted by the API in a Django JSONField, so + unknown provider sections must fail fast here instead of blowing up + at persist time. Registered sections cannot hit this path — they go + through ``model_dump(mode="json", ...)`` which already coerces.""" + + def test_set_inside_unknown_section_is_rejected(self): + # ``set`` is a common trap: ``yaml.safe_load`` never produces it, + # but a hand-built dict might. + normalized, errors = validate_and_normalize_scan_config( + {"future_provider": {"values": {1, 2, 3}}} + ) + assert normalized == {} + assert errors + assert errors[0]["path"] == "future_provider" + assert "JSON-serializable" in errors[0]["message"] + + def test_json_safe_unknown_section_is_still_preserved(self): + payload = {"future_provider": {"nested": {"k": [1, 2, 3]}}} + normalized, errors = validate_and_normalize_scan_config(payload) + assert errors == [] + assert normalized == payload + + +class Test_Backward_Compatible_Wrapper: + def test_valid_payload_yields_no_errors(self): + assert ( + validate_scan_config( + {"aws": {"excluded_checks": ["s3_bucket_public_access"]}} + ) + == [] + ) + + def test_invalid_payload_yields_only_the_errors(self): + errors = validate_scan_config({"aws": {"excluded_checks": ["", ""]}}) + assert errors + assert all(set(err) == {"path", "message"} for err in errors) + + def test_unknown_exclusion_yields_the_semantic_error(self): + assert validate_scan_config( + {"aws": {"excluded_services": ["not_a_real_aws_service"]}} + ) == [ + { + "path": "aws.excluded_services[0]", + "message": ( + "Unknown service 'not_a_real_aws_service' for provider 'aws'." + ), + } + ] + + def test_non_mapping_root_matches_new_contract(self): + assert validate_scan_config(None) == [ + { + "path": "", + "message": "Scan config must be a mapping with provider sections.", + } + ] diff --git a/tests/lib/scan/scan_exclusions_test.py b/tests/lib/scan/scan_exclusions_test.py new file mode 100644 index 0000000000..2c8f260c20 --- /dev/null +++ b/tests/lib/scan/scan_exclusions_test.py @@ -0,0 +1,178 @@ +"""Coverage for ``Scan`` constructor exclusion semantics. + +The Scan class is the single execution entry point used by both the CLI +and the API worker. Its exclusion validation must: + +- Reject duplicates and unknown identifiers with actionable errors. +- Validate excluded checks against the FULL provider catalog so a global + configuration can exclude a valid check that is not part of a scoped + run (see the SDK acceptance criteria for scan-configuration exclusions). +- Refuse a configuration that would leave nothing to execute. +- Produce a deterministic, sorted final scope. + +The catalog dependencies (``CheckMetadata.get_bulk``, ``Compliance.get_bulk``, +``list_services``, ``load_checks_to_execute``) are patched so tests stay +focused on the exclusion logic and avoid walking the provider package tree. +""" + +from unittest.mock import MagicMock, patch + +import pytest + +from prowler.lib.scan.exceptions.exceptions import ( + ScanInvalidCheckError, + ScanInvalidServiceError, +) +from prowler.lib.scan.scan import Scan +from tests.providers.aws.utils import set_mocked_aws_provider + +# The provider catalog for these tests: three checks spread across two +# services (``accessanalyzer`` and ``s3``). Keeps assertions readable. +PROVIDER_CATALOG = { + "accessanalyzer_enabled", + "s3_bucket_encryption_enabled", + "s3_bucket_public_access", +} +PROVIDER_SERVICES = ["accessanalyzer", "s3"] + + +@pytest.fixture +def scan_provider(): + provider = set_mocked_aws_provider() + metadata = MagicMock() + metadata.Categories = [] + bulk = {check: metadata for check in PROVIDER_CATALOG} + + with ( + patch( + "prowler.lib.scan.scan.CheckMetadata.get_bulk", + return_value=bulk, + ), + patch("prowler.lib.scan.scan.Compliance.get_bulk", return_value={}), + patch( + "prowler.lib.scan.scan.update_checks_metadata_with_compliance", + side_effect=lambda _compliance, checks: checks, + ), + patch( + "prowler.lib.scan.scan.load_checks_to_execute", + side_effect=lambda **kwargs: set(kwargs["check_list"] or PROVIDER_CATALOG), + ), + patch( + "prowler.lib.scan.scan.list_services", + return_value=PROVIDER_SERVICES, + ), + ): + yield provider + + +class Test_Exclusion_No_Ops: + def test_none_lists_are_no_ops(self, scan_provider): + scan = Scan(scan_provider, excluded_checks=None, excluded_services=None) + assert scan.checks_to_execute == sorted(PROVIDER_CATALOG) + + def test_empty_lists_are_no_ops(self, scan_provider): + scan = Scan(scan_provider, excluded_checks=[], excluded_services=[]) + assert scan.checks_to_execute == sorted(PROVIDER_CATALOG) + + +class Test_Excluded_Checks: + def test_valid_check_is_removed_from_the_scope(self, scan_provider): + scan = Scan( + scan_provider, + excluded_checks=["s3_bucket_public_access"], + ) + assert scan.checks_to_execute == sorted( + PROVIDER_CATALOG - {"s3_bucket_public_access"} + ) + + def test_excluded_check_may_be_outside_the_selected_scope(self, scan_provider): + # ``s3_bucket_public_access`` is not in the explicitly selected + # ``checks`` list but is still a valid provider check, so the + # global exclusion must be accepted and be a no-op for this run. + scan = Scan( + scan_provider, + checks=["accessanalyzer_enabled"], + excluded_checks=["s3_bucket_public_access"], + ) + assert scan.checks_to_execute == ["accessanalyzer_enabled"] + + def test_unknown_check_is_rejected(self, scan_provider): + with pytest.raises(ScanInvalidCheckError): + Scan(scan_provider, excluded_checks=["not_a_real_check"]) + + def test_duplicate_checks_are_rejected(self, scan_provider): + with pytest.raises(ScanInvalidCheckError): + Scan( + scan_provider, + excluded_checks=[ + "s3_bucket_public_access", + "s3_bucket_public_access", + ], + ) + + +class Test_Excluded_Services: + def test_service_exclusion_removes_every_check_in_the_service(self, scan_provider): + scan = Scan(scan_provider, excluded_services=["s3"]) + assert scan.checks_to_execute == ["accessanalyzer_enabled"] + + def test_unknown_service_is_rejected(self, scan_provider): + with pytest.raises(ScanInvalidServiceError): + Scan(scan_provider, excluded_services=["not_a_real_service"]) + + def test_duplicate_services_are_rejected(self, scan_provider): + with pytest.raises(ScanInvalidServiceError): + Scan(scan_provider, excluded_services=["s3", "s3"]) + + +class Test_Combined_Exclusions: + def test_selected_checks_plus_excluded_checks_and_services(self, scan_provider): + scan = Scan( + scan_provider, + checks=["accessanalyzer_enabled", "s3_bucket_encryption_enabled"], + excluded_checks=["s3_bucket_public_access"], + excluded_services=["s3"], + ) + # The explicit ``checks`` selection is narrowed by both the + # excluded_checks (drops nothing extra here) and excluded_services + # (drops every s3 check), leaving accessanalyzer alone. + assert scan.checks_to_execute == ["accessanalyzer_enabled"] + + def test_result_is_sorted_and_deterministic(self, scan_provider): + scan = Scan( + scan_provider, + excluded_checks=["s3_bucket_public_access"], + ) + assert scan.checks_to_execute == sorted(scan.checks_to_execute) + + +class Test_Empty_Final_Scope_Is_Rejected: + def test_excluding_every_service_is_rejected(self, scan_provider): + with pytest.raises(ScanInvalidCheckError): + Scan(scan_provider, excluded_services=PROVIDER_SERVICES) + + def test_excluding_every_check_is_rejected(self, scan_provider): + with pytest.raises(ScanInvalidCheckError): + Scan(scan_provider, excluded_checks=sorted(PROVIDER_CATALOG)) + + +class Test_Already_Empty_Scope_Does_Not_Blame_Exclusions: + """When a positive filter (severity, categories, checks that resolve + to nothing) leaves the scope empty *before* exclusions run, the + exclusion pass must not falsely claim to be the cause. Otherwise the + real reason (empty selection) is masked by a misleading error.""" + + def test_empty_initial_scope_with_valid_exclusions_does_not_raise( + self, scan_provider + ): + # Force ``load_checks_to_execute`` to return an empty scope while + # keeping the exclusion inputs valid against the provider catalog. + with patch( + "prowler.lib.scan.scan.load_checks_to_execute", + return_value=set(), + ): + scan = Scan( + scan_provider, + excluded_checks=["s3_bucket_public_access"], + ) + assert scan.checks_to_execute == [] diff --git a/ui/actions/organizations/organizations.adapter.test.ts b/ui/actions/organizations/organizations.adapter.test.ts index 07cff60c0c..41e0b8d89f 100644 --- a/ui/actions/organizations/organizations.adapter.test.ts +++ b/ui/actions/organizations/organizations.adapter.test.ts @@ -11,6 +11,7 @@ import { buildOrgTreeData, getOuIdsForSelectedAccounts, getSelectableAccountIds, + getSelectableAccountIdsForTarget, } from "./organizations.adapter"; const discoveryFixture: DiscoveryResult = { @@ -164,6 +165,68 @@ describe("buildAccountLookup", () => { }); }); +describe("getSelectableAccountIdsForTarget", () => { + it("scopes selection to accounts under a target OU, including nested OUs", () => { + // ou-parent contains ou-child (holds 111...) and the blocked 222... + const scoped = getSelectableAccountIdsForTarget( + discoveryFixture, + "ou-parent", + ); + + // Only the selectable descendant is returned; blocked 222... is excluded, + // and 333... (under the root, outside the OU) is not included. + expect(scoped).toEqual(["111111111111"]); + }); + + it("scopes selection to a leaf OU", () => { + const scoped = getSelectableAccountIdsForTarget( + discoveryFixture, + "ou-child", + ); + + expect(scoped).toEqual(["111111111111"]); + }); + + it("includes the deployment account even when it lives outside the target OU", () => { + // Deployment (management) account 333... sits under the root, but gets the + // role via DeployLocalRole, so it must be pre-selected alongside the OU. + const scoped = getSelectableAccountIdsForTarget( + discoveryFixture, + "ou-child", + "333333333333", + ); + + expect(scoped).toEqual(["111111111111", "333333333333"]); + }); + + it("does not include a deployment account that is not selectable", () => { + // 222... is blocked, so even as the deployment account it stays unselected. + const scoped = getSelectableAccountIdsForTarget( + discoveryFixture, + "ou-child", + "222222222222", + ); + + expect(scoped).toEqual(["111111111111"]); + }); + + it("returns every selectable account for a root target (whole organization)", () => { + const scoped = getSelectableAccountIdsForTarget(discoveryFixture, "r-root"); + + expect(scoped).toEqual(["111111111111", "333333333333"]); + }); + + it("falls back to all selectable accounts for an empty or unknown target", () => { + expect(getSelectableAccountIdsForTarget(discoveryFixture, "")).toEqual([ + "111111111111", + "333333333333", + ]); + expect( + getSelectableAccountIdsForTarget(discoveryFixture, "ou-does-not-exist"), + ).toEqual(["111111111111", "333333333333"]); + }); +}); + describe("getOuIdsForSelectedAccounts", () => { it("collects all ancestor OUs for selected accounts without duplicates", () => { const ouIds = getOuIdsForSelectedAccounts(discoveryFixture, [ diff --git a/ui/actions/organizations/organizations.adapter.ts b/ui/actions/organizations/organizations.adapter.ts index e120a1a3dc..2b0c7b5a0b 100644 --- a/ui/actions/organizations/organizations.adapter.ts +++ b/ui/actions/organizations/organizations.adapter.ts @@ -109,6 +109,71 @@ export function buildAccountLookup( return map; } +/** + * Returns the selectable account IDs that fall under a deployment target + * (an OU or root ID), optionally including the deployment account itself. + * + * The StackSet only rolls the role out to member accounts beneath the chosen + * target, and the deployment (management or delegated administrator) account + * gets the role via DeployLocalRole even though it usually lives outside that + * target. Pre-selecting exactly those accounts keeps the confirmation step in + * sync with what was actually deployed. + * + * Falls back to every selectable account when the target is empty or is not + * part of this discovery (e.g. a root ID), preserving the whole-organization + * default. + */ +export function getSelectableAccountIdsForTarget( + result: DiscoveryResult, + targetId: string, + deploymentAccountId?: string, +): string[] { + const selectableAccountIds = getSelectableAccountIds(result); + const normalizedTarget = targetId.trim(); + + if (!normalizedTarget) { + return selectableAccountIds; + } + + const isKnownOu = result.organizational_units.some( + (ou) => ou.id === normalizedTarget, + ); + + // Only a specific OU narrows the selection. A root ID (whole org) or an + // unknown target keeps the whole-organization default. + if (!isKnownOu) { + return selectableAccountIds; + } + + // Collect the target OU plus all of its nested descendant OUs. + const scopeIds = new Set([normalizedTarget]); + let addedNewOu = true; + while (addedNewOu) { + addedNewOu = false; + for (const ou of result.organizational_units) { + if (!scopeIds.has(ou.id) && scopeIds.has(ou.parent_id)) { + scopeIds.add(ou.id); + addedNewOu = true; + } + } + } + + const selectableSet = new Set(selectableAccountIds); + const scopedIds = new Set(); + + for (const account of result.accounts) { + if (scopeIds.has(account.parent_id) && selectableSet.has(account.id)) { + scopedIds.add(account.id); + } + } + + if (deploymentAccountId && selectableSet.has(deploymentAccountId)) { + scopedIds.add(deploymentAccountId); + } + + return selectableAccountIds.filter((id) => scopedIds.has(id)); +} + /** * Given selected account IDs, returns OU IDs that are ancestors of selected accounts. */ diff --git a/ui/app/(auth)/invitation/accept/accept-invitation-client.tsx b/ui/app/(auth)/invitation/accept/accept-invitation-client.tsx index cec9c8fc55..c412eaf48c 100644 --- a/ui/app/(auth)/invitation/accept/accept-invitation-client.tsx +++ b/ui/app/(auth)/invitation/accept/accept-invitation-client.tsx @@ -10,6 +10,7 @@ import { getInvitationErrorDisplay, INVITATION_ERROR_FLOW, } from "@/app/(auth)/invitation/_lib/invitation-errors"; +import { AuthBrand } from "@/components/auth/oss/auth-brand"; import { Button } from "@/components/shadcn"; type AcceptState = @@ -69,6 +70,7 @@ export function AcceptInvitationClient({ return (
+ {/* No token */} {state.kind === "no-token" && (
diff --git a/ui/app/(auth)/layout.tsx b/ui/app/(auth)/layout.tsx index b0a9bc2061..27fff93fe5 100644 --- a/ui/app/(auth)/layout.tsx +++ b/ui/app/(auth)/layout.tsx @@ -5,7 +5,6 @@ import { Metadata, Viewport } from "next"; import { connection } from "next/server"; import { ReactNode, Suspense } from "react"; -import { PublicAuthShell } from "@/components/auth/oss/public-auth-shell"; import { RuntimePublicConfig } from "@/components/runtime-config/runtime-public-config"; import { NavigationProgress, Toaster } from "@/components/shadcn"; import { fontMono, fontSans } from "@/config/fonts"; @@ -66,7 +65,7 @@ export default async function AuthLayout({ - {children} + {children} {gtmId && } diff --git a/ui/app/(prowler)/_overview/_components/lighthouse-overview-banner.test.tsx b/ui/app/(prowler)/_overview/_components/lighthouse-overview-banner.test.tsx index 62bb053f3a..255ac50c63 100644 --- a/ui/app/(prowler)/_overview/_components/lighthouse-overview-banner.test.tsx +++ b/ui/app/(prowler)/_overview/_components/lighthouse-overview-banner.test.tsx @@ -1,39 +1,52 @@ import { render, screen } from "@testing-library/react"; import { describe, expect, it } from "vitest"; +import { LIGHTHOUSE_OVERVIEW_BANNER_HREF } from "../_lib/lighthouse-banner"; import { LighthouseOverviewBanner } from "./lighthouse-overview-banner"; -const REMEDIATION_PROMPT = - "Find and guide me to remediate which actually matters. What do I have to do today to be secure?"; -const REMEDIATION_HREF = `/lighthouse?prompt=${encodeURIComponent( - REMEDIATION_PROMPT, -)}` as const; - describe("LighthouseOverviewBanner", () => { - it("renders Toni copy and links to Lighthouse when connected", () => { + it("renders Toni copy and opens a prompted chat when connected", () => { // Given / When - render(); + render( + , + ); // Then const link = screen.getByRole("link", { - name: /Find and remediate which actually matters\./, + name: /Find and remediate what actually matters\./, }); - expect(link).toHaveAttribute("href", REMEDIATION_HREF); + expect(link).toHaveAttribute("href", LIGHTHOUSE_OVERVIEW_BANNER_HREF.CHAT); expect(link).toHaveTextContent("Lighthouse AI"); - expect(link).toHaveTextContent( - "Find and remediate which actually matters.", - ); + expect(link).toHaveTextContent("Find and remediate what actually matters."); }); it("links to Lighthouse settings when no connected configuration exists", () => { // Given / When - render(); + render( + , + ); // Then expect( screen.getByRole("link", { - name: /Find and remediate which actually matters\./, + name: /Find and remediate what actually matters\./, }), ).toHaveAttribute("href", "/lighthouse/settings"); }); + + it("isolates its stacking so content never paints over the sticky navbar", () => { + // Given / When: the banner's inner z-10 must stay scoped to the card — + // without isolation it ties the sticky header's z-10 and wins by DOM order + render( + , + ); + + // Then + const card = screen + .getByRole("link", { name: /Find and remediate what actually matters\./ }) + .querySelector("[data-slot='card']"); + expect(card).toHaveClass("isolate"); + }); }); diff --git a/ui/app/(prowler)/_overview/_components/lighthouse-overview-banner.tsx b/ui/app/(prowler)/_overview/_components/lighthouse-overview-banner.tsx index eb2e0539fb..8bd8dea414 100644 --- a/ui/app/(prowler)/_overview/_components/lighthouse-overview-banner.tsx +++ b/ui/app/(prowler)/_overview/_components/lighthouse-overview-banner.tsx @@ -66,7 +66,9 @@ export function LighthouseOverviewBanner({ @@ -116,7 +118,7 @@ export function LighthouseOverviewBanner({ Lighthouse AI

- Find and remediate which actually matters. + Find and remediate what actually matters.

diff --git a/ui/app/(prowler)/_overview/_lib/lighthouse-banner.test.ts b/ui/app/(prowler)/_overview/_lib/lighthouse-banner.test.ts index 7668cfafac..baf5e8206f 100644 --- a/ui/app/(prowler)/_overview/_lib/lighthouse-banner.test.ts +++ b/ui/app/(prowler)/_overview/_lib/lighthouse-banner.test.ts @@ -16,7 +16,7 @@ const LIGHTHOUSE_OVERVIEW_CHAT_HREF = `/lighthouse?prompt=${encodeURIComponent( )}`; describe("resolveLighthouseOverviewBannerHref", () => { - it("routes to Lighthouse chat when any v2 configuration is connected", () => { + it("opens Lighthouse with the remediation prompt when any v2 configuration is connected", () => { // Given / When const href = resolveLighthouseOverviewBannerHref([ configuration("openai", false), diff --git a/ui/app/(prowler)/_overview/_lib/lighthouse-banner.ts b/ui/app/(prowler)/_overview/_lib/lighthouse-banner.ts index a80522bf6c..b0b0d1a08c 100644 --- a/ui/app/(prowler)/_overview/_lib/lighthouse-banner.ts +++ b/ui/app/(prowler)/_overview/_lib/lighthouse-banner.ts @@ -2,8 +2,9 @@ import type { LighthouseV2Configuration } from "@/app/(prowler)/lighthouse/_type import { LIGHTHOUSE_ROUTE } from "@/lib/lighthouse-routes"; import type { ServerActionResult } from "@/types/server-actions"; +// Prefilled in the composer when the overview banner opens Lighthouse. export const LIGHTHOUSE_OVERVIEW_PROMPT = - "Find and guide me to remediate which actually matters. What do I have to do today to be secure?"; + "Find and guide me to remediate what actually matters. What do I have to do today to be secure?"; export const LIGHTHOUSE_OVERVIEW_BANNER_HREF = { CHAT: `${LIGHTHOUSE_ROUTE.CHAT}?prompt=${encodeURIComponent(LIGHTHOUSE_OVERVIEW_PROMPT)}`, diff --git a/ui/app/(prowler)/layout.tsx b/ui/app/(prowler)/layout.tsx index 801b1ea28f..f657f32199 100644 --- a/ui/app/(prowler)/layout.tsx +++ b/ui/app/(prowler)/layout.tsx @@ -16,6 +16,7 @@ import { RuntimePublicConfig } from "@/components/runtime-config/runtime-public- import { NavigationProgress } from "@/components/shadcn/navigation-progress"; import { Toaster } from "@/components/shadcn/toast"; import { TaskPollingWatcher } from "@/components/shared/task-polling-watcher"; +import { GlobalSidePanel } from "@/components/side-panel"; import { fontMono, fontSans } from "@/config/fonts"; import { siteConfig } from "@/config/site"; import { isCloud } from "@/lib/shared/env"; @@ -107,6 +108,9 @@ export default async function RootLayout({ )} {children} + {/* Always mounted: it hosts the detail (finding/resource) views in + every deployment; the AI tab inside is cloud-gated on its own. */} + {/* Resumes persisted background-task polling (e.g. cross-provider PDF generation) so completion toasts survive hard reloads. */} diff --git a/ui/app/(prowler)/lighthouse/_components/chat/composer.tsx b/ui/app/(prowler)/lighthouse/_components/chat/composer.tsx index 62d39cfb05..b9f854dbdb 100644 --- a/ui/app/(prowler)/lighthouse/_components/chat/composer.tsx +++ b/ui/app/(prowler)/lighthouse/_components/chat/composer.tsx @@ -118,7 +118,9 @@ function ChatComposer({ aria-label="Message" value={input} onChange={(event) => onInputChange(event.target.value)} - disabled={!canSend} + // Typing stays available while a response streams (sending is what is + // gated, via canSend); only a disconnected provider blocks the input. + disabled={!selectedConfigurationConnected} placeholder={ selectedConfigurationConnected ? "Ask a question" diff --git a/ui/app/(prowler)/lighthouse/_components/chat/empty-state.tsx b/ui/app/(prowler)/lighthouse/_components/chat/empty-state.tsx index 021be182e2..b84c71b3b7 100644 --- a/ui/app/(prowler)/lighthouse/_components/chat/empty-state.tsx +++ b/ui/app/(prowler)/lighthouse/_components/chat/empty-state.tsx @@ -5,6 +5,7 @@ import { type ReactNode, type SubmitEvent } from "react"; import { LighthouseIconWithAura } from "@/components/icons"; import { Button } from "@/components/shadcn/button/button"; +import { cn } from "@/lib/utils"; import { ChatComposerPanel } from "./composer"; import { DecryptedText } from "./decrypted-text"; @@ -45,34 +46,58 @@ interface ChatEmptyStateProps { onInputChange: (value: string) => void; onSubmit: (event: SubmitEvent) => void; onSubmitText: (text: string) => Promise; + footer?: ReactNode; + // Side-panel variant: smaller logo and static (non-animated) copy — the + // decrypt animation reflows multi-line text in narrow widths. + compact?: boolean; } export function ChatEmptyState({ onInputChange, + footer, + compact = false, ...composerPanelProps }: ChatEmptyStateProps) { return (
- +
-

- +

+ {compact ? ( + "Find and remediate what actually matters." + ) : ( + + )}

-

- +

+ {compact ? ( + "What do you want to know today?" + ) : ( + + )}

@@ -101,6 +126,7 @@ export function ChatEmptyState({ ); })}
+ {footer ?
{footer}
: null}
); diff --git a/ui/app/(prowler)/lighthouse/_components/chat/lighthouse-chat-store-provider.tsx b/ui/app/(prowler)/lighthouse/_components/chat/lighthouse-chat-store-provider.tsx new file mode 100644 index 0000000000..55866304bb --- /dev/null +++ b/ui/app/(prowler)/lighthouse/_components/chat/lighthouse-chat-store-provider.tsx @@ -0,0 +1,41 @@ +"use client"; + +import { createContext, type ReactNode, useContext } from "react"; +import { useStore } from "zustand"; + +import type { + LighthouseChatState, + LighthouseChatStore, +} from "@/app/(prowler)/lighthouse/_lib/chat-store"; + +const LighthouseChatStoreContext = createContext( + null, +); + +interface LighthouseChatStoreProviderProps { + store: LighthouseChatStore; + children: ReactNode; +} + +export function LighthouseChatStoreProvider({ + store, + children, +}: LighthouseChatStoreProviderProps) { + return ( + + {children} + + ); +} + +export function useLighthouseChatStore( + selector: (state: LighthouseChatState) => T, +): T { + const store = useContext(LighthouseChatStoreContext); + if (!store) { + throw new Error( + "useLighthouseChatStore must be used within LighthouseChatStoreProvider", + ); + } + return useStore(store, selector); +} diff --git a/ui/app/(prowler)/lighthouse/_components/chat/lighthouse-v2-chat-page.test.tsx b/ui/app/(prowler)/lighthouse/_components/chat/lighthouse-v2-chat-page.test.tsx index 5642762901..b6c5890db4 100644 --- a/ui/app/(prowler)/lighthouse/_components/chat/lighthouse-v2-chat-page.test.tsx +++ b/ui/app/(prowler)/lighthouse/_components/chat/lighthouse-v2-chat-page.test.tsx @@ -3,10 +3,18 @@ import userEvent from "@testing-library/user-event"; import { type ReactNode } from "react"; import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; +import { + getOrCreatePanelChatStore, + resetPanelChatStoreForTests, +} from "@/app/(prowler)/lighthouse/_lib/panel-chat-store"; import { LIGHTHOUSE_V2_SESSIONS_CHANGED_EVENT, notifyLighthouseV2SessionArchived, } from "@/app/(prowler)/lighthouse/_lib/session-events"; +import { + type MockEventSource, + stubEventSource, +} from "@/app/(prowler)/lighthouse/_lib/testing/event-source-mock"; import type { LighthouseV2Configuration, LighthouseV2Message, @@ -16,51 +24,8 @@ import type { import { LighthouseV2ChatPage } from "./lighthouse-v2-chat-page"; -// Controllable EventSource mock: records each instance so tests can drive -// named SSE events and connection failures, while still being a vi.fn so -// `expect(EventSource).toHaveBeenCalledWith(...)` keeps working. -interface MockEventSource { - url: string; - readyState: number; - onerror: ((event: Event) => void) | null; - listeners: Map>; - addEventListener: (type: string, cb: EventListener) => void; - close: ReturnType; - emit: (type: string, data: unknown) => void; - fail: (readyState: number) => void; -} - let eventSources: MockEventSource[] = []; -function stubEventSource() { - eventSources = []; - const EventSourceMock = vi.fn(function (this: MockEventSource, url: string) { - this.url = url; - this.readyState = 0; - this.onerror = null; - this.listeners = new Map(); - this.addEventListener = (type: string, cb: EventListener) => { - const set = this.listeners.get(type) ?? new Set(); - set.add(cb); - this.listeners.set(type, set); - }; - this.close = vi.fn(() => { - this.readyState = 2; - }); - this.emit = (type: string, data: unknown) => { - const event = new MessageEvent(type, { data: JSON.stringify(data) }); - this.listeners.get(type)?.forEach((cb) => cb(event)); - }; - this.fail = (readyState: number) => { - this.readyState = readyState; - this.onerror?.(new Event("error")); - }; - eventSources.push(this); - }); - Object.assign(EventSourceMock, { CONNECTING: 0, OPEN: 1, CLOSED: 2 }); - vi.stubGlobal("EventSource", EventSourceMock); -} - const { createSessionMock, getMessagesMock, @@ -142,11 +107,8 @@ describe("LighthouseV2ChatPage", () => { getMessagesMock.mockReset(); sendMessageMock.mockReset(); updateConfigurationMock.mockReset(); - // The mock never fires "open": the client must POST the message without - // waiting for it (the backend sends no bytes until the worker emits, which - // only happens after the POST). This is the regression guard for the - // open-gate deadlock. - stubEventSource(); + resetPanelChatStoreForTests(); + eventSources = stubEventSource(); createSessionMock.mockResolvedValue({ data: { @@ -185,6 +147,62 @@ describe("LighthouseV2ChatPage", () => { ).toHaveAttribute("href", "/lighthouse/settings"); }); + it("renders the empty-state headline with correct wording", () => { + // Given / When + renderPage(); + + // Then + expect( + screen.getByText("Find and remediate what actually matters."), + ).toBeInTheDocument(); + }); + + it("continues using the panel chat store on the full-page surface", () => { + // Given: the panel owns an in-progress new chat with a draft + const panelStore = getOrCreatePanelChatStore({ + configurations, + modelsByProvider, + supportedProviders, + }); + panelStore.getState().setInput("Draft from the side panel"); + + // When + renderPage(); + + // Then: the page owns the same live store, not a stale server snapshot + const input = screen.getByRole("textbox", { name: "Message" }); + expect(input).toHaveValue("Draft from the side panel"); + act(() => panelStore.getState().setInput("Updated after navigation")); + expect(input).toHaveValue("Updated after navigation"); + }); + + it("enables session URL sync after claiming a new panel chat", async () => { + // Given: the panel owns a new chat before full-page navigation + const user = userEvent.setup(); + getOrCreatePanelChatStore({ + configurations, + modelsByProvider, + supportedProviders, + }); + const replaceStateSpy = vi.spyOn(window.history, "replaceState"); + renderPage(); + + // When: the first page message creates its session + await user.type( + screen.getByRole("textbox", { name: "Message" }), + ["Summarize findings", "{Enter}"].join(""), + ); + + // Then: the claimed panel store now follows the full-page URL contract + await waitFor(() => + expect(replaceStateSpy).toHaveBeenCalledWith( + window.history.state, + "", + "/lighthouse?session=session-1", + ), + ); + }); + it("shows the current OpenAI model without a selector when OpenAI is the only connected provider", () => { // Given / When renderPage({ @@ -208,7 +226,7 @@ describe("LighthouseV2ChatPage", () => { expect(within(currentModel).getByText("GPT-5.1")).toBeInTheDocument(); }); - it("defaults to gpt-5.5 when OpenAI has no remembered model", () => { + it("defaults to gpt-5.6-terra when OpenAI has no remembered model", () => { // Given / When renderPage({ configurations: [ @@ -216,15 +234,20 @@ describe("LighthouseV2ChatPage", () => { { ...configurations[1], connected: false }, ], modelsByProvider: { - openai: [model("gpt-4.1", "GPT-4.1"), model("gpt-5.5", "GPT-5.5")], + openai: [ + model("gpt-4.1", "GPT-4.1"), + model("gpt-5.6-terra", "GPT-5.6 Terra"), + ], bedrock: [model("anthropic.claude-4")], "openai-compatible": [model("llama-3.3")], }, }); // Then - const currentModel = screen.getByLabelText("Current model: OpenAI GPT-5.5"); - expect(within(currentModel).getByText("GPT-5.5")).toBeInTheDocument(); + const currentModel = screen.getByLabelText( + "Current model: OpenAI GPT-5.6 Terra", + ); + expect(within(currentModel).getByText("GPT-5.6 Terra")).toBeInTheDocument(); }); it("uses the AWS onboarding quick prompt instead of the docs prompt", async () => { @@ -251,7 +274,7 @@ describe("LighthouseV2ChatPage", () => { it("prefills the overview remediation prompt without starting a conversation", () => { // Given const initialPrompt = - "Find and guide me to remediate which actually matters. What do I have to do today to be secure?"; + "Find and guide me to remediate what actually matters. What do I have to do today to be secure?"; // When renderPage({ initialPrompt }); @@ -688,6 +711,28 @@ describe("LighthouseV2ChatPage", () => { expect(screen.getByText("Existing answer")).toBeInTheDocument(); }); + it("lets the user draft the next message while a response is streaming, without sending it", async () => { + // Given: a message is in flight (spinner replaces the send button) + const user = userEvent.setup(); + renderPage(); + await user.type( + screen.getByRole("textbox", { name: "Message" }), + ["Summarize findings", "{Enter}"].join(""), + ); + await waitFor(() => expect(sendMessageMock).toHaveBeenCalledTimes(1)); + expect( + screen.getByRole("status", { name: "Generating response" }), + ).toBeInTheDocument(); + + // When: the user types a follow-up and presses Enter mid-stream + const input = screen.getByRole("textbox", { name: "Message" }); + await user.type(input, ["Next question", "{Enter}"].join("")); + + // Then: the draft is kept in the input and no second message is sent + expect(input).toHaveValue("Next question"); + expect(sendMessageMock).toHaveBeenCalledTimes(1); + }); + it("surfaces a connection error when the stream closes without retrying", async () => { // Given const user = userEvent.setup(); diff --git a/ui/app/(prowler)/lighthouse/_components/chat/lighthouse-v2-chat-page.tsx b/ui/app/(prowler)/lighthouse/_components/chat/lighthouse-v2-chat-page.tsx index 3062874649..14b1e18b9d 100644 --- a/ui/app/(prowler)/lighthouse/_components/chat/lighthouse-v2-chat-page.tsx +++ b/ui/app/(prowler)/lighthouse/_components/chat/lighthouse-v2-chat-page.tsx @@ -1,61 +1,33 @@ "use client"; -import { type SubmitEvent, useRef, useState } from "react"; +import { useState } from "react"; import { - createLighthouseV2Session, - getLighthouseV2Messages, - sendLighthouseV2Message, - updateLighthouseV2Configuration, -} from "@/app/(prowler)/lighthouse/_actions"; + createLighthouseChatStore, + type LighthouseChatStore, +} from "@/app/(prowler)/lighthouse/_lib/chat-store"; import { - Conversation, - ConversationContent, - ConversationScrollButton, -} from "@/app/(prowler)/lighthouse/_components/ai-elements/conversation"; + getPanelChatStoreForSession, + isPanelChatStore, +} from "@/app/(prowler)/lighthouse/_lib/panel-chat-store"; import { - createInitialLighthouseV2StreamState, - type LighthouseV2StreamState, - reduceLighthouseV2Event, -} from "@/app/(prowler)/lighthouse/_lib/event-reducer"; -import { - buildOptimisticMessage, - buildSessionTitle, -} from "@/app/(prowler)/lighthouse/_lib/messages"; -import { - buildLighthouseV2ModelSelectionValue, - type LighthouseV2ModelSelection, - parseLighthouseV2ModelSelectionValue, -} from "@/app/(prowler)/lighthouse/_lib/model-selection"; -import { - LIGHTHOUSE_V2_NEW_CHAT_EVENT, - LIGHTHOUSE_V2_SESSION_ARCHIVED_EVENT, - notifyLighthouseV2SessionsChanged, + onLighthouseV2NewChat, + onLighthouseV2SessionArchived, } from "@/app/(prowler)/lighthouse/_lib/session-events"; -import { parseStreamEvent } from "@/app/(prowler)/lighthouse/_lib/stream-event-parser"; -import { buildLighthouseV2StreamUrl } from "@/app/(prowler)/lighthouse/_lib/stream-url"; -import { - LIGHTHOUSE_V2_PROVIDER_TYPE, - LIGHTHOUSE_V2_SSE_EVENT, - type LighthouseV2Configuration, - type LighthouseV2Message, - type LighthouseV2ProviderType, - type LighthouseV2SSEEvent, - type LighthouseV2SupportedModel, - type LighthouseV2SupportedProvider, +import type { + LighthouseV2Configuration, + LighthouseV2Message, + LighthouseV2ProviderType, + LighthouseV2SupportedModel, + LighthouseV2SupportedProvider, } from "@/app/(prowler)/lighthouse/_types"; -import { Card } from "@/components/shadcn"; -import { - Combobox, - type ComboboxGroup, -} from "@/components/shadcn/combobox/combobox"; import { useMountEffect } from "@/hooks/use-mount-effect"; -import { ProviderIcon } from "../config/provider-icon"; -import { ChatComposerPanel } from "./composer"; -import { ChatEmptyState } from "./empty-state"; -import { MessageBubble } from "./message-bubble"; -import { StreamingAssistantMessage } from "./streaming-message"; +import { LighthouseChatStoreProvider } from "./lighthouse-chat-store-provider"; +import { + LIGHTHOUSE_CHAT_SURFACE, + LighthouseV2ChatView, +} from "./lighthouse-v2-chat-view"; interface LighthouseV2ChatPageProps { configurations: LighthouseV2Configuration[]; @@ -79,572 +51,61 @@ export function LighthouseV2ChatPage({ initialPrompt, initialError, }: LighthouseV2ChatPageProps) { - const eventSourceRef = useRef(null); - const connectedConfigurations = configurations.filter( - (configuration) => configuration.connected === true, - ); - const initialModelSelection = resolveInitialModelSelection( - connectedConfigurations, - modelsByProvider, - ); - const [selectedModelSelection, setSelectedModelSelection] = - useState(initialModelSelection); - const [modelPreferenceSaving, setModelPreferenceSaving] = useState(false); - const [activeSessionId, setActiveSessionId] = useState( - initialSessionId ?? null, - ); - // Mirror for window listeners registered on mount, whose closures would - // otherwise keep the first render's activeSessionId. - const activeSessionIdRef = useRef(initialSessionId ?? null); - const [messages, setMessages] = useState(initialMessages); - const [input, setInput] = useState(initialPrompt ?? ""); - const [feedback, setFeedback] = useState(initialError ?? null); - const [blockedByConflict, setBlockedByConflict] = useState(false); - const [isSubmitting, setIsSubmitting] = useState(false); - const [lastSubmittedText, setLastSubmittedText] = useState( - null, - ); - const [streamState, setStreamState] = useState(() => - createInitialLighthouseV2StreamState(), - ); - const selectedConfiguration = selectedModelSelection - ? connectedConfigurations.find( - (configuration) => - configuration.providerType === selectedModelSelection.providerType, - ) - : undefined; - const modelSelectorGroups = buildModelSelectorGroups( - connectedConfigurations, - modelsByProvider, - supportedProviders, - ); - const showStaticOpenAIModel = - isOnlyConnectedProvider( - connectedConfigurations, - LIGHTHOUSE_V2_PROVIDER_TYPE.OPENAI, - ) && - selectedModelSelection?.providerType === LIGHTHOUSE_V2_PROVIDER_TYPE.OPENAI; - const selectedModelLabel = selectedModelSelection - ? getModelSelectionLabel(selectedModelSelection, modelsByProvider) - : "No model selected"; - const selectedProviderName = selectedModelSelection - ? getProviderDisplayName( - selectedModelSelection.providerType, - supportedProviders, - ) - : "OpenAI"; - const selectedModelValue = selectedModelSelection - ? buildLighthouseV2ModelSelectionValue( - selectedModelSelection.providerType, - selectedModelSelection.modelId, - ) - : ""; - - const canSend = - selectedConfiguration?.connected === true && - Boolean(selectedModelSelection?.modelId) && - !streamState.activeTaskId && - !blockedByConflict && - !isSubmitting; - - const refreshMessages = async (sessionId: string): Promise => { - const result = await getLighthouseV2Messages(sessionId); - // The fetch is async, so a reset (new chat, or archiving this session) can - // land while it is in flight. Drop the stale result instead of repopulating - // a chat that no longer points at this session. - if (sessionId !== activeSessionIdRef.current) return false; - if ("data" in result) { - setMessages(result.data); - return true; - } - return false; - }; - - const closeStream = () => { - eventSourceRef.current?.close(); - eventSourceRef.current = null; - }; - - const handleTerminalEvent = async ( - sessionId: string, - event: LighthouseV2SSEEvent, - ) => { - if ( - event.type === LIGHTHOUSE_V2_SSE_EVENT.MESSAGE_END || - event.type === LIGHTHOUSE_V2_SSE_EVENT.ERROR - ) { - closeStream(); - setBlockedByConflict(false); - if (event.type === LIGHTHOUSE_V2_SSE_EVENT.ERROR) { - setFeedback(event.detail || "Agent run failed."); - } - const refreshed = await refreshMessages(sessionId); - if (refreshed) { - setStreamState(createInitialLighthouseV2StreamState()); - } - notifyLighthouseV2SessionsChanged(); - } - }; - - const startStream = (streamUrl: string, sessionId: string) => { - closeStream(); - const source = new EventSource(streamUrl); - eventSourceRef.current = source; - - const applyEvent = (event: LighthouseV2SSEEvent) => { - setStreamState((current) => reduceLighthouseV2Event(current, event)); - void handleTerminalEvent(sessionId, event); - }; - - source.addEventListener("message.delta", (event) => - applyEvent( - parseStreamEvent(event, LIGHTHOUSE_V2_SSE_EVENT.MESSAGE_DELTA), - ), + // Navigation from the side panel transfers its live store so drafts, + // streamed output and the open EventSource continue without a snapshot gap. + // Direct/session-mismatched navigation builds the normal page-owned store. + const [store] = useState(() => { + const panelStore = + initialPrompt === undefined + ? getPanelChatStoreForSession(initialSessionId) + : null; + return ( + panelStore ?? + createLighthouseChatStore({ + config: { configurations, modelsByProvider, supportedProviders }, + syncUrlToSession: true, + initialSessionId, + initialMessages, + initialInput: initialPrompt, + initialError, + }) ); - source.addEventListener("tool_call.start", (event) => - applyEvent( - parseStreamEvent(event, LIGHTHOUSE_V2_SSE_EVENT.TOOL_CALL_START), - ), - ); - source.addEventListener("tool_call.end", (event) => - applyEvent( - parseStreamEvent(event, LIGHTHOUSE_V2_SSE_EVENT.TOOL_CALL_END), - ), - ); - source.addEventListener("message.end", (event) => - applyEvent(parseStreamEvent(event, LIGHTHOUSE_V2_SSE_EVENT.MESSAGE_END)), - ); - source.addEventListener("error", (event) => { - if (event instanceof MessageEvent) { - applyEvent(parseStreamEvent(event, LIGHTHOUSE_V2_SSE_EVENT.ERROR)); - } - }); - // The browser fires `onerror` both on a transient drop (it auto-reconnects) - // and on a non-retryable failure such as a 401/404 on the SSE GET. Only the - // latter leaves the source CLOSED, so surface a connection error there and - // treat everything else as a reconnect. - source.onerror = () => { - if (eventSourceRef.current !== source) return; - if (source.readyState === EventSource.CLOSED) { - closeStream(); - setFeedback("Unable to connect to the response stream."); - } - setStreamState((current) => - reduceLighthouseV2Event(current, { type: "disconnect" }), - ); - }; - }; + }); + const reusesPanelStore = isPanelChatStore(store); - const ensureSession = async (text: string) => { - if (activeSessionId) { - return activeSessionId; + // A reused panel store returns to panel URL semantics when the page leaves; + // a page-owned store closes its EventSource as before. + useMountEffect(() => { + if (reusesPanelStore) { + store.getState().setSessionUrlSyncEnabled(true); } - - const title = buildSessionTitle(text); - const result = await createLighthouseV2Session(title); - if ("error" in result) { - setFeedback(result.error); - return null; - } - - // Update the URL in place (not router.push) so the force-dynamic server - // component is NOT re-run mid-submit. A re-run would change `key` in - // page.tsx and remount this component, tearing down the open EventSource. - replaceLighthouseV2SessionUrl(result.data.id); - setActiveSessionId(result.data.id); - activeSessionIdRef.current = result.data.id; - notifyLighthouseV2SessionsChanged(); - return result.data.id; - }; - - const submitMessage = async (text: string) => { - const trimmedText = text.trim(); - if (!trimmedText) return; - if (!selectedModelSelection) { - setFeedback("Select a model before sending a message."); - return; - } - if (!canSend) return; - - setIsSubmitting(true); - try { - const sessionId = await ensureSession(trimmedText); - if (!sessionId) return; - - const provisionalTaskId = `pending-${Date.now()}`; - setFeedback(null); - setBlockedByConflict(false); - setLastSubmittedText(trimmedText); - setInput(""); - setMessages((current) => [ - ...current, - buildOptimisticMessage("user", trimmedText), - ]); - setStreamState(createInitialLighthouseV2StreamState(provisionalTaskId)); - - // Subscribe to the same-origin SSE proxy BEFORE sending the message: the - // backend has no replay buffer, so the listener must be attached before - // the worker starts emitting. - startStream(buildLighthouseV2StreamUrl(sessionId), sessionId); - - const result = await sendLighthouseV2Message({ - sessionId, - text: trimmedText, - provider: selectedModelSelection.providerType, - model: selectedModelSelection.modelId, - }); - - if ("error" in result) { - closeStream(); - setStreamState(createInitialLighthouseV2StreamState()); - setFeedback(result.error); - if (result.status === 409) { - setBlockedByConflict(true); - await refreshMessages(sessionId); - } + return () => { + if (reusesPanelStore) { + store.getState().setSessionUrlSyncEnabled(false); return; } - - setStreamState((current) => - current.activeTaskId === provisionalTaskId - ? { ...current, activeTaskId: result.data.task.id } - : current, - ); - notifyLighthouseV2SessionsChanged(); - } finally { - setIsSubmitting(false); - } - }; - - const handleModelValueChange = (value: string) => { - const selection = parseLighthouseV2ModelSelectionValue(value); - if (!selection) return; - void handleModelSelectionChange(selection); - }; - - const handleModelSelectionChange = async ( - selection: LighthouseV2ModelSelection, - ) => { - // The selection drives the model used for the next message, so it stays - // applied even if persisting it as the provider's default model fails — - // reverting it would make a connected provider unusable when the save 4xxs. - setSelectedModelSelection(selection); - setFeedback(null); - - const configId = connectedConfigurations.find( - (configuration) => configuration.providerType === selection.providerType, - )?.id; - if (!configId) return; - - setModelPreferenceSaving(true); - - const result = await updateLighthouseV2Configuration(configId, { - defaultModel: selection.modelId, - }); - - setModelPreferenceSaving(false); - - if ("error" in result) { - setFeedback(result.error); - } - }; - - const handleSubmit = (event: SubmitEvent) => { - event.preventDefault(); - void submitMessage(input); - }; - - // Close any open EventSource when the chat unmounts (e.g. route/session change). - useMountEffect(() => { - return () => closeStream(); + store.getState().destroy(); + }; }); - const resetToNewChat = () => { - closeStream(); - setActiveSessionId(null); - activeSessionIdRef.current = null; - setMessages([]); - setInput(""); - setFeedback(null); - setBlockedByConflict(false); - setIsSubmitting(false); - setLastSubmittedText(null); - setStreamState(createInitialLighthouseV2StreamState()); - replaceLighthouseV2SessionUrl(null); - }; - // The sidebar "+" can't rely on routing to reset the latest conversation (its // URL was set via replaceState, invisible to Next's router), so reset in place. useMountEffect(() => { - window.addEventListener(LIGHTHOUSE_V2_NEW_CHAT_EVENT, resetToNewChat); - return () => - window.removeEventListener(LIGHTHOUSE_V2_NEW_CHAT_EVENT, resetToNewChat); + const unsubscribeNewChat = onLighthouseV2NewChat(() => + store.getState().resetToNewChat(), + ); + const unsubscribeSessionArchived = onLighthouseV2SessionArchived( + (sessionId) => store.getState().handleSessionArchived(sessionId), + ); + return () => { + unsubscribeNewChat(); + unsubscribeSessionArchived(); + }; }); - // Archiving deletes the session; when it's the open one, fall back to a new - // chat instead of leaving a dead conversation and its URL on screen. - useMountEffect(() => { - const handleSessionArchived = (event: Event) => { - const archivedId = (event as CustomEvent<{ sessionId: string }>).detail - ?.sessionId; - if (archivedId && archivedId === activeSessionIdRef.current) { - resetToNewChat(); - } - }; - - window.addEventListener( - LIGHTHOUSE_V2_SESSION_ARCHIVED_EVENT, - handleSessionArchived, - ); - return () => - window.removeEventListener( - LIGHTHOUSE_V2_SESSION_ARCHIVED_EVENT, - handleSessionArchived, - ); - }); - - const hasLiveAssistantActivity = - Boolean(streamState.activeTaskId) || - Boolean(streamState.assistantText) || - streamState.toolCalls.length > 0; - const hasConversation = messages.length > 0 || hasLiveAssistantActivity; - - const composerPanelProps = { - feedback, - canRetry: - streamState.status === "disconnected" && lastSubmittedText !== null, - onRetry: () => - lastSubmittedText ? void submitMessage(lastSubmittedText) : undefined, - onDismissFeedback: () => setFeedback(null), - canSend, - input, - isStreaming: Boolean(streamState.activeTaskId), - modelSelector: showStaticOpenAIModel ? ( - - ) : ( -
- -
- ), - selectedConfigurationConnected: selectedConfiguration?.connected === true, - onInputChange: setInput, - onSubmit: handleSubmit, - onSubmitText: submitMessage, - }; - return ( - - {hasConversation ? ( -
-
- - - {messages.map((message) => ( - - ))} - {hasLiveAssistantActivity && ( - - )} - - - -
-
-
-
- -
-
-
- ) : ( - - )} - - ); -} - -function replaceLighthouseV2SessionUrl(sessionId: string | null) { - const url = sessionId - ? `/lighthouse?session=${encodeURIComponent(sessionId)}` - : "/lighthouse"; - - window.history.replaceState(window.history.state, "", url); -} - -function CurrentModelDisplay({ - provider, - providerName, - modelName, -}: { - provider: LighthouseV2ProviderType; - providerName: string; - modelName: string; -}) { - return ( -
- - - {providerName} - - - {modelName} - -
- ); -} - -// Fixed precedence used to pick which connected provider opens the chat. Any -// provider outside this list keeps its relative order behind these. -const LIGHTHOUSE_V2_PROVIDER_PRIORITY = [ - LIGHTHOUSE_V2_PROVIDER_TYPE.OPENAI, - LIGHTHOUSE_V2_PROVIDER_TYPE.BEDROCK, - LIGHTHOUSE_V2_PROVIDER_TYPE.OPENAI_COMPATIBLE, -] as const; - -// Fallback model per provider when the configuration has no remembered model. -const LIGHTHOUSE_V2_PREFERRED_DEFAULT_MODEL: Partial< - Record -> = { - [LIGHTHOUSE_V2_PROVIDER_TYPE.OPENAI]: "gpt-5.5", -}; - -function resolveInitialModelSelection( - connectedConfigurations: LighthouseV2Configuration[], - modelsByProvider: Record< - LighthouseV2ProviderType, - LighthouseV2SupportedModel[] - >, -): LighthouseV2ModelSelection | null { - const priorityIndex = (providerType: LighthouseV2ProviderType) => { - const index = LIGHTHOUSE_V2_PROVIDER_PRIORITY.indexOf(providerType); - return index === -1 ? LIGHTHOUSE_V2_PROVIDER_PRIORITY.length : index; - }; - // Stable sort keeps providers outside the priority list in their original order. - const orderedConfigurations = [...connectedConfigurations].sort( - (a, b) => priorityIndex(a.providerType) - priorityIndex(b.providerType), - ); - - for (const configuration of orderedConfigurations) { - const providerModels = modelsByProvider[configuration.providerType] ?? []; - if (providerModels.length === 0) continue; - // Prefer the provider's remembered model when it is still supported, then - // the provider's preferred default, then the first supported model. - const rememberedModel = providerModels.find( - (model) => model.id === configuration.defaultModel, - ); - const preferredModel = providerModels.find( - (model) => - model.id === - LIGHTHOUSE_V2_PREFERRED_DEFAULT_MODEL[configuration.providerType], - ); - return { - providerType: configuration.providerType, - modelId: (rememberedModel ?? preferredModel ?? providerModels[0]).id, - }; - } - - return null; -} - -function buildModelSelectorGroups( - connectedConfigurations: LighthouseV2Configuration[], - modelsByProvider: Record< - LighthouseV2ProviderType, - LighthouseV2SupportedModel[] - >, - supportedProviders: LighthouseV2SupportedProvider[], -): ComboboxGroup[] { - const groups: ComboboxGroup[] = []; - - for (const provider of supportedProviders) { - const configuration = connectedConfigurations.find( - (item) => item.providerType === provider.id, - ); - if (!configuration) continue; - - const options = (modelsByProvider[configuration.providerType] ?? []).map( - (model) => ({ - value: buildLighthouseV2ModelSelectionValue( - configuration.providerType, - model.id, - ), - label: model.name, - }), - ); - - if (options.length === 0) continue; - - groups.push({ - heading: provider.name, - options, - }); - } - - return groups; -} - -function isOnlyConnectedProvider( - connectedConfigurations: LighthouseV2Configuration[], - providerType: LighthouseV2ProviderType, -) { - return ( - connectedConfigurations.length === 1 && - connectedConfigurations[0]?.providerType === providerType - ); -} - -function getModelSelectionLabel( - selection: LighthouseV2ModelSelection, - modelsByProvider: Record< - LighthouseV2ProviderType, - LighthouseV2SupportedModel[] - >, -) { - return ( - modelsByProvider[selection.providerType]?.find( - (model) => model.id === selection.modelId, - )?.name ?? selection.modelId - ); -} - -function getProviderDisplayName( - providerType: LighthouseV2ProviderType, - supportedProviders: LighthouseV2SupportedProvider[], -) { - return ( - supportedProviders.find((provider) => provider.id === providerType)?.name ?? - providerType + + + ); } diff --git a/ui/app/(prowler)/lighthouse/_components/chat/lighthouse-v2-chat-view.tsx b/ui/app/(prowler)/lighthouse/_components/chat/lighthouse-v2-chat-view.tsx new file mode 100644 index 0000000000..0060dbe44e --- /dev/null +++ b/ui/app/(prowler)/lighthouse/_components/chat/lighthouse-v2-chat-view.tsx @@ -0,0 +1,339 @@ +"use client"; + +import { type ReactNode, type SubmitEvent } from "react"; + +import { + Conversation, + ConversationContent, + ConversationScrollButton, +} from "@/app/(prowler)/lighthouse/_components/ai-elements/conversation"; +import { selectLighthouseChatCanSend } from "@/app/(prowler)/lighthouse/_lib/chat-store"; +import { LIGHTHOUSE_V2_STREAM_STATUS } from "@/app/(prowler)/lighthouse/_lib/event-reducer"; +import { + buildLighthouseV2ModelSelectionValue, + type LighthouseV2ModelSelection, + parseLighthouseV2ModelSelectionValue, +} from "@/app/(prowler)/lighthouse/_lib/model-selection"; +import { + LIGHTHOUSE_V2_PROVIDER_TYPE, + type LighthouseV2Configuration, + type LighthouseV2ProviderType, + type LighthouseV2SupportedModel, + type LighthouseV2SupportedProvider, +} from "@/app/(prowler)/lighthouse/_types"; +import { Card } from "@/components/shadcn"; +import { + Combobox, + type ComboboxGroup, +} from "@/components/shadcn/combobox/combobox"; +import { Skeleton } from "@/components/shadcn/skeleton/skeleton"; + +import { ProviderIcon } from "../config/provider-icon"; +import { ChatComposerPanel } from "./composer"; +import { ChatEmptyState } from "./empty-state"; +import { useLighthouseChatStore } from "./lighthouse-chat-store-provider"; +import { MessageBubble } from "./message-bubble"; +import { StreamingAssistantMessage } from "./streaming-message"; + +export const LIGHTHOUSE_CHAT_SURFACE = { + PAGE: "page", + PANEL: "panel", +} as const; + +export type LighthouseChatSurface = + (typeof LIGHTHOUSE_CHAT_SURFACE)[keyof typeof LIGHTHOUSE_CHAT_SURFACE]; + +interface LighthouseV2ChatViewProps { + surface: LighthouseChatSurface; + emptyStateFooter?: ReactNode; +} + +export function LighthouseV2ChatView({ + surface, + emptyStateFooter, +}: LighthouseV2ChatViewProps) { + // Whole-store subscription is intentional: the view renders most of the state and selectLighthouseChatCanSend takes full state. + const state = useLighthouseChatStore((current) => current); + const { + config, + messages, + streamState, + input, + feedback, + isLoadingSession, + lastSubmittedText, + selectedModelSelection, + modelPreferenceSaving, + setInput, + dismissFeedback, + selectModel, + submitMessage, + } = state; + const { modelsByProvider, supportedProviders } = config; + const connectedConfigurations = config.configurations.filter( + (configuration) => configuration.connected === true, + ); + + const selectedConfiguration = selectedModelSelection + ? connectedConfigurations.find( + (configuration) => + configuration.providerType === selectedModelSelection.providerType, + ) + : undefined; + const modelSelectorGroups = buildModelSelectorGroups( + connectedConfigurations, + modelsByProvider, + supportedProviders, + ); + const showStaticOpenAIModel = + isOnlyConnectedProvider( + connectedConfigurations, + LIGHTHOUSE_V2_PROVIDER_TYPE.OPENAI, + ) && + selectedModelSelection?.providerType === LIGHTHOUSE_V2_PROVIDER_TYPE.OPENAI; + const selectedModelLabel = selectedModelSelection + ? getModelSelectionLabel(selectedModelSelection, modelsByProvider) + : "No model selected"; + const selectedProviderName = selectedModelSelection + ? getProviderDisplayName( + selectedModelSelection.providerType, + supportedProviders, + ) + : "OpenAI"; + const selectedModelValue = selectedModelSelection + ? buildLighthouseV2ModelSelectionValue( + selectedModelSelection.providerType, + selectedModelSelection.modelId, + ) + : ""; + + const canSend = selectLighthouseChatCanSend(state); + + const handleModelValueChange = (value: string) => { + const selection = parseLighthouseV2ModelSelectionValue(value); + if (!selection) return; + void selectModel(selection); + }; + + const handleSubmit = (event: SubmitEvent) => { + event.preventDefault(); + void submitMessage(input); + }; + + const hasLiveAssistantActivity = + Boolean(streamState.activeTaskId) || + Boolean(streamState.assistantText) || + streamState.toolCalls.length > 0; + const hasConversation = messages.length > 0 || hasLiveAssistantActivity; + + const composerPanelProps = { + feedback, + canRetry: + streamState.status === LIGHTHOUSE_V2_STREAM_STATUS.DISCONNECTED && + lastSubmittedText !== null, + onRetry: () => + lastSubmittedText ? void submitMessage(lastSubmittedText) : undefined, + onDismissFeedback: dismissFeedback, + canSend, + input, + isStreaming: Boolean(streamState.activeTaskId), + modelSelector: showStaticOpenAIModel ? ( + + ) : ( +
+ +
+ ), + selectedConfigurationConnected: selectedConfiguration?.connected === true, + onInputChange: setInput, + onSubmit: handleSubmit, + onSubmitText: submitMessage, + }; + + const chatBody = isLoadingSession ? ( + + ) : hasConversation ? ( +
+
+ + + {messages.map((message) => ( + + ))} + {hasLiveAssistantActivity && ( + + )} + + + +
+
+
+
+ +
+
+
+ ) : ( + + ); + + if (surface === LIGHTHOUSE_CHAT_SURFACE.PAGE) { + return ( + + {chatBody} + + ); + } + + return ( +
+ {chatBody} +
+ ); +} + +function SessionLoadingState() { + return ( +
+ + + + +
+ ); +} + +interface CurrentModelDisplayProps { + provider: LighthouseV2ProviderType; + providerName: string; + modelName: string; +} + +function CurrentModelDisplay({ + provider, + providerName, + modelName, +}: CurrentModelDisplayProps) { + return ( +
+ + + {providerName} + + + {modelName} + +
+ ); +} + +function buildModelSelectorGroups( + connectedConfigurations: LighthouseV2Configuration[], + modelsByProvider: Record< + LighthouseV2ProviderType, + LighthouseV2SupportedModel[] + >, + supportedProviders: LighthouseV2SupportedProvider[], +): ComboboxGroup[] { + const groups: ComboboxGroup[] = []; + + for (const provider of supportedProviders) { + const configuration = connectedConfigurations.find( + (item) => item.providerType === provider.id, + ); + if (!configuration) continue; + + const options = (modelsByProvider[configuration.providerType] ?? []).map( + (model) => ({ + value: buildLighthouseV2ModelSelectionValue( + configuration.providerType, + model.id, + ), + label: model.name, + }), + ); + + if (options.length === 0) continue; + + groups.push({ + heading: provider.name, + options, + }); + } + + return groups; +} + +function isOnlyConnectedProvider( + connectedConfigurations: LighthouseV2Configuration[], + providerType: LighthouseV2ProviderType, +) { + return ( + connectedConfigurations.length === 1 && + connectedConfigurations[0]?.providerType === providerType + ); +} + +function getModelSelectionLabel( + selection: LighthouseV2ModelSelection, + modelsByProvider: Record< + LighthouseV2ProviderType, + LighthouseV2SupportedModel[] + >, +) { + return ( + modelsByProvider[selection.providerType]?.find( + (model) => model.id === selection.modelId, + )?.name ?? selection.modelId + ); +} + +function getProviderDisplayName( + providerType: LighthouseV2ProviderType, + supportedProviders: LighthouseV2SupportedProvider[], +) { + return ( + supportedProviders.find((provider) => provider.id === providerType)?.name ?? + providerType + ); +} diff --git a/ui/app/(prowler)/lighthouse/_components/chat/message-bubble.test.tsx b/ui/app/(prowler)/lighthouse/_components/chat/message-bubble.test.tsx index 8b6855392f..ab11429774 100644 --- a/ui/app/(prowler)/lighthouse/_components/chat/message-bubble.test.tsx +++ b/ui/app/(prowler)/lighthouse/_components/chat/message-bubble.test.tsx @@ -51,7 +51,7 @@ describe("MessageBubble", () => { // Given const orderedMessage = buildAssistantMessage([ textPart("part-1", "Voy a buscar los findings por severidad"), - toolCallPart("part-2", "prowler_app_search_security_findings"), + toolCallPart("part-2", "prowler_search_security_findings"), textPart("part-3", "Ahora voy a buscar en los criticos"), ]); diff --git a/ui/app/(prowler)/lighthouse/_components/config/configuration-form.tsx b/ui/app/(prowler)/lighthouse/_components/config/configuration-form.tsx index 5bb6780f33..6ce5bab8b2 100644 --- a/ui/app/(prowler)/lighthouse/_components/config/configuration-form.tsx +++ b/ui/app/(prowler)/lighthouse/_components/config/configuration-form.tsx @@ -24,6 +24,7 @@ import { trimToNullable, } from "@/app/(prowler)/lighthouse/_lib/config"; import { formatLastChecked } from "@/app/(prowler)/lighthouse/_lib/format"; +import { notifyLighthouseV2ConfigurationsChanged } from "@/app/(prowler)/lighthouse/_lib/session-events"; import { type LighthouseV2Configuration, type LighthouseV2ConfigurationUpdateInput, @@ -138,6 +139,7 @@ export function LighthouseV2ConfigurationForm({ form.reset(getFormDefaults(result.data)); onConfigurationSaved(result.data); + notifyLighthouseV2ConfigurationsChanged(); if (shouldTestAfterSave) { await runConnectionTest(result.data.id); } @@ -178,6 +180,7 @@ export function LighthouseV2ConfigurationForm({ setDeleteOpen(false); form.reset(EMPTY_FORM_VALUES); onConfigurationDeleted(configuration.id); + notifyLighthouseV2ConfigurationsChanged(); } catch { onFeedback({ title: "Configuration not removed", diff --git a/ui/app/(prowler)/lighthouse/_components/panel/lighthouse-panel-chat-skeleton.tsx b/ui/app/(prowler)/lighthouse/_components/panel/lighthouse-panel-chat-skeleton.tsx new file mode 100644 index 0000000000..4a75f6c1e9 --- /dev/null +++ b/ui/app/(prowler)/lighthouse/_components/panel/lighthouse-panel-chat-skeleton.tsx @@ -0,0 +1,71 @@ +import { Skeleton } from "@/components/shadcn/skeleton/skeleton"; +import { cn } from "@/lib/utils"; + +// 1:1 skeleton of the panel chat empty state (logo, headline, composer, +// suggestion chips, recent chats). Kept in its own file with light imports: +// the side-panel shell uses it as the Suspense fallback while the real +// (lazy) chat bundle downloads, so it must not pull that bundle in. +export function LighthousePanelChatSkeleton() { + return ( +
+ {/* Lighthouse logo */} + + + {/* Headline + subline */} +
+ + +
+ + {/* Composer: textarea, then model selector + send button row */} +
+ +
+ + +
+
+ + {/* "Try Lighthouse AI for..." suggestion chips */} +
+ +
+ + + + +
+
+ + {/* Recent chats: label, search + new-chat row, session rows */} +
+ +
+ + +
+
+ + + +
+
+
+ ); +} + +interface SessionRowSkeletonProps { + titleWidth: string; +} + +function SessionRowSkeleton({ titleWidth }: SessionRowSkeletonProps) { + return ( +
+ + +
+ ); +} diff --git a/ui/app/(prowler)/lighthouse/_components/panel/lighthouse-panel-chat.test.tsx b/ui/app/(prowler)/lighthouse/_components/panel/lighthouse-panel-chat.test.tsx new file mode 100644 index 0000000000..f903bd87e2 --- /dev/null +++ b/ui/app/(prowler)/lighthouse/_components/panel/lighthouse-panel-chat.test.tsx @@ -0,0 +1,464 @@ +import { act, render, screen, waitFor, within } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { type ReactNode } from "react"; +import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; + +import { resetPanelChatStoreForTests } from "@/app/(prowler)/lighthouse/_lib/panel-chat-store"; +import { notifyLighthouseV2ConfigurationsChanged } from "@/app/(prowler)/lighthouse/_lib/session-events"; +import { stubEventSource } from "@/app/(prowler)/lighthouse/_lib/testing/event-source-mock"; +import type { + LighthouseV2Configuration, + LighthouseV2Session, + LighthouseV2SupportedModel, +} from "@/app/(prowler)/lighthouse/_types"; + +import { + LighthousePanelChat, + resetPanelChatConfigCacheForTests, +} from "./lighthouse-panel-chat"; +import { LighthousePanelHeaderActions } from "./lighthouse-panel-header-actions"; + +const { + getConfigurationsMock, + getSupportedProvidersMock, + getSupportedModelsMock, + getSessionsMock, + archiveSessionMock, + createSessionMock, + getMessagesMock, + sendMessageMock, + updateConfigurationMock, +} = vi.hoisted(() => ({ + getConfigurationsMock: vi.fn(), + getSupportedProvidersMock: vi.fn(), + getSupportedModelsMock: vi.fn(), + getSessionsMock: vi.fn(), + archiveSessionMock: vi.fn(), + createSessionMock: vi.fn(), + getMessagesMock: vi.fn(), + sendMessageMock: vi.fn(), + updateConfigurationMock: vi.fn(), +})); + +vi.mock("@/app/(prowler)/lighthouse/_actions", () => ({ + getLighthouseV2Configurations: getConfigurationsMock, + getLighthouseV2SupportedProviders: getSupportedProvidersMock, + getLighthouseV2SupportedModels: getSupportedModelsMock, + getLighthouseV2Sessions: getSessionsMock, + archiveLighthouseV2Session: archiveSessionMock, + createLighthouseV2Session: createSessionMock, + getLighthouseV2Messages: getMessagesMock, + sendLighthouseV2Message: sendMessageMock, + updateLighthouseV2Configuration: updateConfigurationMock, +})); + +// Streamdown pulls in shiki/wasm syntax highlighting that doesn't run under +// jsdom; render its text passthrough so message bodies are still assertable. +vi.mock("streamdown", () => ({ + Streamdown: ({ children }: { children: ReactNode }) => <>{children}, + defaultRehypePlugins: { katex: undefined, harden: undefined }, +})); + +const configurations: LighthouseV2Configuration[] = [ + { + id: "config-openai", + providerType: "openai", + baseUrl: null, + defaultModel: "gpt-5.1", + businessContext: "Production account", + connected: true, + connectionLastCheckedAt: "2026-06-24T10:00:00Z", + insertedAt: "2026-06-24T09:00:00Z", + updatedAt: "2026-06-24T10:00:00Z", + }, +]; + +describe("LighthousePanelChat", () => { + beforeEach(() => { + vi.stubGlobal( + "ResizeObserver", + class ResizeObserver { + observe = vi.fn(); + unobserve = vi.fn(); + disconnect = vi.fn(); + }, + ); + Object.defineProperty(Element.prototype, "scrollIntoView", { + configurable: true, + value: vi.fn(), + }); + getConfigurationsMock.mockReset(); + getSupportedProvidersMock.mockReset(); + getSupportedModelsMock.mockReset(); + getSessionsMock.mockReset(); + archiveSessionMock.mockReset(); + getMessagesMock.mockReset(); + stubEventSource(); + resetPanelChatStoreForTests(); + resetPanelChatConfigCacheForTests(); + + getConfigurationsMock.mockResolvedValue({ data: configurations }); + getSupportedProvidersMock.mockResolvedValue({ + data: [ + { id: "openai", name: "OpenAI" }, + { id: "bedrock", name: "AWS Bedrock" }, + { id: "openai-compatible", name: "OpenAI Compatible" }, + ], + }); + getSupportedModelsMock.mockResolvedValue({ data: [model("gpt-5.1")] }); + getSessionsMock.mockResolvedValue({ data: [] }); + getMessagesMock.mockResolvedValue({ data: [] }); + }); + + afterEach(() => { + vi.unstubAllGlobals(); + }); + + it("shows a loading skeleton while the config loads", () => { + // Given: a config fetch that never resolves within the assertion window + getConfigurationsMock.mockReturnValue(new Promise(() => {})); + + // When + render(); + + // Then + expect(screen.getByLabelText("Loading Lighthouse AI")).toBeInTheDocument(); + }); + + it("shows the error state with a Retry that reloads the config", async () => { + // Given + getConfigurationsMock.mockResolvedValueOnce({ + error: "Something went wrong.", + status: 500, + }); + const user = userEvent.setup(); + render(); + expect(await screen.findByRole("alert")).toHaveTextContent( + "Something went wrong.", + ); + + // When: retrying after the backend recovers + await user.click(screen.getByRole("button", { name: "Retry" })); + + // Then + expect( + await screen.findByRole("textbox", { name: "Message" }), + ).toBeInTheDocument(); + }); + + it("shows an in-panel connect CTA instead of redirecting when no LLM is connected", async () => { + // Given + getConfigurationsMock.mockResolvedValue({ + data: [{ ...configurations[0], connected: false }], + }); + + // When + render(); + + // Then + expect( + await screen.findByRole("link", { name: "Connect an LLM provider" }), + ).toHaveAttribute("href", "/lighthouse/settings"); + }); + + it("renders the chat composer and recent chats once ready", async () => { + // Given + getSessionsMock.mockResolvedValue({ + data: [session("session-1", "Counting critical findings")], + }); + + // When + render(); + + // Then: composer is live and the empty state lists recent chats + expect( + await screen.findByRole("textbox", { name: "Message" }), + ).toBeInTheDocument(); + expect(screen.getByText("Recent chats")).toBeInTheDocument(); + expect( + await screen.findByText("Counting critical findings"), + ).toBeInTheDocument(); + }); + + it("opens a recent chat in place without navigating", async () => { + // Given + const user = userEvent.setup(); + const replaceStateSpy = vi.spyOn(window.history, "replaceState"); + getSessionsMock.mockResolvedValue({ + data: [session("session-1", "Counting critical findings")], + }); + getMessagesMock.mockResolvedValue({ + data: [ + { + id: "message-1", + role: "assistant", + model: null, + tokenUsage: null, + insertedAt: "2026-06-25T10:00:00Z", + parts: [ + { + id: "message-1-part", + type: "text", + content: "There are 3 critical findings.", + toolCallOutcome: null, + insertedAt: "2026-06-25T10:00:00Z", + updatedAt: "2026-06-25T10:00:00Z", + }, + ], + }, + ], + }); + render(); + + // When + await user.click( + await screen.findByRole("button", { + name: /^Counting critical findings/, + }), + ); + + // Then: the conversation loads in the panel and the URL never changes + expect( + await screen.findByText("There are 3 critical findings."), + ).toBeInTheDocument(); + expect(replaceStateSpy).not.toHaveBeenCalled(); + replaceStateSpy.mockRestore(); + }); + + it("explains why a new chat is unavailable before the first message", async () => { + // Given + const user = userEvent.setup(); + render(); + const newChatButton = screen.getByRole("button", { name: "New chat" }); + + // When + const disabledTrigger = newChatButton.parentElement; + + // Then + expect(newChatButton).toBeDisabled(); + expect(disabledTrigger).toHaveClass("cursor-not-allowed"); + await user.hover(disabledTrigger!); + expect(await screen.findByRole("tooltip")).toHaveTextContent( + "Send a message before starting a new chat", + ); + }); + + it("opens the active panel conversation on the full-page chat route", async () => { + // Given: the panel starts on a new chat and exposes the full-page action + const user = userEvent.setup(); + getSessionsMock.mockResolvedValue({ + data: [session("session-1", "Counting critical findings")], + }); + render( + <> + + + , + ); + const fullPageLink = await screen.findByRole("link", { + name: "Open Lighthouse AI full page", + }); + expect(fullPageLink).toHaveAttribute("href", "/lighthouse"); + + // When: an existing conversation becomes active in the panel + await user.click( + await screen.findByRole("button", { + name: /^Counting critical findings/, + }), + ); + + // Then: full-page navigation carries the active session in the URL + expect(fullPageLink).toHaveAttribute( + "href", + "/lighthouse?session=session-1", + ); + }); + + it("starts a new chat from the panel header", async () => { + // Given: an existing conversation is open in the panel + const user = userEvent.setup(); + getSessionsMock.mockResolvedValue({ + data: [session("session-1", "Counting critical findings")], + }); + getMessagesMock.mockResolvedValue({ + data: [ + { + id: "message-1", + role: "assistant", + model: null, + tokenUsage: null, + insertedAt: "2026-06-25T10:00:00Z", + parts: [ + { + id: "message-1-part", + type: "text", + content: "There are 3 critical findings.", + toolCallOutcome: null, + insertedAt: "2026-06-25T10:00:00Z", + updatedAt: "2026-06-25T10:00:00Z", + }, + ], + }, + ], + }); + render( + <> +
+ +
+ + , + ); + await screen.findByRole("textbox", { name: "Message" }); + const panelHeader = screen.getByLabelText("Panel header actions"); + const newChatButton = within(panelHeader).getByRole("button", { + name: "New chat", + }); + expect(newChatButton).toBeDisabled(); + + await user.click( + await screen.findByRole("button", { + name: /^Counting critical findings/, + }), + ); + expect( + await screen.findByText("There are 3 critical findings."), + ).toBeInTheDocument(); + expect(newChatButton).toBeEnabled(); + + // When + await user.click(newChatButton); + + // Then + expect( + screen.queryByText("There are 3 critical findings."), + ).not.toBeInTheDocument(); + expect(screen.getByRole("textbox", { name: "Message" })).toHaveValue(""); + expect(newChatButton).toBeDisabled(); + }); + + it("caches the loaded config so a remount skips the skeleton", async () => { + // Given: a first mount that loads successfully + const { unmount } = render(); + await screen.findByRole("textbox", { name: "Message" }); + unmount(); + getConfigurationsMock.mockClear(); + + // When + render(); + + // Then: ready immediately, no refetch + await waitFor(() => + expect( + screen.getByRole("textbox", { name: "Message" }), + ).toBeInTheDocument(), + ); + expect(getConfigurationsMock).not.toHaveBeenCalled(); + }); + + it("reloads models after a transient model-loading failure", async () => { + // Given: configuration loads, but the first model request fails + getSupportedModelsMock.mockResolvedValueOnce({ + error: "Models are temporarily unavailable.", + status: 500, + }); + const { unmount } = render(); + await screen.findByRole("textbox", { name: "Message" }); + unmount(); + getSupportedModelsMock.mockClear(); + + // When: the panel reopens after the model endpoint recovers + render(); + + // Then: the partial ready state is not reused as a successful cache entry + await waitFor(() => expect(getSupportedModelsMock).toHaveBeenCalled()); + }); + + it("removes an archived chat from the recent chats list", async () => { + // Given: one recent chat + const user = userEvent.setup(); + getSessionsMock.mockResolvedValue({ + data: [session("session-1", "Counting critical findings")], + }); + archiveSessionMock.mockResolvedValue({ data: { id: "session-1" } }); + render(); + await screen.findByText("Counting critical findings"); + + // When: archiving it from the panel (hover action + confirm dialog) + getSessionsMock.mockResolvedValue({ data: [] }); + await user.click( + screen.getByRole("button", { + name: "Archive Counting critical findings", + }), + ); + await user.click( + within(await screen.findByRole("dialog")).getByRole("button", { + name: "Archive", + }), + ); + + // Then: the archived chat leaves the list + await waitFor(() => + expect( + screen.queryByText("Counting critical findings"), + ).not.toBeInTheDocument(), + ); + }); + + it("swaps the connect CTA for the chat once a provider is connected", async () => { + // Given: no connected provider yet + getConfigurationsMock.mockResolvedValueOnce({ + data: [{ ...configurations[0], connected: false }], + }); + render(); + await screen.findByRole("link", { name: "Connect an LLM provider" }); + + // When: a provider gets connected on the settings page + act(() => notifyLighthouseV2ConfigurationsChanged()); + + // Then: the panel reloads into the live chat + expect( + await screen.findByRole("textbox", { name: "Message" }), + ).toBeInTheDocument(); + }); + + it("drops the cached config when configurations change while unmounted", async () => { + // Given: a cached config from a previous mount + const { unmount } = render(); + await screen.findByRole("textbox", { name: "Message" }); + unmount(); + getConfigurationsMock.mockClear(); + + // When: config CRUD happens with the panel closed, then it reopens + notifyLighthouseV2ConfigurationsChanged(); + render(); + + // Then: the stale cache is gone and the config reloads + expect( + await screen.findByRole("textbox", { name: "Message" }), + ).toBeInTheDocument(); + expect(getConfigurationsMock).toHaveBeenCalled(); + }); +}); + +function model(id: string, name = id): LighthouseV2SupportedModel { + return { + id, + name, + maxInputTokens: null, + maxOutputTokens: null, + supportsFunctionCalling: null, + supportsVision: null, + supportsReasoning: null, + }; +} + +function session(id: string, title: string): LighthouseV2Session { + return { + id, + title, + isArchived: false, + insertedAt: "2026-06-24T10:00:00Z", + updatedAt: "2026-06-24T10:00:00Z", + }; +} diff --git a/ui/app/(prowler)/lighthouse/_components/panel/lighthouse-panel-chat.tsx b/ui/app/(prowler)/lighthouse/_components/panel/lighthouse-panel-chat.tsx new file mode 100644 index 0000000000..acecf64b32 --- /dev/null +++ b/ui/app/(prowler)/lighthouse/_components/panel/lighthouse-panel-chat.tsx @@ -0,0 +1,340 @@ +"use client"; + +import Link from "next/link"; +import { useState } from "react"; + +import { + archiveLighthouseV2Session, + getLighthouseV2Sessions, +} from "@/app/(prowler)/lighthouse/_actions"; +import { LighthouseV2SessionHistory } from "@/app/(prowler)/lighthouse/_components/history"; +import type { LighthouseChatConfig } from "@/app/(prowler)/lighthouse/_lib/chat-store"; +import { + LIGHTHOUSE_CHAT_CONFIG_STATUS, + loadLighthouseChatConfig, +} from "@/app/(prowler)/lighthouse/_lib/load-chat-config"; +import { + resetPanelChatMessageState, + setPanelChatMessageState, +} from "@/app/(prowler)/lighthouse/_lib/panel-chat-message-state"; +import { + getOrCreatePanelChatStore, + resetPanelChatStore, +} from "@/app/(prowler)/lighthouse/_lib/panel-chat-store"; +import { + notifyLighthouseV2SessionArchived, + onLighthouseV2ConfigurationsChanged, + onLighthouseV2NewChat, + onLighthouseV2SessionArchived, + onLighthouseV2SessionsChanged, +} from "@/app/(prowler)/lighthouse/_lib/session-events"; +import type { LighthouseV2Session } from "@/app/(prowler)/lighthouse/_types"; +import { LighthouseIconWithAura } from "@/components/icons"; +import { Button } from "@/components/shadcn/button/button"; +import { useMountEffect } from "@/hooks/use-mount-effect"; +import { LIGHTHOUSE_ROUTE } from "@/lib/lighthouse-routes"; + +import { + LighthouseChatStoreProvider, + useLighthouseChatStore, +} from "../chat/lighthouse-chat-store-provider"; +import { + LIGHTHOUSE_CHAT_SURFACE, + LighthouseV2ChatView, +} from "../chat/lighthouse-v2-chat-view"; +import { LighthousePanelChatSkeleton } from "./lighthouse-panel-chat-skeleton"; + +const PANEL_CHAT_STATUS = { + LOADING: "loading", + ERROR: "error", + NOT_CONFIGURED: "not-configured", + READY: "ready", +} as const; + +interface PanelChatLoadingState { + status: typeof PANEL_CHAT_STATUS.LOADING; +} + +interface PanelChatErrorState { + status: typeof PANEL_CHAT_STATUS.ERROR; + message: string; +} + +interface PanelChatNotConfiguredState { + status: typeof PANEL_CHAT_STATUS.NOT_CONFIGURED; +} + +interface PanelChatReadyState { + status: typeof PANEL_CHAT_STATUS.READY; + config: LighthouseChatConfig; + modelsError?: string; +} + +type PanelChatState = + | PanelChatLoadingState + | PanelChatErrorState + | PanelChatNotConfiguredState + | PanelChatReadyState; + +// Config cache: the panel loads its configs/models lazily on first open (never +// in the layout, so pages don't pay for a panel most sessions never open); the +// cache makes every later mount — reopen, drawer AI tab — instant. +let cachedReadyState: PanelChatReadyState | null = null; + +export function resetPanelChatConfigCacheForTests(): void { + cachedReadyState = null; + resetPanelChatMessageState(); +} + +// Config CRUD happens on the settings route while the global panel can remain +// mounted. Invalidate at module scope so the open panel reloads in place and a +// later open rebuilds cache and store against the new configuration. +if (typeof window !== "undefined") { + onLighthouseV2ConfigurationsChanged(() => { + cachedReadyState = null; + resetPanelChatMessageState(); + resetPanelChatStore(); + }); +} + +export function LighthousePanelChat() { + const [state, setState] = useState( + () => cachedReadyState ?? { status: PANEL_CHAT_STATUS.LOADING }, + ); + + const load = async () => { + setState({ status: PANEL_CHAT_STATUS.LOADING }); + const next = await loadPanelChatState(); + if ( + next.status === PANEL_CHAT_STATUS.READY && + next.modelsError === undefined + ) { + cachedReadyState = next; + } else { + cachedReadyState = null; + } + setState(next); + }; + + useMountEffect(() => { + if (state.status !== PANEL_CHAT_STATUS.READY) { + void load(); + } + // The module-scope listener above already invalidated cache and store + // (registration order); reload so an open panel refreshes in place. + return onLighthouseV2ConfigurationsChanged(() => void load()); + }); + + if (state.status === PANEL_CHAT_STATUS.LOADING) { + return ; + } + if (state.status === PANEL_CHAT_STATUS.ERROR) { + return ( + void load()} /> + ); + } + if (state.status === PANEL_CHAT_STATUS.NOT_CONFIGURED) { + return ; + } + return ( + + ); +} + +interface PanelChatReadyProps { + config: LighthouseChatConfig; + modelsError?: string; +} + +function PanelChatReady({ config, modelsError }: PanelChatReadyProps) { + const [store] = useState(() => + getOrCreatePanelChatStore(config, { initialError: modelsError }), + ); + const [sessions, setSessions] = useState([]); + + const refreshSessions = async () => { + try { + const result = await getLighthouseV2Sessions(); + if ("data" in result) { + setSessions(result.data); + } + } catch { + // Best-effort refresh: swallow transport-level failures so a rejected + // server action never escapes the mount effect as an unhandled error. + } + }; + + useMountEffect(() => { + void refreshSessions(); + const syncPanelChatState = () => { + const chatState = store.getState(); + setPanelChatMessageState({ + hasMessages: chatState.messages.length > 0, + activeSessionId: chatState.activeSessionId, + }); + }; + syncPanelChatState(); + const unsubscribeChatStore = store.subscribe(syncPanelChatState); + const unsubscribeSessionsChanged = onLighthouseV2SessionsChanged(() => { + void refreshSessions(); + }); + const unsubscribeNewChat = onLighthouseV2NewChat(() => + store.getState().resetToNewChat(), + ); + // Archiving from any surface (sidebar, popover) must reset the panel chat + // when its open session is the archived one, and drop the archived chat + // from the "Recent chats" list. + const unsubscribeSessionArchived = onLighthouseV2SessionArchived( + (sessionId) => { + store.getState().handleSessionArchived(sessionId); + void refreshSessions(); + }, + ); + return () => { + unsubscribeChatStore(); + unsubscribeSessionsChanged(); + unsubscribeNewChat(); + unsubscribeSessionArchived(); + // A partial config cannot be reused after the model endpoint recovers: + // this store captured the incomplete model list at creation time. + if (modelsError) resetPanelChatStore(); + }; + }); + + return ( + +
+
+ 0 ? ( +
+ + Recent chats + + +
+ ) : undefined + } + /> +
+
+
+ ); +} + +interface PanelChatSessionsProps { + sessions: LighthouseV2Session[]; + onAfterSelect?: () => void; +} + +function PanelChatSessions({ + sessions, + onAfterSelect, +}: PanelChatSessionsProps) { + const [search, setSearch] = useState(""); + const activeSessionId = useLighthouseChatStore( + (state) => state.activeSessionId, + ); + const isOnNewChat = useLighthouseChatStore( + (state) => state.activeSessionId === null && state.messages.length === 0, + ); + const openSession = useLighthouseChatStore((state) => state.openSession); + const resetToNewChat = useLighthouseChatStore( + (state) => state.resetToNewChat, + ); + + const handleArchiveSession = async (sessionId: string) => { + try { + const result = await archiveLighthouseV2Session(sessionId); + if ("data" in result) { + // Resets this chat when its open session is archived, and prompts + // every session list (sidebar included) to refresh. + notifyLighthouseV2SessionArchived(sessionId); + } + } catch { + // Archiving is recoverable from the list; ignore transient failures. + } + }; + + return ( + { + resetToNewChat(); + onAfterSelect?.(); + }} + onOpenSession={(sessionId) => { + void openSession(sessionId); + onAfterSelect?.(); + }} + onArchiveSession={(sessionId) => void handleArchiveSession(sessionId)} + /> + ); +} + +interface PanelChatErrorProps { + message: string; + onRetry: () => void; +} + +function PanelChatError({ message, onRetry }: PanelChatErrorProps) { + return ( +
+

+ {message} +

+ +
+ ); +} + +function PanelChatConnectCta() { + return ( +
+ +
+

+ Lighthouse AI is not set up yet +

+

+ Connect an LLM provider to start asking questions about your cloud + security posture. +

+
+ +
+ ); +} + +async function loadPanelChatState(): Promise { + try { + const result = await loadLighthouseChatConfig(); + if (result.status === LIGHTHOUSE_CHAT_CONFIG_STATUS.ERROR) { + return { status: PANEL_CHAT_STATUS.ERROR, message: result.message }; + } + if (result.status === LIGHTHOUSE_CHAT_CONFIG_STATUS.NOT_CONFIGURED) { + return { status: PANEL_CHAT_STATUS.NOT_CONFIGURED }; + } + return { + status: PANEL_CHAT_STATUS.READY, + config: result.config, + modelsError: result.modelsError, + }; + } catch { + return { + status: PANEL_CHAT_STATUS.ERROR, + message: "Could not load Lighthouse AI. Try again shortly.", + }; + } +} diff --git a/ui/app/(prowler)/lighthouse/_components/panel/lighthouse-panel-header-actions.tsx b/ui/app/(prowler)/lighthouse/_components/panel/lighthouse-panel-header-actions.tsx new file mode 100644 index 0000000000..9f71449626 --- /dev/null +++ b/ui/app/(prowler)/lighthouse/_components/panel/lighthouse-panel-header-actions.tsx @@ -0,0 +1,74 @@ +"use client"; + +import { Maximize2, Plus } from "lucide-react"; +import Link from "next/link"; +import { useSyncExternalStore } from "react"; + +import { + getPanelChatActiveSessionId, + getPanelChatHasMessages, + subscribePanelChatHasMessages, +} from "@/app/(prowler)/lighthouse/_lib/panel-chat-message-state"; +import { notifyLighthouseV2NewChat } from "@/app/(prowler)/lighthouse/_lib/session-events"; +import { Button } from "@/components/shadcn/button/button"; +import { + Tooltip, + TooltipContent, + TooltipTrigger, +} from "@/components/shadcn/tooltip"; +import { LIGHTHOUSE_ROUTE } from "@/lib/lighthouse-routes"; +import { cn } from "@/lib/utils"; + +export function LighthousePanelHeaderActions() { + const hasMessages = useSyncExternalStore( + subscribePanelChatHasMessages, + getPanelChatHasMessages, + () => false, + ); + const activeSessionId = useSyncExternalStore( + subscribePanelChatHasMessages, + getPanelChatActiveSessionId, + () => null, + ); + const fullPageHref = activeSessionId + ? `${LIGHTHOUSE_ROUTE.CHAT}?session=${encodeURIComponent(activeSessionId)}` + : LIGHTHOUSE_ROUTE.CHAT; + + return ( + <> + + + + + + + + {hasMessages + ? "New chat" + : "Send a message before starting a new chat"} + + + + + + + Open full page + + + ); +} diff --git a/ui/app/(prowler)/lighthouse/_lib/chat-store.test.ts b/ui/app/(prowler)/lighthouse/_lib/chat-store.test.ts new file mode 100644 index 0000000000..16fc3c3e8e --- /dev/null +++ b/ui/app/(prowler)/lighthouse/_lib/chat-store.test.ts @@ -0,0 +1,503 @@ +import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; + +import { + createLighthouseChatStore, + selectLighthouseChatCanSend, +} from "@/app/(prowler)/lighthouse/_lib/chat-store"; +import { + type MockEventSource, + stubEventSource, +} from "@/app/(prowler)/lighthouse/_lib/testing/event-source-mock"; +import type { + LighthouseV2Configuration, + LighthouseV2Message, + LighthouseV2SupportedModel, + LighthouseV2SupportedProvider, +} from "@/app/(prowler)/lighthouse/_types"; + +const { + createSessionMock, + getMessagesMock, + sendMessageMock, + updateConfigurationMock, +} = vi.hoisted(() => ({ + createSessionMock: vi.fn(), + getMessagesMock: vi.fn(), + sendMessageMock: vi.fn(), + updateConfigurationMock: vi.fn(), +})); + +vi.mock("@/app/(prowler)/lighthouse/_actions", () => ({ + createLighthouseV2Session: createSessionMock, + getLighthouseV2Messages: getMessagesMock, + sendLighthouseV2Message: sendMessageMock, + updateLighthouseV2Configuration: updateConfigurationMock, +})); + +const configurations: LighthouseV2Configuration[] = [ + { + id: "config-openai", + providerType: "openai", + baseUrl: null, + defaultModel: "gpt-5.1", + businessContext: "Production account", + connected: true, + connectionLastCheckedAt: "2026-06-24T10:00:00Z", + insertedAt: "2026-06-24T09:00:00Z", + updatedAt: "2026-06-24T10:00:00Z", + }, +]; + +const modelsByProvider = { + openai: [model("gpt-5.1")], + bedrock: [], + "openai-compatible": [], +}; + +const supportedProviders: LighthouseV2SupportedProvider[] = [ + { id: "openai", name: "OpenAI" }, + { id: "bedrock", name: "AWS Bedrock" }, + { id: "openai-compatible", name: "OpenAI Compatible" }, +]; + +const config = { configurations, modelsByProvider, supportedProviders }; + +let eventSources: MockEventSource[] = []; + +describe("createLighthouseChatStore", () => { + beforeEach(() => { + createSessionMock.mockReset(); + getMessagesMock.mockReset(); + sendMessageMock.mockReset(); + updateConfigurationMock.mockReset(); + eventSources = stubEventSource(); + + createSessionMock.mockResolvedValue({ + data: { + id: "session-1", + title: "Summarize findings", + isArchived: false, + insertedAt: "2026-06-24T10:00:00Z", + updatedAt: "2026-06-24T10:00:00Z", + }, + }); + getMessagesMock.mockResolvedValue({ data: [] }); + sendMessageMock.mockResolvedValue({ + data: { + task: { id: "task-1", name: "lighthouse-run", state: "executing" }, + }, + }); + window.history.replaceState(null, "", "/"); + }); + + afterEach(() => { + vi.unstubAllGlobals(); + }); + + it("resolves the connected provider's remembered model on creation", () => { + // Given / When + const store = makeStore(); + + // Then + expect(store.getState().selectedModelSelection).toEqual({ + providerType: "openai", + modelId: "gpt-5.1", + }); + expect(selectLighthouseChatCanSend(store.getState())).toBe(true); + }); + + it("creates a session and subscribes to the stream before sending (no replay buffer)", async () => { + // Given + const store = makeStore(); + + // When + await store.getState().submitMessage("Summarize findings"); + + // Then + expect(createSessionMock).toHaveBeenCalledWith("Summarize findings"); + expect(EventSource).toHaveBeenCalledWith( + "/api/lighthouse/v2/sessions/session-1/event-stream", + ); + const eventSourceOrder = vi.mocked(EventSource).mock.invocationCallOrder[0]; + const sendOrder = sendMessageMock.mock.invocationCallOrder[0]; + expect(eventSourceOrder).toBeLessThan(sendOrder); + // The optimistic user message renders immediately and the task id is live. + expect(store.getState().messages.at(-1)?.parts[0]?.content).toEqual({ + text: "Summarize findings", + }); + expect(store.getState().streamState.activeTaskId).toBe("task-1"); + }); + + it("does not touch the URL when syncUrlToSession is off (panel surface)", async () => { + // Given + const store = makeStore({ syncUrlToSession: false }); + const replaceStateSpy = vi.spyOn(window.history, "replaceState"); + + // When + await store.getState().submitMessage("Summarize findings"); + + // Then + expect(store.getState().activeSessionId).toBe("session-1"); + expect(replaceStateSpy).not.toHaveBeenCalled(); + }); + + it("writes the session URL in place when syncUrlToSession is on (page surface)", async () => { + // Given + const store = makeStore({ syncUrlToSession: true }); + const replaceStateSpy = vi.spyOn(window.history, "replaceState"); + + // When + await store.getState().submitMessage("Summarize findings"); + + // Then + expect(replaceStateSpy).toHaveBeenCalledWith( + null, + "", + "/lighthouse?session=session-1", + ); + }); + + it("reloads persisted messages and closes the stream on message.end", async () => { + // Given + const store = makeStore(); + await store.getState().submitMessage("Summarize findings"); + getMessagesMock.mockResolvedValue({ + data: [message("message-1", "assistant", "Persisted answer")], + }); + + // When + eventSources[0].emit("message.end", { message_id: "message-1" }); + await vi.waitFor(() => + expect(getMessagesMock).toHaveBeenCalledWith("session-1"), + ); + + // Then + await vi.waitFor(() => + expect(store.getState().messages[0]?.parts[0]?.content).toBe( + "Persisted answer", + ), + ); + expect(eventSources[0].close).toHaveBeenCalled(); + expect(store.getState().streamState.activeTaskId).toBeNull(); + }); + + it("blocks sending and refreshes messages on a 409 conflict", async () => { + // Given + const store = makeStore(); + sendMessageMock.mockResolvedValue({ + error: "Another run is in progress.", + status: 409, + }); + + // When + await store.getState().submitMessage("Summarize findings"); + + // Then + expect(store.getState().blockedByConflict).toBe(true); + expect(store.getState().feedback).toBe("Another run is in progress."); + expect(getMessagesMock).toHaveBeenCalledWith("session-1"); + expect(eventSources[0].close).toHaveBeenCalled(); + expect(selectLighthouseChatCanSend(store.getState())).toBe(false); + }); + + it("reconciles the optimistic message when the send fails without a conflict", async () => { + // Given: the backend rejects the message with a plain failure + const store = makeStore(); + sendMessageMock.mockResolvedValue({ error: "Send failed.", status: 500 }); + + // When + await store.getState().submitMessage("Summarize findings"); + + // Then: feedback surfaces without blocking, and the optimistic user + // message is reconciled against the server (it was never persisted) + expect(store.getState().feedback).toBe("Send failed."); + expect(store.getState().blockedByConflict).toBe(false); + expect(getMessagesMock).toHaveBeenCalledWith("session-1"); + expect(store.getState().messages).toHaveLength(0); + }); + + it("drops a failed send once the chat points at another session", async () => { + // Given: a send still in flight + const store = makeStore(); + let resolveSend: (value: unknown) => void = () => {}; + sendMessageMock.mockReturnValueOnce( + new Promise((resolve) => { + resolveSend = resolve; + }), + ); + const submitting = store.getState().submitMessage("Summarize findings"); + await vi.waitFor(() => expect(sendMessageMock).toHaveBeenCalled()); + + // When: the user opens another session before the send fails + await store.getState().openSession("session-9"); + resolveSend({ error: "Send failed.", status: 500 }); + await submitting; + + // Then: the dead submission's failure never surfaces in the new session + expect(store.getState().activeSessionId).toBe("session-9"); + expect(store.getState().feedback).toBeNull(); + }); + + it("keeps a fast follow-up intact when the terminal refresh resolves late", async () => { + // Given: a completed run whose terminal message refresh is still in flight + const store = makeStore(); + await store.getState().submitMessage("First question"); + let resolveRefresh: (value: unknown) => void = () => {}; + getMessagesMock.mockReturnValueOnce( + new Promise((resolve) => { + resolveRefresh = resolve; + }), + ); + eventSources[0].emit("message.end", { message_id: "message-1" }); + await vi.waitFor(() => + expect(getMessagesMock).toHaveBeenCalledWith("session-1"), + ); + + // When: the user sends a follow-up before that refresh resolves + sendMessageMock.mockResolvedValue({ + data: { + task: { id: "task-2", name: "lighthouse-run", state: "executing" }, + }, + }); + await store.getState().submitMessage("Follow-up question"); + resolveRefresh({ data: [message("message-1", "assistant", "Answer")] }); + await new Promise((resolve) => setTimeout(resolve, 0)); + + // Then: the stale snapshot erases neither the new optimistic message nor + // the follow-up's task id + expect(store.getState().streamState.activeTaskId).toBe("task-2"); + expect(store.getState().messages.at(-1)?.parts[0]?.content).toEqual({ + text: "Follow-up question", + }); + }); + + it("abandons an in-flight submit after destroy", async () => { + // Given: destroy fires while the session create is still in flight + const store = makeStore({ syncUrlToSession: true }); + const replaceStateSpy = vi.spyOn(window.history, "replaceState"); + let resolveCreate: (value: unknown) => void = () => {}; + createSessionMock.mockReturnValueOnce( + new Promise((resolve) => { + resolveCreate = resolve; + }), + ); + const submitting = store.getState().submitMessage("Summarize findings"); + store.getState().destroy(); + + // When + resolveCreate({ + data: { + id: "session-1", + title: "Summarize findings", + isArchived: false, + insertedAt: "2026-06-24T10:00:00Z", + updatedAt: "2026-06-24T10:00:00Z", + }, + }); + await submitting; + + // Then: no URL rewrite on whatever page is now open, no orphan stream + expect(replaceStateSpy).not.toHaveBeenCalled(); + expect(eventSources).toHaveLength(0); + }); + + it("does not replace a session opened while a new session is being created", async () => { + // Given: creating the first session is still in flight + const store = makeStore(); + let resolveCreate: (value: unknown) => void = () => {}; + createSessionMock.mockReturnValueOnce( + new Promise((resolve) => { + resolveCreate = resolve; + }), + ); + const submitting = store.getState().submitMessage("Summarize findings"); + await vi.waitFor(() => expect(createSessionMock).toHaveBeenCalled()); + + // When: the user opens another conversation before creation resolves + await store.getState().openSession("session-9"); + resolveCreate({ + data: { + id: "session-1", + title: "Summarize findings", + isArchived: false, + insertedAt: "2026-06-24T10:00:00Z", + updatedAt: "2026-06-24T10:00:00Z", + }, + }); + await submitting; + + // Then: the stale creation cannot replace or submit into the open chat + expect(store.getState().activeSessionId).toBe("session-9"); + expect(sendMessageMock).not.toHaveBeenCalled(); + }); + + it("does not revive a session creation cancelled by a new-chat reset", async () => { + // Given: creating the first session is still in flight + const store = makeStore(); + let resolveCreate: (value: unknown) => void = () => {}; + createSessionMock.mockReturnValueOnce( + new Promise((resolve) => { + resolveCreate = resolve; + }), + ); + const submitting = store.getState().submitMessage("Summarize findings"); + await vi.waitFor(() => expect(createSessionMock).toHaveBeenCalled()); + + // When: the user resets to a new chat before creation resolves + store.getState().resetToNewChat(); + resolveCreate({ + data: { + id: "session-1", + title: "Summarize findings", + isArchived: false, + insertedAt: "2026-06-24T10:00:00Z", + updatedAt: "2026-06-24T10:00:00Z", + }, + }); + await submitting; + + // Then + expect(store.getState().activeSessionId).toBeNull(); + expect(sendMessageMock).not.toHaveBeenCalled(); + }); + + it("opens an existing session client-side without navigation", async () => { + // Given + const store = makeStore({ syncUrlToSession: false }); + const replaceStateSpy = vi.spyOn(window.history, "replaceState"); + getMessagesMock.mockResolvedValue({ + data: [message("message-1", "assistant", "Old answer")], + }); + + // When + await store.getState().openSession("session-9"); + + // Then + expect(store.getState().activeSessionId).toBe("session-9"); + expect(store.getState().messages[0]?.parts[0]?.content).toBe("Old answer"); + expect(replaceStateSpy).not.toHaveBeenCalled(); + }); + + it("drops a stale openSession result when the chat was reset meanwhile", async () => { + // Given: opening a session whose message fetch is still in flight + const store = makeStore(); + let resolveLoad: (value: unknown) => void = () => {}; + getMessagesMock.mockReturnValueOnce( + new Promise((resolve) => { + resolveLoad = resolve; + }), + ); + const opening = store.getState().openSession("session-9"); + + // When: the user starts a new chat before the fetch resolves + store.getState().resetToNewChat(); + resolveLoad({ data: [message("message-1", "assistant", "Stale answer")] }); + await opening; + + // Then: the stale messages never repopulate the reset chat + expect(store.getState().activeSessionId).toBeNull(); + expect(store.getState().messages).toHaveLength(0); + }); + + it("resets to a new chat and closes any open stream", async () => { + // Given + const store = makeStore(); + await store.getState().submitMessage("Summarize findings"); + expect(store.getState().activeSessionId).toBe("session-1"); + + // When + store.getState().resetToNewChat(); + + // Then + expect(eventSources[0].close).toHaveBeenCalled(); + expect(store.getState().activeSessionId).toBeNull(); + expect(store.getState().messages).toHaveLength(0); + expect(store.getState().streamState.activeTaskId).toBeNull(); + }); + + it("resets only when the archived session is the active one", async () => { + // Given + const store = makeStore(); + await store.getState().submitMessage("Summarize findings"); + + // When / Then: an unrelated session leaves the conversation intact + store.getState().handleSessionArchived("session-other"); + expect(store.getState().activeSessionId).toBe("session-1"); + + // When / Then: archiving the active session resets in place + store.getState().handleSessionArchived("session-1"); + expect(store.getState().activeSessionId).toBeNull(); + }); + + it("closes the stream on destroy", async () => { + // Given + const store = makeStore(); + await store.getState().submitMessage("Summarize findings"); + + // When + store.getState().destroy(); + + // Then + expect(eventSources[0].close).toHaveBeenCalled(); + }); + + it("surfaces a connection error when the stream closes terminally", async () => { + // Given + const store = makeStore(); + await store.getState().submitMessage("Summarize findings"); + + // When: the EventSource fails terminally (e.g. 401/404 on the SSE GET) + eventSources[0].fail(2 /* EventSource.CLOSED */); + + // Then + expect(store.getState().feedback).toBe( + "Unable to connect to the response stream.", + ); + }); +}); + +function makeStore( + overrides?: Partial[0]>, +) { + return createLighthouseChatStore({ + config, + syncUrlToSession: false, + ...overrides, + }); +} + +function model(id: string, name = id): LighthouseV2SupportedModel { + return { + id, + name, + maxInputTokens: null, + maxOutputTokens: null, + supportsFunctionCalling: null, + supportsVision: null, + supportsReasoning: null, + }; +} + +function message( + id: string, + role: LighthouseV2Message["role"], + content: string, +): LighthouseV2Message { + return { + id, + role, + model: null, + tokenUsage: null, + insertedAt: "2026-06-25T10:00:00Z", + parts: [ + { + id: `${id}-part`, + type: "text", + content, + toolCallOutcome: null, + insertedAt: "2026-06-25T10:00:00Z", + updatedAt: "2026-06-25T10:00:00Z", + }, + ], + }; +} diff --git a/ui/app/(prowler)/lighthouse/_lib/chat-store.ts b/ui/app/(prowler)/lighthouse/_lib/chat-store.ts new file mode 100644 index 0000000000..0cdbde9478 --- /dev/null +++ b/ui/app/(prowler)/lighthouse/_lib/chat-store.ts @@ -0,0 +1,485 @@ +import { createStore, type StoreApi } from "zustand/vanilla"; + +import { + createLighthouseV2Session, + getLighthouseV2Messages, + sendLighthouseV2Message, + updateLighthouseV2Configuration, +} from "@/app/(prowler)/lighthouse/_actions"; +import { + createInitialLighthouseV2StreamState, + type LighthouseV2StreamState, + reduceLighthouseV2Event, +} from "@/app/(prowler)/lighthouse/_lib/event-reducer"; +import { + buildOptimisticMessage, + buildSessionTitle, +} from "@/app/(prowler)/lighthouse/_lib/messages"; +import type { LighthouseV2ModelSelection } from "@/app/(prowler)/lighthouse/_lib/model-selection"; +import { notifyLighthouseV2SessionsChanged } from "@/app/(prowler)/lighthouse/_lib/session-events"; +import { parseStreamEvent } from "@/app/(prowler)/lighthouse/_lib/stream-event-parser"; +import { buildLighthouseV2StreamUrl } from "@/app/(prowler)/lighthouse/_lib/stream-url"; +import { + LIGHTHOUSE_V2_PROVIDER_TYPE, + LIGHTHOUSE_V2_SSE_EVENT, + type LighthouseV2Configuration, + type LighthouseV2Message, + type LighthouseV2ProviderType, + type LighthouseV2SSEEvent, + type LighthouseV2SupportedModel, + type LighthouseV2SupportedProvider, +} from "@/app/(prowler)/lighthouse/_types"; + +export interface LighthouseChatConfig { + configurations: LighthouseV2Configuration[]; + modelsByProvider: Record< + LighthouseV2ProviderType, + LighthouseV2SupportedModel[] + >; + supportedProviders: LighthouseV2SupportedProvider[]; +} + +export interface CreateLighthouseChatStoreOptions { + config: LighthouseChatConfig; + // The /lighthouse page mirrors the active session into the URL via + // replaceState; other surfaces (side panel, drawers) must never touch it. + syncUrlToSession: boolean; + initialSessionId?: string; + initialMessages?: LighthouseV2Message[]; + initialInput?: string; + initialError?: string; +} + +export interface LighthouseChatState { + config: LighthouseChatConfig; + activeSessionId: string | null; + messages: LighthouseV2Message[]; + streamState: LighthouseV2StreamState; + input: string; + feedback: string | null; + blockedByConflict: boolean; + isSubmitting: boolean; + isLoadingSession: boolean; + lastSubmittedText: string | null; + selectedModelSelection: LighthouseV2ModelSelection | null; + modelPreferenceSaving: boolean; + setSessionUrlSyncEnabled: (enabled: boolean) => void; + setInput: (value: string) => void; + dismissFeedback: () => void; + selectModel: (selection: LighthouseV2ModelSelection) => Promise; + submitMessage: (text: string) => Promise; + openSession: (sessionId: string) => Promise; + resetToNewChat: () => void; + handleSessionArchived: (sessionId: string) => void; + destroy: () => void; +} + +export type LighthouseChatStore = StoreApi; + +export function selectLighthouseChatCanSend( + state: LighthouseChatState, +): boolean { + const selectedConfiguration = state.config.configurations.find( + (configuration) => + configuration.connected === true && + configuration.providerType === state.selectedModelSelection?.providerType, + ); + return ( + selectedConfiguration?.connected === true && + Boolean(state.selectedModelSelection?.modelId) && + !state.streamState.activeTaskId && + !state.blockedByConflict && + !state.isSubmitting + ); +} + +export function createLighthouseChatStore( + options: CreateLighthouseChatStoreOptions, +): LighthouseChatStore { + const { config } = options; + const connectedConfigurations = config.configurations.filter( + (configuration) => configuration.connected === true, + ); + // The EventSource lives in this closure (never in state): it isn't + // serializable, no render depends on it, and here it survives the consuming + // component unmounting — the reason this factory exists. + let eventSource: EventSource | null = null; + // Set by destroy(): async flows check it after each await so a torn-down + // store never rewrites the URL of another page or opens an orphan stream. + let destroyed = false; + // User-driven session changes invalidate async session creation. Comparing + // only activeSessionId is insufficient because both the initial chat and a + // later reset intentionally use null. + let sessionIntentVersion = 0; + let syncUrlToSession = options.syncUrlToSession; + + const syncSessionUrl = (sessionId: string | null) => { + if (!syncUrlToSession) return; + const url = sessionId + ? `/lighthouse?session=${encodeURIComponent(sessionId)}` + : "/lighthouse"; + window.history.replaceState(window.history.state, "", url); + }; + + return createStore()((set, get) => { + const closeStream = () => { + eventSource?.close(); + eventSource = null; + }; + + const refreshMessages = async ( + sessionId: string, + shouldApply: () => boolean = () => true, + ): Promise => { + const result = await getLighthouseV2Messages(sessionId); + // The fetch is async, so a reset (new chat, or archiving this session) + // can land while it is in flight. Drop the stale result instead of + // repopulating a chat that no longer points at this session. + if (sessionId !== get().activeSessionId || !shouldApply()) return false; + if ("data" in result) { + set({ messages: result.data }); + return true; + } + return false; + }; + + const handleTerminalEvent = async ( + sessionId: string, + event: LighthouseV2SSEEvent, + ) => { + if ( + event.type === LIGHTHOUSE_V2_SSE_EVENT.MESSAGE_END || + event.type === LIGHTHOUSE_V2_SSE_EVENT.ERROR + ) { + closeStream(); + set({ blockedByConflict: false }); + if (event.type === LIGHTHOUSE_V2_SSE_EVENT.ERROR) { + set({ feedback: event.detail || "Agent run failed." }); + } + // A fast follow-up can start while this refresh is in flight; applying + // it would erase the new optimistic message and provisional task id. + const noNewerSubmission = () => + !get().isSubmitting && !get().streamState.activeTaskId; + const refreshed = await refreshMessages(sessionId, noNewerSubmission); + if (refreshed) { + set({ streamState: createInitialLighthouseV2StreamState() }); + } + notifyLighthouseV2SessionsChanged(); + } + }; + + const startStream = (streamUrl: string, sessionId: string) => { + closeStream(); + const source = new EventSource(streamUrl); + eventSource = source; + + const applyEvent = (event: LighthouseV2SSEEvent) => { + set((current) => ({ + streamState: reduceLighthouseV2Event(current.streamState, event), + })); + void handleTerminalEvent(sessionId, event); + }; + + source.addEventListener("message.delta", (event) => + applyEvent( + parseStreamEvent(event, LIGHTHOUSE_V2_SSE_EVENT.MESSAGE_DELTA), + ), + ); + source.addEventListener("tool_call.start", (event) => + applyEvent( + parseStreamEvent(event, LIGHTHOUSE_V2_SSE_EVENT.TOOL_CALL_START), + ), + ); + source.addEventListener("tool_call.end", (event) => + applyEvent( + parseStreamEvent(event, LIGHTHOUSE_V2_SSE_EVENT.TOOL_CALL_END), + ), + ); + source.addEventListener("message.end", (event) => + applyEvent( + parseStreamEvent(event, LIGHTHOUSE_V2_SSE_EVENT.MESSAGE_END), + ), + ); + source.addEventListener("error", (event) => { + if (event instanceof MessageEvent) { + applyEvent(parseStreamEvent(event, LIGHTHOUSE_V2_SSE_EVENT.ERROR)); + } + }); + // The browser fires `onerror` both on a transient drop (it auto-reconnects) + // and on a non-retryable failure such as a 401/404 on the SSE GET. Only the + // latter leaves the source CLOSED, so surface a connection error there and + // treat everything else as a reconnect. + source.onerror = () => { + if (eventSource !== source) return; + if (source.readyState === EventSource.CLOSED) { + closeStream(); + set({ feedback: "Unable to connect to the response stream." }); + } + set((current) => ({ + streamState: reduceLighthouseV2Event(current.streamState, { + type: "disconnect", + }), + })); + }; + }; + + const ensureSession = async (text: string) => { + const existingSessionId = get().activeSessionId; + if (existingSessionId) { + return existingSessionId; + } + + const intentVersion = sessionIntentVersion; + const title = buildSessionTitle(text); + const result = await createLighthouseV2Session(title); + if (destroyed || intentVersion !== sessionIntentVersion) return null; + if ("error" in result) { + set({ feedback: result.error }); + return null; + } + + // Update the URL in place (not router.push) so the force-dynamic server + // component is NOT re-run mid-submit. A re-run would change `key` in + // page.tsx and remount the chat, tearing down the open EventSource. + syncSessionUrl(result.data.id); + set({ activeSessionId: result.data.id }); + notifyLighthouseV2SessionsChanged(); + return result.data.id; + }; + + return { + config, + activeSessionId: options.initialSessionId ?? null, + messages: options.initialMessages ?? [], + streamState: createInitialLighthouseV2StreamState(), + input: options.initialInput ?? "", + feedback: options.initialError ?? null, + blockedByConflict: false, + isSubmitting: false, + isLoadingSession: false, + lastSubmittedText: null, + selectedModelSelection: resolveInitialModelSelection( + connectedConfigurations, + config.modelsByProvider, + ), + modelPreferenceSaving: false, + + setSessionUrlSyncEnabled: (enabled) => { + syncUrlToSession = enabled; + }, + + setInput: (value) => set({ input: value }), + + dismissFeedback: () => set({ feedback: null }), + + selectModel: async (selection) => { + // The selection drives the model used for the next message, so it stays + // applied even if persisting it as the provider's default model fails — + // reverting it would make a connected provider unusable when the save 4xxs. + set({ selectedModelSelection: selection, feedback: null }); + + const configId = connectedConfigurations.find( + (configuration) => + configuration.providerType === selection.providerType, + )?.id; + if (!configId) return; + + set({ modelPreferenceSaving: true }); + + const result = await updateLighthouseV2Configuration(configId, { + defaultModel: selection.modelId, + }); + + set({ modelPreferenceSaving: false }); + + if ("error" in result) { + set({ feedback: result.error }); + } + }, + + submitMessage: async (text) => { + const trimmedText = text.trim(); + if (!trimmedText) return; + if (!get().selectedModelSelection) { + set({ feedback: "Select a model before sending a message." }); + return; + } + if (!selectLighthouseChatCanSend(get())) return; + + set({ isSubmitting: true }); + try { + const sessionId = await ensureSession(trimmedText); + if (!sessionId || destroyed) return; + + const selection = get().selectedModelSelection; + if (!selection) return; + + const provisionalTaskId = `pending-${Date.now()}`; + set((current) => ({ + feedback: null, + blockedByConflict: false, + lastSubmittedText: trimmedText, + input: "", + messages: [ + ...current.messages, + buildOptimisticMessage("user", trimmedText), + ], + streamState: + createInitialLighthouseV2StreamState(provisionalTaskId), + })); + + // Subscribe to the same-origin SSE proxy BEFORE sending the message: + // the backend has no replay buffer, so the listener must be attached + // before the worker starts emitting. + startStream(buildLighthouseV2StreamUrl(sessionId), sessionId); + + const result = await sendLighthouseV2Message({ + sessionId, + text: trimmedText, + provider: selection.providerType, + model: selection.modelId, + }); + if (destroyed) return; + + if ("error" in result) { + // Stale guard: the chat may point at another session by now, so + // this failure must not clobber its stream state or feedback. + if (get().activeSessionId !== sessionId) return; + closeStream(); + set({ + streamState: createInitialLighthouseV2StreamState(), + feedback: result.error, + }); + if (result.status === 409) { + set({ blockedByConflict: true }); + } + // Reconcile the optimistic user message against the server on any + // failure — it may or may not have been persisted. + await refreshMessages(sessionId); + return; + } + + set((current) => ({ + streamState: + current.streamState.activeTaskId === provisionalTaskId + ? { ...current.streamState, activeTaskId: result.data.task.id } + : current.streamState, + })); + notifyLighthouseV2SessionsChanged(); + } finally { + set({ isSubmitting: false }); + } + }, + + openSession: async (sessionId) => { + if (get().activeSessionId === sessionId) return; + sessionIntentVersion += 1; + closeStream(); + set({ + activeSessionId: sessionId, + messages: [], + input: "", + feedback: null, + blockedByConflict: false, + isSubmitting: false, + isLoadingSession: true, + lastSubmittedText: null, + streamState: createInitialLighthouseV2StreamState(), + }); + syncSessionUrl(sessionId); + + const result = await getLighthouseV2Messages(sessionId); + // Stale guard: a reset or another openSession can land mid-fetch. + if (get().activeSessionId !== sessionId) return; + if ("data" in result) { + set({ messages: result.data, isLoadingSession: false }); + } else { + set({ feedback: result.error, isLoadingSession: false }); + } + }, + + resetToNewChat: () => { + sessionIntentVersion += 1; + closeStream(); + set({ + activeSessionId: null, + messages: [], + input: "", + feedback: null, + blockedByConflict: false, + isSubmitting: false, + isLoadingSession: false, + lastSubmittedText: null, + streamState: createInitialLighthouseV2StreamState(), + }); + syncSessionUrl(null); + }, + + handleSessionArchived: (sessionId) => { + // Archiving deletes the session; when it's the open one, fall back to a + // new chat instead of leaving a dead conversation on screen. + if (sessionId === get().activeSessionId) { + get().resetToNewChat(); + } + }, + + destroy: () => { + destroyed = true; + closeStream(); + }, + }; + }); +} + +// Fixed precedence used to pick which connected provider opens the chat. Any +// provider outside this list keeps its relative order behind these. +const LIGHTHOUSE_V2_PROVIDER_PRIORITY = [ + LIGHTHOUSE_V2_PROVIDER_TYPE.OPENAI, + LIGHTHOUSE_V2_PROVIDER_TYPE.BEDROCK, + LIGHTHOUSE_V2_PROVIDER_TYPE.OPENAI_COMPATIBLE, +] as const; + +// Fallback model per provider when the configuration has no remembered model. +const LIGHTHOUSE_V2_PREFERRED_DEFAULT_MODEL: Partial< + Record +> = { + [LIGHTHOUSE_V2_PROVIDER_TYPE.OPENAI]: "gpt-5.6-terra", +}; + +function resolveInitialModelSelection( + connectedConfigurations: LighthouseV2Configuration[], + modelsByProvider: Record< + LighthouseV2ProviderType, + LighthouseV2SupportedModel[] + >, +): LighthouseV2ModelSelection | null { + const priorityIndex = (providerType: LighthouseV2ProviderType) => { + const index = LIGHTHOUSE_V2_PROVIDER_PRIORITY.indexOf(providerType); + return index === -1 ? LIGHTHOUSE_V2_PROVIDER_PRIORITY.length : index; + }; + // Stable sort keeps providers outside the priority list in their original order. + const orderedConfigurations = [...connectedConfigurations].sort( + (a, b) => priorityIndex(a.providerType) - priorityIndex(b.providerType), + ); + + for (const configuration of orderedConfigurations) { + const providerModels = modelsByProvider[configuration.providerType] ?? []; + if (providerModels.length === 0) continue; + // Prefer the provider's remembered model when it is still supported, then + // the provider's preferred default, then the first supported model. + const rememberedModel = providerModels.find( + (model) => model.id === configuration.defaultModel, + ); + const preferredModel = providerModels.find( + (model) => + model.id === + LIGHTHOUSE_V2_PREFERRED_DEFAULT_MODEL[configuration.providerType], + ); + return { + providerType: configuration.providerType, + modelId: (rememberedModel ?? preferredModel ?? providerModels[0]).id, + }; + } + + return null; +} diff --git a/ui/app/(prowler)/lighthouse/_lib/event-reducer.test.ts b/ui/app/(prowler)/lighthouse/_lib/event-reducer.test.ts index 4f284a5656..89c62ef8bc 100644 --- a/ui/app/(prowler)/lighthouse/_lib/event-reducer.test.ts +++ b/ui/app/(prowler)/lighthouse/_lib/event-reducer.test.ts @@ -62,7 +62,7 @@ describe("event-reducer", () => { state = reduceLighthouseV2Event(state, { type: "tool_call.start", toolCallId: "tool-1", - toolName: "prowler_app_search_security_findings", + toolName: "prowler_search_security_findings", }); state = reduceLighthouseV2Event(state, { type: "tool_call.end", @@ -84,7 +84,7 @@ describe("event-reducer", () => { { id: "tool-1", type: "tool_call", - name: "prowler_app_search_security_findings", + name: "prowler_search_security_findings", status: "completed", outcome: "success", }, diff --git a/ui/app/(prowler)/lighthouse/_lib/load-chat-config.ts b/ui/app/(prowler)/lighthouse/_lib/load-chat-config.ts new file mode 100644 index 0000000000..becc488a34 --- /dev/null +++ b/ui/app/(prowler)/lighthouse/_lib/load-chat-config.ts @@ -0,0 +1,86 @@ +import { + getLighthouseV2Configurations, + getLighthouseV2SupportedModels, + getLighthouseV2SupportedProviders, +} from "@/app/(prowler)/lighthouse/_actions"; +import type { LighthouseChatConfig } from "@/app/(prowler)/lighthouse/_lib/chat-store"; +import { loadLighthouseV2ConnectedModels } from "@/app/(prowler)/lighthouse/_lib/model-loading"; + +export const LIGHTHOUSE_CHAT_CONFIG_STATUS = { + ERROR: "error", + NOT_CONFIGURED: "not-configured", + READY: "ready", +} as const; + +interface LighthouseChatConfigError { + status: typeof LIGHTHOUSE_CHAT_CONFIG_STATUS.ERROR; + message: string; +} + +interface LighthouseChatConfigNotConfigured { + status: typeof LIGHTHOUSE_CHAT_CONFIG_STATUS.NOT_CONFIGURED; +} + +interface LighthouseChatConfigReady { + status: typeof LIGHTHOUSE_CHAT_CONFIG_STATUS.READY; + config: LighthouseChatConfig; + modelsError?: string; +} + +export type LighthouseChatConfigResult = + | LighthouseChatConfigError + | LighthouseChatConfigNotConfigured + | LighthouseChatConfigReady; + +// Shared by the /lighthouse server page and the client panel: both need the +// same configurations + providers + connected-models bundle. Server actions +// are callable from either context. Rejections propagate to the caller. +export async function loadLighthouseChatConfig(): Promise { + const [configurationsResult, supportedProvidersResult] = await Promise.all([ + getLighthouseV2Configurations(), + getLighthouseV2SupportedProviders(), + ]); + if ("error" in configurationsResult) { + return { + status: LIGHTHOUSE_CHAT_CONFIG_STATUS.ERROR, + message: configurationsResult.error, + }; + } + if ("error" in supportedProvidersResult) { + return { + status: LIGHTHOUSE_CHAT_CONFIG_STATUS.ERROR, + message: supportedProvidersResult.error, + }; + } + + const configurations = configurationsResult.data; + const hasConnectedProvider = configurations.some( + (configuration) => configuration.connected === true, + ); + if (!hasConnectedProvider) { + return { status: LIGHTHOUSE_CHAT_CONFIG_STATUS.NOT_CONFIGURED }; + } + + const { modelsByProvider, failedModelProviders } = + await loadLighthouseV2ConnectedModels( + configurations, + getLighthouseV2SupportedModels, + ); + // Surface (rather than silently swallow to []) connected providers whose + // models failed to load, so their empty list reads as a real backend + // failure. Disconnected providers are never fetched (see model-loading.ts). + const modelsError = + failedModelProviders.length > 0 + ? `Could not load available models for: ${failedModelProviders.join(", ")}. Try again shortly.` + : undefined; + + return { + status: LIGHTHOUSE_CHAT_CONFIG_STATUS.READY, + config: { + configurations, + modelsByProvider, + supportedProviders: supportedProvidersResult.data, + }, + modelsError, + }; +} diff --git a/ui/app/(prowler)/lighthouse/_lib/panel-chat-message-state.ts b/ui/app/(prowler)/lighthouse/_lib/panel-chat-message-state.ts new file mode 100644 index 0000000000..55661c55bd --- /dev/null +++ b/ui/app/(prowler)/lighthouse/_lib/panel-chat-message-state.ts @@ -0,0 +1,44 @@ +type PanelChatMessageStateListener = () => void; + +const listeners = new Set(); +let hasMessages = false; +let activeSessionId: string | null = null; + +export function getPanelChatHasMessages(): boolean { + return hasMessages; +} + +export function subscribePanelChatHasMessages( + listener: PanelChatMessageStateListener, +): () => void { + listeners.add(listener); + return () => listeners.delete(listener); +} + +export function getPanelChatActiveSessionId(): string | null { + return activeSessionId; +} + +interface PanelChatMessageState { + hasMessages: boolean; + activeSessionId: string | null; +} + +export function setPanelChatMessageState( + nextState: PanelChatMessageState, +): void { + if ( + hasMessages === nextState.hasMessages && + activeSessionId === nextState.activeSessionId + ) { + return; + } + + hasMessages = nextState.hasMessages; + activeSessionId = nextState.activeSessionId; + listeners.forEach((listener) => listener()); +} + +export function resetPanelChatMessageState(): void { + setPanelChatMessageState({ hasMessages: false, activeSessionId: null }); +} diff --git a/ui/app/(prowler)/lighthouse/_lib/panel-chat-store.ts b/ui/app/(prowler)/lighthouse/_lib/panel-chat-store.ts new file mode 100644 index 0000000000..d435fde2b2 --- /dev/null +++ b/ui/app/(prowler)/lighthouse/_lib/panel-chat-store.ts @@ -0,0 +1,57 @@ +import { + createLighthouseChatStore, + type LighthouseChatConfig, + type LighthouseChatStore, +} from "@/app/(prowler)/lighthouse/_lib/chat-store"; + +// Module-level singleton: the global side panel keeps the same conversation +// while switching between Details and Lighthouse AI, across route navigation +// and panel closes. The full-page route can reuse it for the same conversation. +let panelChatStore: LighthouseChatStore | null = null; + +interface PanelChatStoreOptions { + initialError?: string; +} + +export function getOrCreatePanelChatStore( + config: LighthouseChatConfig, + options?: PanelChatStoreOptions, +): LighthouseChatStore { + if (!panelChatStore) { + panelChatStore = createLighthouseChatStore({ + config, + syncUrlToSession: false, + initialError: options?.initialError, + }); + } + return panelChatStore; +} + +// Lets the full-page surface reuse the singleton only when both surfaces point +// at the same conversation. This is intentionally a pure lookup: React may +// run state initializers twice in Strict Mode. +export function getPanelChatStoreForSession( + initialSessionId?: string, +): LighthouseChatStore | null { + if (!panelChatStore) return null; + const expectedSessionId = initialSessionId ?? null; + if (panelChatStore.getState().activeSessionId !== expectedSessionId) { + return null; + } + return panelChatStore; +} + +export function isPanelChatStore(store: LighthouseChatStore): boolean { + return panelChatStore === store; +} + +// The config is captured in the store's closure at creation, so a +// configuration change must tear the singleton down and rebuild it. +export function resetPanelChatStore(): void { + panelChatStore?.getState().destroy(); + panelChatStore = null; +} + +export function resetPanelChatStoreForTests(): void { + resetPanelChatStore(); +} diff --git a/ui/app/(prowler)/lighthouse/_lib/session-events.ts b/ui/app/(prowler)/lighthouse/_lib/session-events.ts index ab23282c49..84c722fcd6 100644 --- a/ui/app/(prowler)/lighthouse/_lib/session-events.ts +++ b/ui/app/(prowler)/lighthouse/_lib/session-events.ts @@ -30,3 +30,43 @@ export function notifyLighthouseV2NewChat() { if (typeof window === "undefined") return; window.dispatchEvent(new Event(LIGHTHOUSE_V2_NEW_CHAT_EVENT)); } + +export const LIGHTHOUSE_V2_CONFIGURATIONS_CHANGED_EVENT = + "lighthouse-v2:configurations-changed"; + +// Fired after provider configuration CRUD so cached chat configs (the panel +// keeps one at module scope) can invalidate and reload. +export function notifyLighthouseV2ConfigurationsChanged() { + if (typeof window === "undefined") return; + window.dispatchEvent(new Event(LIGHTHOUSE_V2_CONFIGURATIONS_CHANGED_EVENT)); +} + +// Typed subscribe helpers: each returns an unsubscribe function so consumers +// never hand-roll addEventListener plus the CustomEvent detail cast. +function subscribe(eventName: string, handler: (event: Event) => void) { + if (typeof window === "undefined") return () => {}; + window.addEventListener(eventName, handler); + return () => window.removeEventListener(eventName, handler); +} + +export function onLighthouseV2SessionsChanged(callback: () => void) { + return subscribe(LIGHTHOUSE_V2_SESSIONS_CHANGED_EVENT, callback); +} + +export function onLighthouseV2SessionArchived( + callback: (sessionId: string) => void, +) { + return subscribe(LIGHTHOUSE_V2_SESSION_ARCHIVED_EVENT, (event) => { + const sessionId = (event as CustomEvent<{ sessionId: string }>).detail + ?.sessionId; + if (sessionId) callback(sessionId); + }); +} + +export function onLighthouseV2NewChat(callback: () => void) { + return subscribe(LIGHTHOUSE_V2_NEW_CHAT_EVENT, callback); +} + +export function onLighthouseV2ConfigurationsChanged(callback: () => void) { + return subscribe(LIGHTHOUSE_V2_CONFIGURATIONS_CHANGED_EVENT, callback); +} diff --git a/ui/app/(prowler)/lighthouse/_lib/testing/event-source-mock.ts b/ui/app/(prowler)/lighthouse/_lib/testing/event-source-mock.ts new file mode 100644 index 0000000000..2672509510 --- /dev/null +++ b/ui/app/(prowler)/lighthouse/_lib/testing/event-source-mock.ts @@ -0,0 +1,49 @@ +import { vi } from "vitest"; + +// Controllable EventSource mock: records each instance so tests can drive +// named SSE events and connection failures, while still being a vi.fn so +// `expect(EventSource).toHaveBeenCalledWith(...)` keeps working. +export interface MockEventSource { + url: string; + readyState: number; + onerror: ((event: Event) => void) | null; + listeners: Map>; + addEventListener: (type: string, cb: EventListener) => void; + close: ReturnType; + emit: (type: string, data: unknown) => void; + fail: (readyState: number) => void; +} + +// The mock never fires "open": the client must POST the message without +// waiting for it (the backend sends no bytes until the worker emits, which +// only happens after the POST). This is the regression guard for the +// open-gate deadlock. +export function stubEventSource(): MockEventSource[] { + const eventSources: MockEventSource[] = []; + const EventSourceMock = vi.fn(function (this: MockEventSource, url: string) { + this.url = url; + this.readyState = 0; + this.onerror = null; + this.listeners = new Map(); + this.addEventListener = (type: string, cb: EventListener) => { + const set = this.listeners.get(type) ?? new Set(); + set.add(cb); + this.listeners.set(type, set); + }; + this.close = vi.fn(() => { + this.readyState = 2; + }); + this.emit = (type: string, data: unknown) => { + const event = new MessageEvent(type, { data: JSON.stringify(data) }); + this.listeners.get(type)?.forEach((cb) => cb(event)); + }; + this.fail = (readyState: number) => { + this.readyState = readyState; + this.onerror?.(new Event("error")); + }; + eventSources.push(this); + }); + Object.assign(EventSourceMock, { CONNECTING: 0, OPEN: 1, CLOSED: 2 }); + vi.stubGlobal("EventSource", EventSourceMock); + return eventSources; +} diff --git a/ui/app/(prowler)/lighthouse/_lib/tool-calls.test.ts b/ui/app/(prowler)/lighthouse/_lib/tool-calls.test.ts index 570e7cc597..5ccc2b4129 100644 --- a/ui/app/(prowler)/lighthouse/_lib/tool-calls.test.ts +++ b/ui/app/(prowler)/lighthouse/_lib/tool-calls.test.ts @@ -11,7 +11,7 @@ describe("getToolCallContent", () => { // Given const content = { tool_call_id: "call_1", - tool_name: "prowler_app_search_security_findings", + tool_name: "prowler_search_security_findings", arguments: { severity: "high" }, result: { count: 3 }, outcome: "success", @@ -23,7 +23,7 @@ describe("getToolCallContent", () => { // Then expect(parsed).toEqual({ toolCallId: "call_1", - toolName: "prowler_app_search_security_findings", + toolName: "prowler_search_security_findings", arguments: { severity: "high" }, result: { count: 3 }, outcome: "success", @@ -53,12 +53,18 @@ describe("getToolCallContent", () => { describe("formatToolName", () => { it("should strip the prowler prefix and title-case", () => { - expect(formatToolName("prowler_app_search_security_findings")).toBe( + expect(formatToolName("prowler_search_security_findings")).toBe( "Search security findings", ); expect(formatToolName("prowler_hub_list_checks")).toBe("List checks"); }); + it("should still strip the legacy prowler_app_ prefix", () => { + expect(formatToolName("prowler_app_search_security_findings")).toBe( + "Search security findings", + ); + }); + it("should humanize prefix-less tools", () => { expect(formatToolName("search_tools")).toBe("Search tools"); }); diff --git a/ui/app/(prowler)/lighthouse/_lib/tool-calls.ts b/ui/app/(prowler)/lighthouse/_lib/tool-calls.ts index 46034102d5..88e478089b 100644 --- a/ui/app/(prowler)/lighthouse/_lib/tool-calls.ts +++ b/ui/app/(prowler)/lighthouse/_lib/tool-calls.ts @@ -1,11 +1,14 @@ import type { LighthouseV2ToolCallContent } from "@/app/(prowler)/lighthouse/_types"; // Prefixes shared by the MCP-sourced tools; stripped for display so a name like -// `prowler_app_search_security_findings` reads as "Search security findings". +// `prowler_search_security_findings` reads as "Search security findings". The +// specific prefixes are listed before the bare `prowler_` catch-all so Hub, Docs, +// and legacy `prowler_app_` records match first and render cleanly. const TOOL_NAME_PREFIXES = [ - "prowler_app_", "prowler_hub_", "prowler_docs_", + "prowler_app_", + "prowler_", ] as const; // Reads the snake_case TOOL_CALL blob the backend persists and normalizes it to diff --git a/ui/app/(prowler)/lighthouse/page.tsx b/ui/app/(prowler)/lighthouse/page.tsx index 625aaf795a..f54be45aed 100644 --- a/ui/app/(prowler)/lighthouse/page.tsx +++ b/ui/app/(prowler)/lighthouse/page.tsx @@ -4,14 +4,12 @@ import { getLighthouseProvidersConfig, isLighthouseConfigured, } from "@/actions/lighthouse-v1/lighthouse"; -import { - getLighthouseV2Configurations, - getLighthouseV2Messages, - getLighthouseV2SupportedModels, - getLighthouseV2SupportedProviders, -} from "@/app/(prowler)/lighthouse/_actions"; +import { getLighthouseV2Messages } from "@/app/(prowler)/lighthouse/_actions"; import { LighthouseV2ChatPage } from "@/app/(prowler)/lighthouse/_components/chat"; -import { loadLighthouseV2ConnectedModels } from "@/app/(prowler)/lighthouse/_lib/model-loading"; +import { + LIGHTHOUSE_CHAT_CONFIG_STATUS, + loadLighthouseChatConfig, +} from "@/app/(prowler)/lighthouse/_lib/load-chat-config"; import { LighthouseIcon } from "@/components/icons/Icons"; import { APP_SIDEBAR_MODE, @@ -36,34 +34,13 @@ export default async function AIChatbot({ typeof params.session === "string" ? params.session : undefined; if (isCloud()) { - const [configurationsResult, supportedProvidersResult] = await Promise.all([ - getLighthouseV2Configurations(), - getLighthouseV2SupportedProviders(), - ]); - const configurations = - "data" in configurationsResult ? configurationsResult.data : []; - const supportedProviders = - "data" in supportedProvidersResult ? supportedProvidersResult.data : []; - const connectedConfigurations = configurations.filter( - (configuration) => configuration.connected === true, - ); - - if (connectedConfigurations.length === 0) { + const chatConfigResult = await loadLighthouseChatConfig(); + // Errors and the not-configured case both land on settings, where the + // user can connect (or fix) a provider. + if (chatConfigResult.status !== LIGHTHOUSE_CHAT_CONFIG_STATUS.READY) { return redirect(LIGHTHOUSE_ROUTE.SETTINGS); } - - const { modelsByProvider, failedModelProviders } = - await loadLighthouseV2ConnectedModels( - configurations, - getLighthouseV2SupportedModels, - ); - // Surface (rather than silently swallow to []) connected providers whose - // models failed to load, so their empty list reads as a real backend - // failure. Disconnected providers are never fetched (see model-loading.ts). - const modelsError = - failedModelProviders.length > 0 - ? `Could not load available models for: ${failedModelProviders.join(", ")}. Try again shortly.` - : undefined; + const { config, modelsError } = chatConfigResult; const initialMessages = activeSessionId ? await getLighthouseV2Messages(activeSessionId) @@ -80,15 +57,15 @@ export default async function AIChatbot({ return ( }> - + {/* [contain:layout] traps streamdown's fixed fullscreen overlay inside the chat area so it never covers the sidebar or navbar. */}
({ + setThemeMock: vi.fn(), + themeState: { current: "light" }, +})); + +vi.mock("next-themes", () => ({ + useTheme: () => ({ theme: themeState.current, setTheme: setThemeMock }), +})); + +describe("ThemeSwitch", () => { + beforeEach(() => { + setThemeMock.mockClear(); + themeState.current = "light"; + }); + + it("exposes an accessible switch reflecting the current mode", () => { + // Given / When + render(); + + // Then + const control = screen.getByRole("switch", { + name: "Switch to dark mode", + }); + expect(control).toHaveAttribute("aria-checked", "true"); + }); + + it("toggles to the opposite theme on click", () => { + // Given + render(); + + // When + fireEvent.click(screen.getByRole("switch")); + + // Then + expect(setThemeMock).toHaveBeenCalledWith("dark"); + }); + + it("renders as a shared ghost icon button, matching the navbar cluster", () => { + // Given / When + render(); + + // Then: same 32px square treatment as the other navbar actions + const control = screen.getByRole("switch"); + expect(control).toHaveClass("size-8"); + expect(control).not.toHaveClass("rounded-full"); + }); +}); diff --git a/ui/components/ThemeSwitch.tsx b/ui/components/ThemeSwitch.tsx index 4979dec541..8de15b3531 100644 --- a/ui/components/ThemeSwitch.tsx +++ b/ui/components/ThemeSwitch.tsx @@ -3,12 +3,12 @@ import { useTheme } from "next-themes"; import { ComponentProps, useSyncExternalStore } from "react"; +import { Button } from "@/components/shadcn/button/button"; import { Tooltip, TooltipContent, TooltipTrigger, } from "@/components/shadcn/tooltip"; -import { cn } from "@/lib/utils"; import { MoonFilledIcon, SunFilledIcon } from "./icons"; @@ -30,24 +30,23 @@ export function ThemeSwitch({ className, ...props }: ThemeSwitchProps) { return ( - + {isLightMode ? "Switch to Dark Mode" : "Switch to Light Mode"} diff --git a/ui/components/auth/oss/auth-brand.tsx b/ui/components/auth/oss/auth-brand.tsx new file mode 100644 index 0000000000..396064c52b --- /dev/null +++ b/ui/components/auth/oss/auth-brand.tsx @@ -0,0 +1,14 @@ +import { ProwlerBrand } from "@/components/icons"; +import { cn } from "@/lib/utils"; + +interface AuthBrandProps { + className?: string; +} + +export const AuthBrand = ({ className }: AuthBrandProps) => { + return ( +
+ +
+ ); +}; diff --git a/ui/components/auth/oss/auth-divider.tsx b/ui/components/auth/oss/auth-divider.tsx index 2d9fe56b30..edf4a1c736 100644 --- a/ui/components/auth/oss/auth-divider.tsx +++ b/ui/components/auth/oss/auth-divider.tsx @@ -2,9 +2,9 @@ import { Separator } from "@/components/shadcn"; export const AuthDivider = () => { return ( -
+
-

OR

+

or

); diff --git a/ui/components/auth/oss/auth-layout.test.tsx b/ui/components/auth/oss/auth-layout.test.tsx new file mode 100644 index 0000000000..dd144243da --- /dev/null +++ b/ui/components/auth/oss/auth-layout.test.tsx @@ -0,0 +1,39 @@ +import { render, screen } from "@testing-library/react"; +import { describe, expect, it } from "vitest"; + +import { AuthLayout } from "./auth-layout"; + +describe("AuthLayout", () => { + it("renders the Prowler brand directly above the form card", () => { + render( + +

form content

+
, + ); + + const brand = screen.getByRole("img", { name: /prowler/i }); + const title = screen.getByText("Sign in"); + + expect( + brand.compareDocumentPosition(title) & Node.DOCUMENT_POSITION_FOLLOWING, + ).toBeTruthy(); + }); + + it("renders the footer outside the form card, below it", () => { + render( + footer link

}> +

form content

+
, + ); + + const content = screen.getByText("form content"); + const footer = screen.getByText("footer link"); + const card = content.parentElement!; + + expect(card.contains(footer)).toBe(false); + expect( + content.compareDocumentPosition(footer) & + Node.DOCUMENT_POSITION_FOLLOWING, + ).toBeTruthy(); + }); +}); diff --git a/ui/components/auth/oss/auth-layout.tsx b/ui/components/auth/oss/auth-layout.tsx index 6d90c28844..8c122aa661 100644 --- a/ui/components/auth/oss/auth-layout.tsx +++ b/ui/components/auth/oss/auth-layout.tsx @@ -2,12 +2,15 @@ import { ReactNode } from "react"; import { ThemeSwitch } from "@/components/ThemeSwitch"; +import { AuthBrand } from "./auth-brand"; + interface AuthLayoutProps { title: string; + footer?: ReactNode; children: ReactNode; } -export const AuthLayout = ({ title, children }: AuthLayoutProps) => { +export const AuthLayout = ({ title, footer, children }: AuthLayoutProps) => { return (
@@ -20,6 +23,8 @@ export const AuthLayout = ({ title, children }: AuthLayoutProps) => { }} >
+ + {/* Auth Form Container */}
{/* Header with Title and Theme Toggle */} @@ -31,6 +36,8 @@ export const AuthLayout = ({ title, children }: AuthLayoutProps) => { {/* Content */} {children}
+ + {footer &&
{footer}
}
); diff --git a/ui/components/auth/oss/public-auth-shell.tsx b/ui/components/auth/oss/public-auth-shell.tsx deleted file mode 100644 index 372d2d866a..0000000000 --- a/ui/components/auth/oss/public-auth-shell.tsx +++ /dev/null @@ -1,18 +0,0 @@ -import type { ReactNode } from "react"; - -import { ProwlerBrand } from "@/components/icons"; - -interface PublicAuthShellProps { - children: ReactNode; -} - -export const PublicAuthShell = ({ children }: PublicAuthShellProps) => { - return ( -
-
- -
- {children} -
- ); -}; diff --git a/ui/components/auth/oss/sign-in-form.tsx b/ui/components/auth/oss/sign-in-form.tsx index 64f3352279..8ebd8f3c0e 100644 --- a/ui/components/auth/oss/sign-in-form.tsx +++ b/ui/components/auth/oss/sign-in-form.tsx @@ -12,7 +12,12 @@ import { AuthDivider } from "@/components/auth/oss/auth-divider"; import { AuthFooterLink } from "@/components/auth/oss/auth-footer-link"; import { AuthLayout } from "@/components/auth/oss/auth-layout"; import { SocialButtons } from "@/components/auth/oss/social-buttons"; -import { Button } from "@/components/shadcn"; +import { + Button, + Tooltip, + TooltipContent, + TooltipTrigger, +} from "@/components/shadcn"; import { useToast } from "@/components/shadcn"; import { CustomInput } from "@/components/shadcn/custom"; import { Form } from "@/components/shadcn/form"; @@ -142,10 +147,19 @@ export const SignInForm = ({ } }; - const title = isSamlMode ? "Sign in with SAML SSO" : "Sign in"; + const title = isSamlMode ? "Sign in with SAML SSO" : "Welcome back"; return ( - + + } + >
-
+
{!isSamlMode && ( )} - + {isSamlMode ? ( + + ) : ( + + + + + Continue with SAML SSO + + )}
- - ); }; diff --git a/ui/components/auth/oss/sign-up-form.tsx b/ui/components/auth/oss/sign-up-form.tsx index ac0282aec0..2127151b73 100644 --- a/ui/components/auth/oss/sign-up-form.tsx +++ b/ui/components/auth/oss/sign-up-form.tsx @@ -161,7 +161,16 @@ export const SignUpForm = ({ }; return ( - + + } + > -
+
)} - - ); }; diff --git a/ui/components/auth/oss/social-buttons.test.tsx b/ui/components/auth/oss/social-buttons.test.tsx new file mode 100644 index 0000000000..80b9ca6e90 --- /dev/null +++ b/ui/components/auth/oss/social-buttons.test.tsx @@ -0,0 +1,34 @@ +import { render, screen } from "@testing-library/react"; +import { describe, expect, it } from "vitest"; + +import { SocialButtons } from "./social-buttons"; + +describe("SocialButtons", () => { + it("renders icon-only provider links that keep their accessible names", () => { + render( + , + ); + + const google = screen.getByRole("link", { name: "Continue with Google" }); + const github = screen.getByRole("link", { name: "Continue with Github" }); + + expect(google.textContent).toBe(""); + expect(github.textContent).toBe(""); + }); + + it("keeps accessible names on disabled providers", () => { + render(); + + expect( + screen.getByRole("button", { name: "Continue with Google" }), + ).toBeDisabled(); + expect( + screen.getByRole("button", { name: "Continue with Github" }), + ).toBeDisabled(); + }); +}); diff --git a/ui/components/auth/oss/social-buttons.tsx b/ui/components/auth/oss/social-buttons.tsx index b08d830fb5..fd52d0cfb7 100644 --- a/ui/components/auth/oss/social-buttons.tsx +++ b/ui/components/auth/oss/social-buttons.tsx @@ -35,12 +35,13 @@ const SocialButton = ({ const button = ( ); if (!isDisabled) { - return button; + return ( + + {button} + {provider.label} + + ); } return ( - {button} + {button} {provider.isOAuthEnabled ? ( diff --git a/ui/components/compliance/compliance-charts/heatmap-chart.test.tsx b/ui/components/compliance/compliance-charts/heatmap-chart.test.tsx new file mode 100644 index 0000000000..409974af77 --- /dev/null +++ b/ui/components/compliance/compliance-charts/heatmap-chart.test.tsx @@ -0,0 +1,36 @@ +import { render, screen } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { describe, expect, it, vi } from "vitest"; + +import { HeatmapChart } from "./heatmap-chart"; + +vi.mock("next-themes", () => ({ + useTheme: () => ({ theme: "light" }), +})); + +describe("HeatmapChart", () => { + it("portals its pointer-positioned tooltip outside layout containers", async () => { + // Given + const user = userEvent.setup(); + const { container } = render( + , + ); + + // When + await user.hover(screen.getByTitle("Identity")); + + // Then: fixed client coordinates resolve against the viewport, not
+ const tooltip = screen.getByRole("tooltip"); + expect(container).not.toContainElement(tooltip); + expect(tooltip.parentElement).toBe(document.body); + }); +}); diff --git a/ui/components/compliance/compliance-charts/heatmap-chart.tsx b/ui/components/compliance/compliance-charts/heatmap-chart.tsx index 2d7ce547b3..f57fa70aee 100644 --- a/ui/components/compliance/compliance-charts/heatmap-chart.tsx +++ b/ui/components/compliance/compliance-charts/heatmap-chart.tsx @@ -2,6 +2,7 @@ import { useTheme } from "next-themes"; import { useState } from "react"; +import { createPortal } from "react-dom"; import { cn } from "@/lib/utils"; import { CategoryData } from "@/types/compliance"; @@ -115,44 +116,49 @@ export const HeatmapChart = ({ categories = [] }: HeatmapChartProps) => {
{/* Custom Tooltip */} - {hoveredItem && ( -
-
- {capitalizeFirstLetter(hoveredItem.name)} -
-
- - Failure Rate: {hoveredItem.failurePercentage}% - -
-
- - Failed: {hoveredItem.failedRequirements}/ - {hoveredItem.totalRequirements} - -
-
- )} +
+ {capitalizeFirstLetter(hoveredItem.name)} +
+
+ + Failure Rate: {hoveredItem.failurePercentage}% + +
+
+ + Failed: {hoveredItem.failedRequirements}/ + {hoveredItem.totalRequirements} + +
+
, + document.body, + ) + : null}
); diff --git a/ui/components/feeds/feeds-client.tsx b/ui/components/feeds/feeds-client.tsx index e91e16b6d6..6b0e45e367 100644 --- a/ui/components/feeds/feeds-client.tsx +++ b/ui/components/feeds/feeds-client.tsx @@ -58,8 +58,9 @@ export function FeedsClient({ feedData, error }: FeedsClientProps) { -
+ {/* Portaled to body:
is a layout container (container queries), + which would otherwise capture this fixed button and scroll it away + with the content. */} + {typeof document !== "undefined" + ? createPortal( +
+ +
, + document.body, + ) + : null} ); } diff --git a/ui/components/findings/table/column-finding-groups.test.tsx b/ui/components/findings/table/column-finding-groups.test.tsx index 5a8ec51e56..ca6e01bb7b 100644 --- a/ui/components/findings/table/column-finding-groups.test.tsx +++ b/ui/components/findings/table/column-finding-groups.test.tsx @@ -409,7 +409,9 @@ describe("column-finding-groups — accessibility of check title cell", () => { expect( screen.queryByRole("button", { name: "Fallback IaC Check" }), ).not.toBeInTheDocument(); - expect(screen.getByText("Fallback IaC Check")).toBeInTheDocument(); + // The title renders as plain (non-clickable) text and the inline-mocked + // tooltip duplicates it, so exactly both copies must exist. + expect(screen.getAllByText("Fallback IaC Check")).toHaveLength(2); expect(onDrillDown).not.toHaveBeenCalled(); }); }); diff --git a/ui/components/findings/table/column-finding-groups.tsx b/ui/components/findings/table/column-finding-groups.tsx index f200a4ed39..e5cd81cb4d 100644 --- a/ui/components/findings/table/column-finding-groups.tsx +++ b/ui/components/findings/table/column-finding-groups.tsx @@ -199,20 +199,29 @@ export function getColumnFindingGroups({ {providerName} ) : null} -
- {canExpand ? ( - + ) : ( + + {group.checkTitle} + + )} + + {group.checkTitle} - - ) : ( - - {group.checkTitle} - - )} + +
); diff --git a/ui/components/findings/table/column-standalone-findings.tsx b/ui/components/findings/table/column-standalone-findings.tsx index 75887ee9b9..30facaf97f 100644 --- a/ui/components/findings/table/column-standalone-findings.tsx +++ b/ui/components/findings/table/column-standalone-findings.tsx @@ -9,6 +9,11 @@ import { SeverityBadge, StatusFindingBadge, } from "@/components/shadcn/table"; +import { + Tooltip, + TooltipContent, + TooltipTrigger, +} from "@/components/shadcn/tooltip"; import { getRegionFlag } from "@/lib/region-flags"; import { getOptionalText } from "@/lib/utils"; import { FindingProps, ProviderType } from "@/types"; @@ -67,11 +72,17 @@ function FindingTitleCell({ finding={finding} defaultOpen={defaultOpen} trigger={ -
-

+ // Single line always: ellipsis beyond the max, full title in the tooltip. + + +

+ {finding.attributes.check_metadata.checktitle} +

+ + {finding.attributes.check_metadata.checktitle} -

-
+ + } /> ); 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 af7e38acc3..1d2716985f 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 @@ -74,6 +74,7 @@ import { ResourceMetadataPanel } from "@/components/shared/resource-metadata-pan import { getFailingForLabel } from "@/lib/date-utils"; import { formatDuration } from "@/lib/date-utils"; import { shouldRefreshAfterTriageUpdate } from "@/lib/finding-triage"; +import { buildFindingAnalysisPrompt } from "@/lib/lighthouse/prompts"; import { getRegionFlag } from "@/lib/region-flags"; import { getRecommendationLinkLabel } from "@/lib/vulnerability-references"; import type { ComplianceOverviewData } from "@/types/compliance"; @@ -471,6 +472,16 @@ export function ResourceDetailDrawerContent({ const overviewStatusExtended = currentResource?.statusExtended || f?.statusExtended; const showOverviewStatusExtended = Boolean(overviewStatusExtended); + const findingAnalysisPrompt = buildFindingAnalysisPrompt({ + findingId: currentResource?.findingId ?? f?.id, + providerUid, + resourceUid, + checkId: currentResource?.checkId ?? checkMeta.checkId, + severity: findingSeverity, + status: findingStatus, + detail: overviewStatusExtended, + risk: f?.risk || checkMeta.risk, + }); const handleDrawerTriageUpdate = async (input: UpdateFindingTriageInput) => { await updateFindingTriage(input); @@ -712,7 +723,10 @@ export function ResourceDetailDrawerContent({ <>
{/* Resource info grid — 4 data columns */} -
+
{/* Row 1: Provider, Resource, Service, Region */}
- - - Resource Finding Details - - View finding details for the selected resource - - - - - Close - - {open && ( - - )} - - + + + ); } diff --git a/ui/components/findings/table/resource-detail-drawer/resource-detail-skeleton.tsx b/ui/components/findings/table/resource-detail-drawer/resource-detail-skeleton.tsx index 938b4bc35a..163a2dbd54 100644 --- a/ui/components/findings/table/resource-detail-drawer/resource-detail-skeleton.tsx +++ b/ui/components/findings/table/resource-detail-drawer/resource-detail-skeleton.tsx @@ -8,7 +8,10 @@ import { Skeleton } from "@/components/shadcn/skeleton/skeleton"; export function ResourceDetailSkeleton() { return (
-
+
{/* Row 1: Provider, Resource, Service, Region */}
diff --git a/ui/components/icons/Icons.tsx b/ui/components/icons/Icons.tsx index a4053e9f55..0a2edb0526 100644 --- a/ui/components/icons/Icons.tsx +++ b/ui/components/icons/Icons.tsx @@ -1129,7 +1129,13 @@ export const LighthouseIcon: React.FC = ({ height, ...props }) => { - const gradientId = (id: string) => (animatedAura ? `${id}_animated` : id); + // Gradient defs are referenced by id, and this icon renders many times per + // page. Duplicate ids resolve against the FIRST instance in the document — + // if that one sits in a display:none subtree (e.g. the desktop sidebar on + // mobile), every other copy loses its fill and turns invisible. useId keeps + // each instance self-contained; strip its delimiters for url(#...) safety. + const uid = React.useId().replace(/[«»:]/g, ""); + const gradientId = (id: string) => `${id}_${uid}`; return ( { + it("keeps gradient ids unique across instances", () => { + // Given: the icon rendered several times on one page (sidebar, navbar, + // overview banner) — with duplicate ids, browsers resolve url(#...) + // against the first instance, which may sit in a display:none subtree + // (the desktop sidebar on mobile) and leave the others unpainted. + const { container } = render( + <> + + + + , + ); + + // Then: every gradient id is unique document-wide + const ids = Array.from( + container.querySelectorAll("linearGradient, radialGradient"), + ).map((gradient) => gradient.id); + expect(ids.length).toBeGreaterThan(0); + expect(new Set(ids).size).toBe(ids.length); + }); + + it("paints every path from its own instance's defs", () => { + // Given + const { container } = render(); + const svg = container.querySelector("svg") as SVGElement; + const localIds = new Set( + Array.from(svg.querySelectorAll("linearGradient, radialGradient")).map( + (gradient) => gradient.id, + ), + ); + + // Then: every fill/stroke reference resolves inside this same svg + const references = Array.from(svg.querySelectorAll("path")) + .flatMap((path) => [ + path.getAttribute("fill"), + path.getAttribute("stroke"), + ]) + .filter((paint): paint is string => paint?.startsWith("url(#") ?? false); + expect(references.length).toBeGreaterThan(0); + for (const reference of references) { + const id = reference.slice("url(#".length, -1); + expect(localIds.has(id)).toBe(true); + } + }); +}); diff --git a/ui/components/integrations/s3/s3-integration-form.test.tsx b/ui/components/integrations/s3/s3-integration-form.test.tsx new file mode 100644 index 0000000000..7186cc42ed --- /dev/null +++ b/ui/components/integrations/s3/s3-integration-form.test.tsx @@ -0,0 +1,244 @@ +import { render, screen } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import type { ComponentProps } from "react"; +import { beforeEach, describe, expect, it, vi } from "vitest"; + +import type { IntegrationProps } from "@/types/integrations"; +import type { ProviderProps } from "@/types/providers"; + +import { S3IntegrationForm } from "./s3-integration-form"; + +const { createIntegrationMock, toastMock, updateIntegrationMock } = vi.hoisted( + () => ({ + createIntegrationMock: vi.fn(), + toastMock: vi.fn(), + updateIntegrationMock: vi.fn(), + }), +); + +vi.mock("@/actions/integrations", () => ({ + createIntegration: createIntegrationMock, + updateIntegration: updateIntegrationMock, +})); + +vi.mock("next-auth/react", () => ({ + useSession: () => ({ + data: { + tenantId: "tenant-id", + }, + }), +})); + +vi.mock("@/components/shadcn", async (importOriginal) => ({ + ...(await importOriginal>()), + useToast: () => ({ + toast: toastMock, + }), +})); + +interface MockEnhancedMultiSelectProps { + onValueChange: (values: string[]) => void; + options: Array<{ value: string }>; +} + +vi.mock("@/components/shadcn/select/enhanced-multi-select", () => ({ + EnhancedMultiSelect: ({ + onValueChange, + options, + }: MockEnhancedMultiSelectProps) => ( + + ), +})); + +vi.mock( + "@/components/providers/workflow/forms/select-credentials-type/aws/credentials-type/aws-role-credentials-form", + () => ({ + AWSRoleCredentialsForm: ({ + templateLinks, + }: { + templateLinks: { cloudformationQuickLink: string }; + }) => ( + + {templateLinks.cloudformationQuickLink} + + ), + }), +); + +vi.mock("@/lib", () => ({ + getAWSCredentialsTemplateLinks: ( + _externalId: string, + _bucketName: string, + _integrationType: string, + bucketAccountId?: string, + ) => ({ + cloudformation: "https://example.com/cloudformation", + terraform: "https://example.com/terraform", + cloudformationQuickLink: `https://example.com/quick-create?bucketAccountId=${bucketAccountId ?? ""}`, + }), +})); + +function createProvider( + provider: ProviderProps["attributes"]["provider"], + uid: string, +): ProviderProps { + return { + id: `${provider}-provider`, + type: "providers", + attributes: { + provider, + is_dynamic: false, + uid, + alias: `${provider} provider`, + status: "completed", + resources: 0, + connection: { + connected: true, + last_checked_at: "2026-07-16T00:00:00Z", + }, + scanner_args: { + only_logs: false, + excluded_checks: [], + aws_retries_max_attempts: 3, + }, + inserted_at: "2026-07-16T00:00:00Z", + updated_at: "2026-07-16T00:00:00Z", + created_by: { + object: "users", + id: "user-1", + }, + }, + relationships: { + secret: { + data: null, + }, + provider_groups: { + meta: { + count: 0, + }, + data: [], + }, + }, + }; +} + +function renderS3IntegrationForm( + props?: Partial>, +) { + return render( + , + ); +} + +const integration: IntegrationProps = { + type: "integrations", + id: "integration-1", + attributes: { + inserted_at: "2026-07-16T00:00:00Z", + updated_at: "2026-07-16T00:00:00Z", + enabled: true, + connected: true, + connection_last_checked_at: "2026-07-16T00:00:00Z", + integration_type: "amazon_s3", + configuration: { + bucket_name: "prowler-reports", + output_directory: "output", + }, + }, + relationships: { + providers: { + data: [{ type: "providers", id: "aws-provider" }], + }, + }, + links: { + self: "/integrations/integration-1", + }, +}; + +describe("S3IntegrationForm", () => { + beforeEach(() => { + createIntegrationMock.mockReset(); + toastMock.mockReset(); + updateIntegrationMock.mockReset(); + }); + + it("should require the bucket owner account ID when it cannot derive one", async () => { + // Given + const user = userEvent.setup(); + renderS3IntegrationForm({ + providers: [createProvider("azure", "subscription-id")], + }); + + // When + await user.type(screen.getByLabelText(/Bucket name/i), "prowler-reports"); + await user.click(screen.getByRole("button", { name: "Next" })); + + // Then + expect( + await screen.findByText( + "Bucket owner account ID is required when no AWS account is selected", + ), + ).toBeVisible(); + expect( + screen.queryByLabelText("CloudFormation quick link"), + ).not.toBeInTheDocument(); + }); + + it("should derive the bucket owner account ID from the selected AWS provider", async () => { + // Given + const user = userEvent.setup(); + renderS3IntegrationForm({ + providers: [createProvider("aws", "123456789012")], + }); + + // When + await user.click( + screen.getByRole("button", { name: "Select first provider" }), + ); + await user.type(screen.getByLabelText(/Bucket name/i), "prowler-reports"); + await user.click(screen.getByRole("button", { name: "Next" })); + + // Then + expect( + await screen.findByLabelText("CloudFormation quick link"), + ).toHaveTextContent("bucketAccountId=123456789012"); + }); + + it("should not show a bucket account field that configuration updates cannot persist", () => { + // When + renderS3IntegrationForm({ + integration, + providers: [createProvider("aws", "123456789012")], + editMode: "configuration", + }); + + // Then + expect( + screen.queryByLabelText(/Bucket owner account ID/i), + ).not.toBeInTheDocument(); + }); + + it("should allow changing the bucket owner account for credential updates", () => { + // When + renderS3IntegrationForm({ + integration, + providers: [createProvider("aws", "123456789012")], + editMode: "credentials", + }); + + // Then + expect( + screen.getByLabelText(/Bucket owner account ID/i), + ).toBeInTheDocument(); + }); +}); diff --git a/ui/components/integrations/s3/s3-integration-form.tsx b/ui/components/integrations/s3/s3-integration-form.tsx index 670e3c9aae..555e7d0d1f 100644 --- a/ui/components/integrations/s3/s3-integration-form.tsx +++ b/ui/components/integrations/s3/s3-integration-form.tsx @@ -4,7 +4,8 @@ import { zodResolver } from "@hookform/resolvers/zod"; import { ArrowLeftIcon, ArrowRightIcon } from "lucide-react"; import { useSession } from "next-auth/react"; import { useState } from "react"; -import { Control, useForm } from "react-hook-form"; +import type { Control } from "react-hook-form"; +import { useForm } from "react-hook-form"; import { createIntegration, updateIntegration } from "@/actions/integrations"; import { @@ -25,13 +26,13 @@ import { import { FormButtons } from "@/components/shadcn/form/form-buttons"; import { EnhancedMultiSelect } from "@/components/shadcn/select/enhanced-multi-select"; import { getAWSCredentialsTemplateLinks } from "@/lib"; -import { AWSCredentialsRole } from "@/types"; +import type { AWSCredentialsRole } from "@/types"; +import type { IntegrationProps } from "@/types/integrations"; import { editS3IntegrationFormSchema, - IntegrationProps, s3IntegrationFormSchema, } from "@/types/integrations"; -import { ProviderProps } from "@/types/providers"; +import type { ProviderProps } from "@/types/providers"; interface S3IntegrationFormProps { integration?: IntegrationProps | null; @@ -41,6 +42,25 @@ interface S3IntegrationFormProps { editMode?: "configuration" | "credentials" | null; // null means creating new } +const getSelectedAWSAccountId = ( + selectedProviderIds: string[], + providers: ProviderProps[], +): string => { + for (const providerId of selectedProviderIds) { + const provider = providers.find(({ id }) => id === providerId); + const uid = provider?.attributes.uid; + if ( + provider?.attributes.provider === "aws" && + uid && + /^\d{12}$/.test(uid) + ) { + return uid; + } + } + + return ""; +}; + export const S3IntegrationForm = ({ integration, providers, @@ -75,6 +95,7 @@ export const S3IntegrationForm = ({ defaultValues: { integration_type: "amazon_s3" as const, bucket_name: integration?.attributes.configuration.bucket_name || "", + bucket_account_id: "", output_directory: integration?.attributes.configuration.output_directory || "output", providers: @@ -94,6 +115,14 @@ export const S3IntegrationForm = ({ }); const isLoading = form.formState.isSubmitting; + const selectedProviderIds = form.watch("providers") || []; + const bucketAccountIdOverride = form.watch("bucket_account_id")?.trim() || ""; + const derivedBucketAccountId = getSelectedAWSAccountId( + selectedProviderIds, + providers, + ); + const resolvedBucketAccountId = + bucketAccountIdOverride || derivedBucketAccountId; const handleNext = async (e: React.FormEvent) => { e.preventDefault(); @@ -103,18 +132,36 @@ export const S3IntegrationForm = ({ return; } - // Validate current step fields for creation flow + // Validate current step fields for creation flow. bucket_account_id is + // validated here, while its input is visible, so a malformed value surfaces + // its error instead of silently blocking the step 1 submit. const stepFields = currentStep === 0 - ? (["bucket_name", "output_directory", "providers"] as const) + ? ([ + "bucket_name", + "output_directory", + "providers", + "bucket_account_id", + ] as const) : // Step 1: No required fields since role_arn and external_id are optional []; const isValid = stepFields.length === 0 || (await form.trigger(stepFields)); - if (isValid) { - setCurrentStep(1); + if (!isValid) { + return; } + + if (!resolvedBucketAccountId) { + form.setError("bucket_account_id", { + message: + "Bucket owner account ID is required when no AWS account is selected", + }); + return; + } + + form.clearErrors("bucket_account_id"); + setCurrentStep(1); }; const handleBack = () => { @@ -255,6 +302,30 @@ export const S3IntegrationForm = ({ } }; + const renderBucketAccountIdField = () => ( +
+ +

+ {derivedBucketAccountId + ? `Leave empty to use selected AWS account ${derivedBucketAccountId}, or enter another bucket owner account ID.` + : "Required because the selected provider does not identify the AWS account that owns the bucket."} +

+
+ ); + const renderStepContent = () => { // If editing credentials, show only credentials form if (isEditingCredentials || currentStep === 1) { @@ -265,17 +336,21 @@ export const S3IntegrationForm = ({ externalId, bucketName, "amazon_s3", + resolvedBucketAccountId, ); return ( - } - setValue={form.setValue as any} - externalId={externalId} - templateLinks={templateLinks} - type="integrations" - integrationType="amazon_s3" - /> +
+ {isEditingCredentials && renderBucketAccountIdField()} + } + setValue={form.setValue as any} + externalId={externalId} + templateLinks={templateLinks} + type="integrations" + integrationType="amazon_s3" + /> +
); } @@ -346,6 +421,8 @@ export const S3IntegrationForm = ({ variant="bordered" isRequired /> + + {!isEditingConfig && renderBucketAccountIdField()}
); diff --git a/ui/components/layout/app-sidebar/app-sidebar-mode-sync.test.tsx b/ui/components/layout/app-sidebar/app-sidebar-mode-sync.test.tsx index 59caaae738..7eb39455b4 100644 --- a/ui/components/layout/app-sidebar/app-sidebar-mode-sync.test.tsx +++ b/ui/components/layout/app-sidebar/app-sidebar-mode-sync.test.tsx @@ -1,6 +1,8 @@ import { render } from "@testing-library/react"; import { beforeEach, describe, expect, it } from "vitest"; +import { useSidePanelStore } from "@/store/side-panel"; + import { useAppSidebarMode } from "./app-sidebar-mode-store"; import { AppSidebarModeSync } from "./app-sidebar-mode-sync"; import { APP_SIDEBAR_MODE } from "./types"; @@ -8,6 +10,7 @@ import { APP_SIDEBAR_MODE } from "./types"; describe("AppSidebarModeSync", () => { beforeEach(() => { useAppSidebarMode.setState({ mode: APP_SIDEBAR_MODE.CHAT }); + useSidePanelStore.setState({ isOpen: false }); }); it("restores the requested sidebar mode when a route mounts", () => { @@ -17,4 +20,26 @@ describe("AppSidebarModeSync", () => { // Then expect(useAppSidebarMode.getState().mode).toBe(APP_SIDEBAR_MODE.BROWSE); }); + + it("keeps the side panel open by default", () => { + // Given + useSidePanelStore.setState({ isOpen: true }); + + // When + render(); + + // Then + expect(useSidePanelStore.getState().isOpen).toBe(true); + }); + + it("closes the side panel when the full-page chat mounts", () => { + // Given + useSidePanelStore.setState({ isOpen: true }); + + // When + render(); + + // Then + expect(useSidePanelStore.getState().isOpen).toBe(false); + }); }); diff --git a/ui/components/layout/app-sidebar/app-sidebar-mode-sync.tsx b/ui/components/layout/app-sidebar/app-sidebar-mode-sync.tsx index c97d84a14d..b4d7887dfb 100644 --- a/ui/components/layout/app-sidebar/app-sidebar-mode-sync.tsx +++ b/ui/components/layout/app-sidebar/app-sidebar-mode-sync.tsx @@ -1,19 +1,29 @@ "use client"; import { useMountEffect } from "@/hooks/use-mount-effect"; +import { useSidePanelStore } from "@/store/side-panel"; import { useAppSidebarMode } from "./app-sidebar-mode-store"; import type { AppSidebarMode } from "./types"; interface AppSidebarModeSyncProps { mode: AppSidebarMode; + // The full-page chat dismisses the side panel: the chat lives in one place + // or the other, never both. + closeSidePanel?: boolean; } -export function AppSidebarModeSync({ mode }: AppSidebarModeSyncProps) { +export function AppSidebarModeSync({ + mode, + closeSidePanel = false, +}: AppSidebarModeSyncProps) { const setMode = useAppSidebarMode((state) => state.setMode); useMountEffect(() => { setMode(mode); + if (closeSidePanel) { + useSidePanelStore.getState().closePanel(); + } }); return null; diff --git a/ui/components/layout/app-sidebar/mobile-app-sidebar.test.tsx b/ui/components/layout/app-sidebar/mobile-app-sidebar.test.tsx index 75aae87f5a..343a6aef72 100644 --- a/ui/components/layout/app-sidebar/mobile-app-sidebar.test.tsx +++ b/ui/components/layout/app-sidebar/mobile-app-sidebar.test.tsx @@ -45,6 +45,17 @@ describe("MobileAppSidebar", () => { expect(openButton).toHaveFocus(); }); + it("hides the trigger based on the viewport, not the narrowed content", () => { + // Given / When + render(); + + // Then: lg: is a container query inside
; the side panel squeezing + // the page must not surface the mobile menu on desktop. + expect(screen.getByRole("button", { name: "Open menu" })).toHaveClass( + "min-[64rem]:hidden", + ); + }); + it("closes after selecting an item from the shared sidebar content", async () => { // Given const user = userEvent.setup(); diff --git a/ui/components/layout/app-sidebar/mobile-app-sidebar.tsx b/ui/components/layout/app-sidebar/mobile-app-sidebar.tsx index abb6fbbfe7..0a3f9afdeb 100644 --- a/ui/components/layout/app-sidebar/mobile-app-sidebar.tsx +++ b/ui/components/layout/app-sidebar/mobile-app-sidebar.tsx @@ -35,7 +35,9 @@ export function MobileAppSidebar() { variant="bare" size="icon-sm" aria-label="Open menu" - className={cn("lg:hidden", open && "invisible")} + // min-[64rem] (not lg:): inside
, lg is a container query and + // the side panel squeezing the page would surface the mobile menu. + className={cn("min-[64rem]:hidden", open && "invisible")} >