diff --git a/georm-macros/src/georm/ir/m2m_relationship.rs b/georm-macros/src/georm/ir/m2m_relationship.rs index b6c6e10..ed08662 100644 --- a/georm-macros/src/georm/ir/m2m_relationship.rs +++ b/georm-macros/src/georm/ir/m2m_relationship.rs @@ -1,3 +1,4 @@ +use crate::georm::sql::{self, FetchKind, SqlDialect}; use quote::quote; #[derive(deluxe::ParseMetaItem, Clone)] @@ -60,7 +61,7 @@ impl From<&M2MRelationshipComplete> for proc_macro2::TokenStream { FROM {} local JOIN {} link ON link.{} = local.{} JOIN {} remote ON link.{} = remote.{} -WHERE local.{} = $1", +WHERE local.{} = {}", value.local.table, value.link.table, value.link.from, @@ -68,15 +69,15 @@ WHERE local.{} = $1", value.remote.table, value.link.to, value.remote.id, - value.local.id + value.local.id, + sql::DIALECT.placeholder(1) ); - quote! { - pub async fn #function<'e, E>(&self, mut executor: E) -> ::sqlx::Result> - where - E: ::sqlx::Executor<'e, Database = ::sqlx::Postgres> - { - ::sqlx::query_as!(#entity, #query, self.get_id()).fetch_all(executor).await - } - } + sql::DIALECT.generate_relation_lookup( + &function, + entity, + &query, + "e! { self.get_id() }, + &FetchKind::Many, + ) } } diff --git a/georm-macros/src/georm/ir/mod.rs b/georm-macros/src/georm/ir/mod.rs index aef2803..eaf4dc9 100644 --- a/georm-macros/src/georm/ir/mod.rs +++ b/georm-macros/src/georm/ir/mod.rs @@ -1,3 +1,4 @@ +use crate::georm::sql::SqlDialect; use quote::quote; pub mod simple_relationship; @@ -156,28 +157,24 @@ impl From<&GeormField> for proc_macro2::TokenStream { proc_macro2::Span::call_site(), ); let entity = &relation.entity; - let return_type = if relation.nullable { - quote! { Option<#entity> } - } else { - quote! { #entity } - }; let query = format!( - "SELECT * FROM {} WHERE {} = $1", - relation.table, relation.remote_id + "SELECT * FROM {} WHERE {} = {}", + relation.table, + relation.remote_id, + crate::georm::sql::DIALECT.placeholder(1) ); let local_ident = &value.field.ident; let fetch = if relation.nullable { - quote! { fetch_optional } + crate::georm::sql::FetchKind::Optional } else { - quote! { fetch_one } + crate::georm::sql::FetchKind::One }; - quote! { - pub async fn #function<'e, E>(&self, mut executor: E) -> ::sqlx::Result<#return_type> - where - E: ::sqlx::Executor<'e, Database = ::sqlx::Postgres> - { - ::sqlx::query_as!(#entity, #query, self.#local_ident).#fetch(executor).await - } - } + crate::georm::sql::DIALECT.generate_relation_lookup( + &function, + entity, + &query, + "e! { self.#local_ident }, + &fetch, + ) } } diff --git a/georm-macros/src/georm/ir/simple_relationship.rs b/georm-macros/src/georm/ir/simple_relationship.rs index 6dfa254..fb25b5e 100644 --- a/georm-macros/src/georm/ir/simple_relationship.rs +++ b/georm-macros/src/georm/ir/simple_relationship.rs @@ -1,3 +1,4 @@ +use crate::georm::sql::{self, FetchKind, SqlDialect}; use quote::quote; pub trait SimpleRelationshipType {} @@ -28,7 +29,12 @@ where T: SimpleRelationshipType + deluxe::ParseMetaItem + Default, { pub fn make_query(&self) -> String { - format!("SELECT * FROM {} WHERE {} = $1", self.table, self.remote_id) + format!( + "SELECT * FROM {} WHERE {} = {}", + self.table, + self.remote_id, + sql::DIALECT.placeholder(1) + ) } pub fn make_function_name(&self) -> syn::Ident { @@ -44,14 +50,13 @@ impl From<&SimpleRelationship> for proc_macro2::TokenStream { let query = value.make_query(); let entity = &value.entity; let function = value.make_function_name(); - quote! { - pub async fn #function<'e, E>(&self, mut executor: E) -> ::sqlx::Result> - where - E: ::sqlx::Executor<'e, Database = ::sqlx::Postgres> - { - ::sqlx::query_as!(#entity, #query, self.get_id()).fetch_optional(executor).await - } - } + sql::DIALECT.generate_relation_lookup( + &function, + entity, + &query, + "e! { self.get_id() }, + &FetchKind::Optional, + ) } } @@ -60,13 +65,12 @@ impl From<&SimpleRelationship> for proc_macro2::TokenStream { let query = value.make_query(); let entity = &value.entity; let function = value.make_function_name(); - quote! { - pub async fn #function<'e, E>(&self, mut executor: E) -> ::sqlx::Result> - where - E: ::sqlx::Executor<'e, Database = ::sqlx::Postgres> - { - ::sqlx::query_as!(#entity, #query, self.get_id()).fetch_all(executor).await - } - } + sql::DIALECT.generate_relation_lookup( + &function, + entity, + &query, + "e! { self.get_id() }, + &FetchKind::Many, + ) } } diff --git a/georm-macros/src/georm/mod.rs b/georm-macros/src/georm/mod.rs index 682f444..aabc173 100644 --- a/georm-macros/src/georm/mod.rs +++ b/georm-macros/src/georm/mod.rs @@ -5,6 +5,7 @@ mod defaultable_struct; mod ir; pub(crate) use ir::GeormField; mod relationships; +mod sql; mod traits; pub(crate) use composite_keys::IdType; @@ -67,18 +68,6 @@ fn generate_from_row_impl( ast: &syn::DeriveInput, fields: &[GeormField], ) -> proc_macro2::TokenStream { - let struct_name = &ast.ident; - let field_idents: Vec<&syn::Ident> = fields.iter().map(|f| &f.ident).collect(); - let field_names: Vec = fields.iter().map(|f| f.ident.to_string()).collect(); - - quote! { - impl<'r> ::sqlx::FromRow<'r, ::sqlx::postgres::PgRow> for #struct_name { - fn from_row(row: &'r ::sqlx::postgres::PgRow) -> ::sqlx::Result { - use ::sqlx::Row; - Ok(Self { - #(#field_idents: row.try_get(#field_names)?),* - }) - } - } - } + use sql::SqlDialect; + sql::DIALECT.generate_from_row(&ast.ident, fields) } diff --git a/georm-macros/src/georm/sql/mod.rs b/georm-macros/src/georm/sql/mod.rs new file mode 100644 index 0000000..b8e3cbf --- /dev/null +++ b/georm-macros/src/georm/sql/mod.rs @@ -0,0 +1,332 @@ +use crate::georm::composite_keys::IdType; +use crate::georm::ir::GeneratedType; +use crate::georm::ir::GeormField; +use quote::quote; + +mod postgres; +pub use postgres::PostgresDialect; + +pub type ActiveDialect = PostgresDialect; +pub const DIALECT: ActiveDialect = PostgresDialect; + +/// How a relation-lookup query should fetch its result. +pub enum FetchKind { + /// `fetch_one`, returns `Entity` + One, + /// `fetch_optional`, returns `Option` + Optional, + /// `fetch_all`, returns `Vec` + Many, +} + +/// Abstracts the SQL dialect-specific parts of Georm's generated code: bind +/// parameter placeholder syntax, and the `sqlx` database/row types used in +/// generated trait bounds and `FromRow` impls. `RETURNING` and +/// `ON CONFLICT ... DO UPDATE` are supported identically by every dialect +/// Georm targets, so query shape itself does not need to vary. +pub trait SqlDialect { + fn placeholder(&self, index: usize) -> String; + fn database_type(&self) -> proc_macro2::TokenStream; + fn row_type(&self) -> proc_macro2::TokenStream; + + fn generate_from_row( + &self, + struct_name: &syn::Ident, + fields: &[GeormField], + ) -> proc_macro2::TokenStream { + let field_idents: Vec<&syn::Ident> = fields.iter().map(|f| &f.ident).collect(); + let field_names: Vec = fields.iter().map(|f| f.ident.to_string()).collect(); + let row = self.row_type(); + quote! { + impl<'r> ::sqlx::FromRow<'r, #row> for #struct_name { + fn from_row(row: &'r #row) -> ::sqlx::Result { + use ::sqlx::Row; + Ok(Self { + #(#field_idents: row.try_get(#field_names)?),* + }) + } + } + } + } + + fn generate_find_all(&self, table: &str) -> proc_macro2::TokenStream { + let find_string = format!("SELECT * FROM {table}"); + let database = self.database_type(); + quote! { + async fn find_all<'e, E>(mut executor: E) -> ::sqlx::Result> + where + E: ::sqlx::Executor<'e, Database = #database> + { + ::sqlx::query_as!(Self, #find_string).fetch_all(executor).await + } + } + } + + fn generate_find(&self, table: &str, id: &IdType) -> proc_macro2::TokenStream { + let database = self.database_type(); + match id { + IdType::Simple { + field_name, + field_type, + } => { + let placeholder = self.placeholder(1); + let find_string = + format!("SELECT * FROM {table} WHERE {field_name} = {placeholder}"); + quote! { + async fn find<'e, E>(mut executor: E, id: &#field_type) -> ::sqlx::Result> + where + E: ::sqlx::Executor<'e, Database = #database> + { + ::sqlx::query_as!(Self, #find_string, id) + .fetch_optional(executor) + .await + } + } + } + IdType::Composite { fields, field_type } => { + let id_match_string = fields + .iter() + .enumerate() + .map(|(i, field)| format!("{} = {}", field.name, self.placeholder(i + 1))) + .collect::>() + .join(" AND "); + let id_members: Vec = + fields.iter().map(|field| field.name.clone()).collect(); + let find_string = format!("SELECT * FROM {table} WHERE {id_match_string}"); + quote! { + async fn find<'e, E>(mut executor: E, id: &#field_type) -> ::sqlx::Result> + where + E: ::sqlx::Executor<'e, Database = #database> + { + ::sqlx::query_as!(Self, #find_string, #(id.#id_members),*) + .fetch_optional(executor) + .await + } + } + } + } + } + + fn generate_create(&self, table: &str, fields: &[GeormField]) -> proc_macro2::TokenStream { + let insert_fields: Vec<&GeormField> = fields + .iter() + .filter(|field| !field.exclude_from_insert()) + .collect(); + let field_names: Vec = insert_fields + .iter() + .map(|field| field.ident.to_string()) + .collect(); + let field_idents: Vec = insert_fields + .iter() + .map(|field| field.ident.clone()) + .collect(); + let placeholders: Vec = (1..=insert_fields.len()) + .map(|i| self.placeholder(i)) + .collect(); + let query = format!( + "INSERT INTO {table} ({}) VALUES ({}) RETURNING *", + field_names.join(", "), + placeholders.join(", ") + ); + let database = self.database_type(); + quote! { + async fn create<'e, E>(&self, mut executor: E) -> ::sqlx::Result + where + E: ::sqlx::Executor<'e, Database = #database> + { + ::sqlx::query_as!( + Self, + #query, + #(self.#field_idents),* + ) + .fetch_one(executor) + .await + } + } + } + + fn generate_update(&self, table: &str, fields: &[GeormField]) -> proc_macro2::TokenStream { + let update_fields: Vec<&GeormField> = fields + .iter() + .filter(|field| !field.is_id && !field.exclude_from_update()) + .collect(); + let update_idents: Vec = update_fields + .iter() + .map(|field| field.ident.clone()) + .collect(); + let id_fields: Vec<&GeormField> = fields.iter().filter(|field| field.is_id).collect(); + let id_idents: Vec = id_fields.iter().map(|f| f.ident.clone()).collect(); + let set_clauses: Vec = update_fields + .iter() + .enumerate() + .map(|(i, field)| format!("{} = {}", field.ident, self.placeholder(i + 1))) + .collect(); + let where_clauses: Vec = id_fields + .iter() + .enumerate() + .map(|(i, field)| { + format!( + "{} = {}", + field.ident, + self.placeholder(update_fields.len() + i + 1) + ) + }) + .collect(); + let query = format!( + "UPDATE {table} SET {} WHERE {} RETURNING *", + set_clauses.join(", "), + where_clauses.join(" AND ") + ); + let database = self.database_type(); + quote! { + async fn update<'e, E>(&self, mut executor: E) -> ::sqlx::Result + where + E: ::sqlx::Executor<'e, Database = #database> + { + ::sqlx::query_as!( + Self, + #query, + #(self.#update_idents),*, + #(self.#id_idents),* + ) + .fetch_one(executor) + .await + } + } + } + + fn generate_upsert( + &self, + table: &str, + fields: &[GeormField], + id: &IdType, + ) -> proc_macro2::TokenStream { + let fields: Vec<&GeormField> = fields + .iter() + .filter(|field| !matches!(field.generated_type, GeneratedType::Always)) + .collect(); + let inputs: Vec = (1..=fields.len()).map(|i| self.placeholder(i)).collect(); + let columns = fields + .iter() + .map(|f| f.ident.to_string()) + .collect::>() + .join(", "); + + let primary_key: proc_macro2::TokenStream = match id { + IdType::Simple { field_name, .. } => quote! {#field_name}, + IdType::Composite { fields, .. } => { + let field_names: Vec = fields.iter().map(|f| f.name.clone()).collect(); + quote! { + #(#field_names),* + } + } + }; + + // For ON CONFLICT DO UPDATE, exclude the ID field from updates + let update_assignments = fields + .iter() + .filter(|f| !f.is_id) + .map(|f| format!("{} = EXCLUDED.{}", f.ident, f.ident)) + .collect::>() + .join(", "); + + let upsert_string = format!( + "INSERT INTO {table} ({columns}) VALUES ({}) ON CONFLICT ({}) DO UPDATE SET {update_assignments} RETURNING *", + inputs.join(", "), + primary_key + ); + + let field_idents: Vec = fields.iter().map(|f| f.ident.clone()).collect(); + let database = self.database_type(); + + quote! { + async fn upsert<'e, E>(&self, mut executor: E) -> ::sqlx::Result + where + E: ::sqlx::Executor<'e, Database = #database> + { + ::sqlx::query_as!( + Self, + #upsert_string, + #(self.#field_idents),* + ) + .fetch_one(executor) + .await + } + } + } + + fn generate_delete(&self, table: &str, id: &IdType) -> proc_macro2::TokenStream { + let where_clause = match id { + IdType::Simple { field_name, .. } => { + format!("{} = {}", field_name, self.placeholder(1)) + } + IdType::Composite { fields, .. } => fields + .iter() + .enumerate() + .map(|(i, field)| format!("{} = {}", field.name, self.placeholder(i + 1))) + .collect::>() + .join(" AND "), + }; + let query_args = match id { + IdType::Simple { .. } => quote! { id }, + IdType::Composite { fields, .. } => { + let fields: Vec = fields.iter().map(|f| f.name.clone()).collect(); + quote! { #(id.#fields), * } + } + }; + let id_type = match id { + IdType::Simple { field_type, .. } => quote! { #field_type }, + IdType::Composite { field_type, .. } => quote! { #field_type }, + }; + let delete_string = format!("DELETE FROM {table} WHERE {where_clause}"); + let database = self.database_type(); + quote! { + async fn delete<'e, E>(&self, mut executor: E) -> ::sqlx::Result + where + E: ::sqlx::Executor<'e, Database = #database> + { + Self::delete_by_id(executor, &self.get_id()).await + } + + async fn delete_by_id<'e, E>(mut executor: E, id: &#id_type) -> ::sqlx::Result + where + E: ::sqlx::Executor<'e, Database = #database> + { + let rows_affected = ::sqlx::query!(#delete_string, #query_args) + .execute(executor) + .await? + .rows_affected(); + Ok(rows_affected) + } + } + } + + /// Wraps a single-parameter relation-lookup query (field-level + /// `#[georm(relation = ...)]`, struct-level `one_to_one`/`one_to_many`, + /// and many-to-many joins) in an async method with the right executor + /// bound. `query` must already contain this dialect's placeholder syntax + /// (built via [`SqlDialect::placeholder`]). + fn generate_relation_lookup( + &self, + function: &syn::Ident, + entity: &syn::Type, + query: &str, + arg: &proc_macro2::TokenStream, + fetch: &FetchKind, + ) -> proc_macro2::TokenStream { + let database = self.database_type(); + let (return_type, fetch_method) = match fetch { + FetchKind::One => (quote! { #entity }, quote! { fetch_one }), + FetchKind::Optional => (quote! { Option<#entity> }, quote! { fetch_optional }), + FetchKind::Many => (quote! { Vec<#entity> }, quote! { fetch_all }), + }; + quote! { + pub async fn #function<'e, E>(&self, mut executor: E) -> ::sqlx::Result<#return_type> + where + E: ::sqlx::Executor<'e, Database = #database> + { + ::sqlx::query_as!(#entity, #query, #arg).#fetch_method(executor).await + } + } + } +} diff --git a/georm-macros/src/georm/sql/postgres.rs b/georm-macros/src/georm/sql/postgres.rs new file mode 100644 index 0000000..876bc90 --- /dev/null +++ b/georm-macros/src/georm/sql/postgres.rs @@ -0,0 +1,18 @@ +use super::SqlDialect; +use quote::quote; + +pub struct PostgresDialect; + +impl SqlDialect for PostgresDialect { + fn placeholder(&self, index: usize) -> String { + format!("${index}") + } + + fn database_type(&self) -> proc_macro2::TokenStream { + quote! { ::sqlx::Postgres } + } + + fn row_type(&self) -> proc_macro2::TokenStream { + quote! { ::sqlx::postgres::PgRow } + } +} diff --git a/georm-macros/src/georm/traits/create.rs b/georm-macros/src/georm/traits/create.rs index 810bebc..444ed03 100644 --- a/georm-macros/src/georm/traits/create.rs +++ b/georm-macros/src/georm/traits/create.rs @@ -1,37 +1,6 @@ use crate::georm::GeormField; -use quote::quote; +use crate::georm::sql::{self, SqlDialect}; pub fn generate_create_query(table_name: &str, fields: &[GeormField]) -> proc_macro2::TokenStream { - let insert_fields: Vec<&GeormField> = fields - .iter() - .filter(|field| !field.exclude_from_insert()) - .collect(); - let field_names: Vec = insert_fields - .iter() - .map(|field| field.ident.to_string()) - .collect(); - let field_idents: Vec = insert_fields - .iter() - .map(|field| field.ident.clone()) - .collect(); - let placeholders: Vec = (1..=insert_fields.len()).map(|i| format!("${i}")).collect(); - let query = format!( - "INSERT INTO {table_name} ({}) VALUES ({}) RETURNING *", - field_names.join(", "), - placeholders.join(", ") - ); - quote! { - async fn create<'e, E>(&self, mut executor: E) -> ::sqlx::Result - where - E: ::sqlx::Executor<'e, Database = ::sqlx::Postgres> - { - ::sqlx::query_as!( - Self, - #query, - #(self.#field_idents),* - ) - .fetch_one(executor) - .await - } - } + sql::DIALECT.generate_create(table_name, fields) } diff --git a/georm-macros/src/georm/traits/delete.rs b/georm-macros/src/georm/traits/delete.rs index 253c5be..aa30ace 100644 --- a/georm-macros/src/georm/traits/delete.rs +++ b/georm-macros/src/georm/traits/delete.rs @@ -1,45 +1,6 @@ use crate::georm::IdType; -use quote::quote; +use crate::georm::sql::{self, SqlDialect}; pub fn generate_delete_query(table: &str, id: &IdType) -> proc_macro2::TokenStream { - let where_clause = match id { - IdType::Simple { field_name, .. } => format!("{} = $1", field_name), - IdType::Composite { fields, .. } => fields - .iter() - .enumerate() - .map(|(i, field)| format!("{} = ${}", field.name, i + 1)) - .collect::>() - .join(" AND "), - }; - let query_args = match id { - IdType::Simple { .. } => quote! { id }, - IdType::Composite { fields, .. } => { - let fields: Vec = fields.iter().map(|f| f.name.clone()).collect(); - quote! { #(id.#fields), * } - } - }; - let id_type = match id { - IdType::Simple { field_type, .. } => quote! { #field_type }, - IdType::Composite { field_type, .. } => quote! { #field_type }, - }; - let delete_string = format!("DELETE FROM {table} WHERE {where_clause}"); - quote! { - async fn delete<'e, E>(&self, mut executor: E) -> ::sqlx::Result - where - E: ::sqlx::Executor<'e, Database = ::sqlx::Postgres> - { - Self::delete_by_id(executor, &self.get_id()).await - } - - async fn delete_by_id<'e, E>(mut executor: E, id: &#id_type) -> ::sqlx::Result - where - E: ::sqlx::Executor<'e, Database = ::sqlx::Postgres> - { - let rows_affected = ::sqlx::query!(#delete_string, #query_args) - .execute(executor) - .await? - .rows_affected(); - Ok(rows_affected) - } - } + sql::DIALECT.generate_delete(table, id) } diff --git a/georm-macros/src/georm/traits/find.rs b/georm-macros/src/georm/traits/find.rs index 92d1f8f..5892df8 100644 --- a/georm-macros/src/georm/traits/find.rs +++ b/georm-macros/src/georm/traits/find.rs @@ -1,56 +1,10 @@ use crate::georm::IdType; -use quote::quote; +use crate::georm::sql::{self, SqlDialect}; pub fn generate_find_all_query(table: &str) -> proc_macro2::TokenStream { - let find_string = format!("SELECT * FROM {table}"); - quote! { - async fn find_all<'e, E>(mut executor: E) -> ::sqlx::Result> - where - E: ::sqlx::Executor<'e, Database = ::sqlx::Postgres> - { - ::sqlx::query_as!(Self, #find_string).fetch_all(executor).await - } - } + sql::DIALECT.generate_find_all(table) } pub fn generate_find_query(table: &str, id: &IdType) -> proc_macro2::TokenStream { - match id { - IdType::Simple { - field_name, - field_type, - } => { - let find_string = format!("SELECT * FROM {table} WHERE {} = $1", field_name); - quote! { - async fn find<'e, E>(mut executor: E, id: &#field_type) -> ::sqlx::Result> - where - E: ::sqlx::Executor<'e, Database = ::sqlx::Postgres> - { - ::sqlx::query_as!(Self, #find_string, id) - .fetch_optional(executor) - .await - } - } - } - IdType::Composite { fields, field_type } => { - let id_match_string = fields - .iter() - .enumerate() - .map(|(i, field)| format!("{} = ${}", field.name, i + 1)) - .collect::>() - .join(" AND "); - let id_members: Vec = - fields.iter().map(|field| field.name.clone()).collect(); - let find_string = format!("SELECT * FROM {table} WHERE {id_match_string}"); - quote! { - async fn find<'e, E>(mut executor: E, id: &#field_type) -> ::sqlx::Result> - where - E: ::sqlx::Executor<'e, Database = ::sqlx::Postgres> - { - ::sqlx::query_as!(Self, #find_string, #(id.#id_members),*) - .fetch_optional(executor) - .await - } - } - } - } + sql::DIALECT.generate_find(table, id) } diff --git a/georm-macros/src/georm/traits/update.rs b/georm-macros/src/georm/traits/update.rs index 27bd544..a056753 100644 --- a/georm-macros/src/georm/traits/update.rs +++ b/georm-macros/src/georm/traits/update.rs @@ -1,45 +1,6 @@ use crate::georm::GeormField; -use quote::quote; +use crate::georm::sql::{self, SqlDialect}; pub fn generate_update_query(table_name: &str, fields: &[GeormField]) -> proc_macro2::TokenStream { - let update_fields: Vec<&GeormField> = fields - .iter() - .filter(|field| !field.is_id && !field.exclude_from_update()) - .collect(); - let update_idents: Vec = update_fields - .iter() - .map(|field| field.ident.clone()) - .collect(); - let id_fields: Vec<&GeormField> = fields.iter().filter(|field| field.is_id).collect(); - let id_idents: Vec = id_fields.iter().map(|f| f.ident.clone()).collect(); - let set_clauses: Vec = update_fields - .iter() - .enumerate() - .map(|(i, field)| format!("{} = ${}", field.ident, i + 1)) - .collect(); - let where_clauses: Vec = id_fields - .iter() - .enumerate() - .map(|(i, field)| format!("{} = ${}", field.ident, update_fields.len() + i + 1)) - .collect(); - let query = format!( - "UPDATE {table_name} SET {} WHERE {} RETURNING *", - set_clauses.join(", "), - where_clauses.join(" AND ") - ); - quote! { - async fn update<'e, E>(&self, mut executor: E) -> ::sqlx::Result - where - E: ::sqlx::Executor<'e, Database = ::sqlx::Postgres> - { - ::sqlx::query_as!( - Self, - #query, - #(self.#update_idents),*, - #(self.#id_idents),* - ) - .fetch_one(executor) - .await - } - } + sql::DIALECT.generate_update(table_name, fields) } diff --git a/georm-macros/src/georm/traits/upsert.rs b/georm-macros/src/georm/traits/upsert.rs index 271f0bd..f6821d0 100644 --- a/georm-macros/src/georm/traits/upsert.rs +++ b/georm-macros/src/georm/traits/upsert.rs @@ -1,60 +1,10 @@ -use crate::georm::{GeormField, IdType, ir::GeneratedType}; -use quote::quote; +use crate::georm::sql::{self, SqlDialect}; +use crate::georm::{GeormField, IdType}; pub fn generate_upsert_query( table: &str, fields: &[GeormField], id: &IdType, ) -> proc_macro2::TokenStream { - let fields: Vec<&GeormField> = fields - .iter() - .filter(|field| !matches!(field.generated_type, GeneratedType::Always)) - .collect(); - let inputs: Vec = (1..=fields.len()).map(|num| format!("${num}")).collect(); - let columns = fields - .iter() - .map(|f| f.ident.to_string()) - .collect::>() - .join(", "); - - let primary_key: proc_macro2::TokenStream = match id { - IdType::Simple { field_name, .. } => quote! {#field_name}, - IdType::Composite { fields, .. } => { - let field_names: Vec = fields.iter().map(|f| f.name.clone()).collect(); - quote! { - #(#field_names),* - } - } - }; - - // For ON CONFLICT DO UPDATE, exclude the ID field from updates - let update_assignments = fields - .iter() - .filter(|f| !f.is_id) - .map(|f| format!("{} = EXCLUDED.{}", f.ident, f.ident)) - .collect::>() - .join(", "); - - let upsert_string = format!( - "INSERT INTO {table} ({columns}) VALUES ({}) ON CONFLICT ({}) DO UPDATE SET {update_assignments} RETURNING *", - inputs.join(", "), - primary_key - ); - - let field_idents: Vec = fields.iter().map(|f| f.ident.clone()).collect(); - - quote! { - async fn upsert<'e, E>(&self, mut executor: E) -> ::sqlx::Result - where - E: ::sqlx::Executor<'e, Database = ::sqlx::Postgres> - { - ::sqlx::query_as!( - Self, - #upsert_string, - #(self.#field_idents),* - ) - .fetch_one(executor) - .await - } - } + sql::DIALECT.generate_upsert(table, fields, id) }