use std::{collections::HashMap, sync::Arc};
use datafusion::prelude::{SessionConfig, SessionContext};
use datafusion_execution::{disk_manager::DiskManagerBuilder, runtime_env::RuntimeEnvBuilder};
use datafusion_expr::col;
use futures::TryStreamExt;
use lance_core::ROW_ID;
use lance_datafusion::exec::SessionContextExt;
use crate::{
Error, Result, Table,
arrow::{SendableRecordBatchStream, SendableRecordBatchStreamExt, SimpleRecordBatchStream},
connect,
database::{CreateTableRequest, Database},
dataloader::permutation::{
shuffle::{Shuffler, ShufflerConfig},
split::{SPLIT_ID_COLUMN, SplitStrategy, Splitter},
util::{TemporaryDirectory, rename_column},
},
query::{ExecutableQuery, QueryBase, Select},
};
pub const SRC_ROW_ID_COL: &str = "row_id";
pub const SPLIT_NAMES_CONFIG_KEY: &str = "split_names";
pub const BASE_VERSION_CONFIG_KEY: &str = "base_version";
pub const BASE_BRANCH_CONFIG_KEY: &str = "base_branch";
pub const DEFAULT_MEMORY_LIMIT: usize = 100 * 1024 * 1024;
#[derive(Debug, Clone, Default)]
enum PermutationDestination {
#[default]
Temporary,
Permanent(Arc<dyn Database>, String),
}
#[derive(Debug, Default)]
pub struct PermutationConfig {
split_strategy: SplitStrategy,
split_names: Option<Vec<String>>,
shuffle_strategy: ShuffleStrategy,
filter: Option<String>,
temp_dir: TemporaryDirectory,
destination: PermutationDestination,
}
#[derive(Debug, Clone, Default)]
pub enum ShuffleStrategy {
Random {
seed: Option<u64>,
clump_size: Option<u64>,
},
#[default]
None,
}
pub struct PermutationBuilder {
config: PermutationConfig,
base_table: Table,
}
impl PermutationBuilder {
pub fn new(base_table: Table) -> Self {
Self {
config: PermutationConfig::default(),
base_table,
}
}
pub fn with_split_strategy(
mut self,
split_strategy: SplitStrategy,
split_names: Option<Vec<String>>,
) -> Self {
self.config.split_strategy = split_strategy;
self.config.split_names = split_names;
self
}
pub fn with_shuffle_strategy(mut self, shuffle_strategy: ShuffleStrategy) -> Self {
self.config.shuffle_strategy = shuffle_strategy;
self
}
pub fn with_filter(mut self, filter: String) -> Self {
self.config.filter = Some(filter);
self
}
pub fn with_temp_dir(mut self, temp_dir: TemporaryDirectory) -> Self {
self.config.temp_dir = temp_dir;
self
}
pub fn persist(mut self, database: Arc<dyn Database>, table_name: String) -> Self {
self.config.destination = PermutationDestination::Permanent(database, table_name);
self
}
async fn sort_by_column(
&self,
data: SendableRecordBatchStream,
column: &str,
) -> Result<SendableRecordBatchStream> {
let memory_limit = std::env::var("LANCEDB_PERM_BUILDER_MEMORY_LIMIT")
.unwrap_or_else(|_| DEFAULT_MEMORY_LIMIT.to_string())
.parse::<usize>()
.unwrap_or_else(|_| {
log::error!(
"Failed to parse LANCEDB_PERM_BUILDER_MEMORY_LIMIT, using default: {}",
DEFAULT_MEMORY_LIMIT
);
DEFAULT_MEMORY_LIMIT
});
let ctx = SessionContext::new_with_config_rt(
SessionConfig::default(),
RuntimeEnvBuilder::new()
.with_memory_limit(memory_limit, 1.0)
.with_disk_manager_builder(
DiskManagerBuilder::default()
.with_mode(self.config.temp_dir.to_disk_manager_mode()),
)
.build_arc()
.unwrap(),
);
let df = ctx
.read_one_shot(data.into_df_stream())
.map_err(|e| Error::Other {
message: format!("Failed to setup sort by {}: {}", column, e),
source: Some(e.into()),
})?;
let df_stream = df
.sort_by(vec![col(column)])
.map_err(|e| Error::Other {
message: format!("Failed to plan sort by {}: {}", column, e),
source: Some(e.into()),
})?
.execute_stream()
.await
.map_err(|e| Error::Other {
message: format!("Failed to sort by {}: {}", column, e),
source: Some(e.into()),
})?;
let column = column.to_string();
let schema = df_stream.schema();
let stream = df_stream.map_err(move |e| Error::Other {
message: format!("Failed to execute sort by {}: {}", column, e),
source: Some(e.into()),
});
Ok(Box::pin(SimpleRecordBatchStream { schema, stream }))
}
fn add_config_metadata(
data: SendableRecordBatchStream,
metadata: HashMap<String, String>,
) -> Result<SendableRecordBatchStream> {
let schema = data.schema().as_ref().clone().with_metadata(metadata);
let schema = Arc::new(schema);
let schema_clone = schema.clone();
let stream = data.map_ok(move |batch| batch.with_schema(schema.clone()).unwrap());
Ok(Box::pin(SimpleRecordBatchStream {
schema: schema_clone,
stream,
}))
}
pub async fn build(mut self) -> Result<Table> {
if let Some(snapshot) = self
.base_table
.base_table()
.snapshot_at_current_version()
.await?
{
self.base_table = Table::from(snapshot);
}
match self.base_table.base_table().get_lsm_write_spec().await {
Ok(Some(_)) => {
return Err(Error::NotSupported {
message: "the data loader does not support tables with an LSM write \
spec: rows that have not been flushed to the base table \
have no row id, so a permutation cannot reference them"
.to_string(),
});
}
Ok(None) => {}
Err(Error::NotSupported { .. }) => {}
Err(err) => return Err(err),
}
let base_version = self.base_table.version().await?;
let base_branch = self.base_table.current_branch();
let mut rows = self.base_table.query().select(Select::columns(&[ROW_ID]));
if let Some(filter) = &self.config.filter {
rows = rows.only_if(filter);
}
let splitter = Splitter::new(
self.config.temp_dir.clone(),
self.config.split_strategy.clone(),
);
let mut needs_sort = !splitter.orders_by_split_id();
rows = splitter.project(rows);
let num_rows = self
.base_table
.count_rows(self.config.filter.clone())
.await? as u64;
let rows = rows.execute().await?;
let rows = if self.base_table.base_table().scan_order_is_deterministic() {
rows
} else {
self.sort_by_column(rows, ROW_ID).await?
};
let split_data = splitter.apply(rows, num_rows).await?;
let shuffled = match self.config.shuffle_strategy {
ShuffleStrategy::None => split_data,
ShuffleStrategy::Random { seed, clump_size } => {
let shuffler = Shuffler::new(ShufflerConfig {
seed,
clump_size,
temp_dir: self.config.temp_dir.clone(),
max_rows_per_file: 10 * 1024 * 1024,
});
shuffler.shuffle(split_data, num_rows).await?
}
};
needs_sort |= !matches!(self.config.shuffle_strategy, ShuffleStrategy::None);
let sorted = if needs_sort {
self.sort_by_column(shuffled, SPLIT_ID_COLUMN).await?
} else {
shuffled
};
let renamed = rename_column(sorted, ROW_ID, SRC_ROW_ID_COL)?;
let mut metadata = HashMap::from([(
BASE_VERSION_CONFIG_KEY.to_string(),
base_version.to_string(),
)]);
if let Some(branch) = &base_branch {
metadata.insert(BASE_BRANCH_CONFIG_KEY.to_string(), branch.clone());
}
if let Some(split_names) = &self.config.split_names {
metadata.insert(
SPLIT_NAMES_CONFIG_KEY.to_string(),
serde_json::to_string(split_names).map_err(|e| Error::Other {
message: format!("Failed to serialize split names: {}", e),
source: Some(e.into()),
})?,
);
}
let streaming_data = Self::add_config_metadata(renamed, metadata)?;
let (name, database) = match &self.config.destination {
PermutationDestination::Permanent(database, table_name) => {
(table_name.as_str(), database.clone())
}
PermutationDestination::Temporary => {
let conn = connect("memory:///").execute().await?;
("permutation", conn.database().clone())
}
};
let create_table_request =
CreateTableRequest::new(name.to_string(), Box::new(streaming_data));
let table = database.create_table(create_table_request).await?;
Ok(Table::new(table, database))
}
}
#[cfg(test)]
mod tests {
use arrow::datatypes::Int32Type;
use lance_datagen::{BatchCount, RowCount};
use crate::{arrow::LanceDbDatagenExt, connect, dataloader::permutation::split::SplitSizes};
use super::*;
#[tokio::test]
async fn test_permutation_table_only_stores_row_id_and_split_id() {
let temp_dir = tempfile::tempdir().unwrap();
let db = connect(temp_dir.path().to_str().unwrap())
.execute()
.await
.unwrap();
let initial_data = lance_datagen::gen_batch()
.col("col_a", lance_datagen::array::step::<Int32Type>())
.col("col_b", lance_datagen::array::step::<Int32Type>())
.into_ldb_stream(RowCount::from(100), BatchCount::from(10));
let data_table = db
.create_table("base_tbl", initial_data)
.execute()
.await
.unwrap();
let permutation_table = PermutationBuilder::new(data_table.clone())
.with_split_strategy(
SplitStrategy::Sequential {
sizes: SplitSizes::Percentages(vec![0.5, 0.5]),
},
None,
)
.with_filter("col_a > 57".to_string())
.build()
.await
.unwrap();
let schema = permutation_table.schema().await.unwrap();
let field_names: Vec<&str> = schema.fields().iter().map(|f| f.name().as_str()).collect();
assert_eq!(
field_names,
vec!["row_id", "split_id"],
"Permutation table should only contain row_id and split_id columns, but found: {:?}",
field_names,
);
}
#[tokio::test]
async fn test_native_scan_order_is_deterministic() {
let temp_dir = tempfile::tempdir().unwrap();
let db = connect(temp_dir.path().to_str().unwrap())
.execute()
.await
.unwrap();
let data = lance_datagen::gen_batch()
.col("col_a", lance_datagen::array::step::<Int32Type>())
.into_ldb_stream(RowCount::from(10), BatchCount::from(1));
let table = db.create_table("t", data).execute().await.unwrap();
assert!(table.base_table().scan_order_is_deterministic());
}
#[cfg(feature = "remote")]
#[tokio::test]
async fn test_remote_permutation_builder_pins_snapshot() {
use std::sync::{
Mutex,
atomic::{AtomicU64, Ordering},
};
use arrow_array::{RecordBatch, UInt64Array};
use arrow_schema::{DataType, Field, Schema};
let row_ids = RecordBatch::try_new(
Arc::new(Schema::new(vec![Field::new(
ROW_ID,
DataType::UInt64,
false,
)])),
vec![Arc::new(UInt64Array::from(vec![100]))],
)
.unwrap();
let mut query_body = Vec::new();
{
let mut writer =
arrow_ipc::writer::FileWriter::try_new(&mut query_body, &row_ids.schema()).unwrap();
writer.write(&row_ids).unwrap();
writer.finish().unwrap();
}
let latest = Arc::new(AtomicU64::new(7));
let expected_snapshot = Arc::new(AtomicU64::new(7));
let planning_versions = Arc::new(Mutex::new(Vec::new()));
let latest_ref = latest.clone();
let expected_snapshot_ref = expected_snapshot.clone();
let planning_versions_ref = planning_versions.clone();
let table = Table::new_with_handler("remote_base", move |request| {
let path = request.url().path();
let body = request
.body()
.and_then(|body| body.as_bytes())
.map(|body| serde_json::from_slice::<serde_json::Value>(body).unwrap());
match path {
"/v1/table/remote_base/describe/" => {
let requested = body.as_ref().and_then(|body| body["version"].as_u64());
let version = requested.unwrap_or_else(|| latest_ref.load(Ordering::SeqCst));
http::Response::builder()
.status(200)
.body(
format!(r#"{{"version":{version},"schema":{{"fields":[]}}}}"#)
.into_bytes(),
)
.unwrap()
}
"/v1/table/remote_base/get_lsm_write_spec/" => http::Response::builder()
.status(200)
.body(br#"{"lsm_write_spec":null}"#.to_vec())
.unwrap(),
"/v1/table/remote_base/count_rows/" => {
let body = body.unwrap();
let version = body["version"].as_u64().unwrap();
assert_eq!(version, expected_snapshot_ref.load(Ordering::SeqCst));
assert_eq!(body["predicate"], "value > 0");
planning_versions_ref.lock().unwrap().push(version);
latest_ref.store(8, Ordering::SeqCst);
http::Response::builder()
.status(200)
.body(b"1".to_vec())
.unwrap()
}
"/v1/table/remote_base/query/" => {
let body = body.unwrap();
let version = body["version"].as_u64().unwrap();
assert_eq!(version, expected_snapshot_ref.load(Ordering::SeqCst));
assert_eq!(body["filter"], "value > 0");
assert_eq!(body["columns"], serde_json::json!([ROW_ID]));
planning_versions_ref.lock().unwrap().push(version);
http::Response::builder()
.status(200)
.header("content-type", "application/vnd.apache.arrow.file")
.body(query_body.clone())
.unwrap()
}
_ => panic!("unexpected request: {path}"),
}
});
let permutation = PermutationBuilder::new(table.clone())
.with_filter("value > 0".to_string())
.build()
.await
.unwrap();
assert_eq!(permutation.count_rows(None).await.unwrap(), 1);
assert_eq!(table.version().await.unwrap(), 8);
expected_snapshot.store(6, Ordering::SeqCst);
table.checkout(6).await.unwrap();
let permutation = PermutationBuilder::new(table.clone())
.with_filter("value > 0".to_string())
.build()
.await
.unwrap();
assert_eq!(permutation.count_rows(None).await.unwrap(), 1);
assert_eq!(table.version().await.unwrap(), 6);
assert_eq!(*planning_versions.lock().unwrap(), vec![7, 7, 6, 6]);
}
#[tokio::test]
async fn test_permutation_records_base_version() {
let temp_dir = tempfile::tempdir().unwrap();
let db = connect(temp_dir.path().to_str().unwrap())
.execute()
.await
.unwrap();
let initial_data = lance_datagen::gen_batch()
.col("col_a", lance_datagen::array::step::<Int32Type>())
.into_ldb_stream(RowCount::from(100), BatchCount::from(2));
let data_table = db
.create_table("base_tbl", initial_data)
.execute()
.await
.unwrap();
let build_version = data_table.version().await.unwrap();
let permutation_table = PermutationBuilder::new(data_table.clone())
.build()
.await
.unwrap();
let recorded = permutation_table
.schema()
.await
.unwrap()
.metadata
.get(BASE_VERSION_CONFIG_KEY)
.expect("permutation should record the base version")
.parse::<u64>()
.unwrap();
assert_eq!(recorded, build_version);
let more_data = lance_datagen::gen_batch()
.col("col_a", lance_datagen::array::step::<Int32Type>())
.into_ldb_stream(RowCount::from(50), BatchCount::from(1));
data_table.add(more_data).execute().await.unwrap();
assert!(data_table.version().await.unwrap() > recorded);
assert_eq!(
permutation_table
.schema()
.await
.unwrap()
.metadata
.get(BASE_VERSION_CONFIG_KEY)
.unwrap()
.parse::<u64>()
.unwrap(),
recorded,
);
}
#[tokio::test]
async fn test_permutation_records_base_branch() {
let temp_dir = tempfile::tempdir().unwrap();
let db = connect(temp_dir.path().to_str().unwrap())
.execute()
.await
.unwrap();
let initial_data = lance_datagen::gen_batch()
.col("col_a", lance_datagen::array::step::<Int32Type>())
.into_ldb_stream(RowCount::from(10), BatchCount::from(1));
let data_table = db
.create_table("base_tbl", initial_data)
.execute()
.await
.unwrap();
let branch = data_table
.create_branch("exp", lance::dataset::refs::Ref::from(("main", 1)))
.await
.unwrap();
let permutation_table = PermutationBuilder::new(branch.clone())
.build()
.await
.unwrap();
let metadata = permutation_table.schema().await.unwrap().metadata.clone();
assert_eq!(
metadata.get(BASE_BRANCH_CONFIG_KEY).map(String::as_str),
Some("exp")
);
let main_permutation = PermutationBuilder::new(data_table.clone())
.build()
.await
.unwrap();
assert!(
!main_permutation
.schema()
.await
.unwrap()
.metadata
.contains_key(BASE_BRANCH_CONFIG_KEY)
);
}
#[tokio::test]
async fn test_build_does_not_pin_the_callers_table() {
let temp_dir = tempfile::tempdir().unwrap();
let db = connect(temp_dir.path().to_str().unwrap())
.execute()
.await
.unwrap();
let initial_data = lance_datagen::gen_batch()
.col("col_a", lance_datagen::array::step::<Int32Type>())
.into_ldb_stream(RowCount::from(100), BatchCount::from(1));
let data_table = db
.create_table("base_tbl", initial_data)
.execute()
.await
.unwrap();
PermutationBuilder::new(data_table.clone())
.build()
.await
.unwrap();
let more_data = lance_datagen::gen_batch()
.col("col_a", lance_datagen::array::step::<Int32Type>())
.into_ldb_stream(RowCount::from(50), BatchCount::from(1));
data_table.add(more_data).execute().await.unwrap();
assert_eq!(data_table.count_rows(None).await.unwrap(), 150);
}
#[tokio::test]
async fn test_permutation_builder() {
let temp_dir = tempfile::tempdir().unwrap();
let db = connect(temp_dir.path().to_str().unwrap())
.execute()
.await
.unwrap();
let initial_data = lance_datagen::gen_batch()
.col("some_value", lance_datagen::array::step::<Int32Type>())
.into_ldb_stream(RowCount::from(100), BatchCount::from(10));
let data_table = db
.create_table("mytbl", initial_data)
.execute()
.await
.unwrap();
let permutation_table = PermutationBuilder::new(data_table.clone())
.with_filter("some_value > 57".to_string())
.with_split_strategy(
SplitStrategy::Random {
seed: Some(42),
sizes: SplitSizes::Percentages(vec![0.05, 0.30]),
clump_size: None,
},
None,
)
.build()
.await
.unwrap();
assert_eq!(permutation_table.count_rows(None).await.unwrap(), 330);
assert_eq!(
permutation_table
.count_rows(Some("split_id = 0".to_string()))
.await
.unwrap(),
47
);
assert_eq!(
permutation_table
.count_rows(Some("split_id = 1".to_string()))
.await
.unwrap(),
283
);
}
#[tokio::test]
async fn test_permutation_rejects_lsm_write_spec() {
use crate::table::LsmWriteSpec;
use arrow_array::{Int32Array, RecordBatchIterator};
use arrow_schema::{DataType, Field, Schema};
let temp_dir = tempfile::tempdir().unwrap();
let db = connect(temp_dir.path().to_str().unwrap())
.execute()
.await
.unwrap();
let schema = Arc::new(Schema::new(vec![Field::new("idx", DataType::Int32, false)]));
let batch = arrow_array::RecordBatch::try_new(
schema.clone(),
vec![Arc::new(Int32Array::from(vec![0, 1, 2, 3]))],
)
.unwrap();
let reader: Box<dyn arrow_array::RecordBatchReader + Send> =
Box::new(RecordBatchIterator::new(vec![Ok(batch)], schema.clone()));
let table = db.create_table("tbl", reader).execute().await.unwrap();
PermutationBuilder::new(table.clone())
.build()
.await
.unwrap();
table.set_unenforced_primary_key(["idx"]).await.unwrap();
table
.set_lsm_write_spec(LsmWriteSpec::unsharded())
.await
.unwrap();
let err = PermutationBuilder::new(table).build().await.unwrap_err();
assert!(
err.to_string().contains("LSM write spec"),
"expected the pre-check to refuse the table, got: {err}"
);
}
}