diff --git a/airflow/providers/cncf/kubernetes/operators/kubernetes_pod.py b/airflow/providers/cncf/kubernetes/operators/kubernetes_pod.py index 056d7fe3f3ae2..ba8542e4ac376 100644 --- a/airflow/providers/cncf/kubernetes/operators/kubernetes_pod.py +++ b/airflow/providers/cncf/kubernetes/operators/kubernetes_pod.py @@ -404,14 +404,14 @@ def execute(self, context: 'Context'): def cleanup(self, pod: k8s.V1Pod, remote_pod: k8s.V1Pod): pod_phase = remote_pod.status.phase if hasattr(remote_pod, 'status') else None + if not self.is_delete_operator_pod: + with _suppress(Exception): + self.patch_already_checked(pod) if pod_phase != PodPhase.SUCCEEDED: if self.log_events_on_failure: with _suppress(Exception): for event in self.pod_manager.read_pod_events(pod).items: self.log.error("Pod Event: %s - %s", event.reason, event.message) - if not self.is_delete_operator_pod: - with _suppress(Exception): - self.patch_already_checked(pod) with _suppress(Exception): self.process_pod_deletion(pod) error_message = get_container_termination_message(remote_pod, self.BASE_CONTAINER_NAME) diff --git a/kubernetes_tests/test_kubernetes_pod_operator.py b/kubernetes_tests/test_kubernetes_pod_operator.py index 58c75f1e37730..25c5bab7639e7 100644 --- a/kubernetes_tests/test_kubernetes_pod_operator.py +++ b/kubernetes_tests/test_kubernetes_pod_operator.py @@ -22,6 +22,7 @@ import sys import textwrap import unittest +from copy import copy from unittest import mock from unittest.mock import ANY, MagicMock @@ -167,8 +168,10 @@ def test_config_path_move(self): ) context = create_context(k) k.execute(context) + expected_pod = copy(self.expected_pod) + expected_pod['metadata']['labels']['already_checked'] = 'True' actual_pod = self.api_client.sanitize_for_serialization(k.pod) - assert self.expected_pod == actual_pod + assert expected_pod == actual_pod def test_working_pod(self): k = KubernetesPodOperator( @@ -760,6 +763,7 @@ def test_full_pod_spec(self): 'kubernetes_pod_operator': 'True', 'task_id': mock.ANY, 'try_number': '1', + 'already_checked': 'True', } assert k.pod.spec.containers[0].env == [k8s.V1EnvVar(name="env_name", value="value")] assert result == {"hello": "world"} @@ -982,19 +986,23 @@ def test_reattach_failing_pod_once(self): client = kube_client.get_kube_client(in_cluster=False) name = "test" namespace = "default" - k = KubernetesPodOperator( - namespace='default', - image="ubuntu:16.04", - cmds=["bash", "-cx"], - arguments=["exit 1"], - labels={"foo": "bar"}, - name="test", - task_id=name, - in_cluster=False, - do_xcom_push=False, - is_delete_operator_pod=False, - termination_grace_period=0, - ) + + def get_op(): + return KubernetesPodOperator( + namespace='default', + image="ubuntu:16.04", + cmds=["bash", "-cx"], + arguments=["exit 1"], + labels={"foo": "bar"}, + name="test", + task_id=name, + in_cluster=False, + do_xcom_push=False, + is_delete_operator_pod=False, + termination_grace_period=0, + ) + + k = get_op() context = create_context(k) @@ -1004,9 +1012,14 @@ def test_reattach_failing_pod_once(self): ) as await_pod_completion_mock: pod_mock = MagicMock() - # we don't want failure because we don't want the pod to be patched as "already_checked" pod_mock.status.phase = 'Succeeded' await_pod_completion_mock.return_value = pod_mock + + # we want to simulate that there was a worker failure and the airflow operator process + # was killed without running the cleanup process. in this case the pod will not be marked as + # already checked + k.cleanup = MagicMock() + k.execute(context) name = k.pod.metadata.name pod = client.read_namespaced_pod(name=name, namespace=namespace) @@ -1014,7 +1027,11 @@ def test_reattach_failing_pod_once(self): pod = client.read_namespaced_pod(name=name, namespace=namespace) assert 'already_checked' not in pod.metadata.labels - # should not call `create_pod`, because there's a pod there it should find + # create a new version of the same operator instance to remove the monkey patching in first + # part of the test + k = get_op() + + # `create_pod` should not be called because there's a pod there it should find # should use the found pod and patch as "already_checked" (in failure block) with mock.patch( "airflow.providers.cncf.kubernetes.utils.pod_manager.PodManager.create_pod" @@ -1025,6 +1042,9 @@ def test_reattach_failing_pod_once(self): assert pod.metadata.labels["already_checked"] == "True" create_mock.assert_not_called() + # recreate op just to ensure we're not relying on any statefulness + k = get_op() + # `create_pod` should be called because though there's still a pod to be found, # it will be `already_checked` with mock.patch( diff --git a/kubernetes_tests/test_kubernetes_pod_operator_backcompat.py b/kubernetes_tests/test_kubernetes_pod_operator_backcompat.py index 8fc4fe1e0fc7f..430387118a577 100644 --- a/kubernetes_tests/test_kubernetes_pod_operator_backcompat.py +++ b/kubernetes_tests/test_kubernetes_pod_operator_backcompat.py @@ -18,6 +18,7 @@ import json import sys import unittest +from copy import copy from unittest import mock from unittest.mock import MagicMock, patch @@ -301,14 +302,16 @@ def test_volume_mount(self): k.execute(context=context) mock_logger.info.assert_any_call('retrieved from mount') actual_pod = self.api_client.sanitize_for_serialization(k.pod) - self.expected_pod['spec']['containers'][0]['args'] = args - self.expected_pod['spec']['containers'][0]['volumeMounts'] = [ + expected_pod = copy(self.expected_pod) + expected_pod['spec']['containers'][0]['args'] = args + expected_pod['spec']['containers'][0]['volumeMounts'] = [ {'name': 'test-volume', 'mountPath': '/tmp/test_volume', 'readOnly': False} ] - self.expected_pod['spec']['volumes'] = [ + expected_pod['spec']['volumes'] = [ {'name': 'test-volume', 'persistentVolumeClaim': {'claimName': 'test-volume'}} ] - assert self.expected_pod == actual_pod + expected_pod['metadata']['labels']['already_checked'] = 'True' + assert expected_pod == actual_pod def test_run_as_user_root(self): security_context = { diff --git a/tests/providers/cncf/kubernetes/operators/test_kubernetes_pod.py b/tests/providers/cncf/kubernetes/operators/test_kubernetes_pod.py index 0faa4ef8e1287..ee2a3cdb4378d 100644 --- a/tests/providers/cncf/kubernetes/operators/test_kubernetes_pod.py +++ b/tests/providers/cncf/kubernetes/operators/test_kubernetes_pod.py @@ -17,26 +17,45 @@ from unittest import mock from unittest.mock import MagicMock +import pendulum import pytest from kubernetes.client import ApiClient, models as k8s from airflow.exceptions import AirflowException -from airflow.models import DAG, DagRun, TaskInstance +from airflow.models import DAG, DagModel, DagRun, TaskInstance from airflow.models.xcom import XCom from airflow.providers.cncf.kubernetes.operators.kubernetes_pod import KubernetesPodOperator, _suppress from airflow.utils import timezone +from airflow.utils.session import create_session from airflow.utils.types import DagRunType +from tests.test_utils import db DEFAULT_DATE = timezone.datetime(2016, 1, 1, 1, 0, 0) +KPO_MODULE = "airflow.providers.cncf.kubernetes.operators.kubernetes_pod" -def create_context(task): - dag = DAG(dag_id="dag") +@pytest.fixture(scope='function', autouse=True) +def clear_db(): + db.clear_db_dags() + db.clear_db_runs() + yield + + +def create_context(task, persist_to_db=False): + dag = task.dag if task.has_dag() else DAG(dag_id="dag") dag_run = DagRun( - run_id=DagRun.generate_run_id(DagRunType.MANUAL, DEFAULT_DATE), run_type=DagRunType.MANUAL + run_id=DagRun.generate_run_id(DagRunType.MANUAL, DEFAULT_DATE), + run_type=DagRunType.MANUAL, + dag_id=dag.dag_id, ) task_instance = TaskInstance(task=task, run_id=dag_run.run_id) task_instance.dag_run = dag_run + if persist_to_db: + with create_session() as session: + session.add(DagModel(dag_id=dag.dag_id)) + session.add(dag_run) + session.add(task_instance) + session.commit() return { "dag": dag, "ts": DEFAULT_DATE.isoformat(), @@ -777,36 +796,8 @@ def test_previous_pods_ignored_for_reattached(self): assert 'already_checked!=True' in kwargs['label_selector'] @mock.patch("airflow.providers.cncf.kubernetes.utils.pod_manager.PodManager.delete_pod") - @mock.patch( - "airflow.providers.cncf.kubernetes.operators.kubernetes_pod" - ".KubernetesPodOperator.patch_already_checked" - ) - def test_mark_created_pod_if_not_deleted(self, mock_patch_already_checked, mock_delete_pod): - """If we aren't deleting pods and have a failure, mark it so we don't reattach to it""" - k = KubernetesPodOperator( - namespace="default", - image="ubuntu:16.04", - name="test", - task_id="task", - is_delete_operator_pod=False, - ) - remote_pod_mock = MagicMock() - remote_pod_mock.status.phase = 'Failed' - self.await_pod_mock.return_value = remote_pod_mock - context = create_context(k) - with pytest.raises(AirflowException): - k.execute(context=context) - mock_patch_already_checked.assert_called_once() - mock_delete_pod.assert_not_called() - - @mock.patch("airflow.providers.cncf.kubernetes.utils.pod_manager.PodManager.delete_pod") - @mock.patch( - "airflow.providers.cncf.kubernetes.operators.kubernetes_pod" - ".KubernetesPodOperator.patch_already_checked" - ) - def test_mark_created_pod_if_not_deleted_during_exception( - self, mock_patch_already_checked, mock_delete_pod - ): + @mock.patch(f"{KPO_MODULE}.KubernetesPodOperator.patch_already_checked") + def test_mark_checked_unexpected_exception(self, mock_patch_already_checked, mock_delete_pod): """If we aren't deleting pods and have an exception, mark it so we don't reattach to it""" k = KubernetesPodOperator( namespace="default", @@ -822,26 +813,28 @@ def test_mark_created_pod_if_not_deleted_during_exception( mock_patch_already_checked.assert_called_once() mock_delete_pod.assert_not_called() + @pytest.mark.parametrize('should_fail', [True, False]) @mock.patch("airflow.providers.cncf.kubernetes.utils.pod_manager.PodManager.delete_pod") - @mock.patch( - "airflow.providers.cncf.kubernetes.operators." - "kubernetes_pod.KubernetesPodOperator.patch_already_checked" - ) - def test_mark_reattached_pod_if_not_deleted(self, mock_patch_already_checked, mock_delete_pod): - """If we aren't deleting pods and have a failure, mark it so we don't reattach to it""" + @mock.patch(f"{KPO_MODULE}.KubernetesPodOperator.patch_already_checked") + def test_mark_checked_if_not_deleted(self, mock_patch_already_checked, mock_delete_pod, should_fail): + """If we aren't deleting pods mark "checked" if the task completes (successful or otherwise)""" + dag = DAG('hello2', start_date=pendulum.now()) k = KubernetesPodOperator( namespace="default", image="ubuntu:16.04", name="test", task_id="task", is_delete_operator_pod=False, + dag=dag, ) remote_pod_mock = MagicMock() - remote_pod_mock.status.phase = 'Failed' + remote_pod_mock.status.phase = 'Failed' if should_fail else 'Succeeded' self.await_pod_mock.return_value = remote_pod_mock - - context = create_context(k) - with pytest.raises(AirflowException): + context = create_context(k, persist_to_db=True) + if should_fail: + with pytest.raises(AirflowException): + k.execute(context=context) + else: k.execute(context=context) mock_patch_already_checked.assert_called_once() mock_delete_pod.assert_not_called()