diff --git a/.github/workflows/integration.yml b/.github/workflows/integration.yml index 5d1ab88a..d8d45e6d 100644 --- a/.github/workflows/integration.yml +++ b/.github/workflows/integration.yml @@ -297,12 +297,14 @@ jobs: # The control plane for vector indexes is not reachable over the wire while # this backend declares no vector search capability, so its tests drive the - # storage layer directly against this job's PostgreSQL. They build their own - # throwaway databases; the connection string is the server, not a database. + # storage layer directly against this job's PostgreSQL. The PutItem create + # race tests hold a create open in an outside transaction, which also needs + # direct access. They build their own throwaway databases; the connection + # string is the server, not a database. - name: Run PostgreSQL storage-level tests env: EXTENDDB_TEST_PG_CONNECTION_STRING: postgresql://postgres:devpass@127.0.0.1:5432 - run: cargo test --release -p extenddb-storage-postgres --test vector_control_plane + run: cargo test --release -p extenddb-storage-postgres --test vector_control_plane --test put_create_race # The daemonized server logs to syslog; dump it so server-side failures # are diagnosable from the job log. diff --git a/crates/storage-postgres/src/data/put_item.rs b/crates/storage-postgres/src/data/put_item.rs index e307df66..35e5b339 100755 --- a/crates/storage-postgres/src/data/put_item.rs +++ b/crates/storage-postgres/src/data/put_item.rs @@ -112,25 +112,28 @@ impl PostgresEngine { .await .map_err(|e| StorageError::Internal(e.to_string()))?; - let old: Option<(serde_json::Value,)> = + let mut old: Option<(serde_json::Value,)> = bind_sk_fetch_optional!(&select_sql, pk_text.as_str(), &sk, &mut *tx)?; - if let Some((ref old_json,)) = old { - let old_item: Item = json_to_item(old_json.clone())?; - match check_condition(condition, &old_item, maps) { - Ok(()) => {} - Err(StorageError::ConditionFailed(_)) => { - return Err(StorageError::ConditionFailed(Some(old_item))); + let mut attempt: u32 = 0; + loop { + if let Some((ref old_json,)) = old { + let old_item: Item = json_to_item(old_json.clone())?; + match check_condition(condition, &old_item, maps) { + Ok(()) => {} + Err(StorageError::ConditionFailed(_)) => { + return Err(StorageError::ConditionFailed(Some(old_item))); + } + Err(e) => return Err(e), } - Err(e) => return Err(e), + // Row exists, condition passed: update in place. + let update_sql = format!( + "UPDATE {ddb_table} SET item_data = $3 WHERE pk = $1 AND {sk_col} = $2" + ); + bind_sk_execute!(&update_sql, pk_text.as_str(), &sk, &item_json, &mut *tx)?; + break; } - // Row exists, condition passed — update in place. - let update_sql = format!( - "UPDATE {ddb_table} SET item_data = $3 WHERE pk = $1 AND {sk_col} = $2" - ); - bind_sk_execute!(&update_sql, pk_text.as_str(), &sk, &item_json, &mut *tx)?; - } else { - // No existing item — condition checks against empty item + // No existing item: the condition checks against an empty item let empty = std::collections::BTreeMap::new(); match check_condition(condition, &empty, maps) { Ok(()) => {} @@ -139,21 +142,25 @@ impl PostgresEngine { } Err(e) => return Err(e), } - // Condition passed against empty — atomic insert, fail if someone beat us. let insert_sql = format!( "INSERT INTO {ddb_table} (pk, {sk_col}, item_data) VALUES ($1, $2, $3) \ ON CONFLICT (pk, {sk_col}) DO NOTHING" ); let result = bind_sk_execute!(&insert_sql, pk_text.as_str(), &sk, &item_json, &mut *tx)?; - if result.rows_affected() == 0 { - // Another transaction inserted between our SELECT and INSERT. - // Fetch the winner to return with ConditionFailed. - let winner: Option<(serde_json::Value,)> = - bind_sk_fetch_optional!(&select_sql, pk_text.as_str(), &sk, &mut *tx)?; - let winner_item = winner.map(|(v,)| json_to_item(v)).transpose()?; - return Err(StorageError::ConditionFailed(winner_item)); + if result.rows_affected() == 1 { + break; + } + // Lost the create race. The winner has committed (the insert + // waited for it if it was in flight). The locking read returns + // its row, and this put overwrites it after checking its + // condition against it. If a delete committed since, the read + // returns no row and the insert is retried. + attempt += 1; + if attempt >= MAX_CREATE_RACE_ATTEMPTS { + return Err(create_race_exhausted(attempt)); } + old = bind_sk_fetch_optional!(&select_sql, pk_text.as_str(), &sk, &mut *tx)?; } // Sync GSI/LSI update within transaction (D-4). @@ -269,30 +276,37 @@ impl PostgresEngine { .await .map_err(|e| StorageError::Internal(e.to_string()))?; - let old: Option<(serde_json::Value,)> = sqlx::query_as(&select_sql) - .bind(pk_text.as_str()) - .fetch_optional(&mut *tx) - .await - .map_err(|e| StorageError::Internal(e.to_string()))?; - - if let Some((ref old_json,)) = old { - let old_item: Item = json_to_item(old_json.clone())?; - match check_condition(condition, &old_item, maps) { - Ok(()) => {} - Err(StorageError::ConditionFailed(_)) => { - return Err(StorageError::ConditionFailed(Some(old_item))); - } - Err(e) => return Err(e), - } - // Row exists, condition passed — update in place. - let update_sql = format!("UPDATE {ddb_table} SET item_data = $2 WHERE pk = $1"); - sqlx::query(&update_sql) + let fetch_old = async |tx: &mut sqlx::PgConnection| { + sqlx::query_as::<_, (serde_json::Value,)>(&select_sql) .bind(pk_text.as_str()) - .bind(&item_json) - .execute(&mut *tx) + .fetch_optional(tx) .await - .map_err(|e| StorageError::Internal(e.to_string()))?; - } else { + .map_err(|e| StorageError::Internal(e.to_string())) + }; + let mut old: Option<(serde_json::Value,)> = fetch_old(&mut tx).await?; + + let mut attempt: u32 = 0; + loop { + if let Some((ref old_json,)) = old { + let old_item: Item = json_to_item(old_json.clone())?; + match check_condition(condition, &old_item, maps) { + Ok(()) => {} + Err(StorageError::ConditionFailed(_)) => { + return Err(StorageError::ConditionFailed(Some(old_item))); + } + Err(e) => return Err(e), + } + // Row exists, condition passed: update in place. + let update_sql = + format!("UPDATE {ddb_table} SET item_data = $2 WHERE pk = $1"); + sqlx::query(&update_sql) + .bind(pk_text.as_str()) + .bind(&item_json) + .execute(&mut *tx) + .await + .map_err(|e| StorageError::Internal(e.to_string()))?; + break; + } let empty = std::collections::BTreeMap::new(); match check_condition(condition, &empty, maps) { Ok(()) => {} @@ -301,7 +315,6 @@ impl PostgresEngine { } Err(e) => return Err(e), } - // Condition passed against empty — atomic insert, fail if someone beat us. let insert_sql = format!( "INSERT INTO {ddb_table} (pk, item_data) VALUES ($1, $2) \ ON CONFLICT (pk) DO NOTHING" @@ -312,16 +325,15 @@ impl PostgresEngine { .execute(&mut *tx) .await .map_err(|e| StorageError::Internal(e.to_string()))?; - if result.rows_affected() == 0 { - // Another transaction inserted between our SELECT and INSERT. - let winner: Option<(serde_json::Value,)> = sqlx::query_as(&select_sql) - .bind(pk_text.as_str()) - .fetch_optional(&mut *tx) - .await - .map_err(|e| StorageError::Internal(e.to_string()))?; - let winner_item = winner.map(|(v,)| json_to_item(v)).transpose()?; - return Err(StorageError::ConditionFailed(winner_item)); + if result.rows_affected() == 1 { + break; } + // Lost the create race: overwrite the winner, as above. + attempt += 1; + if attempt >= MAX_CREATE_RACE_ATTEMPTS { + return Err(create_race_exhausted(attempt)); + } + old = fetch_old(&mut tx).await?; } // Sync GSI/LSI update within transaction (D-4). @@ -466,3 +478,17 @@ impl PostgresEngine { json_opt.map(json_to_item).transpose() } } + +/// Bound on insert retries when a put to a missing item keeps losing the +/// create race to a winner that is deleted again before the re-read. Same +/// bound as `UpdateItem`. +const MAX_CREATE_RACE_ATTEMPTS: u32 = 5; + +/// The error after `MAX_CREATE_RACE_ATTEMPTS` lost create races. Like the one +/// `UpdateItem` returns, it is an internal error (HTTP 500). +fn create_race_exhausted(attempt: u32) -> StorageError { + StorageError::Internal(format!( + "PutItem could not create the item after {attempt} attempts: each insert lost the \ + create race to a concurrent writer" + )) +} diff --git a/crates/storage-postgres/tests/put_create_race.rs b/crates/storage-postgres/tests/put_create_race.rs new file mode 100644 index 00000000..8b38f012 --- /dev/null +++ b/crates/storage-postgres/tests/put_create_race.rs @@ -0,0 +1,841 @@ +// Copyright 2026 ExtendDB contributors +// SPDX-License-Identifier: Apache-2.0 +//! Storage-level tests for a `PutItem` that races another writer to create +//! the same item. +//! +//! Amazon DynamoDB applies such puts one after another: the later put +//! overwrites the item, after its own condition is checked against it. These +//! tests hold the competing create open in an outside transaction, so the put +//! deterministically loses the insert and has to order itself after the winner. +//! The four gate tests, on hash and on hash and range tables, also delete the +//! winner before the put re-reads it, so the put retries its insert, and they +//! bound those retries. The last test runs the race on a table with an LSI and +//! a stream, and checks that both get the winner as the old image. +//! +//! Each test builds its own throwaway database, applies the shipped migrations +//! to it, and drops it when it passes. A failing test leaves its database behind +//! on purpose, named `eddb_putr_*`, so the state that failed can be inspected. +//! +//! Requires `EXTENDDB_TEST_PG_CONNECTION_STRING`, a base URL with no database +//! component (for example `postgresql://postgres@127.0.0.1:5432`), pointing at a +//! server whose role may create and drop databases. Without it every test here +//! reports a skip and passes, the same convention the wire suites use. + +use std::collections::BTreeMap; +use std::time::Duration; + +use extenddb_core::expression::{self, Expr, ExpressionMaps}; +use extenddb_core::types::{ + AttributeDefinition, AttributeValue, BillingMode, CreateTableInput, Item, KeySchemaElement, + KeyType, LsiInput, Projection, ProjectionType, ReturnValuesOnConditionCheckFailure, + ScalarAttributeType, StreamRecord, StreamSpecification, StreamViewType, TableKeyInfo, +}; +use extenddb_storage::error::StorageError; +use extenddb_storage::{DataEngine, StreamCapture, TableEngine, TransactWriteOp}; +use extenddb_storage_postgres::{PostgresConfig, PostgresEngine}; +use sqlx::postgres::PgPoolOptions; +use sqlx::{Connection, PgConnection, PgPool, Postgres, Transaction}; + +const ACCOUNT: &str = "123456789012"; +const REGION: &str = "us-east-1"; +const TABLE: &str = "t_put_create_race"; + +struct Scratch { + engine: PostgresEngine, + db: PgPool, + admin: PgPool, + db_name: String, +} + +impl Scratch { + async fn cleanup(self) { + let Scratch { + engine, + db, + admin, + db_name, + } = self; + drop(engine); + db.close().await; + sqlx::query(&format!( + "DROP DATABASE IF EXISTS \"{db_name}\" WITH (FORCE)" + )) + .execute(&admin) + .await + .expect("drop the scratch database"); + admin.close().await; + } +} + +fn base_conn() -> Option { + let conn = std::env::var("EXTENDDB_TEST_PG_CONNECTION_STRING").ok()?; + (!conn.trim().is_empty()).then(|| conn.trim_end_matches('/').to_owned()) +} + +fn skip(test: &str) { + eprintln!( + "SKIP {test}: EXTENDDB_TEST_PG_CONNECTION_STRING is not set, so there is no PostgreSQL \ + to build a scratch catalog in." + ); +} + +async fn scratch() -> Scratch { + let base = base_conn().expect("caller checks base_conn() first"); + let db_name = format!("eddb_putr_{}", uuid::Uuid::new_v4().simple())[..24].to_owned(); + let admin = PgPoolOptions::new() + .max_connections(1) + .connect(&format!("{base}/postgres")) + .await + .expect("connect to the postgres maintenance database"); + sqlx::query(&format!("CREATE DATABASE \"{db_name}\"")) + .execute(&admin) + .await + .expect("create the scratch database"); + let url = format!("{base}/{db_name}"); + let db = PgPoolOptions::new() + .max_connections(2) + .connect(&url) + .await + .expect("connect to the scratch database"); + for sql in [ + include_str!("../migrations/001_schema.sql"), + include_str!("../migrations/002_vector_indexes.sql"), + include_str!("../data_migrations/001_data_schema.sql"), + include_str!("../data_migrations/002_gsi_pending.sql"), + include_str!("../data_migrations/003_idempotency_account_scope.sql"), + include_str!("../data_migrations/004_vector_index_state.sql"), + ] { + sqlx::raw_sql(sql) + .execute(&db) + .await + .expect("apply a shipped migration"); + } + sqlx::query("UPDATE settings SET value = '0' WHERE key = 'control_plane_delay_seconds'") + .execute(&db) + .await + .expect("pin the control-plane delay to zero"); + sqlx::query("INSERT INTO accounts (account_id, account_name) VALUES ($1, $2)") + .bind(ACCOUNT) + .bind(format!("acct-{db_name}")) + .execute(&db) + .await + .expect("seed the account row"); + let engine = PostgresEngine::new( + &PostgresConfig { + connection_string: url, + pool_size: 10, + max_item_size_bytes: 400_000, + }, + REGION, + ) + .await + .expect("open a PostgresEngine on the scratch database"); + Scratch { + engine, + db, + admin, + db_name, + } +} + +fn s_key(name: &str, key_type: KeyType) -> (KeySchemaElement, AttributeDefinition) { + ( + KeySchemaElement { + attribute_name: name.to_owned(), + key_type, + }, + AttributeDefinition { + attribute_name: name.to_owned(), + attribute_type: ScalarAttributeType::S, + }, + ) +} + +/// Create the test table, keyed on `pk` and, when `range` is set, also on `sk`. +async fn table(s: &Scratch, range: bool) -> TableKeyInfo { + let mut keys = vec![s_key("pk", KeyType::Hash)]; + if range { + keys.push(s_key("sk", KeyType::Range)); + } + let (key_schema, attribute_definitions) = keys.into_iter().unzip(); + s.engine + .create_table( + ACCOUNT, + CreateTableInput { + table_name: TABLE.to_owned(), + key_schema, + attribute_definitions, + billing_mode: Some(BillingMode::PayPerRequest), + ..Default::default() + }, + ) + .await + .expect("create the table"); + s.engine + .table_key_info(ACCOUNT, TABLE) + .await + .expect("read the key info") +} + +fn item(range: bool, v: &str) -> Item { + let mut item = BTreeMap::from([ + ("pk".to_owned(), AttributeValue::S("c".to_owned())), + ("v".to_owned(), AttributeValue::S(v.to_owned())), + ]); + if range { + item.insert("sk".to_owned(), AttributeValue::S("1".to_owned())); + } + item +} + +fn condition(text: &str) -> Expr { + let tokens = expression::tokenize(text).expect("tokenize"); + expression::parse_condition(&tokens).expect("parse") +} + +async fn data_table(db: &PgPool) -> String { + let id: String = + sqlx::query_scalar("SELECT table_id FROM tables WHERE account_id = $1 AND table_name = $2") + .bind(ACCOUNT) + .bind(TABLE) + .fetch_one(db) + .await + .expect("look up the table id"); + format!("\"_ddb_{id}\"") +} + +/// Begin an outside transaction that has created `winner` but not committed: a +/// create in flight. Returns it and its backend pid. +async fn create_in_flight( + db: &PgPool, + table: &str, + winner: &Item, +) -> (Transaction<'static, Postgres>, i32) { + let mut creator = db.begin().await.expect("begin the outside transaction"); + let pid: i32 = sqlx::query_scalar("SELECT pg_backend_pid()") + .fetch_one(&mut *creator) + .await + .expect("read the backend pid"); + let data = serde_json::to_value(winner).expect("serialize the winner"); + let sql = if winner.contains_key("sk") { + format!("INSERT INTO {table} (pk, sk_s, item_data) VALUES ('c', '1', $1)") + } else { + format!("INSERT INTO {table} (pk, item_data) VALUES ('c', $1)") + }; + sqlx::query(&sql) + .bind(data) + .execute(&mut *creator) + .await + .expect("insert the winner"); + (creator, pid) +} + +/// Wait until a backend waits for a lock of kind `event` (`transactionid` or +/// `advisory`) that the backend `blocker` holds. Fails after 5 s, so the caller +/// can report what the put did instead. +async fn wait_until_blocked_by(db: &PgPool, blocker: i32, event: &str) -> Result<(), String> { + for _ in 0..500 { + let waiting: i64 = sqlx::query_scalar( + "SELECT count(*) FROM pg_stat_activity \ + WHERE datname = current_database() AND wait_event_type = 'Lock' \ + AND wait_event = $2 AND $1 = ANY(pg_blocking_pids(pid))", + ) + .bind(blocker) + .bind(event) + .fetch_one(db) + .await + .expect("read pg_stat_activity"); + if waiting > 0 { + return Ok(()); + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + Err(format!( + "no backend waited on the {event} lock of backend {blocker}" + )) +} + +/// Run `put` (ReturnValues ALL_OLD) while the winner's create is in flight, +/// commit the create once the put waits on it, and return the put's result. +async fn put_after_winner( + s: &Scratch, + key_info: &TableKeyInfo, + range: bool, + condition: Option<&Expr>, +) -> Result, StorageError> { + let table = data_table(&s.db).await; + let (creator, creator_pid) = create_in_flight(&s.db, &table, &item(range, "winner")).await; + let maps = ExpressionMaps::default(); + let put = s + .engine + .put_item(key_info, item(range, "loser"), true, condition, &maps, None); + let commit = async { + let waited = wait_until_blocked_by(&s.db, creator_pid, "transactionid").await; + creator.commit().await.expect("commit the create"); + waited + }; + let (result, waited) = + tokio::time::timeout(Duration::from_secs(30), async { tokio::join!(put, commit) }) + .await + .expect("the put finished"); + if let Err(e) = waited { + panic!("{e}; the put returned {result:?}"); + } + result +} + +async fn stored_v(s: &Scratch) -> String { + let table = data_table(&s.db).await; + sqlx::query_scalar(&format!( + "SELECT item_data->'v'->>'S' FROM {table} WHERE pk = 'c'" + )) + .fetch_one(&s.db) + .await + .expect("read the stored item") +} + +async fn overwrites_the_winner(test: &str, range: bool) { + if base_conn().is_none() { + return skip(test); + } + let s = scratch().await; + let key_info = table(&s, range).await; + let old = put_after_winner(&s, &key_info, range, None) + .await + .expect("the put succeeds"); + assert_eq!( + old, + Some(item(range, "winner")), + "ALL_OLD returns the winner" + ); + assert_eq!( + stored_v(&s).await, + "loser", + "the later put overwrote the winner" + ); + s.cleanup().await; +} + +#[tokio::test] +async fn a_put_that_loses_the_create_race_overwrites_the_winner() { + overwrites_the_winner( + "a_put_that_loses_the_create_race_overwrites_the_winner", + false, + ) + .await; +} + +#[tokio::test] +async fn a_put_that_loses_the_create_race_overwrites_the_winner_on_a_range_table() { + overwrites_the_winner( + "a_put_that_loses_the_create_race_overwrites_the_winner_on_a_range_table", + true, + ) + .await; +} + +#[tokio::test] +async fn a_lost_create_race_checks_the_condition_against_the_winner() { + // The condition holds for both the missing item and the winner, so the + // put succeeds after the winner instead of failing on the lost insert. + let test = "a_lost_create_race_checks_the_condition_against_the_winner"; + if base_conn().is_none() { + return skip(test); + } + let s = scratch().await; + let key_info = table(&s, false).await; + let cond = condition("attribute_not_exists(gone)"); + let old = put_after_winner(&s, &key_info, false, Some(&cond)) + .await + .expect("the condition holds for the winner too"); + assert_eq!(old, Some(item(false, "winner"))); + assert_eq!(stored_v(&s).await, "loser"); + s.cleanup().await; +} + +#[tokio::test] +async fn a_lost_create_race_fails_a_condition_the_winner_breaks() { + // `attribute_not_exists(pk)` holds for the missing item but not for the + // winner: the put fails its condition and returns the winner it saw. + let test = "a_lost_create_race_fails_a_condition_the_winner_breaks"; + if base_conn().is_none() { + return skip(test); + } + let s = scratch().await; + let key_info = table(&s, false).await; + let cond = condition("attribute_not_exists(pk)"); + match put_after_winner(&s, &key_info, false, Some(&cond)).await { + Err(StorageError::ConditionFailed(old)) => { + assert_eq!(old, Some(item(false, "winner"))); + } + other => panic!("expected ConditionFailed, got {other:?}"), + } + assert_eq!(stored_v(&s).await, "winner", "the failed put wrote nothing"); + s.cleanup().await; +} + +// The tests below take the arm where the winner is gone again when the put +// re-reads it. A BEFORE INSERT trigger parks the put's insert on an advisory +// lock, the gate. While the insert is parked, an outside transaction commits a +// winner, and a second one, the locker, locks it FOR UPDATE. When the gate +// opens, the insert loses at once: the winner is committed, and a row lock does +// not make an insert wait. The put's locking re-read then waits on the locker, +// which deletes the winner and commits, so the re-read returns no row. + +/// The advisory lock key of the gate. +const GATE: i64 = 0x5075_7452; + +/// Open a connection to the scratch database outside its pools. Returns it and +/// its backend pid. +async fn connect(s: &Scratch) -> (PgConnection, i32) { + let base = base_conn().expect("caller checks base_conn() first"); + let mut conn = PgConnection::connect(&format!("{base}/{}", s.db_name)) + .await + .expect("connect to the scratch database"); + let pid: i32 = sqlx::query_scalar("SELECT pg_backend_pid()") + .fetch_one(&mut conn) + .await + .expect("read the backend pid"); + (conn, pid) +} + +/// Make every insert into `table` pass the gate first, unless its transaction +/// sets `test.bypass_gate`. Passing takes the gate shared and releases it, so +/// an insert parks while a session holds the gate exclusively. +async fn install_gate(s: &Scratch, table: &str) { + sqlx::raw_sql(&format!( + "CREATE FUNCTION pass_gate() RETURNS trigger LANGUAGE plpgsql AS $$ BEGIN \ + IF current_setting('test.bypass_gate', true) = 'on' THEN RETURN NEW; END IF; \ + PERFORM pg_advisory_lock_shared({GATE}); \ + PERFORM pg_advisory_unlock_shared({GATE}); \ + RETURN NEW; \ + END $$; \ + CREATE TRIGGER pass_gate BEFORE INSERT ON {table} \ + FOR EACH ROW EXECUTE FUNCTION pass_gate();" + )) + .execute(&s.db) + .await + .expect("install the gate trigger"); +} + +/// Commit the item with `v` as its value, past the gate. +async fn commit_winner(s: &Scratch, table: &str, range: bool, v: &str) { + let mut tx = s.db.begin().await.expect("begin the winner"); + sqlx::query("SET LOCAL test.bypass_gate = 'on'") + .execute(&mut *tx) + .await + .expect("bypass the gate"); + let sql = if range { + format!("INSERT INTO {table} (pk, sk_s, item_data) VALUES ('c', '1', $1)") + } else { + format!("INSERT INTO {table} (pk, item_data) VALUES ('c', $1)") + }; + sqlx::query(&sql) + .bind(serde_json::to_value(item(range, v)).expect("serialize the winner")) + .execute(&mut *tx) + .await + .expect("insert the winner"); + tx.commit().await.expect("commit the winner"); +} + +/// With the put's insert parked at the gate, make it lose to a winner `v` that +/// is deleted before the put re-reads it. With `rearm`, the gate closes again +/// behind the insert, so the put's next insert parks too. +async fn lose_to_a_vanishing_winner( + s: &Scratch, + table: &str, + range: bool, + gate: &mut PgConnection, + v: &str, + rearm: bool, +) -> Result<(), String> { + commit_winner(s, table, range, v).await; + let (mut locker, locker_pid) = connect(s).await; + sqlx::query("BEGIN") + .execute(&mut locker) + .await + .expect("begin the locker"); + let locked: Option<(serde_json::Value,)> = sqlx::query_as(&format!( + "SELECT item_data FROM {table} WHERE pk = 'c' FOR UPDATE" + )) + .fetch_optional(&mut locker) + .await + .expect("lock the winner"); + assert!(locked.is_some(), "the locker holds the committed winner"); + sqlx::query("SELECT pg_advisory_unlock($1)") + .bind(GATE) + .execute(&mut *gate) + .await + .expect("open the gate"); + if rearm { + // Granted only after the parked insert has passed the gate. + sqlx::query("SELECT pg_advisory_lock($1)") + .bind(GATE) + .execute(&mut *gate) + .await + .expect("close the gate again"); + } + wait_until_blocked_by(&s.db, locker_pid, "transactionid").await?; + sqlx::query(&format!("DELETE FROM {table} WHERE pk = 'c'")) + .execute(&mut locker) + .await + .expect("delete the winner"); + sqlx::query("COMMIT") + .execute(&mut locker) + .await + .expect("commit the delete"); + Ok(()) +} + +/// Install the gate on a fresh table, keyed on `pk` and, when `range` is set, +/// also on `sk`, and close the gate. Returns the table's key info, its data +/// table, and the session that holds the gate. +async fn gated_table(s: &Scratch, range: bool) -> (TableKeyInfo, String, PgConnection, i32) { + let key_info = table(s, range).await; + let data = data_table(&s.db).await; + install_gate(s, &data).await; + let (mut gate, gate_pid) = connect(s).await; + sqlx::query("SELECT pg_advisory_lock($1)") + .bind(GATE) + .execute(&mut gate) + .await + .expect("close the gate"); + (key_info, data, gate, gate_pid) +} + +/// Release every advisory lock `gate` holds. A driver that stops early calls +/// this, so that a parked insert goes on and the put can report its result. +async fn open_gate_fully(gate: &mut PgConnection) { + sqlx::query("SELECT pg_advisory_unlock_all()") + .execute(gate) + .await + .expect("open the gate"); +} + +async fn deleted_winner_is_retried(test: &str, range: bool) { + if base_conn().is_none() { + return skip(test); + } + let s = scratch().await; + let (key_info, data, mut gate, gate_pid) = gated_table(&s, range).await; + let maps = ExpressionMaps::default(); + let put = s + .engine + .put_item(&key_info, item(range, "loser"), true, None, &maps, None); + let driver = async { + let steps = async { + wait_until_blocked_by(&s.db, gate_pid, "advisory").await?; + lose_to_a_vanishing_winner(&s, &data, range, &mut gate, "winner", false).await + } + .await; + open_gate_fully(&mut gate).await; + steps + }; + let (result, driven) = + tokio::time::timeout(Duration::from_secs(30), async { tokio::join!(put, driver) }) + .await + .expect("the put finished"); + if let Err(e) = driven { + panic!("{e}; the put returned {result:?}"); + } + assert_eq!( + result.expect("the retried insert succeeds"), + None, + "the winner is gone, so there is no old image" + ); + assert_eq!(stored_v(&s).await, "loser"); + s.cleanup().await; +} + +#[tokio::test] +async fn a_put_whose_create_race_winner_is_deleted_creates_the_item() { + deleted_winner_is_retried( + "a_put_whose_create_race_winner_is_deleted_creates_the_item", + false, + ) + .await; +} + +#[tokio::test] +async fn a_put_whose_create_race_winner_is_deleted_creates_the_item_on_a_range_table() { + deleted_winner_is_retried( + "a_put_whose_create_race_winner_is_deleted_creates_the_item_on_a_range_table", + true, + ) + .await; +} + +async fn create_race_churn_gives_up(test: &str, range: bool) { + if base_conn().is_none() { + return skip(test); + } + let s = scratch().await; + let (key_info, data, mut gate, gate_pid) = gated_table(&s, range).await; + let maps = ExpressionMaps::default(); + let put = s + .engine + .put_item(&key_info, item(range, "loser"), true, None, &maps, None); + let driver = async { + let steps = async { + for k in 1..=4 { + wait_until_blocked_by(&s.db, gate_pid, "advisory").await?; + lose_to_a_vanishing_winner(&s, &data, range, &mut gate, &format!("w{k}"), true) + .await?; + } + // The fifth insert loses to a winner that stays: the put gives up. + wait_until_blocked_by(&s.db, gate_pid, "advisory").await?; + commit_winner(&s, &data, range, "w5").await; + Ok::<(), String>(()) + } + .await; + open_gate_fully(&mut gate).await; + steps + }; + let (result, driven) = + tokio::time::timeout(Duration::from_secs(30), async { tokio::join!(put, driver) }) + .await + .expect("the put finished"); + if let Err(e) = driven { + panic!("{e}; the put returned {result:?}"); + } + match result { + Err(StorageError::Internal(m)) => assert!(m.contains("after 5 attempts"), "{m}"), + other => panic!("expected Internal after 5 inserts, got {other:?}"), + } + assert_eq!( + stored_v(&s).await, + "w5", + "the put that gave up wrote nothing" + ); + s.cleanup().await; +} + +#[tokio::test] +async fn a_put_that_keeps_losing_the_create_race_gives_up_after_five_inserts() { + create_race_churn_gives_up( + "a_put_that_keeps_losing_the_create_race_gives_up_after_five_inserts", + false, + ) + .await; +} + +#[tokio::test] +async fn a_put_that_keeps_losing_the_create_race_gives_up_after_five_inserts_on_a_range_table() { + create_race_churn_gives_up( + "a_put_that_keeps_losing_the_create_race_gives_up_after_five_inserts_on_a_range_table", + true, + ) + .await; +} + +const LSI_TABLE: &str = "t_put_race_lsi_stream"; + +/// Create a (pk, sk) table with LSI `lsi1` on `lsk` and a stream. +async fn lsi_stream_table(s: &Scratch) -> TableKeyInfo { + // The scratch database holds the catalog and data schemas together, so it + // has the catalog's copy of `stream_shards`. Drop its foreign key to + // `tables`, which the data database does not have. + sqlx::query("ALTER TABLE stream_shards DROP CONSTRAINT stream_shards_table_id_fkey") + .execute(&s.db) + .await + .expect("match the data schema"); + let (pk, pk_def) = s_key("pk", KeyType::Hash); + let (sk, sk_def) = s_key("sk", KeyType::Range); + let (lsk, lsk_def) = s_key("lsk", KeyType::Range); + s.engine + .create_table( + ACCOUNT, + CreateTableInput { + table_name: LSI_TABLE.to_owned(), + key_schema: vec![pk.clone(), sk], + attribute_definitions: vec![pk_def, sk_def, lsk_def], + billing_mode: Some(BillingMode::PayPerRequest), + local_secondary_indexes: Some(vec![LsiInput { + index_name: "lsi1".to_owned(), + key_schema: vec![pk, lsk], + projection: Projection { + projection_type: ProjectionType::All, + non_key_attributes: None, + }, + }]), + stream_specification: Some(StreamSpecification { + stream_enabled: true, + stream_view_type: Some(StreamViewType::NewAndOldImages), + }), + ..Default::default() + }, + ) + .await + .expect("create the table"); + s.engine + .table_key_info(ACCOUNT, LSI_TABLE) + .await + .expect("read the key info") +} + +fn range_item(pk: &str, extra: &[(&str, &str)]) -> Item { + let mut item = BTreeMap::from([ + ("pk".to_owned(), AttributeValue::S(pk.to_owned())), + ("sk".to_owned(), AttributeValue::S("1".to_owned())), + ]); + for (name, value) in extra { + item.insert((*name).to_owned(), AttributeValue::S((*value).to_owned())); + } + item +} + +#[tokio::test] +async fn a_lost_create_race_gives_the_lsi_and_the_stream_the_winner() { + // A write transaction creates `c`, with its LSI row and stream record, and + // then waits on `d`, which an outside transaction holds. A PutItem of `c` + // meanwhile loses the insert and waits for that create. Once it commits, + // the put must replace the winner everywhere: one LSI row for the put's + // `lsk`, none for the winner's, and a MODIFY record whose old image is the + // winner. + let test = "a_lost_create_race_gives_the_lsi_and_the_stream_the_winner"; + if base_conn().is_none() { + return skip(test); + } + let s = scratch().await; + let maps = ExpressionMaps::default(); + let key_info = lsi_stream_table(&s).await; + s.engine + .put_item(&key_info, range_item("d", &[]), false, None, &maps, None) + .await + .expect("seed d"); + let table_id: String = + sqlx::query_scalar("SELECT table_id FROM tables WHERE account_id = $1 AND table_name = $2") + .bind(ACCOUNT) + .bind(LSI_TABLE) + .fetch_one(&s.db) + .await + .expect("look up the table id"); + let index_id: String = sqlx::query_scalar( + "SELECT index_id FROM indexes WHERE table_id = $1 AND index_name = 'lsi1'", + ) + .bind(&table_id) + .fetch_one(&s.db) + .await + .expect("look up the index id"); + let (table, lsi) = ( + format!("\"_ddb_{table_id}\""), + format!("\"_ddb_{index_id}\""), + ); + + let mut holder = s.db.begin().await.expect("begin the outside transaction"); + let holder_pid: i32 = sqlx::query_scalar("SELECT pg_backend_pid()") + .fetch_one(&mut *holder) + .await + .expect("read the backend pid"); + sqlx::query(&format!("SELECT 1 FROM {table} WHERE pk = 'd' FOR UPDATE")) + .execute(&mut *holder) + .await + .expect("lock d"); + + let capture = StreamCapture { + view_type: StreamViewType::NewAndOldImages, + user_identity: None, + region: REGION.into(), + }; + let winner = range_item("c", &[("lsk", "l-winner"), ("v", "winner")]); + let loser = range_item("c", &[("lsk", "l-loser"), ("v", "loser")]); + let new_d = range_item("d", &[("v", "creator")]); + let stream_put = |item| TransactWriteOp::Put { + key_info: &key_info, + item, + condition: None, + maps: &maps, + return_values_on_ccf: ReturnValuesOnConditionCheckFailure::None, + stream: Some(capture.clone()), + }; + let creator_ops = [stream_put(&winner), stream_put(&new_d)]; + + let creator = s.engine.transact_write_items(&creator_ops, None); + let put = async { + // The creator has inserted c and waits on d. + wait_until_blocked_by(&s.db, holder_pid, "transactionid").await?; + Ok::<_, String>( + s.engine + .put_item(&key_info, loser.clone(), true, None, &maps, Some(&capture)) + .await, + ) + }; + let release = async { + // The put waits on the creator's insert of c before d is released. + let blocked = async { + for _ in 0..500 { + let n: i64 = sqlx::query_scalar( + "SELECT count(*) FROM pg_stat_activity WHERE datname = current_database() \ + AND wait_event_type = 'Lock' AND wait_event = 'transactionid' \ + AND NOT ($1 = ANY(pg_blocking_pids(pid)))", + ) + .bind(holder_pid) + .fetch_one(&s.db) + .await + .expect("read pg_stat_activity"); + if n > 0 { + return Ok(()); + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + Err("the put never waited on the creator's insert".to_owned()) + } + .await; + holder.rollback().await.expect("release d"); + blocked + }; + let (created, put_result, released) = tokio::time::timeout(Duration::from_secs(30), async { + tokio::join!(creator, put, release) + }) + .await + .expect("both writes finished"); + created.expect("the create commits"); + let put_result = put_result.unwrap_or_else(|e| panic!("{e}")); + if let Err(e) = released { + panic!("{e}; the put returned {put_result:?}"); + } + assert_eq!( + put_result.expect("the put commits after the winner"), + Some(winner.clone()), + "ALL_OLD returns the winner" + ); + + let lsks: Vec = sqlx::query_scalar(&format!( + "SELECT item_data->'lsk'->>'S' FROM {lsi} WHERE pk = 'c'" + )) + .fetch_all(&s.db) + .await + .expect("read the LSI rows"); + assert_eq!(lsks, ["l-loser"], "only the put's LSI row is left"); + let v: String = sqlx::query_scalar(&format!( + "SELECT item_data->'v'->>'S' FROM {table} WHERE pk = 'c'" + )) + .fetch_one(&s.db) + .await + .expect("read c"); + assert_eq!(v, "loser"); + + let records: Vec<(serde_json::Value,)> = sqlx::query_as( + "SELECT record_data FROM stream_records WHERE table_id = $1 ORDER BY sequence_number", + ) + .bind(&table_id) + .fetch_all(&s.db) + .await + .expect("read the stream records"); + let c_events: Vec<(String, Option)> = records + .into_iter() + .map(|(data,)| serde_json::from_value::(data).expect("a stream record")) + .filter(|r| r.dynamodb.keys.get("pk") == Some(&AttributeValue::S("c".to_owned()))) + .map(|r| (format!("{:?}", r.event_name), r.dynamodb.old_image)) + .collect(); + assert_eq!( + c_events, + [ + ("Insert".to_owned(), None), + ("Modify".to_owned(), Some(winner.clone())) + ], + "the create, then the put that replaced the winner" + ); + + s.cleanup().await; +} diff --git a/docs/design/02-high-level-design.md b/docs/design/02-high-level-design.md index a08cf641..3a680e61 100755 --- a/docs/design/02-high-level-design.md +++ b/docs/design/02-high-level-design.md @@ -390,6 +390,7 @@ Read-modify-write operations (UpdateItem, PutItem with conditions, DeleteItem wi - **Atomicity:** The condition check and the write happen against the same snapshot. - **Serialization:** Concurrent updates to the same item are serialized by PostgreSQL's row lock, not by any in-memory mutex. - **No TOCTOU races:** Another request cannot modify the item between the condition check and the write. +- **Missing items:** A missing item has no row to lock, so `INSERT ... ON CONFLICT DO NOTHING` decides which writer creates it. A PutItem or UpdateItem that loses re-reads the winner `FOR UPDATE`, checks its condition against it, and writes after it. If the winner was deleted in the meantime, it retries the insert, up to 5 inserts in total. There is no in-memory locking (no `Mutex`, `RwLock`, or similar) on the data path. All contention is managed by PostgreSQL. diff --git a/tests/test_put_item_create_race.py b/tests/test_put_item_create_race.py new file mode 100644 index 00000000..752628d2 --- /dev/null +++ b/tests/test_put_item_create_race.py @@ -0,0 +1,387 @@ +# Copyright 2026 ExtendDB contributors +# SPDX-License-Identifier: Apache-2.0 + +"""Concurrent PutItem calls that create the same new item. + +Amazon DynamoDB applies the puts one after another. Each one succeeds unless +its own condition fails against the item before it: a put that finds the item +already created overwrites it, and with ReturnValues ALL_OLD it returns the +item it replaced. So the old images of one round form a single chain, from no +item to the final one. The tests cover the request shapes that make a put +read the current item first: a condition, ALL_OLD, ReturnConsumedCapacity, a +GSI, an LSI, and a stream, and BatchWriteItem puts to a stream table. On index +and stream tables they also check that only the final item stays in the index, +and that each stream record's old image is the record before it. +""" + +from __future__ import annotations + +import os +import threading +import time +import uuid + +import boto3 +import pytest +from botocore.config import Config +from botocore.exceptions import ClientError + +from conftest import scoped_table + +WRITERS = 8 +ROUNDS = 15 + + +# TEMPORARY: on MongoDB, a put without a condition on a table with no index and +# no stream returns HTTP 500 when it loses the create race. The MongoDB +# write-race fix repairs it. That fix lands separately. +# TODO: remove this marker once the MongoDB write-race fix is on main. +# Only an assertion counts as the expected failure. The marker is not strict, +# because the race does not fire in every run: these tests still pass now and +# then on MongoDB without the fix. The MongoDB test runner sets +# EXTENDDB_TEST_MONGODB_CONTAINER. +XFAIL_UNTIL_MONGODB_FIX = pytest.mark.xfail( + bool(os.environ.get("EXTENDDB_TEST_MONGODB_CONTAINER", "").strip()), + reason="MongoDB returns HTTP 500 for a lost create race, fixed by the MongoDB write-race fix", + raises=AssertionError, + strict=False, +) + + +@pytest.fixture(scope="module") +def raw_client(endpoint_url): + """A client that never retries, so every failure is seen.""" + kwargs: dict = { + "service_name": "dynamodb", + "region_name": os.environ.get("AWS_DEFAULT_REGION", "us-east-1"), + "config": Config( + retries={"total_max_attempts": 1, "mode": "standard"}, + max_pool_connections=WRITERS * 2, + ), + } + if endpoint_url: + kwargs["endpoint_url"] = endpoint_url + if endpoint_url.startswith("https://"): + kwargs["verify"] = False + return boto3.client(**kwargs) + + +@pytest.fixture(scope="module") +def streams_client(endpoint_url): + kwargs: dict = { + "service_name": "dynamodbstreams", + "region_name": os.environ.get("AWS_DEFAULT_REGION", "us-east-1"), + } + if endpoint_url: + kwargs["endpoint_url"] = endpoint_url + if endpoint_url.startswith("https://"): + kwargs["verify"] = False + return boto3.client(**kwargs) + + +def _s(name: str) -> dict: + return {"AttributeName": name, "AttributeType": "S"} + + +def _k(name: str, kind: str) -> dict: + return {"AttributeName": name, "KeyType": kind} + + +@pytest.fixture(scope="module") +def hash_table(dynamodb_client): + with scoped_table(dynamodb_client) as name: + yield name + + +@pytest.fixture(scope="module") +def range_table(dynamodb_client): + with scoped_table( + dynamodb_client, + attribute_definitions=[_s("pk"), _s("sk")], + key_schema=[_k("pk", "HASH"), _k("sk", "RANGE")], + ) as name: + yield name + + +@pytest.fixture(scope="module") +def gsi_table(dynamodb_client): + with scoped_table( + dynamodb_client, + attribute_definitions=[_s("pk"), _s("gk")], + GlobalSecondaryIndexes=[ + { + "IndexName": "gk-index", + "KeySchema": [_k("gk", "HASH")], + "Projection": {"ProjectionType": "ALL"}, + } + ], + ) as name: + yield name + + +@pytest.fixture(scope="module") +def lsi_table(dynamodb_client): + with scoped_table( + dynamodb_client, + attribute_definitions=[_s("pk"), _s("sk"), _s("lk")], + key_schema=[_k("pk", "HASH"), _k("sk", "RANGE")], + LocalSecondaryIndexes=[ + { + "IndexName": "lk-index", + "KeySchema": [_k("pk", "HASH"), _k("lk", "RANGE")], + "Projection": {"ProjectionType": "ALL"}, + } + ], + ) as name: + yield name + + +@pytest.fixture(scope="module") +def stream_table(dynamodb_client): + with scoped_table( + dynamodb_client, + StreamSpecification={"StreamEnabled": True, "StreamViewType": "NEW_AND_OLD_IMAGES"}, + ) as name: + yield name + + +def _key(table_kind: str, pk: str) -> dict: + key = {"pk": {"S": pk}} + if table_kind == "range": + key["sk"] = {"S": "s"} + return key + + +def _race( + client, table: str, table_kind: str, extra: dict, batch: bool = False +) -> tuple[str, list]: + """WRITERS clients put the same new key at once, each with its own `w`. + + Each writer's item also carries index keys unique to it: `gk` across the + whole table, and `lk` within the key. With `batch`, each writer sends a + BatchWriteItem with one PutRequest instead of a PutItem. + + Returns the key and, per writer, the old `w` it replaced (None when it + created the item or used BatchWriteItem), or the error when the put failed. + """ + pk = f"k-{uuid.uuid4().hex}" + start = threading.Barrier(WRITERS, timeout=30) + out: list = [None] * WRITERS + errors: list[BaseException] = [] + + def run(i: int): + item = { + **_key(table_kind, pk), + "w": {"S": f"w{i}"}, + "gk": {"S": f"g-{pk}-{i}"}, + "lk": {"S": f"l{i}"}, + } + try: + start.wait() + if batch: + r = client.batch_write_item(RequestItems={table: [{"PutRequest": {"Item": item}}]}) + left = r.get("UnprocessedItems") + out[i] = ("unprocessed", left) if left else ("ok", None) + return + r = client.put_item(TableName=table, Item=item, **extra) + out[i] = ("ok", r.get("Attributes", {}).get("w", {}).get("S")) + except ClientError as e: + out[i] = ("error", e.response["Error"]["Code"]) + except BaseException as e: # noqa: BLE001 - surfaced below + errors.append(e) + + threads = [threading.Thread(target=run, args=(i,)) for i in range(WRITERS)] + for t in threads: + t.start() + for t in threads: + t.join() + if errors: + raise errors[0] + return pk, out + + +def _final(client, table: str, table_kind: str, pk: str) -> dict: + return client.get_item(TableName=table, Key=_key(table_kind, pk), ConsistentRead=True)[ + "Item" + ] + + +def _final_w(client, table: str, table_kind: str, pk: str) -> str: + return _final(client, table, table_kind, pk)["w"]["S"] + + +def _assert_all_succeed(out: list): + failed = [o for o in out if o[0] != "ok"] + assert not failed, f"{len(failed)} of {WRITERS} puts failed: {failed}" + + +def _assert_one_chain(out: list, final: str): + """The old images lead from the final item back to no item, once each.""" + prev = {f"w{i}": old for i, (_, old) in enumerate(out)} + seen, cur = 0, final + while cur is not None and seen <= WRITERS: + seen += 1 + cur = prev[cur] + assert seen == WRITERS, f"old images do not form one chain: {prev}, final {final}" + + +@XFAIL_UNTIL_MONGODB_FIX +@pytest.mark.parametrize("table_kind", ["hash", "range"]) +def test_racing_puts_with_all_old_all_succeed( + request, dynamodb_client, raw_client, table_kind +): + table = request.getfixturevalue(f"{table_kind}_table") + for _ in range(ROUNDS): + pk, out = _race(raw_client, table, table_kind, {"ReturnValues": "ALL_OLD"}) + _assert_all_succeed(out) + _assert_one_chain(out, _final_w(dynamodb_client, table, table_kind, pk)) + + +@XFAIL_UNTIL_MONGODB_FIX +def test_racing_puts_that_return_consumed_capacity_all_succeed( + dynamodb_client, raw_client, hash_table +): + extra = {"ReturnConsumedCapacity": "TOTAL"} + for _ in range(ROUNDS): + pk, out = _race(raw_client, hash_table, "hash", extra) + _assert_all_succeed(out) + assert _final_w(dynamodb_client, hash_table, "hash", pk) in { + f"w{i}" for i in range(WRITERS) + } + + +def test_racing_puts_check_their_condition_against_the_winner( + dynamodb_client, raw_client, range_table +): + """A condition that every item satisfies never fails, whoever created the item.""" + extra = {"ConditionExpression": "attribute_not_exists(zz)", "ReturnValues": "ALL_OLD"} + for _ in range(ROUNDS): + pk, out = _race(raw_client, range_table, "range", extra) + _assert_all_succeed(out) + _assert_one_chain(out, _final_w(dynamodb_client, range_table, "range", pk)) + + +def test_racing_conditional_creates_have_one_winner(raw_client, hash_table): + """Control: with attribute_not_exists(pk), exactly one put creates the item.""" + for _ in range(ROUNDS): + _, out = _race( + raw_client, hash_table, "hash", {"ConditionExpression": "attribute_not_exists(pk)"} + ) + codes = sorted(o[1] if o[0] == "error" else "ok" for o in out) + assert codes == ["ConditionalCheckFailedException"] * (WRITERS - 1) + ["ok"], codes + + +def _gsi_keys(client, table: str) -> dict[str, list[str]]: + """Every `gk` in the GSI, grouped by the base key it points at.""" + by_pk: dict[str, list[str]] = {} + kwargs = {"TableName": table, "IndexName": "gk-index"} + while True: + resp = client.scan(**kwargs) + for item in resp["Items"]: + by_pk.setdefault(item["pk"]["S"], []).append(item["gk"]["S"]) + if "LastEvaluatedKey" not in resp: + return {pk: sorted(gks) for pk, gks in by_pk.items()} + kwargs["ExclusiveStartKey"] = resp["LastEvaluatedKey"] + + +def test_racing_puts_on_a_gsi_table_leave_only_the_final_entry( + dynamodb_client, raw_client, gsi_table +): + """The GSI holds the final item's entry and none of the items it replaced.""" + want: dict[str, list[str]] = {} + for _ in range(ROUNDS): + pk, out = _race(raw_client, gsi_table, "hash", {}) + _assert_all_succeed(out) + want[pk] = [_final(dynamodb_client, gsi_table, "hash", pk)["gk"]["S"]] + # The GSI is eventually consistent on Amazon DynamoDB, so poll. + deadline = time.monotonic() + 30 + while (got := _gsi_keys(dynamodb_client, gsi_table)) != want and time.monotonic() < deadline: + time.sleep(0.5) + wrong = {pk: (want[pk], got.get(pk)) for pk in want if got.get(pk) != want[pk]} + assert not wrong and got.keys() == want.keys(), f"stale GSI entries (want, got): {wrong}" + + +def test_racing_puts_on_an_lsi_table_leave_only_the_final_entry( + dynamodb_client, raw_client, lsi_table +): + """A strongly consistent LSI query returns the final item's entry only.""" + for _ in range(ROUNDS): + pk, out = _race(raw_client, lsi_table, "range", {"ReturnValues": "ALL_OLD"}) + _assert_all_succeed(out) + final = _final(dynamodb_client, lsi_table, "range", pk) + _assert_one_chain(out, final["w"]["S"]) + items = dynamodb_client.query( + TableName=lsi_table, + IndexName="lk-index", + ConsistentRead=True, + KeyConditionExpression="pk = :pk", + ExpressionAttributeValues={":pk": {"S": pk}}, + )["Items"] + assert [i["lk"]["S"] for i in items] == [final["lk"]["S"]], (final, items) + + +def _stream_records(streams_client, stream_arn: str) -> dict[str, list[dict]]: + """Every record in the stream, grouped by key, in sequence order.""" + by_pk: dict[str, list[dict]] = {} + shards = streams_client.describe_stream(StreamArn=stream_arn)["StreamDescription"]["Shards"] + for shard in shards: + it = streams_client.get_shard_iterator( + StreamArn=stream_arn, ShardId=shard["ShardId"], ShardIteratorType="TRIM_HORIZON" + )["ShardIterator"] + empty = 0 + for _ in range(50): + resp = streams_client.get_records(ShardIterator=it, Limit=1000) + for r in resp.get("Records", []): + by_pk.setdefault(r["dynamodb"]["Keys"]["pk"]["S"], []).append(r) + it = resp.get("NextShardIterator") + empty = 0 if resp.get("Records") else empty + 1 + if not it or empty >= 3: + break + for records in by_pk.values(): + records.sort(key=lambda r: int(r["dynamodb"]["SequenceNumber"])) + return by_pk + + +def _assert_stream_chains(dynamodb_client, streams_client, table: str, pks: list[str]): + """Per key: one INSERT, then a MODIFY per later put, each with the item + before it as its old image.""" + arn = dynamodb_client.describe_table(TableName=table)["Table"]["LatestStreamArn"] + deadline = time.monotonic() + 30 + while True: + by_pk = _stream_records(streams_client, arn) + if all(len(by_pk.get(pk, [])) >= WRITERS for pk in pks) or time.monotonic() > deadline: + break + time.sleep(1) + for pk in pks: + records = by_pk.get(pk, []) + names = [r["eventName"] for r in records] + assert names == ["INSERT"] + ["MODIFY"] * (WRITERS - 1), (pk, names) + images = [ + (r["dynamodb"].get("OldImage", {}).get("w"), r["dynamodb"]["NewImage"]["w"]) + for r in records + ] + for (_, before), (old, _) in zip(images, images[1:]): + assert old == before, f"old image is not the record before it: {pk} {images}" + + +def test_racing_puts_on_a_stream_table_all_succeed( + dynamodb_client, raw_client, streams_client, stream_table +): + pks = [] + for _ in range(ROUNDS): + pk, out = _race(raw_client, stream_table, "hash", {}) + _assert_all_succeed(out) + pks.append(pk) + _assert_stream_chains(dynamodb_client, streams_client, stream_table, pks) + + +def test_racing_batch_puts_on_a_stream_table_all_succeed( + dynamodb_client, raw_client, streams_client, stream_table +): + """BatchWriteItem puts read the current item for the stream record too.""" + pks = [] + for _ in range(ROUNDS): + pk, out = _race(raw_client, stream_table, "hash", {}, batch=True) + _assert_all_succeed(out) + pks.append(pk) + _assert_stream_chains(dynamodb_client, streams_client, stream_table, pks)