use spark_connect::column::{col, lit};
use spark_connect::dataframe::DataFrame;
use spark_connect::functions as f;
use spark_connect::plan::JoinType;
use spark_connect::session::SparkSession;
fn should_run() -> bool {
std::env::var("SPARK_REMOTE").is_ok()
}
fn session() -> SparkSession {
let url = std::env::var("SPARK_REMOTE").unwrap_or_else(|_| "sc://localhost:15002".to_string());
SparkSession::builder()
.remote(&url)
.get_or_create()
.expect("session")
}
fn df4(s: &SparkSession) -> DataFrame {
s.sql("SELECT * FROM VALUES (1,'a'),(2,'b'),(2,'b'),(3,'c') AS t(id, name)")
.expect("df4")
}
#[test]
fn dataframe_transforms_and_setops() {
if !should_run() {
return;
}
let s = session();
let df = df4(&s);
assert_eq!(df.where_(col("id").gt(lit(1))).count().unwrap(), 3);
let se = df.select_expr(vec!["id + 1 AS x"]);
let xs: Vec<i64> = se
.collect()
.unwrap()
.iter()
.map(|r| r.get(0).unwrap().as_i64().unwrap())
.collect();
assert_eq!(xs, vec![2, 3, 3, 4]);
assert_eq!(df.drop_duplicates(None).count().unwrap(), 3);
assert_eq!(df.drop_duplicates(Some(vec!["id"])).count().unwrap(), 3);
let other = s
.sql("SELECT * FROM VALUES (2,'b'),(9,'z') AS t(id, name)")
.unwrap();
assert_eq!(df.union_all(&other).count().unwrap(), 6);
assert_eq!(df.intersect_all(&other).count().unwrap(), 1); assert_eq!(df.except_all(&other).count().unwrap(), 3); assert_eq!(df.union_by_name(&other).count().unwrap(), 6);
let right = s
.sql("SELECT * FROM VALUES (1,100),(2,200) AS t(id, v)")
.unwrap();
assert_eq!(
df.join_using(&right, vec!["id".to_string()], JoinType::Inner)
.count()
.unwrap(),
3
);
assert_eq!(
df.sort_within_partitions(vec![col("id").expression().clone()])
.count()
.unwrap(),
4
);
assert_eq!(
df.repartition_by_range(2, vec![col("id").expression().clone()])
.count()
.unwrap(),
4
);
assert_eq!(df.to_schema(vec!["id", "name"]).count().unwrap(), 4);
let m = df.melt(vec!["id"], Some(vec!["name"]), "var", "val");
assert_eq!(m.count().unwrap(), 4);
}
#[test]
fn dataframe_metadata_json_semantics() {
if !should_run() {
return;
}
let s = session();
let df = df4(&s);
assert_eq!(df.take(2).unwrap().len(), 2);
let json = df.to_json().unwrap();
assert_eq!(json.len(), 4);
assert!(json
.iter()
.all(|j| j.starts_with('{') && j.contains("id") && j.contains("name")));
assert!(df.exists().unwrap());
let one = s.sql("SELECT 42 AS v").unwrap();
assert_eq!(
one.scalar().unwrap(),
Some(spark_connect::row::Value::Integer(42))
);
assert!(df.same_semantics(&df4(&s)).unwrap());
let _ = df.semantic_hash().unwrap();
assert!(!df.is_streaming());
let _ = df.is_local();
let _ = df.input_files().unwrap();
let observed = df.observe("obs", vec![f::count(lit(1)).expression().clone()]);
let _ = observed.collect().unwrap();
let cached = df.cache().unwrap();
let _ = cached.count().unwrap();
let _ = cached.is_cached().unwrap();
let _ = cached.storage_level().unwrap();
df.create_or_replace_temp_view("cov_df_view").unwrap();
assert_eq!(s.table("cov_df_view").unwrap().count().unwrap(), 4);
let _ = df.create_or_replace_global_temp_view("cov_df_gview");
let parts = df.random_split(vec![0.5, 0.5], Some(1));
assert_eq!(parts.len(), 2);
}