use std::sync::Arc;
use crate::catalog::providers::CatalogProvider;
use crate::cnf::ConfigMap;
use crate::dbs::Session;
use crate::kvs::{Datastore, LockType, TransactionType};
async fn guarded_ds(limit: u64) -> (Arc<Datastore>, Session) {
let config = ConfigMap::empty().with_key_value("transaction_max_write_keys", limit.to_string());
let ds = Datastore::builder().with_config(config).build_with_path("memory").await.unwrap();
{
let tx = ds.transaction(TransactionType::Write, LockType::Optimistic).await.unwrap();
tx.ensure_ns_db(None, "test", "test").await.unwrap();
tx.commit().await.unwrap();
}
(Arc::new(ds), Session::owner().with_ns("test").with_db("test"))
}
async fn run(ds: &Datastore, ses: &Session, sql: &str) {
let res = ds.execute(sql, ses, None).await.unwrap();
for (i, r) in res.into_iter().enumerate() {
if let Err(e) = r.result {
panic!("statement {i} failed: {e}");
}
}
}
async fn run_capture_err(ds: &Datastore, ses: &Session, sql: &str) -> Option<String> {
let res = ds.execute(sql, ses, None).await.unwrap();
res.into_iter().find_map(|r| r.result.err().map(|e| e.to_string()))
}
async fn setup_cascade(ds: &Datastore, ses: &Session, chunks: usize, words: usize) {
run(ds, ses, "DEFINE ANALYZER az TOKENIZERS blank FILTERS lowercase;").await;
run(
ds,
ses,
"DEFINE TABLE document;
DEFINE TABLE chunk;
DEFINE FIELD document ON chunk TYPE record<document> REFERENCE ON DELETE CASCADE;
DEFINE INDEX chunk_text ON chunk FIELDS text FULLTEXT ANALYZER az BM25;",
)
.await;
run(ds, ses, "CREATE document:doc1;").await;
for ci in 0..chunks {
let text: String = (0..words).map(|w| format!("c{ci}w{w}")).collect::<Vec<_>>().join(" ");
run(ds, ses, &format!("CREATE chunk:c{ci} SET document = document:doc1, text = '{text}';"))
.await;
}
}
async fn assert_intact(ds: &Datastore, ses: &Session, chunks: usize) {
let res = ds
.execute(
"SELECT count() FROM chunk GROUP ALL;
SELECT count() FROM document GROUP ALL;
SELECT id FROM chunk WHERE text @@ 'c0w0';",
ses,
None,
)
.await
.unwrap();
let values: Vec<String> = res.into_iter().map(|r| format!("{:?}", r.result.unwrap())).collect();
assert!(
values[0].contains(&format!("Int({chunks})")),
"chunk count changed after rolled-back delete: {}",
values[0]
);
assert!(values[1].contains("Int(1)"), "document missing after rolled-back delete");
assert!(values[2].contains("c0"), "full-text index lost chunk c0: {}", values[2]);
}
#[tokio::test(flavor = "multi_thread")]
async fn write_guard_aborts_oversized_cascade_atomically() {
let (ds, ses) = guarded_ds(200).await;
setup_cascade(&ds, &ses, 10, 50).await;
let err = run_capture_err(&ds, &ses, "DELETE document:doc1;")
.await
.expect("over-limit delete should fail");
assert!(err.contains("maximum number of key writes (200)"), "unexpected error message: {err}");
assert_intact(&ds, &ses, 10).await;
}
#[tokio::test(flavor = "multi_thread")]
async fn write_guard_allows_within_limit() {
let (ds, ses) = guarded_ds(100_000).await;
setup_cascade(&ds, &ses, 10, 50).await;
run(&ds, &ses, "DELETE document:doc1;").await;
let res = ds.execute("SELECT count() FROM chunk GROUP ALL;", &ses, None).await.unwrap();
let v = format!("{:?}", res.into_iter().next().unwrap().result.unwrap());
assert!(v.contains("Int(0)"), "cascade should have removed all chunks: {v}");
}
#[tokio::test(flavor = "multi_thread")]
async fn write_guard_applies_in_begin_blocks() {
let (ds, ses) = guarded_ds(200).await;
setup_cascade(&ds, &ses, 10, 50).await;
let err = run_capture_err(&ds, &ses, "BEGIN; DELETE document:doc1; COMMIT;")
.await
.expect("over-limit delete inside a block should fail");
assert!(err.contains("maximum number of key writes"), "unexpected error: {err}");
assert_intact(&ds, &ses, 10).await;
}
#[tokio::test(flavor = "multi_thread")]
async fn write_guard_charges_range_deletes() {
let ds = Datastore::new("memory").await.unwrap();
let tx = ds
.transaction(TransactionType::Write, LockType::Optimistic)
.await
.unwrap()
.with_write_keys_limit(std::num::NonZeroU64::new(2));
let range = |a: &str, b: &str| a.as_bytes().to_vec()..b.as_bytes().to_vec();
tx.delr(range("za", "zb")).await.unwrap();
tx.delr(range("zb", "zc")).await.unwrap();
let err = tx.delr(range("zc", "zd")).await.expect_err("third range delete should trip");
assert!(
err.to_string().contains("maximum number of key writes (2)"),
"unexpected error: {err}"
);
tx.cancel().await.unwrap();
}
#[tokio::test(flavor = "multi_thread")]
async fn write_guard_counts_changefeed_writes() {
let (ds, ses) = guarded_ds(4).await;
run(&ds, &ses, "DEFINE TABLE t CHANGEFEED 1h; DEFINE TABLE u CHANGEFEED 1h;").await;
let errs: Vec<String> = ds
.execute("BEGIN; CREATE |t:2| RETURN NONE; CREATE |u:2| RETURN NONE; COMMIT;", &ses, None)
.await
.unwrap()
.into_iter()
.filter_map(|r| r.result.err().map(|e| e.to_string()))
.collect();
assert!(
errs.iter().any(|e| e.contains("maximum number of key writes (4)")),
"expected the guard error among the block errors: {errs:?}"
);
let res = ds.execute("SELECT count() FROM t GROUP ALL;", &ses, None).await.unwrap();
let v = format!("{:?}", res.into_iter().next().unwrap().result.unwrap());
assert!(
v.contains("Int(0)") || v == "Array(Array([]))",
"rolled-back create left records: {v}"
);
}
#[tokio::test(flavor = "multi_thread")]
async fn write_guard_applies_to_record_access_clauses() {
let (ds, ses) = guarded_ds(5).await;
run(&ds, &ses, "DEFINE TABLE t;").await;
let clause: crate::expr::Expr =
crate::syn::expr("{ CREATE |t:50| RETURN NONE; RETURN 1; }").unwrap().into();
let err =
ds.evaluate(&clause, &ses, None).await.expect_err("an over-limit access clause must fail");
assert!(
err.to_string().contains("maximum number of key writes (5)"),
"unexpected error: {err}"
);
let res = ds.execute("SELECT count() FROM t GROUP ALL;", &ses, None).await.unwrap();
let v = format!("{:?}", res.into_iter().next().unwrap().result.unwrap());
assert!(
v.contains("Int(0)") || v == "Array(Array([]))",
"rolled-back clause left records: {v}"
);
}
#[tokio::test(flavor = "multi_thread")]
async fn write_guard_applies_to_external_transactions() {
let (ds, ses) = guarded_ds(200).await;
setup_cascade(&ds, &ses, 10, 50).await;
let tx = Arc::new(ds.transaction(TransactionType::Write, LockType::Optimistic).await.unwrap());
let results = ds
.execute_with_transaction("DELETE document:doc1;", &ses, None, Arc::clone(&tx))
.await
.unwrap();
let errs: Vec<String> =
results.into_iter().filter_map(|r| r.result.err().map(|e| e.to_string())).collect();
assert!(
errs.iter().any(|e| e.contains("maximum number of key writes (200)")),
"expected the guard error on the external transaction: {errs:?}"
);
tx.cancel().await.unwrap();
assert_intact(&ds, &ses, 10).await;
}
#[tokio::test(flavor = "multi_thread")]
async fn write_guard_poisons_external_transactions_on_trip() {
let (ds, ses) = guarded_ds(200).await;
setup_cascade(&ds, &ses, 10, 50).await;
let tx = Arc::new(ds.transaction(TransactionType::Write, LockType::Optimistic).await.unwrap());
let results = ds
.execute_with_transaction("DELETE document:doc1;", &ses, None, Arc::clone(&tx))
.await
.unwrap();
assert!(
results.into_iter().any(|r| r.result.is_err()),
"the over-limit statement should have failed"
);
let err = tx.commit().await.expect_err("commit of a poisoned transaction must fail");
assert!(
err.to_string().contains("maximum number of key writes (200)"),
"unexpected commit error: {err}"
);
assert_intact(&ds, &ses, 10).await;
}
#[tokio::test(flavor = "multi_thread")]
async fn write_guard_applies_to_api_handlers() {
use crate::api::request::ApiRequest;
use crate::catalog::ApiMethod;
let (ds, ses) = guarded_ds(5).await;
run(
&ds,
&ses,
r#"
DEFINE TABLE t;
DEFINE API "/spam" FOR get PERMISSIONS FULL THEN {
CREATE |t:50| RETURN NONE;
{ status: 200 };
};
"#,
)
.await;
let req = ApiRequest {
method: ApiMethod::Get,
request_id: "issue-715-guard".to_string(),
..Default::default()
};
let err = ds
.invoke_api_handler("test", "test", "spam", &ses, req)
.await
.expect_err("an over-limit API handler must fail");
assert!(
err.to_string().contains("maximum number of key writes (5)"),
"unexpected error: {err}"
);
let res = ds.execute("SELECT count() FROM t GROUP ALL;", &ses, None).await.unwrap();
let v = format!("{:?}", res.into_iter().next().unwrap().result.unwrap());
assert!(
v.contains("Int(0)") || v == "Array(Array([]))",
"rolled-back handler left records: {v}"
);
}
#[tokio::test(flavor = "multi_thread")]
async fn write_guard_holds_under_concurrent_writes() {
let ds = Datastore::new("memory").await.unwrap();
let tx = ds
.transaction(TransactionType::Write, LockType::Optimistic)
.await
.unwrap()
.with_write_keys_limit(std::num::NonZeroU64::new(2));
let key = |s: &str| s.as_bytes().to_vec();
let (ka, kb, kc) = (key("zz-a"), key("zz-b"), key("zz-c"));
let val = vec![1u8];
let res = futures::try_join!(tx.set(&ka, &val), tx.set(&kb, &val), tx.set(&kc, &val),);
let err = res.expect_err("three concurrent writes must not fit a limit of two");
assert!(
err.to_string().contains("maximum number of key writes (2)"),
"unexpected error: {err}"
);
tx.cancel().await.unwrap();
}
#[test]
fn write_guard_error_is_a_query_error() {
let err = crate::err::Error::TransactionWriteKeysExceeded {
limit: 5,
};
let types_err = crate::err::into_types_error(err);
assert_eq!(
types_err.kind_str(),
"Query",
"guard error must classify as a query error, got: {types_err:?}"
);
}
#[tokio::test(flavor = "multi_thread")]
async fn write_guard_ignores_internal_transactions() {
let (ds, ses) = guarded_ds(5).await;
{
let tx = ds.transaction(TransactionType::Write, LockType::Optimistic).await.unwrap();
for i in 0..100u32 {
let key = format!("zz-guard-test-{i}").into_bytes();
tx.set(&key, &vec![1u8]).await.unwrap();
}
tx.commit().await.unwrap();
}
let err = run_capture_err(&ds, &ses, "CREATE |t:10| RETURN NONE;")
.await
.expect("over-limit statement should fail");
assert!(err.contains("maximum number of key writes (5)"), "unexpected error: {err}");
}