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)这样的类型,最大65536
318            let str_size = if let Some(params) = type_params {
319                let raw = params.base10_parse().unwrap_or(32);
320                if raw > 65536 {
321                    65536
322                } else {
323                    raw
324                }
325            } else {
326                32
327            };
328            (
329                quote!(remdb::types::DataType::VarChar),
330                str_size,
331                quote!(Some(#str_size as usize)),
332            )
333        } else if type_name == "text" {
334            // 处理text类型,字段大小固定为TextStorage的大小(264字节)
335            // 文本内容通过TextStorage动态分配,不直接存储在记录缓冲区中
336            let text_storage_size = 264; // size_of::<TextStorage>()
337            (
338                quote!(remdb::types::DataType::Text),
339                text_storage_size,
340                quote!(None),
341            )
342        } else if type_name == "vector" {
343            // 处理vector(2)这样的向量类型
344            let dim = if let Some(params) = type_params {
345                params.base10_parse().unwrap_or(128)
346            } else {
347                128
348            };
349            (
350                quote!(remdb::types::DataType::Vector),
351                dim * 4,
352                quote!(None),
353            ) // 向量每个维度4字节
354        } else {
355            (quote!(remdb::types::DataType::Int32), 4, quote!(None))
356        };
357
358        // 计算对齐要求
359        let alignment = if type_name == "u64" || type_name == "f64" || type_name == "i64" {
360            8
361        } else if type_name == "i32" || type_name == "u32" || type_name == "f32" {
362            4
363        } else if type_name == "i16" || type_name == "u16" {
364            2
365        } else {
366            1
367        };
368
369        // 调整偏移量以满足对齐要求
370        offset = ((offset + alignment - 1) / alignment) * alignment;
371
372        // 确定约束字段值
373        let is_primary_key = field_name == primary_key;
374        let primary_key_val = is_primary_key;
375        let not_null_val = is_primary_key; // 主键字段默认为非空
376        let unique_val = is_primary_key;
377
378        // 检查是否为自增主键:
379        // 1. 整数主键默认自增
380        // 2. 可以显式指定AUTOINCREMENT
381        let is_integer_type =
382            type_name == "i32" || type_name == "i64" || type_name == "u32" || type_name == "u64";
383        let auto_increment_val = is_primary_key && is_integer_type;
384
385        // 生成向量元数据(仅向量类型字段需要)
386        let vector_metadata_code = if type_name == "vector" {
387            let dim = if let Some(params) = type_params {
388                params.base10_parse::<u16>().unwrap_or(128)
389            } else {
390                128u16
391            };
392            quote! {
393                Some(remdb::types::VectorMetadata {
394                    dimension: #dim,
395                    distance_type: remdb::types::DistanceType::L2,
396                    index_type: remdb::types::VectorIndexType::HNSW,
397                    compression_enabled: false,
398                    compression_scheme: 0,
399                    compression_level: 3,
400                    // HNSW默认参数
401                    hnsw_m: 16,
402                    hnsw_ef_construction: 200,
403                    hnsw_ef_search: 128,
404                    // IVF默认参数
405                    ivf_nlist: 1024,
406                    ivf_nprobe: 16,
407                })
408            }
409        } else {
410            quote! { None }
411        };
412
413        // 生成字段定义
414        let field_def = quote! {
415            remdb::types::FieldDef {
416                name: stringify!(#field_name).to_string(),
417                data_type: #data_type,
418                size: #size_val as usize, // 确保是usize类型
419                string_length: #string_length,
420                offset: #offset as usize,  // 确保是usize类型
421                primary_key: #primary_key_val,
422                not_null: #not_null_val,
423                unique: #unique_val,
424                auto_increment: #auto_increment_val,
425                default_value: None,
426                vector_metadata: #vector_metadata_code,
427                json_metadata: None,
428            }
429        };
430
431        field_defs.push(field_def);
432
433        // 确定主键和二级索引的字段索引
434        if field_name == primary_key {
435            primary_key_index = i;
436        }
437
438        if let Some(secondary_field) = secondary_index {
439            if field_name == secondary_field {
440                secondary_key_index = Some(i);
441            }
442        }
443
444        // 更新偏移量和记录大小
445        offset += size_val;
446        record_size = offset;
447    }
448
449    // 确保整个记录满足最大对齐要求(8字节对齐)
450    let max_alignment = 8;
451    record_size = ((record_size + max_alignment - 1) / max_alignment) * max_alignment;
452
453    // 将max_records转换为usize
454    let max_records_usize = max_records.base10_parse::<usize>().unwrap_or(100);
455
456    // 确定索引类型
457    let index_type = match secondary_index_type.as_ref() {
458        Some(ty) if ty == "btree" => quote!(remdb::types::IndexType::BTree),
459        Some(ty) if ty == "hash" => quote!(remdb::types::IndexType::Hash),
460        Some(ty) if ty == "ttree" => quote!(remdb::types::IndexType::TTree),
461        Some(ty) if ty == "sortedarray" => quote!(remdb::types::IndexType::SortedArray),
462        _ => quote!(remdb::types::IndexType::BTree),
463    };
464
465    // 生成secondary_index代码
466    let secondary_index_code = match secondary_key_index {
467        Some(index) => quote! { Some(vec![#index as usize]) },
468        None => quote! { None },
469    };
470
471    // 生成代码:返回一个静态TableDef变量,使用LazyLock延迟初始化
472    let output = quote! {
473        #[allow(non_upper_case_globals)]
474        pub static #name: std::sync::LazyLock<remdb::types::TableDef> = std::sync::LazyLock::new(|| {
475            remdb::types::TableDef {
476                id: 0,
477                name: stringify!(#name).to_string(),
478                fields: vec![#(#field_defs,)*],
479                primary_key: vec![#primary_key_index as usize],
480                secondary_index: #secondary_index_code,
481                secondary_index_type: #index_type,
482                record_size: #record_size as usize,
483                max_records: #max_records_usize,
484                version: 1,
485                created_at: 0,
486                updated_at: 0,
487            }
488        });
489    };
490
491    output.into()
492}
493
494#[proc_macro]
495pub fn database(input: proc_macro::TokenStream) -> proc_macro::TokenStream {
496    // 解析输入参数
497    let args = parse_macro_input!(input as DatabaseArgs);
498    let name = &args.name;
499    let tables = &args.tables;
500    let low_power = args.low_power;
501    let default_max_records = args.default_max_records;
502    let total_memory = args.total_memory;
503
504    // 处理low_power_max_records,转换为Option<usize>
505    let low_power_max_records = match args.low_power_max_records {
506        Some(val) => quote! { Some(#val) },
507        None => quote! { None },
508    };
509
510    // 生成代码:返回一个静态DbConfig变量,使用LazyLock延迟初始化
511    let output = quote! {
512        #[allow(non_upper_case_globals)]
513        pub static #name: std::sync::LazyLock<remdb::config::DbConfig> = std::sync::LazyLock::new(|| {
514            remdb::config::DbConfig {
515                tables: vec![#( std::sync::LazyLock::force(&#tables).clone(), )*],
516                total_memory: #total_memory,
517                low_power_mode_supported: #low_power,
518                low_power_max_records: #low_power_max_records,
519                default_max_records: #default_max_records,
520                memory_allocator: unsafe {
521                    // 使用默认的内存分配器实现,这里返回一个空指针的静态引用
522                    static mut DEFAULT_ALLOCATOR: remdb::config::DefaultMemoryAllocator = remdb::config::DefaultMemoryAllocator;
523                    &mut DEFAULT_ALLOCATOR
524                },
525                // 日志相关配置
526                wal_config: remdb::config::WALConfig {
527                    log_path: "wal",
528                    log_mode: remdb::config::LogMode::Sync,
529                    checkpoint_interval_ms: 60000, // 默认60秒
530                    log_file_size_limit: 16 * 1024 * 1024, // 默认16MB
531                    log_prealloc_size: 1 * 1024 * 1024, // 默认1MB预分配
532                    log_segment_size: 16 * 1024 * 1024, // 默认16MB分段
533                    retained_checkpoints: 3, // 保留3个检查点
534                    max_consecutive_invalid: 100,
535                    skip_threshold: 1000,
536                    skip_block_size: 1024 * 1024,
537                    max_skip_attempts: 3,
538                    compression_type: remdb::config::WALCompressionType::None,
539                    compression_level: 3
540                },
541                // 时序数据默认配置
542                time_series_defaults: remdb::time_series::TimeSeriesConfig::DEFAULT,
543                // PubSub配置(可选)
544                #[cfg(feature = "pubsub")]
545                pubsub_config: None,
546                // HA相关配置(可选)
547                #[cfg(feature = "ha")]
548                ha_config: Some(remdb::ha::HAConfig {
549                    node_id: 1, // 默认节点ID为1
550                    ha_role: remdb::ha::HARole::Auto,
551                    replication_mode: remdb::ha::ReplicationMode::Async,
552                    heartbeat_interval_ms: 1000, // 默认1秒
553                    failure_detection_ms: 3000, // 默认3秒
554                    sync_timeout_ms: 2000, // 默认2秒
555                    master_address: None,
556                    master_port: None,
557                    replication_port: 5556,
558                }),
559                // Model Worker配置
560                model_worker_config: remdb::config::ModelWorkerConfig::DEFAULT,
561            }
562        });
563    };
564
565    output.into()
566}