diff --git a/impls/src/postgres_store.rs b/impls/src/postgres_store.rs index 41dd921..6821716 100644 --- a/impls/src/postgres_store.rs +++ b/impls/src/postgres_store.rs @@ -12,10 +12,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::{Client, NoTls, Socket, Transaction, error}; +use tokio_postgres::types::ToSql; +use tokio_postgres::{NoTls, Row, Socket, Statement, error}; use log::{debug, info, warn}; @@ -66,6 +68,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 +171,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 +210,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 +239,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 +369,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 +387,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 +450,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,20 +477,21 @@ 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) + const STMT: &str = "INSERT INTO vss_db (user_token, store_id, key, value, version, created_at, last_updated_at) + VALUES ($1, $2, $3, $4, $5, $6, $7) 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); + SET value = EXCLUDED.value, version = EXCLUDED.version, last_updated_at = EXCLUDED.last_updated_at"; let num_rows = transaction .execute( - &stmt, + STMT, &[ &vss_record.user_token, &vss_record.store_id, &vss_record.key, &vss_record.value, + &i64::from(INITIAL_RECORD_VERSION), &vss_record.created_at, &vss_record.last_updated_at, ], @@ -414,19 +504,20 @@ 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); + const STMT: &str = "INSERT INTO vss_db (user_token, store_id, key, value, version, created_at, last_updated_at) + VALUES ($1, $2, $3, $4, $5, $6, $7) + ON CONFLICT DO NOTHING"; let num_rows = transaction .execute( - &stmt, + STMT, &[ &vss_record.user_token, &vss_record.store_id, &vss_record.key, &vss_record.value, + &i64::from(INITIAL_RECORD_VERSION), &vss_record.created_at, &vss_record.last_updated_at, ], @@ -439,7 +530,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 +555,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 +567,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 +580,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 +601,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 +622,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 +680,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 +688,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 +724,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 +780,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(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 +800,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,7 +834,9 @@ where #[cfg(test)] mod tests { - use super::{DUMMY_MIGRATION, MIGRATIONS, decode_page_token, drop_database, encode_page_token}; + use super::{ + Client, DUMMY_MIGRATION, MIGRATIONS, decode_page_token, drop_database, encode_page_token, + }; use crate::postgres_store::PostgresPlaintextBackend; use api::define_kv_store_tests; use api::kv_store::KvStore; @@ -753,9 +846,12 @@ mod tests { use bytes::Bytes; use tokio::sync::OnceCell; - use tokio_postgres::NoTls; + use tokio_postgres::{NoTls, SimpleQueryMessage}; - const POSTGRES_ENDPOINT: &str = "postgresql://postgres:postgres@localhost:5432"; + const POSTGRES_ENDPOINT: &str = match option_env!("POSTGRES_ENDPOINT") { + Some(endpoint) => endpoint, + None => "postgresql://postgres:postgres@localhost:5432", + }; const DEFAULT_DB: &str = "postgres"; const MIGRATIONS_START: usize = 0; const MIGRATIONS_END: usize = MIGRATIONS.len(); @@ -1146,6 +1242,56 @@ mod tests { drop_database(POSTGRES_ENDPOINT, DEFAULT_DB, vss_db, NoTls).await.unwrap(); } + #[tokio::test] + async fn prepared_statements_should_be_reused_within_and_across_transactions() { + async fn statement_usage(conn: &Client, stmt: &str) -> (String, u64) { + // Inspect the same session without preparing another statement for the inspection. + let messages = conn + .uncached_client + .simple_query( + "SELECT name, statement, generic_plans + custom_plans AS uses + FROM pg_prepared_statements", + ) + .await + .unwrap(); + let statements: Vec<_> = messages + .into_iter() + .filter_map(|message| match message { + SimpleQueryMessage::Row(row) if row.get("statement") == Some(stmt) => Some(( + row.get("name").unwrap().to_owned(), + row.get("uses").unwrap().parse::().unwrap(), + )), + _ => None, + }) + .collect(); + assert_eq!(statements.len(), 1, "expected one prepared statement: {statements:?}"); + statements.into_iter().next().unwrap() + } + + let mut conn = Client::connect(POSTGRES_ENDPOINT, DEFAULT_DB, NoTls).await.unwrap(); + let stmt = "SELECT $1::BIGINT"; + + let mut transaction = conn.transaction().await.unwrap(); + assert_eq!(transaction.execute(stmt, &[&1_i64]).await.unwrap(), 1); + assert_eq!(transaction.execute(stmt, &[&2_i64]).await.unwrap(), 1); + transaction.commit().await.unwrap(); + let (name, uses) = statement_usage(&conn, stmt).await; + assert_eq!(uses, 2, "statement should be reused within a transaction"); + + let mut transaction = conn.transaction().await.unwrap(); + assert_eq!(transaction.execute(stmt, &[&3_i64]).await.unwrap(), 1); + transaction.commit().await.unwrap(); + assert_eq!(statement_usage(&conn, stmt).await, (name.clone(), 3)); + + let rows = conn.query(stmt, &[&4_i64]).await.unwrap(); + assert_eq!(rows[0].get::<_, i64>(0), 4); + assert_eq!(statement_usage(&conn, stmt).await, (name.clone(), 4)); + + let row = conn.query_opt(stmt, &[&5_i64]).await.unwrap().unwrap(); + assert_eq!(row.get::<_, i64>(0), 5); + assert_eq!(statement_usage(&conn, stmt).await, (name, 5)); + } + #[test] fn page_token_roundtrips() { let token = encode_page_token(12345);