This is an automated email from the ASF dual-hosted git repository.
dabla pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/airflow.git
The following commit(s) were added to refs/heads/main by this push:
new 67853607ab2 Add async Variable methods (aget, aset, adelete, akeys)
(#72329)
67853607ab2 is described below
commit 67853607ab217962e82d66c4e0ce0b1b6148f9b6
Author: David Blain <[email protected]>
AuthorDate: Mon Oct 5 19:42:46 2026 +0200
Add async Variable methods (aget, aset, adelete, akeys) (#72329)
* Add async Variable methods (aget, aset, adelete, akeys) with tests
Fix Variable.akeys: wrapping an async call inside lazy_object_proxy.Proxy
is invalid because await cannot be used inside a plain lambda. The method
now awaits _async_get_variable_keys directly.
Co-authored-by: Copilot <[email protected]>
* Check if secret backend has aget_variable method available, it not the
run the sync one via asyncio.to_thread
* Simplify async Variable tests to use mock_supervisor_comms.asend directly
Avoid importing and reassigning AsyncMock in every test; the fixture's
asend is already an AsyncMock, so tests can set return_value/side_effect
on it directly, per top-level import convention.
Co-authored-by: Copilot <[email protected]>
* Simplify async Variable tests in test_context.py to avoid inline
AsyncMock imports
Same cleanup as test_variables.py: set return_value/side_effect directly
on the fixture's mock_supervisor_comms.asend instead of importing and
reassigning AsyncMock per test; move json import to the top of the file.
Co-authored-by: Copilot <[email protected]>
* Drop redundant imports in async Variable code and tests
json is already imported at module level in context.py, and
_VARIABLE_KEYS_PAGE_SIZE is now imported once at the top of each test
module.
Co-Authored-By: Claude Opus 5.5 <[email protected]>
* Test masking of async Variable reads, including cache hits
Mirror the sync VariableAccessor masking tests for _async_get_variable so
_async_mask_and_deserialize_variable and the SecretCache hit path are
covered.
Co-Authored-By: Claude Opus 5.5 <[email protected]>
* Test the default fallback of Variable.aget
Co-Authored-By: Claude Opus 5.5 <[email protected]>
* Test the write-conflict warning of async Variable.set
Co-Authored-By: Claude Opus 5.5 <[email protected]>
* Assert the GetVariable request in the async Variable.get test
Co-Authored-By: Claude Fable 5.1 <[email protected]>
* Test that async Variable reads do not fall through after a deny
_async_get_variable re-raises AirflowSecretsBackendAccessDenied instead of
asking the next secrets backend. The sync twin is covered in
test_secrets.py;
this pins the async one.
Co-Authored-By: Claude Fable 5.1 <[email protected]>
---------
Co-authored-by: Copilot <[email protected]>
Co-authored-by: Claude Opus 5.5 <[email protected]>
---
task-sdk/src/airflow/sdk/definitions/variable.py | 55 +++-
task-sdk/src/airflow/sdk/execution_time/context.py | 155 +++++++++
.../tests/task_sdk/definitions/test_variables.py | 139 +++++++-
.../tests/task_sdk/execution_time/test_context.py | 350 ++++++++++++++++++++-
4 files changed, 695 insertions(+), 4 deletions(-)
diff --git a/task-sdk/src/airflow/sdk/definitions/variable.py
b/task-sdk/src/airflow/sdk/definitions/variable.py
index f1db76c3fae..a660debf81c 100644
--- a/task-sdk/src/airflow/sdk/definitions/variable.py
+++ b/task-sdk/src/airflow/sdk/definitions/variable.py
@@ -24,7 +24,7 @@ from typing import Any
import attrs
from airflow.sdk.definitions._internal.types import NOTSET
-from airflow.sdk.log import mask_secret
+from airflow.sdk.log import amask_secret, mask_secret
log = logging.getLogger(__name__)
@@ -58,12 +58,37 @@ class Variable:
return default
raise
+ @classmethod
+ async def aget(cls, key: str, default: Any = NOTSET, deserialize_json:
bool = False):
+ from airflow.sdk.exceptions import AirflowRuntimeError, ErrorType
+ from airflow.sdk.execution_time.context import _async_get_variable
+
+ try:
+ return await _async_get_variable(key,
deserialize_json=deserialize_json)
+ except AirflowRuntimeError as e:
+ if e.error.error == ErrorType.VARIABLE_NOT_FOUND and default is
not NOTSET:
+ await amask_secret(default, name=key)
+ return default
+ raise
+
@classmethod
def set(cls, key: str, value: Any, description: str | None = None,
serialize_json: bool = False) -> None:
from airflow.sdk.execution_time.context import _set_variable
_set_variable(key, value, description, serialize_json=serialize_json)
+ @classmethod
+ async def aset(
+ cls,
+ key: str,
+ value: Any,
+ description: str | None = None,
+ serialize_json: bool = False,
+ ) -> None:
+ from airflow.sdk.execution_time.context import _async_set_variable
+
+ await _async_set_variable(key, value, description,
serialize_json=serialize_json)
+
@classmethod
def keys(cls, prefix: str | None = None) -> Sequence[str]:
"""
@@ -88,8 +113,36 @@ class Variable:
return lazy_object_proxy.Proxy(lambda:
_get_variable_keys(prefix=prefix))
+ @classmethod
+ async def akeys(cls, prefix: str | None = None) -> Sequence[str]:
+ """
+ Return Variable keys that start with the given prefix.
+
+ Unlike :meth:`keys`, the result is **not** lazily evaluated — the async
+ call is awaited immediately and the resolved list is returned.
+
+ .. note::
+ Only keys stored in the metadata database are returned — secrets
backends
+ are **not** consulted. This asymmetry with :meth:`aget` (which
does consult
+ secrets backends) is a deliberate design decision: most secrets
backends
+ either do not expose a listing API at all, or do so inefficiently
and
+ without prefix filtering. See
+ https://github.com/apache/airflow/issues/61166 for context.
+
+ :param prefix: Optional key prefix to filter by. If None, all keys are
returned.
+ """
+ from airflow.sdk.execution_time.context import _async_get_variable_keys
+
+ return await _async_get_variable_keys(prefix=prefix)
+
@classmethod
def delete(cls, key: str) -> None:
from airflow.sdk.execution_time.context import _delete_variable
_delete_variable(key=key)
+
+ @classmethod
+ async def adelete(cls, key: str) -> None:
+ from airflow.sdk.execution_time.context import _async_delete_variable
+
+ await _async_delete_variable(key=key)
diff --git a/task-sdk/src/airflow/sdk/execution_time/context.py
b/task-sdk/src/airflow/sdk/execution_time/context.py
index fdb696ec906..1333fadf79c 100644
--- a/task-sdk/src/airflow/sdk/execution_time/context.py
+++ b/task-sdk/src/airflow/sdk/execution_time/context.py
@@ -16,6 +16,7 @@
# under the License.
from __future__ import annotations
+import asyncio
import collections
import contextlib
import functools
@@ -343,6 +344,23 @@ def _mask_and_deserialize_variable(raw: str, key: str,
deserialize_json: bool) -
return val
+async def _async_mask_and_deserialize_variable(raw: str, key: str,
deserialize_json: bool) -> Any:
+ await amask_secret(raw, key)
+ if not deserialize_json:
+ return raw
+ val = json.loads(raw)
+ if isinstance(val, str):
+ await amask_secret(val, key)
+ elif isinstance(val, dict):
+ # Masked by the dict's own inner key names, which is what ``add_mask``
uses.
+ await amask_secret(val)
+ elif isinstance(val, list):
+ # Pass the Variable's key so list elements inherit the Variable's
sensitivity
+ # instead of being added to the global mask patterns.
+ await amask_secret(val, key)
+ return val
+
+
def _get_variable(key: str, deserialize_json: bool) -> Any:
from airflow.sdk.execution_time.cache import SecretCache
from airflow.sdk.execution_time.supervisor import
ensure_secrets_backend_loaded
@@ -382,6 +400,52 @@ def _get_variable(key: str, deserialize_json: bool) -> Any:
)
+async def _async_get_variable(key: str, deserialize_json: bool) -> Any:
+ from airflow.sdk.execution_time.cache import SecretCache
+ from airflow.sdk.execution_time.supervisor import
ensure_secrets_backend_loaded
+
+ # Check cache first
+ try:
+ var_val = SecretCache.get_variable(key)
+ if var_val is not None:
+ return await _async_mask_and_deserialize_variable(var_val, key,
deserialize_json)
+ except SecretCache.NotPresentException:
+ pass # Continue to check backends
+
+ backends = ensure_secrets_backend_loaded()
+
+ # Iterate over backends if not in cache (or expired)
+ for secrets_backend in backends:
+ try:
+ async_method = getattr(secrets_backend, "aget_variable", None)
+ if async_method is not None:
+ var_val = await async_method(key=key)
+ else:
+ var_val = await
asyncio.to_thread(secrets_backend.get_variable, key=key)
+ if var_val is not None:
+ # Save raw value before deserialization to maintain cache
consistency
+ SecretCache.save_variable(key, var_val)
+ return await _async_mask_and_deserialize_variable(var_val,
key, deserialize_json)
+ except AirflowSecretsBackendAccessDenied:
+ # Authoritative deny — must NOT fall through to a less-restrictive
backend.
+ raise
+ except Exception:
+ log.exception(
+ "Unable to retrieve variable from secrets backend (%s).
Checking subsequent secrets backend.",
+ type(secrets_backend).__name__,
+ )
+
+ # If no backend found the variable, raise a not found error (mirrors
_get_connection)
+ from airflow.sdk.exceptions import AirflowRuntimeError, ErrorType
+
+ raise AirflowRuntimeError(
+ ErrorResponse(
+ error=ErrorType.VARIABLE_NOT_FOUND,
+ detail={"message": f"Variable {key} not found"},
+ )
+ )
+
+
_VARIABLE_KEYS_PAGE_SIZE = 1000
@@ -406,6 +470,27 @@ def _get_variable_keys(prefix: str | None = None) ->
list[str]:
return all_keys
+async def _async_get_variable_keys(prefix: str | None = None) -> list[str]:
+ from airflow.sdk.exceptions import AirflowRuntimeError
+ from airflow.sdk.execution_time.task_runner import SUPERVISOR_COMMS
+
+ all_keys: list[str] = []
+ offset = 0
+ while True:
+ msg = await SUPERVISOR_COMMS.asend(
+ GetVariableKeys(prefix=prefix, limit=_VARIABLE_KEYS_PAGE_SIZE,
offset=offset)
+ )
+ if isinstance(msg, ErrorResponse):
+ raise AirflowRuntimeError(msg)
+ if not isinstance(msg, VariableKeysResult):
+ raise TypeError(f"Unexpected response type for GetVariableKeys:
{type(msg).__name__}")
+ all_keys.extend(msg.keys)
+ if len(msg.keys) < _VARIABLE_KEYS_PAGE_SIZE:
+ break
+ offset += len(msg.keys)
+ return all_keys
+
+
def _set_variable(key: str, value: Any, description: str | None = None,
serialize_json: bool = False) -> None:
# TODO: This should probably be moved to a separate module like
`airflow.sdk.execution_time.comms`
# or `airflow.sdk.execution_time.variable`
@@ -454,6 +539,59 @@ def _set_variable(key: str, value: Any, description: str |
None = None, serializ
SecretCache.invalidate_variable(key)
+async def _async_set_variable(
+ key: str,
+ value: Any,
+ description: str | None = None,
+ serialize_json: bool = False,
+) -> None:
+ # TODO: This should probably be moved to a separate module like
`airflow.sdk.execution_time.comms`
+ # or `airflow.sdk.execution_time.variable`
+ # A reason to not move it to `airflow.sdk.execution_time.comms` is that
it
+ # will make that module depend on Task SDK, which is not ideal because
we intend to
+ # keep Task SDK as a separate package than execution time mods.
+ from airflow.sdk.execution_time.cache import SecretCache
+ from airflow.sdk.execution_time.secrets.execution_api import (
+ ExecutionAPISecretsBackend,
+ )
+ from airflow.sdk.execution_time.supervisor import
ensure_secrets_backend_loaded
+ from airflow.sdk.execution_time.task_runner import SUPERVISOR_COMMS
+
+ # check for write conflicts on the worker
+ for secrets_backend in ensure_secrets_backend_loaded():
+ if isinstance(secrets_backend, ExecutionAPISecretsBackend):
+ continue
+ try:
+ var_val = await asyncio.to_thread(secrets_backend.get_variable,
key=key)
+ if var_val is not None:
+ _backend_name = type(secrets_backend).__name__
+ log.warning(
+ "The variable %s is defined in the %s secrets backend,
which takes "
+ "precedence over reading from the API Server. The value
from the API Server will be "
+ "updated, but to read it you have to delete the
conflicting variable "
+ "from %s",
+ key,
+ _backend_name,
+ _backend_name,
+ )
+ except Exception:
+ log.exception(
+ "Unable to retrieve variable from secrets backend (%s).
Checking subsequent secrets backend.",
+ type(secrets_backend).__name__,
+ )
+
+ try:
+ if serialize_json:
+ value = json.dumps(value, indent=2)
+ except Exception as e:
+ log.exception(e)
+
+ await SUPERVISOR_COMMS.asend(PutVariable(key=key, value=value,
description=description))
+
+ # Invalidate cache after setting the variable
+ SecretCache.invalidate_variable(key)
+
+
def _delete_variable(key: str) -> None:
# TODO: This should probably be moved to a separate module like
`airflow.sdk.execution_time.comms`
# or `airflow.sdk.execution_time.variable`
@@ -471,6 +609,23 @@ def _delete_variable(key: str) -> None:
SecretCache.invalidate_variable(key)
+async def _async_delete_variable(key: str) -> None:
+ # TODO: This should probably be moved to a separate module like
`airflow.sdk.execution_time.comms`
+ # or `airflow.sdk.execution_time.variable`
+ # A reason to not move it to `airflow.sdk.execution_time.comms` is that
it
+ # will make that module depend on Task SDK, which is not ideal because
we intend to
+ # keep Task SDK as a separate package than execution time mods.
+ from airflow.sdk.execution_time.cache import SecretCache
+ from airflow.sdk.execution_time.task_runner import SUPERVISOR_COMMS
+
+ msg = await SUPERVISOR_COMMS.asend(DeleteVariable(key=key))
+ if TYPE_CHECKING:
+ assert isinstance(msg, OKResponse)
+
+ # Invalidate cache after deleting the variable
+ SecretCache.invalidate_variable(key)
+
+
class ConnectionAccessor:
"""Wrapper to access Connection entries in template."""
diff --git a/task-sdk/tests/task_sdk/definitions/test_variables.py
b/task-sdk/tests/task_sdk/definitions/test_variables.py
index 29b6ac0cb97..a3847a20157 100644
--- a/task-sdk/tests/task_sdk/definitions/test_variables.py
+++ b/task-sdk/tests/task_sdk/definitions/test_variables.py
@@ -29,11 +29,13 @@ from airflow.sdk.exceptions import AirflowRuntimeError,
ErrorType
from airflow.sdk.execution_time.comms import (
DeleteVariable,
ErrorResponse,
+ GetVariable,
GetVariableKeys,
PutVariable,
VariableKeysResult,
VariableResult,
)
+from airflow.sdk.execution_time.context import _VARIABLE_KEYS_PAGE_SIZE
from airflow.sdk.execution_time.secrets import
DEFAULT_SECRETS_SEARCH_PATH_WORKERS
from tests_common.test_utils.config import conf_vars
@@ -155,8 +157,6 @@ class TestVariableKeys:
def test_keys_paginates_when_results_exceed_page_size(self,
mock_supervisor_comms):
# Simulate two full pages followed by a short page (signals end).
- from airflow.sdk.execution_time.context import _VARIABLE_KEYS_PAGE_SIZE
-
page1 = [f"k{i}" for i in range(_VARIABLE_KEYS_PAGE_SIZE)]
page2 = [f"k{i}" for i in range(_VARIABLE_KEYS_PAGE_SIZE,
_VARIABLE_KEYS_PAGE_SIZE * 2)]
page3 = ["last_key"]
@@ -204,6 +204,141 @@ class TestVariableKeys:
list(results)
+class TestAsyncVariables:
+ @pytest.mark.asyncio
+ @pytest.mark.parametrize(
+ ("deserialize_json", "value", "expected_value"),
+ [
+ pytest.param(False, "my_value", "my_value", id="simple-value"),
+ pytest.param(
+ True,
+ '{"key": "value", "number": 42, "flag": true}',
+ {"key": "value", "number": 42, "flag": True},
+ id="deser-object-value",
+ ),
+ ],
+ )
+ async def test_avar_get(self, deserialize_json, value, expected_value,
mock_supervisor_comms):
+ mock_supervisor_comms.asend.return_value =
VariableResult(key="my_key", value=value)
+
+ var = await Variable.aget(key="my_key",
deserialize_json=deserialize_json)
+
+ assert var == expected_value
+ mock_supervisor_comms.asend.assert_any_call(GetVariable(key="my_key"))
+
+ @pytest.mark.asyncio
+ @patch("airflow.sdk.definitions.variable.amask_secret")
+ async def test_avar_get_returns_default_when_not_found(self,
mock_amask_secret, mock_supervisor_comms):
+ mock_supervisor_comms.asend.return_value = ErrorResponse(
+ error=ErrorType.VARIABLE_NOT_FOUND, detail={"message": "Variable
my_key not found"}
+ )
+
+ var = await Variable.aget(key="my_key", default="default_value")
+
+ assert var == "default_value"
+ mock_amask_secret.assert_awaited_once_with("default_value",
name="my_key")
+
+ @pytest.mark.asyncio
+ async def test_avar_get_raises_when_not_found_without_default(self,
mock_supervisor_comms):
+ mock_supervisor_comms.asend.return_value = ErrorResponse(
+ error=ErrorType.VARIABLE_NOT_FOUND, detail={"message": "Variable
my_key not found"}
+ )
+
+ with pytest.raises(AirflowRuntimeError) as exc_info:
+ await Variable.aget(key="my_key")
+
+ assert exc_info.value.error.error == ErrorType.VARIABLE_NOT_FOUND
+
+ @pytest.mark.asyncio
+ @pytest.mark.parametrize(
+ ("key", "value", "description", "serialize_json"),
+ [
+ pytest.param("key", "value", "description", False,
id="simple-value"),
+ pytest.param(
+ "key2",
+ {"hi": "there", "hello": 42, "flag": True},
+ "description2",
+ True,
+ id="serialize-json-value",
+ ),
+ ],
+ )
+ async def test_avar_set(self, key, value, description, serialize_json,
mock_supervisor_comms):
+ mock_supervisor_comms.asend.return_value = None
+
+ await Variable.aset(key=key, value=value, description=description,
serialize_json=serialize_json)
+
+ expected_value = value
+ if serialize_json:
+ expected_value = json.dumps(value, indent=2)
+
+ mock_supervisor_comms.asend.assert_called_once_with(
+ PutVariable(key=key, value=expected_value, description=description)
+ )
+
+ @pytest.mark.asyncio
+ async def test_avar_delete(self, mock_supervisor_comms):
+ mock_supervisor_comms.asend.return_value = None
+
+ await Variable.adelete(key="my_key")
+
+
mock_supervisor_comms.asend.assert_called_once_with(DeleteVariable(key="my_key"))
+
+
+class TestAsyncVariableKeys:
+ @pytest.mark.asyncio
+ @pytest.mark.parametrize(
+ ("prefix", "keys"),
+ [
+ pytest.param(None, ["prod_db", "prod_api", "dev_debug"], id="all"),
+ pytest.param("prod_", ["prod_db", "prod_api"], id="with-prefix"),
+ pytest.param("nonexistent_", [], id="empty-result"),
+ ],
+ )
+ async def test_akeys(self, prefix, keys, mock_supervisor_comms):
+ mock_supervisor_comms.asend.return_value =
VariableKeysResult(keys=keys, total_entries=len(keys))
+
+ results = await Variable.akeys(prefix=prefix)
+
+ mock_supervisor_comms.asend.assert_called_once_with(
+ GetVariableKeys(prefix=prefix, limit=1000, offset=0)
+ )
+ assert list(results) == keys
+
+ @pytest.mark.asyncio
+ async def test_akeys_paginates_when_results_exceed_page_size(self,
mock_supervisor_comms):
+ page1 = [f"k{i}" for i in range(_VARIABLE_KEYS_PAGE_SIZE)]
+ page2 = [f"k{i}" for i in range(_VARIABLE_KEYS_PAGE_SIZE,
_VARIABLE_KEYS_PAGE_SIZE * 2)]
+ page3 = ["last_key"]
+ total = _VARIABLE_KEYS_PAGE_SIZE * 2 + 1
+ mock_supervisor_comms.asend.side_effect = [
+ VariableKeysResult(keys=page1, total_entries=total),
+ VariableKeysResult(keys=page2, total_entries=total),
+ VariableKeysResult(keys=page3, total_entries=total),
+ ]
+
+ results = await Variable.akeys(prefix=None)
+
+ assert results == page1 + page2 + page3
+ assert mock_supervisor_comms.asend.call_count == 3
+
+ @pytest.mark.asyncio
+ async def test_akeys_raises_on_error_response(self, mock_supervisor_comms):
+ mock_supervisor_comms.asend.return_value = ErrorResponse(
+ error=ErrorType.GENERIC_ERROR, detail={"message": "boom"}
+ )
+
+ with pytest.raises(AirflowRuntimeError):
+ await Variable.akeys(prefix="x_")
+
+ @pytest.mark.asyncio
+ async def test_akeys_raises_on_unexpected_response_type(self,
mock_supervisor_comms):
+ mock_supervisor_comms.asend.return_value = VariableResult(key="x",
value="y")
+
+ with pytest.raises(TypeError, match="Unexpected response type"):
+ await Variable.akeys(prefix="x_")
+
+
class TestVariableFromSecrets:
def test_var_get_from_secrets_found(self, mock_supervisor_comms, tmp_path):
"""Tests getting a variable from secrets backend."""
diff --git a/task-sdk/tests/task_sdk/execution_time/test_context.py
b/task-sdk/tests/task_sdk/execution_time/test_context.py
index 3c8c5924423..cd518ae22cd 100644
--- a/task-sdk/tests/task_sdk/execution_time/test_context.py
+++ b/task-sdk/tests/task_sdk/execution_time/test_context.py
@@ -17,6 +17,7 @@
from __future__ import annotations
+import json
from datetime import datetime, timedelta, timezone as dt_timezone
from typing import TYPE_CHECKING
from unittest import mock
@@ -45,7 +46,12 @@ from airflow.sdk.definitions.asset import (
)
from airflow.sdk.definitions.connection import Connection
from airflow.sdk.definitions.variable import Variable
-from airflow.sdk.exceptions import AirflowNotFoundException,
AirflowRuntimeError, ErrorType
+from airflow.sdk.exceptions import (
+ AirflowNotFoundException,
+ AirflowRuntimeError,
+ AirflowSecretsBackendAccessDenied,
+ ErrorType,
+)
from airflow.sdk.execution_time.comms import (
AssetEventDagRunReferenceResult,
AssetEventResult,
@@ -62,6 +68,7 @@ from airflow.sdk.execution_time.comms import (
DeleteAssetStateStoreByName,
DeleteAssetStateStoreByUri,
DeleteTaskStateStore,
+ DeleteVariable,
ErrorResponse,
GetAssetByName,
GetAssetByUri,
@@ -71,16 +78,20 @@ from airflow.sdk.execution_time.comms import (
GetAssetStateStoreByUri,
GetDagRun,
GetTaskStateStore,
+ GetVariableKeys,
GetXCom,
OKResponse,
+ PutVariable,
SetAssetStateStoreByName,
SetAssetStateStoreByUri,
SetTaskStateStore,
TaskStateStoreResult,
+ VariableKeysResult,
VariableResult,
XComResult,
)
from airflow.sdk.execution_time.context import (
+ _VARIABLE_KEYS_PAGE_SIZE,
NEVER_EXPIRE,
AssetStateStoreAccessor,
AssetStateStoreAccessors,
@@ -93,7 +104,11 @@ from airflow.sdk.execution_time.context import (
TriggeringAssetEventsAccessor,
VariableAccessor,
_AssetRefResolutionMixin,
+ _async_delete_variable,
_async_get_connection,
+ _async_get_variable,
+ _async_get_variable_keys,
+ _async_set_variable,
_convert_variable_result_to_variable,
_get_connection,
_process_connection_result_conn,
@@ -1312,6 +1327,339 @@ class TestAsyncGetConnection:
mock_supervisor_comms.asend.assert_awaited()
+class TestAsyncVariableContext:
+ """Test async variable context functions: _async_get_variable,
_async_set_variable, etc."""
+
+ @pytest.mark.asyncio
+ @pytest.mark.parametrize(
+ ("deserialize_json", "value", "expected_value"),
+ [
+ pytest.param(False, "my_value", "my_value", id="simple-value"),
+ pytest.param(
+ True,
+ '{"key": "value", "number": 42, "flag": true}',
+ {"key": "value", "number": 42, "flag": True},
+ id="deser-object-value",
+ ),
+ ],
+ )
+ async def test_async_get_variable_from_api(
+ self, deserialize_json, value, expected_value, mock_supervisor_comms
+ ):
+ """_async_get_variable fetches from the Execution API via
ExecutionAPISecretsBackend (async asend)."""
+ mock_supervisor_comms.asend.return_value =
VariableResult(key="my_key", value=value)
+
+ result = await _async_get_variable("my_key",
deserialize_json=deserialize_json)
+
+ assert result == expected_value
+ mock_supervisor_comms.asend.assert_called()
+
+ @pytest.mark.asyncio
+ async def test_async_get_variable_from_secrets_backend(self,
mock_supervisor_comms):
+ """_async_get_variable returns the value from a secrets backend
without calling asend."""
+
+ class MockBackend:
+ def get_variable(self, key: str):
+ return "backend_value"
+
+ with patch(
+
"airflow.sdk.execution_time.supervisor.ensure_secrets_backend_loaded",
autospec=True
+ ) as mock_load:
+ mock_load.return_value = [MockBackend()]
+ result = await _async_get_variable("my_key",
deserialize_json=False)
+
+ assert result == "backend_value"
+ # ExecutionAPISecretsBackend (which uses sync send) must not have been
called
+ mock_supervisor_comms.send.assert_not_called()
+
+ @pytest.mark.asyncio
+ async def test_async_get_variable_not_found_raises(self,
mock_supervisor_comms):
+ """_async_get_variable raises AirflowRuntimeError when the variable
does not exist."""
+ with patch(
+
"airflow.sdk.execution_time.supervisor.ensure_secrets_backend_loaded",
autospec=True
+ ) as mock_load:
+ mock_load.return_value = []
+
+ with pytest.raises(AirflowRuntimeError) as exc_info:
+ await _async_get_variable("missing_key",
deserialize_json=False)
+
+ assert exc_info.value.error.error == ErrorType.VARIABLE_NOT_FOUND
+
+ @pytest.mark.asyncio
+ async def test_async_get_variable_does_not_fall_through_after_deny(self,
mock_supervisor_comms):
+ """An authoritative deny from the Execution API raises; the next
backend is never asked."""
+ mock_supervisor_comms.asend.return_value = ErrorResponse(
+ error=ErrorType.PERMISSION_DENIED,
+ detail={"key": "denied_var", "status_code": 403},
+ )
+
+ later_backend = MagicMock(name="LaterBackend")
+ # the dispatcher prefers aget_variable when present, so spy on both
+ later_backend.aget_variable =
mock.AsyncMock(return_value="leaked-value")
+ later_backend.get_variable = MagicMock(return_value="leaked-value")
+
+ with patch(
+
"airflow.sdk.execution_time.supervisor.ensure_secrets_backend_loaded",
autospec=True
+ ) as mock_load:
+ mock_load.return_value = [ExecutionAPISecretsBackend(),
later_backend]
+
+ with pytest.raises(AirflowSecretsBackendAccessDenied,
match="variable 'denied_var'"):
+ await _async_get_variable("denied_var", deserialize_json=False)
+
+ later_backend.aget_variable.assert_not_awaited()
+ later_backend.get_variable.assert_not_called()
+
+ @pytest.mark.asyncio
+ @mock.patch("airflow.sdk.execution_time.context.amask_secret")
+ async def test_async_var_json_masks_from_cache(self, mock_amask_secret,
mock_supervisor_comms):
+ """SecretCache hit path applies the same masking as the backends
path."""
+ from airflow.sdk.execution_time.cache import SecretCache
+
+ raw_json = '{"password": "s3cr3t", "host": "db.example.com"}'
+ with mock.patch.object(SecretCache, "get_variable",
return_value=raw_json):
+ val = await _async_get_variable("db_config", deserialize_json=True)
+
+ assert val == {"password": "s3cr3t", "host": "db.example.com"}
+ mock_amask_secret.assert_any_await(raw_json, "db_config")
+ mock_amask_secret.assert_any_await({"password": "s3cr3t", "host":
"db.example.com"})
+ # served from the cache, no backend was asked
+ mock_supervisor_comms.asend.assert_not_awaited()
+
+ @pytest.mark.asyncio
+ @mock.patch("airflow.sdk.execution_time.context.amask_secret")
+ async def test_async_var_value_masks_secret(self, mock_amask_secret,
mock_supervisor_comms):
+ """A plain value is masked under the variable's key name."""
+ mock_supervisor_comms.asend.return_value =
VariableResult(key="my_password", value="s3cr3t")
+
+ val = await _async_get_variable("my_password", deserialize_json=False)
+
+ assert val == "s3cr3t"
+ mock_amask_secret.assert_awaited_once_with("s3cr3t", "my_password")
+
+ @pytest.mark.asyncio
+ @mock.patch("airflow.sdk.execution_time.context.amask_secret")
+ async def test_async_var_json_masks_raw_string_and_dict_values(
+ self, mock_amask_secret, mock_supervisor_comms
+ ):
+ """Both the raw JSON string and the deserialized dict's sensitive
fields are masked."""
+ raw_json = '{"password": "s3cr3t", "host": "db.example.com"}'
+ mock_supervisor_comms.asend.return_value =
VariableResult(key="db_config", value=raw_json)
+
+ val = await _async_get_variable("db_config", deserialize_json=True)
+
+ assert val == {"password": "s3cr3t", "host": "db.example.com"}
+ mock_amask_secret.assert_any_await(raw_json, "db_config")
+ mock_amask_secret.assert_any_await({"password": "s3cr3t", "host":
"db.example.com"})
+
+ @pytest.mark.asyncio
+ @mock.patch("airflow.sdk.execution_time.context.amask_secret")
+ async def test_async_var_json_masks_list_values(self, mock_amask_secret,
mock_supervisor_comms):
+ """A JSON list is handed to the masker whole, under the variable's
key."""
+ raw_json = '[{"password": "s3cr3t"}, {"password": "s3cr3t2"}]'
+ mock_supervisor_comms.asend.return_value =
VariableResult(key="db_configs", value=raw_json)
+
+ val = await _async_get_variable("db_configs", deserialize_json=True)
+
+ assert val == [{"password": "s3cr3t"}, {"password": "s3cr3t2"}]
+ mock_amask_secret.assert_any_await(raw_json, "db_configs")
+ mock_amask_secret.assert_any_await([{"password": "s3cr3t"},
{"password": "s3cr3t2"}], "db_configs")
+
+ @pytest.mark.asyncio
+ @mock.patch("airflow.sdk.execution_time.context.amask_secret")
+ async def test_async_var_json_sensitive_key_masks_raw_json(
+ self, mock_amask_secret, mock_supervisor_comms
+ ):
+ """A sensitive variable key masks the entire raw JSON string."""
+ raw_json = '{"endpoint": "https://api.example.com", "token": "abc123"}'
+ mock_supervisor_comms.asend.return_value =
VariableResult(key="my_secret", value=raw_json)
+
+ val = await _async_get_variable("my_secret", deserialize_json=True)
+
+ assert val == {"endpoint": "https://api.example.com", "token":
"abc123"}
+ mock_amask_secret.assert_any_await(raw_json, "my_secret")
+ mock_amask_secret.assert_any_await({"endpoint":
"https://api.example.com", "token": "abc123"})
+
+ @pytest.mark.asyncio
+ @mock.patch("airflow.sdk.execution_time.context.amask_secret")
+ async def test_async_var_json_string_value_masks_both_forms(
+ self, mock_amask_secret, mock_supervisor_comms
+ ):
+ """A JSON string value masks both the quoted raw and the unquoted
value."""
+ mock_supervisor_comms.asend.return_value =
VariableResult(key="my_token", value='"s3cr3t"')
+
+ val = await _async_get_variable("my_token", deserialize_json=True)
+
+ assert val == "s3cr3t"
+ assert mock_amask_secret.await_count == 2
+ mock_amask_secret.assert_any_await('"s3cr3t"', "my_token")
+ mock_amask_secret.assert_any_await("s3cr3t", "my_token")
+
+ @pytest.mark.asyncio
+ @mock.patch("airflow.sdk.execution_time.context.amask_secret")
+ async def test_async_var_json_list_value_does_not_over_mask(
+ self, mock_amask_secret, mock_supervisor_comms
+ ):
+ """A non-sensitive list variable is never masked anonymously."""
+ raw_json = '["us-east-1", "eu-west-1"]'
+ mock_supervisor_comms.asend.return_value =
VariableResult(key="aws_regions", value=raw_json)
+
+ val = await _async_get_variable("aws_regions", deserialize_json=True)
+
+ assert val == ["us-east-1", "eu-west-1"]
+ mock_amask_secret.assert_any_await(raw_json, "aws_regions")
+ mock_amask_secret.assert_any_await(["us-east-1", "eu-west-1"],
"aws_regions")
+ # never anonymously -- that is what would mask the elements globally
+ assert mock.call(["us-east-1", "eu-west-1"]) not in
mock_amask_secret.await_args_list
+
+ @pytest.mark.asyncio
+ @mock.patch("airflow.sdk.execution_time.context.amask_secret")
+ async def test_async_var_json_invalid_json_raises(self, mock_amask_secret,
mock_supervisor_comms):
+ """Invalid JSON raises JSONDecodeError; the raw value is still masked
before the error."""
+ from airflow.sdk.execution_time.cache import SecretCache
+
+ raw = "not-valid-json"
+ with mock.patch.object(SecretCache, "get_variable", return_value=raw):
+ with pytest.raises(json.JSONDecodeError):
+ await _async_get_variable("bad_var", deserialize_json=True)
+
+ mock_amask_secret.assert_awaited_once_with(raw, "bad_var")
+
+ @pytest.mark.asyncio
+ @pytest.mark.parametrize(
+ ("key", "value", "description", "serialize_json"),
+ [
+ pytest.param("key", "value", "description", False,
id="simple-value"),
+ pytest.param(
+ "key2",
+ {"hi": "there", "hello": 42, "flag": True},
+ "description2",
+ True,
+ id="serialize-json-value",
+ ),
+ ],
+ )
+ async def test_async_set_variable(self, key, value, description,
serialize_json, mock_supervisor_comms):
+ """_async_set_variable sends PutVariable via asend."""
+ mock_supervisor_comms.asend.return_value = None
+
+ with patch(
+
"airflow.sdk.execution_time.supervisor.ensure_secrets_backend_loaded",
autospec=True
+ ) as mock_load:
+ mock_load.return_value = []
+ await _async_set_variable(key, value, description,
serialize_json=serialize_json)
+
+ expected_value = json.dumps(value, indent=2) if serialize_json else
value
+ mock_supervisor_comms.asend.assert_called_once_with(
+ PutVariable(key=key, value=expected_value, description=description)
+ )
+
+ @pytest.mark.asyncio
+ async def test_async_set_variable_warns_on_conflicting_backend(self,
mock_supervisor_comms):
+ """A worker-side backend that already holds the key gets a warning;
the write still goes through."""
+ mock_supervisor_comms.asend.return_value = None
+
+ class ConflictingBackend:
+ def get_variable(self, key: str):
+ return "backend_value"
+
+ class EmptyBackend:
+ def get_variable(self, key: str):
+ return None
+
+ class FailingBackend:
+ def get_variable(self, key: str):
+ raise RuntimeError("backend down")
+
+ execution_api_backend =
mock.create_autospec(ExecutionAPISecretsBackend, instance=True)
+
+ with (
+ patch(
+
"airflow.sdk.execution_time.supervisor.ensure_secrets_backend_loaded",
autospec=True
+ ) as mock_load,
+ patch("airflow.sdk.execution_time.context.log") as mock_log,
+ ):
+ mock_load.return_value = [
+ ConflictingBackend(),
+ EmptyBackend(),
+ FailingBackend(),
+ execution_api_backend,
+ ]
+ await _async_set_variable("my_key", "new_value")
+
+ mock_log.warning.assert_called_once()
+ warning_args = mock_log.warning.call_args.args
+ assert warning_args[1:] == ("my_key", "ConflictingBackend",
"ConflictingBackend")
+ mock_log.exception.assert_called_once()
+ assert mock_log.exception.call_args.args[1] == "FailingBackend"
+ # the Execution API is the write target, not a conflict
+ execution_api_backend.get_variable.assert_not_called()
+ mock_supervisor_comms.asend.assert_called_once_with(
+ PutVariable(key="my_key", value="new_value", description=None)
+ )
+
+ @pytest.mark.asyncio
+ async def test_async_delete_variable(self, mock_supervisor_comms):
+ """_async_delete_variable sends DeleteVariable via asend."""
+ mock_supervisor_comms.asend.return_value = OKResponse(ok=True)
+
+ await _async_delete_variable("my_key")
+
+
mock_supervisor_comms.asend.assert_called_once_with(DeleteVariable(key="my_key"))
+
+ @pytest.mark.asyncio
+ @pytest.mark.parametrize(
+ ("prefix", "keys"),
+ [
+ pytest.param(None, ["prod_db", "prod_api", "dev_debug"], id="all"),
+ pytest.param("prod_", ["prod_db", "prod_api"], id="with-prefix"),
+ pytest.param("nonexistent_", [], id="empty-result"),
+ ],
+ )
+ async def test_async_get_variable_keys(self, prefix, keys,
mock_supervisor_comms):
+ """_async_get_variable_keys fetches all keys matching the prefix in
one page."""
+ mock_supervisor_comms.asend.return_value =
VariableKeysResult(keys=keys, total_entries=len(keys))
+
+ result = await _async_get_variable_keys(prefix=prefix)
+
+ assert result == keys
+ mock_supervisor_comms.asend.assert_called_once_with(
+ GetVariableKeys(prefix=prefix, limit=1000, offset=0)
+ )
+
+ @pytest.mark.asyncio
+ async def test_async_get_variable_keys_paginates(self,
mock_supervisor_comms):
+ """_async_get_variable_keys accumulates results across multiple
pages."""
+ page1 = [f"k{i}" for i in range(_VARIABLE_KEYS_PAGE_SIZE)]
+ page2 = ["last_key"]
+ mock_supervisor_comms.asend.side_effect = [
+ VariableKeysResult(keys=page1,
total_entries=_VARIABLE_KEYS_PAGE_SIZE + 1),
+ VariableKeysResult(keys=page2,
total_entries=_VARIABLE_KEYS_PAGE_SIZE + 1),
+ ]
+
+ result = await _async_get_variable_keys(prefix=None)
+
+ assert result == page1 + page2
+ assert mock_supervisor_comms.asend.call_count == 2
+
+ @pytest.mark.asyncio
+ async def test_async_get_variable_keys_raises_on_error(self,
mock_supervisor_comms):
+ """_async_get_variable_keys raises AirflowRuntimeError on an
ErrorResponse."""
+ mock_supervisor_comms.asend.return_value = ErrorResponse(
+ error=ErrorType.GENERIC_ERROR, detail={"message": "boom"}
+ )
+
+ with pytest.raises(AirflowRuntimeError):
+ await _async_get_variable_keys(prefix="x_")
+
+ @pytest.mark.asyncio
+ async def test_async_get_variable_keys_raises_on_unexpected_response(self,
mock_supervisor_comms):
+ """_async_get_variable_keys raises TypeError for an unrecognised
response type."""
+ mock_supervisor_comms.asend.return_value = VariableResult(key="x",
value="y")
+
+ with pytest.raises(TypeError, match="Unexpected response type"):
+ await _async_get_variable_keys(prefix="x_")
+
+
class TestSecretsBackend:
"""Test that connection resolution uses the backend chain correctly."""