Skip to main content

sz_rust_cli/
context_builder.rs

1// SPDX-License-Identifier: Apache-2.0
2// Copyright (c) 2024-2026 SZ-Rust Team
3//
4//! 模板上下文构建器
5//!
6//! 对应 design.md 第 2.2.2.5 节,从 CLI 参数构建 Tera 渲染上下文。
7
8use tera::Context;
9
10use crate::error::CliError;
11use crate::field_parser::{Field, FieldParser};
12
13/// `make:plugin` 命令参数
14#[derive(Debug, Clone)]
15pub struct PluginCommandArgs {
16    /// 模板类型(必填,如 "crud" 或 "master-slave")
17    pub template: String,
18    /// 插件名称(必填)
19    pub name: String,
20    /// 表名(可选,默认取插件名 snake_case)
21    pub table: Option<String>,
22    /// 字段定义(可选,如 "id:i32:pk,name:String")
23    pub fields: Option<String>,
24    /// 是否强制覆盖已存在目录
25    pub force: bool,
26    /// 输出目录(可选,默认 `plugins/<name>/`)
27    pub output: Option<String>,
28    /// 主表名(主从模板,可选)
29    pub master: Option<String>,
30    /// 从表名(主从模板,可选)
31    pub slave: Option<String>,
32    /// 主表字段定义(主从模板,可选)
33    pub master_fields: Option<String>,
34    /// 从表字段定义(主从模板,可选)
35    pub slave_fields: Option<String>,
36    /// 外键字段名(主从模板,可选)
37    pub foreign_key: Option<String>,
38}
39
40/// 模板上下文构建器
41pub struct TemplateContextBuilder {
42    args: PluginCommandArgs,
43}
44
45impl TemplateContextBuilder {
46    /// 创建新的上下文构建器
47    pub fn new(args: PluginCommandArgs) -> Self {
48        Self { args }
49    }
50
51    /// 构建 CRUD 模板上下文
52    ///
53    /// 含 6 个必需变量:plugin_name, table_name, class_name, fields, module_path, template_version
54    /// 外加 generated_at, primary_key_name, primary_key_type
55    pub fn build(&self) -> Result<Context, CliError> {
56        let table_name = self
57            .args
58            .table
59            .clone()
60            .unwrap_or_else(|| to_snake_case(&self.args.name));
61
62        let class_name = to_pascal_case(&table_name);
63
64        let fields_str = self
65            .args
66            .fields
67            .as_deref()
68            .unwrap_or("id:i32:pk,name:String");
69        let fields = FieldParser::parse(fields_str)?;
70
71        let (pk_name, pk_type) = find_primary_key(&fields);
72
73        let module_path = format!("plugins::{}", to_snake_case(&self.args.name));
74
75        let fields_json: Vec<serde_json::Value> = fields
76            .iter()
77            .map(|f| {
78                serde_json::json!({
79                    "name": f.name,
80                    "rust_type": f.rust_type,
81                    "sql_type": f.sql_type,
82                    "is_nullable": f.is_nullable,
83                    "is_primary_key": f.is_primary_key,
84                    "is_indexed": f.is_indexed,
85                })
86            })
87            .collect();
88
89        let mut ctx = Context::new();
90        ctx.insert("plugin_name", &self.args.name);
91        ctx.insert("table_name", &table_name);
92        ctx.insert("class_name", &class_name);
93        ctx.insert("fields", &fields_json);
94        ctx.insert("module_path", &module_path);
95        ctx.insert("template_type", &self.args.template);
96        ctx.insert("template_version", "1.0.0");
97        ctx.insert("generated_at", &current_timestamp());
98        ctx.insert("primary_key_name", &pk_name);
99        ctx.insert("primary_key_type", &pk_type);
100
101        Ok(ctx)
102    }
103
104    /// 构建主从模板上下文
105    ///
106    /// 额外含 master_table, slave_table, foreign_key, master_fields, slave_fields
107    pub fn build_master_slave(&self) -> Result<Context, CliError> {
108        let master_table = self.args.master.clone().ok_or_else(|| {
109            CliError::Generic("--master is required for master-slave template".to_string())
110        })?;
111
112        let slave_table = self.args.slave.clone().ok_or_else(|| {
113            CliError::Generic("--slave is required for master-slave template".to_string())
114        })?;
115
116        if master_table == slave_table {
117            return Err(CliError::MasterSlaveSame);
118        }
119
120        let foreign_key = self.args.foreign_key.clone().ok_or_else(|| {
121            CliError::Generic("--foreign-key is required for master-slave template".to_string())
122        })?;
123
124        let master_fields_str = self
125            .args
126            .master_fields
127            .as_deref()
128            .unwrap_or("id:i32:pk,name:String");
129        let slave_fields_str = self.args.slave_fields.as_deref().unwrap_or("id:i32:pk");
130
131        let master_fields = FieldParser::parse(master_fields_str)?;
132        let slave_fields = FieldParser::parse(slave_fields_str)?;
133
134        crate::validator::InputValidator::validate_foreign_key(&foreign_key, &slave_fields)?;
135
136        let (master_pk_name, master_pk_type) = find_primary_key(&master_fields);
137
138        let master_fields_json: Vec<serde_json::Value> = fields_to_json(&master_fields);
139        let slave_fields_json: Vec<serde_json::Value> = fields_to_json(&slave_fields);
140
141        let master_class_name = to_pascal_case(&master_table);
142        let slave_class_name = to_pascal_case(&slave_table);
143        let module_path = format!("plugins::{}", to_snake_case(&self.args.name));
144
145        let mut ctx = Context::new();
146        ctx.insert("plugin_name", &self.args.name);
147        ctx.insert("table_name", &master_table);
148        ctx.insert("class_name", &master_class_name);
149        ctx.insert("fields", &master_fields_json);
150        ctx.insert("module_path", &module_path);
151        ctx.insert("template_type", &self.args.template);
152        ctx.insert("template_version", "1.0.0");
153        ctx.insert("generated_at", &current_timestamp());
154        ctx.insert("primary_key_name", &master_pk_name);
155        ctx.insert("primary_key_type", &master_pk_type);
156
157        ctx.insert("master_table", &master_table);
158        ctx.insert("slave_table", &slave_table);
159        ctx.insert("master_class_name", &master_class_name);
160        ctx.insert("slave_class_name", &slave_class_name);
161        ctx.insert("master_fields", &master_fields_json);
162        ctx.insert("slave_fields", &slave_fields_json);
163        ctx.insert("foreign_key", &foreign_key);
164
165        Ok(ctx)
166    }
167
168    /// 返回参数引用
169    pub fn args(&self) -> &PluginCommandArgs {
170        &self.args
171    }
172}
173
174/// snake_case 转换
175fn to_snake_case(s: &str) -> String {
176    let mut result = String::new();
177    for (i, ch) in s.chars().enumerate() {
178        if ch.is_uppercase() {
179            if i > 0 {
180                result.push('_');
181            }
182            result.push(ch.to_ascii_lowercase());
183        } else if ch == '-' {
184            result.push('_');
185        } else {
186            result.push(ch);
187        }
188    }
189    result
190}
191
192/// PascalCase 转换
193fn to_pascal_case(s: &str) -> String {
194    let mut result = String::new();
195    let mut next_upper = true;
196    for ch in s.chars() {
197        if ch == '_' || ch == '-' || ch == ' ' {
198            next_upper = true;
199        } else if next_upper {
200            result.push(ch.to_ascii_uppercase());
201            next_upper = false;
202        } else {
203            result.push(ch);
204        }
205    }
206    result
207}
208
209/// 查找主键字段
210fn find_primary_key(fields: &[Field]) -> (String, String) {
211    for f in fields {
212        if f.is_primary_key {
213            return (f.name.clone(), f.rust_type.clone());
214        }
215    }
216    ("id".to_string(), "i32".to_string())
217}
218
219/// 字段列表转 JSON
220fn fields_to_json(fields: &[Field]) -> Vec<serde_json::Value> {
221    fields
222        .iter()
223        .map(|f| {
224            serde_json::json!({
225                "name": f.name,
226                "rust_type": f.rust_type,
227                "sql_type": f.sql_type,
228                "is_nullable": f.is_nullable,
229                "is_primary_key": f.is_primary_key,
230                "is_indexed": f.is_indexed,
231            })
232        })
233        .collect()
234}
235
236/// 当前时间戳
237fn current_timestamp() -> String {
238    chrono::Utc::now()
239        .format("%Y-%m-%d %H:%M:%S UTC")
240        .to_string()
241}
242
243#[cfg(test)]
244mod tests {
245    use super::*;
246
247    fn make_crud_args() -> PluginCommandArgs {
248        PluginCommandArgs {
249            template: "crud".to_string(),
250            name: "user-management".to_string(),
251            table: Some("users".to_string()),
252            fields: Some("id:i32:pk,name:String,age:i32".to_string()),
253            force: false,
254            output: None,
255            master: None,
256            slave: None,
257            master_fields: None,
258            slave_fields: None,
259            foreign_key: None,
260        }
261    }
262
263    #[test]
264    fn test_build_crud_context() {
265        let builder = TemplateContextBuilder::new(make_crud_args());
266        let ctx = builder.build().unwrap();
267
268        assert_eq!(ctx.get("plugin_name").unwrap(), "user-management");
269        assert_eq!(ctx.get("table_name").unwrap(), "users");
270        assert_eq!(ctx.get("class_name").unwrap(), "Users");
271        assert_eq!(ctx.get("template_type").unwrap(), "crud");
272        assert!(ctx.get("generated_at").is_some());
273        assert_eq!(ctx.get("primary_key_name").unwrap(), "id");
274        assert_eq!(ctx.get("primary_key_type").unwrap(), "i32");
275    }
276
277    #[test]
278    fn test_build_crud_default_table() {
279        let mut args = make_crud_args();
280        args.table = None;
281        let builder = TemplateContextBuilder::new(args);
282        let ctx = builder.build().unwrap();
283        assert_eq!(ctx.get("table_name").unwrap(), "user_management");
284    }
285
286    #[test]
287    fn test_build_crud_default_fields() {
288        let mut args = make_crud_args();
289        args.fields = None;
290        let builder = TemplateContextBuilder::new(args);
291        let ctx = builder.build().unwrap();
292        assert!(ctx.get("fields").is_some());
293    }
294
295    #[test]
296    fn test_build_master_slave_context() {
297        let args = PluginCommandArgs {
298            template: "master-slave".to_string(),
299            name: "order-plugin".to_string(),
300            table: None,
301            fields: None,
302            force: false,
303            output: None,
304            master: Some("users".to_string()),
305            slave: Some("orders".to_string()),
306            master_fields: Some("id:i32:pk,name:String".to_string()),
307            slave_fields: Some("id:i32:pk,user_id:i32,total:f64".to_string()),
308            foreign_key: Some("user_id".to_string()),
309        };
310        let builder = TemplateContextBuilder::new(args);
311        let ctx = builder.build_master_slave().unwrap();
312
313        assert_eq!(ctx.get("master_table").unwrap(), "users");
314        assert_eq!(ctx.get("slave_table").unwrap(), "orders");
315        assert_eq!(ctx.get("master_class_name").unwrap(), "Users");
316        assert_eq!(ctx.get("slave_class_name").unwrap(), "Orders");
317        assert_eq!(ctx.get("foreign_key").unwrap(), "user_id");
318    }
319
320    #[test]
321    fn test_build_master_slave_same_table() {
322        let args = PluginCommandArgs {
323            template: "master-slave".to_string(),
324            name: "test".to_string(),
325            table: None,
326            fields: None,
327            force: false,
328            output: None,
329            master: Some("users".to_string()),
330            slave: Some("users".to_string()),
331            master_fields: Some("id:i32:pk".to_string()),
332            slave_fields: Some("id:i32:pk".to_string()),
333            foreign_key: Some("id".to_string()),
334        };
335        let builder = TemplateContextBuilder::new(args);
336        let result = builder.build_master_slave();
337        assert!(result.is_err());
338        assert!(matches!(result.unwrap_err(), CliError::MasterSlaveSame));
339    }
340
341    #[test]
342    fn test_build_master_slave_fk_not_found() {
343        let args = PluginCommandArgs {
344            template: "master-slave".to_string(),
345            name: "test".to_string(),
346            table: None,
347            fields: None,
348            force: false,
349            output: None,
350            master: Some("users".to_string()),
351            slave: Some("orders".to_string()),
352            master_fields: Some("id:i32:pk,name:String".to_string()),
353            slave_fields: Some("id:i32:pk,total:f64".to_string()),
354            foreign_key: Some("user_id".to_string()),
355        };
356        let builder = TemplateContextBuilder::new(args);
357        let result = builder.build_master_slave();
358        assert!(result.is_err());
359        assert!(matches!(
360            result.unwrap_err(),
361            CliError::ForeignKeyNotFound(_)
362        ));
363    }
364
365    #[test]
366    fn test_to_snake_case() {
367        assert_eq!(to_snake_case("UserManagement"), "user_management");
368        assert_eq!(to_snake_case("user-management"), "user_management");
369        assert_eq!(to_snake_case("user_management"), "user_management");
370    }
371
372    #[test]
373    fn test_to_pascal_case() {
374        assert_eq!(to_pascal_case("users"), "Users");
375        assert_eq!(to_pascal_case("user_orders"), "UserOrders");
376        assert_eq!(to_pascal_case("user-orders"), "UserOrders");
377        assert_eq!(to_pascal_case("UserOrders"), "UserOrders");
378    }
379
380    #[test]
381    fn test_find_primary_key() {
382        let fields = vec![
383            Field {
384                name: "name".to_string(),
385                rust_type: "String".to_string(),
386                sql_type: "VARCHAR(255)".to_string(),
387                is_nullable: false,
388                is_primary_key: false,
389                is_indexed: false,
390            },
391            Field {
392                name: "id".to_string(),
393                rust_type: "i64".to_string(),
394                sql_type: "BIGINT".to_string(),
395                is_nullable: false,
396                is_primary_key: true,
397                is_indexed: false,
398            },
399        ];
400        let (name, ty) = find_primary_key(&fields);
401        assert_eq!(name, "id");
402        assert_eq!(ty, "i64");
403    }
404
405    #[test]
406    fn test_find_primary_key_default() {
407        let fields = vec![Field {
408            name: "name".to_string(),
409            rust_type: "String".to_string(),
410            sql_type: "VARCHAR(255)".to_string(),
411            is_nullable: false,
412            is_primary_key: false,
413            is_indexed: false,
414        }];
415        let (name, ty) = find_primary_key(&fields);
416        assert_eq!(name, "id");
417        assert_eq!(ty, "i32");
418    }
419
420    #[test]
421    fn test_build_master_slave_missing_master() {
422        let args = PluginCommandArgs {
423            template: "master-slave".to_string(),
424            name: "test".to_string(),
425            table: None,
426            fields: None,
427            force: false,
428            output: None,
429            master: None,
430            slave: Some("orders".to_string()),
431            master_fields: None,
432            slave_fields: None,
433            foreign_key: None,
434        };
435        let builder = TemplateContextBuilder::new(args);
436        let result = builder.build_master_slave();
437        assert!(result.is_err());
438        assert!(result
439            .unwrap_err()
440            .to_string()
441            .contains("--master is required"));
442    }
443
444    #[test]
445    fn test_build_master_slave_missing_slave() {
446        let args = PluginCommandArgs {
447            template: "master-slave".to_string(),
448            name: "test".to_string(),
449            table: None,
450            fields: None,
451            force: false,
452            output: None,
453            master: Some("users".to_string()),
454            slave: None,
455            master_fields: None,
456            slave_fields: None,
457            foreign_key: None,
458        };
459        let builder = TemplateContextBuilder::new(args);
460        let result = builder.build_master_slave();
461        assert!(result.is_err());
462        assert!(result
463            .unwrap_err()
464            .to_string()
465            .contains("--slave is required"));
466    }
467
468    #[test]
469    fn test_build_master_slave_missing_foreign_key() {
470        let args = PluginCommandArgs {
471            template: "master-slave".to_string(),
472            name: "test".to_string(),
473            table: None,
474            fields: None,
475            force: false,
476            output: None,
477            master: Some("users".to_string()),
478            slave: Some("orders".to_string()),
479            master_fields: None,
480            slave_fields: None,
481            foreign_key: None,
482        };
483        let builder = TemplateContextBuilder::new(args);
484        let result = builder.build_master_slave();
485        assert!(result.is_err());
486        assert!(result
487            .unwrap_err()
488            .to_string()
489            .contains("--foreign-key is required"));
490    }
491
492    #[test]
493    fn test_args_accessor() {
494        let args = make_crud_args();
495        let builder = TemplateContextBuilder::new(args.clone());
496        assert_eq!(builder.args().name, args.name);
497        assert_eq!(builder.args().template, args.template);
498    }
499
500    #[test]
501    fn test_to_snake_case_lowercase_input() {
502        assert_eq!(to_snake_case("users"), "users");
503        assert_eq!(to_snake_case(""), "");
504    }
505
506    #[test]
507    fn test_to_pascal_case_empty() {
508        assert_eq!(to_pascal_case(""), "");
509        assert_eq!(to_pascal_case("a"), "A");
510    }
511
512    #[test]
513    fn test_fields_to_json_empty() {
514        let json = fields_to_json(&[]);
515        assert!(json.is_empty());
516    }
517}