ormer-derive 0.1.9

An ORM framework with a usage style similar to Linq, supporting Turso(SQLite), PostgresQL, MySQL
Documentation
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
use proc_macro2::TokenStream;
use quote::quote;
use syn::{DeriveInput, Expr, ExprLit, Lit, Meta};

pub fn derive_model(input: DeriveInput) -> TokenStream {
    let name = &input.ident;
    let where_name = syn::Ident::new(&format!("{name}Where"), name.span());

    // 提取表名
    let table_name = extract_table_name(&input);

    // 检查是否为元组结构体(用于包装现有模型)
    let is_tuple_struct = matches!(&input.data, syn::Data::Struct(data) if matches!(&data.fields, syn::Fields::Unnamed(_)));

    if is_tuple_struct {
        return derive_model_tuple_wrapper(&input, name, &where_name, table_name);
    }

    // 提取字段(普通命名字段结构体)
    let fields = match &input.data {
        syn::Data::Struct(data) => match &data.fields {
            syn::Fields::Named(fields) => &fields.named,
            _ => panic!("Model must have named fields or be a tuple struct wrapper"),
        },
        _ => panic!("Model must be a struct"),
    };

    // 提取主键字段列表(支持复合主键)
    let primary_keys: Vec<_> = fields
        .iter()
        .filter_map(|f| {
            for attr in &f.attrs {
                if attr.path().is_ident("primary") {
                    let field_name = f.ident.as_ref().unwrap().clone();
                    // 检查是否有 (auto) 参数
                    let is_auto = if let Meta::List(list) = &attr.meta {
                        list.tokens.to_string().contains("auto")
                    } else {
                        false
                    };
                    return Some((field_name, is_auto));
                }
            }
            None
        })
        .collect();

    // 至少需要一个主键
    if primary_keys.is_empty() {
        panic!("Model must have at least one #[primary] field");
    }

    // 检查是否有多个主键且标记了 auto(只有第一个主键可以是 auto)
    let auto_count = primary_keys.iter().filter(|(_, is_auto)| *is_auto).count();
    if auto_count > 1 {
        panic!("Only one primary key field can have #[primary(auto)]");
    }

    // 获取第一个主键(用于向后兼容)
    let primary_key_field = &primary_keys[0].0;
    let is_auto_increment = primary_keys[0].1;

    // 生成主键列名列表(支持复合主键)
    let primary_key_field_names: Vec<_> = primary_keys
        .iter()
        .map(|(field_name, _)| {
            quote! { stringify!(#field_name) }
        })
        .collect();

    // 生成主键值获取(支持复合主键)
    let primary_key_values: Vec<_> = primary_keys
        .iter()
        .map(|(field_name, _)| {
            quote! { ::ormer::Value::from(self.#field_name.clone()) }
        })
        .collect();

    // 生成字段名列表
    let field_names: Vec<String> = fields
        .iter()
        .map(|f| f.ident.as_ref().unwrap().to_string())
        .collect();

    let field_names_lit = field_names.iter().map(|name| {
        quote! { #name }
    });

    // 生成字段元数据 (COLUMN_SCHEMA)
    let column_schema_entries = fields.iter().map(|f| {
        let field_name = f.ident.as_ref().unwrap();
        let field_type = &f.ty;
        let type_str = quote! { #field_type }.to_string();

        // 检查是否是主键字段
        let is_primary = f.attrs.iter().any(|attr| attr.path().is_ident("primary"));

        // 检查是否是自增主键(只有主键字段才可能是自增)
        let field_is_auto_increment = if is_primary { is_auto_increment } else { false };

        // 检查是否是 Option<T>
        let is_nullable = type_str.starts_with("Option <");

        // 提取基础 Rust 类型
        let rust_type = if is_nullable {
            type_str
                .trim_start_matches("Option <")
                .trim_end_matches(">")
                .trim()
                .to_string()
        } else {
            type_str
        };

        // 检查 unique 属性
        let unique_group = extract_unique_group(f);

        // 检查 index 属性
        let is_indexed = f.attrs.iter().any(|attr| attr.path().is_ident("index"));

        // 检查 foreign 属性
        let foreign_key = extract_foreign_key(f);

        // 对于枚举类型支持,我们无法在编译时检测类型是否实现了 ModelEnum,
        // 因此采用简单策略:总是传递 None,数据库后端会根据 rust_type 字符串判断
        // 如果未来需要支持,可以考虑使用 specialization 或宏魔法
        let enum_variants = quote! {
            None
        };

        quote! {
            ::ormer::model::ColumnSchema {
                name: stringify!(#field_name),
                rust_type: #rust_type,
                is_primary: #is_primary,
                is_auto_increment: #field_is_auto_increment,
                is_nullable: #is_nullable,
                unique_group: #unique_group,
                is_indexed: #is_indexed,
                foreign_key: #foreign_key,
                enum_variants: #enum_variants,
            }
        }
    });

    // 生成 from_row 实现
    let from_row_fields = fields.iter().map(|f| {
        let field_name = f.ident.as_ref().unwrap();
        quote! {
            #field_name: row.get(stringify!(#field_name))?
        }
    });

    // 生成 from_row_values 实现(按顺序从行值中读取)
    let from_row_values_fields = fields.iter().enumerate().map(|(i, f)| {
        let field_name = f.ident.as_ref().unwrap();
        let field_type = &f.ty;
        quote! {
            #field_name: <#field_type as ::ormer::FromRowValues>::from_row_values(
                &values[#i..#i+1]
            )?
        }
    });

    // 生成 field_values 实现
    let field_names_for_values = fields.iter().map(|f| {
        let field_name = f.ident.as_ref().unwrap();
        quote! {
            ::ormer::Value::from(self.#field_name.clone())
        }
    });

    // 生成 Where 结构体的字段
    // 为所有字段生成类型化列代理
    let where_fields = fields.iter().map(|f| {
        let field_name = f.ident.as_ref().unwrap();
        let field_type = &f.ty;
        quote! {
            pub #field_name: ::ormer::query::builder::TypedColumn<#field_type>
        }
    });

    // 生成 Where 的 Default 实现
    let where_default_fields = fields.iter().map(|f| {
        let field_name = f.ident.as_ref().unwrap();
        quote! {
            #field_name: ::ormer::query::builder::TypedColumn::new(stringify!(#field_name))
        }
    });

    quote! {
        // 生成 Where 结构体
        pub struct #where_name {
            #(#where_fields),*
        }

        impl Default for #where_name {
            fn default() -> Self {
                Self {
                    #(#where_default_fields),*
                }
            }
        }

        impl ::ormer::Model for #name {
            const TABLE_NAME: &'static str = #table_name;
            const COLUMNS: &'static [&'static str] = &[#(#field_names_lit),*];
            const COLUMN_SCHEMA: &'static [::ormer::model::ColumnSchema] = &[#(#column_schema_entries),*];

            type QueryBuilder = ::ormer::Select<Self>;
            type Where = #where_name;

            fn query() -> Self::QueryBuilder {
                ::ormer::Select::new()
            }

            fn select() -> Self::QueryBuilder {
                ::ormer::Select::new()
            }

            fn from_row(row: &::ormer::Row) -> anyhow::Result<Self> {
                Ok(Self {
                    #(#from_row_fields),*
                })
            }

            fn from_row_values(values: &[::ormer::Value]) -> anyhow::Result<Self> {
                if values.len() < Self::COLUMNS.len() {
                    return Err(anyhow::anyhow!(
                        "Expected {} values for {}", Self::COLUMNS.len(), stringify!(#name)
                    ));
                }
                Ok(Self {
                    #(#from_row_values_fields),*
                })
            }

            fn field_values(&self) -> Vec<::ormer::Value> {
                vec![
                    #(#field_names_for_values),*
                ]
            }

            fn primary_key_columns() -> &'static [&'static str] {
                &[#(#primary_key_field_names),*]
            }

            fn primary_key_values(&self) -> Vec<::ormer::Value> {
                vec![#(#primary_key_values),*]
            }

            // 保持向后兼容的旧方法(已废弃)
            fn primary_key_column() -> &'static str {
                stringify!(#primary_key_field)
            }

            fn primary_key_value(&self) -> ::ormer::Value {
                ::ormer::Value::from(self.#primary_key_field.clone())
            }
        }

        // 生成 inherent 方法,使得不需要 import Model trait 也能调用
        impl #name {
            pub fn select() -> ::ormer::Select<Self> {
                ::ormer::Select::new()
            }

            pub fn query() -> ::ormer::Select<Self> {
                ::ormer::Select::new()
            }
        }
    }
}

/// 为元组结构体包装模型生成实现(例如:struct NewUser(User);)
fn derive_model_tuple_wrapper(
    input: &DeriveInput,
    name: &syn::Ident,
    _where_name: &syn::Ident,
    table_name: String,
) -> TokenStream {
    // 提取元组结构体中的内部类型
    let inner_type = match &input.data {
        syn::Data::Struct(data) => match &data.fields {
            syn::Fields::Unnamed(fields) => {
                if fields.unnamed.len() != 1 {
                    panic!("Tuple struct wrapper must have exactly one field");
                }
                &fields.unnamed[0].ty
            }
            _ => panic!("Expected unnamed fields"),
        },
        _ => panic!("Expected struct"),
    };

    // 生成代码:元组结构体包装器将委托给内部类型的所有 Model 功能,但使用自定义表名
    quote! {
        impl ::ormer::Model for #name {
            const TABLE_NAME: &'static str = #table_name;
            const COLUMNS: &'static [&'static str] = <#inner_type as ::ormer::Model>::COLUMNS;
            const COLUMN_SCHEMA: &'static [::ormer::model::ColumnSchema] = <#inner_type as ::ormer::Model>::COLUMN_SCHEMA;

            type QueryBuilder = ::ormer::Select<Self>;
            type Where = <#inner_type as ::ormer::Model>::Where;

            fn query() -> Self::QueryBuilder {
                ::ormer::Select::new()
            }

            fn select() -> Self::QueryBuilder {
                ::ormer::Select::new()
            }

            fn from_row(row: &::ormer::Row) -> anyhow::Result<Self> {
                let inner = <#inner_type as ::ormer::Model>::from_row(row)?;
                Ok(#name(inner))
            }

            fn from_row_values(values: &[::ormer::Value]) -> anyhow::Result<Self> {
                let inner = <#inner_type as ::ormer::Model>::from_row_values(values)?;
                Ok(#name(inner))
            }

            fn field_values(&self) -> Vec<::ormer::Value> {
                self.0.field_values()
            }

            fn primary_key_columns() -> &'static [&'static str] {
                <#inner_type as ::ormer::Model>::primary_key_columns()
            }

            fn primary_key_values(&self) -> Vec<::ormer::Value> {
                self.0.primary_key_values()
            }

            fn primary_key_column() -> &'static str {
                <#inner_type as ::ormer::Model>::primary_key_column()
            }

            fn primary_key_value(&self) -> ::ormer::Value {
                self.0.primary_key_value()
            }
        }

        // 生成 inherent 方法
        impl #name {
            pub fn select() -> ::ormer::Select<Self> {
                ::ormer::Select::new()
            }

            pub fn query() -> ::ormer::Select<Self> {
                ::ormer::Select::new()
            }
        }

        // 为包装器类型实现 Into<InnerType> 和 From<InnerType>
        impl From<#inner_type> for #name {
            fn from(inner: #inner_type) -> Self {
                #name(inner)
            }
        }

        impl #name {
            pub fn into_inner(self) -> #inner_type {
                self.0
            }

            pub fn inner(&self) -> &#inner_type {
                &self.0
            }
        }
    }
}

fn extract_table_name(input: &DeriveInput) -> String {
    // 查找 #[table = "name"] 属性
    for attr in &input.attrs {
        if attr.path().is_ident("table") {
            if let Meta::NameValue(meta) = &attr.meta {
                if let syn::Expr::Lit(expr) = &meta.value {
                    if let Lit::Str(lit) = &expr.lit {
                        return lit.value();
                    }
                }
            }
        }
    }

    // 默认使用结构体名的蛇形形式
    to_snake_case(&input.ident.to_string())
}

fn to_snake_case(s: &str) -> String {
    let mut result = String::new();
    for (i, c) in s.chars().enumerate() {
        if c.is_uppercase() {
            if i > 0 {
                result.push('_');
            }
            result.push(c.to_lowercase().next().unwrap());
        } else {
            result.push(c);
        }
    }
    result
}

/// 提取 unique 属性的 group 值
fn extract_unique_group(field: &syn::Field) -> proc_macro2::TokenStream {
    for attr in &field.attrs {
        if attr.path().is_ident("unique") {
            // 检查是否有 group 参数
            if let Meta::List(list) = &attr.meta {
                // 解析 tokens 查找 group = N
                let tokens_str = list.tokens.to_string();
                if tokens_str.contains("group") {
                    // 尝试提取 group 值
                    if let Ok(Meta::NameValue(meta)) = syn::parse2(list.tokens.clone()) {
                        if let Expr::Lit(ExprLit {
                            lit: Lit::Int(lit_int),
                            ..
                        }) = &meta.value
                        {
                            let group_value: i32 = lit_int.base10_parse().unwrap_or(0);
                            return quote! { Some(#group_value) };
                        }
                    }
                }
            }
            // 没有 group 参数,使用 0 作为默认组
            return quote! { Some(0) };
        }
    }
    // 没有 unique 属性
    quote! { None }
}

/// 提取 foreign 属性的外键信息
/// 支持两种语法:
/// - #[foreign(Type)] - 新语法,自动关联到目标 model 的主键
/// - #[foreign(Type.field)] - 旧语法,显式指定字段
fn extract_foreign_key(field: &syn::Field) -> proc_macro2::TokenStream {
    for attr in &field.attrs {
        if attr.path().is_ident("foreign") {
            if let Meta::List(list) = &attr.meta {
                let tokens_str = list.tokens.to_string();

                // 尝试解析为 Type.field 格式(旧语法)
                let parts: Vec<&str> = tokens_str.split('.').collect();
                if parts.len() == 2 {
                    let ref_type = parts[0].trim();
                    let ref_field = parts[1].trim();
                    let ref_type_ident = syn::Ident::new(ref_type, proc_macro2::Span::call_site());

                    // 使用目标模型的实际表名,而不是简单转换
                    return quote! {
                        Some(::ormer::model::ForeignKeyInfo {
                            ref_table: <#ref_type_ident as ::ormer::Model>::TABLE_NAME,
                            ref_column: #ref_field,
                            ref_column_fn: None,
                        })
                    };
                } else if parts.len() == 1 {
                    // 新语法:只传递类型,自动关联到目标 model 的主键
                    let ref_type = parts[0].trim();
                    let ref_type_ident = syn::Ident::new(ref_type, proc_macro2::Span::call_site());

                    // 使用函数指针在运行时获取目标模型的主键字段名(避免在常量上下文中调用非 const 函数)
                    // 创建一个辅助函数来返回主键列名
                    let pk_fn_name = syn::Ident::new(
                        &format!("__{}_primary_key_column", ref_type),
                        proc_macro2::Span::call_site(),
                    );

                    return quote! {
                        {
                            fn #pk_fn_name() -> &'static str {
                                <#ref_type_ident as ::ormer::Model>::primary_key_columns()[0]
                            }
                            Some(::ormer::model::ForeignKeyInfo {
                                ref_table: <#ref_type_ident as ::ormer::Model>::TABLE_NAME,
                                ref_column: "",
                                ref_column_fn: Some(#pk_fn_name),
                            })
                        }
                    };
                }
            }
        }
    }
    // 没有 foreign 属性
    quote! { None }
}