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
4 changes: 2 additions & 2 deletions python/pyspark/sql/connect/client/artifact.py
Original file line number Diff line number Diff line change
Expand Up @@ -168,7 +168,7 @@ def __init__(
user_id: Optional[str],
session_id: str,
channel: grpc.Channel,
metadata: Iterable[Tuple[str, str]],
metadata: List[Tuple[str, str]],
add_artifacts_timeout: Optional[float] = None,
artifact_status_timeout: Optional[float] = None,
):
Expand All @@ -177,7 +177,7 @@ def __init__(
self._user_context.user_id = user_id
self._stub = grpc_lib.SparkConnectServiceStub(channel)
self._session_id = session_id
self._metadata = metadata
self._metadata: List[Tuple[str, str]] = list(metadata)
self._add_artifacts_timeout = add_artifacts_timeout
self._artifact_status_timeout = artifact_status_timeout

Expand Down
2 changes: 1 addition & 1 deletion python/pyspark/sql/connect/client/core.py
Original file line number Diff line number Diff line change
Expand Up @@ -400,7 +400,7 @@ def _effective_channel_options(self) -> List[Tuple[str, Any]]:
options.append((key, value))
return options

def metadata(self) -> Iterable[Tuple[str, str]]:
def metadata(self) -> List[Tuple[str, str]]:
"""
Builds the GRPC specific metadata list to be injected into the request. All
parameters will be converted to metadata except ones that are explicitly used
Expand Down
10 changes: 6 additions & 4 deletions python/pyspark/sql/connect/client/reattach.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,7 @@
from threading import RLock
import uuid
from collections.abc import Generator
from typing import Optional, Any, Iterator, Iterable, Tuple, Callable, cast, ClassVar
from typing import Optional, Any, Iterator, Iterable, List, Tuple, Callable, cast, ClassVar
from concurrent.futures import Future, ThreadPoolExecutor
import os
import weakref
Expand Down Expand Up @@ -71,7 +71,7 @@ def __init__(
request: pb2.ExecutePlanRequest,
stub: grpc_lib.SparkConnectServiceStub,
retrying: Callable[[], Retrying],
metadata: Iterable[Tuple[str, str]],
metadata: List[Tuple[str, str]],
reattachable_execute_plan_timeout: Optional[float] = None,
reattach_execute_timeout: Optional[float] = None,
):
Expand Down Expand Up @@ -109,12 +109,14 @@ def __init__(
# Initial iterator comes from ExecutePlan request.
# Note: This is not retried, because no error would ever be thrown here, and GRPC will only
# throw error on first self._has_next().
self._metadata = metadata
# Convert metadata to a list to ensure it remains re-iterable across all RPCs
# (ReattachExecute, ReleaseExecute), so auth headers are always present.
self._metadata: List[Tuple[str, str]] = list(metadata)
Comment thread
anupamme marked this conversation as resolved.
with disable_gc():
self._iterator: Optional[Iterator[pb2.ExecutePlanResponse]] = iter(
self._stub.ExecutePlan(
self._initial_request,
metadata=metadata,
metadata=self._metadata,
timeout=self._reattachable_execute_plan_timeout,
)
)
Expand Down
43 changes: 43 additions & 0 deletions python/pyspark/sql/tests/connect/client/test_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -113,16 +113,23 @@ def __init__(self, execute_ops=None, attach_ops=None):
self.release_calls = 0
self.release_until_calls = 0
self.attach_calls = 0
# Metadata recorded per call (each entry is the list passed as metadata=)
self.execute_metadata = []
self.attach_metadata = []
self.release_metadata = []

def ExecutePlan(self, *args, **kwargs):
self.execute_calls += 1
self.execute_metadata.append(list(kwargs.get("metadata", [])))
return self._execute_ops

def ReattachExecute(self, *args, **kwargs):
self.attach_calls += 1
self.attach_metadata.append(list(kwargs.get("metadata", [])))
return self._attach_ops

def ReleaseExecute(self, req: proto.ReleaseExecuteRequest, *args, **kwargs):
self.release_metadata.append(list(kwargs.get("metadata", [])))
if req.HasField("release_all"):
self.release_calls += 1
elif req.HasField("release_until"):
Expand Down Expand Up @@ -1186,6 +1193,42 @@ def raise_with_sql_state():
self.assertEqual(err.getErrorClass(), expected_error_class)
self.assertEqual(err.getSqlState(), expected_sql_state)

def test_generator_metadata_preserved_across_rpcs(self):
# A single-use generator passed as metadata must not be exhausted before
# ReattachExecute and ReleaseExecute calls; list() in __init__ prevents this.
expected_header = ("x-auth-token", "secret")

def non_fatal():
raise TestException("Non Fatal", grpc.StatusCode.UNAVAILABLE)

stub = self._stub_with(
[self.response, non_fatal], [self.response, self.finished]
)
metadata_gen = (x for x in [expected_header])

ite = ExecutePlanResponseReattachableIterator(
self.request, stub, self.retrying, metadata_gen
)
for _ in ite:
pass

def check():
self.assertEqual(1, stub.attach_calls)
self.assertEqual(1, stub.release_calls)
# All three RPC types must receive the header. If list(metadata) on
# line 114 runs before ExecutePlan (which it does), but ExecutePlan
# still uses the raw `metadata` parameter instead of self._metadata,
# then passing a generator would leave execute_metadata[0] empty.
self.assertEqual(1, len(stub.execute_metadata))
self.assertIn(expected_header, stub.execute_metadata[0])
self.assertEqual(1, len(stub.attach_metadata))
self.assertIn(expected_header, stub.attach_metadata[0])
self.assertGreater(len(stub.release_metadata), 0)
for meta in stub.release_metadata:
self.assertIn(expected_header, meta)

eventually(timeout=1, catch_assertions=True)(check)()


if __name__ == "__main__":
from pyspark.testing import main
Expand Down