Skip to content
Open
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
5 changes: 5 additions & 0 deletions airflow-core/docs/core-concepts/dags.rst
Original file line number Diff line number Diff line change
Expand Up @@ -721,6 +721,11 @@ You can either do this all inside of the Dag bundle, with a standard filesystem
package1/__init__.py
package1/functions.py

Template files used by your tasks, such as ``.sql`` or ``.sh`` files, can be packaged in the same zip file.
They are found relative to the Dag file's folder inside the archive, as they would be in an unpacked Dag bundle.
``template_searchpath`` entries can also point to a folder inside the archive, for example
``os.path.join(os.path.dirname(__file__), "sql")``.

Note that packaged Dags come with some caveats:

* They cannot be used if you have pickling enabled for serialization
Expand Down
82 changes: 80 additions & 2 deletions task-sdk/src/airflow/sdk/definitions/_internal/templater.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,13 +20,17 @@
import datetime
import logging
import os
from collections.abc import Collection, Iterable, Iterator, Sequence
import posixpath
import zipfile
from collections.abc import Callable, Collection, Iterable, Iterator, Sequence
from dataclasses import dataclass
from pathlib import Path, PurePosixPath
from typing import TYPE_CHECKING, Any

import jinja2
import jinja2.nativetypes
import jinja2.sandbox
from jinja2.loaders import split_template_path

from airflow.sdk import ObjectStoragePath
from airflow.sdk.definitions._internal.mixins import ResolveMixin
Expand Down Expand Up @@ -374,6 +378,80 @@ def ts_nodash_with_tz_filter(value: datetime.date | datetime.time | None) -> str
}


def _find_zip_root(path: str) -> tuple[str, str] | None:
"""Return ``(archive, prefix)`` when *path* is a zip archive or a directory inside one."""
candidate = Path(path)
existing = next((p for p in (candidate, *candidate.parents) if os.path.exists(p)), None)
if existing is None or not os.path.isfile(existing) or not zipfile.is_zipfile(existing):
return None
prefix = candidate.relative_to(existing).as_posix()
if ".." in PurePosixPath(prefix).parts:
return None
return os.fspath(existing), "" if prefix == "." else prefix


class _ZipArchiveLoader(jinja2.BaseLoader):
"""Load templates from a directory inside a zip archive, reading members without extracting them."""

def __init__(self, archive: str, prefix: str, encoding: str = "utf-8") -> None:
self.archive = archive
self.prefix = prefix
self.encoding = encoding

def get_source(
self, environment: jinja2.Environment, template: str
) -> tuple[str, str, Callable[[], bool]]:
member = posixpath.join(self.prefix, *split_template_path(template))
mtime = os.path.getmtime(self.archive)
# Opened per lookup: a handle kept across the task supervisor's fork would share its file offset.
with zipfile.ZipFile(self.archive) as zf:
try:
source = zf.read(member).decode(self.encoding)
except KeyError:
raise jinja2.TemplateNotFound(template) from None

def is_up_to_date() -> bool:
try:
return os.path.getmtime(self.archive) == mtime
except OSError:
return False

return source, os.path.join(self.archive, member), is_up_to_date


class _SearchPathLoader(jinja2.BaseLoader):
"""Try each search path's loader in order, reporting the search paths like ``FileSystemLoader``."""

def __init__(self, searchpath: list[str], loaders: list[jinja2.BaseLoader]) -> None:
self.searchpath = searchpath
self.loaders = loaders

def get_source(
self, environment: jinja2.Environment, template: str
) -> tuple[str, str | None, Callable[[], bool] | None]:
for loader in self.loaders:
try:
return loader.get_source(environment, template)
except jinja2.TemplateNotFound:
continue
plural = "path" if len(self.searchpath) == 1 else "paths"
paths = ", ".join(repr(path) for path in self.searchpath)
raise jinja2.TemplateNotFound(template, f"{template!r} not found in search {plural}: {paths}")


def _build_template_loader(searchpath: list[str]) -> jinja2.BaseLoader:
zip_roots = [None if os.path.isdir(path) else _find_zip_root(path) for path in searchpath]
if not any(zip_roots):
return jinja2.FileSystemLoader(searchpath)
return _SearchPathLoader(
searchpath,
[
_ZipArchiveLoader(*zip_root) if zip_root else jinja2.FileSystemLoader(path)
for path, zip_root in zip(searchpath, zip_roots, strict=True)
],
)


def create_template_env(
*,
native: bool = False,
Expand All @@ -391,7 +469,7 @@ def create_template_env(
"cache_size": 0,
}
if searchpath:
jinja_env_options["loader"] = jinja2.FileSystemLoader(searchpath)
jinja_env_options["loader"] = _build_template_loader(searchpath)
if jinja_environment_kwargs:
jinja_env_options.update(jinja_environment_kwargs)

Expand Down
30 changes: 30 additions & 0 deletions task-sdk/tests/task_sdk/bases/test_operator.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,8 +20,10 @@
import asyncio
import copy
import logging
import os
import uuid
import warnings
import zipfile
from datetime import UTC, date, datetime, timedelta
from typing import NamedTuple
from unittest import mock
Expand Down Expand Up @@ -120,6 +122,17 @@ def __init__(self, arg1: str = "", arg2: str = "", **kwargs):
)


class MockFileTemplateOperator(BaseOperator):
"""Operator with a ``.sql`` template file nested in a dict, like ``BigQueryInsertJobOperator``."""

template_fields = ("configuration",)
template_ext = (".sql",)

def __init__(self, configuration: dict, **kwargs):
super().__init__(**kwargs)
self.configuration = configuration


class TestBaseOperator:
def setup_method(self, method):
MockOperator.start_from_trigger = False
Expand Down Expand Up @@ -757,6 +770,23 @@ def fn_to_template(**kwargs):
task.render_template_fields({})
assert task.arg2 == "foo_barbarbar"

@pytest.mark.parametrize("op_native", [None, True], ids=["dag-env", "operator-native-env"])
def test_render_template_fields_reads_template_file_from_zipped_dag(self, tmp_path, op_native):
archive = tmp_path / "dags.zip"
with zipfile.ZipFile(archive, "w") as zf:
zf.writestr("test_sql/test.sql", "SELECT column_a FROM {{ table }}")
with DAG("zipped_dag", schedule=None, start_date=DEFAULT_DATE) as dag:
task = MockFileTemplateOperator(
task_id="op1",
configuration={"query": {"query": "test_sql/test.sql"}},
render_template_as_native_obj=op_native,
)
dag.fileloc = os.path.join(archive, "test.py")

task.render_template_fields(context={"table": "test"})

assert task.configuration == {"query": {"query": "SELECT column_a FROM test"}}

@pytest.mark.parametrize("content", [object(), uuid.uuid4()])
def test_render_template_fields_no_change(self, content):
"""Tests if non-templatable types remain unchanged."""
Expand Down
119 changes: 118 additions & 1 deletion task-sdk/tests/task_sdk/definitions/_internal/test_templater.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,14 +17,22 @@

from __future__ import annotations

import os
import zipfile
from datetime import UTC, datetime
from pathlib import Path
from unittest.mock import MagicMock, NonCallableMagicMock

import jinja2
import pytest

from airflow.sdk import DAG, ObjectStoragePath
from airflow.sdk.definitions._internal.templater import LiteralValue, SandboxedEnvironment, Templater
from airflow.sdk.definitions._internal.templater import (
LiteralValue,
SandboxedEnvironment,
Templater,
create_template_env,
)


class TestTemplater:
Expand Down Expand Up @@ -300,6 +308,115 @@ def test_do_render_template_fields_renders_list_values(self):
assert parent.items == ["first", "second"]


def _write_zip(path: Path, members: dict[str, str]) -> Path:
with zipfile.ZipFile(path, "w") as zf:
for name, content in members.items():
zf.writestr(name, content)
return path


def _write_template_dir(path: Path, templates: dict[str, str]) -> Path:
path.mkdir()
for name, content in templates.items():
(path / name).write_text(content)
return path


class TestCreateTemplateEnvZip:
@pytest.mark.parametrize(
("archive_name", "searchpath_parts", "template_name"),
[
pytest.param("dags.zip", (), "sql/query.sql", id="archive"),
pytest.param("dags.zip", ("sql",), "query.sql", id="directory-in-archive"),
pytest.param("DAGS.ZIP", (), "sql/query.sql", id="uppercase-extension"),
],
)
def test_loads_template_from_zip(self, tmp_path, archive_name, searchpath_parts, template_name):
archive = _write_zip(tmp_path / archive_name, {"sql/query.sql": "SELECT {{ x }}"})

env = create_template_env(searchpath=[os.path.join(archive, *searchpath_parts)])

assert env.get_template(template_name).render(x=1) == "SELECT 1"

@pytest.mark.parametrize(
"searchpath_parts",
[
pytest.param(("templates",), id="directory"),
pytest.param(("templates.zip",), id="directory-named-like-archive"),
pytest.param(("not_a_zip.zip", "sql"), id="non-zip-file"),
pytest.param(("missing", "sql"), id="missing-path"),
],
)
def test_keeps_filesystem_loader_without_zip(self, tmp_path, searchpath_parts):
(tmp_path / "templates").mkdir()
(tmp_path / "templates.zip").mkdir()
(tmp_path / "not_a_zip.zip").write_text("plain text")
searchpath = os.path.join(tmp_path, *searchpath_parts)

env = create_template_env(searchpath=[searchpath])

assert type(env.loader) is jinja2.FileSystemLoader
assert env.loader.searchpath == [searchpath]

@pytest.mark.parametrize(("zip_first", "expected"), [(True, "from zip"), (False, "from directory")])
def test_searches_zip_and_directory_in_order(self, tmp_path, zip_first, expected):
archive = str(_write_zip(tmp_path / "dags.zip", {"query.sql": "from zip"}))
directory = str(_write_template_dir(tmp_path / "templates", {"query.sql": "from directory"}))

env = create_template_env(searchpath=[archive, directory] if zip_first else [directory, archive])

assert env.get_template("query.sql").render() == expected

def test_template_not_found_lists_search_paths(self, tmp_path):
archive = str(_write_zip(tmp_path / "dags.zip", {"query.sql": "SELECT 1"}))
directory = str(_write_template_dir(tmp_path / "templates", {}))
env = create_template_env(searchpath=[archive, directory])

with pytest.raises(jinja2.TemplateNotFound) as ctx:
env.get_template("missing.sql")

assert ctx.value.name == "missing.sql"
assert ctx.value.message == f"'missing.sql' not found in search paths: {archive!r}, {directory!r}"

@pytest.mark.parametrize(
("searchpath_parts", "template_name"),
[
pytest.param(("sql",), "../secret.sql", id="name-leaves-directory"),
pytest.param((), "../outside.sql", id="name-leaves-archive"),
pytest.param(("..",), "outside.sql", id="searchpath-leaves-archive"),
],
)
def test_does_not_escape_search_root(self, tmp_path, searchpath_parts, template_name):
archive = _write_zip(
tmp_path / "dags.zip",
{"sql/query.sql": "SELECT 1", "secret.sql": "secret", "../outside.sql": "outside"},
)
env = create_template_env(searchpath=[os.path.join(archive, *searchpath_parts)])

with pytest.raises(jinja2.TemplateNotFound):
env.get_template(template_name)

def test_unreadable_archive_raises_instead_of_falling_through(self, tmp_path):
archive = _write_zip(tmp_path / "dags.zip", {"query.sql": "from zip"})
directory = _write_template_dir(tmp_path / "templates", {"query.sql": "from directory"})
env = create_template_env(searchpath=[str(archive), str(directory)])
archive.write_bytes(b"no longer a zip archive")

with pytest.raises(zipfile.BadZipFile):
env.get_template("query.sql")

def test_up_to_date_check_tracks_archive_mtime(self, tmp_path):
archive = _write_zip(tmp_path / "dags.zip", {"query.sql": "SELECT 1"})
env = create_template_env(searchpath=[str(archive)])

_, filename, is_up_to_date = env.loader.get_source(env, "query.sql")

assert filename == os.path.join(archive, "query.sql")
assert is_up_to_date()
os.utime(archive, (0, 0))
assert not is_up_to_date()


@pytest.fixture
def env():
return SandboxedEnvironment(undefined=jinja2.StrictUndefined, cache_size=0)
Expand Down
32 changes: 32 additions & 0 deletions task-sdk/tests/task_sdk/importers/test_zip_importer.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,7 @@

import os
import py_compile
import textwrap
import zipfile
from pathlib import Path
from types import SimpleNamespace
Expand Down Expand Up @@ -181,6 +182,37 @@ def test_import_zip_archive_with_dags(self, mock_bundle):
assert dags[0].bundle_name == "test_bundle"
assert len(errors) == 0

def test_imported_dag_resolves_template_file_from_archive(self, mock_bundle):
zip_path = mock_bundle.path / "templated.zip"
with zipfile.ZipFile(zip_path, "w") as z:
z.writestr(
"pkg/templated_dag.py",
textwrap.dedent(
"""\
from airflow.sdk import DAG, BaseOperator

class SqlOperator(BaseOperator):
template_fields = ("sql",)
template_ext = (".sql",)

def __init__(self, sql, **kwargs):
super().__init__(**kwargs)
self.sql = sql

with DAG("zip_templated_dag"):
SqlOperator(task_id="t", sql="sql/query.sql")
"""
),
)
z.writestr("pkg/sql/query.sql", "SELECT 1")

dags, errors = _import_all(ZipImporter(), mock_bundle)
assert errors == []
(dag,) = dags
dag.resolve_template_files()

assert dag.get_task("t").sql == "SELECT 1"

def test_import_zip_archive_with_pyc_dag(self, mock_bundle, tmp_path):
source_file = tmp_path / "compiled_dag.py"
source_file.write_text("from airflow.sdk import DAG\ndag = DAG('zip_pyc_dag')\n")
Expand Down