486 lines
17 KiB
Rust
486 lines
17 KiB
Rust
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());
|
||
}
|
||
}
|