use arrow_array::{Float32Array, Int32Array, RecordBatch};
use arrow_schema::{DataType, Field as ArrowField, Schema as ArrowSchema};
use futures::{StreamExt, TryStreamExt};
use lance::dataset::NewColumnTransform;
use std::sync::Arc;
use super::MaterializedView;
use super::refresh::RefreshMode;
use crate::connect;
use crate::connection::Connection;
use crate::query::{ExecutableQuery, QueryBase, Select};
use crate::table::{CompactionOptions, OptimizeAction, Table};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum SrcOp {
AppendNew,
DeleteEven,
UpdateOddScore,
Compact,
AddColumn,
MergeDropLargest,
MergeUpsert,
}
const ALL_OPS: [SrcOp; 7] = [
SrcOp::AppendNew,
SrcOp::DeleteEven,
SrcOp::UpdateOddScore,
SrcOp::Compact,
SrcOp::AddColumn,
SrcOp::MergeDropLargest,
SrcOp::MergeUpsert,
];
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum Shape {
Identity,
Filtered,
Limited,
}
impl Shape {
fn filter(&self) -> Option<&'static str> {
match self {
Self::Identity | Self::Limited => None,
Self::Filtered => Some("score > 50"),
}
}
fn matches(&self, score: f32) -> bool {
match self {
Self::Identity | Self::Limited => true,
Self::Filtered => score > 50.0,
}
}
fn limit(&self) -> Option<usize> {
match self {
Self::Limited => Some(4),
_ => None,
}
}
}
struct Case {
conn: Connection,
source: Table,
view: MaterializedView,
shape: Shape,
next_id: i32,
added_columns: u32,
}
fn rows_batch(ids: &[i32]) -> RecordBatch {
let scores: Vec<f32> = ids.iter().map(|id| (*id * 10) as f32).collect();
RecordBatch::try_new(
Arc::new(ArrowSchema::new(vec![
ArrowField::new("id", DataType::Int32, true),
ArrowField::new("score", DataType::Float32, true),
])),
vec![
Arc::new(Int32Array::from(ids.to_vec())),
Arc::new(Float32Array::from(scores)),
],
)
.unwrap()
}
fn merge_batch(ids: &[i32]) -> RecordBatch {
let scores: Vec<f32> = ids.iter().map(|id| (*id * 10 + 5) as f32).collect();
RecordBatch::try_new(
Arc::new(ArrowSchema::new(vec![
ArrowField::new("id", DataType::Int32, true),
ArrowField::new("score", DataType::Float32, true),
])),
vec![
Arc::new(Int32Array::from(ids.to_vec())),
Arc::new(Float32Array::from(scores)),
],
)
.unwrap()
}
impl Case {
async fn new(shape: Shape) -> Self {
let conn = connect("memory://").execute().await.unwrap();
let source = conn
.create_table("src", rows_batch(&[1, 2, 3, 4]))
.write_options(crate::materialized_view::tests::stable_row_ids())
.execute()
.await
.unwrap();
let mut builder = conn
.create_materialized_view("view", "src")
.select([("id", "id"), ("score", "score")]);
if let Some(filter) = shape.filter() {
builder = builder.only_if(filter);
}
if let Some(limit) = shape.limit() {
builder = builder.limit(limit as u64);
}
let view = builder.execute().await.unwrap();
Self {
conn,
source,
view,
shape,
next_id: 100,
added_columns: 0,
}
}
async fn apply(&mut self, op: SrcOp) {
match op {
SrcOp::AppendNew => {
let ids = vec![self.next_id, self.next_id + 101, self.next_id + 202];
self.next_id += 303;
self.source.add(rows_batch(&ids)).execute().await.unwrap();
}
SrcOp::DeleteEven => {
self.source.delete("id % 2 = 0").await.unwrap();
}
SrcOp::UpdateOddScore => {
self.source
.update()
.column("score", "-1.0")
.only_if("id % 2 = 1")
.execute()
.await
.unwrap();
}
SrcOp::Compact => {
self.source
.optimize(OptimizeAction::Compact {
options: CompactionOptions::default(),
remap_options: None,
})
.await
.unwrap();
}
SrcOp::MergeDropLargest => {
let mut ids = self.source_ids().await;
ids.sort_unstable();
ids.pop();
if ids.is_empty() {
return;
}
let batch = rows_batch(&ids);
let reader =
arrow_array::RecordBatchIterator::new(vec![Ok(batch.clone())], batch.schema());
let mut merge = self.source.merge_insert(&["id"]);
merge.when_not_matched_by_source_delete(None);
merge.execute(Box::new(reader)).await.unwrap();
}
SrcOp::MergeUpsert => {
let mut ids = self.source_ids().await;
ids.sort_unstable();
let existing = ids.first().copied().unwrap_or(self.next_id);
let fresh = self.next_id;
self.next_id += 1;
let batch = merge_batch(&[existing, fresh]);
let reader =
arrow_array::RecordBatchIterator::new(vec![Ok(batch.clone())], batch.schema());
let mut merge = self.source.merge_insert(&["id"]);
merge
.when_matched_update_all(None)
.when_not_matched_insert_all();
merge.execute(Box::new(reader)).await.unwrap();
}
SrcOp::AddColumn => {
self.added_columns += 1;
let field = ArrowField::new(
format!("extra_{}", self.added_columns),
DataType::Int32,
true,
);
self.source
.add_columns()
.transform(NewColumnTransform::AllNulls(Arc::new(ArrowSchema::new(
vec![field],
))))
.execute()
.await
.unwrap();
}
}
}
async fn source_ids(&self) -> Vec<i32> {
read_rows(
self.source
.query()
.select(Select::columns(&["id", "score"])),
)
.await
.into_iter()
.map(|(id, _)| id)
.collect()
}
async fn oracle(&self) -> Vec<(i32, i32)> {
let mut rows = read_rows(
self.source
.query()
.select(Select::columns(&["id", "score"])),
)
.await
.into_iter()
.filter(|(_, score)| self.shape.matches(*score as f32))
.collect::<Vec<_>>();
rows.sort_unstable();
rows
}
async fn view_rows(&self) -> Vec<(i32, i32)> {
let mut rows = read_rows(
self.view
.table()
.query()
.select(Select::columns(&["id", "score"])),
)
.await;
rows.sort_unstable();
rows
}
async fn check(&self, label: &str) -> Result<(), String> {
let expected = self.oracle().await;
let actual = self.view_rows().await;
let Some(cap) = self.shape.limit() else {
if expected != actual {
return Err(format!(
"{label}: view diverged from oracle\n expected: {expected:?}\n actual: {actual:?}"
));
}
return Ok(());
};
if actual.len() > cap {
return Err(format!(
"{label}: view holds {} rows, over its cap of {cap}: {actual:?}",
actual.len()
));
}
let mut unique = actual.clone();
unique.dedup();
if unique.len() != actual.len() {
return Err(format!("{label}: view holds a row twice: {actual:?}"));
}
if let Some(stray) = actual.iter().find(|row| !expected.contains(row)) {
return Err(format!(
"{label}: view holds {stray:?}, which the definition does not select: {expected:?}"
));
}
if actual.len() < cap.min(expected.len()) {
return Err(format!(
"{label}: view holds {} of {} selectable rows under a cap of {cap}: {actual:?}",
actual.len(),
expected.len()
));
}
Ok(())
}
}
async fn read_rows(query: impl ExecutableQuery) -> Vec<(i32, i32)> {
let batches = query
.execute()
.await
.unwrap()
.try_collect::<Vec<_>>()
.await
.unwrap();
batches
.iter()
.flat_map(|batch| {
let ids = batch["id"].as_any().downcast_ref::<Int32Array>().unwrap();
let scores = batch["score"]
.as_any()
.downcast_ref::<Float32Array>()
.unwrap();
(0..batch.num_rows())
.map(|i| (ids.value(i), scores.value(i) as i32))
.collect::<Vec<_>>()
})
.collect()
}
async fn run_sequence(ops: &[SrcOp], shape: Shape) -> Result<(), String> {
let label = format!("{shape:?} {ops:?}");
let mut case = Case::new(shape).await;
case.view
.refresh()
.execute()
.await
.map_err(|e| format!("{label}: initial refresh failed: {e}"))?;
case.check(&format!("{label} (initial)")).await?;
for (step, op) in ops.iter().enumerate() {
case.apply(*op).await;
case.view
.refresh()
.execute()
.await
.map_err(|e| format!("{label}: refresh at step {step} failed: {e}"))?;
case.check(&format!("{label} (step {step}, {op:?})"))
.await?;
}
case.view
.refresh()
.full(true)
.execute()
.await
.map_err(|e| format!("{label}: final full refresh failed: {e}"))?;
case.check(&format!("{label} (final rebuild)")).await?;
let _ = &case.conn;
Ok(())
}
fn all_sequences(max_len: u32) -> Vec<Vec<SrcOp>> {
let mut sequences = Vec::new();
for len in 1..=max_len {
for mut index in 0..ALL_OPS.len().pow(len) {
let mut ops = Vec::with_capacity(len as usize);
for _ in 0..len {
ops.push(ALL_OPS[index % ALL_OPS.len()]);
index /= ALL_OPS.len();
}
sequences.push(ops);
}
}
sequences
}
async fn run_exhaustive(max_len: u32) {
let mut cases = Vec::new();
for shape in [Shape::Identity, Shape::Filtered, Shape::Limited] {
for ops in all_sequences(max_len) {
cases.push((ops, shape));
}
}
let failures: Vec<String> = futures::stream::iter(cases)
.map(|(ops, shape)| async move { run_sequence(&ops, shape).await.err() })
.buffer_unordered(8)
.filter_map(|failure| async move { failure })
.collect()
.await;
assert!(
failures.is_empty(),
"{} sequences diverged; first: {}",
failures.len(),
failures[0]
);
}
#[tokio::test(flavor = "multi_thread")]
async fn differential_exhaustive() {
run_exhaustive(3).await;
}
#[tokio::test(flavor = "multi_thread")]
#[ignore = "longer sweep; run manually"]
async fn differential_exhaustive_deep() {
run_exhaustive(4).await;
}
#[tokio::test(flavor = "multi_thread")]
async fn differential_named_regressions() {
let mut case = Case::new(Shape::Identity).await;
case.view.refresh().execute().await.unwrap();
case.apply(SrcOp::AppendNew).await;
let result = case.view.refresh().execute().await.unwrap();
assert_eq!(result.mode, RefreshMode::Incremental);
case.check("append stays incremental").await.unwrap();
let mut case = Case::new(Shape::Identity).await;
case.view.refresh().execute().await.unwrap();
case.apply(SrcOp::AddColumn).await;
let result = case.view.refresh().execute().await.unwrap();
assert_eq!(result.mode, RefreshMode::Incremental);
assert_eq!(result.rows_written, 0);
let mut case = Case::new(Shape::Identity).await;
case.view.refresh().execute().await.unwrap();
case.apply(SrcOp::AppendNew).await;
case.view.refresh().execute().await.unwrap();
case.apply(SrcOp::Compact).await;
let result = case.view.refresh().execute().await.unwrap();
assert_eq!(result.mode, RefreshMode::Incremental);
assert_eq!(result.rows_written, 0);
case.check("compaction alone").await.unwrap();
case.apply(SrcOp::AppendNew).await;
let result = case.view.refresh().execute().await.unwrap();
assert_eq!(result.mode, RefreshMode::Incremental);
assert_eq!(result.rows_written, 3);
case.check("compact then append").await.unwrap();
let mut case = Case::new(Shape::Filtered).await;
case.apply(SrcOp::AppendNew).await;
case.view.refresh().execute().await.unwrap();
let before = case.view_rows().await.len();
case.apply(SrcOp::UpdateOddScore).await;
case.view.refresh().execute().await.unwrap();
let after = case.view_rows().await.len();
assert!(
after < before,
"no view-resident row was evicted ({before} -> {after}); the fixture \
no longer exercises the filtered-update transition"
);
case.check("update crosses the filter").await.unwrap();
}
async fn concurrency_oracle(conn: &Connection) -> Vec<i32> {
let batches: Vec<RecordBatch> = conn
.open_table("src")
.execute()
.await
.unwrap()
.query()
.select(Select::columns(&["id"]))
.execute()
.await
.unwrap()
.try_collect()
.await
.unwrap();
let mut ids = Vec::new();
for batch in &batches {
let column = batch
.column(0)
.as_any()
.downcast_ref::<Int32Array>()
.unwrap();
for i in 0..batch.num_rows() {
if column.value(i) > 1 {
ids.push(column.value(i));
}
}
}
ids.sort_unstable();
ids
}
async fn concurrency_view_ids(conn: &Connection) -> Vec<i32> {
let batches: Vec<RecordBatch> = conn
.open_table("mv")
.execute()
.await
.unwrap()
.query()
.select(Select::columns(&["id"]))
.execute()
.await
.unwrap()
.try_collect()
.await
.unwrap();
let mut ids = Vec::new();
for batch in &batches {
let column = batch
.column(0)
.as_any()
.downcast_ref::<Int32Array>()
.unwrap();
for i in 0..batch.num_rows() {
ids.push(column.value(i));
}
}
ids.sort_unstable();
ids
}
#[tokio::test]
#[ignore = "spawned as a child process by the concurrency cases"]
async fn cross_process_refresh_child() {
let Ok(dir) = std::env::var("MV_RACE_DIR") else {
return;
};
let dir = std::path::PathBuf::from(dir);
let tag = std::env::var("MV_RACE_TAG").unwrap();
let conn = connect(dir.to_str().unwrap()).execute().await.unwrap();
let table = conn.open_table("mv").execute().await.unwrap();
let _ = table.schema().await.unwrap();
let _ = table.count_rows(None).await.unwrap();
let source = conn.open_table("src").execute().await.unwrap();
let _ = source.count_rows(None).await.unwrap();
let view = MaterializedView::from_table(table).await.unwrap();
std::fs::write(dir.join(format!("ready-{tag}")), b"1").unwrap();
while !dir.join("START").exists() {
std::thread::sleep(std::time::Duration::from_millis(2));
}
let outcome = match view.refresh().execute().await {
Ok(result) => format!("committed rows={}", result.rows_written),
Err(err) if is_commit_conflict(&err) => "conflicted".to_string(),
Err(err) => format!("failed {err}"),
};
std::fs::write(dir.join(format!("outcome-{tag}")), outcome).unwrap();
}
fn is_commit_conflict(err: &crate::Error) -> bool {
let text = err.to_string();
text.contains("Retryable commit conflict") || text.contains("preempted by concurrent")
}
#[tokio::test]
async fn concurrent_refreshes_hold_each_row_once() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().to_str().unwrap().to_string();
let conn = connect(&path).execute().await.unwrap();
conn.create_table("src", rows_batch(&[1, 2, 3, 4]))
.write_options(crate::materialized_view::tests::stable_row_ids())
.execute()
.await
.unwrap();
let view = conn
.create_materialized_view("mv", "src")
.select([("id", "id"), ("score", "score")])
.only_if("id > 1")
.execute()
.await
.unwrap();
view.refresh().execute().await.unwrap();
let ids: Vec<i32> = (100..200_100).collect();
conn.open_table("src")
.execute()
.await
.unwrap()
.add(rows_batch(&ids))
.execute()
.await
.unwrap();
let tags = ["a", "b"];
let exe = std::env::current_exe().unwrap();
let children: Vec<std::process::Child> = tags
.iter()
.map(|tag| {
std::process::Command::new(&exe)
.args([
"--exact",
"materialized_view::differential::cross_process_refresh_child",
"--ignored",
"--nocapture",
])
.env("MV_RACE_DIR", dir.path())
.env("MV_RACE_SYNC", dir.path())
.env("MV_RACE_PEERS", "2")
.env("MV_RACE_TAG", tag)
.stdout(std::process::Stdio::null())
.stderr(std::process::Stdio::null())
.spawn()
.unwrap()
})
.collect();
let deadline = std::time::Instant::now() + std::time::Duration::from_secs(180);
while tags
.iter()
.any(|tag| !dir.path().join(format!("ready-{tag}")).exists())
{
assert!(
std::time::Instant::now() < deadline,
"children never became ready"
);
std::thread::sleep(std::time::Duration::from_millis(10));
}
std::fs::write(dir.path().join("START"), b"1").unwrap();
for (tag, mut child) in tags.iter().zip(children) {
let status = loop {
match child.try_wait().unwrap() {
Some(status) => break status,
None if std::time::Instant::now() >= deadline => {
child.kill().unwrap();
panic!("child {tag} never finished");
}
None => std::thread::sleep(std::time::Duration::from_millis(10)),
}
};
assert!(status.success(), "child {tag} exited {status}");
}
let outcomes: Vec<String> = tags
.iter()
.map(|tag| {
std::fs::read_to_string(dir.path().join(format!("outcome-{tag}")))
.unwrap_or_else(|_| panic!("child {tag} recorded no outcome"))
})
.collect();
for tag in tags {
assert!(
dir.path().join(format!("planned-{tag}")).exists(),
"child {tag} never reached the commit boundary, so nothing was synchronized"
);
}
let committed = outcomes.iter().filter(|o| o.contains("committed")).count();
let conflicted = outcomes.iter().filter(|o| o.contains("conflicted")).count();
assert_eq!(
(committed, conflicted),
(1, 1),
"exactly one refresh may win the generation both planned: {outcomes:?}"
);
let expected = concurrency_oracle(&conn).await;
let actual = concurrency_view_ids(&conn).await;
assert_eq!(
actual.len(),
expected.len(),
"the view holds {} rows, the oracle {}: a losing refresh left rows behind",
actual.len(),
expected.len()
);
assert_eq!(actual, expected, "the view does not match the oracle");
}