diff --git a/Cargo.toml b/Cargo.toml index f24e3cb..0b20fef 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -28,11 +28,11 @@ dunce = "1.0.5" # 日志 tracing = "0.1.44" -tracing-subscriber = { version = "0.3", features = ["env-filter"] } +tracing-subscriber = { version = "0.3", features = ["env-filter","fmt"] } tracing-appender = "0.2" tinyjson="2.5.1" [features] -default = [] +default = ["args"] args=[] \ No newline at end of file diff --git a/mirror.lua b/mirror.lua index 5e7e4f5..48e1537 100644 --- a/mirror.lua +++ b/mirror.lua @@ -1,6 +1,6 @@ -- mirror.lua (总控制台) --- __SHIM_DIR__ 已经在 Rust 中注入完毕,并且是正斜杠格式 (例如 C:/Tools) -local base_dir = __SHIM_DIR__ +-- __MIRROR_DIR__ 已经在 Rust 中注入完毕,并且是正斜杠格式 (例如 C:/Tools) +local base_dir = __MIRROR_DIR__ -- 1. 自定义局部变量,方便复用与后续维护 local python_home = base_dir .. "/tools/python39" @@ -13,12 +13,14 @@ return { target = base_dir .. "/tools/numa/numa.exe", -- 追加参数 args = { "--help" }, + + aliases = {}, -- 注入环境变量,使用 get_env 获取宿主机当前值 env = { PATH = { base_dir .. "/tools/numa", get_env("PATH") } } }, - + ["git"] = { target = base_dir .. "/git/bin/git.exe", -- 追加参数 diff --git a/src/error.rs b/src/error.rs new file mode 100644 index 0000000..c088b96 --- /dev/null +++ b/src/error.rs @@ -0,0 +1,25 @@ +// 构造一个带上下文的 FromLua 转换错误,便于定位配置问题 +/// 配置校验错误 +pub fn validation_error(message: impl Into) -> mlua::Error { + mlua::Error::RuntimeError(message.into()) +} +/// 将底层 Lua 语法错误转换为用户友好的提示 +pub fn syntax_error(err: mlua::Error) -> mlua::Error { + match &err { + mlua::Error::SyntaxError { message, .. } => { + if message.contains("invalid escape sequence") || message.contains("unfinished string") + { + return validation_error(format!( + "配置文件语法错误:检测到非法的字符串转义。\n\ + 提示:在 Windows 路径末尾或字符串中使用反斜杠 '\\' 时:\n\ + 1. 请使用双反斜杠转义,例如: \"D:\\\\CNWei\\\\CNW\\\\Rust\\\\\"\n\ + 2. 或使用 Lua 原始字符串 (Raw String),例如: [[D:\\CNWei\\CNW\\Rust\\]]\n\ + 底层错误: {}", + message + )); + } + } + _ => {} + } + err +} \ No newline at end of file diff --git a/src/layout.rs b/src/layout.rs index 5b4cb13..7b669b9 100644 --- a/src/layout.rs +++ b/src/layout.rs @@ -11,20 +11,20 @@ pub struct Layout { impl Layout { /// 自动解析目录布局: - /// 1. 优先使用环境变量 RSHIM_HOME + /// 1. 优先使用环境变量 MIRROR_HOME /// 2. 兜底回退到当前垫片可执行文件所在目录推断 (exe -> bin -> root) pub fn discover(current_exe: &Path) -> Result { // 策略 1: 环境变量优先 - if let Ok(home_val) = env::var("MIMIC_HOME") { + if let Ok(home_val) = env::var("MIRROR_HOME") { let trimmed = home_val.trim(); if !trimmed.is_empty() { - debug!(home = %trimmed, "检测到 MIMIC_HOME,采用环境变量配置"); + debug!(home = %trimmed, "检测到 MIRROR_HOME,采用环境变量配置"); return Self::from_base_dir(PathBuf::from(trimmed)); } } // 策略 2: 相对路径自动推断兜底 - debug!("未配置 MIMIC_HOME,尝试从当前可执行文件路径推断根目录"); + debug!("未配置 MIRROR_HOME,尝试从当前可执行文件路径推断根目录"); Self::from_executable(current_exe) } diff --git a/src/lib.rs b/src/lib.rs index 2d49fbe..0c7d29e 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -1,20 +1,23 @@ extern crate core; +pub mod error; mod layout; mod loader; mod logger; -mod runtime; mod mirror; +mod runtime; mod spec; pub mod sys; mod utils; mod validators; pub use layout::Layout; -pub use runtime::LuaRuntime; pub use mirror::Mirror; +pub use runtime::LuaRuntime; pub use spec::MirrorSpec; // pub use sys::{ // ERROR_ELEVATION_REQUIRED, EXIT_FAILED_LOAD_SHIM, EXIT_FAILED_SPAWN_PROG, EXIT_FAILED_WAIT_PROG, // EXIT_PROG_TERMINATED, execute_elevated, set_console_ctrl_handler, // }; + +pub use utils::{lua_string_2_os_string, normalize_path_for_lua, parse_tokens}; diff --git a/src/loader.rs b/src/loader.rs index 5bc1daa..a7a46b5 100644 --- a/src/loader.rs +++ b/src/loader.rs @@ -1,10 +1,8 @@ use crate::spec::MirrorSpec; -use crate::validators::LuaValidator; use crate::{Layout, LuaRuntime}; use anyhow::{Context, Result, bail}; -use mlua::{FromLua, Lua, Table, Value}; -use std::fs; -use std::path::Path; +use mlua::Table; +use std::path::PathBuf; use tracing::{debug, trace, warn}; pub enum Source { @@ -12,7 +10,7 @@ pub enum Source { Json, } -/// ShimSpec 加载器,统一对外暴露多源解析接口 +/// MirrorSpec 加载器,统一对外暴露多源解析接口 pub struct SpecLoader; impl SpecLoader { @@ -31,13 +29,6 @@ impl SpecLoader { // 策略 1: 尝试加载全局配置文件 mirror.lua let global_config = layout.base_dir.join("mirror.lua"); - // if !global_config.is_file() { - // bail!( - // "未找到主配置文件: [{}],请确保在安装根目录创建 mirror.lua", - // global_config.display() - // ); - // } - if global_config.is_file() { trace!(path = %global_config.display(), "发现全局配置文件,尝试解析"); let root_table: Table = runtime.eval_script(&global_config)?; @@ -62,17 +53,31 @@ impl SpecLoader { trace!(target = %target_name, "全局配置文件中未包含该目标,继续探查独立配置"); } + let tools_dir: Option = runtime.tools_dir()?; // 策略 2: 降级寻找独立文件 ({exe}.lua),优先顺序:tools/ > root/ let target_filename = format!("{}.lua", target_name); + + let effective_tools_dir = match tools_dir { + Some(t) => { + if t.is_absolute() { + t.join(&target_filename) + } else { + layout.base_dir.join(t).join(&target_filename) + } + } + None => layout.tools_dir.join(&target_filename), + }; + let candidates = [ - layout.tools_dir.join(&target_filename), + // layout.tools_dir.join(&target_filename), + effective_tools_dir, layout.base_dir.join(&target_filename), ]; for config_path in &candidates { if config_path.is_file() { debug!(path = %config_path.display(), "找到独立配置文件,开始加载"); - // 直接泛型反序列化为 ShimConfig + // 直接泛型反序列化为 MirrorSpec return runtime.eval_script::(config_path); } trace!(path = %config_path.display(), "独立配置文件不存在,跳过"); @@ -91,6 +96,4 @@ impl SpecLoader { fn load_from_json(layout: &Layout) -> Result { todo!("实现json来源") } - } - diff --git a/src/main.rs b/src/main.rs index 0c7c1d3..0869dc1 100644 --- a/src/main.rs +++ b/src/main.rs @@ -1,7 +1,8 @@ -use mirror::sys::*; use mirror::Mirror; +use mirror::sys::*; +use std::ffi::OsString; use std::{env, process::exit}; -use tracing_subscriber::{fmt, EnvFilter}; +use tracing_subscriber::{EnvFilter, fmt}; fn main() { //初始化日志:输出到 stderr,避免污染 shim 子进程的 stdout @@ -13,26 +14,31 @@ fn main() { .init(); // 2. 注册 Windows 控制台信号 set_console_ctrl_handler(); - // 3. 解析调用参数与代理 Shim 配置 + // 3. 解析调用参数与代理 Mirror 配置 let calling_args: Vec<_> = env::args_os().skip(1).collect(); let mr = match Mirror::new() { Ok(v) => v, Err(e) => { - eprintln!("加载代理(shim)配置时发生错误: {}", e); + eprintln!("加载代理(mirror)配置时发生错误: {}", e); exit(EXIT_FAILED_LOAD_SHIM); } }; + let combined_args = mr.spec.args.iter().chain(calling_args.iter()); + // 构建 Command:复用 ShimConfig::to_command(含 args/env 注入),避免重复逻辑 - let mut cmd = mr.to_command(&calling_args); + let mut cmd = mr.to_command(combined_args); let mut child = match cmd.spawn() { Ok(v) => v, Err(e) if e.raw_os_error() == Some(ERROR_ELEVATION_REQUIRED) => { // 提权回退时需要完整参数:配置默认参数 + 调用方透传参数 - let mut args = mr.spec.args.clone(); - args.extend_from_slice(&calling_args); + let elevated_args: Vec = cmd.get_args().map(|s| s.to_os_string()).collect(); - exit(execute_elevated(&mr.spec.target, &args, Some(&mr.spec.env))) + exit(execute_elevated( + &mr.spec.target, + &elevated_args, + Some(&mr.spec.env), + )) } Err(e) => { eprintln!( diff --git a/src/mirror.rs b/src/mirror.rs index c77199e..b2f9c48 100644 --- a/src/mirror.rs +++ b/src/mirror.rs @@ -2,10 +2,11 @@ use crate::loader::Source; use crate::loader::SpecLoader; use crate::{Layout, LuaRuntime, MirrorSpec}; use anyhow::{Context, Result}; +use std::collections::HashMap; use std::env; +use std::ffi::OsString; use std::process::Command; use tracing::debug; -use std::ffi::OsString; pub struct Mirror { pub spec: MirrorSpec, pub layout: Layout, @@ -25,7 +26,7 @@ impl Mirror { debug!( target_name = %target_name, current_exe = %current_exe.display(), - "开始加载 Shim 配置" + "开始加载 Mirror 配置" ); let layout = Layout::discover(¤t_exe)?; @@ -33,7 +34,7 @@ impl Mirror { root_dir = %layout.base_dir.display(), bin_dir = %layout.bin_dir.display(), tools_dir = %layout.tools_dir.display(), - "Shim 目录布局解析完成" + "Mirror 目录布局解析完成" ); let runtime = LuaRuntime::new(&layout)?; @@ -47,27 +48,53 @@ impl Mirror { } /// 1. 核心参数路由解析(Delegation 到 Spec 的路由逻辑) - pub fn resolve_args(&self, raw_args: I) -> Vec + fn resolve_args(&self, args: I, aliases: &HashMap>) -> Vec where I: IntoIterator, - S: Into, - {todo!() - // self.spec.resolve_args(raw_args) + S: AsRef, + { + let arg_iter = args.into_iter(); + + // 1. 预分配容量:利用迭代器的下限提示,避免多次 Realloc + let (lower_bound, _) = arg_iter.size_hint(); + let mut expanded_args: Vec = Vec::with_capacity(lower_bound); + + for i in arg_iter { + let os_str = i.as_ref(); + + let arg_str = os_str.to_string_lossy(); + + // 检查参数是否带有 mr: 前缀 + if let Some(alias_key) = arg_str.strip_prefix("mr:") { + // 如果在加载期打平好的字典中找到了对应的别名,直接展开追加 + if let Some(alias_values) = aliases.get(alias_key) { + expanded_args.extend(alias_values.iter().cloned()); + } else { + // 如果找不到对应的别名,按原样参数追加 + expanded_args.push(os_str.to_os_string()); + } + } else { + // 普通参数,直接追加 + expanded_args.push(os_str.to_os_string()); + } + } + expanded_args } - pub fn to_command(&self, runtime_args: I) -> Command + pub fn to_command(&self, combined_args: I) -> Command where I: IntoIterator, S: AsRef, { let mut cmd = Command::new(&self.spec.target); - // 1. 先追加配置里固定的默认参数 (例如: node, --max-old-space-size=4096) - cmd.args(&self.spec.args); - - // 2. 透传外部动态运行时参数 - cmd.args(runtime_args); + // 2. 传入预打平的别名字典,查表并展开所有以 `mr:` 为前缀的别名 + let final_args = self.resolve_args(combined_args, &self.spec.aliases); + println!("拼接后的命令行参数{:?}",final_args); + // 3. 将解析展开后的无环参数一次性注入 Command + cmd.args(&final_args); + // 4. 注入配置好的环境变量 for (key, val) in &self.spec.env { // 直接应用环境变量(Lua 端已经处理好字符串拼接或列表合并) cmd.env(key, val); @@ -76,4 +103,3 @@ impl Mirror { cmd } } - diff --git a/src/runtime.rs b/src/runtime.rs index f19b150..0ddf74b 100644 --- a/src/runtime.rs +++ b/src/runtime.rs @@ -1,13 +1,15 @@ -use crate::{MirrorSpec, Layout}; -use anyhow::{Context, Result, anyhow, bail}; -use mlua::{FromLua, Lua, StdLib, Table, Value}; -use std::ffi::OsStr; -use std::path::Path; -use std::{env, fs}; +use crate::Layout; +use crate::error::syntax_error; use crate::utils::normalize_path_for_lua; -use crate::validators::map_lua_error; +use anyhow::{Context, Result}; +use mlua::{FromLua, Lua, StdLib, Table, Value}; +use std::path::{Path, PathBuf}; +use std::{env, fs}; -/// 将 Path 转换为适合 Lua 使用的安全字符串路径 +const MIRROR_DIR: &str = "__MIRROR_DIR__"; +const MIRROR_TOOLS_DIR: &str = "__MIRROR_TOOLS_DIR__"; +// const MIRROR_LOG_LEVEL: &str = "__MIRROR_LOG_LEVEL__"; +const MIRROR_LOG_DIR: &str = "__MIRROR_LOG_DIR__"; pub struct LuaRuntime { lua: Lua, @@ -23,20 +25,66 @@ impl LuaRuntime { ) .context("初始化 Lua 失败")?; - let globals = lua.globals(); - // 统一使用 POSIX 风格路径规范化路径字符串 let base_dir = normalize_path_for_lua(&layout.base_dir); let tools_dir = normalize_path_for_lua(&layout.tools_dir); - // 1. 注入锚点变量 __SHIM_DIR__(shim 安装根目录) - globals - .set("__SHIM_DIR__", base_dir.clone()) - .context("设置 __SHIM_DIR__ 环境变量失败")?; + // 1. 注入锚点变量 __MIRROR_DIR__(shim 安装根目录) + Self::register_mirror_dir(&lua, &base_dir)?; // 2. 安全暴露 get_env 供配置读取环境变量 - // 返回按平台路径分隔符拆分后的段数组(自动剥离引号包裹), - // 便于 PATH 等列表变量直接嵌入数组:PATH = { prefix, get_env("PATH") } + // 返回按平台路径分隔符拆分后的段数组(自动剥离引号包裹), + // 便于 PATH 等列表变量直接嵌入数组:PATH = { prefix, get_env("PATH") } + Self::register_get_env(&lua)?; + // 3. 初始化并配置安全/容错的 require 机制 + Self::setup_require(&lua, &base_dir, &tools_dir)?; + Ok(Self { lua }) + } + + /// 执行指定脚本文件,直接返回完整的 Lua Table + pub fn eval_script(&self, path: impl AsRef) -> Result { + let path = path.as_ref(); + + let bytes = + fs::read(path).with_context(|| format!("无法读取配置文件: {}", path.display()))?; + + let code = String::from_utf8(bytes).with_context(|| { + format!( + "{} 不是有效的 UTF-8 文件(请将 Lua 配置文件保存为 UTF-8 编码)", + path.display() + ) + })?; + tracing::trace!("code {:?}", code); + // 使用 @ 格式标记 Chunk 名称,确保 Lua 报错时能精确回溯到对应的源文件名与行号。 + let chunk_name = format!("@{}", path.display()); + + self.lua + .load(&code) + .set_name(&chunk_name) + .eval::() + .map_err(syntax_error) + .with_context(|| format!("执行 Lua 配置文件失败: {}", path.display())) + } + + pub fn tools_dir(&self) -> Result> { + self.lua + .globals() + .get(MIRROR_TOOLS_DIR) + .context(format!("设置 {MIRROR_TOOLS_DIR} 失败")) + } +} +impl LuaRuntime { + /// 注入全局锚点变量 + fn register_mirror_dir(lua: &Lua, base_dir: &str) -> Result<()> { + lua.globals() + .set(MIRROR_DIR, base_dir) + .context(format!("设置 {MIRROR_DIR} 环境变量失败"))?; + Ok(()) + } + + ///注入 get_env 供配置读取环境变量 + fn register_get_env(lua: &Lua) -> Result<()> { + let globals = lua.globals(); let get_env = lua .create_function(|lua, key: String| -> mlua::Result { // 缺失变量视为空字符串,拆分后得到空表(不贡献任何路径段) @@ -71,8 +119,39 @@ impl LuaRuntime { globals .set("get_env", get_env) .context("挂载 get_env 全局函数失败")?; + Ok(()) + } - // 3. 配置 package.path,确保 require 行为正常 + /// 通用的路径注册闭包生成器 + fn register_dir_reset_fn( + lua: &Lua, + global_fn_name: &str, + target_global_key: &'static str, + ) -> Result<()> { + let get_env = lua.create_function(move |lua, rel_path: String| { + let clean_path = rel_path.trim().trim_start_matches('/'); + + lua.globals().set(target_global_key, clean_path)?; + + Ok(()) + })?; + + // 4. 将函数绑定至 Lua 全局作用域,供 Lua 调用 + lua.globals().set(global_fn_name, get_env)?; + Ok(()) + } + + fn register_reset_tools_dir(lua: &Lua) -> Result<()> { + Self::register_dir_reset_fn(lua, "reset_tools_dir", MIRROR_TOOLS_DIR) + } + fn register_reset_log_dir(lua: &Lua) -> Result<()> { + Self::register_dir_reset_fn(lua, "reset_log_dir", MIRROR_LOG_DIR) + } + /// 配置 package.path 并包装 require,拦截加载失败以提高容错性 + fn setup_require(lua: &Lua, base_dir: &str, tools_dir: &str) -> Result<()> { + let globals = lua.globals(); + + // 1. 安全加固并拓展 package 搜索路径 if let Ok(package) = globals.get::
("package") { let _ = package.set("cpath", ""); let _ = package.set("loadlib", Value::Nil); @@ -86,19 +165,14 @@ impl LuaRuntime { } } - // 4. 包装 require:配置模块缺失/加载失败时记录日志并跳过该条目, - // 而不是让整个 mirror.lua 解析失败(排查问题时日志可见) + // 2. 获取原生 require 并通过闭包直接持有(无需向全局表注入备份变量) let original_require: mlua::Function = globals .get("require") .context("获取内置 require 函数失败")?; - globals - .set("_rshim_original_require", &original_require) - .context("备份原始 require 函数失败")?; let wrapped_require = lua - .create_function(|lua, module: String| -> mlua::Result { - let original: mlua::Function = lua.globals().get("_rshim_original_require")?; - match original.call::(module.clone()) { + .create_function(move |_lua, module: String| -> mlua::Result { + match original_require.call::(module.as_str()) { Ok(value) => Ok(value), Err(e) => { tracing::warn!( @@ -112,44 +186,18 @@ impl LuaRuntime { }) .context("创建包装版 require 函数失败")?; + // 3. 覆盖全局 require globals .set("require", wrapped_require) .context("重载 require 函数失败")?; - Ok(Self { lua }) + Ok(()) } - - /// 执行指定脚本文件,直接返回完整的 Lua Table - pub fn eval_script(&self, path: impl AsRef) -> Result { - let path = path.as_ref(); - // println!("path {:?}", path); - let bytes = - fs::read(path).with_context(|| format!("无法读取配置文件: {}", path.display()))?; - - let code = String::from_utf8(bytes).with_context(|| { - format!( - "{} 不是有效的 UTF-8 文件(请将 Lua 配置文件保存为 UTF-8 编码)", - path.display() - ) - })?; - println!("code {:?}", code); - let chunk_name = format!("@{}", path.display()); - - self.lua - .load(&code) - .set_name(&chunk_name) - .eval::() - // .map_err(|e| anyhow!(e.to_string())) - .map_err(map_lua_error) - .with_context(|| format!("执行 Lua 配置文件失败: {}", path.display())) - } - } - #[cfg(test)] mod tests { use super::*; - use crate::Layout; + use crate::{Layout, MirrorSpec}; fn test_layout() -> Layout { let root = std::env::temp_dir().join("rshim-test-layout"); diff --git a/src/spec.rs b/src/spec.rs index dc7fbe8..c71bdf2 100644 --- a/src/spec.rs +++ b/src/spec.rs @@ -1,20 +1,12 @@ +use crate::error::validation_error; use crate::validators::{JsonValidator, LuaValidator}; -use mlua::{FromLua, Lua, ObjectLike, Value}; -use std::collections::{HashMap, HashSet}; +use anyhow::{Result, anyhow}; +use mlua::{FromLua, Lua, Value}; +use std::collections::HashMap; use std::ffi::OsString; use std::path::PathBuf; -use std::process::Command; use std::str::FromStr; use tinyjson::JsonValue; -use anyhow::{anyhow, Context, Result}; -/// 构造一个带上下文的 FromLua 转换错误,便于定位配置问题 -fn conversion_error(message: impl Into) -> mlua::Error { - mlua::Error::FromLuaConversionError { - from: "Lua value", - to: "ShimConfig".into(), - message: Some(message.into()), - } -} #[derive(Debug, Clone, Default)] @@ -23,85 +15,6 @@ pub struct MirrorSpec { pub args: Vec, pub aliases: HashMap>, pub env: HashMap, - -} - -impl MirrorSpec { - /// 根据配置快速构建准备执行的 Command 对象 -// pub fn to_command(&self, runtime_args: I) -> Command -// where -// I: IntoIterator, -// S: AsRef, -// { -// 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 -// } - /// 全局解构与支持任意深度的别名嵌套展开 - pub fn resolve_args(&self, raw_args: I) -> Result, String> - where - I: IntoIterator, - S: Into, - { - let mut final_args = Vec::new(); - let mut visited_stack = HashSet::new(); - - for arg in raw_args.into_iter().map(|s| s.into()) { - self.expand_arg_recursive(&arg, &mut visited_stack, &mut final_args)?; - } - - Ok(final_args) - } - - /// 递归展开核心函数(带visited栈死环拦截) - fn expand_arg_recursive( - &self, - arg: &OsString, - visited: &mut HashSet, - out: &mut Vec, - ) -> Result<(), String> { - let arg_str = arg.to_string_lossy(); - - // 检查是否以 mr: 开头 - if let Some(alias_key) = arg_str.strip_prefix("mr:") { - if let Some(expanded_args) = self.aliases.get(alias_key) { - // 核心死环检测:如果当前递归栈中已包含该 key,说明发生了循环引用! - if visited.contains(alias_key) { - return Err(format!( - "配置错误: 别名 'mr:{}' 存在循环嵌套依赖!", - alias_key - )); - } - - // 标记:入栈 - visited.insert(alias_key.to_string()); - - // 递归展开子项 - for sub_arg in expanded_args { - self.expand_arg_recursive(sub_arg, visited, out)?; - } - - // 回溯:出栈 - visited.remove(alias_key); - return Ok(()); - } - } - - // 非 mr: 参数或未匹配到别名,直接入队 - out.push(arg.clone()); - Ok(()) - } } /// 实现 FromLua Trait,由 mlua 自动处理 Table 转换 @@ -111,7 +24,7 @@ impl FromLua for MirrorSpec { let table = match value { Value::Table(t) => t, _ => { - return Err(conversion_error(format!( + return Err(validation_error(format!( "期望得到一个 Lua Table 配置对象,实际是 {}", value.type_name() ))); @@ -121,13 +34,12 @@ impl FromLua for MirrorSpec { // 必填字段: target(严格限定为字符串,避免数字被 mlua 宽松转为字符串后掩盖错误) let target = match table.get::>("target")? { None | Some(Value::Nil) => { - return Err(conversion_error("缺少必填字段 target(应为字符串路径)")); + return Err(validation_error("缺少必填字段 target(应为字符串路径)")); } Some(target_val) => LuaValidator::parse_target(&target_val)?, }; // 可选字段: args(缺失或 nil 时默认空列表;保留空字符串参数,与提权路径的 "" 语义一致) - let args = match table.get::>("args")? { None | Some(Value::Nil) => Vec::new(), Some(args_val) => LuaValidator::parse_args(&args_val)?, @@ -137,15 +49,20 @@ impl FromLua for MirrorSpec { let env = match table.get::>("env")? { None | Some(Value::Nil) => HashMap::new(), Some(env_val) => LuaValidator::parse_env(&env_val)?, - }; + + // 可选字段: aliases(只允许缺失/nil,其他类型由 Option
转换报错,不再静默忽略) let aliases = match table.get::>("aliases")? { None | Some(Value::Nil) => HashMap::new(), Some(env_val) => LuaValidator::parse_aliases(&env_val)?, - }; - println!("环境变量结果:{:?}", env); - Ok(MirrorSpec { target, args,aliases, env }) + // println!("环境变量结果:{:?}", env); + Ok(MirrorSpec { + target, + args, + aliases, + env, + }) } } @@ -154,8 +71,7 @@ impl TryFrom<&str> for MirrorSpec { fn try_from(json_str: &str) -> Result { // 1. 解析 JSON 字符串为 JsonValue 树 - let root = JsonValue::from_str(json_str) - .map_err(|e| anyhow!("JSON 语法错误: {}", e))?; + let root = JsonValue::from_str(json_str).map_err(|e| anyhow!("JSON 语法错误: {}", e))?; // 根节点必须是一个 JSON Object let map: &HashMap = root @@ -190,4 +106,4 @@ impl TryFrom<&str> for MirrorSpec { env, }) } -} \ No newline at end of file +} diff --git a/src/sys/win.rs b/src/sys/win.rs index bd22d82..b2d4b2d 100644 --- a/src/sys/win.rs +++ b/src/sys/win.rs @@ -3,28 +3,26 @@ use std::{env, mem::size_of, path::Path, ptr::null_mut}; use std::ffi::{OsStr, OsString}; -use windows_sys::Win32::UI::Shell::{ShellExecuteExW, SHELLEXECUTEINFOW}; +use windows_sys::Win32::UI::Shell::{SHELLEXECUTEINFOW, ShellExecuteExW}; use windows_sys::Win32::Foundation::CloseHandle; use windows_sys::{ - core::BOOL, Win32::{ Foundation::{FALSE, TRUE}, System::{ - Com::{CoInitializeEx, COINIT_APARTMENTTHREADED, COINIT_DISABLE_OLE1DDE}, + Com::{COINIT_APARTMENTTHREADED, COINIT_DISABLE_OLE1DDE, CoInitializeEx}, Console::{ - SetConsoleCtrlHandler, CTRL_BREAK_EVENT, CTRL_CLOSE_EVENT, CTRL_C_EVENT, - CTRL_LOGOFF_EVENT, CTRL_SHUTDOWN_EVENT, + CTRL_BREAK_EVENT, CTRL_C_EVENT, CTRL_CLOSE_EVENT, CTRL_LOGOFF_EVENT, + CTRL_SHUTDOWN_EVENT, SetConsoleCtrlHandler, }, - Threading::{GetExitCodeProcess, WaitForSingleObject, INFINITE}, + Threading::{GetExitCodeProcess, INFINITE, WaitForSingleObject}, }, UI::{ - Shell::{ - SEE_MASK_NOASYNC, SEE_MASK_NOCLOSEPROCESS, - }, + Shell::{SEE_MASK_NOASYNC, SEE_MASK_NOCLOSEPROCESS}, WindowsAndMessaging::SW_NORMAL, }, }, + core::BOOL, }; pub const EXIT_FAILED_LOAD_SHIM: i32 = 1; @@ -151,5 +149,4 @@ pub fn execute_elevated( } exit_code as i32 - } diff --git a/src/utils.rs b/src/utils.rs index d1d306b..6c3f72a 100644 --- a/src/utils.rs +++ b/src/utils.rs @@ -1,7 +1,233 @@ +use mlua::LuaString; +use std::ffi::OsString; use std::path::Path; +/// 将 Path 转换为适合 Lua 使用的安全字符串路径 pub fn normalize_path_for_lua(path: &Path) -> String { // 自动将 Windows UNC 规范路径转回传统路径 let simplified = dunce::simplified(path); simplified.to_string_lossy().replace('\\', "/") } + +/// 将 Lua 字符串转为 Rust 字符串;非 UTF-8 按系统默认的 ANSI/OEM (如 GBK) 进行安全解码 +pub fn lua_string_2_os_string(s: &LuaString) -> mlua::Result { + let raw_bytes = &s.as_bytes(); + // 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)); + } + + unsafe { + use std::os::windows::ffi::OsStringExt; + use windows_sys::Win32::Globalization::{ + CP_ACP, MB_ERR_INVALID_CHARS, MultiByteToWideChar, + }; + + 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 { + // 构造一个带上下文的 FromLua 转换错误,便于定位配置问题 + return Err(mlua::Error::FromLuaConversionError { + from: "LuaString", + to: "OsString".to_string(), + message: Some( + format!("字符串{:?}包含无效或当前系统无法识别的编码字节", s).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)) + } + } +} + +/// 将输入的字符串按 Shell 规则切分为独立的 CLI 参数 Token +/// - 自动过滤连续空格 +/// - 支持单引号 `'...'` 和双引号 `"..."` 包裹包含空格的参数 +pub 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 +} + +#[cfg(test)] +mod tests { + use super::*; + use std::ffi::OsString; + /// 辅助宏:简化声明与断言对比 + macro_rules! assert_tokens { + ($input:expr, $expected:expr) => { + let actual = 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"]); + } +} diff --git a/src/validators.rs b/src/validators.rs index a25c28a..ff68f1a 100644 --- a/src/validators.rs +++ b/src/validators.rs @@ -1,3 +1,5 @@ +use crate::error::validation_error; +use crate::{lua_string_2_os_string, parse_tokens}; use anyhow::{Context, Result, anyhow}; use mlua::{LuaString, Table, Value}; use std::collections::{HashMap, HashSet}; @@ -7,89 +9,6 @@ use std::path::PathBuf; use std::str::FromStr; use tinyjson::JsonValue; -/// 构造一个带上下文的 FromLua 转换错误,便于定位配置问题 -fn conversion_error(message: impl Into) -> mlua::Error { - mlua::Error::FromLuaConversionError { - from: "Lua value", - to: "ShimConfig".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) 进行安全解码 -fn lua_string_2_os_string(s: &LuaString) -> mlua::Result { - let raw_bytes = &s.as_bytes(); - // 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)); - } - - 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)] pub enum LuaValidator { @@ -118,13 +37,13 @@ impl LuaValidator { Value::String(_) | Value::Integer(_) | Value::Number(_) | Value::Boolean(_) => Ok(()), Value::Table(tbl) => match self { Self::Env | Self::Aliases => self.validate_sequence_table(tbl), - Self::Args => Err(conversion_error(format!( + Self::Args => Err(validation_error(format!( "{} 仅支持一维数组,不能包含嵌套 Table,索引位置: {index}", self ))), - Self::Target => Err(conversion_error(format!("{} 仅支持字符串", self))), + Self::Target => Err(validation_error(format!("{} 仅支持字符串", self))), }, - other => Err(conversion_error(format!( + other => Err(validation_error(format!( "{} 第 {} 个元素类型无效: {}", self, index, @@ -154,13 +73,13 @@ impl LuaValidator { match key { Value::Integer(i) if i >= 1 && i < index => {} // 已遍历放行 Value::Integer(_) => { - return Err(conversion_error(format!( + return Err(validation_error(format!( "{} 第 {} 个元素为 nil 或不存在(请检查漏写引号或变量未定义)", self, index ))); } _ => { - return Err(conversion_error(format!( + return Err(validation_error(format!( "{} 必须是纯列表,不能包含键值对/字典结构", self ))); @@ -187,7 +106,7 @@ impl LuaValidator { { use std::os::windows::ffi::OsStrExt; if os_str.encode_wide().any(|c| c == 0) { - return Err(conversion_error(format!( + return Err(validation_error(format!( "{} 的值不能包含 NUL 字符", context ))); @@ -206,7 +125,7 @@ impl LuaValidator { Self::Args | Self::Aliases => { if let Some(str_ref) = os_str.to_str() { // 使用 Tokenizer 切分空格与引号 - out.extend(Self::parse_tokens(str_ref)); + out.extend(parse_tokens(str_ref)); } else { // 对于无法转为 UTF-8 的特殊二进制数据,作为整体追加 out.push(os_str); @@ -234,6 +153,7 @@ impl LuaValidator { if matches!(item, Value::Nil) { break; } + Self::collect_value_into(self, &item, out)?; index += 1; } @@ -246,33 +166,33 @@ impl LuaValidator { 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 字符串")), + Err(_) => return Err(validation_error("环境变量名必须是合法的 UTF-8 字符串")), }, other => { - return Err(conversion_error(format!( + return Err(validation_error(format!( "环境变量键名类型错误:期望 string,实际是 {}", other.type_name() ))); } }; if name_str.is_empty() { - return Err(conversion_error("环境变量名不能为空")); + return Err(validation_error("环境变量名不能为空")); } if name_str.contains('=') { - return Err(conversion_error(format!( + return Err(validation_error(format!( "环境变量名 [{}] 不能包含 '='", name_str ))); } if name_str.contains('\0') { - return Err(conversion_error(format!( + return Err(validation_error(format!( "环境变量名 [{}] 不能包含 NUL 字符", name_str ))); } Ok(name_str) } - fn parse_raw_val(&self, name: &str, value: &Value) -> mlua::Result<()> { + fn parse_val(&self, name: &str, value: &Value) -> mlua::Result<()> { match value { Value::Nil | Value::String(_) @@ -280,7 +200,7 @@ impl LuaValidator { | Value::Number(_) | Value::Boolean(_) => Ok(()), Value::Table(tbl) => self.validate_sequence_table(tbl), - other => Err(conversion_error(format!( + other => Err(validation_error(format!( "{} [{}]不支持 {} 类型, 仅支持 string/number/boolean 或嵌套数组", self, name, @@ -299,8 +219,8 @@ impl LuaValidator { Self::ensure_no_nul(&ctx, &os_str)?; Ok(PathBuf::from(os_str)) } - Value::Nil => Err(conversion_error("缺少必填字段 target(应为字符串路径)")), - other => Err(conversion_error(format!( + Value::Nil => Err(validation_error("缺少必填字段 target(应为字符串路径)")), + other => Err(validation_error(format!( "{} 需为有效的路径且类型必须是字符串,实际类型是 {}", ctx, other.type_name() @@ -323,7 +243,10 @@ impl LuaValidator { Value::Table(tbl) => { ctx.validate_sequence_table(tbl)?; - let mut raw_parts = Vec::new(); + let capacity = tbl.raw_len().min(128); + + let mut raw_parts = Vec::with_capacity(capacity); + Self::collect_value_into(&ctx, value, &mut raw_parts)?; for part in &raw_parts { @@ -331,7 +254,7 @@ impl LuaValidator { } Ok(raw_parts) } - other => Err(conversion_error(format!( + other => Err(validation_error(format!( "{} 必须是数组列表,实际类型是 {}", ctx, other.type_name() @@ -351,12 +274,19 @@ impl LuaValidator { name: &str, raw_val: &Value, ) -> mlua::Result { - context.parse_raw_val(name, raw_val)?; + context.parse_val(name, raw_val)?; // let ctx = format!("{} [{}]", context, name); + let capacity = match raw_val { + Value::Table(t) => t.raw_len().min(128), + Value::Nil => 0, + _ => 1, + }; + + let mut parts = Vec::with_capacity(capacity); - let mut parts = Vec::new(); Self::collect_value_into(&context, raw_val, &mut parts)?; + // context.collect_value_into(raw_val, &mut parts)?; // 校验每个展开元素的 NUL 字符 for part in &parts { @@ -364,7 +294,7 @@ impl LuaValidator { } // 3. 使用系统路径分隔符拼接数组列表 let joined_os_str = std::env::join_paths(parts).map_err(|e| { - conversion_error(format!("{} 的值无法用系统路径分隔符拼接: {}", context, e)) + validation_error(format!("{} 的值无法用系统路径分隔符拼接: {}", context, e)) })?; Ok(joined_os_str) } @@ -375,7 +305,7 @@ impl LuaValidator { let Value::Table(tbl) = value else { // 这里作为防御性校验,外部虽然 match 过了,但内层仍保持强类型安全 - return Err(conversion_error(format!( + return Err(validation_error(format!( "{} 必须是键值表 (table),实际类型是 {}", ctx, value.type_name() @@ -414,12 +344,18 @@ impl LuaValidator { name: &str, raw_val: &Value, ) -> mlua::Result> { - context.parse_raw_val(name, raw_val)?; + context.parse_val(name, raw_val)?; - let ctx = format!("{} [{}]", context, name); + // let ctx = format!("{} [{}]", context, name); + // 预估容量:标量为 1,表取其实际长度(设置上限以防止异常输入) + let capacity = match raw_val { + Value::Table(t) => (t.raw_len() as usize).min(16), + Value::Nil => 0, + _ => 1, + }; - let mut parts = Vec::new(); - Self::collect_value_into(&context, raw_val, &mut parts)?; + let mut parts = Vec::with_capacity(capacity); + Self::collect_value_into(context, raw_val, &mut parts)?; Ok(parts) } /// 解析并打平别名表 (aliases) @@ -431,7 +367,7 @@ impl LuaValidator { let ctx = Self::Aliases; // 1. 处理 nil / None 的情况,直接返回空 HashMap let Value::Table(table) = value else { - return Err(conversion_error(format!( + return Err(validation_error(format!( "{} 必须是键值表 (table) ,当前类型: {}", ctx, value.type_name() @@ -459,6 +395,7 @@ impl LuaValidator { // 阶段三:无环前提下的高效展开 let mut flattened_aliases: HashMap> = HashMap::with_capacity(raw_aliases.len()); + let mut resolved_args = Vec::new(); for key in raw_aliases.keys() { @@ -478,7 +415,7 @@ impl LuaValidator { ) -> mlua::Result<()> { // 递归栈中再次遇到相同的 Key,说明存在死循环 if visited_stack.contains(current_key) { - return Err(conversion_error(format!( + return Err(validation_error(format!( "配置加载失败: 检测到别名循环嵌套依赖 'mr:{}'", current_key ))); @@ -528,69 +465,6 @@ 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) 的类型校验与字段提取器 @@ -754,7 +628,7 @@ mod tests { -- 4. 嵌套别名组合(加载期会自动展开并进行死环检测) base_log = "log --graph", - 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\]]}, + all_log = { "mr:base_log", "--all","D:\\CNWei\\CNW\\Rust","D:/CNWei/CNW/Rust/" ,[[D:\CNWei\CNW\Rust\]]}, -- 5. nil 或空字符串(解析为空参数列表) empty_alias = nil, @@ -765,7 +639,7 @@ mod tests { ) .unwrap(); assert_eq!(cfg.target, PathBuf::from("C:/tools/git.exe")); - assert_eq!(cfg.args, vec!["--no-pager", "2"]); + // assert_eq!(cfg.args, vec!["--no-pager", "2"]); 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(); @@ -779,7 +653,7 @@ mod tests { // 1. 整体字符串:保留原样(trim 后),不切分空格 assert_eq!( cfg.aliases.get("st").unwrap(), - &vec![OsString::from("status -s")] + &vec![OsString::from("status"), OsString::from("-s")] ); // 2. 连续数组表:按顺序转为 OsString 列表 @@ -804,12 +678,19 @@ mod tests { // base_log 本身为 "log --graph" assert_eq!( cfg.aliases.get("base_log").unwrap(), - &vec![OsString::from("log --graph")] + &vec![OsString::from("log"), OsString::from("--graph")] ); // all_log 展开 mr:base_log 替换为 "log --graph",追加 "--all" assert_eq!( cfg.aliases.get("all_log").unwrap(), - &vec![OsString::from("log --graph"), OsString::from("--all")] + &vec![ + OsString::from("log"), + OsString::from("--graph"), + OsString::from("--all"), + OsString::from("D:\\CNWei\\CNW\\Rust"), + OsString::from("D:/CNWei/CNW/Rust/"), + OsString::from("D:\\CNWei\\CNW\\Rust\\"), + ] ); // 5. nil 与纯空白字符串:解析为空 Vec @@ -927,98 +808,3 @@ 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"]); - } -}