refactor(validators): 优化别名解析流程

- 新增 parse_tokens 按空格切分字符串
- 其他优化
This commit is contained in:
2026-08-26 17:47:21 +08:00
parent 9e1051737a
commit 219a6b3d51
5 changed files with 288 additions and 73 deletions

10
Cargo.lock generated
View File

@@ -88,15 +88,6 @@ version = "1.17.0"
source = "registry+https://github.com/rust-lang/crates.io-index" source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "9e5e8f6c15a24b9a3ee5efec809ccd006d3b30e8b3bb63c39af737c7f87daa1d" checksum = "9e5e8f6c15a24b9a3ee5efec809ccd006d3b30e8b3bb63c39af737c7f87daa1d"
[[package]]
name = "encoding_rs"
version = "0.8.35"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "75030f3c4f45dafd7586dd6780965a8c7e8e285a5ecb86713e63a79c5b2766f3"
dependencies = [
"cfg-if",
]
[[package]] [[package]]
name = "find-msvc-tools" name = "find-msvc-tools"
version = "0.1.11" version = "0.1.11"
@@ -170,7 +161,6 @@ version = "0.1.0"
dependencies = [ dependencies = [
"anyhow", "anyhow",
"dunce", "dunce",
"encoding_rs",
"mlua", "mlua",
"tinyjson", "tinyjson",
"tracing", "tracing",

View File

@@ -22,6 +22,7 @@ windows-sys = { version = "0.61.2", features = [
"Win32_System_Threading", "Win32_System_Threading",
"Win32_UI_Shell", "Win32_UI_Shell",
"Win32_UI_WindowsAndMessaging", "Win32_UI_WindowsAndMessaging",
"Win32_Globalization"
] } ] }
dunce = "1.0.5" dunce = "1.0.5"
@@ -29,7 +30,6 @@ dunce = "1.0.5"
tracing = "0.1.44" tracing = "0.1.44"
tracing-subscriber = { version = "0.3", features = ["env-filter"] } tracing-subscriber = { version = "0.3", features = ["env-filter"] }
tracing-appender = "0.2" tracing-appender = "0.2"
encoding_rs = "0.8.35"
tinyjson="2.5.1" tinyjson="2.5.1"

View File

@@ -15,7 +15,7 @@ fn main() {
set_console_ctrl_handler(); set_console_ctrl_handler();
// 3. 解析调用参数与代理 Shim 配置 // 3. 解析调用参数与代理 Shim 配置
let calling_args: Vec<_> = env::args_os().skip(1).collect(); let calling_args: Vec<_> = env::args_os().skip(1).collect();
let shim = match Mirror::load() { let mr = match Mirror::new() {
Ok(v) => v, Ok(v) => v,
Err(e) => { Err(e) => {
eprintln!("加载代理(shim)配置时发生错误: {}", e); eprintln!("加载代理(shim)配置时发生错误: {}", e);
@@ -23,21 +23,21 @@ fn main() {
} }
}; };
// 构建 Command复用 ShimConfig::to_command含 args/env 注入),避免重复逻辑 // 构建 Command复用 ShimConfig::to_command含 args/env 注入),避免重复逻辑
let mut cmd = shim.to_command(&calling_args); let mut cmd = mr.to_command(&calling_args);
let mut child = match cmd.spawn() { let mut child = match cmd.spawn() {
Ok(v) => v, Ok(v) => v,
Err(e) if e.raw_os_error() == Some(ERROR_ELEVATION_REQUIRED) => { Err(e) if e.raw_os_error() == Some(ERROR_ELEVATION_REQUIRED) => {
// 提权回退时需要完整参数:配置默认参数 + 调用方透传参数 // 提权回退时需要完整参数:配置默认参数 + 调用方透传参数
let mut args = shim.args.clone(); let mut args = mr.spec.args.clone();
args.extend_from_slice(&calling_args); args.extend_from_slice(&calling_args);
exit(execute_elevated(&shim.target, &args, Some(&shim.env))) exit(execute_elevated(&mr.spec.target, &args, Some(&mr.spec.env)))
} }
Err(e) => { Err(e) => {
eprintln!( eprintln!(
"启动目标程序 [{}] 时发生错误: {}", "启动目标程序 [{}] 时发生错误: {}",
shim.target.to_string_lossy(), mr.spec.target.to_string_lossy(),
e e
); );
exit(EXIT_FAILED_SPAWN_PROG); exit(EXIT_FAILED_SPAWN_PROG);
@@ -49,7 +49,7 @@ fn main() {
Err(e) => { Err(e) => {
eprintln!( eprintln!(
"等待目标程序 [{}] 执行完毕时发生错误: {}", "等待目标程序 [{}] 执行完毕时发生错误: {}",
shim.target.to_string_lossy(), mr.spec.target.to_string_lossy(),
e e
); );
exit(EXIT_FAILED_WAIT_PROG); exit(EXIT_FAILED_WAIT_PROG);

View File

@@ -5,6 +5,7 @@ use std::ffi::OsStr;
use std::path::Path; use std::path::Path;
use std::{env, fs}; use std::{env, fs};
use crate::utils::normalize_path_for_lua; use crate::utils::normalize_path_for_lua;
use crate::validators::map_lua_error;
/// 将 Path 转换为适合 Lua 使用的安全字符串路径 /// 将 Path 转换为适合 Lua 使用的安全字符串路径
@@ -139,6 +140,7 @@ impl LuaRuntime {
.set_name(&chunk_name) .set_name(&chunk_name)
.eval::<T>() .eval::<T>()
// .map_err(|e| anyhow!(e.to_string())) // .map_err(|e| anyhow!(e.to_string()))
.map_err(map_lua_error)
.with_context(|| format!("执行 Lua 配置文件失败: {}", path.display())) .with_context(|| format!("执行 Lua 配置文件失败: {}", path.display()))
} }

View File

@@ -15,10 +15,28 @@ fn conversion_error(message: impl Into<String>) -> mlua::Error {
message: Some(message.into()), message: Some(message.into()),
} }
} }
/// 将底层 Lua 语法错误转换为用户友好的提示
pub fn map_lua_error(err: mlua::Error) -> mlua::Error {
match &err {
mlua::Error::SyntaxError { message, .. } => {
if message.contains("invalid escape sequence") || message.contains("unfinished string") {
return mlua::Error::RuntimeError(format!(
"配置文件语法错误:检测到非法的字符串转义。\n\
提示:在 Windows 路径末尾或字符串中使用反斜杠 '\\' 时:\n\
1. 请使用双反斜杠转义,例如: \"D:\\\\CNWei\\\\CNW\\\\Rust\\\\\"\n\
2. 或使用 Lua 原始字符串 (Raw String),例如: [[D:\\CNWei\\CNW\\Rust\\]]\n\
底层错误: {}",
message
));
}
}
_ => {}
}
err
}
/// 将 Lua 字符串转为 Rust 字符串;非 UTF-8 按系统默认的 ANSI/OEM (如 GBK) 进行安全解码 /// 将 Lua 字符串转为 Rust 字符串;非 UTF-8 按系统默认的 ANSI/OEM (如 GBK) 进行安全解码
fn lua_string_2_os_string(s: &LuaString) -> mlua::Result<OsString> { fn lua_string_2_os_string(s: &LuaString) -> mlua::Result<OsString> {
let raw_bytes = &s.as_bytes().to_vec(); let raw_bytes = &s.as_bytes();
// 1. Unix 平台:直接零拷贝透传原始字节(无损支持任意编码) // 1. Unix 平台:直接零拷贝透传原始字节(无损支持任意编码)
#[cfg(unix)] #[cfg(unix)]
{ {
@@ -33,16 +51,44 @@ fn lua_string_2_os_string(s: &LuaString) -> mlua::Result<OsString> {
return Ok(OsString::from(utf8_str)); return Ok(OsString::from(utf8_str));
} }
// 非 UTF-8 时,按系统默认的 ANSI/OEM (如 GBK) 进行安全解码 unsafe {
let (cow, _, had_errors) = encoding_rs::GBK.decode(raw_bytes); use windows_sys::Win32::Globalization::{MultiByteToWideChar, CP_ACP, MB_ERR_INVALID_CHARS};
if had_errors { use std::os::windows::ffi::OsStringExt;
return Err(mlua::Error::RuntimeError(
"路径包含无效/不支持的字节编码".into(), if raw_bytes.is_empty() {
)); return Ok(OsString::new());
} }
Ok(OsString::from(cow.as_ref()))
let len = MultiByteToWideChar(
CP_ACP,
MB_ERR_INVALID_CHARS,
raw_bytes.as_ptr(),
raw_bytes.len() as i32,
std::ptr::null_mut(),
0,
);
if len <= 0 {
return Err(mlua::Error::FromLuaConversionError {
from: "LuaString",
to: "OsString".to_string(),
message: Some("字符串包含无效或当前系统无法识别的编码字节".to_string()),
});
}
let mut buf = vec![0u16; len as usize];
MultiByteToWideChar(
CP_ACP,
MB_ERR_INVALID_CHARS,
raw_bytes.as_ptr(),
raw_bytes.len() as i32,
buf.as_mut_ptr(),
len,
);
Ok(OsString::from_wide(&buf))
} }
} }}
/// Lua 值校验器:针对不同上下文定义校验规则 /// Lua 值校验器:针对不同上下文定义校验规则
#[derive(Debug, Clone, Copy, PartialEq, Eq)] #[derive(Debug, Clone, Copy, PartialEq, Eq)]
@@ -151,21 +197,34 @@ impl LuaValidator {
} }
/// 【纯粹数据转换与展开】将已通过校验的 Value 递归解析为 OsString 动态数组 /// 【纯粹数据转换与展开】将已通过校验的 Value 递归解析为 OsString 动态数组
fn collect_value_into( fn collect_value_into(&self, value: &Value, out: &mut Vec<OsString>) -> mlua::Result<()> {
context: &impl fmt::Display,
value: &Value,
out: &mut Vec<OsString>,
) -> mlua::Result<()> {
match value { match value {
Value::Nil => {} Value::Nil => {}
Value::String(s) => out.push(lua_string_2_os_string(s)?), Value::String(s) => {
let os_str = lua_string_2_os_string(s)?;
match self {
Self::Args | Self::Aliases => {
if let Some(str_ref) = os_str.to_str() {
// 使用 Tokenizer 切分空格与引号
out.extend(Self::parse_tokens(str_ref));
} else {
// 对于无法转为 UTF-8 的特殊二进制数据,作为整体追加
out.push(os_str);
}
}
Self::Env => {
out.push(os_str);
}
Self::Target => {}
}
}
Value::Integer(i) => out.push(OsString::from(i.to_string())), Value::Integer(i) => out.push(OsString::from(i.to_string())),
Value::Number(n) => { Value::Number(n) => {
tracing::warn!(%context, value = %n, "浮点数将按十进制格式转换为字符串"); tracing::warn!(%self, value = %n, "浮点数将按十进制格式转换为字符串");
out.push(OsString::from(n.to_string())); out.push(OsString::from(n.to_string()));
} }
Value::Boolean(b) => { Value::Boolean(b) => {
tracing::warn!(%context, value = %b, "布尔值将转换为字符串"); tracing::warn!(%self, value = %b, "布尔值将转换为字符串");
out.push(OsString::from(b.to_string())); out.push(OsString::from(b.to_string()));
} }
Value::Table(tbl) => { Value::Table(tbl) => {
@@ -175,7 +234,7 @@ impl LuaValidator {
if matches!(item, Value::Nil) { if matches!(item, Value::Nil) {
break; break;
} }
Self::collect_value_into(context, &item, out)?; Self::collect_value_into(self, &item, out)?;
index += 1; index += 1;
} }
} }
@@ -183,7 +242,36 @@ impl LuaValidator {
} }
Ok(()) Ok(())
} }
fn parse_name(name: &Value) -> mlua::Result<String> {
let name_str = match name {
Value::String(s) => match s.to_str() {
Ok(str_ref) => str_ref.to_string(),
Err(_) => return Err(conversion_error("环境变量名必须是合法的 UTF-8 字符串")),
},
other => {
return Err(conversion_error(format!(
"环境变量键名类型错误:期望 string实际是 {}",
other.type_name()
)));
}
};
if name_str.is_empty() {
return Err(conversion_error("环境变量名不能为空"));
}
if name_str.contains('=') {
return Err(conversion_error(format!(
"环境变量名 [{}] 不能包含 '='",
name_str
)));
}
if name_str.contains('\0') {
return Err(conversion_error(format!(
"环境变量名 [{}] 不能包含 NUL 字符",
name_str
)));
}
Ok(name_str)
}
fn parse_raw_val(&self, name: &str, value: &Value) -> mlua::Result<()> { fn parse_raw_val(&self, name: &str, value: &Value) -> mlua::Result<()> {
match value { match value {
Value::Nil Value::Nil
@@ -254,34 +342,7 @@ impl LuaValidator {
/// 校验环境变量名的合法性Windows 约束:非空、不含 '='、不含 NUL /// 校验环境变量名的合法性Windows 约束:非空、不含 '='、不含 NUL
fn parse_env_name(name: &Value) -> mlua::Result<String> { fn parse_env_name(name: &Value) -> mlua::Result<String> {
let name_str = match name { Self::parse_name(name)
Value::String(s) => match s.to_str() {
Ok(str_ref) => str_ref.to_string(),
Err(_) => return Err(conversion_error("环境变量名必须是合法的 UTF-8 字符串")),
},
other => {
return Err(conversion_error(format!(
"环境变量键名类型错误:期望 string实际是 {}",
other.type_name()
)));
}
};
if name_str.is_empty() {
return Err(conversion_error("环境变量名不能为空"));
}
if name_str.contains('=') {
return Err(conversion_error(format!(
"环境变量名 [{}] 不能包含 '='",
name_str
)));
}
if name_str.contains('\0') {
return Err(conversion_error(format!(
"环境变量名 [{}] 不能包含 NUL 字符",
name_str
)));
}
Ok(name_str)
} }
/// 解析、校验并拼接单个环境变量键值对 (`key`, `value`) -> (`OsString`) /// 解析、校验并拼接单个环境变量键值对 (`key`, `value`) -> (`OsString`)
@@ -292,10 +353,10 @@ impl LuaValidator {
) -> mlua::Result<OsString> { ) -> mlua::Result<OsString> {
context.parse_raw_val(name, raw_val)?; context.parse_raw_val(name, raw_val)?;
let ctx = format!("{} [{}]", context, name); // let ctx = format!("{} [{}]", context, name);
let mut parts = Vec::new(); let mut parts = Vec::new();
Self::collect_value_into(&ctx, raw_val, &mut parts)?; Self::collect_value_into(&context, raw_val, &mut parts)?;
// 校验每个展开元素的 NUL 字符 // 校验每个展开元素的 NUL 字符
for part in &parts { for part in &parts {
@@ -303,7 +364,7 @@ impl LuaValidator {
} }
// 3. 使用系统路径分隔符拼接数组列表 // 3. 使用系统路径分隔符拼接数组列表
let joined_os_str = std::env::join_paths(parts).map_err(|e| { let joined_os_str = std::env::join_paths(parts).map_err(|e| {
conversion_error(format!("{} 的值无法用系统路径分隔符拼接: {}", ctx, e)) conversion_error(format!("{} 的值无法用系统路径分隔符拼接: {}", context, e))
})?; })?;
Ok(joined_os_str) Ok(joined_os_str)
} }
@@ -344,6 +405,10 @@ impl LuaValidator {
} }
impl LuaValidator { impl LuaValidator {
/// 校验环境变量名的合法性Windows 约束:非空、不含 '='、不含 NUL
fn parse_aliases_name(name: &Value) -> mlua::Result<String> {
Self::parse_name(name)
}
fn parse_aliases_val( fn parse_aliases_val(
context: &LuaValidator, context: &LuaValidator,
name: &str, name: &str,
@@ -354,7 +419,7 @@ impl LuaValidator {
let ctx = format!("{} [{}]", context, name); let ctx = format!("{} [{}]", context, name);
let mut parts = Vec::new(); let mut parts = Vec::new();
Self::collect_value_into(&ctx, raw_val, &mut parts)?; Self::collect_value_into(&context, raw_val, &mut parts)?;
Ok(parts) Ok(parts)
} }
/// 解析并打平别名表 (aliases) /// 解析并打平别名表 (aliases)
@@ -378,7 +443,7 @@ impl LuaValidator {
for pair in table.pairs::<Value, Value>() { for pair in table.pairs::<Value, Value>() {
let (raw_key, raw_val) = pair?; let (raw_key, raw_val) = pair?;
let key = Self::parse_env_name(&raw_key)?; let key = Self::parse_aliases_name(&raw_key)?;
let val = Self::parse_aliases_val(&ctx, &key, &raw_val)?; let val = Self::parse_aliases_val(&ctx, &key, &raw_val)?;
raw_aliases.insert(key, val); raw_aliases.insert(key, val);
@@ -463,6 +528,69 @@ impl LuaValidator {
} }
} }
} }
/// 将输入的字符串按 Shell 规则切分为独立的 CLI 参数 Token
/// - 自动过滤连续空格
/// - 支持单引号 `'...'` 和双引号 `"..."` 包裹包含空格的参数
fn parse_tokens(input: &str) -> Vec<OsString> {
let mut tokens = Vec::new();
let mut current_token = String::new();
let mut in_quote: Option<char> = None;
let mut chars = input.chars().peekable();
while let Some(ch) = chars.next() {
match (ch, in_quote) {
// 处理转义字符 (例如 \")
('\\', quote) => {
if let Some(&next_ch) = chars.peek() {
let should_escape = if cfg!(windows) {
// Windows 策略:只有在转义引号、反斜杠本身时才剥离 \
// (如果在双引号内部,空格也不应该被 \ 转义)
next_ch == '"' || next_ch == '\'' || next_ch == '\\'
} else {
// Unix 策略:标准 Shell 转义(引号、反斜杠、空格等)
next_ch == '"' || next_ch == '\'' || next_ch == '\\' || next_ch.is_whitespace()
};
if should_escape {
chars.next(); // 消耗掉下一个字符
current_token.push(next_ch);
} else {
// 保留 Windows 路径分隔符或未知转义中的 \
current_token.push('\\');
}
} else {
// 结尾孤立的 \
current_token.push('\\');
}
}
// 遇到引号:开启或关闭引号包裹
('"' | '\'', None) => {
in_quote = Some(ch);
}
('"' | '\'', Some(q)) if q == ch => {
in_quote = None;
}
// 引号外部遇到空白字符:切分出一个完整的 Token
(ch, None) if ch.is_whitespace() => {
if !current_token.is_empty() {
tokens.push(OsString::from(std::mem::take(&mut current_token)));
}
}
// 其他字符或引号内部字符:直接追加
(ch, _) => {
current_token.push(ch);
}
}
}
// 收尾最后一个 Token
if !current_token.is_empty() {
tokens.push(OsString::from(current_token));
}
tokens
}
} }
/// 专用于 JSON (tinyjson) 的类型校验与字段提取器 /// 专用于 JSON (tinyjson) 的类型校验与字段提取器
@@ -612,7 +740,7 @@ mod tests {
HOME = "C:/tools/home", HOME = "C:/tools/home",
CONST = 3, CONST = 3,
BOOL = true, BOOL = true,
} },
aliases = { aliases = {
-- 1. 整体字符串(会自动 trim 首尾空格,保留为单个参数) -- 1. 整体字符串(会自动 trim 首尾空格,保留为单个参数)
st = "status -s", st = "status -s",
@@ -626,7 +754,7 @@ mod tests {
-- 4. 嵌套别名组合(加载期会自动展开并进行死环检测) -- 4. 嵌套别名组合(加载期会自动展开并进行死环检测)
base_log = "log --graph", base_log = "log --graph",
all_log = { "mr:base_log", "--all" }, all_log = { "mr:base_log", "--all","D:\\CNWei\\CNW\\Rust\\rshim\\target\\debug\\build\\mlua-sys-f33759261acaca16\\out\\lib","D:/CNWei/CNW/Rust/" ,[[D:\CNWei\CNW\Rust\]]},
-- 5. nil 或空字符串(解析为空参数列表) -- 5. nil 或空字符串(解析为空参数列表)
empty_alias = nil, empty_alias = nil,
@@ -799,3 +927,98 @@ mod tests {
assert!(parse(r#"return { target = "t.exe", args = "-B" }"#).is_err()); assert!(parse(r#"return { target = "t.exe", args = "-B" }"#).is_err());
} }
} }
#[cfg(test)]
mod tests2 {
use super::*;
use std::ffi::OsString;
/// 辅助宏:简化声明与断言对比
macro_rules! assert_tokens {
($input:expr, $expected:expr) => {
let actual = LuaValidator::parse_tokens($input);
let expected_os: Vec<OsString> = $expected.into_iter().map(OsString::from).collect();
assert_eq!(
actual, expected_os,
"\n测试输入: {:?}\n期望输出: {:?}\n实际输出: {:?}",
$input, expected_os, actual
);
};
}
#[test]
fn test_parse_tokens_basic_split() {
// 场景 1基础多参数拆分空格分隔
assert_tokens!("cargo run --verbose", vec!["cargo", "run", "--verbose"]);
assert_tokens!("git status", vec!["git", "status"]);
}
#[test]
fn test_parse_tokens_multi_alias_ref() {
// 场景 2多别名混合与组合引用
assert_tokens!(
"mr:run --bin mr:base_flags",
vec!["mr:run", "--bin", "mr:base_flags"]
);
assert_tokens!(
"mr:app1 mr:app2 --flag",
vec!["mr:app1", "mr:app2", "--flag"]
);
}
#[test]
fn test_parse_tokens_continuous_whitespaces() {
// 场景 3连续多空格与制表符过滤
assert_tokens!(
"mr:run --bin \t my_app",
vec!["mr:run", "--bin", "my_app"]
);
assert_tokens!(" cargo build ", vec!["cargo", "build"]);
}
#[test]
fn test_parse_tokens_double_quotes() {
// 场景 4双引号包裹包含空格的参数
assert_tokens!(
"git commit -m \"fix a bug\"",
vec!["git", "commit", "-m", "fix a bug"]
);
assert_tokens!("echo \"hello world\"", vec!["echo", "hello world"]);
}
#[test]
fn test_parse_tokens_single_quotes() {
// 场景 5单引号包裹包含空格的参数
assert_tokens!(
"gcc -O2 'my file.c' -o app",
vec!["gcc", "-O2", "my file.c", "-o", "app"]
);
assert_tokens!(
"python 'script with space.py'",
vec!["python", "script with space.py"]
);
}
#[test]
fn test_parse_tokens_escaped_characters() {
// 场景 6反斜杠转义字符
assert_tokens!("echo hello\\ world", vec!["echo", "hello world"]);
assert_tokens!(
"echo \"hello \\\"world\\\"\"",
vec!["echo", "hello \"world\""]
);
}
#[test]
fn test_parse_tokens_single_scalar_and_edge_cases() {
// 场景 7单标量参数与边界情况
assert_tokens!("git", vec!["git"]);
assert_tokens!("8000", vec!["8000"]);
assert_tokens!("", Vec::<&str>::new());
assert_tokens!(" ", Vec::<&str>::new());
}
#[test]
fn test_parse_tokens_unclosed_quotes() {
// 场景 8未闭合引号的容错处理会尽量追加到当前 Token 中)
assert_tokens!("echo \"hello world", vec!["echo", "hello world"]);
}
}