From 885c15a5f315e399941acd464b71d94b8f0f05c9 Mon Sep 17 00:00:00 2001 From: Pathways-on-Cloud Team Date: Sun, 20 Sep 2026 18:32:45 -0700 Subject: [PATCH] Add `PathwaysJobSet.with_custom_libtpu` to override libtpu at runtime on GKE. Testing a custom `libtpu.so` with Pathways on GKE currently requires rebuilding the server and proxy images, which bake libtpu at `/lib/libtpu.so`. `with_custom_libtpu(uri)` adds an init container that downloads a `libtpu.so` (or a `.whl`/`.zip` containing one) from `gs://` or `http(s)://` into an `emptyDir` volume and `subPath`-mounts it over `/lib/libtpu.so` in the RM, proxy and worker containers (also setting `TPU_LIBRARY_PATH`), so the servers `dlopen` the custom library without any image rebuild. Calling it again replaces the override. PiperOrigin-RevId: 984946941 --- pathwaysutils/experimental/gke/jobset.py | 125 +++++++- .../test/experimental/gke/jobset_test.py | 276 ++++++++++++++---- 2 files changed, 339 insertions(+), 62 deletions(-) diff --git a/pathwaysutils/experimental/gke/jobset.py b/pathwaysutils/experimental/gke/jobset.py index 28bee85..9772ccf 100644 --- a/pathwaysutils/experimental/gke/jobset.py +++ b/pathwaysutils/experimental/gke/jobset.py @@ -18,7 +18,8 @@ import logging import math import time -from typing import TYPE_CHECKING, Any, Mapping, Sequence +from typing import Any, Mapping, Sequence, TYPE_CHECKING + import yaml try: @@ -50,6 +51,33 @@ PATHWAYS_RM_PORT = 29001 PATHWAYS_WORKER_PORT = 29005 +CUSTOM_LIBTPU_VOLUME_NAME = "custom-libtpu" +CUSTOM_LIBTPU_INIT_CONTAINER_NAME = "fetch-custom-libtpu" +CUSTOM_LIBTPU_INIT_IMAGE = "gcr.io/google.com/cloudsdktool/cloud-sdk:slim" +CUSTOM_LIBTPU_SO_PATH = "/lib/libtpu.so" +_CUSTOM_LIBTPU_CONTAINERS = ("pathways-rm", "pathways-proxy", "pathways-worker") +_CUSTOM_LIBTPU_STAGING_DIR = "/tmp/libtpu" +# Runs in the init container: downloads argv[1] and leaves `libtpu.so` at +# argv[2], extracting it first if argv[1] is a wheel/zip. +_FETCH_CUSTOM_LIBTPU_SCRIPT = """ +import os, shutil, subprocess, sys, urllib.request, zipfile +uri, so_path = sys.argv[1], sys.argv[2] +artifact = so_path + ".download" +if uri.startswith("gs://"): + subprocess.run(["gcloud", "storage", "cp", uri, artifact], check=True) +else: + urllib.request.urlretrieve(uri, artifact) +if uri.endswith((".whl", ".zip")): + with zipfile.ZipFile(artifact) as archive: + member = next(n for n in archive.namelist() if n.endswith("libtpu.so")) + with archive.open(member) as src, open(so_path, "wb") as dst: + shutil.copyfileobj(src, dst) + os.remove(artifact) +else: + os.replace(artifact, so_path) +os.chmod(so_path, 0o755) +""" + MACHINE_TYPE_TO_TPU_VERSION_MAP = { "tpu7x-standard-4t": "tpu7x", "tpu7x": "tpu7x", @@ -86,7 +114,7 @@ def _format_image(image: str, default_tag: str) -> str: if "@" in image: return image last_slash = image.rfind("/") - if ":" in image[last_slash + 1:]: + if ":" in image[last_slash + 1 :]: return image return f"{image}:{default_tag}" @@ -135,14 +163,16 @@ def __init__( topology: TPU topology (e.g., "2x2"). num_slices: Number of slices. max_restarts: Maximum number of restarts for the JobSet. - max_slice_restarts: Maximum number of slice restarts (defaults to 1_000_000 in headless and SPS mode). + max_slice_restarts: Maximum number of slice restarts (defaults to + 1_000_000 in headless and SPS mode). termination_grace_period_seconds: Optional termination grace period. pathways_version: Version tag for Pathways images. jobset_api_version: API version of JobSet. elastic_slices: Number of elastic slices. labels: Optional labels for the JobSet. annotations: Optional annotations for the JobSet. - shared_pathways_service: Whether to run only RM for Shared Pathways Service. + shared_pathways_service: Whether to run only RM for Shared Pathways + Service. pathways_rm_and_worker_image: Base Docker image for Resource Manager and Worker containers. pathways_proxy_image: Base Docker image for Proxy container. @@ -251,7 +281,8 @@ def _build_head_job_template( instance_type: TPU instance type (e.g., "tpuv5:2x2"). image_tag: Version tag for Pathways images. elastic_slices: Number of elastic slices. - shared_pathways_service: Whether to run only RM for Shared Pathways Service. + shared_pathways_service: Whether to run only RM for Shared Pathways + Service. pathways_rm_and_worker_image: Base Docker image for Resource Manager. pathways_proxy_image: Base Docker image for Proxy container. @@ -689,6 +720,90 @@ def add_gcsfuse( return self + def with_custom_libtpu( + self, uri: str, init_image: str = CUSTOM_LIBTPU_INIT_IMAGE + ) -> "PathwaysJobSet": + """Overrides `/lib/libtpu.so` in the Pathways containers at runtime. + + An init container downloads `uri` into an `emptyDir` volume that is then + mounted over `/lib/libtpu.so` (and pointed to by `TPU_LIBRARY_PATH`) in the + RM, proxy and worker containers, so a custom libtpu can be used without + rebuilding the Pathways images. Calling this again replaces the override. + + Only effective with images that load libtpu dynamically from + `TPU_LIBRARY_PATH` (the OSS-built Pathways images); images that link the + TPU runtime statically ignore it. + + Args: + uri: `gs://` or `http(s)://` URI of a `libtpu.so`, or of a `.whl`/`.zip` + containing one. + init_image: Image used to download the artifact; must provide `python3` + and, for `gs://` URIs, `gcloud`. + + Returns: + This `PathwaysJobSet`, for chaining. + """ + if not uri.startswith(("gs://", "http://", "https://")): + raise ValueError(f"Unsupported custom libtpu URI: {uri!r}") + + for job_template in (self._head_job_template, self._worker_job_template): + pod_spec = job_template.spec.template.spec + self._add_volume_to_pod_spec( + pod_spec, + client.V1Volume( + name=CUSTOM_LIBTPU_VOLUME_NAME, + empty_dir=client.V1EmptyDirVolumeSource(), + ), + ) + others = [ + c + for c in pod_spec.init_containers or [] + if c.name != CUSTOM_LIBTPU_INIT_CONTAINER_NAME + ] + fetch_container = client.V1Container( + name=CUSTOM_LIBTPU_INIT_CONTAINER_NAME, + image=init_image, + command=[ + "python3", + "-c", + _FETCH_CUSTOM_LIBTPU_SCRIPT, + uri, + f"{_CUSTOM_LIBTPU_STAGING_DIR}/libtpu.so", + ], + volume_mounts=[ + client.V1VolumeMount( + name=CUSTOM_LIBTPU_VOLUME_NAME, + mount_path=_CUSTOM_LIBTPU_STAGING_DIR, + ) + ], + ) + pod_spec.init_containers = [fetch_container] + others + + # The RM and proxy may run as sidecar init containers (e.g. imported YAML). + for container in (pod_spec.containers or []) + others: + if container.name not in _CUSTOM_LIBTPU_CONTAINERS: + continue + container.volume_mounts = [ + m + for m in container.volume_mounts or [] + if m.mount_path != CUSTOM_LIBTPU_SO_PATH + ] + [ + client.V1VolumeMount( + name=CUSTOM_LIBTPU_VOLUME_NAME, + mount_path=CUSTOM_LIBTPU_SO_PATH, + sub_path="libtpu.so", + ) + ] + container.env = [ + e for e in container.env or [] if e.name != "TPU_LIBRARY_PATH" + ] + [ + client.V1EnvVar( + name="TPU_LIBRARY_PATH", value=CUSTOM_LIBTPU_SO_PATH + ) + ] + + return self + def _compile_config(self) -> dict[str, Any]: """Compiles the JobSet configuration into a dictionary.""" with client.ApiClient() as api_client: diff --git a/pathwaysutils/test/experimental/gke/jobset_test.py b/pathwaysutils/test/experimental/gke/jobset_test.py index 41fd382..154d61d 100644 --- a/pathwaysutils/test/experimental/gke/jobset_test.py +++ b/pathwaysutils/test/experimental/gke/jobset_test.py @@ -106,7 +106,6 @@ def get_all_containers_by_name( return matches - class PathwaysJobSetTest(parameterized.TestCase): def _create_jobset( @@ -255,7 +254,9 @@ def test_worker_job_small_topology_completions(self): @parameterized.parameters(True, False) def test_add_gcsfuse_read_only(self, read_only): pw_jobset = self._create_jobset(topology="2x2", num_slices=1) - bucket_hash = int(hashlib.md5("my-bucket".encode()).hexdigest(), 16) % (10**8) + bucket_hash = int(hashlib.md5("my-bucket".encode()).hexdigest(), 16) % ( + 10**8 + ) expected_vol_name = f"gcsfuse-{bucket_hash}" pw_jobset.add_gcsfuse( @@ -269,18 +270,26 @@ def test_add_gcsfuse_read_only(self, read_only): rm_container = helper.containers["pathways-head"]["pathways-rm"] mounts = {m["name"]: m for m in rm_container.get("volumeMounts", [])} self.assertIn(expected_vol_name, mounts) - self.assertEqual(mounts[expected_vol_name].get("readOnly", False), read_only) + self.assertEqual( + mounts[expected_vol_name].get("readOnly", False), read_only + ) @parameterized.named_parameters( ("all", "all", ["pathways-rm", "pathways-proxy", "pathways-worker"]), ("worker", "pathways-worker", ["pathways-worker"]), - ("explicit", ["pathways-worker", "pathways-rm"], ["pathways-worker", "pathways-rm"]), + ( + "explicit", + ["pathways-worker", "pathways-rm"], + ["pathways-worker", "pathways-rm"], + ), ) def test_add_gcsfuse_container_filtering( self, containers_param, expected_containers ): pw_jobset = self._create_jobset(topology="2x2", num_slices=1) - bucket_hash = int(hashlib.md5("my-bucket".encode()).hexdigest(), 16) % (10**8) + bucket_hash = int(hashlib.md5("my-bucket".encode()).hexdigest(), 16) % ( + 10**8 + ) expected_vol_name = f"gcsfuse-{bucket_hash}" pw_jobset.add_gcsfuse( @@ -290,20 +299,33 @@ def test_add_gcsfuse_container_filtering( ) helper = JobSetManifestHelper(pw_jobset.to_dict()) - all_possible_containers = ["pathways-rm", "pathways-proxy", "pathways-worker"] + all_possible_containers = [ + "pathways-rm", + "pathways-proxy", + "pathways-worker", + ] for c_name in all_possible_containers: matches = helper.get_all_containers_by_name(c_name) self.assertNotEmpty(matches) for job_name, container in matches: - has_mount = any(m["mountPath"] == "/gcs/data" for m in container.get("volumeMounts", [])) + has_mount = any( + m["mountPath"] == "/gcs/data" + for m in container.get("volumeMounts", []) + ) if c_name in expected_containers: - self.assertTrue(has_mount, f"Expected {c_name} in {job_name} to have mount") + self.assertTrue( + has_mount, f"Expected {c_name} in {job_name} to have mount" + ) else: - self.assertFalse(has_mount, f"Expected {c_name} in {job_name} NOT to have mount") + self.assertFalse( + has_mount, f"Expected {c_name} in {job_name} NOT to have mount" + ) def test_add_gcsfuse_volumes_and_annotations(self): pw_jobset = self._create_jobset(topology="2x2", num_slices=1) - bucket_hash = int(hashlib.md5("my-bucket".encode()).hexdigest(), 16) % (10**8) + bucket_hash = int(hashlib.md5("my-bucket".encode()).hexdigest(), 16) % ( + 10**8 + ) expected_vol_name = f"gcsfuse-{bucket_hash}" pw_jobset.add_gcsfuse( @@ -318,17 +340,37 @@ def test_add_gcsfuse_volumes_and_annotations(self): vol = helper.volumes["pathways-worker"][expected_vol_name] self.assertEqual(vol["csi"]["driver"], "gcsfuse.csi.storage.gke.io") self.assertEqual(vol["csi"]["volumeAttributes"]["bucketName"], "my-bucket") - self.assertEqual(helper.job_metadatas["pathways-worker"].get("annotations", {}).get("gke-gcsfuse/volumes"), "true") - self.assertEqual(helper.pod_metadatas["pathways-worker"].get("annotations", {}).get("gke-gcsfuse/volumes"), "true") + self.assertEqual( + helper.job_metadatas["pathways-worker"] + .get("annotations", {}) + .get("gke-gcsfuse/volumes"), + "true", + ) + self.assertEqual( + helper.pod_metadatas["pathways-worker"] + .get("annotations", {}) + .get("gke-gcsfuse/volumes"), + "true", + ) # Head job should NOT have volume or annotations self.assertNotIn(expected_vol_name, helper.volumes["pathways-head"]) - self.assertNotEqual(helper.job_metadatas["pathways-head"].get("annotations", {}).get("gke-gcsfuse/volumes"), "true") - self.assertNotEqual(helper.pod_metadatas["pathways-head"].get("annotations", {}).get("gke-gcsfuse/volumes"), "true") + self.assertNotEqual( + helper.job_metadatas["pathways-head"] + .get("annotations", {}) + .get("gke-gcsfuse/volumes"), + "true", + ) + self.assertNotEqual( + helper.pod_metadatas["pathways-head"] + .get("annotations", {}) + .get("gke-gcsfuse/volumes"), + "true", + ) def test_add_gcsfuse_handles_none_metadata(self): pw_jobset = self._create_jobset(topology="2x2", num_slices=1) - + # Force metadata to be None to simulate imported templates or raw specs without metadata pw_jobset._head_job_template.metadata = None pw_jobset._head_job_template.spec.template.metadata = None @@ -342,23 +384,43 @@ def test_add_gcsfuse_handles_none_metadata(self): bucket="my-bucket", ) helper = JobSetManifestHelper(pw_jobset.to_dict()) - - self.assertEqual(helper.job_metadatas["pathways-head"].get("annotations", {}).get("gke-gcsfuse/volumes"), "true") - self.assertEqual(helper.pod_metadatas["pathways-head"].get("annotations", {}).get("gke-gcsfuse/volumes"), "true") - self.assertEqual(helper.job_metadatas["pathways-worker"].get("annotations", {}).get("gke-gcsfuse/volumes"), "true") - self.assertEqual(helper.pod_metadatas["pathways-worker"].get("annotations", {}).get("gke-gcsfuse/volumes"), "true") + + self.assertEqual( + helper.job_metadatas["pathways-head"] + .get("annotations", {}) + .get("gke-gcsfuse/volumes"), + "true", + ) + self.assertEqual( + helper.pod_metadatas["pathways-head"] + .get("annotations", {}) + .get("gke-gcsfuse/volumes"), + "true", + ) + self.assertEqual( + helper.job_metadatas["pathways-worker"] + .get("annotations", {}) + .get("gke-gcsfuse/volumes"), + "true", + ) + self.assertEqual( + helper.pod_metadatas["pathways-worker"] + .get("annotations", {}) + .get("gke-gcsfuse/volumes"), + "true", + ) def test_add_gcsfuse_preserves_existing_metadata(self): pw_jobset = self._create_jobset(topology="2x2", num_slices=1) - + # Pre-populate metadata, annotations, and labels pw_jobset._head_job_template.metadata = client.V1ObjectMeta( labels={"existing-job-label": "value"}, - annotations={"existing-job-anno": "value"} + annotations={"existing-job-anno": "value"}, ) pw_jobset._head_job_template.spec.template.metadata = client.V1ObjectMeta( labels={"existing-pod-label": "value"}, - annotations={"existing-pod-anno": "value"} + annotations={"existing-pod-anno": "value"}, ) pw_jobset.add_gcsfuse( @@ -370,14 +432,26 @@ def test_add_gcsfuse_preserves_existing_metadata(self): # Verify existing annotations and labels are preserved, and new annotation is added job_meta = helper.job_metadatas["pathways-head"] - self.assertEqual(job_meta.get("labels", {}).get("existing-job-label"), "value") - self.assertEqual(job_meta.get("annotations", {}).get("existing-job-anno"), "value") - self.assertEqual(job_meta.get("annotations", {}).get("gke-gcsfuse/volumes"), "true") + self.assertEqual( + job_meta.get("labels", {}).get("existing-job-label"), "value" + ) + self.assertEqual( + job_meta.get("annotations", {}).get("existing-job-anno"), "value" + ) + self.assertEqual( + job_meta.get("annotations", {}).get("gke-gcsfuse/volumes"), "true" + ) pod_meta = helper.pod_metadatas["pathways-head"] - self.assertEqual(pod_meta.get("labels", {}).get("existing-pod-label"), "value") - self.assertEqual(pod_meta.get("annotations", {}).get("existing-pod-anno"), "value") - self.assertEqual(pod_meta.get("annotations", {}).get("gke-gcsfuse/volumes"), "true") + self.assertEqual( + pod_meta.get("labels", {}).get("existing-pod-label"), "value" + ) + self.assertEqual( + pod_meta.get("annotations", {}).get("existing-pod-anno"), "value" + ) + self.assertEqual( + pod_meta.get("annotations", {}).get("gke-gcsfuse/volumes"), "true" + ) def test_add_gcsfuse_preserves_existing_volumes(self): pw_jobset = self._create_jobset(topology="2x2", num_slices=1) @@ -403,28 +477,37 @@ def test_add_gcsfuse_preserves_existing_volumes(self): def test_add_colocated_python_handles_none_volumes(self): pw_jobset = self._create_jobset(topology="2x2", num_slices=1) - + # Force volumes to be None pw_jobset._worker_job_template.spec.template.spec.volumes = None # Should not crash and should correctly add volume - pw_jobset.add_colocated_python(image="gcr.io/my-project/colocated-python:custom") + pw_jobset.add_colocated_python( + image="gcr.io/my-project/colocated-python:custom" + ) helper = JobSetManifestHelper(pw_jobset.to_dict()) - + self.assertIn("shared-memory", helper.volumes["pathways-worker"]) def test_add_colocated_python_sidecar(self): pw_jobset = self._create_jobset(topology="2x2", num_slices=1) - pw_jobset.add_colocated_python(image="gcr.io/my-project/colocated-python:custom") + pw_jobset.add_colocated_python( + image="gcr.io/my-project/colocated-python:custom" + ) helper = JobSetManifestHelper(pw_jobset.to_dict()) - self.assertIn("colocated-python-sidecar", helper.containers["pathways-worker"]) + self.assertIn( + "colocated-python-sidecar", helper.containers["pathways-worker"] + ) sidecar = helper.containers["pathways-worker"]["colocated-python-sidecar"] - self.assertEqual(sidecar["image"], "gcr.io/my-project/colocated-python:custom") + self.assertEqual( + sidecar["image"], "gcr.io/my-project/colocated-python:custom" + ) self.assertTrue( any( - m["name"] == "shared-memory" and m["mountPath"] == "/tmp/shared-memory" + m["name"] == "shared-memory" + and m["mountPath"] == "/tmp/shared-memory" for m in sidecar["volumeMounts"] ) ) @@ -438,23 +521,33 @@ def test_add_colocated_python_sidecar(self): def test_add_colocated_python_preserves_init_containers(self): pw_jobset = self._create_jobset(topology="2x2", num_slices=1) - + # Pre-populate init container on worker pod worker_spec = pw_jobset._worker_job_template.spec.template.spec - existing_init = client.V1Container(name="existing-init-container", image="ubuntu:latest") + existing_init = client.V1Container( + name="existing-init-container", image="ubuntu:latest" + ) worker_spec.init_containers = [existing_init] - pw_jobset.add_colocated_python(image="gcr.io/my-project/colocated-python:custom") + pw_jobset.add_colocated_python( + image="gcr.io/my-project/colocated-python:custom" + ) helper = JobSetManifestHelper(pw_jobset.to_dict()) # Verify both exist - self.assertIn("existing-init-container", helper.init_containers["pathways-worker"]) - self.assertIn("colocated-python-sidecar", helper.containers["pathways-worker"]) + self.assertIn( + "existing-init-container", helper.init_containers["pathways-worker"] + ) + self.assertIn( + "colocated-python-sidecar", helper.containers["pathways-worker"] + ) def test_add_colocated_python_volume_default(self): pw_jobset = self._create_jobset(topology="2x2", num_slices=1) - pw_jobset.add_colocated_python(image="gcr.io/my-project/colocated-python:custom") + pw_jobset.add_colocated_python( + image="gcr.io/my-project/colocated-python:custom" + ) helper = JobSetManifestHelper(pw_jobset.to_dict()) self.assertIn("shared-memory", helper.volumes["pathways-worker"]) @@ -464,14 +557,17 @@ def test_add_colocated_python_volume_default(self): def test_add_colocated_python_worker_mount(self): pw_jobset = self._create_jobset(topology="2x2", num_slices=1) - pw_jobset.add_colocated_python(image="gcr.io/my-project/colocated-python:custom") + pw_jobset.add_colocated_python( + image="gcr.io/my-project/colocated-python:custom" + ) helper = JobSetManifestHelper(pw_jobset.to_dict()) self.assertIn("pathways-worker", helper.containers["pathways-worker"]) worker_container = helper.containers["pathways-worker"]["pathways-worker"] self.assertTrue( any( - m["name"] == "shared-memory" and m["mountPath"] == "/tmp/shared-memory" + m["name"] == "shared-memory" + and m["mountPath"] == "/tmp/shared-memory" for m in worker_container["volumeMounts"] ) ) @@ -493,7 +589,9 @@ def test_add_colocated_python_custom_shm(self): ) helper = JobSetManifestHelper(pw_jobset.to_dict()) - self.assertIn("colocated-python-sidecar", helper.containers["pathways-worker"]) + self.assertIn( + "colocated-python-sidecar", helper.containers["pathways-worker"] + ) sidecar = helper.containers["pathways-worker"]["colocated-python-sidecar"] self.assertTrue( any( @@ -540,7 +638,8 @@ def test_apply_create( ): """Tests deploying a JobSet to GKE via Kubernetes CustomObjectsApi. - Modeled after official GKE JobSet unit test patterns (e.g. gke_callbacks_test.py). + Modeled after official GKE JobSet unit test patterns (e.g. + gke_callbacks_test.py). """ mock_api = mock_custom_objects_api.return_value # Mock GET to return 404 (not exists). @@ -636,7 +735,9 @@ def test_export_import_roundtrip(self): topology="2x2", num_slices=1, ) - pw_jobset.add_colocated_python(image="gcr.io/my-project/colocated-python:custom") + pw_jobset.add_colocated_python( + image="gcr.io/my-project/colocated-python:custom" + ) pw_jobset.add_gcsfuse( containers="pathways-worker", mount_path="/tmp/gcs", bucket="my-bucket" ) @@ -679,15 +780,17 @@ def test_configurable_images(self): c.name: c.image for c in pw_jobset.head_job_template.spec.template.spec.containers } - self.assertEqual(head_containers["pathways-rm"], "gcr.io/custom/server:v2.0") + self.assertEqual( + head_containers["pathways-rm"], "gcr.io/custom/server:v2.0" + ) self.assertEqual( head_containers["pathways-proxy"], "gcr.io/custom/proxy:tagged" ) # Verify worker job container. - worker_image = ( - pw_jobset.worker_job_template.spec.template.spec.containers[0].image - ) + worker_image = pw_jobset.worker_job_template.spec.template.spec.containers[ + 0 + ].image self.assertEqual(worker_image, "gcr.io/custom/server:v2.0") # Verify roundtrip import. @@ -792,15 +895,15 @@ def test_direct_mutation(self): job["name"]: job for job in config["spec"]["replicatedJobs"] } self.assertEqual( - replicated_jobs["pathways-head"]["template"]["spec"]["template"]["spec"][ - "activeDeadlineSeconds" - ], + replicated_jobs["pathways-head"]["template"]["spec"]["template"][ + "spec" + ]["activeDeadlineSeconds"], 100, ) self.assertEqual( - replicated_jobs["pathways-worker"]["template"]["spec"]["template"]["spec"][ - "activeDeadlineSeconds" - ], + replicated_jobs["pathways-worker"]["template"]["spec"]["template"][ + "spec" + ]["activeDeadlineSeconds"], 200, ) @@ -845,6 +948,65 @@ def test_shared_pathways_service(self): self.assertNotIn("pathways-proxy", helper.containers["pathways-head"]) self.assertLen(pod_spec["containers"], 1) + def test_with_custom_libtpu(self): + uri = "https://example.com/libtpu-0.0.48-cp312-cp312-manylinux_2_31.whl" + pw_jobset = self._create_jobset(topology="2x2", num_slices=1) + self.assertIs(pw_jobset.with_custom_libtpu(uri), pw_jobset) + helper = JobSetManifestHelper(pw_jobset.to_dict()) + + for job_name in ("pathways-head", "pathways-worker"): + self.assertIn("emptyDir", helper.volumes[job_name]["custom-libtpu"]) + fetch = helper.init_containers[job_name]["fetch-custom-libtpu"] + self.assertEqual(fetch["command"][:2], ["python3", "-c"]) + self.assertEqual(fetch["command"][-2:], [uri, "/tmp/libtpu/libtpu.so"]) + self.assertEqual( + fetch["volumeMounts"], + [{"name": "custom-libtpu", "mountPath": "/tmp/libtpu"}], + ) + + for name in ("pathways-rm", "pathways-proxy", "pathways-worker"): + matches = helper.get_all_containers_by_name(name) + self.assertLen(matches, 1, name) + _, container = matches[0] + self.assertIn( + { + "name": "custom-libtpu", + "mountPath": "/lib/libtpu.so", + "subPath": "libtpu.so", + }, + container["volumeMounts"], + ) + self.assertIn( + {"name": "TPU_LIBRARY_PATH", "value": "/lib/libtpu.so"}, + container["env"], + ) + + def test_with_custom_libtpu_replaces_previous_override(self): + pw_jobset = self._create_jobset(topology="2x2", num_slices=1) + pw_jobset.with_custom_libtpu("gs://bucket/first/libtpu.so") + pw_jobset.with_custom_libtpu("gs://bucket/second/libtpu.so") + helper = JobSetManifestHelper(pw_jobset.to_dict()) + + self.assertLen(helper.pod_specs["pathways-worker"]["initContainers"], 1) + fetch = helper.init_containers["pathways-worker"]["fetch-custom-libtpu"] + self.assertIn("gs://bucket/second/libtpu.so", fetch["command"]) + worker = helper.containers["pathways-worker"]["pathways-worker"] + self.assertLen( + [e for e in worker["env"] if e["name"] == "TPU_LIBRARY_PATH"], 1 + ) + self.assertLen( + [ + m + for m in worker["volumeMounts"] + if m["mountPath"] == "/lib/libtpu.so" + ], + 1, + ) + + def test_with_custom_libtpu_rejects_unsupported_uri(self): + with self.assertRaisesRegex(ValueError, "Unsupported"): + self._create_jobset().with_custom_libtpu("/local/libtpu.so") + if __name__ == "__main__": absltest.main()