"""Lambda test module.""" import os from unittest.mock import MagicMock import boto3 import pytest from moto import mock_aws from update_dns_firewall import app from update_dns_firewall import config as Config @pytest.fixture() def aws_credentials(): """Mock AWS Credentials for moto.""" os.environ["AWS_ACCESS_KEY_ID"] = "testing" os.environ["AWS_SECRET_ACCESS_KEY"] = "testing" os.environ["AWS_SECURITY_TOKEN"] = "testing" os.environ["AWS_SESSION_TOKEN"] = "testing" os.environ["AWS_DEFAULT_REGION"] = "us-east-1" @pytest.fixture() def s3_client(aws_credentials): """Return an S3 client.""" with mock_aws(): yield boto3.client("s3") @pytest.fixture() def route53resolver_client(aws_credentials): """Return an Route53 Resolver client.""" with mock_aws(): yield boto3.client("route53resolver") @pytest.fixture def mock_route53resolver_client(monkeypatch): """ Mock Route53 Resolver client. The moto library currently does not support Route53 Resolver, so we mock the client instead. """ mock_client = MagicMock() monkeypatch.setattr("boto3.client", lambda *args, **kwargs: mock_client) return mock_client @pytest.fixture() def valid_domain_list(): """Provides valid domain blocks list.""" return [f"domain{i}.com" for i in range(2003)] @pytest.fixture() def s3_bucket_name(): """Test bucket name.""" return "test-s3-bucket" @pytest.fixture() def s3_bucket(s3_client, s3_bucket_name): """Create bucket for domain list storage.""" _ = s3_client.create_bucket( Bucket=s3_bucket_name, ) return s3_bucket_name @pytest.fixture() def valid_s3_object(s3_client, s3_bucket, valid_domain_list): """Create a file in s3 containing domain list""" key = "domain_list.txt" _ = s3_client.put_object(Bucket=s3_bucket, Key=key, Body="\n".join(valid_domain_list).encode("utf-8")) return key def test_get_domains(s3_client, s3_bucket_name, valid_s3_object, valid_domain_list): domain_list = app.get_domains(s3_bucket_name, valid_s3_object) assert domain_list == valid_domain_list def test_update_firewall_domain_list_chunks(mock_route53resolver_client, valid_domain_list): domain_list_id = Config.DOMAIN_LIST_ID region = Config.DEFAULT_REGION domains = valid_domain_list app.update_firewall_domain_list(domain_list_id, domains, region) assert mock_route53resolver_client.update_firewall_domains.call_count == 3 calls = mock_route53resolver_client.update_firewall_domains.call_args_list assert calls[0][1]["Operation"] == "REPLACE" assert calls[1][1]["Operation"] == "ADD" assert calls[2][1]["Operation"] == "ADD" assert len(calls[0][1]["Domains"]) == 1000 assert len(calls[1][1]["Domains"]) == 1000 assert len(calls[2][1]["Domains"]) == 3