1use proc_macro::TokenStream;
2use quote::{format_ident, quote};
3use syn::parse::Parser;
4use syn::{parse_macro_input, Attribute, Field, Fields, Ident, ItemStruct, Meta, Type};
5
6#[proc_macro_derive(
7 Model,
8 attributes(model, key, autoincrement, unique, index, has_many, belongs_to, many_to_many, dbkit)
9)]
10pub fn derive_model(_input: TokenStream) -> TokenStream {
11 TokenStream::from(quote! {
12 compile_error!("dbkit: use #[model] instead of #[derive(Model)]");
13 })
14}
15
16#[proc_macro_derive(DbEnum, attributes(dbkit))]
17pub fn derive_db_enum(input: TokenStream) -> TokenStream {
18 let input = parse_macro_input!(input as syn::ItemEnum);
19 match expand_db_enum(input) {
20 Ok(tokens) => tokens,
21 Err(err) => err.to_compile_error().into(),
22 }
23}
24
25#[proc_macro_attribute]
26pub fn model(attr: TokenStream, item: TokenStream) -> TokenStream {
27 let input = parse_macro_input!(item as ItemStruct);
28 let args = parse_macro_input!(attr with syn::punctuated::Punctuated::<Meta, syn::Token![,]>::parse_terminated);
29 let args = parse_model_args(args);
30 match expand_model(args, input) {
31 Ok(tokens) => tokens,
32 Err(err) => err.to_compile_error().into(),
33 }
34}
35
36#[derive(Debug, Clone, Copy, PartialEq, Eq)]
37enum RelationKind {
38 HasMany,
39 BelongsTo,
40 ManyToMany,
41}
42
43struct RelationInfo {
44 field: Field,
45 param_ident: Ident,
46 state_mod_ident: Ident,
47 child_type: Type,
48 kind: RelationKind,
49 belongs_to_key: Option<Ident>,
50 belongs_to_ref: Option<Ident>,
51 many_to_many_through: Option<Ident>,
52 many_to_many_left_key: Option<Ident>,
53 many_to_many_right_key: Option<Ident>,
54}
55
56struct ScalarFieldInfo {
57 field: Field,
58 ident: Ident,
59 ty: Type,
60 column_name: String,
61 is_key: bool,
62 is_autoincrement: bool,
63}
64
65#[derive(Default)]
66struct ModelArgs {
67 table: Option<String>,
68 schema: Option<String>,
69}
70
71fn expand_model(args: ModelArgs, input: ItemStruct) -> syn::Result<TokenStream> {
72 if !input.generics.params.is_empty() {
73 return Err(syn::Error::new_spanned(
74 input.generics,
75 "dbkit: #[model] does not support generics yet",
76 ));
77 }
78
79 let struct_ident = input.ident;
80 let model_ident = format_ident!("{}Model", struct_ident);
81 let insert_ident = format_ident!("{}Insert", struct_ident);
82 let vis = input.vis;
83
84 let table_name = args.table.unwrap_or_else(|| to_snake_case(&struct_ident.to_string()));
85 let schema_name = args.schema;
86
87 let mut primary_keys: Vec<(Ident, Type, String)> = Vec::new();
88 let mut relation_fields = Vec::new();
89 let mut output_fields = Vec::new();
90 let mut insert_fields = Vec::new();
91 let mut scalar_fields = Vec::new();
92
93 let struct_attrs = filter_struct_attrs(&input.attrs);
94
95 let fields = match input.fields {
96 Fields::Named(named) => named.named,
97 _ => {
98 return Err(syn::Error::new_spanned(
99 struct_ident,
100 "dbkit: #[model] requires a struct with named fields",
101 ))
102 }
103 };
104
105 for field in fields {
106 let field_ident = field
107 .ident
108 .clone()
109 .ok_or_else(|| syn::Error::new_spanned(&field, "dbkit: unnamed field"))?;
110
111 let is_relation =
112 has_attr(&field.attrs, "has_many") || has_attr(&field.attrs, "belongs_to") || has_attr(&field.attrs, "many_to_many");
113
114 let is_key = has_attr(&field.attrs, "key");
115 let is_autoincrement = has_attr(&field.attrs, "autoincrement");
116
117 if is_relation {
118 if parse_field_column_name(&field.attrs)?.is_some() {
119 return Err(syn::Error::new_spanned(
120 &field,
121 "dbkit: `#[dbkit(column = \"...\")]` is only supported on scalar fields",
122 ));
123 }
124 let (kind, child_type) = relation_type(&field)?;
125 let state_mod_ident = format_ident!("{}_{}_state", to_snake_case(&struct_ident.to_string()), field_ident);
126 let param_ident = format_ident!("{}Rel", to_camel_case(&field_ident.to_string()));
127 let (belongs_to_key, belongs_to_ref) = if kind == RelationKind::BelongsTo {
128 let (key, references) = parse_belongs_to_args(&field.attrs)?;
129 (Some(key), Some(references))
130 } else {
131 (None, None)
132 };
133 let (many_to_many_through, many_to_many_left_key, many_to_many_right_key) = if kind == RelationKind::ManyToMany {
134 let (through, left_key, right_key) = parse_many_to_many_args(&field.attrs)?;
135 (Some(through), Some(left_key), Some(right_key))
136 } else {
137 (None, None, None)
138 };
139
140 relation_fields.push(RelationInfo {
141 field: field.clone(),
142 param_ident: param_ident.clone(),
143 state_mod_ident,
144 child_type,
145 kind,
146 belongs_to_key,
147 belongs_to_ref,
148 many_to_many_through,
149 many_to_many_left_key,
150 many_to_many_right_key,
151 });
152
153 let cleaned_field = Field {
154 attrs: filter_field_attrs(&field.attrs),
155 ty: syn::parse_quote!(#param_ident),
156 ..field
157 };
158 output_fields.push(cleaned_field);
159 continue;
160 }
161
162 let column_name = parse_field_column_name(&field.attrs)?.unwrap_or_else(|| field_ident.to_string());
163
164 if is_key {
165 primary_keys.push((field_ident.clone(), field.ty.clone(), column_name.clone()));
166 }
167
168 let cleaned_field = Field {
169 attrs: filter_field_attrs(&field.attrs),
170 ..field.clone()
171 };
172 output_fields.push(cleaned_field.clone());
173
174 if !(is_key && is_autoincrement) {
175 insert_fields.push(cleaned_field.clone());
176 }
177
178 scalar_fields.push(ScalarFieldInfo {
179 field: cleaned_field,
180 ident: field_ident,
181 ty: field.ty.clone(),
182 column_name,
183 is_key,
184 is_autoincrement,
185 });
186 }
187
188 let table_expr = if let Some(schema) = schema_name {
189 quote!(::dbkit::Table::new(#table_name).with_schema(#schema))
190 } else {
191 quote!(::dbkit::Table::new(#table_name))
192 };
193
194 if relation_fields.iter().any(|rel| rel.kind == RelationKind::ManyToMany) && primary_keys.len() != 1 {
195 return Err(syn::Error::new_spanned(
196 struct_ident,
197 "dbkit: many-to-many requires exactly one #[key] on the parent model",
198 ));
199 }
200
201 let generics_with_defaults = relation_fields
202 .iter()
203 .map(|rel| {
204 let ident = &rel.param_ident;
205 let state_mod = &rel.state_mod_ident;
206 quote!(#ident: #state_mod::State = ::dbkit::NotLoaded)
207 })
208 .collect::<Vec<_>>();
209
210 let impl_generics_params = relation_fields
211 .iter()
212 .map(|rel| {
213 let ident = &rel.param_ident;
214 let state_mod = &rel.state_mod_ident;
215 quote!(#ident: #state_mod::State)
216 })
217 .collect::<Vec<_>>();
218
219 let generic_idents = relation_fields.iter().map(|rel| &rel.param_ident).collect::<Vec<_>>();
220
221 let struct_generics = if generics_with_defaults.is_empty() {
222 quote!()
223 } else {
224 quote!(<#(#generics_with_defaults),*>)
225 };
226
227 let impl_generics = if impl_generics_params.is_empty() {
228 quote!()
229 } else {
230 quote!(<#(#impl_generics_params),*>)
231 };
232
233 let struct_type_args = if generic_idents.is_empty() {
234 quote!()
235 } else {
236 quote!(<#(#generic_idents),*>)
237 };
238
239 let columns = output_fields
240 .iter()
241 .filter(|field| !is_relation_field(field, &relation_fields))
242 .map(|field| {
243 let ident = field.ident.as_ref().expect("field ident");
244 let name = scalar_column_name(&scalar_fields, ident).expect("scalar field column name");
245 let ty = option_inner_type(&field.ty).unwrap_or_else(|| field.ty.clone());
246 quote!(pub const #ident: ::dbkit::Column<#struct_ident, #ty> = ::dbkit::Column::new(Self::TABLE, #name);)
247 })
248 .collect::<Vec<_>>();
249
250 let column_refs = output_fields
251 .iter()
252 .filter(|field| !is_relation_field(field, &relation_fields))
253 .map(|field| {
254 let ident = field.ident.as_ref().expect("field ident");
255 quote!(Self::#ident.as_ref())
256 })
257 .collect::<Vec<_>>();
258
259 let columns_const = quote!(
260 pub const COLUMNS: &'static [::dbkit::ColumnRef] = &[#(#column_refs),*];
261 );
262
263 let primary_key_refs = primary_keys
264 .iter()
265 .map(|(ident, _, _)| quote!(Self::#ident.as_ref()))
266 .collect::<Vec<_>>();
267
268 let primary_keys_const = if primary_keys.is_empty() {
269 quote!(
270 pub const PRIMARY_KEYS: &'static [::dbkit::ColumnRef] = &[];
271 )
272 } else {
273 quote!(pub const PRIMARY_KEYS: &'static [::dbkit::ColumnRef] = &[#(#primary_key_refs),*];)
274 };
275
276 let insert_values = insert_fields.iter().map(|field| {
277 let ident = field.ident.as_ref().expect("field ident");
278 quote!(insert = insert.value(Self::#ident, values.#ident);)
279 });
280 let insert_field_idents = insert_fields
281 .iter()
282 .map(|field| field.ident.as_ref().expect("field ident"))
283 .collect::<Vec<_>>();
284
285 let active_ident = format_ident!("{}Active", struct_ident);
286
287 let active_fields = scalar_fields.iter().map(|field| {
288 let ident = &field.ident;
289 let vis = &field.field.vis;
290 let ty = option_inner_type(&field.ty).unwrap_or_else(|| field.ty.clone());
291 quote!(#vis #ident: ::dbkit::ActiveValue<#ty>)
292 });
293
294 let active_from_model = scalar_fields.iter().map(|field| {
295 let ident = &field.ident;
296 if option_inner_type(&field.ty).is_some() {
297 quote!(#ident: ::dbkit::ActiveValue::unchanged_option(#ident))
298 } else {
299 quote!(#ident: ::dbkit::ActiveValue::unchanged(#ident))
300 }
301 });
302
303 let active_destructure = scalar_fields.iter().map(|field| field.ident.clone()).collect::<Vec<_>>();
304
305 let active_insert_steps = scalar_fields.iter().map(|field| {
306 let ident = &field.ident;
307 let name = ident.to_string();
308 let ty = option_inner_type(&field.ty).unwrap_or_else(|| field.ty.clone());
309 let is_option = option_inner_type(&field.ty).is_some();
310 let required = !field.is_autoincrement && !is_option;
311 let required_check = if required {
312 quote!(return Err(::dbkit::Error::Decode(format!("missing required field: {}", #name)));)
313 } else {
314 quote!()
315 };
316 quote!(
317 match #ident {
318 ::dbkit::ActiveValue::Unset => {
319 #required_check
320 }
321 ::dbkit::ActiveValue::Set(value) => {
322 insert = insert.value(#struct_ident::#ident, value);
323 }
324 ::dbkit::ActiveValue::Unchanged(value) => {
325 insert = insert.value(#struct_ident::#ident, value);
326 }
327 ::dbkit::ActiveValue::UnchangedNull => {
328 insert = insert.value(#struct_ident::#ident, None::<#ty>);
329 }
330 ::dbkit::ActiveValue::Null => {
331 insert = insert.value(#struct_ident::#ident, None::<#ty>);
332 }
333 }
334 )
335 });
336
337 let active_insert_fn = quote!(
338 pub async fn insert(
339 self,
340 ex: &(impl ::dbkit::Executor + Send + Sync),
341 ) -> Result<#struct_ident, ::dbkit::Error> {
342 let Self { #(#active_destructure,)* } = self;
343 let mut insert = ::dbkit::Insert::new(#struct_ident::TABLE);
344 #(#active_insert_steps)*
345 let insert = insert.returning_all();
346 let row = ::dbkit::InsertExt::one(insert, ex).await?;
347 row.ok_or(::dbkit::Error::NotFound)
348 }
349 );
350
351 let pk_idents = primary_keys.iter().map(|(ident, _, _)| ident.clone()).collect::<Vec<_>>();
352
353 let active_update_fn = if !primary_keys.is_empty() {
354 let pk_vars = primary_keys
355 .iter()
356 .enumerate()
357 .map(|(idx, _)| format_ident!("pk_value_{}", idx))
358 .collect::<Vec<_>>();
359 let pk_extracts = primary_keys.iter().zip(pk_vars.iter()).map(|((ident, _, _), var)| {
360 let pk_name = ident.to_string();
361 quote!(
362 let #var = match #ident {
363 ::dbkit::ActiveValue::Set(value) | ::dbkit::ActiveValue::Unchanged(value) => value,
364 ::dbkit::ActiveValue::Null | ::dbkit::ActiveValue::Unset | ::dbkit::ActiveValue::UnchangedNull => {
365 return Err(::dbkit::Error::Decode(format!(
366 "missing required field: {}",
367 #pk_name
368 )));
369 }
370 };
371 )
372 });
373 let pk_filters = primary_keys
374 .iter()
375 .zip(pk_vars.iter())
376 .map(|((ident, _, _), var)| quote!(update = update.filter(#struct_ident::#ident.eq(#var));));
377 let update_steps = scalar_fields.iter().filter(|field| !field.is_key).map(|field| {
378 let ident = &field.ident;
379 let ty = option_inner_type(&field.ty).unwrap_or_else(|| field.ty.clone());
380 quote!(
381 match #ident {
382 ::dbkit::ActiveValue::Unset => {}
383 ::dbkit::ActiveValue::Set(value) => {
384 update = update.set(#struct_ident::#ident, value);
385 any_set = true;
386 }
387 ::dbkit::ActiveValue::Unchanged(_) | ::dbkit::ActiveValue::UnchangedNull => {}
388 ::dbkit::ActiveValue::Null => {
389 update = update.set(#struct_ident::#ident, None::<#ty>);
390 any_set = true;
391 }
392 }
393 )
394 });
395 quote!(
396 pub async fn update(
397 self,
398 ex: &(impl ::dbkit::Executor + Send + Sync),
399 ) -> Result<#struct_ident, ::dbkit::Error> {
400 let Self { #(#active_destructure,)* } = self;
401 #(#pk_extracts)*
402 let mut update = ::dbkit::Update::new(#struct_ident::TABLE);
403 let mut any_set = false;
404 #(#update_steps)*
405 if !any_set {
406 return Err(::dbkit::Error::Decode("no fields set for update".to_string()));
407 }
408 #(#pk_filters)*
409 let update = update.returning_all();
410 let mut rows = ::dbkit::UpdateExt::all(update, ex).await?;
411 rows.pop().ok_or(::dbkit::Error::NotFound)
412 }
413 )
414 } else {
415 quote!()
416 };
417
418 let active_delete_fn = if !primary_keys.is_empty() {
419 let pk_vars = primary_keys
420 .iter()
421 .enumerate()
422 .map(|(idx, _)| format_ident!("pk_value_{}", idx))
423 .collect::<Vec<_>>();
424 let pk_extracts = primary_keys.iter().zip(pk_vars.iter()).map(|((ident, _, _), var)| {
425 let pk_name = ident.to_string();
426 quote!(
427 let #var = match #ident {
428 ::dbkit::ActiveValue::Set(value) | ::dbkit::ActiveValue::Unchanged(value) => value,
429 ::dbkit::ActiveValue::Null | ::dbkit::ActiveValue::Unset | ::dbkit::ActiveValue::UnchangedNull => {
430 return Err(::dbkit::Error::Decode(format!(
431 "missing required field: {}",
432 #pk_name
433 )));
434 }
435 };
436 )
437 });
438 let pk_filters = primary_keys
439 .iter()
440 .zip(pk_vars.iter())
441 .map(|((ident, _, _), var)| quote!(delete = delete.filter(#struct_ident::#ident.eq(#var));));
442 quote!(
443 pub async fn delete(
444 self,
445 ex: &(impl ::dbkit::Executor + Send + Sync),
446 ) -> Result<u64, ::dbkit::Error> {
447 let Self { #(#pk_idents,)* .. } = self;
448 #(#pk_extracts)*
449 let mut delete = ::dbkit::Delete::new(#struct_ident::TABLE);
450 #(#pk_filters)*
451 ::dbkit::DeleteExt::execute(delete, ex).await
452 }
453 )
454 } else {
455 quote!()
456 };
457
458 let active_save_flag_checks = scalar_fields.iter().map(|field| {
459 let ident = &field.ident;
460 quote!(
461 match &#ident {
462 ::dbkit::ActiveValue::Unchanged(_) | ::dbkit::ActiveValue::UnchangedNull => {
463 any_loaded = true;
464 }
465 ::dbkit::ActiveValue::Set(_) | ::dbkit::ActiveValue::Null => {
466 any_changed = true;
467 }
468 ::dbkit::ActiveValue::Unset => {}
469 }
470 )
471 });
472
473 let active_save_model_fields = scalar_fields.iter().map(|field| {
474 let ident = &field.ident;
475 let name = ident.to_string();
476 if option_inner_type(&field.ty).is_some() {
477 quote!(
478 #ident: match #ident {
479 ::dbkit::ActiveValue::Set(value) | ::dbkit::ActiveValue::Unchanged(value) => Some(value),
480 ::dbkit::ActiveValue::Null | ::dbkit::ActiveValue::UnchangedNull => None,
481 ::dbkit::ActiveValue::Unset => {
482 return Err(::dbkit::Error::Decode(format!(
483 "missing required field: {}",
484 #name
485 )));
486 }
487 },
488 )
489 } else {
490 quote!(
491 #ident: match #ident {
492 ::dbkit::ActiveValue::Set(value) | ::dbkit::ActiveValue::Unchanged(value) => value,
493 ::dbkit::ActiveValue::Null
494 | ::dbkit::ActiveValue::Unset
495 | ::dbkit::ActiveValue::UnchangedNull => {
496 return Err(::dbkit::Error::Decode(format!(
497 "missing required field: {}",
498 #name
499 )));
500 }
501 },
502 )
503 }
504 });
505
506 let active_save_relation_defaults = relation_fields.iter().map(|rel| {
507 let ident = rel.field.ident.as_ref().expect("field ident");
508 quote!(#ident: Default::default(),)
509 });
510
511 let active_save_update_branch = if !primary_keys.is_empty() {
512 quote!(return Self { #(#active_destructure,)* }.update(ex).await;)
513 } else {
514 quote!(
515 return Err(::dbkit::Error::Decode(
516 "update requires primary key".to_string(),
517 ));
518 )
519 };
520
521 let active_save_fn = quote!(
522 pub async fn save(
523 self,
524 ex: &(impl ::dbkit::Executor + Send + Sync),
525 ) -> Result<#struct_ident, ::dbkit::Error> {
526 let Self { #(#active_destructure,)* } = self;
527 let mut any_loaded = false;
528 let mut any_changed = false;
529 #(#active_save_flag_checks)*
530
531 if any_loaded {
532 if any_changed {
533 #active_save_update_branch
534 }
535 let model = #struct_ident {
536 #(#active_save_model_fields)*
537 #(#active_save_relation_defaults)*
538 };
539 return Ok(model);
540 }
541
542 Self { #(#active_destructure,)* }.insert(ex).await
543 }
544 );
545
546 let model_delete_impl = if !primary_keys.is_empty() {
547 let pk_filters = primary_keys
548 .iter()
549 .map(|(ident, _, _)| quote!(delete = delete.filter(Self::#ident.eq(#ident));));
550 quote!(
551 impl #impl_generics ::dbkit::ModelDelete for #model_ident #struct_type_args {
552 fn delete<'e, E>(self, ex: &'e E) -> ::dbkit::executor::BoxFuture<'e, Result<u64, ::dbkit::Error>>
553 where
554 E: ::dbkit::Executor + Send + Sync + 'e,
555 {
556 let Self { #(#pk_idents,)* .. } = self;
557 let mut delete = ::dbkit::Delete::new(Self::TABLE);
558 #(#pk_filters)*
559 ::dbkit::DeleteExt::execute(delete, ex)
560 }
561 }
562 )
563 } else {
564 quote!()
565 };
566
567 let into_active_fn = quote!(
568 pub fn into_active(self) -> #active_ident {
569 let Self { #(#active_destructure,)* .. } = self;
570 #active_ident {
571 #(#active_from_model,)*
572 }
573 }
574 );
575
576 let primary_key_const = if primary_keys.len() == 1 {
577 let (_, ty, name) = primary_keys.first().expect("primary key length checked");
578 Some(quote!(pub const PRIMARY_KEY: ::dbkit::Column<#struct_ident, #ty> = ::dbkit::Column::new(Self::TABLE, #name);))
579 } else {
580 None
581 };
582
583 let by_id_fn = if primary_keys.len() == 1 {
584 let (ident, ty, _) = primary_keys.first().expect("primary key length checked");
585 Some(quote!(
586 pub fn by_id(id: #ty) -> ::dbkit::Select<#struct_ident> {
587 Self::query().filter(Self::#ident.eq(id)).limit(1)
588 }
589 ))
590 } else {
591 None
592 };
593
594 let any_state_ident = format_ident!("{}AnyState", struct_ident);
595
596 let relation_state_modules = relation_fields.iter().map(|rel| {
597 let state_mod = &rel.state_mod_ident;
598 let (sealed_impl, state_impl) = match rel.kind {
599 RelationKind::HasMany | RelationKind::ManyToMany => (
600 quote!(
601 impl<T> Sealed for Vec<T> {}
602 ),
603 quote!(
604 impl<T> State for Vec<T> {}
605 ),
606 ),
607 RelationKind::BelongsTo => (
608 quote!(
609 impl<T> Sealed for Option<T> {}
610 ),
611 quote!(
612 impl<T> State for Option<T> {}
613 ),
614 ),
615 };
616 quote!(
617 pub mod #state_mod {
618 mod sealed {
619 pub trait Sealed {}
620 impl Sealed for ::dbkit::NotLoaded {}
621 #sealed_impl
622 }
623 pub trait State: sealed::Sealed {}
624 impl State for ::dbkit::NotLoaded {}
625 #state_impl
626 }
627 )
628 });
629
630 let relation_methods = relation_fields.iter().map(|rel| {
631 let field_ident = rel.field.ident.as_ref().expect("field ident");
632 let method_ident = format_ident!("{}_loaded", field_ident);
633 let item_ident = format_ident!("{}Item", to_camel_case(&field_ident.to_string()));
634 let loaded_type: Type = match rel.kind {
635 RelationKind::HasMany | RelationKind::ManyToMany => syn::parse_quote!(Vec<#item_ident>),
636 RelationKind::BelongsTo => syn::parse_quote!(Option<#item_ident>),
637 };
638
639 let mut other_params = Vec::new();
640 let mut type_params = Vec::new();
641 for other in &relation_fields {
642 if other.field.ident == rel.field.ident {
643 type_params.push(quote!(#loaded_type));
644 } else {
645 let ident = &other.param_ident;
646 let state_mod = &other.state_mod_ident;
647 other_params.push(quote!(#ident: #state_mod::State));
648 type_params.push(quote!(#ident));
649 }
650 }
651
652 let mut impl_params = Vec::new();
653 impl_params.push(quote!(#item_ident));
654 impl_params.extend(other_params);
655
656 let impl_generics = if impl_params.is_empty() {
657 quote!()
658 } else {
659 quote!(<#(#impl_params),*>)
660 };
661 let type_args = if type_params.is_empty() {
662 quote!()
663 } else {
664 quote!(<#(#type_params),*>)
665 };
666
667 let (return_ty, body) = match rel.kind {
668 RelationKind::HasMany | RelationKind::ManyToMany => (quote!(&[#item_ident]), quote!(&self.#field_ident)),
669 RelationKind::BelongsTo => (quote!(Option<&#item_ident>), quote!(self.#field_ident.as_ref())),
670 };
671
672 quote!(
673 impl #impl_generics #model_ident #type_args {
674 pub fn #method_ident(&self) -> #return_ty {
675 #body
676 }
677 }
678 )
679 });
680
681 let model_value_arms = output_fields
682 .iter()
683 .filter(|field| !is_relation_field(field, &relation_fields))
684 .map(|field| {
685 let ident = field.ident.as_ref().expect("field ident");
686 let name = scalar_column_name(&scalar_fields, ident).expect("scalar field column name");
687 quote!(#name => Some(self.#ident.clone().into()),)
688 });
689
690 let model_value_impl = quote!(
691 impl #impl_generics ::dbkit::ModelValue for #model_ident #struct_type_args {
692 fn column_value(&self, column: ::dbkit::ColumnRef) -> Option<::dbkit::Value> {
693 if column.table.name != Self::TABLE.name {
694 return None;
695 }
696 match column.name {
697 #(#model_value_arms)*
698 _ => None,
699 }
700 }
701 }
702 );
703
704 let from_row_generics = relation_fields.iter().map(|rel| {
705 let ident = &rel.param_ident;
706 let state_mod = &rel.state_mod_ident;
707 quote!(#ident: #state_mod::State + Default)
708 });
709
710 let from_row_impl_generics = if relation_fields.is_empty() {
711 quote!(<'r>)
712 } else {
713 quote!(<'r, #(#from_row_generics),*>)
714 };
715
716 let from_row_fields = output_fields.iter().map(|field| {
717 let ident = field.ident.as_ref().expect("field ident");
718 if is_relation_field(field, &relation_fields) {
719 quote!(#ident: Default::default())
720 } else {
721 let name = scalar_column_name(&scalar_fields, ident).expect("scalar field column name");
722 quote!(#ident: ::dbkit::sqlx::Row::try_get(row, #name)?)
723 }
724 });
725
726 let from_row_impl = quote!(
727 impl #from_row_impl_generics ::dbkit::sqlx::FromRow<'r, ::dbkit::sqlx::postgres::PgRow>
728 for #model_ident #struct_type_args
729 {
730 fn from_row(row: &'r ::dbkit::sqlx::postgres::PgRow) -> Result<Self, ::dbkit::sqlx::Error> {
731 Ok(Self {
732 #(#from_row_fields,)*
733 })
734 }
735 }
736 );
737
738 let joined_from_row_fields = output_fields.iter().map(|field| {
739 let ident = field.ident.as_ref().expect("field ident");
740 if is_relation_field(field, &relation_fields) {
741 quote!(#ident: Default::default())
742 } else {
743 let name = scalar_column_name(&scalar_fields, ident).expect("scalar field column name");
744 quote!(
745 #ident: {
746 let column = format!("{}{}", prefix, #name);
747 ::dbkit::sqlx::Row::try_get(row, column.as_str())?
748 }
749 )
750 }
751 });
752
753 let joined_pk_checks = if primary_keys.is_empty() {
754 if let Some(first_field) = scalar_fields.first() {
755 let name = &first_field.column_name;
756 let ty = option_inner_type(&first_field.ty).unwrap_or_else(|| first_field.ty.clone());
757 quote!(
758 let value: Option<#ty> = {
759 let column = format!("{}{}", prefix, #name);
760 ::dbkit::sqlx::Row::try_get(row, column.as_str())?
761 };
762 Ok(value.is_some())
763 )
764 } else {
765 quote!(Ok(false))
766 }
767 } else {
768 let checks = primary_keys.iter().map(|(_, ty, name)| {
769 let ty = option_inner_type(ty).unwrap_or_else(|| ty.clone());
770 quote!(
771 let value: Option<#ty> = {
772 let column = format!("{}{}", prefix, #name);
773 ::dbkit::sqlx::Row::try_get(row, column.as_str())?
774 };
775 if value.is_some() {
776 return Ok(true);
777 }
778 )
779 });
780 quote!(
781 #(#checks)*
782 Ok(false)
783 )
784 };
785
786 let joined_model_impl = quote!(
787 impl #from_row_impl_generics ::dbkit::JoinedModel for #model_ident #struct_type_args {
788 fn joined_columns() -> &'static [::dbkit::ColumnRef] {
789 Self::COLUMNS
790 }
791
792 fn joined_primary_keys() -> &'static [::dbkit::ColumnRef] {
793 Self::PRIMARY_KEYS
794 }
795
796 fn joined_from_row_prefixed(
797 row: &::dbkit::sqlx::postgres::PgRow,
798 prefix: &str,
799 ) -> Result<Self, ::dbkit::sqlx::Error> {
800 Ok(Self {
801 #(#joined_from_row_fields,)*
802 })
803 }
804
805 fn joined_row_has_pk(
806 row: &::dbkit::sqlx::postgres::PgRow,
807 prefix: &str,
808 ) -> Result<bool, ::dbkit::sqlx::Error> {
809 #joined_pk_checks
810 }
811 }
812 );
813
814 let set_relation_impls = relation_fields.iter().map(|rel| {
815 let field_ident = rel.field.ident.as_ref().expect("field ident");
816 let child_type = &rel.child_type;
817 let item_ident = format_ident!("{}Item", to_camel_case(&field_ident.to_string()));
818 let (value_ty, rel_ty) = match rel.kind {
819 RelationKind::HasMany => (quote!(Vec<#item_ident>), quote!(::dbkit::rel::HasMany<#struct_ident, #child_type>)),
820 RelationKind::ManyToMany => {
821 let through = rel.many_to_many_through.as_ref().expect("many-to-many through");
822 (
823 quote!(Vec<#item_ident>),
824 quote!(::dbkit::rel::ManyToMany<#struct_ident, #child_type, #through>),
825 )
826 }
827 RelationKind::BelongsTo => (
828 quote!(Option<#item_ident>),
829 quote!(::dbkit::rel::BelongsTo<#struct_ident, #child_type>),
830 ),
831 };
832
833 let mut other_params = Vec::new();
834 let mut type_params = Vec::new();
835 for other in &relation_fields {
836 if other.field.ident == rel.field.ident {
837 type_params.push(value_ty.clone());
838 } else {
839 let ident = &other.param_ident;
840 let state_mod = &other.state_mod_ident;
841 other_params.push(quote!(#ident: #state_mod::State));
842 type_params.push(quote!(#ident));
843 }
844 }
845
846 let mut impl_params = Vec::new();
847 impl_params.push(quote!(#item_ident));
848 impl_params.extend(other_params);
849
850 let impl_generics = if impl_params.is_empty() {
851 quote!()
852 } else {
853 quote!(<#(#impl_params),*>)
854 };
855 let type_args = if type_params.is_empty() {
856 quote!()
857 } else {
858 quote!(<#(#type_params),*>)
859 };
860
861 quote!(
862 impl #impl_generics ::dbkit::SetRelation<#rel_ty, #value_ty> for #model_ident #type_args {
863 fn set_relation(&mut self, _rel: #rel_ty, value: #value_ty) -> Result<(), ::dbkit::Error> {
864 self.#field_ident = value;
865 Ok(())
866 }
867 }
868 )
869 });
870
871 let get_relation_impls = relation_fields.iter().map(|rel| {
872 let field_ident = rel.field.ident.as_ref().expect("field ident");
873 let child_type = &rel.child_type;
874 let item_ident = format_ident!("{}Item", to_camel_case(&field_ident.to_string()));
875 let (value_ty, rel_ty) = match rel.kind {
876 RelationKind::HasMany => (quote!(Vec<#item_ident>), quote!(::dbkit::rel::HasMany<#struct_ident, #child_type>)),
877 RelationKind::ManyToMany => {
878 let through = rel.many_to_many_through.as_ref().expect("many-to-many through");
879 (
880 quote!(Vec<#item_ident>),
881 quote!(::dbkit::rel::ManyToMany<#struct_ident, #child_type, #through>),
882 )
883 }
884 RelationKind::BelongsTo => (
885 quote!(Option<#item_ident>),
886 quote!(::dbkit::rel::BelongsTo<#struct_ident, #child_type>),
887 ),
888 };
889
890 let mut other_params = Vec::new();
891 let mut type_params = Vec::new();
892 for other in &relation_fields {
893 if other.field.ident == rel.field.ident {
894 type_params.push(value_ty.clone());
895 } else {
896 let ident = &other.param_ident;
897 let state_mod = &other.state_mod_ident;
898 other_params.push(quote!(#ident: #state_mod::State));
899 type_params.push(quote!(#ident));
900 }
901 }
902
903 let mut impl_params = Vec::new();
904 impl_params.push(quote!(#item_ident));
905 impl_params.extend(other_params);
906
907 let impl_generics = if impl_params.is_empty() {
908 quote!()
909 } else {
910 quote!(<#(#impl_params),*>)
911 };
912 let type_args = if type_params.is_empty() {
913 quote!()
914 } else {
915 quote!(<#(#type_params),*>)
916 };
917
918 quote!(
919 impl #impl_generics ::dbkit::GetRelation<#rel_ty, #value_ty> for #model_ident #type_args {
920 fn get_relation(&self, _rel: #rel_ty) -> Option<&#value_ty> {
921 Some(&self.#field_ident)
922 }
923
924 fn get_relation_mut(&mut self, _rel: #rel_ty) -> Option<&mut #value_ty> {
925 Some(&mut self.#field_ident)
926 }
927 }
928 )
929 });
930
931 let load_method = quote!(
932 pub async fn load<Rel>(
933 self,
934 rel: Rel,
935 ex: &(impl ::dbkit::Executor + Send + Sync),
936 ) -> Result<<Self as ::dbkit::LoadRelation<Rel>>::Out, ::dbkit::Error>
937 where
938 Self: ::dbkit::LoadRelation<Rel>,
939 {
940 ::dbkit::LoadRelation::load_relation(self, rel, ex).await
941 }
942 );
943
944 let load_relation_impls = relation_fields.iter().map(|rel| {
945 let field_ident = rel.field.ident.as_ref().expect("field ident");
946 let child_type = &rel.child_type;
947 let rel_type = match rel.kind {
948 RelationKind::HasMany => quote!(::dbkit::rel::HasMany<#struct_ident, #child_type>),
949 RelationKind::BelongsTo => quote!(::dbkit::rel::BelongsTo<#struct_ident, #child_type>),
950 RelationKind::ManyToMany => {
951 let through = rel.many_to_many_through.as_ref().expect("many-to-many through");
952 quote!(::dbkit::rel::ManyToMany<#struct_ident, #child_type, #through>)
953 }
954 };
955 let loaded_type = match rel.kind {
956 RelationKind::HasMany | RelationKind::ManyToMany => quote!(Vec<#child_type>),
957 RelationKind::BelongsTo => quote!(Option<#child_type>),
958 };
959 let loader_fn = match rel.kind {
960 RelationKind::HasMany => quote!(::dbkit::runtime::load_selectin_has_many),
961 RelationKind::ManyToMany => quote!(::dbkit::runtime::load_selectin_many_to_many),
962 RelationKind::BelongsTo => quote!(::dbkit::runtime::load_selectin_belongs_to),
963 };
964
965 let mut other_params = Vec::new();
966 let mut type_params = Vec::new();
967 let mut out_params = Vec::new();
968 for other in &relation_fields {
969 if other.field.ident == rel.field.ident {
970 type_params.push(quote!(::dbkit::NotLoaded));
971 out_params.push(loaded_type.clone());
972 } else {
973 let ident = &other.param_ident;
974 let state_mod = &other.state_mod_ident;
975 other_params.push(quote!(#ident: #state_mod::State + Send + 'static));
976 type_params.push(quote!(#ident));
977 out_params.push(quote!(#ident));
978 }
979 }
980
981 let impl_generics = if other_params.is_empty() {
982 quote!()
983 } else {
984 quote!(<#(#other_params),*>)
985 };
986 let type_args = if type_params.is_empty() {
987 quote!()
988 } else {
989 quote!(<#(#type_params),*>)
990 };
991 let out_type = if out_params.is_empty() {
992 quote!(#model_ident)
993 } else {
994 quote!(#model_ident<#(#out_params),*>)
995 };
996 let out_construct = if out_params.is_empty() {
997 quote!(#model_ident)
998 } else {
999 quote!(#model_ident::<#(#out_params),*>)
1000 };
1001
1002 let destructure_fields = output_fields.iter().map(|field| {
1003 let ident = field.ident.as_ref().expect("field ident");
1004 if ident == field_ident {
1005 quote!(#ident: _)
1006 } else {
1007 quote!(#ident)
1008 }
1009 });
1010
1011 let build_fields = output_fields.iter().map(|field| {
1012 let ident = field.ident.as_ref().expect("field ident");
1013 if ident == field_ident {
1014 quote!(#ident: Default::default())
1015 } else {
1016 quote!(#ident)
1017 }
1018 });
1019
1020 quote!(
1021 impl #impl_generics ::dbkit::LoadRelation<#rel_type> for #model_ident #type_args {
1022 type Out = #out_type;
1023
1024 fn load_relation<'e, E>(
1025 self,
1026 rel: #rel_type,
1027 ex: &'e E,
1028 ) -> ::dbkit::executor::BoxFuture<'e, Result<Self::Out, ::dbkit::Error>>
1029 where
1030 E: ::dbkit::Executor + Send + Sync + 'e,
1031 {
1032 Box::pin(async move {
1033 let Self { #(#destructure_fields,)* } = self;
1034 let mut out = #out_construct {
1035 #(#build_fields,)*
1036 };
1037 let mut rows = vec![out];
1038 #loader_fn(ex, &mut rows, rel, &::dbkit::load::NoLoad).await?;
1039 Ok(rows.pop().expect("loaded row"))
1040 })
1041 }
1042 }
1043 )
1044 });
1045
1046 let relation_consts = relation_fields.iter().filter_map(|rel| {
1047 let field_ident = rel.field.ident.as_ref().expect("field ident");
1048 let child_type = &rel.child_type;
1049 match rel.kind {
1050 RelationKind::HasMany => Some(quote!(
1051 pub const #field_ident: ::dbkit::rel::HasMany<#struct_ident, #child_type> =
1052 ::dbkit::rel::HasMany::new(
1053 <#child_type as ::dbkit::rel::BelongsToSpec<#struct_ident>>::PARENT_TABLE,
1054 <#child_type as ::dbkit::rel::BelongsToSpec<#struct_ident>>::CHILD_TABLE,
1055 <#child_type as ::dbkit::rel::BelongsToSpec<#struct_ident>>::PARENT_KEY,
1056 <#child_type as ::dbkit::rel::BelongsToSpec<#struct_ident>>::CHILD_KEY,
1057 );
1058 )),
1059 RelationKind::BelongsTo => {
1060 let key = rel.belongs_to_key.as_ref().expect("belongs_to key");
1061 let references = rel.belongs_to_ref.as_ref().expect("belongs_to references");
1062 Some(quote!(
1063 pub const #field_ident: ::dbkit::rel::BelongsTo<#struct_ident, #child_type> =
1064 ::dbkit::rel::BelongsTo::new(
1065 Self::TABLE,
1066 #child_type::TABLE,
1067 Self::#key.as_ref(),
1068 #child_type::#references.as_ref(),
1069 );
1070 ))
1071 }
1072 RelationKind::ManyToMany => {
1073 let through = rel.many_to_many_through.as_ref().expect("many-to-many through");
1074 let left_key = rel.many_to_many_left_key.as_ref().expect("many-to-many left_key");
1075 let right_key = rel.many_to_many_right_key.as_ref().expect("many-to-many right_key");
1076 let parent_pk = primary_keys.first().map(|(ident, _, _)| ident).expect("many-to-many parent pk");
1077 Some(quote!(
1078 pub const #field_ident: ::dbkit::rel::ManyToMany<#struct_ident, #child_type, #through> =
1079 ::dbkit::rel::ManyToMany::new(
1080 Self::TABLE,
1081 #child_type::TABLE,
1082 #through::TABLE,
1083 Self::#parent_pk.as_ref(),
1084 #child_type::PRIMARY_KEY.as_ref(),
1085 #through::#left_key.as_ref(),
1086 #through::#right_key.as_ref(),
1087 );
1088 ))
1089 }
1090 }
1091 });
1092
1093 let belongs_to_specs = relation_fields.iter().filter_map(|rel| {
1094 if rel.kind != RelationKind::BelongsTo {
1095 return None;
1096 }
1097 let parent_type = &rel.child_type;
1098 let key = rel.belongs_to_key.as_ref().expect("belongs_to key");
1099 let references = rel.belongs_to_ref.as_ref().expect("belongs_to references");
1100 Some(quote!(
1101 impl #impl_generics ::dbkit::rel::BelongsToSpec<#parent_type> for #model_ident #struct_type_args {
1102 const CHILD_TABLE: ::dbkit::Table = Self::TABLE;
1103 const PARENT_TABLE: ::dbkit::Table = #parent_type::TABLE;
1104 const CHILD_KEY: ::dbkit::ColumnRef = Self::#key.as_ref();
1105 const PARENT_KEY: ::dbkit::ColumnRef = #parent_type::#references.as_ref();
1106 }
1107 ))
1108 });
1109
1110 let apply_load_impls = relation_fields.iter().flat_map(|rel| {
1111 let child_type = &rel.child_type;
1112 let rel_type = match rel.kind {
1113 RelationKind::HasMany => quote!(::dbkit::rel::HasMany<#struct_ident, #child_type>),
1114 RelationKind::BelongsTo => quote!(::dbkit::rel::BelongsTo<#struct_ident, #child_type>),
1115 RelationKind::ManyToMany => {
1116 let through = rel.many_to_many_through.as_ref().expect("many-to-many through");
1117 quote!(::dbkit::rel::ManyToMany<#struct_ident, #child_type, #through>)
1118 }
1119 };
1120
1121 let loaded_child = quote!(<Nested as ::dbkit::load::ApplyLoad<#child_type>>::Out2);
1122 let loaded_param = match rel.kind {
1123 RelationKind::HasMany | RelationKind::ManyToMany => quote!(Vec<#loaded_child>),
1124 RelationKind::BelongsTo => quote!(Option<#loaded_child>),
1125 };
1126
1127 let mut out_params = Vec::new();
1128 for other in &relation_fields {
1129 if other.field.ident == rel.field.ident {
1130 out_params.push(loaded_param.clone());
1131 } else {
1132 let ident = &other.param_ident;
1133 out_params.push(quote!(#ident));
1134 }
1135 }
1136
1137 let model_type = if generic_idents.is_empty() {
1138 quote!(#model_ident)
1139 } else {
1140 quote!(#model_ident<#(#generic_idents),*>)
1141 };
1142 let out_type = if out_params.is_empty() {
1143 quote!(#model_ident)
1144 } else {
1145 quote!(#model_ident<#(#out_params),*>)
1146 };
1147
1148 let mut apply_generics = Vec::new();
1149 apply_generics.push(quote!(Nested));
1150 apply_generics.extend(impl_generics_params.iter().cloned());
1151 let apply_generics = if apply_generics.is_empty() {
1152 quote!()
1153 } else {
1154 quote!(<#(#apply_generics),*>)
1155 };
1156
1157 let mut items = Vec::new();
1158 for strategy in ["SelectIn", "Joined"] {
1159 let load_ty = if strategy == "SelectIn" {
1160 quote!(::dbkit::load::SelectIn<#rel_type, Nested>)
1161 } else {
1162 quote!(::dbkit::load::Joined<#rel_type, Nested>)
1163 };
1164 items.push(quote!(
1165 impl #apply_generics ::dbkit::load::ApplyLoad<#model_type> for #load_ty
1166 where
1167 Nested: ::dbkit::load::ApplyLoad<#child_type>,
1168 {
1169 type Out2 = #out_type;
1170 }
1171 ));
1172 }
1173 items.into_iter()
1174 });
1175
1176 let run_load_impls = relation_fields.iter().flat_map(|rel| {
1177 let child_type = &rel.child_type;
1178 let through = rel.many_to_many_through.as_ref();
1179 let rel_type = match rel.kind {
1180 RelationKind::HasMany => quote!(::dbkit::rel::HasMany<#struct_ident, #child_type>),
1181 RelationKind::BelongsTo => quote!(::dbkit::rel::BelongsTo<#struct_ident, #child_type>),
1182 RelationKind::ManyToMany => {
1183 let through = through.expect("many-to-many through");
1184 quote!(::dbkit::rel::ManyToMany<#struct_ident, #child_type, #through>)
1185 }
1186 };
1187
1188 let loaded_child = quote!(<Nested as ::dbkit::load::ApplyLoad<#child_type>>::Out2);
1189 let loaded_param = match rel.kind {
1190 RelationKind::HasMany | RelationKind::ManyToMany => quote!(Vec<#loaded_child>),
1191 RelationKind::BelongsTo => quote!(Option<#loaded_child>),
1192 };
1193
1194 let mut out_params = Vec::new();
1195 for other in &relation_fields {
1196 if other.field.ident == rel.field.ident {
1197 out_params.push(loaded_param.clone());
1198 } else {
1199 let ident = &other.param_ident;
1200 out_params.push(quote!(#ident));
1201 }
1202 }
1203
1204 let out_type = if out_params.is_empty() {
1205 quote!(#model_ident)
1206 } else {
1207 quote!(#model_ident<#(#out_params),*>)
1208 };
1209
1210 let mut apply_generics = Vec::new();
1211 apply_generics.push(quote!(Nested));
1212 for other in &relation_fields {
1213 if other.field.ident == rel.field.ident {
1214 continue;
1215 }
1216 let ident = &other.param_ident;
1217 let state_mod = &other.state_mod_ident;
1218 apply_generics.push(quote!(#ident: #state_mod::State + Send + 'static));
1219 }
1220 let apply_generics = if apply_generics.is_empty() {
1221 quote!()
1222 } else {
1223 quote!(<#(#apply_generics),*>)
1224 };
1225
1226 let (child_bounds, loader_fn) = match rel.kind {
1227 RelationKind::HasMany => (
1228 quote!(#loaded_child: ::dbkit::ModelValue + for<'r> ::dbkit::sqlx::FromRow<'r, ::dbkit::sqlx::postgres::PgRow> + Send + Unpin,),
1229 quote!(::dbkit::runtime::load_selectin_has_many),
1230 ),
1231 RelationKind::ManyToMany => {
1232 let through = through.expect("many-to-many through");
1233 (
1234 quote!(
1235 #loaded_child: ::dbkit::ModelValue + Clone + for<'r> ::dbkit::sqlx::FromRow<'r, ::dbkit::sqlx::postgres::PgRow> + Send + Unpin,
1236 #through: ::dbkit::ModelValue + for<'r> ::dbkit::sqlx::FromRow<'r, ::dbkit::sqlx::postgres::PgRow> + Send + Unpin,
1237 ),
1238 quote!(::dbkit::runtime::load_selectin_many_to_many),
1239 )
1240 }
1241 RelationKind::BelongsTo => (
1242 quote!(#loaded_child: ::dbkit::ModelValue + Clone + for<'r> ::dbkit::sqlx::FromRow<'r, ::dbkit::sqlx::postgres::PgRow> + Send + Unpin,),
1243 quote!(::dbkit::runtime::load_selectin_belongs_to),
1244 ),
1245 };
1246
1247 let joined_loader_fn = match rel.kind {
1248 RelationKind::HasMany => quote!(::dbkit::runtime::load_joined_has_many),
1249 RelationKind::ManyToMany => quote!(::dbkit::runtime::load_joined_many_to_many),
1250 RelationKind::BelongsTo => quote!(::dbkit::runtime::load_joined_belongs_to),
1251 };
1252
1253 let mut items = Vec::new();
1254 for (strategy, loader) in [
1255 ("SelectIn", loader_fn),
1256 ("Joined", joined_loader_fn),
1257 ] {
1258 let load_ty = if strategy == "SelectIn" {
1259 quote!(::dbkit::load::SelectIn<#rel_type, Nested>)
1260 } else {
1261 quote!(::dbkit::load::Joined<#rel_type, Nested>)
1262 };
1263 let out_bound = if strategy == "SelectIn" {
1264 quote!(::dbkit::ModelValue + ::dbkit::SetRelation<#rel_type, #loaded_param>)
1265 } else {
1266 quote!(::dbkit::GetRelation<#rel_type, #loaded_param>)
1267 };
1268
1269 items.push(quote!(
1270 impl #apply_generics ::dbkit::runtime::RunLoad<#out_type> for #load_ty
1271 where
1272 Nested: ::dbkit::load::ApplyLoad<#child_type> + ::dbkit::runtime::RunLoads<#loaded_child> + Sync,
1273 #out_type: #out_bound,
1274 #child_bounds
1275 {
1276 fn run<'e, E>(
1277 &'e self,
1278 ex: &'e E,
1279 rows: &'e mut [#out_type],
1280 ) -> ::dbkit::executor::BoxFuture<'e, Result<(), ::dbkit::Error>>
1281 where
1282 E: ::dbkit::Executor + Send + Sync + 'e,
1283 {
1284 #loader(ex, rows, self.rel.clone(), &self.nested)
1285 }
1286 }
1287 ));
1288 }
1289 items.into_iter()
1290 });
1291
1292 let output = quote! {
1293 #(#struct_attrs)*
1294 #[derive(Debug, Clone)]
1295 #vis struct #model_ident #struct_generics {
1296 #(#output_fields,)*
1297 }
1298
1299 #vis type #struct_ident = #model_ident;
1300
1301 #(#relation_state_modules)*
1302
1303 #vis trait #any_state_ident {}
1304 impl #impl_generics #any_state_ident for #model_ident #struct_type_args {}
1305
1306 impl #impl_generics #model_ident #struct_type_args {
1307 pub const TABLE: ::dbkit::Table = #table_expr;
1308 #(#columns)*
1309 #columns_const
1310 #primary_key_const
1311 #primary_keys_const
1312 #(#relation_consts)*
1313
1314 pub fn query() -> ::dbkit::Select<#struct_ident> {
1315 ::dbkit::Select::new(Self::TABLE)
1316 }
1317
1318 #by_id_fn
1319
1320 pub fn insert(values: #insert_ident) -> ::dbkit::Insert<#struct_ident> {
1321 let mut insert = ::dbkit::Insert::new(Self::TABLE);
1322 #(#insert_values)*
1323 insert
1324 }
1325
1326 pub fn insert_many(values: Vec<#insert_ident>) -> ::dbkit::Insert<#struct_ident> {
1327 let mut insert = ::dbkit::Insert::new(Self::TABLE);
1328 for value in values {
1329 insert = insert.row(|row| {
1330 let mut row = row;
1331 #(
1332 row = row.value(Self::#insert_field_idents, value.#insert_field_idents);
1333 )*
1334 row
1335 });
1336 }
1337 insert
1338 }
1339
1340 pub fn update() -> ::dbkit::Update<#struct_ident> {
1341 ::dbkit::Update::new(Self::TABLE)
1342 }
1343
1344 pub fn delete() -> ::dbkit::Delete {
1345 ::dbkit::Delete::new(Self::TABLE)
1346 }
1347
1348 pub fn new_active() -> #active_ident {
1349 #active_ident::new()
1350 }
1351
1352 #into_active_fn
1353 #load_method
1354 }
1355
1356 #[derive(Debug, Clone)]
1357 #vis struct #insert_ident {
1358 #(#insert_fields,)*
1359 }
1360
1361 #[derive(Debug, Clone, Default)]
1362 #vis struct #active_ident {
1363 #(#active_fields,)*
1364 }
1365
1366 impl #active_ident {
1367 pub fn new() -> Self {
1368 Self::default()
1369 }
1370
1371 #active_insert_fn
1372 #active_update_fn
1373 #active_delete_fn
1374 #active_save_fn
1375 }
1376
1377 #(#relation_methods)*
1378 #model_value_impl
1379 #from_row_impl
1380 #joined_model_impl
1381 #(#set_relation_impls)*
1382 #(#get_relation_impls)*
1383 #(#load_relation_impls)*
1384 #(#belongs_to_specs)*
1385 #(#apply_load_impls)*
1386 #(#run_load_impls)*
1387 #model_delete_impl
1388 };
1389
1390 Ok(output.into())
1391}
1392
1393fn parse_model_args(args: syn::punctuated::Punctuated<Meta, syn::Token![,]>) -> ModelArgs {
1394 let mut out = ModelArgs::default();
1395 for meta in args {
1396 if let Meta::NameValue(nv) = meta {
1397 if nv.path.is_ident("table") {
1398 if let Some(value) = extract_lit_str(&nv.value) {
1399 out.table = Some(value);
1400 }
1401 } else if nv.path.is_ident("schema") {
1402 if let Some(value) = extract_lit_str(&nv.value) {
1403 out.schema = Some(value);
1404 }
1405 }
1406 }
1407 }
1408 out
1409}
1410
1411fn parse_belongs_to_args(attrs: &[Attribute]) -> syn::Result<(Ident, Ident)> {
1412 for attr in attrs {
1413 if !attr.path().is_ident("belongs_to") {
1414 continue;
1415 }
1416 let args = attr.parse_args_with(syn::punctuated::Punctuated::<Meta, syn::Token![,]>::parse_terminated)?;
1417 let mut key = None;
1418 let mut references = None;
1419 for meta in args {
1420 if let Meta::NameValue(nv) = meta {
1421 if nv.path.is_ident("key") {
1422 key = extract_ident(&nv.value);
1423 } else if nv.path.is_ident("references") {
1424 references = extract_ident(&nv.value);
1425 }
1426 }
1427 }
1428 if let (Some(key), Some(references)) = (key, references) {
1429 return Ok((key, references));
1430 }
1431 }
1432 Err(syn::Error::new(
1433 proc_macro2::Span::call_site(),
1434 "dbkit: #[belongs_to] requires key = <field> and references = <field>",
1435 ))
1436}
1437
1438fn parse_many_to_many_args(attrs: &[Attribute]) -> syn::Result<(Ident, Ident, Ident)> {
1439 for attr in attrs {
1440 if !attr.path().is_ident("many_to_many") {
1441 continue;
1442 }
1443 let args = attr.parse_args_with(syn::punctuated::Punctuated::<Meta, syn::Token![,]>::parse_terminated)?;
1444 let mut through = None;
1445 let mut left_key = None;
1446 let mut right_key = None;
1447 for meta in args {
1448 if let Meta::NameValue(nv) = meta {
1449 if nv.path.is_ident("through") {
1450 through = extract_ident(&nv.value);
1451 } else if nv.path.is_ident("left_key") {
1452 left_key = extract_ident(&nv.value);
1453 } else if nv.path.is_ident("right_key") {
1454 right_key = extract_ident(&nv.value);
1455 }
1456 }
1457 }
1458 if let (Some(through), Some(left_key), Some(right_key)) = (through, left_key, right_key) {
1459 return Ok((through, left_key, right_key));
1460 }
1461 }
1462 Err(syn::Error::new(
1463 proc_macro2::Span::call_site(),
1464 "dbkit: #[many_to_many] requires through = <Model>, left_key = <field>, right_key = <field>",
1465 ))
1466}
1467
1468fn extract_lit_str(expr: &syn::Expr) -> Option<String> {
1469 if let syn::Expr::Lit(syn::ExprLit {
1470 lit: syn::Lit::Str(lit), ..
1471 }) = expr
1472 {
1473 Some(lit.value())
1474 } else {
1475 None
1476 }
1477}
1478
1479fn extract_ident(expr: &syn::Expr) -> Option<Ident> {
1480 if let syn::Expr::Path(path) = expr {
1481 path.path.get_ident().cloned()
1482 } else {
1483 None
1484 }
1485}
1486
1487fn parse_field_column_name(attrs: &[Attribute]) -> syn::Result<Option<String>> {
1488 let mut column_name = None;
1489 for attr in attrs {
1490 if !attr.path().is_ident("dbkit") {
1491 continue;
1492 }
1493 attr.parse_nested_meta(|meta| {
1494 if meta.path.is_ident("column") {
1495 if column_name.is_some() {
1496 return Err(meta.error("dbkit: duplicate field column rename"));
1497 }
1498 let lit: syn::LitStr = meta.value()?.parse()?;
1499 column_name = Some(lit.value());
1500 return Ok(());
1501 }
1502 Err(meta.error("dbkit: unsupported field option; expected `column`"))
1503 })?;
1504 }
1505 Ok(column_name)
1506}
1507
1508fn scalar_column_name<'a>(fields: &'a [ScalarFieldInfo], ident: &Ident) -> Option<&'a str> {
1509 fields
1510 .iter()
1511 .find(|field| field.ident == *ident)
1512 .map(|field| field.column_name.as_str())
1513}
1514
1515fn option_inner_type(ty: &Type) -> Option<Type> {
1516 let path = match ty {
1517 Type::Path(path) => path,
1518 _ => return None,
1519 };
1520 let segment = path.path.segments.last()?;
1521 if segment.ident != "Option" {
1522 return None;
1523 }
1524 let args = match &segment.arguments {
1525 syn::PathArguments::AngleBracketed(args) => args,
1526 _ => return None,
1527 };
1528 let inner = args.args.first()?;
1529 match inner {
1530 syn::GenericArgument::Type(inner_ty) => Some(inner_ty.clone()),
1531 _ => None,
1532 }
1533}
1534
1535fn has_attr(attrs: &[Attribute], name: &str) -> bool {
1536 attrs.iter().any(|attr| attr.path().is_ident(name))
1537}
1538
1539fn filter_struct_attrs(attrs: &[Attribute]) -> Vec<Attribute> {
1540 let mut kept = Vec::new();
1541 for attr in attrs {
1542 if is_model_attr(attr) {
1543 continue;
1544 }
1545 if attr.path().is_ident("derive") {
1546 if let Ok(mut paths) = attr.parse_args_with(syn::punctuated::Punctuated::<syn::Path, syn::Token![,]>::parse_terminated) {
1547 paths = paths
1548 .into_iter()
1549 .filter(|path| !path.segments.last().map(|seg| seg.ident == "Model").unwrap_or(false))
1550 .collect();
1551 if paths.is_empty() {
1552 continue;
1553 }
1554 let new_attr = quote!(#[derive(#paths)]);
1555 let parsed = syn::Attribute::parse_outer.parse2(new_attr).expect("derive attr");
1556 kept.extend(parsed);
1557 continue;
1558 }
1559 }
1560 kept.push(attr.clone());
1561 }
1562 kept
1563}
1564
1565fn filter_field_attrs(attrs: &[Attribute]) -> Vec<Attribute> {
1566 attrs.iter().filter(|attr| !is_field_orm_attr(attr)).cloned().collect()
1567}
1568
1569fn is_field_orm_attr(attr: &Attribute) -> bool {
1570 let name = attr.path().get_ident().map(|ident| ident.to_string());
1571 matches!(
1572 name.as_deref(),
1573 Some("key")
1574 | Some("autoincrement")
1575 | Some("unique")
1576 | Some("index")
1577 | Some("has_many")
1578 | Some("belongs_to")
1579 | Some("many_to_many")
1580 | Some("dbkit")
1581 )
1582}
1583
1584fn is_model_attr(attr: &Attribute) -> bool {
1585 attr.path().is_ident("model")
1586}
1587
1588fn relation_type(field: &Field) -> syn::Result<(RelationKind, Type)> {
1589 let kind = if has_attr(&field.attrs, "has_many") {
1590 RelationKind::HasMany
1591 } else if has_attr(&field.attrs, "belongs_to") {
1592 RelationKind::BelongsTo
1593 } else if has_attr(&field.attrs, "many_to_many") {
1594 RelationKind::ManyToMany
1595 } else {
1596 return Err(syn::Error::new_spanned(field, "dbkit: missing relation attribute"));
1597 };
1598
1599 let child_type = match &field.ty {
1600 Type::Path(path) => {
1601 let segment = path
1602 .path
1603 .segments
1604 .last()
1605 .ok_or_else(|| syn::Error::new_spanned(&field.ty, "dbkit: invalid type"))?;
1606 let expected = match kind {
1607 RelationKind::HasMany => "HasMany",
1608 RelationKind::BelongsTo => "BelongsTo",
1609 RelationKind::ManyToMany => "ManyToMany",
1610 };
1611 if segment.ident != expected {
1612 return Err(syn::Error::new_spanned(
1613 &segment.ident,
1614 format!("dbkit: expected {} marker type", expected),
1615 ));
1616 }
1617 match &segment.arguments {
1618 syn::PathArguments::AngleBracketed(args) => {
1619 let ty = args.args.iter().find_map(|arg| match arg {
1620 syn::GenericArgument::Type(ty) => Some(ty.clone()),
1621 _ => None,
1622 });
1623 ty.ok_or_else(|| syn::Error::new_spanned(&segment, "dbkit: missing type"))?
1624 }
1625 _ => return Err(syn::Error::new_spanned(&segment.arguments, "dbkit: expected generic argument")),
1626 }
1627 }
1628 _ => return Err(syn::Error::new_spanned(&field.ty, "dbkit: relation marker must be a type path")),
1629 };
1630
1631 Ok((kind, child_type))
1632}
1633
1634fn is_relation_field(field: &Field, rels: &[RelationInfo]) -> bool {
1635 rels.iter().any(|rel| rel.field.ident == field.ident)
1636}
1637
1638fn to_snake_case(name: &str) -> String {
1639 let chars: Vec<char> = name.chars().collect();
1640 let mut out = String::with_capacity(name.len() + (name.len() / 4));
1641
1642 for (idx, &ch) in chars.iter().enumerate() {
1643 let prev = idx.checked_sub(1).and_then(|i| chars.get(i)).copied();
1644 let next = chars.get(idx + 1).copied();
1645
1646 if ch.is_uppercase() {
1647 let prev_is_lower_or_digit = prev.map(|p| p.is_lowercase() || p.is_ascii_digit()).unwrap_or(false);
1648 let prev_is_upper = prev.map(|p| p.is_uppercase()).unwrap_or(false);
1649 let next_is_lower = next.map(|n| n.is_lowercase()).unwrap_or(false);
1650 let leading_upper_pair = idx == 1 && prev_is_upper && next_is_lower;
1651 let needs_separator = idx > 0 && (prev_is_lower_or_digit || (prev_is_upper && next_is_lower && !leading_upper_pair));
1652
1653 if needs_separator && !out.ends_with('_') {
1654 out.push('_');
1655 }
1656 for lower in ch.to_lowercase() {
1657 out.push(lower);
1658 }
1659 continue;
1660 }
1661
1662 out.push(ch);
1663 }
1664
1665 out
1666}
1667
1668fn to_camel_case(name: &str) -> String {
1669 let mut out = String::new();
1670 let mut uppercase_next = true;
1671 for ch in name.chars() {
1672 if ch == '_' {
1673 uppercase_next = true;
1674 continue;
1675 }
1676 if uppercase_next {
1677 for up in ch.to_uppercase() {
1678 out.push(up);
1679 }
1680 uppercase_next = false;
1681 } else {
1682 out.push(ch);
1683 }
1684 }
1685 out
1686}
1687
1688#[derive(Default)]
1693struct DbEnumArgs {
1694 type_name: Option<String>,
1695 rename_all: Option<String>,
1696}
1697
1698#[derive(Clone, Copy)]
1699enum DbEnumRenameAll {
1700 AsIs,
1701 SnakeCase,
1702 LowerCase,
1703 UpperCase,
1704 ScreamingSnakeCase,
1705}
1706
1707fn expand_db_enum(input: syn::ItemEnum) -> syn::Result<TokenStream> {
1708 if !input.generics.params.is_empty() {
1709 return Err(syn::Error::new_spanned(
1710 input.generics,
1711 "dbkit: #[derive(DbEnum)] does not support generics",
1712 ));
1713 }
1714
1715 let args = parse_db_enum_args(&input.attrs)?;
1716 let type_name = args
1717 .type_name
1718 .ok_or_else(|| syn::Error::new_spanned(&input.ident, "dbkit: DbEnum requires #[dbkit(type_name = \"...\")]"))?;
1719 let rename_rule = parse_db_enum_rename_all(args.rename_all.as_deref())?;
1720
1721 let enum_ident = input.ident.clone();
1722
1723 let mut as_db_arms = Vec::new();
1724 let mut from_db_arms = Vec::new();
1725 let mut expected_values = Vec::new();
1726 let mut seen_db_names: std::collections::BTreeMap<String, syn::Ident> = std::collections::BTreeMap::new();
1727
1728 for variant in input.variants.iter() {
1729 if !matches!(variant.fields, syn::Fields::Unit) {
1730 return Err(syn::Error::new_spanned(
1731 &variant.fields,
1732 "dbkit: DbEnum only supports unit variants",
1733 ));
1734 }
1735
1736 let variant_ident = &variant.ident;
1737 let explicit = parse_db_enum_variant_rename(&variant.attrs)?;
1738 let db_name = match explicit {
1739 Some(value) => value,
1740 None => apply_db_enum_rename_rule(&variant.ident.to_string(), rename_rule),
1741 };
1742 if let Some(first_variant) = seen_db_names.get(&db_name) {
1743 return Err(syn::Error::new_spanned(
1744 variant_ident,
1745 format!(
1746 "dbkit: duplicate DbEnum wire name `{}` for variants `{}` and `{}`",
1747 db_name, first_variant, variant_ident
1748 ),
1749 ));
1750 }
1751 seen_db_names.insert(db_name.clone(), variant_ident.clone());
1752 let db_name_lit = syn::LitStr::new(&db_name, variant.ident.span());
1753 expected_values.push(db_name);
1754
1755 as_db_arms.push(quote!(Self::#variant_ident => #db_name_lit,));
1756 from_db_arms.push(quote!(#db_name_lit => Ok(Self::#variant_ident),));
1757 }
1758
1759 if as_db_arms.is_empty() {
1760 return Err(syn::Error::new_spanned(enum_ident, "dbkit: DbEnum requires at least one variant"));
1761 }
1762
1763 let type_name_lit = syn::LitStr::new(&type_name, proc_macro2::Span::call_site());
1764 let expected_lit = syn::LitStr::new(&expected_values.join(", "), proc_macro2::Span::call_site());
1765
1766 let tokens = quote! {
1767 impl #enum_ident {
1768 pub const DB_TYPE_NAME: &'static str = #type_name_lit;
1769
1770 pub fn as_db_str(&self) -> &'static str {
1771 match self {
1772 #(#as_db_arms)*
1773 }
1774 }
1775 }
1776
1777 impl ::std::str::FromStr for #enum_ident {
1778 type Err = String;
1779
1780 fn from_str(value: &str) -> Result<Self, Self::Err> {
1781 match value {
1782 #(#from_db_arms)*
1783 _ => Err(format!(
1784 "dbkit: invalid value `{}` for enum {} (expected one of: {})",
1785 value,
1786 stringify!(#enum_ident),
1787 #expected_lit
1788 )),
1789 }
1790 }
1791 }
1792
1793 impl From<#enum_ident> for ::dbkit::Value {
1794 fn from(value: #enum_ident) -> Self {
1795 ::dbkit::Value::Enum {
1796 type_name: #type_name_lit,
1797 value: value.as_db_str().to_string(),
1798 }
1799 }
1800 }
1801
1802 impl ::dbkit::sqlx::Type<::dbkit::sqlx::Postgres> for #enum_ident {
1803 fn type_info() -> ::dbkit::sqlx::postgres::PgTypeInfo {
1804 ::dbkit::sqlx::postgres::PgTypeInfo::with_name(#type_name_lit)
1805 }
1806
1807 fn compatible(ty: &::dbkit::sqlx::postgres::PgTypeInfo) -> bool {
1808 *ty == ::dbkit::sqlx::postgres::PgTypeInfo::with_name(#type_name_lit)
1809 || <&str as ::dbkit::sqlx::Type<::dbkit::sqlx::Postgres>>::compatible(ty)
1810 }
1811 }
1812
1813 impl<'q> ::dbkit::sqlx::Encode<'q, ::dbkit::sqlx::Postgres> for #enum_ident {
1814 fn encode_by_ref(
1815 &self,
1816 buf: &mut ::dbkit::sqlx::postgres::PgArgumentBuffer,
1817 ) -> Result<::dbkit::sqlx::encode::IsNull, ::dbkit::sqlx::error::BoxDynError> {
1818 <&str as ::dbkit::sqlx::Encode<'q, ::dbkit::sqlx::Postgres>>::encode(self.as_db_str(), buf)
1819 }
1820
1821 fn produces(&self) -> Option<::dbkit::sqlx::postgres::PgTypeInfo> {
1822 Some(::dbkit::sqlx::postgres::PgTypeInfo::with_name(#type_name_lit))
1823 }
1824
1825 fn size_hint(&self) -> usize {
1826 self.as_db_str().len()
1827 }
1828 }
1829
1830 impl<'r> ::dbkit::sqlx::Decode<'r, ::dbkit::sqlx::Postgres> for #enum_ident {
1831 fn decode(value: ::dbkit::sqlx::postgres::PgValueRef<'r>) -> Result<Self, ::dbkit::sqlx::error::BoxDynError> {
1832 let value = <&str as ::dbkit::sqlx::Decode<'r, ::dbkit::sqlx::Postgres>>::decode(value)?;
1833 <Self as ::std::str::FromStr>::from_str(value).map_err(|err| err.into())
1834 }
1835 }
1836 };
1837
1838 Ok(TokenStream::from(tokens))
1839}
1840
1841fn parse_db_enum_args(attrs: &[Attribute]) -> syn::Result<DbEnumArgs> {
1842 let mut args = DbEnumArgs::default();
1843
1844 for attr in attrs {
1845 if !attr.path().is_ident("dbkit") {
1846 continue;
1847 }
1848 attr.parse_nested_meta(|meta| {
1849 if meta.path.is_ident("type_name") {
1850 let lit: syn::LitStr = meta.value()?.parse()?;
1851 args.type_name = Some(lit.value());
1852 return Ok(());
1853 }
1854 if meta.path.is_ident("rename_all") {
1855 let lit: syn::LitStr = meta.value()?.parse()?;
1856 args.rename_all = Some(lit.value());
1857 return Ok(());
1858 }
1859 Err(meta.error("dbkit: unsupported DbEnum option; expected `type_name` or `rename_all`"))
1860 })?;
1861 }
1862
1863 Ok(args)
1864}
1865
1866fn parse_db_enum_variant_rename(attrs: &[Attribute]) -> syn::Result<Option<String>> {
1867 let mut rename = None;
1868
1869 for attr in attrs {
1870 if !attr.path().is_ident("dbkit") {
1871 continue;
1872 }
1873 attr.parse_nested_meta(|meta| {
1874 if meta.path.is_ident("rename") {
1875 let lit: syn::LitStr = meta.value()?.parse()?;
1876 rename = Some(lit.value());
1877 return Ok(());
1878 }
1879 Err(meta.error("dbkit: unsupported DbEnum variant option; expected `rename`"))
1880 })?;
1881 }
1882
1883 Ok(rename)
1884}
1885
1886fn parse_db_enum_rename_all(value: Option<&str>) -> syn::Result<DbEnumRenameAll> {
1887 match value {
1888 None => Ok(DbEnumRenameAll::AsIs),
1889 Some("snake_case") => Ok(DbEnumRenameAll::SnakeCase),
1890 Some("lowercase") => Ok(DbEnumRenameAll::LowerCase),
1891 Some("UPPERCASE") => Ok(DbEnumRenameAll::UpperCase),
1892 Some("SCREAMING_SNAKE_CASE") => Ok(DbEnumRenameAll::ScreamingSnakeCase),
1893 Some(other) => Err(syn::Error::new(
1894 proc_macro2::Span::call_site(),
1895 format!(
1896 "dbkit: unsupported rename_all strategy `{}` for DbEnum; supported values: snake_case, lowercase, UPPERCASE, SCREAMING_SNAKE_CASE",
1897 other
1898 ),
1899 )),
1900 }
1901}
1902
1903fn apply_db_enum_rename_rule(value: &str, rule: DbEnumRenameAll) -> String {
1904 match rule {
1905 DbEnumRenameAll::AsIs => value.to_string(),
1906 DbEnumRenameAll::SnakeCase => to_snake_case(value),
1907 DbEnumRenameAll::LowerCase => value.to_lowercase(),
1908 DbEnumRenameAll::UpperCase => value.to_uppercase(),
1909 DbEnumRenameAll::ScreamingSnakeCase => to_snake_case(value).to_uppercase(),
1910 }
1911}
1912
1913#[cfg(test)]
1914mod tests {
1915 use super::{apply_db_enum_rename_rule, to_snake_case, DbEnumRenameAll};
1916
1917 #[test]
1918 fn snake_case_respects_acronym_word_boundaries() {
1919 assert_eq!(to_snake_case("HTTPWebhook"), "http_webhook");
1920 assert_eq!(to_snake_case("OAuthToken"), "oauth_token");
1921 assert_eq!(to_snake_case("XMLHttpRequest"), "xml_http_request");
1922 assert_eq!(to_snake_case("WebhookHTTP"), "webhook_http");
1923 }
1924
1925 #[test]
1926 fn screaming_snake_case_respects_acronym_word_boundaries() {
1927 assert_eq!(
1928 apply_db_enum_rename_rule("HTTPWebhook", DbEnumRenameAll::ScreamingSnakeCase),
1929 "HTTP_WEBHOOK"
1930 );
1931 assert_eq!(
1932 apply_db_enum_rename_rule("XMLHttpRequest", DbEnumRenameAll::ScreamingSnakeCase),
1933 "XML_HTTP_REQUEST"
1934 );
1935 }
1936}