From 50d8534f084c826f916136bbe2ca90c313a81a10 Mon Sep 17 00:00:00 2001 From: Andy Twigg Date: Sat, 11 Jul 2026 03:06:56 -0700 Subject: [PATCH] Increase JAX IFRT proxy client connection timeout to 10 minutes. PiperOrigin-RevId: 946138871 --- pathwaysutils/proxy_backend.py | 18 ++++++++--- pathwaysutils/test/proxy_backend_test.py | 41 ++++++++++++++++++++++++ 2 files changed, 55 insertions(+), 4 deletions(-) diff --git a/pathwaysutils/proxy_backend.py b/pathwaysutils/proxy_backend.py index cf1f806..2507a8a 100644 --- a/pathwaysutils/proxy_backend.py +++ b/pathwaysutils/proxy_backend.py @@ -13,17 +13,27 @@ # limitations under the License. """Register the IFRT Proxy as a backend for JAX.""" +import os import jax from jax.extend import backend from jax.extend.backend import ifrt_proxy def register_backend_factory() -> None: + """Registers the IFRT Proxy backend factory with JAX.""" + + def make_client(): + options = ifrt_proxy.ClientConnectionOptions() + timeout_secs = os.environ.get("PATHWAYS_PROXY_CONNECTION_TIMEOUT_SECS") + if timeout_secs: + options.connection_timeout_in_seconds = int(timeout_secs) + return ifrt_proxy.get_client( + jax.config.read("jax_backend_target"), + options, + ) + backend.register_backend_factory( "proxy", - lambda: ifrt_proxy.get_client( - jax.config.read("jax_backend_target"), - ifrt_proxy.ClientConnectionOptions(), - ), + make_client, priority=-1, ) diff --git a/pathwaysutils/test/proxy_backend_test.py b/pathwaysutils/test/proxy_backend_test.py index fb8ad8c..13fce79 100644 --- a/pathwaysutils/test/proxy_backend_test.py +++ b/pathwaysutils/test/proxy_backend_test.py @@ -13,6 +13,7 @@ # limitations under the License. """Tests for the proxy backend module.""" +import os from unittest import mock from absl.testing import absltest @@ -54,6 +55,46 @@ def test_proxy_backend_registration(self): proxy_backend.register_backend_factory() self.assertIn("proxy", backend.backends()) + def test_proxy_backend_registration_with_timeout(self): + mock_get_client = self.enter_context( + mock.patch.object( + ifrt_proxy, + "get_client", + return_value=mock.MagicMock(), + ) + ) + self.enter_context( + mock.patch.dict( + os.environ, {"PATHWAYS_PROXY_CONNECTION_TIMEOUT_SECS": "42"} + ) + ) + proxy_backend.register_backend_factory() + self.assertIn("proxy", backend.backends()) + mock_get_client.assert_called_once() + args, _ = mock_get_client.call_args + self.assertEqual(args[0], "grpc://localhost:12345") + options = args[1] + self.assertEqual(options.connection_timeout_in_seconds, 42) + + def test_proxy_backend_registration_without_timeout(self): + mock_get_client = self.enter_context( + mock.patch.object( + ifrt_proxy, + "get_client", + return_value=mock.MagicMock(), + ) + ) + self.enter_context(mock.patch.dict(os.environ)) + os.environ.pop("PATHWAYS_PROXY_CONNECTION_TIMEOUT_SECS", None) + + proxy_backend.register_backend_factory() + self.assertIn("proxy", backend.backends()) + mock_get_client.assert_called_once() + args, _ = mock_get_client.call_args + self.assertEqual(args[0], "grpc://localhost:12345") + options = args[1] + self.assertNotEqual(options.connection_timeout_in_seconds, 42) + if __name__ == "__main__": absltest.main()