Skip to content
Merged
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
12 changes: 2 additions & 10 deletions task-sdk/src/airflow/sdk/definitions/variable.py
Original file line number Diff line number Diff line change
Expand Up @@ -60,13 +60,9 @@ def get(cls, key: str, default: Any = NOTSET, deserialize_json: bool = False):

@classmethod
def set(cls, key: str, value: Any, description: str | None = None, serialize_json: bool = False) -> None:
from airflow.sdk.exceptions import AirflowRuntimeError
from airflow.sdk.execution_time.context import _set_variable

try:
return _set_variable(key, value, description, serialize_json=serialize_json)
except AirflowRuntimeError as e:
log.exception(e)
_set_variable(key, value, description, serialize_json=serialize_json)

@classmethod
def keys(cls, prefix: str | None = None) -> Sequence[str]:
Expand Down Expand Up @@ -94,10 +90,6 @@ def keys(cls, prefix: str | None = None) -> Sequence[str]:

@classmethod
def delete(cls, key: str) -> None:
from airflow.sdk.exceptions import AirflowRuntimeError
from airflow.sdk.execution_time.context import _delete_variable

try:
_delete_variable(key=key)
except AirflowRuntimeError as e:
log.exception(e)
_delete_variable(key=key)
17 changes: 14 additions & 3 deletions task-sdk/tests/task_sdk/definitions/test_variables.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,7 +25,15 @@

from airflow.sdk import Variable
from airflow.sdk.configuration import initialize_secrets_backends
from airflow.sdk.execution_time.comms import GetVariableKeys, PutVariable, VariableKeysResult, VariableResult
from airflow.sdk.exceptions import AirflowRuntimeError, ErrorType
from airflow.sdk.execution_time.comms import (
DeleteVariable,
ErrorResponse,
GetVariableKeys,
PutVariable,
VariableKeysResult,
VariableResult,
)
from airflow.sdk.execution_time.secrets import DEFAULT_SECRETS_SEARCH_PATH_WORKERS

from tests_common.test_utils.config import conf_vars
Expand Down Expand Up @@ -89,6 +97,11 @@ def test_var_set(self, key, value, description, serialize_json, mock_supervisor_
),
)

def test_var_delete(self, mock_supervisor_comms):
Variable.delete(key="my_key")

mock_supervisor_comms.send.assert_called_once_with(msg=DeleteVariable(key="my_key"))


class TestVariableKeys:
@pytest.mark.parametrize(
Expand Down Expand Up @@ -171,8 +184,6 @@ def test_keys_paginates_when_results_exceed_page_size(self, mock_supervisor_comm
)

def test_keys_raises_on_error_response(self, mock_supervisor_comms):
from airflow.sdk.exceptions import AirflowRuntimeError, ErrorType
from airflow.sdk.execution_time.comms import ErrorResponse

mock_supervisor_comms.send.return_value = ErrorResponse(
error=ErrorType.GENERIC_ERROR, detail={"message": "boom"}
Expand Down
Loading