use std::{
any::TypeId, future::Future, marker::PhantomData, mem, panic::AssertUnwindSafe, pin::pin,
task::Poll,
};
use saddle_admission::DbRequestPermit;
use saddle_runtime::db_finalizer::{
DbQueryPoll, DbQueryTransition, DbTransitionRequest, drive_db_finalizer_with_transition,
};
use sqlx::MySql;
use crate::{
Database,
c6_query_optional::{
QueryOptionalContractError, QueryOptionalParameterShape, sealed::ParameterShape as _,
},
};
pub trait StaticWriteOperation: Send + 'static {
type Parameters: QueryOptionalParameterShape;
const OPERATION: &'static str;
const SQL: &'static str;
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub struct WriteCreditDemand {
connections: u32,
operations: u32,
}
impl WriteCreditDemand {
const ONE: Self = Self {
connections: 1,
operations: 1,
};
pub const fn connections(self) -> u32 {
self.connections
}
pub const fn operations(self) -> u32 {
self.operations
}
pub const fn merge_route_max(self, other: Self) -> Self {
Self {
connections: if self.connections > other.connections {
self.connections
} else {
other.connections
},
operations: if self.operations > other.operations {
self.operations
} else {
other.operations
},
}
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub struct WriteLayout {
operation: TypeId,
parameter_bytes: usize,
credits: WriteCreditDemand,
}
impl WriteLayout {
pub const fn parameter_bytes(self) -> usize {
self.parameter_bytes
}
pub const fn credits(self) -> WriteCreditDemand {
self.credits
}
pub fn belongs_to<O: StaticWriteOperation>(self) -> bool {
self.operation == TypeId::of::<O>()
}
}
#[derive(Clone, Copy, Debug)]
pub struct WriteOperationProof<O: StaticWriteOperation> {
layout: WriteLayout,
_operation: PhantomData<fn() -> O>,
}
impl<O: StaticWriteOperation> WriteOperationProof<O> {
pub fn bind() -> Result<Self, QueryOptionalContractError> {
validate_operation(O::OPERATION)?;
validate_sql(O::SQL)?;
let parameter_bytes = mem::size_of::<O::Parameters>();
if parameter_bytes == 0 {
return Err(QueryOptionalContractError::EmptyShape);
}
Ok(Self {
layout: WriteLayout {
operation: TypeId::of::<O>(),
parameter_bytes,
credits: WriteCreditDemand::ONE,
},
_operation: PhantomData,
})
}
pub const fn layout(&self) -> WriteLayout {
self.layout
}
pub fn invocation(self, parameters: O::Parameters) -> WriteInvocation<O> {
WriteInvocation {
layout: self.layout,
parameters,
_operation: PhantomData,
}
}
}
pub struct WriteInvocation<O: StaticWriteOperation> {
layout: WriteLayout,
parameters: O::Parameters,
_operation: PhantomData<fn() -> O>,
}
impl<O: StaticWriteOperation> WriteInvocation<O> {
pub const fn layout(&self) -> WriteLayout {
self.layout
}
}
pub struct WriteExecution<O: StaticWriteOperation> {
permit: DbRequestPermit,
invocation: WriteInvocation<O>,
}
impl<O: StaticWriteOperation> WriteExecution<O> {
#[doc(hidden)]
pub fn from_compiled_handoff(permit: DbRequestPermit, invocation: WriteInvocation<O>) -> Self {
Self { permit, invocation }
}
fn into_parts(self) -> (DbRequestPermit, WriteInvocation<O>) {
(self.permit, self.invocation)
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub struct ManagedWriteResult {
rows_affected: u64,
last_insert_id: u64,
}
impl ManagedWriteResult {
pub const fn rows_affected(self) -> u64 {
self.rows_affected
}
pub const fn last_insert_id(self) -> u64 {
self.last_insert_id
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum WriteExecutionError {
ConnectionUnavailable,
WriteFailed,
Cancelled,
Shutdown,
FinalizerFailed,
}
impl WriteExecutionError {
pub const fn code(self) -> &'static str {
match self {
Self::ConnectionUnavailable => "db.connection_unavailable",
Self::WriteFailed => "db.write_failed",
Self::Cancelled => "db.write_cancelled",
Self::Shutdown => "db.write_shutdown",
Self::FinalizerFailed => "db.finalizer_failed",
}
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum TransactionDecision {
Commit,
Rollback,
}
pub struct TransactionInvocation<O: StaticWriteOperation> {
write: WriteInvocation<O>,
decision: TransactionDecision,
}
impl<O: StaticWriteOperation> TransactionInvocation<O> {
pub const fn layout(&self) -> WriteLayout {
self.write.layout
}
pub const fn decision(&self) -> TransactionDecision {
self.decision
}
}
#[derive(Debug)]
pub struct TransactionOperationProof<O: StaticWriteOperation> {
write: WriteOperationProof<O>,
}
impl<O: StaticWriteOperation> TransactionOperationProof<O> {
pub fn bind() -> Result<Self, QueryOptionalContractError> {
WriteOperationProof::bind().map(|write| Self { write })
}
pub const fn layout(&self) -> WriteLayout {
self.write.layout()
}
pub fn invocation(
self,
parameters: O::Parameters,
decision: TransactionDecision,
) -> TransactionInvocation<O> {
TransactionInvocation {
write: self.write.invocation(parameters),
decision,
}
}
}
pub struct TransactionExecution<O: StaticWriteOperation> {
permit: DbRequestPermit,
invocation: TransactionInvocation<O>,
}
impl<O: StaticWriteOperation> TransactionExecution<O> {
#[doc(hidden)]
pub fn from_compiled_handoff(
permit: DbRequestPermit,
invocation: TransactionInvocation<O>,
) -> Self {
Self { permit, invocation }
}
fn into_parts(self) -> (DbRequestPermit, TransactionInvocation<O>) {
(self.permit, self.invocation)
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum TransactionExecutionError {
ConnectionUnavailable,
BeginFailed,
WriteFailed,
BusinessRollback,
CommitFailed,
RollbackFailed,
Cancelled,
Shutdown,
FinalizerFailed,
}
impl TransactionExecutionError {
pub const fn code(self) -> &'static str {
match self {
Self::ConnectionUnavailable => "db.connection_unavailable",
Self::BeginFailed => "db.transaction_begin_failed",
Self::WriteFailed => "db.write_failed",
Self::BusinessRollback => "db.transaction_business_rollback",
Self::CommitFailed => "db.transaction_commit_failed",
Self::RollbackFailed => "db.transaction_rollback_failed",
Self::Cancelled => "db.transaction_cancelled",
Self::Shutdown => "db.transaction_shutdown",
Self::FinalizerFailed => "db.finalizer_failed",
}
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub struct WriteTransactionFinalizationProof {
cancel_external_io_awaits: u8,
shutdown_external_io_awaits: u8,
panic_external_io_awaits: u8,
single_connection: bool,
nested_transactions: bool,
releases_pool_size_before_permit: bool,
normal_return_requires_termination_bound: bool,
}
impl WriteTransactionFinalizationProof {
const PRODUCTION: Self = Self {
cancel_external_io_awaits: 0,
shutdown_external_io_awaits: 0,
panic_external_io_awaits: 0,
single_connection: true,
nested_transactions: false,
releases_pool_size_before_permit: true,
normal_return_requires_termination_bound: true,
};
pub const fn cancel_external_io_awaits(self) -> u8 {
self.cancel_external_io_awaits
}
pub const fn shutdown_external_io_awaits(self) -> u8 {
self.shutdown_external_io_awaits
}
pub const fn panic_external_io_awaits(self) -> u8 {
self.panic_external_io_awaits
}
pub const fn single_connection(self) -> bool {
self.single_connection
}
pub const fn nested_transactions(self) -> bool {
self.nested_transactions
}
pub const fn releases_pool_size_before_permit(self) -> bool {
self.releases_pool_size_before_permit
}
pub const fn normal_return_requires_termination_bound(self) -> bool {
self.normal_return_requires_termination_bound
}
}
#[doc(hidden)]
pub const fn write_transaction_finalization_proof() -> WriteTransactionFinalizationProof {
WriteTransactionFinalizationProof::PRODUCTION
}
impl Database {
#[doc(hidden)]
pub async fn execute_write<O, C, S>(
&self,
execution: WriteExecution<O>,
cancel: C,
shutdown: S,
) -> Result<ManagedWriteResult, WriteExecutionError>
where
O: StaticWriteOperation,
C: Future + Unpin + Send + 'static,
S: Future + Unpin + Send + 'static,
{
let pool = self.pool.clone();
drive_db_finalizer_with_transition(cancel, shutdown, move |transition| {
execute_write_with_transition(pool, execution, transition)
})
.await
.map_err(|_| WriteExecutionError::FinalizerFailed)
.and_then(|result| result)
}
#[doc(hidden)]
pub async fn execute_transaction<O, C, S>(
&self,
execution: TransactionExecution<O>,
cancel: C,
shutdown: S,
) -> Result<ManagedWriteResult, TransactionExecutionError>
where
O: StaticWriteOperation,
C: Future + Unpin + Send + 'static,
S: Future + Unpin + Send + 'static,
{
let pool = self.pool.clone();
drive_db_finalizer_with_transition(cancel, shutdown, move |transition| {
execute_transaction_with_transition(pool, execution, transition)
})
.await
.map_err(|_| TransactionExecutionError::FinalizerFailed)
.and_then(|result| result)
}
}
async fn execute_write_with_transition<O, C, S>(
pool: sqlx::MySqlPool,
execution: WriteExecution<O>,
mut transition: DbQueryTransition<C, S>,
) -> Result<
saddle_runtime::db_finalizer::DbFinalizingOutput<
Result<ManagedWriteResult, WriteExecutionError>,
impl Future<Output = ()> + Send + 'static,
>,
saddle_admission::AdmissionError,
>
where
O: StaticWriteOperation,
C: Future + Unpin + Send + 'static,
S: Future + Unpin + Send + 'static,
{
let (permit, invocation) = execution.into_parts();
let mut connection = pool.try_acquire();
let (value, physical) = if let Some(connection) = connection.as_mut() {
let query = async {
invocation
.parameters
.bind(sqlx::query::<MySql>(O::SQL))
.execute(&mut **connection)
.await
};
let mut query = pin!(query);
let outcome = poll_operation(&mut transition, &permit, query.as_mut()).await;
match outcome {
Ok(DbQueryPoll::Ready(Ok(result))) => (
Ok(ManagedWriteResult {
rows_affected: result.rows_affected(),
last_insert_id: result.last_insert_id(),
}),
PhysicalFinalization::ReturnToPool,
),
Ok(DbQueryPoll::Ready(Err(_))) | Err(()) => (
Err(WriteExecutionError::WriteFailed),
PhysicalFinalization::PoisonDiscard,
),
Ok(DbQueryPoll::Transition(DbTransitionRequest::Cancel)) => (
Err(WriteExecutionError::Cancelled),
PhysicalFinalization::PoisonDiscard,
),
Ok(DbQueryPoll::Transition(DbTransitionRequest::Shutdown)) => (
Err(WriteExecutionError::Shutdown),
PhysicalFinalization::PoisonDiscard,
),
}
} else {
(
Err(WriteExecutionError::ConnectionUnavailable),
PhysicalFinalization::ReturnToPool,
)
};
let finalizer = physical_finalizer(connection, physical);
transition.begin_finalizing(value, permit, finalizer)
}
async fn execute_transaction_with_transition<O, C, S>(
pool: sqlx::MySqlPool,
execution: TransactionExecution<O>,
mut transition: DbQueryTransition<C, S>,
) -> Result<
saddle_runtime::db_finalizer::DbFinalizingOutput<
Result<ManagedWriteResult, TransactionExecutionError>,
impl Future<Output = ()> + Send + 'static,
>,
saddle_admission::AdmissionError,
>
where
O: StaticWriteOperation,
C: Future + Unpin + Send + 'static,
S: Future + Unpin + Send + 'static,
{
let (permit, invocation) = execution.into_parts();
let mut connection = pool.try_acquire();
let (value, physical) = if let Some(connection) = connection.as_mut() {
let transaction = async {
sqlx::query("BEGIN")
.execute(&mut **connection)
.await
.map_err(|_| TransactionExecutionError::BeginFailed)?;
let write = invocation
.write
.parameters
.bind(sqlx::query::<MySql>(O::SQL))
.execute(&mut **connection)
.await
.map_err(|_| TransactionExecutionError::WriteFailed)?;
let result = ManagedWriteResult {
rows_affected: write.rows_affected(),
last_insert_id: write.last_insert_id(),
};
match invocation.decision {
TransactionDecision::Commit => {
sqlx::query("COMMIT")
.execute(&mut **connection)
.await
.map_err(|_| TransactionExecutionError::CommitFailed)?;
Ok(result)
}
TransactionDecision::Rollback => {
sqlx::query("ROLLBACK")
.execute(&mut **connection)
.await
.map_err(|_| TransactionExecutionError::RollbackFailed)?;
Err(TransactionExecutionError::BusinessRollback)
}
}
};
let mut transaction = pin!(transaction);
let outcome = poll_operation(&mut transition, &permit, transaction.as_mut()).await;
match outcome {
Ok(DbQueryPoll::Ready(Ok(result))) => (Ok(result), PhysicalFinalization::ReturnToPool),
Ok(DbQueryPoll::Ready(Err(TransactionExecutionError::BusinessRollback))) => (
Err(TransactionExecutionError::BusinessRollback),
PhysicalFinalization::ReturnToPool,
),
Ok(DbQueryPoll::Ready(Err(error))) => (Err(error), PhysicalFinalization::PoisonDiscard),
Ok(DbQueryPoll::Transition(DbTransitionRequest::Cancel)) => (
Err(TransactionExecutionError::Cancelled),
PhysicalFinalization::PoisonDiscard,
),
Ok(DbQueryPoll::Transition(DbTransitionRequest::Shutdown)) => (
Err(TransactionExecutionError::Shutdown),
PhysicalFinalization::PoisonDiscard,
),
Err(()) => (
Err(TransactionExecutionError::WriteFailed),
PhysicalFinalization::PoisonDiscard,
),
}
} else {
(
Err(TransactionExecutionError::ConnectionUnavailable),
PhysicalFinalization::ReturnToPool,
)
};
let finalizer = physical_finalizer(connection, physical);
transition.begin_finalizing(value, permit, finalizer)
}
async fn poll_operation<C, S, Q>(
transition: &mut DbQueryTransition<C, S>,
permit: &DbRequestPermit,
mut query: std::pin::Pin<&mut Q>,
) -> Result<DbQueryPoll<Q::Output>, ()>
where
C: Future + Unpin,
S: Future + Unpin,
Q: Future,
{
std::future::poll_fn(|context| {
match std::panic::catch_unwind(AssertUnwindSafe(|| {
transition.poll_query(permit, query.as_mut(), context)
})) {
Ok(Poll::Ready(output)) => Poll::Ready(Ok(output)),
Ok(Poll::Pending) => Poll::Pending,
Err(_) => Poll::Ready(Err(())),
}
})
.await
}
async fn physical_finalizer(
mut connection: Option<sqlx::pool::PoolConnection<MySql>>,
physical: PhysicalFinalization,
) {
if let Some(mut connection) = connection.take() {
match physical {
PhysicalFinalization::ReturnToPool => connection.return_to_pool().await,
PhysicalFinalization::PoisonDiscard => drop(connection.detach()),
}
}
}
#[derive(Clone, Copy)]
enum PhysicalFinalization {
ReturnToPool,
PoisonDiscard,
}
fn validate_operation(operation: &str) -> Result<(), QueryOptionalContractError> {
if operation.is_empty()
|| operation.len() > 128
|| !operation
.bytes()
.all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'.' | b'_' | b'-'))
{
return Err(QueryOptionalContractError::InvalidOperation);
}
Ok(())
}
fn validate_sql(sql: &str) -> Result<(), QueryOptionalContractError> {
if sql.trim().is_empty() || sql.len() > 65_536 {
return Err(QueryOptionalContractError::InvalidSql);
}
Ok(())
}
#[cfg(test)]
mod tests {
use std::{
env,
future::Future,
io,
io::{Read, Write},
net::{Shutdown, TcpListener, TcpStream},
pin::Pin,
sync::{Arc, Mutex},
task::{Context, Poll},
thread,
};
use saddle_admission::{
AdmissionError, DbCreditProfile, DbRouteCreditDemand, DbRouteResources, EntryIoAuditPlan,
EntryReadPoll, ManagedBytes, ManagedResponse, OfficialTokioEntryIoAttemptOutcome,
OfficialTokioRegistrationProfile, ProcessLedger, RequestMemory, ResourceConfig,
ResponseWritePoll,
};
use saddle_core::ComponentLifecycle;
use saddle_observability::{Observer, ObserverConfig};
use sqlx::Connection;
use super::*;
use crate::{
DatabaseConfig,
c6_query_optional::{DbPair, DbU64, ManagedDbField, ManagedQueryParameters, sealed},
};
struct Upsert;
impl StaticWriteOperation for Upsert {
type Parameters = ManagedQueryParameters<DbPair<DbU64, DbU64>>;
const OPERATION: &'static str = "orders.upsert";
const SQL: &'static str = "INSERT INTO saddle_c1_write (id, value_number) VALUES (?, ?) ON DUPLICATE KEY UPDATE value_number = VALUES(value_number)";
}
#[derive(Clone, Copy)]
struct PanicField(u8);
impl sealed::Field for PanicField {
fn bind<'q>(
&'q self,
_: sqlx::query::Query<'q, MySql, sqlx::mysql::MySqlArguments>,
) -> sqlx::query::Query<'q, MySql, sqlx::mysql::MySqlArguments> {
let _ = self.0;
panic!("generated bind panic")
}
fn decode(_: &sqlx::mysql::MySqlRow, _: &mut usize) -> std::result::Result<Self, ()> {
unreachable!()
}
}
impl ManagedDbField for PanicField {}
struct PanicWrite;
impl StaticWriteOperation for PanicWrite {
type Parameters = ManagedQueryParameters<PanicField>;
const OPERATION: &'static str = "orders.panic";
const SQL: &'static str = "SELECT ?";
}
struct LockWrite;
impl StaticWriteOperation for LockWrite {
type Parameters = ManagedQueryParameters<DbPair<DbU64, DbU64>>;
const OPERATION: &'static str = "orders.lock";
const SQL: &'static str = "UPDATE saddle_c1_write SET value_number = ? WHERE id = ?";
}
struct MissingWrite;
impl StaticWriteOperation for MissingWrite {
type Parameters = ManagedQueryParameters<DbU64>;
const OPERATION: &'static str = "orders.missing";
const SQL: &'static str = "UPDATE saddle_c1_missing SET value_number = ? WHERE id = 1";
}
struct EntryConnection;
struct ReadyRead;
struct ReadyWrite;
impl EntryReadPoll<EntryConnection> for ReadyRead {
fn poll_read(
&mut self,
_: &mut EntryConnection,
memory: &RequestMemory,
_: &mut Context<'_>,
) -> Poll<Result<ManagedBytes, AdmissionError>> {
Poll::Ready(memory.try_bytes(&[]))
}
}
impl ResponseWritePoll<EntryConnection> for ReadyWrite {
fn poll_write(
&mut self,
_: &mut EntryConnection,
_: &ManagedResponse,
_: &mut Context<'_>,
) -> Poll<Result<(), AdmissionError>> {
Poll::Ready(Ok(()))
}
}
struct FixedSignal {
pending_polls: Option<u8>,
}
impl FixedSignal {
const fn pending() -> Self {
Self {
pending_polls: None,
}
}
const fn after_one_poll() -> Self {
Self {
pending_polls: Some(1),
}
}
}
impl Future for FixedSignal {
type Output = ();
fn poll(mut self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll<()> {
match self.pending_polls {
None => Poll::Pending,
Some(0) => Poll::Ready(()),
Some(remaining) => {
self.pending_polls = Some(remaining - 1);
context.waker().wake_by_ref();
Poll::Pending
}
}
}
}
fn resource_config(
registration: OfficialTokioRegistrationProfile,
task_reserve: usize,
) -> ResourceConfig {
let process_state_reserve = ProcessLedger::minimum_process_state_reserve_with_waiters(1, 1)
.unwrap()
+ ProcessLedger::official_tokio_state_reserve(registration).unwrap();
ResourceConfig {
managed_capacity: 4096,
entry_reserve: 64,
framework_reserve: 1024 * 1024,
task_reserve,
process_state_reserve,
system_estimate: 1024 * 1024,
safety_margin: 1024 * 1024,
process_limit: 4096
+ 64
+ 1024 * 1024
+ task_reserve
+ process_state_reserve
+ 2 * 1024 * 1024,
max_active_requests: 1,
}
}
fn entry_plan() -> EntryIoAuditPlan {
EntryIoAuditPlan::locked_linux_x86_64_tokio_1_53_1(64, 64, &[]).unwrap()
}
async fn admitted_permit<T, F, B>(
task_reserve: usize,
run: F,
) -> (T, saddle_admission::DbCreditSnapshot)
where
T: Send + 'static,
F: FnOnce(DbRequestPermit) -> B + Send + Unpin + 'static,
B: Future<Output = T> + Send + 'static,
{
let registration = OfficialTokioRegistrationProfile {
listener: 1,
transport_connections: 1,
runtime_fixed: 1,
};
let ledger =
ProcessLedger::new_with_waiters(resource_config(registration, task_reserve), 1)
.unwrap();
let runtime_domain = ledger.prepare_official_tokio_domain(registration).unwrap();
let allocation = ledger
.prepare_process_allocation_profile(usize::MAX)
.unwrap();
let db_domain = ledger
.prepare_db_domain(DbCreditProfile {
connections: 1,
operations: 1,
})
.unwrap();
let demand = DbRouteCreditDemand::new(1, 1).unwrap();
let result = Arc::new(Mutex::new(None));
let result_for_task = result.clone();
let outcome = ledger.attempt_official_tokio_entry_io(
&runtime_domain,
DbRouteResources::required(&db_domain, demand),
4096,
task_reserve,
entry_plan(),
entry_plan(),
|_| (EntryConnection, ReadyRead, ReadyWrite),
move |_, permit, memory| {
let operation = run(permit.unwrap());
let response = memory.try_response(&[]).unwrap();
async move {
*result_for_task.lock().unwrap() = Some(operation.await);
response
}
},
);
let ready = match outcome {
OfficialTokioEntryIoAttemptOutcome::Ready(ready) => ready,
_ => panic!("fixed DB resources must admit"),
};
let (envelope, task_slot) = ready.into_runtime_parts();
let envelope_size = std::mem::size_of_val(&envelope);
let envelope_align = std::mem::align_of_val(&envelope);
assert!(envelope_size <= task_reserve);
assert!(envelope_align.is_power_of_two());
eprintln!(
"C1 envelope_size={envelope_size} envelope_align={envelope_align} tested_reserve={task_reserve}"
);
tokio::spawn(envelope).await.unwrap().unwrap();
drop(task_slot);
let snapshot = db_domain.snapshot().unwrap();
drop(db_domain);
drop(runtime_domain);
assert!(!allocation.finish().unwrap().breached);
assert_eq!(ledger.try_shutdown().unwrap().active_accounts, 0);
let result = Arc::try_unwrap(result)
.ok()
.unwrap()
.into_inner()
.unwrap()
.unwrap();
(result, snapshot)
}
async fn run_write<O, C, S>(
database: Database,
invocation: WriteInvocation<O>,
cancel: C,
shutdown: S,
) -> Result<ManagedWriteResult, WriteExecutionError>
where
O: StaticWriteOperation,
O::Parameters: Unpin,
C: Future<Output = ()> + Unpin + Send + 'static,
S: Future<Output = ()> + Unpin + Send + 'static,
{
let (result, snapshot) = admitted_permit(1024 * 1024, move |permit| {
let execution = WriteExecution::from_compiled_handoff(permit, invocation);
async move { database.execute_write(execution, cancel, shutdown).await }
})
.await;
assert_eq!(snapshot.connections_in_use, 0);
assert_eq!(snapshot.operations_in_use, 0);
result
}
async fn run_transaction<O, C, S>(
database: Database,
invocation: TransactionInvocation<O>,
cancel: C,
shutdown: S,
) -> Result<ManagedWriteResult, TransactionExecutionError>
where
O: StaticWriteOperation,
O::Parameters: Unpin,
C: Future<Output = ()> + Unpin + Send + 'static,
S: Future<Output = ()> + Unpin + Send + 'static,
{
let (result, snapshot) = admitted_permit(1024 * 1024, move |permit| {
let execution = TransactionExecution::from_compiled_handoff(permit, invocation);
async move {
database
.execute_transaction(execution, cancel, shutdown)
.await
}
})
.await;
assert_eq!(snapshot.connections_in_use, 0);
assert_eq!(snapshot.operations_in_use, 0);
result
}
async fn database(url: &str) -> Database {
Database::connect(
DatabaseConfig::new(url).max_connections(1),
Observer::with_writer(ObserverConfig::default(), io::sink()).unwrap(),
)
.await
.unwrap()
}
fn parameters(id: u64, value: u64) -> ManagedQueryParameters<DbPair<DbU64, DbU64>> {
ManagedQueryParameters(DbPair(DbU64(id), DbU64(value)))
}
fn commit_response_cut_proxy(database_url: &str) -> (String, thread::JoinHandle<()>) {
let backend = database_url
.split_once('@')
.and_then(|(_, suffix)| suffix.split_once('/'))
.map(|(authority, _)| authority)
.expect("test database URL has an authority");
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let proxy = listener.local_addr().unwrap();
let proxy_url = database_url.replacen(backend, &proxy.to_string(), 1);
let backend = backend.to_owned();
let worker = thread::spawn(move || {
loop {
let (client, _) = listener.accept().unwrap();
let server = TcpStream::connect(&backend).unwrap();
let mut request_client = client.try_clone().unwrap();
let mut request_server = server.try_clone().unwrap();
let response_client = client.try_clone().unwrap();
let response_server = server.try_clone().unwrap();
let response = thread::spawn(move || {
let _ = io::copy(&mut &response_server, &mut &response_client);
});
let mut cut_commit_response = false;
let mut buffer = [0_u8; 4096];
loop {
let count = request_client.read(&mut buffer).unwrap();
if count == 0 {
break;
}
request_server.write_all(&buffer[..count]).unwrap();
if buffer[..count]
.windows(b"COMMIT".len())
.any(|window| window.eq_ignore_ascii_case(b"COMMIT"))
{
cut_commit_response = true;
let _ = request_client.shutdown(Shutdown::Both);
let _ = request_server.shutdown(Shutdown::Both);
break;
}
}
let _ = response.join();
if cut_commit_response {
break;
}
}
});
(proxy_url, worker)
}
#[test]
fn fixed_contract_has_one_credit_and_no_nested_transaction() {
let write = WriteOperationProof::<Upsert>::bind().unwrap();
assert_eq!(write.layout().credits(), WriteCreditDemand::ONE);
let transaction = TransactionOperationProof::<Upsert>::bind().unwrap();
assert_eq!(transaction.layout().credits(), WriteCreditDemand::ONE);
let finalization = write_transaction_finalization_proof();
assert!(finalization.single_connection());
assert!(!finalization.nested_transactions());
assert_eq!(finalization.cancel_external_io_awaits(), 0);
assert!(finalization.releases_pool_size_before_permit());
}
#[tokio::test]
async fn real_mariadb_write_commit_rollback_cancel_shutdown_and_panic_close() {
let Ok(url) = env::var("SADDLE_TEST_DATABASE_URL") else {
eprintln!("skipping C1 write/transaction: SADDLE_TEST_DATABASE_URL is not set");
return;
};
let setup = database(&url).await;
sqlx::query(
"CREATE TABLE IF NOT EXISTS saddle_c1_write (id BIGINT UNSIGNED PRIMARY KEY, value_number BIGINT UNSIGNED NOT NULL)",
)
.execute(&setup.pool)
.await
.unwrap();
sqlx::query("DELETE FROM saddle_c1_write")
.execute(&setup.pool)
.await
.unwrap();
setup.shutdown().await.unwrap();
let write_db = database(&url).await;
let result = run_write(
write_db.clone(),
WriteOperationProof::<Upsert>::bind()
.unwrap()
.invocation(parameters(1, 10)),
FixedSignal::pending(),
FixedSignal::pending(),
)
.await
.unwrap();
assert_eq!(result.rows_affected(), 1);
assert_eq!((write_db.pool.size(), write_db.pool.num_idle()), (1, 1));
write_db.shutdown().await.unwrap();
let error_db = database(&url).await;
let error = run_write(
error_db.clone(),
WriteOperationProof::<MissingWrite>::bind()
.unwrap()
.invocation(ManagedQueryParameters(DbU64(10))),
FixedSignal::pending(),
FixedSignal::pending(),
)
.await
.unwrap_err();
assert_eq!(error, WriteExecutionError::WriteFailed);
assert_eq!((error_db.pool.size(), error_db.pool.num_idle()), (0, 0));
error_db.shutdown().await.unwrap();
let commit_db = database(&url).await;
run_transaction(
commit_db.clone(),
TransactionOperationProof::<Upsert>::bind()
.unwrap()
.invocation(parameters(2, 20), TransactionDecision::Commit),
FixedSignal::pending(),
FixedSignal::pending(),
)
.await
.unwrap();
assert_eq!((commit_db.pool.size(), commit_db.pool.num_idle()), (1, 1));
commit_db.shutdown().await.unwrap();
let (ambiguous_url, proxy) = commit_response_cut_proxy(&url);
let ambiguous_db = database(&ambiguous_url).await;
let ambiguous = run_transaction(
ambiguous_db.clone(),
TransactionOperationProof::<Upsert>::bind()
.unwrap()
.invocation(parameters(4, 40), TransactionDecision::Commit),
FixedSignal::pending(),
FixedSignal::pending(),
)
.await
.unwrap_err();
assert_eq!(ambiguous, TransactionExecutionError::CommitFailed);
assert_eq!(
(ambiguous_db.pool.size(), ambiguous_db.pool.num_idle()),
(0, 0)
);
ambiguous_db.shutdown().await.unwrap();
proxy.join().unwrap();
let rollback_db = database(&url).await;
let rollback = run_transaction(
rollback_db.clone(),
TransactionOperationProof::<Upsert>::bind()
.unwrap()
.invocation(parameters(3, 30), TransactionDecision::Rollback),
FixedSignal::pending(),
FixedSignal::pending(),
)
.await
.unwrap_err();
assert_eq!(rollback, TransactionExecutionError::BusinessRollback);
assert_eq!(
(rollback_db.pool.size(), rollback_db.pool.num_idle()),
(1, 1)
);
rollback_db.shutdown().await.unwrap();
let panic_db = database(&url).await;
let panic = run_write(
panic_db.clone(),
WriteOperationProof::<PanicWrite>::bind()
.unwrap()
.invocation(ManagedQueryParameters(PanicField(0))),
FixedSignal::pending(),
FixedSignal::pending(),
)
.await
.unwrap_err();
assert_eq!(panic, WriteExecutionError::WriteFailed);
assert_eq!((panic_db.pool.size(), panic_db.pool.num_idle()), (0, 0));
panic_db.shutdown().await.unwrap();
for (cancel, shutdown, expected) in [
(
FixedSignal::after_one_poll(),
FixedSignal::pending(),
TransactionExecutionError::Cancelled,
),
(
FixedSignal::pending(),
FixedSignal::after_one_poll(),
TransactionExecutionError::Shutdown,
),
] {
let blocked = database(&url).await;
let mut blocker = sqlx::mysql::MySqlConnection::connect(&url).await.unwrap();
sqlx::query("BEGIN").execute(&mut blocker).await.unwrap();
sqlx::query("UPDATE saddle_c1_write SET value_number = 99 WHERE id = 1")
.execute(&mut blocker)
.await
.unwrap();
let error = run_transaction(
blocked.clone(),
TransactionOperationProof::<LockWrite>::bind()
.unwrap()
.invocation(parameters(100, 1), TransactionDecision::Commit),
cancel,
shutdown,
)
.await
.unwrap_err();
assert_eq!(error, expected);
assert_eq!((blocked.pool.size(), blocked.pool.num_idle()), (0, 0));
sqlx::query("ROLLBACK").execute(&mut blocker).await.unwrap();
blocked.shutdown().await.unwrap();
}
let verify = database(&url).await;
let rows: Vec<(u64, u64)> =
sqlx::query_as("SELECT id, value_number FROM saddle_c1_write ORDER BY id")
.fetch_all(&verify.pool)
.await
.unwrap();
assert_eq!(rows, vec![(1, 10), (2, 20)]);
verify.shutdown().await.unwrap();
}
}