sqlite-functions 0.1.2

A small toolkit for authoring Rust functions for SQLite
Documentation
use sqlite_functions::rusqlite::functions::{Aggregate, Context, FunctionFlags};
use sqlite_functions::rusqlite::{Connection, Error, Result};
use sqlite_functions::{Arity, FunctionOptions, register_aggregate, register_scalar};

#[test]
fn options_map_arity_and_flags() {
    let defaults = FunctionOptions::new("example", Arity::Exact(2));
    assert_eq!(defaults.name(), "example");
    assert_eq!(defaults.arity(), Arity::Exact(2));
    assert_eq!(defaults.flags().bits(), FunctionFlags::SQLITE_UTF8.bits());

    let configured = FunctionOptions::new("many", Arity::Variadic)
        .deterministic()
        .innocuous()
        .direct_only();
    assert_eq!(configured.arity(), Arity::Variadic);
    assert!(
        configured
            .flags()
            .contains(FunctionFlags::SQLITE_DETERMINISTIC)
    );
    assert!(configured.flags().contains(FunctionFlags::SQLITE_INNOCUOUS));
    assert!(
        configured
            .flags()
            .contains(FunctionFlags::SQLITE_DIRECTONLY)
    );
}

#[test]
fn scalar_registration_propagates_null_and_errors() -> Result<()> {
    let connection = Connection::open_in_memory()?;
    register_scalar(
        &connection,
        FunctionOptions::new("double_text", Arity::Exact(1)).deterministic(),
        |ctx| {
            let value = ctx.get::<Option<String>>(0)?;
            Ok(value.map(|value| format!("{value}{value}")))
        },
    )?;

    let doubled: String = connection.query_row("SELECT double_text('ab')", [], |row| row.get(0))?;
    assert_eq!(doubled, "abab");

    let null: Option<String> =
        connection.query_row("SELECT double_text(NULL)", [], |row| row.get(0))?;
    assert_eq!(null, None);

    let error = connection
        .query_row::<String, _, _>("SELECT double_text(42)", [], |row| row.get(0))
        .expect_err("integer input should be rejected");
    assert!(matches!(error, Error::SqliteFailure(_, _)));
    Ok(())
}

struct SumSquares;

impl Aggregate<i64, Option<i64>> for SumSquares {
    fn init(&self, _ctx: &mut Context<'_>) -> Result<i64> {
        Ok(0)
    }

    fn step(&self, ctx: &mut Context<'_>, total: &mut i64) -> Result<()> {
        let value = ctx.get::<i64>(0)?;
        *total += value * value;
        Ok(())
    }

    fn finalize(&self, _ctx: &mut Context<'_>, total: Option<i64>) -> Result<Option<i64>> {
        Ok(total)
    }
}

#[test]
fn aggregate_registration_keeps_state() -> Result<()> {
    let connection = Connection::open_in_memory()?;
    register_aggregate(
        &connection,
        FunctionOptions::new("sum_squares", Arity::Exact(1)).deterministic(),
        SumSquares,
    )?;
    let result: i64 = connection.query_row(
        "SELECT sum_squares(value) FROM (SELECT 2 AS value UNION ALL SELECT 3)",
        [],
        |row| row.get(0),
    )?;
    assert_eq!(result, 13);
    Ok(())
}

#[test]
fn deterministic_innocuous_function_can_back_an_index() -> Result<()> {
    let connection = Connection::open_in_memory()?;
    register_scalar(
        &connection,
        FunctionOptions::new("text_length", Arity::Exact(1))
            .deterministic()
            .innocuous(),
        |ctx| {
            let value = ctx.get::<String>(0)?;
            Ok(i64::from(!value.is_empty()))
        },
    )?;
    connection.execute_batch(
        "CREATE TABLE values_(value TEXT NOT NULL);
         CREATE INDEX values_length ON values_(text_length(value));",
    )?;
    Ok(())
}