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 @@ -176,6 +176,9 @@ def execute(self, context: Context) -> dict:
graph_id=self.graph_id,
waiter_delay=self.waiter_delay,
waiter_max_attempts=self.waiter_max_attempts,
region_name=self.region_name,
verify=self.verify,
botocore_config=self.botocore_config,
),
method_name="execute_complete",
)
Expand Down Expand Up @@ -321,6 +324,9 @@ def execute(self, context: Context) -> dict:
vpc_id=self.vpc_id,
waiter_delay=self.waiter_delay,
waiter_max_attempts=self.waiter_max_attempts,
region_name=self.region_name,
verify=self.verify,
botocore_config=self.botocore_config,
),
method_name="execute_complete",
kwargs={"vpc_id": self.vpc_id},
Expand Down Expand Up @@ -425,6 +431,9 @@ def execute(self, context: Context) -> None:
endpoint_id=endpoint_id,
waiter_delay=self.waiter_delay,
waiter_max_attempts=self.waiter_max_attempts,
region_name=self.region_name,
verify=self.verify,
botocore_config=self.botocore_config,
),
method_name="execute_complete",
)
Expand Down Expand Up @@ -520,6 +529,9 @@ def execute(self, context: Context):
graph_id=self.graph_id,
waiter_delay=self.waiter_delay,
waiter_max_attempts=self.waiter_max_attempts,
region_name=self.region_name,
verify=self.verify,
botocore_config=self.botocore_config,
),
method_name="execute_complete",
)
Expand Down Expand Up @@ -732,6 +744,9 @@ def execute(self, context: Context) -> dict:
graph_id=self.graph_id,
waiter_delay=self.waiter_delay,
waiter_max_attempts=self.waiter_max_attempts,
region_name=self.region_name,
verify=self.verify,
botocore_config=self.botocore_config,
),
method_name="defer_wait_for_task",
kwargs={"import_task_id": import_task_id},
Expand Down Expand Up @@ -773,6 +788,9 @@ def defer_wait_for_task(
waiter_delay=self.waiter_delay,
waiter_max_attempts=self.waiter_max_attempts,
aws_conn_id=self.aws_conn_id,
region_name=self.region_name,
verify=self.verify,
botocore_config=self.botocore_config,
),
method_name="execute_complete",
kwargs={"graph_id": graph_id},
Expand Down Expand Up @@ -914,6 +932,9 @@ def execute(self, context: Context) -> dict:
waiter_delay=self.waiter_delay,
waiter_max_attempts=self.waiter_max_attempts,
aws_conn_id=self.aws_conn_id,
region_name=self.region_name,
verify=self.verify,
botocore_config=self.botocore_config,
),
method_name="execute_complete",
)
Expand Down Expand Up @@ -1002,6 +1023,9 @@ def execute(self, context: Context) -> dict:
waiter_delay=self.waiter_delay,
waiter_max_attempts=self.waiter_max_attempts,
aws_conn_id=self.aws_conn_id,
region_name=self.region_name,
verify=self.verify,
botocore_config=self.botocore_config,
),
method_name="execute_complete",
)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -1187,7 +1187,8 @@ def test_execute_complete_success(self):

class TestNeptuneCancelImportTaskOperator:
@mock.patch.object(NeptuneAnalyticsHook, "conn")
def test_init_defaults(self, mock_conn):
@mock.patch.object(NeptuneAnalyticsHook, "get_waiter")
def test_init_defaults(self, mock_get_waiter, mock_conn):
mock_conn.cancel_import_task.return_value = {
"taskId": TASK_ID,
"graphId": GRAPH_ID,
Expand Down Expand Up @@ -1275,3 +1276,187 @@ def test_execute_complete_success(self):
result = operator.execute_complete(None, event)

assert result == {"import_task_id": TASK_ID}


class TestNeptuneDeferForwardsHookParams:
"""Regression tests: the operator's region_name/verify/botocore_config must be
forwarded to the trigger, otherwise the triggerer (a separate process) rebuilds
the async hook with region_name=None and can fail with NoRegionError or a
wrong-region ResourceNotFoundException while the synchronous path succeeds.
"""

REGION = "eu-west-1"
BOTOCORE_CONFIG = {"read_timeout": 42}
VERIFY = False

@mock.patch("airflow.providers.amazon.aws.operators.neptune_analytics.NeptuneGraphLink.persist")
@mock.patch.object(NeptuneAnalyticsHook, "conn")
def test_create_graph_forwards_hook_params(self, mock_conn, mock_persist):
mock_conn.create_graph.return_value = {"id": GRAPH_ID, "status": "CREATING"}
operator = NeptuneCreateGraphOperator(
task_id="test_task",
graph_name=GRAPH_NAME,
vector_search_config={"test": 123},
provisioned_memory=16,
deferrable=True,
region_name=self.REGION,
verify=self.VERIFY,
botocore_config=self.BOTOCORE_CONFIG,
)
with pytest.raises(TaskDeferred) as exc_info:
operator.execute(None)

trigger = exc_info.value.trigger
assert trigger.region_name == self.REGION
assert trigger.verify == self.VERIFY
assert trigger.botocore_config == self.BOTOCORE_CONFIG

@mock.patch("airflow.providers.amazon.aws.operators.neptune_analytics.VpcEndpointLink.persist")
@mock.patch.object(NeptuneAnalyticsHook, "conn")
def test_create_private_endpoint_forwards_hook_params(self, mock_conn, mock_persist):
mock_conn.create_private_graph_endpoint.return_value = {
"status": "CREATING",
"vpcEndpointId": ENDPOINT_ID,
"vpcId": VPC_ID,
}
mock_conn.get_private_graph_endpoint.return_value = {"vpcEndpointId": ENDPOINT_ID}
operator = NeptuneCreatePrivateGraphEndpointOperator(
task_id="test_task",
graph_identifier=GRAPH_ID,
deferrable=True,
region_name=self.REGION,
verify=self.VERIFY,
botocore_config=self.BOTOCORE_CONFIG,
)
with pytest.raises(TaskDeferred) as exc_info:
operator.execute(None)

trigger = exc_info.value.trigger
assert trigger.region_name == self.REGION
assert trigger.verify == self.VERIFY
assert trigger.botocore_config == self.BOTOCORE_CONFIG

@mock.patch.object(NeptuneAnalyticsHook, "conn")
def test_delete_private_endpoint_forwards_hook_params(self, mock_conn):
mock_conn.delete_private_graph_endpoint.return_value = {
"status": "DELETING",
"vpcEndpointId": ENDPOINT_ID,
}
operator = NeptuneDeletePrivateGraphEndpointOperator(
task_id="test_task",
graph_identifier=GRAPH_ID,
vpc_id=VPC_ID,
deferrable=True,
region_name=self.REGION,
verify=self.VERIFY,
botocore_config=self.BOTOCORE_CONFIG,
)
with pytest.raises(TaskDeferred) as exc_info:
operator.execute(None)

trigger = exc_info.value.trigger
assert trigger.region_name == self.REGION
assert trigger.verify == self.VERIFY
assert trigger.botocore_config == self.BOTOCORE_CONFIG

@mock.patch.object(NeptuneAnalyticsHook, "conn")
def test_delete_graph_forwards_hook_params(self, mock_conn):
operator = NeptuneDeleteGraphOperator(
task_id="test_task",
graph_id=GRAPH_ID,
skip_snapshot=True,
deferrable=True,
region_name=self.REGION,
verify=self.VERIFY,
botocore_config=self.BOTOCORE_CONFIG,
)
with pytest.raises(TaskDeferred) as exc_info:
operator.execute(None)

trigger = exc_info.value.trigger
assert trigger.region_name == self.REGION
assert trigger.verify == self.VERIFY
assert trigger.botocore_config == self.BOTOCORE_CONFIG

@mock.patch("airflow.providers.amazon.aws.operators.neptune_analytics.NeptuneImportTaskLink.persist")
@mock.patch.object(NeptuneAnalyticsHook, "conn")
def test_start_import_forwards_hook_params(self, mock_conn, mock_persist):
mock_conn.start_import_task.return_value = {"taskId": TASK_ID}
operator = NeptuneStartImportTaskOperator(
task_id="test_task",
graph_identifier=GRAPH_ID,
role_arn=ROLE_ARN,
source=SOURCE_S3_URI,
deferrable=True,
region_name=self.REGION,
verify=self.VERIFY,
botocore_config=self.BOTOCORE_CONFIG,
)
with pytest.raises(TaskDeferred) as exc_info:
operator.execute(None)

trigger = exc_info.value.trigger
assert trigger.region_name == self.REGION
assert trigger.verify == self.VERIFY
assert trigger.botocore_config == self.BOTOCORE_CONFIG

@mock.patch.object(NeptuneAnalyticsHook, "conn")
def test_cancel_import_forwards_hook_params(self, mock_conn):
mock_conn.cancel_import_task.return_value = {"status": "CANCELLING"}
operator = NeptuneCancelImportTaskOperator(
task_id="test_task",
import_task_id=TASK_ID,
deferrable=True,
region_name=self.REGION,
verify=self.VERIFY,
botocore_config=self.BOTOCORE_CONFIG,
)
with pytest.raises(TaskDeferred) as exc_info:
operator.execute(None)

trigger = exc_info.value.trigger
assert trigger.region_name == self.REGION
assert trigger.verify == self.VERIFY
assert trigger.botocore_config == self.BOTOCORE_CONFIG

@mock.patch("airflow.providers.amazon.aws.operators.neptune_analytics.NeptuneImportTaskLink.persist")
@mock.patch("airflow.providers.amazon.aws.operators.neptune_analytics.NeptuneGraphLink.persist")
@mock.patch.object(NeptuneAnalyticsHook, "conn")
def test_create_graph_with_import_forwards_hook_params(
self, mock_conn, mock_graph_persist, mock_task_persist
):
mock_conn.create_graph_using_import_task.return_value = {
"graphId": GRAPH_ID,
"taskId": TASK_ID,
"status": "CREATING",
}
operator = NeptuneCreateGraphWithImportOperator(
task_id="test_task",
graph_name=GRAPH_NAME,
vector_search_config={"test": 123},
source=SOURCE_S3_URI,
role_arn=ROLE_ARN,
deferrable=True,
region_name=self.REGION,
verify=self.VERIFY,
botocore_config=self.BOTOCORE_CONFIG,
)
# First defer: graph availability
with pytest.raises(TaskDeferred) as exc_info:
operator.execute(None)
trigger = exc_info.value.trigger
assert trigger.region_name == self.REGION
assert trigger.verify == self.VERIFY
assert trigger.botocore_config == self.BOTOCORE_CONFIG

# Second defer: import task completion (via defer_wait_for_task)
with pytest.raises(TaskDeferred) as exc_info:
operator.defer_wait_for_task(
None,
event={"status": "success", "graph_id": GRAPH_ID},
import_task_id=TASK_ID,
)
trigger = exc_info.value.trigger
assert trigger.region_name == self.REGION
assert trigger.verify == self.VERIFY
assert trigger.botocore_config == self.BOTOCORE_CONFIG
Original file line number Diff line number Diff line change
@@ -0,0 +1,94 @@
# Licensed to the Apache Software Foundation (ASF) under one
# or more contributor license agreements. See the NOTICE file
# distributed with this work for additional information
# regarding copyright ownership. The ASF licenses this file
# to you under the Apache License, Version 2.0 (the
# "License"); you may not use this file except in compliance
# with the License. You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing,
# software distributed under the License is distributed on an
# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
# KIND, either express or implied. See the License for the
# specific language governing permissions and limitations
# under the License.
from __future__ import annotations

from unittest import mock

import boto3
import botocore
import pytest

from airflow.providers.amazon.aws.hooks.neptune_analytics import NeptuneAnalyticsHook


class TestNeptuneAnalyticsCustomWaiters:
def test_service_waiters(self):
"""The custom waiter file must be discovered and loaded by the hook.

Regression test: the waiter file must be named to match the hook's
``client_type`` (``neptune-graph``) so that ``AwsGenericHook.waiter_path``
resolves it. If the file name does not match, ``list_waiters`` silently
falls back to botocore's official waiters and the custom
``import_task_cancelled`` acceptors become dead code.
"""
assert "import_task_cancelled" in NeptuneAnalyticsHook().list_waiters()

def test_custom_import_task_cancelled_waiter_is_used(self):
"""The custom (not botocore-official) waiter definition must be selected.

The custom waiter treats a fast-completing import that reaches SUCCEEDED
(or FAILED) before the cancel takes effect as *success*, whereas
botocore's official ImportTaskCancelled waiter treats anything other than
CANCELLING/CANCELLED as failure. This asserts the custom acceptors win.
"""
hook = NeptuneAnalyticsHook()
assert hook.waiter_path is not None
assert "import_task_cancelled" in hook._list_custom_waiters()


class TestNeptuneAnalyticsImportTaskCancelledWaiter:
WAITER_NAME = "import_task_cancelled"
TASK_ID = "t-abc123"

@pytest.fixture(autouse=True)
def mock_conn(self, monkeypatch):
self.client = boto3.client("neptune-graph", region_name="us-east-1")
monkeypatch.setattr(NeptuneAnalyticsHook, "conn", self.client)

@pytest.fixture
def mock_get_import_task(self):
with mock.patch.object(self.client, "get_import_task") as mock_getter:
yield mock_getter

@pytest.mark.parametrize("state", ["CANCELLED", "SUCCEEDED", "FAILED"])
def test_import_task_cancelled_success_states(self, state, mock_get_import_task):
# SUCCEEDED / FAILED can happen when the import finishes before the cancel
# takes effect. The custom waiter must treat these as success, not failure.
mock_get_import_task.return_value = {"status": state}

NeptuneAnalyticsHook().get_waiter(self.WAITER_NAME).wait(
taskIdentifier=self.TASK_ID,
WaiterConfig={"Delay": 0.01, "MaxAttempts": 3},
)

def test_import_task_cancelled_failure_state(self, mock_get_import_task):
mock_get_import_task.return_value = {"status": "ERROR_ENCOUNTERED"}
with pytest.raises(botocore.exceptions.WaiterError):
NeptuneAnalyticsHook().get_waiter(self.WAITER_NAME).wait(
taskIdentifier=self.TASK_ID,
WaiterConfig={"Delay": 0.01, "MaxAttempts": 3},
)

def test_import_task_cancelled_wait(self, mock_get_import_task):
cancelling = {"status": "CANCELLING"}
cancelled = {"status": "CANCELLED"}
mock_get_import_task.side_effect = [cancelling, cancelling, cancelled]

NeptuneAnalyticsHook().get_waiter(self.WAITER_NAME).wait(
taskIdentifier=self.TASK_ID,
WaiterConfig={"Delay": 0.01, "MaxAttempts": 3},
)