From 219a6b3d515cdf26d19d7ff7d3863819408c3507 Mon Sep 17 00:00:00 2001 From: CNWei Date: Wed, 26 Aug 2026 17:47:21 +0800 Subject: [PATCH] =?UTF-8?q?refactor(validators):=20=E4=BC=98=E5=8C=96?= =?UTF-8?q?=E5=88=AB=E5=90=8D=E8=A7=A3=E6=9E=90=E6=B5=81=E7=A8=8B?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 新增 parse_tokens 按空格切分字符串 - 其他优化 --- Cargo.lock | 10 -- Cargo.toml | 2 +- src/main.rs | 12 +- src/runtime.rs | 2 + src/validators.rs | 335 ++++++++++++++++++++++++++++++++++++++-------- 5 files changed, 288 insertions(+), 73 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index 4cd5c96..e8f1b39 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -88,15 +88,6 @@ version = "1.17.0" source = "registry+https://github.com/rust-lang/crates.io-index" 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]] name = "find-msvc-tools" version = "0.1.11" @@ -170,7 +161,6 @@ version = "0.1.0" dependencies = [ "anyhow", "dunce", - "encoding_rs", "mlua", "tinyjson", "tracing", diff --git a/Cargo.toml b/Cargo.toml index 4da076e..f24e3cb 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -22,6 +22,7 @@ windows-sys = { version = "0.61.2", features = [ "Win32_System_Threading", "Win32_UI_Shell", "Win32_UI_WindowsAndMessaging", + "Win32_Globalization" ] } dunce = "1.0.5" @@ -29,7 +30,6 @@ dunce = "1.0.5" tracing = "0.1.44" tracing-subscriber = { version = "0.3", features = ["env-filter"] } tracing-appender = "0.2" -encoding_rs = "0.8.35" tinyjson="2.5.1" diff --git a/src/main.rs b/src/main.rs index c879c80..0c7c1d3 100644 --- a/src/main.rs +++ b/src/main.rs @@ -15,7 +15,7 @@ fn main() { set_console_ctrl_handler(); // 3. 解析调用参数与代理 Shim 配置 let calling_args: Vec<_> = env::args_os().skip(1).collect(); - let shim = match Mirror::load() { + let mr = match Mirror::new() { Ok(v) => v, Err(e) => { eprintln!("加载代理(shim)配置时发生错误: {}", e); @@ -23,21 +23,21 @@ fn main() { } }; // 构建 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() { Ok(v) => v, 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); - exit(execute_elevated(&shim.target, &args, Some(&shim.env))) + exit(execute_elevated(&mr.spec.target, &args, Some(&mr.spec.env))) } Err(e) => { eprintln!( "启动目标程序 [{}] 时发生错误: {}", - shim.target.to_string_lossy(), + mr.spec.target.to_string_lossy(), e ); exit(EXIT_FAILED_SPAWN_PROG); @@ -49,7 +49,7 @@ fn main() { Err(e) => { eprintln!( "等待目标程序 [{}] 执行完毕时发生错误: {}", - shim.target.to_string_lossy(), + mr.spec.target.to_string_lossy(), e ); exit(EXIT_FAILED_WAIT_PROG); diff --git a/src/runtime.rs b/src/runtime.rs index 5ee43c6..f19b150 100644 --- a/src/runtime.rs +++ b/src/runtime.rs @@ -5,6 +5,7 @@ use std::ffi::OsStr; use std::path::Path; use std::{env, fs}; use crate::utils::normalize_path_for_lua; +use crate::validators::map_lua_error; /// 将 Path 转换为适合 Lua 使用的安全字符串路径 @@ -139,6 +140,7 @@ impl LuaRuntime { .set_name(&chunk_name) .eval::() // .map_err(|e| anyhow!(e.to_string())) + .map_err(map_lua_error) .with_context(|| format!("执行 Lua 配置文件失败: {}", path.display())) } diff --git a/src/validators.rs b/src/validators.rs index 63b39e1..a25c28a 100644 --- a/src/validators.rs +++ b/src/validators.rs @@ -15,10 +15,28 @@ fn conversion_error(message: impl Into) -> mlua::Error { 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) 进行安全解码 fn lua_string_2_os_string(s: &LuaString) -> mlua::Result { - let raw_bytes = &s.as_bytes().to_vec(); + let raw_bytes = &s.as_bytes(); // 1. Unix 平台:直接零拷贝透传原始字节(无损支持任意编码) #[cfg(unix)] { @@ -33,16 +51,44 @@ fn lua_string_2_os_string(s: &LuaString) -> mlua::Result { 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( - "路径包含无效/不支持的字节编码".into(), - )); - } - Ok(OsString::from(cow.as_ref())) + unsafe { + use windows_sys::Win32::Globalization::{MultiByteToWideChar, CP_ACP, MB_ERR_INVALID_CHARS}; + use std::os::windows::ffi::OsStringExt; + + if raw_bytes.is_empty() { + return Ok(OsString::new()); + } + + 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 值校验器:针对不同上下文定义校验规则 #[derive(Debug, Clone, Copy, PartialEq, Eq)] @@ -151,21 +197,34 @@ impl LuaValidator { } /// 【纯粹数据转换与展开】将已通过校验的 Value 递归解析为 OsString 动态数组 - fn collect_value_into( - context: &impl fmt::Display, - value: &Value, - out: &mut Vec, - ) -> mlua::Result<()> { + fn collect_value_into(&self, value: &Value, out: &mut Vec) -> mlua::Result<()> { match value { 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::Number(n) => { - tracing::warn!(%context, value = %n, "浮点数将按十进制格式转换为字符串"); + tracing::warn!(%self, value = %n, "浮点数将按十进制格式转换为字符串"); out.push(OsString::from(n.to_string())); } Value::Boolean(b) => { - tracing::warn!(%context, value = %b, "布尔值将转换为字符串"); + tracing::warn!(%self, value = %b, "布尔值将转换为字符串"); out.push(OsString::from(b.to_string())); } Value::Table(tbl) => { @@ -175,7 +234,7 @@ impl LuaValidator { if matches!(item, Value::Nil) { break; } - Self::collect_value_into(context, &item, out)?; + Self::collect_value_into(self, &item, out)?; index += 1; } } @@ -183,7 +242,36 @@ impl LuaValidator { } Ok(()) } - + fn parse_name(name: &Value) -> mlua::Result { + 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<()> { match value { Value::Nil @@ -254,34 +342,7 @@ impl LuaValidator { /// 校验环境变量名的合法性(Windows 约束:非空、不含 '='、不含 NUL) fn parse_env_name(name: &Value) -> mlua::Result { - 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) + Self::parse_name(name) } /// 解析、校验并拼接单个环境变量键值对 (`key`, `value`) -> (`OsString`) @@ -292,10 +353,10 @@ impl LuaValidator { ) -> mlua::Result { context.parse_raw_val(name, raw_val)?; - let ctx = format!("{} [{}]", context, name); + // let ctx = format!("{} [{}]", context, name); let mut parts = Vec::new(); - Self::collect_value_into(&ctx, raw_val, &mut parts)?; + Self::collect_value_into(&context, raw_val, &mut parts)?; // 校验每个展开元素的 NUL 字符 for part in &parts { @@ -303,7 +364,7 @@ impl LuaValidator { } // 3. 使用系统路径分隔符拼接数组列表 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) } @@ -344,6 +405,10 @@ impl LuaValidator { } impl LuaValidator { + /// 校验环境变量名的合法性(Windows 约束:非空、不含 '='、不含 NUL) + fn parse_aliases_name(name: &Value) -> mlua::Result { + Self::parse_name(name) + } fn parse_aliases_val( context: &LuaValidator, name: &str, @@ -354,7 +419,7 @@ impl LuaValidator { let ctx = format!("{} [{}]", context, name); 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) } /// 解析并打平别名表 (aliases) @@ -378,7 +443,7 @@ impl LuaValidator { for pair in table.pairs::() { 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)?; raw_aliases.insert(key, val); @@ -463,6 +528,69 @@ impl LuaValidator { } } } + + /// 将输入的字符串按 Shell 规则切分为独立的 CLI 参数 Token + /// - 自动过滤连续空格 + /// - 支持单引号 `'...'` 和双引号 `"..."` 包裹包含空格的参数 + fn parse_tokens(input: &str) -> Vec { + let mut tokens = Vec::new(); + let mut current_token = String::new(); + let mut in_quote: Option = 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) 的类型校验与字段提取器 @@ -612,7 +740,7 @@ mod tests { HOME = "C:/tools/home", CONST = 3, BOOL = true, - } + }, aliases = { -- 1. 整体字符串(会自动 trim 首尾空格,保留为单个参数) st = "status -s", @@ -626,7 +754,7 @@ mod tests { -- 4. 嵌套别名组合(加载期会自动展开并进行死环检测) 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 或空字符串(解析为空参数列表) empty_alias = nil, @@ -799,3 +927,98 @@ mod tests { 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 = $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"]); + } +}