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
Original file line number Diff line number Diff line change
Expand Up @@ -50,7 +50,7 @@ def test_resolve_self_hostname_missing_env(self):
with self.assertRaises(ValueError):
k8s_context.resolve_self_hostname()

def test_k8s_jax_context_pathways(self):
def test_jax_context_pathways(self):
envs = {
"JAX_PLATFORMS": "proxy",
"JAX_BACKEND_TARGET": "0.0.0.0:8000",
Expand All @@ -61,15 +61,15 @@ def test_k8s_jax_context_pathways(self):
k8s_context.K8sJaxContext().initialize()
mock_pw.initialize.assert_called_once()

def test_k8s_jax_context_mcjax(self):
def test_jax_context_mcjax(self):
envs = {}
mock_jax = mock.MagicMock()
with mock.patch.dict(os.environ, envs, clear=True):
with mock.patch.dict(sys.modules, {"jax": mock_jax}):
k8s_context.K8sJaxContext().initialize()
mock_jax.distributed.initialize.assert_called_once()

def test_k8s_discovery_context_register(self):
def test_discovery_context_register(self):
envs = {
"JOBSET_NAME": "myjobset",
"REPLICATED_JOB_NAME": "worker",
Expand All @@ -95,7 +95,68 @@ def test_k8s_discovery_context_register(self):
b"pod-meta",
)

def test_k8s_process_context(self):
@mock.patch(
"tunix.experimental.distributed.runtime.contexts.k8s_context.discovery.connect"
)
@mock.patch(
"tunix.experimental.distributed.runtime.contexts.k8s_context.discovery.DiscoveryServer"
)
def test_discovery_context_connect(self, mock_server_cls, mock_connect):
envs = {
"JOBSET_NAME": "myjobset",
"REPLICATED_JOB_NAME": "worker",
"JOB_INDEX": "0",
"POD_INDEX": "1",
}
port = 8888
args = argparse.Namespace(
discovery_port=port,
discovery_addrs=f"door:{port}",
discovery_id="door",
)

mock_server = mock_server_cls.return_value

with mock.patch.dict(os.environ, envs):
with k8s_context.K8sDiscoveryContext(args) as disc_ctx:
on_client_connected = lambda cid, h, p, m, rec: None
on_client_disconnected = lambda cid, h, p, r: None

disc_ctx.on_connect(
on_client_connected=on_client_connected,
on_client_disconnected=on_client_disconnected,
)
mock_server.on_connect.assert_called_once_with(
on_client_connected=on_client_connected,
on_client_disconnected=on_client_disconnected,
)
mock_server.start.assert_called_once_with(port)

on_connected = lambda epoch, rec: None
on_disconnected = lambda epoch, r: None

client = disc_ctx.connect(
b"pod-meta",
client_id="door",
on_connected=on_connected,
on_disconnected=on_disconnected,
)
mock_connect.assert_called_once_with(
"door-proc-0-0.door:8888",
"myjobset-worker-0-1.myjobset",
port,
b"pod-meta",
client_id="door",
on_connected=on_connected,
on_disconnected=on_disconnected,
)
self.assertEqual(disc_ctx._client, mock_connect.return_value)

mock_connect.return_value.stop.assert_called_once()
mock_server.stop.assert_called_once()
self.assertIsNone(disc_ctx._client)

def test_process_context(self):
args = argparse.Namespace(
discovery_port=portpicker.pick_unused_port(),
discovery_addrs="door:8888",
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -32,9 +32,7 @@ def test_resolve_discovery_address(self):
@mock.patch(
"tunix.experimental.distributed.runtime.discovery.discovery.grpc.server"
)
def test_local_discovery_context_lifecycle_and_registration(
self, mock_grpc_server
):
def test_discovery_context_register(self, mock_grpc_server):
port = portpicker.pick_unused_port()
args = argparse.Namespace(
discovery_port=port,
Expand All @@ -56,7 +54,61 @@ def test_local_discovery_context_lifecycle_and_registration(

self.assertFalse(disc_ctx._server.is_started())

def test_local_process_context(self):
@mock.patch(
"tunix.experimental.distributed.runtime.contexts.local_context.discovery.connect"
)
@mock.patch(
"tunix.experimental.distributed.runtime.contexts.local_context.discovery.DiscoveryServer"
)
def test_discovery_context_connect(self, mock_server_cls, mock_connect):
port = 8888
args = argparse.Namespace(
discovery_port=port,
discovery_addrs=f"leader:{port}",
discovery_id="worker-0",
)

mock_server = mock_server_cls.return_value

with local_context.LocalDiscoveryContext(args) as disc_ctx:
on_client_connected = lambda cid, h, p, m, rec: None
on_client_disconnected = lambda cid, h, p, r: None

disc_ctx.on_connect(
on_client_connected=on_client_connected,
on_client_disconnected=on_client_disconnected,
)
mock_server.on_connect.assert_called_once_with(
on_client_connected=on_client_connected,
on_client_disconnected=on_client_disconnected,
)
mock_server.start.assert_called_once_with(port)

on_connected = lambda epoch, rec: None
on_disconnected = lambda epoch, r: None

client = disc_ctx.connect(
b"my-metadata",
client_id="worker-0",
on_connected=on_connected,
on_disconnected=on_disconnected,
)
mock_connect.assert_called_once_with(
"localhost:8888",
"localhost",
port,
b"my-metadata",
client_id="worker-0",
on_connected=on_connected,
on_disconnected=on_disconnected,
)
self.assertEqual(disc_ctx._client, mock_connect.return_value)

mock_connect.return_value.stop.assert_called_once()
mock_server.stop.assert_called_once()
self.assertIsNone(disc_ctx._client)

def test_process_context(self):
args = argparse.Namespace(
discovery_port=portpicker.pick_unused_port(),
discovery_addrs="leader:9999",
Expand Down
159 changes: 155 additions & 4 deletions tests/experimental/distributed/runtime/discovery/discovery_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,19 +25,26 @@

class DiscoveryTest(absltest.TestCase):

def test_start_unconfigured_mode_raises(self):
server = discovery.DiscoveryServer()
with self.assertRaises(RuntimeError):
server.start(8888)

def test_start_with_zero_port_raises(self):
server = discovery.DiscoveryServer()
server.on_register(lambda h, p, m: None)
with self.assertRaises(ValueError):
server.start(0, lambda h, p, m: None)
server.start(0)

@mock.patch.object(grpc, "server")
def test_start_twice_raises(self, mock_grpc_server):
server = discovery.DiscoveryServer()
port = 8888
server.start(port, lambda h, p, m: None)
server.on_register(lambda h, p, m: None)
server.start(port)
try:
with self.assertRaises(RuntimeError):
server.start(port, lambda h, p, m: None)
server.start(port)
finally:
server.stop()

Expand All @@ -52,7 +59,8 @@ def callback(hostname, p, metadata):
received["port"] = p
received["metadata"] = metadata

server.start(port, callback)
server.on_register(callback)
server.start(port)
try:
self.assertTrue(server.is_started())
mock_grpc_server.return_value.add_insecure_port.assert_called_once_with(
Expand Down Expand Up @@ -102,6 +110,149 @@ def test_register_non_retryable_error_raises(self, mock_channel):
with self.assertRaises(RuntimeError):
discovery.register("localhost:9999", "node-0", 1234, b"meta")

def test_connect_initial_connection(self):
port = portpicker.pick_unused_port()
server = discovery.DiscoveryServer(heartbeat_sec=1)
server_connected_events = []

server.on_connect(
on_client_connected=lambda cid, h, p, m, rec: server_connected_events.append(
(cid, h, p, m, rec)
)
)
server.start(port, heartbeat_sec=1)

client_connected_events = []
try:
client = discovery.connect(
f"localhost:{port}",
"node-0",
1234,
b"meta-data",
client_id="node-0",
on_connected=lambda epoch, rec: client_connected_events.append(
(epoch, rec)
),
)

self.assertEqual(len(client_connected_events), 1)
epoch, is_reconnect = client_connected_events[0]
self.assertFalse(is_reconnect)

self.assertEqual(len(server_connected_events), 1)
cid, h, p, m, is_rec = server_connected_events[0]
self.assertEqual(cid, "node-0")
self.assertFalse(is_rec)

client.stop()
finally:
server.stop()

def test_connect_reconnect_on_server_restart(self):
port = portpicker.pick_unused_port()
server = discovery.DiscoveryServer(heartbeat_sec=1)
server_connected_events = []

server.on_connect(
on_client_connected=lambda cid, h, p, m, rec: server_connected_events.append(
(cid, h, p, m, rec)
)
)
server.start(port, heartbeat_sec=1)

client_connected_events = []
client_disconnected_events = []
reconnected_event = threading.Event()

def on_connected(epoch, rec):
client_connected_events.append((epoch, rec))
if rec:
reconnected_event.set()

try:
client = discovery.connect(
f"localhost:{port}",
"node-0",
1234,
b"meta-data",
client_id="node-0",
on_connected=on_connected,
on_disconnected=lambda epoch, reason: client_disconnected_events.append(
(epoch, reason)
),
)

initial_epoch = client_connected_events[0][0]

# Simulate server restart by changing servicer epoch
server._servicer._server_epoch = "rebooted-epoch-1234"

# Wait for heartbeat loop to detect epoch mismatch and reconnect
self.assertTrue(reconnected_event.wait(timeout=5.0))

self.assertGreaterEqual(len(client_disconnected_events), 1)
self.assertEqual(client_disconnected_events[0][0], initial_epoch)
self.assertEqual(client_disconnected_events[0][1], "epoch_mismatch")

self.assertGreaterEqual(len(client_connected_events), 2)
new_epoch, is_reconnected = client_connected_events[1]
self.assertTrue(is_reconnected)
self.assertEqual(new_epoch, "rebooted-epoch-1234")

self.assertGreaterEqual(len(server_connected_events), 2)
self.assertTrue(server_connected_events[1][4]) # is_reconnect=True

client.stop()
finally:
server.stop()

def test_server_lease_eviction_on_heartbeat_timeout(self):
port = portpicker.pick_unused_port()
server = discovery.DiscoveryServer(heartbeat_sec=1)
server_disconnected_events = []
evicted_event = threading.Event()

def on_disconnected(cid, h, p, reason):
server_disconnected_events.append((cid, h, p, reason))
evicted_event.set()

server.on_connect(on_client_disconnected=on_disconnected)
server.start(port, heartbeat_sec=1)

try:
client = discovery.connect(
f"localhost:{port}",
"node-0",
1234,
b"meta-data",
client_id="node-0",
)
# Stop client heartbeat thread prematurely to simulate crashed/dead client
client._stop_event.set()

# Wait for server eviction loop (threshold = 3 * heartbeat_sec)
self.assertTrue(evicted_event.wait(timeout=5.0))

self.assertGreaterEqual(len(server_disconnected_events), 1)
cid, h, p, reason = server_disconnected_events[0]
self.assertEqual(cid, "node-0")
self.assertEqual(reason, "heartbeat_timeout")

client.stop()
finally:
server.stop()

def test_mode_mutual_exclusion(self):
server = discovery.DiscoveryServer()
server.on_register(lambda h, p, m: None)
with self.assertRaises(RuntimeError):
server.on_connect(lambda cid, h, p, m, rec: None)

server2 = discovery.DiscoveryServer()
server2.on_connect(lambda cid, h, p, m, rec: None)
with self.assertRaises(RuntimeError):
server2.on_register(lambda h, p, m: None)


if __name__ == "__main__":
absltest.main()
Loading
Loading