Skip to content
59 changes: 58 additions & 1 deletion python/pyspark/sql/connect/client/core.py
Original file line number Diff line number Diff line change
Expand Up @@ -709,6 +709,9 @@ def __init__(
# cleanup ml cache if possible
atexit.register(self._cleanup_ml_cache)

self.global_user_context_extensions = []
self.global_user_context_extensions_lock = threading.Lock()

@property
def _stub(self) -> grpc_lib.SparkConnectServiceStub:
if self.is_closed:
Expand Down Expand Up @@ -1240,6 +1243,24 @@ def token(self) -> Optional[str]:
"""
return self._builder.token

def _update_request_with_user_context_extensions(
self,
req: Union[
pb2.AnalyzePlanRequest,
pb2.ConfigRequest,
pb2.ExecutePlanRequest,
pb2.FetchErrorDetailsRequest,
pb2.InterruptRequest,
],
) -> None:
with self.global_user_context_extensions_lock:
for _, extension in self.global_user_context_extensions:
req.user_context.extensions.append(extension)
if not hasattr(self.thread_local, "user_context_extensions"):
return
for _, extension in self.thread_local.user_context_extensions:
req.user_context.extensions.append(extension)

def _execute_plan_request_with_metadata(
self, operation_id: Optional[str] = None
) -> pb2.ExecutePlanRequest:
Expand Down Expand Up @@ -1270,6 +1291,7 @@ def _execute_plan_request_with_metadata(
messageParameters={"arg_name": "operation_id", "origin": str(ve)},
)
req.operation_id = operation_id
self._update_request_with_user_context_extensions(req)
return req

def _analyze_plan_request_with_metadata(self) -> pb2.AnalyzePlanRequest:
Expand All @@ -1280,6 +1302,7 @@ def _analyze_plan_request_with_metadata(self) -> pb2.AnalyzePlanRequest:
req.client_type = self._builder.userAgent
if self._user_id:
req.user_context.user_id = self._user_id
self._update_request_with_user_context_extensions(req)
return req

def _analyze(self, method: str, **kwargs: Any) -> AnalyzeResult:
Expand Down Expand Up @@ -1694,6 +1717,7 @@ def _config_request_with_metadata(self) -> pb2.ConfigRequest:
req.client_type = self._builder.userAgent
if self._user_id:
req.user_context.user_id = self._user_id
self._update_request_with_user_context_extensions(req)
return req

def get_configs(self, *keys: str) -> Tuple[Optional[str], ...]:
Expand Down Expand Up @@ -1770,6 +1794,7 @@ def _interrupt_request(
)
if self._user_id:
req.user_context.user_id = self._user_id
self._update_request_with_user_context_extensions(req)
return req

def interrupt_all(self) -> Optional[List[str]]:
Expand Down Expand Up @@ -1868,6 +1893,38 @@ def _throw_if_invalid_tag(self, tag: str) -> None:
messageParameters={"arg_name": "Spark Connect tag", "arg_value": tag},
)

def add_threadlocal_user_context_extension(self, extension: any_pb2.Any) -> str:
if not hasattr(self.thread_local, "user_context_extensions"):
self.thread_local.user_context_extensions = list()
extension_id = "threadlocal_" + str(uuid.uuid4())
self.thread_local.user_context_extensions.append((extension_id, extension))
return extension_id

def add_global_user_context_extension(self, extension: any_pb2.Any) -> str:
extension_id = "global_" + str(uuid.uuid4())
with self.global_user_context_extensions_lock:
self.global_user_context_extensions.append((extension_id, extension))
return extension_id

def remove_user_context_extension(self, extension_id: str) -> None:
if extension_id.find("threadlocal_") == 0:
if not hasattr(self.thread_local, "user_context_extensions"):
return
self.thread_local.user_context_extensions = list(
filter(lambda ex: ex[0] != extension_id, self.thread_local.user_context_extensions)
)
elif extension_id.find("global_") == 0:
with self.global_user_context_extensions_lock:
self.global_user_context_extensions = list(
filter(lambda ex: ex[0] != extension_id, self.global_user_context_extensions)
)

def clear_user_context_extensions(self) -> None:
if hasattr(self.thread_local, "user_context_extensions"):
self.thread_local.user_context_extensions = list()
with self.global_user_context_extensions_lock:
self.global_user_context_extensions = list()

def _handle_error(self, error: Exception) -> NoReturn:
"""
Handle errors that occur during RPC calls.
Expand Down Expand Up @@ -1908,7 +1965,7 @@ def _fetch_enriched_error(self, info: "ErrorInfo") -> Optional[pb2.FetchErrorDet
req.client_observed_server_side_session_id = self._server_session_id
if self._user_id:
req.user_context.user_id = self._user_id

self._update_request_with_user_context_extensions(req)
try:
return self._stub.FetchErrorDetails(req, metadata=self._builder.metadata())
except grpc.RpcError:
Expand Down
55 changes: 55 additions & 0 deletions python/pyspark/sql/connect/session.py
Original file line number Diff line number Diff line change
Expand Up @@ -42,6 +42,7 @@
ClassVar,
)

import google.protobuf.any_pb2 as any_pb2
import numpy as np
import pandas as pd
import pyarrow as pa
Expand Down Expand Up @@ -894,6 +895,60 @@ def clearTags(self) -> None:

clearTags.__doc__ = PySparkSession.clearTags.__doc__

def addThreadlocalUserContextExtension(self, extension: any_pb2.Any) -> str:
Comment thread
cookiedough77 marked this conversation as resolved.
Outdated
"""
Add a user context extension to the current session in the current thread.
It will be sent in the UserContext of every request sent from the current thread, until
it is removed with removeUserContextExtension using the returned id.

Parameters
----------
extension: any_pb2.Any
Protobuf Any message to add as the extension to UserContext.

Returns
-------
str
Id that can be used with removeUserContextExtension to remove the extension.
"""
return self.client.add_threadlocal_user_context_extension(extension)

def addGlobalUserContextExtension(self, extension: any_pb2.Any) -> str:
"""
Add a user context extension to the current session, globally.
It will be sent in the UserContext of every request, until it is removed with
removeUserContextExtension using the returned id. It will precede any threadlocal extension.

Parameters
----------
extension: any_pb2.Any
Protobuf Any message to add as the extension to UserContext.

Returns
-------
str
Id that can be used with removeUserContextExtension to remove the extension.
"""
return self.client.add_global_user_context_extension(extension)

def removeUserContextExtension(self, extension_id: str) -> None:
"""
Remove a user context extension previously added by addThreadlocalUserContextExtension.

Parameters
----------
extension_id: str
id returned by addThreadlocalUserContextExtension.
"""
self.client.remove_user_context_extension(extension_id)

def clearUserContextExtensions(self) -> None:
"""
Clear all user context extensions previously added by addGlobalUserContextExtension and
addThreadlocalUserContextExtension
"""
self.client.clear_user_context_extensions()

def stop(self) -> None:
"""
Release the current session and close the GRPC connection to the Spark Connect server.
Expand Down
91 changes: 91 additions & 0 deletions python/pyspark/sql/tests/connect/client/test_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -136,9 +136,11 @@ class MockService:
def __init__(self, session_id: str):
self._session_id = session_id
self.req = None
self.client_user_context_extensions = []

def ExecutePlan(self, req: proto.ExecutePlanRequest, metadata):
self.req = req
self.client_user_context_extensions = req.user_context.extensions
resp = proto.ExecutePlanResponse()
resp.session_id = self._session_id
resp.operation_id = req.operation_id
Expand All @@ -159,12 +161,14 @@ def ExecutePlan(self, req: proto.ExecutePlanRequest, metadata):

def Interrupt(self, req: proto.InterruptRequest, metadata):
self.req = req
self.client_user_context_extensions = req.user_context.extensions
resp = proto.InterruptResponse()
resp.session_id = self._session_id
return resp

def Config(self, req: proto.ConfigRequest, metadata):
self.req = req
self.client_user_context_extensions = req.user_context.extensions
resp = proto.ConfigResponse()
resp.session_id = self._session_id
if req.operation.HasField("get"):
Expand Down Expand Up @@ -229,6 +233,93 @@ def userId(self) -> Optional[str]:

self.assertEqual(client._user_id, "abc")

def test_user_context_extension(self):
client = SparkConnectClient("sc://foo/", use_reattachable_execute=False)
mock = MockService(client._session_id)
client._stub = mock

exlocal = any_pb2.Any()
exlocal.Pack(wrappers_pb2.StringValue(value="abc"))
exlocal2 = any_pb2.Any()
exlocal2.Pack(wrappers_pb2.StringValue(value="def"))
exglobal = any_pb2.Any()
exglobal.Pack(wrappers_pb2.StringValue(value="ghi"))
exglobal2 = any_pb2.Any()
exglobal2.Pack(wrappers_pb2.StringValue(value="jkl"))

exlocal_id = client.add_threadlocal_user_context_extension(exlocal)
exglobal_id = client.add_global_user_context_extension(exglobal)

mock.client_user_context_extensions = []
command = proto.Command()
client.execute_command(command)
self.assertTrue(exlocal in mock.client_user_context_extensions)
self.assertTrue(exglobal in mock.client_user_context_extensions)
self.assertFalse(exlocal2 in mock.client_user_context_extensions)
self.assertFalse(exglobal2 in mock.client_user_context_extensions)

client.add_threadlocal_user_context_extension(exlocal2)

mock.client_user_context_extensions = []
plan = proto.Plan()
client.semantic_hash(plan) # use semantic_hash to test analyze
self.assertTrue(exlocal in mock.client_user_context_extensions)
self.assertTrue(exglobal in mock.client_user_context_extensions)
self.assertTrue(exlocal2 in mock.client_user_context_extensions)
self.assertFalse(exglobal2 in mock.client_user_context_extensions)

client.add_global_user_context_extension(exglobal2)

mock.client_user_context_extensions = []
client.interrupt_all()
self.assertTrue(exlocal in mock.client_user_context_extensions)
self.assertTrue(exglobal in mock.client_user_context_extensions)
self.assertTrue(exlocal2 in mock.client_user_context_extensions)
self.assertTrue(exglobal2 in mock.client_user_context_extensions)

client.remove_user_context_extension(exlocal_id)

mock.client_user_context_extensions = []
client.get_configs("foo", "bar")
self.assertFalse(exlocal in mock.client_user_context_extensions)
self.assertTrue(exglobal in mock.client_user_context_extensions)
self.assertTrue(exlocal2 in mock.client_user_context_extensions)
self.assertTrue(exglobal2 in mock.client_user_context_extensions)

client.remove_user_context_extension(exglobal_id)

mock.client_user_context_extensions = []
command = proto.Command()
client.execute_command(command)
self.assertFalse(exlocal in mock.client_user_context_extensions)
self.assertFalse(exglobal in mock.client_user_context_extensions)
self.assertTrue(exlocal2 in mock.client_user_context_extensions)
self.assertTrue(exglobal2 in mock.client_user_context_extensions)

client.clear_user_context_extensions()

mock.client_user_context_extensions = []
plan = proto.Plan()
client.semantic_hash(plan) # use semantic_hash to test analyze
self.assertFalse(exlocal in mock.client_user_context_extensions)
self.assertFalse(exglobal in mock.client_user_context_extensions)
self.assertFalse(exlocal2 in mock.client_user_context_extensions)
self.assertFalse(exglobal2 in mock.client_user_context_extensions)

mock.client_user_context_extensions = []
client.interrupt_all()
self.assertFalse(exlocal in mock.client_user_context_extensions)
self.assertFalse(exglobal in mock.client_user_context_extensions)
self.assertFalse(exlocal2 in mock.client_user_context_extensions)
self.assertFalse(exglobal2 in mock.client_user_context_extensions)

mock.client_user_context_extensions = []
client.get_configs("foo", "bar")
self.assertFalse(exlocal in mock.client_user_context_extensions)
self.assertFalse(exglobal in mock.client_user_context_extensions)
self.assertFalse(exlocal2 in mock.client_user_context_extensions)
self.assertFalse(exglobal2 in mock.client_user_context_extensions)

def test_interrupt_all(self):
client = SparkConnectClient("sc://foo/;token=bar", use_reattachable_execute=False)
mock = MockService(client._session_id)
Expand Down