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