diff --git a/impls/src/const_utils.rs b/impls/src/const_utils.rs new file mode 100644 index 0000000..81a29cf --- /dev/null +++ b/impls/src/const_utils.rs @@ -0,0 +1,46 @@ +/// Concatenates constant string expressions into a single constant string. +macro_rules! const_concat_str { + ($($string:expr),* $(,)?) => {{ + $(const _: &str = $string;)* + const LEN: usize = 0 $(+ $string.len())*; + const BYTES: [u8; LEN] = + $crate::const_utils::const_concat_str_inner::(&[$($string.as_bytes()),*]); + match std::str::from_utf8(&BYTES) { + Ok(string) => string, + Err(_) => panic!("concatenated string was not valid UTF-8"), + } + }}; +} + +pub(crate) use const_concat_str; + +/// Copies constant byte slices into a single constant byte array. +pub(crate) const fn const_concat_str_inner(slices: &[&[u8]]) -> [u8; LEN] { + let mut bytes = [0; LEN]; + let mut base = 0; + let mut i = 0; + while i < slices.len() { + let slice = slices[i]; + let mut j = 0; + while j < slice.len() { + bytes[base + j] = slice[j]; + j += 1; + } + base += slice.len(); + i += 1; + } + assert!(base == LEN, "invalid concatenated string length"); + bytes +} + +#[cfg(test)] +mod tests { + use super::const_concat_str; + + #[test] + fn concatenates_const_string_expressions() { + const NAME: &str = const_concat_str!("HUGE", " ", "MAN"); + const GREETING: &str = const_concat_str!("Hello ", NAME, "!!"); + assert_eq!(GREETING, "Hello HUGE MAN!!"); + } +} diff --git a/impls/src/lib.rs b/impls/src/lib.rs index d58b84e..d83e1cd 100644 --- a/impls/src/lib.rs +++ b/impls/src/lib.rs @@ -11,6 +11,7 @@ #![deny(rustdoc::private_intra_doc_links)] #![deny(missing_docs)] +mod const_utils; mod migrations; /// Contains [PostgreSQL](https://www.postgresql.org/) based backend implementation for VSS. pub mod postgres_store; diff --git a/impls/src/postgres_store.rs b/impls/src/postgres_store.rs index 765099e..162771b 100644 --- a/impls/src/postgres_store.rs +++ b/impls/src/postgres_store.rs @@ -1,7 +1,8 @@ +use crate::const_utils::const_concat_str; use crate::migrations::*; use api::error::VssError; -use api::kv_store::{KvStore, GLOBAL_VERSION_KEY, INITIAL_RECORD_VERSION}; +use api::kv_store::{KvStore, GLOBAL_VERSION_KEY}; use api::types::{ DeleteObjectRequest, DeleteObjectResponse, GetObjectRequest, GetObjectResponse, KeyValue, ListKeyVersionsRequest, ListKeyVersionsResponse, PutObjectRequest, PutObjectResponse, @@ -12,10 +13,12 @@ use chrono::Utc; use native_tls::TlsConnector; use postgres_native_tls::MakeTlsConnector; use std::cmp::min; +use std::collections::HashMap; use std::io::{self, Error, ErrorKind}; use tokio::sync::Mutex; use tokio_postgres::tls::{MakeTlsConnect, TlsConnect}; -use tokio_postgres::{error, Client, NoTls, Socket, Transaction}; +use tokio_postgres::types::ToSql; +use tokio_postgres::{error, NoTls, Row, Socket, Statement}; use log::{debug, info, warn}; @@ -34,6 +37,7 @@ const KEY_COLUMN: &str = "key"; const VALUE_COLUMN: &str = "value"; const VERSION_COLUMN: &str = "version"; const SORT_ORDER_COLUMN: &str = "sort_order"; +const INITIAL_RECORD_VERSION_STR: &str = "1"; /// Page token is the `sort_order` value of the last item in the previous page, /// encoded as a decimal string. @@ -66,6 +70,93 @@ pub const MAX_PUT_REQUEST_ITEM_COUNT: usize = 1000; const POOL_SIZE: usize = 10; +/// A simple query string -> prepared [`Statement`] cache. +/// +/// Prepared statements can only be used with the connection that created them. +struct StatementCache { + statements: HashMap<&'static str, Statement>, +} + +impl StatementCache { + fn new() -> Self { + Self { statements: HashMap::new() } + } + + async fn get_or_prepare( + &mut self, client: &C, query: &'static str, + ) -> Result { + if let Some(statement) = self.statements.get(query) { + return Ok(statement.clone()); + } + + let statement = client.prepare(query).await?; + self.statements.insert(query, statement.clone()); + Ok(statement) + } +} + +/// A [`tokio_postgres::Client`] connection that transparently prepares and caches any query statements. +struct Client { + uncached_client: tokio_postgres::Client, + statement_cache: StatementCache, +} + +impl Client { + async fn connect(postgres_endpoint: &str, db_name: &str, tls: T) -> Result + where + T: MakeTlsConnect + Clone + Send + Sync + 'static, + T::Stream: Send + Sync, + T::TlsConnect: Send, + <>::TlsConnect as TlsConnect>::Future: Send, + { + let client = make_db_connection(postgres_endpoint, db_name, tls).await?; + let statement_cache = StatementCache::new(); + Ok(Self { uncached_client: client, statement_cache }) + } + + async fn query( + &mut self, query: &'static str, params: &[&(dyn ToSql + Sync)], + ) -> Result, tokio_postgres::Error> { + let statement = self.statement_cache.get_or_prepare(&self.uncached_client, query).await?; + self.uncached_client.query(&statement, params).await + } + + async fn query_opt( + &mut self, query: &'static str, params: &[&(dyn ToSql + Sync)], + ) -> Result, tokio_postgres::Error> { + let statement = self.statement_cache.get_or_prepare(&self.uncached_client, query).await?; + self.uncached_client.query_opt(&statement, params).await + } + + async fn transaction(&mut self) -> Result, tokio_postgres::Error> { + let transaction = self.uncached_client.transaction().await?; + Ok(Transaction { transaction, statement_cache: &mut self.statement_cache }) + } +} + +/// A [`tokio_postgres::Transaction`] that transparently prepares and caches any query statements. +struct Transaction<'a> { + transaction: tokio_postgres::Transaction<'a>, + statement_cache: &'a mut StatementCache, +} + +impl Transaction<'_> { + async fn execute( + &mut self, query: &'static str, params: &[&(dyn ToSql + Sync)], + ) -> Result { + let statement = self.statement_cache.get_or_prepare(&self.transaction, query).await?; + self.transaction.execute(&statement, params).await + } + + async fn commit(self) -> Result<(), tokio_postgres::Error> { + self.transaction.commit().await + } + + async fn rollback(self) -> Result<(), tokio_postgres::Error> { + self.transaction.rollback().await + } +} + struct SmallPool { connections: [Mutex; POOL_SIZE], endpoint: String, @@ -82,16 +173,16 @@ where { async fn new(postgres_endpoint: &str, vss_db: &str, tls: T) -> Result { let connections = [ - Mutex::new(make_db_connection(postgres_endpoint, vss_db, tls.clone()).await?), - Mutex::new(make_db_connection(postgres_endpoint, vss_db, tls.clone()).await?), - Mutex::new(make_db_connection(postgres_endpoint, vss_db, tls.clone()).await?), - Mutex::new(make_db_connection(postgres_endpoint, vss_db, tls.clone()).await?), - Mutex::new(make_db_connection(postgres_endpoint, vss_db, tls.clone()).await?), - Mutex::new(make_db_connection(postgres_endpoint, vss_db, tls.clone()).await?), - Mutex::new(make_db_connection(postgres_endpoint, vss_db, tls.clone()).await?), - Mutex::new(make_db_connection(postgres_endpoint, vss_db, tls.clone()).await?), - Mutex::new(make_db_connection(postgres_endpoint, vss_db, tls.clone()).await?), - Mutex::new(make_db_connection(postgres_endpoint, vss_db, tls.clone()).await?), + Mutex::new(Client::connect(postgres_endpoint, vss_db, tls.clone()).await?), + Mutex::new(Client::connect(postgres_endpoint, vss_db, tls.clone()).await?), + Mutex::new(Client::connect(postgres_endpoint, vss_db, tls.clone()).await?), + Mutex::new(Client::connect(postgres_endpoint, vss_db, tls.clone()).await?), + Mutex::new(Client::connect(postgres_endpoint, vss_db, tls.clone()).await?), + Mutex::new(Client::connect(postgres_endpoint, vss_db, tls.clone()).await?), + Mutex::new(Client::connect(postgres_endpoint, vss_db, tls.clone()).await?), + Mutex::new(Client::connect(postgres_endpoint, vss_db, tls.clone()).await?), + Mutex::new(Client::connect(postgres_endpoint, vss_db, tls.clone()).await?), + Mutex::new(Client::connect(postgres_endpoint, vss_db, tls.clone()).await?), ]; let pool = SmallPool { @@ -121,11 +212,11 @@ where } async fn ensure_connected(&self, client: &mut Client) -> Result<(), Error> { - if client.is_closed() || client.check_connection().await.is_err() { + if client.uncached_client.is_closed() + || client.uncached_client.check_connection().await.is_err() + { debug!("Rotating connection to the postgres database"); - let new_client = - make_db_connection(&self.endpoint, &self.db_name, self.tls.clone()).await?; - *client = new_client; + *client = Client::connect(&self.endpoint, &self.db_name, self.tls.clone()).await?; } Ok(()) } @@ -150,7 +241,7 @@ pub type PostgresTlsBackend = PostgresBackend; async fn make_db_connection( postgres_endpoint: &str, db_name: &str, tls: T, -) -> Result +) -> Result where T: MakeTlsConnect + Clone + Send + Sync + 'static, T::Stream: Send + Sync, @@ -280,7 +371,7 @@ where async fn migrate_vss_database(&self, migrations: &[&str]) -> Result<(usize, usize), Error> { let mut conn = self.pool.get().await?; // Get the next migration to be applied. - let migration_start = match conn.query_one(GET_VERSION_STMT, &[]).await { + let migration_start = match conn.uncached_client.query_one(GET_VERSION_STMT, &[]).await { Ok(row) => { let i: i32 = row.get(DB_VERSION_COLUMN); usize::try_from(i).expect("The column should always contain unsigned integers") @@ -298,10 +389,10 @@ where }, }; - let tx = conn - .transaction() - .await - .map_err(|e| Error::new(ErrorKind::Other, format!("Transaction start error: {}", e)))?; + let tx = + conn.uncached_client.transaction().await.map_err(|e| { + Error::new(ErrorKind::Other, format!("Transaction start error: {}", e)) + })?; if migration_start == migrations.len() { // No migrations needed, we are done @@ -361,14 +452,14 @@ where #[cfg(test)] async fn get_schema_version(&self) -> usize { let conn = self.pool.get().await.unwrap(); - let row = conn.query_one(GET_VERSION_STMT, &[]).await.unwrap(); + let row = conn.uncached_client.query_one(GET_VERSION_STMT, &[]).await.unwrap(); usize::try_from(row.get::<&str, i32>(DB_VERSION_COLUMN)).unwrap() } #[cfg(test)] async fn get_upgrades_list(&self) -> Vec { let conn = self.pool.get().await.unwrap(); - let rows = conn.query(GET_MIGRATION_LOG_STMT, &[]).await.unwrap(); + let rows = conn.uncached_client.query(GET_MIGRATION_LOG_STMT, &[]).await.unwrap(); rows.iter() .map(|row| usize::try_from(row.get::<&str, i32>(MIGRATION_LOG_COLUMN)).unwrap()) .collect() @@ -388,15 +479,19 @@ where } async fn execute_non_conditional_upsert( - &self, transaction: &Transaction<'_>, vss_record: &VssDbRecord, + &self, transaction: &mut Transaction<'_>, vss_record: &VssDbRecord, ) -> io::Result { - let stmt = format!("INSERT INTO vss_db (user_token, store_id, key, value, version, created_at, last_updated_at) - VALUES ($1, $2, $3, $4, {}, $5, $6) - ON CONFLICT (user_token, store_id, key) DO UPDATE - SET value = EXCLUDED.value, version = {}, last_updated_at = EXCLUDED.last_updated_at", INITIAL_RECORD_VERSION, INITIAL_RECORD_VERSION); + #[rustfmt::skip] + const STMT: &str = const_concat_str!( + "INSERT INTO vss_db (user_token, store_id, key, value, version, created_at, last_updated_at) ", + "VALUES ($1, $2, $3, $4, ", INITIAL_RECORD_VERSION_STR, ", $5, $6) ", + "ON CONFLICT (user_token, store_id, key) DO UPDATE ", + "SET value = EXCLUDED.value, version = ", INITIAL_RECORD_VERSION_STR, + ", last_updated_at = EXCLUDED.last_updated_at", + ); let num_rows = transaction .execute( - &stmt, + STMT, &[ &vss_record.user_token, &vss_record.store_id, @@ -414,14 +509,17 @@ where } async fn execute_conditional_insert( - &self, transaction: &Transaction<'_>, vss_record: &VssDbRecord, + &self, transaction: &mut Transaction<'_>, vss_record: &VssDbRecord, ) -> io::Result { - let stmt = format!("INSERT INTO vss_db (user_token, store_id, key, value, version, created_at, last_updated_at) - VALUES ($1, $2, $3, $4, {}, $5, $6) - ON CONFLICT DO NOTHING", INITIAL_RECORD_VERSION); + #[rustfmt::skip] + const STMT: &str = const_concat_str!( + "INSERT INTO vss_db (user_token, store_id, key, value, version, created_at, last_updated_at) ", + "VALUES ($1, $2, $3, $4, ", INITIAL_RECORD_VERSION_STR, ", $5, $6) ", + "ON CONFLICT DO NOTHING", + ); let num_rows = transaction .execute( - &stmt, + STMT, &[ &vss_record.user_token, &vss_record.store_id, @@ -439,7 +537,7 @@ where } async fn execute_conditional_update( - &self, transaction: &Transaction<'_>, vss_record: &VssDbRecord, + &self, transaction: &mut Transaction<'_>, vss_record: &VssDbRecord, ) -> io::Result { let stmt = "UPDATE vss_db SET value = $1, version = $2, last_updated_at = $3 WHERE user_token = $4 AND store_id = $5 AND key = $6 AND version = $7"; @@ -464,7 +562,7 @@ where } async fn execute_put_object_query( - &self, transaction: &Transaction<'_>, vss_record: &VssDbRecord, + &self, transaction: &mut Transaction<'_>, vss_record: &VssDbRecord, ) -> io::Result { if vss_record.version == -1 { self.execute_non_conditional_upsert(transaction, vss_record).await @@ -476,7 +574,7 @@ where } async fn execute_non_conditional_delete( - &self, transaction: &Transaction<'_>, vss_record: &VssDbRecord, + &self, transaction: &mut Transaction<'_>, vss_record: &VssDbRecord, ) -> io::Result { let stmt = "DELETE FROM vss_db WHERE user_token = $1 AND store_id = $2 AND key = $3"; let num_rows = transaction @@ -489,7 +587,7 @@ where } async fn execute_conditional_delete( - &self, transaction: &Transaction<'_>, vss_record: &VssDbRecord, + &self, transaction: &mut Transaction<'_>, vss_record: &VssDbRecord, ) -> io::Result { let stmt = "DELETE FROM vss_db WHERE user_token = $1 AND store_id = $2 AND key = $3 AND version = $4"; let num_rows = transaction @@ -510,7 +608,7 @@ where } async fn execute_delete_object_query( - &self, transaction: &Transaction<'_>, vss_record: &VssDbRecord, + &self, transaction: &mut Transaction<'_>, vss_record: &VssDbRecord, ) -> io::Result { if vss_record.version == -1 { self.execute_non_conditional_delete(transaction, vss_record).await @@ -531,7 +629,7 @@ where async fn get( &self, user_token: String, request: GetObjectRequest, ) -> Result { - let conn = self.pool.get().await?; + let mut conn = self.pool.get().await?; let stmt = "SELECT key, value, version FROM vss_db WHERE user_token = $1 AND store_id = $2 AND key = $3"; let row = conn .query_opt(stmt, &[&user_token, &request.store_id, &request.key]) @@ -589,7 +687,7 @@ where } let mut conn = self.pool.get().await?; - let transaction = conn + let mut transaction = conn .transaction() .await .map_err(|e| Error::new(ErrorKind::Other, format!("Transaction start error: {}", e)))?; @@ -597,12 +695,12 @@ where let mut batch_results = Vec::new(); for vss_record in &vss_put_records { - let num_rows = self.execute_put_object_query(&transaction, vss_record).await?; + let num_rows = self.execute_put_object_query(&mut transaction, vss_record).await?; batch_results.push(num_rows); } for vss_record in &vss_delete_records { - let num_rows = self.execute_delete_object_query(&transaction, vss_record).await?; + let num_rows = self.execute_delete_object_query(&mut transaction, vss_record).await?; batch_results.push(num_rows); } @@ -633,12 +731,12 @@ where let vss_record = self.build_vss_record(user_token, store_id, key_value); let mut conn = self.pool.get().await?; - let transaction = conn + let mut transaction = conn .transaction() .await .map_err(|e| Error::new(ErrorKind::Other, format!("Transaction start error: {}", e)))?; - let num_rows = self.execute_delete_object_query(&transaction, &vss_record).await?; + let num_rows = self.execute_delete_object_query(&mut transaction, &vss_record).await?; if num_rows == 0 { transaction.rollback().await.map_err(|e| { @@ -689,14 +787,14 @@ where // Fetch one extra to determine if there are more pages. let fetch_limit = limit + 1; - let conn = self.pool.get().await?; + let mut conn = self.pool.get().await?; let key_like = format!("{}%", key_prefix.as_deref().unwrap_or_default()); let rows = if let Some(ref token) = page_token { let page_sort_order = decode_page_token(token)?; let stmt = "SELECT key, version, sort_order FROM vss_db WHERE user_token = $1 AND store_id = $2 AND sort_order < $3 AND key LIKE $4 AND key != $5 ORDER BY sort_order DESC LIMIT $6"; - let params: Vec<&(dyn tokio_postgres::types::ToSql + Sync)> = vec![ + let params: Vec<&(dyn ToSql + Sync)> = vec![ &user_token, &store_id, &page_sort_order, @@ -709,7 +807,7 @@ where .map_err(|e| Error::new(ErrorKind::Other, format!("Query error: {}", e)))? } else { let stmt = "SELECT key, version, sort_order FROM vss_db WHERE user_token = $1 AND store_id = $2 AND key LIKE $3 AND key != $4 ORDER BY sort_order DESC LIMIT $5"; - let params: Vec<&(dyn tokio_postgres::types::ToSql + Sync)> = + let params: Vec<&(dyn ToSql + Sync)> = vec![&user_token, &store_id, &key_like, &GLOBAL_VERSION_KEY, &fetch_limit]; conn.query(stmt, ¶ms) .await @@ -743,19 +841,26 @@ where #[cfg(test)] mod tests { - use super::{decode_page_token, drop_database, encode_page_token, DUMMY_MIGRATION, MIGRATIONS}; + use super::{ + decode_page_token, drop_database, encode_page_token, Client, DUMMY_MIGRATION, + INITIAL_RECORD_VERSION_STR, MIGRATIONS, + }; use crate::postgres_store::PostgresPlaintextBackend; use api::define_kv_store_tests; - use api::kv_store::KvStore; + use api::kv_store::{KvStore, INITIAL_RECORD_VERSION}; use api::types::{ DeleteObjectRequest, GetObjectRequest, KeyValue, ListKeyVersionsRequest, PutObjectRequest, }; use bytes::Bytes; + use std::sync::LazyLock; use tokio::sync::OnceCell; use tokio_postgres::NoTls; - const POSTGRES_ENDPOINT: &str = "postgresql://postgres:postgres@localhost:5432"; + static POSTGRES_ENDPOINT: LazyLock = LazyLock::new(|| { + std::env::var("POSTGRES_ENDPOINT") + .unwrap_or_else(|_| "postgresql://postgres:postgres@localhost:5432".to_string()) + }); const DEFAULT_DB: &str = "postgres"; const MIGRATIONS_START: usize = 0; const MIGRATIONS_END: usize = MIGRATIONS.len(); @@ -788,8 +893,8 @@ mod tests { let vss_db = "postgres_kv_store_tests"; START .get_or_init(|| async { - let _ = drop_database(POSTGRES_ENDPOINT, DEFAULT_DB, vss_db, NoTls).await; - let store = PostgresPlaintextBackend::new(POSTGRES_ENDPOINT, DEFAULT_DB, vss_db) + let _ = drop_database(&POSTGRES_ENDPOINT, DEFAULT_DB, vss_db, NoTls).await; + let store = PostgresPlaintextBackend::new(&POSTGRES_ENDPOINT, DEFAULT_DB, vss_db) .await .unwrap(); let (start, end) = store.migrate_vss_database(MIGRATIONS).await.unwrap(); @@ -798,7 +903,7 @@ mod tests { }) .await; let store = - PostgresPlaintextBackend::new(POSTGRES_ENDPOINT, DEFAULT_DB, vss_db).await.unwrap(); + PostgresPlaintextBackend::new(&POSTGRES_ENDPOINT, DEFAULT_DB, vss_db).await.unwrap(); let (start, end) = store.migrate_vss_database(MIGRATIONS).await.unwrap(); assert_eq!(start, MIGRATIONS_END); assert_eq!(end, MIGRATIONS_END); @@ -811,19 +916,21 @@ mod tests { #[should_panic(expected = "We do not allow downgrades")] async fn panic_on_downgrade() { let vss_db = "panic_on_downgrade_test"; - let _ = drop_database(POSTGRES_ENDPOINT, DEFAULT_DB, vss_db, NoTls).await; + let _ = drop_database(&POSTGRES_ENDPOINT, DEFAULT_DB, vss_db, NoTls).await; { let mut migrations = MIGRATIONS.to_vec(); migrations.push(DUMMY_MIGRATION); - let store = - PostgresPlaintextBackend::new(POSTGRES_ENDPOINT, DEFAULT_DB, vss_db).await.unwrap(); + let store = PostgresPlaintextBackend::new(&POSTGRES_ENDPOINT, DEFAULT_DB, vss_db) + .await + .unwrap(); let (start, end) = store.migrate_vss_database(&migrations).await.unwrap(); assert_eq!(start, MIGRATIONS_START); assert_eq!(end, MIGRATIONS_END + 1); }; { - let store = - PostgresPlaintextBackend::new(POSTGRES_ENDPOINT, DEFAULT_DB, vss_db).await.unwrap(); + let store = PostgresPlaintextBackend::new(&POSTGRES_ENDPOINT, DEFAULT_DB, vss_db) + .await + .unwrap(); let _ = store.migrate_vss_database(MIGRATIONS).await.unwrap(); }; } @@ -831,10 +938,11 @@ mod tests { #[tokio::test] async fn new_migrations_increments_upgrades() { let vss_db = "new_migrations_increments_upgrades_test"; - let _ = drop_database(POSTGRES_ENDPOINT, DEFAULT_DB, vss_db, NoTls).await; + let _ = drop_database(&POSTGRES_ENDPOINT, DEFAULT_DB, vss_db, NoTls).await; { - let store = - PostgresPlaintextBackend::new(POSTGRES_ENDPOINT, DEFAULT_DB, vss_db).await.unwrap(); + let store = PostgresPlaintextBackend::new(&POSTGRES_ENDPOINT, DEFAULT_DB, vss_db) + .await + .unwrap(); let (start, end) = store.migrate_vss_database(MIGRATIONS).await.unwrap(); assert_eq!(start, MIGRATIONS_START); assert_eq!(end, MIGRATIONS_END); @@ -842,8 +950,9 @@ mod tests { assert_eq!(store.get_schema_version().await, MIGRATIONS_END); }; { - let store = - PostgresPlaintextBackend::new(POSTGRES_ENDPOINT, DEFAULT_DB, vss_db).await.unwrap(); + let store = PostgresPlaintextBackend::new(&POSTGRES_ENDPOINT, DEFAULT_DB, vss_db) + .await + .unwrap(); let (start, end) = store.migrate_vss_database(MIGRATIONS).await.unwrap(); assert_eq!(start, MIGRATIONS_END); assert_eq!(end, MIGRATIONS_END); @@ -854,8 +963,9 @@ mod tests { let mut migrations = MIGRATIONS.to_vec(); migrations.push(DUMMY_MIGRATION); { - let store = - PostgresPlaintextBackend::new(POSTGRES_ENDPOINT, DEFAULT_DB, vss_db).await.unwrap(); + let store = PostgresPlaintextBackend::new(&POSTGRES_ENDPOINT, DEFAULT_DB, vss_db) + .await + .unwrap(); let (start, end) = store.migrate_vss_database(&migrations).await.unwrap(); assert_eq!(start, MIGRATIONS_END); assert_eq!(end, MIGRATIONS_END + 1); @@ -866,8 +976,9 @@ mod tests { migrations.push(DUMMY_MIGRATION); migrations.push(DUMMY_MIGRATION); { - let store = - PostgresPlaintextBackend::new(POSTGRES_ENDPOINT, DEFAULT_DB, vss_db).await.unwrap(); + let store = PostgresPlaintextBackend::new(&POSTGRES_ENDPOINT, DEFAULT_DB, vss_db) + .await + .unwrap(); let (start, end) = store.migrate_vss_database(&migrations).await.unwrap(); assert_eq!(start, MIGRATIONS_END + 1); assert_eq!(end, MIGRATIONS_END + 3); @@ -879,21 +990,22 @@ mod tests { }; { - let store = - PostgresPlaintextBackend::new(POSTGRES_ENDPOINT, DEFAULT_DB, vss_db).await.unwrap(); + let store = PostgresPlaintextBackend::new(&POSTGRES_ENDPOINT, DEFAULT_DB, vss_db) + .await + .unwrap(); let list = store.get_upgrades_list().await; assert_eq!(list, [MIGRATIONS_START, MIGRATIONS_END, MIGRATIONS_END + 1]); let version = store.get_schema_version().await; assert_eq!(version, MIGRATIONS_END + 3); } - drop_database(POSTGRES_ENDPOINT, DEFAULT_DB, vss_db, NoTls).await.unwrap(); + drop_database(&POSTGRES_ENDPOINT, DEFAULT_DB, vss_db, NoTls).await.unwrap(); } #[tokio::test] async fn supports_objects_up_to_non_large_object_threshold() { let vss_db = "supports_objects_up_to_non_large_object_threshold"; - let _ = drop_database(POSTGRES_ENDPOINT, DEFAULT_DB, vss_db, NoTls).await; + let _ = drop_database(&POSTGRES_ENDPOINT, DEFAULT_DB, vss_db, NoTls).await; const MAXIMUM_SUPPORTED_VALUE_SIZE: usize = 1024 * 1024 * 1024; const PROTOCOL_OVERHEAD_MARGIN: usize = 150; @@ -903,8 +1015,9 @@ mod tests { let kv = KeyValue { key: "k1".into(), version: 0, value: Bytes::from(large_value.clone()) }; { - let store = - PostgresPlaintextBackend::new(POSTGRES_ENDPOINT, DEFAULT_DB, vss_db).await.unwrap(); + let store = PostgresPlaintextBackend::new(&POSTGRES_ENDPOINT, DEFAULT_DB, vss_db) + .await + .unwrap(); let (start, end) = store.migrate_vss_database(MIGRATIONS).await.unwrap(); assert_eq!(start, MIGRATIONS_START); assert_eq!(end, MIGRATIONS_END); @@ -953,17 +1066,18 @@ mod tests { .unwrap(); }; - drop_database(POSTGRES_ENDPOINT, DEFAULT_DB, vss_db, NoTls).await.unwrap(); + drop_database(&POSTGRES_ENDPOINT, DEFAULT_DB, vss_db, NoTls).await.unwrap(); } #[tokio::test] async fn list_orders_by_sort_order_desc() { let vss_db = "list_orders_by_sort_order_desc"; - let _ = drop_database(POSTGRES_ENDPOINT, DEFAULT_DB, vss_db, NoTls).await; + let _ = drop_database(&POSTGRES_ENDPOINT, DEFAULT_DB, vss_db, NoTls).await; { - let store = - PostgresPlaintextBackend::new(POSTGRES_ENDPOINT, DEFAULT_DB, vss_db).await.unwrap(); + let store = PostgresPlaintextBackend::new(&POSTGRES_ENDPOINT, DEFAULT_DB, vss_db) + .await + .unwrap(); let (start, end) = store.migrate_vss_database(MIGRATIONS).await.unwrap(); assert_eq!(start, MIGRATIONS_START); assert_eq!(end, MIGRATIONS_END); @@ -1014,17 +1128,18 @@ mod tests { assert_eq!(all_keys, vec!["b_key", "a_key", "c_key"]); } - drop_database(POSTGRES_ENDPOINT, DEFAULT_DB, vss_db, NoTls).await.unwrap(); + drop_database(&POSTGRES_ENDPOINT, DEFAULT_DB, vss_db, NoTls).await.unwrap(); } #[tokio::test] async fn list_zero_page_size_should_return_only_global_version() { let vss_db = "list_zero_page_size_should_return_only_global_version"; - let _ = drop_database(POSTGRES_ENDPOINT, DEFAULT_DB, vss_db, NoTls).await; + let _ = drop_database(&POSTGRES_ENDPOINT, DEFAULT_DB, vss_db, NoTls).await; { - let store = - PostgresPlaintextBackend::new(POSTGRES_ENDPOINT, DEFAULT_DB, vss_db).await.unwrap(); + let store = PostgresPlaintextBackend::new(&POSTGRES_ENDPOINT, DEFAULT_DB, vss_db) + .await + .unwrap(); let (start, end) = store.migrate_vss_database(MIGRATIONS).await.unwrap(); assert_eq!(start, MIGRATIONS_START); assert_eq!(end, MIGRATIONS_END); @@ -1051,17 +1166,18 @@ mod tests { assert_eq!(resp.next_page_token.filter(|t| !t.is_empty()), None); } - drop_database(POSTGRES_ENDPOINT, DEFAULT_DB, vss_db, NoTls).await.unwrap(); + drop_database(&POSTGRES_ENDPOINT, DEFAULT_DB, vss_db, NoTls).await.unwrap(); } #[tokio::test] async fn list_should_return_empty_page_token_when_exact_fit() { let vss_db = "list_should_return_empty_page_token_when_exact_fit"; - let _ = drop_database(POSTGRES_ENDPOINT, DEFAULT_DB, vss_db, NoTls).await; + let _ = drop_database(&POSTGRES_ENDPOINT, DEFAULT_DB, vss_db, NoTls).await; { - let store = - PostgresPlaintextBackend::new(POSTGRES_ENDPOINT, DEFAULT_DB, vss_db).await.unwrap(); + let store = PostgresPlaintextBackend::new(&POSTGRES_ENDPOINT, DEFAULT_DB, vss_db) + .await + .unwrap(); let (start, end) = store.migrate_vss_database(MIGRATIONS).await.unwrap(); assert_eq!(start, MIGRATIONS_START); assert_eq!(end, MIGRATIONS_END); @@ -1089,17 +1205,18 @@ mod tests { assert_eq!(resp.next_page_token.filter(|t| !t.is_empty()), None); } - drop_database(POSTGRES_ENDPOINT, DEFAULT_DB, vss_db, NoTls).await.unwrap(); + drop_database(&POSTGRES_ENDPOINT, DEFAULT_DB, vss_db, NoTls).await.unwrap(); } #[tokio::test] async fn list_should_return_empty_page_token_on_last_non_empty_page() { let vss_db = "list_should_return_empty_page_token_on_last_non_empty_page"; - let _ = drop_database(POSTGRES_ENDPOINT, DEFAULT_DB, vss_db, NoTls).await; + let _ = drop_database(&POSTGRES_ENDPOINT, DEFAULT_DB, vss_db, NoTls).await; { - let store = - PostgresPlaintextBackend::new(POSTGRES_ENDPOINT, DEFAULT_DB, vss_db).await.unwrap(); + let store = PostgresPlaintextBackend::new(&POSTGRES_ENDPOINT, DEFAULT_DB, vss_db) + .await + .unwrap(); let (start, end) = store.migrate_vss_database(MIGRATIONS).await.unwrap(); assert_eq!(start, MIGRATIONS_START); assert_eq!(end, MIGRATIONS_END); @@ -1143,7 +1260,27 @@ mod tests { assert!(second_page.global_version.is_none()); } - drop_database(POSTGRES_ENDPOINT, DEFAULT_DB, vss_db, NoTls).await.unwrap(); + drop_database(&POSTGRES_ENDPOINT, DEFAULT_DB, vss_db, NoTls).await.unwrap(); + } + + #[test] + fn initial_record_version_string_matches_numeric_value() { + assert_eq!(&INITIAL_RECORD_VERSION.to_string(), INITIAL_RECORD_VERSION_STR); + } + + #[tokio::test] + async fn prepared_statements_are_reused_across_transactions() { + let mut conn = Client::connect(&POSTGRES_ENDPOINT, DEFAULT_DB, NoTls).await.unwrap(); + let stmt = "SELECT $1::BIGINT"; + + let mut transaction = conn.transaction().await.unwrap(); + transaction.execute(stmt, &[&1_i64]).await.unwrap(); + transaction.commit().await.unwrap(); + assert_eq!(conn.statement_cache.statements.len(), 1); + + let rows = conn.query(stmt, &[&2_i64]).await.unwrap(); + assert_eq!(rows[0].get::<_, i64>(0), 2); + assert_eq!(conn.statement_cache.statements.len(), 1); } #[test]