feat: support request-level insert WAL skipping (#9088)

* refactor: add skip_wal fields to internal write requests

Signed-off-by: WenyXu <wenymedia@gmail.com>

* chore(deps): update greptime-proto for insert skip_wal

Signed-off-by: WenyXu <wenymedia@gmail.com>

* feat(mito): support request-level WAL skipping

Signed-off-by: WenyXu <wenymedia@gmail.com>

* feat(metric-engine): handle request-level WAL policies

Signed-off-by: WenyXu <wenymedia@gmail.com>

* feat: propagate insert WAL policy through query context

Signed-off-by: WenyXu <wenymedia@gmail.com>

* feat: support session-level insert WAL policy via SET

Signed-off-by: WenyXu <wenymedia@gmail.com>

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

Signed-off-by: WenyXu <wenymedia@gmail.com>

* test(metric-engine): simplify WAL policy coverage

Signed-off-by: WenyXu <wenymedia@gmail.com>

* refactor(servers): simplify gRPC hint extraction

Signed-off-by: WenyXu <wenymedia@gmail.com>

* test: flatten Mito and Metric WAL scenario orchestration

Signed-off-by: WenyXu <wenymedia@gmail.com>

* refactor: separate WAL and memtable-only mutations

Signed-off-by: WenyXu <wenymedia@gmail.com>

* refactor: carry skip-WAL policy in table insert requests

Signed-off-by: WenyXu <wenymedia@gmail.com>

* test: flatten skip-WAL policy cases

Signed-off-by: WenyXu <wenymedia@gmail.com>

* refactor: clarify WAL notifier naming

Signed-off-by: WenyXu <wenymedia@gmail.com>

* chore: pin merged skip-WAL proto revision

Signed-off-by: WenyXu <wenymedia@gmail.com>

---------

Signed-off-by: WenyXu <wenymedia@gmail.com>
This commit is contained in:
Weny Xu
2026-09-10 11:40:51 +00:00
committed by GitHub
parent 536f42e4a2
commit fa128adb8e
46 changed files with 1750 additions and 98 deletions
Generated
+1 -1
View File
@@ -6097,7 +6097,7 @@ dependencies = [
[[package]]
name = "greptime-proto"
version = "0.1.0"
source = "git+https://github.com/GreptimeTeam/greptime-proto.git?rev=32f467fa2ba2b3588a58381a24af83de09fbb00a#32f467fa2ba2b3588a58381a24af83de09fbb00a"
source = "git+https://github.com/GreptimeTeam/greptime-proto.git?rev=b74667d17827481dbb6222fa0326163627211557#b74667d17827481dbb6222fa0326163627211557"
dependencies = [
"prost 0.14.1",
"prost-types 0.14.1",
+1 -1
View File
@@ -159,7 +159,7 @@ fs2 = "0.4"
fst = "0.4.7"
futures = "0.3"
futures-util = "0.3"
greptime-proto = { git = "https://github.com/GreptimeTeam/greptime-proto.git", rev = "32f467fa2ba2b3588a58381a24af83de09fbb00a" }
greptime-proto = { git = "https://github.com/GreptimeTeam/greptime-proto.git", rev = "b74667d17827481dbb6222fa0326163627211557" }
hex = "0.4"
hostname = "0.4.0"
http = "1"
+1
View File
@@ -2018,6 +2018,7 @@ mod tests {
assert!(RegionServerInner::is_ingest_request(&RegionRequest::Put(
RegionPutRequest {
skip_wal: false,
rows: rows(),
hint: None,
partition_expr_version: None,
+1
View File
@@ -124,6 +124,7 @@ impl flow_server::Flow for FlowService {
api::v1::region::InsertRequest {
region_id: insert.region_id,
rows: insert.rows,
skip_wal: false,
partition_expr_version: insert.partition_expr_version,
}
})
+1
View File
@@ -636,6 +636,7 @@ mod test {
.handle_request(
logical_region_id,
RegionRequest::Put(RegionPutRequest {
skip_wal: false,
rows: Rows {
schema: row_schema_with_tags(&["job"]),
rows: build_rows(1, 5),
@@ -599,6 +599,7 @@ mod tests {
.handle_request(
logical_region_id,
RegionRequest::Put(RegionPutRequest {
skip_wal: false,
rows: api::v1::Rows { schema, rows },
hint: None,
partition_expr_version: None,
+1
View File
@@ -90,6 +90,7 @@ mod tests {
let schema = row_schema_with_tags(&["job"]);
let rows = build_rows(1, 10);
let request = RegionRequest::Put(RegionPutRequest {
skip_wal: false,
rows: Rows { schema, rows },
hint: None,
partition_expr_version: None,
+367 -27
View File
@@ -142,6 +142,8 @@ impl MetricEngineInner {
let data_region_id = to_data_region_id(physical_region_id);
let primary_key_encoding = self.get_primary_key_encoding(data_region_id)?;
// TODO(weny): Consolidate validation and merging to avoid redundant request traversals,
// while ensuring the entire batch is validated before writing.
// Validate all requests
self.validate_batch_requests(physical_region_id, &mut requests)
.await?;
@@ -176,6 +178,18 @@ impl MetricEngineInner {
physical_region_id: RegionId,
requests: &mut [(RegionId, RegionPutRequest)],
) -> Result<()> {
let skip_wal = requests
.first()
.is_some_and(|(_, request)| request.skip_wal);
ensure!(
requests
.iter()
.all(|(_, request)| request.skip_wal == skip_wal),
InvalidRequestSnafu {
region_id: physical_region_id,
reason: "inconsistent WAL policy in batch"
}
);
for (logical_region_id, request) in requests {
self.verify_rows(
*logical_region_id,
@@ -194,6 +208,9 @@ impl MetricEngineInner {
physical_region_id: RegionId,
requests: Vec<(RegionId, RegionPutRequest)>,
) -> 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;
@@ -233,6 +250,7 @@ impl MetricEngineInner {
}
let merged_request = RegionPutRequest {
skip_wal,
rows: Rows {
schema,
rows: merged_rows,
@@ -256,6 +274,9 @@ impl MetricEngineInner {
data_region_id: RegionId,
requests: Vec<(RegionId, RegionPutRequest)>,
) -> 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()));
@@ -291,6 +312,7 @@ impl MetricEngineInner {
};
let merged_request = RegionPutRequest {
skip_wal,
rows: final_rows,
hint: None,
partition_expr_version: merged_version,
@@ -769,25 +791,272 @@ mod tests {
use common_function::utils::partition_expr_version;
use common_query::prelude::{greptime_native_histogram, greptime_timestamp, greptime_value};
use common_recordbatch::RecordBatches;
use datatypes::arrow::array::{Float64Array, TimestampMillisecondArray};
use datatypes::prelude::ConcreteDataType;
use datatypes::schema::{ColumnDefaultConstraint, ColumnSchema};
use datatypes::value::Value as PartitionValue;
use partition::expr::col;
use store_api::metadata::ColumnMetadata;
use store_api::metric_engine_consts::{
DATA_SCHEMA_TABLE_ID_COLUMN_NAME, DATA_SCHEMA_TSID_COLUMN_NAME, PRIMARY_KEY_ENCODING,
DATA_SCHEMA_TABLE_ID_COLUMN_NAME, DATA_SCHEMA_TSID_COLUMN_NAME, METRIC_ENGINE_NAME,
PHYSICAL_TABLE_METADATA_KEY, PRIMARY_KEY_ENCODING,
};
use store_api::path_utils::table_dir;
use store_api::region_engine::RegionEngine;
use store_api::region_request::{
EnterStagingRequest, RegionRequest, StagingPartitionDirective,
EnterStagingRequest, PathType, RegionCloseRequest, RegionOpenRequest, RegionRequest,
StagingPartitionDirective,
};
use store_api::storage::ScanRequest;
use store_api::storage::consts::PRIMARY_KEY_COLUMN_NAME;
use super::*;
use crate::engine::MetricEngine;
use crate::test_util::{self, TestEnv};
async fn scan_timestamp_values(engine: &MetricEngine, region_id: RegionId) -> Vec<(i64, f64)> {
let stream = engine
.scan_to_stream(region_id, ScanRequest::default())
.await
.unwrap();
let batches = RecordBatches::try_collect(stream).await.unwrap();
let mut rows = Vec::new();
for batch in batches.iter() {
let batch = batch.df_record_batch();
let timestamp_index = batch.schema().index_of(greptime_timestamp()).unwrap();
let value_index = batch.schema().index_of(greptime_value()).unwrap();
let timestamps = batch
.column(timestamp_index)
.as_any()
.downcast_ref::<TimestampMillisecondArray>()
.unwrap();
let values = batch
.column(value_index)
.as_any()
.downcast_ref::<Float64Array>()
.unwrap();
rows.extend(
timestamps
.values()
.iter()
.copied()
.zip(values.values().iter().copied()),
);
}
rows.sort_unstable_by_key(|(timestamp, _)| *timestamp);
rows
}
#[tokio::test]
async fn test_batch_partition_versions() {
check_batch_partition_versions("sparse").await;
check_batch_partition_versions("dense").await;
}
async fn check_batch_partition_versions(encoding: &str) {
let env = TestEnv::new().await;
let physical_region_id = env.default_physical_region_id();
let logical_region_id = env.default_logical_region_id();
env.create_physical_region(
physical_region_id,
&TestEnv::default_table_dir(),
vec![(PRIMARY_KEY_ENCODING.to_string(), encoding.to_string())],
)
.await;
create_logical_region_with_tags(&env, physical_region_id, logical_region_id, &["job"])
.await;
let build_requests = |versions: [Option<u64>; 3]| {
versions
.into_iter()
.map(|partition_expr_version| {
(
logical_region_id,
RegionPutRequest {
skip_wal: false,
rows: Rows {
schema: test_util::row_schema_with_tags(&["job"]),
rows: test_util::build_rows(1, 1),
},
hint: None,
partition_expr_version,
},
)
})
.collect::<Vec<_>>()
};
// Conflicting explicit versions must fail before any data is written.
let err = env
.metric()
.inner
.put_regions_batch_single_physical(
physical_region_id,
build_requests([None, Some(10), Some(11)]),
)
.await
.unwrap_err();
assert!(
err.to_string()
.contains("inconsistent partition expr version")
);
assert!(
scan_timestamp_values(&env.metric(), logical_region_id)
.await
.is_empty()
);
for (versions, expected) in [
([None, None, None], None),
([None, Some(7), None], Some(7)),
([Some(7), None, Some(7)], Some(7)),
] {
let mut requests = build_requests(versions);
let engine = env.metric();
engine
.inner
.validate_batch_requests(physical_region_id, &mut requests)
.await
.unwrap();
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_batch_recovery() {
check_put_skip_wal_batch_recovery("sparse", false).await;
check_put_skip_wal_batch_recovery("sparse", true).await;
check_put_skip_wal_batch_recovery("dense", false).await;
check_put_skip_wal_batch_recovery("dense", true).await;
}
async fn check_put_skip_wal_batch_recovery(encoding: &str, skip_wal: bool) {
let env = TestEnv::new().await;
let engine = env.metric();
engine.inner.flush_task.stop().await.unwrap();
let physical_region_id = env.default_physical_region_id();
let logical_region_id = env.default_logical_region_id();
env.create_physical_region(
physical_region_id,
&TestEnv::default_table_dir(),
vec![(PRIMARY_KEY_ENCODING.to_string(), encoding.to_string())],
)
.await;
create_logical_region_with_tags(&env, physical_region_id, logical_region_id, &["job"])
.await;
let metadata_before = engine.get_metadata(logical_region_id).await.unwrap();
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 to verify merge order.
let rows = [0, timestamp]
.into_iter()
.map(|timestamp| Row {
values: vec![
Value {
value_data: Some(ValueData::TimestampMillisecondValue(timestamp)),
},
Value {
value_data: Some(ValueData::F64Value(value)),
},
Value {
value_data: Some(ValueData::StringValue("tag_0".to_string())),
},
],
})
.collect();
(
logical_region_id,
RegionPutRequest {
rows: Rows {
schema: test_util::row_schema_with_tags(&["job"]),
rows,
},
hint: None,
partition_expr_version: None,
skip_wal,
},
)
});
let affected_rows = engine.inner.put_regions_batch(requests).await.unwrap();
assert_eq!(affected_rows, 6);
assert_eq!(
scan_timestamp_values(&engine, logical_region_id).await,
vec![(0, 30.0), (1, 10.0), (2, 20.0), (3, 30.0)]
);
// Neither data nor metadata has an SST to hide missing WAL.
for region_id in [
to_data_region_id(physical_region_id),
crate::utils::to_metadata_region_id(physical_region_id),
] {
let stat = env.mito().region_statistic(region_id).unwrap();
assert!(stat.memtable_size > 0);
assert_eq!(stat.sst_num, 0);
}
engine
.handle_request(
physical_region_id,
RegionRequest::Close(RegionCloseRequest {
flush_on_close: false,
}),
)
.await
.unwrap();
// Recreate the wrapper as well, discarding its metadata cache.
let reopened = MetricEngine::try_new(env.mito(), Default::default()).unwrap();
reopened.inner.flush_task.stop().await.unwrap();
reopened
.handle_request(
physical_region_id,
RegionRequest::Open(RegionOpenRequest {
engine: METRIC_ENGINE_NAME.to_string(),
table_dir: TestEnv::default_table_dir(),
path_type: PathType::Bare,
options: [
(PHYSICAL_TABLE_METADATA_KEY.to_string(), String::new()),
(PRIMARY_KEY_ENCODING.to_string(), encoding.to_string()),
]
.into_iter()
.collect(),
skip_wal_replay: false,
checkpoint: None,
requirements: Default::default(),
}),
)
.await
.unwrap();
let recovered_metadata = reopened.get_metadata(logical_region_id).await.unwrap();
assert_eq!(
metadata_before.column_metadatas,
recovered_metadata.column_metadatas
);
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}, skip_wal={skip_wal}"
);
}
fn assert_merged_schema(rows: &Rows, expect_sparse: bool) {
let column_names: HashSet<String> = rows
.schema
@@ -874,6 +1143,45 @@ mod tests {
.unwrap()
}
fn check_batch_merge_wal_policy(
env: &TestEnv,
physical_region_id: RegionId,
mut requests: Vec<(RegionId, RegionPutRequest)>,
expect_sparse: bool,
skip_wal: bool,
) {
for (_, request) in &mut requests {
request.skip_wal = skip_wal;
}
let (merged_request, affected_rows) = if expect_sparse {
let (merged_request, affected_rows) = env
.metric()
.inner
.merge_sparse_batch(physical_region_id, requests)
.unwrap();
let hint = merged_request
.hint
.as_ref()
.expect("missing sparse write hint");
assert_eq!(
hint.primary_key_encoding,
PrimaryKeyEncodingProto::Sparse as i32
);
(merged_request, affected_rows)
} else {
let (merged_request, affected_rows) = env
.metric()
.inner
.merge_dense_batch(to_data_region_id(physical_region_id), requests)
.unwrap();
assert!(merged_request.hint.is_none());
(merged_request, affected_rows)
};
assert_merged_schema(&merged_request.rows, expect_sparse);
assert_eq!(merged_request.skip_wal, skip_wal);
assert_eq!(affected_rows, 5);
}
async fn run_batch_write_with_schema_variants(
env: &TestEnv,
physical_region_id: RegionId,
@@ -921,6 +1229,7 @@ mod tests {
(
logical_region_1,
RegionPutRequest {
skip_wal: false,
rows: Rows {
schema: schema_1.clone(),
rows: rows_1,
@@ -932,6 +1241,7 @@ mod tests {
(
logical_region_2,
RegionPutRequest {
skip_wal: false,
rows: Rows {
schema: schema_2.clone(),
rows: rows_2,
@@ -943,32 +1253,41 @@ mod tests {
]
};
let merged_request = if expect_sparse {
let (merged_request, _) = env
.metric()
.inner
.merge_sparse_batch(physical_region_id, build_requests())
.unwrap();
let hint = merged_request
.hint
.as_ref()
.expect("missing sparse write hint");
assert_eq!(
hint.primary_key_encoding,
PrimaryKeyEncodingProto::Sparse as i32
);
merged_request
} else {
let (merged_request, _) = env
.metric()
.inner
.merge_dense_batch(data_region_id, build_requests())
.unwrap();
assert!(merged_request.hint.is_none());
merged_request
};
check_batch_merge_wal_policy(
env,
physical_region_id,
build_requests(),
expect_sparse,
false,
);
check_batch_merge_wal_policy(
env,
physical_region_id,
build_requests(),
expect_sparse,
true,
);
assert_merged_schema(&merged_request.rows, expect_sparse);
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
.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()
@@ -1095,6 +1414,7 @@ mod tests {
let schema = test_util::row_schema_with_tags(&["job"]);
let rows = test_util::build_rows(1, 5);
let request = RegionRequest::Put(RegionPutRequest {
skip_wal: false,
rows: Rows { schema, rows },
hint: None,
partition_expr_version: None,
@@ -1170,6 +1490,7 @@ mod tests {
let schema = test_util::row_schema_with_tags(columns);
let rows = test_util::build_rows(3, 100);
let request = RegionRequest::Put(RegionPutRequest {
skip_wal: false,
rows: Rows { schema, rows },
hint: None,
partition_expr_version: None,
@@ -1193,6 +1514,7 @@ mod tests {
let schema = test_util::row_schema_with_tags(&["abc"]);
let rows = test_util::build_rows(1, 100);
let request = RegionRequest::Put(RegionPutRequest {
skip_wal: false,
rows: Rows { schema, rows },
hint: None,
partition_expr_version: None,
@@ -1214,6 +1536,7 @@ mod tests {
let schema = test_util::row_schema_with_tags(&["def"]);
let rows = test_util::build_rows(1, 100);
let request = RegionRequest::Put(RegionPutRequest {
skip_wal: false,
rows: Rows { schema, rows },
hint: None,
partition_expr_version: None,
@@ -1274,6 +1597,7 @@ mod tests {
(
logical_region_1,
RegionPutRequest {
skip_wal: false,
rows: Rows {
schema: schema.clone(),
rows: rows1,
@@ -1285,6 +1609,7 @@ mod tests {
(
logical_region_2,
RegionPutRequest {
skip_wal: false,
rows: Rows {
schema: schema.clone(),
rows: rows2,
@@ -1296,6 +1621,7 @@ mod tests {
(
logical_region_3,
RegionPutRequest {
skip_wal: false,
rows: Rows {
schema: schema.clone(),
rows: rows3,
@@ -1347,6 +1673,7 @@ mod tests {
(
logical_region_1,
RegionPutRequest {
skip_wal: false,
rows: Rows {
schema: schema.clone(),
rows: test_util::build_rows(1, 3),
@@ -1358,6 +1685,7 @@ mod tests {
(
nonexistent_region,
RegionPutRequest {
skip_wal: false,
rows: Rows {
schema: schema.clone(),
rows: test_util::build_rows(1, 2),
@@ -1369,6 +1697,7 @@ mod tests {
(
logical_region_2,
RegionPutRequest {
skip_wal: false,
rows: Rows {
schema: schema.clone(),
rows: test_util::build_rows(1, 5),
@@ -1407,6 +1736,7 @@ mod tests {
let requests = vec![(
physical_region_id,
RegionPutRequest {
skip_wal: false,
rows: Rows {
schema,
rows: test_util::build_rows(1, 1),
@@ -1441,6 +1771,7 @@ mod tests {
(
logical_region_id,
RegionPutRequest {
skip_wal: false,
rows: Rows {
schema: schema.clone(),
rows: test_util::build_rows(1, 1),
@@ -1452,6 +1783,7 @@ mod tests {
(
physical_region_id,
RegionPutRequest {
skip_wal: false,
rows: Rows {
schema,
rows: test_util::build_rows(1, 1),
@@ -1487,6 +1819,7 @@ mod tests {
let requests = vec![(
logical_region_id,
RegionPutRequest {
skip_wal: false,
rows: Rows {
schema,
rows: test_util::build_rows(1, 5),
@@ -1565,6 +1898,7 @@ mod tests {
.handle_request(
logical_region_id,
RegionRequest::Put(RegionPutRequest {
skip_wal: false,
rows,
hint: None,
partition_expr_version: Some(1),
@@ -1607,6 +1941,7 @@ mod tests {
.handle_request(
logical_region_id,
RegionRequest::Put(RegionPutRequest {
skip_wal: false,
rows: rows.clone(),
hint: None,
partition_expr_version: Some(expected_version.wrapping_add(1)),
@@ -1621,6 +1956,7 @@ mod tests {
.handle_request(
logical_region_id,
RegionRequest::Put(RegionPutRequest {
skip_wal: false,
rows: rows.clone(),
hint: None,
partition_expr_version: None,
@@ -1635,6 +1971,7 @@ mod tests {
.handle_request(
logical_region_id,
RegionRequest::Put(RegionPutRequest {
skip_wal: false,
rows,
hint: None,
partition_expr_version: Some(expected_version),
@@ -1703,6 +2040,7 @@ mod tests {
.handle_request(
logical_region_id,
RegionRequest::Put(RegionPutRequest {
skip_wal: false,
rows: Rows { schema, rows },
hint: None,
partition_expr_version: None,
@@ -1759,6 +2097,7 @@ mod tests {
.handle_request(
logical_region_id,
RegionRequest::Put(RegionPutRequest {
skip_wal: false,
rows: Rows { schema, rows },
hint: None,
partition_expr_version: None,
@@ -1813,6 +2152,7 @@ mod tests {
.handle_request(
logical_region_id,
RegionRequest::Put(RegionPutRequest {
skip_wal: false,
rows: Rows { schema, rows },
hint: None,
partition_expr_version: None,
+1
View File
@@ -500,6 +500,7 @@ mod test {
let schema = test_util::row_schema_with_tags(&["job"]);
let put = |rows| {
RegionRequest::Put(RegionPutRequest {
skip_wal: false,
rows: Rows {
schema: schema.clone(),
rows: test_util::build_rows(1, rows),
+33
View File
@@ -585,6 +585,8 @@ impl MetadataRegion {
};
RegionPutRequest {
// Metadata must remain recoverable regardless of the user write policy.
skip_wal: false,
rows,
hint: None,
partition_expr_version: None,
@@ -819,6 +821,37 @@ mod test {
use crate::test_util::TestEnv;
use crate::utils::to_metadata_region_id;
#[test]
fn test_metadata_put_always_writes_wal() {
// Both metadata put paths use this constructor rather than forwarding
// a user insert request, so its WAL policy must always be independent.
for entries in [
vec![],
vec![("region", "")],
vec![("region", ""), ("column", "metadata")],
] {
let request = MetadataRegion::build_put_request_from_iter(
entries
.iter()
.map(|(key, value)| (key.to_string(), value.to_string())),
);
assert!(!request.skip_wal);
assert!(request.hint.is_none());
assert!(request.partition_expr_version.is_none());
let expected_rows = entries
.into_iter()
.map(|(key, value)| {
row(vec![
ValueData::TimestampMillisecondValue(0),
ValueData::StringValue(key.to_string()),
ValueData::StringValue(value.to_string()),
])
})
.collect::<Vec<_>>();
assert_eq!(request.rows.rows, expected_rows);
}
}
#[test]
fn test_concat_table_key() {
let region_id = RegionId::new(1234, 7844);
+1
View File
@@ -3411,6 +3411,7 @@ async fn test_alter_time_index_widen_sparse_compaction() {
};
let put_sparse = |rows| {
RegionRequest::Put(RegionPutRequest {
skip_wal: false,
rows,
hint: Some(WriteHint {
primary_key_encoding: api::v1::PrimaryKeyEncoding::Sparse.into(),
@@ -952,6 +952,7 @@ async fn test_apply_staging_manifest_preserves_unflushed_memtable_with_format(fl
.handle_request(
region_id,
RegionRequest::Put(RegionPutRequest {
skip_wal: false,
rows: unflushed_rows,
hint: None,
partition_expr_version: Some(expected_version),
+386
View File
@@ -732,6 +732,7 @@ async fn test_absent_and_invalid_columns_with_format(flat_format: bool) {
.handle_request(
region_id,
RegionRequest::Put(RegionPutRequest {
skip_wal: false,
rows,
hint: None,
partition_expr_version: None,
@@ -1320,3 +1321,388 @@ async fn test_all_index_metas_list_all_types_with_format(flat_format: bool, expe
assert_eq!(expect_format, debug_format);
}
#[tokio::test]
async fn test_request_skip_wal_recovery() {
check_request_skip_wal_recovery(false, false).await;
check_request_skip_wal_recovery(false, true).await;
check_request_skip_wal_recovery(true, false).await;
check_request_skip_wal_recovery(true, true).await;
}
async fn check_request_skip_wal_recovery(flat_format: bool, skip_wal: bool) {
let mut env = TestEnv::new().await;
let engine = env
.create_engine(MitoConfig {
default_flat_format: flat_format,
..Default::default()
})
.await;
let region_id = RegionId::new(1, 1);
let request = CreateRequestBuilder::new().build();
let table_dir = request.table_dir.clone();
let schema = rows_schema(&request);
engine
.handle_request(region_id, RegionRequest::Create(request))
.await
.unwrap();
let affected = engine
.handle_request(
region_id,
RegionRequest::Put(RegionPutRequest {
rows: Rows {
schema,
rows: build_rows_for_key("a", 0, 4, 0),
},
hint: None,
partition_expr_version: None,
skip_wal,
}),
)
.await
.unwrap();
assert_eq!(affected.affected_rows, 4);
let current = engine
.get_region(region_id)
.unwrap()
.version_control
.current();
assert_eq!(current.committed_sequence, 4);
assert_eq!(current.last_entry_id, u64::from(!skip_wal));
let stream = engine
.scan_to_stream(region_id, ScanRequest::default())
.await
.unwrap();
let before = RecordBatches::try_collect(stream).await.unwrap();
assert_eq!(before.iter().map(|b| b.num_rows()).sum::<usize>(), 4);
reopen_region(&engine, region_id, table_dir, false, HashMap::new()).await;
let stream = engine
.scan_to_stream(region_id, ScanRequest::default())
.await
.unwrap();
let after = RecordBatches::try_collect(stream).await.unwrap();
assert_eq!(
after.iter().map(|b| b.num_rows()).sum::<usize>(),
if skip_wal { 0 } else { 4 }
);
}
#[tokio::test]
async fn test_request_skip_wal_flush_watermarks() {
check_request_skip_wal_flush_watermarks(false).await;
check_request_skip_wal_flush_watermarks(true).await;
}
async fn check_request_skip_wal_flush_watermarks(flat_format: bool) {
let mut env = TestEnv::new().await;
let engine = env
.create_engine(MitoConfig {
default_flat_format: flat_format,
..Default::default()
})
.await;
let region_id = RegionId::new(1, 1);
let request = CreateRequestBuilder::new().build();
let schema = rows_schema(&request);
engine
.handle_request(region_id, RegionRequest::Create(request))
.await
.unwrap();
// Establish a nonzero flushed baseline, skip one write, then resume WAL.
for (round, skip_wal) in [false, true, false].into_iter().enumerate() {
let region = engine.get_region(region_id).unwrap();
let before = region.version_control.current();
let flushed_entry_id = engine
.region_statistic(region_id)
.unwrap()
.manifest
.data_flushed_entry_id();
let affected = engine
.handle_request(
region_id,
RegionRequest::Put(RegionPutRequest {
rows: Rows {
schema: schema.clone(),
rows: build_rows_for_key("a", 0, 4, 0),
},
hint: None,
partition_expr_version: None,
skip_wal,
}),
)
.await
.unwrap();
assert_eq!(affected.affected_rows, 4);
let after_write = region.version_control.current();
let expected_entry_id = before.last_entry_id + u64::from(!skip_wal);
assert_eq!(after_write.last_entry_id, expected_entry_id);
assert_eq!(
after_write.committed_sequence,
before.committed_sequence + 4
);
assert_eq!(after_write.version.flushed_entry_id, flushed_entry_id);
assert_eq!(
engine
.region_statistic(region_id)
.unwrap()
.manifest
.data_flushed_entry_id(),
flushed_entry_id
);
assert_eq!(
after_write.version.flushed_sequence,
before.version.flushed_sequence
);
flush_region(&engine, region_id, None).await;
let after_flush = region.version_control.current();
assert_eq!(after_flush.last_entry_id, expected_entry_id);
assert_eq!(after_flush.version.flushed_sequence, (round as u64 + 1) * 4);
assert_eq!(
engine
.region_statistic(region_id)
.unwrap()
.manifest
.data_flushed_entry_id(),
expected_entry_id
);
if skip_wal {
assert_eq!(expected_entry_id, flushed_entry_id);
} else {
assert!(expected_entry_id > flushed_entry_id);
}
}
}
/// Returns `(committed_sequence, last_entry_id)` from one region version snapshot.
/// Returns `None` if the region is not open.
fn region_write_watermarks(engine: &MitoEngine, region_id: RegionId) -> Option<(u64, u64)> {
let region = engine.find_region(region_id)?;
let current = region.version_control.current();
Some((current.committed_sequence, current.last_entry_id))
}
#[tokio::test]
async fn test_request_skip_wal_mixed_batch_recovery() {
check_request_skip_wal_mixed_batch_recovery(false, false, false).await;
check_request_skip_wal_mixed_batch_recovery(false, false, true).await;
check_request_skip_wal_mixed_batch_recovery(false, true, false).await;
check_request_skip_wal_mixed_batch_recovery(false, true, true).await;
check_request_skip_wal_mixed_batch_recovery(true, false, false).await;
check_request_skip_wal_mixed_batch_recovery(true, false, true).await;
check_request_skip_wal_mixed_batch_recovery(true, true, false).await;
check_request_skip_wal_mixed_batch_recovery(true, true, true).await;
}
async fn check_request_skip_wal_mixed_batch_recovery(
flat_format: bool,
skip_wal: bool,
flush_before_reopen: bool,
) {
use crate::region_write_ctx::RegionWriteCtx;
use crate::request::OptionOutputTx;
use crate::test_util::LogStoreImpl;
use crate::wal::Wal;
let mut env = TestEnv::new().await;
let engine = env
.create_engine(MitoConfig {
default_flat_format: flat_format,
..Default::default()
})
.await;
let region_id = RegionId::new(1, 1);
let request = CreateRequestBuilder::new().build();
let table_dir = request.table_dir.clone();
let schema = rows_schema(&request);
engine
.handle_request(region_id, RegionRequest::Create(request))
.await
.unwrap();
let region = engine.get_region(region_id).unwrap();
let LogStoreImpl::RaftEngine(store) = env.get_log_store().unwrap() else {
panic!("expected the default local WAL");
};
let wal = Wal::new(store);
// Assemble one real worker write context deterministically, instead
// of relying on concurrently submitted requests landing in one batch.
let mut ctx = RegionWriteCtx::new(
region_id,
&region.version_control,
region.provider.clone(),
None,
);
let mut receivers = Vec::with_capacity(4);
for (index, (skip, timestamp)) in [(false, 0), (skip_wal, 0), (false, 0), (skip_wal, 1)]
.into_iter()
.enumerate()
{
let (tx, rx) = tokio::sync::oneshot::channel();
receivers.push(rx);
ctx.push_mutation(
api::v1::OpType::Put as i32,
Some(Rows {
schema: schema.clone(),
// Timestamp zero ends with a WAL-backed update; timestamp one
// ends with a skipped update. Grouping mutations must not
// change either winner chosen by the original sequences.
rows: build_rows_for_key("a", timestamp, timestamp + 1, index * 10)
.into_iter()
.chain(build_rows_for_key(
"a",
index + 1,
index + 2,
index * 10 + 1,
))
.collect(),
}),
None,
OptionOutputTx::from(tx),
None,
skip,
);
}
let mut writer = wal.writer();
ctx.add_wal_entry(&mut writer).unwrap();
let response = writer.write_to_wal().await.unwrap();
assert_eq!(response.last_entry_ids.get(&region_id), Some(&1));
ctx.write_memtable().await;
ctx.publish_sequence_and_entry_id();
drop(ctx);
for rx in receivers {
assert_eq!(rx.await.unwrap().unwrap(), 2);
}
assert_eq!(region_write_watermarks(&engine, region_id), Some((8, 1)));
assert_eq!(
request_skip_wal_values(&engine, region_id).await,
vec![
(0, 20.0),
(1000, 30.0),
(2000, 11.0),
(3000, 21.0),
(4000, 31.0)
]
);
// Inspect persisted WAL mutations, not just an encoder or counter.
let mut reader = wal.wal_entry_reader(&region.provider, region_id, None);
let entries = reader
.read(&region.provider, 1)
.unwrap()
.try_collect::<Vec<_>>()
.await
.unwrap();
assert_eq!(entries.len(), 1);
assert_eq!(entries[0].0, 1);
assert_eq!(
entries[0]
.1
.mutations
.iter()
.map(|m| m.sequence)
.collect::<Vec<_>>(),
if skip_wal {
vec![1, 5]
} else {
vec![1, 3, 5, 7]
}
);
assert!(entries[0].1.bulk_entries.is_empty());
drop(region);
if flush_before_reopen {
flush_region(&engine, region_id, None).await;
let current = engine
.get_region(region_id)
.unwrap()
.version_control
.current();
assert_eq!(
(
current.version.flushed_sequence,
current.version.flushed_entry_id
),
(8, 1)
);
}
// A no-flush close discards memtables. Only WAL-backed rows recover
// unless an explicit flush has already persisted all requests.
reopen_region(&engine, region_id, table_dir, true, HashMap::new()).await;
let loses_skipped_rows = skip_wal && !flush_before_reopen;
let mut expected_values = if loses_skipped_rows {
vec![(0, 20.0), (1000, 1.0), (3000, 21.0)]
} else {
vec![
(0, 20.0),
(1000, 30.0),
(2000, 11.0),
(3000, 21.0),
(4000, 31.0),
]
};
assert_eq!(
request_skip_wal_values(&engine, region_id).await,
expected_values
);
let recovered_sequence = if loses_skipped_rows { 6 } else { 8 };
assert_eq!(
region_write_watermarks(&engine, region_id),
Some((recovered_sequence, 1))
);
// A subsequent default request still writes WAL, even after a
// trailing skipped request or a flush with sequence/entry-id gaps.
put_rows(
&engine,
region_id,
Rows {
schema,
rows: build_rows_for_key("a", 8, 9, 8),
},
)
.await;
assert_eq!(
region_write_watermarks(&engine, region_id),
Some((recovered_sequence + 1, 2))
);
expected_values.push((8000, 8.0));
assert_eq!(
request_skip_wal_values(&engine, region_id).await,
expected_values
);
}
async fn request_skip_wal_values(engine: &MitoEngine, region_id: RegionId) -> Vec<(i64, f64)> {
use datatypes::arrow::array::{Float64Array, TimestampMillisecondArray};
let stream = engine
.scan_to_stream(region_id, ScanRequest::default())
.await
.unwrap();
let batches = RecordBatches::try_collect(stream).await.unwrap();
let mut values = batches
.iter()
.flat_map(|batch| {
let batch = batch.df_record_batch();
let timestamps = batch
.column(2)
.as_any()
.downcast_ref::<TimestampMillisecondArray>()
.unwrap();
let values = batch
.column(1)
.as_any()
.downcast_ref::<Float64Array>()
.unwrap();
timestamps
.values()
.iter()
.copied()
.zip(values.values().iter().copied())
})
.collect::<Vec<_>>();
values.sort_unstable_by_key(|(timestamp, _)| *timestamp);
values
}
+2
View File
@@ -286,6 +286,7 @@ async fn test_write_during_region_editing_is_queued() {
.handle_request(
region_id,
RegionRequest::Put(RegionPutRequest {
skip_wal: false,
rows,
hint: None,
partition_expr_version: None,
@@ -400,6 +401,7 @@ async fn test_stalled_write_fails_fast_if_region_closed_during_editing() {
.handle_request(
region_id,
RegionRequest::Put(RegionPutRequest {
skip_wal: false,
rows,
hint: None,
partition_expr_version: None,
+2
View File
@@ -740,6 +740,7 @@ async fn test_region_write_buffer_does_not_stall_follower_write() {
engine.handle_request(
region_id,
RegionRequest::Put(RegionPutRequest {
skip_wal: false,
rows: Rows {
schema,
rows: build_rows_for_key("follower", 2, 4, 0),
@@ -935,6 +936,7 @@ async fn test_region_write_buffer_rejects_only_full_region_queue() {
.handle_request(
hot_region_id,
RegionRequest::Put(RegionPutRequest {
skip_wal: false,
rows: Rows {
schema: stalled_schema,
rows: build_rows_for_key("hot", 2, 1026, 2),
+1
View File
@@ -478,6 +478,7 @@ async fn test_engine_open_readonly_with_format(flat_format: bool) {
.handle_request(
region_id,
RegionRequest::Put(RegionPutRequest {
skip_wal: false,
rows: rows.clone(),
hint: None,
partition_expr_version: None,
+1
View File
@@ -1126,6 +1126,7 @@ async fn test_two_phase_series_scan() {
};
let put = |rows| {
RegionRequest::Put(RegionPutRequest {
skip_wal: false,
rows,
hint: Some(WriteHint {
primary_key_encoding: api::v1::PrimaryKeyEncoding::Sparse.into(),
@@ -114,6 +114,7 @@ async fn test_set_role_state_gracefully_with_format(flat_format: bool) {
.handle_request(
region_id,
RegionRequest::Put(RegionPutRequest {
skip_wal: false,
rows: rows.clone(),
hint: None,
partition_expr_version: None,
@@ -208,6 +209,7 @@ async fn test_write_downgrading_region_with_format(flat_format: bool) {
.handle_request(
region_id,
RegionRequest::Put(RegionPutRequest {
skip_wal: false,
rows: rows.clone(),
hint: None,
partition_expr_version: None,
+3
View File
@@ -522,6 +522,7 @@ async fn test_close_region_skip_wal_rejects_writes_queued_after_close() {
.handle_request(
region_id,
RegionRequest::Put(RegionPutRequest {
skip_wal: false,
rows: Rows {
schema: rows_schema(&request),
rows: build_rows(3, 4),
@@ -553,6 +554,7 @@ async fn test_close_region_skip_wal_rejects_writes_queued_after_close() {
.handle_request(
region_id,
RegionRequest::Put(RegionPutRequest {
skip_wal: false,
rows: Rows {
schema: rows_schema(&request_cloned),
rows: build_rows(4, 5),
@@ -579,6 +581,7 @@ async fn test_close_region_skip_wal_rejects_writes_queued_after_close() {
.handle_request(
region_id,
RegionRequest::Put(RegionPutRequest {
skip_wal: false,
rows: Rows {
schema: rows_schema(&request_cloned),
rows: build_rows(5, 6),
+7
View File
@@ -285,6 +285,7 @@ async fn test_staging_reject_all_writes_rejects_put() {
.handle_request(
region_id,
RegionRequest::Put(RegionPutRequest {
skip_wal: false,
rows,
hint: None,
partition_expr_version: None,
@@ -355,6 +356,7 @@ async fn test_staging_write_partition_expr_version_with_format(flat_format: bool
.handle_request(
region_id,
RegionRequest::Put(RegionPutRequest {
skip_wal: false,
rows: bad_rows,
hint: None,
partition_expr_version: Some(origin_version),
@@ -376,6 +378,7 @@ async fn test_staging_write_partition_expr_version_with_format(flat_format: bool
.handle_request(
region_id,
RegionRequest::Put(RegionPutRequest {
skip_wal: false,
rows: compat_rows,
hint: None,
partition_expr_version: None,
@@ -393,6 +396,7 @@ async fn test_staging_write_partition_expr_version_with_format(flat_format: bool
.handle_request(
region_id,
RegionRequest::Put(RegionPutRequest {
skip_wal: false,
rows: ok_rows,
hint: None,
partition_expr_version: Some(expected_version),
@@ -440,6 +444,7 @@ async fn test_staging_write_partition_expr_version_with_format(flat_format: bool
.handle_request(
region_id,
RegionRequest::Put(RegionPutRequest {
skip_wal: false,
rows: exit_rows,
hint: None,
partition_expr_version: Some(origin_version),
@@ -457,6 +462,7 @@ async fn test_staging_write_partition_expr_version_with_format(flat_format: bool
.handle_request(
region_id,
RegionRequest::Put(RegionPutRequest {
skip_wal: false,
rows: compat_rows,
hint: None,
partition_expr_version: None,
@@ -483,6 +489,7 @@ async fn test_staging_write_partition_expr_version_with_format(flat_format: bool
.handle_request(
region_id,
RegionRequest::Put(RegionPutRequest {
skip_wal: false,
rows: commit_rows,
hint: None,
partition_expr_version: Some(expected_version),
+1
View File
@@ -953,6 +953,7 @@ where
OptionOutputTx::none(),
// We should respect the sequence in WAL during replay.
Some(mutation.sequence),
false,
);
}
+202 -24
View File
@@ -89,12 +89,14 @@ pub(crate) struct RegionWriteCtx {
/// We keep [WalEntry] instead of mutations to avoid taking mutations
/// out of the context to construct the wal entry when we write to the wal.
wal_entry: WalEntry,
/// Mutations that skip WAL, paired with their write notifiers.
memtable_mutations: Vec<(Mutation, WriteNotify)>,
/// Wal options of the region being written to.
provider: Provider,
/// Notifiers to send write results to waiters.
///
/// The i-th notify is for i-th mutation.
notifiers: Vec<WriteNotify>,
/// The i-th notify is for the i-th mutation in `wal_entry`.
wal_notifiers: Vec<WriteNotify>,
/// Notifiers for bulk requests.
bulk_notifiers: Vec<WriteNotify>,
/// Pending bulk write requests
@@ -133,8 +135,9 @@ impl RegionWriteCtx {
next_sequence: committed_sequence + 1,
next_entry_id: last_entry_id + 1,
wal_entry: WalEntry::default(),
memtable_mutations: Vec::new(),
provider,
notifiers: Vec::new(),
wal_notifiers: Vec::new(),
bulk_notifiers: vec![],
failed: false,
put_num: 0,
@@ -153,21 +156,28 @@ impl RegionWriteCtx {
write_hint: Option<WriteHint>,
tx: OptionOutputTx,
sequence: Option<SequenceNumber>,
skip_wal: bool,
) {
if let Some(sequence) = sequence {
self.next_sequence = sequence;
}
let num_rows = rows.as_ref().map(|rows| rows.rows.len()).unwrap_or(0);
self.wal_entry.mutations.push(Mutation {
let mutation = Mutation {
op_type,
sequence: self.next_sequence,
rows,
write_hint,
});
};
// Assign sequences before routing so concurrent memtable writes retain
// their logical order regardless of the WAL policy.
let notify = WriteNotify::new(tx, num_rows);
// Notifiers are 1:1 map to mutations.
self.notifiers.push(notify);
if skip_wal {
self.memtable_mutations.push((mutation, notify));
} else {
self.wal_entry.mutations.push(mutation);
self.wal_notifiers.push(notify);
}
// Increase sequence number.
self.next_sequence += num_rows as u64;
@@ -206,13 +216,19 @@ impl RegionWriteCtx {
/// Returns whether writes in this context should skip WAL.
pub(crate) fn skip_wal(&self) -> bool {
self.provider == Provider::Noop || self.version.options.skip_wal
self.provider == Provider::Noop
|| self.version.options.skip_wal
|| (self.wal_entry.mutations.is_empty() && self.wal_entry.bulk_entries.is_empty())
}
/// Sets error and marks all write operations are failed.
pub(crate) fn set_error(&mut self, err: Arc<Error>) {
// Set error for all notifiers.
for notify in &mut self.notifiers {
for notify in self
.wal_notifiers
.iter_mut()
.chain(self.memtable_mutations.iter_mut().map(|(_, notify)| notify))
{
notify.err = Some(err.clone());
}
for notify in &mut self.bulk_notifiers {
@@ -241,7 +257,7 @@ impl RegionWriteCtx {
/// Consumes mutations and writes them into mutable memtable.
pub(crate) async fn write_memtable(&mut self) {
debug_assert_eq!(self.notifiers.len(), self.wal_entry.mutations.len());
debug_assert_eq!(self.wal_notifiers.len(), self.wal_entry.mutations.len());
if self.failed {
return;
@@ -254,34 +270,38 @@ impl RegionWriteCtx {
None
};
let mutations = mem::take(&mut self.wal_entry.mutations)
let mut mutations = mem::take(&mut self.wal_entry.mutations)
.into_iter()
.enumerate()
.filter_map(|(i, mutation)| {
.zip(&mut self.wal_notifiers)
.chain(
self.memtable_mutations
.iter_mut()
// Keep notifiers in the context until all writes complete.
.map(|(mutation, notify)| (mem::take(mutation), notify)),
)
.filter_map(|(mutation, notify)| {
let kvs = KeyValues::new(&self.version.metadata, mutation)?;
Some((i, kvs))
Some((notify, kvs))
})
.collect::<Vec<_>>();
if mutations.len() == 1 {
if let Err(err) = mutable_memtable.write(&mutations[0].1) {
self.notifiers[mutations[0].0].err = Some(Arc::new(err));
mutations[0].0.err = Some(Arc::new(err));
}
} else {
let mut tasks = FuturesUnordered::new();
for (i, kvs) in mutations {
for (notify, kvs) in mutations {
let mutable = mutable_memtable.clone();
// use tokio runtime to schedule tasks.
tasks.push(common_runtime::spawn_blocking_global(move || {
(i, mutable.write(&kvs))
}));
let task = common_runtime::spawn_blocking_global(move || mutable.write(&kvs));
tasks.push(async move { (notify, task.await) });
}
while let Some(result) = tasks.next().await {
// first unwrap the result from `spawn` above
let (i, result) = result.unwrap();
if let Err(err) = result {
self.notifiers[i].err = Some(Arc::new(err));
while let Some((notify, result)) = tasks.next().await {
// First unwrap the result from `spawn` above.
if let Err(err) = result.unwrap() {
notify.err = Some(Arc::new(err));
}
}
}
@@ -523,6 +543,7 @@ mod tests {
use common_recordbatch::DfRecordBatch;
use datatypes::arrow::array::{ArrayRef, TimestampMillisecondArray};
use datatypes::arrow::datatypes::{DataType, Field, Schema};
use prost::Message;
use store_api::logstore::provider::Provider;
use tokio::sync::oneshot;
@@ -531,6 +552,163 @@ mod tests {
use crate::memtable::bulk::part::BulkPart;
use crate::test_util::version_util::VersionControlBuilder;
#[test]
fn test_request_skip_wal_preserves_sequences_and_other_writes() {
// Ablate only the request flag: the workload and sequence allocation stay identical.
check_request_skip_wal_preserves_sequences_and_other_writes(false);
check_request_skip_wal_preserves_sequences_and_other_writes(true);
}
fn check_request_skip_wal_preserves_sequences_and_other_writes(skip_wal: bool) {
let builder = VersionControlBuilder::new();
let region_id = builder.region_id();
let version_control = Arc::new(builder.build());
let mut ctx = RegionWriteCtx::new(
region_id,
&version_control,
Provider::raft_engine_provider(region_id.as_u64()),
None,
);
for (op_type, skip, num_rows) in [
(OpType::Put, skip_wal, 2),
(OpType::Put, false, 3),
(OpType::Delete, false, 1),
] {
ctx.push_mutation(
op_type as i32,
Some(Rows {
schema: vec![],
rows: vec![api::v1::Row::default(); num_rows],
}),
None,
OptionOutputTx::none(),
None,
skip,
);
}
// Unequal row counts make a notifier/mutation pairing mismatch visible.
for (mutation, notify) in ctx
.wal_entry
.mutations
.iter()
.zip(&ctx.wal_notifiers)
.chain(
ctx.memtable_mutations
.iter()
.map(|(mutation, notify)| (mutation, notify)),
)
{
assert_eq!(mutation.rows.as_ref().unwrap().rows.len(), notify.num_rows);
}
assert!(ctx.push_bulk(OptionOutputTx::none(), new_bulk_part(), None));
assert!(!ctx.skip_wal());
assert_eq!(ctx.next_sequence, 9);
assert_eq!(ctx.wal_entry.bulk_entries.len(), 1);
assert_eq!(ctx.bulk_parts[0].sequence, 7);
let sequences: Vec<_> = ctx.wal_entry.mutations.iter().map(|m| m.sequence).collect();
assert_eq!(sequences, if skip_wal { vec![3, 6] } else { vec![1, 3, 6] });
// Check the actual WAL bytes after routing the mutations.
let encoded = crate::wal::encoder::WalEntryEncoder::new().encode_to_vec(&ctx.wal_entry);
let decoded = WalEntry::decode(encoded.as_slice()).unwrap();
assert_eq!(
decoded
.mutations
.iter()
.map(|m| m.sequence)
.collect::<Vec<_>>(),
sequences
);
assert_eq!(decoded.bulk_entries, ctx.wal_entry.bulk_entries);
assert_eq!(ctx.wal_entry.mutations.len(), if skip_wal { 2 } else { 3 });
assert_eq!(ctx.memtable_mutations.len(), usize::from(skip_wal));
if skip_wal {
assert_eq!(ctx.memtable_mutations[0].0.sequence, 1);
}
assert_eq!(
ctx.wal_entry.mutations.last().unwrap().op_type,
OpType::Delete as i32
);
}
#[test]
fn test_internal_delete_respects_skip_wal_flag() {
check_internal_delete_respects_skip_wal_flag(false);
check_internal_delete_respects_skip_wal_flag(true);
}
fn check_internal_delete_respects_skip_wal_flag(skip_wal: bool) {
let builder = VersionControlBuilder::new();
let region_id = builder.region_id();
let version_control = Arc::new(builder.build());
let mut ctx = RegionWriteCtx::new(
region_id,
&version_control,
Provider::raft_engine_provider(region_id.as_u64()),
None,
);
ctx.push_mutation(
OpType::Delete as i32,
Some(Rows {
schema: vec![],
rows: vec![api::v1::Row::default(); 2],
}),
None,
OptionOutputTx::none(),
None,
skip_wal,
);
assert_eq!(ctx.skip_wal(), skip_wal);
assert_eq!(ctx.next_sequence, 3);
assert_eq!(ctx.delete_num, 2);
assert_eq!(ctx.wal_entry.mutations.len(), usize::from(!skip_wal));
assert_eq!(ctx.memtable_mutations.len(), usize::from(skip_wal));
let encoded = crate::wal::encoder::WalEntryEncoder::new().encode_to_vec(&ctx.wal_entry);
let decoded = WalEntry::decode(encoded.as_slice()).unwrap();
if skip_wal {
assert!(decoded.mutations.is_empty());
} else {
assert_eq!(decoded, ctx.wal_entry);
}
}
#[test]
fn test_all_request_skip_wal_keeps_entry_id_and_propagates_errors() {
let builder = VersionControlBuilder::new();
let region_id = builder.region_id();
let version_control = Arc::new(builder.build());
let mut ctx = RegionWriteCtx::new(
region_id,
&version_control,
Provider::raft_engine_provider(region_id.as_u64()),
None,
);
let (tx, rx) = oneshot::channel();
ctx.push_mutation(
OpType::Put as i32,
Some(Rows {
schema: vec![],
rows: vec![api::v1::Row::default(); 2],
}),
None,
OptionOutputTx::from(tx),
None,
true,
);
assert!(ctx.skip_wal());
assert!(ctx.wal_entry.mutations.is_empty());
assert_eq!(ctx.memtable_mutations.len(), 1);
assert_eq!(ctx.next_entry_id(), 1);
assert_eq!(ctx.next_sequence, 3);
ctx.set_error(Arc::new(
UnexpectedSnafu {
reason: "wal failed".to_string(),
}
.build(),
));
drop(ctx);
assert!(rx.blocking_recv().unwrap().is_err());
}
#[test]
fn test_set_error_marks_bulk_notifiers_failed() {
let builder = VersionControlBuilder::new();
+29
View File
@@ -76,6 +76,8 @@ pub struct WriteRequest {
pub name_to_index: HashMap<String, usize>,
/// Whether each column has null.
pub has_null: Vec<bool>,
/// Whether this insert should skip WAL. Never applies to deletes.
pub skip_wal: bool,
/// Write hint.
pub hint: Option<WriteHint>,
/// Region metadata on the time of this request is created.
@@ -137,11 +139,18 @@ impl WriteRequest {
name_to_index,
has_null,
hint: None,
skip_wal: false,
region_metadata,
partition_expr_version: None,
})
}
/// Sets the request-level WAL policy.
pub fn with_skip_wal(mut self, skip_wal: bool) -> Self {
self.skip_wal = skip_wal;
self
}
/// Sets the write hint.
pub fn with_hint(mut self, hint: Option<WriteHint>) -> Self {
self.hint = hint;
@@ -672,6 +681,7 @@ impl WorkerRequest {
let mut write_request =
WriteRequest::new(region_id, OpType::Put, v.rows, region_metadata.clone())?
.with_hint(v.hint)
.with_skip_wal(v.skip_wal)
.with_partition_expr_version(v.partition_expr_version);
if write_request.primary_key_encoding() == PrimaryKeyEncoding::Dense
&& let Some(region_metadata) = &region_metadata
@@ -1873,6 +1883,25 @@ mod tests {
);
}
#[test]
fn test_delete_request_defaults_to_writing_wal() {
let (request, _receiver) = WorkerRequest::try_from_region_request(
RegionId::new(1, 1),
RegionRequest::Delete(store_api::region_request::RegionDeleteRequest {
rows: Rows::default(),
hint: None,
partition_expr_version: None,
}),
None,
)
.unwrap();
let WorkerRequest::Write(request) = request else {
panic!("expected a write request");
};
assert_eq!(request.request.op_type, OpType::Delete);
assert!(!request.request.skip_wal);
}
#[test]
fn test_write_request_metadata() {
let rows = Rows {
+1
View File
@@ -1351,6 +1351,7 @@ pub async fn put_rows(engine: &MitoEngine, region_id: RegionId, rows: Rows) {
.handle_request(
region_id,
RegionRequest::Put(RegionPutRequest {
skip_wal: false,
rows,
hint: None,
partition_expr_version: None,
+44 -4
View File
@@ -451,6 +451,7 @@ impl<S> RegionWorkerLoop<S> {
sender_req.request.hint,
sender_req.sender,
None,
sender_req.request.skip_wal,
);
}
}
@@ -613,14 +614,21 @@ async fn write_wal<S: LogStore>(
region_ctxs: &mut HashMap<RegionId, RegionWriteCtx>,
) -> bool {
let mut wal_writer = wal.writer();
let mut has_wal_entries = false;
for region_ctx in region_ctxs.values_mut() {
if region_ctx.skip_wal() {
continue;
}
if let Err(e) = region_ctx.add_wal_entry(&mut wal_writer).map_err(Arc::new) {
region_ctx.set_error(e);
} else {
has_wal_entries = true;
}
}
// All-skipped batches should not touch the log store, even with an empty append.
if !has_wal_entries {
return true;
}
match wal_writer.write_to_wal().await.map_err(Arc::new) {
Ok(response) => {
for (region_id, region_ctx) in region_ctxs.iter_mut() {
@@ -880,6 +888,7 @@ mod tests {
fn new_region_ctx(
region_id: RegionId,
skip_wal: bool,
) -> (RegionWriteCtx, oneshot::Receiver<Result<AffectedRows>>) {
let version_control = Arc::new(VersionControlBuilder::new().build());
let mut ctx = RegionWriteCtx::new(
@@ -908,10 +917,41 @@ mod tests {
None,
OptionOutputTx::from(tx),
None,
skip_wal,
);
(ctx, rx)
}
#[tokio::test]
async fn test_request_skip_wal_does_not_append_empty_batch() {
// Only change the request flag. A failing log store demonstrates that
// the all-skipped path never invokes append_batch, including empty appends.
check_request_skip_wal_does_not_append_empty_batch(false).await;
check_request_skip_wal_does_not_append_empty_batch(true).await;
}
async fn check_request_skip_wal_does_not_append_empty_batch(skip_wal: bool) {
let region_id = RegionId::new(1, 1);
let wal = Wal::new(Arc::new(MockLogStore {
fail_append: true,
..Default::default()
}));
let (ctx, rx) = new_region_ctx(region_id, skip_wal);
let version_control = ctx.version_control().clone();
let mut contexts = HashMap::from([(region_id, ctx)]);
assert_eq!(write_wal(&wal, &mut contexts).await, skip_wal);
if skip_wal {
let ctx = contexts.get_mut(&region_id).unwrap();
assert_eq!(ctx.next_entry_id(), 1);
ctx.write_memtable().await;
ctx.publish_sequence_and_entry_id();
assert_eq!(version_control.committed_sequence(), 1);
assert_eq!(version_control.current().last_entry_id, 0);
}
drop(contexts);
assert_eq!(rx.await.unwrap().is_ok(), skip_wal);
}
#[tokio::test]
async fn test_write_wal_skips_region_failed_to_build_entry() {
let failing_region = RegionId::new(1, 1);
@@ -922,10 +962,10 @@ mod tests {
}));
let mut region_ctxs = HashMap::new();
let (ctx, failing_rx) = new_region_ctx(failing_region);
let (ctx, failing_rx) = new_region_ctx(failing_region, false);
let failing_committed_sequence = ctx.version_control().committed_sequence();
region_ctxs.insert(failing_region, ctx);
let (ctx, ok_rx) = new_region_ctx(ok_region);
let (ctx, ok_rx) = new_region_ctx(ok_region, false);
let ok_committed_sequence = ctx.version_control().committed_sequence();
region_ctxs.insert(ok_region, ctx);
let entry_id = region_ctxs[&ok_region].next_entry_id();
@@ -1026,7 +1066,7 @@ mod tests {
}));
let mut region_ctxs = HashMap::new();
let (ctx, rx) = new_region_ctx(failing_region);
let (ctx, rx) = new_region_ctx(failing_region, false);
region_ctxs.insert(failing_region, ctx);
// Writing an empty batch to the WAL succeeds, the failed region must not panic
@@ -1047,7 +1087,7 @@ mod tests {
}));
let mut region_ctxs = HashMap::new();
let (ctx, rx) = new_region_ctx(region_id);
let (ctx, rx) = new_region_ctx(region_id, false);
region_ctxs.insert(region_id, ctx);
assert!(!write_wal(&wal, &mut region_ctxs).await);
+24 -2
View File
@@ -265,6 +265,8 @@ impl Inserter {
accommodate_existing_schema: bool,
is_single_value: bool,
) -> Result<Output> {
let skip_wal = ctx.skip_wal();
// remove empty requests
requests.inserts.retain(|req| {
req.rows
@@ -297,7 +299,7 @@ impl Inserter {
instant_table_ids,
self.partition_manager.as_ref(),
)
.convert(requests)
.convert(requests, skip_wal)
.await?;
self.do_request(inserts, &table_infos, &ctx).await
@@ -311,6 +313,8 @@ impl Inserter {
statement_executor: &StatementExecutor,
physical_table: String,
) -> Result<Output> {
let skip_wal = ctx.skip_wal();
// remove empty requests
requests.inserts.retain(|req| {
req.rows
@@ -343,7 +347,7 @@ impl Inserter {
.map(|info| (info.name.clone(), info.clone()))
.collect::<HashMap<_, _>>();
let inserts = RowToRegion::new(name_to_info, instant_table_ids, &self.partition_manager)
.convert(requests)
.convert(requests, skip_wal)
.await?;
self.do_request(inserts, &table_infos, &ctx).await
@@ -1743,6 +1747,24 @@ mod tests {
assert!(table_is_native_histogram(&table));
}
#[test]
fn test_skip_wal_does_not_change_table_options() {
check_skip_wal_does_not_change_table_options(false);
check_skip_wal_does_not_change_table_options(true);
}
fn check_skip_wal_does_not_change_table_options(skip_wal: bool) {
let ctx = Arc::new(QueryContext::with(
DEFAULT_CATALOG_NAME,
DEFAULT_SCHEMA_NAME,
));
ctx.set_skip_wal(skip_wal);
let mut options = Default::default();
fill_table_options_for_create(&mut options, &AutoCreateTableType::Physical, &ctx);
assert!(!options.contains_key(session::hints::INSERT_SKIP_WAL_HINT));
assert!(!options.contains_key("skip_wal"));
}
#[test]
fn test_last_non_null_create_options_preserve_default_without_append_mode() {
let ctx = Arc::new(QueryContext::with(
@@ -34,6 +34,7 @@ impl<'a> Partitioner<'a> {
&self,
table_info: &TableInfo,
rows: Rows,
skip_wal: bool,
) -> Result<Vec<InsertRequest>> {
let table_id = table_info.table_id();
let requests = self
@@ -44,6 +45,7 @@ impl<'a> Partitioner<'a> {
.into_iter()
.map(
|(region_number, (rows, partition_expr_version))| InsertRequest {
skip_wal,
region_id: RegionId::new(table_id, region_number).into(),
rows: Some(rows),
partition_expr_version: partition_expr_version
@@ -45,6 +45,7 @@ impl<'a> RowToRegion<'a> {
pub async fn convert(
&self,
requests: RowInsertRequests,
skip_wal: bool,
) -> Result<InstantAndNormalInsertRequests> {
let mut region_request = Vec::with_capacity(requests.inserts.len());
let mut instant_request = Vec::with_capacity(requests.inserts.len());
@@ -55,7 +56,7 @@ impl<'a> RowToRegion<'a> {
let table_id = table_info.table_id();
let requests = Partitioner::new(self.partition_manager)
.partition_insert_requests(table_info, rows)
.partition_insert_requests(table_info, rows, skip_wal)
.await?;
if self.instant_table_ids.contains(&table_id) {
@@ -81,3 +82,79 @@ impl<'a> RowToRegion<'a> {
.context(TableNotFoundSnafu { table_name })
}
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use api::v1::helper::tag_column_schema;
use api::v1::value::ValueData;
use api::v1::{ColumnDataType, Row, RowInsertRequest, Rows, Value};
use super::*;
use crate::tests::{
create_partition_rule_manager, new_test_table_info, prepare_mocked_backend,
};
#[tokio::test]
async fn test_partitioned_insert_skip_wal_normal_and_instant() {
check_partitioned_insert_skip_wal(false, false).await;
check_partitioned_insert_skip_wal(false, true).await;
check_partitioned_insert_skip_wal(true, false).await;
check_partitioned_insert_skip_wal(true, true).await;
}
async fn check_partitioned_insert_skip_wal(instant: bool, skip_wal: bool) {
let backend = prepare_mocked_backend().await;
let partition_manager = create_partition_rule_manager(backend).await;
let table_info = Arc::new(new_test_table_info(1, "table_1", [1, 2, 3].into_iter()));
let instant_table_ids = if instant {
HashSet::from_iter([1])
} else {
HashSet::default()
};
let converter = RowToRegion::new(
HashMap::from_iter([("table_1".to_string(), table_info)]),
instant_table_ids,
&partition_manager,
);
let requests = RowInsertRequests {
inserts: vec![RowInsertRequest {
table_name: "table_1".to_string(),
rows: Some(Rows {
schema: vec![tag_column_schema("a", ColumnDataType::Int32)],
rows: [1, 11, 101]
.into_iter()
.map(|value| Row {
values: vec![Value {
value_data: Some(ValueData::I32Value(value)),
}],
})
.collect(),
}),
}],
};
let result = converter.convert(requests, skip_wal).await.unwrap();
let (selected, other) = if instant {
(result.instant_requests, result.normal_requests)
} else {
(result.normal_requests, result.instant_requests)
};
assert!(other.requests.is_empty());
assert_eq!(selected.requests.len(), 3);
assert!(
selected
.requests
.iter()
.all(|request| request.skip_wal == skip_wal)
);
assert_eq!(
selected
.requests
.iter()
.map(|request| request.rows.as_ref().unwrap().rows.len())
.sum::<usize>(),
3
);
}
}
@@ -154,7 +154,7 @@ impl<'a> StatementToRegion<'a> {
}
let requests = Partitioner::new(self.partition_manager)
.partition_insert_requests(&table_info, Rows { schema, rows })
.partition_insert_requests(&table_info, Rows { schema, rows }, query_ctx.skip_wal())
.await?;
let requests = RegionInsertRequests { requests };
if table_info.is_ttl_instant_table() {
@@ -46,7 +46,7 @@ impl<'a> TableToRegion<'a> {
let rows = Rows { schema, rows };
let requests = Partitioner::new(self.partition_manager)
.partition_insert_requests(self.table_info, rows)
.partition_insert_requests(self.table_info, rows, request.skip_wal)
.await?;
let requests = RegionInsertRequests { requests };
@@ -84,6 +84,11 @@ mod tests {
#[tokio::test]
async fn test_insert_request_table_to_region() {
check_insert_request_table_to_region(false).await;
check_insert_request_table_to_region(true).await;
}
async fn check_insert_request_table_to_region(skip_wal: bool) {
// region to datanode placement:
// 1 -> 1
// 2 -> 2
@@ -100,12 +105,13 @@ mod tests {
let converter = TableToRegion::new(&table_info, &partition_manager);
let table_request = build_table_request(Arc::new(Int32Vector::from(vec![
let mut table_request = build_table_request(Arc::new(Int32Vector::from(vec![
Some(1),
None,
Some(11),
Some(101),
])));
table_request.skip_wal = skip_wal;
let versions = partition_manager
.find_physical_partition_info(1)
.await
@@ -127,21 +133,26 @@ mod tests {
let region_request = region_id_to_region_requests.remove(&region_id).unwrap();
assert_eq!(
region_request,
build_region_request(vec![Some(101)], region_id, versions[&region_id])
build_region_request(vec![Some(101)], region_id, versions[&region_id], skip_wal)
);
let region_id = RegionId::new(1, 2).as_u64();
let region_request = region_id_to_region_requests.remove(&region_id).unwrap();
assert_eq!(
region_request,
build_region_request(vec![Some(11)], region_id, versions[&region_id])
build_region_request(vec![Some(11)], region_id, versions[&region_id], skip_wal)
);
let region_id = RegionId::new(1, 3).as_u64();
let region_request = region_id_to_region_requests.remove(&region_id).unwrap();
assert_eq!(
region_request,
build_region_request(vec![Some(1), None], region_id, versions[&region_id])
build_region_request(
vec![Some(1), None],
region_id,
versions[&region_id],
skip_wal
)
);
}
@@ -151,6 +162,7 @@ mod tests {
schema_name: DEFAULT_SCHEMA_NAME.to_string(),
table_name: "table_1".to_string(),
columns_values: HashMap::from([("a".to_string(), vector)]),
skip_wal: false,
}
}
@@ -158,8 +170,10 @@ mod tests {
rows: Vec<Option<i32>>,
region_id: u64,
version: Option<u64>,
skip_wal: bool,
) -> RegionInsertRequest {
RegionInsertRequest {
skip_wal,
region_id,
rows: Some(Rows {
schema: vec![tag_column_schema("a", ColumnDataType::Int32)],
+2 -1
View File
@@ -63,7 +63,7 @@ use query::QueryEngineRef;
use query::parser::QueryStatement;
use session::context::{Channel, QueryContextBuilder, QueryContextRef};
use session::table_name::table_idents_to_full_name;
use set::{set_query_timeout, set_read_preference};
use set::{set_query_timeout, set_read_preference, set_skip_wal};
use snafu::{OptionExt, ResultExt, ensure};
use sql::ast::ObjectNamePartExt;
use sql::statements::OptionMap;
@@ -534,6 +534,7 @@ impl StatementExecutor {
match var_name.as_str() {
"READ_PREFERENCE" => set_read_preference(set_var.value, query_ctx)?,
"SKIP_WAL" => set_skip_wal(set_var.value, query_ctx)?,
"@@TIME_ZONE" | "@@SESSION.TIME_ZONE" | "TIMEZONE" | "TIME_ZONE" => {
set_timezone(set_var.value, query_ctx)?
@@ -487,6 +487,7 @@ impl StatementExecutor {
schema_name: req.schema_name.clone(),
table_name: req.table_name.clone(),
columns_values,
skip_wal: query_ctx.skip_wal(),
},
query_ctx.clone(),
));
+79
View File
@@ -38,6 +38,35 @@ lazy_static! {
static ref PG_TIME_INPUT_REGEX: Regex = Regex::new(r"^(\d+)(ms|s|min|h|d)$").unwrap();
}
/// Sets the session WAL policy for ordinary inserts without changing table options.
pub fn set_skip_wal(exprs: Vec<Expr>, ctx: QueryContextRef) -> Result<()> {
let [Expr::Value(value)] = exprs.as_slice() else {
return NotSupportedSnafu {
feat: "SET skip_wal requires exactly one boolean value",
}
.fail();
};
let skip_wal = match &value.value {
Value::Boolean(value) => *value,
Value::SingleQuotedString(value) | Value::DoubleQuotedString(value) => {
value.parse::<bool>().map_err(|_| {
NotSupportedSnafu {
feat: format!("Invalid skip_wal value {value:?}: expected true or false"),
}
.build()
})?
}
_ => {
return NotSupportedSnafu {
feat: "SET skip_wal requires true or false",
}
.fail();
}
};
ctx.set_skip_wal(skip_wal);
Ok(())
}
pub fn set_read_preference(exprs: Vec<Expr>, ctx: QueryContextRef) -> Result<()> {
let read_preference_expr = exprs.first().context(NotSupportedSnafu {
feat: "No read preference find in set variable statement",
@@ -381,8 +410,58 @@ fn parse_pg_query_timeout_input(input: &str) -> Result<u64> {
#[cfg(test)]
mod test {
use session::Session;
use session::context::{Channel, QueryContext};
use sql::ast::{Expr, Value};
use super::set_skip_wal;
use crate::statement::set::parse_pg_query_timeout_input;
#[test]
fn test_set_skip_wal() {
let ctx = QueryContext::arc();
for value in [true, false] {
set_skip_wal(vec![Expr::Value(Value::Boolean(value).into())], ctx.clone()).unwrap();
assert_eq!(ctx.skip_wal(), value);
set_skip_wal(
vec![Expr::Value(
Value::SingleQuotedString(value.to_string()).into(),
)],
ctx.clone(),
)
.unwrap();
assert_eq!(ctx.skip_wal(), value);
}
ctx.set_skip_wal(true);
for values in [
vec![],
vec![Expr::Value(Value::Number("1".to_string(), false).into())],
vec![Expr::Value(
Value::SingleQuotedString("invalid".to_string()).into(),
)],
vec![Expr::Value(Value::Boolean(false).into()); 2],
] {
assert!(set_skip_wal(values, ctx.clone()).is_err());
assert!(ctx.skip_wal());
}
}
#[test]
fn test_set_skip_wal_session_isolation() {
for channel in [Channel::Mysql, Channel::Postgres] {
let session = Session::new(None, channel, Default::default(), 0);
let other = Session::new(None, channel, Default::default(), 1);
assert!(!session.new_query_context().skip_wal());
set_skip_wal(
vec![Expr::Value(Value::Boolean(true).into())],
session.new_query_context(),
)
.unwrap();
assert!(session.new_query_context().skip_wal());
assert!(!other.new_query_context().skip_wal());
}
}
#[test]
fn test_parse_pg_query_timeout_input() {
assert!(parse_pg_query_timeout_input("").is_err());
+1
View File
@@ -392,6 +392,7 @@ impl DatafusionQueryEngine {
schema_name,
table_name,
columns_values: column_vectors,
skip_wal: query_ctx.skip_wal(),
};
self.state
+77 -7
View File
@@ -21,6 +21,7 @@ use common_catalog::consts::{DEFAULT_CATALOG_NAME, DEFAULT_SCHEMA_NAME};
use common_catalog::parse_catalog_and_schema_from_db_string;
use common_error::ext::ErrorExt;
use session::context::{Channel, QueryContextBuilder, QueryContextRef};
use session::hints::INSERT_SKIP_WAL_HINT;
use snafu::{OptionExt, ResultExt};
use tonic::Status;
use tonic::metadata::MetadataMap;
@@ -28,6 +29,7 @@ use tonic::metadata::MetadataMap;
use crate::error::Error::UnsupportedAuthScheme;
use crate::error::{AuthSnafu, InvalidParameterSnafu, NotFoundAuthHeaderSnafu, Result};
use crate::grpc::TonicResult;
use crate::hint_headers;
use crate::http::AUTHORIZATION_HEADER;
use crate::http::header::constants::GREPTIME_DB_HEADER_NAME;
use crate::metrics::METRIC_AUTH_FAILURE;
@@ -45,13 +47,26 @@ pub fn create_query_context_from_grpc_metadata(
)
};
Ok(Arc::new(
QueryContextBuilder::default()
.current_catalog(catalog)
.current_schema(schema)
.channel(Channel::Grpc)
.build(),
))
let ctx = QueryContextBuilder::default()
.current_catalog(catalog)
.current_schema(schema)
.channel(Channel::Grpc)
.build();
// OTEL Arrow uses ordinary inserts. Accept only its request-level WAL hint,
// leaving unrelated hints and reserved internal extensions unchanged.
if let Some((key, value)) = hint_headers::extract_hints(headers)
.into_iter()
.find(|(key, _)| key == INSERT_SKIP_WAL_HINT)
{
let skip_wal = value.parse::<bool>().map_err(|_| {
InvalidParameterSnafu {
reason: format!("Invalid {key} hint: expected true or false, got {value:?}"),
}
.build()
})?;
ctx.set_skip_wal(skip_wal);
}
Ok(Arc::new(ctx))
}
/// Helper function to extract a header from the metadata map.
@@ -161,3 +176,58 @@ pub async fn auth(
.inc();
})
}
#[cfg(test)]
mod tests {
use session::hints::{HINTS_KEY, REMOTE_QUERY_ID_EXTENSION_KEY, RESERVED_EXTENSION_KEYS};
use super::*;
#[test]
fn test_arrow_insert_hint_does_not_accept_reserved_extensions() {
let mut headers = MetadataMap::new();
assert_eq!(
create_query_context_from_grpc_metadata(&headers)
.unwrap()
.extension(INSERT_SKIP_WAL_HINT),
None
);
for (value, expected) in [("true", true), ("false", false)] {
let mut hints = format!("insert_skip_wal={value},ttl=7d");
for key in RESERVED_EXTENSION_KEYS {
hints.push_str(&format!(",{key}=external"));
}
headers.insert(HINTS_KEY, hints.parse().unwrap());
let ctx = create_query_context_from_grpc_metadata(&headers).unwrap();
assert_eq!(ctx.skip_wal(), expected);
assert_eq!(ctx.extension(INSERT_SKIP_WAL_HINT), None);
assert_eq!(ctx.extension("ttl"), None);
for key in RESERVED_EXTENSION_KEYS {
if key == REMOTE_QUERY_ID_EXTENSION_KEY {
// The builder generates this ID; external hints must not replace it.
assert!(ctx.remote_query_id().is_some());
assert_ne!(ctx.extension(key), Some("external"));
} else {
assert_eq!(ctx.extension(key), None);
}
}
}
// Only the first matching hint is parsed and applied.
for (hints, expected) in [
("insert_skip_wal=true,insert_skip_wal=false", true),
("insert_skip_wal=false,insert_skip_wal=true", false),
("insert_skip_wal=true,insert_skip_wal=invalid", true),
] {
headers.insert(HINTS_KEY, hints.parse().unwrap());
let ctx = create_query_context_from_grpc_metadata(&headers).unwrap();
assert_eq!(ctx.skip_wal(), expected);
}
for value in ["", "TRUE", "1", "invalid"] {
headers.insert(
HINTS_KEY,
format!("insert_skip_wal={value}").parse().unwrap(),
);
assert!(create_query_context_from_grpc_metadata(&headers).is_err());
}
}
}
+105 -23
View File
@@ -38,7 +38,7 @@ use common_telemetry::{debug, error, tracing, warn};
use common_time::timezone::parse_timezone;
use futures_util::StreamExt;
use session::context::{Channel, QueryContextBuilder, QueryContextRef};
use session::hints::{READ_PREFERENCE_HINT, is_reserved_extension_key};
use session::hints::{INSERT_SKIP_WAL_HINT, READ_PREFERENCE_HINT, is_reserved_extension_key};
use snafu::{OptionExt, ResultExt};
use tokio::sync::mpsc;
use tokio::sync::mpsc::error::TrySendError;
@@ -202,7 +202,7 @@ pub fn get_request_type(request: &GreptimeRequest) -> &'static str {
pub(crate) fn create_query_context(
channel: Channel,
header: Option<&RequestHeader>,
mut extensions: Vec<(String, String)>,
extensions: Vec<(String, String)>,
snapshot_seqs: HashMap<u64, u64>,
) -> Result<QueryContextRef> {
let (catalog, schema) = header
@@ -240,29 +240,36 @@ pub(crate) fn create_query_context(
.channel(channel)
.snapshot_seqs(Arc::new(RwLock::new(snapshot_seqs)));
if let Some(x) = extensions
.iter()
.position(|(k, _)| k == READ_PREFERENCE_HINT)
{
let (k, v) = extensions.swap_remove(x);
let Ok(read_preference) = ReadPreference::from_str(&v) else {
return UnknownHintSnafu {
hint: format!("{k}={v}"),
}
.fail();
};
ctx_builder = ctx_builder.read_preference(read_preference);
}
for (key, value) in extensions {
if is_reserved_extension_key(&key) {
debug!(
key = key.as_str(),
"Ignoring reserved external query context extension key"
);
continue;
match key.as_str() {
READ_PREFERENCE_HINT => {
let Ok(read_preference) = ReadPreference::from_str(&value) else {
return UnknownHintSnafu {
hint: format!("{key}={value}"),
}
.fail();
};
ctx_builder = ctx_builder.read_preference(read_preference);
}
INSERT_SKIP_WAL_HINT => {
let skip_wal = value.parse::<bool>().map_err(|_| {
UnknownHintSnafu {
hint: format!("{key}={value}"),
}
.build()
})?;
ctx_builder = ctx_builder.skip_wal(skip_wal);
}
_ if is_reserved_extension_key(&key) => {
debug!(
key = key.as_str(),
"Ignoring reserved external query context extension key"
);
}
_ => {
ctx_builder = ctx_builder.set_extension(key, value);
}
}
ctx_builder = ctx_builder.set_extension(key, value);
}
Ok(ctx_builder.build().into())
}
@@ -322,6 +329,81 @@ mod tests {
use super::*;
use crate::error::{ExecuteGrpcRequestSnafu, InvalidParameterSnafu};
#[test]
fn test_create_query_context_typed_skip_wal() {
let ctx = create_query_context(Channel::Grpc, None, vec![], HashMap::new()).unwrap();
assert!(!ctx.skip_wal());
let legacy = create_query_context(
Channel::Grpc,
None,
vec![("skip_wal".to_string(), "true".to_string())],
HashMap::new(),
)
.unwrap();
assert!(!legacy.skip_wal());
assert_eq!(legacy.extension("skip_wal"), Some("true"));
for (value, expected) in [("true", true), ("false", false)] {
let ctx = create_query_context(
Channel::Grpc,
None,
vec![(INSERT_SKIP_WAL_HINT.to_string(), value.to_string())],
HashMap::new(),
)
.unwrap();
assert_eq!(ctx.skip_wal(), expected);
assert_eq!(ctx.extension(INSERT_SKIP_WAL_HINT), None);
}
for value in ["", "TRUE", "1", "invalid"] {
assert!(
create_query_context(
Channel::Grpc,
None,
vec![(INSERT_SKIP_WAL_HINT.to_string(), value.to_string())],
HashMap::new()
)
.is_err()
);
}
let ctx = create_query_context(
Channel::Grpc,
None,
vec![
(INSERT_SKIP_WAL_HINT.to_string(), "true".to_string()),
(INSERT_SKIP_WAL_HINT.to_string(), "false".to_string()),
],
HashMap::new(),
)
.unwrap();
assert!(!ctx.skip_wal());
assert_eq!(ctx.extension(INSERT_SKIP_WAL_HINT), None);
}
#[test]
fn test_create_query_context_read_preference_duplicates() {
for (values, valid) in [
(["leader", "LEADER"], true),
(["invalid", "leader"], false),
(["leader", "invalid"], false),
] {
let result = create_query_context(
Channel::Grpc,
None,
values
.into_iter()
.map(|value| (READ_PREFERENCE_HINT.to_string(), value.to_string()))
.collect(),
HashMap::new(),
);
if valid {
let ctx = result.unwrap();
assert!(matches!(ctx.read_preference(), ReadPreference::Leader));
assert_eq!(ctx.extension(READ_PREFERENCE_HINT), None);
} else {
assert!(result.is_err());
}
}
}
#[test]
fn test_create_query_context() {
let header = RequestHeader {
+16
View File
@@ -59,6 +59,22 @@ mod tests {
use super::*;
#[test]
fn test_extract_skip_wal_hint() {
use session::hints::INSERT_SKIP_WAL_HINT;
let mut headers = HeaderMap::new();
headers.insert(HINTS_KEY, HeaderValue::from_static("insert_skip_wal=true"));
let mut metadata = MetadataMap::new();
metadata.insert(
HINTS_KEY,
MetadataValue::from_static("insert_skip_wal=true"),
);
let expected = vec![(INSERT_SKIP_WAL_HINT.to_string(), "true".to_string())];
assert_eq!(extract_hints(&headers), expected);
assert_eq!(extract_hints(&metadata), expected);
}
#[test]
fn test_extract_hints_with_full_header_map() {
let mut headers = HeaderMap::new();
+3 -1
View File
@@ -117,6 +117,7 @@ mod client_ip;
use crate::prom_remote_write::validation::PromValidationMode;
mod hints;
mod read_preference;
mod skip_wal;
#[cfg(any(test, feature = "testing"))]
pub mod test_helpers;
@@ -1059,7 +1060,8 @@ impl HttpServer {
.layer(middleware::from_fn(client_ip::log_error_with_client_ip))
.layer(middleware::from_fn(
read_preference::extract_read_preference,
)),
))
.layer(middleware::from_fn(skip_wal::extract_skip_wal)),
);
// Debug handlers are part of the complete router; the API listener hides
+5
View File
@@ -44,6 +44,7 @@ pub mod constants {
pub const GREPTIME_DB_HEADER_METRICS: &str = "x-greptime-metrics";
pub const GREPTIME_DB_HEADER_NAME: &str = "x-greptime-db-name";
pub const GREPTIME_DB_HEADER_READ_PREFERENCE: &str = "x-greptime-read-preference";
pub const GREPTIME_INSERT_SKIP_WAL_HEADER_NAME: &str = "x-greptime-insert-skip-wal";
pub const GREPTIME_TIMEZONE_HEADER_NAME: &str = "x-greptime-timezone";
pub const GREPTIME_DB_HEADER_ERROR_CODE: &str = common_error::GREPTIME_DB_HEADER_ERROR_CODE;
@@ -94,6 +95,10 @@ pub static GREPTIME_TIMEZONE_HEADER_NAME: HeaderName =
pub static GREPTIME_DB_HEADER_READ_PREFERENCE: HeaderName =
HeaderName::from_static(constants::GREPTIME_DB_HEADER_READ_PREFERENCE);
/// Request-level WAL policy, independent of the table-level skip_wal option.
pub static GREPTIME_INSERT_SKIP_WAL_HEADER_NAME: HeaderName =
HeaderName::from_static(constants::GREPTIME_INSERT_SKIP_WAL_HEADER_NAME);
pub static CONTENT_TYPE_PROTOBUF_STR: &str = "application/x-protobuf";
pub static CONTENT_TYPE_PROTOBUF: HeaderValue = HeaderValue::from_static(CONTENT_TYPE_PROTOBUF_STR);
pub static CONTENT_ENCODING_SNAPPY: HeaderValue = HeaderValue::from_static("snappy");
+56
View File
@@ -0,0 +1,56 @@
// Copyright 2023 Greptime Team
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
use axum::body::Body;
use axum::http::Request;
use axum::middleware::Next;
use axum::response::{IntoResponse, Response};
use session::context::QueryContext;
use crate::error::InvalidParameterSnafu;
use crate::http::header::GREPTIME_INSERT_SKIP_WAL_HEADER_NAME;
use crate::http::result::error_result::ErrorResponse;
/// Extract the request-level WAL policy from the dedicated HTTP header.
pub async fn extract_skip_wal(mut request: Request<Body>, next: Next) -> Response {
let skip_wal = match request.headers().get(&GREPTIME_INSERT_SKIP_WAL_HEADER_NAME) {
None => false,
Some(value) => match value
.to_str()
.ok()
.and_then(|value| value.parse::<bool>().ok())
{
Some(skip_wal) => skip_wal,
None => {
return (
http::StatusCode::BAD_REQUEST,
ErrorResponse::from_error(
InvalidParameterSnafu {
reason: format!(
"{} must be true or false",
GREPTIME_INSERT_SKIP_WAL_HEADER_NAME
),
}
.build(),
),
)
.into_response();
}
},
};
if let Some(query_ctx) = request.extensions_mut().get_mut::<QueryContext>() {
query_ctx.set_skip_wal(skip_wal);
}
next.run(request).await
}
@@ -980,3 +980,102 @@ async fn get_body(response: Response) -> Bytes {
.await
.unwrap()
}
#[tokio::test]
async fn test_http_skip_wal_header() {
struct SkipWalSqlHandler {
inner: ServerSqlQueryHandlerRef,
observed: AtomicUsize,
}
#[async_trait]
impl SqlQueryHandler for SkipWalSqlHandler {
async fn do_query(&self, query: &str, ctx: QueryContextRef) -> Vec<Result<Output>> {
self.observed
.store(usize::from(ctx.skip_wal()), Ordering::Relaxed);
self.inner.do_query(query, ctx).await
}
async fn do_analyze_stream_query(
&self,
query: &str,
ctx: QueryContextRef,
) -> Result<Output> {
self.inner.do_analyze_stream_query(query, ctx).await
}
async fn do_exec_plan(
&self,
plan: LogicalPlan,
stmt: Option<Statement>,
ctx: QueryContextRef,
) -> Result<Output> {
self.inner.do_exec_plan(plan, stmt, ctx).await
}
async fn do_promql_query(
&self,
query: &PromQuery,
ctx: QueryContextRef,
) -> Vec<Result<Output>> {
self.inner.do_promql_query(query, ctx).await
}
async fn do_describe(
&self,
stmt: Statement,
ctx: QueryContextRef,
) -> Result<Option<DescribeResult>> {
self.inner.do_describe(stmt, ctx).await
}
async fn is_valid_schema(&self, catalog: &str, schema: &str) -> Result<bool> {
self.inner.is_valid_schema(catalog, schema).await
}
}
let handler = Arc::new(SkipWalSqlHandler {
inner: create_testing_sql_query_handler(MemTable::default_numbers_table()),
observed: AtomicUsize::new(2),
});
let server = HttpServerBuilder::new(HttpOptions::default())
.with_sql_handler(handler.clone())
.build();
let client = TestClient::new(server.build(server.make_app()).unwrap()).await;
for (header, expected) in [
(Some(("x-greptime-insert-skip-wal", "true")), Some(true)),
(Some(("x-greptime-insert-skip-wal", "false")), Some(false)),
(None, Some(false)),
(Some(("x-greptime-insert-skip-wal", "yes")), None),
(Some(("x-greptime-insert-skip-wal", "1")), None),
(Some(("x-greptime-insert-skip-wal", "")), None),
(Some(("x-greptime-skip-wal", "true")), Some(false)),
(Some(("x-greptime-hints", "skip_wal=true")), Some(false)),
(
Some(("x-greptime-hints", "insert_skip_wal=true")),
Some(false),
),
] {
// Sentinel also proves invalid values are rejected before the SQL handler.
handler.observed.store(2, Ordering::Relaxed);
let mut request = client.get("/v1/sql?sql=SELECT%201");
if let Some((key, value)) = header {
request = request.header(key, value);
}
let response = request.send().await;
match expected {
Some(expected) => {
assert_eq!(response.status(), StatusCode::OK, "{header:?}");
assert_eq!(
handler.observed.load(Ordering::Relaxed),
usize::from(expected),
"{header:?}"
);
}
None => {
assert_eq!(response.status(), StatusCode::BAD_REQUEST, "{header:?}");
assert_eq!(handler.observed.load(Ordering::Relaxed), 2);
}
}
}
}
+49
View File
@@ -153,6 +153,15 @@ impl QueryContextBuilder {
self
}
pub fn skip_wal(mut self, skip_wal: bool) -> Self {
self.mutable_session_data
.get_or_insert_default()
.write()
.unwrap()
.skip_wal = skip_wal;
self
}
pub fn read_preference(mut self, read_preference: ReadPreference) -> Self {
self.mutable_session_data
.get_or_insert_default()
@@ -348,6 +357,15 @@ impl QueryContext {
self.mutable_session_data.write().unwrap().timezone = timezone;
}
/// Returns whether ordinary inserts in this request should skip WAL.
pub fn skip_wal(&self) -> bool {
self.mutable_session_data.read().unwrap().skip_wal
}
pub fn set_skip_wal(&self, skip_wal: bool) {
self.mutable_session_data.write().unwrap().skip_wal = skip_wal;
}
pub fn read_preference(&self) -> ReadPreference {
self.mutable_session_data.read().unwrap().read_preference
}
@@ -739,6 +757,37 @@ mod test {
assert_eq!(fork.current_schema(), "private");
}
#[test]
fn test_skip_wal_default_builder_and_fork() {
let default_context = QueryContext::with(DEFAULT_CATALOG_NAME, "public");
assert!(!default_context.skip_wal());
assert!(!QueryContextBuilder::default().build().skip_wal());
let context = QueryContextBuilder::default().skip_wal(true).build();
assert!(context.skip_wal());
let fork = context.fork();
assert!(fork.skip_wal());
fork.set_skip_wal(false);
assert!(context.skip_wal());
assert!(!fork.skip_wal());
context.set_skip_wal(false);
fork.set_skip_wal(true);
assert!(!context.skip_wal());
assert!(fork.skip_wal());
}
#[test]
fn test_skip_wal_is_not_serialized_in_query_context() {
let context = QueryContextBuilder::default().skip_wal(true).build();
let api_context: api::v1::QueryContext = context.into();
assert!(
!api_context
.extensions
.contains_key(crate::hints::INSERT_SKIP_WAL_HINT)
);
let restored: QueryContext = api_context.into();
assert!(!restored.skip_wal());
}
#[test]
fn test_api_query_context_roundtrip_with_sequences() {
let api_ctx = api::v1::QueryContext {
+3
View File
@@ -23,6 +23,9 @@ pub const SUPPORT_FLIGHT_METRICS_BEFORE_BATCH_EXTENSION_KEY: &str =
"query.support_flight_metrics_before_batch";
pub const LIVE_ANALYZE_METRICS_EXTENSION_KEY: &str = "query.live_analyze_metrics";
/// Skip WAL for this insert only; never persisted as a table option.
pub const INSERT_SKIP_WAL_HINT: &str = "insert_skip_wal";
pub const READ_PREFERENCE_HINT: &str = "read_preference";
pub const RESERVED_EXTENSION_KEYS: [&str; 4] = [
REMOTE_QUERY_ID_EXTENSION_KEY,
+3
View File
@@ -60,6 +60,8 @@ pub(crate) struct MutableInner {
timezone: Timezone,
query_timeout: Option<Duration>,
read_preference: ReadPreference,
/// Request-level WAL policy for ordinary inserts.
skip_wal: bool,
#[debug(skip)]
pub(crate) cursors: HashMap<String, Arc<RecordBatchStreamCursor>>,
/// Warning messages for MySQL SHOW WARNINGS support
@@ -74,6 +76,7 @@ impl Default for MutableInner {
timezone: get_timezone(None).clone(),
query_timeout: None,
read_preference: ReadPreference::Leader,
skip_wal: false,
cursors: HashMap::with_capacity(0),
warnings: VecDeque::new(),
}
+34
View File
@@ -221,6 +221,7 @@ fn make_region_puts(inserts: InsertRequests) -> Result<Vec<(RegionId, RegionRequ
RegionRequest::Put(RegionPutRequest {
rows,
hint: None,
skip_wal: r.skip_wal,
partition_expr_version: r.partition_expr_version.map(|v| v.value),
}),
)
@@ -524,6 +525,9 @@ pub struct RegionPutRequest {
pub rows: Rows,
/// Write hint.
pub hint: Option<WriteHint>,
/// Skip WAL for this insert without changing region options.
/// Metadata writes must not inherit this option from user inserts.
pub skip_wal: bool,
/// Partition expression version for the region.
pub partition_expr_version: Option<u64>,
}
@@ -1864,6 +1868,36 @@ mod tests {
use super::*;
use crate::metadata::RegionMetadataBuilder;
#[test]
fn test_make_region_puts_preserves_skip_wal() {
let region_id = RegionId::new(42, 3);
let rows = Rows::default();
let requests = make_region_puts(InsertRequests {
requests: [false, true, false]
.into_iter()
.map(|skip_wal| api::v1::region::InsertRequest {
region_id: region_id.as_u64(),
rows: Some(rows.clone()),
partition_expr_version: Some(api::v1::PartitionExprVersion { value: 7 }),
skip_wal,
})
.collect(),
})
.unwrap();
assert_eq!(3, requests.len());
for ((id, request), skip_wal) in requests.into_iter().zip([false, true, false]) {
assert_eq!(region_id, id);
let RegionRequest::Put(request) = request else {
panic!("expected a put request");
};
assert_eq!(rows, request.rows);
assert_eq!(skip_wal, request.skip_wal);
assert_eq!(Some(7), request.partition_expr_version);
assert!(request.hint.is_none());
}
}
#[test]
fn test_make_region_compact_with_time_range() {
let requests = make_region_compact(CompactRequest {
+2
View File
@@ -714,6 +714,8 @@ pub struct InsertRequest {
pub schema_name: String,
pub table_name: String,
pub columns_values: HashMap<String, VectorRef>,
/// Whether this insert should skip WAL.
pub skip_wal: bool,
}
/// Delete (by primary key) request