1 Commits

Author SHA1 Message Date
c7999a90f6 fix(config): 优化 lua 数据校验逻辑
- 新增 ValueValidator 优化结构化校验逻辑
2026-08-21 17:29:40 +08:00

View File

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