refactor(mirror),feat(runtime): 支持在 mirror.lua 中通过 tools_dir 设置自定义目录;优化参数解析与 Command 构建逻辑

- 新增 Lua 动态配置 tools 目录并返回完整路径
- 将 `resolve_args` 和 `to_command` 改为接受泛型 `IntoIterator`,解耦参数来源
- 统一使用 `chain` 零拷贝合并预设参数与运行时参数,消除多余堆分配
- 提权回退分支改为直接复用 `cmd.get_args()`,修复提权时 `mr:` 别名未展开的缺陷
- 规范参数与变量命名,完善代码注释
This commit is contained in:
2026-08-28 17:21:29 +08:00
parent 219a6b3d51
commit 3db8945e16
13 changed files with 526 additions and 488 deletions

View File

@@ -28,11 +28,11 @@ 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","fmt"] }
tracing-appender = "0.2" tracing-appender = "0.2"
tinyjson="2.5.1" tinyjson="2.5.1"
[features] [features]
default = [] default = ["args"]
args=[] args=[]

View File

@@ -1,6 +1,6 @@
-- mirror.lua (总控制台) -- mirror.lua (总控制台)
-- __SHIM_DIR__ 已经在 Rust 中注入完毕,并且是正斜杠格式 (例如 C:/Tools) -- __MIRROR_DIR__ 已经在 Rust 中注入完毕,并且是正斜杠格式 (例如 C:/Tools)
local base_dir = __SHIM_DIR__ local base_dir = __MIRROR_DIR__
-- 1. 自定义局部变量,方便复用与后续维护 -- 1. 自定义局部变量,方便复用与后续维护
local python_home = base_dir .. "/tools/python39" local python_home = base_dir .. "/tools/python39"
@@ -13,12 +13,14 @@ return {
target = base_dir .. "/tools/numa/numa.exe", target = base_dir .. "/tools/numa/numa.exe",
-- 追加参数 -- 追加参数
args = { "--help" }, args = { "--help" },
aliases = {},
-- 注入环境变量,使用 get_env 获取宿主机当前值 -- 注入环境变量,使用 get_env 获取宿主机当前值
env = { env = {
PATH = { base_dir .. "/tools/numa", get_env("PATH") } PATH = { base_dir .. "/tools/numa", get_env("PATH") }
} }
}, },
["git"] = { ["git"] = {
target = base_dir .. "/git/bin/git.exe", target = base_dir .. "/git/bin/git.exe",
-- 追加参数 -- 追加参数

25
src/error.rs Normal file
View File

@@ -0,0 +1,25 @@
// 构造一个带上下文的 FromLua 转换错误,便于定位配置问题
/// 配置校验错误
pub fn validation_error(message: impl Into<String>) -> 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
}

View File

@@ -11,20 +11,20 @@ pub struct Layout {
impl Layout { impl Layout {
/// 自动解析目录布局: /// 自动解析目录布局:
/// 1. 优先使用环境变量 RSHIM_HOME /// 1. 优先使用环境变量 MIRROR_HOME
/// 2. 兜底回退到当前垫片可执行文件所在目录推断 (exe -> bin -> root) /// 2. 兜底回退到当前垫片可执行文件所在目录推断 (exe -> bin -> root)
pub fn discover(current_exe: &Path) -> Result<Self> { pub fn discover(current_exe: &Path) -> Result<Self> {
// 策略 1: 环境变量优先 // 策略 1: 环境变量优先
if let Ok(home_val) = env::var("MIMIC_HOME") { if let Ok(home_val) = env::var("MIRROR_HOME") {
let trimmed = home_val.trim(); let trimmed = home_val.trim();
if !trimmed.is_empty() { if !trimmed.is_empty() {
debug!(home = %trimmed, "检测到 MIMIC_HOME采用环境变量配置"); debug!(home = %trimmed, "检测到 MIRROR_HOME采用环境变量配置");
return Self::from_base_dir(PathBuf::from(trimmed)); return Self::from_base_dir(PathBuf::from(trimmed));
} }
} }
// 策略 2: 相对路径自动推断兜底 // 策略 2: 相对路径自动推断兜底
debug!("未配置 MIMIC_HOME尝试从当前可执行文件路径推断根目录"); debug!("未配置 MIRROR_HOME尝试从当前可执行文件路径推断根目录");
Self::from_executable(current_exe) Self::from_executable(current_exe)
} }

View File

@@ -1,20 +1,23 @@
extern crate core; extern crate core;
pub mod error;
mod layout; mod layout;
mod loader; mod loader;
mod logger; mod logger;
mod runtime;
mod mirror; mod mirror;
mod runtime;
mod spec; mod spec;
pub mod sys; pub mod sys;
mod utils; mod utils;
mod validators; mod validators;
pub use layout::Layout; pub use layout::Layout;
pub use runtime::LuaRuntime;
pub use mirror::Mirror; pub use mirror::Mirror;
pub use runtime::LuaRuntime;
pub use spec::MirrorSpec; pub use spec::MirrorSpec;
// pub use sys::{ // pub use sys::{
// ERROR_ELEVATION_REQUIRED, EXIT_FAILED_LOAD_SHIM, EXIT_FAILED_SPAWN_PROG, EXIT_FAILED_WAIT_PROG, // ERROR_ELEVATION_REQUIRED, EXIT_FAILED_LOAD_SHIM, EXIT_FAILED_SPAWN_PROG, EXIT_FAILED_WAIT_PROG,
// EXIT_PROG_TERMINATED, execute_elevated, set_console_ctrl_handler, // EXIT_PROG_TERMINATED, execute_elevated, set_console_ctrl_handler,
// }; // };
pub use utils::{lua_string_2_os_string, normalize_path_for_lua, parse_tokens};

View File

@@ -1,10 +1,8 @@
use crate::spec::MirrorSpec; use crate::spec::MirrorSpec;
use crate::validators::LuaValidator;
use crate::{Layout, LuaRuntime}; use crate::{Layout, LuaRuntime};
use anyhow::{Context, Result, bail}; use anyhow::{Context, Result, bail};
use mlua::{FromLua, Lua, Table, Value}; use mlua::Table;
use std::fs; use std::path::PathBuf;
use std::path::Path;
use tracing::{debug, trace, warn}; use tracing::{debug, trace, warn};
pub enum Source { pub enum Source {
@@ -12,7 +10,7 @@ pub enum Source {
Json, Json,
} }
/// ShimSpec 加载器,统一对外暴露多源解析接口 /// MirrorSpec 加载器,统一对外暴露多源解析接口
pub struct SpecLoader; pub struct SpecLoader;
impl SpecLoader { impl SpecLoader {
@@ -31,13 +29,6 @@ impl SpecLoader {
// 策略 1: 尝试加载全局配置文件 mirror.lua // 策略 1: 尝试加载全局配置文件 mirror.lua
let global_config = layout.base_dir.join("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() { if global_config.is_file() {
trace!(path = %global_config.display(), "发现全局配置文件,尝试解析"); trace!(path = %global_config.display(), "发现全局配置文件,尝试解析");
let root_table: Table = runtime.eval_script(&global_config)?; let root_table: Table = runtime.eval_script(&global_config)?;
@@ -62,17 +53,31 @@ impl SpecLoader {
trace!(target = %target_name, "全局配置文件中未包含该目标,继续探查独立配置"); trace!(target = %target_name, "全局配置文件中未包含该目标,继续探查独立配置");
} }
let tools_dir: Option<PathBuf> = runtime.tools_dir()?;
// 策略 2: 降级寻找独立文件 ({exe}.lua)优先顺序tools/ > root/ // 策略 2: 降级寻找独立文件 ({exe}.lua)优先顺序tools/ > root/
let target_filename = format!("{}.lua", target_name); 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 = [ let candidates = [
layout.tools_dir.join(&target_filename), // layout.tools_dir.join(&target_filename),
effective_tools_dir,
layout.base_dir.join(&target_filename), layout.base_dir.join(&target_filename),
]; ];
for config_path in &candidates { for config_path in &candidates {
if config_path.is_file() { if config_path.is_file() {
debug!(path = %config_path.display(), "找到独立配置文件,开始加载"); debug!(path = %config_path.display(), "找到独立配置文件,开始加载");
// 直接泛型反序列化为 ShimConfig // 直接泛型反序列化为 MirrorSpec
return runtime.eval_script::<MirrorSpec>(config_path); return runtime.eval_script::<MirrorSpec>(config_path);
} }
trace!(path = %config_path.display(), "独立配置文件不存在,跳过"); trace!(path = %config_path.display(), "独立配置文件不存在,跳过");
@@ -91,6 +96,4 @@ impl SpecLoader {
fn load_from_json(layout: &Layout) -> Result<MirrorSpec> { fn load_from_json(layout: &Layout) -> Result<MirrorSpec> {
todo!("实现json来源") todo!("实现json来源")
} }
} }

View File

@@ -1,7 +1,8 @@
use mirror::sys::*;
use mirror::Mirror; use mirror::Mirror;
use mirror::sys::*;
use std::ffi::OsString;
use std::{env, process::exit}; use std::{env, process::exit};
use tracing_subscriber::{fmt, EnvFilter}; use tracing_subscriber::{EnvFilter, fmt};
fn main() { fn main() {
//初始化日志:输出到 stderr避免污染 shim 子进程的 stdout //初始化日志:输出到 stderr避免污染 shim 子进程的 stdout
@@ -13,26 +14,31 @@ fn main() {
.init(); .init();
// 2. 注册 Windows 控制台信号 // 2. 注册 Windows 控制台信号
set_console_ctrl_handler(); set_console_ctrl_handler();
// 3. 解析调用参数与代理 Shim 配置 // 3. 解析调用参数与代理 Mirror 配置
let calling_args: Vec<_> = env::args_os().skip(1).collect(); let calling_args: Vec<_> = env::args_os().skip(1).collect();
let mr = match Mirror::new() { let mr = match Mirror::new() {
Ok(v) => v, Ok(v) => v,
Err(e) => { Err(e) => {
eprintln!("加载代理(shim)配置时发生错误: {}", e); eprintln!("加载代理(mirror)配置时发生错误: {}", e);
exit(EXIT_FAILED_LOAD_SHIM); exit(EXIT_FAILED_LOAD_SHIM);
} }
}; };
let combined_args = mr.spec.args.iter().chain(calling_args.iter());
// 构建 Command复用 ShimConfig::to_command含 args/env 注入),避免重复逻辑 // 构建 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() { 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 = mr.spec.args.clone(); let elevated_args: Vec<OsString> = cmd.get_args().map(|s| s.to_os_string()).collect();
args.extend_from_slice(&calling_args);
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) => { Err(e) => {
eprintln!( eprintln!(

View File

@@ -2,10 +2,11 @@ use crate::loader::Source;
use crate::loader::SpecLoader; use crate::loader::SpecLoader;
use crate::{Layout, LuaRuntime, MirrorSpec}; use crate::{Layout, LuaRuntime, MirrorSpec};
use anyhow::{Context, Result}; use anyhow::{Context, Result};
use std::collections::HashMap;
use std::env; use std::env;
use std::ffi::OsString;
use std::process::Command; use std::process::Command;
use tracing::debug; use tracing::debug;
use std::ffi::OsString;
pub struct Mirror { pub struct Mirror {
pub spec: MirrorSpec, pub spec: MirrorSpec,
pub layout: Layout, pub layout: Layout,
@@ -25,7 +26,7 @@ impl Mirror {
debug!( debug!(
target_name = %target_name, target_name = %target_name,
current_exe = %current_exe.display(), current_exe = %current_exe.display(),
"开始加载 Shim 配置" "开始加载 Mirror 配置"
); );
let layout = Layout::discover(&current_exe)?; let layout = Layout::discover(&current_exe)?;
@@ -33,7 +34,7 @@ impl Mirror {
root_dir = %layout.base_dir.display(), root_dir = %layout.base_dir.display(),
bin_dir = %layout.bin_dir.display(), bin_dir = %layout.bin_dir.display(),
tools_dir = %layout.tools_dir.display(), tools_dir = %layout.tools_dir.display(),
"Shim 目录布局解析完成" "Mirror 目录布局解析完成"
); );
let runtime = LuaRuntime::new(&layout)?; let runtime = LuaRuntime::new(&layout)?;
@@ -47,27 +48,53 @@ impl Mirror {
} }
/// 1. 核心参数路由解析Delegation 到 Spec 的路由逻辑) /// 1. 核心参数路由解析Delegation 到 Spec 的路由逻辑)
pub fn resolve_args<I, S>(&self, raw_args: I) -> Vec<OsString> fn resolve_args<I, S>(&self, args: I, aliases: &HashMap<String, Vec<OsString>>) -> Vec<OsString>
where where
I: IntoIterator<Item = S>, I: IntoIterator<Item = S>,
S: Into<OsString>, S: AsRef<std::ffi::OsStr>,
{todo!() {
// self.spec.resolve_args(raw_args) let arg_iter = args.into_iter();
// 1. 预分配容量:利用迭代器的下限提示,避免多次 Realloc
let (lower_bound, _) = arg_iter.size_hint();
let mut expanded_args: Vec<OsString> = 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<I, S>(&self, runtime_args: I) -> Command pub fn to_command<I, S>(&self, combined_args: I) -> Command
where where
I: IntoIterator<Item = S>, I: IntoIterator<Item = S>,
S: AsRef<std::ffi::OsStr>, S: AsRef<std::ffi::OsStr>,
{ {
let mut cmd = Command::new(&self.spec.target); let mut cmd = Command::new(&self.spec.target);
// 1. 先追加配置里固定的默认参数 (例如: node, --max-old-space-size=4096) // 2. 传入预打平的别名字典,查表并展开所有以 `mr:` 为前缀的别名
cmd.args(&self.spec.args); let final_args = self.resolve_args(combined_args, &self.spec.aliases);
println!("拼接后的命令行参数{:?}",final_args);
// 2. 透传外部动态运行时参数 // 3. 将解析展开后的无环参数一次性注入 Command
cmd.args(runtime_args); cmd.args(&final_args);
// 4. 注入配置好的环境变量
for (key, val) in &self.spec.env { for (key, val) in &self.spec.env {
// 直接应用环境变量Lua 端已经处理好字符串拼接或列表合并) // 直接应用环境变量Lua 端已经处理好字符串拼接或列表合并)
cmd.env(key, val); cmd.env(key, val);
@@ -76,4 +103,3 @@ impl Mirror {
cmd cmd
} }
} }

View File

@@ -1,13 +1,15 @@
use crate::{MirrorSpec, Layout}; use crate::Layout;
use anyhow::{Context, Result, anyhow, bail}; use crate::error::syntax_error;
use mlua::{FromLua, Lua, StdLib, Table, Value};
use std::ffi::OsStr;
use std::path::Path;
use std::{env, fs};
use crate::utils::normalize_path_for_lua; 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 { pub struct LuaRuntime {
lua: Lua, lua: Lua,
@@ -23,20 +25,66 @@ impl LuaRuntime {
) )
.context("初始化 Lua 失败")?; .context("初始化 Lua 失败")?;
let globals = lua.globals();
// 统一使用 POSIX 风格路径规范化路径字符串 // 统一使用 POSIX 风格路径规范化路径字符串
let base_dir = normalize_path_for_lua(&layout.base_dir); let base_dir = normalize_path_for_lua(&layout.base_dir);
let tools_dir = normalize_path_for_lua(&layout.tools_dir); let tools_dir = normalize_path_for_lua(&layout.tools_dir);
// 1. 注入锚点变量 __SHIM_DIR__shim 安装根目录) // 1. 注入锚点变量 __MIRROR_DIR__shim 安装根目录)
globals Self::register_mirror_dir(&lua, &base_dir)?;
.set("__SHIM_DIR__", base_dir.clone())
.context("设置 __SHIM_DIR__ 环境变量失败")?;
// 2. 安全暴露 get_env 供配置读取环境变量 // 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<T: FromLua>(&self, path: impl AsRef<Path>) -> Result<T> {
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);
// 使用 @<path> 格式标记 Chunk 名称,确保 Lua 报错时能精确回溯到对应的源文件名与行号。
let chunk_name = format!("@{}", path.display());
self.lua
.load(&code)
.set_name(&chunk_name)
.eval::<T>()
.map_err(syntax_error)
.with_context(|| format!("执行 Lua 配置文件失败: {}", path.display()))
}
pub fn tools_dir(&self) -> Result<Option<PathBuf>> {
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 let get_env = lua
.create_function(|lua, key: String| -> mlua::Result<Table> { .create_function(|lua, key: String| -> mlua::Result<Table> {
// 缺失变量视为空字符串,拆分后得到空表(不贡献任何路径段) // 缺失变量视为空字符串,拆分后得到空表(不贡献任何路径段)
@@ -71,8 +119,39 @@ impl LuaRuntime {
globals globals
.set("get_env", get_env) .set("get_env", get_env)
.context("挂载 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::<Table>("package") { if let Ok(package) = globals.get::<Table>("package") {
let _ = package.set("cpath", ""); let _ = package.set("cpath", "");
let _ = package.set("loadlib", Value::Nil); let _ = package.set("loadlib", Value::Nil);
@@ -86,19 +165,14 @@ impl LuaRuntime {
} }
} }
// 4. 包装 require:配置模块缺失/加载失败时记录日志并跳过该条目, // 2. 获取原生 require 并通过闭包直接持有(无需向全局表注入备份变量)
// 而不是让整个 mirror.lua 解析失败(排查问题时日志可见)
let original_require: mlua::Function = globals let original_require: mlua::Function = globals
.get("require") .get("require")
.context("获取内置 require 函数失败")?; .context("获取内置 require 函数失败")?;
globals
.set("_rshim_original_require", &original_require)
.context("备份原始 require 函数失败")?;
let wrapped_require = lua let wrapped_require = lua
.create_function(|lua, module: String| -> mlua::Result<Value> { .create_function(move |_lua, module: String| -> mlua::Result<Value> {
let original: mlua::Function = lua.globals().get("_rshim_original_require")?; match original_require.call::<Value>(module.as_str()) {
match original.call::<Value>(module.clone()) {
Ok(value) => Ok(value), Ok(value) => Ok(value),
Err(e) => { Err(e) => {
tracing::warn!( tracing::warn!(
@@ -112,44 +186,18 @@ impl LuaRuntime {
}) })
.context("创建包装版 require 函数失败")?; .context("创建包装版 require 函数失败")?;
// 3. 覆盖全局 require
globals globals
.set("require", wrapped_require) .set("require", wrapped_require)
.context("重载 require 函数失败")?; .context("重载 require 函数失败")?;
Ok(Self { lua }) Ok(())
} }
/// 执行指定脚本文件,直接返回完整的 Lua Table
pub fn eval_script<T: FromLua>(&self, path: impl AsRef<Path>) -> Result<T> {
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::<T>()
// .map_err(|e| anyhow!(e.to_string()))
.map_err(map_lua_error)
.with_context(|| format!("执行 Lua 配置文件失败: {}", path.display()))
}
} }
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::*; use super::*;
use crate::Layout; use crate::{Layout, MirrorSpec};
fn test_layout() -> Layout { fn test_layout() -> Layout {
let root = std::env::temp_dir().join("rshim-test-layout"); let root = std::env::temp_dir().join("rshim-test-layout");

View File

@@ -1,20 +1,12 @@
use crate::error::validation_error;
use crate::validators::{JsonValidator, LuaValidator}; use crate::validators::{JsonValidator, LuaValidator};
use mlua::{FromLua, Lua, ObjectLike, Value}; use anyhow::{Result, anyhow};
use std::collections::{HashMap, HashSet}; use mlua::{FromLua, Lua, Value};
use std::collections::HashMap;
use std::ffi::OsString; use std::ffi::OsString;
use std::path::PathBuf; use std::path::PathBuf;
use std::process::Command;
use std::str::FromStr; use std::str::FromStr;
use tinyjson::JsonValue; use tinyjson::JsonValue;
use anyhow::{anyhow, Context, Result};
/// 构造一个带上下文的 FromLua 转换错误,便于定位配置问题
fn conversion_error(message: impl Into<String>) -> mlua::Error {
mlua::Error::FromLuaConversionError {
from: "Lua value",
to: "ShimConfig".into(),
message: Some(message.into()),
}
}
#[derive(Debug, Clone, Default)] #[derive(Debug, Clone, Default)]
@@ -23,85 +15,6 @@ pub struct MirrorSpec {
pub args: Vec<OsString>, pub args: Vec<OsString>,
pub aliases: HashMap<String, Vec<OsString>>, pub aliases: HashMap<String, Vec<OsString>>,
pub env: HashMap<String, OsString>, pub env: HashMap<String, OsString>,
}
impl MirrorSpec {
/// 根据配置快速构建准备执行的 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
// }
/// 全局解构与支持任意深度的别名嵌套展开
pub fn resolve_args<I, S>(&self, raw_args: I) -> Result<Vec<OsString>, String>
where
I: IntoIterator<Item = S>,
S: Into<OsString>,
{
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<String>,
out: &mut Vec<OsString>,
) -> 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 转换 /// 实现 FromLua Trait由 mlua 自动处理 Table 转换
@@ -111,7 +24,7 @@ impl FromLua for MirrorSpec {
let table = match value { let table = match value {
Value::Table(t) => t, Value::Table(t) => t,
_ => { _ => {
return Err(conversion_error(format!( return Err(validation_error(format!(
"期望得到一个 Lua Table 配置对象,实际是 {}", "期望得到一个 Lua Table 配置对象,实际是 {}",
value.type_name() value.type_name()
))); )));
@@ -121,13 +34,12 @@ impl FromLua for MirrorSpec {
// 必填字段: target严格限定为字符串避免数字被 mlua 宽松转为字符串后掩盖错误) // 必填字段: target严格限定为字符串避免数字被 mlua 宽松转为字符串后掩盖错误)
let target = match table.get::<Option<Value>>("target")? { let target = match table.get::<Option<Value>>("target")? {
None | Some(Value::Nil) => { None | Some(Value::Nil) => {
return Err(conversion_error("缺少必填字段 target应为字符串路径")); return Err(validation_error("缺少必填字段 target应为字符串路径"));
} }
Some(target_val) => LuaValidator::parse_target(&target_val)?, Some(target_val) => LuaValidator::parse_target(&target_val)?,
}; };
// 可选字段: args缺失或 nil 时默认空列表;保留空字符串参数,与提权路径的 "" 语义一致) // 可选字段: args缺失或 nil 时默认空列表;保留空字符串参数,与提权路径的 "" 语义一致)
let args = match table.get::<Option<Value>>("args")? { let args = match table.get::<Option<Value>>("args")? {
None | Some(Value::Nil) => Vec::new(), None | Some(Value::Nil) => Vec::new(),
Some(args_val) => LuaValidator::parse_args(&args_val)?, Some(args_val) => LuaValidator::parse_args(&args_val)?,
@@ -137,15 +49,20 @@ impl FromLua for MirrorSpec {
let env = match table.get::<Option<Value>>("env")? { let env = match table.get::<Option<Value>>("env")? {
None | Some(Value::Nil) => HashMap::new(), None | Some(Value::Nil) => HashMap::new(),
Some(env_val) => LuaValidator::parse_env(&env_val)?, Some(env_val) => LuaValidator::parse_env(&env_val)?,
}; };
// 可选字段: aliases只允许缺失/nil其他类型由 Option<Table> 转换报错,不再静默忽略)
let aliases = match table.get::<Option<Value>>("aliases")? { let aliases = match table.get::<Option<Value>>("aliases")? {
None | Some(Value::Nil) => HashMap::new(), None | Some(Value::Nil) => HashMap::new(),
Some(env_val) => LuaValidator::parse_aliases(&env_val)?, Some(env_val) => LuaValidator::parse_aliases(&env_val)?,
}; };
println!("环境变量结果:{:?}", env); // println!("环境变量结果:{:?}", env);
Ok(MirrorSpec { target, args,aliases, env }) Ok(MirrorSpec {
target,
args,
aliases,
env,
})
} }
} }
@@ -154,8 +71,7 @@ impl TryFrom<&str> for MirrorSpec {
fn try_from(json_str: &str) -> Result<Self> { fn try_from(json_str: &str) -> Result<Self> {
// 1. 解析 JSON 字符串为 JsonValue 树 // 1. 解析 JSON 字符串为 JsonValue 树
let root = JsonValue::from_str(json_str) let root = JsonValue::from_str(json_str).map_err(|e| anyhow!("JSON 语法错误: {}", e))?;
.map_err(|e| anyhow!("JSON 语法错误: {}", e))?;
// 根节点必须是一个 JSON Object // 根节点必须是一个 JSON Object
let map: &HashMap<String, JsonValue> = root let map: &HashMap<String, JsonValue> = root
@@ -190,4 +106,4 @@ impl TryFrom<&str> for MirrorSpec {
env, env,
}) })
} }
} }

View File

@@ -3,28 +3,26 @@ use std::{env, mem::size_of, path::Path, ptr::null_mut};
use std::ffi::{OsStr, OsString}; 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::Win32::Foundation::CloseHandle;
use windows_sys::{ use windows_sys::{
core::BOOL,
Win32::{ Win32::{
Foundation::{FALSE, TRUE}, Foundation::{FALSE, TRUE},
System::{ System::{
Com::{CoInitializeEx, COINIT_APARTMENTTHREADED, COINIT_DISABLE_OLE1DDE}, Com::{COINIT_APARTMENTTHREADED, COINIT_DISABLE_OLE1DDE, CoInitializeEx},
Console::{ Console::{
SetConsoleCtrlHandler, CTRL_BREAK_EVENT, CTRL_CLOSE_EVENT, CTRL_C_EVENT, CTRL_BREAK_EVENT, CTRL_C_EVENT, CTRL_CLOSE_EVENT, CTRL_LOGOFF_EVENT,
CTRL_LOGOFF_EVENT, CTRL_SHUTDOWN_EVENT, CTRL_SHUTDOWN_EVENT, SetConsoleCtrlHandler,
}, },
Threading::{GetExitCodeProcess, WaitForSingleObject, INFINITE}, Threading::{GetExitCodeProcess, INFINITE, WaitForSingleObject},
}, },
UI::{ UI::{
Shell::{ Shell::{SEE_MASK_NOASYNC, SEE_MASK_NOCLOSEPROCESS},
SEE_MASK_NOASYNC, SEE_MASK_NOCLOSEPROCESS,
},
WindowsAndMessaging::SW_NORMAL, WindowsAndMessaging::SW_NORMAL,
}, },
}, },
core::BOOL,
}; };
pub const EXIT_FAILED_LOAD_SHIM: i32 = 1; pub const EXIT_FAILED_LOAD_SHIM: i32 = 1;
@@ -151,5 +149,4 @@ pub fn execute_elevated(
} }
exit_code as i32 exit_code as i32
} }

View File

@@ -1,7 +1,233 @@
use mlua::LuaString;
use std::ffi::OsString;
use std::path::Path; use std::path::Path;
/// 将 Path 转换为适合 Lua 使用的安全字符串路径
pub fn normalize_path_for_lua(path: &Path) -> String { pub fn normalize_path_for_lua(path: &Path) -> String {
// 自动将 Windows UNC 规范路径转回传统路径 // 自动将 Windows UNC 规范路径转回传统路径
let simplified = dunce::simplified(path); let simplified = dunce::simplified(path);
simplified.to_string_lossy().replace('\\', "/") simplified.to_string_lossy().replace('\\', "/")
} }
/// 将 Lua 字符串转为 Rust 字符串;非 UTF-8 按系统默认的 ANSI/OEM (如 GBK) 进行安全解码
pub fn lua_string_2_os_string(s: &LuaString) -> mlua::Result<OsString> {
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<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
}
#[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<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"]);
}
}

View File

@@ -1,3 +1,5 @@
use crate::error::validation_error;
use crate::{lua_string_2_os_string, parse_tokens};
use anyhow::{Context, Result, anyhow}; use anyhow::{Context, Result, anyhow};
use mlua::{LuaString, Table, Value}; use mlua::{LuaString, Table, Value};
use std::collections::{HashMap, HashSet}; use std::collections::{HashMap, HashSet};
@@ -7,89 +9,6 @@ use std::path::PathBuf;
use std::str::FromStr; use std::str::FromStr;
use tinyjson::JsonValue; use tinyjson::JsonValue;
/// 构造一个带上下文的 FromLua 转换错误,便于定位配置问题
fn conversion_error(message: impl Into<String>) -> 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<OsString> {
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 值校验器:针对不同上下文定义校验规则 /// Lua 值校验器:针对不同上下文定义校验规则
#[derive(Debug, Clone, Copy, PartialEq, Eq)] #[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum LuaValidator { pub enum LuaValidator {
@@ -118,13 +37,13 @@ impl LuaValidator {
Value::String(_) | Value::Integer(_) | Value::Number(_) | Value::Boolean(_) => Ok(()), Value::String(_) | Value::Integer(_) | Value::Number(_) | Value::Boolean(_) => Ok(()),
Value::Table(tbl) => match self { Value::Table(tbl) => match self {
Self::Env | Self::Aliases => self.validate_sequence_table(tbl), Self::Env | Self::Aliases => self.validate_sequence_table(tbl),
Self::Args => Err(conversion_error(format!( Self::Args => Err(validation_error(format!(
"{} 仅支持一维数组,不能包含嵌套 Table,索引位置: {index}", "{} 仅支持一维数组,不能包含嵌套 Table,索引位置: {index}",
self 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, self,
index, index,
@@ -154,13 +73,13 @@ impl LuaValidator {
match key { match key {
Value::Integer(i) if i >= 1 && i < index => {} // 已遍历放行 Value::Integer(i) if i >= 1 && i < index => {} // 已遍历放行
Value::Integer(_) => { Value::Integer(_) => {
return Err(conversion_error(format!( return Err(validation_error(format!(
"{}{} 个元素为 nil 或不存在(请检查漏写引号或变量未定义)", "{}{} 个元素为 nil 或不存在(请检查漏写引号或变量未定义)",
self, index self, index
))); )));
} }
_ => { _ => {
return Err(conversion_error(format!( return Err(validation_error(format!(
"{} 必须是纯列表,不能包含键值对/字典结构", "{} 必须是纯列表,不能包含键值对/字典结构",
self self
))); )));
@@ -187,7 +106,7 @@ impl LuaValidator {
{ {
use std::os::windows::ffi::OsStrExt; use std::os::windows::ffi::OsStrExt;
if os_str.encode_wide().any(|c| c == 0) { if os_str.encode_wide().any(|c| c == 0) {
return Err(conversion_error(format!( return Err(validation_error(format!(
"{} 的值不能包含 NUL 字符", "{} 的值不能包含 NUL 字符",
context context
))); )));
@@ -206,7 +125,7 @@ impl LuaValidator {
Self::Args | Self::Aliases => { Self::Args | Self::Aliases => {
if let Some(str_ref) = os_str.to_str() { if let Some(str_ref) = os_str.to_str() {
// 使用 Tokenizer 切分空格与引号 // 使用 Tokenizer 切分空格与引号
out.extend(Self::parse_tokens(str_ref)); out.extend(parse_tokens(str_ref));
} else { } else {
// 对于无法转为 UTF-8 的特殊二进制数据,作为整体追加 // 对于无法转为 UTF-8 的特殊二进制数据,作为整体追加
out.push(os_str); out.push(os_str);
@@ -234,6 +153,7 @@ impl LuaValidator {
if matches!(item, Value::Nil) { if matches!(item, Value::Nil) {
break; break;
} }
Self::collect_value_into(self, &item, out)?; Self::collect_value_into(self, &item, out)?;
index += 1; index += 1;
} }
@@ -246,33 +166,33 @@ impl LuaValidator {
let name_str = match name { let name_str = match name {
Value::String(s) => match s.to_str() { Value::String(s) => match s.to_str() {
Ok(str_ref) => str_ref.to_string(), Ok(str_ref) => str_ref.to_string(),
Err(_) => return Err(conversion_error("环境变量名必须是合法的 UTF-8 字符串")), Err(_) => return Err(validation_error("环境变量名必须是合法的 UTF-8 字符串")),
}, },
other => { other => {
return Err(conversion_error(format!( return Err(validation_error(format!(
"环境变量键名类型错误:期望 string实际是 {}", "环境变量键名类型错误:期望 string实际是 {}",
other.type_name() other.type_name()
))); )));
} }
}; };
if name_str.is_empty() { if name_str.is_empty() {
return Err(conversion_error("环境变量名不能为空")); return Err(validation_error("环境变量名不能为空"));
} }
if name_str.contains('=') { if name_str.contains('=') {
return Err(conversion_error(format!( return Err(validation_error(format!(
"环境变量名 [{}] 不能包含 '='", "环境变量名 [{}] 不能包含 '='",
name_str name_str
))); )));
} }
if name_str.contains('\0') { if name_str.contains('\0') {
return Err(conversion_error(format!( return Err(validation_error(format!(
"环境变量名 [{}] 不能包含 NUL 字符", "环境变量名 [{}] 不能包含 NUL 字符",
name_str name_str
))); )));
} }
Ok(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 { match value {
Value::Nil Value::Nil
| Value::String(_) | Value::String(_)
@@ -280,7 +200,7 @@ impl LuaValidator {
| Value::Number(_) | Value::Number(_)
| Value::Boolean(_) => Ok(()), | Value::Boolean(_) => Ok(()),
Value::Table(tbl) => self.validate_sequence_table(tbl), Value::Table(tbl) => self.validate_sequence_table(tbl),
other => Err(conversion_error(format!( other => Err(validation_error(format!(
"{} [{}]不支持 {} 类型, 仅支持 string/number/boolean 或嵌套数组", "{} [{}]不支持 {} 类型, 仅支持 string/number/boolean 或嵌套数组",
self, self,
name, name,
@@ -299,8 +219,8 @@ impl LuaValidator {
Self::ensure_no_nul(&ctx, &os_str)?; Self::ensure_no_nul(&ctx, &os_str)?;
Ok(PathBuf::from(os_str)) Ok(PathBuf::from(os_str))
} }
Value::Nil => Err(conversion_error("缺少必填字段 target应为字符串路径")), Value::Nil => Err(validation_error("缺少必填字段 target应为字符串路径")),
other => Err(conversion_error(format!( other => Err(validation_error(format!(
"{} 需为有效的路径且类型必须是字符串,实际类型是 {}", "{} 需为有效的路径且类型必须是字符串,实际类型是 {}",
ctx, ctx,
other.type_name() other.type_name()
@@ -323,7 +243,10 @@ impl LuaValidator {
Value::Table(tbl) => { Value::Table(tbl) => {
ctx.validate_sequence_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)?; Self::collect_value_into(&ctx, value, &mut raw_parts)?;
for part in &raw_parts { for part in &raw_parts {
@@ -331,7 +254,7 @@ impl LuaValidator {
} }
Ok(raw_parts) Ok(raw_parts)
} }
other => Err(conversion_error(format!( other => Err(validation_error(format!(
"{} 必须是数组列表,实际类型是 {}", "{} 必须是数组列表,实际类型是 {}",
ctx, ctx,
other.type_name() other.type_name()
@@ -351,12 +274,19 @@ impl LuaValidator {
name: &str, name: &str,
raw_val: &Value, raw_val: &Value,
) -> mlua::Result<OsString> { ) -> mlua::Result<OsString> {
context.parse_raw_val(name, raw_val)?; context.parse_val(name, raw_val)?;
// let ctx = format!("{} [{}]", context, name); // 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)?; Self::collect_value_into(&context, raw_val, &mut parts)?;
// context.collect_value_into(raw_val, &mut parts)?;
// 校验每个展开元素的 NUL 字符 // 校验每个展开元素的 NUL 字符
for part in &parts { for part in &parts {
@@ -364,7 +294,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!("{} 的值无法用系统路径分隔符拼接: {}", context, e)) validation_error(format!("{} 的值无法用系统路径分隔符拼接: {}", context, e))
})?; })?;
Ok(joined_os_str) Ok(joined_os_str)
} }
@@ -375,7 +305,7 @@ impl LuaValidator {
let Value::Table(tbl) = value else { let Value::Table(tbl) = value else {
// 这里作为防御性校验,外部虽然 match 过了,但内层仍保持强类型安全 // 这里作为防御性校验,外部虽然 match 过了,但内层仍保持强类型安全
return Err(conversion_error(format!( return Err(validation_error(format!(
"{} 必须是键值表 (table),实际类型是 {}", "{} 必须是键值表 (table),实际类型是 {}",
ctx, ctx,
value.type_name() value.type_name()
@@ -414,12 +344,18 @@ impl LuaValidator {
name: &str, name: &str,
raw_val: &Value, raw_val: &Value,
) -> mlua::Result<Vec<OsString>> { ) -> mlua::Result<Vec<OsString>> {
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(); let mut parts = Vec::with_capacity(capacity);
Self::collect_value_into(&context, raw_val, &mut parts)?; Self::collect_value_into(context, raw_val, &mut parts)?;
Ok(parts) Ok(parts)
} }
/// 解析并打平别名表 (aliases) /// 解析并打平别名表 (aliases)
@@ -431,7 +367,7 @@ impl LuaValidator {
let ctx = Self::Aliases; let ctx = Self::Aliases;
// 1. 处理 nil / None 的情况,直接返回空 HashMap // 1. 处理 nil / None 的情况,直接返回空 HashMap
let Value::Table(table) = value else { let Value::Table(table) = value else {
return Err(conversion_error(format!( return Err(validation_error(format!(
"{} 必须是键值表 (table) ,当前类型: {}", "{} 必须是键值表 (table) ,当前类型: {}",
ctx, ctx,
value.type_name() value.type_name()
@@ -459,6 +395,7 @@ impl LuaValidator {
// 阶段三:无环前提下的高效展开 // 阶段三:无环前提下的高效展开
let mut flattened_aliases: HashMap<String, Vec<OsString>> = let mut flattened_aliases: HashMap<String, Vec<OsString>> =
HashMap::with_capacity(raw_aliases.len()); HashMap::with_capacity(raw_aliases.len());
let mut resolved_args = Vec::new(); let mut resolved_args = Vec::new();
for key in raw_aliases.keys() { for key in raw_aliases.keys() {
@@ -478,7 +415,7 @@ impl LuaValidator {
) -> mlua::Result<()> { ) -> mlua::Result<()> {
// 递归栈中再次遇到相同的 Key说明存在死循环 // 递归栈中再次遇到相同的 Key说明存在死循环
if visited_stack.contains(current_key) { if visited_stack.contains(current_key) {
return Err(conversion_error(format!( return Err(validation_error(format!(
"配置加载失败: 检测到别名循环嵌套依赖 'mr:{}'", "配置加载失败: 检测到别名循环嵌套依赖 'mr:{}'",
current_key current_key
))); )));
@@ -528,69 +465,6 @@ 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) 的类型校验与字段提取器
@@ -754,7 +628,7 @@ mod tests {
-- 4. 嵌套别名组合(加载期会自动展开并进行死环检测) -- 4. 嵌套别名组合(加载期会自动展开并进行死环检测)
base_log = "log --graph", 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 或空字符串(解析为空参数列表) -- 5. nil 或空字符串(解析为空参数列表)
empty_alias = nil, empty_alias = nil,
@@ -765,7 +639,7 @@ mod tests {
) )
.unwrap(); .unwrap();
assert_eq!(cfg.target, PathBuf::from("C:/tools/git.exe")); 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")); assert_eq!(cfg.env.get("HOME").unwrap().to_str(), Some("C:/tools/home"));
// PATH 前缀来自配置,随后附加 get_env("PATH") 拆出的宿主 PATH 段 // PATH 前缀来自配置,随后附加 get_env("PATH") 拆出的宿主 PATH 段
let path = cfg.env.get("PATH").unwrap().to_str().unwrap(); let path = cfg.env.get("PATH").unwrap().to_str().unwrap();
@@ -779,7 +653,7 @@ mod tests {
// 1. 整体字符串保留原样trim 后),不切分空格 // 1. 整体字符串保留原样trim 后),不切分空格
assert_eq!( assert_eq!(
cfg.aliases.get("st").unwrap(), cfg.aliases.get("st").unwrap(),
&vec![OsString::from("status -s")] &vec![OsString::from("status"), OsString::from("-s")]
); );
// 2. 连续数组表:按顺序转为 OsString 列表 // 2. 连续数组表:按顺序转为 OsString 列表
@@ -804,12 +678,19 @@ mod tests {
// base_log 本身为 "log --graph" // base_log 本身为 "log --graph"
assert_eq!( assert_eq!(
cfg.aliases.get("base_log").unwrap(), 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" // all_log 展开 mr:base_log 替换为 "log --graph",追加 "--all"
assert_eq!( assert_eq!(
cfg.aliases.get("all_log").unwrap(), 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 // 5. nil 与纯空白字符串:解析为空 Vec
@@ -927,98 +808,3 @@ 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"]);
}
}