"""Lambda test module.""" import pytest from update_waf_ipsets import app from update_waf_ipsets import config as Config @pytest.fixture() def patch_sts_client(sts_client, monkeypatch): """Patch the global STS client with a mocked one.""" monkeypatch.setattr(app, "sts_client", sts_client) return sts_client @pytest.fixture() def patch_s3_client(s3_client, monkeypatch): """Patch the global S3 client with a mocked one.""" monkeypatch.setattr(app, "s3_client", s3_client) return s3_client def test_assume_source_account_role(patch_sts_client): creds = app.assume_source_account_role( account_id="123456789012", role_name=Config.DEFAULT_ROLE_NAME, external_id="TEST_ID" ) assert set(["aws_access_key_id", "aws_secret_access_key", "aws_session_token"]) == set(creds.keys()) def test_get_ip_addresses_valid(patch_s3_client, s3_bucket_name, valid_s3_object, valid_ip_list): ip_list = app.get_ip_addresses(s3_bucket_name, valid_s3_object) assert ip_list == valid_ip_list def test_get_ip_addresses_invalid(patch_s3_client, s3_bucket_name, invalid_s3_object): with pytest.raises(ValueError) as exc_info: _ = app.get_ip_addresses(s3_bucket_name, invalid_s3_object) assert "does not appear to be an IPv4 or IPv6 network" in str(exc_info.value) def test_get_ip_set_by_name_cloudfront(wafv2_client, cloudfront_ip_set): ip_set = app.get_ip_set_by_name(wafv2_client, cloudfront_ip_set["Name"], "CLOUDFRONT") assert ip_set["ARN"].startswith(f"arn:aws:wafv2:us-east-1:123456789012:global/ipset/{Config.DEFAULT_IP_SET_NAME}") assert ip_set["Name"] == Config.DEFAULT_IP_SET_NAME def test_get_ip_set_by_name_regional(wafv2_client, regional_ip_set): ip_set = app.get_ip_set_by_name(wafv2_client, regional_ip_set["Name"], "REGIONAL") assert ip_set["ARN"].startswith(f"arn:aws:wafv2:us-east-1:123456789012:regional/ipset/{Config.DEFAULT_IP_SET_NAME}") assert ip_set["Name"] == Config.DEFAULT_IP_SET_NAME def test_get_ip_set_by_name_empty(wafv2_client): with pytest.raises(ValueError) as exc_info: _ = app.get_ip_set_by_name(wafv2_client, "non-existent-set", "REGIONAL") assert "IP Set non-existent-set is not found." in str(exc_info.value) def test_update_waf_ip_set_cloudfront(wafv2_client, cloudfront_ip_set, valid_ip_list): app.update_waf_ip_set(wafv2_client, valid_ip_list, "CLOUDFRONT", cloudfront_ip_set["Name"]) ip_set = wafv2_client.get_ip_set(Name=cloudfront_ip_set["Name"], Id=cloudfront_ip_set["Id"], Scope="CLOUDFRONT") assert ip_set["IPSet"]["Addresses"] == valid_ip_list def test_update_waf_ip_set_regional(wafv2_client, regional_ip_set, valid_ip_list): app.update_waf_ip_set(wafv2_client, valid_ip_list, "REGIONAL", regional_ip_set["Name"]) ip_set = wafv2_client.get_ip_set(Name=regional_ip_set["Name"], Id=regional_ip_set["Id"], Scope="REGIONAL") assert ip_set["IPSet"]["Addresses"] == valid_ip_list assert ip_set["IPSet"]["Description"] == "Test set" def test_update_waf_ip_set_regional_empty_desc(wafv2_client, regional_ip_set_empty_desc, valid_ip_list): app.update_waf_ip_set(wafv2_client, valid_ip_list, "REGIONAL", regional_ip_set_empty_desc["Name"]) ip_set = wafv2_client.get_ip_set( Name=regional_ip_set_empty_desc["Name"], Id=regional_ip_set_empty_desc["Id"], Scope="REGIONAL" ) assert ip_set["IPSet"]["Addresses"] == valid_ip_list assert "Description" not in ip_set["IPSet"]