diff --git a/cfa/cloudops/_cloudclient.py b/cfa/cloudops/_cloudclient.py index e62bd9c2..3f6492c0 100644 --- a/cfa/cloudops/_cloudclient.py +++ b/cfa/cloudops/_cloudclient.py @@ -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 @@ -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 @@ -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. diff --git a/changelog.md b/changelog.md index 6a02a2a2..e70898de 100644 --- a/changelog.md +++ b/changelog.md @@ -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 diff --git a/pyproject.toml b/pyproject.toml index bd080532..dab2a62d 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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"} diff --git a/tests/test_cloudclient.py b/tests/test_cloudclient.py index 1666c382..bad014f8 100644 --- a/tests/test_cloudclient.py +++ b/tests/test_cloudclient.py @@ -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, diff --git a/uv.lock b/uv.lock index 5c04e86a..1dfecd7b 100644 --- a/uv.lock +++ b/uv.lock @@ -489,7 +489,7 @@ wheels = [ [[package]] name = "cfa-cloudops" -version = "1.1.1" +version = "1.2.0" source = { editable = "." } dependencies = [ { name = "anyio" },