Skip to content
Open
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
23 changes: 21 additions & 2 deletions airflow-core/src/airflow/serialization/decoders.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]):
Expand Down
102 changes: 101 additions & 1 deletion airflow-core/tests/unit/jobs/test_triggerer_job.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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):
"""
Expand Down Expand Up @@ -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):
Expand Down