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
2 changes: 1 addition & 1 deletion generated/known_airflow_exceptions.txt
Original file line number Diff line number Diff line change
Expand Up @@ -279,7 +279,7 @@ providers/google/src/airflow/providers/google/cloud/sensors/cloud_composer.py::4
providers/google/src/airflow/providers/google/cloud/sensors/cloud_storage_transfer_service.py::1
providers/google/src/airflow/providers/google/cloud/sensors/dataflow.py::8
providers/google/src/airflow/providers/google/cloud/sensors/dataform.py::3
providers/google/src/airflow/providers/google/cloud/sensors/datafusion.py::2
providers/google/src/airflow/providers/google/cloud/sensors/datafusion.py::1
providers/google/src/airflow/providers/google/cloud/sensors/dataplex.py::7
providers/google/src/airflow/providers/google/cloud/sensors/dataproc.py::6
providers/google/src/airflow/providers/google/cloud/sensors/dataproc_metastore.py::2
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -28,7 +28,7 @@
from urllib.parse import quote, urlencode, urljoin

import google.auth
from aiohttp import ClientSession
from aiohttp import ClientResponseError, ClientSession
from gcloud.aio.auth import AioSession, Token
from google.api_core.retry import exponential_sleep_generator
from googleapiclient.discovery import Resource, build
Expand Down Expand Up @@ -614,6 +614,11 @@ async def _get_link(self, url: str, session):
try:
pipeline = await session_aio.get(url=url, headers=headers)
break
except ClientResponseError as exc:
if exc.status == 404:
await asyncio.sleep(time_to_wait)
else:
raise
except ValueError as exc:
if "404" in str(exc):
await asyncio.sleep(time_to_wait)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -22,7 +22,9 @@
from collections.abc import Iterable, Sequence
from typing import TYPE_CHECKING

from airflow.providers.common.compat.sdk import AirflowException, AirflowNotFoundException, BaseSensorOperator
from requests.exceptions import HTTPError

from airflow.providers.common.compat.sdk import AirflowException, BaseSensorOperator
from airflow.providers.google.cloud.hooks.datafusion import DataFusionHook
from airflow.providers.google.common.hooks.base_google import PROVIDE_PROJECT_ID

Expand Down Expand Up @@ -113,11 +115,8 @@ def poke(self, context: Context) -> bool:
namespace=self.namespace,
)
pipeline_status = pipeline_workflow.get("status")
except AirflowNotFoundException:
message = "Specified Pipeline ID was not found."
raise AirflowException(message)
except AirflowException:
pass # Because the pipeline may not be visible in system yet
except HTTPError:
pass # A newly started pipeline run may not be visible in CDAP yet
if pipeline_status is not None:
if self.failure_statuses and pipeline_status in self.failure_statuses:
message = (
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,7 @@
from unittest import mock

import pytest
from aiohttp import ClientResponseError
from requests.exceptions import HTTPError

from airflow.providers.google.cloud.hooks.datafusion import DataFusionAsyncHook, DataFusionHook
Expand Down Expand Up @@ -645,6 +646,43 @@ def test_cdap_program_id(self, pipeline_type, expected_program_id):


class TestDataFusionHookAsynch:
@pytest.mark.asyncio
@mock.patch(HOOK_STR.format("asyncio.sleep"), new_callable=mock.AsyncMock)
@mock.patch(HOOK_STR.format("AioSession"))
@mock.patch(HOOK_STR.format("Token"))
async def test_get_link_retries_after_404(self, mock_token, mock_aio_session, mock_sleep, hook_async):
mock_token.return_value.__aenter__.return_value.get = mock.AsyncMock(return_value="token")
response = MockAiohttpClientResponse(payload={"status": "RUNNING"})
mock_aio_session.return_value.get = mock.AsyncMock(
side_effect=[
ClientResponseError(request_info=mock.Mock(), history=(), status=404),
response,
]
)

result = await hook_async._get_link(url=CONSTRUCTED_PIPELINE_URL, session=session)

assert result is response
assert mock_aio_session.return_value.get.await_count == 2
mock_sleep.assert_awaited_once()

@pytest.mark.asyncio
@mock.patch(HOOK_STR.format("asyncio.sleep"), new_callable=mock.AsyncMock)
@mock.patch(HOOK_STR.format("AioSession"))
@mock.patch(HOOK_STR.format("Token"))
async def test_get_link_propagates_non_404_error(
self, mock_token, mock_aio_session, mock_sleep, hook_async
):
mock_token.return_value.__aenter__.return_value.get = mock.AsyncMock(return_value="token")
error = ClientResponseError(request_info=mock.Mock(), history=(), status=500)
mock_aio_session.return_value.get = mock.AsyncMock(side_effect=error)

with pytest.raises(ClientResponseError) as ctx:
await hook_async._get_link(url=CONSTRUCTED_PIPELINE_URL, session=session)

assert ctx.value is error
mock_sleep.assert_not_awaited()

@pytest.mark.asyncio
@mock.patch(HOOK_STR.format("DataFusionAsyncHook._get_link"))
async def test_async_get_pipeline_should_execute_successfully(self, mocked_link, hook_async):
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -20,8 +20,9 @@
from unittest import mock

import pytest
from requests.exceptions import HTTPError, RequestException

from airflow.providers.common.compat.sdk import AirflowException, AirflowNotFoundException
from airflow.providers.common.compat.sdk import AirflowException
from airflow.providers.google.cloud.hooks.datafusion import PipelineStates
from airflow.providers.google.cloud.sensors.datafusion import CloudDataFusionPipelineStateSensor

Expand Down Expand Up @@ -99,9 +100,9 @@ def test_assertion(self, mock_hook):
task.poke(mock.MagicMock())

@mock.patch("airflow.providers.google.cloud.sensors.datafusion.DataFusionHook")
def test_not_found_exception(self, mock_hook):
def test_pipeline_not_visible_yet(self, mock_hook):
mock_hook.return_value.get_instance.return_value = {"apiEndpoint": INSTANCE_URL}
mock_hook.return_value.get_pipeline_workflow.side_effect = AirflowNotFoundException()
mock_hook.return_value.get_pipeline_workflow.side_effect = HTTPError()

task = CloudDataFusionPipelineStateSensor(
task_id="test_task_id",
Expand All @@ -116,8 +117,28 @@ def test_not_found_exception(self, mock_hook):
impersonation_chain=IMPERSONATION_CHAIN,
)

with pytest.raises(
AirflowException,
match="Specified Pipeline ID was not found.",
):
assert task.poke(mock.MagicMock()) is False

@mock.patch("airflow.providers.google.cloud.sensors.datafusion.DataFusionHook")
def test_other_request_error_is_not_suppressed(self, mock_hook):
mock_hook.return_value.get_instance.return_value = {"apiEndpoint": INSTANCE_URL}
error = RequestException("Retrieving a pipeline state failed with code 500")
mock_hook.return_value.get_pipeline_workflow.side_effect = error

task = CloudDataFusionPipelineStateSensor(
task_id="test_task_id",
pipeline_name=PIPELINE_NAME,
pipeline_id=PIPELINE_ID,
project_id=PROJECT_ID,
expected_statuses={PipelineStates.COMPLETED},
failure_statuses=FAILURE_STATUSES,
instance_name=INSTANCE_NAME,
location=LOCATION,
gcp_conn_id=GCP_CONN_ID,
impersonation_chain=IMPERSONATION_CHAIN,
)

with pytest.raises(RequestException) as ctx:
task.poke(mock.MagicMock())

assert ctx.value is error