diff --git a/prowler/prowler/_core/prowler_client/__init__.py b/prowler/prowler/_core/prowler_client/__init__.py index 7afa21f4..897dabb9 100644 --- a/prowler/prowler/_core/prowler_client/__init__.py +++ b/prowler/prowler/_core/prowler_client/__init__.py @@ -7,7 +7,7 @@ ProwlerClient, ProwlerClientConsumedError, ) -from .contracts import AwsServiceSelector +from .contracts import AwsServiceSelector, AzureServiceSelector, ServiceSelector from .credentials import ( CredentialCleanupError, TemporaryCredentialLease, @@ -39,9 +39,11 @@ "OutputWorkspaceCleanupError", "OutputWorkspacePreparationError", "AwsServiceSelector", + "AzureServiceSelector", "ProwlerClient", "ProwlerClientConsumedError", "ProwlerClientFactory", + "ServiceSelector", "TemporaryCredentialLease", "TemporaryCredentialLeaseFactory", "TemporaryOutputWorkspace", diff --git a/prowler/prowler/_core/prowler_client/client.py b/prowler/prowler/_core/prowler_client/client.py index a65cc245..8e31c3eb 100644 --- a/prowler/prowler/_core/prowler_client/client.py +++ b/prowler/prowler/_core/prowler_client/client.py @@ -12,12 +12,16 @@ ValidatedCommandRequest, ) from prowler.models.configs.config_loader import ProwlerConfig -from prowler.models.provider_inputs import AwsProviderInput, ProviderInput +from prowler.models.provider_inputs import ( + AwsProviderInput, + AzureProviderInput, + ProviderInput, +) from .contracts import ( - AwsServiceSelector, CliEnginePort, OutputWorkspaceFactoryPort, + ServiceSelector, ) from .credentials import CredentialCleanupError from .output_workspace import ( @@ -101,7 +105,7 @@ def run( self, check_filters: Sequence[str] = (), *, - service_selector: AwsServiceSelector | None = None, + service_selector: ServiceSelector | None = None, ) -> CommandResult: """Run one assessment and capture its controlled OCSF artifact.""" with self._consumption_lock: @@ -120,12 +124,22 @@ def run( filters = tuple(check_filters) if any(not isinstance(item, str) or not item.strip() for item in filters): raise ValueError("check filters must be nonblank strings") - if service_selector not in (None, "iam", "s3", "ec2"): - raise ValueError("unsupported AWS service selector") - if service_selector is not None and not isinstance( + if service_selector not in (None, "iam", "s3", "ec2", "storage"): + raise ValueError("unsupported service selector") + if service_selector in ("s3", "ec2") and not isinstance( provider, AwsProviderInput ): raise ValueError("AWS service selector requires an AWS provider") + if service_selector == "storage" and not isinstance( + provider, AzureProviderInput + ): + raise ValueError("Azure service selector requires an Azure provider") + if service_selector == "iam" and not isinstance( + provider, AwsProviderInput | AzureProviderInput + ): + raise ValueError( + "IAM service selector requires an AWS or Azure provider" + ) _safe_log(logging.INFO, "Preparing Prowler output workspace") try: diff --git a/prowler/prowler/_core/prowler_client/contracts.py b/prowler/prowler/_core/prowler_client/contracts.py index 70842d82..d4a2b125 100644 --- a/prowler/prowler/_core/prowler_client/contracts.py +++ b/prowler/prowler/_core/prowler_client/contracts.py @@ -10,6 +10,8 @@ from prowler._core.cli_engine.contracts import EnvironmentValue AwsServiceSelector = Literal["iam", "s3", "ec2"] +AzureServiceSelector = Literal["iam", "storage"] +ServiceSelector = AwsServiceSelector | AzureServiceSelector class CliEnginePort(Protocol): diff --git a/prowler/prowler/_core/prowler_client/factory.py b/prowler/prowler/_core/prowler_client/factory.py index 4684960c..2133cca2 100644 --- a/prowler/prowler/_core/prowler_client/factory.py +++ b/prowler/prowler/_core/prowler_client/factory.py @@ -9,10 +9,10 @@ from .client import ProwlerClient from .contracts import ( - AwsServiceSelector, CliEngineFactoryPort, CredentialLeaseFactoryPort, OutputWorkspaceFactoryPort, + ServiceSelector, ) from .credentials import TemporaryCredentialLeaseFactory from .output_workspace import TemporaryOutputWorkspaceFactory @@ -47,7 +47,7 @@ def run( provider: ProviderInput, *, check_filters: Sequence[str] = (), - service_selector: AwsServiceSelector | None = None, + service_selector: ServiceSelector | None = None, ) -> CommandResult: """Create a client and synchronously run one assessment.""" return self.create(config, provider).run( diff --git a/prowler/prowler/contracts/__init__.py b/prowler/prowler/contracts/__init__.py index 58756356..64f50588 100644 --- a/prowler/prowler/contracts/__init__.py +++ b/prowler/prowler/contracts/__init__.py @@ -7,7 +7,12 @@ AwsS3Contract, AwsServiceContract, ) -from .azure import AzureBaseContract +from .azure import ( + AzureBaseContract, + AzureIamContract, + AzureServiceContract, + AzureStorageContract, +) from .base import ( BaseProwlerContract, ContractExecutionOutcome, @@ -35,6 +40,9 @@ "AwsS3Contract", "AwsServiceContract", "AzureBaseContract", + "AzureIamContract", + "AzureServiceContract", + "AzureStorageContract", "GcpBaseContract", "KubernetesBaseContract", "ContractDispatcher", diff --git a/prowler/prowler/contracts/aws.py b/prowler/prowler/contracts/aws.py index 202b0e88..8a45e59e 100644 --- a/prowler/prowler/contracts/aws.py +++ b/prowler/prowler/contracts/aws.py @@ -1,6 +1,5 @@ """Executable complete-scope and service-specific AWS contracts.""" -from collections.abc import Sequence from dataclasses import replace from typing import ClassVar @@ -9,7 +8,6 @@ from prowler.models.findings import ( OcsfDecodeError, OcsfMappingError, - OpenAevFinding, map_command_result_with_evidence, ) from prowler.models.provider_inputs import AwsProviderInput, ProviderInput @@ -28,17 +26,6 @@ class AwsBaseContract(BaseProwlerContract): label = "Prowler AWS" check_filters = () - @staticmethod - def _aws_findings( - findings: Sequence[OpenAevFinding], - ) -> tuple[OpenAevFinding, ...]: - """Retain ordered AWS findings with one canonical provider spelling.""" - return tuple( - finding.model_copy(update={"cloud_provider": "aws"}) - for finding in findings - if finding.cloud_provider.casefold() == "aws" - ) - def execute( self, config: ProwlerConfig, provider: ProviderInput ) -> ContractExecutionOutcome: @@ -46,7 +33,7 @@ def execute( outcome = super().execute(config, provider) if outcome.error is not None or outcome.command_result.return_code != 0: return outcome - return replace(outcome, findings=self._aws_findings(outcome.findings)) + return replace(outcome, findings=self._provider_findings(outcome.findings)) class AwsServiceContract(AwsBaseContract): @@ -88,7 +75,7 @@ def execute( raw_output_bytes=mapping.raw_output_bytes, raw_preview=mapping.raw_preview, ) - return replace(outcome, findings=self._aws_findings(outcome.findings)) + return replace(outcome, findings=self._provider_findings(outcome.findings)) class AwsIamContract(AwsServiceContract): diff --git a/prowler/prowler/contracts/azure.py b/prowler/prowler/contracts/azure.py index df39df4b..53dd8d0f 100644 --- a/prowler/prowler/contracts/azure.py +++ b/prowler/prowler/contracts/azure.py @@ -1,12 +1,18 @@ -"""Executable CHK.008 complete-scope Azure base contract.""" +"""Executable complete-scope and service-specific Azure contracts.""" from dataclasses import replace from typing import ClassVar +from prowler._core.prowler_client import AzureServiceSelector from prowler.models.configs.config_loader import ProwlerConfig -from prowler.models.provider_inputs import ProviderInput +from prowler.models.findings import ( + OcsfDecodeError, + OcsfMappingError, + map_command_result_with_evidence, +) +from prowler.models.provider_inputs import AzureProviderInput, ProviderInput -from .base import BaseProwlerContract, ContractExecutionOutcome +from .base import BaseProwlerContract, ContractExecutionOutcome, RouteFamily class AzureBaseContract(BaseProwlerContract): @@ -16,7 +22,7 @@ class AzureBaseContract(BaseProwlerContract): external_id: ClassVar[str] = "prowler:azure" route_name: ClassVar[str] = "azure" provider = "azure" - family = "base" + family: ClassVar[RouteFamily] = "base" label = "Prowler Azure" check_filters = () @@ -27,9 +33,66 @@ def execute( outcome = super().execute(config, provider) if outcome.error is not None or outcome.command_result.return_code != 0: return outcome - findings = tuple( - finding.model_copy(update={"cloud_provider": "azure"}) - for finding in outcome.findings - if finding.cloud_provider.casefold() == "azure" + return replace(outcome, findings=self._provider_findings(outcome.findings)) + + +class AzureServiceContract(AzureBaseContract): + """Execute one route-owned Azure service selector through the CHK.004 seam.""" + + family = "service" + service_selector: ClassVar[AzureServiceSelector] + + def safe_request_info(self, provider: ProviderInput | None) -> dict[str, object]: + """Identify the service through safe route metadata, never form input.""" + info = super().safe_request_info(provider) + info["filters"] = f"service={self.service_selector}" + return info + + def execute( + self, config: ProwlerConfig, provider: ProviderInput + ) -> ContractExecutionOutcome: + """Reject invalid route combinations before one service-specific client call.""" + if self.service_selector not in ("iam", "storage"): + raise ValueError("unsupported Azure service selector") + if not isinstance(provider, AzureProviderInput): + raise ValueError("Azure service selector requires an Azure provider") + result = self._client_factory.run( + config, + provider, + check_filters=self.check_filters, + service_selector=self.service_selector, ) - return replace(outcome, findings=findings) + if result.error is not None or result.return_code != 0: + return ContractExecutionOutcome(command_result=result, error=result.error) + try: + mapping = map_command_result_with_evidence(result) + except (OcsfDecodeError, OcsfMappingError) as error: + return ContractExecutionOutcome(command_result=result, error=error) + outcome = ContractExecutionOutcome( + command_result=result, + findings=mapping.findings, + raw_record_count=mapping.raw_record_count, + raw_output_bytes=mapping.raw_output_bytes, + raw_preview=mapping.raw_preview, + ) + return replace(outcome, findings=self._provider_findings(outcome.findings)) + + +class AzureIamContract(AzureServiceContract): + """Run only Prowler Azure IAM checks.""" + + contract_id = "b559fad9-b928-5bfb-bf31-10479e0bf8c1" + external_id = "prowler:azure/iam" + route_name = "azure/iam" + label = "Prowler Azure IAM" + service_selector = "iam" + + +class AzureStorageContract(AzureServiceContract): + """Run only Prowler Azure Storage checks.""" + + contract_id = "f056c4d2-6ee8-5554-8f23-b67cf0ff4e58" + external_id = "prowler:azure/storage" + route_name = "azure/storage" + label = "Prowler Azure Storage" + service_selector = "storage" diff --git a/prowler/prowler/contracts/base.py b/prowler/prowler/contracts/base.py index 47ce02fc..8499ea20 100644 --- a/prowler/prowler/contracts/base.py +++ b/prowler/prowler/contracts/base.py @@ -21,7 +21,7 @@ ) from prowler._core.cli_engine import CommandResult -from prowler._core.prowler_client import AwsServiceSelector, ProwlerClientFactory +from prowler._core.prowler_client import ProwlerClientFactory, ServiceSelector from prowler.models.configs.config_loader import ProwlerConfig from prowler.models.findings import ( OcsfDecodeError, @@ -65,7 +65,7 @@ def run( provider: ProviderInput, *, check_filters: Sequence[str] = (), - service_selector: AwsServiceSelector | None = None, + service_selector: ServiceSelector | None = None, ) -> CommandResult: """Run one assessment and return the exact command result.""" @@ -208,6 +208,16 @@ def output_payload( ], } + def _provider_findings( + self, findings: Sequence[OpenAevFinding] + ) -> tuple[OpenAevFinding, ...]: + """Retain source order and normalize findings for this contract provider.""" + return tuple( + finding.model_copy(update={"cloud_provider": self.provider}) + for finding in findings + if finding.cloud_provider.casefold() == self.provider + ) + @staticmethod def output_trace_config() -> dict[str, object]: """Return the common flattened-field trace contract.""" diff --git a/prowler/prowler/contracts/gcp.py b/prowler/prowler/contracts/gcp.py index 9636cbe2..31c7c042 100644 --- a/prowler/prowler/contracts/gcp.py +++ b/prowler/prowler/contracts/gcp.py @@ -27,9 +27,4 @@ def execute( outcome = super().execute(config, provider) if outcome.error is not None or outcome.command_result.return_code != 0: return outcome - findings = tuple( - finding.model_copy(update={"cloud_provider": "gcp"}) - for finding in outcome.findings - if finding.cloud_provider.casefold() == "gcp" - ) - return replace(outcome, findings=findings) + return replace(outcome, findings=self._provider_findings(outcome.findings)) diff --git a/prowler/prowler/contracts/kubernetes.py b/prowler/prowler/contracts/kubernetes.py index 310064a0..94a794ff 100644 --- a/prowler/prowler/contracts/kubernetes.py +++ b/prowler/prowler/contracts/kubernetes.py @@ -1,11 +1,9 @@ """Executable CHK.010 complete-scope Kubernetes base contract.""" -from collections.abc import Sequence from dataclasses import replace from typing import ClassVar from prowler.models.configs.config_loader import ProwlerConfig -from prowler.models.findings import OpenAevFinding from prowler.models.provider_inputs import ProviderInput from .base import BaseProwlerContract, ContractExecutionOutcome @@ -22,17 +20,6 @@ class KubernetesBaseContract(BaseProwlerContract): label = "Prowler Kubernetes" check_filters = () - @staticmethod - def _provider_findings( - findings: Sequence[OpenAevFinding], - ) -> tuple[OpenAevFinding, ...]: - """Retain source order while normalizing matching provider labels.""" - return tuple( - finding.model_copy(update={"cloud_provider": "kubernetes"}) - for finding in findings - if finding.cloud_provider.casefold() == "kubernetes" - ) - def execute( self, config: ProwlerConfig, provider: ProviderInput ) -> ContractExecutionOutcome: diff --git a/prowler/prowler/contracts/registry.py b/prowler/prowler/contracts/registry.py index d4174526..5b939c64 100644 --- a/prowler/prowler/contracts/registry.py +++ b/prowler/prowler/contracts/registry.py @@ -8,7 +8,7 @@ from pyoaev.contracts.contract_config import prepare_contracts from .aws import AwsBaseContract, AwsEc2Contract, AwsIamContract, AwsS3Contract -from .azure import AzureBaseContract +from .azure import AzureBaseContract, AzureIamContract, AzureStorageContract from .base import BaseProwlerContract from .catalog import ROUTE_CATALOG from .gcp import GcpBaseContract @@ -90,5 +90,7 @@ def contracts(self) -> list[dict[str, object]]: AwsIamContract, AwsS3Contract, AwsEc2Contract, + AzureIamContract, + AzureStorageContract, ) ) diff --git a/prowler/tests/behaviour/chk001_catalog_scaffold/test_chk001_catalog_scaffold_bdd.py b/prowler/tests/behaviour/chk001_catalog_scaffold/test_chk001_catalog_scaffold_bdd.py index c74c5456..a503f4cb 100644 --- a/prowler/tests/behaviour/chk001_catalog_scaffold/test_chk001_catalog_scaffold_bdd.py +++ b/prowler/tests/behaviour/chk001_catalog_scaffold/test_chk001_catalog_scaffold_bdd.py @@ -86,6 +86,8 @@ def _then_base_contracts_are_registered(config: ConfigLoader, helper: Mock) -> N str(stable_contract_id("aws/iam")), str(stable_contract_id("aws/s3")), str(stable_contract_id("aws/ec2")), + str(stable_contract_id("azure/iam")), + str(stable_contract_id("azure/storage")), ] callback = helper.listen.call_args.kwargs["message_callback"] assert callable(callback) diff --git a/prowler/tests/behaviour/chk007_base_aws_provider/test_chk007_base_aws_provider_bdd.py b/prowler/tests/behaviour/chk007_base_aws_provider/test_chk007_base_aws_provider_bdd.py index f9cda7ea..e5c94257 100644 --- a/prowler/tests/behaviour/chk007_base_aws_provider/test_chk007_base_aws_provider_bdd.py +++ b/prowler/tests/behaviour/chk007_base_aws_provider/test_chk007_base_aws_provider_bdd.py @@ -98,6 +98,8 @@ def test_default_registration_identity_fields_and_outputs() -> None: str(stable_contract_id("aws/iam")), str(stable_contract_id("aws/s3")), str(stable_contract_id("aws/ec2")), + str(stable_contract_id("azure/iam")), + str(stable_contract_id("azure/storage")), ] assert UUID(serialized[0]["contract_id"]) == expected_id assert expected_id.version == 5 diff --git a/prowler/tests/behaviour/chk008_base_azure_provider/test_chk008_base_azure_provider_bdd.py b/prowler/tests/behaviour/chk008_base_azure_provider/test_chk008_base_azure_provider_bdd.py index 33c17907..63af92e4 100644 --- a/prowler/tests/behaviour/chk008_base_azure_provider/test_chk008_base_azure_provider_bdd.py +++ b/prowler/tests/behaviour/chk008_base_azure_provider/test_chk008_base_azure_provider_bdd.py @@ -85,11 +85,11 @@ def run( def test_default_registration_identity_fields_and_outputs() -> None: - """Azure remains second as the canonical registry grows through CHK.011.""" + """Azure remains second as the canonical registry grows through CHK.012.""" serialized = DEFAULT_PROWLER_CONTRACTS.contracts() expected_id = stable_contract_id("azure") - assert len(serialized) == 7 + assert len(serialized) == 9 assert [item["contract_id"] for item in serialized] == [ str(stable_contract_id("aws")), str(expected_id), @@ -98,6 +98,8 @@ def test_default_registration_identity_fields_and_outputs() -> None: str(stable_contract_id("aws/iam")), str(stable_contract_id("aws/s3")), str(stable_contract_id("aws/ec2")), + str(stable_contract_id("azure/iam")), + str(stable_contract_id("azure/storage")), ] assert UUID(serialized[1]["contract_id"]) == expected_id assert expected_id.version == 5 diff --git a/prowler/tests/behaviour/chk009_base_gcp_provider/test_chk009_base_gcp_provider_bdd.py b/prowler/tests/behaviour/chk009_base_gcp_provider/test_chk009_base_gcp_provider_bdd.py index 5154a8f6..88046cac 100644 --- a/prowler/tests/behaviour/chk009_base_gcp_provider/test_chk009_base_gcp_provider_bdd.py +++ b/prowler/tests/behaviour/chk009_base_gcp_provider/test_chk009_base_gcp_provider_bdd.py @@ -89,11 +89,11 @@ def run( def test_default_registration_identity_fields_and_outputs() -> None: - """GCP remains third as the canonical registry grows through CHK.011.""" + """GCP remains third as the canonical registry grows through CHK.012.""" serialized = DEFAULT_PROWLER_CONTRACTS.contracts() expected_id = stable_contract_id("gcp") - assert len(serialized) == 7 + assert len(serialized) == 9 assert [item["contract_id"] for item in serialized] == [ str(stable_contract_id("aws")), str(stable_contract_id("azure")), @@ -102,6 +102,8 @@ def test_default_registration_identity_fields_and_outputs() -> None: str(stable_contract_id("aws/iam")), str(stable_contract_id("aws/s3")), str(stable_contract_id("aws/ec2")), + str(stable_contract_id("azure/iam")), + str(stable_contract_id("azure/storage")), ] assert UUID(serialized[2]["contract_id"]) == expected_id assert expected_id.version == 5 diff --git a/prowler/tests/behaviour/chk010_base_kubernetes_provider/test_chk010_base_kubernetes_provider_bdd.py b/prowler/tests/behaviour/chk010_base_kubernetes_provider/test_chk010_base_kubernetes_provider_bdd.py index ac09560c..f9bebedd 100644 --- a/prowler/tests/behaviour/chk010_base_kubernetes_provider/test_chk010_base_kubernetes_provider_bdd.py +++ b/prowler/tests/behaviour/chk010_base_kubernetes_provider/test_chk010_base_kubernetes_provider_bdd.py @@ -89,11 +89,11 @@ def run( def test_default_registration_identity_fields_and_outputs() -> None: - """The four base routes remain first in the CHK.011 executable surface.""" + """The four base routes remain first in the CHK.012 executable surface.""" serialized = DEFAULT_PROWLER_CONTRACTS.contracts() expected_id = stable_contract_id("kubernetes") - assert len(serialized) == 7 + assert len(serialized) == 9 assert [item["contract_id"] for item in serialized] == [ str(stable_contract_id("aws")), str(stable_contract_id("azure")), @@ -102,6 +102,8 @@ def test_default_registration_identity_fields_and_outputs() -> None: str(stable_contract_id("aws/iam")), str(stable_contract_id("aws/s3")), str(stable_contract_id("aws/ec2")), + str(stable_contract_id("azure/iam")), + str(stable_contract_id("azure/storage")), ] assert UUID(serialized[3]["contract_id"]) == expected_id assert expected_id.version == 5 diff --git a/prowler/tests/behaviour/chk011_service_aws_assessment/test_chk011_service_aws_assessment_bdd.py b/prowler/tests/behaviour/chk011_service_aws_assessment/test_chk011_service_aws_assessment_bdd.py index 4aea98da..85149fa9 100644 --- a/prowler/tests/behaviour/chk011_service_aws_assessment/test_chk011_service_aws_assessment_bdd.py +++ b/prowler/tests/behaviour/chk011_service_aws_assessment/test_chk011_service_aws_assessment_bdd.py @@ -95,15 +95,23 @@ def test_route_selects_exact_service_once( assert factory.calls[0][2:] == ((), service) -def test_registry_has_exact_seven_canonical_contracts_without_selector_fields() -> None: +def test_registry_has_exact_nine_canonical_contracts_without_selector_fields() -> None: """The public surface is ordered, stable, labelled, and not user-selectable.""" serialized = DEFAULT_PROWLER_CONTRACTS.contracts() - routes = ("aws", "azure", "gcp", "kubernetes", *(item[0] for item in _ROUTES)) + routes = ( + "aws", + "azure", + "gcp", + "kubernetes", + *(item[0] for item in _ROUTES), + "azure/iam", + "azure/storage", + ) assert [item["contract_id"] for item in serialized] == [ str(stable_contract_id(route)) for route in routes ] - for item, (route, service) in zip(serialized[4:], _ROUTES, strict=True): + for item, (route, service) in zip(serialized[4:7], _ROUTES, strict=True): content = json.loads(item["contract_content"]) assert service.upper() in content["label"]["en"] assert tuple(field["key"] for field in content["fields"]) == ( @@ -570,7 +578,7 @@ def test_existing_check_filter_argv_is_unchanged( assert "--services" not in engine.requests[0].arguments -def test_service_selector_rejects_non_aws_before_engine() -> None: +def test_aws_service_selector_rejects_azure_before_engine() -> None: """Provider/selector combinations are validated before any CLI call.""" engine = _Engine(b"[]") factory = ProwlerClientFactory(_EngineFactory(engine), _NoLeaseFactory()) @@ -584,6 +592,6 @@ def test_service_selector_rejects_non_aws_before_engine() -> None: ) with pytest.raises(ValueError, match="AWS service selector"): - factory.run(ProwlerConfig(), provider, service_selector="iam") + factory.run(ProwlerConfig(), provider, service_selector="s3") assert engine.requests == [] diff --git a/prowler/tests/behaviour/chk012_service_azure_assessment/__init__.py b/prowler/tests/behaviour/chk012_service_azure_assessment/__init__.py new file mode 100644 index 00000000..84b94d62 --- /dev/null +++ b/prowler/tests/behaviour/chk012_service_azure_assessment/__init__.py @@ -0,0 +1 @@ +"""CHK.012 Azure service assessment behavior package.""" diff --git a/prowler/tests/behaviour/chk012_service_azure_assessment/chk012_service_azure_assessment.feature b/prowler/tests/behaviour/chk012_service_azure_assessment/chk012_service_azure_assessment.feature new file mode 100644 index 00000000..3623afb6 --- /dev/null +++ b/prowler/tests/behaviour/chk012_service_azure_assessment/chk012_service_azure_assessment.feature @@ -0,0 +1,57 @@ +Feature: CHK.012 service-specific Azure assessments + OpenAEV operators can run one supported Azure service assessment without changing + the existing Prowler check-filter behavior or exposing routing controls as input. + + Scenario Outline: A canonical Azure service route selects exactly one Prowler service + Given the executable Prowler registry contains the canonical route "" + When a valid Azure assessment is executed through "" + Then the Prowler client is called exactly once with service "" + And no check filter is supplied + + Examples: + | route | service | + | azure/iam | iam | + | azure/storage | storage | + + Scenario: The nine executable contracts have stable canonical identities + Given the executable Prowler registry + Then it contains the four base routes, three AWS routes, azure/iam, and azure/storage + And every Azure service route has service-identifying labels and no selector field + + Scenario: Unsupported service combinations stop before the client + Given an Azure service contract is paired with an unsupported selector + When the invalid contract execution is attempted + Then the request is rejected before the Prowler client is called + + Scenario: Service findings preserve CHK.005 mapping order in shared outputs + Given ordered mapped and non-Azure OCSF records for azure/storage + When the azure/storage assessment completes + Then mapped Azure findings retain their order in Text and Vulnerability outputs + And the same mapped findings appear in the dynamic Rich trace + + Scenario Outline: Runtime dispatch emits exact service argv without real execution + Given a fake CLI engine and valid Azure runtime message for "" + When the injector processes the runtime message + Then exactly one fake CLI request contains "--services" followed by "" + And credentials, stderr, and path canaries are absent from the callback + + Examples: + | route | service | + | azure/iam | iam | + | azure/storage | storage | + +# ---- Constraints identified ---- + Scenario: Existing check filters retain their dedicated CLI flag + Given a fake Azure Prowler client request with an existing check filter + When that request is rendered without a service selector + Then the check filter still uses "-c" and no "--services" argument is emitted + + Scenario: An Azure service selector cannot target another provider + Given a fake AWS Prowler client request with Azure service "storage" + When the client request is validated + Then it is rejected before the fake CLI engine is called + + Scenario: Azure inputs receive structural validation only + Given an Azure service form whose provider fields are nonblank + When the service contract parses the form + Then the values are accepted without cloud access or semantic identifier lookup diff --git a/prowler/tests/behaviour/chk012_service_azure_assessment/conftest.py b/prowler/tests/behaviour/chk012_service_azure_assessment/conftest.py new file mode 100644 index 00000000..a5d438dd --- /dev/null +++ b/prowler/tests/behaviour/chk012_service_azure_assessment/conftest.py @@ -0,0 +1,111 @@ +"""Local deterministic CHK.012 fixtures.""" + +from dataclasses import dataclass, field +from typing import Any + +import pytest + + +@dataclass(frozen=True) +class RecordedLog: + """One AppLogger-compatible call captured without formatting side effects.""" + + level: str + message: str + metadata: dict[str, object] | None = None + exc_info: bool | None = None + + +class _RecordingLocalLogger: + """Capture the direct standard-library ERROR path used by the injector.""" + + def __init__(self, events: list[RecordedLog]) -> None: + self._events = events + + def error( + self, + message: str, + *, + exc_info: bool, + extra: dict[str, object], + ) -> None: + """Record safe ERROR metadata in the same shape accepted by AppLogger.""" + attributes = extra.get("attributes") + metadata = dict(attributes) if isinstance(attributes, dict) else None + self._events.append(RecordedLog("error", message, metadata, exc_info)) + + +@dataclass +class RecordingLogger: + """Record the lifecycle logger surface exposed by the OpenAEV helper.""" + + events: list[RecordedLog] = field(default_factory=list) + + def __post_init__(self) -> None: + """Attach the direct ERROR surface to the shared event stream.""" + self.local_logger = _RecordingLocalLogger(self.events) + + def debug(self, message: str, metadata: dict[str, object]) -> None: + """Record one DEBUG lifecycle event.""" + self.events.append(RecordedLog("debug", message, metadata)) + + def info(self, message: str, metadata: dict[str, object] | None = None) -> None: + """Record one INFO lifecycle event.""" + self.events.append(RecordedLog("info", message, metadata)) + + def warning(self, message: str) -> None: + """Record one WARNING lifecycle event.""" + self.events.append(RecordedLog("warning", message)) + + +@pytest.fixture +def azure_form() -> dict[str, object]: + """Return structurally valid placeholder-only Azure form input.""" + return { + "azure_tenant_id": "TENANT-CANARY", + "azure_client_id": "CLIENT-ID-CANARY", + "azure_client_secret": "CLIENT-SECRET-CANARY\nFORM-CANARY", + "azure_subscription_id": "SUBSCRIPTION-CANARY", + "azure_provider": "PROVIDER-CANARY", + } + + +@pytest.fixture +def azure_ocsf_record_factory() -> Any: + """Build one complete minimal OCSF record for Azure service tests.""" + + def build( + title: str, + *, + provider: str = "azure", + status: str = "FAIL", + ) -> dict[str, Any]: + if status == "PASS": + record_status, status_code = "New", "PASS" + elif status == "MUTED": + record_status, status_code = "Suppressed", "FAIL" + else: + record_status, status_code = "New", status + return { + "finding_info": { + "uid": f"check-{title}", + "title": title, + "desc": f"Description {title}", + }, + "status": record_status, + "status_code": status_code, + "severity": "High", + "resources": [{"uid": f"asset-{title}", "name": f"Asset {title}"}], + "cloud": { + "provider": provider, + "region": "westeurope", + "account": {"uid": "subscription-123"}, + }, + "unmapped": {"compliance": ["cis", "nis2"]}, + "remediation": { + "desc": f"Remediate {title}", + "references": ["https://example.invalid/remediation"], + }, + } + + return build diff --git a/prowler/tests/behaviour/chk012_service_azure_assessment/test_chk012_service_azure_assessment_bdd.py b/prowler/tests/behaviour/chk012_service_azure_assessment/test_chk012_service_azure_assessment_bdd.py new file mode 100644 index 00000000..414a80cd --- /dev/null +++ b/prowler/tests/behaviour/chk012_service_azure_assessment/test_chk012_service_azure_assessment_bdd.py @@ -0,0 +1,613 @@ +"""Raw pytest executable contract for CHK.012.""" + +from __future__ import annotations + +import json +from dataclasses import dataclass, field +from pathlib import Path +from typing import Any, cast + +import pytest +from pyoaev.configuration import ConfigLoaderOAEV + +from prowler._core.cli_engine import ( + CommandResult, + ExecutionSpecification, + OutputSpecification, + ValidatedCommandRequest, +) +from prowler._core.prowler_client import ( + OUTPUT_ARTIFACT_FILENAME, + ProwlerClientFactory, +) +from prowler.contracts import DEFAULT_PROWLER_CONTRACTS, stable_contract_id +from prowler.models.configs.config_loader import ( + ConfigLoader, + InjectorConfig, + ProwlerConfig, +) +from prowler.models.provider_inputs import AwsProviderInput + +from .conftest import RecordingLogger + +_ROUTES = (("azure/iam", "iam"), ("azure/storage", "storage")) +_ASSESSMENT_RECEIVED = "[PROWLER_INJECTOR] - Assessment received" +_RECEPTION_ACKNOWLEDGED = "[PROWLER_INJECTOR] - Reception acknowledged" +_CONTRACT_RESOLVED = "[PROWLER_INJECTOR] - Contract resolved" +_ASSESSMENT_VALIDATED = "[PROWLER_INJECTOR] - Assessment input validated" +_EXECUTION_STARTED = "[PROWLER_INJECTOR] - Assessment execution starting" +_ASSESSMENT_SUCCEEDED = "[PROWLER_INJECTOR] - Assessment completed" +_ASSESSMENT_FAILED = "[PROWLER_INJECTOR] - Assessment failed" +_CALLBACK_COMPLETED = "[PROWLER_INJECTOR] - Assessment callback completed" + + +def _specification(arguments: tuple[str, ...] = ()) -> ExecutionSpecification: + return ExecutionSpecification( + executable="/fake/prowler", + arguments=arguments, + environment=(), + working_directory=None, + input_bytes=b"", + output=OutputSpecification(parser="raw"), + timeout_seconds=1.0, + maximum_accepted_output_bytes=1024, + ) + + +@dataclass +class _ClientFactory: + result: CommandResult + calls: list[tuple[Any, Any, tuple[str, ...], object]] = field(default_factory=list) + + def run( + self, + config: Any, + provider: Any, + *, + check_filters: Any = (), + service_selector: object = None, + ) -> CommandResult: + self.calls.append((config, provider, tuple(check_filters), service_selector)) + return self.result + + +def _contract(route: str) -> Any: + return DEFAULT_PROWLER_CONTRACTS.resolve(str(stable_contract_id(route))) + + +@pytest.mark.parametrize(("route", "service"), _ROUTES) +def test_route_selects_exact_service_once( + route: str, service: str, azure_form: dict[str, object] +) -> None: + """A route-owned selector crosses the client seam once without check filters.""" + factory = _ClientFactory(CommandResult(specification=_specification())) + contract = _contract(route) + contract._client_factory = factory + + contract.execute(ProwlerConfig(), contract.parse_input(azure_form)) + + assert len(factory.calls) == 1 + assert factory.calls[0][2:] == ((), service) + + +def test_registry_has_exact_nine_canonical_contracts_without_selector_fields() -> None: + """The public surface is ordered, stable, labelled, and not user-selectable.""" + serialized = DEFAULT_PROWLER_CONTRACTS.contracts() + routes = ( + "aws", + "azure", + "gcp", + "kubernetes", + "aws/iam", + "aws/s3", + "aws/ec2", + *(item[0] for item in _ROUTES), + ) + + assert [item["contract_id"] for item in serialized] == [ + str(stable_contract_id(route)) for route in routes + ] + for item, (route, service) in zip(serialized[7:], _ROUTES, strict=True): + content = json.loads(item["contract_content"]) + assert service.casefold() in content["label"]["en"].casefold() + assert tuple(field["key"] for field in content["fields"]) == ( + "azure_tenant_id", + "azure_client_id", + "azure_client_secret", + "azure_subscription_id", + "azure_provider", + ) + assert all(route in output["labels"] for output in content["outputs"]) + + +def test_unsupported_contract_selector_is_rejected_pre_client( + azure_form: dict[str, object], +) -> None: + """Invalid route metadata cannot consume the CHK.004 seam.""" + import prowler.contracts as contract_api + + service_contract = cast(Any, contract_api.__dict__["AzureServiceContract"]) + + class InvalidServiceContract(service_contract): + contract_id = str(stable_contract_id("azure/iam")) + external_id = "prowler:azure/iam" + route_name = "azure/iam" + label = "Invalid" + service_selector = "keyvault" + + factory = _ClientFactory(CommandResult(specification=_specification())) + contract = InvalidServiceContract(factory) + + with pytest.raises(ValueError, match="unsupported Azure service selector"): + contract.execute(ProwlerConfig(), contract.parse_input(azure_form)) + + assert factory.calls == [] + + +@pytest.mark.parametrize( + "field_name", + ( + "azure_tenant_id", + "azure_client_id", + "azure_client_secret", + "azure_subscription_id", + "azure_provider", + ), +) +def test_blank_provider_field_is_rejected_pre_client( + azure_form: dict[str, object], field_name: str +) -> None: + """Service routes inherit structural nonblank validation before execution.""" + factory = _ClientFactory(CommandResult(specification=_specification())) + contract = _contract("azure/iam") + contract._client_factory = factory + + with pytest.raises(ValueError): + contract.parse_input({**azure_form, field_name: " \t"}) + + assert factory.calls == [] + + +def test_nonblank_provider_fields_receive_no_semantic_or_live_validation() -> None: + """Opaque nonblank identifiers parse locally without cloud access.""" + contract = _contract("azure/storage") + provider = contract.parse_input( + { + "azure_tenant_id": "!", + "azure_client_id": "?", + "azure_client_secret": "#", + "azure_subscription_id": "$", + "azure_provider": "%", + } + ) + + assert provider.azure_tenant_id == "!" + assert provider.azure_subscription_id == "$" + + +def test_mapping_order_shared_outputs_and_dynamic_trace( + azure_form: dict[str, object], azure_ocsf_record_factory: Any +) -> None: + """Keep ordered CHK.005 output while tracing the same dynamic findings.""" + records = [ + azure_ocsf_record_factory("first", provider="Azure", status="PASS"), + azure_ocsf_record_factory("excluded", provider="aws", status="FAIL"), + azure_ocsf_record_factory("second", status="FAIL"), + ] + artifact = json.dumps(records).encode() + factory = _ClientFactory( + CommandResult( + specification=_specification(), + return_code=0, + stdout=b"\x1b[32mconsole output is not OCSF JSON\x1b[0m", + parsed=artifact, + ) + ) + contract = _contract("azure/storage") + contract._client_factory = factory + provider = contract.parse_input(azure_form) + + outcome = contract.execute(ProwlerConfig(), provider) + payload = contract.output_payload(outcome.findings) + trace = contract.render_trace( + provider, + outcome.findings, + 1, + raw_record_count=outcome.raw_record_count, + raw_output_bytes=outcome.raw_output_bytes, + raw_preview=outcome.raw_preview, + ) + + assert tuple(item.value for item in outcome.findings) == ("first", "second") + assert tuple(json.loads(item)["value"] for item in payload["findings"]) == ( + "first", + "second", + ) + assert tuple(item["name"] for item in payload["vulnerabilities"]) == ("second",) + assert all(name in trace for name in ("first", "second")) + raw_section_index = trace.index("[PROWLER] Raw OCSF evidence (bounded preview)") + assert "excluded" not in trace[:raw_section_index] + assert trace.index("Prowler Findings") < raw_section_index + assert outcome.raw_record_count == 3 + assert outcome.raw_output_bytes == len(artifact) + assert "azure/storage" in trace + assert "service=storage" in trace + canaries = ( + azure_form["azure_tenant_id"], + azure_form["azure_client_id"], + azure_form["azure_client_secret"], + ) + assert all(str(marker) not in trace for marker in canaries) + + +@dataclass +class _Engine: + payload: bytes + requests: list[ValidatedCommandRequest] = field(default_factory=list) + + def run(self, request: ValidatedCommandRequest) -> CommandResult: + self.requests.append(request) + arguments = tuple(request.arguments) + output_directory = Path(arguments[arguments.index("--output-directory") + 1]) + (output_directory / OUTPUT_ARTIFACT_FILENAME).write_bytes(self.payload) + specification = ExecutionSpecification.from_request(request) + return CommandResult( + specification=ExecutionSpecification( + executable=specification.executable, + arguments=(*specification.arguments, "PROCESS-ARGV-CANARY"), + environment=( + *specification.environment, + ("PROCESS-ENV-NAME-CANARY", "PROCESS-ENV-VALUE-CANARY"), + ), + working_directory="/tmp/PROCESS-WORKDIR-CANARY", # noqa: S108 + input_bytes=specification.input_bytes, + output=specification.output, + timeout_seconds=specification.timeout_seconds, + maximum_accepted_output_bytes=( + specification.maximum_accepted_output_bytes + ), + ), + return_code=0, + stdout=b"\x1b[31mPROCESS-STDOUT-CANARY\x1b[0m", + stderr=b"PROCESS-STDERR-CANARY", + ) + + +@dataclass +class _EngineFactory: + engine: _Engine + calls: int = 0 + + def create(self) -> _Engine: + self.calls += 1 + return self.engine + + +@dataclass +class _NoLeaseFactory: + calls: int = 0 + + def create(self, secret: Any, *, suffix: str) -> Any: + del secret, suffix + self.calls += 1 + raise AssertionError("Azure must not create credential files") + + +class _InjectApi: + def __init__(self) -> None: + self.events: list[tuple[str, str, dict[str, Any]]] = [] + + def execution_reception(self, *, inject_id: str, data: dict[str, Any]) -> None: + self.events.append(("reception", inject_id, data)) + + def execution_callback(self, *, inject_id: str, data: dict[str, Any]) -> None: + self.events.append(("callback", inject_id, data)) + + +class _Helper: + def __init__(self) -> None: + self.api = type("Api", (), {})() + self.api.inject = _InjectApi() + self.injector_logger = RecordingLogger() + + +def _config() -> ConfigLoader: + return ConfigLoader.model_construct( + openaev=ConfigLoaderOAEV( + url="http://127.0.0.1:8080", token="runtime-placeholder" + ), + injector=InjectorConfig(id="injector-test"), + prowler=ProwlerConfig(executable_path="/fake/prowler"), + ) + + +def _message(route: str, service: str, content: dict[str, object]) -> dict[str, object]: + return { + "injection": { + "inject_id": f"inject-{service}", + "injector_contract_id": str(stable_contract_id(route)), + "inject_content": content, + } + } + + +@pytest.mark.parametrize(("route", "service"), _ROUTES) +def test_runtime_one_call_exact_service_argv_outputs_and_canaries( + route: str, + service: str, + azure_form: dict[str, object], + azure_ocsf_record_factory: Any, +) -> None: + """The full runtime uses one fake execution and preserves safe dynamic output.""" + from prowler.injector import ProwlerInjector + + callback_name = f"CALLBACK-CANARY-{service}" + finding_name = f"FINDING-CANARY-{service}" + records = [ + azure_ocsf_record_factory(callback_name, status="PASS"), + azure_ocsf_record_factory(finding_name, status="FAIL"), + azure_ocsf_record_factory("excluded", provider="aws", status="FAIL"), + ] + engine = _Engine(json.dumps(records).encode()) + engine_factory = _EngineFactory(engine) + leases = _NoLeaseFactory() + contract = _contract(route) + contract._client_factory = ProwlerClientFactory(engine_factory, leases) + helper = _Helper() + ProwlerInjector(_config(), helper).process_message( + _message(route, service, azure_form) + ) + + assert engine_factory.calls == 1 + assert len(engine.requests) == 1 + assert tuple(engine.requests[0].arguments) == ( + "azure", + "--sp-env-auth", + "--subscription-id", + "SUBSCRIPTION-CANARY", + "--azure-region", + "PROVIDER-CANARY", + "--services", + service, + "--output-directory", + engine.requests[0].arguments[ + engine.requests[0].arguments.index("--output-directory") + 1 + ], + "--output-filename", + "findings", + "-z", + "--only-logs", + "--no-color", + "-M", + "json-ocsf", + ) + assert "-c" not in engine.requests[0].arguments + assert tuple(event[0] for event in helper.api.inject.events) == ( + "reception", + "callback", + ) + callback = helper.api.inject.events[1][2] + assert callback["execution_status"] == "SUCCESS" + mapped_names = (callback_name, finding_name) + structured = json.loads(callback["execution_output_structured"]) + assert tuple(json.loads(item)["value"] for item in structured["findings"]) == ( + mapped_names + ) + assert tuple(item["name"] for item in structured["vulnerabilities"]) == ( + finding_name, + ) + assert all(name in callback["execution_message"] for name in mapped_names) + raw_section_index = callback["execution_message"].index( + "[PROWLER] Raw OCSF evidence (bounded preview)" + ) + assert "excluded" not in callback["execution_message"][:raw_section_index] + assert str(azure_form["azure_subscription_id"]) in callback["execution_message"] + assert str(azure_form["azure_provider"]) in callback["execution_message"] + callback_excluded_canaries = ( + azure_form["azure_tenant_id"], + azure_form["azure_client_id"], + azure_form["azure_client_secret"], + "FORM-CANARY", + "PROCESS-ARGV-CANARY", + "PROCESS-ENV-NAME-CANARY", + "PROCESS-ENV-VALUE-CANARY", + "PROCESS-STDOUT-CANARY", + "PROCESS-STDERR-CANARY", + "/tmp/PROCESS-WORKDIR-CANARY", # noqa: S108 - deliberate leak canary + "EXCEPTION-FIELD-CANARY", + "EXCEPTION-VALUE-CANARY", + ) + assert all( + str(marker) not in json.dumps(callback) for marker in callback_excluded_canaries + ) + assert leases.calls == 0 + + logs = helper.injector_logger.events + assert tuple((event.level, event.message) for event in logs) == ( + ("info", _ASSESSMENT_RECEIVED), + ("debug", _RECEPTION_ACKNOWLEDGED), + ("debug", _CONTRACT_RESOLVED), + ("debug", _ASSESSMENT_VALIDATED), + ("info", _EXECUTION_STARTED), + ("info", _ASSESSMENT_SUCCEEDED), + ("debug", _CALLBACK_COMPLETED), + ) + assert [event.metadata["stage"] for event in logs if event.metadata] == [ + "message_reception", + "reception_acknowledged", + "contract_resolution", + "input_validation", + "assessment_execution", + "assessment_completion", + "callback", + ] + assert all(event.metadata is not None for event in logs) + assert all( + event.metadata["inject_id"] == f"inject-{service}" + for event in logs + if event.metadata + ) + assert all( + 0 <= event.metadata["elapsed_ms"] <= 86_400_000 + for event in logs + if event.metadata + ) + identifier = str(stable_contract_id(route)) + for event in logs[2:]: + assert event.metadata is not None + assert event.metadata["contract_id"] == identifier + assert event.metadata["route"] == route + assert event.metadata["provider"] == "azure" + for event in logs[3:]: + assert event.metadata is not None + assert event.metadata["azure_tenant_id_present"] is True + assert event.metadata["azure_client_id_present"] is True + assert event.metadata["azure_client_secret_present"] is True + assert event.metadata["azure_subscription_id"] == "SUBSCRIPTION-CANARY" + assert event.metadata["azure_provider"] == "PROVIDER-CANARY" + success_metadata = logs[5].metadata + assert success_metadata is not None + assert success_metadata["status"] == "SUCCESS" + assert success_metadata["finding_count"] == 2 + assert success_metadata["vulnerability_count"] == 1 + assert success_metadata["raw_record_count"] == 3 + assert success_metadata["raw_output_bytes"] == len(engine.payload) + callback_metadata = logs[6].metadata + assert callback_metadata is not None + assert callback_metadata["assessment_status"] == "SUCCESS" + assert callback_metadata["delivery_status"] == "SUCCESS" + assert "status" not in callback_metadata + assert "attempted_status" not in callback_metadata + log_excluded_canaries = ( + *callback_excluded_canaries, + callback_name, + finding_name, + ) + assert all(str(marker) not in repr(logs) for marker in log_excluded_canaries) + + invalid_helper = _Helper() + ProwlerInjector(_config(), invalid_helper).process_message( + _message( + route, + service, + { + **azure_form, + "azure_subscription_id": " ", + "EXCEPTION-FIELD-CANARY": "EXCEPTION-VALUE-CANARY", + }, + ) + ) + + assert tuple(event[0] for event in invalid_helper.api.inject.events) == ( + "reception", + "callback", + ) + invalid_callback = invalid_helper.api.inject.events[1][2] + assert invalid_callback["execution_status"] == "ERROR" + assert all( + str(marker) not in json.dumps(invalid_callback) + for marker in log_excluded_canaries + ) + assert engine_factory.calls == 1 + + invalid_logs = invalid_helper.injector_logger.events + assert tuple((event.level, event.message) for event in invalid_logs) == ( + ("info", _ASSESSMENT_RECEIVED), + ("debug", _RECEPTION_ACKNOWLEDGED), + ("debug", _CONTRACT_RESOLVED), + ("error", _ASSESSMENT_FAILED), + ("debug", _CALLBACK_COMPLETED), + ) + assert [event.metadata["stage"] for event in invalid_logs if event.metadata] == [ + "message_reception", + "reception_acknowledged", + "contract_resolution", + "input_validation", + "callback", + ] + assert all(event.metadata is not None for event in invalid_logs) + assert all( + event.metadata["inject_id"] == f"inject-{service}" + for event in invalid_logs + if event.metadata + ) + assert all( + 0 <= event.metadata["elapsed_ms"] <= 86_400_000 + for event in invalid_logs + if event.metadata + ) + for event in invalid_logs[2:]: + assert event.metadata is not None + assert event.metadata["contract_id"] == identifier + assert event.metadata["route"] == route + assert event.metadata["provider"] == "azure" + failure_metadata = invalid_logs[3].metadata + assert failure_metadata is not None + assert failure_metadata["status"] == "ERROR" + assert failure_metadata["stage"] == "input_validation" + assert failure_metadata["failure_kind"] == "invalid_input" + assert failure_metadata["failure_summary"] == "The assessment input was invalid." + assert failure_metadata["operator_guidance"] == ( + "Correct the listed assessment fields and retry." + ) + assert failure_metadata["issues"] == [ + { + "location": ["azure", "azure_subscription_id"], + "type": "value_error", + }, + { + "location": ["azure", "unrecognized_field"], + "type": "extra_forbidden", + }, + ] + assert failure_metadata["issues_truncated"] is True + assert invalid_logs[3].exc_info is False + invalid_callback_metadata = invalid_logs[4].metadata + assert invalid_callback_metadata is not None + assert invalid_callback_metadata["assessment_status"] == "ERROR" + assert invalid_callback_metadata["delivery_status"] == "SUCCESS" + assert "status" not in invalid_callback_metadata + assert "attempted_status" not in invalid_callback_metadata + assert all( + str(marker) not in repr(invalid_logs) for marker in log_excluded_canaries + ) + + +def test_existing_check_filter_argv_is_unchanged( + azure_form: dict[str, object], azure_ocsf_record_factory: Any +) -> None: + """The selector seam does not repurpose or remove CHK.004 check filtering.""" + engine = _Engine(json.dumps([azure_ocsf_record_factory("existing")]).encode()) + factory = ProwlerClientFactory(_EngineFactory(engine), _NoLeaseFactory()) + provider = _contract("azure").parse_input(azure_form) + + factory.run( + ProwlerConfig(executable_path="/fake/prowler"), + provider, + check_filters=("check-one",), + ) + + check_index = engine.requests[0].arguments.index("-c") + assert tuple(engine.requests[0].arguments[check_index : check_index + 2]) == ( + "-c", + "check-one", + ) + assert "--services" not in engine.requests[0].arguments + + +def test_azure_service_selector_rejects_aws_before_engine() -> None: + """Provider/selector combinations are validated before any CLI call.""" + engine = _Engine(b"[]") + factory = ProwlerClientFactory(_EngineFactory(engine), _NoLeaseFactory()) + provider = AwsProviderInput( + provider="aws", + aws_access_key_id="access", + aws_secret_access_key="secret", + aws_account_id="123456789012", + aws_region="eu-west-1", + ) + + with pytest.raises(ValueError, match="Azure service selector"): + factory.run(ProwlerConfig(), provider, service_selector=cast(Any, "storage")) + + assert engine.requests == [] diff --git a/prowler/tests/unit/chk006_executable_base/test_outputs_registry_runtime.py b/prowler/tests/unit/chk006_executable_base/test_outputs_registry_runtime.py index 31c62620..473769e4 100644 --- a/prowler/tests/unit/chk006_executable_base/test_outputs_registry_runtime.py +++ b/prowler/tests/unit/chk006_executable_base/test_outputs_registry_runtime.py @@ -398,7 +398,7 @@ def test_runtime_start_logs_one_fixed_listener_event() -> None: assert message == _LISTENER_START assert metadata["injector_id"] == "injector-test" assert metadata["injector_name"] == "Prowler" - assert metadata["registered_contract_count"] == 7 + assert metadata["registered_contract_count"] == 9 assert metadata["configured_executable_path"] == "/usr/local/bin/prowler" assert set(metadata) == { "injector_id", @@ -1317,7 +1317,7 @@ def test_runtime_resolved_contract_uses_renderer_for_safe_error( def test_default_registry_and_daemon_config_register_executable_routes() -> None: - """CHK.011 adds three AWS services after the four canonical base routes.""" + """CHK.012 adds two Azure services after the CHK.011 executable routes.""" subject = _subject() contracts = subject.DEFAULT_PROWLER_CONTRACTS.contracts() assert [item["contract_id"] for item in contracts] == [ @@ -1328,6 +1328,8 @@ def test_default_registry_and_daemon_config_register_executable_routes() -> None str(subject.stable_contract_id("aws/iam")), str(subject.stable_contract_id("aws/s3")), str(subject.stable_contract_id("aws/ec2")), + str(subject.stable_contract_id("azure/iam")), + str(subject.stable_contract_id("azure/storage")), ] daemon = _config().to_daemon_config() assert daemon.get("injector_contracts") == contracts