use std::{
env, io,
sync::{Arc, Mutex},
time::Duration,
};
use saddle_core::{
ApplicationId, CallContext, ComponentLifecycle, ErrorKind, ModuleId, OperationId, SaddleError,
ServiceId, SpanId, TraceId,
};
use saddle_db::{Database, DatabaseConfig, DbValue, Statement};
use saddle_observability::{Observer, ObserverConfig};
use serde_json::Value;
use sqlx::{Connection, mysql::MySqlConnection};
#[derive(Clone, Default)]
struct Capture(Arc<Mutex<Vec<u8>>>);
impl io::Write for Capture {
fn write(&mut self, bytes: &[u8]) -> io::Result<usize> {
self.0.lock().unwrap().extend_from_slice(bytes);
Ok(bytes.len())
}
fn flush(&mut self) -> io::Result<()> {
Ok(())
}
}
fn context() -> CallContext {
CallContext::new(
ApplicationId::from("test"),
ModuleId::from("orders"),
ServiceId::from("orders"),
OperationId::from("create"),
TraceId::from_u128(0x1234),
SpanId::from_u64(0x5678),
)
}
fn statement(operation: impl AsRef<str>, sql: impl AsRef<str>) -> Statement {
Statement::new(operation, sql).unwrap()
}
trait TestStatementExt {
fn test_bind(self, value: impl Into<DbValue>) -> Self;
}
impl TestStatementExt for Statement {
fn test_bind(self, value: impl Into<DbValue>) -> Self {
self.bind(value).unwrap()
}
}
#[tokio::test]
async fn rejects_server_packet_limit_above_driver_allocation_contract() {
let Ok(url) = env::var("SADDLE_TEST_OVERSIZED_PACKET_DATABASE_URL") else {
eprintln!(
"skipping oversized packet configuration test: SADDLE_TEST_OVERSIZED_PACKET_DATABASE_URL is not set"
);
return;
};
let observer = Observer::with_writer(ObserverConfig::default(), io::sink()).unwrap();
let error = match Database::connect(DatabaseConfig::new(url), observer).await {
Ok(_) => panic!("server above the inbound packet contract was accepted"),
Err(error) => error,
};
assert_eq!(error.code(), "db.invalid_config");
}
#[tokio::test]
async fn mysql_pool_query_write_commit_rollback_and_trace() {
let Ok(url) = env::var("SADDLE_TEST_DATABASE_URL") else {
eprintln!("skipping real MySQL test: SADDLE_TEST_DATABASE_URL is not set");
return;
};
let capture = Capture::default();
let observer = Observer::with_writer(ObserverConfig::default(), capture.clone()).unwrap();
let database = Database::connect(
DatabaseConfig::new(url.clone()).max_connections(1),
observer.clone(),
)
.await
.unwrap();
database.write(&context(), statement("test.create", "CREATE TABLE IF NOT EXISTS saddle_v1_db_test (id BIGINT UNSIGNED PRIMARY KEY, value_text VARCHAR(64) NOT NULL)" )).await.unwrap();
database
.write(
&context(),
statement("test.clear", "DELETE FROM saddle_v1_db_test"),
)
.await
.unwrap();
database
.transaction(&context(), "orders.commit", |transaction| {
Box::pin(async move {
transaction
.write(
statement(
"orders.insert",
"INSERT INTO saddle_v1_db_test(id, value_text) VALUES (?, ?)",
)
.test_bind(1_u64)
.test_bind("committed-secret"),
)
.await?;
let row = transaction
.query_optional(
statement(
"orders.find",
"SELECT id, value_text FROM saddle_v1_db_test WHERE id = ?",
)
.test_bind(1_u64),
)
.await?
.expect("inserted row");
assert_eq!(row.u64("id")?, 1);
assert_eq!(row.string("value_text")?, "committed-secret");
Ok(())
})
})
.await
.unwrap();
let business_error = database
.transaction(&context(), "orders.rollback", |transaction| {
Box::pin(async move {
transaction
.write(
statement(
"orders.insert",
"INSERT INTO saddle_v1_db_test(id, value_text) VALUES (?, ?)",
)
.test_bind(2_u64)
.test_bind("rolled-back-secret"),
)
.await?;
Err::<(), _>(SaddleError::new(
ErrorKind::Business,
"order.rejected",
"order rejected",
))
})
})
.await
.unwrap_err();
assert_eq!(business_error.code(), "order.rejected");
let committed = database
.query_optional(
&context(),
statement(
"orders.find",
"SELECT id FROM saddle_v1_db_test WHERE id = ?",
)
.test_bind(1_u64),
)
.await
.unwrap();
let rolled_back = database
.query_optional(
&context(),
statement(
"orders.find",
"SELECT id FROM saddle_v1_db_test WHERE id = ?",
)
.test_bind(2_u64),
)
.await
.unwrap();
assert!(committed.is_some());
assert!(rolled_back.is_none());
let (inserted, inserted_rx) = tokio::sync::oneshot::channel();
let cancelled_database = database.clone();
let cancelled_context = context();
let cancelled = tokio::spawn(async move {
cancelled_database
.transaction(&cancelled_context, "orders.cancel", |transaction| {
Box::pin(async move {
transaction
.write(
statement(
"orders.insert",
"INSERT INTO saddle_v1_db_test(id, value_text) VALUES (?, ?)",
)
.test_bind(3_u64)
.test_bind("cancelled-secret"),
)
.await?;
let _ = inserted.send(());
std::future::pending::<saddle_core::Result<()>>().await
})
})
.await
});
inserted_rx.await.unwrap();
cancelled.abort();
assert!(cancelled.await.unwrap_err().is_cancelled());
let cancelled_row = database
.query_optional(
&context(),
statement(
"orders.find",
"SELECT id FROM saddle_v1_db_test WHERE id = ?",
)
.test_bind(3_u64),
)
.await
.unwrap();
assert!(cancelled_row.is_none());
let oversized_parameter = statement(
"orders.insert",
"INSERT INTO saddle_v1_db_test(id, value_text) VALUES (4, ?)",
)
.bind(vec![0_u8; saddle_db::MAX_PARAMETER_BYTES + 1])
.unwrap_err();
assert_eq!(oversized_parameter.code(), "db.invalid_statement");
let oversized_field = database
.query_optional(
&context(),
statement(
"test.oversized_field",
format!(
"SELECT REPEAT('x', {}) AS payload",
saddle_db::MAX_FIELD_BYTES + 1
),
),
)
.await
.unwrap_err();
assert_eq!(oversized_field.code(), "db.result_limit_exceeded");
let oversized_transaction_field = database
.transaction(&context(), "orders.oversized", |transaction| {
Box::pin(async move {
transaction
.query_optional(statement(
"test.oversized_field",
format!(
"SELECT REPEAT('x', {}) AS payload",
saddle_db::MAX_FIELD_BYTES + 1
),
))
.await?;
Ok(())
})
})
.await
.unwrap_err();
assert_eq!(
oversized_transaction_field.code(),
"db.result_limit_exceeded"
);
let oversized_result = database
.query_all(
&context(),
statement(
"test.oversized_result",
format!(
"WITH RECURSIVE seq(n) AS (SELECT 1 UNION ALL SELECT n + 1 FROM seq WHERE n < 9) SELECT REPEAT('x', {}) AS payload FROM seq",
saddle_db::MAX_FIELD_BYTES
),
),
)
.await
.unwrap_err();
assert_eq!(oversized_result.code(), "db.result_limit_exceeded");
let (shutdown_inserted, shutdown_inserted_rx) = tokio::sync::oneshot::channel();
let shutdown_database = database.clone();
let shutdown_context = context();
let shutdown_cancelled = tokio::spawn(async move {
shutdown_database
.transaction(&shutdown_context, "orders.shutdown_cancel", |transaction| {
Box::pin(async move {
transaction
.write(
statement(
"orders.insert",
"INSERT INTO saddle_v1_db_test(id, value_text) VALUES (?, ?)",
)
.test_bind(5_u64)
.test_bind("shutdown-cancelled-secret"),
)
.await?;
let _ = shutdown_inserted.send(());
std::future::pending::<saddle_core::Result<()>>().await
})
})
.await
});
shutdown_inserted_rx.await.unwrap();
shutdown_cancelled.abort();
database.shutdown().await.unwrap();
assert!(shutdown_cancelled.await.unwrap_err().is_cancelled());
let verification = Database::connect(
DatabaseConfig::new(url.clone()).max_connections(1),
observer.clone(),
)
.await
.unwrap();
let shutdown_cancelled_row = verification
.query_optional(
&context(),
statement(
"orders.find",
"SELECT id FROM saddle_v1_db_test WHERE id = ?",
)
.test_bind(5_u64),
)
.await
.unwrap();
assert!(shutdown_cancelled_row.is_none());
verification.shutdown().await.unwrap();
let reconnecting = Database::connect(
DatabaseConfig::new(url.clone())
.max_connections(1)
.acquire_timeout(Duration::from_millis(500)),
observer.clone(),
)
.await
.unwrap();
let connection_id = reconnecting
.query_optional(
&context(),
statement("test.connection_id", "SELECT CONNECTION_ID() AS id"),
)
.await
.unwrap()
.unwrap()
.u64("id")
.unwrap();
let mut administrator = MySqlConnection::connect(&url).await.unwrap();
sqlx::query("SET GLOBAL max_allowed_packet = 16777216")
.execute(&mut administrator)
.await
.unwrap();
sqlx::query("KILL ?")
.bind(connection_id)
.execute(&mut administrator)
.await
.unwrap();
let reconnect_error = reconnecting
.query_optional(&context(), statement("test.reconnect", "SELECT 1 AS value"))
.await
.unwrap_err();
assert_eq!(reconnect_error.code(), "db.connection_unavailable");
sqlx::query("SET GLOBAL max_allowed_packet = 8388608")
.execute(&mut administrator)
.await
.unwrap();
administrator.close().await.unwrap();
reconnecting.shutdown().await.unwrap();
observer.flush().await.unwrap();
let output = String::from_utf8(capture.0.lock().unwrap().clone()).unwrap();
assert!(!output.contains("committed-secret"));
assert!(!output.contains("rolled-back-secret"));
assert!(!output.contains("cancelled-secret"));
assert!(!output.contains("shutdown-cancelled-secret"));
assert!(!output.contains("INSERT INTO"));
let records: Vec<Value> = output
.lines()
.map(|line| serde_json::from_str(line).unwrap())
.collect();
assert!(
records
.iter()
.all(|record| record["trace_id"] == context().trace_id().to_string())
);
for operation in [
"begin",
"commit",
"rollback",
"orders.insert",
"orders.find",
"orders.cancel",
"orders.shutdown_cancel",
] {
assert!(
records
.iter()
.any(|record| record["operation"] == operation),
"missing operation {operation}"
);
}
assert!(records.iter().any(|record| {
record["operation"] == "orders.cancel" && record["error_code"] == "db.transaction_cancelled"
}));
}