Skip to content
Merged
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
7 changes: 6 additions & 1 deletion airflow/providers/amazon/aws/transfers/s3_to_sql.py
Original file line number Diff line number Diff line change
Expand Up @@ -44,6 +44,8 @@ class S3ToSqlOperator(BaseOperator):
:param s3_bucket: reference to a specific S3 bucket
:param s3_key: reference to a specific S3 key
:param sql_conn_id: reference to a specific SQL database. Must be of type DBApiHook
:param sql_hook_params: Extra config params to be passed to the underlying hook.
Should match the desired hook constructor params.
:param aws_conn_id: reference to a specific S3 / AWS connection
:param column_list: list of column names to use in the insert SQL.
:param commit_every: The maximum number of rows to insert in one
Expand Down Expand Up @@ -83,6 +85,7 @@ def __init__(
commit_every: int = 1000,
schema: str | None = None,
sql_conn_id: str = "sql_default",
sql_hook_params: dict | None = None,
aws_conn_id: str = "aws_default",
**kwargs,
) -> None:
Expand All @@ -96,6 +99,7 @@ def __init__(
self.column_list = column_list
self.commit_every = commit_every
self.parser = parser
self.sql_hook_params = sql_hook_params

def execute(self, context: Context) -> None:
self.log.info("Loading %s to SQL table %s...", self.s3_key, self.table)
Expand All @@ -120,7 +124,8 @@ def execute(self, context: Context) -> None:
@cached_property
def db_hook(self):
self.log.debug("Get connection for %s", self.sql_conn_id)
hook = BaseHook.get_hook(self.sql_conn_id)
conn = BaseHook.get_connection(self.sql_conn_id)
hook = conn.get_hook(hook_params=self.sql_hook_params)
if not callable(getattr(hook, "insert_rows", None)):
raise AirflowException(
"This hook is not supported. The hook class must have an `insert_rows` method."
Expand Down
17 changes: 14 additions & 3 deletions tests/providers/amazon/aws/transfers/test_s3_to_sql.py
Original file line number Diff line number Diff line change
Expand Up @@ -76,7 +76,7 @@ def mock_bad_hook(self):
return bad_hook

@patch("airflow.providers.amazon.aws.transfers.s3_to_sql.NamedTemporaryFile")
@patch("airflow.providers.amazon.aws.transfers.s3_to_sql.BaseHook")
@patch("airflow.models.connection.Connection.get_hook")
@patch("airflow.providers.amazon.aws.transfers.s3_to_sql.S3Hook.get_key")
def test_execute(self, mock_get_key, mock_hook, mock_tempfile, mock_parser):

Expand All @@ -93,7 +93,7 @@ def test_execute(self, mock_get_key, mock_hook, mock_tempfile, mock_parser):

mock_parser.assert_called_once_with(mock_tempfile.return_value.__enter__.return_value.name)

mock_hook.get_hook.return_value.insert_rows.assert_called_once_with(
mock_hook.return_value.insert_rows.assert_called_once_with(
table=self.s3_to_sql_transfer_kwargs["table"],
schema=self.s3_to_sql_transfer_kwargs["schema"],
target_fields=self.s3_to_sql_transfer_kwargs["column_list"],
Expand All @@ -102,13 +102,24 @@ def test_execute(self, mock_get_key, mock_hook, mock_tempfile, mock_parser):
)

@patch("airflow.providers.amazon.aws.transfers.s3_to_sql.NamedTemporaryFile")
@patch("airflow.providers.amazon.aws.transfers.s3_to_sql.BaseHook.get_hook", return_value=mock_bad_hook)
@patch("airflow.models.connection.Connection.get_hook", return_value=mock_bad_hook)
@patch("airflow.providers.amazon.aws.transfers.s3_to_sql.S3Hook.get_key")
def test_execute_with_bad_hook(self, mock_get_key, mock_bad_hook, mock_tempfile, mock_parser):

with pytest.raises(AirflowException):
S3ToSqlOperator(parser=mock_parser, **self.s3_to_sql_transfer_kwargs).execute({})

def test_hook_params(self, mock_parser):
op = S3ToSqlOperator(
parser=mock_parser,
sql_hook_params={
"log_sql": False,
},
**self.s3_to_sql_transfer_kwargs,
)
hook = op.db_hook
assert hook.log_sql == op.sql_hook_params["log_sql"]

def teardown_method(self):
with create_session() as session:
(
Expand Down