use std::collections::BTreeMap;
use std::str::FromStr;
use tonic::metadata::MetadataValue;
use tonic::transport::{Channel, ClientTlsConfig};
use tonic::{Request, Response, Streaming};
use uuid::Uuid;
use spark_connect_proto::spark_connect_service_client::SparkConnectServiceClient;
use spark_connect_proto::{
AnalyzePlanRequest, AnalyzePlanResponse, ArtifactStatusesRequest, ConfigRequest,
ConfigResponse, ExecutePlanRequest, ExecutePlanResponse, FetchErrorDetailsRequest,
FetchErrorDetailsResponse, InterruptRequest, InterruptResponse, ReattachExecuteRequest,
ReleaseExecuteRequest, ReleaseExecuteResponse, ReleaseSessionRequest, UserContext,
};
use crate::artifact::{build_artifact_request_stream, FileArtifact, NamedArtifact};
use crate::channel::{ChannelBuilder, GRPC_MAX_MESSAGE_LENGTH_DEFAULT};
use crate::error::{Result, SparkError};
use crate::reattach::ExecutePlanResponseReattachableIterator;
use crate::retries::{RetryPolicy, RetryPolicyState};
const READY_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(10);
const AUTHORIZATION_HEADER: &str = "authorization";
async fn wait_ready(grpc: &mut tonic::client::Grpc<Channel>) -> Result<()> {
match tokio::time::timeout(READY_TIMEOUT, grpc.ready()).await {
Ok(Ok(())) => Ok(()),
Ok(Err(e)) => Err(SparkError::from_grpc_status(tonic::Status::unavailable(
format!("channel not ready: {e}"),
))),
Err(_elapsed) => Err(SparkError::from_grpc_status(tonic::Status::unavailable(
"channel not ready: connection timed out",
))),
}
}
pub struct SparkConnectClient {
channel: Channel,
stub: SparkConnectServiceClient<Channel>,
session_id: String,
user_id: Option<String>,
metadata: Vec<(String, String)>,
user_agent: String,
retry_policy: RetryPolicy,
}
impl SparkConnectClient {
pub async fn connect(builder: &ChannelBuilder) -> Result<Self> {
let session_id = match builder.session_id()? {
Some(id) => id,
None => Uuid::new_v4().to_string(),
};
let user_agent = builder.user_agent()?;
let endpoint = builder.endpoint();
let channel = if builder.use_ssl() {
let tls_config = ClientTlsConfig::new()
.domain_name(builder.host())
.with_native_roots();
tonic::transport::Channel::from_shared(format!("https://{endpoint}"))
.map_err(|e| SparkError::connect_msg(format!("Invalid endpoint: {}", e)))?
.tls_config(tls_config)
.map_err(|e| SparkError::connect_msg(format!("Failed to configure TLS: {}", e)))?
.connect_timeout(READY_TIMEOUT)
.connect_lazy()
} else {
tonic::transport::Channel::from_shared(format!("http://{endpoint}"))
.map_err(|e| SparkError::connect_msg(format!("Invalid endpoint: {}", e)))?
.connect_timeout(READY_TIMEOUT)
.connect_lazy()
};
let stub = SparkConnectServiceClient::new(channel.clone())
.max_decoding_message_size(GRPC_MAX_MESSAGE_LENGTH_DEFAULT)
.max_encoding_message_size(GRPC_MAX_MESSAGE_LENGTH_DEFAULT);
let mut metadata: Vec<(String, String)> = builder.metadata();
if let Some(tok) = builder.token() {
if !builder.use_ssl() && !builder.is_loopback() {
return Err(SparkError::value_msg(format!(
"Refusing to send the authentication token to '{}' over an insecure \
connection. Set 'use_ssl=true' in the connection string to enable TLS \
(a token is only allowed without TLS when connecting to localhost).",
builder.host()
)));
}
metadata.push((AUTHORIZATION_HEADER.to_string(), format!("Bearer {}", tok)));
}
Ok(Self {
channel,
stub,
session_id,
user_id: builder.user_id().map(String::from),
metadata,
user_agent,
retry_policy: RetryPolicy::default(),
})
}
pub fn session_id(&self) -> &str {
&self.session_id
}
pub fn with_new_session_id(&self) -> Self {
SparkConnectClient {
channel: self.channel.clone(),
stub: self.stub.clone(),
session_id: Uuid::new_v4().to_string(),
user_id: self.user_id.clone(),
metadata: self.metadata.clone(),
user_agent: self.user_agent.clone(),
retry_policy: self.retry_policy.clone(),
}
}
pub fn user_id(&self) -> Option<&str> {
self.user_id.as_deref()
}
pub fn set_retry_policy(&mut self, policy: RetryPolicy) {
self.retry_policy = policy;
}
pub async fn execute_plan(
&self,
request: ExecutePlanRequest,
) -> Result<Streaming<ExecutePlanResponse>> {
let mut req = Request::new(request);
self._attach_metadata(&mut req);
let resp = self.stub.clone().execute_plan(req).await;
resp.map(Response::into_inner)
.map_err(SparkError::from_grpc_status)
}
pub async fn execute_plan_raw(&self, request: Vec<u8>) -> Result<Streaming<Vec<u8>>> {
let mut grpc = tonic::client::Grpc::new(self.channel.clone())
.max_decoding_message_size(GRPC_MAX_MESSAGE_LENGTH_DEFAULT)
.max_encoding_message_size(GRPC_MAX_MESSAGE_LENGTH_DEFAULT);
wait_ready(&mut grpc).await?;
let mut req = Request::new(request);
self._attach_metadata(&mut req);
let path = tonic::codegen::http::uri::PathAndQuery::from_static(
"/spark.connect.SparkConnectService/ExecutePlan",
);
let resp = grpc
.server_streaming(req, path, crate::bytes_codec::BytesCodec)
.await
.map_err(SparkError::from_grpc_status)?;
Ok(resp.into_inner())
}
pub async fn reattach_execute_raw(&self, request: Vec<u8>) -> Result<Streaming<Vec<u8>>> {
let mut grpc = tonic::client::Grpc::new(self.channel.clone())
.max_decoding_message_size(GRPC_MAX_MESSAGE_LENGTH_DEFAULT)
.max_encoding_message_size(GRPC_MAX_MESSAGE_LENGTH_DEFAULT);
wait_ready(&mut grpc).await?;
let mut req = Request::new(request);
self._attach_metadata(&mut req);
let path = tonic::codegen::http::uri::PathAndQuery::from_static(
"/spark.connect.SparkConnectService/ReattachExecute",
);
let resp = grpc
.server_streaming(req, path, crate::bytes_codec::BytesCodec)
.await
.map_err(SparkError::from_grpc_status)?;
Ok(resp.into_inner())
}
pub async fn reattach_execute(
&self,
request: ReattachExecuteRequest,
) -> Result<Streaming<ExecutePlanResponse>> {
let mut req = Request::new(request);
self._attach_metadata(&mut req);
let resp = self.stub.clone().reattach_execute(req).await;
resp.map(Response::into_inner)
.map_err(SparkError::from_grpc_status)
}
pub async fn release_execute(
&self,
request: ReleaseExecuteRequest,
) -> Result<ReleaseExecuteResponse> {
let mut req = Request::new(request);
self._attach_metadata(&mut req);
let resp = self.stub.clone().release_execute(req).await;
resp.map(Response::into_inner)
.map_err(SparkError::from_grpc_status)
}
pub async fn analyze_plan_raw(&self, request: Vec<u8>) -> Result<Vec<u8>> {
let mut grpc = tonic::client::Grpc::new(self.channel.clone())
.max_decoding_message_size(GRPC_MAX_MESSAGE_LENGTH_DEFAULT)
.max_encoding_message_size(GRPC_MAX_MESSAGE_LENGTH_DEFAULT);
wait_ready(&mut grpc).await?;
let mut req = Request::new(request);
self._attach_metadata(&mut req);
let path = tonic::codegen::http::uri::PathAndQuery::from_static(
"/spark.connect.SparkConnectService/AnalyzePlan",
);
let resp = grpc
.unary(req, path, crate::bytes_codec::BytesCodec)
.await
.map_err(SparkError::from_grpc_status)?;
Ok(resp.into_inner())
}
pub async fn analyze_plan(&self, request: AnalyzePlanRequest) -> Result<AnalyzePlanResponse> {
self.with_retry(|| {
let mut req = Request::new(request.clone());
self._attach_metadata(&mut req);
let mut stub = self.stub.clone();
async move { stub.analyze_plan(req).await.map(Response::into_inner) }
})
.await
}
async fn with_retry<T, Fut, F>(&self, mut op: F) -> Result<T>
where
F: FnMut() -> Fut,
Fut: std::future::Future<Output = std::result::Result<T, tonic::Status>>,
{
let mut state = RetryPolicyState::new(self.retry_policy.clone());
loop {
match op().await {
Ok(v) => return Ok(v),
Err(status) => {
if self.retry_policy.can_retry(&status) {
if let Some(wait_ms) = state.next_attempt(None) {
tokio::time::sleep(std::time::Duration::from_millis(wait_ms)).await;
continue;
}
}
return Err(SparkError::from_grpc_status(status));
}
}
}
}
pub async fn execute_plan_reattachable(
&self,
request: ExecutePlanRequest,
) -> Result<ReattachableResponseStream> {
let iter = ExecutePlanResponseReattachableIterator::new(request);
let start = iter.request().clone();
let stream = self
.with_retry(|| {
let mut req = Request::new(start.clone());
self._attach_metadata(&mut req);
let mut stub = self.stub.clone();
async move { stub.execute_plan(req).await.map(Response::into_inner) }
})
.await?;
let retry_policy = self.retry_policy.clone();
Ok(ReattachableResponseStream {
stub: self.stub.clone(),
user_agent: self.user_agent.clone(),
user_id: self.user_id.clone(),
metadata: self.metadata.clone(),
retry_state: RetryPolicyState::new(retry_policy.clone()),
retry_policy,
iter,
stream,
done: false,
})
}
pub async fn fetch_error_details(
&self,
request: FetchErrorDetailsRequest,
) -> Result<FetchErrorDetailsResponse> {
let mut req = Request::new(request);
self._attach_metadata(&mut req);
let resp = self.stub.clone().fetch_error_details(req).await;
resp.map(Response::into_inner)
.map_err(SparkError::from_grpc_status)
}
pub async fn config(&self, request: ConfigRequest) -> Result<ConfigResponse> {
self.with_retry(|| {
let mut req = Request::new(request.clone());
self._attach_metadata(&mut req);
let mut stub = self.stub.clone();
async move { stub.config(req).await.map(Response::into_inner) }
})
.await
}
pub async fn interrupt(&self, request: InterruptRequest) -> Result<InterruptResponse> {
self.with_retry(|| {
let mut req = Request::new(request.clone());
self._attach_metadata(&mut req);
let mut stub = self.stub.clone();
async move { stub.interrupt(req).await.map(Response::into_inner) }
})
.await
}
pub async fn get_configs(&self, keys: &[&str]) -> Result<Vec<Option<String>>> {
let mut request = ConfigRequest::default();
request.session_id = self.session_id.clone();
request.user_context = Some(UserContext {
user_id: self.user_id.clone().unwrap_or_default(),
..Default::default()
});
let mut operation = spark_connect_proto::config_request::Operation::default();
let mut get = spark_connect_proto::config_request::Get::default();
get.keys = keys.iter().map(|k| k.to_string()).collect();
operation.op_type = Some(spark_connect_proto::config_request::operation::OpType::Get(
get,
));
request.operation = Some(operation);
let response = self.config(request).await?;
let mut config_dict = BTreeMap::new();
for pair in response.pairs {
if let Some(value) = pair.value {
config_dict.insert(pair.key, value);
}
}
Ok(keys.iter().map(|k| config_dict.get(*k).cloned()).collect())
}
pub async fn set_config(&self, key: &str, value: &str) -> Result<()> {
let mut request = ConfigRequest::default();
request.session_id = self.session_id.clone();
request.user_context = Some(UserContext {
user_id: self.user_id.clone().unwrap_or_default(),
..Default::default()
});
let mut set = spark_connect_proto::config_request::Set::default();
set.pairs.push(spark_connect_proto::KeyValue {
key: key.to_string(),
value: Some(value.to_string()),
});
let mut operation = spark_connect_proto::config_request::Operation::default();
operation.op_type = Some(spark_connect_proto::config_request::operation::OpType::Set(
set,
));
request.operation = Some(operation);
let _ = self.config(request).await?;
Ok(())
}
pub async fn get_config_with_defaults(
&self,
pairs: &[(&str, Option<&str>)],
) -> Result<Vec<Option<String>>> {
let mut request = ConfigRequest::default();
request.session_id = self.session_id.clone();
request.user_context = Some(UserContext {
user_id: self.user_id.clone().unwrap_or_default(),
..Default::default()
});
let mut get_with_default = spark_connect_proto::config_request::GetWithDefault::default();
get_with_default.pairs = pairs
.iter()
.map(|(key, default_value)| spark_connect_proto::KeyValue {
key: key.to_string(),
value: default_value.map(|v| v.to_string()),
})
.collect();
let mut operation = spark_connect_proto::config_request::Operation::default();
operation.op_type = Some(
spark_connect_proto::config_request::operation::OpType::GetWithDefault(
get_with_default,
),
);
request.operation = Some(operation);
let response = self.config(request).await?;
let mut config_dict = BTreeMap::new();
for pair in response.pairs {
if let Some(value) = pair.value {
config_dict.insert(pair.key, value);
}
}
Ok(pairs
.iter()
.map(|(k, _)| config_dict.get(*k).cloned())
.collect())
}
pub async fn unset_config(&self, key: &str) -> Result<()> {
let mut request = ConfigRequest::default();
request.session_id = self.session_id.clone();
request.user_context = Some(UserContext {
user_id: self.user_id.clone().unwrap_or_default(),
..Default::default()
});
let mut unset = spark_connect_proto::config_request::Unset::default();
unset.keys = vec![key.to_string()];
let mut operation = spark_connect_proto::config_request::Operation::default();
operation.op_type =
Some(spark_connect_proto::config_request::operation::OpType::Unset(unset));
request.operation = Some(operation);
let _ = self.config(request).await?;
Ok(())
}
pub async fn get_configs_all(&self) -> Result<std::collections::HashMap<String, String>> {
let mut request = ConfigRequest::default();
request.session_id = self.session_id.clone();
request.user_context = Some(UserContext {
user_id: self.user_id.clone().unwrap_or_default(),
..Default::default()
});
let mut operation = spark_connect_proto::config_request::Operation::default();
let get_all = spark_connect_proto::config_request::GetAll::default();
operation.op_type =
Some(spark_connect_proto::config_request::operation::OpType::GetAll(get_all));
request.operation = Some(operation);
let response = self.config(request).await?;
let mut config_dict = std::collections::HashMap::new();
for pair in response.pairs {
if let Some(value) = pair.value {
config_dict.insert(pair.key, value);
}
}
Ok(config_dict)
}
pub async fn is_config_modifiable(&self, key: &str) -> Result<bool> {
let mut request = ConfigRequest::default();
request.session_id = self.session_id.clone();
request.user_context = Some(UserContext {
user_id: self.user_id.clone().unwrap_or_default(),
..Default::default()
});
let mut is_modifiable = spark_connect_proto::config_request::IsModifiable::default();
is_modifiable.keys = vec![key.to_string()];
let mut operation = spark_connect_proto::config_request::Operation::default();
operation.op_type = Some(
spark_connect_proto::config_request::operation::OpType::IsModifiable(is_modifiable),
);
request.operation = Some(operation);
let response = self.config(request).await?;
if let Some(pair) = response.pairs.first() {
if let Some(value) = &pair.value {
return Ok(value == "true");
}
}
Ok(false)
}
pub async fn interrupt_all(&self) -> Result<Vec<String>> {
let mut request = InterruptRequest::default();
request.session_id = self.session_id.clone();
request.user_context = Some(UserContext {
user_id: self.user_id.clone().unwrap_or_default(),
..Default::default()
});
request.interrupt_type = 1;
let response = self.interrupt(request).await?;
Ok(response.interrupted_ids)
}
pub async fn interrupt_tag(&self, tag: &str) -> Result<Vec<String>> {
let mut request = InterruptRequest::default();
request.session_id = self.session_id.clone();
request.user_context = Some(UserContext {
user_id: self.user_id.clone().unwrap_or_default(),
..Default::default()
});
request.interrupt_type = 2;
request.interrupt =
Some(spark_connect_proto::interrupt_request::Interrupt::OperationTag(tag.to_string()));
let response = self.interrupt(request).await?;
Ok(response.interrupted_ids)
}
pub async fn interrupt_operation(&self, operation_id: &str) -> Result<Vec<String>> {
let mut request = InterruptRequest::default();
request.session_id = self.session_id.clone();
request.user_context = Some(UserContext {
user_id: self.user_id.clone().unwrap_or_default(),
..Default::default()
});
request.interrupt_type = 3;
request.interrupt = Some(
spark_connect_proto::interrupt_request::Interrupt::OperationId(
operation_id.to_string(),
),
);
let response = self.interrupt(request).await?;
Ok(response.interrupted_ids)
}
pub async fn release_session(&self) -> Result<()> {
let mut request = ReleaseSessionRequest::default();
request.session_id = self.session_id.clone();
request.user_context = Some(UserContext {
user_id: self.user_id.clone().unwrap_or_default(),
..Default::default()
});
let mut req = Request::new(request);
self._attach_metadata(&mut req);
let resp = self.stub.clone().release_session(req).await;
resp.map(|_| ()).map_err(SparkError::from_grpc_status)
}
pub async fn add_artifacts(
&self,
paths: &[&str],
_pyfile: bool,
_archive: bool,
_file: bool,
) -> Result<()> {
let mut artifacts = Vec::new();
for path in paths {
let path_obj = std::path::Path::new(path);
let file_name = path_obj
.file_name()
.and_then(|n| n.to_str())
.ok_or_else(|| {
SparkError::connect_msg(format!("Invalid artifact path: {}", path))
})?;
let artifact_name = if _pyfile
&& (file_name.ends_with(".py")
|| file_name.ends_with(".zip")
|| file_name.ends_with(".egg")
|| file_name.ends_with(".jar"))
{
format!("pyfiles/{}", file_name)
} else if _archive
&& (file_name.ends_with(".zip")
|| file_name.ends_with(".jar")
|| file_name.ends_with(".tar.gz")
|| file_name.ends_with(".tgz")
|| file_name.ends_with(".tar"))
{
format!("archives/{}", file_name)
} else if _file {
format!("files/{}", file_name)
} else if file_name.ends_with(".jar") {
format!("jars/{}", file_name)
} else {
return Err(SparkError::connect_msg(format!(
"Unsupported artifact type: {}",
path
)));
};
artifacts.push(NamedArtifact::new(
artifact_name,
Box::new(FileArtifact::new(path)),
));
}
if artifacts.is_empty() {
return Ok(());
}
let requests = build_artifact_request_stream(self.session_id.clone(), artifacts)?;
let request_stream = futures::stream::iter(requests);
let mut req = Request::new(request_stream);
self._attach_metadata(&mut req);
let response = self
.stub
.clone()
.add_artifacts(req)
.await
.map_err(SparkError::from_grpc_status)?
.into_inner();
for summary in response.artifacts {
if !summary.is_crc_successful {
return Err(SparkError::connect_msg(format!(
"CRC check failed for artifact: {}",
summary.name
)));
}
}
Ok(())
}
pub async fn add_named_artifact(&self, name: &str, local_path: &str) -> Result<()> {
let artifacts = vec![NamedArtifact::new(
name.to_string(),
Box::new(FileArtifact::new(local_path)),
)];
let requests = build_artifact_request_stream(self.session_id.clone(), artifacts)?;
let request_stream = futures::stream::iter(requests);
let mut req = Request::new(request_stream);
self._attach_metadata(&mut req);
let response = self
.stub
.clone()
.add_artifacts(req)
.await
.map_err(SparkError::from_grpc_status)?
.into_inner();
for summary in response.artifacts {
if !summary.is_crc_successful {
return Err(SparkError::connect_msg(format!(
"CRC check failed for artifact: {}",
summary.name
)));
}
}
Ok(())
}
pub async fn artifact_status(
&self,
names: &[&str],
) -> Result<std::collections::HashMap<String, bool>> {
let mut request = ArtifactStatusesRequest::default();
request.session_id = self.session_id.clone();
request.user_context = Some(UserContext {
user_id: self.user_id.clone().unwrap_or_default(),
..Default::default()
});
request.names = names.iter().map(|n| n.to_string()).collect();
let mut req = Request::new(request);
self._attach_metadata(&mut req);
let response = self
.stub
.clone()
.artifact_status(req)
.await
.map_err(SparkError::from_grpc_status)?
.into_inner();
let mut result = std::collections::HashMap::new();
for (name, status) in response.statuses {
result.insert(name, status.exists);
}
Ok(result)
}
fn _attach_metadata<T>(&self, req: &mut Request<T>) {
let metadata = req.metadata_mut();
if let Ok(header_value) = MetadataValue::from_str(&self.user_agent) {
let _ = metadata.insert("user-agent", header_value);
}
if let Some(user_id) = &self.user_id {
if let Ok(header_value) = MetadataValue::from_str(user_id) {
let _ = metadata.insert("user_id", header_value);
}
}
for (k, v) in &self.metadata {
if let Ok(header_value) = MetadataValue::from_str(v) {
if let Ok(key) = tonic::metadata::MetadataKey::from_bytes(k.as_bytes()) {
let _ = metadata.insert(key, header_value);
}
}
}
}
}
pub struct ReattachableResponseStream {
stub: SparkConnectServiceClient<Channel>,
user_agent: String,
user_id: Option<String>,
metadata: Vec<(String, String)>,
retry_policy: RetryPolicy,
retry_state: RetryPolicyState,
iter: ExecutePlanResponseReattachableIterator,
stream: Streaming<ExecutePlanResponse>,
done: bool,
}
impl ReattachableResponseStream {
fn attach_metadata<T>(&self, req: &mut Request<T>) {
let metadata = req.metadata_mut();
if let Ok(v) = MetadataValue::from_str(&self.user_agent) {
let _ = metadata.insert("user-agent", v);
}
if let Some(user_id) = &self.user_id {
if let Ok(v) = MetadataValue::from_str(user_id) {
let _ = metadata.insert("user_id", v);
}
}
for (k, v) in &self.metadata {
if let Ok(v) = MetadataValue::from_str(v) {
if let Ok(key) = tonic::metadata::MetadataKey::from_bytes(k.as_bytes()) {
let _ = metadata.insert(key, v);
}
}
}
}
pub async fn message(&mut self) -> Result<Option<ExecutePlanResponse>> {
if self.done {
return Ok(None);
}
loop {
match self.stream.message().await {
Ok(Some(resp)) => {
self.iter.set_last_response_id(&resp).await;
if self.iter.is_completed().await {
self.done = true;
}
return Ok(Some(resp));
}
Ok(None) => {
self.done = true;
return Ok(None);
}
Err(status) => {
if self.iter.is_completed().await || !self.retry_policy.can_retry(&status) {
return Err(SparkError::from_grpc_status(status));
}
loop {
let reattach = self.iter.create_reattach_request().await;
let mut req = Request::new(reattach);
self.attach_metadata(&mut req);
match self.stub.reattach_execute(req).await {
Ok(r) => {
self.stream = r.into_inner();
break;
}
Err(e) => {
if self.retry_policy.can_retry(&e) {
if let Some(w) = self.retry_state.next_attempt(None) {
tokio::time::sleep(std::time::Duration::from_millis(w))
.await;
continue;
}
}
return Err(SparkError::from_grpc_status(e));
}
}
}
}
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_session_id_generation() {
let builder = ChannelBuilder::parse("sc://localhost").unwrap();
let client = SparkConnectClient::connect(&builder).await.unwrap();
let session_id = client.session_id();
assert!(!session_id.is_empty());
assert!(Uuid::parse_str(session_id).is_ok());
}
#[tokio::test]
async fn test_user_id_passed_through() {
let builder = ChannelBuilder::parse("sc://localhost/;user_id=test_user").unwrap();
let client = SparkConnectClient::connect(&builder).await.unwrap();
assert_eq!(client.user_id(), Some("test_user"));
}
#[tokio::test]
async fn test_metadata_headers() {
let builder = ChannelBuilder::parse("sc://localhost/;custom_header=custom_value").unwrap();
let client = SparkConnectClient::connect(&builder).await.unwrap();
assert!(client.user_agent.contains("spark/"));
}
#[test]
fn test_tls_enabled_channel_builder() {
let builder = ChannelBuilder::parse("sc://example.com/;use_ssl=true").unwrap();
assert!(builder.use_ssl());
assert!(builder.secure());
}
#[tokio::test]
async fn test_token_bearer_in_metadata() {
let builder =
ChannelBuilder::parse("sc://localhost/;use_ssl=true;token=my_token_123").unwrap();
assert!(builder.secure());
assert_eq!(builder.token(), Some("my_token_123".to_string()));
let client = SparkConnectClient::connect(&builder).await.unwrap();
let mut req = Request::new(());
client._attach_metadata(&mut req);
assert_eq!(
req.metadata()
.get("authorization")
.map(|v| v.to_str().unwrap()),
Some("Bearer my_token_123"),
"expected an `Authorization: Bearer` header on the request, got {:?}",
req.metadata()
);
}
#[tokio::test]
async fn test_token_without_ssl_to_remote_host_is_rejected() {
let builder = ChannelBuilder::parse("sc://example.com/;token=SECRET").unwrap();
match SparkConnectClient::connect(&builder).await {
Ok(_) => panic!("expected a token over cleartext to a remote host to be rejected"),
Err(e) => assert!(
e.to_string().contains("use_ssl=true"),
"expected the error to point at use_ssl=true, got: {e}"
),
}
}
#[tokio::test]
async fn test_token_without_ssl_to_localhost_is_allowed() {
for url in [
"sc://localhost/;token=SECRET",
"sc://127.0.0.1/;token=SECRET",
"sc://[::1]/;token=SECRET",
] {
let builder = ChannelBuilder::parse(url).unwrap();
let client = SparkConnectClient::connect(&builder)
.await
.unwrap_or_else(|e| panic!("{url} should be allowed, got: {e}"));
assert!(
client
.metadata
.iter()
.any(|(k, v)| k == "authorization" && v == "Bearer SECRET"),
"{url}: expected the bearer header to be attached, got {:?}",
client.metadata
);
}
}
#[tokio::test]
async fn test_get_configs_request_structure() {
let builder = ChannelBuilder::parse("sc://localhost").unwrap();
let client = SparkConnectClient::connect(&builder).await.unwrap();
let _result = client.get_configs(&["spark.sql.shuffle.partitions"]);
}
#[tokio::test]
async fn test_set_config_request_structure() {
let builder = ChannelBuilder::parse("sc://localhost").unwrap();
let client = SparkConnectClient::connect(&builder).await.unwrap();
let _result = client.set_config("spark.sql.shuffle.partitions", "200");
}
#[tokio::test]
async fn test_unset_config_request_structure() {
let builder = ChannelBuilder::parse("sc://localhost").unwrap();
let client = SparkConnectClient::connect(&builder).await.unwrap();
let _result = client.unset_config("spark.sql.shuffle.partitions");
}
#[tokio::test]
async fn test_interrupt_all_request_structure() {
let builder = ChannelBuilder::parse("sc://localhost").unwrap();
let client = SparkConnectClient::connect(&builder).await.unwrap();
let _result = client.interrupt_all();
}
#[tokio::test]
async fn test_interrupt_tag_request_structure() {
let builder = ChannelBuilder::parse("sc://localhost").unwrap();
let client = SparkConnectClient::connect(&builder).await.unwrap();
let _result = client.interrupt_tag("my-tag");
}
#[tokio::test]
async fn test_interrupt_operation_request_structure() {
let builder = ChannelBuilder::parse("sc://localhost").unwrap();
let client = SparkConnectClient::connect(&builder).await.unwrap();
let _result = client.interrupt_operation("operation-123");
}
#[tokio::test]
async fn test_release_session_request_structure() {
let builder = ChannelBuilder::parse("sc://localhost").unwrap();
let client = SparkConnectClient::connect(&builder).await.unwrap();
let _result = client.release_session();
}
#[tokio::test]
async fn test_add_artifacts_request_structure() {
let builder = ChannelBuilder::parse("sc://localhost").unwrap();
let client = SparkConnectClient::connect(&builder).await.unwrap();
let _result = client.add_artifacts(&[], false, false, false).await;
assert!(_result.is_ok());
}
#[tokio::test]
async fn test_get_config_with_defaults_request_structure() {
let builder = ChannelBuilder::parse("sc://localhost").unwrap();
let client = SparkConnectClient::connect(&builder).await.unwrap();
let pairs = vec![
("spark.sql.shuffle.partitions", Some("200")),
("spark.sql.adaptive.enabled", None),
];
let _result = client.get_config_with_defaults(&pairs);
}
}