diff --git a/airflow-core/src/airflow/callbacks/callback_requests.py b/airflow-core/src/airflow/callbacks/callback_requests.py index 4d2ed18b36dd5..a4692df2356fb 100644 --- a/airflow-core/src/airflow/callbacks/callback_requests.py +++ b/airflow-core/src/airflow/callbacks/callback_requests.py @@ -16,7 +16,7 @@ # under the License. from __future__ import annotations -from typing import TYPE_CHECKING, Annotated, Literal +from typing import TYPE_CHECKING, Annotated, Any, Literal from pydantic import BaseModel, Field @@ -84,6 +84,8 @@ class DagCallbackRequest(BaseCallbackRequest): run_id: str is_failure_callback: bool | None = True """Flag to determine whether it is a Failure Callback or Success Callback""" + dag_run: dict[str, Any] | None = None + """Serialized dag_run information to be included in the callback context""" type: Literal["DagCallbackRequest"] = "DagCallbackRequest" diff --git a/airflow-core/src/airflow/dag_processing/processor.py b/airflow-core/src/airflow/dag_processing/processor.py index dd41c6c8057b3..3c3144ba051e5 100644 --- a/airflow-core/src/airflow/dag_processing/processor.py +++ b/airflow-core/src/airflow/dag_processing/processor.py @@ -217,13 +217,16 @@ def _execute_dag_callbacks(dagbag: DagBag, request: DagCallbackRequest, log: Fil return callbacks = callbacks if isinstance(callbacks, list) else [callbacks] - # TODO:We need a proper context object! context: Context = { "dag": dag, "run_id": request.run_id, "reason": request.msg, } + # Only add dag_run to context if it's provided + if request.dag_run is not None: + context["dag_run"] = request.dag_run + for callback in callbacks: log.info( "Executing on_%s dag callback", diff --git a/airflow-core/src/airflow/jobs/scheduler_job_runner.py b/airflow-core/src/airflow/jobs/scheduler_job_runner.py index 53f91991e89d3..319fbfa72dbe8 100644 --- a/airflow-core/src/airflow/jobs/scheduler_job_runner.py +++ b/airflow-core/src/airflow/jobs/scheduler_job_runner.py @@ -1856,6 +1856,7 @@ def _schedule_dag_run( bundle_version=dag_run.bundle_version, is_failure_callback=True, msg="timed_out", + dag_run=dag_run.serialize_for_callback(), ) dag_run.notify_dagrun_state_changed() diff --git a/airflow-core/src/airflow/models/dagrun.py b/airflow-core/src/airflow/models/dagrun.py index b910e00f9b5b0..bc537e97281ae 100644 --- a/airflow-core/src/airflow/models/dagrun.py +++ b/airflow-core/src/airflow/models/dagrun.py @@ -1188,6 +1188,7 @@ def recalculate(self) -> _UnfinishedStates: bundle_version=self.bundle_version, is_failure_callback=True, msg="task_failure", + dag_run=self.serialize_for_callback(), ) # Check if the max_consecutive_failed_dag_runs has been provided and not 0 @@ -1217,6 +1218,7 @@ def recalculate(self) -> _UnfinishedStates: bundle_version=self.bundle_version, is_failure_callback=False, msg="success", + dag_run=self.serialize_for_callback(), ) if (deadline := dag.deadline) and isinstance(deadline.reference, DeadlineReference.TYPES.DAGRUN): @@ -1240,6 +1242,7 @@ def recalculate(self) -> _UnfinishedStates: bundle_version=self.bundle_version, is_failure_callback=True, msg="all_tasks_deadlocked", + dag_run=self.serialize_for_callback(), ) # finally, if the leaves aren't done, the dag is still running @@ -1356,6 +1359,7 @@ def handle_dag_callback(self, dag: SDKDAG, success: bool = True, reason: str = " "dag": dag, "run_id": str(self.run_id), "reason": reason, + "dag_run": self, } callbacks = dag.on_success_callback if success else dag.on_failure_callback @@ -2013,6 +2017,29 @@ def _get_log_template(log_template_id: int | None, session: Session = NEW_SESSIO def _get_partial_task_ids(dag: DAG | None) -> list[str] | None: return dag.task_ids if dag and dag.partial else None + def serialize_for_callback(self) -> dict[str, Any]: + """ + Serialize DagRun object into a dictionary for callback requests. + + This method creates a serialized representation of the DagRun that can be + safely passed to subprocesses without requiring database access. + + :return: Dictionary containing serialized DagRun information + """ + return { + "dag_id": self.dag_id, + "run_id": self.run_id, + "state": self.state, + "logical_date": self.logical_date.isoformat() if self.logical_date else None, + "start_date": self.start_date.isoformat() if self.start_date else None, + "end_date": self.end_date.isoformat() if self.end_date else None, + "conf": self.conf, + "run_type": self.run_type, + "run_after": self.run_after.isoformat() if self.run_after else None, + "data_interval_start": self.data_interval_start.isoformat() if self.data_interval_start else None, + "data_interval_end": self.data_interval_end.isoformat() if self.data_interval_end else None, + } + class DagRunNote(Base): """For storage of arbitrary notes concerning the dagrun instance.""" diff --git a/airflow-core/tests/unit/callbacks/test_callback_requests.py b/airflow-core/tests/unit/callbacks/test_callback_requests.py index 08c1a9a666d43..0ef6d0d7d9801 100644 --- a/airflow-core/tests/unit/callbacks/test_callback_requests.py +++ b/airflow-core/tests/unit/callbacks/test_callback_requests.py @@ -114,3 +114,269 @@ def test_is_failure_callback_property( ) assert request.is_failure_callback == expected_is_failure + + +class TestDagCallbackRequest: + """Test the DagCallbackRequest class with the new dag_run field.""" + + def test_dag_callback_request_with_dag_run(self): + """Test DagCallbackRequest creation with dag_run field.""" + dag_run_data = { + "dag_id": "test_dag", + "run_id": "test_run_2024-01-01T00:00:00+00:00", + "state": "success", + "logical_date": "2024-01-01T00:00:00+00:00", + "start_date": "2024-01-01T00:00:00+00:00", + "end_date": "2024-01-01T01:00:00+00:00", + "run_type": "manual", + "run_after": "2024-01-01T00:00:00+00:00", + "conf": {"key": "value"}, + "data_interval_start": "2024-01-01T00:00:00+00:00", + "data_interval_end": "2024-01-01T01:00:00+00:00", + } + + request = DagCallbackRequest( + filepath="test_dag.py", + dag_id="test_dag", + run_id="test_run_2024-01-01T00:00:00+00:00", + is_failure_callback=False, + bundle_name="testing", + bundle_version=None, + msg="success", + dag_run=dag_run_data, + ) + + assert request.dag_run == dag_run_data + assert request.dag_run["dag_id"] == "test_dag" + assert request.dag_run["run_id"] == "test_run_2024-01-01T00:00:00+00:00" + assert request.dag_run["state"] == "success" + assert request.dag_run["conf"]["key"] == "value" + + def test_dag_callback_request_without_dag_run(self): + """Test DagCallbackRequest creation without dag_run field (backward compatibility).""" + request = DagCallbackRequest( + filepath="test_dag.py", + dag_id="test_dag", + run_id="test_run", + is_failure_callback=False, + bundle_name="testing", + bundle_version=None, + msg="success", + ) + + assert request.dag_run is None + assert request.dag_id == "test_dag" + assert request.run_id == "test_run" + + def test_dag_callback_request_serialization_with_dag_run(self): + """Test DagCallbackRequest serialization and deserialization with dag_run field.""" + dag_run_data = { + "dag_id": "test_dag", + "run_id": "test_run_2024-01-01T00:00:00+00:00", + "state": "success", + "logical_date": "2024-01-01T00:00:00+00:00", + "start_date": "2024-01-01T00:00:00+00:00", + "end_date": "2024-01-01T01:00:00+00:00", + "run_type": "manual", + "run_after": "2024-01-01T00:00:00+00:00", + "conf": {"key": "value", "nested": {"inner": "data"}}, + "data_interval_start": "2024-01-01T00:00:00+00:00", + "data_interval_end": "2024-01-01T01:00:00+00:00", + } + + original_request = DagCallbackRequest( + filepath="test_dag.py", + dag_id="test_dag", + run_id="test_run_2024-01-01T00:00:00+00:00", + is_failure_callback=False, + bundle_name="testing", + bundle_version=None, + msg="success", + dag_run=dag_run_data, + ) + + # Serialize to JSON + json_str = original_request.to_json() + + # Deserialize from JSON + deserialized_request = DagCallbackRequest.from_json(json_str) + + # Verify all fields are preserved + assert deserialized_request == original_request + assert deserialized_request.dag_run == dag_run_data + assert deserialized_request.dag_run["conf"]["nested"]["inner"] == "data" + + def test_dag_callback_request_serialization_without_dag_run(self): + """Test DagCallbackRequest serialization and deserialization without dag_run field.""" + original_request = DagCallbackRequest( + filepath="test_dag.py", + dag_id="test_dag", + run_id="test_run", + is_failure_callback=True, + bundle_name="testing", + bundle_version=None, + msg="task_failure", + ) + + # Serialize to JSON + json_str = original_request.to_json() + + # Deserialize from JSON + deserialized_request = DagCallbackRequest.from_json(json_str) + + # Verify all fields are preserved + assert deserialized_request == original_request + assert deserialized_request.dag_run is None + + def test_dag_callback_request_with_none_dag_run(self): + """Test DagCallbackRequest with explicitly None dag_run field.""" + request = DagCallbackRequest( + filepath="test_dag.py", + dag_id="test_dag", + run_id="test_run", + is_failure_callback=False, + bundle_name="testing", + bundle_version=None, + msg="success", + dag_run=None, + ) + + assert request.dag_run is None + + def test_dag_callback_request_with_minimal_dag_run_data(self): + """Test DagCallbackRequest with minimal dag_run data.""" + minimal_dag_run_data = { + "dag_id": "test_dag", + "run_id": "test_run", + "state": "success", + } + + request = DagCallbackRequest( + filepath="test_dag.py", + dag_id="test_dag", + run_id="test_run", + is_failure_callback=False, + bundle_name="testing", + bundle_version=None, + msg="success", + dag_run=minimal_dag_run_data, + ) + + assert request.dag_run == minimal_dag_run_data + assert request.dag_run["dag_id"] == "test_dag" + assert request.dag_run["run_id"] == "test_run" + assert request.dag_run["state"] == "success" + + def test_dag_callback_request_with_null_values_in_dag_run(self): + """Test DagCallbackRequest with null values in dag_run data.""" + dag_run_data_with_nulls = { + "dag_id": "test_dag", + "run_id": "test_run", + "state": "success", + "logical_date": None, + "start_date": "2024-01-01T00:00:00+00:00", + "end_date": None, + "run_type": "manual", + "run_after": "2024-01-01T00:00:00+00:00", + "conf": None, + "data_interval_start": None, + "data_interval_end": None, + } + + request = DagCallbackRequest( + filepath="test_dag.py", + dag_id="test_dag", + run_id="test_run", + is_failure_callback=False, + bundle_name="testing", + bundle_version=None, + msg="success", + dag_run=dag_run_data_with_nulls, + ) + + assert request.dag_run == dag_run_data_with_nulls + assert request.dag_run["logical_date"] is None + assert request.dag_run["end_date"] is None + assert request.dag_run["conf"] is None + + def test_dag_callback_request_equality_with_dag_run(self): + """Test DagCallbackRequest equality comparison with dag_run field.""" + dag_run_data = { + "dag_id": "test_dag", + "run_id": "test_run", + "state": "success", + "conf": {"key": "value"}, + } + + request1 = DagCallbackRequest( + filepath="test_dag.py", + dag_id="test_dag", + run_id="test_run", + is_failure_callback=False, + bundle_name="testing", + bundle_version=None, + msg="success", + dag_run=dag_run_data, + ) + + request2 = DagCallbackRequest( + filepath="test_dag.py", + dag_id="test_dag", + run_id="test_run", + is_failure_callback=False, + bundle_name="testing", + bundle_version=None, + msg="success", + dag_run=dag_run_data, + ) + + request3 = DagCallbackRequest( + filepath="test_dag.py", + dag_id="test_dag", + run_id="test_run", + is_failure_callback=False, + bundle_name="testing", + bundle_version=None, + msg="success", + dag_run={"dag_id": "different_dag", "run_id": "test_run", "state": "success"}, + ) + + assert request1 == request2 + assert request1 != request3 + assert request2 != request3 + + def test_dag_callback_request_equality_without_dag_run(self): + """Test DagCallbackRequest equality comparison without dag_run field.""" + request1 = DagCallbackRequest( + filepath="test_dag.py", + dag_id="test_dag", + run_id="test_run", + is_failure_callback=False, + bundle_name="testing", + bundle_version=None, + msg="success", + ) + + request2 = DagCallbackRequest( + filepath="test_dag.py", + dag_id="test_dag", + run_id="test_run", + is_failure_callback=False, + bundle_name="testing", + bundle_version=None, + msg="success", + ) + + request3 = DagCallbackRequest( + filepath="test_dag.py", + dag_id="test_dag", + run_id="test_run", + is_failure_callback=False, + bundle_name="testing", + bundle_version=None, + msg="success", + dag_run={"dag_id": "test_dag", "run_id": "test_run"}, + ) + + assert request1 == request2 + assert request1 != request3 # Different because one has dag_run and the other doesn't diff --git a/airflow-core/tests/unit/dag_processing/test_processor.py b/airflow-core/tests/unit/dag_processing/test_processor.py index 6c6f31a1ff8d0..8dbbc040b8c37 100644 --- a/airflow-core/tests/unit/dag_processing/test_processor.py +++ b/airflow-core/tests/unit/dag_processing/test_processor.py @@ -43,6 +43,7 @@ DagFileParseRequest, DagFileParsingResult, DagFileProcessorProcess, + _execute_dag_callbacks, _execute_task_callbacks, _parse_file, _pre_import_airflow_modules, @@ -819,3 +820,214 @@ def fake_collect_dags(self, *args, **kwargs): _execute_task_callbacks(dagbag, request, log) assert call_count == 2 + + def test_execute_dag_callbacks_with_dag_run_data(self, spy_agency): + """Test _execute_dag_callbacks with dag_run data in the request.""" + called = False + context_received = None + + def on_success(context): + nonlocal called, context_received + called = True + context_received = context + + with DAG(dag_id="test_dag", on_success_callback=on_success) as dag: + BaseOperator(task_id="test_task") + + def fake_collect_dags(self, *args, **kwargs): + self.dags[dag.dag_id] = dag + + spy_agency.spy_on(DagBag.collect_dags, call_fake=fake_collect_dags, owner=DagBag) + + dagbag = DagBag() + dagbag.collect_dags() + + # Create serialized dag_run data + dag_run_data = { + "dag_id": "test_dag", + "run_id": "test_run_2024-01-01T00:00:00+00:00", + "state": "success", + "logical_date": "2024-01-01T00:00:00+00:00", + "start_date": "2024-01-01T00:00:00+00:00", + "end_date": "2024-01-01T01:00:00+00:00", + "run_type": "manual", + "run_after": "2024-01-01T00:00:00+00:00", + "conf": {"key": "value"}, + "data_interval_start": "2024-01-01T00:00:00+00:00", + "data_interval_end": "2024-01-02T00:00:00+00:00", + } + + request = DagCallbackRequest( + filepath="test.py", + msg="DAG succeeded", + dag_id="test_dag", + run_id="test_run_2024-01-01T00:00:00+00:00", + bundle_name="testing", + bundle_version=None, + is_failure_callback=False, + dag_run=dag_run_data, + ) + + log = structlog.get_logger() + _execute_dag_callbacks(dagbag, request, log) + + assert called is True + assert context_received is not None + assert context_received["dag"] == dag + assert context_received["run_id"] == "test_run_2024-01-01T00:00:00+00:00" + assert context_received["reason"] == "DAG succeeded" + + # Verify dag_run data is present and accessible + assert "dag_run" in context_received + dag_run_context = context_received["dag_run"] + assert dag_run_context["dag_id"] == "test_dag" + assert dag_run_context["run_id"] == "test_run_2024-01-01T00:00:00+00:00" + assert dag_run_context["state"] == "success" + assert dag_run_context["logical_date"] == "2024-01-01T00:00:00+00:00" + assert dag_run_context["conf"]["key"] == "value" + + def test_execute_dag_callbacks_without_dag_run_data(self, spy_agency): + """Test _execute_dag_callbacks without dag_run data (backward compatibility).""" + called = False + context_received = None + + def on_failure(context): + nonlocal called, context_received + called = True + context_received = context + + with DAG(dag_id="test_dag", on_failure_callback=on_failure) as dag: + BaseOperator(task_id="test_task") + + def fake_collect_dags(self, *args, **kwargs): + self.dags[dag.dag_id] = dag + + spy_agency.spy_on(DagBag.collect_dags, call_fake=fake_collect_dags, owner=DagBag) + + dagbag = DagBag() + dagbag.collect_dags() + + request = DagCallbackRequest( + filepath="test.py", + msg="DAG failed", + dag_id="test_dag", + run_id="test_run", + bundle_name="testing", + bundle_version=None, + is_failure_callback=True, + # No dag_run field + ) + + log = structlog.get_logger() + _execute_dag_callbacks(dagbag, request, log) + + assert called is True + assert context_received is not None + assert context_received["dag"] == dag + assert context_received["run_id"] == "test_run" + assert context_received["reason"] == "DAG failed" + + # Verify dag_run is not present when not provided + assert "dag_run" not in context_received + + def test_execute_dag_callbacks_with_none_dag_run_data(self, spy_agency): + """Test _execute_dag_callbacks with explicitly None dag_run data.""" + called = False + context_received = None + + def on_success(context): + nonlocal called, context_received + called = True + context_received = context + + with DAG(dag_id="test_dag", on_success_callback=on_success) as dag: + BaseOperator(task_id="test_task") + + def fake_collect_dags(self, *args, **kwargs): + self.dags[dag.dag_id] = dag + + spy_agency.spy_on(DagBag.collect_dags, call_fake=fake_collect_dags, owner=DagBag) + + dagbag = DagBag() + dagbag.collect_dags() + + request = DagCallbackRequest( + filepath="test.py", + msg="DAG succeeded", + dag_id="test_dag", + run_id="test_run", + bundle_name="testing", + bundle_version=None, + is_failure_callback=False, + dag_run=None, # Explicitly None + ) + + log = structlog.get_logger() + _execute_dag_callbacks(dagbag, request, log) + + assert called is True + assert context_received is not None + assert context_received["dag"] == dag + assert context_received["run_id"] == "test_run" + assert context_received["reason"] == "DAG succeeded" + + # Verify dag_run is not present when explicitly None + assert "dag_run" not in context_received + + def test_execute_dag_callbacks_with_minimal_dag_run_data(self, spy_agency): + """Test _execute_dag_callbacks with minimal dag_run data.""" + called = False + context_received = None + + def on_success(context): + nonlocal called, context_received + called = True + context_received = context + + with DAG(dag_id="test_dag", on_success_callback=on_success) as dag: + BaseOperator(task_id="test_task") + + def fake_collect_dags(self, *args, **kwargs): + self.dags[dag.dag_id] = dag + + spy_agency.spy_on(DagBag.collect_dags, call_fake=fake_collect_dags, owner=DagBag) + + dagbag = DagBag() + dagbag.collect_dags() + + # Minimal dag_run data + minimal_dag_run_data = { + "dag_id": "test_dag", + "run_id": "test_run", + "state": "success", + } + + request = DagCallbackRequest( + filepath="test.py", + msg="DAG succeeded", + dag_id="test_dag", + run_id="test_run", + bundle_name="testing", + bundle_version=None, + is_failure_callback=False, + dag_run=minimal_dag_run_data, + ) + + log = structlog.get_logger() + _execute_dag_callbacks(dagbag, request, log) + + assert called is True + assert context_received is not None + assert context_received["dag"] == dag + assert context_received["run_id"] == "test_run" + assert context_received["reason"] == "DAG succeeded" + + # Verify minimal dag_run data is present + assert "dag_run" in context_received + dag_run_context = context_received["dag_run"] + assert dag_run_context["dag_id"] == "test_dag" + assert dag_run_context["run_id"] == "test_run" + assert dag_run_context["state"] == "success" + # Verify optional fields are not present + assert "logical_date" not in dag_run_context + assert "conf" not in dag_run_context diff --git a/airflow-core/tests/unit/models/test_dagrun.py b/airflow-core/tests/unit/models/test_dagrun.py index cf4dafb4ee094..93663d5912d7c 100644 --- a/airflow-core/tests/unit/models/test_dagrun.py +++ b/airflow-core/tests/unit/models/test_dagrun.py @@ -2813,3 +2813,256 @@ def my_teardown(): "tg_2.my_teardown": "skipped", "tg_2.my_work": "skipped", } + + def test_dagrun_callback_context_has_dag_run(self, dag_maker, session): + """Test that DAG callbacks have access to dag_run in their context.""" + + # Track what we receive in the callback context + received_context = {} + + def on_success_callable(context): + nonlocal received_context + received_context = context.copy() + # Verify dag_run is present and accessible + assert "dag_run" in context, "dag_run should be present in DAG callback context" + assert context["dag_run"]["dag_id"] == "test_dagrun_callback_context_has_dag_run" + assert context["dag_run"]["run_id"] == context["run_id"] + # Test accessing dag_run properties + assert context["dag_run"]["state"] == "success" + assert context["dag_run"]["logical_date"] is not None + + with dag_maker( + dag_id="test_dagrun_callback_context_has_dag_run", + on_success_callback=on_success_callable, + ) as dag: + EmptyOperator(task_id="task1") + + dag_run = dag_maker.create_dagrun() + + # Execute the callback + dag_run.handle_dag_callback(dag, success=True, reason="test") + + # Verify the context contains the expected fields + assert received_context["dag"] == dag + assert received_context["run_id"] == str(dag_run.run_id) + assert received_context["reason"] == "test" + assert received_context["dag_run"]["dag_id"] == dag.dag_id + assert received_context["dag_run"]["run_id"] == str(dag_run.run_id) + + def test_dagrun_callback_context_missing_dag_run_bug(self, dag_maker, session): + """Test that demonstrates the current Airflow 3.0 behavior where dag_run is available in DAG callbacks.""" + + # This test documents the current behavior in Airflow 3.0 + received_context = {} + + def on_success_callable(context): + nonlocal received_context + received_context = context.copy() + # In Airflow 3.0, dag_run is available in DAG callbacks as a dictionary + assert "dag" in context, "dag should be present in DAG callback context" + assert "run_id" in context, "run_id should be present in DAG callback context" + assert "reason" in context, "reason should be present in DAG callback context" + assert "dag_run" in context, "dag_run should be present in DAG callback context" + # dag_run is now available as a dictionary with serialized information + + with dag_maker( + dag_id="test_dagrun_callback_context_missing_dag_run_bug", + on_success_callback=on_success_callable, + ) as dag: + EmptyOperator(task_id="task1") + + dag_run = dag_maker.create_dagrun() + + # Execute the callback + dag_run.handle_dag_callback(dag, success=True, reason="test") + + # Verify the context contains the expected fields + assert received_context["dag"] == dag + assert received_context["run_id"] == str(dag_run.run_id) + assert received_context["reason"] == "test" + assert received_context["dag_run"]["dag_id"] == dag.dag_id + assert received_context["dag_run"]["run_id"] == str(dag_run.run_id) + + def test_serialize_for_callback_with_all_fields(self, dag_maker, session): + """Test serialize_for_callback method with all fields populated.""" + with dag_maker( + dag_id="test_serialize_for_callback", + schedule="@daily", + ): + EmptyOperator(task_id="task1") + + # Create a DagRun with all fields populated + logical_date = timezone.datetime(2024, 1, 1, tzinfo=timezone.utc) + start_date = timezone.datetime(2024, 1, 1, 1, 0, 0, tzinfo=timezone.utc) + end_date = timezone.datetime(2024, 1, 1, 2, 0, 0, tzinfo=timezone.utc) + run_after = timezone.datetime(2024, 1, 1, 0, 0, 0, tzinfo=timezone.utc) + data_interval_start = timezone.datetime(2024, 1, 1, 0, 0, 0, tzinfo=timezone.utc) + data_interval_end = timezone.datetime(2024, 1, 2, 0, 0, 0, tzinfo=timezone.utc) + + dag_run = dag_maker.create_dagrun( + logical_date=logical_date, + start_date=start_date, + end_date=end_date, + run_after=run_after, + data_interval_start=data_interval_start, + data_interval_end=data_interval_end, + conf={"key": "value", "nested": {"inner": "data"}}, + run_type="manual", + ) + + # Serialize the DagRun + serialized = dag_run.serialize_for_callback() + + # Verify all fields are correctly serialized + assert serialized["dag_id"] == "test_serialize_for_callback" + assert serialized["run_id"] == str(dag_run.run_id) + assert serialized["state"] == dag_run.state + assert serialized["logical_date"] == "2024-01-01T00:00:00+00:00" + assert serialized["start_date"] == "2024-01-01T01:00:00+00:00" + assert serialized["end_date"] == "2024-01-01T02:00:00+00:00" + assert serialized["run_type"] == "manual" + assert serialized["run_after"] == "2024-01-01T00:00:00+00:00" + assert serialized["conf"] == {"key": "value", "nested": {"inner": "data"}} + assert serialized["data_interval_start"] == "2024-01-01T00:00:00+00:00" + assert serialized["data_interval_end"] == "2024-01-02T00:00:00+00:00" + + def test_serialize_for_callback_with_none_fields(self, dag_maker, session): + """Test serialize_for_callback method with None fields.""" + with dag_maker( + dag_id="test_serialize_for_callback_none", + schedule=None, + ): + EmptyOperator(task_id="task1") + + # Create a DagRun with some None fields + dag_run = dag_maker.create_dagrun( + logical_date=None, + start_date=None, + end_date=None, + run_after=None, + data_interval_start=None, + data_interval_end=None, + conf=None, + run_type="manual", + ) + + # Serialize the DagRun + serialized = dag_run.serialize_for_callback() + + # Verify None fields are handled correctly + assert serialized["dag_id"] == "test_serialize_for_callback_none" + assert serialized["run_id"] == str(dag_run.run_id) + assert serialized["state"] == dag_run.state + assert serialized["logical_date"] is None + assert serialized["start_date"] is None + assert serialized["end_date"] is None + assert serialized["run_type"] == "manual" + assert serialized["run_after"] is None + assert serialized["conf"] is None + assert serialized["data_interval_start"] is None + assert serialized["data_interval_end"] is None + + def test_serialize_for_callback_with_minimal_fields(self, dag_maker, session): + """Test serialize_for_callback method with minimal required fields.""" + with dag_maker( + dag_id="test_serialize_for_callback_minimal", + ): + EmptyOperator(task_id="task1") + + # Create a DagRun with minimal fields + dag_run = dag_maker.create_dagrun() + + # Serialize the DagRun + serialized = dag_run.serialize_for_callback() + + # Verify minimal fields are present + assert serialized["dag_id"] == "test_serialize_for_callback_minimal" + assert serialized["run_id"] == str(dag_run.run_id) + assert serialized["state"] == dag_run.state + assert "logical_date" in serialized + assert "start_date" in serialized + assert "end_date" in serialized + assert "run_type" in serialized + assert "run_after" in serialized + assert "conf" in serialized + assert "data_interval_start" in serialized + assert "data_interval_end" in serialized + + def test_serialize_for_callback_creates_immutable_copy(self, dag_maker, session): + """Test that serialize_for_callback creates an immutable copy of the data.""" + with dag_maker( + dag_id="test_serialize_for_callback_immutable", + ): + EmptyOperator(task_id="task1") + + dag_run = dag_maker.create_dagrun( + conf={"mutable": "value"}, + ) + + # Serialize the DagRun + serialized = dag_run.serialize_for_callback() + + # Modify the original DagRun + dag_run.conf["mutable"] = "modified" + dag_run.state = "failed" + + # Verify the serialized data is unchanged + assert serialized["conf"]["mutable"] == "value" + assert serialized["state"] != "failed" + + def test_serialize_for_callback_with_complex_conf(self, dag_maker, session): + """Test serialize_for_callback method with complex configuration.""" + complex_conf = { + "string": "value", + "number": 42, + "boolean": True, + "list": [1, 2, 3], + "dict": {"nested": "value", "deep": {"level": 3}}, + "null": None, + } + + with dag_maker( + dag_id="test_serialize_for_callback_complex", + ): + EmptyOperator(task_id="task1") + + dag_run = dag_maker.create_dagrun(conf=complex_conf) + + # Serialize the DagRun + serialized = dag_run.serialize_for_callback() + + # Verify complex configuration is preserved + assert serialized["conf"] == complex_conf + assert serialized["conf"]["string"] == "value" + assert serialized["conf"]["number"] == 42 + assert serialized["conf"]["boolean"] is True + assert serialized["conf"]["list"] == [1, 2, 3] + assert serialized["conf"]["dict"]["deep"]["level"] == 3 + assert serialized["conf"]["null"] is None + + def test_serialize_for_callback_round_trip_json(self, dag_maker, session): + """Test that serialized data can be round-tripped through JSON.""" + import json + + with dag_maker( + dag_id="test_serialize_for_callback_json", + ): + EmptyOperator(task_id="task1") + + dag_run = dag_maker.create_dagrun( + conf={"key": "value"}, + run_type="scheduled", + ) + + # Serialize the DagRun + serialized = dag_run.serialize_for_callback() + + # Convert to JSON and back + json_str = json.dumps(serialized) + deserialized = json.loads(json_str) + + # Verify the data is preserved + assert deserialized == serialized + assert deserialized["dag_id"] == "test_serialize_for_callback_json" + assert deserialized["conf"]["key"] == "value" + assert deserialized["run_type"] == "scheduled" diff --git a/task-sdk/src/airflow/sdk/definitions/context.py b/task-sdk/src/airflow/sdk/definitions/context.py index 082ad36202ec2..2a2e43a6d0f73 100644 --- a/task-sdk/src/airflow/sdk/definitions/context.py +++ b/task-sdk/src/airflow/sdk/definitions/context.py @@ -39,7 +39,7 @@ class Context(TypedDict, total=False): conn: Any dag: DAG - dag_run: DagRunProtocol + dag_run: DagRunProtocol | dict[str, Any] data_interval_end: DateTime | None data_interval_start: DateTime | None outlet_events: OutletEventAccessorsProtocol