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."""
 

Reply via email to