diff --git a/airflow-core/src/airflow/serialization/definitions/dag.py b/airflow-core/src/airflow/serialization/definitions/dag.py index 8ed0fee2ccabd..158a4fccdaea7 100644 --- a/airflow-core/src/airflow/serialization/definitions/dag.py +++ b/airflow-core/src/airflow/serialization/definitions/dag.py @@ -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 diff --git a/airflow-core/src/airflow/serialization/schema.json b/airflow-core/src/airflow/serialization/schema.json index 872c3a1331ee3..5fcb36cc2218a 100644 --- a/airflow-core/src/airflow/serialization/schema.json +++ b/airflow-core/src/airflow/serialization/schema.json @@ -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", diff --git a/airflow-core/src/airflow/serialization/serialized_objects.py b/airflow-core/src/airflow/serialization/serialized_objects.py index f42e51e28f7d0..ff0f769d0da6b 100644 --- a/airflow-core/src/airflow/serialization/serialized_objects.py +++ b/airflow-core/src/airflow/serialization/serialized_objects.py @@ -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. @@ -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() @@ -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( diff --git a/airflow-core/tests/unit/serialization/test_dag_serialization.py b/airflow-core/tests/unit/serialization/test_dag_serialization.py index 7852c25dc5ee8..253a98e4c4380 100644 --- a/airflow-core/tests/unit/serialization/test_dag_serialization.py +++ b/airflow-core/tests/unit/serialization/test_dag_serialization.py @@ -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 = { @@ -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"), [ diff --git a/scripts/in_container/run_schema_defaults_check.py b/scripts/in_container/run_schema_defaults_check.py index 7afcbd9ce546f..16686e0f54cef 100755 --- a/scripts/in_container/run_schema_defaults_check.py +++ b/scripts/in_container/run_schema_defaults_check.py @@ -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" ) diff --git a/task-sdk/src/airflow/sdk/definitions/dag.py b/task-sdk/src/airflow/sdk/definitions/dag.py index cadfe6978300e..60d54c213ab20 100644 --- a/task-sdk/src/airflow/sdk/definitions/dag.py +++ b/task-sdk/src/airflow/sdk/definitions/dag.py @@ -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", @@ -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()) + + @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)] diff --git a/task-sdk/tests/task_sdk/definitions/test_dag.py b/task-sdk/tests/task_sdk/definitions/test_dag.py index 9b76816886c76..a89cafff673b3 100644 --- a/task-sdk/tests/task_sdk/definitions/test_dag.py +++ b/task-sdk/tests/task_sdk/definitions/test_dag.py @@ -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 @@ -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"):