use std::sync::Arc;
use arrow_array::{FixedSizeListArray, Int32Array, RecordBatch, RecordBatchIterator};
use arrow_schema::{DataType, Field as ArrowField, Schema as ArrowSchema};
use lance_arrow::FixedSizeListArrayExt;
use lance_index::IndexType;
use lance_index::scalar::{BuiltinIndexType, ScalarIndexParams};
use lance_linalg::distance::MetricType;
use lance_testing::datagen::generate_random_array;
use crate::Dataset;
use crate::dataset::transaction::{Operation, Transaction};
use crate::index::DatasetIndexExt;
use crate::index::vector::VectorIndexParams;
pub const ROWS_PER_FRAGMENT: i32 = 512;
pub const DIMENSION: i32 = 16;
pub const NUM_PARTITIONS: u32 = 4;
fn vector_field(name: &str) -> ArrowField {
ArrowField::new(
name,
DataType::FixedSizeList(
Arc::new(ArrowField::new("item", DataType::Float32, true)),
DIMENSION,
),
false,
)
}
fn random_vectors(rows: i32) -> Arc<FixedSizeListArray> {
Arc::new(
FixedSizeListArray::try_new_from_values(
generate_random_array(rows as usize * DIMENSION as usize),
DIMENSION,
)
.unwrap(),
)
}
pub async fn write_vector_payload_dataset(uri: &str) -> Dataset {
let schema = Arc::new(ArrowSchema::new(vec![
vector_field("vec"),
ArrowField::new("payload", DataType::Int32, false),
]));
let batch = RecordBatch::try_new(
schema.clone(),
vec![
random_vectors(ROWS_PER_FRAGMENT),
Arc::new(Int32Array::from_iter_values(0..ROWS_PER_FRAGMENT)),
],
)
.unwrap();
let reader = RecordBatchIterator::new(vec![Ok(batch)], schema);
Dataset::write(reader, uri, None).await.unwrap()
}
pub async fn write_two_vector_column_dataset(uri: &str) -> Dataset {
let schema = Arc::new(ArrowSchema::new(vec![
vector_field("vec"),
vector_field("payload_vec"),
]));
let batch = RecordBatch::try_new(
schema.clone(),
vec![
random_vectors(ROWS_PER_FRAGMENT),
random_vectors(ROWS_PER_FRAGMENT),
],
)
.unwrap();
let reader = RecordBatchIterator::new(vec![Ok(batch)], schema);
Dataset::write(reader, uri, None).await.unwrap()
}
pub async fn append_vector_payload_rows(dataset: &mut Dataset, rows: i32) {
let schema = Arc::new(ArrowSchema::new(vec![
vector_field("vec"),
ArrowField::new("payload", DataType::Int32, false),
]));
let batch = RecordBatch::try_new(
schema.clone(),
vec![
random_vectors(rows),
Arc::new(Int32Array::from_iter_values(
ROWS_PER_FRAGMENT..ROWS_PER_FRAGMENT + rows,
)),
],
)
.unwrap();
let reader = RecordBatchIterator::new(vec![Ok(batch)], schema);
dataset.append(reader, None).await.unwrap();
}
pub async fn create_ivf_pq_index(dataset: &mut Dataset, column: &str) {
let params = VectorIndexParams::ivf_pq(NUM_PARTITIONS as usize, 8, 2, MetricType::L2, 50);
dataset
.create_index(&[column], IndexType::Vector, None, ¶ms, true)
.await
.unwrap();
}
pub async fn create_btree_index(dataset: &mut Dataset, column: &str, name: Option<&str>) {
let params = ScalarIndexParams::for_builtin(BuiltinIndexType::BTree);
dataset
.create_index(
&[column],
IndexType::BTree,
name.map(str::to_string),
¶ms,
true,
)
.await
.unwrap();
}
pub async fn write_three_int_column_dataset(uri: &str) -> Dataset {
let schema = Arc::new(ArrowSchema::new(vec![
ArrowField::new("a", DataType::Int32, false),
ArrowField::new("b", DataType::Int32, false),
ArrowField::new("carried", DataType::Int32, false),
]));
let column = || Arc::new(Int32Array::from_iter_values(0..64)) as _;
let batch = RecordBatch::try_new(schema.clone(), vec![column(), column(), column()]).unwrap();
let reader = RecordBatchIterator::new(vec![Ok(batch)], schema);
Dataset::write(reader, uri, None).await.unwrap()
}
pub async fn commit_synthetic_covered_index(
dataset: &mut Dataset,
name: &str,
keyed: &str,
carried: &str,
) -> (i32, i32) {
let keyed_id = dataset.schema().field_id(keyed).unwrap();
let carried_id = dataset.schema().field_id(carried).unwrap();
let covered = lance_table::format::IndexMetadata {
uuid: uuid::Uuid::new_v4(),
name: name.to_string(),
fields: vec![keyed_id, carried_id],
covering_fields: vec![carried_id],
dataset_version: dataset.manifest.version,
fragment_bitmap: Some(dataset.fragment_bitmap.as_ref().clone()),
index_details: None,
index_version: 0,
created_at: Some(chrono::Utc::now()),
base_id: None,
files: None,
};
dataset
.apply_commit(
Transaction::new(
dataset.manifest.version,
Operation::CreateIndex {
new_indices: vec![covered],
removed_indices: vec![],
},
None,
),
&Default::default(),
&Default::default(),
)
.await
.unwrap();
(keyed_id, carried_id)
}
pub async fn append_three_int_column_rows(dataset: &mut Dataset, rows: i32) {
let schema = Arc::new(ArrowSchema::new(vec![
ArrowField::new("a", DataType::Int32, false),
ArrowField::new("b", DataType::Int32, false),
ArrowField::new("carried", DataType::Int32, false),
]));
let column = || Arc::new(Int32Array::from_iter_values(64..64 + rows)) as _;
let batch = RecordBatch::try_new(schema.clone(), vec![column(), column(), column()]).unwrap();
let reader = RecordBatchIterator::new(vec![Ok(batch)], schema);
dataset.append(reader, None).await.unwrap();
}
pub async fn declare_covering(dataset: &mut Dataset, keyed: &str, carried: &str) -> (i32, i32) {
let keyed_id = dataset.schema().field_id(keyed).unwrap();
let carried_id = dataset.schema().field_id(carried).unwrap();
let current = dataset.load_indices().await.unwrap();
let plain = current
.iter()
.find(|idx| idx.fields == vec![keyed_id])
.cloned()
.unwrap_or_else(|| panic!("no index keyed on '{keyed}' to declare covering on"));
let covered = lance_table::format::IndexMetadata {
fields: vec![keyed_id, carried_id],
covering_fields: vec![carried_id],
..plain.clone()
};
dataset
.apply_commit(
Transaction::new(
dataset.manifest.version,
Operation::CreateIndex {
new_indices: vec![covered],
removed_indices: vec![plain],
},
None,
),
&Default::default(),
&Default::default(),
)
.await
.unwrap();
(keyed_id, carried_id)
}