refactor: 迁移 winapi 到 windows-sys,修复配置解析漏洞
- 依赖替换为 windows-sys 0.61,main.rs 全面适配新 API - 配置解析错误显式传播:env/args 类型错误、数组空洞、非法键不再静默 吞错 - 修正空环境变量与空参数语义,补充 UTF-8 与路径拼接校验 - require 容错移至 Rust 侧,模块加载失败记录日志并跳过 - 新增配置解析与运行时单元测试(19 个)
This commit is contained in:
405
src/config.rs
405
src/config.rs
@@ -1,29 +1,135 @@
|
||||
use mlua::{FromLua, Lua, Table, Value};
|
||||
use std::collections::HashMap;
|
||||
use std::path::PathBuf;
|
||||
use std::process::Command;
|
||||
use mlua::{FromLua, Lua, Table, Value};
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
/// 构造一个带上下文的 FromLua 转换错误,便于定位配置问题
|
||||
fn conversion_error(message: impl Into<String>) -> mlua::Error {
|
||||
mlua::Error::FromLuaConversionError {
|
||||
from: "Lua value".into(),
|
||||
to: "ShimConfig".into(),
|
||||
message: Some(message.into()),
|
||||
}
|
||||
}
|
||||
|
||||
/// 将 Lua 字符串转为 Rust 字符串;非 UTF-8 字节替换为 U+FFFD 并告警
|
||||
fn lua_string_to_string(s: mlua::LuaString) -> String {
|
||||
match s.to_str() {
|
||||
Ok(str_val) => str_val.to_string(),
|
||||
Err(_) => {
|
||||
tracing::warn!("配置字符串包含非 UTF-8 字节,已替换为 U+FFFD");
|
||||
s.to_string_lossy()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 按连续整数下标遍历表的序列部分;存在空洞或非序列键时返回错误,避免静默截断
|
||||
fn for_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::<Value, Value>() {
|
||||
let (key, _) = pair?;
|
||||
match key {
|
||||
Value::Integer(i) if i >= 1 && i < index => {}
|
||||
_ => {
|
||||
return Err(conversion_error(format!(
|
||||
"数组只能包含连续的整数下标 [1..{}],发现非序列键或空洞",
|
||||
index - 1
|
||||
)));
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 校验环境变量名的合法性(Windows 约束:非空、不含 '='、不含 NUL)
|
||||
fn validate_env_key(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_strings(value: Value, out: &mut Vec<String>) -> mlua::Result<()> {
|
||||
match value {
|
||||
// 1. 字符串
|
||||
Value::String(s) => out.push(lua_string_to_string(s)),
|
||||
// 2. 整数与浮点数
|
||||
Value::Integer(i) => out.push(i.to_string()),
|
||||
Value::Number(n) => {
|
||||
tracing::warn!(value = %n, "环境变量中的浮点数将按十进制格式转换为字符串");
|
||||
out.push(n.to_string());
|
||||
}
|
||||
// 3. 布尔值
|
||||
Value::Boolean(b) => {
|
||||
tracing::warn!(value = %b, "环境变量中的布尔值将转换为字符串");
|
||||
out.push(b.to_string());
|
||||
}
|
||||
// 4. 表/数组:必须是连续整数下标的纯序列,递归解包(支持任意深度的嵌套数组)
|
||||
Value::Table(tbl) => for_each_sequence_item(&tbl, |item| collect_env_strings(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_path: PathBuf,
|
||||
pub args: Option<Vec<String>>,
|
||||
pub envs: Option<HashMap<String, String>>,
|
||||
pub target: PathBuf,
|
||||
pub args: Vec<String>,
|
||||
pub env: HashMap<String, String>,
|
||||
}
|
||||
|
||||
impl ShimConfig {
|
||||
/// 根据配置快速构建准备执行的 Command 对象
|
||||
pub fn to_command(&self) -> Command {
|
||||
let mut cmd = Command::new(&self.target_path);
|
||||
pub fn to_command<I, S>(&self, runtime_args: I) -> Command
|
||||
where
|
||||
I: IntoIterator<Item = S>,
|
||||
S: AsRef<std::ffi::OsStr>,
|
||||
{
|
||||
let mut cmd = Command::new(&self.target);
|
||||
|
||||
if let Some(args) = &self.args {
|
||||
cmd.args(args);
|
||||
}
|
||||
// 1. 先追加配置里固定的默认参数 (例如: node, --max-old-space-size=4096)
|
||||
cmd.args(&self.args);
|
||||
|
||||
if let Some(envs) = &self.envs {
|
||||
for (key, val) in envs {
|
||||
// 直接应用环境变量(Lua 端已经处理好字符串拼接或列表合并)
|
||||
cmd.env(key, val);
|
||||
}
|
||||
// 2. 透传外部动态运行时参数
|
||||
cmd.args(runtime_args);
|
||||
|
||||
for (key, val) in &self.env {
|
||||
// 直接应用环境变量(Lua 端已经处理好字符串拼接或列表合并)
|
||||
cmd.env(key, val);
|
||||
}
|
||||
|
||||
cmd
|
||||
@@ -33,60 +139,221 @@ impl ShimConfig {
|
||||
/// 实现 FromLua Trait,由 mlua 自动处理 Table 转换
|
||||
impl FromLua for ShimConfig {
|
||||
fn from_lua(value: Value, _lua: &Lua) -> mlua::Result<Self> {
|
||||
match value {
|
||||
Value::Table(table) => {
|
||||
let path_str: String = table.get("path")?;
|
||||
let args: Option<Vec<String>> = table.get("args")?;
|
||||
// 解析 env Table
|
||||
let mut envs_map = HashMap::new();
|
||||
if let Ok(env_table) = table.get::<Table>("env") {
|
||||
// 获取当前系统的路径分隔符(Windows 为 ";",Unix 为 ":")
|
||||
#[cfg(windows)]
|
||||
let sep = ";";
|
||||
#[cfg(not(windows))]
|
||||
let sep = ":";
|
||||
|
||||
for pair in env_table.pairs::<String, Value>() {
|
||||
let (k, v) = pair?;
|
||||
match v {
|
||||
// 情况 1: 普通字符串,如 HOME = "C:/path" -> 直接覆盖
|
||||
Value::String(s) => {
|
||||
envs_map.insert(k, s.to_str()?.to_string());
|
||||
}
|
||||
// 情况 2: 数组 Table,如 PATH = { bin_dir, get_env("PATH") }
|
||||
Value::Table(arr) => {
|
||||
let paths: Vec<String> = arr
|
||||
// 将 arr 作为序列(数组)处理,每个元素转为 String
|
||||
.sequence_values::<String>()
|
||||
.filter_map(|r| r.ok())
|
||||
.filter(|s| !s.is_empty()) // 过滤空串,防止生成不必要的连续 ;;
|
||||
.collect();
|
||||
|
||||
let combined = paths.join(sep);
|
||||
envs_map.insert(k, combined);
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let envs = if envs_map.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(envs_map)
|
||||
};
|
||||
|
||||
Ok(ShimConfig {
|
||||
target_path: PathBuf::from(path_str),
|
||||
args,
|
||||
envs,
|
||||
})
|
||||
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::<Value>("target")? {
|
||||
Value::String(s) => lua_string_to_string(s),
|
||||
Value::Nil => {
|
||||
return Err(conversion_error("缺少必填字段 target(应为字符串路径)"));
|
||||
}
|
||||
other => {
|
||||
return Err(conversion_error(format!(
|
||||
"target 必须是字符串,实际是 {}",
|
||||
other.type_name()
|
||||
)));
|
||||
}
|
||||
};
|
||||
// 可选字段: args(缺失或 nil 时默认空列表;保留空字符串参数,与提权路径的 "" 语义一致)
|
||||
let args = match table.get::<Option<Value>>("args")? {
|
||||
None => Vec::new(),
|
||||
Some(Value::Table(tbl)) => {
|
||||
let mut args = Vec::new();
|
||||
for_each_sequence_item(&tbl, |item| match item {
|
||||
Value::String(s) => {
|
||||
args.push(lua_string_to_string(s));
|
||||
Ok(())
|
||||
}
|
||||
other => Err(conversion_error(format!(
|
||||
"args 数组元素必须是字符串,实际是 {}",
|
||||
other.type_name()
|
||||
))),
|
||||
})?;
|
||||
args
|
||||
}
|
||||
Some(other) => {
|
||||
return Err(conversion_error(format!(
|
||||
"args 必须是字符串数组,实际是 {}",
|
||||
other.type_name()
|
||||
)));
|
||||
}
|
||||
};
|
||||
// 可选字段: env(只允许缺失/nil,其他类型由 Option<Table> 转换报错,不再静默忽略)
|
||||
let env_table: Option<Table> = table.get("env")?;
|
||||
println!("env_table:{:?}", env_table);
|
||||
let mut env = HashMap::new();
|
||||
if let Some(env_table) = env_table {
|
||||
for pair in env_table.pairs::<String, Value>() {
|
||||
let (key, value) = pair?;
|
||||
validate_env_key(&key)?;
|
||||
|
||||
let mut parts = Vec::new();
|
||||
collect_env_strings(value, &mut parts)?;
|
||||
|
||||
// 空数组/空串也显式设置(空值表示清空该变量),
|
||||
// 与“未配置该变量(继承宿主环境)”相区分
|
||||
let joined_os_str =
|
||||
std::env::join_paths(parts.iter().map(PathBuf::from)).map_err(|e| {
|
||||
conversion_error(format!(
|
||||
"环境变量 [{}] 的值无法用系统路径分隔符拼接: {}",
|
||||
key, e
|
||||
))
|
||||
})?;
|
||||
println!("环境变量拼接结果:{:?}", joined_os_str);
|
||||
|
||||
let joined_str = joined_os_str.into_string().map_err(|_| {
|
||||
conversion_error(format!("环境变量 [{}] 的值不是合法文本", key))
|
||||
})?;
|
||||
if joined_str.contains('\0') {
|
||||
return Err(conversion_error(format!(
|
||||
"环境变量 [{}] 的值不能包含 NUL 字符",
|
||||
key
|
||||
)));
|
||||
}
|
||||
env.insert(key, joined_str);
|
||||
}
|
||||
_ => Err(mlua::Error::FromLuaConversionError {
|
||||
from: value.type_name(),
|
||||
to: "ShimConfig".into(),
|
||||
message: Some("Expected a Lua table".to_string()),
|
||||
}),
|
||||
}
|
||||
println!("环境变量结果:{:?}", env);
|
||||
Ok(ShimConfig {
|
||||
target: PathBuf::from(target),
|
||||
args,
|
||||
env,
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
fn parse(src: &str) -> mlua::Result<ShimConfig> {
|
||||
let lua = Lua::new();
|
||||
let value = lua.load(src).eval::<Value>()?;
|
||||
ShimConfig::from_lua(value, &lua)
|
||||
}
|
||||
|
||||
#[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" },
|
||||
HOME = "C:/tools/home",
|
||||
}
|
||||
}
|
||||
"#,
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(cfg.target, PathBuf::from("C:/tools/git.exe"));
|
||||
assert_eq!(cfg.args, vec!["--no-pager"]);
|
||||
assert_eq!(
|
||||
cfg.env.get("HOME").map(String::as_str),
|
||||
Some("C:/tools/home")
|
||||
);
|
||||
assert_eq!(
|
||||
cfg.env.get("PATH").map(String::as_str),
|
||||
Some("C:/tools/git/bin;C:/Windows")
|
||||
);
|
||||
}
|
||||
|
||||
#[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").map(String::as_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").map(String::as_str), Some(""));
|
||||
}
|
||||
|
||||
#[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());
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user