From aa2e7f99cd6eead076749df8c20fb97124877f67 Mon Sep 17 00:00:00 2001 From: shaealh Date: Mon, 27 Apr 2026 20:36:51 -0700 Subject: [PATCH] Fix trigger kwargs decoding for asset watcher triggers --- .../src/airflow/serialization/decoders.py | 23 +++- .../tests/unit/jobs/test_triggerer_job.py | 102 +++++++++++++++++- 2 files changed, 122 insertions(+), 3 deletions(-) diff --git a/airflow-core/src/airflow/serialization/decoders.py b/airflow-core/src/airflow/serialization/decoders.py index 890a1296a4a00..78ae241813c24 100644 --- a/airflow-core/src/airflow/serialization/decoders.py +++ b/airflow-core/src/airflow/serialization/decoders.py @@ -90,9 +90,28 @@ def smart_decode_trigger_kwargs(d): """ from airflow.serialization.serialized_objects import BaseSerialization - if not isinstance(d, dict) or Encoding.TYPE not in d: + d = _normalize_stringified_encoding_keys(d) + if not isinstance(d, dict): return d - return BaseSerialization.deserialize(d) + if Encoding.TYPE in d and Encoding.VAR in d: + return BaseSerialization.deserialize(d) + return {k: smart_decode_trigger_kwargs(v) for k, v in d.items()} + + +def _normalize_stringified_encoding_keys(value): + if isinstance(value, list): + return [_normalize_stringified_encoding_keys(v) for v in value] + if not isinstance(value, dict): + return value + + if str(Encoding.TYPE) in value and str(Encoding.VAR) in value: + return { + Encoding.TYPE if k == str(Encoding.TYPE) else Encoding.VAR if k == str(Encoding.VAR) else k: ( + _normalize_stringified_encoding_keys(v) + ) + for k, v in value.items() + } + return {k: _normalize_stringified_encoding_keys(v) for k, v in value.items()} def _decode_asset(var: dict[str, Any]): diff --git a/airflow-core/tests/unit/jobs/test_triggerer_job.py b/airflow-core/tests/unit/jobs/test_triggerer_job.py index fa4d68f5fff6a..fa3e8b861498b 100644 --- a/airflow-core/tests/unit/jobs/test_triggerer_job.py +++ b/airflow-core/tests/unit/jobs/test_triggerer_job.py @@ -30,7 +30,7 @@ import uuid from collections.abc import AsyncIterator from socket import socket, socketpair -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING, Any, ClassVar from unittest import mock from unittest.mock import ANY, AsyncMock, MagicMock, patch @@ -124,6 +124,33 @@ pytestmark = pytest.mark.db_test +class SerializedKwargsTrigger(BaseTrigger): + constructed: ClassVar[list[SerializedKwargsTrigger]] = [] + + def __init__( + self, + *, + apply_function_args: tuple[Any, ...], + apply_function_kwargs: dict[str, Any], + ) -> None: + super().__init__() + self.apply_function_args = apply_function_args + self.apply_function_kwargs = apply_function_kwargs + self.constructed.append(self) + + def serialize(self) -> tuple[str, dict[str, Any]]: + return ( + "tests.unit.jobs.test_triggerer_job.SerializedKwargsTrigger", + { + "apply_function_args": self.apply_function_args, + "apply_function_kwargs": self.apply_function_kwargs, + }, + ) + + async def run(self) -> AsyncIterator[TriggerEvent]: + yield TriggerEvent(True) + + @pytest.fixture(autouse=True) def _force_bare_fork(monkeypatch): """ @@ -1647,6 +1674,79 @@ async def test_create_triggers_emits_creation_duration_metric( await runner.cleanup_finished_triggers() + @pytest.mark.asyncio + @patch( + "airflow.jobs.triggerer_job_runner.TriggerRunner.get_trigger_by_classpath", + return_value=SerializedKwargsTrigger, + ) + async def test_trigger_kwargs_cleanup_decodes_stringified_encoding_keys( + self, mock_get_trigger_by_classpath, session + ): + import json + + from airflow.models.crypto import get_fernet + from airflow.sdk.serde import serialize + + SerializedKwargsTrigger.constructed.clear() + trigger_orm = Trigger( + classpath="tests.unit.jobs.test_triggerer_job.SerializedKwargsTrigger", + kwargs={}, + ) + legacy_stringified_kwargs = { + "apply_function_args": {"Encoding.TYPE": "tuple", "Encoding.VAR": []}, + "apply_function_kwargs": { + "Encoding.TYPE": "dict", + "Encoding.VAR": { + "action_filter": ["create"], + "data_filter": { + "Encoding.TYPE": "dict", + "Encoding.VAR": {"BatchId": None}, + }, + }, + }, + } + trigger_orm.encrypted_kwargs = ( + get_fernet().encrypt(json.dumps(serialize(legacy_stringified_kwargs)).encode()).decode() + ) + session.add(trigger_orm) + session.commit() + + stored_kwargs = trigger_orm.kwargs + assert stored_kwargs == { + "apply_function_args": {"Encoding.TYPE": "tuple", "Encoding.VAR": []}, + "apply_function_kwargs": { + "Encoding.TYPE": "dict", + "Encoding.VAR": { + "action_filter": ["create"], + "data_filter": { + "Encoding.TYPE": "dict", + "Encoding.VAR": {"BatchId": None}, + }, + }, + }, + } + + runner = TriggerRunner() + runner.to_create.append( + workloads.RunTrigger.model_construct( + id=trigger_orm.id, + ti=None, + classpath=trigger_orm.classpath, + encrypted_kwargs=trigger_orm.encrypted_kwargs, + ) + ) + + await runner.create_triggers() + + trigger_instance = SerializedKwargsTrigger.constructed[-1] + assert trigger_instance.apply_function_args == () + assert trigger_instance.apply_function_kwargs == { + "action_filter": ["create"], + "data_filter": {"BatchId": None}, + } + runner.triggers[trigger_orm.id]["task"].cancel() + await runner.cleanup_finished_triggers() + @pytest.mark.asyncio @patch("airflow.sdk.execution_time.task_runner.SUPERVISOR_COMMS", create=True) async def test_sync_state_to_supervisor(self, supervisor_builder):