From 6d2a25fda7b49423258c0cd9ebc2d0119a24b085 Mon Sep 17 00:00:00 2001 From: Sean Ghaeli Date: Mon, 13 Jul 2026 22:45:53 +0000 Subject: [PATCH] Fix Neptune Analytics deferrable path: load custom waiter + forward hook params to triggers The example_neptune_analytics system test fails deferrable-only. Two bugs: 1. Custom waiter never loaded: hook client_type is 'neptune-graph' so waiter_path derives 'neptune-graph.json', but the file shipped as 'neptune_analytics.json'. get_waiter fell through to botocore's official ImportTaskCancelled waiter, which fails on any non-CANCELLING/CANCELLED status -- so a fast import reaching SUCCEEDED before cancel raised NeptuneImportTaskCancellationFailedError. Renamed the file to match. 2. region_name/verify/botocore_config were not forwarded to the deferred triggers, so the triggerer rebuilt the async hook with region_name=None (sync path passes because it uses self.hook). Forwarded them to all 8 defer() sites. Adds regression tests: a waiter-loads test and TestNeptuneDeferForwardsHookParams. --- .../amazon/aws/operators/neptune_analytics.py | 24 +++ ...tune_analytics.json => neptune-graph.json} | 0 .../aws/operators/test_neptune_analytics.py | 187 +++++++++++++++++- .../aws/waiters/test_neptune_analytics.py | 94 +++++++++ 4 files changed, 304 insertions(+), 1 deletion(-) rename providers/amazon/src/airflow/providers/amazon/aws/waiters/{neptune_analytics.json => neptune-graph.json} (100%) create mode 100644 providers/amazon/tests/unit/amazon/aws/waiters/test_neptune_analytics.py diff --git a/providers/amazon/src/airflow/providers/amazon/aws/operators/neptune_analytics.py b/providers/amazon/src/airflow/providers/amazon/aws/operators/neptune_analytics.py index 134d13df982c5..bdc4cf2040f54 100644 --- a/providers/amazon/src/airflow/providers/amazon/aws/operators/neptune_analytics.py +++ b/providers/amazon/src/airflow/providers/amazon/aws/operators/neptune_analytics.py @@ -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", ) @@ -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}, @@ -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", ) @@ -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", ) @@ -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}, @@ -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}, @@ -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", ) @@ -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", ) diff --git a/providers/amazon/src/airflow/providers/amazon/aws/waiters/neptune_analytics.json b/providers/amazon/src/airflow/providers/amazon/aws/waiters/neptune-graph.json similarity index 100% rename from providers/amazon/src/airflow/providers/amazon/aws/waiters/neptune_analytics.json rename to providers/amazon/src/airflow/providers/amazon/aws/waiters/neptune-graph.json diff --git a/providers/amazon/tests/unit/amazon/aws/operators/test_neptune_analytics.py b/providers/amazon/tests/unit/amazon/aws/operators/test_neptune_analytics.py index 472e386888c31..75c4ca8055a56 100644 --- a/providers/amazon/tests/unit/amazon/aws/operators/test_neptune_analytics.py +++ b/providers/amazon/tests/unit/amazon/aws/operators/test_neptune_analytics.py @@ -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, @@ -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 diff --git a/providers/amazon/tests/unit/amazon/aws/waiters/test_neptune_analytics.py b/providers/amazon/tests/unit/amazon/aws/waiters/test_neptune_analytics.py new file mode 100644 index 0000000000000..39f09e2ae7ca8 --- /dev/null +++ b/providers/amazon/tests/unit/amazon/aws/waiters/test_neptune_analytics.py @@ -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}, + )