Skip to main content

remdb_macros/
lib.rs

1mod codegen;
2mod ddl_parser;
3
4use proc_macro::TokenStream;
5use quote::quote;
6use syn::parse_macro_input;
7
8#[proc_macro]
9pub fn define_schema(input: TokenStream) -> TokenStream {
10    let input = parse_macro_input!(input as syn::LitStr);
11    let schema = input.value();
12
13    match ddl_parser::parse_ddl(&schema) {
14        Ok(table_defs) => codegen::generate_code(table_defs),
15        Err(e) => {
16            panic!("Failed to parse DDL: {}", e);
17        }
18    }
19}
20
21#[proc_macro_derive(MemdbTable, attributes(memdb_schema))]
22pub fn derive_memdb_table(input: TokenStream) -> TokenStream {
23    let derive_input = parse_macro_input!(input as syn::DeriveInput);
24
25    // 查找memdb_schema属性
26    let mut ddl = String::new();
27
28    for attr in &derive_input.attrs {
29        if attr.path().is_ident("memdb_schema") {
30            // 使用正确的syn 2.0 API解析属性
31            attr.parse_nested_meta(|meta| {
32                if meta.path.is_ident("ddl") {
33                    let lit = meta.value()?;
34                    let lit_str = lit.parse::<syn::LitStr>()?;
35                    ddl = lit_str.value();
36                }
37                Ok(())
38            })
39            .unwrap();
40        }
41    }
42
43    if ddl.is_empty() {
44        panic!("memdb_schema attribute with ddl parameter is required");
45    }
46
47    // 解析DDL并生成代码
48    match ddl_parser::parse_ddl(&ddl) {
49        Ok(table_defs) => codegen::generate_code(table_defs),
50        Err(e) => {
51            panic!("Failed to parse DDL: {}", e);
52        }
53    }
54}
55
56use syn::parse::{Parse, ParseStream};
57use syn::{Ident, LitInt, Token};
58
59// 字段定义
60struct Field {
61    name: Ident,
62    #[allow(dead_code)]
63    colon: Token![:],
64    // 自定义类型解析,支持 str(32) 这种语法
65    type_name: Ident,
66    type_params: Option<LitInt>,
67}
68
69impl Parse for Field {
70    fn parse(input: ParseStream) -> syn::Result<Self> {
71        let name = input.parse()?;
72        let colon = input.parse()?;
73
74        // 解析类型名称
75        let type_name = input.parse()?;
76
77        // 检查是否有括号参数,如 str(32)
78        let type_params = if input.peek(syn::token::Paren) {
79            let content;
80            syn::parenthesized!(content in input);
81            let params = content.parse()?;
82            Some(params)
83        } else {
84            None
85        };
86
87        Ok(Self {
88            name,
89            colon,
90            type_name,
91            type_params,
92        })
93    }
94}
95
96// 表定义结构
97struct TableArgs {
98    name: Ident,
99    max_records: LitInt,
100    primary_key: Ident,
101    secondary_index: Option<Ident>,
102    secondary_index_type: Option<Ident>,
103    fields: Vec<Field>,
104}
105
106impl Parse for TableArgs {
107    fn parse(input: ParseStream) -> syn::Result<Self> {
108        // 解析表名
109        let name = input.parse()?;
110
111        // 解析逗号
112        let _comma1: Token![,] = input.parse()?;
113
114        // 解析最大记录数
115        let max_records = input.parse()?;
116
117        // 解析逗号
118        let _comma2: Token![,] = input.parse()?;
119
120        // 解析primary_key
121        let _primary_key_keyword: Ident = input.parse()?;
122        let _colon1: Token![:] = input.parse()?;
123        let primary_key = input.parse()?;
124
125        // 解析secondary_index(可选)
126        let mut secondary_index = None;
127        let mut secondary_index_type = None;
128
129        // 检查primary_key之后是否有逗号
130        if input.peek(Token![,]) {
131            let _comma3: Token![,] = input.parse()?;
132        }
133
134        // 解析secondary_index、secondary_index_type和fields关键字
135        loop {
136            // 检查下一个标记
137            let next = input.lookahead1();
138            if next.peek(Ident) {
139                let param_name = input.parse::<Ident>()?;
140                if param_name == "secondary_index" {
141                    let _colon: Token![:] = input.parse()?;
142                    secondary_index = Some(input.parse()?);
143
144                    // 解析逗号
145                    if input.peek(Token![,]) {
146                        let _comma4: Token![,] = input.parse()?;
147                    }
148                } else if param_name == "secondary_index_type" {
149                    let _colon: Token![:] = input.parse()?;
150                    secondary_index_type = Some(input.parse()?);
151
152                    // 解析逗号
153                    if input.peek(Token![,]) {
154                        let _comma5: Token![,] = input.parse()?;
155                    }
156                } else if param_name == "fields" {
157                    let _colon_fields: Token![:] = input.parse()?;
158                    break;
159                } else {
160                    return Err(syn::Error::new(param_name.span(), format!("expected 'secondary_index', 'secondary_index_type' or 'fields' keyword, got '{}'", param_name)));
161                }
162            } else {
163                return Err(next.error());
164            }
165        }
166
167        // 解析fields块
168        let content;
169        syn::braced!(content in input);
170
171        // 解析fields块内的内容
172        let mut fields = Vec::new();
173        while !content.is_empty() {
174            // 解析字段
175            let field = content.parse::<Field>()?;
176            fields.push(field);
177
178            // 如果还有逗号,解析它
179            if content.peek(Token![,]) {
180                content.parse::<Token![,]>()?;
181            }
182        }
183
184        Ok(Self {
185            name,
186            max_records,
187            primary_key,
188            secondary_index,
189            secondary_index_type,
190            fields,
191        })
192    }
193}
194
195// 数据库定义结构,解析数据库名和表列表
196struct DatabaseArgs {
197    name: Ident,
198    tables: Vec<Ident>,
199    low_power: bool,
200    low_power_max_records: Option<usize>,
201    default_max_records: usize,
202    total_memory: usize,
203}
204
205impl Parse for DatabaseArgs {
206    fn parse(input: ParseStream) -> syn::Result<Self> {
207        // 解析数据库名
208        let name = input.parse()?;
209
210        // 解析逗号
211        let _comma: Token![,] = input.parse()?;
212
213        // 解析tables关键字
214        let _tables: Ident = input.parse()?;
215
216        // 解析冒号
217        let _colon: Token![:] = input.parse()?;
218
219        // 解析表列表
220        let content;
221        syn::bracketed!(content in input);
222
223        let mut tables = Vec::new();
224        while !content.is_empty() {
225            // 解析表名
226            let table = content.parse::<Ident>()?;
227            tables.push(table);
228
229            // 如果还有逗号,解析它
230            if content.peek(Token![,]) {
231                content.parse::<Token![,]>()?;
232            }
233        }
234
235        // 解析可选的low_power参数
236        let mut low_power = false;
237        let mut low_power_max_records = None;
238        let mut default_max_records = 100000; // 默认值
239        let mut total_memory = 65536; // 默认64KB
240
241        // 检查是否还有更多参数
242        while !input.is_empty() {
243            // 解析逗号
244            let _comma: Token![,] = input.parse()?;
245
246            // 解析参数名
247            let param_name = input.parse::<Ident>()?;
248
249            // 解析冒号
250            let _colon: Token![:] = input.parse()?;
251
252            if param_name == "low_power" {
253                // 解析布尔值
254                let lit_bool = input.parse::<syn::LitBool>()?;
255                low_power = lit_bool.value;
256            } else if param_name == "low_power_max_records" {
257                // 解析数字
258                let lit_int = input.parse::<syn::LitInt>()?;
259                low_power_max_records = Some(lit_int.base10_parse().unwrap_or(0));
260            } else if param_name == "default_max_records" {
261                // 解析数字
262                let lit_int = input.parse::<syn::LitInt>()?;
263                default_max_records = lit_int.base10_parse().unwrap_or(100000);
264            } else if param_name == "total_memory" {
265                // 解析数字
266                let lit_int = input.parse::<syn::LitInt>()?;
267                total_memory = lit_int.base10_parse().unwrap_or(65536);
268            }
269        }
270
271        Ok(Self {
272            name,
273            tables,
274            low_power,
275            low_power_max_records,
276            default_max_records,
277            total_memory,
278        })
279    }
280}
281
282#[proc_macro]
283pub fn table(input: proc_macro::TokenStream) -> proc_macro::TokenStream {
284    // 解析输入参数
285    let args = parse_macro_input!(input as TableArgs);
286    let name = &args.name;
287    let max_records = &args.max_records;
288    let primary_key = &args.primary_key;
289    let secondary_index = &args.secondary_index;
290    let secondary_index_type = &args.secondary_index_type;
291    let fields = &args.fields;
292
293    // 生成字段定义
294    let mut offset = 0;
295    let mut field_defs = Vec::new();
296    let mut record_size = 0;
297    let mut primary_key_index = 0usize;
298    let mut secondary_key_index: Option<usize> = None;
299
300    for (i, field) in fields.iter().enumerate() {
301        let field_name = &field.name;
302        let type_name = &field.type_name;
303        let type_params = &field.type_params;
304
305        // 确定数据类型和大小
306        let (data_type, size_val, string_length) = if type_name == "i32" {
307            (quote!(remdb::types::DataType::Int32), 4, quote!(None))
308        } else if type_name == "i8" {
309            (quote!(remdb::types::DataType::Int8), 1, quote!(None))
310        } else if type_name == "u64" {
311            (quote!(remdb::types::DataType::UInt64), 8, quote!(None))
312        } else if type_name == "f64" {
313            (quote!(remdb::types::DataType::Float64), 8, quote!(None))
314        } else if type_name == "bool" {
315            (quote!(remdb::types::DataType::Bool), 1, quote!(None))
316        } else if type_name == "str" {
317            // 处理str(32)这样的类型
318            let str_size = if let Some(params) = type_params {
319                params.base10_parse().unwrap_or(32)
320            } else {
321                32
322            };
323            (
324                quote!(remdb::types::DataType::VarChar),
325                str_size,
326                quote!(Some(#str_size as usize)),
327            )
328        } else if type_name == "vector" {
329            // 处理vector(2)这样的向量类型
330            let dim = if let Some(params) = type_params {
331                params.base10_parse().unwrap_or(128)
332            } else {
333                128
334            };
335            (
336                quote!(remdb::types::DataType::Vector),
337                dim * 4,
338                quote!(None),
339            ) // 向量每个维度4字节
340        } else {
341            (quote!(remdb::types::DataType::Int32), 4, quote!(None))
342        };
343
344        // 计算对齐要求
345        let alignment = if type_name == "u64" || type_name == "f64" || type_name == "i64" {
346            8
347        } else if type_name == "i32" || type_name == "u32" || type_name == "f32" {
348            4
349        } else if type_name == "i16" || type_name == "u16" {
350            2
351        } else {
352            1
353        };
354
355        // 调整偏移量以满足对齐要求
356        offset = ((offset + alignment - 1) / alignment) * alignment;
357
358        // 确定约束字段值
359        let is_primary_key = field_name == primary_key;
360        let primary_key_val = is_primary_key;
361        let not_null_val = is_primary_key; // 主键字段默认为非空
362        let unique_val = is_primary_key;
363
364        // 检查是否为自增主键:
365        // 1. 整数主键默认自增
366        // 2. 可以显式指定AUTOINCREMENT
367        let is_integer_type =
368            type_name == "i32" || type_name == "i64" || type_name == "u32" || type_name == "u64";
369        let auto_increment_val = is_primary_key && is_integer_type;
370
371        // 生成向量元数据(仅向量类型字段需要)
372        let vector_metadata_code = if type_name == "vector" {
373            let dim = if let Some(params) = type_params {
374                params.base10_parse::<u16>().unwrap_or(128)
375            } else {
376                128u16
377            };
378            quote! {
379                Some(remdb::types::VectorMetadata {
380                    dimension: #dim,
381                    distance_type: remdb::types::DistanceType::L2,
382                    index_type: remdb::types::VectorIndexType::HNSW,
383                    compression_enabled: false,
384                    compression_scheme: 0,
385                    compression_level: 3,
386                    // HNSW默认参数
387                    hnsw_m: 16,
388                    hnsw_ef_construction: 200,
389                    hnsw_ef_search: 128,
390                    // IVF默认参数
391                    ivf_nlist: 1024,
392                    ivf_nprobe: 16,
393                })
394            }
395        } else {
396            quote! { None }
397        };
398
399        // 生成字段定义
400        let field_def = quote! {
401            remdb::types::FieldDef {
402                name: stringify!(#field_name).to_string(),
403                data_type: #data_type,
404                size: #size_val as usize, // 确保是usize类型
405                string_length: #string_length,
406                offset: #offset as usize,  // 确保是usize类型
407                primary_key: #primary_key_val,
408                not_null: #not_null_val,
409                unique: #unique_val,
410                auto_increment: #auto_increment_val,
411                default_value: None,
412                vector_metadata: #vector_metadata_code,
413                json_metadata: None,
414            }
415        };
416
417        field_defs.push(field_def);
418
419        // 确定主键和二级索引的字段索引
420        if field_name == primary_key {
421            primary_key_index = i;
422        }
423
424        if let Some(secondary_field) = secondary_index {
425            if field_name == secondary_field {
426                secondary_key_index = Some(i);
427            }
428        }
429
430        // 更新偏移量和记录大小
431        offset += size_val;
432        record_size = offset;
433    }
434
435    // 确保整个记录满足最大对齐要求(8字节对齐)
436    let max_alignment = 8;
437    record_size = ((record_size + max_alignment - 1) / max_alignment) * max_alignment;
438
439    // 将max_records转换为usize
440    let max_records_usize = max_records.base10_parse::<usize>().unwrap_or(100);
441
442    // 确定索引类型
443    let index_type = match secondary_index_type.as_ref() {
444        Some(ty) if ty == "btree" => quote!(remdb::types::IndexType::BTree),
445        Some(ty) if ty == "hash" => quote!(remdb::types::IndexType::Hash),
446        Some(ty) if ty == "ttree" => quote!(remdb::types::IndexType::TTree),
447        Some(ty) if ty == "sortedarray" => quote!(remdb::types::IndexType::SortedArray),
448        _ => quote!(remdb::types::IndexType::BTree),
449    };
450
451    // 生成secondary_index代码
452    let secondary_index_code = match secondary_key_index {
453        Some(index) => quote! { Some(vec![#index as usize]) },
454        None => quote! { None },
455    };
456
457    // 生成代码:返回一个静态TableDef变量,使用LazyLock延迟初始化
458    let output = quote! {
459        #[allow(non_upper_case_globals)]
460        pub static #name: std::sync::LazyLock<remdb::types::TableDef> = std::sync::LazyLock::new(|| {
461            remdb::types::TableDef {
462                id: 0,
463                name: stringify!(#name).to_string(),
464                fields: vec![#(#field_defs,)*],
465                primary_key: vec![#primary_key_index as usize],
466                secondary_index: #secondary_index_code,
467                secondary_index_type: #index_type,
468                record_size: #record_size as usize,
469                max_records: #max_records_usize,
470                version: 1,
471                created_at: 0,
472                updated_at: 0,
473            }
474        });
475    };
476
477    output.into()
478}
479
480#[proc_macro]
481pub fn database(input: proc_macro::TokenStream) -> proc_macro::TokenStream {
482    // 解析输入参数
483    let args = parse_macro_input!(input as DatabaseArgs);
484    let name = &args.name;
485    let tables = &args.tables;
486    let low_power = args.low_power;
487    let default_max_records = args.default_max_records;
488    let total_memory = args.total_memory;
489
490    // 处理low_power_max_records,转换为Option<usize>
491    let low_power_max_records = match args.low_power_max_records {
492        Some(val) => quote! { Some(#val) },
493        None => quote! { None },
494    };
495
496    // 生成代码:返回一个静态DbConfig变量,使用LazyLock延迟初始化
497    let output = quote! {
498        #[allow(non_upper_case_globals)]
499        pub static #name: std::sync::LazyLock<remdb::config::DbConfig> = std::sync::LazyLock::new(|| {
500            remdb::config::DbConfig {
501                tables: vec![#( std::sync::LazyLock::force(&#tables).clone(), )*],
502                total_memory: #total_memory,
503                low_power_mode_supported: #low_power,
504                low_power_max_records: #low_power_max_records,
505                default_max_records: #default_max_records,
506                memory_allocator: unsafe {
507                    // 使用默认的内存分配器实现,这里返回一个空指针的静态引用
508                    static mut DEFAULT_ALLOCATOR: remdb::config::DefaultMemoryAllocator = remdb::config::DefaultMemoryAllocator;
509                    &mut DEFAULT_ALLOCATOR
510                },
511                // 日志相关配置
512                wal_config: remdb::config::WALConfig {
513                    log_path: "wal",
514                    log_mode: remdb::config::LogMode::Sync,
515                    checkpoint_interval_ms: 60000, // 默认60秒
516                    log_file_size_limit: 16 * 1024 * 1024, // 默认16MB
517                    log_prealloc_size: 1 * 1024 * 1024, // 默认1MB预分配
518                    log_segment_size: 16 * 1024 * 1024, // 默认16MB分段
519                    retained_checkpoints: 3, // 保留3个检查点
520                    max_consecutive_invalid: 100,
521                    skip_threshold: 1000,
522                    skip_block_size: 1024 * 1024,
523                    max_skip_attempts: 3,
524                    compression_type: remdb::config::WALCompressionType::None,
525                    compression_level: 3
526                },
527                // 时序数据默认配置
528                time_series_defaults: remdb::time_series::TimeSeriesConfig::DEFAULT,
529                // PubSub配置(可选)
530                #[cfg(feature = "pubsub")]
531                pubsub_config: None,
532                // HA相关配置(可选)
533                #[cfg(feature = "ha")]
534                ha_config: Some(remdb::ha::HAConfig {
535                    node_id: 1, // 默认节点ID为1
536                    ha_role: remdb::ha::HARole::Auto,
537                    replication_mode: remdb::ha::ReplicationMode::Async,
538                    heartbeat_interval_ms: 1000, // 默认1秒
539                    failure_detection_ms: 3000, // 默认3秒
540                    sync_timeout_ms: 2000, // 默认2秒
541                    master_address: None,
542                    master_port: None,
543                    replication_port: 5556,
544                }),
545                // Model Worker配置
546                model_worker_config: remdb::config::ModelWorkerConfig::DEFAULT,
547            }
548        });
549    };
550
551    output.into()
552}