use proc_macro2::TokenStream;
use quote::{format_ident, quote};
use syn::Ident;
use crate::from_row::FieldInfo;
#[derive(Clone, Copy)]
pub enum Driver {
Postgres,
Libsql,
Rusqlite,
}
impl Driver {
fn read(self, row: &Ident, info: &FieldInfo) -> syn::Result<TokenStream> {
let ty = &info.field_ty;
let idx = info.idx_usize;
Ok(match self {
Self::Postgres => quote! { #row.try_get::<usize, #ty>(#idx) },
Self::Libsql => {
let Ok(idx) = i32::try_from(idx) else {
return Err(syn::Error::new_spanned(
&info.name,
"too many fields to index a libsql row",
));
};
quote! { #row.get::<#ty>(#idx) }
},
Self::Rusqlite => quote! { #row.get::<usize, #ty>(#idx) },
})
}
}
fn field_init(
core: &TokenStream,
driver: Driver,
row: &Ident,
info: &FieldInfo,
) -> syn::Result<TokenStream> {
let name = &info.name;
let name_str = &info.name_str;
let idx = info.idx_usize;
let read = driver.read(row, info)?;
let mapping = quote! {
|e| #core::error::DbCoreError::RowMapping(
format!("column {} ({}): {}", #idx, #name_str, e)
)
};
let Some(func) = &info.with_fn else {
return Ok(quote! { #name: #read.map_err(#mapping)? });
};
let func_ident: syn::Ident = syn::parse_str(func)
.map_err(|e| syn::Error::new_spanned(&info.name, format!("invalid function name: {e}")))?;
Ok(quote! {
#name: {
let raw = #read.map_err(#mapping)?;
#func_ident(raw).map_err(#mapping)?
}
})
}
fn decoder(
core: &TokenStream,
driver: Driver,
name: &Ident,
fields: &[FieldInfo],
) -> syn::Result<TokenStream> {
let row = format_ident!(
"row_{}",
match driver {
Driver::Postgres => "pg",
Driver::Libsql => "libsql",
Driver::Rusqlite => "rusqlite",
}
);
let inits: Vec<TokenStream> = fields
.iter()
.map(|info| field_init(core, driver, &row, info))
.collect::<syn::Result<_>>()?;
Ok(quote! {
|#row| { Ok(#name { #(#inits,)* }) }
})
}
pub fn emit_from_row_impl(
core: &TokenStream,
name: &Ident,
column_names: &[&str],
fields: &[FieldInfo],
) -> syn::Result<TokenStream> {
let pg = decoder(core, Driver::Postgres, name, fields)?;
let libsql = decoder(core, Driver::Libsql, name, fields)?;
let rusqlite = decoder(core, Driver::Rusqlite, name, fields)?;
Ok(quote! {
#core::impl_derived_from_row! {
#name, &[#(#column_names),*],
postgres = #pg,
libsql = #libsql,
rusqlite = #rusqlite,
}
})
}