Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
18 changes: 14 additions & 4 deletions pathwaysutils/proxy_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
)
41 changes: 41 additions & 0 deletions pathwaysutils/test/proxy_backend_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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()
Loading