fix(metric-engine): require uniform WAL policy in batch puts

Signed-off-by: WenyXu <wenymedia@gmail.com>
This commit is contained in:
WenyXu
2026-09-10 03:11:28 +00:00
parent e14377ed39
commit 82c59c576e
+145 -213
View File
@@ -143,76 +143,23 @@ impl MetricEngineInner {
let primary_key_encoding = self.get_primary_key_encoding(data_region_id)?;
// Validate all requests
let partition_expr_version = self
.validate_batch_requests(physical_region_id, &mut requests)
self.validate_batch_requests(physical_region_id, &mut requests)
.await?;
// Keep the common uniform-policy path allocation-free beyond the merge.
if requests
.iter()
.all(|(_, request)| request.skip_wal == requests[0].1.skip_wal)
{
let (request, affected_rows) = self.merge_batch(
physical_region_id,
primary_key_encoding,
requests,
partition_expr_version,
)?;
self.data_region
.write_data(data_region_id, RegionRequest::Put(request))
.await?;
return Ok(affected_rows);
}
// Merge requests according to encoding strategy
let (merged_request, total_affected_rows) = match primary_key_encoding {
PrimaryKeyEncoding::Sparse => self.merge_sparse_batch(physical_region_id, requests)?,
PrimaryKeyEncoding::Dense => self.merge_dense_batch(data_region_id, requests)?,
};
// Only merge consecutive requests with the same WAL policy, preserving
// write order even when the same logical region occurs more than once.
// Prepare every batch before writing so a merge error cannot partially
// apply this physical-region group.
let batches = Self::split_batch_by_wal_policy(requests);
let mut merged_requests = Vec::with_capacity(batches.len());
let mut total_affected_rows = 0;
for requests in batches {
let (request, affected_rows) = self.merge_batch(
physical_region_id,
primary_key_encoding,
requests,
partition_expr_version,
)?;
merged_requests.push(request);
total_affected_rows += affected_rows;
}
for request in merged_requests {
self.data_region
.write_data(data_region_id, RegionRequest::Put(request))
.await?;
}
// Write once to the physical region
self.data_region
.write_data(data_region_id, RegionRequest::Put(merged_request))
.await?;
Ok(total_affected_rows)
}
fn split_batch_by_wal_policy(
requests: Vec<(RegionId, RegionPutRequest)>,
) -> Vec<Vec<(RegionId, RegionPutRequest)>> {
let run_count = requests
.chunk_by(|a, b| a.1.skip_wal == b.1.skip_wal)
.count();
let mut batches = Vec::with_capacity(run_count);
let mut requests = requests.into_iter();
while let Some(first) = requests.next() {
let remaining_run_len = requests
.as_slice()
.iter()
.take_while(|(_, request)| request.skip_wal == first.1.skip_wal)
.count();
let mut batch = Vec::with_capacity(remaining_run_len + 1);
batch.push(first);
batch.extend(requests.by_ref().take(remaining_run_len));
batches.push(batch);
}
batches
}
/// Get primary key encoding for a data region.
fn get_primary_key_encoding(&self, data_region_id: RegionId) -> Result<PrimaryKeyEncoding> {
let state = self.state.read().unwrap();
@@ -223,47 +170,12 @@ impl MetricEngineInner {
})
}
/// Validates all requests and returns the physical batch's explicit partition version.
/// Validates all requests in a batch.
async fn validate_batch_requests(
&self,
physical_region_id: RegionId,
requests: &mut [(RegionId, RegionPutRequest)],
) -> Result<Option<u64>> {
// Preserve the original physical-batch version boundary before splitting
// by WAL policy. A conflict must fail before any data batch is written.
let mut merged_version = None;
for (_, request) in requests.iter() {
if let Some(version) = request.partition_expr_version {
ensure!(
merged_version.is_none_or(|merged| merged == version),
InvalidRequestSnafu {
region_id: physical_region_id,
reason: "inconsistent partition expr version in batch"
}
);
merged_version = Some(version);
}
}
for (logical_region_id, request) in requests {
self.verify_rows(
*logical_region_id,
physical_region_id,
&mut request.rows,
true,
)
.await?;
}
Ok(merged_version)
}
/// Merges one WAL-policy run and attaches the physical batch's write options.
fn merge_batch(
&self,
physical_region_id: RegionId,
encoding: PrimaryKeyEncoding,
requests: Vec<(RegionId, RegionPutRequest)>,
partition_expr_version: Option<u64>,
) -> Result<(RegionPutRequest, AffectedRows)> {
) -> Result<()> {
let skip_wal = requests
.first()
.is_some_and(|(_, request)| request.skip_wal);
@@ -276,28 +188,16 @@ impl MetricEngineInner {
reason: "inconsistent WAL policy in batch"
}
);
let (rows, hint) = match encoding {
PrimaryKeyEncoding::Sparse => (
self.merge_sparse_batch(physical_region_id, requests)?,
Some(WriteHint {
primary_key_encoding: PrimaryKeyEncodingProto::Sparse.into(),
}),
),
PrimaryKeyEncoding::Dense => (
self.merge_dense_batch(to_data_region_id(physical_region_id), requests)?,
None,
),
};
let affected_rows = rows.rows.len() as AffectedRows;
Ok((
RegionPutRequest {
rows,
hint,
skip_wal,
partition_expr_version,
},
affected_rows,
))
for (logical_region_id, request) in requests {
self.verify_rows(
*logical_region_id,
physical_region_id,
&mut request.rows,
true,
)
.await?;
}
Ok(())
}
/// Merges multiple requests using sparse primary key encoding.
@@ -305,11 +205,29 @@ impl MetricEngineInner {
&self,
physical_region_id: RegionId,
requests: Vec<(RegionId, RegionPutRequest)>,
) -> Result<Rows> {
) -> Result<(RegionPutRequest, AffectedRows)> {
let skip_wal = requests
.first()
.is_some_and(|(_, request)| request.skip_wal);
let total_rows: usize = requests.iter().map(|(_, req)| req.rows.rows.len()).sum();
let mut modified_requests = Vec::with_capacity(requests.len());
let mut total_affected_rows: AffectedRows = 0;
let mut merged_version: Option<u64> = None;
for (logical_region_id, mut request) in requests {
if let Some(request_version) = request.partition_expr_version {
if let Some(merged_version) = merged_version {
ensure!(
merged_version == request_version,
InvalidRequestSnafu {
region_id: physical_region_id,
reason: "inconsistent partition expr version in batch"
}
);
} else {
merged_version = Some(request_version);
}
}
self.modify_rows(
physical_region_id,
logical_region_id.table_id(),
@@ -317,6 +235,8 @@ impl MetricEngineInner {
PrimaryKeyEncoding::Sparse,
)?;
let row_count = request.rows.rows.len();
total_affected_rows += row_count as AffectedRows;
modified_requests.push(request.rows);
}
@@ -327,10 +247,19 @@ impl MetricEngineInner {
merged_rows.extend(Self::align_rows_to_schema(rows, &schema));
}
Ok(Rows {
schema,
rows: merged_rows,
})
let merged_request = RegionPutRequest {
skip_wal,
rows: Rows {
schema,
rows: merged_rows,
},
hint: Some(WriteHint {
primary_key_encoding: PrimaryKeyEncodingProto::Sparse.into(),
}),
partition_expr_version: merged_version,
};
Ok((merged_request, total_affected_rows))
}
/// Merges multiple requests using dense primary key encoding.
@@ -342,13 +271,17 @@ impl MetricEngineInner {
&self,
data_region_id: RegionId,
requests: Vec<(RegionId, RegionPutRequest)>,
) -> Result<Rows> {
) -> Result<(RegionPutRequest, AffectedRows)> {
let skip_wal = requests
.first()
.is_some_and(|(_, request)| request.skip_wal);
// Build union schema from all requests
let merged_schema =
Self::build_union_schema(requests.iter().map(|(_, req)| req.rows.schema.as_slice()));
// Align all rows to the merged schema and collect table_ids
let (merged_rows, table_ids) = Self::align_requests_to_schema(requests, &merged_schema);
let (merged_rows, table_ids, merged_version) =
Self::align_requests_to_schema(requests, &merged_schema)?;
// Batch-modify all rows (add __table_id and __tsid columns)
let final_rows = {
@@ -376,7 +309,14 @@ impl MetricEngineInner {
)?
};
Ok(final_rows)
let merged_request = RegionPutRequest {
skip_wal,
rows: final_rows,
hint: None,
partition_expr_version: merged_version,
};
Ok((merged_request, table_ids.len() as AffectedRows))
}
fn build_union_schema<'a>(
@@ -399,20 +339,34 @@ impl MetricEngineInner {
fn align_requests_to_schema(
requests: Vec<(RegionId, RegionPutRequest)>,
merged_schema: &[ColumnSchema],
) -> (Vec<Row>, Vec<TableId>) {
) -> Result<(Vec<Row>, Vec<TableId>, Option<u64>)> {
// Pre-calculate total capacity
let total_rows: usize = requests.iter().map(|(_, req)| req.rows.rows.len()).sum();
let mut merged_rows = Vec::with_capacity(total_rows);
let mut table_ids = Vec::with_capacity(total_rows);
let mut merged_version: Option<u64> = None;
for (logical_region_id, request) in requests {
if let Some(request_version) = request.partition_expr_version {
if let Some(merged_version) = merged_version {
ensure!(
merged_version == request_version,
InvalidRequestSnafu {
region_id: logical_region_id,
reason: "inconsistent partition expr version in batch"
}
);
} else {
merged_version = Some(request_version);
}
}
let table_id = logical_region_id.table_id();
let row_count = request.rows.rows.len();
merged_rows.extend(Self::align_rows_to_schema(request.rows, merged_schema));
table_ids.extend(std::iter::repeat_n(table_id, row_count));
}
(merged_rows, table_ids)
Ok((merged_rows, table_ids, merged_version))
}
fn align_rows_to_schema(rows: Rows, merged_schema: &[ColumnSchema]) -> Vec<Row> {
@@ -892,7 +846,7 @@ mod tests {
}
#[tokio::test]
async fn test_mixed_wal_batch_partition_versions() {
async fn test_batch_partition_versions() {
for encoding in ["sparse", "dense"] {
let env = TestEnv::new().await;
let physical_region_id = env.default_physical_region_id();
@@ -908,7 +862,7 @@ mod tests {
let build_requests = |versions: [Option<u64>; 3]| {
versions
.into_iter()
.zip([false, true, false])
.zip([false; 3])
.map(|(partition_expr_version, skip_wal)| {
(
logical_region_id,
@@ -925,8 +879,7 @@ mod tests {
})
.collect::<Vec<_>>()
};
// The first run has no version and could otherwise be written
// before a later run reveals the conflicting explicit versions.
// Conflicting explicit versions must fail before any data is written.
let err = env
.metric()
.inner
@@ -952,25 +905,32 @@ mod tests {
([Some(7), None, Some(7)], Some(7)),
] {
let mut requests = build_requests(versions);
let actual = env
.metric()
let engine = env.metric();
engine
.inner
.validate_batch_requests(physical_region_id, &mut requests)
.await
.unwrap();
assert_eq!(actual, expected);
for ((_, request), original_version) in requests.iter().zip(versions) {
assert_eq!(request.partition_expr_version, original_version);
let (merged, _) = match encoding {
"sparse" => engine
.inner
.merge_sparse_batch(physical_region_id, requests),
"dense" => engine
.inner
.merge_dense_batch(to_data_region_id(physical_region_id), requests),
_ => unreachable!(),
}
.unwrap();
assert_eq!(merged.partition_expr_version, expected);
}
}
}
#[tokio::test]
async fn test_put_skip_wal_mixed_batch_recovery() {
async fn test_put_skip_wal_batch_recovery() {
for encoding in ["sparse", "dense"] {
// Paired runs differ only in the middle request's WAL policy.
for middle_skip_wal in [false, true] {
// Paired runs differ only in the batch's WAL policy.
for skip_wal in [false, true] {
let env = TestEnv::new().await;
let engine = env.metric();
engine.inner.flush_task.stop().await.unwrap();
@@ -991,13 +951,14 @@ mod tests {
.await;
let metadata_before = engine.get_metadata(logical_region_id).await.unwrap();
let requests = [false, middle_skip_wal, false].into_iter().enumerate().map(
|(index, skip_wal)| {
let requests = [skip_wal; 3]
.into_iter()
.enumerate()
.map(|(index, skip_wal)| {
let timestamp = index as i64 + 1;
let value = timestamp as f64 * 10.0;
// Every request updates the same key at timestamp zero and
// also inserts a distinct key, exposing both reordering and
// accidental WAL-policy propagation to neighboring requests.
// also inserts a distinct key to verify merge order.
let rows = [0, timestamp]
.into_iter()
.map(|timestamp| Row {
@@ -1030,8 +991,7 @@ mod tests {
skip_wal,
},
)
},
);
});
let affected_rows = engine.inner.put_regions_batch(requests).await.unwrap();
assert_eq!(affected_rows, 6);
assert_eq!(
@@ -1086,63 +1046,20 @@ mod tests {
metadata_before.column_metadatas,
recovered_metadata.column_metadatas
);
let expected = if middle_skip_wal {
vec![(0, 30.0), (1, 10.0), (3, 30.0)]
let expected = if skip_wal {
vec![]
} else {
vec![(0, 30.0), (1, 10.0), (2, 20.0), (3, 30.0)]
};
assert_eq!(
scan_timestamp_values(&reopened, logical_region_id).await,
expected,
"encoding={encoding}, middle_skip_wal={middle_skip_wal}"
"encoding={encoding}, skip_wal={skip_wal}"
);
}
}
}
#[test]
fn test_split_batch_by_wal_policy_preserves_order() {
let region_id = RegionId::new(1024, 1);
let requests = [false, false, true, false, true, true]
.into_iter()
.enumerate()
.map(|(index, skip_wal)| {
(
region_id,
RegionPutRequest {
rows: Rows::default(),
hint: None,
partition_expr_version: Some(index as u64),
skip_wal,
},
)
})
.collect();
let batches = MetricEngineInner::split_batch_by_wal_policy(requests);
let actual: Vec<_> = batches
.iter()
.map(|batch| {
(
batch[0].1.skip_wal,
batch
.iter()
.map(|(_, request)| request.partition_expr_version.unwrap())
.collect::<Vec<_>>(),
)
})
.collect();
assert_eq!(
actual,
vec![
(false, vec![0, 1]),
(true, vec![2]),
(false, vec![3]),
(true, vec![4, 5])
]
);
assert!(MetricEngineInner::split_batch_by_wal_policy(vec![]).is_empty());
}
fn assert_merged_schema(rows: &Rows, expect_sparse: bool) {
let column_names: HashSet<String> = rows
.schema
@@ -1305,11 +1222,17 @@ mod tests {
} else {
PrimaryKeyEncoding::Dense
};
let (merged_request, _) = env
.metric()
.inner
.merge_batch(physical_region_id, encoding, build_requests(), None)
.unwrap();
let merge = |requests| match encoding {
PrimaryKeyEncoding::Sparse => env
.metric()
.inner
.merge_sparse_batch(physical_region_id, requests),
PrimaryKeyEncoding::Dense => env
.metric()
.inner
.merge_dense_batch(data_region_id, requests),
};
let (merged_request, _) = merge(build_requests()).unwrap();
if expect_sparse {
assert_eq!(
merged_request.hint.as_ref().unwrap().primary_key_encoding,
@@ -1325,25 +1248,34 @@ mod tests {
let mut requests = build_requests();
for (_, request) in &mut requests {
request.skip_wal = skip_wal;
request.partition_expr_version = Some(7);
}
let (merged, affected_rows) = env
.metric()
.inner
.merge_batch(physical_region_id, encoding, requests, Some(7))
.unwrap();
let (merged, affected_rows) = merge(requests).unwrap();
assert_eq!(merged.skip_wal, skip_wal);
assert_eq!(merged.partition_expr_version, Some(7));
assert_eq!(affected_rows, 5);
}
let mut mixed_requests = build_requests();
mixed_requests[1].1.skip_wal = true;
assert!(
env.metric()
for policies in [[false, true], [true, false]] {
let mut mixed_requests = build_requests();
for ((_, request), skip_wal) in mixed_requests.iter_mut().zip(policies) {
request.skip_wal = skip_wal;
}
let err = env
.metric()
.inner
.merge_batch(physical_region_id, encoding, mixed_requests, None)
.is_err()
);
.put_regions_batch(mixed_requests.into_iter())
.await
.unwrap_err();
assert!(err.to_string().contains("inconsistent WAL policy in batch"));
for logical_region_id in [logical_region_1, logical_region_2] {
assert!(
scan_timestamp_values(&env.metric(), logical_region_id)
.await
.is_empty()
);
}
}
let affected_rows = env
.metric()