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
8 changes: 5 additions & 3 deletions airflow/providers/amazon/aws/operators/ecs.py
Original file line number Diff line number Diff line change
Expand Up @@ -49,9 +49,11 @@
class EcsBaseOperator(BaseOperator):
"""This is the base operator for all Elastic Container Service operators."""

def __init__(self, **kwargs):
self.aws_conn_id = kwargs.get('aws_conn_id', DEFAULT_CONN_ID)
self.region = kwargs.get('region')
def __init__(
self, *, aws_conn_id: Optional[str] = DEFAULT_CONN_ID, region: Optional[str] = None, **kwargs
):
self.aws_conn_id = aws_conn_id
self.region = region
super().__init__(**kwargs)

@cached_property
Expand Down
8 changes: 5 additions & 3 deletions airflow/providers/amazon/aws/sensors/ecs.py
Original file line number Diff line number Diff line change
Expand Up @@ -45,9 +45,11 @@ def _check_failed(current_state, target_state, failure_states):
class EcsBaseSensor(BaseSensorOperator):
"""Contains general sensor behavior for Elastic Container Service."""

def __init__(self, **kwargs):
self.aws_conn_id = kwargs.get('aws_conn_id', DEFAULT_CONN_ID)
self.region = kwargs.get('region')
def __init__(
self, *, aws_conn_id: Optional[str] = DEFAULT_CONN_ID, region: Optional[str] = None, **kwargs
):
self.aws_conn_id = aws_conn_id
self.region = region
super().__init__(**kwargs)

@cached_property
Expand Down
39 changes: 39 additions & 0 deletions tests/providers/amazon/aws/operators/test_ecs.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,6 +29,7 @@
from airflow.providers.amazon.aws.exceptions import EcsOperatorError, EcsTaskFailToStart
from airflow.providers.amazon.aws.hooks.ecs import EcsHook
from airflow.providers.amazon.aws.operators.ecs import (
DEFAULT_CONN_ID,
EcsBaseOperator,
EcsCreateClusterOperator,
EcsDeleteClusterOperator,
Expand Down Expand Up @@ -84,6 +85,44 @@
}
],
}
NOTSET = type("ArgumentNotSet", (), {"__str__": lambda self: "argument-not-set"})


class TestEcsBaseOperator:
"""Test Base ECS Operator."""

@pytest.mark.parametrize("aws_conn_id", [None, NOTSET, "aws_test_conn"])
@pytest.mark.parametrize("region_name", [None, NOTSET, "ca-central-1"])
def test_initialise_operator(self, aws_conn_id, region_name):
"""Test initialize operator."""
op_kw = {"aws_conn_id": aws_conn_id, "region": region_name}
op_kw = {k: v for k, v in op_kw.items() if v is not NOTSET}
op = EcsBaseOperator(task_id="test_ecs_base", **op_kw)

assert op.aws_conn_id == (aws_conn_id if aws_conn_id is not NOTSET else DEFAULT_CONN_ID)
assert op.region == (region_name if region_name is not NOTSET else None)

@mock.patch("airflow.providers.amazon.aws.operators.ecs.EcsHook")
@pytest.mark.parametrize("aws_conn_id", [None, NOTSET, "aws_test_conn"])
@pytest.mark.parametrize("region_name", [None, NOTSET, "ca-central-1"])
def test_hook_and_client(self, mock_ecs_hook_cls, aws_conn_id, region_name):
"""Test initialize ``EcsHook`` and ``boto3.client``."""
mock_ecs_hook = mock_ecs_hook_cls.return_value
mock_conn = mock.MagicMock()
type(mock_ecs_hook).conn = mock.PropertyMock(return_value=mock_conn)

op_kw = {"aws_conn_id": aws_conn_id, "region": region_name}
op_kw = {k: v for k, v in op_kw.items() if v is not NOTSET}
op = EcsBaseOperator(task_id="test_ecs_base_hook_client", **op_kw)

hook = op.hook
assert op.hook is hook
mock_ecs_hook_cls.assert_called_once_with(aws_conn_id=op.aws_conn_id, region_name=op.region)

client = op.client
mock_ecs_hook_cls.assert_called_once_with(aws_conn_id=op.aws_conn_id, region_name=op.region)
assert client == mock_conn
assert op.client is client


@pytest.mark.skipif(mock_ecs is None, reason="mock_ecs package not present")
Expand Down