#![allow(clippy::unwrap_used, clippy::expect_used)]
use std::collections::BTreeMap;
use std::fs::File;
use std::net::{TcpListener, TcpStream};
use std::path::Path;
use std::process::{Child, Command as StdCommand, Stdio};
use std::time::{Duration, Instant};
use arrow::array::{Array, Float64Array, Int64Array};
use assert_cmd::Command;
use parquet::arrow::arrow_reader::ParquetRecordBatchReaderBuilder;
const QUERY: &str = "SELECT k, SUM(v) AS sv FROM t GROUP BY k ORDER BY k";
fn oxide() -> Command {
Command::cargo_bin("oxide").unwrap()
}
#[test]
fn tui_rejects_non_interactive_terminal() {
oxide()
.arg("tui")
.assert()
.code(2)
.stderr("oxide tui needs an interactive terminal; use oxide sql or oxide explain for non-interactive output\n");
}
fn gen_data(dir: &Path, rows: u64) {
let assert = oxide()
.args([
"gen-data",
"--rows",
&rows.to_string(),
"--out",
dir.to_str().unwrap(),
"--row-group-rows",
"4096",
])
.assert()
.success();
let stdout = String::from_utf8(assert.get_output().stdout.clone()).unwrap();
assert!(stdout.contains(&format!("wrote {rows} rows")), "{stdout}");
}
fn expected_sums(dir: &Path) -> BTreeMap<Option<i64>, f64> {
let file = File::open(dir.join("t.parquet")).unwrap();
let reader = ParquetRecordBatchReaderBuilder::try_new(file)
.unwrap()
.build()
.unwrap();
let mut sums: BTreeMap<Option<i64>, f64> = BTreeMap::new();
for batch in reader {
let batch = batch.unwrap();
let k = batch
.column(batch.schema().index_of("k").unwrap())
.as_any()
.downcast_ref::<Int64Array>()
.unwrap()
.clone();
let v = batch
.column(batch.schema().index_of("v").unwrap())
.as_any()
.downcast_ref::<Float64Array>()
.unwrap()
.clone();
for row in 0..batch.num_rows() {
if v.is_null(row) {
continue;
}
let key = (!k.is_null(row)).then(|| k.value(row));
*sums.entry(key).or_insert(0.0) += v.value(row);
}
}
sums
}
fn parse_table(stdout: &str) -> BTreeMap<Option<i64>, f64> {
let mut rows = BTreeMap::new();
for line in stdout.lines() {
let cells: Vec<&str> = line
.strip_prefix('|')
.and_then(|l| l.strip_suffix('|'))
.map(|l| l.split('|').map(str::trim).collect())
.unwrap_or_default();
if cells.len() != 2 || cells[0] == "k" {
continue;
}
let key = (!cells[0].is_empty()).then(|| cells[0].parse::<i64>().unwrap());
let sum = cells[1].parse::<f64>().unwrap();
rows.insert(key, sum);
}
rows
}
#[test]
fn gen_data_sql_and_explain_are_correct_end_to_end() {
let dir = tempfile::tempdir().unwrap();
gen_data(dir.path(), 20_000);
let table_arg = format!("t={}", dir.path().display());
let assert = oxide()
.args(["sql", "-q", QUERY, "--table", &table_arg])
.assert()
.success();
let stdout = String::from_utf8(assert.get_output().stdout.clone()).unwrap();
let printed = parse_table(&stdout);
let expected = expected_sums(dir.path());
assert_eq!(printed.len(), expected.len(), "{stdout}");
for (key, sum) in &expected {
let got = printed
.get(key)
.unwrap_or_else(|| panic!("missing {key:?}"));
assert_eq!(got, sum, "group {key:?}\n{stdout}");
}
let assert = oxide()
.args([
"sql", "-q", QUERY, "--table", &table_arg, "--target", "cuda",
])
.assert()
.success();
let gpu_stdout = String::from_utf8(assert.get_output().stdout.clone()).unwrap();
assert_eq!(parse_table(&gpu_stdout), expected);
let assert = oxide()
.args([
"explain",
"-q",
"SELECT k, SUM(v) FROM t WHERE k >= 2 AND v < 4.0 GROUP BY k",
"--table",
&table_arg,
"--target",
"cuda",
])
.assert()
.success();
let plan = String::from_utf8(assert.get_output().stdout.clone()).unwrap();
assert!(plan.contains("GpuAggregateExec[cuda]"), "{plan}");
assert!(plan.contains("GpuFilterExec[cuda]"), "{plan}");
let assert = oxide()
.args([
"explain",
"-q",
"SELECT *, l2_distance(emb, [1.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0]) AS d FROM t",
"--table",
&table_arg,
"--target",
"metal",
])
.assert()
.success();
let plan = String::from_utf8(assert.get_output().stdout.clone()).unwrap();
assert!(
plan.contains("GpuVectorDistanceExec[metal]: l2(emb) AS d, dim=8"),
"{plan}"
);
}
struct ClusterGuard {
scheduler: Child,
worker: Option<Child>,
}
impl Drop for ClusterGuard {
fn drop(&mut self) {
if let Some(worker) = &mut self.worker {
let _ = worker.kill();
let _ = worker.wait();
}
let _ = self.scheduler.kill();
let _ = self.scheduler.wait();
}
}
fn free_port() -> u16 {
TcpListener::bind("127.0.0.1:0")
.unwrap()
.local_addr()
.unwrap()
.port()
}
fn wait_for_port(port: u16, what: &str) {
let deadline = Instant::now() + Duration::from_secs(30);
while Instant::now() < deadline {
if TcpStream::connect(("127.0.0.1", port)).is_ok() {
return;
}
std::thread::sleep(Duration::from_millis(100));
}
panic!("{what} did not open port {port} within 30s");
}
#[test]
fn cluster_sql_matches_embedded_sql() {
let dir = tempfile::tempdir().unwrap();
gen_data(dir.path(), 600_000);
let table_arg = format!("t={}", dir.path().display());
let embedded = oxide()
.args(["sql", "-q", QUERY, "--table", &table_arg])
.assert()
.success();
let embedded_stdout = String::from_utf8(embedded.get_output().stdout.clone()).unwrap();
let scheduler_port = free_port();
let scheduler = StdCommand::new(env!("CARGO_BIN_EXE_oxide-scheduler"))
.args(["--port", &scheduler_port.to_string()])
.env("OXIDE_CLUSTER_BACKEND", "cuda")
.stdout(Stdio::null())
.stderr(File::create(dir.path().join("scheduler.log")).unwrap())
.spawn()
.unwrap();
let mut cluster = ClusterGuard {
scheduler,
worker: None,
};
wait_for_port(scheduler_port, "oxide-scheduler");
let (flight_port, grpc_port) = (free_port(), free_port());
cluster.worker = Some(
StdCommand::new(env!("CARGO_BIN_EXE_oxide-worker"))
.args([
"--scheduler-host",
"127.0.0.1",
"--scheduler-port",
&scheduler_port.to_string(),
"--port",
&flight_port.to_string(),
"--grpc-port",
&grpc_port.to_string(),
"--concurrent-tasks",
"2",
"--work-dir",
dir.path().join("shuffle").to_str().unwrap(),
])
.stdout(Stdio::null())
.stderr(File::create(dir.path().join("worker.log")).unwrap())
.spawn()
.unwrap(),
);
wait_for_port(grpc_port, "oxide-worker");
let url = format!("df://127.0.0.1:{scheduler_port}");
let deadline = Instant::now() + Duration::from_secs(60);
let cluster_stdout = loop {
let output = oxide()
.args(["sql", "-q", QUERY, "--table", &table_arg, "--cluster", &url])
.timeout(Duration::from_secs(30))
.output()
.unwrap();
if output.status.success() {
break String::from_utf8(output.stdout).unwrap();
}
assert!(
Instant::now() < deadline,
"cluster query kept failing:\n{}\nscheduler log:\n{}\nworker log:\n{}",
String::from_utf8_lossy(&output.stderr),
std::fs::read_to_string(dir.path().join("scheduler.log")).unwrap_or_default(),
std::fs::read_to_string(dir.path().join("worker.log")).unwrap_or_default(),
);
std::thread::sleep(Duration::from_millis(500));
};
assert_eq!(cluster_stdout, embedded_stdout);
}
#[test]
fn sql_into_a_closed_pipe_exits_zero_without_panicking() {
let dir = tempfile::tempdir().unwrap();
gen_data(dir.path(), 50_000);
let bin = assert_cmd::cargo::cargo_bin("oxide");
let mut sql = StdCommand::new(&bin)
.args(["sql", "-q", "SELECT id, k, v FROM t ORDER BY id", "--table"])
.arg(format!("t={}", dir.path().display()))
.stdout(Stdio::piped())
.stderr(Stdio::piped())
.spawn()
.unwrap();
let head = StdCommand::new("head")
.arg("-1")
.stdin(Stdio::from(sql.stdout.take().unwrap()))
.stdout(Stdio::piped())
.spawn()
.unwrap();
let head_out = head.wait_with_output().unwrap();
let sql_out = sql.wait_with_output().unwrap();
let stderr = String::from_utf8_lossy(&sql_out.stderr);
assert!(
!stderr.contains("panicked"),
"a closed pipe is not a crash:\n{stderr}"
);
assert_eq!(
sql_out.status.code(),
Some(0),
"a closed pipe is a clean exit; stderr:\n{stderr}"
);
assert_eq!(
String::from_utf8_lossy(&head_out.stdout).lines().count(),
1,
"head took its one line"
);
}
#[test]
fn explain_into_a_closed_pipe_exits_zero_without_panicking() {
let dir = tempfile::tempdir().unwrap();
gen_data(dir.path(), 1_000);
let bin = assert_cmd::cargo::cargo_bin("oxide");
let mut explain = StdCommand::new(&bin)
.args(["explain", "-q", QUERY, "--table"])
.arg(format!("t={}", dir.path().display()))
.stdout(Stdio::piped())
.stderr(Stdio::piped())
.spawn()
.unwrap();
let head = StdCommand::new("head")
.arg("-1")
.stdin(Stdio::from(explain.stdout.take().unwrap()))
.stdout(Stdio::piped())
.spawn()
.unwrap();
head.wait_with_output().unwrap();
let out = explain.wait_with_output().unwrap();
let stderr = String::from_utf8_lossy(&out.stderr);
assert!(!stderr.contains("panicked"), "{stderr}");
assert_eq!(out.status.code(), Some(0), "stderr:\n{stderr}");
}
fn parse_csv(stdout: &str) -> BTreeMap<Option<i64>, i64> {
let mut lines = stdout.lines();
assert_eq!(lines.next(), Some("k,n"), "csv header\n{stdout}");
lines
.map(|line| {
let (k, n) = line.split_once(',').unwrap_or_else(|| panic!("{line:?}"));
let key = (!k.is_empty()).then(|| k.parse::<i64>().unwrap());
(key, n.parse::<i64>().unwrap())
})
.collect()
}
fn parse_json(stdout: &str) -> BTreeMap<Option<i64>, i64> {
let trimmed = stdout.trim();
let inner = trimmed
.strip_prefix('[')
.and_then(|s| s.strip_suffix(']'))
.unwrap_or_else(|| panic!("not a JSON array: {trimmed}"));
if inner.is_empty() {
return BTreeMap::new();
}
inner
.split("},{")
.map(|object| {
let object = object.trim_matches(['{', '}']);
let mut row = (None, None);
for field in object.split(',') {
let (name, value) = field.split_once(':').unwrap();
match name.trim_matches('"') {
"k" => row.0 = Some(value.parse::<i64>().unwrap()),
"n" => row.1 = Some(value.parse::<i64>().unwrap()),
other => panic!("unexpected field {other}"),
}
}
(row.0, row.1.expect("every row has a count"))
})
.collect()
}
const COUNT_QUERY: &str = "SELECT k, COUNT(v) AS n FROM t GROUP BY k ORDER BY k";
#[test]
fn output_formats_and_batch_sizes_agree_on_the_same_rows() {
let dir = tempfile::tempdir().unwrap();
gen_data(dir.path(), 20_000);
let table_arg = format!("t={}", dir.path().display());
let run = |args: &[&str]| {
let assert = oxide()
.args(["sql", "-q", COUNT_QUERY, "--table", &table_arg])
.args(args)
.assert()
.success();
String::from_utf8(assert.get_output().stdout.clone()).unwrap()
};
let csv = parse_csv(&run(&["--output", "csv"]));
assert!(!csv.is_empty());
assert_eq!(parse_json(&run(&["--output", "json"])), csv);
let table = run(&[]);
assert_eq!(run(&["--output", "table"]), table);
assert!(table.starts_with('+'), "{table}");
for rows in ["1", "7", "1000000"] {
assert_eq!(
parse_csv(&run(&["--output", "csv", "--batch-size", rows])),
csv,
"--batch-size {rows}"
);
}
}
#[test]
fn an_empty_result_keeps_its_csv_header_and_json_array() {
let dir = tempfile::tempdir().unwrap();
gen_data(dir.path(), 4_096);
let table_arg = format!("t={}", dir.path().display());
let run = |format: &str| {
let assert = oxide()
.args([
"sql",
"-q",
"SELECT k, COUNT(v) AS n FROM t WHERE k = -1 GROUP BY k",
"--table",
&table_arg,
"--output",
format,
])
.assert()
.success();
String::from_utf8(assert.get_output().stdout.clone()).unwrap()
};
assert_eq!(run("csv"), "k,n\n");
assert_eq!(run("json"), "[]\n");
}
#[test]
fn a_zero_batch_size_is_refused_rather_than_returning_nothing() {
let assert = oxide()
.args(["sql", "-q", "SELECT 1", "--batch-size", "0"])
.assert()
.failure();
let stderr = String::from_utf8(assert.get_output().stderr.clone()).unwrap();
assert!(stderr.contains("at least 1 row"), "{stderr}");
}
#[test]
fn an_unknown_output_format_names_the_ones_that_exist() {
let assert = oxide()
.args(["sql", "-q", "SELECT 1", "--output", "yaml"])
.assert()
.failure();
let stderr = String::from_utf8(assert.get_output().stderr.clone()).unwrap();
assert!(stderr.contains("table, json, csv"), "{stderr}");
}
#[test]
fn the_flags_appear_in_help() {
let assert = oxide().args(["sql", "--help"]).assert().success();
let help = String::from_utf8(assert.get_output().stdout.clone()).unwrap();
for flag in ["--batch-size", "--output", "--target", "--cluster"] {
assert!(help.contains(flag), "oxide sql --help lacks {flag}\n{help}");
}
let assert = Command::cargo_bin("oxide-worker")
.unwrap()
.arg("--help")
.assert()
.success();
let help = String::from_utf8(assert.get_output().stdout.clone()).unwrap();
assert!(help.contains("--backend"), "{help}");
assert!(help.contains("OXIDE_BACKEND"), "{help}");
}
#[test]
fn a_worker_refuses_a_backend_this_machine_does_not_have() {
let availability = oxidelake_device::HardwareDetector::availability();
let absent = if !availability.cuda {
"cuda"
} else if !availability.metal {
"metal"
} else {
return;
};
let assert = Command::cargo_bin("oxide-worker")
.unwrap()
.args(["--backend", absent])
.assert()
.failure();
let stderr = String::from_utf8(assert.get_output().stderr.clone()).unwrap();
assert!(stderr.contains(absent), "{stderr}");
assert!(stderr.contains("explicitly requested"), "{stderr}");
}
#[test]
fn a_worker_rejects_an_unknown_backend_name() {
let assert = Command::cargo_bin("oxide-worker")
.unwrap()
.args(["--backend", "tpu"])
.assert()
.failure();
let stderr = String::from_utf8(assert.get_output().stderr.clone()).unwrap();
assert!(stderr.contains("cpu, cuda, metal"), "{stderr}");
}
#[test]
fn explain_names_why_a_node_stayed_on_the_cpu() {
let dir = tempfile::tempdir().unwrap();
gen_data(dir.path(), 4_096);
let table_arg = format!("t={}", dir.path().display());
let assert = oxide()
.args([
"explain",
"-q",
"SELECT k, AVG(v) FROM t GROUP BY k",
"--table",
&table_arg,
"--target",
"cuda",
])
.assert()
.success();
let stdout = String::from_utf8(assert.get_output().stdout.clone()).unwrap();
assert!(stdout.contains("placement notes (target cuda)"), "{stdout}");
assert!(stdout.contains("AggregateExec:"), "{stdout}");
assert!(stdout.contains("sum, count, min, max"), "{stdout}");
let assert = oxide()
.args([
"explain",
"-q",
"SELECT k, SUM(v) FROM t WHERE k >= 2 GROUP BY k",
"--table",
&table_arg,
"--target",
"cuda",
])
.assert()
.success();
let stdout = String::from_utf8(assert.get_output().stdout.clone()).unwrap();
assert!(stdout.contains("GpuAggregateExec[cuda]"), "{stdout}");
assert!(!stdout.contains("placement notes"), "{stdout}");
}
#[test]
fn a_query_logs_one_line_with_rows_and_elapsed() {
let dir = tempfile::tempdir().unwrap();
gen_data(dir.path(), 4_096);
let table_arg = format!("t={}", dir.path().display());
let assert = oxide()
.env("RUST_LOG", "oxidelake_runtime=info")
.args([
"sql",
"-q",
"SELECT k, SUM(v) FROM t WHERE k >= 2 GROUP BY k",
"--table",
&table_arg,
"--target",
"cuda",
])
.assert()
.success();
let stderr = String::from_utf8(assert.get_output().stderr.clone()).unwrap();
assert!(stderr.contains("query finished"), "{stderr}");
assert!(stderr.contains("mode=embedded/cuda"), "{stderr}");
assert!(stderr.contains("rows="), "{stderr}");
assert!(stderr.contains("elapsed_ms="), "{stderr}");
assert!(stderr.contains("fallback_batches="), "{stderr}");
let stdout = String::from_utf8(assert.get_output().stdout.clone()).unwrap();
assert!(!stdout.contains("query finished"), "{stdout}");
}
#[test]
#[cfg_attr(feature = "metrics", ignore = "the flag is honoured in this build")]
fn a_metrics_port_without_the_feature_is_refused() {
let assert = Command::cargo_bin("oxide-worker")
.unwrap()
.args(["--metrics-port", "19999"])
.assert()
.failure();
let stderr = String::from_utf8(assert.get_output().stderr.clone()).unwrap();
assert!(
stderr.contains("--metrics-port needs the `metrics` feature"),
"{stderr}"
);
assert!(stderr.contains("--features metrics"), "{stderr}");
}