diff --git a/airflow-core/docs/core-concepts/dags.rst b/airflow-core/docs/core-concepts/dags.rst index 724619aae9ff6..f67711183eed2 100644 --- a/airflow-core/docs/core-concepts/dags.rst +++ b/airflow-core/docs/core-concepts/dags.rst @@ -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 diff --git a/task-sdk/src/airflow/sdk/definitions/_internal/templater.py b/task-sdk/src/airflow/sdk/definitions/_internal/templater.py index cfe4a6100e482..f5e4ae6d06d33 100644 --- a/task-sdk/src/airflow/sdk/definitions/_internal/templater.py +++ b/task-sdk/src/airflow/sdk/definitions/_internal/templater.py @@ -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 @@ -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, @@ -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) diff --git a/task-sdk/tests/task_sdk/bases/test_operator.py b/task-sdk/tests/task_sdk/bases/test_operator.py index d2060362f2b94..cd317cc635a0c 100644 --- a/task-sdk/tests/task_sdk/bases/test_operator.py +++ b/task-sdk/tests/task_sdk/bases/test_operator.py @@ -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 @@ -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 @@ -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.""" diff --git a/task-sdk/tests/task_sdk/definitions/_internal/test_templater.py b/task-sdk/tests/task_sdk/definitions/_internal/test_templater.py index 6018ba5a01a34..6877277634c7d 100644 --- a/task-sdk/tests/task_sdk/definitions/_internal/test_templater.py +++ b/task-sdk/tests/task_sdk/definitions/_internal/test_templater.py @@ -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: @@ -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) diff --git a/task-sdk/tests/task_sdk/importers/test_zip_importer.py b/task-sdk/tests/task_sdk/importers/test_zip_importer.py index 3d4a52bc696d8..05911a941950e 100644 --- a/task-sdk/tests/task_sdk/importers/test_zip_importer.py +++ b/task-sdk/tests/task_sdk/importers/test_zip_importer.py @@ -20,6 +20,7 @@ import os import py_compile +import textwrap import zipfile from pathlib import Path from types import SimpleNamespace @@ -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")