diff --git a/crates/libs/lib-lualib/src/lua_sqlx.rs b/crates/libs/lib-lualib/src/lua_sqlx.rs index 83e6551..36c86ba 100644 --- a/crates/libs/lib-lualib/src/lua_sqlx.rs +++ b/crates/libs/lib-lualib/src/lua_sqlx.rs @@ -3,13 +3,13 @@ use std::time::Duration; use dashmap::DashMap; use lazy_static::lazy_static; -use sqlx::types::Uuid; +use sqlx::types::{Json, Uuid}; use sqlx::{ Column, ColumnIndex, Database, MySql, MySqlPool, PgPool, Postgres, Row, Sqlite, SqlitePool, TypeInfo, ValueRef, migrate::MigrateDatabase, mysql::MySqlRow, - postgres::{PgPoolOptions, PgRow}, + postgres::{PgPoolOptions, PgRow, PgValueRef}, sqlite::SqliteRow, types::chrono::{NaiveDate, NaiveDateTime, NaiveTime}, }; @@ -100,6 +100,51 @@ impl DatabasePool { QueryParams::Text(value) => query.bind(value.as_str()), QueryParams::Json(value) => query.bind(value), QueryParams::Bytes(value) => query.bind(value), + QueryParams::PgBoolArray(_) + | QueryParams::PgInt2Array(_) + | QueryParams::PgInt4Array(_) + | QueryParams::PgInt8Array(_) + | QueryParams::PgFloat4Array(_) + | QueryParams::PgFloat8Array(_) + | QueryParams::PgTextArray(_) + | QueryParams::PgBytesArray(_) + | QueryParams::PgUuidArray(_) + | QueryParams::PgJsonArray(_) => { + return Err(sqlx::Error::Configuration( + "PostgreSQL array parameter used with a non-PostgreSQL connection".into(), + )); + } + }; + } + Ok(query) + } + + fn make_pg_query<'a>( + sql: &'a str, + binds: &'a [QueryParams], + ) -> Result< + sqlx::query::Query<'a, Postgres, ::Arguments<'a>>, + sqlx::Error, + > { + let mut query = sqlx::query(sql); + for bind in binds { + query = match bind { + QueryParams::Bool(value) => query.bind(*value), + QueryParams::Int(value) => query.bind(*value), + QueryParams::Float(value) => query.bind(*value), + QueryParams::Text(value) => query.bind(value.as_str()), + QueryParams::Json(value) => query.bind(value), + QueryParams::Bytes(value) => query.bind(value), + QueryParams::PgBoolArray(value) => query.bind(value.as_slice()), + QueryParams::PgInt2Array(value) => query.bind(value.as_slice()), + QueryParams::PgInt4Array(value) => query.bind(value.as_slice()), + QueryParams::PgInt8Array(value) => query.bind(value.as_slice()), + QueryParams::PgFloat4Array(value) => query.bind(value.as_slice()), + QueryParams::PgFloat8Array(value) => query.bind(value.as_slice()), + QueryParams::PgTextArray(value) => query.bind(value.as_slice()), + QueryParams::PgBytesArray(value) => query.bind(value.as_slice()), + QueryParams::PgUuidArray(value) => query.bind(value.as_slice()), + QueryParams::PgJsonArray(value) => query.bind(value.as_slice()), }; } Ok(query) @@ -113,7 +158,7 @@ impl DatabasePool { Ok(DatabaseResponse::MysqlRows(rows)) } DatabasePool::Postgres(pool) => { - let query = Self::make_query(&request.sql, &request.binds)?; + let query = Self::make_pg_query(&request.sql, &request.binds)?; let rows = query.fetch_all(pool).await?; Ok(DatabaseResponse::PgRows(rows)) } @@ -142,7 +187,7 @@ impl DatabasePool { DatabasePool::Postgres(pool) => { let mut transaction = pool.begin().await?; for request in requests { - let query = Self::make_query(&request.sql, &request.binds)?; + let query = Self::make_pg_query(&request.sql, &request.binds)?; query.execute(&mut *transaction).await?; } transaction.commit().await?; @@ -191,6 +236,16 @@ enum QueryParams { Text(String), Json(serde_json::Value), Bytes(Vec), + PgBoolArray(Vec>), + PgInt2Array(Vec>), + PgInt4Array(Vec>), + PgInt8Array(Vec>), + PgFloat4Array(Vec>), + PgFloat8Array(Vec>), + PgTextArray(Vec>), + PgBytesArray(Vec>>), + PgUuidArray(Vec>), + PgJsonArray(Vec>>), } #[derive(Debug, Clone)] @@ -350,6 +405,10 @@ fn get_query_param(state: LuaState, i: i32) -> Result { } } LuaValue::Table(val) => { + if val.getmetafield(cstr!("__sqlx_array")).is_some() { + return get_pg_array_param(&val, &options); + } + let mut buffer = Vec::new(); if let Err(err) = encode_table(&mut buffer, &val, 0, false, &options) { drop(buffer); @@ -375,6 +434,176 @@ fn get_query_param(state: LuaState, i: i32) -> Result { Ok(res) } +fn collect_pg_array( + values: &LuaTable, + expected: &str, + mut convert: F, +) -> Result>, String> +where + F: for<'a> FnMut(LuaValue<'a>) -> Option, +{ + let (is_array, len) = values.array_len(); + if !is_array && values.iter().next().is_some() { + return Err("sqlx.array values must be a Lua sequence".to_string()); + } + + let mut result = Vec::with_capacity(len); + for (index, value) in values.array_iter().enumerate() { + match value { + LuaValue::LightUserData(ptr) if ptr.is_null() => result.push(None), + value => match convert(value) { + Some(value) => result.push(Some(value)), + None => { + return Err(format!( + "sqlx.array element {} must be {}", + index + 1, + expected + )); + } + }, + } + } + Ok(result) +} + +fn lua_value_to_json(value: LuaValue<'_>, options: &JsonOptions) -> Option { + match value { + LuaValue::Nil => Some(serde_json::Value::Null), + LuaValue::LightUserData(ptr) if ptr.is_null() => Some(serde_json::Value::Null), + LuaValue::Boolean(value) => Some(serde_json::Value::Bool(value)), + LuaValue::Integer(value) => Some(serde_json::Value::Number(value.into())), + LuaValue::Number(value) => { + serde_json::Number::from_f64(value).map(serde_json::Value::Number) + } + LuaValue::String(value) => Some(serde_json::Value::String( + String::from_utf8_lossy(value).into_owned(), + )), + LuaValue::Table(value) => { + let mut buffer = Vec::new(); + encode_table(&mut buffer, &value, 0, false, options).ok()?; + serde_json::from_slice(&buffer).ok() + } + _ => None, + } +} + +fn get_pg_array_param( + wrapper: &LuaTable, + options: &JsonOptions, +) -> Result { + let type_name = { + let field = wrapper.rawget("type"); + match &field.value { + LuaValue::String(value) => String::from_utf8_lossy(value).into_owned(), + _ => return Err("sqlx.array type must be a string".to_string()), + } + }; + + let values_field = wrapper.rawget("values"); + let values = match &values_field.value { + LuaValue::Table(value) => value, + _ => return Err("sqlx.array values must be a table".to_string()), + }; + + let normalized = type_name.trim().to_ascii_lowercase(); + let element_type = normalized.strip_suffix("[]").unwrap_or(&normalized); + + match element_type { + "bool" | "boolean" => Ok(QueryParams::PgBoolArray(collect_pg_array( + values, + "a boolean", + |value| match value { + LuaValue::Boolean(value) => Some(value), + _ => None, + }, + )?)), + "int2" | "smallint" => Ok(QueryParams::PgInt2Array(collect_pg_array( + values, + "an int2 integer", + |value| match value { + LuaValue::Integer(value) => i16::try_from(value).ok(), + _ => None, + }, + )?)), + "int4" | "integer" | "int" => Ok(QueryParams::PgInt4Array(collect_pg_array( + values, + "an int4 integer", + |value| match value { + LuaValue::Integer(value) => i32::try_from(value).ok(), + _ => None, + }, + )?)), + "int8" | "bigint" => Ok(QueryParams::PgInt8Array(collect_pg_array( + values, + "an int8 integer", + |value| match value { + LuaValue::Integer(value) => Some(value), + _ => None, + }, + )?)), + "float4" | "real" => Ok(QueryParams::PgFloat4Array(collect_pg_array( + values, + "a number", + |value| match value { + LuaValue::Integer(value) => Some(value as f32), + LuaValue::Number(value) => Some(value as f32), + _ => None, + }, + )?)), + "float8" | "double" | "double precision" => { + Ok(QueryParams::PgFloat8Array(collect_pg_array( + values, + "a number", + |value| match value { + LuaValue::Integer(value) => Some(value as f64), + LuaValue::Number(value) => Some(value), + _ => None, + }, + )?)) + } + "text" | "varchar" | "char" | "character varying" | "name" => { + Ok(QueryParams::PgTextArray(collect_pg_array( + values, + "a string", + |value| match value { + LuaValue::String(value) => { + Some(String::from_utf8_lossy(value).into_owned()) + } + _ => None, + }, + )?)) + } + "bytea" | "bytes" => Ok(QueryParams::PgBytesArray(collect_pg_array( + values, + "a string", + |value| match value { + LuaValue::String(value) => Some(value.to_vec()), + _ => None, + }, + )?)), + "uuid" => Ok(QueryParams::PgUuidArray(collect_pg_array( + values, + "a UUID string", + |value| match value { + LuaValue::String(value) => std::str::from_utf8(value) + .ok() + .and_then(|value| Uuid::parse_str(value).ok()), + _ => None, + }, + )?)), + "json" | "jsonb" => { + let values = collect_pg_array(values, "a JSON value", |value| { + lua_value_to_json(value, options).map(Json) + })?; + Ok(QueryParams::PgJsonArray(values)) + } + _ => Err(format!( + "sqlx.array unsupported PostgreSQL element type: {}", + type_name + )), + } +} + extern "C-unwind" fn query(state: LuaState) -> i32 { let mut args = LuaArgs::new(1); let conn = laux::lua_touserdata::(state, args.iter_arg()) @@ -810,6 +1039,238 @@ where Ok(1) } +fn push_optional_array(state: LuaState, values: Vec>, mut push_value: F) +where + F: FnMut(LuaState, T), +{ + let table = LuaTable::new(state, values.len(), 0); + for (index, value) in values.into_iter().enumerate() { + match value { + Some(value) => push_value(state, value), + None => laux::lua_pushlightuserdata(state, std::ptr::null_mut()), + } + table.rawseti(index + 1); + } +} + +fn process_pg_array_value( + state: LuaState, + row_table: &LuaTable, + column_name: &str, + type_name: &str, + value: PgValueRef<'_>, +) -> Result { + macro_rules! decode_array { + ($value_type:ty) => {{ + > as sqlx::decode::Decode>::decode(value) + .map_err(|err| format!("{} decode error: {}", column_name, err))? + }}; + } + + match type_name { + "BOOL[]" => { + let decoded = decode_array!(bool); + row_table.insert_x(column_name, || { + push_optional_array(state, decoded, laux::lua_push) + }); + } + "INT2[]" => { + let decoded = decode_array!(i16); + row_table.insert_x(column_name, || { + push_optional_array(state, decoded, |state, value| { + laux::lua_push(state, value as i64) + }) + }); + } + "INT4[]" => { + let decoded = decode_array!(i32); + row_table.insert_x(column_name, || { + push_optional_array(state, decoded, |state, value| { + laux::lua_push(state, value as i64) + }) + }); + } + "INT8[]" => { + let decoded = decode_array!(i64); + row_table.insert_x(column_name, || { + push_optional_array(state, decoded, laux::lua_push) + }); + } + "FLOAT4[]" => { + let decoded = decode_array!(f32); + row_table.insert_x(column_name, || { + push_optional_array(state, decoded, |state, value| { + laux::lua_push(state, value as f64) + }) + }); + } + "FLOAT8[]" => { + let decoded = decode_array!(f64); + row_table.insert_x(column_name, || { + push_optional_array(state, decoded, laux::lua_push) + }); + } + "TEXT[]" | "VARCHAR[]" | "CHAR[]" | "NAME[]" => { + let decoded = decode_array!(String); + row_table.insert_x(column_name, || { + push_optional_array(state, decoded, laux::lua_push) + }); + } + "BYTEA[]" => { + let decoded = decode_array!(Vec); + row_table.insert_x(column_name, || { + push_optional_array(state, decoded, |state, value| { + laux::lua_push(state, value.as_slice()) + }) + }); + } + "UUID[]" => { + let decoded = decode_array!(Uuid); + row_table.insert_x(column_name, || { + push_optional_array(state, decoded, |state, value| { + laux::lua_push(state, value.to_string()) + }) + }); + } + "JSON[]" | "JSONB[]" => { + let decoded = decode_array!(Json); + let decoded: Result>, String> = decoded + .into_iter() + .map(|value| { + value + .map(|value| { + serde_json::to_string(&value.0).map_err(|err| { + format!("{} JSON encode error: {}", column_name, err) + }) + }) + .transpose() + }) + .collect(); + let decoded = decoded?; + row_table.insert_x(column_name, || { + push_optional_array(state, decoded, laux::lua_push) + }); + } + _ => return Ok(false), + } + + Ok(true) +} + +fn insert_pg_scalar_value( + row_table: &LuaTable, + column_name: &str, + db_type: DbType, + value: PgValueRef<'_>, +) -> Result<(), String> { + macro_rules! decode_value { + ($value_type:ty) => {{ + <$value_type as sqlx::decode::Decode>::decode(value) + .map_err(|err| format!("{} decode error: {}", column_name, err))? + }}; + } + + match db_type { + DbType::Int8 => row_table.insert(column_name, decode_value!(i8)), + DbType::UInt8 => row_table.insert(column_name, decode_value!(i8) as u8), + DbType::Int16 => row_table.insert(column_name, decode_value!(i16)), + DbType::UInt16 => row_table.insert(column_name, decode_value!(i16) as u16), + DbType::Int32 => row_table.insert(column_name, decode_value!(i32)), + DbType::UInt32 => row_table.insert(column_name, decode_value!(i32) as u32), + DbType::Int64 => row_table.insert(column_name, decode_value!(i64)), + DbType::UInt64 => row_table.insert(column_name, decode_value!(i64) as u64), + DbType::Float32 => row_table.insert(column_name, decode_value!(f32)), + DbType::Float64 => row_table.insert(column_name, decode_value!(f64)), + DbType::Text => row_table.insert(column_name, decode_value!(&str)), + DbType::Bool => row_table.insert(column_name, decode_value!(bool)), + DbType::Timestamp => { + let value = decode_value!(NaiveDateTime); + row_table.insert(column_name, value.format("%Y-%m-%d %H:%M:%S").to_string()) + } + DbType::Date => { + let value = decode_value!(NaiveDate); + row_table.insert(column_name, value.format("%Y-%m-%d").to_string()) + } + DbType::Time => { + let value = decode_value!(NaiveTime); + row_table.insert(column_name, value.format("%H:%M:%S").to_string()) + } + DbType::Uuid => row_table.insert(column_name, decode_value!(Uuid).to_string()), + DbType::Bytes => row_table.insert(column_name, decode_value!(&[u8])), + DbType::Json => { + let value = decode_value!(serde_json::Value); + let json = serde_json::to_string(&value) + .map_err(|err| format!("{} JSON encode error: {}", column_name, err))?; + row_table.insert(column_name, json) + } + DbType::Null => row_table.insert(column_name, LuaNil {}), + DbType::UnsupportedDecimal => { + return Err(format!("Unsupported decimal type for column '{}'", column_name)); + } + DbType::UnsupportedTimeWithTz => { + return Err(format!( + "Unsupported time with time zone type for column '{}'", + column_name + )); + } + DbType::Unknown => { + if let Ok(bytes) = <&[u8] as sqlx::decode::Decode>::decode(value) { + row_table.insert(column_name, bytes) + } else { + row_table.insert(column_name, LuaNil {}) + } + } + }; + + Ok(()) +} + +fn process_pg_rows(state: LuaState, rows: &[PgRow]) -> Result { + let table = LuaTable::new(state, rows.len(), 0); + if rows.is_empty() { + return Ok(1); + } + + let column_info: Vec<_> = rows[0] + .columns() + .iter() + .enumerate() + .map(|(index, column)| { + let type_name = column.type_info().name(); + (index, column.name(), type_name, DbType::from_name(type_name)) + }) + .collect(); + + for (row_index, row) in rows.iter().enumerate() { + let row_table = LuaTable::new(state, 0, row.len()); + for (index, column_name, type_name, db_type) in &column_info { + let value = row + .try_get_raw(*index) + .map_err(|err| format!("{} decode error: {}", column_name, err))?; + + if value.is_null() { + row_table.insert(*column_name, LuaNil {}); + continue; + } + + if process_pg_array_value( + state, + &row_table, + column_name, + type_name, + value.clone(), + )? { + continue; + } + + insert_pg_scalar_value(&row_table, column_name, *db_type, value)?; + } + table.rawseti(row_index + 1); + } + + Ok(1) +} + extern "C-unwind" fn find_connection(state: LuaState) -> i32 { let name = laux::lua_get::<&str>(state, 1); match DATABASE_CONNECTIONSS.get(name) { @@ -845,7 +1306,7 @@ extern "C-unwind" fn decode(state: LuaState) -> i32 { match *result { DatabaseResponse::PgRows(rows) => { - return process_rows::(state, &rows) + return process_pg_rows(state, &rows) .map_err(|e| { push_lua_table!( state, diff --git a/lualib/sqlx.lua b/lualib/sqlx.lua index 9123118..0d5d665 100644 --- a/lualib/sqlx.lua +++ b/lualib/sqlx.lua @@ -17,6 +17,23 @@ moon.register_protocol { ---@class SqlX local M = {} +local pg_array_metatable = { + __sqlx_array = true, +} + +--- Create an explicitly typed one-dimensional PostgreSQL array parameter +---@param element_type string PostgreSQL element type, e.g. "int8", "text", "uuid", or "jsonb" +---@param values table Array values +---@return table +function M.array(element_type, values) + assert(type(element_type) == "string", "sqlx.array element_type must be a string") + assert(type(values) == "table", "sqlx.array values must be a table") + return setmetatable({ + type = element_type, + values = values, + }, pg_array_metatable) +end + --- Connect to a database --- Supported database types: MySQL (mysql://), PostgreSQL (postgres://), SQLite (sqlite://) --- For SQLite, the database will be automatically created if it doesn't exist