use mlua::{FromLua, Lua, LuaString, Table, Value}; use std::collections::HashMap; use std::ffi::OsString; use std::fmt; use std::fmt::Display; use std::path::PathBuf; use std::process::Command; /// 构造一个带上下文的 FromLua 转换错误,便于定位配置问题 fn conversion_error(message: impl Into) -> mlua::Error { mlua::Error::FromLuaConversionError { from: "Lua value", to: "ShimConfig".into(), message: Some(message.into()), } } /// 将 Lua 字符串转为 Rust 字符串;非 UTF-8 按系统默认的 ANSI/OEM (如 GBK) 进行安全解码 fn lua_string_2_os_string(s: LuaString) -> mlua::Result { let raw_bytes = &s.as_bytes().to_vec(); // 1. Unix 平台:直接零拷贝透传原始字节(无损支持任意编码) #[cfg(unix)] { use std::os::unix::ffi::OsStrExt; Ok(OsStr::from_bytes(raw_bytes).to_os_string()) } // 2. Windows 平台:优先按 UTF-8 解码,失败则按当前系统本地代码页 (ANSI/GBK) 转换 #[cfg(windows)] { // 先尝试标准的 UTF-8 if let Ok(utf8_str) = std::str::from_utf8(raw_bytes) { return Ok(OsString::from(utf8_str)); } // 非 UTF-8 时,按系统默认的 ANSI/OEM (如 GBK) 进行安全解码 let (cow, _, had_errors) = encoding_rs::GBK.decode(raw_bytes); if had_errors { return Err(mlua::Error::RuntimeError( "Path contains invalid/unsupported byte encoding".into(), )); } Ok(OsString::from(cow.as_ref())) } } /// 按连续整数下标遍历表的序列部分;存在空洞或非序列键时返回错误,避免静默截断 fn validate_each_sequence_item( tbl: &Table, mut f: impl FnMut(Value) -> mlua::Result<()>, ) -> mlua::Result<()> { let mut index = 1i64; loop { let item: Value = tbl.raw_get(index)?; if matches!(item, Value::Nil) { break; } f(item)?; index += 1; } // 校验剩余键:只允许已经被遍历的连续整数键 for pair in tbl.pairs::() { let (key, _) = pair?; match key { Value::Integer(i) if i >= 1 && i < index => {} _ => { return Err(conversion_error(format!( "数组只能包含连续的整数下标 [1..{}],发现非序列键或空洞", index - 1 ))); } } } Ok(()) } /// 值分流与上下文分发枚举 #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum ValueShunt { /// 命令行参数元素:仅支持一维基础标量,禁止嵌套 Table Args, /// 环境变量值:支持基础标量及多维嵌套 Table(递归展平) Env, } impl Display for ValueShunt { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { match self { Self::Args => write!(f, "命令行参数 (args)"), Self::Env => write!(f, "环境变量 (env)"), } } } impl ValueShunt { fn expected_types_desc(&self) -> &str { match self { Self::Args => "string/number/boolean(不支持嵌套数组)", Self::Env => "string/number/boolean 或包含这些类型的数组/嵌套数组", } } /// 递归/标量收集实现 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 => { validate_each_sequence_item(&tbl, |item| self.collect_into(item, out))?; } other => { return Err(conversion_error(format!( "{} 不支持类型 {}(仅支持 {})", self, other.type_name(), self.expected_types_desc() ))); } } Ok(()) } } /// 校验环境变量名的合法性(Windows 约束:非空、不含 '='、不含 NUL) fn validate_env_var_name(key: &str) -> mlua::Result<()> { if key.is_empty() { return Err(conversion_error("环境变量名不能为空")); } if key.contains('=') { return Err(conversion_error(format!( "环境变量名 [{}] 不能包含 '='", key ))); } if key.contains('\0') { return Err(conversion_error(format!( "环境变量名 [{}] 不能包含 NUL 字符", key ))); } Ok(()) } /// 递归将任意 Lua Value 展开为扁平的字符串片段列表 fn collect_env_segments(value: Value, out: &mut Vec) -> mlua::Result<()> { match value { // 1. 字符串 Value::String(s) => out.push(lua_string_2_os_string(s)?), // 2. 整数与浮点数 Value::Integer(i) => out.push(OsString::from(i.to_string())), Value::Number(n) => { tracing::warn!(value = %n, "环境变量中的浮点数将按十进制格式转换为字符串"); out.push(OsString::from(n.to_string())); } // 3. 布尔值 Value::Boolean(b) => { tracing::warn!(value = %b, "环境变量中的布尔值将转换为字符串"); out.push(OsString::from(b.to_string())); } // 4. 表/数组:必须是连续整数下标的纯序列,递归解包(支持任意深度的嵌套数组) Value::Table(tbl) => { validate_each_sequence_item(&tbl, |item| collect_env_segments(item, out))? } // 5. 安全忽 // 略 nil Value::Nil => {} // 6. 无法转为环境变量的非法类型(函数、协程、UserData 等):报错而非静默忽略 other => { return Err(conversion_error(format!( "环境变量值不支持类型 {}(仅支持 string/number/boolean/数组)", other.type_name() ))); } } Ok(()) } #[derive(Debug, Clone, Default)] pub struct ShimConfig { pub target: PathBuf, pub args: Vec, pub env: HashMap, } impl ShimConfig { /// 根据配置快速构建准备执行的 Command 对象 pub fn to_command(&self, runtime_args: I) -> Command where I: IntoIterator, S: AsRef, { let mut cmd = Command::new(&self.target); // 1. 先追加配置里固定的默认参数 (例如: node, --max-old-space-size=4096) cmd.args(&self.args); // 2. 透传外部动态运行时参数 cmd.args(runtime_args); for (key, val) in &self.env { // 直接应用环境变量(Lua 端已经处理好字符串拼接或列表合并) cmd.env(key, val); } cmd } } /// 实现 FromLua Trait,由 mlua 自动处理 Table 转换 impl FromLua for ShimConfig { fn from_lua(value: Value, _lua: &Lua) -> mlua::Result { // 脚本返回必须是一个 Table 变体 let table = match value { Value::Table(t) => t, _ => { return Err(conversion_error(format!( "期望得到一个 Lua Table 配置对象,实际是 {}", value.type_name() ))); } }; // 必填字段: target(严格限定为字符串,避免数字被 mlua 宽松转为字符串后掩盖错误) let target = match table.get::("target")? { 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(); // 可选字段: args(缺失或 nil 时默认空列表;保留空字符串参数,与提权路径的 "" 语义一致) match table.get::>("args")? { None | Some(Value::Nil) => {} Some(Value::Table(tbl)) => { validate_each_sequence_item(&tbl, |item| { ValueShunt::Args.collect_into(item, &mut args) })?; } Some(other) => { return Err(conversion_error(format!( "args 必须是数组列表,实际类型是 {}", other.type_name() ))); } }; // 可选字段: 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?; validate_env_var_name(&name)?; let mut parts = Vec::new(); ValueShunt::Env.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) => { return Err(conversion_error(format!( "env 必须是键值表 (table),实际类型是 {}", other.type_name() ))); } } println!("环境变量结果:{:?}", env); Ok(ShimConfig { target: PathBuf::from(target), args, env, }) } } #[cfg(test)] mod tests { use super::*; fn parse(src: &str) -> mlua::Result { let lua = Lua::new(); // 模拟 runtime 注入的 get_env:返回按平台分隔符拆分的段数组(空变量返回空表) let get_env = lua .create_function(|lua, key: String| -> mlua::Result
{ let value = std::env::var(key).unwrap_or_default(); let segments: Vec = if value.is_empty() { Vec::new() } else { std::env::split_paths(std::ffi::OsStr::new(&value)) .map(|p| p.to_string_lossy().into_owned()) .collect() }; lua.create_sequence_from(segments) }) .unwrap(); lua.globals().set("get_env", get_env).unwrap(); let value = lua.load(src).eval::()?; let t = ShimConfig::from_lua(value, &lua); println!("读取出的数据:{:?}", t.clone()?); t } #[test] fn parses_basic_config() { let cfg = parse( r#" return { target = "C:/tools/git.exe", args = { "--no-pager" }, env = { PATH = { "C:/tools/git/bin", "C:/Windows",get_env("PATH") }, HOME = "C:/tools/home", CONST = 3, BOOL = true, MUT ={ A=3, B=true } } } "#, ) .unwrap(); assert_eq!(cfg.target, PathBuf::from("C:/tools/git.exe")); assert_eq!(cfg.args, vec!["--no-pager"]); assert_eq!(cfg.env.get("HOME").unwrap().to_str(), Some("C:/tools/home")); // PATH 前缀来自配置,随后附加 get_env("PATH") 拆出的宿主 PATH 段 let path = cfg.env.get("PATH").unwrap().to_str().unwrap(); assert!( path == "C:/tools/git/bin;C:/Windows" || path.starts_with("C:/tools/git/bin;C:/Windows;"), "unexpected PATH: {path}" ); } #[test] fn missing_args_and_env_are_empty() { let cfg = parse(r#"return { target = "t.exe" }"#).unwrap(); assert!(cfg.args.is_empty()); assert!(cfg.env.is_empty()); } #[test] fn keeps_empty_args() { let cfg = parse(r#"return { target = "t.exe", args = { "" } }"#).unwrap(); assert_eq!(cfg.args, vec![""]); } #[test] fn keeps_empty_env_value() { let cfg = parse(r#"return { target = "t.exe", env = { FOO = "" } }"#).unwrap(); assert_eq!(cfg.env.get("FOO").unwrap().to_str(), Some("")); } #[test] fn empty_array_clears_env_var() { let cfg = parse(r#"return { target = "t.exe", env = { PATH = {} } }"#).unwrap(); assert_eq!(cfg.env.get("PATH").unwrap().to_str(), Some("")); } #[test] fn expands_nested_env_array() { // get_env("PATH") 现在返回拆分后的段数组,嵌套表应被递归展开 let cfg = parse( r#"return { target = "t.exe", env = { PATH = { "C:/a", { "C:/b", "C:/c" } } } }"#, ) .unwrap(); assert_eq!( cfg.env.get("PATH").unwrap().to_str(), Some("C:/a;C:/b;C:/c") ); } #[test] fn rejects_wrong_env_type() { assert!(parse(r#"return { target = "t.exe", env = "PATH=C:/x" }"#).is_err()); } #[test] fn rejects_sparse_env_array() { assert!(parse(r#"return { target = "t.exe", env = { P = { "a", nil, "b" } } }"#).is_err()); } #[test] fn rejects_mixed_key_env_array() { assert!(parse(r#"return { target = "t.exe", env = { P = { a = "b" } } }"#).is_err()); } #[test] fn rejects_invalid_env_key() { assert!(parse(r#"return { target = "t.exe", env = { ["FOO=1"] = "x" } }"#).is_err()); assert!(parse(r#"return { target = "t.exe", env = { [""] = "x" } }"#).is_err()); } #[test] fn rejects_nul_in_env_value() { assert!(parse(r#"return { target = "t.exe", env = { P = { string.char(0) } } }"#).is_err()); } #[test] fn rejects_quote_in_env_value() { // Windows 的 join_paths 对含双引号的路径元素返回错误 assert!( parse(r#"return { target = "t.exe", env = { P = { string.char(34) } } }"#).is_err() ); } #[test] fn rejects_unsupported_env_value_type() { assert!(parse(r#"return { target = "t.exe", env = { F = function() end } }"#).is_err()); } #[test] fn missing_target_is_error() { assert!(parse(r#"return { args = { "x" } }"#).is_err()); } #[test] fn rejects_non_string_target() { assert!(parse(r#"return { target = 123 }"#).is_err()); assert!(parse(r#"return { target = false }"#).is_err()); } #[test] fn rejects_non_string_args_element() { assert!(parse(r#"return { target = "t.exe", args = { 1 } }"#).is_err()); } #[test] fn rejects_sparse_args() { assert!(parse(r#"return { target = "t.exe", args = { "a", nil, "b" } }"#).is_err()); } #[test] fn rejects_non_table_args() { assert!(parse(r#"return { target = "t.exe", args = "-B" }"#).is_err()); } }