diff --git a/src/metric-engine/src/engine/put.rs b/src/metric-engine/src/engine/put.rs index 404b463c3a..850703b030 100644 --- a/src/metric-engine/src/engine/put.rs +++ b/src/metric-engine/src/engine/put.rs @@ -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> { - 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 { 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> { - // 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, - ) -> 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 { + ) -> 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 = 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 { + ) -> 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, Vec) { + ) -> Result<(Vec, Vec, Option)> { // 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 = 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 { @@ -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; 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::>() }; - // 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::>(), - ) - }) - .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 = 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()