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(())
}