diff --git a/task-sdk/src/airflow/sdk/definitions/variable.py b/task-sdk/src/airflow/sdk/definitions/variable.py index a379022a90918..f1db76c3fae60 100644 --- a/task-sdk/src/airflow/sdk/definitions/variable.py +++ b/task-sdk/src/airflow/sdk/definitions/variable.py @@ -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]: @@ -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) diff --git a/task-sdk/tests/task_sdk/definitions/test_variables.py b/task-sdk/tests/task_sdk/definitions/test_variables.py index 6e94ccf503f8c..29b6ac0cb97ca 100644 --- a/task-sdk/tests/task_sdk/definitions/test_variables.py +++ b/task-sdk/tests/task_sdk/definitions/test_variables.py @@ -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 @@ -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( @@ -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"}