fix(shim): 优化配置解析、跨平台编码与 UAC 提权逻辑

- 使用 `OsString`/`PathBuf` 替代 `String`,实现跨平台无损编码与命令行参数透传
- 完善 `FromLua` 转换与边界校验,拦截空洞 (`nil`) 及非整数键
- 修复环境变量拼接时由于内置分隔符导致的 `join_paths` 报错
- 优化 Windows 下 UAC 提权执行逻辑,改用宽字符 API (`ShellExecuteExW`)
This commit is contained in:
2026-08-18 11:15:30 +08:00
parent 5e69a6a980
commit e3cb065b35
7 changed files with 370 additions and 90 deletions

View File

@@ -1,25 +1,43 @@
use mlua::{FromLua, Lua, Table, Value};
use mlua::{FromLua, Lua, LuaString, ObjectLike, Table, Value};
use std::collections::HashMap;
use std::ffi::OsString;
use std::path::PathBuf;
use std::process::Command;
/// 构造一个带上下文的 FromLua 转换错误,便于定位配置问题
fn conversion_error(message: impl Into<String>) -> mlua::Error {
mlua::Error::FromLuaConversionError {
from: "Lua value".into(),
from: "Lua value",
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 lua_string_to_os_string(s: LuaString) -> Result<OsString, mlua::Error> {
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()))
}
}
@@ -74,20 +92,20 @@ fn validate_env_key(key: &str) -> mlua::Result<()> {
}
/// 递归将任意 Lua Value 展开为扁平的字符串片段列表
fn collect_env_strings(value: Value, out: &mut Vec<String>) -> mlua::Result<()> {
fn collect_env_strings(value: Value, out: &mut Vec<OsString>) -> mlua::Result<()> {
match value {
// 1. 字符串
Value::String(s) => out.push(lua_string_to_string(s)),
Value::String(s) => out.push(lua_string_to_os_string(s)?),
// 2. 整数与浮点数
Value::Integer(i) => out.push(i.to_string()),
Value::Integer(i) => out.push(OsString::from(i.to_string())),
Value::Number(n) => {
tracing::warn!(value = %n, "环境变量中的浮点数将按十进制格式转换为字符串");
out.push(n.to_string());
out.push(OsString::from(n.to_string()));
}
// 3. 布尔值
Value::Boolean(b) => {
tracing::warn!(value = %b, "环境变量中的布尔值将转换为字符串");
out.push(b.to_string());
out.push(OsString::from(b.to_string()));
}
// 4. 表/数组:必须是连续整数下标的纯序列,递归解包(支持任意深度的嵌套数组)
Value::Table(tbl) => for_each_sequence_item(&tbl, |item| collect_env_strings(item, out))?,
@@ -108,8 +126,8 @@ fn collect_env_strings(value: Value, out: &mut Vec<String>) -> mlua::Result<()>
#[derive(Debug, Clone, Default)]
pub struct ShimConfig {
pub target: PathBuf,
pub args: Vec<String>,
pub env: HashMap<String, String>,
pub args: Vec<OsString>,
pub env: HashMap<String, OsString>,
}
impl ShimConfig {
@@ -139,6 +157,7 @@ impl ShimConfig {
/// 实现 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,
_ => {
@@ -151,25 +170,41 @@ impl FromLua for ShimConfig {
// 必填字段: target严格限定为字符串避免数字被 mlua 宽松转为字符串后掩盖错误)
let target = match table.get::<Value>("target")? {
Value::String(s) => lua_string_to_string(s),
Value::String(s) => lua_string_to_os_string(s)?,
Value::Nil => {
return Err(conversion_error("缺少必填字段 target应为字符串路径"));
}
other => {
return Err(conversion_error(format!(
"target 必须是字符串,实际是 {}",
"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
// }
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));
args.push(lua_string_to_os_string(s)?);
Ok(())
}
other => Err(conversion_error(format!(
@@ -187,7 +222,7 @@ impl FromLua for ShimConfig {
}
};
// 可选字段: env只允许缺失/nil其他类型由 Option<Table> 转换报错,不再静默忽略)
let env_table: Option<Table> = table.get("env")?;
let env_table = table.get::<Option<Table>>("env")?;
println!("env_table:{:?}", env_table);
let mut env = HashMap::new();
if let Some(env_table) = env_table {
@@ -200,25 +235,36 @@ impl FromLua for ShimConfig {
// 空数组/空串也显式设置(空值表示清空该变量),
// 与“未配置该变量(继承宿主环境)”相区分
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))
let joined_os_str = std::env::join_paths(parts).map_err(|e| {
conversion_error(format!(
"环境变量 [{}] 的值无法用系统路径分隔符拼接: {}",
key, e
))
})?;
if joined_str.contains('\0') {
return Err(conversion_error(format!(
"环境变量 [{}] 的值不能包含 NUL 字符",
key
)));
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
)));
}
}
env.insert(key, joined_str);
#[cfg(windows)]
{
use std::os::windows::ffi::OsStrExt;
if joined_os_str.encode_wide().any(|c| c == 0) {
return Err(conversion_error(format!(
"环境变量 [{}] 的值不能包含 NUL 字符",
key
)));
}
}
env.insert(key, joined_os_str);
}
}
println!("环境变量结果:{:?}", env);
@@ -236,8 +282,26 @@ mod tests {
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>()?;
ShimConfig::from_lua(value, &lua)
let t = ShimConfig::from_lua(value, &lua);
println!("读取出的数据:{:?}", t.clone()?);
t
}
#[test]
@@ -248,8 +312,10 @@ mod tests {
target = "C:/tools/git.exe",
args = { "--no-pager" },
env = {
PATH = { "C:/tools/git/bin", "C:/Windows" },
PATH = { "C:/tools/git/bin", "C:/Windows",get_env("PATH") },
HOME = "C:/tools/home",
CONST = 3,
BOOL = true
}
}
"#,
@@ -257,13 +323,13 @@ mod tests {
.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")
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}"
);
}
@@ -283,13 +349,26 @@ mod tests {
#[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(""));
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").map(String::as_str), Some(""));
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]