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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
22 changes: 17 additions & 5 deletions kubeflow/spark/backends/kubernetes/backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -457,13 +457,13 @@ def get_connect_url(
if port is None:
port_str = os.environ.get("SPARK_CONNECT_LOCAL_PORT")
port = int(port_str) if port_str else random.randint(15002, 16002)
# Prefer pod when available (bypasses Service/EndpointSlice); then try svc names
# Prefer service-based port-forward (more stable via kube-proxy); fall back to pod
candidates: list[tuple[str, str]] = []
for svc in [info.service_name, f"{info.name}-server", f"{info.name}-svc"]:
if svc and not any(c[1] == svc for c in candidates):
candidates.append(("svc", svc))
if info.driver_pod_name:
candidates.append(("pod", info.driver_pod_name))
for svc in [f"{info.name}-svc", info.service_name, f"{info.name}-server"]:
if svc and not any(c[0] == "svc" and c[1] == svc for c in candidates):
candidates.append(("svc", svc))
seen: set[str] = set()
for kind, target in candidates:
key = f"{kind}/{target}"
Expand All @@ -475,6 +475,8 @@ def get_connect_url(
cmd = [
"kubectl",
"port-forward",
"--address",
"127.0.0.1",
key,
f"{port}:{constants.SPARK_CONNECT_PORT}",
"-n",
Expand Down Expand Up @@ -596,9 +598,19 @@ def connect(
while time.monotonic() - probe_start < delay_sec:
# Check if port-forward process died
if pf_proc is not None and pf_proc.poll() is not None:
logger.warning("Port-forward died during gRPC ready wait, restarting...")
stderr_b = pf_proc.stderr.read() if pf_proc.stderr else b""
stderr_str = (
stderr_b.decode("utf-8", errors="replace").strip() if stderr_b else ""
)
logger.warning(
"Port-forward died during gRPC ready wait, restarting... "
"(exit=%s, stderr=%s)",
pf_proc.returncode,
stderr_str,
)
connect_url, pf_proc = self.get_connect_url(info, local_port=local_port)
local_port = int(connect_url.split(":")[-1]) if pf_proc else None
time.sleep(2) # Allow new port-forward to stabilize
# Verify port is still reachable
if local_port and not self._wait_for_connect_port(
"127.0.0.1", local_port, timeout_sec=1, interval_sec=0.5
Expand Down
10 changes: 9 additions & 1 deletion kubeflow/spark/backends/kubernetes/backend_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -202,6 +202,7 @@ def mock_get_response(name: str) -> dict:
server_status=models.SparkV1alpha1SparkConnectServerStatus(
pod_name=f"{name}-0",
pod_ip="10.0.0.5",
service_name=f"{name}-server-svc",
),
).to_dict()
elif name == SPARK_CONNECT_PROVISIONING:
Expand Down Expand Up @@ -725,12 +726,19 @@ def test_get_connect_url(kubernetes_backend, test_case):
patch(
"kubeflow.spark.backends.kubernetes.backend.subprocess.Popen",
return_value=mock_popen,
),
) as mock_popen_cls,
patch("kubeflow.spark.backends.kubernetes.backend.time.sleep"),
patch.object(kubernetes_backend, "_wait_for_connect_port", return_value=True),
):
url, proc = kubernetes_backend.get_connect_url(info)

# Verify service is tried before pod (service-first ordering)
first_call_args = mock_popen_cls.call_args[0][0]
first_call_key = first_call_args[4]
assert first_call_key.startswith("svc/"), (
f"Expected service-first port-forward, got {first_call_key}"
)

if "url_contains" in test_case.expected_output:
assert test_case.expected_output["url_contains"] in url
else:
Expand Down