Skip to content
Draft
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
2 changes: 2 additions & 0 deletions airflow-core/src/airflow/serialization/definitions/dag.py
Original file line number Diff line number Diff line change
Expand Up @@ -122,6 +122,8 @@ class SerializedDAG:
fail_fast: bool = False
has_on_failure_callback: bool = False
has_on_success_callback: bool = False
# Derived, never authored -- see ``DAG.is_mixed_language_dag``.
is_mixed_language_dag: bool = False
is_paused_upon_creation: bool | None = None
max_active_runs: int = 16
max_active_tasks: int = 16
Expand Down
3 changes: 2 additions & 1 deletion airflow-core/src/airflow/serialization/schema.json
Original file line number Diff line number Diff line change
Expand Up @@ -226,7 +226,8 @@
"edge_info": { "$ref": "#/definitions/edge_info" },
"dag_dependencies": { "$ref": "#/definitions/dag_dependencies" },
"disable_bundle_versioning": {"type": "boolean" },
"rerun_with_latest_version": {"type": ["boolean", "null"], "default": null}
"rerun_with_latest_version": {"type": ["boolean", "null"], "default": null},
"is_mixed_language_dag": { "type": "boolean", "default": false }
},
"required": [
"dag_id",
Expand Down
11 changes: 11 additions & 0 deletions airflow-core/src/airflow/serialization/serialized_objects.py
Original file line number Diff line number Diff line change
Expand Up @@ -1760,6 +1760,11 @@ def serialize_dag(cls, dag: DAG) -> dict:
if dag.has_on_failure_callback:
serialized_dag["has_on_failure_callback"] = True

# Likewise only stored when True. On a Python Dag this is derived from the presence of
# @task.stub tasks; a Lang-SDK producer sets it on its own Dag instead.
if dag.is_mixed_language_dag:
serialized_dag["is_mixed_language_dag"] = True

# TODO: Move this logic to a better place -- ideally before serializing contents of default_args.
# There is some duplication with this and SerializedBaseOperator.partial_kwargs serialization.
# Ideally default_args goes through same logic as fields of SerializedBaseOperator.
Expand Down Expand Up @@ -1886,6 +1891,8 @@ def _deserialize_dag_internal(
if "has_on_failure_callback" in encoded_dag:
dag.has_on_failure_callback = True

dag.is_mixed_language_dag = encoded_dag.get("is_mixed_language_dag") is True

dag.deadline = encoded_dag.get("deadline")

keys_to_set_none = dag.get_serialized_fields() - encoded_dag.keys() - cls._CONSTRUCTOR_PARAMS.keys()
Expand Down Expand Up @@ -2307,6 +2314,10 @@ def __getattr__(self, name: str, /) -> Any:
def timetable(self) -> Timetable:
return decode_timetable(self.data["dag"]["timetable"])

@property
def is_mixed_language_dag(self) -> bool:
return self.data["dag"].get("is_mixed_language_dag") is True

@property
def has_task_concurrency_limits(self) -> bool:
return any(
Expand Down
80 changes: 80 additions & 0 deletions airflow-core/tests/unit/serialization/test_dag_serialization.py
Original file line number Diff line number Diff line change
Expand Up @@ -1573,6 +1573,7 @@ def test_dag_serialized_fields_with_schema(self):
"has_on_failure_callback",
"dag_dependencies",
"params",
"is_mixed_language_dag",
}

keys_for_backwards_compat: set = {
Expand Down Expand Up @@ -2407,6 +2408,85 @@ def test_dag_on_failure_callback_roundtrip(self, passed_failure_callback, expect

assert deserialized_dag.has_on_failure_callback is expected_value

def _serialize_simple_dag(self, dag_id):
dag = DAG(dag_id=dag_id, schedule=None)
BaseOperator(task_id="simple_task", dag=dag, start_date=datetime(2019, 8, 1))
return DagSerialization.to_dict(dag)

def test_is_mixed_language_dag_absent_without_stub_tasks(self):
"""A pure-Python Dag omits the key entirely, so every reader falls back to the False default."""
serialized_dag = self._serialize_simple_dag("test_is_mixed_language_dag_absent")

assert "is_mixed_language_dag" not in serialized_dag["dag"]
assert DagSerialization.from_dict(serialized_dag).is_mixed_language_dag is False
assert LazyDeserializedDAG(data=serialized_dag).is_mixed_language_dag is False

def test_is_mixed_language_dag_derived_from_stub_tasks(self):
"""A Dag holding a @task.stub task is mixed-language, and serialization records that."""
from airflow.providers.standard.decorators.stub import stub

def run_on_golang(): ...

with DAG(
dag_id="test_is_mixed_language_dag_derived",
schedule=None,
start_date=datetime(2019, 8, 1),
) as dag:
BaseOperator(task_id="python_task", start_date=datetime(2019, 8, 1))
stub(run_on_golang, queue="golang")()

assert dag.is_mixed_language_dag is True

serialized_dag = DagSerialization.to_dict(dag)

assert serialized_dag["dag"]["is_mixed_language_dag"] is True
assert DagSerialization.from_dict(serialized_dag).is_mixed_language_dag is True
assert LazyDeserializedDAG(data=serialized_dag).is_mixed_language_dag is True

@pytest.mark.parametrize("flag", [True, False])
def test_is_mixed_language_dag_roundtrip(self, flag):
"""A producer-supplied flag validates against the schema and survives deserialization."""
serialized_dag = self._serialize_simple_dag(f"test_is_mixed_language_dag_roundtrip_{flag}")
serialized_dag["dag"]["is_mixed_language_dag"] = flag
DagSerialization.validate_schema(serialized_dag)

assert DagSerialization.from_dict(serialized_dag).is_mixed_language_dag is flag
assert LazyDeserializedDAG(data=serialized_dag).is_mixed_language_dag is flag

def test_is_mixed_language_dag_not_author_settable(self):
"""
The flag is derived, never authored. Neither the constructor nor attribute assignment may set
it, so a Dag can never claim to be mixed-language without actually holding a stub task.
"""
assert "is_mixed_language_dag" not in {a.name for a in attrs.fields(DAG)}
assert "is_mixed_language_dag" not in DAG.get_serialized_fields()

with pytest.raises(TypeError, match="is_mixed_language_dag"):
DAG(dag_id="test_is_mixed_language_dag_rejected", schedule=None, is_mixed_language_dag=True)

dag = DAG(dag_id="test_is_mixed_language_dag_not_leaked", schedule=None)
BaseOperator(task_id="simple_task", dag=dag, start_date=datetime(2019, 8, 1))

with pytest.raises(AttributeError, match="can not be set"):
dag.is_mixed_language_dag = True

serialized_dag = DagSerialization.to_dict(dag)

assert "is_mixed_language_dag" not in serialized_dag["dag"]
assert DagSerialization.from_dict(serialized_dag).is_mixed_language_dag is False

@pytest.mark.parametrize("raw", ["false", "true", 0, 1, None, [], {"a": 1}])
def test_is_mixed_language_dag_fails_closed_on_non_boolean(self, raw):
"""
The read path runs no schema validation, so a malformed producer value must fall back to False.
Truthiness would read "false" and 1 as True and wrongly discard a native Lang-SDK Dag.
"""
serialized_dag = self._serialize_simple_dag("test_is_mixed_language_dag_non_boolean")
serialized_dag["dag"]["is_mixed_language_dag"] = raw

assert DagSerialization.from_dict(serialized_dag).is_mixed_language_dag is False
assert LazyDeserializedDAG(data=serialized_dag).is_mixed_language_dag is False

@pytest.mark.parametrize(
("dag_arg", "conf_arg", "expected"),
[
Expand Down
8 changes: 5 additions & 3 deletions scripts/in_container/run_schema_defaults_check.py
Original file line number Diff line number Diff line change
Expand Up @@ -226,12 +226,14 @@ def compare_dag_defaults() -> list[str]:
# Check for schema defaults that don't have corresponding server defaults
for field_name, schema_value in schema_defaults.items():
if field_name not in server_defaults:
# Some schema fields are computed properties (like has_on_*_callback)
computed_properties = {
# Fields SerializedDAG carries but deliberately keeps out of get_serialized_fields(): computed
# ones (has_on_*_callback) and ones only a non-Python producer ever sets.
fields_outside_serialized_fields = {
"has_on_success_callback",
"has_on_failure_callback",
"is_mixed_language_dag",
}
if field_name not in computed_properties:
if field_name not in fields_outside_serialized_fields:
errors.append(
f"DAG schema has default for '{field_name}' = {schema_value!r} but no corresponding server default"
)
Expand Down
20 changes: 20 additions & 0 deletions task-sdk/src/airflow/sdk/definitions/dag.py
Original file line number Diff line number Diff line change
Expand Up @@ -86,6 +86,9 @@

TAG_MAX_LEN = 100

# ``@task.stub`` lives in providers.standard, which task-sdk must not import.
STUB_OPERATOR_NAME = "@task.stub"

__all__ = [
"DAG",
"dag",
Expand Down Expand Up @@ -806,6 +809,23 @@ def tasks(self, val):
def task_ids(self) -> list[str]:
return list(self.task_dict)

@property
def is_mixed_language_dag(self) -> bool:
"""
Whether this Dag pairs Python with a Lang-SDK runtime, i.e. holds any ``@task.stub`` task.

Always derived from the Dag's tasks, never author-set, so that a Lang-SDK Dag processor can
tell a Dag that merely backs stub tasks from one authored natively in that language.
"""
return any(task.operator_name == STUB_OPERATOR_NAME for task in self.task_dict.values())

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Operator name should not be used for logical conditions. Add a flag instead.


@is_mixed_language_dag.setter
def is_mixed_language_dag(self, val):
raise AttributeError(
"DAG.is_mixed_language_dag can not be set. It is derived from whether the Dag has any "
"@task.stub task."
)

@property
def teardowns(self) -> list[Operator]:
return [task for task in self.tasks if getattr(task, "is_teardown", None)]
Expand Down
37 changes: 37 additions & 0 deletions task-sdk/tests/task_sdk/definitions/test_dag.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,6 +37,7 @@
)
from airflow.sdk.bases.operator import BaseOperator
from airflow.sdk.bases.timetable import BaseTimetable
from airflow.sdk.definitions.dag import STUB_OPERATOR_NAME
from airflow.sdk.definitions.param import DagParam, ParamsDict
from airflow.sdk.exceptions import AirflowDagCycleException, DuplicateTaskIdFound, RemovedInAirflow4Warning
from airflow.utils.types import DagRunType
Expand Down Expand Up @@ -660,6 +661,42 @@ def test__tags_mutable():
assert test_dag.tags == expected_tags


class _FakeStubOperator(BaseOperator):
"""Stands in for providers.standard's ``@task.stub``, which task-sdk must not import."""

custom_operator_name = STUB_OPERATOR_NAME


@pytest.mark.parametrize(
("task_classes", "expected"),
[
pytest.param([], False, id="no-tasks"),
pytest.param([BaseOperator], False, id="python-only"),
pytest.param([_FakeStubOperator], True, id="stub-only"),
pytest.param([BaseOperator, _FakeStubOperator], True, id="mixed"),
pytest.param([_FakeStubOperator, _FakeStubOperator], True, id="several-stubs"),
],
)
def test_is_mixed_language_dag_derived_from_stub_tasks(task_classes, expected):
with DAG(dag_id="derived") as test_dag:
for i, task_class in enumerate(task_classes):
task_class(task_id=f"task_{i}")

assert test_dag.is_mixed_language_dag is expected


def test_is_mixed_language_dag_cannot_be_set():
"""It is derived from the Dag's tasks, so no author-facing way to set it may exist."""
with pytest.raises(TypeError, match="is_mixed_language_dag"):
DAG(dag_id="rejected", is_mixed_language_dag=True)

test_dag = DAG(dag_id="not_settable")
with pytest.raises(AttributeError, match="can not be set"):
test_dag.is_mixed_language_dag = True

assert test_dag.is_mixed_language_dag is False


def test_create_dag_while_active_context():
"""Test that we can safely create a Dag whilst a Dag is activated via ``with dag1:``."""
with DAG(dag_id="simple_dag"):
Expand Down
Loading