"""Tests for frontend_rollback.""" import unittest from unittest.mock import MagicMock, patch import pytest from botocore.exceptions import ClientError ENV_VARS = { "APP_NAME": "test-app", "ENV": "qa", "NUMBER_OF_DEPLOYMENTS_TO_ROLLBACK": "1", } CUSTOM_BUCKET_ENV_VARS = { "APP_NAME": "test-app", "ENV": "qa", "NUMBER_OF_DEPLOYMENTS_TO_ROLLBACK": "1", "S3_BUCKET": "app.qa.theorchard.io", } TEST_SCENARIOS = [ pytest.param(ENV_VARS, "qa-orcd-cdn", id="default_bucket"), pytest.param( CUSTOM_BUCKET_ENV_VARS, "app.qa.theorchard.io", id="custom_bucket" ), ] ACCOUNT_ID = "123456789012" def _client_error(code, message=""): return ClientError( {"Error": {"Code": code, "Message": message}}, "test_operation", ) @patch.dict("os.environ", ENV_VARS) class TestRevertObject(unittest.TestCase): """Tests for the revert_object function.""" def _import_module(self): """Import after env vars are patched so module-level globals resolve.""" import importlib import frontend_rollback as mod importlib.reload(mod) return mod def test_successful_revert(self): """Verify a successful revert returns the new version ID.""" mod = self._import_module() client = MagicMock() client.get_object.return_value = {} client.list_object_versions.return_value = { "Versions": [ {"Key": "test-app/index.html", "VersionId": "v2"}, {"Key": "test-app/index.html", "VersionId": "v1"}, ] } client.copy_object.return_value = {"VersionId": "v3"} result = mod.revert_object( client, "qa-orcd-cdn", "test-app/index.html", ACCOUNT_ID ) self.assertEqual(result, "v3") client.copy_object.assert_called_once() call_kwargs = client.copy_object.call_args[1] self.assertEqual(call_kwargs["CopySource"]["VersionId"], "v1") self.assertEqual(call_kwargs["ExpectedBucketOwner"], ACCOUNT_ID) def test_passes_expected_bucket_owner_to_all_calls(self): """Verify ExpectedBucketOwner is passed to get, list, and copy calls.""" mod = self._import_module() client = MagicMock() client.get_object.return_value = {} client.list_object_versions.return_value = { "Versions": [ {"Key": "test-app/index.html", "VersionId": "v2"}, {"Key": "test-app/index.html", "VersionId": "v1"}, ] } client.copy_object.return_value = {"VersionId": "v3"} mod.revert_object( client, "qa-orcd-cdn", "test-app/index.html", ACCOUNT_ID ) get_kwargs = client.get_object.call_args[1] self.assertEqual(get_kwargs["ExpectedBucketOwner"], ACCOUNT_ID) list_kwargs = client.list_object_versions.call_args[1] self.assertEqual(list_kwargs["ExpectedBucketOwner"], ACCOUNT_ID) copy_kwargs = client.copy_object.call_args[1] self.assertEqual(copy_kwargs["ExpectedBucketOwner"], ACCOUNT_ID) def test_skips_missing_key(self): """Verify None is returned when the key does not exist (NoSuchKey).""" mod = self._import_module() client = MagicMock() client.get_object.side_effect = _client_error("NoSuchKey") result = mod.revert_object( client, "qa-orcd-cdn", "test-app/index-sme.html", ACCOUNT_ID ) self.assertIsNone(result) client.list_object_versions.assert_not_called() def test_skips_404(self): """Verify None is returned when the key does not exist (404).""" mod = self._import_module() client = MagicMock() client.get_object.side_effect = _client_error("404") result = mod.revert_object( client, "qa-orcd-cdn", "test-app/index-sme.html", ACCOUNT_ID ) self.assertIsNone(result) def test_raises_on_other_client_error(self): """Verify SystemExit is raised for non-404 client errors.""" mod = self._import_module() client = MagicMock() client.get_object.side_effect = _client_error( "AccessDenied", "Forbidden" ) with self.assertRaises(SystemExit): mod.revert_object( client, "qa-orcd-cdn", "test-app/index.html", ACCOUNT_ID ) def test_skips_when_not_enough_versions(self): """Verify None returned when too few versions exist to roll back.""" mod = self._import_module() client = MagicMock() client.get_object.return_value = {} client.list_object_versions.return_value = { "Versions": [ {"Key": "test-app/index.html", "VersionId": "v1"}, ] } result = mod.revert_object( client, "qa-orcd-cdn", "test-app/index.html", ACCOUNT_ID ) self.assertIsNone(result) client.copy_object.assert_not_called() def test_skips_when_no_versions_key(self): """Verify None is returned when Versions key is absent from response.""" mod = self._import_module() client = MagicMock() client.get_object.return_value = {} client.list_object_versions.return_value = {} result = mod.revert_object( client, "qa-orcd-cdn", "test-app/index.html", ACCOUNT_ID ) self.assertIsNone(result) def test_skips_when_key_mismatch(self): """Verify None returned when response key mismatches requested key.""" mod = self._import_module() client = MagicMock() client.get_object.return_value = {} client.list_object_versions.return_value = { "Versions": [ {"Key": "test-app/OTHER-FILE", "VersionId": "v2"}, {"Key": "test-app/OTHER-FILE", "VersionId": "v1"}, ] } result = mod.revert_object( client, "qa-orcd-cdn", "test-app/index.html", ACCOUNT_ID ) self.assertIsNone(result) client.copy_object.assert_not_called() def _import_module(env_vars): import importlib with patch.dict("os.environ", env_vars, clear=True): import frontend_rollback as mod importlib.reload(mod) return mod def _make_client(mock_boto_client, index_pages=None, js_pages=None): """Set up mock S3 and STS clients with paginator support.""" s3_client = MagicMock() sts_client = MagicMock() sts_client.get_caller_identity.return_value = {"Account": ACCOUNT_ID} def client_factory(service): if service == "sts": return sts_client return s3_client mock_boto_client.side_effect = client_factory # Default: prefix exists s3_client.list_objects_v2.return_value = {"KeyCount": 1} def get_paginator(operation): paginator = MagicMock() if operation != "list_objects_v2": return paginator def paginate(**kwargs): prefix = kwargs.get("Prefix", "") if prefix.endswith("/index") and index_pages is not None: return iter(index_pages) if not prefix.endswith("/index") and js_pages is not None: return iter(js_pages) return iter([]) paginator.paginate.side_effect = paginate return paginator s3_client.get_paginator.side_effect = get_paginator return s3_client @pytest.mark.parametrize("env_vars,expected_bucket", TEST_SCENARIOS) @patch("boto3.client") def test_aborts_when_folder_missing( mock_boto_client, env_vars, expected_bucket ): """Verify SystemExit is raised when the app folder is missing in S3.""" mod = _import_module(env_vars) client = _make_client(mock_boto_client) client.list_objects_v2.return_value = {"KeyCount": 0} with pytest.raises(SystemExit, match="No objects found"): mod.main() @pytest.mark.parametrize("env_vars,expected_bucket", TEST_SCENARIOS) @patch("boto3.client") def test_aborts_when_no_index_html(mock_boto_client, env_vars, expected_bucket): """Verify SystemExit is raised when no index*.html is found.""" mod = _import_module(env_vars) _make_client( mock_boto_client, index_pages=[ { "Contents": [{"Key": "test-app/index.js"}], } ], js_pages=[], ) with pytest.raises(SystemExit, match="No index"): mod.main() @pytest.mark.parametrize("env_vars,expected_bucket", TEST_SCENARIOS) @patch("boto3.client") def test_aborts_when_no_index_html_empty_pages( mock_boto_client, env_vars, expected_bucket ): """Verify SystemExit is raised when paginator returns no pages.""" mod = _import_module(env_vars) _make_client( mock_boto_client, index_pages=[], js_pages=[], ) with pytest.raises(SystemExit, match="No index"): mod.main() @pytest.mark.parametrize("env_vars,expected_bucket", TEST_SCENARIOS) @patch("boto3.client") def test_aborts_when_no_js_files(mock_boto_client, env_vars, expected_bucket): """Verify SystemExit is raised when no .js files exist under the prefix.""" mod = _import_module(env_vars) _make_client( mock_boto_client, index_pages=[ { "Contents": [{"Key": "test-app/index.html"}], } ], js_pages=[ { "Contents": [ {"Key": "test-app/index.html"}, {"Key": "test-app/manifest.json"}, ], } ], ) with pytest.raises(SystemExit, match="No .js files"): mod.main() @pytest.mark.parametrize("env_vars,expected_bucket", TEST_SCENARIOS) @patch("boto3.client") def test_finds_index_html_across_pages( mock_boto_client, env_vars, expected_bucket ): """index*.html in the second page should still pass validation.""" mod = _import_module(env_vars) client = _make_client( mock_boto_client, index_pages=[ {"Contents": [{"Key": "test-app/index.js"}]}, {"Contents": [{"Key": "test-app/index.html"}]}, ], js_pages=[ { "Contents": [{"Key": "test-app/bundle.js"}], } ], ) client.get_object.side_effect = _client_error("NoSuchKey") mod.main() assert client.get_object.call_count == len(mod.S3_OBJECTS_TO_REVERT) @pytest.mark.parametrize("env_vars,expected_bucket", TEST_SCENARIOS) @patch("boto3.client") def test_finds_js_across_pages(mock_boto_client, env_vars, expected_bucket): """A .js file in the second page should still pass validation.""" mod = _import_module(env_vars) client = _make_client( mock_boto_client, index_pages=[ { "Contents": [{"Key": "test-app/index.html"}], } ], js_pages=[ {"Contents": [{"Key": "test-app/style.css"}]}, {"Contents": [{"Key": "test-app/bundle.js"}]}, ], ) client.get_object.side_effect = _client_error("NoSuchKey") mod.main() assert client.get_object.call_count == len(mod.S3_OBJECTS_TO_REVERT) @pytest.mark.parametrize("env_vars,expected_bucket", TEST_SCENARIOS) @patch("boto3.client") def test_happy_path_processes_all_objects( mock_boto_client, env_vars, expected_bucket ): """Verify all S3_OBJECTS_TO_REVERT are processed on the happy path.""" mod = _import_module(env_vars) client = _make_client( mock_boto_client, index_pages=[ { "Contents": [{"Key": "test-app/index.html"}], } ], js_pages=[ { "Contents": [{"Key": "test-app/bundle.js"}], } ], ) client.get_object.side_effect = _client_error("NoSuchKey") mod.main() assert client.get_object.call_count == len(mod.S3_OBJECTS_TO_REVERT) @pytest.mark.parametrize("env_vars,expected_bucket", TEST_SCENARIOS) @patch("boto3.client") def test_uses_correct_bucket(mock_boto_client, env_vars, expected_bucket): """Verify that the correct bucket is used for each scenario.""" mod = _import_module(env_vars) assert mod.S3_BUCKETS == [expected_bucket] client = _make_client( mock_boto_client, index_pages=[ { "Contents": [{"Key": "test-app/index.html"}], } ], js_pages=[ { "Contents": [{"Key": "test-app/bundle.js"}], } ], ) client.get_object.side_effect = _client_error("NoSuchKey") mod.main() # Verify that the expected bucket was used in list_objects_v2 calls list_calls = [call for call in client.list_objects_v2.call_args_list] for call in list_calls: assert call[1]["Bucket"] == expected_bucket if __name__ == "__main__": unittest.main()