Skip to content
Closed
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
Original file line number Diff line number Diff line change
Expand Up @@ -44,6 +44,8 @@ class AzureBatchHook(BaseHook):

:param azure_batch_conn_id: :ref:`Azure Batch connection id<howto/connection:azure_batch>`
of a service principal which will be used to start the container instance.
:param batch_max_retries: The number of times a request to the Batch service is retried
before it is considered failed. Default is 3.
"""

conn_name_attr = "azure_batch_conn_id"
Expand Down Expand Up @@ -74,9 +76,10 @@ def get_ui_field_behaviour(cls) -> dict[str, Any]:
},
}

def __init__(self, azure_batch_conn_id: str = default_conn_name) -> None:
def __init__(self, azure_batch_conn_id: str = default_conn_name, batch_max_retries: int = 3) -> None:
super().__init__()
self.conn_id = azure_batch_conn_id
self.batch_max_retries = batch_max_retries

def _get_field(self, extras, name):
return get_field(
Expand Down Expand Up @@ -117,6 +120,7 @@ def get_conn(self) -> BatchClient:
batch_client = BatchClient(
endpoint=batch_account_url,
credential=credential,
retry_total=self.batch_max_retries,
)
return batch_client

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -57,8 +57,8 @@ class AzureBatchOperator(BaseOperator):
:param batch_start_task: A Task specified to run on each Compute Node as it joins the Pool.
The Task runs when the Compute Node is added to the Pool or
when the Compute Node is restarted.
:param batch_max_retries: The number of times to retry this batch operation before it's
considered a failed operation. Default is 3
:param batch_max_retries: The number of times a request to the Batch service is retried
before it is considered failed. Default is 3
:param batch_task_resource_files: A list of files that the Batch service will
download to the Compute Node before running the command line.
:param batch_task_output_files: A list of files that the Batch service will upload
Expand Down Expand Up @@ -184,7 +184,7 @@ def __init__(
@cached_property
def hook(self) -> AzureBatchHook:
"""Create and return an AzureBatchHook (cached)."""
return AzureBatchHook(self.azure_batch_conn_id)
return AzureBatchHook(self.azure_batch_conn_id, batch_max_retries=self.batch_max_retries)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

In deferrable mode, AzureBatchTrigger.run() creates a separate AzureBatchHook without passing batch_max_retries. I think we should also add on it.

hook = AzureBatchHook(
azure_batch_conn_id=self.azure_batch_conn_id,
)


def _check_inputs(self) -> Any:
if not self.vm_publisher:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -67,6 +67,20 @@ def test_connection_and_client(self):
assert isinstance(conn, BatchClient)
assert hook.connection is conn, "`connection` property should be cached"

@pytest.mark.parametrize(
("hook_kwargs", "expected_retry_total"),
[
pytest.param({}, 3, id="default"),
pytest.param({"batch_max_retries": 7}, 7, id="explicit"),
],
)
def test_batch_max_retries_is_applied_to_the_client(self, hook_kwargs, expected_retry_total):
# Asserted on a real client: ``retry_total`` reaches azure-core through ``**kwargs``.
hook = AzureBatchHook(azure_batch_conn_id=self.test_vm_conn_id, **hook_kwargs)
client = hook.get_conn()
assert client._config.retry_policy.total_retries == expected_retry_total
assert client._config.retry_policy in client._client._pipeline._impl_policies

@mock.patch(f"{MODULE}.get_sync_default_azure_credential")
def test_fallback_to_default_azure_credential_when_name_and_key_is_not_provided(
self, mock_get_default_credential, create_mock_connections
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -146,6 +146,35 @@ def setup_test_cases(self, mocked_batch_client, create_mock_connections):
timeout=2,
)

@pytest.mark.parametrize(
("operator_kwargs", "expected_retry_total"),
[
pytest.param({}, 3, id="default"),
pytest.param({"batch_max_retries": 7}, 7, id="explicit"),
],
)
def test_batch_max_retries_reaches_the_batch_client(
self, mocked_batch_client, operator_kwargs, expected_retry_total
):
operator = AzureBatchOperator(
task_id=TASK_ID,
batch_pool_id=BATCH_POOL_ID,
batch_pool_vm_size=BATCH_VM_SIZE,
batch_job_id=BATCH_JOB_ID,
batch_task_id=BATCH_TASK_ID,
vm_publisher=self.test_vm_publisher,
vm_offer=self.test_vm_offer,
vm_sku=self.test_vm_sku,
vm_node_agent_sku_id=self.test_node_agent_sku,
sku_starts_with=self.test_vm_sku,
batch_task_command_line="echo hello",
azure_batch_conn_id=self.test_vm_conn_id,
target_dedicated_nodes=1,
**operator_kwargs,
)
operator.hook.get_conn()
assert mocked_batch_client.call_args.kwargs["retry_total"] == expected_retry_total

@mock.patch.object(AzureBatchHook, "wait_for_all_node_state")
def test_execute_without_failures(self, wait_mock):
wait_mock.return_value = True
Expand Down