|
13 | 13 | # limitations under the License. |
14 | 14 | """Tests for the proxy backend module.""" |
15 | 15 |
|
| 16 | +import os |
16 | 17 | from unittest import mock |
17 | 18 |
|
18 | 19 | from absl.testing import absltest |
@@ -54,6 +55,46 @@ def test_proxy_backend_registration(self): |
54 | 55 | proxy_backend.register_backend_factory() |
55 | 56 | self.assertIn("proxy", backend.backends()) |
56 | 57 |
|
| 58 | + def test_proxy_backend_registration_with_timeout(self): |
| 59 | + mock_get_client = self.enter_context( |
| 60 | + mock.patch.object( |
| 61 | + ifrt_proxy, |
| 62 | + "get_client", |
| 63 | + return_value=mock.MagicMock(), |
| 64 | + ) |
| 65 | + ) |
| 66 | + self.enter_context( |
| 67 | + mock.patch.dict( |
| 68 | + os.environ, {"PATHWAYS_PROXY_CONNECTION_TIMEOUT_SECS": "42"} |
| 69 | + ) |
| 70 | + ) |
| 71 | + proxy_backend.register_backend_factory() |
| 72 | + self.assertIn("proxy", backend.backends()) |
| 73 | + mock_get_client.assert_called_once() |
| 74 | + args, _ = mock_get_client.call_args |
| 75 | + self.assertEqual(args[0], "grpc://localhost:12345") |
| 76 | + options = args[1] |
| 77 | + self.assertEqual(options.connection_timeout_in_seconds, 42) |
| 78 | + |
| 79 | + def test_proxy_backend_registration_without_timeout(self): |
| 80 | + mock_get_client = self.enter_context( |
| 81 | + mock.patch.object( |
| 82 | + ifrt_proxy, |
| 83 | + "get_client", |
| 84 | + return_value=mock.MagicMock(), |
| 85 | + ) |
| 86 | + ) |
| 87 | + self.enter_context(mock.patch.dict(os.environ)) |
| 88 | + os.environ.pop("PATHWAYS_PROXY_CONNECTION_TIMEOUT_SECS", None) |
| 89 | + |
| 90 | + proxy_backend.register_backend_factory() |
| 91 | + self.assertIn("proxy", backend.backends()) |
| 92 | + mock_get_client.assert_called_once() |
| 93 | + args, _ = mock_get_client.call_args |
| 94 | + self.assertEqual(args[0], "grpc://localhost:12345") |
| 95 | + options = args[1] |
| 96 | + self.assertNotEqual(options.connection_timeout_in_seconds, 42) |
| 97 | + |
57 | 98 |
|
58 | 99 | if __name__ == "__main__": |
59 | 100 | absltest.main() |
0 commit comments