Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
157 changes: 35 additions & 122 deletions rust/lance/src/dataset/transaction.rs
Original file line number Diff line number Diff line change
Expand Up @@ -2549,59 +2549,44 @@ impl Transaction {
Ok(())
}

fn duplicate_rewrite_error(version: u64, fragment_id: u64) -> Error {
Error::commit_conflict_source(
version,
format!(
"a rewrite operation attempts to replace fragment id={} more than once",
fragment_id
)
.into(),
)
}

fn handle_rewrite_fragments(
final_fragments: &mut Vec<Fragment>,
groups: &[RewriteGroup],
fragment_id: &mut u64,
version: u64,
_next_row_id: Option<&u64>,
) -> Result<()> {
// Index fragment positions up front so each rewrite group is located in
// O(1). Scanning and splicing/retaining the fragment list per group is
// O(groups * fragments) and dominates commit time on manifests with
// millions of fragments.
let positions: HashMap<u64, usize> = final_fragments
.iter()
.enumerate()
.map(|(position, fragment)| (fragment.id, position))
.collect();

// Contiguous groups have their new fragments spliced in at the position
// of their first old fragment; non-contiguous groups have their new
// fragments appended at the end.
let mut removed = vec![false; final_fragments.len()];
let mut replacements: HashMap<usize, Vec<Fragment>> = HashMap::new();
let mut appended: Vec<Fragment> = Vec::new();

for group in groups {
let start = *positions.get(&group.old_fragments[0].id).ok_or_else(|| {
Error::commit_conflict_source(
version,
format!(
"dataset does not contain a fragment a rewrite operation wants to replace: id={}",
group.old_fragments[0].id
)
.into(),
)
})?;

let contiguous = start + group.old_fragments.len() <= final_fragments.len()
&& group
.old_fragments
// If the old fragments are contiguous, find the range
let replace_range = {
let start = final_fragments
.iter()
.enumerate()
.all(|(i, old)| final_fragments[start + i].id == old.id);
.find(|(_, f)| f.id == group.old_fragments[0].id)
.ok_or_else(|| {
Error::commit_conflict_source(
version,
format!(
"dataset does not contain a fragment a rewrite operation wants to replace: id={}",
group.old_fragments[0].id
)
.into(),
)
})?
.0;

// Verify old_fragments matches contiguous range
let mut i = 1;
loop {
if i == group.old_fragments.len() {
break Some(start..start + i);
}
if final_fragments[start + i].id != group.old_fragments[i].id {
break None;
}
i += 1;
}
};

let new_fragments = Self::fragments_with_ids(group.new_fragments.clone(), fragment_id)
.collect::<Vec<_>>();
Expand All @@ -2610,39 +2595,17 @@ impl Transaction {
// (recalc_versions_for_rewritten_fragments) which preserves version information
// from the original fragments. We don't modify it here.

if contiguous {
for (i, old_fragment) in group.old_fragments.iter().enumerate() {
let position = start + i;
if removed[position] {
return Err(Self::duplicate_rewrite_error(version, old_fragment.id));
}
removed[position] = true;
}
replacements.insert(start, new_fragments);
if let Some(replace_range) = replace_range {
// Efficiently path using slice
final_fragments.splice(replace_range, new_fragments);
} else {
for old_fragment in group.old_fragments.iter() {
if let Some(&position) = positions.get(&old_fragment.id) {
if removed[position] {
return Err(Self::duplicate_rewrite_error(version, old_fragment.id));
}
removed[position] = true;
}
// Slower path for non-contiguous ranges
for fragment in group.old_fragments.iter() {
final_fragments.retain(|f| f.id != fragment.id);
}
appended.extend(new_fragments);
}
}

let mut rewritten = Vec::with_capacity(final_fragments.len() + appended.len());
for (position, fragment) in std::mem::take(final_fragments).into_iter().enumerate() {
if let Some(new_fragments) = replacements.remove(&position) {
rewritten.extend(new_fragments);
}
if !removed[position] {
rewritten.push(fragment);
final_fragments.extend(new_fragments);
}
}
rewritten.extend(appended);
*final_fragments = rewritten;
Ok(())
}

Expand Down Expand Up @@ -3564,56 +3527,6 @@ mod tests {
assert_eq!(final_fragments, expected_fragments);
}

#[test]
fn test_rewrite_fragments_missing_fragment() {
let mut final_fragments: Vec<Fragment> = (0..4).map(Fragment::new).collect();
let rewrite_groups = vec![RewriteGroup {
old_fragments: vec![Fragment::new(7)],
new_fragments: vec![Fragment::new(0)],
}];

let mut fragment_id = 10;
let result = Transaction::handle_rewrite_fragments(
&mut final_fragments,
&rewrite_groups,
&mut fragment_id,
42,
None,
);

let error = result.unwrap_err();
assert!(matches!(error, Error::CommitConflict { .. }));
assert!(error.to_string().contains("id=7"));
}

#[test]
fn test_rewrite_fragments_duplicate_fragment() {
let mut final_fragments: Vec<Fragment> = (0..4).map(Fragment::new).collect();
let rewrite_groups = vec![
RewriteGroup {
old_fragments: vec![Fragment::new(1), Fragment::new(2)],
new_fragments: vec![Fragment::new(0)],
},
RewriteGroup {
old_fragments: vec![Fragment::new(2), Fragment::new(3)],
new_fragments: vec![Fragment::new(0)],
},
];

let mut fragment_id = 10;
let result = Transaction::handle_rewrite_fragments(
&mut final_fragments,
&rewrite_groups,
&mut fragment_id,
42,
None,
);

let error = result.unwrap_err();
assert!(matches!(error, Error::CommitConflict { .. }));
assert!(error.to_string().contains("id=2"));
}

#[test]
fn test_merge_fragments_valid() {
use arrow_schema::{DataType, Field as ArrowField, Schema as ArrowSchema};
Expand Down
Loading