diff --git a/products/warehouse_sources/backend/temporal/data_imports/sources/SOURCES.md b/products/warehouse_sources/backend/temporal/data_imports/sources/SOURCES.md index b7f62456de42..e82b70e3d577 100644 --- a/products/warehouse_sources/backend/temporal/data_imports/sources/SOURCES.md +++ b/products/warehouse_sources/backend/temporal/data_imports/sources/SOURCES.md @@ -106,6 +106,7 @@ the row lists both. | aws_budgets | HTTP | requests | ✅ | | aws_cloudtrail | HTTP | requests | ✅ | | aws_compute_optimizer | HTTP | requests | ✅ | +| aws_config | HTTP | requests | ✅ | | aws_cost_anomaly_detection | HTTP | requests | ✅ | | aws_cost_explorer | HTTP | requests | ✅ | | aws_glue_data_catalog | HTTP | requests | ✅ | @@ -922,7 +923,6 @@ doesn't conflict with concurrent PRs. - automox - aws_athena - aws_cloudformation -- aws_config - aws_connect - aws_cost_and_usage_report - aws_guardduty diff --git a/products/warehouse_sources/backend/temporal/data_imports/sources/aws_config/aws_config.py b/products/warehouse_sources/backend/temporal/data_imports/sources/aws_config/aws_config.py new file mode 100644 index 000000000000..6c67ef578c0e --- /dev/null +++ b/products/warehouse_sources/backend/temporal/data_imports/sources/aws_config/aws_config.py @@ -0,0 +1,239 @@ +import re +import json +import datetime as dt +from collections.abc import Generator +from typing import TYPE_CHECKING, Any + +import requests +from botocore.auth import SigV4Auth +from botocore.awsrequest import AWSRequest +from botocore.credentials import Credentials +from tenacity import retry, retry_if_exception_type, stop_after_attempt, wait_exponential_jitter + +from posthog.dataclasses import frozen + +from products.warehouse_sources.backend.temporal.data_imports.sources.aws_config.settings import ( + AWS_CONFIG_ENDPOINTS, + ERROR_MESSAGES, + TARGET_PREFIXES, + AwsConfigEndpoint, +) +from products.warehouse_sources.backend.temporal.data_imports.sources.common.http import make_tracked_session +from products.warehouse_sources.backend.temporal.data_imports.sources.common.http.transport import BoundedRetry +from products.warehouse_sources.backend.temporal.data_imports.sources.common.typings import SourceResponse + +if TYPE_CHECKING: + from products.warehouse_sources.backend.temporal.data_imports.sources.common.resumable import ResumableSourceManager + from products.warehouse_sources.backend.temporal.data_imports.sources.generated_configs.awsconfig import ( + AwsConfigSourceConfig, + ) + +_REGION = re.compile(r"[a-z]{2}(?:-[a-z]+)+-\d+") +_CAMEL_BOUNDARY = re.compile(r"(?<=[a-z0-9])(?=[A-Z])|(?<=[A-Z])(?=[A-Z][a-z])") +_THROTTLE_CODES = {"Throttling", "ThrottlingException", "TooManyRequestsException", "RequestLimitExceeded"} +_TIMESTAMP_COLUMNS = {"configuration_item_capture_time", "resource_creation_time", "last_update_requested_time"} + +# Config reads use POST, so the transport must retry this method for transient HTTP failures. +TRANSPORT_RETRY = BoundedRetry( + total=3, + backoff_factor=1, + status_forcelist=(429, 500, 502, 503, 504), + allowed_methods=frozenset({"POST"}), + raise_on_status=False, +) + + +@frozen +class AwsConfigResumeConfig: + next_token: str | None = None + complete: bool = False + + +class AwsConfigError(Exception): + def __init__(self, code: str, message: str) -> None: + super().__init__(f"AWS Config request failed: {code} - {message}") + self.code = code + + +class AwsConfigThrottledError(AwsConfigError): + pass + + +def error_for_response(response: requests.Response) -> AwsConfigError: + try: + body = response.json() + except ValueError: + body = {} + if not isinstance(body, dict): + body = {} + raw_code = response.headers.get("x-amzn-ErrorType") or body.get("__type") or body.get("code") + code = str(raw_code or f"HTTP {response.status_code}").split(":", 1)[0].rsplit("#", 1)[-1] + message = str(body.get("message") or body.get("Message") or response.text)[:500] + # HTTP retries belong to the tracked transport. Only body-level throttles need another retry policy. + error_class = AwsConfigThrottledError if response.status_code == 400 and code in _THROTTLE_CODES else AwsConfigError + return error_class(code, message) + + +class AwsConfigClient: + def __init__(self, config: "AwsConfigSourceConfig", api_version: str) -> None: + if not config.aws_access_key_id or not config.aws_secret_access_key: + raise ValueError("Enter both an AWS access key ID and a secret access key.") + self.region = config.region or "us-east-1" + if not _REGION.fullmatch(self.region): + raise ValueError("Enter an AWS region such as us-east-1.") + if api_version not in TARGET_PREFIXES: + raise ValueError(f"Unsupported AWS Config API version: {api_version}") + self.target_prefix = TARGET_PREFIXES[api_version] + suffix = "amazonaws.com.cn" if self.region.startswith("cn-") else "amazonaws.com" + self.url = f"https://config.{self.region}.{suffix}/" + self._signer = SigV4Auth( + Credentials(config.aws_access_key_id, config.aws_secret_access_key, config.aws_session_token or None), + "config", + self.region, + ) + self._session = make_tracked_session( + retry=TRANSPORT_RETRY, + redact_values=tuple(value for value in (config.aws_secret_access_key, config.aws_session_token) if value), + ) + + def close(self) -> None: + self._session.close() + + @retry( + retry=retry_if_exception_type(AwsConfigThrottledError), + stop=stop_after_attempt(5), + wait=wait_exponential_jitter(initial=1, max=30), + reraise=True, + ) + def request(self, operation: str, payload: dict[str, Any]) -> dict[str, Any]: + body = json.dumps(payload).encode("utf-8") + request = AWSRequest( + method="POST", + url=self.url, + data=body, + headers={ + "Content-Type": "application/x-amz-json-1.1", + "X-Amz-Target": f"{self.target_prefix}.{operation}", + }, + ) + self._signer.add_auth(request) + response = self._session.post( + self.url, data=body, headers=dict(request.headers), timeout=60, allow_redirects=False + ) + if response.status_code != 200: + raise error_for_response(response) + parsed = response.json() + if not isinstance(parsed, dict): + raise ValueError("AWS Config returned an invalid response. Try the sync again.") + return parsed + + +def request_payload(endpoint: AwsConfigEndpoint, *, probe: bool = False) -> dict[str, Any]: + payload: dict[str, Any] = {} + if endpoint.page_size is not None: + payload["Limit"] = 1 if probe else endpoint.page_size + if endpoint.expression is not None: + payload["Expression"] = endpoint.expression + return payload + + +def normalize_row(item: dict[str, Any], region: str) -> dict[str, Any]: + row = {_CAMEL_BOUNDARY.sub("_", key).lower(): value for key, value in item.items()} + row["region"] = region + for key in _TIMESTAMP_COLUMNS & row.keys(): + value = row[key] + if isinstance(value, int | float) and not isinstance(value, bool): + row[key] = dt.datetime.fromtimestamp(value, tz=dt.UTC) + elif isinstance(value, str) and value: + row[key] = dt.datetime.fromisoformat(value.replace("Z", "+00:00")) + return row + + +def get_rows( + config: "AwsConfigSourceConfig", + endpoint: AwsConfigEndpoint, + api_version: str, + manager: "ResumableSourceManager[AwsConfigResumeConfig]", +) -> Generator[list[dict[str, Any]]]: + resume = manager.load_state() if manager.can_resume() else None + if resume is not None and resume.complete: + manager.clear_state() + return + next_token = resume.next_token if resume else None + restarting = next_token is not None + client = AwsConfigClient(config, api_version) + try: + while True: + payload = request_payload(endpoint) + if next_token: + payload["NextToken"] = next_token + try: + body = client.request(endpoint.operation, payload) + except AwsConfigError as error: + if restarting and error.code == "InvalidNextTokenException": + # A fresh retry must restart the full-refresh destination as well as extraction. + manager.clear_state() + raise + restarting = False + rows = [] + for item in body.get(endpoint.result_key) or []: + if endpoint.expression is not None: + item = json.loads(item) + if not isinstance(item, dict): + raise ValueError("AWS Config returned an invalid resource. Try the sync again.") + rows.append(normalize_row(item, client.region)) + token = body.get("NextToken") or None + if token is not None and (not isinstance(token, str) or token == next_token): + raise ValueError("AWS Config returned an invalid page token. Try the sync again.") + next_token = token + manager.save_state(AwsConfigResumeConfig(next_token=next_token, complete=next_token is None)) + if rows: + yield rows + manager.safe_point() + if next_token is None: + break + manager.clear_state() + finally: + client.close() + + +def validate_credentials( + config: "AwsConfigSourceConfig", api_version: str, schema_name: str | None = None +) -> tuple[bool, str | None]: + endpoint = AWS_CONFIG_ENDPOINTS.get(schema_name or "config_rules") + if endpoint is None: + return False, f"Unknown AWS Config table: {schema_name}" + try: + client = AwsConfigClient(config, api_version) + except ValueError as error: + return False, str(error) + try: + client.request(endpoint.operation, request_payload(endpoint, probe=True)) + except AwsConfigError as error: + if error.code in {"AccessDenied", "AccessDeniedException"}: + if schema_name is None: + return True, None + return False, f"Grant config:{endpoint.operation} to this IAM user or role to sync this table." + return False, ERROR_MESSAGES.get(error.code, "Could not read AWS Config. Check the region and try again.") + except (requests.RequestException, ValueError): + return False, "Could not reach the AWS Config API. Check the region and try again." + finally: + client.close() + return True, None + + +def aws_config_source( + config: "AwsConfigSourceConfig", + endpoint: str, + api_version: str, + manager: "ResumableSourceManager[AwsConfigResumeConfig]", +) -> SourceResponse: + endpoint_config = AWS_CONFIG_ENDPOINTS.get(endpoint) + if endpoint_config is None: + raise ValueError(f"Unknown AWS Config table: {endpoint}") + return SourceResponse( + name=endpoint, + items=lambda: get_rows(config, endpoint_config, api_version, manager), + primary_keys=list(endpoint_config.primary_key), + sort_mode=None, + ) diff --git a/products/warehouse_sources/backend/temporal/data_imports/sources/aws_config/canonical_descriptions.py b/products/warehouse_sources/backend/temporal/data_imports/sources/aws_config/canonical_descriptions.py new file mode 100644 index 000000000000..9baf60bd2c0f --- /dev/null +++ b/products/warehouse_sources/backend/temporal/data_imports/sources/aws_config/canonical_descriptions.py @@ -0,0 +1,63 @@ +from products.warehouse_sources.backend.temporal.data_imports.sources.common.canonical_descriptions import ( + CanonicalDescriptions, +) + +CANONICAL_DESCRIPTIONS: CanonicalDescriptions = { + "resources": { + "description": "Current configurations of recorded AWS resources in the selected region. Excludes deleted resources.", + "docs_url": "https://docs.aws.amazon.com/config/latest/APIReference/API_SelectResourceConfig.html", + "columns": { + "account_id": "AWS account that owns the resource.", + "aws_region": "AWS region reported for the resource.", + "region": "AWS region selected for this import.", + "arn": "Amazon Resource Name of the resource.", + "resource_id": "Resource identifier assigned by its AWS service.", + "resource_type": "AWS resource type, such as AWS::EC2::Instance.", + "resource_name": "Resource name, when the service provides one.", + "configuration": "Configuration properties recorded for the resource.", + "configuration_item_capture_time": "Time when AWS Config recorded this configuration.", + "resource_creation_time": "Time when the resource was created, when available.", + "tags": "Tags attached to the resource.", + }, + }, + "config_rules": { + "description": "AWS Config rules and their evaluation settings in the selected region.", + "docs_url": "https://docs.aws.amazon.com/config/latest/APIReference/API_DescribeConfigRules.html", + "columns": { + "config_rule_arn": "Amazon Resource Name of the rule.", + "config_rule_id": "Identifier assigned to the rule by AWS Config.", + "config_rule_name": "Name of the AWS Config rule.", + "config_rule_state": "Current state of the rule.", + "description": "Description of the rule.", + "scope": "Resource types, identifiers, or tags that restrict the rule's evaluations.", + "source": "Rule owner, identifier, and evaluation triggers.", + "input_parameters": "Parameters passed to the rule as a JSON string.", + "maximum_execution_frequency": "Maximum frequency for periodic evaluations.", + "region": "AWS region selected for this import.", + }, + }, + "rule_compliance": { + "description": "Current compliance status for AWS Config rules, including counts of noncompliant resources.", + "docs_url": "https://docs.aws.amazon.com/config/latest/APIReference/API_DescribeComplianceByConfigRule.html", + "columns": { + "config_rule_name": "Name of the evaluated AWS Config rule.", + "compliance": "Compliance status and a capped count of contributing resources, with an indicator when the cap is exceeded.", + "region": "AWS region selected for this import.", + }, + }, + "conformance_packs": { + "description": "Conformance packs and their deployment settings in the selected region.", + "docs_url": "https://docs.aws.amazon.com/config/latest/APIReference/API_DescribeConformancePacks.html", + "columns": { + "conformance_pack_arn": "Amazon Resource Name of the conformance pack.", + "conformance_pack_id": "Identifier assigned to the conformance pack.", + "conformance_pack_name": "Name of the conformance pack.", + "conformance_pack_input_parameters": "Parameters used to deploy the conformance pack.", + "delivery_s3_bucket": "S3 bucket used for the conformance pack template.", + "delivery_s3_key_prefix": "S3 prefix used for the conformance pack template.", + "last_update_requested_time": "Time when the most recent update was requested.", + "created_by": "AWS service that created the conformance pack.", + "region": "AWS region selected for this import.", + }, + }, +} diff --git a/products/warehouse_sources/backend/temporal/data_imports/sources/aws_config/settings.py b/products/warehouse_sources/backend/temporal/data_imports/sources/aws_config/settings.py new file mode 100644 index 000000000000..b5dbe3fd601f --- /dev/null +++ b/products/warehouse_sources/backend/temporal/data_imports/sources/aws_config/settings.py @@ -0,0 +1,68 @@ +from posthog.dataclasses import frozen + +CONFIG_API_VERSION = "2014-11-12" +TARGET_PREFIXES = {CONFIG_API_VERSION: "StarlingDoveService"} + +RESOURCE_EXPRESSION = ( + "SELECT accountId, awsRegion, arn, resourceId, resourceType, resourceName, " + "configuration, configurationItemCaptureTime, resourceCreationTime, tags" +) + + +@frozen +class AwsConfigEndpoint: + operation: str + result_key: str + primary_key: tuple[str, ...] + description: str + page_size: int | None = None + expression: str | None = None + + +AWS_CONFIG_ENDPOINTS: dict[str, AwsConfigEndpoint] = { + "resources": AwsConfigEndpoint( + operation="SelectResourceConfig", + result_key="Results", + primary_key=("account_id", "aws_region", "resource_type", "resource_id"), + description="Current configurations of recorded resources in the selected region. Excludes deleted resources.", + page_size=100, + expression=RESOURCE_EXPRESSION, + ), + "config_rules": AwsConfigEndpoint( + operation="DescribeConfigRules", + result_key="ConfigRules", + primary_key=("config_rule_arn",), + description="AWS Config rules, their evaluation settings, and resource scopes.", + ), + "rule_compliance": AwsConfigEndpoint( + operation="DescribeComplianceByConfigRule", + result_key="ComplianceByConfigRules", + primary_key=("region", "config_rule_name"), + description="Current compliance status and counts of noncompliant resources for each rule.", + ), + "conformance_packs": AwsConfigEndpoint( + operation="DescribeConformancePacks", + result_key="ConformancePackDetails", + primary_key=("conformance_pack_arn",), + description="Conformance packs and their deployment settings in the selected region.", + page_size=20, + ), +} + +ENDPOINTS = tuple(AWS_CONFIG_ENDPOINTS) +ENDPOINT_DESCRIPTIONS = {name: endpoint.description for name, endpoint in AWS_CONFIG_ENDPOINTS.items()} + +ERROR_MESSAGES = { + "AccessDenied": "AWS denied access. Grant the config read permission for the selected table to this IAM user or role.", + "AccessDeniedException": "AWS denied access. Grant the config read permission for the selected table to this IAM user or role.", + "UnrecognizedClientException": "AWS rejected the credentials. Check the access key ID, secret access key, and session token.", + "InvalidClientTokenId": "AWS rejected the access key ID. Check that the key is active.", + "InvalidSignatureException": "AWS rejected the signature. Check the secret access key and session token.", + "SignatureDoesNotMatch": "AWS rejected the signature. Check the secret access key and session token.", + "ExpiredTokenException": "The AWS session token has expired. Enter new temporary credentials.", + "ExpiredToken": "The AWS session token has expired. Enter new temporary credentials.", + "MissingAuthenticationToken": "AWS did not receive valid credentials. Enter the access key ID and secret access key.", + "OptInRequired": "Enable AWS Config in the selected account and region, then try again.", + "SubscriptionRequiredException": "Enable AWS Config in the selected account and region, then try again.", + "NoAvailableConfigurationRecorderException": "Set up an AWS Config recorder in the selected region, then try again.", +} diff --git a/products/warehouse_sources/backend/temporal/data_imports/sources/aws_config/source.py b/products/warehouse_sources/backend/temporal/data_imports/sources/aws_config/source.py index bb2f43841dd9..b7ada904f248 100644 --- a/products/warehouse_sources/backend/temporal/data_imports/sources/aws_config/source.py +++ b/products/warehouse_sources/backend/temporal/data_imports/sources/aws_config/source.py @@ -1,8 +1,37 @@ from typing import cast -from products.warehouse_sources.backend.facade.source_config import DataWarehouseSourceCategory, SourceConfig -from products.warehouse_sources.backend.temporal.data_imports.sources.common.base import FieldType, SimpleSource +from products.warehouse_sources.backend.facade.source_config import ( + DataWarehouseSourceCategory, + ReleaseStatus, + SourceConfig, + SourceFieldInputConfig, + SourceFieldInputConfigType, +) +from products.warehouse_sources.backend.temporal.data_imports.sources.aws_config.aws_config import ( + AwsConfigResumeConfig, + aws_config_source, + validate_credentials as validate_aws_config_credentials, +) +from products.warehouse_sources.backend.temporal.data_imports.sources.aws_config.canonical_descriptions import ( + CANONICAL_DESCRIPTIONS, +) +from products.warehouse_sources.backend.temporal.data_imports.sources.aws_config.settings import ( + CONFIG_API_VERSION, + ENDPOINT_DESCRIPTIONS, + ENDPOINTS, + ERROR_MESSAGES, +) +from products.warehouse_sources.backend.temporal.data_imports.sources.common.base import FieldType, ResumableSource +from products.warehouse_sources.backend.temporal.data_imports.sources.common.canonical_descriptions import ( + CanonicalDescriptions, +) from products.warehouse_sources.backend.temporal.data_imports.sources.common.registry import SourceRegistry +from products.warehouse_sources.backend.temporal.data_imports.sources.common.resumable import ResumableSourceManager +from products.warehouse_sources.backend.temporal.data_imports.sources.common.schema import ( + SourceSchema, + build_endpoint_schemas, +) +from products.warehouse_sources.backend.temporal.data_imports.sources.common.typings import SourceInputs, SourceResponse from products.warehouse_sources.backend.temporal.data_imports.sources.generated_configs.awsconfig import ( AwsConfigSourceConfig, ) @@ -10,18 +39,111 @@ @SourceRegistry.register -class AwsConfigSource(SimpleSource[AwsConfigSourceConfig]): +class AwsConfigSource(ResumableSource[AwsConfigSourceConfig, AwsConfigResumeConfig]): + lists_tables_without_credentials = True + supported_versions = (CONFIG_API_VERSION,) + default_version = CONFIG_API_VERSION + api_docs_url = "https://docs.aws.amazon.com/config/latest/APIReference/Welcome.html" + @property def source_type(self) -> ExternalDataSourceType: return ExternalDataSourceType.AWSCONFIG + def get_non_retryable_errors(self) -> dict[str, str | None]: + return {f"AWS Config request failed: {code}": message for code, message in ERROR_MESSAGES.items()} + + def get_canonical_descriptions(self) -> CanonicalDescriptions: + return CANONICAL_DESCRIPTIONS + + def get_schemas( + self, + config: AwsConfigSourceConfig, + team_id: int, + with_counts: bool = False, + names: list[str] | None = None, + force_refresh: bool = False, + api_version: str | None = None, + ) -> list[SourceSchema]: + return build_endpoint_schemas(ENDPOINTS, {}, names, descriptions=ENDPOINT_DESCRIPTIONS) + + def validate_credentials( + self, + config: AwsConfigSourceConfig, + team_id: int, + schema_name: str | None = None, + api_version: str | None = None, + ) -> tuple[bool, str | None]: + return validate_aws_config_credentials(config, self.resolve_api_version(api_version), schema_name) + + def get_resumable_source_manager(self, inputs: SourceInputs) -> ResumableSourceManager[AwsConfigResumeConfig]: + return ResumableSourceManager(inputs, AwsConfigResumeConfig) + + def source_for_pipeline( + self, + config: AwsConfigSourceConfig, + resumable_source_manager: ResumableSourceManager[AwsConfigResumeConfig], + inputs: SourceInputs, + ) -> SourceResponse: + return aws_config_source( + config, inputs.schema_name, self.resolve_api_version(inputs.api_version), resumable_source_manager + ) + @property def get_source_config(self) -> SourceConfig: return SourceConfig( name=ExternalDataSourceType.AWSCONFIG, category=DataWarehouseSourceCategory.ENGINEERING___MONITORING, - label="Amazon Web Services (AWS Config)", + label="AWS Config", + caption="""Sync resource configurations, rules, compliance status, and conformance packs from one AWS region. + +Grant these IAM permissions for the tables you want to sync: +- `config:SelectResourceConfig` +- `config:DescribeConfigRules` +- `config:DescribeComplianceByConfigRule` +- `config:DescribeConformancePacks` + +Enable an AWS Config recorder for resource inventory. The resources table contains current configurations and excludes deleted resources. +Add a session token if you use temporary credentials. Replace temporary credentials when they expire.""", iconPath="/static/services/aws_config.png", - fields=cast(list[FieldType], []), - unreleasedSource=True, + docsUrl="https://docs.aws.amazon.com/config/latest/APIReference/Welcome.html", + releaseStatus=ReleaseStatus.ALPHA, + keywords=["aws", "config", "compliance", "inventory"], + fields=cast( + list[FieldType], + [ + SourceFieldInputConfig( + name="aws_access_key_id", + label="AWS access key ID", + type=SourceFieldInputConfigType.TEXT, + required=True, + placeholder="AKIA...", + secret=False, + ), + SourceFieldInputConfig( + name="aws_secret_access_key", + label="AWS secret access key", + type=SourceFieldInputConfigType.PASSWORD, + required=True, + placeholder="", + secret=True, + ), + SourceFieldInputConfig( + name="aws_session_token", + label="AWS session token", + type=SourceFieldInputConfigType.PASSWORD, + required=False, + placeholder="Only needed for temporary credentials", + secret=True, + ), + SourceFieldInputConfig( + name="region", + label="AWS region", + type=SourceFieldInputConfigType.TEXT, + required=False, + placeholder="us-east-1", + caption="Defaults to us-east-1. Only the selected region is synced.", + secret=False, + ), + ], + ), ) diff --git a/products/warehouse_sources/backend/temporal/data_imports/sources/aws_config/tests/test_aws_config.py b/products/warehouse_sources/backend/temporal/data_imports/sources/aws_config/tests/test_aws_config.py new file mode 100644 index 000000000000..cdaab2857c1e --- /dev/null +++ b/products/warehouse_sources/backend/temporal/data_imports/sources/aws_config/tests/test_aws_config.py @@ -0,0 +1,372 @@ +import json +import datetime as dt +from collections.abc import Iterable, Iterator +from typing import Any, cast + +import pytest +from unittest.mock import MagicMock, patch + +import requests +from tenacity import wait_none + +from products.warehouse_sources.backend.temporal.data_imports.sources.aws_config import aws_config +from products.warehouse_sources.backend.temporal.data_imports.sources.aws_config.aws_config import ( + AwsConfigClient, + AwsConfigError, + AwsConfigResumeConfig, + AwsConfigThrottledError, + aws_config_source, + validate_credentials, +) +from products.warehouse_sources.backend.temporal.data_imports.sources.common.resumable import ResumableSourceManager +from products.warehouse_sources.backend.temporal.data_imports.sources.generated_configs.awsconfig import ( + AwsConfigSourceConfig, +) + + +def make_config(**overrides: Any) -> AwsConfigSourceConfig: + return AwsConfigSourceConfig.from_dict( + { + "aws_access_key_id": "AKIAEXAMPLE", + "aws_secret_access_key": "example-secret", + "aws_session_token": "example-session", + "region": "eu-west-1", + **overrides, + } + ) + + +def make_response(body: Any, status: int = 200, headers: dict[str, str] | None = None) -> requests.Response: + response = requests.Response() + response.status_code = status + response._content = json.dumps(body).encode() + response.headers.update(headers or {}) + return response + + +@pytest.fixture +def session() -> Iterator[MagicMock]: + with patch.object(aws_config, "make_tracked_session") as factory: + factory.return_value.post.return_value = make_response({}) + yield factory.return_value + + +def make_manager(state: AwsConfigResumeConfig | None = None) -> MagicMock: + manager = MagicMock(spec=ResumableSourceManager) + manager.can_resume.return_value = state is not None + manager.load_state.return_value = state + return manager + + +@pytest.mark.parametrize( + "region,token,host,signing_region", + [ + ("eu-west-1", "example-session", "config.eu-west-1.amazonaws.com", "eu-west-1"), + ("cn-north-1", None, "config.cn-north-1.amazonaws.com.cn", "cn-north-1"), + ("us-gov-west-1", None, "config.us-gov-west-1.amazonaws.com", "us-gov-west-1"), + (None, None, "config.us-east-1.amazonaws.com", "us-east-1"), + ], +) +def test_signed_tracked_request(region: str | None, token: str | None, host: str, signing_region: str) -> None: + with patch.object(aws_config, "make_tracked_session") as factory: + factory.return_value.post.return_value = make_response({"ConfigRules": []}) + client = AwsConfigClient(make_config(region=region, aws_session_token=token), "2014-11-12") + assert client.request("DescribeConfigRules", {"NextToken": "page-2"}) == {"ConfigRules": []} + + call = factory.return_value.post.call_args + assert call.args == (f"https://{host}/",) + assert json.loads(call.kwargs["data"]) == {"NextToken": "page-2"} + headers = call.kwargs["headers"] + assert headers["X-Amz-Target"] == "StarlingDoveService.DescribeConfigRules" + assert headers["Content-Type"] == "application/x-amz-json-1.1" + assert headers["Authorization"].startswith("AWS4-HMAC-SHA256 Credential=AKIAEXAMPLE/") + assert f"/{signing_region}/config/aws4_request" in headers["Authorization"] + assert "x-amz-target" in headers["Authorization"] + assert headers["X-Amz-Date"] + assert headers.get("X-Amz-Security-Token") == token + assert call.kwargs["allow_redirects"] is False + assert call.kwargs["timeout"] == 60 + assert "example-secret" in factory.call_args.kwargs["redact_values"] + if token: + assert token in factory.call_args.kwargs["redact_values"] + retry_policy = factory.call_args.kwargs["retry"] + assert retry_policy.is_retry("POST", 503) + assert retry_policy.is_retry("POST", 429) + assert not retry_policy.is_retry("POST", 400) + + +@pytest.mark.parametrize( + "endpoint,result_key,operation,item,expected_key,limit", + [ + ( + "resources", + "Results", + "SelectResourceConfig", + '{"accountId":"111111111111","awsRegion":"eu-west-1","resourceType":"AWS::EC2::Instance","resourceId":"i-example","configuration":{"state":{"name":"running"}},"configurationItemCaptureTime":"2025-01-01T00:00:00Z"}', + "resource_id", + 100, + ), + ( + "config_rules", + "ConfigRules", + "DescribeConfigRules", + {"ConfigRuleArn": "arn:example:rule"}, + "config_rule_arn", + None, + ), + ( + "rule_compliance", + "ComplianceByConfigRules", + "DescribeComplianceByConfigRule", + {"ConfigRuleName": "example-rule", "Compliance": {"ComplianceType": "NON_COMPLIANT"}}, + "config_rule_name", + None, + ), + ( + "conformance_packs", + "ConformancePackDetails", + "DescribeConformancePacks", + {"ConformancePackArn": "arn:example:pack", "LastUpdateRequestedTime": 1735689600}, + "conformance_pack_arn", + 20, + ), + ], +) +def test_full_refresh_paginates_through_empty_and_terminal_pages( + session: MagicMock, endpoint: str, result_key: str, operation: str, item: Any, expected_key: str, limit: int | None +) -> None: + session.post.side_effect = [ + make_response({result_key: [item], "NextToken": "second"}), + make_response({result_key: [], "NextToken": "third"}), + make_response({result_key: [item]}), + ] + manager = make_manager() + response = aws_config_source(make_config(), endpoint, "2014-11-12", manager) + batches = list(cast(Iterable[Any], response.items())) + + assert len(batches) == 2 + assert all(expected_key in batch[0] for batch in batches) + assert all(batch[0]["region"] == "eu-west-1" for batch in batches) + assert all(key in batches[0][0] for key in response.primary_keys or []) + if endpoint == "resources": + assert batches[0][0]["configuration"] == {"state": {"name": "running"}} + assert batches[0][0]["configuration_item_capture_time"] == dt.datetime(2025, 1, 1, tzinfo=dt.UTC) + if endpoint == "conformance_packs": + assert batches[0][0]["last_update_requested_time"] == dt.datetime(2025, 1, 1, tzinfo=dt.UTC) + if endpoint == "rule_compliance": + assert batches[0][0]["compliance"] == {"ComplianceType": "NON_COMPLIANT"} + for call, token in zip(session.post.call_args_list, [None, "second", "third"], strict=True): + payload = json.loads(call.kwargs["data"]) + expected: dict[str, Any] = {} + if limit is not None: + expected["Limit"] = limit + if endpoint == "resources": + expected["Expression"] = ( + "SELECT accountId, awsRegion, arn, resourceId, resourceType, resourceName, " + "configuration, configurationItemCaptureTime, resourceCreationTime, tags" + ) + if token: + expected["NextToken"] = token + assert payload == expected + assert call.kwargs["headers"]["X-Amz-Target"] == f"StarlingDoveService.{operation}" + assert [call.args[0] for call in manager.save_state.call_args_list] == [ + AwsConfigResumeConfig(next_token="second"), + AwsConfigResumeConfig(next_token="third"), + AwsConfigResumeConfig(complete=True), + ] + assert manager.safe_point.call_count == 3 + manager.clear_state.assert_called_once() + session.close.assert_called_once() + + +def test_resume_uses_saved_token(session: MagicMock) -> None: + session.post.return_value = make_response({"ConfigRules": [{"ConfigRuleArn": "arn:example:rule"}]}) + manager = make_manager(AwsConfigResumeConfig(next_token="saved-token")) + response = aws_config_source(make_config(), "config_rules", "2014-11-12", manager) + assert len(list(cast(Iterable[Any], response.items()))) == 1 + assert json.loads(session.post.call_args.kwargs["data"]) == {"NextToken": "saved-token"} + manager.clear_state.assert_called_once() + + +def test_expired_resume_token_clears_state_and_fails_attempt(session: MagicMock) -> None: + session.post.return_value = make_response({"__type": "InvalidNextTokenException"}, 400) + manager = make_manager(AwsConfigResumeConfig(next_token="expired-token")) + response = aws_config_source(make_config(), "config_rules", "2014-11-12", manager) + + with pytest.raises(AwsConfigError, match="InvalidNextTokenException"): + list(cast(Iterable[Any], response.items())) + + assert json.loads(session.post.call_args.kwargs["data"]) == {"NextToken": "expired-token"} + manager.clear_state.assert_called_once() + session.close.assert_called_once() + + +def test_checkpoint_is_staged_before_yield_and_session_closes_on_interruption(session: MagicMock) -> None: + session.post.return_value = make_response( + {"ConfigRules": [{"ConfigRuleArn": "arn:example:rule"}], "NextToken": "next"} + ) + manager = make_manager() + rows = aws_config.get_rows(make_config(), aws_config.AWS_CONFIG_ENDPOINTS["config_rules"], "2014-11-12", manager) + next(rows) + manager.save_state.assert_called_once_with(AwsConfigResumeConfig(next_token="next")) + manager.safe_point.assert_not_called() + rows.close() + session.close.assert_called_once() + manager.clear_state.assert_not_called() + + +def test_completed_resume_does_not_repeat_requests(session: MagicMock) -> None: + response = aws_config_source( + make_config(), "config_rules", "2014-11-12", make_manager(AwsConfigResumeConfig(complete=True)) + ) + assert list(cast(Iterable[Any], response.items())) == [] + session.post.assert_not_called() + + +@pytest.mark.parametrize("token", ["saved-token", 123]) +def test_invalid_page_token_does_not_advance_checkpoint(session: MagicMock, token: Any) -> None: + session.post.return_value = make_response({"ConfigRules": [], "NextToken": token}) + manager = make_manager(AwsConfigResumeConfig(next_token="saved-token")) + with pytest.raises(ValueError, match="invalid page token"): + list(cast(Iterable[Any], aws_config_source(make_config(), "config_rules", "2014-11-12", manager).items())) + manager.save_state.assert_not_called() + session.close.assert_called_once() + + +@pytest.mark.parametrize("body", [{"Results": ["not JSON"]}, {"Results": ["[]"]}, []]) +def test_malformed_results_do_not_advance_checkpoint(session: MagicMock, body: Any) -> None: + session.post.return_value = make_response(body) + manager = make_manager() + with pytest.raises(ValueError): + list(cast(Iterable[Any], aws_config_source(make_config(), "resources", "2014-11-12", manager).items())) + manager.save_state.assert_not_called() + session.close.assert_called_once() + + +@pytest.mark.parametrize( + "status,body,headers,code,throttled", + [ + (400, {"__type": "com.amazonaws.config#ThrottlingException"}, {}, "ThrottlingException", True), + (400, {"code": "Throttling"}, {}, "Throttling", True), + (429, {"__type": "ThrottlingException"}, {}, "ThrottlingException", False), + (503, {"__type": "ServiceUnavailableException"}, {}, "ServiceUnavailableException", False), + ( + 403, + {"Message": "Denied"}, + {"x-amzn-ErrorType": "AccessDeniedException:http"}, + "AccessDeniedException", + False, + ), + (400, {"__type": "UnrecognizedClientException"}, {}, "UnrecognizedClientException", False), + (502, [], {}, "HTTP 502", False), + ], +) +def test_error_classification(status: int, body: Any, headers: dict[str, str], code: str, throttled: bool) -> None: + error = aws_config.error_for_response(make_response(body, status, headers)) + assert error.code == code + assert isinstance(error, AwsConfigThrottledError) is throttled + + +def test_body_throttles_retry_without_sleep(session: MagicMock) -> None: + session.post.side_effect = [ + make_response({"__type": "ThrottlingException", "message": "Rate exceeded"}, 400), + make_response({"ConfigRules": []}), + ] + client = AwsConfigClient(make_config(), "2014-11-12") + with patch.object(AwsConfigClient.request.retry, "wait", wait_none()): # type: ignore[attr-defined] + assert client.request("DescribeConfigRules", {}) == {"ConfigRules": []} + assert session.post.call_count == 2 + + +@pytest.mark.parametrize("schema_name", [None, "config_rules"]) +@pytest.mark.parametrize("code", ["AccessDenied", "AccessDeniedException"]) +def test_access_denied_only_blocks_selected_table(session: MagicMock, schema_name: str | None, code: str) -> None: + session.post.return_value = make_response({"__type": code, "message": "Access denied"}, 400) + valid, reason = validate_credentials(make_config(), "2014-11-12", schema_name) + assert valid is (schema_name is None) + assert reason is None if schema_name is None else "config:DescribeConfigRules" in (reason or "") + session.post.assert_called_once() + session.close.assert_called_once() + + +@pytest.mark.parametrize( + "code,message", + [ + ("UnrecognizedClientException", "Check the access key ID"), + ("InvalidSignatureException", "Check the secret access key"), + ("ExpiredTokenException", "expired"), + ("OptInRequired", "Enable AWS Config"), + ("SubscriptionRequiredException", "Enable AWS Config"), + ("InternalError", "try again"), + ], +) +def test_credential_errors_are_actionable(session: MagicMock, code: str, message: str) -> None: + session.post.return_value = make_response({"__type": code, "message": "Vendor error"}, 400) + valid, reason = validate_credentials(make_config(), "2014-11-12") + assert not valid + assert message in (reason or "") + session.post.assert_called_once() + + +@pytest.mark.parametrize( + "schema,operation,limit", + [ + (None, "DescribeConfigRules", None), + ("resources", "SelectResourceConfig", 1), + ("conformance_packs", "DescribeConformancePacks", 1), + ], +) +def test_credential_probe_is_one_request( + session: MagicMock, schema: str | None, operation: str, limit: int | None +) -> None: + assert validate_credentials(make_config(), "2014-11-12", schema) == (True, None) + session.post.assert_called_once() + call = session.post.call_args + assert call.kwargs["headers"]["X-Amz-Target"] == f"StarlingDoveService.{operation}" + assert json.loads(call.kwargs["data"]).get("Limit") == limit + + +@pytest.mark.parametrize("region", ["https://example.com", "us-east-1.example.com/path", "us-east-1\n", "../us-east-1"]) +def test_invalid_region_fails_before_transport(session: MagicMock, region: str) -> None: + valid, reason = validate_credentials(make_config(region=region), "2014-11-12") + assert not valid + assert reason == "Enter an AWS region such as us-east-1." + session.post.assert_not_called() + + +@pytest.mark.parametrize("field", ["aws_access_key_id", "aws_secret_access_key"]) +def test_missing_credentials_fail_before_transport(session: MagicMock, field: str) -> None: + valid, reason = validate_credentials(make_config(**{field: ""}), "2014-11-12") + assert not valid + assert "Enter both" in (reason or "") + session.post.assert_not_called() + + +def test_unknown_table_and_version_fail_before_transport(session: MagicMock) -> None: + assert validate_credentials(make_config(), "2014-11-12", "unknown") == (False, "Unknown AWS Config table: unknown") + with pytest.raises(ValueError, match="Unknown AWS Config table"): + aws_config_source(make_config(), "unknown", "2014-11-12", make_manager()) + valid, reason = validate_credentials(make_config(), "invalid-version") + assert not valid + assert "Unsupported AWS Config API version" in (reason or "") + session.post.assert_not_called() + + +def test_network_failure_closes_probe_session(session: MagicMock) -> None: + session.post.side_effect = requests.Timeout() + valid, reason = validate_credentials(make_config(), "2014-11-12") + assert not valid + assert "try again" in (reason or "") + session.close.assert_called_once() + + +def test_new_invalid_token_is_not_silently_restarted(session: MagicMock) -> None: + session.post.return_value = make_response({"__type": "InvalidNextTokenException"}, 400) + with pytest.raises(AwsConfigError, match="InvalidNextTokenException"): + list( + cast( + Iterable[Any], + aws_config_source(make_config(), "config_rules", "2014-11-12", make_manager()).items(), + ) + ) + session.post.assert_called_once() diff --git a/products/warehouse_sources/backend/temporal/data_imports/sources/aws_config/tests/test_aws_config_source.py b/products/warehouse_sources/backend/temporal/data_imports/sources/aws_config/tests/test_aws_config_source.py new file mode 100644 index 000000000000..89d5a5b24442 --- /dev/null +++ b/products/warehouse_sources/backend/temporal/data_imports/sources/aws_config/tests/test_aws_config_source.py @@ -0,0 +1,28 @@ +import pytest + +from products.warehouse_sources.backend.temporal.data_imports.sources.aws_config.aws_config import AwsConfigError +from products.warehouse_sources.backend.temporal.data_imports.sources.aws_config.source import AwsConfigSource + + +@pytest.mark.parametrize( + "code,expected", + [ + ("AccessDeniedException", "Grant"), + ("UnrecognizedClientException", "Check"), + ("InvalidSignatureException", "signature"), + ("ExpiredTokenException", "expired"), + ("SubscriptionRequiredException", "Enable AWS Config"), + ("ThrottlingException", None), + ("HTTP 503", None), + ], +) +def test_sync_errors_match_actionable_terminal_messages(code: str, expected: str | None) -> None: + error = str(AwsConfigError(code, "Example AWS response")) + messages = [ + message for pattern, message in AwsConfigSource().get_non_retryable_errors().items() if pattern in error + ] + if expected is None: + assert messages == [] + else: + assert messages + assert all(expected in (message or "") for message in messages) diff --git a/products/warehouse_sources/backend/temporal/data_imports/sources/generated_configs/awsconfig.py b/products/warehouse_sources/backend/temporal/data_imports/sources/generated_configs/awsconfig.py index b89b025fe2c8..0ea5ba884c5b 100644 --- a/products/warehouse_sources/backend/temporal/data_imports/sources/generated_configs/awsconfig.py +++ b/products/warehouse_sources/backend/temporal/data_imports/sources/generated_configs/awsconfig.py @@ -6,4 +6,7 @@ @config.config class AwsConfigSourceConfig(config.Config): - pass + aws_access_key_id: str + aws_secret_access_key: str + aws_session_token: str | None = None + region: str | None = None