diff --git a/dapr/aio/clients/health.py b/dapr/aio/clients/health.py index 9ab66ebba..1733a25ea 100644 --- a/dapr/aio/clients/health.py +++ b/dapr/aio/clients/health.py @@ -38,8 +38,14 @@ async def wait_for_sidecar(): connector = aiohttp.TCPConnector(ssl=ssl_context) async with aiohttp.ClientSession(connector=connector) as session: while True: + # Bound each request by the time left: aiohttp's default total timeout + # (300 s) lets one unanswered request outlive the deadline. + request_timeout_seconds = max((start + timeout) - time.time(), 1.0) + request_timeout = aiohttp.ClientTimeout(total=request_timeout_seconds) try: - async with session.get(health_url, headers=headers) as response: + async with session.get( + health_url, headers=headers, timeout=request_timeout + ) as response: if 200 <= response.status < 300: break except aiohttp.ClientError as e: diff --git a/dapr/clients/health.py b/dapr/clients/health.py index 8e1002292..876c2d659 100644 --- a/dapr/clients/health.py +++ b/dapr/clients/health.py @@ -43,9 +43,15 @@ def wait_for_sidecar(): start = time.time() while True: + # Bound each request by the time left: without a timeout, an endpoint that + # accepts the connection but never answers blocks past the deadline forever. + request_timeout = max((start + timeout) - time.time(), 1.0) try: req = urllib.request.Request(health_url, headers=headers) - with urllib.request.urlopen(req, context=DaprHealth.get_ssl_context()) as response: + ssl_context = DaprHealth.get_ssl_context() + with urllib.request.urlopen( + req, context=ssl_context, timeout=request_timeout + ) as response: if 200 <= response.status < 300: break except urllib.error.URLError as e: diff --git a/tests/clients/test_healthcheck.py b/tests/clients/test_healthcheck.py index c5b49aee9..897e70e3f 100644 --- a/tests/clients/test_healthcheck.py +++ b/tests/clients/test_healthcheck.py @@ -13,6 +13,8 @@ limitations under the License. """ +import socket +import threading import time import unittest from unittest.mock import MagicMock, patch @@ -75,3 +77,28 @@ def test_wait_for_sidecar_timeout(self, mock_urlopen): self.assertGreaterEqual(time.time() - start, 2.5) self.assertGreater(mock_urlopen.call_count, 1) + + @patch.object(settings, 'DAPR_HEALTH_TIMEOUT', '1') + def test_wait_for_sidecar_timeout_when_endpoint_never_responds(self): + # The listener never calls accept(): the TCP connect succeeds through the + # backlog, but the HTTP request never gets a response. + with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as listener: + listener.bind(('127.0.0.1', 0)) + listener.listen() + port = listener.getsockname()[1] + errors: list[Exception] = [] + + def wait() -> None: + try: + DaprHealth.wait_for_sidecar() + except Exception as e: + errors.append(e) + + with patch.object(settings, 'DAPR_HTTP_ENDPOINT', f'http://127.0.0.1:{port}'): + waiter = threading.Thread(target=wait, daemon=True) + waiter.start() + waiter.join(timeout=10) + + self.assertFalse(waiter.is_alive(), 'wait_for_sidecar() is still blocked') + self.assertEqual(len(errors), 1) + self.assertIsInstance(errors[0], TimeoutError) diff --git a/tests/clients/test_healthcheck_async.py b/tests/clients/test_healthcheck_async.py index 5f497bb8d..66f6a8d22 100644 --- a/tests/clients/test_healthcheck_async.py +++ b/tests/clients/test_healthcheck_async.py @@ -14,6 +14,7 @@ """ import asyncio +import socket import time import unittest from unittest.mock import AsyncMock, MagicMock, patch @@ -192,6 +193,23 @@ async def test_multiple_health_checks_concurrent(self, mock_get): # Verify multiple calls were made self.assertGreaterEqual(mock_get.call_count, 3) + @patch.object(settings, 'DAPR_HEALTH_TIMEOUT', '1') + async def test_wait_for_sidecar_timeout_when_endpoint_never_responds(self): + # The listener never calls accept(): the TCP connect succeeds through the + # backlog, but the HTTP request never gets a response. + with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as listener: + listener.bind(('127.0.0.1', 0)) + listener.listen() + port = listener.getsockname()[1] + + with patch.object(settings, 'DAPR_HTTP_ENDPOINT', f'http://127.0.0.1:{port}'): + start = time.time() + with self.assertRaises(TimeoutError) as raised: + await asyncio.wait_for(DaprHealth.wait_for_sidecar(), timeout=10) + + self.assertLess(time.time() - start, 10) + self.assertIn('Dapr health check timed out', str(raised.exception)) + if __name__ == '__main__': unittest.main()