From 5e69a6a98019bf2405b1177febbf969d4ce8b217 Mon Sep 17 00:00:00 2001 From: CNWei Date: Fri, 14 Aug 2026 20:09:34 +0800 Subject: [PATCH] =?UTF-8?q?=20refactor:=20=E8=BF=81=E7=A7=BB=20winapi=20?= =?UTF-8?q?=E5=88=B0=20windows-sys=EF=BC=8C=E4=BF=AE=E5=A4=8D=E9=85=8D?= =?UTF-8?q?=E7=BD=AE=E8=A7=A3=E6=9E=90=E6=BC=8F=E6=B4=9E?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 依赖替换为 windows-sys 0.61,main.rs 全面适配新 API - 配置解析错误显式传播:env/args 类型错误、数组空洞、非法键不再静默 吞错 - 修正空环境变量与空参数语义,补充 UTF-8 与路径拼接校验 - require 容错移至 Rust 侧,模块加载失败记录日志并跳过 - 新增配置解析与运行时单元测试(19 个) --- .gitignore | 1 + Cargo.toml | 39 +++-- build.rs | 9 +- shims.lua | 11 +- src/config.rs | 405 ++++++++++++++++++++++++++++++++++++++++--------- src/env.rs | 57 ------- src/error.rs | 8 +- src/layout.rs | 45 ++++++ src/lib.rs | 5 +- src/logger.rs | 39 +++++ src/main.rs | 133 ++++++++-------- src/runtime.rs | 130 ++++++++++++---- src/shim.rs | 87 ++++++++--- 13 files changed, 691 insertions(+), 278 deletions(-) delete mode 100644 src/env.rs create mode 100644 src/layout.rs create mode 100644 src/logger.rs diff --git a/.gitignore b/.gitignore index bd74b9e..bba9b17 100644 --- a/.gitignore +++ b/.gitignore @@ -1,4 +1,5 @@ /target /.vscode +/.idea *.exe ./Cargo.lock diff --git a/Cargo.toml b/Cargo.toml index 1c4fdba..977a0ed 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,32 +1,31 @@ [package] name = "rshim" version = "0.1.0" -authors = ["anonymous "] edition = "2024" - -# See more keys and their definitions at https://doc.rust-lang.org/cargo/reference/manifest.html - +rust-version = "1.85" +license = "MIT OR Unlicense" +description = "A fast, safe Rust shim launcher for Scoop" [profile.release] opt-level = "z" panic = "abort" + [dependencies] -fs-err = "3.3.1" -unicode-bom = "2.0.3" -thiserror={version = "2.0.20"} +thiserror = "2.0.20" + mlua = { version = "0.12.0", features = ["lua54", "vendored"] } -winapi = { version = "0.3", features = [ - "wincon", - "consoleapi", - "minwindef", - "shellapi", - "winuser", - "synchapi", - "combaseapi", - "winbase", - "processthreadsapi", - "objbase", - "impl-default" +windows-sys = { version = "0.61.2", features = [ + "Win32_Foundation", + "Win32_System_Com", + "Win32_System_Console", + "Win32_System_Registry", + "Win32_System_Threading", + "Win32_UI_Shell", + "Win32_UI_WindowsAndMessaging", ] } +dunce = "1.0.5" - +# 日志 +tracing = "0.1.44" +tracing-subscriber = { version = "0.3", features = ["env-filter"] } +tracing-appender = "0.2" diff --git a/build.rs b/build.rs index 5eee996..fb4c4c0 100644 --- a/build.rs +++ b/build.rs @@ -18,18 +18,17 @@ fn main() { // 在 target/debug/ 或 target/release/ 下创建 bin2 目录 let output_dir = target_dir.join(&profile).join("bin2"); - // cargo:rerun-if-changed -> 当指定文件变化时,重新运行 xxx println!("cargo:rerun-if-changed=build.rs"); println!("cargo:rerun-if-changed=shims.lua"); println!("cargo:rustc-env=OUTPUT_DIR={}", output_dir.display()); // 创建目录 - let dirs = ["bin", "tools"]; - for dir in &dirs { - let path = output_dir.join(dir); + let subdirs = ["bin", "tools"]; + for subdir in &subdirs { + let path = output_dir.join(subdir); if !path.exists() { - fs::create_dir_all(&path).expect(&format!("Failed to create {} directory", dir)); + fs::create_dir_all(&path).expect(&format!("Failed to create {} directory", subdir)); println!("Created: {:?}", path); } } diff --git a/shims.lua b/shims.lua index ff17b3d..6f20176 100644 --- a/shims.lua +++ b/shims.lua @@ -10,7 +10,7 @@ return { -- 1. 标准相对路径 + 正斜杠拼接 (最推荐,绿色便携) --------------------------------------------------- ["git"] = { - path = base_dir .. "/git/bin/git.exe", + target = base_dir .. "/git/bin/git.exe", -- 追加参数 args = { "--no-pager" }, -- 注入环境变量,使用 get_env 获取宿主机当前值 @@ -24,7 +24,7 @@ return { --------------------------------------------------- ["python"] = { -- 在 [[]] 内部,\ 不需要写成 \\,直接粘贴即可 - path = [[C:\Python310\python.exe]], + target = [[C:\Python310\python.exe]], args = { "-B" }, -- 空字典也是合法的,等同于不设置 env = { @@ -46,14 +46,15 @@ return { -- 3. 极简参数覆盖 (没有 args 和 env) --------------------------------------------------- ["curl"] = { - path = base_dir .. "/curl/curl.exe" + target = base_dir .. "/curl/curl.exe" }, --------------------------------------------------- -- 4. 模块化路由 (得益于 Rust 注入的 package.path) - -- require 能够直接在当前目录 或 conf.d/ 目录下寻找 node.lua + -- require 能够直接在当前目录 或 tools/ 目录下寻找 node.lua + -- 模块缺失时由 Rust 侧记录日志并跳过该条目(见 runtime.rs) --------------------------------------------------- ["node"] = require("node"), ["npm"] = require("npm") -} \ No newline at end of file +} diff --git a/src/config.rs b/src/config.rs index 68dbb3d..3166845 100644 --- a/src/config.rs +++ b/src/config.rs @@ -1,29 +1,135 @@ +use mlua::{FromLua, Lua, Table, Value}; use std::collections::HashMap; use std::path::PathBuf; use std::process::Command; -use mlua::{FromLua, Lua, Table, Value}; -#[derive(Debug, Clone)] +/// 构造一个带上下文的 FromLua 转换错误,便于定位配置问题 +fn conversion_error(message: impl Into) -> mlua::Error { + mlua::Error::FromLuaConversionError { + from: "Lua value".into(), + to: "ShimConfig".into(), + message: Some(message.into()), + } +} + +/// 将 Lua 字符串转为 Rust 字符串;非 UTF-8 字节替换为 U+FFFD 并告警 +fn lua_string_to_string(s: mlua::LuaString) -> String { + match s.to_str() { + Ok(str_val) => str_val.to_string(), + Err(_) => { + tracing::warn!("配置字符串包含非 UTF-8 字节,已替换为 U+FFFD"); + s.to_string_lossy() + } + } +} + +/// 按连续整数下标遍历表的序列部分;存在空洞或非序列键时返回错误,避免静默截断 +fn for_each_sequence_item( + tbl: &Table, + mut f: impl FnMut(Value) -> mlua::Result<()>, +) -> mlua::Result<()> { + let mut index = 1i64; + loop { + let item: Value = tbl.raw_get(index)?; + if matches!(item, Value::Nil) { + break; + } + f(item)?; + index += 1; + } + // 校验剩余键:只允许已经被遍历的连续整数键 + for pair in tbl.pairs::() { + let (key, _) = pair?; + match key { + Value::Integer(i) if i >= 1 && i < index => {} + _ => { + return Err(conversion_error(format!( + "数组只能包含连续的整数下标 [1..{}],发现非序列键或空洞", + index - 1 + ))); + } + } + } + Ok(()) +} + +/// 校验环境变量名的合法性(Windows 约束:非空、不含 '='、不含 NUL) +fn validate_env_key(key: &str) -> mlua::Result<()> { + if key.is_empty() { + return Err(conversion_error("环境变量名不能为空")); + } + if key.contains('=') { + return Err(conversion_error(format!( + "环境变量名 [{}] 不能包含 '='", + key + ))); + } + if key.contains('\0') { + return Err(conversion_error(format!( + "环境变量名 [{}] 不能包含 NUL 字符", + key + ))); + } + Ok(()) +} + +/// 递归将任意 Lua Value 展开为扁平的字符串片段列表 +fn collect_env_strings(value: Value, out: &mut Vec) -> mlua::Result<()> { + match value { + // 1. 字符串 + Value::String(s) => out.push(lua_string_to_string(s)), + // 2. 整数与浮点数 + Value::Integer(i) => out.push(i.to_string()), + Value::Number(n) => { + tracing::warn!(value = %n, "环境变量中的浮点数将按十进制格式转换为字符串"); + out.push(n.to_string()); + } + // 3. 布尔值 + Value::Boolean(b) => { + tracing::warn!(value = %b, "环境变量中的布尔值将转换为字符串"); + out.push(b.to_string()); + } + // 4. 表/数组:必须是连续整数下标的纯序列,递归解包(支持任意深度的嵌套数组) + Value::Table(tbl) => for_each_sequence_item(&tbl, |item| collect_env_strings(item, out))?, + // 5. 安全忽 + // 略 nil + Value::Nil => {} + // 6. 无法转为环境变量的非法类型(函数、协程、UserData 等):报错而非静默忽略 + other => { + return Err(conversion_error(format!( + "环境变量值不支持类型 {}(仅支持 string/number/boolean/数组)", + other.type_name() + ))); + } + } + Ok(()) +} + +#[derive(Debug, Clone, Default)] pub struct ShimConfig { - pub target_path: PathBuf, - pub args: Option>, - pub envs: Option>, + pub target: PathBuf, + pub args: Vec, + pub env: HashMap, } impl ShimConfig { /// 根据配置快速构建准备执行的 Command 对象 - pub fn to_command(&self) -> Command { - let mut cmd = Command::new(&self.target_path); + pub fn to_command(&self, runtime_args: I) -> Command + where + I: IntoIterator, + S: AsRef, + { + let mut cmd = Command::new(&self.target); - if let Some(args) = &self.args { - cmd.args(args); - } + // 1. 先追加配置里固定的默认参数 (例如: node, --max-old-space-size=4096) + cmd.args(&self.args); - if let Some(envs) = &self.envs { - for (key, val) in envs { - // 直接应用环境变量(Lua 端已经处理好字符串拼接或列表合并) - cmd.env(key, val); - } + // 2. 透传外部动态运行时参数 + cmd.args(runtime_args); + + for (key, val) in &self.env { + // 直接应用环境变量(Lua 端已经处理好字符串拼接或列表合并) + cmd.env(key, val); } cmd @@ -33,60 +139,221 @@ impl ShimConfig { /// 实现 FromLua Trait,由 mlua 自动处理 Table 转换 impl FromLua for ShimConfig { fn from_lua(value: Value, _lua: &Lua) -> mlua::Result { - match value { - Value::Table(table) => { - let path_str: String = table.get("path")?; - let args: Option> = table.get("args")?; - // 解析 env Table - let mut envs_map = HashMap::new(); - if let Ok(env_table) = table.get::("env") { - // 获取当前系统的路径分隔符(Windows 为 ";",Unix 为 ":") - #[cfg(windows)] - let sep = ";"; - #[cfg(not(windows))] - let sep = ":"; - - for pair in env_table.pairs::() { - let (k, v) = pair?; - match v { - // 情况 1: 普通字符串,如 HOME = "C:/path" -> 直接覆盖 - Value::String(s) => { - envs_map.insert(k, s.to_str()?.to_string()); - } - // 情况 2: 数组 Table,如 PATH = { bin_dir, get_env("PATH") } - Value::Table(arr) => { - let paths: Vec = arr - // 将 arr 作为序列(数组)处理,每个元素转为 String - .sequence_values::() - .filter_map(|r| r.ok()) - .filter(|s| !s.is_empty()) // 过滤空串,防止生成不必要的连续 ;; - .collect(); - - let combined = paths.join(sep); - envs_map.insert(k, combined); - } - _ => {} - } - } - } - - let envs = if envs_map.is_empty() { - None - } else { - Some(envs_map) - }; - - Ok(ShimConfig { - target_path: PathBuf::from(path_str), - args, - envs, - }) + let table = match value { + Value::Table(t) => t, + _ => { + return Err(conversion_error(format!( + "期望得到一个 Lua Table 配置对象,实际是 {}", + value.type_name() + ))); + } + }; + + // 必填字段: target(严格限定为字符串,避免数字被 mlua 宽松转为字符串后掩盖错误) + let target = match table.get::("target")? { + Value::String(s) => lua_string_to_string(s), + Value::Nil => { + return Err(conversion_error("缺少必填字段 target(应为字符串路径)")); + } + other => { + return Err(conversion_error(format!( + "target 必须是字符串,实际是 {}", + other.type_name() + ))); + } + }; + // 可选字段: args(缺失或 nil 时默认空列表;保留空字符串参数,与提权路径的 "" 语义一致) + let args = match table.get::>("args")? { + None => Vec::new(), + Some(Value::Table(tbl)) => { + let mut args = Vec::new(); + for_each_sequence_item(&tbl, |item| match item { + Value::String(s) => { + args.push(lua_string_to_string(s)); + Ok(()) + } + other => Err(conversion_error(format!( + "args 数组元素必须是字符串,实际是 {}", + other.type_name() + ))), + })?; + args + } + Some(other) => { + return Err(conversion_error(format!( + "args 必须是字符串数组,实际是 {}", + other.type_name() + ))); + } + }; + // 可选字段: env(只允许缺失/nil,其他类型由 Option
转换报错,不再静默忽略) + let env_table: Option
= table.get("env")?; + println!("env_table:{:?}", env_table); + let mut env = HashMap::new(); + if let Some(env_table) = env_table { + for pair in env_table.pairs::() { + let (key, value) = pair?; + validate_env_key(&key)?; + + let mut parts = Vec::new(); + collect_env_strings(value, &mut parts)?; + + // 空数组/空串也显式设置(空值表示清空该变量), + // 与“未配置该变量(继承宿主环境)”相区分 + let joined_os_str = + std::env::join_paths(parts.iter().map(PathBuf::from)).map_err(|e| { + conversion_error(format!( + "环境变量 [{}] 的值无法用系统路径分隔符拼接: {}", + key, e + )) + })?; + println!("环境变量拼接结果:{:?}", joined_os_str); + + let joined_str = joined_os_str.into_string().map_err(|_| { + conversion_error(format!("环境变量 [{}] 的值不是合法文本", key)) + })?; + if joined_str.contains('\0') { + return Err(conversion_error(format!( + "环境变量 [{}] 的值不能包含 NUL 字符", + key + ))); + } + env.insert(key, joined_str); } - _ => Err(mlua::Error::FromLuaConversionError { - from: value.type_name(), - to: "ShimConfig".into(), - message: Some("Expected a Lua table".to_string()), - }), } + println!("环境变量结果:{:?}", env); + Ok(ShimConfig { + target: PathBuf::from(target), + args, + env, + }) } -} \ No newline at end of file +} + +#[cfg(test)] +mod tests { + use super::*; + + fn parse(src: &str) -> mlua::Result { + let lua = Lua::new(); + let value = lua.load(src).eval::()?; + ShimConfig::from_lua(value, &lua) + } + + #[test] + fn parses_basic_config() { + let cfg = parse( + r#" + return { + target = "C:/tools/git.exe", + args = { "--no-pager" }, + env = { + PATH = { "C:/tools/git/bin", "C:/Windows" }, + HOME = "C:/tools/home", + } + } + "#, + ) + .unwrap(); + assert_eq!(cfg.target, PathBuf::from("C:/tools/git.exe")); + assert_eq!(cfg.args, vec!["--no-pager"]); + assert_eq!( + cfg.env.get("HOME").map(String::as_str), + Some("C:/tools/home") + ); + assert_eq!( + cfg.env.get("PATH").map(String::as_str), + Some("C:/tools/git/bin;C:/Windows") + ); + } + + #[test] + fn missing_args_and_env_are_empty() { + let cfg = parse(r#"return { target = "t.exe" }"#).unwrap(); + assert!(cfg.args.is_empty()); + assert!(cfg.env.is_empty()); + } + + #[test] + fn keeps_empty_args() { + let cfg = parse(r#"return { target = "t.exe", args = { "" } }"#).unwrap(); + assert_eq!(cfg.args, vec![""]); + } + + #[test] + fn keeps_empty_env_value() { + let cfg = parse(r#"return { target = "t.exe", env = { FOO = "" } }"#).unwrap(); + assert_eq!(cfg.env.get("FOO").map(String::as_str), Some("")); + } + + #[test] + fn empty_array_clears_env_var() { + let cfg = parse(r#"return { target = "t.exe", env = { PATH = {} } }"#).unwrap(); + assert_eq!(cfg.env.get("PATH").map(String::as_str), Some("")); + } + + #[test] + fn rejects_wrong_env_type() { + assert!(parse(r#"return { target = "t.exe", env = "PATH=C:/x" }"#).is_err()); + } + + #[test] + fn rejects_sparse_env_array() { + assert!(parse(r#"return { target = "t.exe", env = { P = { "a", nil, "b" } } }"#).is_err()); + } + + #[test] + fn rejects_mixed_key_env_array() { + assert!(parse(r#"return { target = "t.exe", env = { P = { a = "b" } } }"#).is_err()); + } + + #[test] + fn rejects_invalid_env_key() { + assert!(parse(r#"return { target = "t.exe", env = { ["FOO=1"] = "x" } }"#).is_err()); + assert!(parse(r#"return { target = "t.exe", env = { [""] = "x" } }"#).is_err()); + } + + #[test] + fn rejects_nul_in_env_value() { + assert!(parse(r#"return { target = "t.exe", env = { P = { string.char(0) } } }"#).is_err()); + } + + #[test] + fn rejects_quote_in_env_value() { + // Windows 的 join_paths 对含双引号的路径元素返回错误 + assert!( + parse(r#"return { target = "t.exe", env = { P = { string.char(34) } } }"#).is_err() + ); + } + + #[test] + fn rejects_unsupported_env_value_type() { + assert!(parse(r#"return { target = "t.exe", env = { F = function() end } }"#).is_err()); + } + + #[test] + fn missing_target_is_error() { + assert!(parse(r#"return { args = { "x" } }"#).is_err()); + } + + #[test] + fn rejects_non_string_target() { + assert!(parse(r#"return { target = 123 }"#).is_err()); + assert!(parse(r#"return { target = false }"#).is_err()); + } + + #[test] + fn rejects_non_string_args_element() { + assert!(parse(r#"return { target = "t.exe", args = { 1 } }"#).is_err()); + } + + #[test] + fn rejects_sparse_args() { + assert!(parse(r#"return { target = "t.exe", args = { "a", nil, "b" } }"#).is_err()); + } + + #[test] + fn rejects_non_table_args() { + assert!(parse(r#"return { target = "t.exe", args = "-B" }"#).is_err()); + } +} diff --git a/src/env.rs b/src/env.rs deleted file mode 100644 index 579fc2c..0000000 --- a/src/env.rs +++ /dev/null @@ -1,57 +0,0 @@ -use crate::error::ShimError; -use std::path::PathBuf; -pub struct ShimEnv { - pub bin_dir: PathBuf, - pub root_dir: PathBuf, - pub tools_dir: PathBuf, - pub target_name: String, -} - -impl ShimEnv { - /// 提取当前代理程序的运行环境信息 - pub fn new(current_exe: PathBuf) -> Result { - let bin_dir = current_exe - .parent() - .ok_or_else(|| { - ShimError::PathResolutionError(format!( - "无法获取程序 [{}] 的父级 bin 目录", - current_exe.display() - )) - })? - .to_path_buf(); - println!("bin_dir目录 {}", bin_dir.display()); - - let root_dir = bin_dir - .parent() - .ok_or_else(|| { - ShimError::PathResolutionError(format!( - "无法获取 bin 目录 [{}] 的父级 root 目录", - bin_dir.display() - )) - })? - .to_path_buf(); - println!("root_dir 目录 {}", root_dir.display()); - - let tools_dir = root_dir.join("tools"); - println!("tools_dir 目录 {}", tools_dir.display()); - - let target_name = current_exe - .file_stem() - .and_then(|s| s.to_str()) - .ok_or_else(|| { - ShimError::PathResolutionError(format!( - "无法从路径 [{}] 提取有效的程序名称", - current_exe.display() - )) - })? - .to_lowercase(); - println!("程序名称 {}", target_name); - - Ok(Self { - bin_dir, - root_dir, - tools_dir, - target_name, - }) - } -} diff --git a/src/error.rs b/src/error.rs index 10f92af..16cb5dc 100644 --- a/src/error.rs +++ b/src/error.rs @@ -3,16 +3,16 @@ use thiserror::Error; // 推荐引入 thiserror 库,若不使用可手动实 #[derive(Debug, Error)] pub enum ShimError { #[error("路径解析失败: {0}")] - PathResolutionError(String), + PathResolution(String), #[error("获取环境信息失败: {0}")] - EnvError(String), + Environment(String), #[error("配置文件未找到: {0}")] - ConfigNotFound(String), + ConfigMissing(String), #[error("Lua 运行时/语法错误 [{file}]: {source}")] - LuaExecutionError { + LuaExecution { file: String, #[source] source: mlua::Error, diff --git a/src/layout.rs b/src/layout.rs new file mode 100644 index 0000000..aa41040 --- /dev/null +++ b/src/layout.rs @@ -0,0 +1,45 @@ +use crate::error::ShimError; +use std::path::{Path, PathBuf}; + +pub struct ShimLayout { + pub root_dir: PathBuf, + pub bin_dir: PathBuf, + pub tools_dir: PathBuf, +} + +impl ShimLayout { + /// 从当前可执行文件解析 shim 安装目录布局 + pub fn from_executable(exe_path: impl AsRef) -> Result { + let exe_path = exe_path.as_ref(); + let bin_dir = exe_path + .parent() + .ok_or_else(|| { + ShimError::PathResolution(format!( + "无法获取程序 [{}] 的父级 bin 目录", + exe_path.display() + )) + })? + .to_path_buf(); + eprintln!("bin_dir目录 {}", bin_dir.display()); + + let root_dir = bin_dir + .parent() + .ok_or_else(|| { + ShimError::PathResolution(format!( + "无法获取 bin 目录 [{}] 的父级 root 目录", + bin_dir.display() + )) + })? + .to_path_buf(); + eprintln!("root_dir 目录 {}", root_dir.display()); + + let tools_dir = root_dir.join("tools"); + eprintln!("tools_dir 目录 {}", tools_dir.display()); + + Ok(Self { + root_dir, + bin_dir, + tools_dir, + }) + } +} diff --git a/src/lib.rs b/src/lib.rs index b91bb3c..11507be 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -1,11 +1,12 @@ mod config; -mod env; mod error; +mod layout; +mod logger; mod runtime; mod shim; pub use config::ShimConfig; -pub use env::ShimEnv; pub use error::ShimError; +pub use layout::ShimLayout; pub use runtime::LuaRuntime; pub use shim::Shim; diff --git a/src/logger.rs b/src/logger.rs new file mode 100644 index 0000000..3abeac0 --- /dev/null +++ b/src/logger.rs @@ -0,0 +1,39 @@ +use std::path::Path; +use tracing_appender::non_blocking::WorkerGuard; +use tracing_subscriber::{EnvFilter, fmt}; + +/// 初始化日志系统,返回的 `_guard` 必须在 main 作用域内保持存活直到程序退出 +pub fn init_file_logger(log_dir: impl AsRef) -> Option { + // 允许通过环境变量动态控制日志级别,如 SHIM_LOG=debug,默认 debug 或 info + let filter = EnvFilter::try_from_env("SHIM_LOG").unwrap_or_else(|_| EnvFilter::new("debug")); + + // 1. 创建按天滚动的日志追加器 (每天生成类似 shim.2026-08-14.log) + let file_appender = tracing_appender::rolling::daily(log_dir, "shim.log"); + + // 2. 包装为非阻塞后台写入(不会拖慢主程序的启动与执行速度) + let (non_blocking, guard) = tracing_appender::non_blocking(file_appender); + + // 3. 构建 Subscriber,只输出到文件,不输出到控制台 + tracing_subscriber::fmt() + .with_env_filter(filter) + .with_writer(non_blocking) // 写入文件 + .with_ansi(false) // 关闭终端彩色转义字符 + .with_target(false) // 隐藏模块前缀(可选) + .init(); + + Some(guard) +} + +// 调用 +// fn main() -> Result<(), Box> { +// // 假设日志存放在安装根目录下的 logs 文件夹 +// // 也可以先快速推导 layout 拿到 log_dir +// let log_dir = "path/to/root_dir/logs"; +// let _guard = init_file_logger(log_dir); +// +// // 此处写你的 Shim 业务逻辑 +// // 业务代码中所有的 debug!/info!/warn! 都会静默写入文件,控制台干干净净 +// let config = Shim::load()?; +// +// Ok(()) +// } diff --git a/src/main.rs b/src/main.rs index e72e009..d17e8e1 100644 --- a/src/main.rs +++ b/src/main.rs @@ -1,36 +1,36 @@ -use std::{ - env, - ffi::CString, - mem::size_of, - path::Path, - process::{Command, exit}, - ptr::null_mut, -}; +use std::{env, ffi::CString, mem::size_of, path::Path, process::exit, ptr::null_mut}; use rshim::Shim; +use tracing_subscriber::{EnvFilter, fmt}; -use winapi::{ - shared::minwindef::{BOOL, DWORD, FALSE, TRUE}, - um::{ - combaseapi::CoInitializeEx, - consoleapi, - objbase::{COINIT_APARTMENTTHREADED, COINIT_DISABLE_OLE1DDE}, - processthreadsapi::GetExitCodeProcess, - shellapi::{SEE_MASK_NOASYNC, SEE_MASK_NOCLOSEPROCESS, SHELLEXECUTEINFOA, ShellExecuteExA}, - synchapi::WaitForSingleObject, - winbase::INFINITE, - wincon, - winuser::SW_NORMAL, +use windows_sys::{ + Win32::{ + Foundation::{FALSE, TRUE}, + System::{ + Com::{COINIT_APARTMENTTHREADED, COINIT_DISABLE_OLE1DDE, CoInitializeEx}, + Console::{ + CTRL_BREAK_EVENT, CTRL_C_EVENT, CTRL_CLOSE_EVENT, CTRL_LOGOFF_EVENT, + CTRL_SHUTDOWN_EVENT, SetConsoleCtrlHandler, + }, + Threading::{GetExitCodeProcess, INFINITE, WaitForSingleObject}, + }, + UI::{ + Shell::{ + SEE_MASK_NOASYNC, SEE_MASK_NOCLOSEPROCESS, SHELLEXECUTEINFOA, ShellExecuteExA, + }, + WindowsAndMessaging::SW_NORMAL, + }, }, + core::BOOL, }; -unsafe extern "system" fn routine_handler(evt: DWORD) -> BOOL { +unsafe extern "system" fn console_ctrl_handler(evt: u32) -> BOOL { match evt { - wincon::CTRL_C_EVENT => TRUE, //eprintln!("ctrl_c handled!"), - wincon::CTRL_BREAK_EVENT => TRUE, //eprintln!("ctrl_break handled!"), - wincon::CTRL_CLOSE_EVENT => TRUE, //eprintln!("ctrl_close handled!"), - wincon::CTRL_LOGOFF_EVENT => TRUE, //eprintln!("ctrl_logoff handled!"), - wincon::CTRL_SHUTDOWN_EVENT => TRUE, //eprintln!("ctrl_shutdown handled!"), + CTRL_C_EVENT => TRUE, //eprintln!("ctrl_c handled!"), + CTRL_BREAK_EVENT => TRUE, //eprintln!("ctrl_break handled!"), + CTRL_CLOSE_EVENT => TRUE, //eprintln!("ctrl_close handled!"), + CTRL_LOGOFF_EVENT => TRUE, //eprintln!("ctrl_logoff handled!"), + CTRL_SHUTDOWN_EVENT => TRUE, //eprintln!("ctrl_shutdown handled!"), other => { eprintln!("未知的系统事件编号: {},未处理!", other); FALSE @@ -45,13 +45,21 @@ const EXIT_PROG_TERMINATED: i32 = 4; const ERROR_ELEVATION_REQUIRED: i32 = 740; fn main() { - let res: BOOL = unsafe { consoleapi::SetConsoleCtrlHandler(Some(routine_handler), TRUE) }; + // 初始化日志:输出到 stderr,避免污染 shim 子进程的 stdout + fmt() + .with_writer(std::io::stderr) + .with_env_filter( + EnvFilter::try_from_env("SHIM_LOG").unwrap_or_else(|_| EnvFilter::new("warn")), + ) + .init(); + + let res: BOOL = unsafe { SetConsoleCtrlHandler(Some(console_ctrl_handler), TRUE) }; if res == FALSE { eprintln!("警告: 注册控制台中断事件处理器失败。"); } let calling_args: Vec<_> = env::args().skip(1).collect(); - let shim = match Shim::init() { + let shim = match Shim::load() { Ok(v) => v, Err(e) => { eprintln!("加载代理(shim)配置时发生错误: {}", e); @@ -59,31 +67,22 @@ fn main() { } }; - let args = if let Some(mut shim_args) = shim.args { - shim_args.extend_from_slice(calling_args.as_slice()); - shim_args - } else { - calling_args - }; - // ======= 【修改位置 1:构建 Command 并注入环境变量】 ======= - let mut cmd_builder = Command::new(&shim.target_path); - cmd_builder.args(&args); + // 构建 Command:复用 ShimConfig::to_command(含 args/env 注入),避免重复逻辑 + let mut cmd = shim.to_command(&calling_args); - // 仅作用于目标子进程,完全 Safe 且隔离 - if let Some(ref envs) = shim.envs { - cmd_builder.envs(envs); - } - let mut cmd = match cmd_builder.spawn() { + // 提权回退时需要完整参数:配置默认参数 + 调用方透传参数 + let mut args = shim.args.clone(); + args.extend_from_slice(&calling_args); + + let mut cmd = match cmd.spawn() { Ok(v) => v, - Err(e) if e.raw_os_error() == Some(ERROR_ELEVATION_REQUIRED) => exit(execute_elevated( - &shim.target_path, - &args, - shim.envs.as_ref(), - )), + Err(e) if e.raw_os_error() == Some(ERROR_ELEVATION_REQUIRED) => { + exit(execute_elevated(&shim.target, &args, Some(&shim.env))) + } Err(e) => { eprintln!( "启动目标程序 [{}] 时发生错误: {}", - shim.target_path.to_string_lossy(), + shim.target.to_string_lossy(), e ); exit(EXIT_FAILED_SPAWN_PROG); @@ -94,7 +93,7 @@ fn main() { Err(e) => { eprintln!( "等待目标程序 [{}] 执行完毕时发生错误: {}", - shim.target_path.to_string_lossy(), + shim.target.to_string_lossy(), e ); exit(EXIT_FAILED_WAIT_PROG); @@ -106,10 +105,10 @@ fn main() { fn execute_elevated( program: &Path, args: &[String], - envs: Option<&std::collections::HashMap>, + env_vars: Option<&std::collections::HashMap>, ) -> i32 { // 若提权启动,在此处将环境变量设置给当前进程(即将弹窗 UAC 的进程,随后会被子进程继承) - if let Some(env_map) = envs { + if let Some(env_map) = env_vars { for (k, v) in env_map { unsafe { env::set_var(k, v); @@ -119,45 +118,45 @@ fn execute_elevated( let runas = CString::new("runas").unwrap(); let program = CString::new(program.to_str().unwrap()).unwrap(); - let mut params = String::new(); + let mut arguments = String::new(); for arg in args.iter() { - params.push(' '); + arguments.push(' '); if arg.len() == 0 { - params.push_str("\"\""); + arguments.push_str("\"\""); } else if arg.find(&[' ', '\t', '"'][..]).is_none() { - params.push_str(&arg); + arguments.push_str(&arg); } else { - params.push('"'); + arguments.push('"'); for c in arg.chars() { match c { - '\\' => params.push_str("\\\\"), - '"' => params.push_str("\\\""), - c => params.push(c), + '\\' => arguments.push_str("\\\\"), + '"' => arguments.push_str("\\\""), + c => arguments.push(c), } } - params.push('"'); + arguments.push('"'); } } - let params = CString::new(¶ms[..]).unwrap(); + let arguments = CString::new(&arguments[..]).unwrap(); let mut info = SHELLEXECUTEINFOA::default(); - info.cbSize = size_of::() as DWORD; + info.cbSize = size_of::() as u32; info.fMask = SEE_MASK_NOASYNC | SEE_MASK_NOCLOSEPROCESS; - info.lpVerb = runas.as_ptr(); - info.lpFile = program.as_ptr(); - info.lpParameters = params.as_ptr(); + info.lpVerb = runas.as_ptr().cast::(); + info.lpFile = program.as_ptr().cast::(); + info.lpParameters = arguments.as_ptr().cast::(); info.nShow = SW_NORMAL; let res = unsafe { CoInitializeEx( null_mut(), - COINIT_APARTMENTTHREADED | COINIT_DISABLE_OLE1DDE, + (COINIT_APARTMENTTHREADED | COINIT_DISABLE_OLE1DDE) as u32, ); ShellExecuteExA(&mut info as *mut _) }; if res == FALSE || info.hProcess == null_mut() { return EXIT_FAILED_SPAWN_PROG; } - let mut code: DWORD = 0; + let mut code: u32 = 0; unsafe { WaitForSingleObject(info.hProcess, INFINITE); if GetExitCodeProcess(info.hProcess, &mut code as *mut _) == FALSE { diff --git a/src/runtime.rs b/src/runtime.rs index 2c33d40..651fb10 100644 --- a/src/runtime.rs +++ b/src/runtime.rs @@ -1,18 +1,13 @@ use crate::error::ShimError; -use crate::{ShimConfig, ShimEnv}; +use crate::{ShimConfig, ShimLayout}; use mlua::{FromLua, Lua, StdLib, Table, Value}; use std::path::Path; use std::{env, fs}; /// 将 Path 转换为适合 Lua 使用的安全字符串路径 fn normalize_path_for_lua(path: &Path) -> String { - let path_str = path.to_string_lossy(); - - // 1. 剥离 Windows UNC 规范路径前缀 (\\?\) - let clean_str = path_str.strip_prefix(r"\\?\").unwrap_or(&path_str); - - // 2. 将反斜杠转换成正斜杠(在非 UNC 路径下,Windows 和 Lua 均完美支持 /) - // 这样既避免了 Lua 字符串转义隐患,又不会破坏 Windows 路径 - clean_str.replace('\\', "/") + // 自动将 Windows UNC 规范路径转回传统路径 + let simplified = dunce::simplified(path); + simplified.to_string_lossy().replace('\\', "/") } pub struct LuaRuntime { @@ -21,66 +16,147 @@ pub struct LuaRuntime { impl LuaRuntime { /// 初始化限定权限的 Lua 沙箱环境 - pub fn new(shim_env: &ShimEnv) -> Result { + pub fn new(layout: &ShimLayout) -> Result { // 只加载安全的标准库,剥离 os / io 等风险模块 let lua = Lua::new_with( StdLib::TABLE | StdLib::STRING | StdLib::MATH | StdLib::PACKAGE, mlua::LuaOptions::default(), ) - .map_err(|e| ShimError::EnvError(format!("初始化 Lua 失败: {}", e)))?; + .map_err(|e| ShimError::Environment(format!("初始化 Lua 失败: {}", e)))?; let globals = lua.globals(); // 统一使用 POSIX 风格路径规范化路径字符串 - let root_dir_str = normalize_path_for_lua(&shim_env.root_dir); - let tools_dir_str = normalize_path_for_lua(&shim_env.tools_dir); + let root_dir = normalize_path_for_lua(&layout.root_dir); + let tools_dir = normalize_path_for_lua(&layout.tools_dir); - // 1. 注入锚点变量 + // 1. 注入锚点变量 __SHIM_DIR__(shim 安装根目录) globals - .set("__SHIM_DIR__", root_dir_str.clone()) - .map_err(|e| ShimError::EnvError(e.to_string()))?; + .set("__SHIM_DIR__", root_dir.clone()) + .map_err(|e| ShimError::Environment(e.to_string()))?; // 2. 安全暴露 get_env 供配置读取环境变量 let get_env = lua .create_function(|_, key: String| -> mlua::Result { Ok(env::var(key).unwrap_or_default()) }) - .map_err(|e| ShimError::EnvError(e.to_string()))?; + .map_err(|e| ShimError::Environment(e.to_string()))?; globals .set("get_env", get_env) - .map_err(|e| ShimError::EnvError(e.to_string()))?; + .map_err(|e| ShimError::Environment(e.to_string()))?; // 3. 配置 package.path,确保 require 行为正常 if let Ok(package) = globals.get::
("package") { + let _ = package.set("cpath", ""); + let _ = package.set("loadlib", Value::Nil); + if let Ok(path) = package.get::("path") { let new_path = format!( "{};{}/?.lua;{}/?/init.lua;{}/?.lua;{}/?/init.lua", - path, root_dir_str, root_dir_str, tools_dir_str, tools_dir_str + path, root_dir, root_dir, tools_dir, tools_dir ); let _ = package.set("path", new_path); } } + // 4. 包装 require:配置模块缺失/加载失败时记录日志并跳过该条目, + // 而不是让整个 shims.lua 解析失败(排查问题时日志可见) + let original_require: mlua::Function = globals + .get("require") + .map_err(|e| ShimError::Environment(format!("获取 require 失败: {}", e)))?; + globals + .set("_rshim_original_require", &original_require) + .map_err(|e| ShimError::Environment(e.to_string()))?; + + 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()) { + Ok(value) => Ok(value), + Err(e) => { + tracing::warn!( + module = %module, + error = %e, + "配置模块加载失败,已跳过该条目(可在独立配置文件中定义)" + ); + Ok(Value::Nil) + } + } + }) + .map_err(|e| ShimError::Environment(e.to_string()))?; + + globals + .set("require", wrapped_require) + .map_err(|e| ShimError::Environment(e.to_string()))?; + Ok(Self { lua }) } /// 执行指定脚本文件,直接返回完整的 Lua Table - pub fn evaluate_lua_script(&self, path: &Path) -> Result { - let code = fs::read_to_string(path)?; + pub fn eval_script(&self, path: impl AsRef) -> Result { + let path = path.as_ref(); + + let bytes = fs::read(path)?; + let code = String::from_utf8(bytes).map_err(|e| { + ShimError::InvalidConfig(format!( + "{} 不是有效的 UTF-8 文件(请将 Lua 配置文件保存为 UTF-8 编码): {}", + path.display(), + e.utf8_error() + )) + })?; + + let chunk_name = format!("@{}", path.display()); self.lua .load(&code) - .set_name(path.to_string_lossy()) - .eval::
() - .map_err(|e| ShimError::LuaExecutionError { + .set_name(&chunk_name) + .eval::() + .map_err(|e| ShimError::LuaExecution { file: path.display().to_string(), source: e, }) } - /// 将 Lua Value 解析转化为 ShimConfig 数据对象 - pub fn parse_config(&self, value: Value) -> Result { - ShimConfig::from_lua(value, &self.lua).map_err(|e| ShimError::InvalidConfig(e.to_string())) + // /// 将 Lua Value 解析转化为 ShimConfig 数据对象 + // pub fn parse_config(&self, value: Value) -> Result { + // ShimConfig::from_lua(value, &self.lua).map_err(|e| ShimError::InvalidConfig(e.to_string())) + // } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::ShimLayout; + + fn test_layout() -> ShimLayout { + let root = std::env::temp_dir().join("rshim-test-layout"); + ShimLayout { + root_dir: root.clone(), + bin_dir: root.join("bin"), + tools_dir: root.join("tools"), + } + } + + #[test] + fn missing_module_require_returns_nil() { + let runtime = LuaRuntime::new(&test_layout()).unwrap(); + let value: Value = runtime + .lua + .load(r#"return require("rshim_test_no_such_module")"#) + .eval() + .unwrap(); + assert!(matches!(value, Value::Nil)); + } + + #[test] + fn builtin_module_require_still_works() { + let runtime = LuaRuntime::new(&test_layout()).unwrap(); + let value: Value = runtime + .lua + .load(r#"return pcall(require, "string")"#) + .eval() + .unwrap(); + assert!(matches!(value, Value::Boolean(true))); } } diff --git a/src/shim.rs b/src/shim.rs index ac81ae3..989a591 100644 --- a/src/shim.rs +++ b/src/shim.rs @@ -1,57 +1,100 @@ use crate::ShimError; -use mlua::Value; +use crate::{LuaRuntime, ShimConfig, ShimLayout}; +use mlua::{Table, Value}; use std::{ env, io::{Error, ErrorKind}, }; -use crate::{ShimConfig, ShimEnv, LuaRuntime}; - +use tracing::{debug, trace, warn}; pub struct Shim; impl Shim { - pub fn init() -> Result { + pub fn load() -> Result { let current_exe = env::current_exe() .map_err(|e| Error::new(ErrorKind::Other, format!("获取代理程序路径失败: {}", e)))?; - println!("当前目录 {}", current_exe.display()); - let shim_env = ShimEnv::new(current_exe)?; - let runtime = LuaRuntime::new(&shim_env)?; + debug!("当前目录 {}", current_exe.display()); - Self::resolve_config(&runtime, &shim_env) + let target_name = current_exe + .file_stem() + .and_then(|s| s.to_str()) + .ok_or_else(|| { + ShimError::PathResolution(format!( + "无法从路径 [{}] 提取有效的程序名称", + current_exe.display() + )) + })? + .to_lowercase(); + debug!( + target_name = %target_name, + current_exe = %current_exe.display(), + "开始加载 Shim 配置" + ); + + let layout = ShimLayout::from_executable(current_exe)?; + debug!( + root_dir = %layout.root_dir.display(), + bin_dir = %layout.bin_dir.display(), + tools_dir = %layout.tools_dir.display(), + "Shim 目录布局解析完成" + ); + + let runtime = LuaRuntime::new(&layout)?; + + Self::resolve_config(&runtime, &layout, &target_name) } - fn resolve_config(runtime: &LuaRuntime, env: &ShimEnv) -> Result { + fn resolve_config( + runtime: &LuaRuntime, + paths: &ShimLayout, + target_name: &str, + ) -> Result { // 策略 1: 尝试加载全局配置文件 shims.lua - let global_config = env.root_dir.join("shims.lua"); + let global_config = paths.root_dir.join("shims.lua"); if global_config.is_file() { - let root_table = runtime.evaluate_lua_script(&global_config)?; + trace!(path = %global_config.display(), "发现全局配置文件,尝试解析"); + let root_table: Table = runtime.eval_script(&global_config)?; // 检查 shims.lua 中是否存在以 target_name 命名的 Table 节点 - if let Ok(target_val) = root_table.get::(env.target_name.as_str()) { - if matches!(target_val, Value::Table(_)) { - return runtime.parse_config(target_val); - } + if root_table + .contains_key(target_name) + .map_err(|e| ShimError::InvalidConfig(format!("检查全局配置失败: {}", e)))? + { + let target_val: ShimConfig = root_table.get(target_name).map_err(|e| { + ShimError::InvalidConfig(format!("解析配置 [{}] 失败: {}", target_name, e)) + })?; + debug!( + target = %target_name, + source = %global_config.display(), + "成功从全局配置文件中匹配到目标工具" + ); + + return Ok(target_val); } // 穿透:若全局配置文件存在但未包含当前程序的 key,继续向下探查 + trace!(target = %target_name, "全局配置文件中未包含该目标,继续探查独立配置"); } // 策略 2: 降级寻找独立文件 ({exe}.lua),优先顺序:tools/ > root/ - let target_filename = format!("{}.lua", env.target_name); + let target_filename = format!("{}.lua", target_name); let candidates = [ - env.tools_dir.join(&target_filename), - env.root_dir.join(&target_filename), + paths.tools_dir.join(&target_filename), + paths.root_dir.join(&target_filename), ]; for config_path in &candidates { if config_path.is_file() { - let table = runtime.evaluate_lua_script(config_path)?; - return runtime.parse_config(Value::Table(table)); + debug!(path = %config_path.display(), "找到独立配置文件,开始加载"); + // 直接泛型反序列化为 ShimConfig + return runtime.eval_script::(config_path); } + trace!(path = %config_path.display(), "独立配置文件不存在,跳过"); } // 策略 3: 所有查找失败,抛出错误 - Err(ShimError::ConfigNotFound(format!( + warn!(target = %target_name, "未找到任何匹配的配置文件"); + Err(ShimError::ConfigMissing(format!( "未找到关于 '{}' 的配置。请检查 shims.lua 或特定的 {}.lua 文件", - env.target_name, env.target_name + target_name, target_name ))) } }