diff --git a/rust/src/storage/db/mod.rs b/rust/src/storage/db/mod.rs index 14f4a2692e..c09c3dba2d 100644 --- a/rust/src/storage/db/mod.rs +++ b/rust/src/storage/db/mod.rs @@ -18,6 +18,14 @@ pub mod rust_db_pool; #[async_trait::async_trait] pub trait DatabasePool { + /// TODO + async fn get_connection(&self) -> Result, anyhow::Error>; +} + +/// A `tokio_postgres` Connection looking thing that we can use on the Rust side to +/// interact with the database +#[async_trait::async_trait] +pub trait DatabaseConnection { /// TODO /// /// Arguments: @@ -32,7 +40,7 @@ pub trait DatabasePool { /// interact with the database #[async_trait::async_trait] pub trait Transaction { - async fn query(&self, sql: &str, args: &[&str]) -> Vec; + async fn query(&self, sql: &str, args: &[&str]) -> Result, anyhow::Error>; async fn commit(self) -> Result<(), anyhow::Error>; } diff --git a/rust/src/storage/db/python_db_pool.rs b/rust/src/storage/db/python_db_pool.rs index a2636716cf..b284515f69 100644 --- a/rust/src/storage/db/python_db_pool.rs +++ b/rust/src/storage/db/python_db_pool.rs @@ -15,7 +15,7 @@ use pyo3::{intern, prelude::*}; -use crate::storage::db::{DatabasePool, Row, Transaction}; +use crate::storage::db::{DatabaseConnection, DatabasePool, Row, Transaction}; /// The database engines we support in the Python side of Synapse #[derive(Copy, Clone, Debug)] @@ -42,14 +42,48 @@ pub struct PythonDatabasePool { #[async_trait::async_trait] impl DatabasePool for PythonDatabasePool { + async fn get_connection(&self) -> Result, anyhow::Error> { + let callback_func = + PyCFunction::new_closure(py, None, None, move |args, _| -> PyResult> { + // TODO: Error if already called + + let py = args.py(); + // We found our `LoggingDatabaseConnection` + let conn_py = args.get_item(0)?; + tx.send(conn_py); + }); + + let execute_fn = self + .database_pool_py + .getattr(intern!(self.database_pool_py.py(), "runWithConnection"))?; + execute_fn.call1((callback_func,))?; + + Ok(Box::new(LoggingDatabaseConnectionWrapper { + database_pool_py: self, + logging_database_connection_py: conn_py, + })) + } +} + +/// Wrapper for a `LoggingDatabaseConnection` from the Python side of Synapse. +pub struct LoggingDatabaseConnectionWrapper { + /// The underlying `DatabasePool` + database_pool_py: Bound<'py, PyAny>, + /// The underlying `LoggingDatabaseConnection` + logging_database_connection_py: Py, +} + +#[async_trait::async_trait] +impl DatabaseConnection for LoggingDatabaseConnectionWrapper { async fn get_transaction( &self, description: &str, ) -> Result, anyhow::Error> { - // Synapse has built-in retry functionality and can call this function multiple - // times under certain failure modes. Normally, everything in the transaction - // happens in the callback but since we have a little bit of a different API - // surface, we instead extract the transaction for us to use outside. + // Synapse's `DatabasePool.new_transaction` has built-in retry functionality and + // can call this function multiple times under certain failure modes. Normally, + // everything in the transaction happens in the callback but since we have a + // little bit of a different API surface, we instead extract the transaction for + // us to use outside. // // Re-using `runInteraction`, means we get all of the logging, metrics, etc for // free. @@ -58,14 +92,23 @@ impl DatabasePool for PythonDatabasePool { // TODO: Error if already called let py = args.py(); - let txn_py = args.get_item(0)?; - txn + let txn_py: LoggingTransactionWrapper = args.get_item(0)?; + tx.send(txn_py); + + // Wait until we see the signal that we're `done_with_txn_rx` }); let execute_fn = self .database_pool_py - .getattr(intern!(self.database_pool_py.py(), "runInteraction"))?; - execute_fn.call1((description, callback_func))?; + .getattr(intern!(self.database_pool_py.py(), "new_transaction"))?; + execute_fn.call1(( + self.logging_database_connection_py, + description, + [], + [], + [], + callback_func, + ))?; Ok(Box::new(txn)) } @@ -93,6 +136,7 @@ fn detect_engine(txn_py: &Bound<'_, PyAny>) -> PyResult { /// Wrapper for a `LoggingTransaction` from the Python side of Synapse. /// +/// TODO: Verify if this is the correct approach or we should use`Bound<'py, PyAny>`: /// Holds no `'py` lifetime so it can be stored and moved freely across threads. /// Use [`execute`](Self::execute) (or other methods) while holding the GIL. pub struct LoggingTransactionWrapper { @@ -118,6 +162,26 @@ impl<'a, 'py> FromPyObject<'a, 'py> for LoggingTransactionWrapper { } } +pub trait ValidDatabaseFieldType {} +pub trait ValidDatabaseReturnType {} +impl ValidDatabaseFieldType for String {} +impl ValidDatabaseFieldType for usize {} +impl ValidDatabaseFieldType for Option {} +impl ValidDatabaseReturnType for (T0,) {} +impl ValidDatabaseReturnType for (T0, T1) {} +impl + ValidDatabaseReturnType for (T0, T1, T2) +{ +} +impl< + T0: ValidDatabaseFieldType, + T1: ValidDatabaseFieldType, + T2: ValidDatabaseFieldType, + T3: ValidDatabaseFieldType, + > ValidDatabaseReturnType for (T0, T1, T2, T3) +{ +} + impl LoggingTransactionWrapper { pub fn execute<'py>( &mut self, @@ -132,15 +196,33 @@ impl LoggingTransactionWrapper { execute_fn.call1((sql, args))?; Ok(()) } + + pub fn fetchall + ValidDatabaseReturnType>( + &mut self, + ) -> anyhow::Result> { + let fetch_fn = self + .logging_transaction_py + .getattr(intern!(self.logging_transaction_py.py(), "fetchall"))?; + Ok(fetch_fn.call0()?.extract()?) + } } #[async_trait::async_trait] impl Transaction for LoggingTransactionWrapper { - async fn query(&self, sql: &str, args: &[&str]) -> Vec { + async fn query(&self, sql: &str, args: &[&str]) -> Result, anyhow::Error> { self.execute(sql, args).await; + let rows = self.fetchall(sql, args).await?; + + Ok(rows) } async fn commit(&self) -> Result<(), anyhow::Error> { - // In Synapse, `commit` is part of `LoggingDatabaseConnection` + // In Synapse, `commit` is part of `LoggingDatabaseConnection` and will be + // called as part of the `new_transaction(...)` machinery we used to create the + // transaction in the first place. + // + // We just need to send the proper signal which will finish the txn callback and + // have it run. + done_with_txn_tx.send(()) } } diff --git a/rust/src/storage/db/rust_db_pool.rs b/rust/src/storage/db/rust_db_pool.rs index c12cd81199..bdf7ed375b 100644 --- a/rust/src/storage/db/rust_db_pool.rs +++ b/rust/src/storage/db/rust_db_pool.rs @@ -20,7 +20,7 @@ use anyhow::Context; use bb8_postgres::PostgresConnectionManager; use postgres_native_tls::MakeTlsConnector; -use crate::storage::db::{DatabasePool, Row, Transaction}; +use crate::storage::db::{DatabaseConnection, DatabasePool, Row, Transaction}; /// Native Rust database access backed by `tokio-postgres` (for use in synapse-rust-apps) pub struct RustDatabasePool { @@ -29,10 +29,7 @@ pub struct RustDatabasePool { #[async_trait::async_trait] impl DatabasePool for RustDatabasePool { - async fn get_transaction( - &self, - _description: &str, - ) -> Result, anyhow::Error> { + async fn get_connection(&self) -> Result, anyhow::Error> { let mut conn = self .db_pool .get() @@ -40,9 +37,22 @@ impl DatabasePool for RustDatabasePool { .await .context("Failed to acquire database connection")?; + Ok(Box::new(RustConnection { connection: conn })) + } +} +pub struct RustConnection<'a> { + connection: bb8::PooledConnection<'a, PostgresConnectionManager>, +} + +impl DatabaseConnection for RustConnection<'_> { + async fn get_transaction( + &self, + _description: &str, + ) -> Result, anyhow::Error> { // TODO: Set repeatable-read isolation level (like Synapse) - let txn = conn + let txn = self + .connection .transaction() // .instrument(tracing::info_span!("start transaction")) .await @@ -58,12 +68,12 @@ struct TokioPostgresTransaction<'a> { #[async_trait::async_trait] impl Transaction for TokioPostgresTransaction<'_> { - async fn query(&self, sql: &str, args: &[&str]) -> Vec { + async fn query(&self, sql: &str, args: &[&str]) -> Result, anyhow::Error> { // TODO: Convert `?` SQL param style to `tokio-postgres` compatible let rows = self.txn.query(sql, args).await?; - rows + Ok(rows) } async fn commit(self) -> Result<(), anyhow::Error> {