diff --git a/python/pyspark/sql/connect/client/artifact.py b/python/pyspark/sql/connect/client/artifact.py index 94879171b41f0..2ab1ee88fe7dd 100644 --- a/python/pyspark/sql/connect/client/artifact.py +++ b/python/pyspark/sql/connect/client/artifact.py @@ -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, ): @@ -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 diff --git a/python/pyspark/sql/connect/client/core.py b/python/pyspark/sql/connect/client/core.py index 43a22c4998f2b..c3628c4d110e2 100644 --- a/python/pyspark/sql/connect/client/core.py +++ b/python/pyspark/sql/connect/client/core.py @@ -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 diff --git a/python/pyspark/sql/connect/client/reattach.py b/python/pyspark/sql/connect/client/reattach.py index f1d06320866e0..0171774ad331b 100644 --- a/python/pyspark/sql/connect/client/reattach.py +++ b/python/pyspark/sql/connect/client/reattach.py @@ -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 @@ -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, ): @@ -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) 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, ) ) diff --git a/python/pyspark/sql/tests/connect/client/test_client.py b/python/pyspark/sql/tests/connect/client/test_client.py index b5bf76d86df48..64216c16e1efe 100644 --- a/python/pyspark/sql/tests/connect/client/test_client.py +++ b/python/pyspark/sql/tests/connect/client/test_client.py @@ -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"): @@ -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