diff --git a/prowler/__main__.py b/prowler/__main__.py index 4435b4185a..cbf8380d58 100644 --- a/prowler/__main__.py +++ b/prowler/__main__.py @@ -10,7 +10,6 @@ from colorama import Fore, Style from colorama import init as colorama_init from prowler.config.config import ( - EXTERNAL_TOOL_PROVIDERS, cloud_api_base_url, csv_file_suffix, get_available_compliance_frameworks, @@ -235,7 +234,7 @@ def prowler(): logger.debug("Loading compliance frameworks from .json files") # Skip compliance frameworks for external-tool providers - if provider not in EXTERNAL_TOOL_PROVIDERS: + if not Provider.is_tool_wrapper_provider(provider): bulk_compliance_frameworks = Compliance.get_bulk(provider) # Complete checks metadata with the compliance framework specification bulk_checks_metadata = update_checks_metadata_with_compliance( @@ -300,7 +299,7 @@ def prowler(): sys.exit() # Skip service and check loading for external-tool providers - if provider not in EXTERNAL_TOOL_PROVIDERS: + if not Provider.is_tool_wrapper_provider(provider): # Import custom checks from folder if checks_folder: custom_checks = parse_checks_from_folder(global_provider, checks_folder) @@ -423,7 +422,7 @@ def prowler(): # Execute checks findings = [] - if provider in EXTERNAL_TOOL_PROVIDERS: + if Provider.is_tool_wrapper_provider(provider): # For external-tool providers, run the scan directly if provider == "llm": diff --git a/prowler/lib/check/models.py b/prowler/lib/check/models.py index fea16a19d2..2864d04b8a 100644 --- a/prowler/lib/check/models.py +++ b/prowler/lib/check/models.py @@ -11,7 +11,6 @@ from typing import Any, Dict, Optional, Set from pydantic.v1 import BaseModel, Field, ValidationError, validator from pydantic.v1.error_wrappers import ErrorWrapper -from prowler.config.config import EXTERNAL_TOOL_PROVIDERS from prowler.lib.check.compliance_models import Compliance from prowler.lib.check.utils import recover_checks_from_provider from prowler.lib.logger import logger @@ -256,7 +255,7 @@ class CheckMetadata(BaseModel): ) if ( value_lower not in VALID_CATEGORIES - and values.get("Provider") not in EXTERNAL_TOOL_PROVIDERS + and not ProviderABC.is_tool_wrapper_provider(values.get("Provider")) ): raise ValueError( f"Invalid category: '{value_lower}'. Must be one of: {', '.join(sorted(VALID_CATEGORIES))}." @@ -285,7 +284,9 @@ class CheckMetadata(BaseModel): raise ValueError("ServiceName must be a non-empty string") check_id = values.get("CheckID") - if check_id and values.get("Provider") not in EXTERNAL_TOOL_PROVIDERS: + if check_id and not ProviderABC.is_tool_wrapper_provider( + values.get("Provider") + ): service_from_check_id = check_id.split("_")[0] if service_name != service_from_check_id: raise ValueError( @@ -301,7 +302,9 @@ class CheckMetadata(BaseModel): if not check_id: raise ValueError("CheckID must be a non-empty string") - if check_id and values.get("Provider") not in EXTERNAL_TOOL_PROVIDERS: + if check_id and not ProviderABC.is_tool_wrapper_provider( + values.get("Provider") + ): if "-" in check_id: raise ValueError( f"CheckID {check_id} contains a hyphen, which is not allowed" @@ -311,7 +314,7 @@ class CheckMetadata(BaseModel): @validator("CheckTitle", pre=True, always=True) def validate_check_title(cls, check_title, values): - if values.get("Provider") not in EXTERNAL_TOOL_PROVIDERS: + if not ProviderABC.is_tool_wrapper_provider(values.get("Provider")): if len(check_title) > 150: raise ValueError( f"CheckTitle must not exceed 150 characters, got {len(check_title)} characters" @@ -324,13 +327,15 @@ class CheckMetadata(BaseModel): @validator("RelatedUrl", pre=True, always=True) def validate_related_url(cls, related_url, values): - if related_url and values.get("Provider") not in EXTERNAL_TOOL_PROVIDERS: + if related_url and not ProviderABC.is_tool_wrapper_provider( + values.get("Provider") + ): raise ValueError("RelatedUrl must be empty. This field is deprecated.") return related_url @validator("Remediation") def validate_recommendation_url(cls, remediation, values): - if values.get("Provider") not in EXTERNAL_TOOL_PROVIDERS: + if not ProviderABC.is_tool_wrapper_provider(values.get("Provider")): url = remediation.Recommendation.Url if url and not url.startswith("https://hub.prowler.com/"): raise ValueError( @@ -343,7 +348,7 @@ class CheckMetadata(BaseModel): provider = values.get("Provider", "").lower() # Non-AWS providers must have an empty CheckType list - if provider != "aws" and provider not in EXTERNAL_TOOL_PROVIDERS: + if provider != "aws" and not ProviderABC.is_tool_wrapper_provider(provider): if check_type: raise ValueError( f"CheckType must be empty for non-AWS providers. Got {check_type} for provider '{provider}'." @@ -369,7 +374,7 @@ class CheckMetadata(BaseModel): @validator("Description", pre=True, always=True) def validate_description(cls, description, values): - if values.get("Provider") not in EXTERNAL_TOOL_PROVIDERS: + if not ProviderABC.is_tool_wrapper_provider(values.get("Provider")): if len(description) > 400: raise ValueError( f"Description must not exceed 400 characters, got {len(description)} characters" @@ -378,7 +383,7 @@ class CheckMetadata(BaseModel): @validator("Risk", pre=True, always=True) def validate_risk(cls, risk, values): - if values.get("Provider") not in EXTERNAL_TOOL_PROVIDERS: + if not ProviderABC.is_tool_wrapper_provider(values.get("Provider")): if len(risk) > 400: raise ValueError( f"Risk must not exceed 400 characters, got {len(risk)} characters" diff --git a/prowler/providers/common/provider.py b/prowler/providers/common/provider.py index fb400e3489..eeea0d2e5a 100644 --- a/prowler/providers/common/provider.py +++ b/prowler/providers/common/provider.py @@ -8,7 +8,10 @@ from argparse import Namespace from importlib import import_module from typing import Any, Optional -from prowler.config.config import load_and_validate_config_file +from prowler.config.config import ( + EXTERNAL_TOOL_PROVIDERS, + load_and_validate_config_file, +) from prowler.lib.logger import logger from prowler.lib.mutelist.mutelist import Mutelist @@ -223,10 +226,13 @@ class Provider(ABC): f"{self.__class__.__name__} has not implemented display_compliance_table()" ) - @property - def is_external_tool_provider(self) -> bool: - """True for providers that delegate scanning to an external tool.""" - return False + # Class-level flag: True for providers that delegate scanning to an external + # tool (e.g. Trivy, promptfoo) and bypass standard check/service loading and + # metadata validation. Subclasses override as `is_external_tool_provider = True`. + # Kept as a class attribute (not a property) so it can be read from the class + # without instantiation — the metadata validators in lib.check.models need to + # decide whether to relax validation before any provider instance exists. + is_external_tool_provider: bool = False # --- End dynamic provider contract methods --- @@ -532,6 +538,21 @@ class Provider(ABC): providers.add(ep.name) return sorted(providers) + @staticmethod + def is_tool_wrapper_provider(provider: str) -> bool: + """Return True if the provider delegates scanning to an external tool. + + Combines the built-in EXTERNAL_TOOL_PROVIDERS frozenset (fast path for + iac/llm/image) with the `is_external_tool_provider` class attribute of + external plug-in providers registered via entry points. This is the + single source of truth consulted by the execution flow and the + CheckMetadata validators. + """ + if provider in EXTERNAL_TOOL_PROVIDERS: + return True + ep_cls = Provider._load_ep_provider(provider) + return bool(ep_cls and getattr(ep_cls, "is_external_tool_provider", False)) + @staticmethod def _load_ep_provider(name: str): """Load an external provider class from entry points, with cache.""" diff --git a/tests/providers/external/test_dynamic_provider_loading.py b/tests/providers/external/test_dynamic_provider_loading.py index 6aaa3c51b6..810c9053ef 100644 --- a/tests/providers/external/test_dynamic_provider_loading.py +++ b/tests/providers/external/test_dynamic_provider_loading.py @@ -116,6 +116,35 @@ class FakeExternalProvider(Provider): pass +class FakeToolWrapperProvider(Provider): + """External provider that declares itself a tool wrapper.""" + + _type = "faketoolwrapper" + is_external_tool_provider = True + + @property + def type(self): + return self._type + + @property + def identity(self): + return MagicMock() + + @property + def session(self): + return MagicMock() + + @property + def audit_config(self): + return {} + + def setup_session(self): + return MagicMock() + + def print_credentials(self): + pass + + class FakeProviderNoHelpText(Provider): """Provider without _cli_help_text.""" @@ -264,6 +293,57 @@ class TestProviderDiscovery: assert help_text["nohelptext"] == "" + +class TestIsToolWrapperProvider: + """Tests for Provider.is_tool_wrapper_provider — the helper that combines the + built-in EXTERNAL_TOOL_PROVIDERS frozenset with the is_external_tool_provider + class attribute of entry-point providers.""" + + def test_returns_true_for_builtin_tool_wrappers(self): + # iac/llm/image are in the EXTERNAL_TOOL_PROVIDERS frozenset (fast path) + assert Provider.is_tool_wrapper_provider("iac") is True + assert Provider.is_tool_wrapper_provider("llm") is True + assert Provider.is_tool_wrapper_provider("image") is True + + def test_returns_false_for_regular_builtin_providers(self): + # Regular built-ins must not be classified as tool wrappers + assert Provider.is_tool_wrapper_provider("aws") is False + assert Provider.is_tool_wrapper_provider("gcp") is False + assert Provider.is_tool_wrapper_provider("github") is False + + @patch("prowler.providers.common.provider.importlib.metadata.entry_points") + def test_returns_true_for_external_provider_declaring_flag(self, mock_ep): + # External plugin explicitly declares is_external_tool_provider = True + mock_ep.return_value = [ + _make_entry_point("faketoolwrapper", "pkg:Cls", "prowler.providers"), + ] + mock_ep.return_value[0].load.return_value = FakeToolWrapperProvider + + assert Provider.is_tool_wrapper_provider("faketoolwrapper") is True + + @patch("prowler.providers.common.provider.importlib.metadata.entry_points") + def test_returns_false_for_external_provider_without_flag(self, mock_ep): + # External plugin without the flag (default False) is treated as regular + mock_ep.return_value = [ + _make_entry_point("fakeexternal", "pkg:Cls", "prowler.providers"), + ] + mock_ep.return_value[0].load.return_value = FakeExternalProvider + + assert Provider.is_tool_wrapper_provider("fakeexternal") is False + + @patch("prowler.providers.common.provider.importlib.metadata.entry_points") + def test_returns_false_for_unknown_provider(self, mock_ep): + mock_ep.return_value = [] + + assert Provider.is_tool_wrapper_provider("does-not-exist") is False + + @patch("prowler.providers.common.provider.importlib.metadata.entry_points") + def test_returns_false_for_none_provider(self, mock_ep): + # Pydantic validators may pass None when values.get("Provider") is missing + mock_ep.return_value = [] + + assert Provider.is_tool_wrapper_provider(None) is False + @patch("prowler.providers.common.provider.importlib.metadata.entry_points") def test_load_ep_provider_handles_load_exception(self, mock_ep): """_load_ep_provider returns None when ep.load() raises."""