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
4 changes: 4 additions & 0 deletions providers/amazon/provider.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -1123,6 +1123,10 @@ logging:
- airflow.providers.amazon.aws.log.s3_task_handler.S3TaskHandler
- airflow.providers.amazon.aws.log.cloudwatch_task_handler.CloudwatchTaskHandler

remote-logging:
- classpath: airflow.providers.amazon.aws.log.cloudwatch_task_handler.CloudWatchRemoteLogIO
scheme: cloudwatch

config:
aws:
description: This section contains settings for Amazon Web Services (AWS) integration.
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@

import contextlib
import copy
import inspect
import json
import logging
import os
Expand All @@ -28,6 +29,7 @@
from functools import cached_property
from pathlib import Path
from typing import TYPE_CHECKING, Any
from urllib.parse import urlsplit

import attrs
import watchtower
Expand Down Expand Up @@ -105,6 +107,40 @@ def _(self):
def _(self):
return self.log_group_arn.split(":")[3]

@classmethod
def from_config(cls) -> CloudWatchRemoteLogIO:
"""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) - {
Comment thread
ferruzzi marked this conversation as resolved.
"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")
url_parts = urlsplit(remote_base_log_folder)
log_group_arn = url_parts.netloc + url_parts.path
if not log_group_arn:
raise ValueError(
"Cannot derive a CloudWatch log group ARN 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")),
"remote_base": remote_base_log_folder,
"delete_local_copy": conf.getboolean("logging", "delete_local_logs"),
"log_group_arn": log_group_arn,
}
| io_kwargs,
)

@cached_property
def hook(self):
"""Returns AwsLogsHook."""
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -1249,6 +1249,12 @@ def get_provider_info():
"airflow.providers.amazon.aws.log.s3_task_handler.S3TaskHandler",
"airflow.providers.amazon.aws.log.cloudwatch_task_handler.CloudwatchTaskHandler",
],
"remote-logging": [
{
"classpath": "airflow.providers.amazon.aws.log.cloudwatch_task_handler.CloudWatchRemoteLogIO",
"scheme": "cloudwatch",
}
],
"config": {
"aws": {
"description": "This section contains settings for Amazon Web Services (AWS) integration.",
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@

import contextlib
import logging
import os
import textwrap
import time
from datetime import datetime as dt, timedelta, timezone
Expand Down Expand Up @@ -84,6 +85,107 @@ def _cleanup_cloudwatch_handlers():
logging._removeHandlerRef(handler_ref)


class TestCloudWatchRemoteLogIOFromConfig:
@conf_vars(
{
("logging", "base_log_folder"): "~/airflow/logs",
("logging", "remote_base_log_folder"): (
"cloudwatch://arn:aws:logs:us-west-2:123456789098:log-group:log_group_name"
),
("logging", "delete_local_logs"): "True",
}
)
def test_from_config(self):
subject = CloudWatchRemoteLogIO.from_config()

assert (
subject.remote_base == "cloudwatch://arn:aws:logs:us-west-2:123456789098:log-group:log_group_name"
)
assert subject.base_log_folder == Path(os.path.expanduser("~/airflow/logs"))
assert subject.delete_local_copy is True
assert subject.log_group_arn == "arn:aws:logs:us-west-2:123456789098:log-group:log_group_name"
assert subject.log_group == "log_group_name"
assert subject.region_name == "us-west-2"

@conf_vars(
{
("logging", "base_log_folder"): "/tmp/airflow/logs",
("logging", "remote_base_log_folder"): (
"cloudwatch://arn:aws:logs:us-west-2:123456789098:log-group:log_group_name"
),
("logging", "delete_local_logs"): "False",
("logging", "remote_task_handler_kwargs"): (
'{"log_stream_name": "custom-stream", "max_bytes": 1024}'
),
}
)
def test_from_config_applies_io_kwargs_and_filters_file_handler_kwargs(self):
subject = CloudWatchRemoteLogIO.from_config()

assert subject.log_stream_name == "custom-stream"
assert subject.delete_local_copy is False
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"):
CloudWatchRemoteLogIO.from_config()

@conf_vars({("logging", "remote_base_log_folder"): "cloudwatch://"})
def test_from_config_rejects_remote_base_without_log_group_arn(self):
with pytest.raises(ValueError, match="log group ARN"):
CloudWatchRemoteLogIO.from_config()

def test_provider_registers_cloudwatch_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("cloudwatch")

assert info is not None
assert (
info.classpath == "airflow.providers.amazon.aws.log.cloudwatch_task_handler.CloudWatchRemoteLogIO"
)

@pytest.mark.parametrize(
"manager_classpath",
[
pytest.param("airflow.providers_manager.ProvidersManager", id="core"),
pytest.param(
"airflow.sdk.providers_manager_runtime.ProvidersManagerTaskRuntime", id="task-runtime"
),
],
)
@conf_vars(
{
("logging", "remote_logging"): "True",
("logging", "remote_base_log_folder"): (
"cloudwatch://arn:aws:logs:us-west-2:123456789098:log-group:log_group_name"
),
("logging", "remote_log_conn_id"): "aws_default",
}
)
def test_resolve_remote_task_log_uses_provider_dispatch_not_local_settings(self, manager_classpath):
factory = pytest.importorskip("airflow._shared.logging.factory")
from airflow._shared.module_loading import import_string
from airflow.configuration import conf

with mock.patch.object(factory, "discover_remote_log_handler", autospec=True) as legacy_discover:
remote_task_log, conn_id = factory.resolve_remote_task_log(
conf=conf,
providers_manager=import_string(manager_classpath)(),
import_string=import_string,
)

assert isinstance(remote_task_log, CloudWatchRemoteLogIO)
assert remote_task_log.log_group_arn == "arn:aws:logs:us-west-2:123456789098:log-group:log_group_name"
assert conn_id == "aws_default"
legacy_discover.assert_not_called()


# We only test this directly on Airflow 3
@pytest.mark.skipif(not AIRFLOW_V_3_0_PLUS, reason="This path only works on Airflow 3")
class TestCloudRemoteLogIO:
Expand Down
Loading