From 32a68faa37e942157dc1673b61e9bac9564386d8 Mon Sep 17 00:00:00 2001 From: Akanksha Gupta Date: Mon, 27 Jul 2026 14:02:17 -0700 Subject: [PATCH] Move Shared Pathways Service tests to GitHub PiperOrigin-RevId: 954816972 --- .../deploy_pathways_service_test.py | 189 ++++ .../shared_pathways_service/gke_utils_test.py | 962 ++++++++++++++++++ .../isc_pathways_test.py | 852 ++++++++++++++++ .../metrics_collector_test.py | 247 +++++ .../run_connect_example_test.py | 135 +++ .../run_workload_test.py | 206 ++++ .../start_vscode_on_cpu_np_test.py | 324 ++++++ .../shared_pathways_service/tpu_specs_test.py | 77 ++ .../validators_test.py | 370 +++++++ 9 files changed, 3362 insertions(+) create mode 100644 pathwaysutils/test/experimental/shared_pathways_service/deploy_pathways_service_test.py create mode 100644 pathwaysutils/test/experimental/shared_pathways_service/gke_utils_test.py create mode 100644 pathwaysutils/test/experimental/shared_pathways_service/isc_pathways_test.py create mode 100644 pathwaysutils/test/experimental/shared_pathways_service/metrics_collector_test.py create mode 100644 pathwaysutils/test/experimental/shared_pathways_service/run_connect_example_test.py create mode 100644 pathwaysutils/test/experimental/shared_pathways_service/run_workload_test.py create mode 100644 pathwaysutils/test/experimental/shared_pathways_service/start_vscode_on_cpu_np_test.py create mode 100644 pathwaysutils/test/experimental/shared_pathways_service/tpu_specs_test.py create mode 100644 pathwaysutils/test/experimental/shared_pathways_service/validators_test.py diff --git a/pathwaysutils/test/experimental/shared_pathways_service/deploy_pathways_service_test.py b/pathwaysutils/test/experimental/shared_pathways_service/deploy_pathways_service_test.py new file mode 100644 index 0000000..a541175 --- /dev/null +++ b/pathwaysutils/test/experimental/shared_pathways_service/deploy_pathways_service_test.py @@ -0,0 +1,189 @@ +"""Unit tests for the deploy_pathways_service script.""" + +from unittest import mock +from absl import flags +from absl.testing import absltest +from absl.testing import parameterized +from pathwaysutils.experimental.shared_pathways_service import deploy_pathways_service + + +class DeployPathwaysServiceTest(parameterized.TestCase): + + @parameterized.named_parameters( + dict( + testcase_name="v5p", + tpu_type="v5p", + expected_machine_type="ct5p-hightpu-4t", + ), + dict( + testcase_name="v5e", + tpu_type="v5e", + expected_machine_type="ct5lp-hightpu-4t", + ), + dict( + testcase_name="v6e", + tpu_type="v6e", + expected_machine_type="ct6e-standard-4t", + ), + dict( + testcase_name="tpu7x", + tpu_type="tpu7x", + expected_machine_type="tpu7x-standard-4t", + ), + ) + def test_get_tpu_config_valid(self, tpu_type, expected_machine_type): + config = deploy_pathways_service.get_tpu_config(tpu_type) + self.assertEqual(config.machine_type, expected_machine_type) + + @parameterized.named_parameters( + dict( + testcase_name="invalid_tpu_type", + tpu_type="invalid", + ), + dict( + testcase_name="empty_tpu_type", + tpu_type="", + ), + dict( + testcase_name="whitespace_tpu_type", + tpu_type=" ", + ), + dict( + testcase_name="v5_tpu_type", + tpu_type="v5", + ), + ) + def test_get_tpu_config_invalid(self, tpu_type): + with self.assertRaises(ValueError): + deploy_pathways_service.get_tpu_config(tpu_type) + + def test_calculate_vms_per_slice_valid(self): + vms = deploy_pathways_service.calculate_vms_per_slice("4x8", 4) + self.assertEqual(vms, 8) + + def test_calculate_vms_per_slice_invalid_format(self): + with self.assertRaises(ValueError): + deploy_pathways_service.calculate_vms_per_slice("4x8x", 4) + + def test_calculate_vms_per_slice_not_divisible(self): + with self.assertRaises(ValueError): + deploy_pathways_service.calculate_vms_per_slice("4x8", 5) + + @mock.patch("pathwaysutils.experimental.shared_pathways_service.deploy_pathways_service.jobset.PathwaysJobSet") + def test_run_deployment(self, mock_jobset_cls): + mock_jobset = mock_jobset_cls.return_value + mock_jobset.to_dict.return_value = {"metadata": {"name": "test-jobset"}} + + # Mock head and worker job templates for mutation + mock_head_job = mock.MagicMock() + mock_head_job.spec.template.spec.containers = [ + mock.MagicMock(name="pathways-rm") + ] + mock_head_job.spec.template.spec.containers[0].name = "pathways-rm" + + mock_worker_job = mock.MagicMock() + mock_worker_job.spec.template.spec.containers = [ + mock.MagicMock(name="pathways-worker") + ] + mock_worker_job.spec.template.spec.containers[0].name = "pathways-worker" + mock_worker_job.spec.template.spec.containers[0].args = [] + + # Sidecar will be added by add_colocated_python, but we mock it here as if it was added + mock_sidecar = mock.MagicMock(name="colocated-python-sidecar") + mock_sidecar.name = "colocated-python-sidecar" + mock_sidecar.env = [] + mock_worker_job.spec.template.spec.init_containers = [mock_sidecar] + + mock_jobset.head_job_template = mock_head_job + mock_jobset.worker_job_template = mock_worker_job + + mock_deploy = mock.MagicMock() + + deploy_pathways_service.run_deployment( + tpu_type="v5e", + topology="4x8", + num_slices=2, + jobset_name="test-jobset", + gcs_bucket="test-bucket", + server_image="custom-server-image", + sidecar_image="custom-sidecar-image", + dry_run=False, + deploy_func=mock_deploy, + ) + + # Verify PathwaysJobSet was instantiated correctly + mock_jobset_cls.assert_called_once_with( + name="test-jobset", + namespace="default", + pathways_dir="test-bucket", + tpu_type="v5e", + topology="4x8", + num_slices=2, + shared_pathways_service=True, + max_slice_restarts=1000000, + ) + + # Verify colocated python was added with correct image and SHM path + mock_jobset.add_colocated_python.assert_called_once_with( + image="custom-sidecar-image", + shm_mount_path="/tmp/sidecar_dir", + ) + + # Verify server images were mutated + self.assertEqual(mock_head_job.spec.template.spec.containers[0].image, "custom-server-image") + self.assertEqual(mock_worker_job.spec.template.spec.containers[0].image, "custom-server-image") + + # Verify extra logging env vars were added to sidecar + self.assertTrue(any(e.name == "LOGLEVEL" and e.value == "DEBUG" for e in mock_sidecar.env)) + + # Verify arg was added to pathways-worker + self.assertIn( + "--cloud_pathways_sidecar_shm_directory=/tmp/sidecar_dir", + mock_worker_job.spec.template.spec.containers[0].args, + ) + + # Verify deploy_func was called with the dict + mock_deploy.assert_called_once_with({"metadata": {"name": "test-jobset"}}) + + def test_run_deployment_worker_backoff_limit(self): + captured_config = {} + + def capture_deploy(config): + nonlocal captured_config + captured_config = config + + deploy_pathways_service.run_deployment( + tpu_type="v5e", + topology="4x8", + num_slices=2, + jobset_name="test-jobset", + gcs_bucket="gs://test-bucket", + server_image=( + "us-docker.pkg.dev/test-project/test-repo/server:test-tag" + ), + sidecar_image=( + "us-docker.pkg.dev/test-project/test-repo/sidecar:test-tag" + ), + dry_run=False, + deploy_func=capture_deploy, + ) + + replicated_jobs = captured_config["spec"]["replicatedJobs"] + worker_job = next( + j for j in replicated_jobs if j["name"] == "pathways-worker" + ) + worker_backoff = worker_job["template"]["spec"]["backoffLimit"] + + # Verify worker backoff limit is set to a large value + self.assertGreaterEqual(worker_backoff, 1000000) + + +if __name__ == "__main__": + FLAGS = flags.FLAGS + FLAGS.jobset_name = "dummy" + FLAGS.jax_version = "dummy" + FLAGS.tpu_type = "v5e" + FLAGS.topology = "4x8" + FLAGS.num_slices = 2 + FLAGS.gcs_bucket = "dummy" + absltest.main() diff --git a/pathwaysutils/test/experimental/shared_pathways_service/gke_utils_test.py b/pathwaysutils/test/experimental/shared_pathways_service/gke_utils_test.py new file mode 100644 index 0000000..1b69dcb --- /dev/null +++ b/pathwaysutils/test/experimental/shared_pathways_service/gke_utils_test.py @@ -0,0 +1,962 @@ +"""Tests for gke_utils.py. +""" + +import io +import json +import socket +import subprocess +from unittest import mock + +from absl.testing import absltest +from pathwaysutils.experimental.shared_pathways_service import gke_utils +import portpicker + + +class GKEUtilsTest(absltest.TestCase): + """Tests for gke_utils.py.""" + + def test_fetch_cluster_credentials_success(self): + """Tests that fetch_cluster_credentials calls gcloud with the correct arguments.""" + mock_run = self.enter_context( + mock.patch.object(subprocess, "run", autospec=True) + ) + gke_utils.fetch_cluster_credentials( + cluster_name="test-cluster", + project_id="test-project", + location="test-zone", + ) + mock_run.assert_called_once_with( + [ + "gcloud", + "container", + "clusters", + "get-credentials", + "--location=test-zone", + "--project=test-project", + "--dns-endpoint", + "--", + "test-cluster", + ], + check=True, + capture_output=True, + text=True, + ) + + def test_fetch_cluster_credentials_failure(self): + """Tests that fetch_cluster_credentials raises an error when gcloud fails.""" + mock_run = self.enter_context( + mock.patch.object(subprocess, "run", autospec=True) + ) + mock_run.side_effect = subprocess.CalledProcessError( + returncode=1, cmd="gcloud", stderr="error" + ) + with self.assertRaises(subprocess.CalledProcessError): + gke_utils.fetch_cluster_credentials( + cluster_name="test-cluster", + project_id="test-project", + location="test-zone", + ) + + def test_validate_k8s_name_valid(self): + gke_utils._validate_k8s_name("valid-name-123") + gke_utils._validate_k8s_name("a") + gke_utils._validate_k8s_name("a-b") + + def test_validate_k8s_name_invalid(self): + with self.assertRaises(ValueError): + gke_utils._validate_k8s_name("-invalid") + with self.assertRaises(ValueError): + gke_utils._validate_k8s_name("invalid-") + with self.assertRaises(ValueError): + gke_utils._validate_k8s_name("Invalid") + with self.assertRaises(ValueError): + gke_utils._validate_k8s_name("invalid_name") + with self.assertRaises(ValueError): + gke_utils._validate_k8s_name("invalid.name") + + def test_deploy_gke_yaml_success(self): + mock_run = self.enter_context( + mock.patch.object(subprocess, "run", autospec=True) + ) + test_yaml = "apiVersion: v1\nkind: Pod\nmetadata:\n name: test" + gke_utils.deploy_gke_yaml(test_yaml) + mock_run.assert_called_once_with( + ["kubectl", "apply", "-f", "-"], + input=test_yaml, + check=True, + capture_output=True, + text=True, + ) + + def test_deploy_gke_yaml_failure(self): + mock_run = self.enter_context( + mock.patch.object(subprocess, "run", autospec=True) + ) + mock_run.side_effect = subprocess.CalledProcessError( + returncode=1, cmd="kubectl apply", stderr="error" + ) + with self.assertRaises(subprocess.CalledProcessError): + gke_utils.deploy_gke_yaml("test_yaml") + + def test_deploy_gke_yaml_create_success(self): + mock_run = self.enter_context( + mock.patch.object(subprocess, "run", autospec=True) + ) + test_yaml = "apiVersion: v1\nkind: Pod\nmetadata:\n name: test" + gke_utils.deploy_gke_yaml(test_yaml, action="create") + mock_run.assert_called_once_with( + ["kubectl", "create", "-f", "-"], + input=test_yaml, + check=True, + capture_output=True, + text=True, + ) + + def test_delete_gke_resource_success(self): + mock_run = self.enter_context( + mock.patch.object(subprocess, "run", autospec=True) + ) + gke_utils.delete_gke_resource("deployment", "test-deploy", "test-ns") + mock_run.assert_called_once_with( + [ + "kubectl", + "delete", + "deployment", + "-n", + "test-ns", + "--ignore-not-found", + "--", + "test-deploy", + ], + check=True, + capture_output=True, + text=True, + ) + + def test_delete_gke_resource_failure(self): + mock_run = self.enter_context( + mock.patch.object(subprocess, "run", autospec=True) + ) + mock_run.side_effect = subprocess.CalledProcessError( + returncode=1, cmd="kubectl delete", stderr="error" + ) + with self.assertRaises(subprocess.CalledProcessError): + gke_utils.delete_gke_resource("deployment", "test-deploy", "test-ns") + + def test_get_pod_from_job_success(self): + """Tests that get_pod_from_job returns the pod name on success.""" + mock_run = self.enter_context( + mock.patch.object( + subprocess, + "run", + autospec=True, + ) + ) + mock_get_pods_result = subprocess.CompletedProcess( + args=["kubectl", "get", "pods"], + returncode=0, + stdout="pod/test-pod-123\n", + ) + mock_run.return_value = mock_get_pods_result + + pod_name = gke_utils.get_pod_from_job("test-proxy-job") + + self.assertEqual(pod_name, "test-pod-123") + mock_run.assert_called_once() + self.assertIn("get", mock_run.call_args[0][0]) + self.assertIn("job-name=test-proxy-job", mock_run.call_args[0][0]) + + def test_get_pod_from_job_failure(self): + """Tests that get_pod_from_job raises an error if kubectl fails.""" + mock_run = self.enter_context( + mock.patch.object( + subprocess, + "run", + autospec=True, + ) + ) + mock_run.side_effect = subprocess.CalledProcessError( + returncode=1, cmd="kubectl get pods", stderr="error" + ) + + with self.assertRaises(subprocess.CalledProcessError): + gke_utils.get_pod_from_job("test-proxy-job") + + def test_check_pod_ready_success(self): + """Tests that check_pod_ready returns the pod name on success.""" + mock_run = self.enter_context( + mock.patch.object( + subprocess, + "run", + autospec=True, + ) + ) + mock_wait_success_result = subprocess.CompletedProcess( + args=["kubectl", "wait"], + returncode=0, + ) + mock_run.return_value = mock_wait_success_result + + pod_name = gke_utils.check_pod_ready("test-pod-123") + + self.assertEqual(pod_name, "test-pod-123") + mock_run.assert_called_once_with( + [ + "kubectl", + "wait", + "--for=condition=Ready", + "--timeout=30s", + "--", + "pod/test-pod-123", + ], + check=True, + capture_output=True, + text=True, + ) + + def test_check_pod_ready_failure(self): + """Tests that check_pod_ready raises a RuntimeError if kubectl wait fails.""" + mock_run = self.enter_context( + mock.patch.object( + subprocess, + "run", + autospec=True, + ) + ) + mock_run.side_effect = subprocess.CalledProcessError( + returncode=1, cmd="kubectl wait", stderr="error" + ) + + with self.assertRaisesRegex( + RuntimeError, "Pod did not become ready: error." + ): + gke_utils.check_pod_ready("test-pod-123") + + def test_get_log_link(self): + cluster = "test-cluster" + project = "test-project" + job_name = "test-job" + log_link = gke_utils.get_log_link( + cluster=cluster, project=project, job_name=job_name + ) + self.assertEqual( + log_link, + r"https://console.cloud.google.com/logs/query;query=resource.type%3D" + r"%22k8s_container%22%0Aresource.labels.cluster_name%3D" + "%22test-cluster%22%0Aresource.labels.namespace_name%3D" + "%22default%22%0Alabels.k8s-pod%2Fjob-name%3A%22test-job%22;" + "duration=PT1H?project=test-project", + ) + + def test_wait_for_pod_success(self): + """Tests that wait_for_pod returns the pod name on success.""" + mock_run = self.enter_context( + mock.patch.object( + subprocess, + "run", + autospec=True, + ) + ) + mock_get_pods_result = subprocess.CompletedProcess( + args=["kubectl", "get", "pods"], + returncode=0, + stdout="pod/test-pod-123\n", + ) + mock_wait_success_result = subprocess.CompletedProcess( + args=["kubectl", "wait"], + returncode=0, + ) + mock_run.side_effect = [mock_get_pods_result, mock_wait_success_result] + + pod_name = gke_utils.wait_for_pod("test-proxy-job") + + self.assertEqual(pod_name, "test-pod-123") + self.assertEqual(mock_run.call_count, 2) + self.assertIn("get", mock_run.call_args_list[0].args[0]) + self.assertIn("wait", mock_run.call_args_list[1].args[0]) + + def test_wait_for_pod_get_pods_fails(self): + """Tests that wait_for_pod raises an error if 'get pods' fails.""" + mock_run = self.enter_context( + mock.patch.object( + subprocess, + "run", + autospec=True, + ) + ) + mock_run.side_effect = subprocess.CalledProcessError( + returncode=1, cmd="kubectl get pods", stderr="error" + ) + + with self.assertRaises(subprocess.CalledProcessError): + gke_utils.wait_for_pod("test-proxy-job") + self.assertEqual(mock_run.call_count, 1) + + def test_wait_for_pod_wait_fails(self): + """Tests that wait_for_pod raises a RuntimeError if 'wait' times out.""" + mock_run = self.enter_context( + mock.patch.object( + subprocess, + "run", + autospec=True, + ) + ) + mock_get_pods_result = subprocess.CompletedProcess( + args=["kubectl", "get", "pods"], + returncode=0, + stdout="pod/test-pod-123\n", + ) + mock_run.side_effect = [ + mock_get_pods_result, + subprocess.TimeoutExpired(cmd="kubectl wait", timeout=30), + ] + + with self.assertRaises(RuntimeError): + gke_utils.wait_for_pod("test-proxy-job") + self.assertEqual(mock_run.call_count, 2) + + def test_enable_port_forwarding_success(self): + """Tests successful port forwarding.""" + # Arrange + mock_create_connection = self.enter_context( + mock.patch.object(socket, "create_connection", autospec=True) + ) + mock_popen = self.enter_context( + mock.patch.object( + subprocess, + "Popen", + autospec=True, + ) + ) + pod_name = "test-pod-123" + mock_pick_port = self.enter_context( + mock.patch.object(portpicker, "pick_unused_port", autospec=True) + ) + mock_pick_port.return_value = 29007 + mock_process = mock_popen.return_value + mock_process.stdout = io.StringIO( + "Forwarding from 127.0.0.1:29007 -> 8080\n" + ) + + # Act + port, process = gke_utils.enable_port_forwarding(pod_name, 8080) + + # Assert + self.assertEqual(port, 29007) + self.assertIs(process, mock_process) + mock_popen.assert_called_once_with( + [ + "kubectl", + "port-forward", + "-n", + "default", + "--address", + "localhost", + "--", + "test-pod-123", + "29007:8080", + ], + stdout=subprocess.PIPE, + stderr=subprocess.PIPE, + text=True, + ) + mock_process.terminate.assert_not_called() + mock_create_connection.assert_called_once_with( + ("localhost", 29007), timeout=30 + ) + + def test_enable_port_forwarding_timeout_raises_error(self): + """Tests that a timeout in port forwarding raises a RuntimeError.""" + # Arrange + mock_create_connection = self.enter_context( + mock.patch.object(socket, "create_connection", autospec=True) + ) + mock_popen = self.enter_context( + mock.patch.object( + subprocess, + "Popen", + autospec=True, + ) + ) + pod_name = "test-pod-123" + mock_pick_port = self.enter_context( + mock.patch.object(portpicker, "pick_unused_port", autospec=True) + ) + mock_pick_port.return_value = 29007 + mock_process = mock_popen.return_value + mock_process.stdout = io.StringIO("") # Empty output simulates a timeout + mock_process.communicate.return_value = ("", "") + mock_process.poll.return_value = None + def terminate_effect(): + mock_process.poll.return_value = 1 + mock_process.terminate.side_effect = terminate_effect + mock_popen.return_value = mock_process + + # Act & Assert + with self.assertRaises(RuntimeError): + gke_utils.enable_port_forwarding(pod_name, 8080) + + mock_create_connection.assert_not_called() + mock_process.terminate.assert_called_once() + + def test_enable_port_forwarding_socket_error_raises_error(self): + """Tests that a socket error during connection check raises an error.""" + # Arrange + mock_create_connection = self.enter_context( + mock.patch.object(socket, "create_connection", autospec=True) + ) + mock_popen = self.enter_context( + mock.patch.object( + subprocess, + "Popen", + autospec=True, + ) + ) + pod_name = "test-pod-123" + mock_pick_port = self.enter_context( + mock.patch.object(portpicker, "pick_unused_port", autospec=True) + ) + mock_pick_port.return_value = 29007 + mock_process = mock_popen.return_value + mock_process.stdout = io.StringIO( + "Forwarding from 127.0.0.1:29007 -> 8080\n" + ) + mock_process.poll.return_value = None + mock_popen.return_value = mock_process + mock_create_connection.side_effect = OSError("Connection failed") + + # Act & Assert + with self.assertRaises(OSError): + gke_utils.enable_port_forwarding(pod_name, 8080) + + mock_process.terminate.assert_called_once() + + def test_stream_pod_logs_success(self): + mock_popen = self.enter_context( + mock.patch.object(subprocess, "Popen", autospec=True) + ) + mock_process = mock_popen.return_value + + process = gke_utils.stream_pod_logs("test-pod-123") + + self.assertIs(process, mock_process) + mock_popen.assert_called_once_with( + ["kubectl", "logs", "-f", "--", "pod/test-pod-123"], + stdout=subprocess.PIPE, + stderr=subprocess.STDOUT, + text=True, + bufsize=1, + ) + + def test_stream_pod_logs_failure(self): + mock_popen = self.enter_context( + mock.patch.object(subprocess, "Popen", autospec=True) + ) + mock_popen.side_effect = Exception("error") + + with self.assertRaisesRegex(Exception, "error"): + gke_utils.stream_pod_logs("test-pod-123") + + def test_deploy_gke_yaml_invalid_action(self): + with self.assertRaisesRegex(ValueError, "Invalid kubectl action:"): + gke_utils.deploy_gke_yaml("test_yaml", action="invalid_action") + + def test_get_pod_from_job_invalid_format_empty(self): + mock_run = self.enter_context( + mock.patch.object(subprocess, "run", autospec=True) + ) + mock_run.return_value = subprocess.CompletedProcess( + args=["kubectl", "get", "pods"], + returncode=0, + stdout="\n", + ) + with self.assertRaisesRegex(RuntimeError, "Failed to get pod name. Expected format:"): + gke_utils.get_pod_from_job("test-job") + + def test_get_pod_from_job_invalid_format_no_prefix(self): + mock_run = self.enter_context( + mock.patch.object(subprocess, "run", autospec=True) + ) + mock_run.return_value = subprocess.CompletedProcess( + args=["kubectl", "get", "pods"], + returncode=0, + stdout="test-pod-123\n", + ) + with self.assertRaisesRegex(RuntimeError, "Failed to get pod name. Expected format:"): + gke_utils.get_pod_from_job("test-job") + + def test_get_pod_from_job_invalid_format_too_many_slashes(self): + mock_run = self.enter_context( + mock.patch.object(subprocess, "run", autospec=True) + ) + mock_run.return_value = subprocess.CompletedProcess( + args=["kubectl", "get", "pods"], + returncode=0, + stdout="pod/test-pod/extra\n", + ) + with self.assertRaisesRegex(RuntimeError, "Failed to get pod name. Expected format:"): + gke_utils.get_pod_from_job("test-job") + + def test_test_remote_connection_success(self): + mock_create = self.enter_context( + mock.patch.object(socket, "create_connection", autospec=True) + ) + gke_utils._test_remote_connection(8080) + mock_create.assert_called_once_with(("localhost", 8080), timeout=30) + + def test_test_remote_connection_timeout(self): + mock_create = self.enter_context( + mock.patch.object(socket, "create_connection", autospec=True) + ) + mock_create.side_effect = socket.timeout("timeout error") + with self.assertRaisesRegex(RuntimeError, "Could not connect to the pod."): + gke_utils._test_remote_connection(8080) + + def test_test_remote_connection_refused(self): + mock_create = self.enter_context( + mock.patch.object(socket, "create_connection", autospec=True) + ) + mock_create.side_effect = ConnectionRefusedError("connection refused") + with self.assertRaisesRegex(RuntimeError, "Could not connect to the pod."): + gke_utils._test_remote_connection(8080) + + def test_enable_port_forwarding_pick_port_fails(self): + self.enter_context( + mock.patch.object(portpicker, "pick_unused_port", side_effect=ValueError("pick failed")) + ) + with self.assertRaisesRegex(ValueError, "pick failed"): + gke_utils.enable_port_forwarding("test-pod", 8080) + + def test_enable_port_forwarding_popen_fails(self): + self.enter_context( + mock.patch.object(portpicker, "pick_unused_port", return_value=12345) + ) + self.enter_context( + mock.patch.object(subprocess, "Popen", side_effect=OSError("Popen failed")) + ) + with self.assertRaisesRegex(OSError, "Popen failed"): + gke_utils.enable_port_forwarding("test-pod", 8080) + + def test_enable_port_forwarding_stdout_none(self): + self.enter_context( + mock.patch.object(portpicker, "pick_unused_port", return_value=12345) + ) + mock_popen = self.enter_context( + mock.patch.object(subprocess, "Popen", autospec=True) + ) + mock_process = mock_popen.return_value + mock_process.stdout = None + mock_process.communicate.return_value = ("stdout", "stderr_out") + + with self.assertRaisesRegex(RuntimeError, "Failed to start port forwarding: stdout not available.\nSTDERR: stderr_out"): + gke_utils.enable_port_forwarding("test-pod", 8080) + mock_process.terminate.assert_called_once() + mock_process.communicate.assert_called_once() + + def test_wait_for_deployment_success(self): + mock_run = self.enter_context( + mock.patch.object(subprocess, "run", autospec=True) + ) + gke_utils.wait_for_deployment("my-deploy", "my-ns") + mock_run.assert_called_once_with( + [ + "kubectl", + "rollout", + "status", + "deployment/my-deploy", + "-n", + "my-ns", + "--timeout=300s", + ], + check=True, + capture_output=True, + text=True, + ) + + def test_wait_for_deployment_failure(self): + mock_run = self.enter_context( + mock.patch.object(subprocess, "run", autospec=True) + ) + mock_run.side_effect = subprocess.CalledProcessError( + returncode=1, cmd="kubectl rollout", stderr="rollout failed" + ) + with self.assertRaisesRegex(RuntimeError, "Deployment did not become ready: rollout failed"): + gke_utils.wait_for_deployment("my-deploy", "my-ns") + + def test_wait_for_service_ip_success_first_try(self): + mock_run = self.enter_context( + mock.patch.object(subprocess, "run", autospec=True) + ) + mock_run.return_value = subprocess.CompletedProcess( + args=["kubectl", "get", "svc"], + returncode=0, + stdout="1.2.3.4\n", + ) + ip = gke_utils.wait_for_service_ip("my-svc", "my-ns", timeout=10) + self.assertEqual(ip, "1.2.3.4") + mock_run.assert_called_once_with( + [ + "kubectl", + "get", + "svc", + "my-svc", + "-n", + "my-ns", + "-o", + "jsonpath={.status.loadBalancer.ingress[0].ip}", + ], + check=True, + capture_output=True, + text=True, + ) + + def test_wait_for_service_ip_success_third_try(self): + mock_run = self.enter_context( + mock.patch.object(subprocess, "run", autospec=True) + ) + mock_run.side_effect = [ + subprocess.CompletedProcess(args=[], returncode=0, stdout=""), + subprocess.CompletedProcess(args=[], returncode=0, stdout=""), + subprocess.CompletedProcess(args=[], returncode=0, stdout="1.2.3.4"), + ] + mock_sleep = self.enter_context( + mock.patch("time.sleep", autospec=True) + ) + ip = gke_utils.wait_for_service_ip("my-svc", "my-ns", timeout=10) + self.assertEqual(ip, "1.2.3.4") + self.assertEqual(mock_run.call_count, 3) + self.assertEqual(mock_sleep.call_count, 2) + + def test_wait_for_service_ip_timeout(self): + mock_run = self.enter_context( + mock.patch.object(subprocess, "run", autospec=True) + ) + mock_run.return_value = subprocess.CompletedProcess( + args=["kubectl", "get", "svc"], + returncode=0, + stdout="", + ) + self.enter_context(mock.patch("time.sleep", autospec=True)) + self.enter_context( + mock.patch("time.time", side_effect=[0, 0, 4, 4, 10]) + ) + with self.assertRaisesRegex( + RuntimeError, "Timeout waiting for service IP for my-svc" + ): + gke_utils.wait_for_service_ip("my-svc", "my-ns", timeout=5) + self.assertGreaterEqual(mock_run.call_count, 2) + + def test_pick_unused_local_port(self): + mock_pick = self.enter_context( + mock.patch.object(portpicker, "pick_unused_port", return_value=12345) + ) + port = gke_utils.pick_unused_local_port() + self.assertEqual(port, 12345) + mock_pick.assert_called_once() + + def test_is_local_port_free(self): + mock_check = self.enter_context( + mock.patch.object(portpicker, "is_port_free", return_value=True) + ) + is_free = gke_utils.is_local_port_free(12345) + self.assertTrue(is_free) + mock_check.assert_called_once_with(12345) + + def test_enable_port_forwarding_with_slash_success(self): + """Tests successful port forwarding when remote_server contains a slash.""" + self.enter_context( + mock.patch.object(socket, "create_connection", autospec=True) + ) + mock_popen = self.enter_context( + mock.patch.object( + subprocess, + "Popen", + autospec=True, + ) + ) + mock_pick_port = self.enter_context( + mock.patch.object(portpicker, "pick_unused_port", autospec=True) + ) + mock_pick_port.return_value = 29007 + mock_process = mock_popen.return_value + mock_process.stdout = io.StringIO( + "Forwarding from 127.0.0.1:29007 -> 8080\n" + ) + + port, process = gke_utils.enable_port_forwarding("svc/test-service", 8080) + + self.assertEqual(port, 29007) + self.assertIs(process, mock_process) + mock_popen.assert_called_once_with( + [ + "kubectl", + "port-forward", + "-n", + "default", + "--address", + "localhost", + "--", + "svc/test-service", + "29007:8080", + ], + stdout=subprocess.PIPE, + stderr=subprocess.PIPE, + text=True, + ) + + def test_enable_port_forwarding_invalid_format(self): + with self.assertRaises(ValueError): + gke_utils.enable_port_forwarding("invalid/svc/name", 8080) + + def test_enable_port_forwarding_invalid_name(self): + with self.assertRaises(ValueError): + gke_utils.enable_port_forwarding("svc/invalid_name", 8080) + + def test_delete_gke_resource_invalid_params(self): + with self.assertRaises(ValueError): + gke_utils.delete_gke_resource("invalid_type", "name-123", "namespace") + with self.assertRaises(ValueError): + gke_utils.delete_gke_resource("deployment", "invalid_name", "namespace") + with self.assertRaises(ValueError): + gke_utils.delete_gke_resource("deployment", "name-123", + "invalid_namespace") + + def test_get_worker_sidecar_image_success_with_jobset_name(self): + mock_run = self.enter_context( + mock.patch.object(subprocess, "run", autospec=True) + ) + pods_json = { + "items": [ + { + "metadata": { + "name": "my-jobset-worker-0-0", + "labels": { + "jobset.sigs.k8s.io/jobset-name": "my-jobset" + } + }, + "spec": { + "initContainers": [ + { + "name": "colocated-python-sidecar", + "image": ( + "us-docker.pkg.dev/cloud-tpu-v2-images/" + "pathways-colocated-python/sidecar:" + "20260423-python_3.12-jax_0.10.0" + ) + } + ] + } + } + ] + } + mock_run.return_value = subprocess.CompletedProcess( + args=["kubectl", "get", "pods"], + returncode=0, + stdout=json.dumps(pods_json), + ) + image = gke_utils.get_worker_sidecar_image( + pathways_service="my-jobset-pathways-head-0-0:8080" + ) + self.assertEqual( + image, + ("us-docker.pkg.dev/cloud-tpu-v2-images/pathways-colocated-python/" + "sidecar:20260423-python_3.12-jax_0.10.0"), + ) + + def test_get_worker_sidecar_image_failure_none(self): + mock_run = self.enter_context( + mock.patch.object(subprocess, "run", autospec=True) + ) + mock_run.return_value = subprocess.CompletedProcess( + args=["kubectl", "get", "pods"], + returncode=0, + stdout="missing sidecar image", + ) + image = gke_utils.get_worker_sidecar_image( + pathways_service="my-jobset-pathways-head-0-0:8080" + ) + self.assertIsNone(image) + + def test_get_worker_sidecar_image_invalid_namespace(self): + with self.assertRaises(ValueError): + gke_utils.get_worker_sidecar_image( + pathways_service="my-jobset-pathways-head-0-0:8080", + namespace="invalid namespace!", + ) + + def test_get_worker_sidecar_image_no_pathways_head_in_hostname(self): + mock_run = self.enter_context( + mock.patch.object(subprocess, "run", autospec=True) + ) + mock_run.return_value = subprocess.CompletedProcess( + args=["kubectl", "get", "pods"], + returncode=0, + stdout="{}", + ) + image = gke_utils.get_worker_sidecar_image( + pathways_service="my-jobset:8080" + ) + self.assertIsNone(image) + + def test_get_worker_sidecar_image_kubectl_error(self): + mock_run = self.enter_context( + mock.patch.object(subprocess, "run", autospec=True) + ) + mock_run.side_effect = subprocess.CalledProcessError( + returncode=1, cmd=["kubectl", "get", "pods"], stderr="kubectl error" + ) + image = gke_utils.get_worker_sidecar_image( + pathways_service="my-jobset-pathways-head-0-0:8080" + ) + self.assertIsNone(image) + + def test_get_worker_sidecar_image_success_with_pod_name_prefix(self): + mock_run = self.enter_context( + mock.patch.object(subprocess, "run", autospec=True) + ) + pods_json = { + "items": [ + { + "metadata": { + "name": "my-jobset-worker-0-0", + "labels": {} + }, + "spec": { + "initContainers": [ + { + "name": "colocated-python-sidecar", + "image": "sidecar-image-url" + } + ] + } + } + ] + } + mock_run.return_value = subprocess.CompletedProcess( + args=["kubectl", "get", "pods"], + returncode=0, + stdout=json.dumps(pods_json), + ) + image = gke_utils.get_worker_sidecar_image( + pathways_service="my-jobset-pathways-head-0-0:8080" + ) + self.assertEqual(image, "sidecar-image-url") + + def test_get_worker_sidecar_image_success_in_containers(self): + mock_run = self.enter_context( + mock.patch.object(subprocess, "run", autospec=True) + ) + pods_json = { + "items": [ + { + "metadata": { + "name": "my-jobset-worker-0-0", + "labels": {} + }, + "spec": { + "containers": [ + { + "name": "colocated-python-sidecar", + "image": "sidecar-image-url" + } + ] + } + } + ] + } + mock_run.return_value = subprocess.CompletedProcess( + args=["kubectl", "get", "pods"], + returncode=0, + stdout=json.dumps(pods_json), + ) + image = gke_utils.get_worker_sidecar_image( + pathways_service="my-jobset-pathways-head-0-0:8080" + ) + self.assertEqual(image, "sidecar-image-url") + + def test_get_worker_sidecar_image_missing_sidecar_container(self): + mock_run = self.enter_context( + mock.patch.object(subprocess, "run", autospec=True) + ) + pods_json = { + "items": [ + { + "metadata": { + "name": "my-jobset-worker-0-0", + "labels": {} + }, + "spec": { + "containers": [ + { + "name": "some-other-container", + "image": "some-image" + } + ] + } + } + ] + } + mock_run.return_value = subprocess.CompletedProcess( + args=["kubectl", "get", "pods"], + returncode=0, + stdout=json.dumps(pods_json), + ) + image = gke_utils.get_worker_sidecar_image( + pathways_service="my-jobset-pathways-head-0-0:8080" + ) + self.assertIsNone(image) + + def test_get_worker_sidecar_image_missing_image_field(self): + mock_run = self.enter_context( + mock.patch.object(subprocess, "run", autospec=True) + ) + pods_json = { + "items": [ + { + "metadata": { + "name": "my-jobset-worker-0-0", + "labels": {} + }, + "spec": { + "containers": [ + { + "name": "colocated-python-sidecar", + } + ] + } + } + ] + } + mock_run.return_value = subprocess.CompletedProcess( + args=["kubectl", "get", "pods"], + returncode=0, + stdout=json.dumps(pods_json), + ) + image = gke_utils.get_worker_sidecar_image( + pathways_service="my-jobset-pathways-head-0-0:8080" + ) + self.assertIsNone(image) + + def test_get_worker_sidecar_image_custom_namespace(self): + mock_run = self.enter_context( + mock.patch.object(subprocess, "run", autospec=True) + ) + mock_run.return_value = subprocess.CompletedProcess( + args=["kubectl", "get", "pods"], + returncode=0, + stdout="{}", + ) + gke_utils.get_worker_sidecar_image( + pathways_service="my-jobset-pathways-head-0-0:8080", + namespace="my-custom-ns", + ) + mock_run.assert_called_once_with( + ["kubectl", "get", "pods", "-n", "my-custom-ns", "-o", "json"], + check=True, + capture_output=True, + text=True, + ) + + +if __name__ == "__main__": + absltest.main() diff --git a/pathwaysutils/test/experimental/shared_pathways_service/isc_pathways_test.py b/pathwaysutils/test/experimental/shared_pathways_service/isc_pathways_test.py new file mode 100644 index 0000000..aa189b3 --- /dev/null +++ b/pathwaysutils/test/experimental/shared_pathways_service/isc_pathways_test.py @@ -0,0 +1,852 @@ +"""Tests for the ISCPathways class. +""" +import io +import os +import subprocess +from unittest import mock + +from absl import flags +from absl.testing import absltest +from absl.testing import parameterized +from pathwaysutils.experimental.shared_pathways_service import isc_pathways + + + +class ISCPathwaysTest(parameterized.TestCase): + """Tests for the ISCPathways class.""" + + def test_wait_for_placement_success(self): + """Tests that _wait_for_placement correctly processes logs.""" + mock_process = mock.create_autospec(subprocess.Popen, instance=True) + mock_process.stdout = io.StringIO( + "Some log\nPlacement info\nTransition slice\nSignaling to RM\nunplaced" + " -> placed\n" + ) + + mock_stream_func = mock.MagicMock() + mock_stream_func.return_value.__enter__.return_value = mock_process + + isc_pathways._wait_for_placement( + "test-pod", + num_slices=1, + stream_logs_func=mock_stream_func, + metrics_collector_inst=mock.Mock(), + ) + + mock_stream_func.assert_called_once_with("test-pod") + + def test_wait_for_placement_timeout(self): + """Tests that _wait_for_placement kills the process on timeout.""" + mock_process = mock.create_autospec(subprocess.Popen, instance=True) + mock_process.stdout = io.StringIO("unplaced -> placed\n") + mock_process.wait.side_effect = subprocess.TimeoutExpired( + cmd="wait", timeout=5 + ) + + mock_stream_func = mock.MagicMock() + mock_stream_func.return_value.__enter__.return_value = mock_process + + isc_pathways._wait_for_placement( + "test-pod", + num_slices=1, + stream_logs_func=mock_stream_func, + metrics_collector_inst=mock.Mock(), + ) + + def test_wait_for_placement_reports_metrics(self): + """Tests that _wait_for_placement reports metrics on success.""" + mock_process = mock.create_autospec(subprocess.Popen, instance=True) + mock_process.stdout = io.StringIO("Some log\nunplaced -> placed\n") + + mock_stream_func = mock.MagicMock() + mock_stream_func.return_value.__enter__.return_value = mock_process + + mock_metrics_collector = mock.Mock() + start_time = 100.0 + + with mock.patch("time.time", return_value=150.0): + isc_pathways._wait_for_placement( + "test-pod", + num_slices=1, + stream_logs_func=mock_stream_func, + metrics_collector_inst=mock_metrics_collector, + start_time=start_time, + ) + + mock_metrics_collector.record_assignment_time.assert_called_once_with(50.0) + mock_metrics_collector.record_successful_request.assert_called_once() + + def test_deploy_pathways_proxy_server_success(self): + mock_deploy_gke_yaml = self.enter_context( + mock.patch.object( + isc_pathways.gke_utils, "deploy_gke_yaml", autospec=True + ) + ) + pathways_service = "test-service:8080" + proxy_name = "test-proxy" + expected_instances = {"tpuv6e:2x2": 2} + gcs_bucket = "test-bucket" + + isc_pathways._deploy_pathways_proxy_server( + pathways_service=pathways_service, + proxy_job_name=proxy_name, + expected_instances=expected_instances, + gcs_scratch_location=gcs_bucket, + proxy_server_image="test-image:latest", + proxy_options=isc_pathways.ProxyOptions(use_insecure_credentials=False), + ) + + mock_deploy_gke_yaml.assert_called_once() + substituted_yaml = mock_deploy_gke_yaml.call_args[0][0] + self.assertIn("name: test-proxy", substituted_yaml) + self.assertIn( + "--resource_manager_address=test-service:8080", substituted_yaml + ) + self.assertIn("--gcs_scratch_location=test-bucket", substituted_yaml) + self.assertIn("--virtual_slices=tpuv6e:2x2,tpuv6e:2x2", substituted_yaml) + self.assertIn("image: test-image:latest", substituted_yaml) + # Extract the env section and check that it doesn't contain any - name: + # entries. + env_section = substituted_yaml.split("env:\n")[1].split("ports:")[0] + self.assertNotIn("- name:", env_section) + + def test_deploy_pathways_proxy_server_with_insecure_credentials_success(self): + mock_deploy_gke_yaml = self.enter_context( + mock.patch.object( + isc_pathways.gke_utils, "deploy_gke_yaml", autospec=True + ) + ) + pathways_service = "test-service:8080" + proxy_name = "test-proxy" + expected_instances = {"tpuv6e:2x2": 2} + gcs_bucket = "test-bucket" + + isc_pathways._deploy_pathways_proxy_server( + pathways_service=pathways_service, + proxy_job_name=proxy_name, + expected_instances=expected_instances, + gcs_scratch_location=gcs_bucket, + proxy_server_image="test-image:latest", + proxy_options=isc_pathways.ProxyOptions(use_insecure_credentials=True), + ) + + mock_deploy_gke_yaml.assert_called_once() + substituted_yaml = mock_deploy_gke_yaml.call_args[0][0] + self.assertIn("env:", substituted_yaml) + self.assertIn( + "- name: IFRT_PROXY_USE_INSECURE_GRPC_CREDENTIALS", substituted_yaml + ) + self.assertIn('value: "true"', substituted_yaml) + + def test_deploy_pathways_proxy_server_with_xla_flags_success(self): + mock_deploy_gke_yaml = self.enter_context( + mock.patch.object( + isc_pathways.gke_utils, "deploy_gke_yaml", autospec=True + ) + ) + pathways_service = "test-service:8080" + proxy_name = "test-proxy" + expected_instances = {"tpuv6e:2x2": 1} + gcs_bucket = "test-bucket" + isc_pathways._deploy_pathways_proxy_server( + pathways_service=pathways_service, + proxy_job_name=proxy_name, + expected_instances=expected_instances, + gcs_scratch_location=gcs_bucket, + proxy_server_image="test-image:latest", + proxy_options=isc_pathways.ProxyOptions( + xla_flags=["--xla_flag1", "--xla_flag2"] + ), + ) + mock_deploy_gke_yaml.assert_called_once() + substituted_yaml = mock_deploy_gke_yaml.call_args[0][0] + self.assertIn("- --xla_flag1", substituted_yaml) + self.assertIn("- --xla_flag2", substituted_yaml) + + def test_deploy_pathways_proxy_server_with_sidecar_success(self): + mock_deploy_gke_yaml = self.enter_context( + mock.patch.object( + isc_pathways.gke_utils, "deploy_gke_yaml", autospec=True + ) + ) + pathways_service = "test-service:8080" + proxy_name = "test-proxy" + expected_instances = {"tpuv6e:2x2": 1} + gcs_bucket = "test-bucket" + isc_pathways._deploy_pathways_proxy_server( + pathways_service=pathways_service, + proxy_job_name=proxy_name, + expected_instances=expected_instances, + gcs_scratch_location=gcs_bucket, + proxy_server_image="test-image:latest", + proxy_options=isc_pathways.ProxyOptions(sidecar=True), + ) + mock_deploy_gke_yaml.assert_called_once() + substituted_yaml = mock_deploy_gke_yaml.call_args[0][0] + self.assertIn("- --sidecar_name=external", substituted_yaml) + + def test_proxy_options_from_list(self): + """Tests ProxyOptions.from_list with varied input formats.""" + # Standard valid input + self.assertTrue( + isc_pathways.ProxyOptions.from_list( + ["use_insecure_credentials:true"] + ).use_insecure_credentials + ) + # Case sensitivity and whitespace + self.assertTrue( + isc_pathways.ProxyOptions.from_list( + [" USE_INSECURE_CREDENTIALS : True "] + ).use_insecure_credentials + ) + # Valid false + self.assertFalse( + isc_pathways.ProxyOptions.from_list( + ["use_insecure_credentials:false"] + ).use_insecure_credentials + ) + # Empty and None + self.assertFalse( + isc_pathways.ProxyOptions.from_list([]).use_insecure_credentials + ) + self.assertFalse( + isc_pathways.ProxyOptions.from_list(None).use_insecure_credentials + ) + # Invalid formats and unknown keys + self.assertFalse( + isc_pathways.ProxyOptions.from_list( + ["invalid_format", "unknown_key:value"] + ).use_insecure_credentials + ) + # Invalid value for known key + self.assertFalse( + isc_pathways.ProxyOptions.from_list( + ["use_insecure_credentials:maybe"] + ).use_insecure_credentials + ) + + def test_proxy_options_from_list_with_sidecar(self): + """Tests ProxyOptions.from_list parsing the sidecar option.""" + self.assertTrue( + isc_pathways.ProxyOptions.from_list(["sidecar:true"]).sidecar + ) + self.assertTrue( + isc_pathways.ProxyOptions.from_list([" SIDECAR : True "]).sidecar + ) + self.assertFalse( + isc_pathways.ProxyOptions.from_list(["sidecar:false"]).sidecar + ) + self.assertFalse( + isc_pathways.ProxyOptions.from_list([]).sidecar + ) + self.assertFalse( + isc_pathways.ProxyOptions.from_list(None).sidecar + ) + + @parameterized.named_parameters( + ( + "standard_valid", + ['xla_flags:"--xla_flag1 --xla_flag2"'], + ["--xla_flag1", "--xla_flag2"], + ), + ("single_flag", ['xla_flags:"--xla_flag1"'], ["--xla_flag1"]), + ( + "no_quotes", + ["xla_flags:--xla_flag1 --xla_flag2"], + ["--xla_flag1", "--xla_flag2"], + ), + ) + def test_proxy_options_from_list_with_xla_flags( + self, input_list, expected_flags + ): + options = isc_pathways.ProxyOptions.from_list(input_list) + self.assertEqual(options.xla_flags, expected_flags) + + def test_proxy_options_from_list_with_xla_flags_failure(self): + with self.assertRaisesRegex( + flags.ValidationError, "must start with '--xla_'" + ): + isc_pathways.ProxyOptions.from_list(["xla_flags:--not_xla_flag"]) + + def test_deploy_pathways_proxy_server_failure(self): + mock_deploy_gke_yaml = self.enter_context( + mock.patch.object( + isc_pathways.gke_utils, "deploy_gke_yaml", autospec=True + ) + ) + + mock_deploy_gke_yaml.side_effect = subprocess.CalledProcessError( + returncode=1, cmd="kubectl", stderr="error" + ) + + with self.assertRaises(subprocess.CalledProcessError): + isc_pathways._deploy_pathways_proxy_server( + pathways_service="service:1234", + proxy_job_name="proxy", + expected_instances={"tpuv6e:2x2": 1}, + gcs_scratch_location="bucket", + proxy_server_image="test-image:latest", + ) + + def test_deploy_pathways_proxy_server_file_not_found(self): + """Tests ValueError when PROXY_FILEPATH is not found. + + Ensures that _deploy_pathways_proxy_server raises a ValueError if the + PROXY_FILEPATH cannot be read. + """ + mock_open = self.enter_context( + mock.patch("builtins.open", autospec=True) + ) + mock_open.side_effect = OSError("File not found") + + with self.assertRaisesRegex(ValueError, "Could not read file:"): + isc_pathways._deploy_pathways_proxy_server( + pathways_service="service:1234", + proxy_job_name="proxy", + expected_instances={"tpuv6e:2x2": 1}, + gcs_scratch_location="bucket", + proxy_server_image="test-image:latest", + ) + + # Ensure open was called with the expected file path. + mock_open.assert_called_once_with(isc_pathways.PROXY_FILEPATH, "r") + + def test_isc_pathways_pod_wait_failure_raises_error(self): + """Tests ISCPathways raises an error if pods do not become ready.""" + # Arrange + mock_deploy = self.enter_context( + mock.patch.object( + isc_pathways, "_deploy_pathways_proxy_server", + autospec=True, + ) + ) + mock_run = self.enter_context(mock.patch("subprocess.run", autospec=True)) + self.enter_context( + mock.patch("urllib.parse.quote", return_value="encoded_filter") + ) + # Simulate 'kubectl get' returning a pod, 'kubectl wait' timing out, + # then 'kubectl delete' succeeding. + mock_get_pods_result = subprocess.CompletedProcess( + args=["kubectl", "get", "pods"], + returncode=0, + stdout="pod/test-pod-123\n", + ) + mock_delete_result = subprocess.CompletedProcess( + args=["kubectl", "delete"], + returncode=0, + ) + mock_run.side_effect = [ + # First call for 'get' + mock_get_pods_result, + # Second call for 'wait' + subprocess.TimeoutExpired(cmd="kubectl wait", timeout=30), + # Third call for 'delete' + mock_delete_result, + ] + + # Act & Assert + with self.assertRaises(RuntimeError) as context: + with isc_pathways._ISCPathways( + cluster="test-cluster", + project="test-project", + region="test-region", + gcs_bucket="test-bucket", + pathways_service="test-service:1234", + expected_tpu_instances={"tpuv5:4x4x4": 1}, + proxy_job_name="test-proxy", + proxy_server_image="test-image:latest", + ): + self.fail( + "ISCPathways context should not be entered because we expect " + "a RuntimeError to be raised." + ) + + self.assertIn("Pod did not become ready", str(context.exception)) + mock_deploy.assert_called_once() + + # Check that subprocess.run was called for get, wait, and delete. + self.assertEqual(mock_run.call_count, 3) + self.assertIn("get", mock_run.call_args_list[0][0][0]) + self.assertIn("wait", mock_run.call_args_list[1][0][0]) + self.assertIn("delete", mock_run.call_args_list[2][0][0]) + + def test_isc_pathways(self): + """Tests the full lifecycle of ISCPathways.""" + # Arrange + self.enter_context( + mock.patch( + "pathwaysutils.experimental.shared_pathways_service.gke_utils.socket.create_connection", + autospec=True, + ) + ) + mock_pick_unused_port = self.enter_context( + mock.patch.object( + isc_pathways.gke_utils.portpicker, "pick_unused_port", autospec=True + ) + ) + mock_deploy = self.enter_context( + mock.patch.object( + isc_pathways, "_deploy_pathways_proxy_server", + autospec=True, + ) + ) + mock_run = self.enter_context(mock.patch("subprocess.run", autospec=True)) + mock_popen = self.enter_context( + mock.patch("subprocess.Popen", autospec=True) + ) + mock_random = self.enter_context( + mock.patch( + "pathwaysutils.experimental.shared_pathways_service.isc_pathways.random", + autospec=True, + ) + ) + self.enter_context( + mock.patch.dict( + os.environ, + { + "USER": "testuser", + "JAX_PLATFORMS": "original_platform", + "JAX_BACKEND_TARGET": "original_target", + }, + ) + ) + # Mock jax.config + mock_jax_config = self.enter_context( + mock.patch( + "pathwaysutils.experimental.shared_pathways_service.isc_pathways.jax.config" + ) + ) + mock_jax_config.jax_platforms = "original_platform_config" + mock_jax_config.jax_backend_target = "original_target_config" + mock_jax_config_update = mock_jax_config.update + mock_clear_backends = self.enter_context( + mock.patch( + "pathwaysutils.experimental.shared_pathways_service.isc_pathways.jax_backend.clear_backends", + autospec=True, + ) + ) + mock_clear_caches = self.enter_context( + mock.patch( + "pathwaysutils.experimental.shared_pathways_service.isc_pathways.jax.clear_caches", + autospec=True, + ) + ) + mock_gc_collect = self.enter_context( + mock.patch( + "pathwaysutils.experimental.shared_pathways_service.isc_pathways.gc.collect", + autospec=True, + ) + ) + mock_metrics_collector = self.enter_context( + mock.patch.object( + isc_pathways.metrics_collector, "MetricsCollector", autospec=True + ) + ) + mock_collector_instance = mock_metrics_collector.return_value + mock_random.choices.return_value = list("abcde") + mock_random.randint.return_value = 29005 + + # Mock for 'kubectl wait' and 'kubectl delete' + mock_success_result = subprocess.CompletedProcess( + args=["kubectl"], + returncode=0, + stdout="", + ) + + # Mock for 'kubectl get pods' + mock_get_pods_result = subprocess.CompletedProcess( + args=["kubectl", "get", "pods"], + returncode=0, + stdout="pod/test-pod-123\n", + ) + + mock_run.side_effect = [ + mock_get_pods_result, # For 'get_pod_from_job' + mock_success_result, # For 'check_pod_ready' + mock_success_result, # For 'kubectl delete job' in __exit__ + ] + + mock_pick_unused_port.return_value = 29007 + + mock_process = mock_popen.return_value + mock_process.stdout = io.StringIO( + f"Forwarding from 127.0.0.1:{mock_pick_unused_port.return_value} ->" + " 8080\n" + ) + proxy_job_name = "test-proxy" + + # Act + with isc_pathways._ISCPathways( + cluster="test-cluster", + project="test-project", + region="test-region", + gcs_bucket="test-bucket", + pathways_service="test-service:1234", + expected_tpu_instances={"tpuv5:4x4x4": 1}, + proxy_job_name=proxy_job_name, + proxy_server_image="test-image:latest", + proxy_options=isc_pathways.ProxyOptions(use_insecure_credentials=False), + collect_service_metrics=True, + ): + # Assertions inside the context + self.assertEqual( + os.environ["JAX_PLATFORMS"], + isc_pathways._JAX_PLATFORM_PROXY, + ) + self.assertEqual( + os.environ["JAX_BACKEND_TARGET"], "grpc://127.0.0.1:29007", + ) + + expected_jax_calls = [ + mock.call("jax_platforms", "proxy"), + mock.call("jax_backend_target", "grpc://127.0.0.1:29007"), + ] + mock_jax_config_update.assert_has_calls( + expected_jax_calls, any_order=True + ) + self.assertEqual( + os.environ.get("JAX_PLATFORMS"), + isc_pathways._JAX_PLATFORM_PROXY, + ) + self.assertEqual( + os.environ.get("JAX_BACKEND_TARGET"), + f"{isc_pathways._JAX_BACKEND_TARGET_HOSTNAME}:{mock_pick_unused_port.return_value}", + ) + + # Assertions outside the context (cleanup) + self.assertEqual( + os.environ.get("JAX_PLATFORMS"), "original_platform" + ) + self.assertEqual( + os.environ.get("JAX_BACKEND_TARGET"), "original_target" + ) + + restoration_calls = [ + mock.call("jax_platforms", "original_platform_config"), + mock.call("jax_backend_target", "original_target_config"), + ] + mock_jax_config_update.assert_has_calls(restoration_calls, any_order=True) + + mock_clear_backends.assert_called_once() + mock_deploy.assert_called_once() + mock_clear_caches.assert_called_once() + mock_gc_collect.assert_called_once() + mock_metrics_collector.assert_called_once_with( + "test-project", "test-cluster", "test-proxy" + ) + mock_collector_instance.record_active_user.assert_not_called() + mock_collector_instance.record_requested_capacity.assert_called_once_with( + 64 + ) + mock_run.assert_called_with( + [ + "kubectl", + "delete", + "job", + "-n", + "default", + "--ignore-not-found", + "--", + proxy_job_name, + ], + check=True, + capture_output=True, + text=True, + ) + + def test_connect_success(self): + """Tests that connect calls the dependencies and yields the manager.""" + # Arrange + self.enter_context(mock.patch.dict(os.environ, {"USER": "testuser"})) + mock_random = self.enter_context( + mock.patch( + "pathwaysutils.experimental.shared_pathways_service.isc_pathways.random", + autospec=True, + ) + ) + mock_random.choices.return_value = list("abcde") + expected_proxy_job_name = "isc-proxy-testuser-abcde" + mock_validate_tpu = self.enter_context( + mock.patch.object( + isc_pathways.validators, "validate_tpu_instances", autospec=True + ) + ) + mock_validate_proxy_image = self.enter_context( + mock.patch.object( + isc_pathways.validators, + "validate_proxy_server_image", + autospec=True, + ) + ) + mock_fetch_creds = self.enter_context( + mock.patch.object( + isc_pathways.gke_utils, "fetch_cluster_credentials", autospec=True + ) + ) + mock_isc_pathways = self.enter_context( + mock.patch.object(isc_pathways, "_ISCPathways", autospec=True) + ) + mock_thread = self.enter_context( + mock.patch("threading.Thread", autospec=True) + ) + + cluster = "test-cluster" + project = "test-project" + region = "test-region" + bucket = "test-bucket" + pathways_service = "test-service:1234" + expected_instances = {"tpuv5:4x4x4": 1} + + mock_manager_instance = ( + mock_isc_pathways.return_value.__enter__.return_value + ) + mock_manager_instance.proxy_pod_name = "test-pod-123" + mock_manager_instance.expected_tpu_instances = expected_instances + mock_manager_instance.metrics_collector = mock.Mock() + mock_manager_instance.start_time = 100.0 + + # Act + with isc_pathways.connect( + cluster=cluster, + project=project, + region=region, + gcs_bucket=bucket, + pathways_service=pathways_service, + expected_tpu_instances=expected_instances, + ) as tm: + # Assert + mock_validate_tpu.assert_called_once_with(expected_instances) + mock_validate_proxy_image.assert_called_once_with( + isc_pathways.DEFAULT_PROXY_IMAGE + ) + mock_fetch_creds.assert_called_once_with( + cluster_name=cluster, project_id=project, location=region + ) + mock_isc_pathways.assert_called_once_with( + cluster=cluster, + project=project, + region=region, + gcs_bucket=bucket, + pathways_service=pathways_service, + expected_tpu_instances=expected_instances, + proxy_job_name=expected_proxy_job_name, + proxy_server_image=isc_pathways.DEFAULT_PROXY_IMAGE, + proxy_options=isc_pathways.ProxyOptions(), + collect_service_metrics=False, + ) + self.assertIs(tm, mock_manager_instance) + + # Verify thread start + mock_thread.assert_called_once_with( + target=isc_pathways._wait_for_placement, + args=( + "test-pod-123", + 1, + isc_pathways.gke_utils.stream_pod_logs, + mock_manager_instance.metrics_collector, + mock_manager_instance.start_time, + mock_manager_instance.total_chips, + ), + daemon=True, + ) + mock_thread.return_value.start.assert_called_once() + + def test_connect_with_non_existent_cluster_raises_error(self): + """Tests that connect raises an error if the cluster doesn't exist.""" + mock_fetch_creds = self.enter_context( + mock.patch.object( + isc_pathways.gke_utils, "fetch_cluster_credentials", autospec=True + ) + ) + mock_fetch_creds.side_effect = subprocess.CalledProcessError( + returncode=1, cmd="gcloud", stderr="cluster not found" + ) + with self.assertRaises(subprocess.CalledProcessError): + with isc_pathways.connect( + cluster="non-existent-cluster", + project="test-project", + region="test-zone", + gcs_bucket="test-bucket", + pathways_service="test-service:1234", + expected_tpu_instances={"tpuv6e:2x2": 1}, + ): + self.fail("ISCPathways context should not be entered.") + + def test_connect_with_tpuv3_raises_error(self): + """Tests that connect raises an error for tpuv3 configurations.""" + with self.assertRaisesRegex( + ValueError, + "Unrecognized instance format: tpuv3:4x4.", + ): + with isc_pathways.connect( + cluster="test-cluster", + project="test-project", + region="test-zone", + gcs_bucket="test-bucket", + pathways_service="test-service:1234", + expected_tpu_instances={ + "tpuv3:4x4": 1 + }, + ): + self.fail("ISCPathways context should not be entered.") + + def test_connect_with_invalid_proxy_image_raises_error(self): + """Tests that connect raises an error for invalid proxy image.""" + with self.assertRaisesRegex( + ValueError, + "Proxy server image cannot be empty.", + ): + with isc_pathways.connect( + cluster="test-cluster", + project="test-project", + region="test-zone", + gcs_bucket="test-bucket", + pathways_service="test-service:1234", + expected_tpu_instances={"tpuv6e:2x2": 1}, + proxy_server_image="", + ): + self.fail("ISCPathways context should not be entered.") + + def test_connect_passes_collect_service_metrics(self): + """Tests that connect passes collect_service_metrics to _ISCPathways.""" + self.enter_context(mock.patch.dict(os.environ, {"USER": "testuser"})) + self.enter_context( + mock.patch.object( + isc_pathways.gke_utils, "fetch_cluster_credentials", autospec=True + ) + ) + mock_isc_pathways = self.enter_context( + mock.patch.object(isc_pathways, "_ISCPathways", autospec=True) + ) + self.enter_context(mock.patch("threading.Thread", autospec=True)) + self.enter_context( + mock.patch.object( + isc_pathways.validators, "validate_tpu_instances", autospec=True + ) + ) + self.enter_context( + mock.patch.object( + isc_pathways.validators, + "validate_proxy_server_image", + autospec=True, + ) + ) + + mock_manager_instance = ( + mock_isc_pathways.return_value.__enter__.return_value + ) + mock_manager_instance.proxy_pod_name = "test-pod-123" + mock_manager_instance.expected_tpu_instances = {"tpuv6e:2x2": 1} + + with isc_pathways.connect( + cluster="test-cluster", + project="test-project", + region="test-region", + gcs_bucket="test-bucket", + pathways_service="test-service:1234", + expected_tpu_instances={"tpuv6e:2x2": 1}, + collect_service_metrics=True, + ): + pass + + mock_isc_pathways.assert_called_once() + _, kwargs = mock_isc_pathways.call_args + self.assertTrue(kwargs["collect_service_metrics"]) + + def test_connect_with_sidecar_validation_success(self): + self.enter_context(mock.patch.dict(os.environ, {"USER": "testuser"})) + self.enter_context( + mock.patch.object( + isc_pathways.gke_utils, "fetch_cluster_credentials", autospec=True + ) + ) + mock_get_sidecar = self.enter_context( + mock.patch.object( + isc_pathways.gke_utils, "get_worker_sidecar_image", autospec=True + ) + ) + mock_get_sidecar.return_value = ( + "us-docker.pkg.dev/.../sidecar:20260423-python_3.12-jax_0.10.0" + ) + mock_validate_versions = self.enter_context( + mock.patch.object( + isc_pathways.validators, + "validate_sidecar_image_versions", + autospec=True, + ) + ) + mock_isc_pathways = self.enter_context( + mock.patch.object(isc_pathways, "_ISCPathways", autospec=True) + ) + self.enter_context(mock.patch("threading.Thread", autospec=True)) + + mock_manager_instance = ( + mock_isc_pathways.return_value.__enter__.return_value + ) + mock_manager_instance.proxy_pod_name = "test-pod-123" + mock_manager_instance.expected_tpu_instances = {"tpuv6e:2x2": 1} + + with isc_pathways.connect( + cluster="test-cluster", + project="test-project", + region="test-region", + gcs_bucket="test-bucket", + pathways_service="test-service:1234", + expected_tpu_instances={"tpuv6e:2x2": 1}, + proxy_options=["sidecar:true"], + ): + pass + + mock_get_sidecar.assert_called_once_with( + pathways_service="test-service:1234" + ) + mock_validate_versions.assert_called_once_with( + "us-docker.pkg.dev/.../sidecar:20260423-python_3.12-jax_0.10.0" + ) + + def test_connect_with_sidecar_validation_mismatch_raises_error(self): + self.enter_context(mock.patch.dict(os.environ, {"USER": "testuser"})) + self.enter_context( + mock.patch.object( + isc_pathways.gke_utils, "fetch_cluster_credentials", autospec=True + ) + ) + mock_get_sidecar = self.enter_context( + mock.patch.object( + isc_pathways.gke_utils, "get_worker_sidecar_image", autospec=True + ) + ) + mock_get_sidecar.return_value = ( + "us-docker.pkg.dev/.../sidecar:20260423-python_3.12-jax_0.10.0" + ) + mock_validate_versions = self.enter_context( + mock.patch.object( + isc_pathways.validators, + "validate_sidecar_image_versions", + autospec=True, + side_effect=ValueError("Python version mismatch"), + ) + ) + + with self.assertRaisesRegex(ValueError, "Python version mismatch"): + with isc_pathways.connect( + cluster="test-cluster", + project="test-project", + region="test-region", + gcs_bucket="test-bucket", + pathways_service="test-service:1234", + expected_tpu_instances={"tpuv6e:2x2": 1}, + proxy_options=["sidecar:true"], + ): + pass + + mock_get_sidecar.assert_called_once_with( + pathways_service="test-service:1234" + ) + mock_validate_versions.assert_called_once_with( + "us-docker.pkg.dev/.../sidecar:20260423-python_3.12-jax_0.10.0" + ) + + +if __name__ == "__main__": + absltest.main() diff --git a/pathwaysutils/test/experimental/shared_pathways_service/metrics_collector_test.py b/pathwaysutils/test/experimental/shared_pathways_service/metrics_collector_test.py new file mode 100644 index 0000000..9f8f3d5 --- /dev/null +++ b/pathwaysutils/test/experimental/shared_pathways_service/metrics_collector_test.py @@ -0,0 +1,247 @@ +"""Unit tests for the MetricsCollector class.""" + +from unittest import mock +from absl.testing import absltest +from pathwaysutils.experimental.shared_pathways_service import metrics_collector + + +class MetricsCollectorTest(absltest.TestCase): + + @mock.patch( + "pathwaysutils.experimental.shared_pathways_service.metrics_collector.monitoring_v3.MetricServiceClient" + ) + def test_record_active_user(self, mock_client_class): + mock_client = mock.Mock() + mock_client_class.return_value = mock_client + + collector = metrics_collector.MetricsCollector( + "gcp_project", "gke_cluster", "proxy_job" + ) + + collector.record_active_user(True) + collector.flush() + + mock_client.create_time_series.assert_called_once() + _, kwargs = mock_client.create_time_series.call_args + self.assertEqual(kwargs["name"], "projects/gcp_project") + time_series = kwargs["time_series"][0] + self.assertEqual( + time_series.metric.type, + "custom.googleapis.com/shared_pathways_service/num_active_users", + ) + self.assertEqual(time_series.points[0].value.int64_value, 1) + self.assertEqual( + time_series.metric.labels["cluster_name"], "gke_cluster" + ) + self.assertEqual(time_series.metric.labels["job_name"], "proxy_job") + + @mock.patch( + "pathwaysutils.experimental.shared_pathways_service.metrics_collector.monitoring_v3.MetricServiceClient" + ) + def test_record_capacity_in_use(self, mock_client_class): + mock_client = mock.Mock() + mock_client_class.return_value = mock_client + + collector = metrics_collector.MetricsCollector("gcp_project", "gke_cluster") + collector.record_capacity_in_use(8) + collector.flush() + + mock_client.create_time_series.assert_called_once() + _, kwargs = mock_client.create_time_series.call_args + time_series = kwargs["time_series"][0] + self.assertEqual( + time_series.metric.type, + "custom.googleapis.com/shared_pathways_service/capacity_in_use", + ) + self.assertEqual(time_series.points[0].value.int64_value, 8) + + @mock.patch( + "pathwaysutils.experimental.shared_pathways_service.metrics_collector.monitoring_v3.MetricServiceClient" + ) + def test_record_assignment_time(self, mock_client_class): + mock_client = mock.Mock() + mock_client_class.return_value = mock_client + + collector = metrics_collector.MetricsCollector("gcp_project", "gke_cluster") + collector.record_assignment_time(12.5) + collector.flush() + + mock_client.create_time_series.assert_called_once() + _, kwargs = mock_client.create_time_series.call_args + time_series = kwargs["time_series"][0] + self.assertEqual( + time_series.metric.type, + "custom.googleapis.com/shared_pathways_service/assignment_time", + ) + self.assertEqual(time_series.points[0].value.double_value, 12.5) + + @mock.patch( + "pathwaysutils.experimental.shared_pathways_service.metrics_collector.monitoring_v3.MetricServiceClient" + ) + def test_record_successful_request(self, mock_client_class): + mock_client = mock.Mock() + mock_client_class.return_value = mock_client + + collector = metrics_collector.MetricsCollector("gcp_project", "gke_cluster") + collector.record_successful_request() + collector.flush() + + mock_client.create_time_series.assert_called_once() + _, kwargs = mock_client.create_time_series.call_args + time_series = kwargs["time_series"][0] + self.assertEqual( + time_series.metric.type, + "custom.googleapis.com/shared_pathways_service/num_successful_reqs", + ) + self.assertEqual(time_series.points[0].value.int64_value, 1) + + @mock.patch( + "pathwaysutils.experimental.shared_pathways_service.metrics_collector.monitoring_v3.MetricServiceClient" + ) + def test_record_user_waiting(self, mock_client_class): + mock_client = mock.Mock() + mock_client_class.return_value = mock_client + + collector = metrics_collector.MetricsCollector("gcp_project", "gke_cluster") + collector.record_user_waiting(True) + collector.flush() + + mock_client.create_time_series.assert_called_once() + _, kwargs = mock_client.create_time_series.call_args + time_series = kwargs["time_series"][0] + self.assertEqual( + time_series.metric.type, + "custom.googleapis.com/shared_pathways_service/num_users_waiting", + ) + self.assertEqual(time_series.points[0].value.int64_value, 1) + + @mock.patch( + "pathwaysutils.experimental.shared_pathways_service.metrics_collector.monitoring_v3.MetricServiceClient" + ) + def test_record_requested_capacity(self, mock_client_class): + mock_client = mock.Mock() + mock_client_class.return_value = mock_client + + collector = metrics_collector.MetricsCollector("gcp_project", "gke_cluster") + collector.record_requested_capacity(64) + collector.flush() + + mock_client.create_time_series.assert_called_once() + _, kwargs = mock_client.create_time_series.call_args + time_series = kwargs["time_series"][0] + self.assertEqual( + time_series.metric.type, + "custom.googleapis.com/shared_pathways_service/requested_capacity", + ) + self.assertEqual(time_series.points[0].value.int64_value, 64) + + @mock.patch( + "pathwaysutils.experimental.shared_pathways_service.metrics_collector.monitoring_v3.MetricServiceClient" + ) + def test_initialize_descriptors(self, mock_client_class): + mock_client = mock.Mock() + mock_client.get_metric_descriptor.side_effect = ( + metrics_collector.exceptions.NotFound("not found") + ) + mock_client_class.return_value = mock_client + + _ = metrics_collector.MetricsCollector("gcp_project", "gke_cluster") + + self.assertEqual(mock_client.create_metric_descriptor.call_count, 6) + + # Verify units for each descriptor + calls = mock_client.create_metric_descriptor.call_args_list + + call_args = [c.kwargs["metric_descriptor"] for c in calls] + + def find_descriptor(name): + for desc in call_args: + if desc["type"].endswith(f"/{name}"): + return desc + return None + + self.assertEqual(find_descriptor("num_active_users")["unit"], "1") + self.assertEqual(find_descriptor("capacity_in_use")["unit"], "chips") + self.assertEqual(find_descriptor("assignment_time")["unit"], "s") + self.assertEqual(find_descriptor("num_successful_reqs")["unit"], "1") + self.assertEqual(find_descriptor("num_users_waiting")["unit"], "1") + self.assertEqual(find_descriptor("requested_capacity")["unit"], "chips") + + @mock.patch( + "pathwaysutils.experimental.shared_pathways_service.metrics_collector.monitoring_v3.MetricServiceClient" + ) + def test_initialize_descriptors_already_exists(self, mock_client_class): + mock_client = mock.Mock() + # By default, get_metric_descriptor returns a mock without raising, + # meaning the metric already exists. + mock_client_class.return_value = mock_client + + _ = metrics_collector.MetricsCollector("gcp_project", "gke_cluster") + + self.assertEqual(mock_client.get_metric_descriptor.call_count, 6) + self.assertEqual(mock_client.create_metric_descriptor.call_count, 0) + + @mock.patch( + "pathwaysutils.experimental.shared_pathways_service.metrics_collector.monitoring_v3.MetricServiceClient" + ) + def test_buffer_queue_and_throttle(self, mock_client_class): + mock_client = mock.Mock() + mock_client_class.return_value = mock_client + + collector = metrics_collector.MetricsCollector("gcp_project", "gke_cluster") + + # Send 2 states back-to-back + collector.record_user_waiting(True) + collector.record_user_waiting(False) + + # Queue should hold both + self.assertLen(collector._buffer["num_users_waiting"], 2) + + # First flush sends [0] + collector.flush() + self.assertEqual(mock_client.create_time_series.call_count, 1) + self.assertLen(collector._buffer["num_users_waiting"], 1) + + # Immediate second flush should do NOTHING due to 10.5s limit + collector.flush() + self.assertEqual(mock_client.create_time_series.call_count, 1) + self.assertLen(collector._buffer["num_users_waiting"], 1) + + # Shift time forward by 11s and flush -> sends remaining [1] + with mock.patch( + "time.time", + return_value=collector._last_sent_time["num_users_waiting"] + 11.0, + ): + collector.flush() + + self.assertEqual(mock_client.create_time_series.call_count, 2) + self.assertNotIn("num_users_waiting", collector._buffer) + + @mock.patch( + "pathwaysutils.experimental.shared_pathways_service.metrics_collector.monitoring_v3.MetricServiceClient" + ) + def test_shutdown_exhausts_queue(self, mock_client_class): + mock_client = mock.Mock() + mock_client_class.return_value = mock_client + collector = metrics_collector.MetricsCollector("gcp_project", "gke_cluster") + collector.record_capacity_in_use(16) + collector.record_capacity_in_use(0) + + current_time = [1000.0] + + def mock_sleep_func(seconds): + current_time[0] += seconds + + def mock_time_func(): + return current_time[0] + + with mock.patch("time.sleep", side_effect=mock_sleep_func) as mock_sleep: + with mock.patch("time.time", side_effect=mock_time_func): + collector._shutdown() + + self.assertEqual(mock_client.create_time_series.call_count, 2) + mock_sleep.assert_called() + + +if __name__ == "__main__": + absltest.main() diff --git a/pathwaysutils/test/experimental/shared_pathways_service/run_connect_example_test.py b/pathwaysutils/test/experimental/shared_pathways_service/run_connect_example_test.py new file mode 100644 index 0000000..8912838 --- /dev/null +++ b/pathwaysutils/test/experimental/shared_pathways_service/run_connect_example_test.py @@ -0,0 +1,135 @@ +from unittest import mock + +from absl.testing import absltest +from absl.testing import flagsaver +import numpy as np + + +class RunConnectExampleTest(absltest.TestCase): + """Tests the logic from the run_connect_example.py script.""" + + @flagsaver.flagsaver( + cluster="random-cluster-name", + project="random-project-id", + region="random-region", + gcs_bucket="random-bucket", + pathways_service="random-pathways-service:1234", + tpu_type="tpuv6e:2x2", + tpu_count=2, + ) + def test_run_connect_example_main(self): + """Tests that the main function calls connect and executes the logic.""" + # Import inside the test to avoid flag parsing errors on module load. + from pathwaysutils.experimental.shared_pathways_service import run_connect_example + + mock_connect = self.enter_context( + mock.patch.object( + run_connect_example.isc_pathways, "connect", autospec=True + ) + ) + mock_pprint = self.enter_context(mock.patch("pprint.pprint", autospec=True)) + + run_connect_example.main(["unused_argv"]) + + mock_connect.assert_called_once_with( + cluster="random-cluster-name", + project="random-project-id", + region="random-region", + gcs_bucket="random-bucket", + pathways_service="random-pathways-service:1234", + expected_tpu_instances={"tpuv6e:2x2": 2}, + proxy_job_name=None, + proxy_server_image=( + "us-docker.pkg.dev/cloud-tpu-v2-images/pathways/proxy_server:latest" + ), + proxy_options=None, + collect_service_metrics=False, + ) + self.assertEqual(mock_pprint.call_count, 2) + np.testing.assert_array_equal( + mock_pprint.call_args_list[0].args[0], np.zeros(5) + ) + np.testing.assert_array_equal( + mock_pprint.call_args_list[1].args[0], np.ones(5) + ) + + @flagsaver.flagsaver( + cluster="random-cluster-name", + project="random-project-id", + region="random-region", + gcs_bucket="random-bucket", + pathways_service="random-pathways-service:1234", + tpu_type="tpuv6e:2x2", + tpu_count=2, + proxy_job_name="test-job-name", + proxy_server_image="test-image", + proxy_options=["use_insecure_credentials:true"], + ) + def test_run_connect_example_main_with_optional_flags(self): + """Tests that main passes optional flags to connect.""" + # Import inside the test to avoid flag parsing errors on module load. + from pathwaysutils.experimental.shared_pathways_service import run_connect_example + + mock_connect = self.enter_context( + mock.patch.object( + run_connect_example.isc_pathways, "connect", autospec=True + ) + ) + self.enter_context(mock.patch("pprint.pprint", autospec=True)) + + run_connect_example.main(["unused_argv"]) + + mock_connect.assert_called_once_with( + cluster="random-cluster-name", + project="random-project-id", + region="random-region", + gcs_bucket="random-bucket", + pathways_service="random-pathways-service:1234", + expected_tpu_instances={"tpuv6e:2x2": 2}, + proxy_job_name="test-job-name", + proxy_server_image="test-image", + proxy_options=["use_insecure_credentials:true"], + collect_service_metrics=False, + ) + + @flagsaver.flagsaver( + cluster="random-cluster-name", + project="random-project-id", + region="random-region", + gcs_bucket="random-bucket", + pathways_service="random-pathways-service:1234", + tpu_type="tpuv6e:2x2", + tpu_count=2, + collect_service_metrics=True, + ) + def test_run_connect_example_main_with_metrics_enabled(self): + """Tests that main passes collect_service_metrics flag to connect.""" + from pathwaysutils.experimental.shared_pathways_service import run_connect_example + + mock_connect = self.enter_context( + mock.patch.object( + run_connect_example.isc_pathways, "connect", autospec=True + ) + ) + self.enter_context(mock.patch("pprint.pprint", autospec=True)) + + run_connect_example.main(["unused_argv"]) + + mock_connect.assert_called_once_with( + cluster="random-cluster-name", + project="random-project-id", + region="random-region", + gcs_bucket="random-bucket", + pathways_service="random-pathways-service:1234", + expected_tpu_instances={"tpuv6e:2x2": 2}, + proxy_job_name=None, + proxy_server_image=( + "us-docker.pkg.dev/cloud-tpu-v2-images/pathways/proxy_server:latest" + ), + proxy_options=None, + collect_service_metrics=True, + ) + + +if __name__ == "__main__": + absltest.main() diff --git a/pathwaysutils/test/experimental/shared_pathways_service/run_workload_test.py b/pathwaysutils/test/experimental/shared_pathways_service/run_workload_test.py new file mode 100644 index 0000000..a7ddffb --- /dev/null +++ b/pathwaysutils/test/experimental/shared_pathways_service/run_workload_test.py @@ -0,0 +1,206 @@ +import contextlib +import subprocess +import sys +import types +from typing import Any, Generator +from unittest import mock + +from absl.testing import absltest +from absl.testing import flagsaver +from pathwaysutils.experimental.shared_pathways_service import run_workload + + +class FakeConnect: + """A fake connection manager that tracks entry and exit.""" + + def __init__(self, **kwargs: Any): + self.kwargs = kwargs + self.entered = False + self.exited = False + + def __enter__(self) -> "FakeConnect": + self.entered = True + return self + + def __exit__( + self, + exc_type: type[BaseException] | None, + exc_val: BaseException | None, + exc_tb: types.TracebackType | None, + ) -> None: + self.exited = True + + +class RunTpuWorkloadTest(absltest.TestCase): + + def test_run_workload_success(self): + fake_instances = [] + + @contextlib.contextmanager + def fake_connect_fn(**kwargs: Any) -> Generator[FakeConnect, None, None]: + fake = FakeConnect(**kwargs) + fake_instances.append(fake) + with fake as f: + yield f + + mock_run = self.enter_context( + mock.patch.object(subprocess, "run", autospec=True) + ) + + run_workload.run_command( + cluster="test-cluster", + project="test-project", + region="test-region", + gcs_bucket="test-bucket", + pathways_service="test-service:1234", + tpu_type="tpuv6e:4x8", + tpu_count=1, + command="echo hello", + connect_fn=fake_connect_fn, + ) + + with self.subTest("Connect function called correctly"): + self.assertLen(fake_instances, 1) + fake = fake_instances[0] + self.assertEqual(fake.kwargs["cluster"], "test-cluster") + self.assertEqual(fake.kwargs["expected_tpu_instances"], {"tpuv6e:4x8": 1}) + + with self.subTest("Context manager lifecycle"): + fake = fake_instances[0] + self.assertTrue(fake.entered) + self.assertTrue(fake.exited) + + with self.subTest("Command executed"): + mock_run.assert_called_once_with( + ["echo", "hello"], check=True, env=mock.ANY + ) + + def test_run_command_runs_command_inside_context(self): + """Verifies that the command is executed while the connection is active.""" + connection_active_during_run = False + + @contextlib.contextmanager + def fake_connect_fn(**kwargs: Any) -> Generator[None, None, None]: + del kwargs + yield None + + def mock_run_side_effect(*args: Any, **kwargs: Any) -> None: + nonlocal connection_active_during_run + connection_active_during_run = True + + mock_run = self.enter_context( + mock.patch.object(subprocess, "run", autospec=True) + ) + mock_run.side_effect = mock_run_side_effect + + run_workload.run_command( + cluster="test-cluster", + project="test-project", + region="test-region", + gcs_bucket="test-bucket", + pathways_service="test-service:1234", + tpu_type="tpuv6e:4x8", + tpu_count=1, + command="echo hello", + connect_fn=fake_connect_fn, + ) + + self.assertTrue(connection_active_during_run) + + def test_run_command_error(self): + + @contextlib.contextmanager + def fake_connect_fn(**kwargs: Any) -> Generator[None, None, None]: + del kwargs + yield None + + mock_run = self.enter_context( + mock.patch.object(subprocess, "run", autospec=True) + ) + mock_run.side_effect = subprocess.CalledProcessError(1, "false") + + with self.assertRaises(subprocess.CalledProcessError): + run_workload.run_command( + cluster="test-cluster", + project="test-project", + region="test-region", + gcs_bucket="test-bucket", + pathways_service="test-service:1234", + tpu_type="tpuv6e:4x8", + tpu_count=1, + command="false", + connect_fn=fake_connect_fn, + ) + + @flagsaver.flagsaver( + cluster="test-cluster", + project="test-project", + region="test-region", + gcs_bucket="test-bucket", + pathways_service="test-service:1234", + tpu_type="tpuv6e:4x8", + tpu_count=1, + command="echo hello", + ) + def test_main_calls_run_command(self): + with mock.patch.object( + run_workload, "run_command", autospec=True + ) as mock_run_command: + run_workload.main(["unused_argv"]) + mock_run_command.assert_called_once_with( + cluster="test-cluster", + project="test-project", + region="test-region", + gcs_bucket="test-bucket", + pathways_service="test-service:1234", + tpu_type="tpuv6e:4x8", + tpu_count=1, + command="echo hello", + proxy_server_image="", + proxy_options=[], + collect_service_metrics=False, + ) + + @flagsaver.flagsaver( + cluster="test-cluster", + project="test-project", + region="test-region", + gcs_bucket="test-bucket", + pathways_service="test-service:1234", + tpu_type="tpuv6e:4x8", + tpu_count=1, + command="echo hello", + collect_service_metrics=True, + ) + def test_main_calls_run_command_with_metrics_enabled(self): + with mock.patch.object( + run_workload, "run_command", autospec=True + ) as mock_run_command: + run_workload.main(["unused_argv"]) + mock_run_command.assert_called_once_with( + cluster="test-cluster", + project="test-project", + region="test-region", + gcs_bucket="test-bucket", + pathways_service="test-service:1234", + tpu_type="tpuv6e:4x8", + tpu_count=1, + command="echo hello", + proxy_server_image="", + proxy_options=[], + collect_service_metrics=True, + ) + + +if __name__ == "__main__": + # Provide dummy values for required flags to satisfy absl verification. + # These are overridden by tests as needed using flagsaver or direct arguments. + sys.argv.extend([ + "--cluster=dummy", + "--project=dummy", + "--region=dummy", + "--gcs_bucket=dummy", + "--pathways_service=dummy:1234", + "--command=dummy", + ]) + absltest.main() diff --git a/pathwaysutils/test/experimental/shared_pathways_service/start_vscode_on_cpu_np_test.py b/pathwaysutils/test/experimental/shared_pathways_service/start_vscode_on_cpu_np_test.py new file mode 100644 index 0000000..28813a7 --- /dev/null +++ b/pathwaysutils/test/experimental/shared_pathways_service/start_vscode_on_cpu_np_test.py @@ -0,0 +1,324 @@ +import os +import random +import string +from unittest import mock + +from absl import app +from absl import flags +from absl.testing import absltest +from absl.testing import flagsaver +from pathwaysutils.experimental.shared_pathways_service import gke_utils +from pathwaysutils.experimental.shared_pathways_service import start_vscode_on_cpu_np + + +class StartVSCodeOnCPUNPTest(absltest.TestCase): + + + @flagsaver.flagsaver( + namespace="my-ns", + image="my-image", + password="password123", + instance_type="cpu-type", + ) + def test_prepare_deployment_yaml_success(self): + template_content = ( + "name: ${NAME}\n" + "namespace: ${NAMESPACE}\n" + "image: ${IMAGE}\n" + "password: ${PASSWORD}\n" + "node: ${INSTANCE_TYPE}\n" + "service: ${SERVICE_NAME}\n" + "port: ${PORT}" + ) + self.enter_context( + mock.patch( + "builtins.open", + new_callable=mock.mock_open, + read_data=template_content, + ) + ) + yaml_content = start_vscode_on_cpu_np._prepare_deployment_yaml( + service_name="test-service", remote_port=9090 + ) + self.assertEqual( + yaml_content, + "name: test-service\n" + "namespace: my-ns\n" + "image: my-image\n" + "password: password123\n" + "node: cpu-type\n" + "service: test-service\n" + "port: 9090", + ) + + def test_prepare_deployment_yaml_missing_file_raises_value_error(self): + self.enter_context(mock.patch("builtins.open", side_effect=OSError())) + with self.assertRaisesRegex(ValueError, "Could not read template file:"): + start_vscode_on_cpu_np._prepare_deployment_yaml( + service_name="test-service", remote_port=9090 + ) + + @flagsaver.flagsaver(namespace="my-ns") + def test_deploy_vscode_success(self): + mock_deploy = self.enter_context( + mock.patch.object(gke_utils, "deploy_gke_yaml", autospec=True) + ) + mock_wait_dep = self.enter_context( + mock.patch.object(gke_utils, "wait_for_deployment", autospec=True) + ) + mock_wait_svc = self.enter_context( + mock.patch.object(gke_utils, "wait_for_service_ip", autospec=True) + ) + mock_wait_svc.return_value = "10.0.0.1" + + start_vscode_on_cpu_np._deploy_vscode("test-service", "dummy-yaml") + + mock_deploy.assert_called_once_with("dummy-yaml", action="create") + mock_wait_dep.assert_called_once_with("test-service", "my-ns") + mock_wait_svc.assert_called_once_with("test-service", "my-ns") + + @flagsaver.flagsaver(namespace="my-ns") + def test_deploy_vscode_service_ip_failure(self): + mock_deploy = self.enter_context( + mock.patch.object(gke_utils, "deploy_gke_yaml", autospec=True) + ) + mock_wait_dep = self.enter_context( + mock.patch.object(gke_utils, "wait_for_deployment", autospec=True) + ) + mock_wait_svc = self.enter_context( + mock.patch.object(gke_utils, "wait_for_service_ip", autospec=True) + ) + mock_wait_svc.side_effect = RuntimeError("Service IP timeout") + + with self.assertLogs(level="WARNING") as log_capture: + start_vscode_on_cpu_np._deploy_vscode("test-service", "dummy-yaml") + + mock_deploy.assert_called_once_with("dummy-yaml", action="create") + mock_wait_dep.assert_called_once_with("test-service", "my-ns") + mock_wait_svc.assert_called_once_with("test-service", "my-ns") + self.assertTrue( + any( + "Could not get service IP" in record.message + for record in log_capture.records + ) + ) + + @flagsaver.flagsaver(namespace="my-ns") + def test_start_port_forwarding_keyboard_interrupt(self): + mock_enable = self.enter_context( + mock.patch.object( + gke_utils, "enable_port_forwarding", autospec=True + ) + ) + mock_process = mock.Mock() + mock_enable.return_value = (8080, mock_process) + mock_sleep = self.enter_context( + mock.patch("time.sleep", side_effect=KeyboardInterrupt()) + ) + + start_vscode_on_cpu_np._start_port_forwarding("test-service", 9090) + + mock_enable.assert_called_once_with( + remote_server="svc/test-service", + server_port=9090, + namespace="my-ns", + ) + mock_sleep.assert_called_once_with(1) + mock_process.terminate.assert_called_once() + mock_process.wait.assert_called_once() + + @flagsaver.flagsaver(namespace="my-ns") + def test_start_port_forwarding_other_exception(self): + mock_enable = self.enter_context( + mock.patch.object( + gke_utils, "enable_port_forwarding", autospec=True + ) + ) + mock_process = mock.Mock() + mock_enable.return_value = (8080, mock_process) + mock_sleep = self.enter_context( + mock.patch("time.sleep", side_effect=Exception("Forwarding crash")) + ) + + start_vscode_on_cpu_np._start_port_forwarding("test-service", 9090) + + mock_enable.assert_called_once_with( + remote_server="svc/test-service", + server_port=9090, + namespace="my-ns", + ) + mock_sleep.assert_called_once_with(1) + mock_process.terminate.assert_called_once() + mock_process.wait.assert_called_once() + + def test_cleanup_gke_resources_success(self): + mock_delete_resource = self.enter_context( + mock.patch.object(gke_utils, "delete_gke_resource", autospec=True) + ) + + start_vscode_on_cpu_np._cleanup_gke_resources("test-service", "my-ns") + + mock_delete_resource.assert_has_calls([ + mock.call("deployment", "test-service", "my-ns"), + mock.call("service", "test-service", "my-ns"), + ]) + + def test_cleanup_gke_resources_ignores_exceptions(self): + mock_delete_resource = self.enter_context( + mock.patch.object(gke_utils, "delete_gke_resource", autospec=True) + ) + mock_delete_resource.side_effect = Exception("delete fail") + + # Should not raise an exception + start_vscode_on_cpu_np._cleanup_gke_resources("test-service", "my-ns") + + mock_delete_resource.assert_has_calls([ + mock.call("deployment", "test-service", "my-ns"), + mock.call("service", "test-service", "my-ns"), + ]) + + def test_main_too_many_args(self): + with self.assertRaises(app.UsageError): + start_vscode_on_cpu_np.main(["script_name", "extra_arg"]) + + @flagsaver.flagsaver(dry_run=True, name="my-vscode") + @mock.patch.dict(os.environ, {"USER": "testuser"}) + def test_main_dry_run(self): + mock_prepare = self.enter_context( + mock.patch.object( + start_vscode_on_cpu_np, "_prepare_deployment_yaml", autospec=True + ) + ) + mock_prepare.return_value = "dummy-yaml" + mock_deploy = self.enter_context( + mock.patch.object( + start_vscode_on_cpu_np, "_deploy_vscode", autospec=True + ) + ) + mock_cleanup = self.enter_context( + mock.patch.object( + start_vscode_on_cpu_np, "_cleanup_gke_resources", autospec=True + ) + ) + + mock_choices = self.enter_context( + mock.patch.object(random, "choices", autospec=True) + ) + mock_choices.return_value = ["a", "b", "c", "d"] + + start_vscode_on_cpu_np.main(["script_name"]) + + mock_prepare.assert_called_once_with("my-vscode-testuser-abcd", 8080) + mock_choices.assert_called_once_with( + string.ascii_lowercase + string.digits, k=4 + ) + mock_deploy.assert_not_called() + mock_cleanup.assert_not_called() + + @flagsaver.flagsaver(dry_run=False, name="my-vscode", namespace="my-ns") + @mock.patch.dict(os.environ, {"USER": "testuser"}) + def test_main_real_run_success(self): + mock_fetch = self.enter_context( + mock.patch.object(gke_utils, "fetch_cluster_credentials", autospec=True) + ) + mock_prepare = self.enter_context( + mock.patch.object( + start_vscode_on_cpu_np, "_prepare_deployment_yaml", autospec=True + ) + ) + mock_prepare.return_value = "dummy-yaml" + mock_deploy = self.enter_context( + mock.patch.object( + start_vscode_on_cpu_np, "_deploy_vscode", autospec=True + ) + ) + mock_forward = self.enter_context( + mock.patch.object( + start_vscode_on_cpu_np, "_start_port_forwarding", autospec=True + ) + ) + mock_cleanup = self.enter_context( + mock.patch.object( + start_vscode_on_cpu_np, "_cleanup_gke_resources", autospec=True + ) + ) + + mock_choices = self.enter_context( + mock.patch.object(random, "choices", autospec=True) + ) + mock_choices.return_value = ["a", "b", "c", "d"] + + start_vscode_on_cpu_np.main(["script_name"]) + + expected_service_name = "my-vscode-testuser-abcd" + mock_prepare.assert_called_once_with(expected_service_name, 8080) + mock_choices.assert_called_once_with( + string.ascii_lowercase + string.digits, k=4 + ) + mock_deploy.assert_called_once_with(expected_service_name, "dummy-yaml") + mock_forward.assert_called_once_with(expected_service_name, 8080) + mock_cleanup.assert_called_once_with(expected_service_name, "my-ns") + mock_fetch.assert_called_once_with( + cluster_name="dummy-cluster", + project_id="dummy-project", + location="dummy-region", + ) + + @flagsaver.flagsaver(dry_run=False, name="my-vscode", namespace="my-ns") + @mock.patch.dict(os.environ, {"USER": "testuser"}) + def test_main_real_run_exception_still_cleans_up(self): + mock_fetch = self.enter_context( + mock.patch.object(gke_utils, "fetch_cluster_credentials", autospec=True) + ) + mock_prepare = self.enter_context( + mock.patch.object( + start_vscode_on_cpu_np, "_prepare_deployment_yaml", autospec=True + ) + ) + mock_prepare.return_value = "dummy-yaml" + mock_deploy = self.enter_context( + mock.patch.object( + start_vscode_on_cpu_np, "_deploy_vscode", autospec=True + ) + ) + mock_deploy.side_effect = RuntimeError("Deploy failed") + mock_forward = self.enter_context( + mock.patch.object( + start_vscode_on_cpu_np, "_start_port_forwarding", autospec=True + ) + ) + mock_cleanup = self.enter_context( + mock.patch.object( + start_vscode_on_cpu_np, "_cleanup_gke_resources", autospec=True + ) + ) + + mock_choices = self.enter_context( + mock.patch.object(random, "choices", autospec=True) + ) + mock_choices.return_value = ["a", "b", "c", "d"] + + with self.assertRaises(RuntimeError): + start_vscode_on_cpu_np.main(["script_name"]) + + expected_service_name = "my-vscode-testuser-abcd" + mock_prepare.assert_called_once() + mock_choices.assert_called_once_with( + string.ascii_lowercase + string.digits, k=4 + ) + mock_deploy.assert_called_once_with(expected_service_name, "dummy-yaml") + mock_forward.assert_not_called() + mock_cleanup.assert_called_once_with(expected_service_name, "my-ns") + mock_fetch.assert_called_once_with( + cluster_name="dummy-cluster", + project_id="dummy-project", + location="dummy-region", + ) + + +if __name__ == "__main__": + FLAGS = flags.FLAGS + FLAGS.cluster = "dummy-cluster" + FLAGS.project = "dummy-project" + FLAGS.region = "dummy-region" + absltest.main() diff --git a/pathwaysutils/test/experimental/shared_pathways_service/tpu_specs_test.py b/pathwaysutils/test/experimental/shared_pathways_service/tpu_specs_test.py new file mode 100644 index 0000000..46b0f2a --- /dev/null +++ b/pathwaysutils/test/experimental/shared_pathways_service/tpu_specs_test.py @@ -0,0 +1,77 @@ +"""Tests for tpu_specs.py.""" + +from absl.testing import absltest +from absl.testing import parameterized +from pathwaysutils.experimental.shared_pathways_service import tpu_specs + + +class TpuSpecsTest(parameterized.TestCase): + + @parameterized.named_parameters( + ("v5e", "v5e", 4, "tpu-v5-lite-podslice", "tpuv5e"), + ("v5p", "v5p", 4, "tpu-v5p-slice", "tpuv5"), + ("v6e", "v6e", 4, "tpu-v6e-slice", "tpuv6e"), + ("tpu7x", "tpu7x", 4, "tpu7x", "tpu7x"), + ) + def test_get_tpu_config( + self, tpu_type, expected_chips, expected_label, expected_prefix + ): + config = tpu_specs.get_tpu_config(tpu_type) + self.assertEqual(config.chips_per_vm, expected_chips) + self.assertEqual(config.accelerator_label, expected_label) + self.assertEqual(config.instance_prefix, expected_prefix) + + def test_get_tpu_config_invalid(self): + with self.assertRaises(ValueError): + tpu_specs.get_tpu_config("invalid_tpu") + + @parameterized.named_parameters( + ("v5e_single", "4x4", 4, 4), + ("v5e_large", "4x8", 4, 8), + ("v5e_multi", "8x8", 4, 16), + ("3d_topology", "2x2x2", 4, 2), + ) + def test_calculate_vms_per_slice(self, topology, chips_per_vm, expected_vms): + self.assertEqual( + tpu_specs.calculate_vms_per_slice(topology, chips_per_vm), expected_vms + ) + + def test_calculate_vms_per_slice_invalid_format(self): + with self.assertRaisesRegex(ValueError, "Invalid topology format"): + tpu_specs.calculate_vms_per_slice("invalid", 4) + + def test_calculate_vms_per_slice_indivisible(self): + with self.assertRaisesRegex(ValueError, "is not divisible by chips_per_vm"): + tpu_specs.calculate_vms_per_slice("2x1", 4) + + @parameterized.named_parameters( + ("v5e", "tpuv5e:4x8", ("v5e", "4x8")), + ("v5p", "tpuv5:8x8", ("v5p", "8x8")), + ("v6e", "tpuv6e:4x4", ("v6e", "4x4")), + ("tpu7x", "tpu7x:4x4", ("tpu7x", "4x4")), + ) + def test_parse_tpu_type_string(self, tpu_type_str, expected): + self.assertEqual(tpu_specs.parse_tpu_type_string(tpu_type_str), expected) + + def test_parse_tpu_type_string_invalid_format(self): + with self.assertRaisesRegex(ValueError, "Invalid tpu_type string"): + tpu_specs.parse_tpu_type_string("tpuv5e-4x8") + + def test_parse_tpu_type_string_invalid_prefix(self): + with self.assertRaisesRegex(ValueError, "Unsupported TPU prefix"): + tpu_specs.parse_tpu_type_string("invalid_prefix:4x8") + + def test_get_tpu_params(self): + params = tpu_specs.get_tpu_params("v5e", "4x8") + expected = { + "ACCELERATOR_LABEL": "tpu-v5-lite-podslice", + "TOPOLOGY": "4x8", + "VMS_PER_SLICE": "8", + "CHIPS_PER_VM": "4", + "INSTANCE_TYPE": "tpuv5e:4x8", + } + self.assertEqual(params, expected) + + +if __name__ == "__main__": + absltest.main() diff --git a/pathwaysutils/test/experimental/shared_pathways_service/validators_test.py b/pathwaysutils/test/experimental/shared_pathways_service/validators_test.py new file mode 100644 index 0000000..56897e1 --- /dev/null +++ b/pathwaysutils/test/experimental/shared_pathways_service/validators_test.py @@ -0,0 +1,370 @@ +"""Tests for validation functions for the Shared Pathways service.""" + +from unittest import mock + +from absl import flags +from absl.testing import absltest +from absl.testing import parameterized + +from pathwaysutils.experimental.shared_pathways_service import validators + + +class ValidatorsTest(parameterized.TestCase): + + @parameterized.named_parameters( + dict(testcase_name="simple_service", service="test-service:1234"), + dict( + testcase_name="complex_hostname", + service="pathways-cluster-pathways-head-0-0.pathways-cluster:8080", + ), + ) + def test_validate_pathways_service_success(self, service): + """Tests that valid pathways service strings pass validation.""" + validators.validate_pathways_service(service) + + @parameterized.named_parameters( + dict( + testcase_name="missing_port", + service="test-service", + expected_regex=( + "pathways_service=test-service is not in the expected format of" + ), + ), + dict( + testcase_name="empty_port", + service="test-service:", + expected_regex=( + "pathways_service=test-service: contains an empty string for the" + " service port." + ), + ), + dict( + testcase_name="empty_service_name", + service=":1234", + expected_regex=( + "pathways_service=:1234 contains an empty string for the service" + " name." + ), + ), + dict( + testcase_name="non_numeric_port", + service="test-service:port", + expected_regex=( + "pathways_service=test-service:port contains a non-numeric" + " service port." + ), + ), + dict( + testcase_name="empty_string", + service="", + expected_regex="No Pathways service found.", + ), + dict( + testcase_name="too_many_parts", + service="test-service:1234:5678", + expected_regex=("pathways_service=test-service:1234:5678 is not in" + " the expected format of"), + ), + ) + def test_validate_pathways_service_failure(self, service, expected_regex): + """Tests that invalid pathways service strings raise a ValueError.""" + with self.assertRaisesRegex(ValueError, expected_regex): + validators.validate_pathways_service(service) + + @parameterized.named_parameters( + dict(testcase_name="tpuv6e_2x2", instance_dict={"tpuv6e:2x2": 4}), + dict(testcase_name="tpuv6e_2x4", instance_dict={"tpuv6e:2x4": 1}), + dict(testcase_name="tpuv6e_2x2x2", instance_dict={"tpuv6e:2x2x2": 1}), + dict(testcase_name="tpuv6e_4x4", instance_dict={"tpuv6e:4x4": 1}), + dict(testcase_name="tpuv5e_2x2", instance_dict={"tpuv5e:2x2": 4}), + dict(testcase_name="tpuv5e_4x4", instance_dict={"tpuv5e:4x4": 1}), + dict(testcase_name="tpuv5_2x2x2", instance_dict={"tpuv5:2x2x2": 1}), + dict(testcase_name="tpuv5_2x2x4", instance_dict={"tpuv5:2x2x4": 1}), + ) + def test_validate_tpu_instances_success(self, instance_dict): + """Tests that valid instance lists pass validation.""" + validators.validate_tpu_instances(instance_dict) + + @parameterized.named_parameters( + dict( + testcase_name="tpuv6e_8", + instance_dict={"tpuv6e-8": 1}, + expected_regex="Unrecognized instance format: tpuv6e-8.", + ), + dict( + testcase_name="tpuv6e_4", + instance_dict={"tpuv6e-4": 1}, + expected_regex="Unrecognized instance format: tpuv6e-4.", + ), + dict( + testcase_name="tpuv6e_16", + instance_dict={"tpuv6e-16": 1}, + expected_regex="Unrecognized instance format: tpuv6e-16.", + ), + dict( + testcase_name="invalid_format", + instance_dict={"invalid-format": 1}, + expected_regex="Unrecognized instance format: invalid-format.", + ), + dict( + testcase_name="ct5lp-hightpu-16t_4x4", + instance_dict={"ct5lp-hightpu-16t:4x4": 1}, + expected_regex="Unrecognized instance format: ct5lp-hightpu-16t:4x4.", + ), + dict( + testcase_name="ct5lp_hightpu_16t_2x2", + instance_dict={"ct5lp-hightpu-16t:2x2": 1}, + expected_regex="Unrecognized instance format: ct5lp-hightpu-16t:2x2.", + ), + dict( + testcase_name="tpuv5p_2x4x4", + instance_dict={"tpuv5p:2x4x4": 1}, + expected_regex="Unrecognized instance format: tpuv5p:2x4x4.", + ), + dict( + testcase_name="ct5lp_8", + instance_dict={"ct5lp-8": 1}, + expected_regex="Unrecognized instance format: ct5lp-8.", + ), + dict( + testcase_name="ct5l_8", + instance_dict={"ct5l-8": 1}, + expected_regex="Unrecognized instance format: ct5l-8.", + ), + dict( + testcase_name="ct5p_2x2", + instance_dict={"ct5p:2x2": 1}, + expected_regex="Unrecognized instance format: ct5p:2x2.", + ), + dict( + testcase_name="ct5p_4", + instance_dict={"ct5p-4": 1}, + expected_regex="Unrecognized instance format: ct5p-4.", + ), + dict( + testcase_name="ct7e_8", + instance_dict={"ct7e-8": 1}, + expected_regex="Unrecognized instance format: ct7e-8.", + ), + dict( + testcase_name="tpuv7e-8", + instance_dict={"tpuv7e-8": 1}, + expected_regex="Unrecognized instance format: tpuv7e-8.", + ), + dict( + testcase_name="ct5lp_1x1x", + instance_dict={"ct5lp:1x1x": 1}, + expected_regex="Unrecognized instance format: ct5lp:1x1x.", + ), + dict( + testcase_name="ct5lp_axb", + instance_dict={"ct5lp:axb": 1}, + expected_regex="Unrecognized instance format: ct5lp:axb.", + ), + dict( + testcase_name="foo_bar", + instance_dict={"foo-bar": 1}, + expected_regex="Unrecognized instance format: foo-bar.", + ), + dict( + testcase_name="empty_dict", + instance_dict={}, + expected_regex="No instances found.", + ), + dict( + testcase_name="empty_key", + instance_dict={"": 2}, + expected_regex=( + r"expected_tpu_instances=\{\'\'\: 2\} contains an empty string" + r" for an instance name." + ), + ), + dict( + testcase_name="tpuv6e_1_dim", + instance_dict={"tpuv6e:1": 1}, + expected_regex="Unrecognized instance format: tpuv6e:1.", + ), + dict( + testcase_name="tpuv6e_4_dims", + instance_dict={"tpuv6e:1x2x3x4": 1}, + expected_regex="Unrecognized instance format: tpuv6e:1x2x3x4.", + ), + dict( + testcase_name="multiple_keys", + instance_dict={"tpuv6e:2x2": 2, "v5e-16": 1}, + expected_regex="Only one machine type is supported at this time.", + ), + ) + def test_validate_tpu_instances_failure( + self, instance_dict, expected_regex + ): + """Tests that invalid TPU instance dictionaries raise a ValueError.""" + with self.assertRaisesRegex(ValueError, expected_regex): + validators.validate_tpu_instances(instance_dict) + + @parameterized.named_parameters( + dict( + testcase_name="with_tag", + image="gcr.io/project/image:tag", + ), + dict( + testcase_name="with_digest", + image="gcr.io/project/image@sha256:12345", + ), + dict( + testcase_name="with_tag_and_digest", + image="gcr.io/project/image:tag@sha256:12345", + ), + dict( + testcase_name="path_with_hyphen", + image="gcr.io/project-id/image-name:tag", + ), + dict( + testcase_name="host_with_region", + image="us-docker.pkg.dev/project/repo/image:tag", + ), + ) + def test_validate_proxy_server_image_success(self, image): + """Tests that valid proxy server image strings pass validation.""" + validators.validate_proxy_server_image(image) + + @parameterized.named_parameters( + dict( + testcase_name="empty_string", + image="", + expected_regex="Proxy server image cannot be empty.", + ), + dict( + testcase_name="whitespace_only", + image=" ", + expected_regex="Proxy server image cannot be empty.", + ), + dict( + testcase_name="no_slash", + image="image:tag", + expected_regex="Proxy server image 'image:tag' must contain '/'.", + ), + dict( + testcase_name="no_tag_or_digest", + image="gcr.io/project/image", + expected_regex=( + "Proxy server image 'gcr.io/project/image' must contain a tag" + " with ':' or a digest with '@'." + ), + ), + ) + def test_validate_proxy_server_image_failure(self, image, expected_regex): + """Tests that invalid proxy server image strings raise a ValueError.""" + with self.assertRaisesRegex(ValueError, expected_regex): + validators.validate_proxy_server_image(image) + + @parameterized.named_parameters( + dict(testcase_name="empty_list", options=[]), + dict(testcase_name="none", options=None), + dict(testcase_name="valid_options", options=["key1:val1", "key2:val2"]), + dict( + testcase_name="with_xla_flags", + options=['xla_flags:"--flag1 --flag2"'], + ), + ) + def test_validate_proxy_options_success(self, options): + validators.validate_proxy_options(options) + + @parameterized.named_parameters( + dict( + testcase_name="no_colon", + options=["invalid_option"], + expected_regex='--proxy_options must be in the format "key:value".', + ), + dict( + testcase_name="empty_key", + options=[":value"], + expected_regex='--proxy_options must be in the format "key:value".', + ), + dict( + testcase_name="empty_value", + options=["key:"], + expected_regex='--proxy_options must be in the format "key:value".', + ), + ) + def test_validate_proxy_options_failure(self, options, expected_regex): + with self.assertRaisesRegex(flags.ValidationError, expected_regex): + validators.validate_proxy_options(options) + + @parameterized.named_parameters( + dict(testcase_name="empty_list", xla_flags=[]), + dict(testcase_name="none", xla_flags=None), + dict( + testcase_name="valid_flags", + xla_flags=[ + "--xla_tpu_scoped_vmem_limit_kib=98304", + "--xla_tpu_use_minor_sharding_for_major_trivial_input=true", + ], + ), + ) + def test_validate_xla_flags_success(self, xla_flags): + validators.validate_xla_flags(xla_flags) + + @parameterized.named_parameters( + dict( + testcase_name="invalid_prefix", + xla_flags=["--not_xla_flag"], + expected_regex="XLA flag '--not_xla_flag' must start with '--xla_'.", + ), + ) + def test_validate_xla_flags_failure(self, xla_flags, expected_regex): + with self.assertRaisesRegex(flags.ValidationError, expected_regex): + validators.validate_xla_flags(xla_flags) + + def test_validate_sidecar_image_versions_success(self): + mock_sys_info = mock.Mock() + mock_sys_info.major = 3 + mock_sys_info.minor = 12 + mock_sys_info.micro = 8 + with mock.patch("sys.version_info", mock_sys_info), mock.patch( + "jax.__version__", "0.10.0" + ): + # Python 3.12 (matches 3.12.8), JAX 0.10.0 (matches 0.10.0) + validators.validate_sidecar_image_versions( + "us-docker.pkg.dev/.../sidecar:20260423-python_3.12-jax_0.10.0" + ) + + # Python omitted, JAX 0.10 (matches 0.10.0) + validators.validate_sidecar_image_versions( + "us-docker.pkg.dev/.../sidecar:20260423-jax_0.10" + ) + + # No version info + validators.validate_sidecar_image_versions( + "us-docker.pkg.dev/.../sidecar:latest" + ) + + def test_validate_sidecar_image_versions_python_mismatch(self): + mock_sys_info = mock.Mock() + mock_sys_info.major = 3 + mock_sys_info.minor = 11 + mock_sys_info.micro = 5 + with mock.patch("sys.version_info", mock_sys_info), mock.patch( + "jax.__version__", "0.10.0" + ): + with self.assertRaisesRegex(ValueError, "Python version mismatch"): + validators.validate_sidecar_image_versions( + "us-docker.pkg.dev/.../sidecar:20260423-python_3.12-jax_0.10.0" + ) + + def test_validate_sidecar_image_versions_jax_mismatch(self): + mock_sys_info = mock.Mock() + mock_sys_info.major = 3 + mock_sys_info.minor = 12 + mock_sys_info.micro = 8 + with mock.patch("sys.version_info", mock_sys_info), mock.patch( + "jax.__version__", "0.9.0" + ): + with self.assertRaisesRegex(ValueError, "JAX version mismatch"): + validators.validate_sidecar_image_versions( + "us-docker.pkg.dev/.../sidecar:20260423-python_3.12-jax_0.10.0" + ) + + +if __name__ == "__main__": + absltest.main()