diff --git a/CHANGELOG.md b/CHANGELOG.md index 4a57a4c1ca..b3ff646bb4 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -4,6 +4,9 @@ ENHANCEMENTS: +BUG FIXES: +* Fix inconsistent ServiceBusClient lifecycle management in deployment_status_updater.py, airlock_request_status_update.py, and runner.py to prevent connection socket and AMQP channel leaks ([#4930](https://github.com/microsoft/AzureTRE/pull/4930)) + ## (0.29.0) (August 14, 2026) **BREAKING CHANGES** * Remove Windows 10 and dsvm image support from Guacamole. ([#4890](https://github.com/microsoft/AzureTRE/issues/4890)) diff --git a/api_app/_version.py b/api_app/_version.py index ae62eb6326..1932d0ba38 100644 --- a/api_app/_version.py +++ b/api_app/_version.py @@ -1 +1 @@ -__version__ = "0.26.5" +__version__ = "0.26.6" diff --git a/api_app/service_bus/airlock_request_status_update.py b/api_app/service_bus/airlock_request_status_update.py index c4d50b80d0..bf631094c9 100644 --- a/api_app/service_bus/airlock_request_status_update.py +++ b/api_app/service_bus/airlock_request_status_update.py @@ -34,43 +34,60 @@ async def receive_messages(self): while True: try: - current_time = time.time() - polling_count += 1 - # Log a heartbeat message every 60 seconds to show the service is still working - if current_time - last_heartbeat_time >= 60: - logger.info(f"Queue reader heartbeat: Polled {config.SERVICE_BUS_STEP_RESULT_QUEUE} queue {polling_count} times in the last minute") - last_heartbeat_time = current_time - polling_count = 0 - async with credentials.get_credential_async_context() as credential: + # We keep a single ServiceBusClient alive across the inner loop to avoid excessive connection + # and reconnection churn. Any fatal connection-related errors or other exceptions + # will propagate out of the inner loop, closing this context manager and recreating the client. async with ServiceBusClient(config.SERVICE_BUS_FULLY_QUALIFIED_NAMESPACE, credential) as service_bus_client: - receiver = service_bus_client.get_queue_receiver(queue_name=config.SERVICE_BUS_STEP_RESULT_QUEUE) - logger.debug(f"Looking for new messages on {config.SERVICE_BUS_STEP_RESULT_QUEUE} queue...") - async with receiver: - received_msgs = await receiver.receive_messages(max_message_count=10, max_wait_time=1) - for msg in received_msgs: - async with AutoLockRenewer() as renewer: - renewer.register(receiver, msg, max_lock_renewal_duration=60) - complete_message = await self.process_message(msg) - if complete_message: - await receiver.complete_message(msg) - else: - # could have been any kind of transient issue, we'll abandon back to the queue, and retry - await receiver.abandon_message(msg) - - await asyncio.sleep(10) - - except OperationTimeoutError: - # Timeout occurred whilst connecting to a session - this is expected and indicates no non-empty sessions are available - logger.debug("No sessions for this process. Will look again...") + client_created_time = time.time() + while True: + try: + # Recreate the client periodically (every hour) to ensure connection freshness + # and avoid holding a potentially stale client open indefinitely. + if time.time() - client_created_time > 3600: + logger.info("ServiceBusClient has been active for 1 hour. Recreating for freshness...") + break + + current_time = time.time() + polling_count += 1 + # Log a heartbeat message every 60 seconds to show the service is still working + if current_time - last_heartbeat_time >= 60: + logger.info(f"Queue reader heartbeat: Polled {config.SERVICE_BUS_STEP_RESULT_QUEUE} queue {polling_count} times in the last minute") + last_heartbeat_time = current_time + polling_count = 0 + + logger.debug(f"Looking for new messages on {config.SERVICE_BUS_STEP_RESULT_QUEUE} queue...") + receiver = service_bus_client.get_queue_receiver(queue_name=config.SERVICE_BUS_STEP_RESULT_QUEUE) + async with receiver: + received_msgs = await receiver.receive_messages(max_message_count=10, max_wait_time=1) + for msg in received_msgs: + async with AutoLockRenewer() as renewer: + renewer.register(receiver, msg, max_lock_renewal_duration=60) + complete_message = await self.process_message(msg) + if complete_message: + await receiver.complete_message(msg) + else: + # could have been any kind of transient issue, we'll abandon back to the queue, and retry + await receiver.abandon_message(msg) + + await asyncio.sleep(10) + + except OperationTimeoutError: + # Timeout occurred whilst connecting - this is expected and indicates no messages are available + logger.debug("No messages for this process. Will look again...") except ServiceBusConnectionError: # Occasionally there will be a transient / network-level error in connecting to SB. logger.info("Unknown Service Bus connection error. Will retry...") + await asyncio.sleep(10) + + except asyncio.CancelledError: + raise except Exception as e: # Catch all other exceptions, log them via .exception to get the stack trace, and reconnect logger.exception(f"Unknown exception. Will retry - {e}") + await asyncio.sleep(10) async def process_message(self, msg): with tracer.start_as_current_span("process_message") as current_span: diff --git a/api_app/service_bus/deployment_status_updater.py b/api_app/service_bus/deployment_status_updater.py index 702f878892..71dd4d628e 100644 --- a/api_app/service_bus/deployment_status_updater.py +++ b/api_app/service_bus/deployment_status_updater.py @@ -43,43 +43,60 @@ async def receive_messages(self): while True: try: - current_time = time.time() - polling_count += 1 - # Log a heartbeat message every 60 seconds to show the service is still working - if current_time - last_heartbeat_time >= 60: - logger.info(f"Queue reader heartbeat: Polled {config.SERVICE_BUS_DEPLOYMENT_STATUS_UPDATE_QUEUE} queue {polling_count} times in the last minute") - last_heartbeat_time = current_time - polling_count = 0 - async with credentials.get_credential_async_context() as credential: - service_bus_client = ServiceBusClient(config.SERVICE_BUS_FULLY_QUALIFIED_NAMESPACE, credential) - - logger.debug(f"Looking for new messages on {config.SERVICE_BUS_DEPLOYMENT_STATUS_UPDATE_QUEUE} queue...") - # max_wait_time=1 -> don't hold the session open after processing of the message has finished - async with service_bus_client.get_queue_receiver(queue_name=config.SERVICE_BUS_DEPLOYMENT_STATUS_UPDATE_QUEUE, max_wait_time=1, session_id=NEXT_AVAILABLE_SESSION) as receiver: - logger.info(f"Got a session containing messages: {receiver.session.session_id}") - async with AutoLockRenewer() as renewer: - renewer.register(receiver, receiver.session, max_lock_renewal_duration=60) - async for msg in receiver: - complete_message = await self.process_message(msg) - if complete_message: - await receiver.complete_message(msg) - else: - # could have been any kind of transient issue, we'll abandon back to the queue, and retry - await receiver.abandon_message(msg) - logger.info(f"Closing session: {receiver.session.session_id}") - - except OperationTimeoutError: - # Timeout occurred whilst connecting to a session - this is expected and indicates no non-empty sessions are available - logger.debug("No sessions for this process. Will look again...") + # We keep a single ServiceBusClient alive across the inner loop to avoid excessive connection + # and reconnection churn, as get_queue_receiver with NEXT_AVAILABLE_SESSION is polled frequently. + # Any fatal connection-related errors or other exceptions (other than OperationTimeoutError) + # will propagate out of the inner loop, closing this context manager and recreating the client. + async with ServiceBusClient(config.SERVICE_BUS_FULLY_QUALIFIED_NAMESPACE, credential) as service_bus_client: + client_created_time = time.time() + while True: + try: + # Recreate the client periodically (every hour) to ensure connection freshness + # and avoid holding a potentially stale client open indefinitely. + if time.time() - client_created_time > 3600: + logger.info("ServiceBusClient has been active for 1 hour. Recreating for freshness...") + break + + current_time = time.time() + polling_count += 1 + # Log a heartbeat message every 60 seconds to show the service is still working + if current_time - last_heartbeat_time >= 60: + logger.info(f"Queue reader heartbeat: Polled {config.SERVICE_BUS_DEPLOYMENT_STATUS_UPDATE_QUEUE} queue {polling_count} times in the last minute") + last_heartbeat_time = current_time + polling_count = 0 + + logger.debug(f"Looking for new messages on {config.SERVICE_BUS_DEPLOYMENT_STATUS_UPDATE_QUEUE} queue...") + # max_wait_time=1 -> don't hold the session open after processing of the message has finished + async with service_bus_client.get_queue_receiver(queue_name=config.SERVICE_BUS_DEPLOYMENT_STATUS_UPDATE_QUEUE, max_wait_time=1, session_id=NEXT_AVAILABLE_SESSION) as receiver: + logger.info(f"Got a session containing messages: {receiver.session.session_id}") + async with AutoLockRenewer() as renewer: + renewer.register(receiver, receiver.session, max_lock_renewal_duration=60) + async for msg in receiver: + complete_message = await self.process_message(msg) + if complete_message: + await receiver.complete_message(msg) + else: + # could have been any kind of transient issue, we'll abandon back to the queue, and retry + await receiver.abandon_message(msg) + logger.info(f"Closing session: {receiver.session.session_id}") + + except OperationTimeoutError: + # Timeout occurred whilst connecting to a session - this is expected and indicates no non-empty sessions are available + logger.debug("No sessions for this process. Will look again...") except ServiceBusConnectionError: # Occasionally there will be a transient / network-level error in connecting to SB. logger.info("Unknown Service Bus connection error. Will retry...") + await asyncio.sleep(10) + + except asyncio.CancelledError: + raise except Exception as e: # Catch all other exceptions, log them via .exception to get the stack trace, and reconnect logger.exception(f"Unknown exception. Will retry - {e}") + await asyncio.sleep(10) async def process_message(self, msg): complete_message = False diff --git a/api_app/tests_ma/test_service_bus/test_airlock_request_status_update.py b/api_app/tests_ma/test_service_bus/test_airlock_request_status_update.py index 6404ba122f..6c5463bf76 100644 --- a/api_app/tests_ma/test_service_bus/test_airlock_request_status_update.py +++ b/api_app/tests_ma/test_service_bus/test_airlock_request_status_update.py @@ -3,13 +3,19 @@ import pytest import time -from mock import AsyncMock, patch +from unittest.mock import AsyncMock, patch from service_bus.airlock_request_status_update import AirlockStatusUpdater from models.domain.events import AirlockNotificationUserData, AirlockFile from models.domain.airlock_request import AirlockRequest, AirlockRequestStatus, AirlockRequestType from models.domain.workspace import Workspace from db.errors import EntityDoesNotExist from resources import strings +from tests_ma.test_service_bus.test_helpers import ( + StopReceiveMessages, + credential_context, + queue_receiver_context, + service_bus_client_context, +) WORKSPACE_ID = "abc000d3-82da-4bfc-b6e9-9a7853ef753e" AIRLOCK_REQUEST_ID = "5dbc15ae-40e1-49a5-834b-595f59d626b7" @@ -104,6 +110,86 @@ def __str__(self): return self.message +async def run_receive_messages_with_mocks(service_bus_client, time_values, client_side_effect=None): + updater = AirlockStatusUpdater() + credential = credential_context() + receiver = queue_receiver_context(receive_messages=True) + service_bus_client.get_queue_receiver.return_value = receiver + + with patch("service_bus.airlock_request_status_update.credentials.get_credential_async_context", return_value=credential), \ + patch("service_bus.airlock_request_status_update.ServiceBusClient", return_value=service_bus_client, side_effect=client_side_effect), \ + patch("service_bus.airlock_request_status_update.time.time", side_effect=time_values), \ + patch("service_bus.airlock_request_status_update.asyncio.sleep", new_callable=AsyncMock): + await updater.receive_messages() + + +async def test_receive_messages_reuses_client_for_multiple_polls(): + service_bus_client = service_bus_client_context() + time_call_count = 0 + client_call_count = 0 + + def time_after_two_polls(): + nonlocal time_call_count + time_call_count += 1 + return 0 if time_call_count <= 5 else 3601 + + def create_client(*args, **kwargs): + nonlocal client_call_count + client_call_count += 1 + if client_call_count == 1: + return service_bus_client + raise StopReceiveMessages() + + with pytest.raises(StopReceiveMessages): + await run_receive_messages_with_mocks(service_bus_client, time_after_two_polls, create_client) + + assert service_bus_client.get_queue_receiver.call_count == 2 + service_bus_client.__aenter__.assert_awaited_once() + service_bus_client.__aexit__.assert_awaited_once() + + +async def test_receive_messages_closes_client_before_hourly_recreation(): + first_client = service_bus_client_context() + + def create_client(*args, **kwargs): + if create_client.called: + raise StopReceiveMessages() + create_client.called = True + return first_client + + create_client.called = False + + with pytest.raises(StopReceiveMessages): + with patch("service_bus.airlock_request_status_update.credentials.get_credential_async_context", return_value=credential_context()), \ + patch("service_bus.airlock_request_status_update.ServiceBusClient", side_effect=create_client), \ + patch("service_bus.airlock_request_status_update.time.time", side_effect=[0, 0, 0, 3601]), \ + patch("service_bus.airlock_request_status_update.asyncio.sleep", new_callable=AsyncMock): + first_client.get_queue_receiver.return_value = queue_receiver_context() + await AirlockStatusUpdater().receive_messages() + + first_client.__aexit__.assert_awaited_once() + assert first_client.get_queue_receiver.call_count == 1 + + +async def test_receive_messages_closes_client_after_receiver_failure(): + service_bus_client = service_bus_client_context() + service_bus_client.get_queue_receiver.side_effect = RuntimeError("receiver failed") + client_call_count = 0 + + def create_client(*args, **kwargs): + nonlocal client_call_count + client_call_count += 1 + if client_call_count == 1: + return service_bus_client + raise StopReceiveMessages() + + with pytest.raises(StopReceiveMessages): + await run_receive_messages_with_mocks(service_bus_client, lambda: 0, create_client) + + service_bus_client.__aenter__.assert_awaited_once() + service_bus_client.__aexit__.assert_awaited_once() + + @patch("event_grid.helpers.EventGridPublisherClient") @patch('service_bus.airlock_request_status_update.AirlockRequestRepository.create') @patch('service_bus.airlock_request_status_update.WorkspaceRepository.create') diff --git a/api_app/tests_ma/test_service_bus/test_deployment_status_update.py b/api_app/tests_ma/test_service_bus/test_deployment_status_update.py index 16a63937aa..f4562eb2cc 100644 --- a/api_app/tests_ma/test_service_bus/test_deployment_status_update.py +++ b/api_app/tests_ma/test_service_bus/test_deployment_status_update.py @@ -1,11 +1,10 @@ import copy import json -from unittest.mock import MagicMock, ANY +from unittest.mock import AsyncMock, MagicMock, ANY, patch from pydantic import TypeAdapter import pytest import uuid -from mock import AsyncMock, patch from tests_ma.test_api.test_routes.test_resource_helpers import FAKE_CREATE_TIMESTAMP, FAKE_UPDATE_TIMESTAMP from models.domain.request_action import RequestAction from models.domain.resource import ResourceType @@ -15,6 +14,12 @@ from models.domain.operation import DeploymentStatusUpdateMessage, Operation, OperationStep, Status from resources import strings from service_bus.deployment_status_updater import DeploymentStatusUpdater +from tests_ma.test_service_bus.test_helpers import ( + StopReceiveMessages, + credential_context, + queue_receiver_context, + service_bus_client_context, +) pytestmark = pytest.mark.asyncio @@ -79,6 +84,89 @@ def __str__(self): return self.message +async def run_receive_messages_with_mocks(service_bus_client, time_values, client_side_effect=None): + credential = credential_context() + receiver = queue_receiver_context(session=True, iterate=True) + service_bus_client.get_queue_receiver.return_value = receiver + renewer = MagicMock() + + with patch("service_bus.deployment_status_updater.credentials.get_credential_async_context", return_value=credential), \ + patch("service_bus.deployment_status_updater.ServiceBusClient", return_value=service_bus_client, side_effect=client_side_effect), \ + patch("service_bus.deployment_status_updater.time.time", side_effect=time_values), \ + patch("service_bus.deployment_status_updater.asyncio.sleep", new_callable=AsyncMock), \ + patch("service_bus.deployment_status_updater.AutoLockRenewer") as auto_lock_renewer: + auto_lock_renewer.return_value.__aenter__ = AsyncMock(return_value=renewer) + auto_lock_renewer.return_value.__aexit__ = AsyncMock(return_value=False) + await DeploymentStatusUpdater().receive_messages() + + +async def test_receive_messages_reuses_client_for_multiple_polls(): + service_bus_client = service_bus_client_context() + time_call_count = 0 + client_call_count = 0 + + def time_after_two_polls(): + nonlocal time_call_count + time_call_count += 1 + return 0 if time_call_count <= 5 else 3601 + + def create_client(*args, **kwargs): + nonlocal client_call_count + client_call_count += 1 + if client_call_count == 1: + return service_bus_client + raise StopReceiveMessages() + + with pytest.raises(StopReceiveMessages): + await run_receive_messages_with_mocks(service_bus_client, time_after_two_polls, create_client) + + assert service_bus_client.get_queue_receiver.call_count == 2 + service_bus_client.__aenter__.assert_awaited_once() + service_bus_client.__aexit__.assert_awaited_once() + + +async def test_receive_messages_closes_client_before_hourly_recreation(): + first_client = service_bus_client_context() + + def create_client(*args, **kwargs): + if create_client.called: + raise StopReceiveMessages() + create_client.called = True + return first_client + + create_client.called = False + + with pytest.raises(StopReceiveMessages): + with patch("service_bus.deployment_status_updater.credentials.get_credential_async_context", return_value=credential_context()), \ + patch("service_bus.deployment_status_updater.ServiceBusClient", side_effect=create_client), \ + patch("service_bus.deployment_status_updater.time.time", side_effect=[0, 0, 0, 3601]), \ + patch("service_bus.deployment_status_updater.asyncio.sleep", new_callable=AsyncMock): + first_client.get_queue_receiver.return_value = queue_receiver_context() + await DeploymentStatusUpdater().receive_messages() + + first_client.__aexit__.assert_awaited_once() + assert first_client.get_queue_receiver.call_count == 1 + + +async def test_receive_messages_closes_client_after_receiver_failure(): + service_bus_client = service_bus_client_context() + service_bus_client.get_queue_receiver.side_effect = RuntimeError("receiver failed") + client_call_count = 0 + + def create_client(*args, **kwargs): + nonlocal client_call_count + client_call_count += 1 + if client_call_count == 1: + return service_bus_client + raise StopReceiveMessages() + + with pytest.raises(StopReceiveMessages): + await run_receive_messages_with_mocks(service_bus_client, lambda: 0, create_client) + + service_bus_client.__aenter__.assert_awaited_once() + service_bus_client.__aexit__.assert_awaited_once() + + def create_sample_workspace_object(workspace_id): return Workspace( id=workspace_id, diff --git a/api_app/tests_ma/test_service_bus/test_helpers.py b/api_app/tests_ma/test_service_bus/test_helpers.py new file mode 100644 index 0000000000..a05c0a5b39 --- /dev/null +++ b/api_app/tests_ma/test_service_bus/test_helpers.py @@ -0,0 +1,32 @@ +from unittest.mock import AsyncMock, MagicMock + + +class StopReceiveMessages(BaseException): + pass + + +def service_bus_client_context(): + client = MagicMock() + client.__aenter__ = AsyncMock(return_value=client) + client.__aexit__ = AsyncMock(return_value=False) + return client + + +def credential_context(): + context = MagicMock() + context.__aenter__ = AsyncMock(return_value=MagicMock()) + context.__aexit__ = AsyncMock(return_value=False) + return context + + +def queue_receiver_context(*, session=False, receive_messages=False, iterate=False): + receiver = MagicMock() + receiver.__aenter__ = AsyncMock(return_value=receiver) + receiver.__aexit__ = AsyncMock(return_value=False) + if session: + receiver.session.session_id = "test_session_id" + if receive_messages: + receiver.receive_messages = AsyncMock(return_value=[]) + if iterate: + receiver.__aiter__.return_value = [] + return receiver diff --git a/api_app/tests_ma/test_service_bus/test_resource_request_sender.py b/api_app/tests_ma/test_service_bus/test_resource_request_sender.py index 5f51770f02..22db82ded6 100644 --- a/api_app/tests_ma/test_service_bus/test_resource_request_sender.py +++ b/api_app/tests_ma/test_service_bus/test_resource_request_sender.py @@ -4,7 +4,7 @@ import uuid from azure.servicebus import ServiceBusMessage -from mock import AsyncMock, patch +from unittest.mock import AsyncMock, patch from resources import strings from models.schemas.resource import ResourcePatch from service_bus.helpers import ( diff --git a/resource_processor/_version.py b/resource_processor/_version.py index e318db3960..af935ca6d9 100644 --- a/resource_processor/_version.py +++ b/resource_processor/_version.py @@ -1 +1 @@ -__version__ = "0.13.6" +__version__ = "0.13.7" diff --git a/resource_processor/tests_rp/test_runner.py b/resource_processor/tests_rp/test_runner.py index 6c9166b017..4f5c0b5e89 100644 --- a/resource_processor/tests_rp/test_runner.py +++ b/resource_processor/tests_rp/test_runner.py @@ -48,8 +48,11 @@ async def test_set_up_config(mock_get_config): async def setup_service_bus_client_and_credential(mock_service_bus_client, mock_default_credential, msi_id): mock_credential = AsyncMock() - mock_default_credential.return_value.__aenter__.return_value = mock_credential + mock_default_credential.return_value.__aenter__ = AsyncMock(return_value=mock_credential) + mock_default_credential.return_value.__aexit__ = AsyncMock(return_value=False) mock_service_bus_client_instance = mock_service_bus_client.return_value + mock_service_bus_client.return_value.__aenter__ = AsyncMock(return_value=mock_service_bus_client_instance) + mock_service_bus_client.return_value.__aexit__ = AsyncMock(return_value=False) return mock_service_bus_client_instance, mock_credential @@ -67,6 +70,12 @@ async def test_runner(mock_receive_message, mock_service_bus_client, mock_defaul mock_service_bus_client.assert_called_once_with("test_namespace", mock_credential) mock_receive_message.assert_called_once_with(mock_service_bus_client_instance, config) + # Verify context manager entered and exited cleanly + mock_default_credential.return_value.__aenter__.assert_awaited_once() + mock_default_credential.return_value.__aexit__.assert_awaited_once() + mock_service_bus_client.return_value.__aenter__.assert_awaited_once() + mock_service_bus_client.return_value.__aexit__.assert_awaited_once() + @pytest.mark.asyncio @patch("vmss_porter.runner.receive_message") @@ -82,6 +91,12 @@ async def test_runner_no_msi_id(mock_receive_message, mock_service_bus_client, m mock_service_bus_client.assert_called_once_with("test_namespace", mock_credential) mock_receive_message.assert_called_once_with(mock_service_bus_client_instance, config) + # Verify context manager entered and exited cleanly + mock_default_credential.return_value.__aenter__.assert_awaited_once() + mock_default_credential.return_value.__aexit__.assert_awaited_once() + mock_service_bus_client.return_value.__aenter__.assert_awaited_once() + mock_service_bus_client.return_value.__aexit__.assert_awaited_once() + @pytest.mark.asyncio @patch("vmss_porter.runner.receive_message") @@ -99,6 +114,28 @@ async def test_runner_exception(mock_receive_message, mock_service_bus_client, m mock_service_bus_client.assert_called_once_with("test_namespace", mock_credential) mock_receive_message.assert_called_once_with(mock_service_bus_client_instance, config) + # Verify context manager entered and exited cleanly, even on exception + mock_default_credential.return_value.__aenter__.assert_awaited_once() + mock_default_credential.return_value.__aexit__.assert_awaited_once() + mock_service_bus_client.return_value.__aenter__.assert_awaited_once() + mock_service_bus_client.return_value.__aexit__.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_default_credentials_closes_real_credential_on_exception(mock_service_bus_client): + mock_credential = AsyncMock() + mock_service_bus_client_instance = mock_service_bus_client.return_value + mock_service_bus_client.return_value.__aenter__.return_value = mock_service_bus_client_instance + + with patch("vmss_porter.runner.DefaultAzureCredential", return_value=mock_credential), \ + patch("vmss_porter.runner.receive_message", side_effect=Exception("Test Exception")): + config = {"vmss_msi_id": "test_msi_id", "service_bus_namespace": "test_namespace"} + + with pytest.raises(Exception, match="Test Exception"): + await runner(0, config) + + mock_credential.close.assert_awaited_once() + @pytest.mark.asyncio @patch("vmss_porter.runner.invoke_porter_action", return_value=True) diff --git a/resource_processor/vmss_porter/runner.py b/resource_processor/vmss_porter/runner.py index 120ececba0..46dc1d233c 100644 --- a/resource_processor/vmss_porter/runner.py +++ b/resource_processor/vmss_porter/runner.py @@ -33,8 +33,10 @@ async def default_credentials(msi_id): Context manager which yields the default credentials. """ credential = DefaultAzureCredential(managed_identity_client_id=msi_id) if msi_id else DefaultAzureCredential() - yield credential - await credential.close() + try: + yield credential + finally: + await credential.close() async def receive_message(service_bus_client, config: dict, keep_running=lambda: True): @@ -100,11 +102,15 @@ async def receive_message(service_bus_client, config: dict, keep_running=lambda: except ServiceBusConnectionError: # Occasionally there will be a transient / network-level error in connecting to SB. logger.info("Unknown Service Bus connection error. Will retry...") + await asyncio.sleep(10) + + except asyncio.CancelledError: + raise except Exception: # Catch all other exceptions, log them via .exception to get the stack trace, sleep, and reconnect - logger.exception("Unknown exception. Will retry...") + await asyncio.sleep(10) async def run_porter(command_parts_list: list, config: dict): @@ -278,8 +284,8 @@ async def get_porter_outputs(msg_body: dict, config: dict): async def runner(process_number: int, config: dict): with tracer.start_as_current_span(process_number): async with default_credentials(config["vmss_msi_id"]) as credential: - service_bus_client = ServiceBusClient(config["service_bus_namespace"], credential) - await receive_message(service_bus_client, config) + async with ServiceBusClient(config["service_bus_namespace"], credential) as service_bus_client: + await receive_message(service_bus_client, config) async def check_runners(processes: list, httpserver: Process, keep_running=lambda: True):