Files
mirror/src/config.rs
CNWei f59040c11e fix(config): 优化配置解析
- 完善 `FromLua` 转换与边界校验,拦截空洞 (`nil`) 及非整数键
2026-08-19 11:34:16 +08:00

486 lines
17 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

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<String>) -> 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<OsString> {
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::<Value, Value>() {
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<OsString>) -> 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<OsString>) -> 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<OsString>,
pub env: HashMap<String, OsString>,
}
impl ShimConfig {
/// 根据配置快速构建准备执行的 Command 对象
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);
// 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<Self> {
// 脚本返回必须是一个 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::<Value>("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::<Option<Value>>("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<Table> 转换报错,不再静默忽略)
let env_table = table.get::<Option<Value>>("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::<String, Value>() {
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<ShimConfig> {
let lua = Lua::new();
// 模拟 runtime 注入的 get_env返回按平台分隔符拆分的段数组空变量返回空表
let get_env = lua
.create_function(|lua, key: String| -> mlua::Result<Table> {
let value = std::env::var(key).unwrap_or_default();
let segments: Vec<String> = 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::<Value>()?;
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());
}
}