diff --git a/providers/google/provider.yaml b/providers/google/provider.yaml index e1e6d841c6e6c..7262c0699c0bd 100644 --- a/providers/google/provider.yaml +++ b/providers/google/provider.yaml @@ -1495,6 +1495,8 @@ logging: remote-logging: - classpath: airflow.providers.google.cloud.log.gcs_task_handler.GCSRemoteLogIO scheme: gs + - classpath: airflow.providers.google.cloud.log.stackdriver_task_handler.StackdriverRemoteLogIO + scheme: stackdriver queues: - airflow.providers.google.event_scheduling.events.pubsub.PubSubMessageQueueEventTriggerContainer diff --git a/providers/google/src/airflow/providers/google/cloud/log/stackdriver_task_handler.py b/providers/google/src/airflow/providers/google/cloud/log/stackdriver_task_handler.py index dd184c230ebf0..6c1772a88490c 100644 --- a/providers/google/src/airflow/providers/google/cloud/log/stackdriver_task_handler.py +++ b/providers/google/src/airflow/providers/google/cloud/log/stackdriver_task_handler.py @@ -20,6 +20,7 @@ import contextlib import copy +import inspect import logging import os import shutil @@ -31,7 +32,7 @@ from logging import getLogRecordFactory from pathlib import Path from typing import TYPE_CHECKING -from urllib.parse import urlencode +from urllib.parse import urlencode, urlsplit import attrs from google.cloud import logging as gcp_logging @@ -41,9 +42,11 @@ from google.cloud.logging_v2.types import ListLogEntriesRequest, ListLogEntriesResponse from airflow.exceptions import AirflowProviderDeprecationWarning +from airflow.providers.common.compat.sdk import conf from airflow.providers.google.cloud.utils.credentials_provider import get_credentials_and_project_id from airflow.providers.google.common.consts import CLIENT_INFO from airflow.providers.google.version_compat import AIRFLOW_V_3_0_PLUS +from airflow.utils.log.file_task_handler import FileTaskHandler from airflow.utils.log.logging_mixin import LoggingMixin try: @@ -94,6 +97,39 @@ class StackdriverRemoteLogIO(LoggingMixin): resource: Resource = _GLOBAL_RESOURCE labels: dict[str, str] | None = None + @classmethod + def from_config(cls) -> StackdriverRemoteLogIO: + """Build the remote log IO from Airflow logging configuration.""" + remote_task_handler_kwargs = conf.getjson("logging", "remote_task_handler_kwargs", fallback={}) + if not isinstance(remote_task_handler_kwargs, dict): + raise ValueError( + "logging/remote_task_handler_kwargs must be a JSON object (a python dict), we got " + f"{type(remote_task_handler_kwargs)}" + ) + # remote_task_handler_kwargs mixes FileTaskHandler kwargs with IO kwargs; only the + # latter belong to this class (same split as airflow_local_settings.py). + fth_params = frozenset(inspect.signature(FileTaskHandler.__init__).parameters) - { + "self", + "base_log_folder", + } + io_kwargs = {k: v for k, v in remote_task_handler_kwargs.items() if k not in fth_params} + remote_base_log_folder = conf.get_mandatory_value("logging", "remote_base_log_folder") + log_name = urlsplit(remote_base_log_folder).path[1:] + if not log_name: + raise ValueError( + "Cannot derive a Stackdriver log name from " + f"logging/remote_base_log_folder: {remote_base_log_folder!r}" + ) + return cls( + **{ + "base_log_folder": os.path.expanduser(conf.get_mandatory_value("logging", "base_log_folder")), + "gcp_log_name": log_name, + "gcp_key_path": conf.get_mandatory_value("logging", "GOOGLE_KEY_PATH", fallback=None), + "delete_local_copy": conf.getboolean("logging", "delete_local_logs"), + } + | io_kwargs, + ) + @cached_property def credentials_and_project(self) -> tuple[Credentials, str]: credentials, project = get_credentials_and_project_id( diff --git a/providers/google/src/airflow/providers/google/get_provider_info.py b/providers/google/src/airflow/providers/google/get_provider_info.py index 8a1024cd71431..7334666783da8 100644 --- a/providers/google/src/airflow/providers/google/get_provider_info.py +++ b/providers/google/src/airflow/providers/google/get_provider_info.py @@ -1713,7 +1713,11 @@ def get_provider_info(): { "classpath": "airflow.providers.google.cloud.log.gcs_task_handler.GCSRemoteLogIO", "scheme": "gs", - } + }, + { + "classpath": "airflow.providers.google.cloud.log.stackdriver_task_handler.StackdriverRemoteLogIO", + "scheme": "stackdriver", + }, ], "queues": [ "airflow.providers.google.event_scheduling.events.pubsub.PubSubMessageQueueEventTriggerContainer" diff --git a/providers/google/tests/unit/google/cloud/log/test_stackdriver_task_handler.py b/providers/google/tests/unit/google/cloud/log/test_stackdriver_task_handler.py index de4fd3e0cdb2c..c66ff24421b9d 100644 --- a/providers/google/tests/unit/google/cloud/log/test_stackdriver_task_handler.py +++ b/providers/google/tests/unit/google/cloud/log/test_stackdriver_task_handler.py @@ -55,6 +55,69 @@ def clean_stackdriver_handlers(): del handler +class TestStackdriverRemoteLogIOFromConfig: + @conf_vars( + { + ("logging", "base_log_folder"): "~/airflow/logs", + ("logging", "remote_base_log_folder"): "stackdriver:///airflow-tasks", + ("logging", "delete_local_logs"): "True", + ("logging", "google_key_path"): "/tmp/google-key.json", + } + ) + def test_from_config(self): + subject = StackdriverRemoteLogIO.from_config() + + assert subject.base_log_folder == Path("~/airflow/logs").expanduser() + assert subject.gcp_log_name == "airflow-tasks" + assert subject.gcp_key_path == "/tmp/google-key.json" + assert subject.delete_local_copy is True + + @conf_vars( + { + ("logging", "base_log_folder"): "/tmp/airflow/logs", + ("logging", "remote_base_log_folder"): "stackdriver:///airflow-tasks", + ("logging", "delete_local_logs"): "False", + ("logging", "remote_task_handler_kwargs"): '{"delete_local_copy": true, "max_bytes": 1024}', + } + ) + def test_from_config_applies_io_kwargs_and_filters_file_handler_kwargs(self): + subject = StackdriverRemoteLogIO.from_config() + + assert subject.delete_local_copy is True + assert not hasattr(subject, "max_bytes") + + @conf_vars({("logging", "remote_task_handler_kwargs"): '["not", "a", "dict"]'}) + def test_from_config_rejects_non_dict_remote_task_handler_kwargs(self): + with pytest.raises(ValueError, match="remote_task_handler_kwargs"): + StackdriverRemoteLogIO.from_config() + + @pytest.mark.parametrize( + "remote_base_log_folder", + [ + pytest.param("stackdriver://", id="scheme-only"), + pytest.param("stackdriver://host", id="no-path"), + ], + ) + def test_from_config_rejects_remote_base_without_log_name(self, remote_base_log_folder): + with conf_vars({("logging", "remote_base_log_folder"): remote_base_log_folder}): + with pytest.raises(ValueError, match="Stackdriver log name"): + StackdriverRemoteLogIO.from_config() + + def test_provider_registers_stackdriver_scheme(self): + from airflow.providers_manager import ProvidersManager + + manager = ProvidersManager() + if not hasattr(manager, "remote_logging_handler_by_scheme"): + pytest.skip("Airflow core does not support remote logging provider dispatch") + + info = manager.remote_logging_handler_by_scheme("stackdriver") + + assert info is not None + assert info.classpath == ( + "airflow.providers.google.cloud.log.stackdriver_task_handler.StackdriverRemoteLogIO" + ) + + class TestStackdriverRemoteLogIO: @pytest.fixture(autouse=True) def _setup(self, tmp_path):