use std::collections::HashMap;
use spark_connect_core::client::ReattachableResponseStream;
use spark_connect_core::error::{Result, SparkError};
use spark_connect_core::runtime::{block_on, get_runtime};
use spark_connect_proto as proto;
use crate::column::Column;
use crate::expression::Expression;
use crate::plan::{AggregateGroupType, JoinType, LogicalPlan, SetOpType};
use crate::row::{Row, Value};
use crate::session::{ExecutionInfo, SparkSession};
use crate::types::DataType;
use crate::udf::CommonInlineUserDefinedFunctionExpression;
#[derive(Clone)]
pub struct DataFrame {
pub(crate) session: SparkSession,
pub(crate) plan: LogicalPlan,
}
pub struct LocalRowIterator {
current_rows: std::vec::IntoIter<Row>,
source: RowSource,
done: bool,
}
enum RowSource {
OnDemand {
session: SparkSession,
stream: ReattachableResponseStream,
execution_info: ExecutionInfo,
execution_recorded: bool,
},
Prefetch {
rx: tokio::sync::mpsc::Receiver<Result<Vec<Row>>>,
},
}
impl LocalRowIterator {
pub(crate) fn new(
session: SparkSession,
stream: ReattachableResponseStream,
prefetch_partitions: bool,
) -> Self {
let source = if prefetch_partitions {
RowSource::Prefetch {
rx: spawn_prefetch(session, stream),
}
} else {
RowSource::OnDemand {
session,
stream,
execution_info: ExecutionInfo::default(),
execution_recorded: false,
}
};
LocalRowIterator {
current_rows: vec![].into_iter(),
source,
done: false,
}
}
fn fetch_next_batch(&mut self) -> Option<Result<Vec<Row>>> {
match &mut self.source {
RowSource::OnDemand {
session,
stream,
execution_info,
execution_recorded,
} => loop {
match block_on(stream.message()) {
Ok(Some(mut resp)) => {
capture_execution(&mut resp, execution_info, session);
if let Some(proto::execute_plan_response::ResponseType::ArrowBatch(batch)) =
resp.response_type
{
return Some(decode_arrow_batch(&batch));
}
}
Ok(None) => {
if !*execution_recorded {
session.record_execution(execution_info.clone());
*execution_recorded = true;
}
return None;
}
Err(e) => {
if !*execution_recorded {
session.record_execution(execution_info.clone());
*execution_recorded = true;
}
return Some(Err(e));
}
}
},
RowSource::Prefetch { rx } => block_on(rx.recv()),
}
}
}
impl Iterator for LocalRowIterator {
type Item = Result<Row>;
fn next(&mut self) -> Option<Self::Item> {
loop {
if let Some(row) = self.current_rows.next() {
return Some(Ok(row));
}
if self.done {
return None;
}
match self.fetch_next_batch() {
Some(Ok(rows)) => self.current_rows = rows.into_iter(),
Some(Err(e)) => {
self.done = true;
return Some(Err(e));
}
None => {
self.done = true;
return None;
}
}
}
}
}
fn spawn_prefetch(
session: SparkSession,
mut stream: ReattachableResponseStream,
) -> tokio::sync::mpsc::Receiver<Result<Vec<Row>>> {
let (tx, rx) = tokio::sync::mpsc::channel::<Result<Vec<Row>>>(1);
get_runtime().spawn(async move {
let mut execution_info = ExecutionInfo::default();
loop {
match stream.message().await {
Ok(Some(mut resp)) => {
capture_execution(&mut resp, &mut execution_info, &session);
if let Some(proto::execute_plan_response::ResponseType::ArrowBatch(batch)) =
resp.response_type
{
match decode_arrow_batch(&batch) {
Ok(rows) => {
if tx.send(Ok(rows)).await.is_err() {
break;
}
}
Err(e) => {
let _ = tx.send(Err(e)).await;
break;
}
}
}
}
Ok(None) => break,
Err(e) => {
let _ = tx.send(Err(e)).await;
break;
}
}
}
session.record_execution(execution_info);
});
rx
}
impl DataFrame {
pub(crate) fn new(session: SparkSession, plan: LogicalPlan) -> Self {
DataFrame { session, plan }
}
pub(crate) fn plan(&self) -> &LogicalPlan {
&self.plan
}
pub fn select<C: Into<Column>>(&self, columns: impl IntoIterator<Item = C>) -> DataFrame {
let columns: Vec<Column> = columns.into_iter().map(Into::into).collect();
let plan = LogicalPlan::Project {
input: Box::new(self.plan.clone()),
columns,
};
DataFrame::new(self.session.clone(), plan)
}
pub fn filter(&self, condition: Column) -> DataFrame {
let plan = LogicalPlan::Filter {
input: Box::new(self.plan.clone()),
condition,
};
DataFrame::new(self.session.clone(), plan)
}
pub fn where_(&self, condition: Column) -> DataFrame {
self.filter(condition)
}
pub fn with_column(&self, name: &str, col: Column) -> DataFrame {
let plan = LogicalPlan::WithColumns {
input: Box::new(self.plan.clone()),
column_names: vec![name.to_string()],
columns: vec![col],
};
DataFrame::new(self.session.clone(), plan)
}
pub fn with_columns(&self, columns: Vec<(String, Column)>) -> DataFrame {
let (names, cols) = columns.into_iter().unzip();
let plan = LogicalPlan::WithColumns {
input: Box::new(self.plan.clone()),
column_names: names,
columns: cols,
};
DataFrame::new(self.session.clone(), plan)
}
pub fn with_column_renamed(&self, existing: &str, new: &str) -> DataFrame {
let mut renames = HashMap::new();
renames.insert(existing.to_string(), new.to_string());
let plan = LogicalPlan::WithColumnsRenamed {
input: Box::new(self.plan.clone()),
renames,
};
DataFrame::new(self.session.clone(), plan)
}
pub fn with_columns_renamed(&self, renames: Vec<(String, String)>) -> DataFrame {
let mut rename_map = HashMap::new();
for (old, new) in renames {
rename_map.insert(old, new);
}
let plan = LogicalPlan::WithColumnsRenamed {
input: Box::new(self.plan.clone()),
renames: rename_map,
};
DataFrame::new(self.session.clone(), plan)
}
pub fn drop(&self, columns: Vec<&str>) -> DataFrame {
let col_names = columns.iter().map(|s| s.to_string()).collect();
let plan = LogicalPlan::Drop {
input: Box::new(self.plan.clone()),
columns: col_names,
};
DataFrame::new(self.session.clone(), plan)
}
pub fn limit(&self, n: i32) -> DataFrame {
let plan = LogicalPlan::Limit {
input: Box::new(self.plan.clone()),
limit: n,
};
DataFrame::new(self.session.clone(), plan)
}
pub fn offset(&self, n: i32) -> DataFrame {
let plan = LogicalPlan::Offset {
input: Box::new(self.plan.clone()),
offset: n,
};
DataFrame::new(self.session.clone(), plan)
}
pub fn tail(&self, n: i32) -> DataFrame {
let plan = LogicalPlan::Tail {
input: Box::new(self.plan.clone()),
limit: n,
};
DataFrame::new(self.session.clone(), plan)
}
pub fn distinct(&self) -> DataFrame {
let plan = LogicalPlan::Deduplicate {
input: Box::new(self.plan.clone()),
all_columns_as_keys: true,
column_names: vec![],
within_watermark: false,
};
DataFrame::new(self.session.clone(), plan)
}
pub fn drop_duplicates(&self, column_names: Option<Vec<&str>>) -> DataFrame {
let all_cols = column_names.is_none();
let cols = column_names
.map(|c| c.iter().map(|s| s.to_string()).collect())
.unwrap_or_default();
let plan = LogicalPlan::Deduplicate {
input: Box::new(self.plan.clone()),
all_columns_as_keys: all_cols,
column_names: cols,
within_watermark: false,
};
DataFrame::new(self.session.clone(), plan)
}
pub fn sort(&self, columns: Vec<Expression>) -> DataFrame {
let plan = LogicalPlan::Sort {
input: Box::new(self.plan.clone()),
order: columns,
is_global: true,
};
DataFrame::new(self.session.clone(), plan)
}
pub fn order_by(&self, columns: Vec<Expression>) -> DataFrame {
self.sort(columns)
}
pub fn join(&self, right: &DataFrame, on: Option<Column>, join_type: JoinType) -> DataFrame {
let plan = LogicalPlan::Join {
left: Box::new(self.plan.clone()),
right: Box::new(right.plan.clone()),
join_type,
on,
using_columns: vec![],
};
DataFrame::new(self.session.clone(), plan)
}
pub fn join_using<S: Into<String>>(
&self,
right: &DataFrame,
using_columns: impl IntoIterator<Item = S>,
join_type: JoinType,
) -> DataFrame {
let plan = LogicalPlan::Join {
left: Box::new(self.plan.clone()),
right: Box::new(right.plan.clone()),
join_type,
on: None,
using_columns: using_columns.into_iter().map(Into::into).collect(),
};
DataFrame::new(self.session.clone(), plan)
}
pub fn nearest_by_join(
&self,
other: &DataFrame,
ranking_expression: Column,
num_results: i32,
mode: &str,
direction: &str,
join_type: &str,
) -> DataFrame {
let plan = LogicalPlan::NearestByJoin {
left: Box::new(self.plan.clone()),
right: Box::new(other.plan.clone()),
ranking_expression: ranking_expression.expression().clone(),
num_results,
join_type: join_type.to_string(),
mode: mode.to_string(),
direction: direction.to_string(),
};
DataFrame::new(self.session.clone(), plan)
}
pub fn cross_join(&self, right: &DataFrame) -> DataFrame {
self.join(right, None, JoinType::Cross)
}
pub fn lateral_join(
&self,
right: &DataFrame,
on: Option<Column>,
join_type: JoinType,
) -> DataFrame {
let plan = LogicalPlan::LateralJoin {
left: Box::new(self.plan.clone()),
right: Box::new(right.plan.clone()),
join_type,
on,
};
DataFrame::new(self.session.clone(), plan)
}
pub fn union(&self, other: &DataFrame) -> DataFrame {
let plan = LogicalPlan::SetOperation {
left: Box::new(self.plan.clone()),
right: Box::new(other.plan.clone()),
set_op_type: SetOpType::Union,
is_all: true,
by_name: false,
allow_missing_columns: false,
};
DataFrame::new(self.session.clone(), plan)
}
pub fn union_by_name(&self, other: &DataFrame) -> DataFrame {
self.union_by_name_opt(other, false)
}
pub fn union_by_name_opt(&self, other: &DataFrame, allow_missing_columns: bool) -> DataFrame {
let plan = LogicalPlan::SetOperation {
left: Box::new(self.plan.clone()),
right: Box::new(other.plan.clone()),
set_op_type: SetOpType::Union,
is_all: true,
by_name: true,
allow_missing_columns,
};
DataFrame::new(self.session.clone(), plan)
}
pub fn intersect(&self, other: &DataFrame) -> DataFrame {
let plan = LogicalPlan::SetOperation {
left: Box::new(self.plan.clone()),
right: Box::new(other.plan.clone()),
set_op_type: SetOpType::Intersect,
is_all: false,
by_name: false,
allow_missing_columns: false,
};
DataFrame::new(self.session.clone(), plan)
}
pub fn subtract(&self, other: &DataFrame) -> DataFrame {
let plan = LogicalPlan::SetOperation {
left: Box::new(self.plan.clone()),
right: Box::new(other.plan.clone()),
set_op_type: SetOpType::Except,
is_all: false,
by_name: false,
allow_missing_columns: false,
};
DataFrame::new(self.session.clone(), plan)
}
pub fn repartition(&self, num_partitions: i32) -> DataFrame {
let plan = LogicalPlan::Repartition {
input: Box::new(self.plan.clone()),
num_partitions,
shuffle: true,
};
DataFrame::new(self.session.clone(), plan)
}
pub fn coalesce(&self, num_partitions: i32) -> DataFrame {
let plan = LogicalPlan::Repartition {
input: Box::new(self.plan.clone()),
num_partitions,
shuffle: false,
};
DataFrame::new(self.session.clone(), plan)
}
pub fn hint<S: Into<String>>(
&self,
name: &str,
parameters: impl IntoIterator<Item = S>,
) -> DataFrame {
let plan = LogicalPlan::Hint {
input: Box::new(self.plan.clone()),
name: name.to_string(),
parameters: parameters.into_iter().map(Into::into).collect(),
};
DataFrame::new(self.session.clone(), plan)
}
pub fn broadcast(&self) -> DataFrame {
self.hint("broadcast", Vec::<String>::new())
}
pub fn to_df(&self, column_names: Vec<&str>) -> DataFrame {
let names = column_names.iter().map(|s| s.to_string()).collect();
let plan = LogicalPlan::ToDF {
input: Box::new(self.plan.clone()),
column_names: names,
};
DataFrame::new(self.session.clone(), plan)
}
pub fn alias(&self, alias: &str) -> DataFrame {
let plan = LogicalPlan::SubqueryAlias {
input: Box::new(self.plan.clone()),
alias: alias.to_string(),
};
DataFrame::new(self.session.clone(), plan)
}
pub fn map_in_pandas(
&self,
func: CommonInlineUserDefinedFunctionExpression,
is_barrier: bool,
) -> DataFrame {
self.map_partitions(func, is_barrier)
}
pub fn map_in_arrow(
&self,
func: CommonInlineUserDefinedFunctionExpression,
is_barrier: bool,
) -> DataFrame {
self.map_partitions(func, is_barrier)
}
fn map_partitions(
&self,
func: CommonInlineUserDefinedFunctionExpression,
is_barrier: bool,
) -> DataFrame {
let plan = LogicalPlan::MapPartitions {
input: Box::new(self.plan.clone()),
func,
is_barrier,
};
DataFrame::new(self.session.clone(), plan)
}
pub fn foreach(&self, func: CommonInlineUserDefinedFunctionExpression) -> Result<()> {
let _ = self.map_partitions(func, false).collect()?;
Ok(())
}
pub fn foreach_partition(&self, func: CommonInlineUserDefinedFunctionExpression) -> Result<()> {
let _ = self.map_partitions(func, false).collect()?;
Ok(())
}
pub fn sample(&self, fraction: f64, seed: Option<i64>) -> DataFrame {
self.sample_opt(fraction, false, seed)
}
pub fn sample_opt(
&self,
fraction: f64,
with_replacement: bool,
seed: Option<i64>,
) -> DataFrame {
let plan = LogicalPlan::Sample {
input: Box::new(self.plan.clone()),
lower_bound: 0.0,
upper_bound: fraction,
with_replacement,
seed,
};
DataFrame::new(self.session.clone(), plan)
}
pub fn group_by<C: Into<Column>>(
&self,
group_cols: impl IntoIterator<Item = C>,
) -> crate::group::GroupedData {
let group_cols: Vec<Column> = group_cols.into_iter().map(Into::into).collect();
crate::group::GroupedData::new(self.clone(), group_cols, AggregateGroupType::GroupBy)
}
pub fn collect(&self) -> Result<Vec<Row>> {
let request = self.build_execute_request()?;
let mut stream = block_on(self.session.client().execute_plan_reattachable(request))?;
let mut rows = vec![];
let mut info = ExecutionInfo::default();
loop {
let resp = block_on(stream.message())?;
let Some(mut resp) = resp else {
break;
};
capture_execution(&mut resp, &mut info, &self.session);
if let Some(proto::execute_plan_response::ResponseType::ArrowBatch(batch)) =
resp.response_type
{
let batch_rows = decode_arrow_batch(&batch)?;
rows.extend(batch_rows);
}
}
self.session.record_execution(info);
Ok(rows)
}
pub fn to_local_iterator(&self, prefetch_partitions: bool) -> Result<LocalRowIterator> {
let request = self.build_execute_request()?;
let stream = block_on(self.session.client().execute_plan_reattachable(request))?;
Ok(LocalRowIterator::new(
self.session.clone(),
stream,
prefetch_partitions,
))
}
pub fn execution_info(&self) -> Result<ExecutionInfo> {
self.session.last_execution_info().ok_or_else(|| {
SparkError::connect_msg("no execution info available; run an action first")
})
}
pub fn collect_record_batches(&self) -> Result<Vec<arrow::record_batch::RecordBatch>> {
let request = self.build_execute_request()?;
let mut stream = block_on(self.session.client().execute_plan_reattachable(request))?;
let mut batches = vec![];
let mut info = ExecutionInfo::default();
loop {
let resp = block_on(stream.message())?;
let Some(mut resp) = resp else {
break;
};
capture_execution(&mut resp, &mut info, &self.session);
if let Some(proto::execute_plan_response::ResponseType::ArrowBatch(batch)) =
resp.response_type
{
let record_batches = decode_arrow_record_batches(&batch)?;
batches.extend(record_batches);
}
}
self.session.record_execution(info);
Ok(batches)
}
pub fn count(&self) -> Result<i64> {
let count_expr = crate::functions::count(Column::new(Expression::Literal(
crate::expression::LiteralExpression::int(1),
)))
.expression()
.clone();
let plan = LogicalPlan::Aggregate {
input: Box::new(self.plan.clone()),
group_type: AggregateGroupType::GroupBy,
grouping_expressions: vec![],
aggregate_expressions: vec![count_expr],
pivot_col: None,
pivot_values: vec![],
grouping_sets: vec![],
};
let agg_df = DataFrame::new(self.session.clone(), plan);
let rows = agg_df.collect()?;
match rows.into_iter().next() {
Some(row) => row.get(0).and_then(|v| v.as_i64()).ok_or_else(|| {
SparkError::connect_msg("count() aggregate returned a non-integer value")
}),
None => Ok(0),
}
}
pub fn show(&self, n: usize) -> Result<()> {
let limited = self.limit(n as i32).collect()?;
for row in limited {
println!("{}", row);
}
Ok(())
}
pub fn schema(&self) -> Result<DataType> {
let request = self.build_analyze_request()?;
let response = block_on(self.session.client().analyze_plan(request))?;
if let Some(proto::analyze_plan_response::Result::Schema(schema)) = response.result {
Ok(DataType::from_proto(&schema.schema.ok_or_else(|| {
SparkError::connect_msg("Schema is missing")
})?)?)
} else {
Err(SparkError::connect_msg(
"Schema analyze failed: no schema in response",
))
}
}
pub fn first(&self) -> Result<Option<Row>> {
self.limit(1).collect().map(|rows| rows.into_iter().next())
}
pub fn head(&self) -> Result<Option<Row>> {
self.first()
}
pub fn take(&self, n: usize) -> Result<Vec<Row>> {
self.limit(n as i32).collect()
}
pub fn is_empty(&self) -> Result<bool> {
self.limit(1).count().map(|c| c == 0)
}
pub fn columns(&self) -> Result<Vec<String>> {
let schema = self.schema()?;
match schema {
DataType::Struct { fields } => Ok(fields.iter().map(|f| f.name.clone()).collect()),
_ => Err(SparkError::connect_msg("Schema is not a struct type")),
}
}
fn build_execute_request(&self) -> Result<proto::ExecutePlanRequest> {
let mut relation = self.plan.to_proto();
assign_plan_ids(&mut relation, &self.session)?;
let mut request = proto::ExecutePlanRequest::default();
request.session_id = self.session.client().session_id().to_string();
request.user_context = Some(proto::UserContext::default());
request.tags = self.session.tags();
let mut plan = proto::Plan::default();
plan.op_type = Some(proto::plan::OpType::Root(relation));
request.plan = Some(plan);
Ok(request)
}
fn build_analyze_request(&self) -> Result<proto::AnalyzePlanRequest> {
let mut relation = self.plan.to_proto();
assign_plan_ids(&mut relation, &self.session)?;
let mut plan = proto::Plan::default();
plan.op_type = Some(proto::plan::OpType::Root(relation));
let mut schema = proto::analyze_plan_request::Schema::default();
schema.plan = Some(plan);
let mut request = proto::AnalyzePlanRequest::default();
request.session_id = self.session.client().session_id().to_string();
request.user_context = Some(proto::UserContext::default());
request.analyze = Some(proto::analyze_plan_request::Analyze::Schema(schema));
Ok(request)
}
pub fn write(&self) -> crate::readwriter::DataFrameWriter {
crate::readwriter::DataFrameWriter::new(self.session.clone(), self.plan.clone())
}
pub fn write_to(&self, table_name: &str) -> crate::readwriter::DataFrameWriterV2 {
crate::readwriter::DataFrameWriterV2::new(
self.session.clone(),
self.plan.clone(),
table_name,
)
}
pub fn merge_into(&self, table: &str, condition: Column) -> crate::merge::MergeIntoWriter {
crate::merge::MergeIntoWriter::new(
self.session.clone(),
self.plan.clone(),
table.to_string(),
condition,
)
}
pub fn write_stream(&self) -> crate::streaming::DataStreamWriter {
crate::streaming::DataStreamWriter::new(self.session.clone(), self.plan.clone())
}
fn memory_and_disk_deser() -> proto::StorageLevel {
proto::StorageLevel {
use_disk: true,
use_memory: true,
use_off_heap: false,
deserialized: true,
replication: 1,
}
}
fn analyze_relation(&self) -> Result<proto::Relation> {
let mut relation = self.plan.to_proto();
assign_plan_ids(&mut relation, &self.session)?;
Ok(relation)
}
fn analyze_request(
&self,
analyze: proto::analyze_plan_request::Analyze,
) -> proto::AnalyzePlanRequest {
proto::AnalyzePlanRequest {
session_id: self.session.client().session_id().to_string(),
user_context: Some(proto::UserContext::default()),
analyze: Some(analyze),
..Default::default()
}
}
pub fn cache(&self) -> Result<DataFrame> {
self.persist(Self::memory_and_disk_deser())
}
pub fn persist(&self, storage_level: proto::StorageLevel) -> Result<DataFrame> {
let persist = proto::analyze_plan_request::Persist {
relation: Some(self.analyze_relation()?),
storage_level: Some(storage_level),
};
let request = self.analyze_request(proto::analyze_plan_request::Analyze::Persist(persist));
block_on(self.session.client().analyze_plan(request))?;
Ok(self.clone())
}
pub fn unpersist(&self, blocking: bool) -> Result<DataFrame> {
let unpersist = proto::analyze_plan_request::Unpersist {
relation: Some(self.analyze_relation()?),
blocking: Some(blocking),
};
let request =
self.analyze_request(proto::analyze_plan_request::Analyze::Unpersist(unpersist));
block_on(self.session.client().analyze_plan(request))?;
Ok(self.clone())
}
pub fn checkpoint(&self) -> Result<DataFrame> {
self.checkpoint_impl(false, true)
}
pub fn local_checkpoint(&self) -> Result<DataFrame> {
self.checkpoint_impl(true, true)
}
fn checkpoint_impl(&self, local: bool, eager: bool) -> Result<DataFrame> {
let mut cmd = proto::CheckpointCommand::default();
cmd.relation = Some(self.plan.to_proto());
cmd.local = local;
cmd.eager = eager;
let responses = execute_command_collect(
&self.session,
proto::command::CommandType::CheckpointCommand(cmd),
)?;
for resp in &responses {
if let Some(proto::execute_plan_response::ResponseType::CheckpointCommandResult(res)) =
&resp.response_type
{
if let Some(rel) = &res.relation {
return Ok(DataFrame::new(
self.session.clone(),
LogicalPlan::CachedRemoteRelation {
relation_id: rel.relation_id.clone(),
},
));
}
}
}
Err(SparkError::connect_msg(
"checkpoint: server returned no CheckpointCommandResult",
))
}
pub fn create_temp_view(&self, name: &str) -> Result<()> {
self.create_view(name, false, false)
}
pub fn create_or_replace_temp_view(&self, name: &str) -> Result<()> {
self.create_view(name, true, false)
}
pub fn create_global_temp_view(&self, name: &str) -> Result<()> {
self.create_view(name, false, true)
}
pub fn create_or_replace_global_temp_view(&self, name: &str) -> Result<()> {
self.create_view(name, true, true)
}
fn create_view(&self, name: &str, replace: bool, global: bool) -> Result<()> {
let mut input = self.plan.to_proto();
assign_plan_ids(&mut input, &self.session)?;
let mut cmd = proto::CreateDataFrameViewCommand::default();
cmd.input = Some(input);
cmd.name = name.to_string();
cmd.is_global = global;
cmd.replace = replace;
execute_command(
&self.session,
proto::command::CommandType::CreateDataframeView(cmd),
)
}
pub fn explain(&self) -> Result<()> {
self.explain_mode("simple")
}
pub fn explain_mode(&self, mode: &str) -> Result<()> {
use proto::analyze_plan_request::explain::ExplainMode;
let explain_mode = match mode.to_lowercase().as_str() {
"simple" => ExplainMode::Simple,
"extended" => ExplainMode::Extended,
"codegen" => ExplainMode::Codegen,
"cost" => ExplainMode::Cost,
"formatted" => ExplainMode::Formatted,
other => {
return Err(SparkError::value(
"UNSUPPORTED_EXPLAIN_MODE",
&[("mode", other)],
))
}
};
let mut relation = self.plan.to_proto();
assign_plan_ids(&mut relation, &self.session)?;
let mut plan = proto::Plan::default();
plan.op_type = Some(proto::plan::OpType::Root(relation));
let mut ex = proto::analyze_plan_request::Explain::default();
ex.plan = Some(plan);
ex.explain_mode = explain_mode as i32;
let mut request = proto::AnalyzePlanRequest::default();
request.session_id = self.session.client().session_id().to_string();
request.user_context = Some(proto::UserContext::default());
request.analyze = Some(proto::analyze_plan_request::Analyze::Explain(ex));
let response = block_on(self.session.client().analyze_plan(request))?;
if let Some(proto::analyze_plan_response::Result::Explain(e)) = response.result {
println!("{}", e.explain_string);
}
Ok(())
}
pub fn with_watermark(&self, time_column: &str, delay_threshold: &str) -> DataFrame {
let plan = LogicalPlan::WithWatermark {
input: Box::new(self.plan.clone()),
time_column: time_column.to_string(),
delay_threshold: delay_threshold.to_string(),
};
DataFrame::new(self.session.clone(), plan)
}
pub fn repartition_by_range(&self, num_partitions: i32, columns: Vec<Expression>) -> DataFrame {
let plan = LogicalPlan::RepartitionByRange {
input: Box::new(self.plan.clone()),
num_partitions: Some(num_partitions),
partition_exprs: columns,
};
DataFrame::new(self.session.clone(), plan)
}
pub fn repartition_by_expressions(
&self,
num_partitions: i32,
columns: Vec<Expression>,
) -> DataFrame {
let plan = LogicalPlan::RepartitionByExpression {
input: Box::new(self.plan.clone()),
num_partitions,
expressions: columns,
};
DataFrame::new(self.session.clone(), plan)
}
pub fn to_schema(&self, column_names: Vec<&str>) -> DataFrame {
self.to_df(column_names)
}
pub fn melt(
&self,
id_vars: Vec<&str>,
value_vars: Option<Vec<&str>>,
var_name: &str,
value_name: &str,
) -> DataFrame {
use crate::column::col;
let ids: Vec<Column> = id_vars.iter().map(|name| col(name)).collect();
let vals: Option<Vec<Column>> =
value_vars.map(|v| v.iter().map(|name| col(name)).collect());
let plan = LogicalPlan::Unpivot {
input: Box::new(self.plan.clone()),
ids,
values: vals,
variable_column_name: var_name.to_string(),
value_column_name: value_name.to_string(),
};
DataFrame::new(self.session.clone(), plan)
}
pub fn input_files(&self) -> Result<Vec<String>> {
let mut relation = self.plan.to_proto();
assign_plan_ids(&mut relation, &self.session)?;
let mut plan = proto::Plan::default();
plan.op_type = Some(proto::plan::OpType::Root(relation));
let mut inp = proto::analyze_plan_request::InputFiles::default();
inp.plan = Some(plan);
let mut request = proto::AnalyzePlanRequest::default();
request.session_id = self.session.client().session_id().to_string();
request.user_context = Some(proto::UserContext::default());
request.analyze = Some(proto::analyze_plan_request::Analyze::InputFiles(inp));
let response = block_on(self.session.client().analyze_plan(request))?;
match response.result {
Some(proto::analyze_plan_response::Result::InputFiles(f)) => Ok(f.files),
_ => Ok(vec![]),
}
}
pub fn observe(&self, name: &str, exprs: Vec<Expression>) -> DataFrame {
let plan = LogicalPlan::Observe {
input: Box::new(self.plan.clone()),
name: name.to_string(),
exprs,
};
DataFrame::new(self.session.clone(), plan)
}
pub fn stat(&self) -> crate::group::StatFunctions {
crate::group::StatFunctions::new(self.clone())
}
pub fn na(&self) -> crate::group::NaFunctions {
crate::group::NaFunctions::new(self.clone())
}
pub fn agg(&self, expressions: Vec<Expression>) -> DataFrame {
let plan = LogicalPlan::Aggregate {
input: Box::new(self.plan.clone()),
group_type: AggregateGroupType::GroupBy,
grouping_expressions: vec![],
aggregate_expressions: expressions,
pivot_col: None,
pivot_values: vec![],
grouping_sets: vec![],
};
DataFrame::new(self.session.clone(), plan)
}
pub fn select_expr(&self, exprs: Vec<&str>) -> DataFrame {
let cols: Vec<Column> = exprs.iter().map(|e| crate::functions::expr(e)).collect();
self.select(cols)
}
pub fn fillna(&self, value: i64, subset: Option<Vec<&str>>) -> DataFrame {
self.fillna_value(crate::row::Value::Long(value), subset)
}
pub fn fillna_double(&self, value: f64, subset: Option<Vec<&str>>) -> DataFrame {
self.fillna_value(crate::row::Value::Double(value), subset)
}
pub fn fillna_string(&self, value: &str, subset: Option<Vec<&str>>) -> DataFrame {
self.fillna_value(crate::row::Value::String(value.to_string()), subset)
}
pub fn fillna_bool(&self, value: bool, subset: Option<Vec<&str>>) -> DataFrame {
self.fillna_value(crate::row::Value::Bool(value), subset)
}
pub fn fillna_value(&self, value: crate::row::Value, subset: Option<Vec<&str>>) -> DataFrame {
let columns = subset
.map(|v| v.iter().map(|s| s.to_string()).collect())
.unwrap_or_default();
let plan = LogicalPlan::NAFill {
input: Box::new(self.plan.clone()),
fill_value: value,
columns,
};
DataFrame::new(self.session.clone(), plan)
}
pub fn fillna_map(&self, pairs: Vec<(String, crate::row::Value)>) -> DataFrame {
let (cols, values): (Vec<String>, Vec<crate::row::Value>) = pairs.into_iter().unzip();
let plan = LogicalPlan::NAFillColumns {
input: Box::new(self.plan.clone()),
cols,
values,
};
DataFrame::new(self.session.clone(), plan)
}
pub fn dropna(
&self,
how: Option<&str>,
thresh: Option<i32>,
subset: Option<Vec<&str>>,
) -> DataFrame {
let how_str = how.unwrap_or("any").to_string();
let columns = subset
.map(|v| v.iter().map(|s| s.to_string()).collect())
.unwrap_or_default();
let plan = LogicalPlan::NADrop {
input: Box::new(self.plan.clone()),
how: how_str,
min_non_null: thresh,
columns,
};
DataFrame::new(self.session.clone(), plan)
}
pub fn replace(
&self,
to_replace: Vec<(String, String)>,
subset: Option<Vec<&str>>,
) -> DataFrame {
let columns = subset
.map(|v| v.iter().map(|s| s.to_string()).collect())
.unwrap_or_default();
let plan = LogicalPlan::NAReplace {
input: Box::new(self.plan.clone()),
replacements: to_replace,
columns,
};
DataFrame::new(self.session.clone(), plan)
}
pub fn describe(&self, columns: Vec<&str>) -> DataFrame {
let col_names = columns.iter().map(|s| s.to_string()).collect();
let plan = LogicalPlan::Describe {
input: Box::new(self.plan.clone()),
columns: col_names,
};
DataFrame::new(self.session.clone(), plan)
}
pub fn summary(&self, percentiles: Vec<&str>) -> DataFrame {
let percs = percentiles.iter().map(|s| s.to_string()).collect();
let plan = LogicalPlan::Summary {
input: Box::new(self.plan.clone()),
percentiles: percs,
};
DataFrame::new(self.session.clone(), plan)
}
pub fn col_regex(&self, col_name: &str) -> DataFrame {
let plan = LogicalPlan::ColRegex {
input: Box::new(self.plan.clone()),
col_name: col_name.to_string(),
};
DataFrame::new(self.session.clone(), plan)
}
pub fn metadata_column(&self, name: &str) -> Column {
Column::new(Expression::ColumnReference(
crate::expression::ColumnReference::new(name).metadata(),
))
}
pub fn rollup<C: Into<Column>>(
&self,
group_cols: impl IntoIterator<Item = C>,
) -> crate::group::GroupedData {
let group_cols: Vec<Column> = group_cols.into_iter().map(Into::into).collect();
crate::group::GroupedData::new(self.clone(), group_cols, AggregateGroupType::Rollup)
}
pub fn cube<C: Into<Column>>(
&self,
group_cols: impl IntoIterator<Item = C>,
) -> crate::group::GroupedData {
let group_cols: Vec<Column> = group_cols.into_iter().map(Into::into).collect();
crate::group::GroupedData::new(self.clone(), group_cols, AggregateGroupType::Cube)
}
pub fn grouping_sets(&self, group_cols: Vec<Vec<Column>>) -> crate::group::GroupedData {
crate::group::GroupedData::new_grouping_sets(self.clone(), group_cols)
}
pub fn sort_within_partitions(&self, columns: Vec<Expression>) -> DataFrame {
let plan = LogicalPlan::Sort {
input: Box::new(self.plan.clone()),
order: columns,
is_global: false,
};
DataFrame::new(self.session.clone(), plan)
}
pub fn drop_duplicates_within_watermark(&self, column_names: Option<Vec<&str>>) -> DataFrame {
let all_cols = column_names.is_none();
let cols = column_names
.map(|c| c.iter().map(|s| s.to_string()).collect())
.unwrap_or_default();
let plan = LogicalPlan::Deduplicate {
input: Box::new(self.plan.clone()),
all_columns_as_keys: all_cols,
column_names: cols,
within_watermark: true,
};
DataFrame::new(self.session.clone(), plan)
}
pub fn transform<F>(&self, f: F) -> DataFrame
where
F: Fn(&DataFrame) -> DataFrame,
{
f(self)
}
pub fn random_split(&self, weights: Vec<f64>, seed: Option<i64>) -> Vec<DataFrame> {
let total: f64 = weights.iter().sum();
let normalized: Vec<f64> = weights.iter().map(|w| w / total).collect();
let mut results = vec![];
let mut cumulative = 0.0;
for weight in normalized {
let upper = cumulative + weight;
let plan = LogicalPlan::Sample {
input: Box::new(self.plan.clone()),
lower_bound: cumulative,
upper_bound: upper,
with_replacement: false,
seed,
};
results.push(DataFrame::new(self.session.clone(), plan));
cumulative = upper;
}
results
}
pub fn print_schema(&self) -> Result<()> {
let schema = self.schema()?;
println!("{}", schema);
Ok(())
}
pub fn storage_level(&self) -> Result<proto::StorageLevel> {
let get = proto::analyze_plan_request::GetStorageLevel {
relation: Some(self.analyze_relation()?),
};
let request =
self.analyze_request(proto::analyze_plan_request::Analyze::GetStorageLevel(get));
let response = block_on(self.session.client().analyze_plan(request))?;
match response.result {
Some(proto::analyze_plan_response::Result::GetStorageLevel(g)) => {
Ok(g.storage_level.unwrap_or_default())
}
_ => Ok(proto::StorageLevel::default()),
}
}
pub fn is_cached(&self) -> Result<bool> {
let level = self.storage_level()?;
Ok(level.use_memory || level.use_disk)
}
pub fn dtypes(&self) -> Result<Vec<(String, String)>> {
let schema = self.schema()?;
match schema {
DataType::Struct { fields } => {
let dtypes = fields
.iter()
.map(|f| (f.name.clone(), f.data_type.to_string()))
.collect();
Ok(dtypes)
}
_ => Err(SparkError::connect_msg("Schema is not a struct type")),
}
}
pub fn semantic_hash(&self) -> Result<i32> {
let mut relation = self.plan.to_proto();
assign_plan_ids(&mut relation, &self.session)?;
let mut plan = proto::Plan::default();
plan.op_type = Some(proto::plan::OpType::Root(relation));
let mut request = proto::AnalyzePlanRequest::default();
request.session_id = self.session.client().session_id().to_string();
request.user_context = Some(proto::UserContext::default());
request.analyze = Some(proto::analyze_plan_request::Analyze::SemanticHash(
proto::analyze_plan_request::SemanticHash { plan: Some(plan) },
));
let resp = block_on(self.session.client().analyze_plan(request))?;
match resp.result {
Some(proto::analyze_plan_response::Result::SemanticHash(h)) => Ok(h.result),
_ => Err(SparkError::connect_msg(
"AnalyzePlan response did not contain a semantic hash",
)),
}
}
pub fn same_semantics(&self, other: &DataFrame) -> Result<bool> {
let mut self_rel = self.plan.to_proto();
assign_plan_ids(&mut self_rel, &self.session)?;
let mut other_rel = other.plan.to_proto();
assign_plan_ids(&mut other_rel, &other.session)?;
let mut target_plan = proto::Plan::default();
target_plan.op_type = Some(proto::plan::OpType::Root(self_rel));
let mut other_plan = proto::Plan::default();
other_plan.op_type = Some(proto::plan::OpType::Root(other_rel));
let mut request = proto::AnalyzePlanRequest::default();
request.session_id = self.session.client().session_id().to_string();
request.user_context = Some(proto::UserContext::default());
request.analyze = Some(proto::analyze_plan_request::Analyze::SameSemantics(
proto::analyze_plan_request::SameSemantics {
target_plan: Some(target_plan),
other_plan: Some(other_plan),
},
));
let resp = block_on(self.session.client().analyze_plan(request))?;
match resp.result {
Some(proto::analyze_plan_response::Result::SameSemantics(r)) => Ok(r.result),
_ => Err(SparkError::connect_msg(
"AnalyzePlan response did not contain a sameSemantics result",
)),
}
}
pub fn to_json(&self) -> Result<Vec<String>> {
let cols: Vec<Column> = self
.columns()?
.iter()
.map(|c| crate::column::col(c))
.collect();
let json_col = crate::functions::to_json(crate::functions::r#struct(cols));
let rows = self.select(vec![json_col]).collect()?;
Ok(rows
.iter()
.map(|r| {
r.get(0)
.and_then(|v| v.as_str())
.unwrap_or_default()
.to_string()
})
.collect())
}
pub fn union_all(&self, other: &DataFrame) -> DataFrame {
let plan = LogicalPlan::SetOperation {
left: Box::new(self.plan.clone()),
right: Box::new(other.plan.clone()),
set_op_type: SetOpType::Union,
is_all: true,
by_name: false,
allow_missing_columns: false,
};
DataFrame::new(self.session.clone(), plan)
}
pub fn except_all(&self, other: &DataFrame) -> DataFrame {
let plan = LogicalPlan::SetOperation {
left: Box::new(self.plan.clone()),
right: Box::new(other.plan.clone()),
set_op_type: SetOpType::Except,
is_all: true,
by_name: false,
allow_missing_columns: false,
};
DataFrame::new(self.session.clone(), plan)
}
pub fn intersect_all(&self, other: &DataFrame) -> DataFrame {
let plan = LogicalPlan::SetOperation {
left: Box::new(self.plan.clone()),
right: Box::new(other.plan.clone()),
set_op_type: SetOpType::Intersect,
is_all: true,
by_name: false,
allow_missing_columns: false,
};
DataFrame::new(self.session.clone(), plan)
}
pub fn unpivot<C: Into<Column>, D: Into<Column>>(
&self,
ids: impl IntoIterator<Item = C>,
values: Option<impl IntoIterator<Item = D>>,
variable_column_name: &str,
value_column_name: &str,
) -> DataFrame {
let ids: Vec<Column> = ids.into_iter().map(Into::into).collect();
let values: Option<Vec<Column>> = values.map(|v| v.into_iter().map(Into::into).collect());
let plan = LogicalPlan::Unpivot {
input: Box::new(self.plan.clone()),
ids,
values,
variable_column_name: variable_column_name.to_string(),
value_column_name: value_column_name.to_string(),
};
DataFrame::new(self.session.clone(), plan)
}
pub fn with_metadata(&self, column_name: &str, metadata: HashMap<String, String>) -> DataFrame {
let metadata_json = serde_json::to_string(&metadata).unwrap_or_else(|_| "{}".to_string());
let plan = LogicalPlan::WithColumnMetadata {
input: Box::new(self.plan.clone()),
column_name: column_name.to_string(),
metadata_json,
};
DataFrame::new(self.session.clone(), plan)
}
pub fn spark_session(&self) -> SparkSession {
self.session.clone()
}
pub fn is_local(&self) -> bool {
matches!(self.plan, LogicalPlan::LocalRelation { .. })
}
pub fn is_streaming(&self) -> bool {
matches!(
self.plan,
LogicalPlan::Read {
is_streaming: true,
..
}
)
}
pub fn to_arrow(&self) -> Result<Vec<u8>> {
record_batches_to_ipc(&self.collect_record_batches()?)
}
#[cfg(feature = "datafusion")]
pub fn to_datafusion(
&self,
ctx: &datafusion::prelude::SessionContext,
) -> Result<datafusion::dataframe::DataFrame> {
record_batches_to_datafusion(ctx, self.collect_record_batches()?)
}
#[cfg(feature = "polars")]
pub fn to_polars(&self) -> Result<polars::frame::DataFrame> {
record_batches_to_polars(&self.collect_record_batches()?)
}
pub fn repartition_by_id(&self, num_partitions: i32, partition_id_col: Column) -> DataFrame {
let direct =
Expression::DirectShufflePartitionId(Box::new(partition_id_col.expression().clone()));
self.repartition_by_expressions(num_partitions, vec![direct])
}
pub fn zip_with_index(&self, index_col_name: &str) -> DataFrame {
let star = Column::new(Expression::UnresolvedStar(None));
let seq = Column::new(Expression::UnresolvedFunction(
crate::expression::UnresolvedFunction::new("distributed_sequence_id", vec![]),
))
.alias(index_col_name);
self.select(vec![star, seq])
}
pub fn to(&self, schema: DataType) -> DataFrame {
let plan = LogicalPlan::ToSchema {
input: Box::new(self.plan.clone()),
schema,
};
DataFrame::new(self.session.clone(), plan)
}
pub fn exists(&self) -> Result<bool> {
self.limit(1).count().map(|c| c > 0)
}
pub fn scalar(&self) -> Result<Option<Value>> {
let rows = self.limit(1).collect()?;
if rows.is_empty() {
return Ok(None);
}
let row = &rows[0];
Ok(row.get(0).cloned())
}
pub fn transpose(&self) -> Result<DataFrame> {
let plan = LogicalPlan::Transpose {
input: Box::new(self.plan.clone()),
index_columns: vec![],
};
Ok(DataFrame::new(self.session.clone(), plan))
}
pub fn transpose_with_index(&self, index_column: Column) -> Result<DataFrame> {
let plan = LogicalPlan::Transpose {
input: Box::new(self.plan.clone()),
index_columns: vec![index_column.expression().clone()],
};
Ok(DataFrame::new(self.session.clone(), plan))
}
pub fn zip(&self, other: &DataFrame) -> Result<DataFrame> {
let plan = LogicalPlan::Zip {
left: Box::new(self.plan.clone()),
right: Box::new(other.plan.clone()),
};
Ok(DataFrame {
plan,
session: self.session.clone(),
})
}
pub fn register_temp_table(&self, name: &str) -> Result<()> {
self.create_temp_view(name)?;
Ok(())
}
pub fn as_table(&self, alias: &str) -> DataFrame {
self.alias(alias)
}
}
pub(crate) fn build_input_relation(
plan: &LogicalPlan,
session: &SparkSession,
) -> Result<proto::Relation> {
let mut relation = plan.to_proto();
assign_plan_ids(&mut relation, session)?;
Ok(relation)
}
pub(crate) fn execute_command(
session: &SparkSession,
command_type: proto::command::CommandType,
) -> Result<()> {
execute_command_collect(session, command_type).map(|_| ())
}
pub(crate) fn execute_command_collect(
session: &SparkSession,
command_type: proto::command::CommandType,
) -> Result<Vec<proto::ExecutePlanResponse>> {
let mut command = proto::Command::default();
command.command_type = Some(command_type);
let mut plan = proto::Plan::default();
plan.op_type = Some(proto::plan::OpType::Command(command));
let mut request = proto::ExecutePlanRequest::default();
request.session_id = session.client().session_id().to_string();
request.user_context = Some(proto::UserContext::default());
request.tags = session.tags();
request.plan = Some(plan);
let mut stream = block_on(session.client().execute_plan_reattachable(request))?;
let mut info = ExecutionInfo::default();
let mut responses = Vec::new();
while let Some(mut resp) = block_on(stream.message())? {
capture_execution(&mut resp, &mut info, session);
responses.push(resp);
}
session.record_execution(info);
Ok(responses)
}
fn capture_execution(
resp: &mut proto::ExecutePlanResponse,
info: &mut ExecutionInfo,
session: &SparkSession,
) {
if let Some(metrics) = resp.metrics.take() {
info.metrics = Some(metrics);
}
if !resp.observed_metrics.is_empty() {
let metrics = std::mem::take(&mut resp.observed_metrics);
session.profiler().accumulate_observed_metrics(&metrics);
info.observed_metrics.extend(metrics);
}
if let Some(proto::execute_plan_response::ResponseType::ExecutionProgress(progress)) =
&resp.response_type
{
session.notify_progress(progress);
}
}
pub(crate) fn assign_plan_ids(
relation: &mut proto::Relation,
session: &SparkSession,
) -> Result<()> {
if let Some(rel_type) = &mut relation.rel_type {
use proto::relation::RelType;
match rel_type {
RelType::Range(_) => {}
RelType::Sql(_) => {}
RelType::LocalRelation(_) => {}
RelType::CachedRemoteRelation(_) => {}
RelType::Project(proj) => {
if let Some(input) = &mut proj.input {
assign_plan_ids(input, session)?;
}
}
RelType::Filter(filter) => {
if let Some(input) = &mut filter.input {
assign_plan_ids(input, session)?;
}
}
RelType::Join(join) => {
if let Some(left) = &mut join.left {
assign_plan_ids(left, session)?;
}
if let Some(right) = &mut join.right {
assign_plan_ids(right, session)?;
}
}
RelType::SetOp(set_op) => {
if let Some(left) = &mut set_op.left_input {
assign_plan_ids(left, session)?;
}
if let Some(right) = &mut set_op.right_input {
assign_plan_ids(right, session)?;
}
}
RelType::Aggregate(agg) => {
if let Some(input) = &mut agg.input {
assign_plan_ids(input, session)?;
}
}
RelType::Sort(sort) => {
if let Some(input) = &mut sort.input {
assign_plan_ids(input, session)?;
}
}
RelType::Limit(limit) => {
if let Some(input) = &mut limit.input {
assign_plan_ids(input, session)?;
}
}
RelType::Offset(offset) => {
if let Some(input) = &mut offset.input {
assign_plan_ids(input, session)?;
}
}
RelType::Tail(tail) => {
if let Some(input) = &mut tail.input {
assign_plan_ids(input, session)?;
}
}
RelType::Deduplicate(dedup) => {
if let Some(input) = &mut dedup.input {
assign_plan_ids(input, session)?;
}
}
RelType::Repartition(repartition) => {
if let Some(input) = &mut repartition.input {
assign_plan_ids(input, session)?;
}
}
RelType::RepartitionByExpression(repart_expr) => {
if let Some(input) = &mut repart_expr.input {
assign_plan_ids(input, session)?;
}
}
RelType::WithColumns(with_cols) => {
if let Some(input) = &mut with_cols.input {
assign_plan_ids(input, session)?;
}
}
RelType::WithColumnsRenamed(with_renamed) => {
if let Some(input) = &mut with_renamed.input {
assign_plan_ids(input, session)?;
}
}
RelType::Drop(drop) => {
if let Some(input) = &mut drop.input {
assign_plan_ids(input, session)?;
}
}
RelType::ToDf(to_df) => {
if let Some(input) = &mut to_df.input {
assign_plan_ids(input, session)?;
}
}
RelType::ToSchema(to_schema) => {
if let Some(input) = &mut to_schema.input {
assign_plan_ids(input, session)?;
}
}
RelType::Hint(hint) => {
if let Some(input) = &mut hint.input {
assign_plan_ids(input, session)?;
}
}
RelType::Unpivot(unpivot) => {
if let Some(input) = &mut unpivot.input {
assign_plan_ids(input, session)?;
}
}
RelType::Sample(sample) => {
if let Some(input) = &mut sample.input {
assign_plan_ids(input, session)?;
}
}
RelType::FillNa(fill_na) => {
if let Some(input) = &mut fill_na.input {
assign_plan_ids(input, session)?;
}
}
RelType::DropNa(drop_na) => {
if let Some(input) = &mut drop_na.input {
assign_plan_ids(input, session)?;
}
}
RelType::Replace(replace) => {
if let Some(input) = &mut replace.input {
assign_plan_ids(input, session)?;
}
}
RelType::Describe(describe) => {
if let Some(input) = &mut describe.input {
assign_plan_ids(input, session)?;
}
}
RelType::Summary(summary) => {
if let Some(input) = &mut summary.input {
assign_plan_ids(input, session)?;
}
}
RelType::SubqueryAlias(sq_alias) => {
if let Some(input) = &mut sq_alias.input {
assign_plan_ids(input, session)?;
}
}
RelType::CachedLocalRelation(_cached) => {
}
RelType::WithWatermark(watermark) => {
if let Some(input) = &mut watermark.input {
assign_plan_ids(input, session)?;
}
}
RelType::Crosstab(stat) => {
if let Some(input) = &mut stat.input {
assign_plan_ids(input, session)?;
}
}
RelType::FreqItems(stat) => {
if let Some(input) = &mut stat.input {
assign_plan_ids(input, session)?;
}
}
RelType::ApproxQuantile(stat) => {
if let Some(input) = &mut stat.input {
assign_plan_ids(input, session)?;
}
}
RelType::Corr(stat) => {
if let Some(input) = &mut stat.input {
assign_plan_ids(input, session)?;
}
}
RelType::Cov(stat) => {
if let Some(input) = &mut stat.input {
assign_plan_ids(input, session)?;
}
}
RelType::SampleBy(stat) => {
if let Some(input) = &mut stat.input {
assign_plan_ids(input, session)?;
}
}
RelType::CollectMetrics(metrics) => {
if let Some(input) = &mut metrics.input {
assign_plan_ids(input, session)?;
}
}
_ => {
}
}
}
if relation.common.is_none() {
relation.common = Some(proto::RelationCommon::default());
}
if let Some(common) = &mut relation.common {
common.plan_id = Some(session.next_plan_id());
}
Ok(())
}
fn decode_arrow_batch(batch: &proto::execute_plan_response::ArrowBatch) -> Result<Vec<Row>> {
use arrow::ipc::reader::StreamReader;
use std::io::Cursor;
if batch.data.is_empty() {
return Ok(vec![]);
}
let cursor = Cursor::new(&batch.data);
let mut reader = StreamReader::try_new(cursor, None).map_err(|e| {
SparkError::connect_msg(format!("Failed to create Arrow stream reader: {}", e))
})?;
let mut rows = vec![];
while let Some(record_batch) = reader
.next()
.transpose()
.map_err(|e| SparkError::connect_msg(format!("Failed to decode Arrow batch: {}", e)))?
{
let schema = record_batch.schema();
let num_rows = record_batch.num_rows();
let num_cols = record_batch.num_columns();
for row_idx in 0..num_rows {
let mut field_names = vec![];
let mut values = vec![];
for col_idx in 0..num_cols {
let field_name = schema.field(col_idx).name().clone();
let column = record_batch.column(col_idx);
let value = arrow_value_at(column.as_ref(), row_idx)?;
field_names.push(field_name);
values.push(value);
}
rows.push(Row::new(field_names, values));
}
}
Ok(rows)
}
fn decode_arrow_record_batches(
batch: &proto::execute_plan_response::ArrowBatch,
) -> Result<Vec<arrow::record_batch::RecordBatch>> {
use arrow::ipc::reader::StreamReader;
use std::io::Cursor;
if batch.data.is_empty() {
return Ok(vec![]);
}
let cursor = Cursor::new(&batch.data);
let mut reader = StreamReader::try_new(cursor, None).map_err(|e| {
SparkError::connect_msg(format!("Failed to create Arrow stream reader: {}", e))
})?;
let mut batches = vec![];
while let Some(record_batch) = reader
.next()
.transpose()
.map_err(|e| SparkError::connect_msg(format!("Failed to decode Arrow batch: {}", e)))?
{
batches.push(record_batch);
}
Ok(batches)
}
fn i128_to_decimal_string(unscaled: i128, scale: i32) -> String {
if scale <= 0 {
return unscaled.to_string();
}
let scale = scale as usize;
let neg = unscaled < 0;
let mut digits = unscaled.unsigned_abs().to_string();
if digits.len() <= scale {
digits = format!("{}{}", "0".repeat(scale - digits.len() + 1), digits);
}
let point = digits.len() - scale;
let s = format!("{}.{}", &digits[..point], &digits[point..]);
if neg {
format!("-{s}")
} else {
s
}
}
fn record_batches_to_ipc(batches: &[arrow::record_batch::RecordBatch]) -> Result<Vec<u8>> {
use arrow::ipc::writer::FileWriter;
let schema = match batches.first() {
Some(b) => b.schema(),
None => std::sync::Arc::new(arrow::datatypes::Schema::empty()),
};
let mut buf: Vec<u8> = Vec::new();
{
let mut writer = FileWriter::try_new(&mut buf, schema.as_ref())
.map_err(|e| SparkError::connect_msg(format!("Arrow IPC writer init failed: {e}")))?;
for batch in batches {
writer
.write(batch)
.map_err(|e| SparkError::connect_msg(format!("Arrow IPC write failed: {e}")))?;
}
writer
.finish()
.map_err(|e| SparkError::connect_msg(format!("Arrow IPC finish failed: {e}")))?;
}
Ok(buf)
}
#[cfg(feature = "datafusion")]
fn record_batches_to_datafusion(
ctx: &datafusion::prelude::SessionContext,
batches: Vec<arrow::record_batch::RecordBatch>,
) -> Result<datafusion::dataframe::DataFrame> {
if batches.is_empty() {
return Err(SparkError::connect_msg(
"Cannot create DataFusion DataFrame from empty result",
));
}
ctx.read_batches(batches)
.map_err(|e| SparkError::connect_msg(format!("Failed to create DataFusion DataFrame: {e}")))
}
#[cfg(feature = "polars")]
fn record_batches_to_polars(
batches: &[arrow::record_batch::RecordBatch],
) -> Result<polars::frame::DataFrame> {
use polars::prelude::{IpcReader, SerReader};
use std::io::Cursor;
if batches.is_empty() {
return Ok(polars::frame::DataFrame::empty());
}
let buf = record_batches_to_ipc(batches)?;
IpcReader::new(Cursor::new(buf))
.finish()
.map_err(|e| SparkError::connect_msg(format!("Failed to create Polars DataFrame: {e}")))
}
fn map_key_to_string(v: Value) -> String {
match v {
Value::String(s) => s,
Value::Bool(b) => b.to_string(),
Value::Byte(x) => x.to_string(),
Value::Short(x) => x.to_string(),
Value::Integer(x) => x.to_string(),
Value::Long(x) => x.to_string(),
Value::Float(x) => x.to_string(),
Value::Double(x) => x.to_string(),
Value::Date(d) => d.to_string(),
Value::Timestamp(t) => t.to_string(),
Value::Decimal { value, .. } => value,
other => format!("{other:?}"),
}
}
pub(crate) fn arrow_value_at(array: &dyn arrow::array::Array, index: usize) -> Result<Value> {
use arrow::array::*;
if array.is_null(index) {
return Ok(Value::Null);
}
if array.as_any().downcast_ref::<NullArray>().is_some() {
return Ok(Value::Null);
}
if let Some(arr) = array.as_any().downcast_ref::<BooleanArray>() {
return Ok(Value::Bool(arr.value(index)));
}
if let Some(arr) = array.as_any().downcast_ref::<Int8Array>() {
return Ok(Value::Byte(arr.value(index)));
}
if let Some(arr) = array.as_any().downcast_ref::<Int16Array>() {
return Ok(Value::Short(arr.value(index)));
}
if let Some(arr) = array.as_any().downcast_ref::<Int32Array>() {
return Ok(Value::Integer(arr.value(index)));
}
if let Some(arr) = array.as_any().downcast_ref::<Int64Array>() {
return Ok(Value::Long(arr.value(index)));
}
if let Some(arr) = array.as_any().downcast_ref::<Float32Array>() {
return Ok(Value::Float(arr.value(index)));
}
if let Some(arr) = array.as_any().downcast_ref::<Float64Array>() {
return Ok(Value::Double(arr.value(index)));
}
if let Some(arr) = array.as_any().downcast_ref::<StringArray>() {
return Ok(Value::String(arr.value(index).to_string()));
}
if let Some(arr) = array.as_any().downcast_ref::<BinaryArray>() {
return Ok(Value::Binary(arr.value(index).to_vec()));
}
if let Some(arr) = array.as_any().downcast_ref::<Date32Array>() {
return Ok(Value::Date(arr.value(index)));
}
if let Some(arr) = array.as_any().downcast_ref::<TimestampMicrosecondArray>() {
return Ok(Value::Timestamp(arr.value(index)));
}
if let Some(arr) = array.as_any().downcast_ref::<UInt8Array>() {
return Ok(Value::Short(arr.value(index) as i16));
}
if let Some(arr) = array.as_any().downcast_ref::<UInt16Array>() {
return Ok(Value::Integer(arr.value(index) as i32));
}
if let Some(arr) = array.as_any().downcast_ref::<UInt32Array>() {
return Ok(Value::Long(arr.value(index) as i64));
}
if let Some(arr) = array.as_any().downcast_ref::<UInt64Array>() {
let val = arr.value(index);
let i64_val = i64::try_from(val).map_err(|_| {
SparkError::connect_msg(format!("UInt64 value {} exceeds i64 range", val))
})?;
return Ok(Value::Long(i64_val));
}
if let Some(arr) = array.as_any().downcast_ref::<Decimal128Array>() {
let scale = arr.scale() as i32;
return Ok(Value::Decimal {
value: i128_to_decimal_string(arr.value(index), scale),
precision: Some(arr.precision() as i32),
scale: Some(scale),
});
}
if let Some(arr) = array.as_any().downcast_ref::<LargeStringArray>() {
return Ok(Value::String(arr.value(index).to_string()));
}
if let Some(arr) = array.as_any().downcast_ref::<LargeBinaryArray>() {
return Ok(Value::Binary(arr.value(index).to_vec()));
}
if let Some(arr) = array.as_any().downcast_ref::<StringViewArray>() {
return Ok(Value::String(arr.value(index).to_string()));
}
if let Some(arr) = array.as_any().downcast_ref::<BinaryViewArray>() {
return Ok(Value::Binary(arr.value(index).to_vec()));
}
if let Some(arr) = array.as_any().downcast_ref::<TimestampSecondArray>() {
return Ok(Value::Timestamp(arr.value(index) * 1_000_000));
}
if let Some(arr) = array.as_any().downcast_ref::<TimestampMillisecondArray>() {
return Ok(Value::Timestamp(arr.value(index) * 1_000));
}
if let Some(arr) = array.as_any().downcast_ref::<TimestampNanosecondArray>() {
return Ok(Value::Timestamp(arr.value(index) / 1_000));
}
if let Some(arr) = array.as_any().downcast_ref::<Date64Array>() {
return Ok(Value::Date((arr.value(index) / 86_400_000) as i32));
}
if let Some(arr) = array.as_any().downcast_ref::<ListArray>() {
let child = arr.value(index);
let mut items = Vec::with_capacity(child.len());
for i in 0..child.len() {
items.push(arrow_value_at(child.as_ref(), i)?);
}
return Ok(Value::List(items));
}
if let Some(arr) = array.as_any().downcast_ref::<StructArray>() {
let is_variant = arr.fields().iter().any(|f| {
f.metadata()
.get("variant")
.map(|v| v == "true")
.unwrap_or(false)
});
if is_variant {
let bin_field = |name: &str| -> Result<Vec<u8>> {
match arr.column_by_name(name) {
Some(col) => match arrow_value_at(col.as_ref(), index)? {
Value::Binary(b) => Ok(b),
Value::Null => Ok(vec![]),
_ => Err(SparkError::connect_msg("variant field is not binary")),
},
None => Ok(vec![]),
}
};
return Ok(Value::Variant {
value: bin_field("value")?,
metadata: bin_field("metadata")?,
});
}
let mut fields = Vec::new();
for (f, col) in arr.fields().iter().zip(arr.columns()) {
fields.push((f.name().clone(), arrow_value_at(col.as_ref(), index)?));
}
return Ok(Value::Struct(fields));
}
if let Some(arr) = array.as_any().downcast_ref::<MapArray>() {
let entries = arr.value(index);
let keys = entries.column(0);
let vals = entries.column(1);
let mut map = std::collections::BTreeMap::new();
for i in 0..entries.len() {
let k = map_key_to_string(arrow_value_at(keys.as_ref(), i)?);
map.insert(k, arrow_value_at(vals.as_ref(), i)?);
}
return Ok(Value::Map(map));
}
if let Some(arr) = array.as_any().downcast_ref::<Decimal256Array>() {
return Ok(Value::Decimal {
value: arr.value_as_string(index),
precision: Some(arr.precision() as i32),
scale: Some(arr.scale() as i32),
});
}
if let Some(arr) = array.as_any().downcast_ref::<FixedSizeBinaryArray>() {
return Ok(Value::Binary(arr.value(index).to_vec()));
}
if let Some(arr) = array.as_any().downcast_ref::<Time64MicrosecondArray>() {
return Ok(Value::String(micros_to_time_string(arr.value(index))));
}
if let Some(arr) = array.as_any().downcast_ref::<Time64NanosecondArray>() {
return Ok(Value::String(micros_to_time_string(
arr.value(index) / 1_000,
)));
}
if let Some(arr) = array.as_any().downcast_ref::<Time32MillisecondArray>() {
return Ok(Value::String(micros_to_time_string(
arr.value(index) as i64 * 1_000,
)));
}
if let Some(arr) = array.as_any().downcast_ref::<Time32SecondArray>() {
return Ok(Value::String(micros_to_time_string(
arr.value(index) as i64 * 1_000_000,
)));
}
if let Some(arr) = array.as_any().downcast_ref::<IntervalYearMonthArray>() {
let months = arr.value(index);
return Ok(Value::String(format!(
"{}-{}",
months / 12,
(months % 12).abs()
)));
}
if let Some(arr) = array.as_any().downcast_ref::<IntervalDayTimeArray>() {
let v = arr.value(index);
return Ok(Value::String(format!(
"{} days {} ms",
v.days, v.milliseconds
)));
}
if let Some(arr) = array.as_any().downcast_ref::<IntervalMonthDayNanoArray>() {
let v = arr.value(index);
return Ok(Value::String(format!(
"{} months {} days {} ns",
v.months, v.days, v.nanoseconds
)));
}
Err(SparkError::connect_msg(format!(
"Unsupported Arrow type {:?} - cannot convert to Value",
array.data_type()
)))
}
fn micros_to_time_string(micros: i64) -> String {
let total_secs = micros.div_euclid(1_000_000);
let us = micros.rem_euclid(1_000_000);
let (h, m, s) = (total_secs / 3600, (total_secs % 3600) / 60, total_secs % 60);
if us == 0 {
format!("{h:02}:{m:02}:{s:02}")
} else {
format!("{h:02}:{m:02}:{s:02}.{us:06}")
}
}
#[cfg(test)]
mod cache_tests {
use super::*;
use prost::Message;
#[test]
fn cache_default_is_memory_and_disk_deser() {
let sl = DataFrame::memory_and_disk_deser();
assert!(sl.use_memory && sl.use_disk && sl.deserialized);
assert!(!sl.use_off_heap);
assert_eq!(sl.replication, 1);
}
#[test]
fn persist_request_carries_storage_level_over_the_wire() {
let persist = proto::analyze_plan_request::Persist {
relation: None,
storage_level: Some(DataFrame::memory_and_disk_deser()),
};
let decoded =
proto::analyze_plan_request::Persist::decode(persist.encode_to_vec().as_slice())
.unwrap();
let sl = decoded
.storage_level
.expect("storage_level must be present");
assert!(sl.use_memory && sl.use_disk && sl.deserialized && sl.replication == 1);
}
#[test]
fn get_storage_level_response_maps_to_is_cached() {
let cached = proto::StorageLevel {
use_memory: true,
..Default::default()
};
let uncached = proto::StorageLevel::default();
assert!(cached.use_memory || cached.use_disk);
assert!(!(uncached.use_memory || uncached.use_disk));
}
#[test]
fn to_local_iterator_builds_same_plan_as_collect() {
let _iter: LocalRowIterator;
}
}
#[cfg(test)]
mod conversion_tests {
use super::*;
use arrow::array::{Int64Array, StringArray};
use arrow::datatypes::{DataType as ArrowDataType, Field, Schema};
use arrow::record_batch::RecordBatch;
use std::sync::Arc;
fn sample_batch() -> RecordBatch {
let schema = Arc::new(Schema::new(vec![
Field::new("id", ArrowDataType::Int64, false),
Field::new("name", ArrowDataType::Utf8, false),
]));
RecordBatch::try_new(
schema,
vec![
Arc::new(Int64Array::from(vec![1, 2, 3])),
Arc::new(StringArray::from(vec!["a", "b", "c"])),
],
)
.unwrap()
}
#[test]
fn to_arrow_ipc_round_trips() {
use arrow::ipc::reader::FileReader;
use std::io::Cursor;
let ipc = record_batches_to_ipc(&[sample_batch()]).expect("ipc encode");
let reader = FileReader::try_new(Cursor::new(ipc), None).expect("ipc decode");
let batches: Vec<_> = reader.map(|b| b.unwrap()).collect();
let total: usize = batches.iter().map(|b| b.num_rows()).sum();
assert_eq!(total, 3, "round-trip must preserve all rows");
assert_eq!(batches[0].num_columns(), 2);
let ids = batches[0]
.column(0)
.as_any()
.downcast_ref::<Int64Array>()
.unwrap();
assert_eq!(ids.values(), &[1, 2, 3]);
}
#[test]
fn to_arrow_ipc_empty_is_valid() {
use arrow::ipc::reader::FileReader;
use std::io::Cursor;
let ipc = record_batches_to_ipc(&[]).expect("empty ipc");
let reader = FileReader::try_new(Cursor::new(ipc), None).expect("empty ipc decode");
assert_eq!(reader.map(|b| b.unwrap().num_rows()).sum::<usize>(), 0);
}
#[cfg(feature = "datafusion")]
#[test]
fn to_datafusion_preserves_rows_and_columns() {
use datafusion::prelude::SessionContext;
use spark_connect_core::runtime::block_on;
let ctx = SessionContext::new();
let df = record_batches_to_datafusion(&ctx, vec![sample_batch()]).expect("to datafusion");
let collected = block_on(df.collect()).expect("collect datafusion");
assert_eq!(collected.iter().map(|b| b.num_rows()).sum::<usize>(), 3);
assert_eq!(collected[0].num_columns(), 2);
}
#[cfg(feature = "datafusion")]
#[test]
fn to_datafusion_empty_errors() {
use datafusion::prelude::SessionContext;
let ctx = SessionContext::new();
assert!(record_batches_to_datafusion(&ctx, vec![]).is_err());
}
#[cfg(feature = "polars")]
#[test]
fn to_polars_preserves_shape() {
let pdf = record_batches_to_polars(&[sample_batch()]).expect("to polars");
assert_eq!(
pdf.height(),
3,
"row count must survive the Arrow-IPC bridge"
);
assert_eq!(
pdf.width(),
2,
"column count must survive the Arrow-IPC bridge"
);
}
#[cfg(feature = "polars")]
#[test]
fn to_polars_empty_is_empty() {
let pdf = record_batches_to_polars(&[]).expect("empty polars");
assert_eq!(pdf.height(), 0);
}
}
#[cfg(test)]
mod plan_construction_tests {
use super::*;
use crate::session::SparkSession;
fn session() -> SparkSession {
SparkSession::builder()
.remote("sc://localhost:15002")
.get_or_create()
.expect("failed to build session")
}
#[test]
fn with_watermark_plan() {
let spark = session();
let df = spark.range(3).unwrap();
let result = df.with_watermark("timestamp", "1 minute");
match &result.plan {
LogicalPlan::WithWatermark {
time_column,
delay_threshold,
..
} => {
assert_eq!(time_column, "timestamp");
assert_eq!(delay_threshold, "1 minute");
}
_ => panic!("expected WithWatermark plan"),
}
}
#[test]
fn with_metadata_plan() {
let spark = session();
let df = spark.range(3).unwrap();
let mut metadata = std::collections::HashMap::new();
metadata.insert("key".to_string(), "value".to_string());
let result = df.with_metadata("col", metadata);
match &result.plan {
LogicalPlan::WithColumnMetadata {
column_name,
metadata_json,
..
} => {
assert_eq!(column_name, "col");
assert!(!metadata_json.is_empty());
}
_ => panic!("expected WithColumnMetadata plan"),
}
}
#[test]
fn random_split_plan() {
let spark = session();
let df = spark.range(10).unwrap();
let dfs = df.random_split(vec![0.7, 0.3], None);
assert_eq!(dfs.len(), 2);
for split_df in &dfs {
match &split_df.plan {
LogicalPlan::Sample {
with_replacement: false,
..
} => {
}
_ => panic!("expected Sample plan"),
}
}
}
#[test]
fn replace_plan() {
let spark = session();
let df = spark.range(3).unwrap();
let replacements = vec![("old".to_string(), "new".to_string())];
let result = df.replace(replacements, Some(vec!["col"]));
match &result.plan {
LogicalPlan::NAReplace { replacements, .. } => {
assert_eq!(replacements.len(), 1);
}
_ => panic!("expected NAReplace plan"),
}
}
#[test]
fn builders_construct_and_serialize() {
use crate::functions::col;
use crate::types::{DataType, StructField};
let spark = session();
let df = spark.range(5).unwrap();
let df2 = spark.range(5).unwrap();
let e = || col("id").expression().clone();
let ser = |d: &DataFrame| {
build_input_relation(d.plan(), &spark).expect("plan serializes to a relation");
};
ser(&df.select(vec![col("id")]));
ser(&df.filter(col("id")));
ser(&df.where_(col("id")));
ser(&df.with_column("x", col("id")));
ser(&df.with_column_renamed("id", "y"));
ser(&df.drop(vec!["id"]));
ser(&df.limit(3));
ser(&df.offset(1));
ser(&df.distinct());
ser(&df.drop_duplicates(Some(vec!["id"])));
ser(&df.sort(vec![e()]));
ser(&df.order_by(vec![e()]));
ser(&df.sort_within_partitions(vec![e()]));
ser(&df.cross_join(&df2));
ser(&df.union(&df2));
ser(&df.union_all(&df2));
ser(&df.union_by_name(&df2));
ser(&df.intersect(&df2));
ser(&df.intersect_all(&df2));
ser(&df.subtract(&df2));
ser(&df.except_all(&df2));
ser(&df.repartition(4));
ser(&df.coalesce(2));
ser(&df.repartition_by_range(3, vec![e()]));
ser(&df.hint("broadcast", Vec::<String>::new()));
ser(&df.to_df(vec!["a"]));
ser(&df.alias("t"));
ser(&df.sample(0.5, Some(1)));
ser(&df.select_expr(vec!["id + 1"]));
ser(&df.col_regex("id"));
ser(&df.describe(vec!["id"]));
ser(&df.summary(vec!["count"]));
ser(&df.as_table("t2"));
ser(&df.to(DataType::Struct {
fields: vec![StructField {
name: "id".to_string(),
data_type: DataType::Long,
nullable: true,
metadata: std::collections::BTreeMap::new(),
}],
}));
ser(&df.unpivot(vec![col("id")], None::<Vec<Column>>, "var", "val"));
ser(&df.melt(vec!["id"], None, "var", "val"));
ser(&df.group_by(vec![col("id")]).agg(vec![e()]));
ser(&df.rollup(vec![col("id")]).agg(vec![e()]));
ser(&df.cube(vec![col("id")]).agg(vec![e()]));
ser(&df.grouping_sets(vec![vec![col("id")]]).agg(vec![e()]));
ser(&df.with_watermark("id", "1 minute"));
let mut md = std::collections::HashMap::new();
md.insert("k".to_string(), "v".to_string());
ser(&df.with_metadata("id", md));
ser(&df.replace(vec![("a".to_string(), "b".to_string())], None));
ser(&df.stat().crosstab("id", "id"));
ser(&df.stat().freq_items(vec!["id"], 0.5));
}
#[test]
fn streaming_reader_and_writer_builders() {
use crate::streaming::Trigger;
let spark = session();
let ser = |d: &DataFrame| {
build_input_relation(d.plan(), &spark).expect("stream plan serializes");
};
ser(&spark
.read_stream()
.format("rate")
.option("rowsPerSecond", "5")
.load(None));
ser(&spark.read_stream().schema("value long").json("/tmp/in"));
ser(&spark.read_stream().parquet("/tmp/in"));
ser(&spark.read_stream().csv("/tmp/in"));
ser(&spark.read_stream().orc("/tmp/in"));
ser(&spark.read_stream().text("/tmp/in"));
ser(&spark.read_stream().format("rate").table("t"));
let base = spark.range(3).unwrap();
for trig in [
Trigger::ProcessingTime("1 second".to_string()),
Trigger::Once,
Trigger::AvailableNow,
Trigger::Continuous("1 second".to_string()),
] {
let _w = base
.write_stream()
.output_mode("append")
.format("console")
.option("k", "v")
.partition_by(vec!["id"])
.cluster_by(vec!["id"])
.query_name("q")
.trigger(trig);
}
}
#[test]
fn column_operations_and_expressions() {
use crate::functions::col;
let a = || col("a");
let b = || col("b");
let exprs = vec![
a().add(b()),
a().sub(b()),
a().mul(b()),
a().div(b()),
a().modulo(b()),
a().and(b()),
a().or(b()),
a().not(),
a().neg(),
a().eq(b()),
a().ne(b()),
a().gt(b()),
a().lt(b()),
a().ge(b()),
a().le(b()),
a().bitwise_and(b()),
a().bitwise_or(b()),
a().bitwise_xor(b()),
a().eq_null_safe(b()),
a().is_null(),
a().is_not_null(),
a().is_nan(),
a().like("x%"),
a().rlike("x.*"),
a().ilike("x%"),
a().contains(b()),
a().startswith(b()),
a().endswith(b()),
a().substr(b(), b()),
a().between(b(), b()),
a().isin(vec![b()]),
a().get_field("f"),
a().get_item(b()),
a().with_field("f", b()),
a().drop_fields(vec!["f"]),
a().asc(),
a().asc_nulls_first(),
a().asc_nulls_last(),
a().desc(),
a().desc_nulls_first(),
a().desc_nulls_last(),
a().alias("x"),
a().name("y"),
a().cast_str("int"),
a().try_cast_str("int"),
a().astype(crate::types::DataType::Integer),
a().when(b(), b()).otherwise(b()),
];
for e in &exprs {
let _ = e.to_proto();
}
}
#[test]
fn exotic_plan_variants_serialize() {
use crate::functions::col;
use crate::types::DataType;
use crate::udf::{CommonInlineUserDefinedFunctionExpression, PythonUDFPayload};
let spark = session();
let ser = |d: &DataFrame| {
build_input_relation(d.plan(), &spark).expect("exotic plan serializes");
};
let df = spark.range(5).unwrap();
let df2 = spark.range(5).unwrap();
ser(&df.zip(&df2).unwrap());
ser(&df.transpose().unwrap());
ser(&df.transpose_with_index(col("id")).unwrap());
ser(&df.nearest_by_join(&df2, col("id"), 5, "inner", "asc", "inner"));
let udf = || {
CommonInlineUserDefinedFunctionExpression::new(
"f".to_string(),
true,
vec![],
PythonUDFPayload::new(DataType::Integer, 200, vec![1, 2, 3], "3.11".to_string()),
)
};
ser(&df.map_in_pandas(udf(), false));
ser(&df.map_in_arrow(udf(), false));
ser(&df.group_by(vec![col("id")]).apply_in_pandas(udf()));
ser(&df.group_by(vec![col("id")]).apply_in_arrow(udf()));
let g1 = df.group_by(vec![col("id")]);
let g2 = df2.group_by(vec![col("id")]);
ser(&g1.cogroup(&g2).apply_in_pandas(udf()));
let udtf_df = spark.tvf().udtf(
"myudtf",
vec![],
Some(DataType::Integer),
300,
vec![1, 2],
"3.11".to_string(),
true,
);
ser(&udtf_df);
}
}
#[cfg(test)]
mod arrow_value_tests {
use super::*;
use arrow::array::*;
use arrow::datatypes::{
i256, DataType as ArrowDataType, Field, Int32Type, IntervalDayTime, IntervalMonthDayNano,
};
use std::sync::Arc;
#[test]
fn primitives_and_signed_ints() {
assert!(matches!(
arrow_value_at(&BooleanArray::from(vec![true]), 0).unwrap(),
Value::Bool(true)
));
assert!(matches!(
arrow_value_at(&Int8Array::from(vec![1i8]), 0).unwrap(),
Value::Byte(1)
));
assert!(matches!(
arrow_value_at(&Int16Array::from(vec![1i16]), 0).unwrap(),
Value::Short(1)
));
assert!(matches!(
arrow_value_at(&Int32Array::from(vec![1i32]), 0).unwrap(),
Value::Integer(1)
));
assert!(matches!(
arrow_value_at(&Int64Array::from(vec![1i64]), 0).unwrap(),
Value::Long(1)
));
assert!(matches!(
arrow_value_at(&Float32Array::from(vec![1.0f32]), 0).unwrap(),
Value::Float(_)
));
assert!(matches!(
arrow_value_at(&Float64Array::from(vec![1.0f64]), 0).unwrap(),
Value::Double(_)
));
assert!(matches!(
arrow_value_at(&StringArray::from(vec!["x"]), 0).unwrap(),
Value::String(_)
));
assert!(matches!(
arrow_value_at(&BinaryArray::from_iter_values([b"x".as_ref()]), 0).unwrap(),
Value::Binary(_)
));
assert!(matches!(
arrow_value_at(&Date32Array::from(vec![1i32]), 0).unwrap(),
Value::Date(1)
));
assert!(matches!(
arrow_value_at(&TimestampMicrosecondArray::from(vec![1i64]), 0).unwrap(),
Value::Timestamp(1)
));
}
#[test]
fn unsigned_ints() {
assert!(matches!(
arrow_value_at(&UInt8Array::from(vec![1u8]), 0).unwrap(),
Value::Short(1)
));
assert!(matches!(
arrow_value_at(&UInt16Array::from(vec![1u16]), 0).unwrap(),
Value::Integer(1)
));
assert!(matches!(
arrow_value_at(&UInt32Array::from(vec![1u32]), 0).unwrap(),
Value::Long(1)
));
assert!(matches!(
arrow_value_at(&UInt64Array::from(vec![1u64]), 0).unwrap(),
Value::Long(1)
));
}
#[test]
fn decimals_128_and_256() {
let d128 = Decimal128Array::from(vec![12345i128])
.with_precision_and_scale(10, 2)
.unwrap();
assert!(matches!(
arrow_value_at(&d128, 0).unwrap(),
Value::Decimal { .. }
));
let d256 = Decimal256Array::from(vec![i256::from_i128(12345)])
.with_precision_and_scale(10, 2)
.unwrap();
assert!(matches!(
arrow_value_at(&d256, 0).unwrap(),
Value::Decimal { .. }
));
}
#[test]
fn large_and_view_bytes() {
assert!(matches!(
arrow_value_at(&LargeStringArray::from_iter_values(["x"]), 0).unwrap(),
Value::String(_)
));
assert!(matches!(
arrow_value_at(&LargeBinaryArray::from_iter_values([b"x".as_ref()]), 0).unwrap(),
Value::Binary(_)
));
assert!(matches!(
arrow_value_at(&StringViewArray::from_iter_values(["x"]), 0).unwrap(),
Value::String(_)
));
assert!(matches!(
arrow_value_at(&BinaryViewArray::from_iter_values([b"x".as_ref()]), 0).unwrap(),
Value::Binary(_)
));
}
#[test]
fn timestamps_and_date64() {
assert!(matches!(
arrow_value_at(&TimestampSecondArray::from(vec![1i64]), 0).unwrap(),
Value::Timestamp(_)
));
assert!(matches!(
arrow_value_at(&TimestampMillisecondArray::from(vec![1i64]), 0).unwrap(),
Value::Timestamp(_)
));
assert!(matches!(
arrow_value_at(&TimestampNanosecondArray::from(vec![1000i64]), 0).unwrap(),
Value::Timestamp(_)
));
assert!(matches!(
arrow_value_at(&Date64Array::from(vec![86_400_000i64]), 0).unwrap(),
Value::Date(_)
));
}
#[test]
fn time_types_render_as_string() {
assert!(matches!(
arrow_value_at(&Time64MicrosecondArray::from(vec![1i64]), 0).unwrap(),
Value::String(_)
));
assert!(matches!(
arrow_value_at(&Time64NanosecondArray::from(vec![1000i64]), 0).unwrap(),
Value::String(_)
));
assert!(matches!(
arrow_value_at(&Time32MillisecondArray::from(vec![1i32]), 0).unwrap(),
Value::String(_)
));
assert!(matches!(
arrow_value_at(&Time32SecondArray::from(vec![1i32]), 0).unwrap(),
Value::String(_)
));
}
#[test]
fn interval_types_render_as_string() {
assert!(matches!(
arrow_value_at(&IntervalYearMonthArray::from(vec![13i32]), 0).unwrap(),
Value::String(_)
));
let dt = IntervalDayTimeArray::from(vec![IntervalDayTime::new(1, 100)]);
assert!(matches!(arrow_value_at(&dt, 0).unwrap(), Value::String(_)));
let mdn = IntervalMonthDayNanoArray::from(vec![IntervalMonthDayNano::new(1, 2, 3)]);
assert!(matches!(arrow_value_at(&mdn, 0).unwrap(), Value::String(_)));
}
#[test]
fn fixed_size_binary() {
let arr = FixedSizeBinaryArray::try_from_iter(vec![vec![1u8, 2u8]].into_iter()).unwrap();
assert!(matches!(arrow_value_at(&arr, 0).unwrap(), Value::Binary(_)));
}
#[test]
fn nested_list_struct_map() {
let list =
ListArray::from_iter_primitive::<Int32Type, _, _>(vec![Some(vec![Some(1), Some(2)])]);
assert!(matches!(arrow_value_at(&list, 0).unwrap(), Value::List(_)));
let field = Arc::new(Field::new("a", ArrowDataType::Int32, false));
let col: ArrayRef = Arc::new(Int32Array::from(vec![1]));
let s = StructArray::from(vec![(field, col)]);
assert!(matches!(arrow_value_at(&s, 0).unwrap(), Value::Struct(_)));
let mut b = MapBuilder::new(None, StringBuilder::new(), Int32Builder::new());
b.keys().append_value("k");
b.values().append_value(1);
b.append(true).unwrap();
let m = b.finish();
assert!(matches!(arrow_value_at(&m, 0).unwrap(), Value::Map(_)));
}
#[test]
fn null_element_and_unsupported_type() {
let with_null = Int32Array::from(vec![None as Option<i32>]);
assert!(matches!(
arrow_value_at(&with_null, 0).unwrap(),
Value::Null
));
let dur = DurationSecondArray::from(vec![1i64]);
assert!(arrow_value_at(&dur, 0).is_err());
}
#[test]
fn map_key_to_string_covers_scalar_arms() {
assert_eq!(map_key_to_string(Value::String("x".to_string())), "x");
assert_eq!(map_key_to_string(Value::Bool(true)), "true");
assert_eq!(map_key_to_string(Value::Byte(1)), "1");
assert_eq!(map_key_to_string(Value::Short(2)), "2");
assert_eq!(map_key_to_string(Value::Integer(3)), "3");
assert_eq!(map_key_to_string(Value::Long(4)), "4");
assert_eq!(map_key_to_string(Value::Float(1.5)), "1.5");
assert_eq!(map_key_to_string(Value::Double(2.5)), "2.5");
assert_eq!(map_key_to_string(Value::Date(5)), "5");
assert_eq!(map_key_to_string(Value::Timestamp(6)), "6");
assert_eq!(
map_key_to_string(Value::Decimal {
value: "7.5".to_string(),
precision: None,
scale: None,
}),
"7.5"
);
let _ = map_key_to_string(Value::List(vec![]));
}
#[test]
fn i128_to_decimal_string_branches() {
assert_eq!(i128_to_decimal_string(12345, 0), "12345");
assert_eq!(i128_to_decimal_string(12345, 2), "123.45");
assert_eq!(i128_to_decimal_string(5, 4), "0.0005");
assert_eq!(i128_to_decimal_string(-5, 4), "-0.0005");
}
#[test]
fn micros_to_time_string_branches() {
assert_eq!(micros_to_time_string(0), "00:00:00");
assert!(micros_to_time_string(1).contains('.'));
}
}