Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -54,8 +54,10 @@ def __init__(
super().__init__(**kwargs)
self.source_aws_conn_id = source_aws_conn_id
self.dest_aws_conn_id = dest_aws_conn_id
self.source_aws_conn_id = source_aws_conn_id
if is_arg_set(dest_aws_conn_id):
self.dest_aws_conn_id = dest_aws_conn_id
else:
self.dest_aws_conn_id = self.source_aws_conn_id

@property
def resolved_dest_aws_conn_id(self) -> str | None:
"""Destination connection id, falling back to source_aws_conn_id when not set."""
if is_arg_set(self.dest_aws_conn_id):
return self.dest_aws_conn_id
return self.source_aws_conn_id
Original file line number Diff line number Diff line change
Expand Up @@ -224,7 +224,9 @@ def _export_entire_data(self):
raise e
finally:
if err is None:
_upload_file_to_s3(f, self.s3_bucket_name, self.s3_key_prefix, self.dest_aws_conn_id)
_upload_file_to_s3(
f, self.s3_bucket_name, self.s3_key_prefix, self.resolved_dest_aws_conn_id
)

def _scan_dynamodb_and_upload_to_s3(self, temp_file: IO, scan_kwargs: dict, table: Any) -> IO:
while True:
Expand All @@ -242,7 +244,9 @@ def _scan_dynamodb_and_upload_to_s3(self, temp_file: IO, scan_kwargs: dict, tabl

# Upload the file to S3 if reach file size limit
if os.path.getsize(temp_file.name) >= self.file_size:
_upload_file_to_s3(temp_file, self.s3_bucket_name, self.s3_key_prefix, self.dest_aws_conn_id)
_upload_file_to_s3(
temp_file, self.s3_bucket_name, self.s3_key_prefix, self.resolved_dest_aws_conn_id
)
temp_file.close()

temp_file = NamedTemporaryFile()
Expand Down
17 changes: 17 additions & 0 deletions providers/amazon/tests/unit/amazon/aws/transfers/test_base.py
Original file line number Diff line number Diff line change
Expand Up @@ -73,3 +73,20 @@ def test_render_template(self, session, clean_dags_dagruns_and_dagbundles, testi
render_template_fields(ti, operator)
assert getattr(operator, "source_aws_conn_id") == "2020-01-01"
assert getattr(operator, "dest_aws_conn_id") == "2020-01-01"

@pytest.mark.parametrize(
("dest_kwargs", "expected"),
[
pytest.param({}, "source-conn", id="fallback-to-source"),
pytest.param({"dest_aws_conn_id": "dest-conn"}, "dest-conn", id="explicit-dest"),
pytest.param({"dest_aws_conn_id": None}, None, id="explicit-none"),
],
)
def test_resolved_dest_aws_conn_id(self, dest_kwargs, expected):
operator = AwsToAwsBaseOperator(
task_id="test_resolved_dest",
dag=self.dag,
source_aws_conn_id="source-conn",
**dest_kwargs,
)
assert operator.resolved_dest_aws_conn_id == expected
1 change: 0 additions & 1 deletion scripts/ci/prek/validate_operators_init_exemptions.txt
Original file line number Diff line number Diff line change
Expand Up @@ -14,7 +14,6 @@ providers/amazon/src/airflow/providers/amazon/aws/operators/s3.py::S3DeleteObjec
providers/amazon/src/airflow/providers/amazon/aws/operators/sagemaker.py::SageMakerCreateNotebookOperator
providers/amazon/src/airflow/providers/amazon/aws/operators/sagemaker.py::SageMakerProcessingOperator
providers/amazon/src/airflow/providers/amazon/aws/operators/step_function.py::StepFunctionStartExecutionOperator
providers/amazon/src/airflow/providers/amazon/aws/transfers/base.py::AwsToAwsBaseOperator
providers/amazon/src/airflow/providers/amazon/aws/transfers/gcs_to_s3.py::GCSToS3Operator
providers/amazon/src/airflow/providers/amazon/aws/transfers/s3_to_redshift.py::S3ToRedshiftOperator
providers/cncf/kubernetes/src/airflow/providers/cncf/kubernetes/operators/pod.py::KubernetesPodOperator
Expand Down