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
7 changes: 4 additions & 3 deletions airflow-core/tests/unit/utils/test_log_handlers.py
Original file line number Diff line number Diff line change
Expand Up @@ -291,8 +291,9 @@ def test_file_task_handler_with_multiple_executors(
else:
path_to_executor_class = executors_mapping.get(executor_name)

with patch(f"{path_to_executor_class}.get_task_log", return_value=([], [])) as mock_get_task_log:
mock_get_task_log.return_value = ([], [])
with patch(
f"{path_to_executor_class}.get_streaming_task_log", return_value=([], [])
) as mock_get_streaming_task_log:
ti = create_task_instance(
dag_id="dag_for_testing_multiple_executors",
task_id="task_for_testing_multiple_executors",
Expand Down Expand Up @@ -326,7 +327,7 @@ def test_file_task_handler_with_multiple_executors(
assert hasattr(file_handler, "read")
file_handler.read(ti)
os.remove(log_filename)
mock_get_task_log.assert_called_once()
mock_get_streaming_task_log.assert_called_once()

if executor_name is None:
mock_get_default_executor.assert_called_once()
Expand Down
1 change: 1 addition & 0 deletions docs/spelling_wordlist.txt
Original file line number Diff line number Diff line change
Expand Up @@ -1621,6 +1621,7 @@ StoredInfoType
storedInfoType
str
Streamable
StreamingLogResponse
strftime
Stringified
stringified
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,7 @@
from airflow.utils.providers_configuration_loader import providers_configuration_loaded

if TYPE_CHECKING:
from airflow._shared.logging.remote import StreamingLogResponse
from airflow.callbacks.base_callback_sink import BaseCallbackSink
from airflow.callbacks.callback_requests import CallbackRequest
from airflow.cli.cli_config import GroupCommand
Expand Down Expand Up @@ -206,6 +207,12 @@ def get_task_log(self, ti: TaskInstance, try_number: int) -> tuple[list[str], li
return self.kubernetes_executor.get_task_log(ti=ti, try_number=try_number)
return [], []

def get_streaming_task_log(self, ti: TaskInstance, try_number: int) -> StreamingLogResponse:
"""Fetch streaming task log from Kubernetes executor."""
if ti.queue == self.kubernetes_executor.kubernetes_queue:
return self.kubernetes_executor.get_streaming_task_log(ti=ti, try_number=try_number)
return [], []

def has_task(self, task_instance: TaskInstance) -> bool:
"""
Check if a task is either queued or running in either celery or kubernetes executor.
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -198,6 +198,27 @@ def test_log_is_fetched_from_k8s_executor_only_for_k8s_queue(self):
k8s_executor_mock.get_task_log.assert_not_called()
assert log == ([], [])

def test_streaming_log_is_fetched_from_k8s_executor_only_for_k8s_queue(self):
celery_executor_mock = mock.MagicMock()
k8s_executor_mock = mock.MagicMock()
cke = CeleryKubernetesExecutor(celery_executor_mock, k8s_executor_mock)
simple_task_instance = mock.MagicMock()
simple_task_instance.queue = KUBERNETES_QUEUE

cke.get_streaming_task_log(ti=simple_task_instance, try_number=1)

k8s_executor_mock.get_streaming_task_log.assert_called_once_with(
ti=simple_task_instance, try_number=1
)

k8s_executor_mock.reset_mock()
simple_task_instance.queue = "test-queue"

log = cke.get_streaming_task_log(ti=simple_task_instance, try_number=1)

k8s_executor_mock.get_streaming_task_log.assert_not_called()
assert log == ([], [])

def test_get_event_buffer(self):
celery_executor_mock = mock.MagicMock()
k8s_executor_mock = mock.MagicMock()
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -30,9 +30,11 @@
import multiprocessing
import time
from collections import Counter, defaultdict
from collections.abc import Iterable
from contextlib import suppress
from dataclasses import dataclass
from datetime import datetime, timedelta
from itertools import chain
from queue import Empty, Queue
from typing import TYPE_CHECKING, Any

Expand Down Expand Up @@ -68,6 +70,7 @@
from kubernetes.client import models as k8s
from sqlalchemy.orm import Session

from airflow._shared.logging.remote import RawLogStream, StreamingLogResponse
from airflow.cli.cli_config import GroupCommand
from airflow.executors import workloads
from airflow.models.taskinstance import TaskInstance
Expand Down Expand Up @@ -736,15 +739,32 @@ def _get_pod_namespace(self, ti: TaskInstance):
return namespace or self.conf.get("kubernetes_executor", "namespace")

def get_task_log(self, ti: TaskInstance, try_number: int) -> tuple[list[str], list[str]]:
messages = []
log = []
messages: list[str] = []
log: list[str] = []
try:
messages, log_streams = self.get_streaming_task_log(ti, try_number)
log = ["\n".join(stream) for stream in log_streams]
except Exception as e:
messages.append(f"Reading from k8s pod logs failed: {e}")
return messages, log or [""]

@staticmethod
def _create_log_stream(logs: Iterable[bytes]) -> RawLogStream:
for line in logs:
yield remove_escape_codes(line.decode())

def get_streaming_task_log(self, ti: TaskInstance, try_number: int) -> StreamingLogResponse:
messages: list[str] = []
log_streams: list[RawLogStream] = []

try:
from airflow.providers.cncf.kubernetes.kube_client import get_kube_client
from airflow.providers.cncf.kubernetes.pod_generator import PodGenerator

client = get_kube_client()

messages.append(f"Attempting to fetch logs from pod {ti.hostname} through kube API")
hostname_desc = f" {ti.hostname}" if ti.hostname else ""
messages.append(f"Attempting to fetch logs from pod{hostname_desc} through kube API")
selector = PodGenerator.build_selector_for_k8s_executor_pod(
dag_id=ti.dag_id,
task_id=ti.task_id,
Expand All @@ -770,13 +790,16 @@ def get_task_log(self, ti: TaskInstance, try_number: int) -> tuple[list[str], li
tail_lines=self.RUNNING_POD_LOG_LINES,
_preload_content=False,
)
for line in res:
log.append(remove_escape_codes(line.decode()))
if log:

log_iter = iter(res)
first_line = next(log_iter, None)
if first_line is not None:
log_streams.append(self._create_log_stream(chain([first_line], log_iter)))
messages.append("Found logs through kube API")
except Exception as e:
messages.append(f"Reading from k8s pod logs failed: {e}")
return messages, ["\n".join(log)]

return messages, log_streams

def try_adopt_task_instances(self, tis: Sequence[TaskInstance]) -> Sequence[TaskInstance]:
with Stats.timer(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,7 @@
from airflow.providers.common.compat.sdk import conf

if TYPE_CHECKING:
from airflow._shared.logging.remote import StreamingLogResponse
from airflow.callbacks.base_callback_sink import BaseCallbackSink
from airflow.callbacks.callback_requests import CallbackRequest
from airflow.cli.cli_config import GroupCommand
Expand Down Expand Up @@ -201,6 +202,12 @@ def get_task_log(self, ti: TaskInstance, try_number: int) -> tuple[list[str], li
return self.kubernetes_executor.get_task_log(ti=ti, try_number=try_number)
return [], []

def get_streaming_task_log(self, ti: TaskInstance, try_number: int) -> StreamingLogResponse:
"""Fetch streaming task log from kubernetes executor."""
if ti.queue == self.kubernetes_executor.kubernetes_queue:
return self.kubernetes_executor.get_streaming_task_log(ti=ti, try_number=try_number)
return [], []

def has_task(self, task_instance: TaskInstance) -> bool:
"""
Check if a task is either queued or running in either local or kubernetes executor.
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -2628,35 +2628,80 @@ def test_kube_config_get_namespace_list(

assert executor.kube_config.multi_namespace_mode_namespace_list == expected_value_in_kube_config

@pytest.mark.db_test
@mock.patch("airflow.providers.cncf.kubernetes.kube_client.get_kube_client")
def test_get_task_log(self, mock_get_kube_client, create_task_instance_of_operator):
def test_get_streaming_task_log(self, mock_get_kube_client):
"""fetch task log from pod"""
mock_kube_client = mock_get_kube_client.return_value

mock_kube_client.read_namespaced_pod_log.return_value = [b"a_", b"b_", b"c_"]
mock_pod = mock.Mock()
mock_pod.metadata.name = "x"
mock_kube_client.list_namespaced_pod.return_value.items = [mock_pod]
ti = create_task_instance_of_operator(EmptyOperator, dag_id="test_k8s_log_dag", task_id="test_task")
ti = mock.MagicMock(
dag_id="test_k8s_log_dag",
task_id="test_task",
map_index=-1,
run_id="test_run",
queued_by_job_id=None,
hostname="",
executor_config={},
)

executor = KubernetesExecutor()
messages, logs = executor.get_task_log(ti=ti, try_number=1)
messages, log_streams = executor.get_streaming_task_log(ti=ti, try_number=1)

mock_kube_client.read_namespaced_pod_log.assert_called_once()
assert messages == [
"Attempting to fetch logs from pod through kube API",
"Attempting to fetch logs from pod through kube API",
"Found logs through kube API",
]
assert logs[0] == "a_\nb_\nc_"
assert list(log_streams[0]) == ["a_", "b_", "c_"]

mock_kube_client.reset_mock()
mock_kube_client.read_namespaced_pod_log.side_effect = Exception("error_fetching_pod_log")

messages, log_streams = executor.get_streaming_task_log(ti=ti, try_number=1)
assert log_streams == []
assert messages == [
"Attempting to fetch logs from pod through kube API",
"Reading from k8s pod logs failed: error_fetching_pod_log",
]

@mock.patch("airflow.providers.cncf.kubernetes.kube_client.get_kube_client")
def test_get_task_log(self, mock_get_kube_client):
"""Fetch legacy task log response from pod."""
mock_kube_client = mock_get_kube_client.return_value
mock_kube_client.read_namespaced_pod_log.return_value = [b"a_", b"b_", b"c_"]
mock_pod = mock.Mock()
mock_pod.metadata.name = "x"
mock_kube_client.list_namespaced_pod.return_value.items = [mock_pod]
ti = mock.MagicMock(
dag_id="test_k8s_log_dag",
task_id="test_task",
map_index=-1,
run_id="test_run",
queued_by_job_id=None,
hostname="",
executor_config={},
)

executor = KubernetesExecutor()
messages, logs = executor.get_task_log(ti=ti, try_number=1)

assert messages == [
"Attempting to fetch logs from pod through kube API",
"Found logs through kube API",
]
assert logs == ["a_\nb_\nc_"]

mock_kube_client.reset_mock()
mock_kube_client.read_namespaced_pod_log.side_effect = Exception("error_fetching_pod_log")

messages, logs = executor.get_task_log(ti=ti, try_number=1)

assert logs == [""]
assert messages == [
"Attempting to fetch logs from pod through kube API",
"Attempting to fetch logs from pod through kube API",
"Reading from k8s pod logs failed: error_fetching_pod_log",
]

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -112,6 +112,27 @@ def test_log_is_fetched_from_k8s_executor_only_for_k8s_queue(self):
assert logs == []
assert messages == []

def test_streaming_log_is_fetched_from_k8s_executor_only_for_k8s_queue(self):
local_executor_mock = mock.MagicMock()
k8s_executor_mock = mock.MagicMock()
local_k8s_exec = LocalKubernetesExecutor(local_executor_mock, k8s_executor_mock)
simple_task_instance = mock.MagicMock()
simple_task_instance.queue = conf.get("local_kubernetes_executor", "kubernetes_queue")

local_k8s_exec.get_streaming_task_log(ti=simple_task_instance, try_number=3)

k8s_executor_mock.get_streaming_task_log.assert_called_once_with(
ti=simple_task_instance, try_number=3
)

k8s_executor_mock.reset_mock()
simple_task_instance.queue = "test-queue"
messages, logs = local_k8s_exec.get_streaming_task_log(ti=simple_task_instance, try_number=3)

k8s_executor_mock.get_streaming_task_log.assert_not_called()
assert logs == []
assert messages == []

def test_send_callback(self):
local_executor_mock = mock.MagicMock()
k8s_executor_mock = mock.MagicMock()
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -74,13 +74,13 @@ def teardown_method(self):
self.clean_up()

@mock.patch(
"airflow.providers.cncf.kubernetes.executors.kubernetes_executor.KubernetesExecutor.get_task_log"
"airflow.providers.cncf.kubernetes.executors.kubernetes_executor.KubernetesExecutor.get_streaming_task_log"
)
@pytest.mark.parametrize("state", [TaskInstanceState.RUNNING, TaskInstanceState.SUCCESS])
@pytest.mark.usefixtures("clean_executor_loader")
def test__read_for_k8s_executor(self, mock_k8s_get_task_log, create_task_instance, state):
"""Test for k8s executor, the log is read from get_task_log method"""
mock_k8s_get_task_log.return_value = ([], [])
def test__read_for_k8s_executor(self, mock_k8s_get_streaming_task_log, create_task_instance, state):
"""Test for k8s executor, the log is read from get_streaming_task_log method."""
mock_k8s_get_streaming_task_log.return_value = ([], [])
executor_name = "KubernetesExecutor"
ti = create_task_instance(
dag_id="dag_for_testing_k8s_executor_log_read",
Expand All @@ -96,9 +96,9 @@ def test__read_for_k8s_executor(self, mock_k8s_get_task_log, create_task_instanc
fth = FileTaskHandler("")
fth._read(ti=ti, try_number=2)
if state == TaskInstanceState.RUNNING:
mock_k8s_get_task_log.assert_called_once_with(ti, 2)
mock_k8s_get_streaming_task_log.assert_called_once_with(ti, 2)
else:
mock_k8s_get_task_log.assert_not_called()
mock_k8s_get_streaming_task_log.assert_not_called()

@pytest.mark.parametrize(
("pod_override", "namespace_to_call"),
Expand Down
Loading