saddle-db 0.1.1

Saddle managed asynchronous database access and transactions
Documentation
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"
    }));
}