diff --git a/src/config.rs b/src/config.rs index 22444f4..e554a4f 100644 --- a/src/config.rs +++ b/src/config.rs @@ -15,7 +15,7 @@ fn conversion_error(message: impl Into) -> mlua::Error { } /// 将 Lua 字符串转为 Rust 字符串;非 UTF-8 按系统默认的 ANSI/OEM (如 GBK) 进行安全解码 -fn lua_string_2_os_string(s: LuaString) -> mlua::Result { +fn lua_string_2_os_string(s: &LuaString) -> mlua::Result { let raw_bytes = &s.as_bytes().to_vec(); // 1. Unix 平台:直接零拷贝透传原始字节(无损支持任意编码) #[cfg(unix)] @@ -44,58 +44,128 @@ fn lua_string_2_os_string(s: LuaString) -> mlua::Result { /// 按连续整数下标遍历表的序列部分;存在空洞或非序列键时返回错误,避免静默截断 -/// 值分流与上下文分发枚举 +/// Lua 值校验器:针对不同上下文定义校验规则 #[derive(Debug, Clone, Copy, PartialEq, Eq)] -pub enum ValueShunt<'a> { +pub enum ValueValidator { + /// 目标程序路径:必须是字符串路径 + Target, /// 命令行参数元素:仅支持一维基础标量,禁止嵌套 Table Args, /// 环境变量值:支持基础标量及多维嵌套 Table(递归展平) - Env(&'a str), + Env, } -impl<'a> Display for ValueShunt<'a> { + +impl fmt::Display for ValueValidator { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { match self { + Self::Target => write!(f, "目标路径 (target)"), Self::Args => write!(f, "命令行参数 (args)"), - Self::Env(name) => write!(f, "环境变量 ({name})"), + Self::Env => write!(f, "环境变量 (env)"), } } } - -impl<'a> ValueShunt<'a> { - fn parse_sequence(&self, tbl: &Table, out: &mut Vec) -> mlua::Result<()> { - self.validate_each_sequence_item(tbl, |item| self.collect_into(item, out)) +impl ValueValidator { + /// 【纯粹校验入口】只进行逻辑与结构判定,零副作用、不产生内存分配 + pub fn validate(&self, value: &Value) -> mlua::Result<()> { + if matches!(self, Self::Target) { + println!("跟踪taget{}", value.type_name()); + } + if matches!(self, Self::Args) { + println!("跟踪args2{}", value.type_name()); + } + if matches!(self, Self::Env) { + println!("跟踪env{}", value.type_name()); + } + match self { + Self::Target => match value { + Value::String(_) => Ok(()), + Value::Nil => Err(conversion_error("缺少必填字段 target(应为字符串路径)")), + other => Err(conversion_error(format!( + "{} 需为有效的路径且类型必须是字符串,实际类型是 {}", + self, + other.type_name() + ))), + }, + Self::Args => match value { + Value::Table(tbl) => self.validate_sequence_table(tbl), + other => Err(conversion_error(format!( + "{} 必须是数组列表,实际类型是 {}", + self, + other.type_name() + ))), + }, + Self::Env => match value { + Value::Nil + | Value::String(_) + | Value::Integer(_) + | Value::Number(_) + | Value::Boolean(_) => Ok(()), + Value::Table(tbl) => self.validate_sequence_table(tbl), + other => Err(conversion_error(format!( + "{} 不支持 {} 类型, 仅支持 string/number/boolean 或嵌套数组", + self, + other.type_name() + ))), + }, + } } - fn validate_each_sequence_item( - &self, - tbl: &Table, - mut f: impl FnMut(Value) -> mlua::Result<()>, - ) -> mlua::Result<()> { + /// 校验 Table 是否为严格连续的纯数组,并递归校验其内部元素 + fn validate_sequence_table(&self, tbl: &Table) -> mlua::Result<()> { let mut index = 1i64; + + // 1. 顺序遍历连续整数索引 1..N loop { let item: Value = tbl.raw_get(index)?; if matches!(item, Value::Nil) { break; } - f(item)?; + match self { + Self::Args => match item { + Value::Table(_) => { + return Err(conversion_error(format!( + "{} 仅支持一维数组,不能包含嵌套 Table,索引位置: {index}", + self + ))); + } + Value::String(_)|Value::Integer(_)|Value::Number(_)|Value::Boolean(_) => {} + _ => {return Err(conversion_error("args异常"))} + }, + Self::Env => match item { + Value::Table(_) => self.validate(&item)?, + _ => {} + }, + _ => {} + } + + // // 元素规则判断 + // match item { + // // Args 规则拦截:禁止嵌套数组 + // Value::Table(_) if matches!(self, Self::Args) => { + // return Err(conversion_error(format!( + // "{} 仅支持一维数组,不能包含嵌套 Table,索引位置{index}", + // self + // ))); + // } + // // 标量或合法嵌套表:调用 validate 递归深度判定 + // _ => self.validate(&item)?, + // } + index += 1; } - // 校验剩余键:只允许已经被遍历的连续整数键 + + // 2. 查漏:校验是否存在空洞索引或字典键 (Key-Value 键值对) for pair in tbl.pairs::() { let (key, _) = pair?; match key { - Value::Integer(i) if i >= 1 && i < index => {} + Value::Integer(i) if i >= 1 && i < index => {} // 已遍历放行 Value::Integer(_) => { return Err(conversion_error(format!( - "{} 第 {} 个元素为 nil 或不存在(请检查是否漏写引号或变量未定义)", + "{} 第 {} 个元素为 nil 或不存在(请检查漏写引号或变量未定义)", self, index ))); } _ => { - // let label = match self { - // ValueShunt::Args => "命令行参数 (args)".to_string(), - // ValueShunt::Env(n) => format!("环境变量 [{n}]"), - // }; return Err(conversion_error(format!( "{} 必须是纯列表,不能包含键值对/字典结构", self @@ -103,52 +173,139 @@ impl<'a> ValueShunt<'a> { } } } + Ok(()) } -} -impl<'a> ValueShunt<'a> { - /// 递归/标量收集实现 - pub fn collect_into(&self, value: Value, out: &mut Vec) -> mlua::Result<()> { - match value { - Value::Nil => {} - Value::String(s) => out.push(lua_string_2_os_string(s)?), - Value::Integer(i) => out.push(OsString::from(i.to_string())), - Value::Number(n) => { - tracing::warn!(context = %self, value = %n, "浮点数将按十进制格式转换为字符串"); - out.push(OsString::from(n.to_string())); - } - Value::Boolean(b) => { - tracing::warn!(context = %self, value = %b, "布尔值将转换为字符串"); - out.push(OsString::from(b.to_string())); - } - // 嵌套 Table 仅在 EnvValue 下允许递归展开(支持 get_env("PATH") 等返回的子表) - // Value::Table(tbl) if *self == Self::Env(_) => { - // self.validate_each_sequence_item(&tbl, |item| self.collect_into(item, out))?; - // } - Value::Table(tbl) => { - match self { - Self::Args => { - return Err(conversion_error(format!( - "{} 不支持 {} 类型, 仅支持 string/number/boolean", - self, - tbl.to_string()? - ))); - } - Self::Env(_) => { - self.validate_each_sequence_item(&tbl, |item| { - self.collect_into(item, out) - })?; - } - } - // self.validate_each_sequence_item(&tbl, |item| self.collect_into(item, out))?; - } + /// 解析并校验 `target` + pub fn parse_target(value: &Value) -> mlua::Result { + let validator = Self::Target; + validator.validate(&value)?; + let ctx = format!("{}", validator); + if let Value::String(s) = value { + let os_str = lua_string_2_os_string(&s)?; + Self::ensure_no_nul(&ctx, &os_str)?; + Ok(PathBuf::from(os_str)) + } else { + unreachable!() + } + } + /// 解析并校验 `args` + pub fn parse_args(value: &Value) -> mlua::Result> { + println!("跟踪args1{}", value.type_name()); + let validator = Self::Args; + validator.validate(&value)?; + let ctx = format!("{}", validator); + + let mut raw_parts = Vec::new(); + collect_value_into(&ctx, value, &mut raw_parts)?; + + for part in &raw_parts { + Self::ensure_no_nul(&ctx, &part)?; + } + Ok(raw_parts) + } + fn parse_env_name(name: &Value) -> mlua::Result { + let name_str = match name { + Value::String(s) => match s.to_str() { + Ok(str_ref) => str_ref.to_string(), + Err(_) => return Err(conversion_error("环境变量名必须是合法的 UTF-8 字符串")), + }, other => { return Err(conversion_error(format!( - "{} 不支持 {} 类型, 仅支持 string/number/boolean 或包含这些类型的数组/嵌套数组", - self, - other.type_name(), + "环境变量键名类型错误:期望 string,实际是 {}", + other.type_name() + ))); + } + }; + if name_str.is_empty() { + return Err(conversion_error("环境变量名不能为空")); + } + if name_str.contains('=') { + return Err(conversion_error(format!( + "环境变量名 [{}] 不能包含 '='", + name_str + ))); + } + if name_str.contains('\0') { + return Err(conversion_error(format!( + "环境变量名 [{}] 不能包含 NUL 字符", + name_str + ))); + } + Ok(name_str) + } + + /// 解析、校验并拼接单个环境变量键值对 (`key`, `value`) -> (`String`, `OsString`) + fn parse_env_val(name: &str, raw_val: &Value) -> mlua::Result { + Self::Env.validate(&raw_val)?; + + let ctx = format!("环境变量 [{}]", name); + let mut parts = Vec::new(); + + collect_value_into(&ctx, raw_val, &mut parts)?; + + // 校验每个展开元素的 NUL 字符 + for part in &parts { + Self::ensure_no_nul(name, part)?; + } + + // 3. 使用系统路径分隔符拼接数组列表 + let joined_os_str = std::env::join_paths(parts).map_err(|e| { + conversion_error(format!("{} 的值无法用系统路径分隔符拼接: {}", ctx, e)) + })?; + Ok(joined_os_str) + } + /// 解析并校验单个环境变量键值对 (`key`, `val`) -> (`String`, `OsString`) + fn parse_env_pair(key: &Value, val: &Value) -> mlua::Result<(String, OsString)> { + // 1. 解析并校验 Key,拿到安全的 String + let name = Self::parse_env_name(key)?; + // 2. 借用 &name 传递给 Value 解析器作为上下文 + let os_val = Self::parse_env_val(&name, val)?; + // 3. 所有权平滑转移,组装元组返回 + Ok((name, os_val)) + } + /// 解析整个 `env` Table,直接返回安全的环境变量 Map + pub fn parse_env(value: &Value) -> mlua::Result> { + let Value::Table(tbl) = value else { + // 这里作为防御性校验,外部虽然 match 过了,但内层仍保持强类型安全 + return Err(conversion_error(format!( + "env 必须是键值表 (table),实际类型是 {}", + value.type_name() + ))); + }; + + let mut env_map = HashMap::new(); + + for pair in tbl.pairs::() { + let (raw_key, raw_val) = pair?; + let (key, val) = Self::parse_env_pair(&raw_key, &raw_val)?; + env_map.insert(key, val); + } + + Ok(env_map) + } + // ================= 3. 底层 NUL 字符跨平台安全检查 ================= + + fn ensure_no_nul(context_desc: &str, os_str: &std::ffi::OsStr) -> mlua::Result<()> { + #[cfg(unix)] + { + use std::os::unix::ffi::OsStrExt; + if os_str.as_bytes().contains(&0) { + return Err(conversion_error(format!( + "{} 的值不能包含 NUL 字符", + context_desc + ))); + } + } + #[cfg(windows)] + { + use std::os::windows::ffi::OsStrExt; + if os_str.encode_wide().any(|c| c == 0) { + return Err(conversion_error(format!( + "{} 的值不能包含 NUL 字符", + context_desc ))); } } @@ -156,38 +313,71 @@ impl<'a> ValueShunt<'a> { } } -/// 校验环境变量名的合法性(Windows 约束:非空、不含 '='、不含 NUL) -fn validate_env_var_name(key_val: &Value, val: &Value) -> mlua::Result { - let key = match key_val { - Value::String(s) => s.to_str()?.to_string(), - - _ => { - return Err(conversion_error(format!( - r#"env 配置解析失败:值: {} (类型: {}) 无效的键值! - 示例 HOME="/home" 或 PATH = {{ "C:/tools/", "D:/Windows"}})"#, - val.to_string()?, - val.type_name() - ))); +/// 【纯粹数据转换与展开】将已通过校验的 Value 递归解析为 OsString 动态数组 +pub fn collect_value_into( + context: &impl fmt::Display, + value: &Value, + out: &mut Vec, +) -> mlua::Result<()> { + match value { + Value::Nil => {} + Value::String(s) => out.push(lua_string_2_os_string(s)?), + Value::Integer(i) => out.push(OsString::from(i.to_string())), + Value::Number(n) => { + tracing::warn!(%context, value = %n, "浮点数将按十进制格式转换为字符串"); + out.push(OsString::from(n.to_string())); } - }; - if key.is_empty() { - return Err(conversion_error("环境变量名不能为空")); + Value::Boolean(b) => { + tracing::warn!(%context, value = %b, "布尔值将转换为字符串"); + out.push(OsString::from(b.to_string())); + } + Value::Table(tbl) => { + let mut index = 1i64; + loop { + let item: Value = tbl.raw_get(index)?; + if matches!(item, Value::Nil) { + break; + } + collect_value_into(context, &item, out)?; + index += 1; + } + } + _ => unreachable!("传入收集器的 Value 应已通过 validate 校验"), } - if key.contains('=') { - return Err(conversion_error(format!( - "环境变量名 [{}] 不能包含 '='", - key - ))); - } - if key.contains('\0') { - return Err(conversion_error(format!( - "环境变量名 [{}] 不能包含 NUL 字符", - key - ))); - } - Ok(key) + Ok(()) } +/// 校验环境变量名的合法性(Windows 约束:非空、不含 '='、不含 NUL) +// fn parse_env_name(name: &Value, val: &Value) -> mlua::Result { +// let name_str = match name { +// Value::String(s) => s.to_str()?.to_string(), +// +// _ => { +// return Err(conversion_error(format!( +// "env 配置解析失败:(对应的值: {}, 类型: {}) 无效的键值!示例 HOME=\"/home\" 或 PATH = {{ \"C:/tools/\", \"D:/Windows\"}})", +// val.to_string()?, +// val.type_name() +// ))); +// } +// }; +// if name_str.is_empty() { +// return Err(conversion_error("环境变量名不能为空")); +// } +// if name_str.contains('=') { +// return Err(conversion_error(format!( +// "环境变量名 [{}] 不能包含 '='", +// name_str +// ))); +// } +// if name_str.contains('\0') { +// return Err(conversion_error(format!( +// "环境变量名 [{}] 不能包含 NUL 字符", +// name_str +// ))); +// } +// Ok(name_str) +// } + /// 递归将任意 Lua Value 展开为扁平的字符串片段列表 // fn collect_env_segments(value: Value, out: &mut Vec) -> mlua::Result<()>{OK(())} @@ -237,93 +427,43 @@ impl FromLua for ShimConfig { }; // 必填字段: target(严格限定为字符串,避免数字被 mlua 宽松转为字符串后掩盖错误) - let target = match table.get::("target")? { - Value::String(s) => lua_string_2_os_string(s)?, - Value::Nil => { + let target = match table.get::>("target")? { + None | Some(Value::Nil) => { return Err(conversion_error("缺少必填字段 target(应为字符串路径)")); } - other => { - return Err(conversion_error(format!( - "target 需为有效的路径且类型必须是字符串,实际类型是 {}", - other.type_name() - ))); - } + Some(target_val) => ValueValidator::parse_target(&target_val)?, + // Value::String(s) => lua_string_2_os_string(&s)?, + // Value::Nil => { + // return Err(conversion_error("缺少必填字段 target(应为字符串路径)")); + // } + // other => { + // return Err(conversion_error(format!( + // "target 需为有效的路径且类型必须是字符串,实际类型是 {}", + // other.type_name() + // ))); + // } }; - let mut args = Vec::new(); + // let mut args = Vec::new(); // 可选字段: args(缺失或 nil 时默认空列表;保留空字符串参数,与提权路径的 "" 语义一致) - match table.get::>("args")? { - None | Some(Value::Nil) => {} - Some(Value::Table(tbl)) => ValueShunt::Args.parse_sequence(&tbl, &mut args)?, - Some(other) => { - return Err(conversion_error(format!( - "args 必须是数组列表,实际类型是 {}", - other.type_name() - ))); - } + let args = match table.get::>("args")? { + None | Some(Value::Nil) => Vec::new(), + Some(args_val) => ValueValidator::parse_args(&args_val)?, }; + // 可选字段: env(只允许缺失/nil,其他类型由 Option 转换报错,不再静默忽略) - let env_table = table.get::>("env")?; - println!("env_table:{:?}", env_table); - let mut env = HashMap::new(); - match env_table { - None | Some(Value::Nil) => {} - Some(Value::Table(tbl)) => { - for pair in tbl.pairs::() { - let (name, value) = pair?; - - let name = validate_env_var_name(&name, &value)?; - println!("==========={name}:{:?}", value); - let mut parts = Vec::new(); - ValueShunt::Env(&name).collect_into(value, &mut parts)?; - - // 空数组/空串也显式设置(空值表示清空该变量), - // 与“未配置该变量(继承宿主环境)”相区分 - let joined_os_str = std::env::join_paths(parts).map_err(|e| { - conversion_error(format!( - "环境变量 [{}] 的值无法用系统路径分隔符拼接: {}", - name, e - )) - })?; - // println!("环境变量拼接结果:{:?}", joined_os_str); - // 检查是否包含非法 NUL 字符 - #[cfg(unix)] - { - use std::os::unix::ffi::OsStrExt; - if joined_os_str.as_bytes().contains(&0) { - return Err(conversion_error(format!( - "环境变量 [{}] 的值不能包含 NUL 字符", - key - ))); - } - } - #[cfg(windows)] - { - use std::os::windows::ffi::OsStrExt; - if joined_os_str.encode_wide().any(|c| c == 0) { - return Err(conversion_error(format!( - "环境变量 [{}] 的值不能包含 NUL 字符", - name - ))); - } - } - - env.insert(name, joined_os_str); - } - } + let env = match table.get::>("env")? { + None | Some(Value::Nil) => HashMap::new(), + Some(ref env_val @ Value::Table(_)) => ValueValidator::parse_env(&env_val)?, Some(other) => { return Err(conversion_error(format!( "env 必须是键值表 (table),实际类型是 {}", other.type_name() ))); } - } + }; println!("环境变量结果:{:?}", env); - Ok(ShimConfig { - target: PathBuf::from(target), - args, - env, - }) + Ok(ShimConfig { target, args, env }) } } @@ -361,12 +501,13 @@ mod tests { r#" return { target = "C:/tools/git.exe", - args = { "--no-pager" }, + args = { "--no-pager"}, env = { PATH = { "C:/tools/git/bin", "C:/Windows", get_env("PATH") }, HOME = "C:/tools/home", CONST = 3, - BOOL = true + BOOL = true, + } } "#,