diff --git a/crates/lance-context-core/src/rollout_store.rs b/crates/lance-context-core/src/rollout_store.rs index dc59693..bc164f7 100644 --- a/crates/lance-context-core/src/rollout_store.rs +++ b/crates/lance-context-core/src/rollout_store.rs @@ -54,9 +54,9 @@ use arrow_array::builder::{ }; use arrow_array::types::Int8Type; use arrow_array::{ - Array, ArrayRef, BooleanArray, DictionaryArray, Float32Array, Int32Array, Int64Array, - Int8Array, LargeBinaryArray, LargeStringArray, ListArray, RecordBatch, RecordBatchIterator, - StringArray, TimestampMicrosecondArray, UInt64Array, + new_null_array, Array, ArrayRef, BooleanArray, DictionaryArray, Float32Array, Int32Array, + Int64Array, Int8Array, LargeBinaryArray, LargeStringArray, ListArray, RecordBatch, + RecordBatchIterator, StringArray, TimestampMicrosecondArray, UInt64Array, }; use arrow_schema::{ArrowError, DataType, Field, Schema, TimeUnit}; use chrono::{DateTime, Utc}; @@ -65,7 +65,9 @@ use lance::dataset::mem_wal::{ DatasetMemWalExt, LsmScanner, ShardManifestStore, ShardSnapshot, ShardWriter, ShardWriterConfig, }; use lance::dataset::optimize::{compact_files, CompactionMetrics, CompactionOptions}; -use lance::dataset::{builder::DatasetBuilder, Dataset, WriteMode, WriteParams}; +use lance::dataset::{ + builder::DatasetBuilder, Dataset, NewColumnTransform, WriteMode, WriteParams, +}; use lance::index::DatasetIndexExt; use lance::io::{ObjectStoreParams, StorageOptionsAccessor}; use lance::{Error as LanceError, Result as LanceResult}; @@ -91,6 +93,14 @@ const DEFAULT_MANIFEST_SCAN_BATCH_SIZE: usize = 16; /// concurrently while collecting observability metrics. const DEFAULT_OBSERVE_CONCURRENCY: usize = 16; +const CLAIM_CHECK_COLUMNS: [&str; 5] = [ + "model_input_string", + "model_output_string", + "rationale", + "problem_text", + "user_metadata", +]; + /// Name of the scalar index on the base table's `id` column. Kept in sync with /// `ContextStore`'s `ID_INDEX_NAME` so both tables index `id` under one name. const ROLLOUT_ID_INDEX_NAME: &str = "id_idx"; @@ -593,6 +603,8 @@ impl RolloutStore { // `add` transparently reopens against the freshly-claimed epoch. self.close().await?; + self.ensure_latest_rollout_schema().await?; + // Resolve each flushed generation to its absolute dataset path and read // all its rows into memory. Record which generation ids we merge so the // drain can remove exactly these and nothing else. @@ -602,6 +614,7 @@ impl RolloutStore { // delete the blob directory after the manifest drain (see below). let mut merged_paths: Vec = Vec::new(); let mut batches: Vec = Vec::new(); + let merge_schema: Arc = Arc::new(self.dataset.schema().into()); for flushed in &manifest.flushed_generations { let gen_uri = format!( "{}/_mem_wal/{}/{}", @@ -612,7 +625,7 @@ impl RolloutStore { let mut stream = gen_dataset.scan().try_into_stream().await?; while let Some(batch) = stream.try_next().await? { if batch.num_rows() > 0 { - batches.push(batch); + batches.push(align_batch_to_schema(batch, merge_schema.clone())?); } } merged_generations.insert(flushed.generation); @@ -621,10 +634,9 @@ impl RolloutStore { // Append the merged rows to the base table. if !batches.is_empty() { - let schema = Arc::new(rollout_schema()); let reader = RecordBatchIterator::new( batches.into_iter().map(Ok::), - schema, + merge_schema, ); let mut params = WriteParams { mode: WriteMode::Append, @@ -698,6 +710,39 @@ impl RolloutStore { Ok(()) } + /// Evolve an older base table to the latest additive rollout schema. + /// + /// Missing nullable columns are added as all-null arrays. Existing unknown + /// columns, type changes, and missing required columns remain hard errors. + async fn ensure_latest_rollout_schema(&mut self) -> LanceResult<()> { + self.dataset.checkout_latest().await?; + + let base_schema: Arc = Arc::new(self.dataset.schema().into()); + let latest_schema = Arc::new(rollout_schema()); + align_batch_to_schema( + RecordBatch::new_empty(base_schema.clone()), + latest_schema.clone(), + )?; + + let missing_fields = latest_schema + .fields() + .iter() + .filter(|field| base_schema.field_with_name(field.name()).is_err()) + .cloned() + .collect::>(); + if !missing_fields.is_empty() { + self.dataset + .add_columns( + NewColumnTransform::AllNulls(Arc::new(Schema::new(missing_fields))), + None, + None, + ) + .await?; + } + + Ok(()) + } + /// Compact the base table's small fragments into larger ones. /// /// Every WAL merge ([`Self::merge_own_shard`]) `append`s a new fragment to @@ -1182,7 +1227,13 @@ impl RolloutStore { ListSource::Wal | ListSource::All => self.wal_shard_snapshots().await?, }; let escaped_id = id.replace('\'', "''"); - let columns = self.non_blob_columns(); + let columns = self + .non_blob_columns() + .into_iter() + .filter(|column| { + source == ListSource::Fragments || !CLAIM_CHECK_COLUMNS.contains(&column.as_str()) + }) + .collect::>(); let refs: Vec<&str> = columns.iter().map(String::as_str).collect(); let scanner = self .lsm_scanner_for_source(source, shard_snapshots) @@ -1828,6 +1879,60 @@ fn is_not_found_error(err: &LanceError) -> bool { text.contains("was not found") || text.contains("not found") || text.contains("NotFound") } +/// Align a flushed-generation batch with the base table schema before append. +/// +/// Nullable columns added after the generation was written are materialized as +/// null arrays. Missing required columns, type changes, and generation-only +/// columns remain hard errors so a merge cannot silently corrupt or discard +/// data. +fn align_batch_to_schema( + batch: RecordBatch, + target_schema: Arc, +) -> LanceResult { + let source_schema = batch.schema(); + + for source_field in source_schema.fields() { + if target_schema.field_with_name(source_field.name()).is_err() { + return Err(ArrowError::SchemaError(format!( + "WAL generation column '{}' does not exist in the base table schema", + source_field.name() + )) + .into()); + } + } + + let columns = target_schema + .fields() + .iter() + .map( + |target_field| match source_schema.column_with_name(target_field.name()) { + Some((index, source_field)) => { + if source_field.data_type() != target_field.data_type() { + return Err(ArrowError::SchemaError(format!( + "WAL generation column '{}' has type {}, expected {}", + target_field.name(), + source_field.data_type(), + target_field.data_type() + )) + .into()); + } + Ok(batch.column(index).clone()) + } + None if target_field.is_nullable() => { + Ok(new_null_array(target_field.data_type(), batch.num_rows())) + } + None => Err(ArrowError::SchemaError(format!( + "required base table column '{}' is missing from the WAL generation", + target_field.name() + )) + .into()), + }, + ) + .collect::>>()?; + + Ok(RecordBatch::try_new(target_schema, columns)?) +} + /// Derive the MemWAL shard UUID a server instance writes to from its stable /// instance id. Deterministic (UUID v5), so the same instance id always maps to /// the same shard across restarts and reopens — the losing/gaining semantics of @@ -2339,6 +2444,78 @@ mod tests { } } + fn pre_claim_check_schema() -> Arc { + Arc::new(Schema::new( + rollout_schema() + .fields() + .iter() + .filter(|field| !CLAIM_CHECK_COLUMNS.contains(&field.name().as_str())) + .cloned() + .collect::>(), + )) + } + + async fn create_empty_dataset(uri: &str, schema: Arc) { + let empty_batch = RecordBatch::new_empty(schema.clone()); + let batches = RecordBatchIterator::new( + vec![Ok::(empty_batch)].into_iter(), + schema, + ); + Dataset::write(batches, uri, None).await.unwrap(); + } + + async fn store_with_legacy_base_and_wal(evolve_base: bool) -> (TempDir, RolloutStore) { + let dir = TempDir::new().unwrap(); + let uri = dir.path().to_string_lossy().to_string(); + let legacy_schema = pre_claim_check_schema(); + create_empty_dataset(&uri, legacy_schema.clone()).await; + + let mut store = RolloutStore::open_with_options( + &uri, + RolloutStoreOptions { + shard_id: Some("pre-claim-check".to_string()), + merge_after_generations: None, + ..Default::default() + }, + ) + .await + .unwrap(); + + let base_batch = store + .records_to_batch(&[assistant_record("legacy-base")]) + .unwrap(); + let base_schema = base_batch.schema(); + let base_reader = RecordBatchIterator::new( + vec![Ok::(base_batch)].into_iter(), + base_schema, + ); + store.dataset.append(base_reader, None).await.unwrap(); + + store.add(&[assistant_record("legacy-wal")]).await.unwrap(); + assert_eq!(flushed_generation_count(&store).await, 1); + store.close().await.unwrap(); + + if evolve_base { + let claim_check_fields = rollout_schema() + .fields() + .iter() + .filter(|field| legacy_schema.field_with_name(field.name()).is_err()) + .cloned() + .collect::>(); + store + .dataset + .add_columns( + NewColumnTransform::AllNulls(Arc::new(Schema::new(claim_check_fields))), + None, + None, + ) + .await + .unwrap(); + } + + (dir, store) + } + fn assert_records_eq(actual: &RolloutRecord, expected: &RolloutRecord) { assert_eq!(actual.id, expected.id); assert_eq!(actual.rollout_id, expected.rollout_id); @@ -2349,6 +2526,11 @@ mod tests { assert_eq!(actual.created_at, expected.created_at); assert_eq!(actual.content, expected.content); assert_eq!(actual.content_type, expected.content_type); + assert_eq!(actual.model_input_string, expected.model_input_string); + assert_eq!(actual.model_output_string, expected.model_output_string); + assert_eq!(actual.rationale, expected.rationale); + assert_eq!(actual.problem_text, expected.problem_text); + assert_eq!(actual.user_metadata, expected.user_metadata); assert_eq!(actual.input_tokens, expected.input_tokens); assert_eq!(actual.output_tokens, expected.output_tokens); assert_eq!(actual.num_input_tokens, expected.num_input_tokens); @@ -3297,6 +3479,113 @@ mod tests { }); } + #[test] + fn cleanup_merges_pre_claim_check_generations_after_schema_evolution() { + let runtime = tokio::runtime::Runtime::new().unwrap(); + runtime.block_on(async { + let (_dir, mut store) = store_with_legacy_base_and_wal(false).await; + + assert_eq!(store.cleanup_own_shard().await.unwrap(), 1); + assert_eq!(flushed_generation_count(&store).await, 0); + let field_paths = store.dataset.schema().field_paths(); + for column in CLAIM_CHECK_COLUMNS { + assert!(field_paths.iter().any(|path| path == column)); + } + + let mut rows = store.list(None, None).await.unwrap(); + rows.sort_by(|left, right| left.id.cmp(&right.id)); + assert_eq!(rows.len(), 2); + assert_eq!(rows[0].id, "legacy-base"); + assert_eq!(rows[1].id, "legacy-wal"); + for row in rows { + assert!(row.model_input_string.is_none()); + assert!(row.model_output_string.is_none()); + assert!(row.rationale.is_none()); + assert!(row.problem_text.is_none()); + assert!(row.user_metadata.is_none()); + } + }); + } + + #[test] + fn latest_schema_alignment_preserves_claim_check_values() { + let runtime = tokio::runtime::Runtime::new().unwrap(); + runtime.block_on(async { + let legacy_dir = TempDir::new().unwrap(); + let legacy_uri = legacy_dir.path().to_string_lossy().to_string(); + create_empty_dataset(&legacy_uri, pre_claim_check_schema()).await; + let mut legacy_store = RolloutStore::open(&legacy_uri).await.unwrap(); + + let current_dir = TempDir::new().unwrap(); + let current_uri = current_dir.path().to_string_lossy().to_string(); + let current_store = RolloutStore::open(¤t_uri).await.unwrap(); + let mut record = assistant_record("current-generation"); + record.model_input_string = Some("input".to_string()); + record.model_output_string = Some("output".to_string()); + record.rationale = Some("reason".to_string()); + record.problem_text = Some("problem".to_string()); + record.user_metadata = Some(r#"{"source":"worker"}"#.to_string()); + let generation_batch = current_store.records_to_batch(&[record]).unwrap(); + + legacy_store.ensure_latest_rollout_schema().await.unwrap(); + let merge_schema: Arc = Arc::new(legacy_store.dataset.schema().into()); + let aligned = align_batch_to_schema(generation_batch, merge_schema.clone()).unwrap(); + let reader = RecordBatchIterator::new( + vec![Ok::(aligned)].into_iter(), + merge_schema, + ); + legacy_store.dataset.append(reader, None).await.unwrap(); + + let merged = legacy_store + .get_by_id_source("current-generation", ListSource::Fragments) + .await + .unwrap() + .unwrap(); + assert_eq!(merged.model_input_string.as_deref(), Some("input")); + assert_eq!(merged.model_output_string.as_deref(), Some("output")); + assert_eq!(merged.rationale.as_deref(), Some("reason")); + assert_eq!(merged.problem_text.as_deref(), Some("problem")); + assert_eq!( + merged.user_metadata.as_deref(), + Some(r#"{"source":"worker"}"#) + ); + }); + } + + #[test] + fn get_by_id_reads_pre_claim_check_base_and_wal_after_schema_evolution() { + let runtime = tokio::runtime::Runtime::new().unwrap(); + runtime.block_on(async { + let (_dir, store) = store_with_legacy_base_and_wal(true).await; + + let mut base = store.get_by_id("legacy-base").await.unwrap().unwrap(); + let wal = store.get_by_id("legacy-wal").await.unwrap().unwrap(); + + assert!(store + .get_by_id_source("legacy-base", ListSource::Fragments) + .await + .unwrap() + .is_some()); + assert!(store + .get_by_id_source("legacy-wal", ListSource::Wal) + .await + .unwrap() + .is_some()); + + for row in [&base, &wal] { + assert!(row.model_input_string.is_none()); + assert!(row.model_output_string.is_none()); + assert!(row.rationale.is_none()); + assert!(row.problem_text.is_none()); + assert!(row.user_metadata.is_none()); + } + + base.id.clone_from(&wal.id); + assert_records_eq(&base, &wal); + assert_eq!(base.binary_payload, wal.binary_payload); + }); + } + #[test] fn create_id_zonemap_index_builds_and_is_idempotent() { // Building the ZoneMap index on `id` must succeed even though the