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
31 changes: 27 additions & 4 deletions cfa/cloudops/_cloudclient.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@
BatchJobConstraints,
BatchMetadataItem,
)
from azure.core.pipeline.transport import RequestsTransport
from azure.keyvault.secrets import SecretClient
from azure.mgmt.batch import models
from azure.mgmt.resource.subscriptions import SubscriptionClient
Expand Down Expand Up @@ -155,12 +156,30 @@ def __init__(
**kwargs,
)

# Create a shared Azure HTTP transport so SDK clients can reuse connections.
self._http_transport = RequestsTransport()
self._http_transport.open()

def _transport() -> RequestsTransport:
return RequestsTransport(
session=self._http_transport.session,
session_owner=False,
)

# get clients
logger.debug("Getting Azure clients and setting other attributes.")
self.batch_mgmt_client = get_batch_management_client(self.cred)
self.compute_mgmt_client = get_compute_management_client(self.cred)
self.batch_service_client = get_batch_service_client(self.cred)
self.blob_service_client = get_blob_service_client(self.cred)
self.batch_mgmt_client = get_batch_management_client(
self.cred, transport=_transport()
)
self.compute_mgmt_client = get_compute_management_client(
self.cred, transport=_transport()
)
self.batch_service_client = get_batch_service_client(
self.cred, transport=_transport()
)
self.blob_service_client = get_blob_service_client(
self.cred, transport=_transport()
)

# set other defaults
self.full_container_name = None
Expand All @@ -169,6 +188,10 @@ def __init__(
self.task_id_ints = False
self.task_id_max = 0

def close(self) -> None:
"""Close the shared HTTP transport used by Azure SDK clients."""
self._http_transport.close()

def check_credentials(self) -> pd.DataFrame:
"""Check credentials and return accessible Azure resource information.

Expand Down
4 changes: 4 additions & 0 deletions changelog.md
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,10 @@ and this project adheres to [Semantic Versioning](https://semver.org/).
The versioning pattern is `major.minor.patch`.

---
## 1.2.0

- reuse HTTP connections across the core Azure SDK clients created by `CloudClient`
- add explicit cleanup for the shared HTTP transport
## 1.1.1

- add notebook-friendly progress bars for task-stat writing and blob download, deletion, and protection operations
Expand Down
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
[project]
name = "cfa.cloudops"
version = "1.1.1"
version = "1.2.0"
description = "Cloud storage, batch, functions, MLOps assistance"
authors = [
{name = "Ryan Raasch", email = "xng3@cdc.gov"}
Expand Down
57 changes: 57 additions & 0 deletions tests/test_cloudclient.py
Original file line number Diff line number Diff line change
Expand Up @@ -433,6 +433,63 @@ def test_cloudclient_init_with_env_credentials(
mock_compute_management_client.assert_called_once()


def test_cloudclient_reuses_http_session_across_azure_clients(
mock_env_vars,
mock_batch_service_client,
mock_batch_management_client,
mock_blob_service_client,
mock_compute_management_client,
):
with patch(
"cfa.cloudops._cloudclient.DefaultCredentialHandler"
) as mock_cred_handler:
mock_cred_handler.return_value = MagicMock()

client = CloudClient(dotenv_path=None, use_sp=False, use_federated=False)

client_factories = [
mock_batch_management_client,
mock_compute_management_client,
mock_batch_service_client,
mock_blob_service_client,
]

transports = [
factory.call_args.kwargs["transport"] for factory in client_factories
]

sessions = [transport.session for transport in transports]

assert all(session is client._http_transport.session for session in sessions)

shared_session = client._http_transport.session

with patch.object(shared_session, "close") as mock_close:
transports[0].close()

mock_close.assert_not_called()


def test_cloudclient_close_closes_shared_http_session(
mock_env_vars,
mock_batch_service_client,
mock_batch_management_client,
mock_blob_service_client,
mock_compute_management_client,
):
with patch(
"cfa.cloudops._cloudclient.DefaultCredentialHandler"
) as mock_cred_handler:
mock_cred_handler.return_value = MagicMock()

client = CloudClient(dotenv_path=None, use_sp=False, use_federated=False)

with patch.object(client._http_transport, "close") as mock_close:
client.close()

mock_close.assert_called_once_with()


def test_cloudclient_init_with_default_credentials(
mock_env_vars,
mock_batch_service_client,
Expand Down
2 changes: 1 addition & 1 deletion uv.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

Loading