Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 3 additions & 1 deletion prowler/prowler/_core/prowler_client/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,7 @@
ProwlerClient,
ProwlerClientConsumedError,
)
from .contracts import AwsServiceSelector
from .contracts import AwsServiceSelector, AzureServiceSelector, ServiceSelector
from .credentials import (
CredentialCleanupError,
TemporaryCredentialLease,
Expand Down Expand Up @@ -39,9 +39,11 @@
"OutputWorkspaceCleanupError",
"OutputWorkspacePreparationError",
"AwsServiceSelector",
"AzureServiceSelector",
"ProwlerClient",
"ProwlerClientConsumedError",
"ProwlerClientFactory",
"ServiceSelector",
"TemporaryCredentialLease",
"TemporaryCredentialLeaseFactory",
"TemporaryOutputWorkspace",
Expand Down
26 changes: 20 additions & 6 deletions prowler/prowler/_core/prowler_client/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand Down Expand Up @@ -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:
Expand All @@ -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:
Expand Down
2 changes: 2 additions & 0 deletions prowler/prowler/_core/prowler_client/contracts.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down
4 changes: 2 additions & 2 deletions prowler/prowler/_core/prowler_client/factory.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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(
Expand Down
10 changes: 9 additions & 1 deletion prowler/prowler/contracts/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,12 @@
AwsS3Contract,
AwsServiceContract,
)
from .azure import AzureBaseContract
from .azure import (
AzureBaseContract,
AzureIamContract,
AzureServiceContract,
AzureStorageContract,
)
from .base import (
BaseProwlerContract,
ContractExecutionOutcome,
Expand Down Expand Up @@ -35,6 +40,9 @@
"AwsS3Contract",
"AwsServiceContract",
"AzureBaseContract",
"AzureIamContract",
"AzureServiceContract",
"AzureStorageContract",
"GcpBaseContract",
"KubernetesBaseContract",
"ContractDispatcher",
Expand Down
17 changes: 2 additions & 15 deletions prowler/prowler/contracts/aws.py
Original file line number Diff line number Diff line change
@@ -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

Expand All @@ -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
Expand All @@ -28,25 +26,14 @@ 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:
"""Map one complete-scope result, preserving ordered AWS findings only."""
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):
Expand Down Expand Up @@ -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):
Expand Down
81 changes: 72 additions & 9 deletions prowler/prowler/contracts/azure.py
Original file line number Diff line number Diff line change
@@ -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):
Expand All @@ -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 = ()

Expand All @@ -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"
14 changes: 12 additions & 2 deletions prowler/prowler/contracts/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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."""

Expand Down Expand Up @@ -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."""
Expand Down
7 changes: 1 addition & 6 deletions prowler/prowler/contracts/gcp.py
Original file line number Diff line number Diff line change
Expand Up @@ -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))
13 changes: 0 additions & 13 deletions prowler/prowler/contracts/kubernetes.py
Original file line number Diff line number Diff line change
@@ -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
Expand All @@ -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:
Expand Down
4 changes: 3 additions & 1 deletion prowler/prowler/contracts/registry.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -90,5 +90,7 @@ def contracts(self) -> list[dict[str, object]]:
AwsIamContract,
AwsS3Contract,
AwsEc2Contract,
AzureIamContract,
AzureStorageContract,
)
)
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Loading
Loading