diff --git a/airflow/providers/google/cloud/hooks/gcs.py b/airflow/providers/google/cloud/hooks/gcs.py index b35f647438136..1547c41794056 100644 --- a/airflow/providers/google/cloud/hooks/gcs.py +++ b/airflow/providers/google/cloud/hooks/gcs.py @@ -1338,7 +1338,15 @@ def _prepare_sync_plan( for current_name in names_to_check: source_blob = source_names_index[current_name] destination_blob = destination_names_index[current_name] - # If the objects are different, save it + # If either object is CMEK-protected, use the Cloud Storage Objects Get API to retrieve them + # so that the crc32c is included + if source_blob.kms_key_name: + source_blob = source_bucket.get_blob(source_blob.name, generation=source_blob.generation) + if destination_blob.kms_key_name: + destination_blob = destination_bucket.get_blob( + destination_blob.name, generation=destination_blob.generation + ) + # if the objects are different, save it if source_blob.crc32c != destination_blob.crc32c: to_rewrite_blobs.add(source_blob) diff --git a/tests/providers/google/cloud/hooks/test_gcs.py b/tests/providers/google/cloud/hooks/test_gcs.py index 7759b7de35823..5501b568609a6 100644 --- a/tests/providers/google/cloud/hooks/test_gcs.py +++ b/tests/providers/google/cloud/hooks/test_gcs.py @@ -25,6 +25,7 @@ from datetime import datetime, timedelta from io import BytesIO from unittest import mock +from unittest.mock import MagicMock import dateutil import pytest @@ -1279,6 +1280,47 @@ def test_should_overwrite_files(self, mock_get_conn, mock_delete, mock_rewrite, ) mock_copy.assert_not_called() + @mock.patch(GCS_STRING.format("GCSHook.copy")) + @mock.patch(GCS_STRING.format("GCSHook.rewrite")) + @mock.patch(GCS_STRING.format("GCSHook.delete")) + @mock.patch(GCS_STRING.format("GCSHook.get_conn")) + def test_should_overwrite_cmek_files(self, mock_get_conn, mock_delete, mock_rewrite, mock_copy): + source_bucket = self._create_bucket(name="SOURCE_BUCKET") + source_bucket.list_blobs.return_value = [ + self._create_blob("FILE_A", "C1", kms_key_name="KMS_KEY_1", generation=1), + self._create_blob("FILE_B", "C1"), + ] + destination_bucket = self._create_bucket(name="DEST_BUCKET") + destination_bucket.list_blobs.return_value = [ + self._create_blob("FILE_A", "C2", kms_key_name="KMS_KEY_2", generation=2), + self._create_blob("FILE_B", "C2"), + ] + mock_get_conn.return_value.bucket.side_effect = [source_bucket, destination_bucket] + self.gcs_hook.sync( + source_bucket="SOURCE_BUCKET", destination_bucket="DEST_BUCKET", allow_overwrite=True + ) + mock_delete.assert_not_called() + source_bucket.get_blob.assert_called_once_with("FILE_A", generation=1) + destination_bucket.get_blob.assert_called_once_with("FILE_A", generation=2) + mock_rewrite.assert_has_calls( + [ + mock.call( + source_bucket="SOURCE_BUCKET", + source_object="FILE_B", + destination_bucket="DEST_BUCKET", + destination_object="FILE_B", + ), + mock.call( + source_bucket="SOURCE_BUCKET", + source_object=source_bucket.get_blob.return_value.name, + destination_bucket="DEST_BUCKET", + destination_object=source_bucket.get_blob.return_value.name.__getitem__.return_value, + ), + ], + any_order=True, + ) + mock_copy.assert_not_called() + @mock.patch(GCS_STRING.format("GCSHook.copy")) @mock.patch(GCS_STRING.format("GCSHook.rewrite")) @mock.patch(GCS_STRING.format("GCSHook.delete")) @@ -1440,11 +1482,20 @@ def test_should_not_overwrite_when_overwrite_is_disabled( mock_rewrite.assert_not_called() mock_copy.assert_not_called() - def _create_blob(self, name: str, crc32: str, bucket=None): + def _create_blob( + self, + name: str, + crc32: str, + bucket: MagicMock | None = None, + kms_key_name: str | None = None, + generation: int = 0, + ): blob = mock.MagicMock(name=f"BLOB:{name}") blob.name = name blob.crc32 = crc32 blob.bucket = bucket + blob.kms_key_name = kms_key_name + blob.generation = generation return blob def _create_bucket(self, name: str):