refactor: 迁移 winapi 到 windows-sys,修复配置解析漏洞

- 依赖替换为 windows-sys 0.61,main.rs 全面适配新 API
  - 配置解析错误显式传播:env/args 类型错误、数组空洞、非法键不再静默
  吞错
  - 修正空环境变量与空参数语义,补充 UTF-8 与路径拼接校验
  - require 容错移至 Rust 侧,模块加载失败记录日志并跳过
  - 新增配置解析与运行时单元测试(19 个)
This commit is contained in:
2026-08-14 20:09:34 +08:00
parent ed5439eaa1
commit 5e69a6a980
13 changed files with 691 additions and 278 deletions

1
.gitignore vendored
View File

@@ -1,4 +1,5 @@
/target
/.vscode
/.idea
*.exe
./Cargo.lock

View File

@@ -1,32 +1,31 @@
[package]
name = "rshim"
version = "0.1.0"
authors = ["anonymous <anonymous@example.com>"]
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"

View File

@@ -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);
}
}

View File

@@ -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,12 +46,13 @@ 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"),

View File

@@ -1,30 +1,136 @@
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<String>) -> 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::<Value, Value>() {
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<String>) -> 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<Vec<String>>,
pub envs: Option<HashMap<String, String>>,
pub target: PathBuf,
pub args: Vec<String>,
pub env: HashMap<String, String>,
}
impl ShimConfig {
/// 根据配置快速构建准备执行的 Command 对象
pub fn to_command(&self) -> Command {
let mut cmd = Command::new(&self.target_path);
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);
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 {
// 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<Self> {
match value {
Value::Table(table) => {
let path_str: String = table.get("path")?;
let args: Option<Vec<String>> = table.get("args")?;
// 解析 env Table
let mut envs_map = HashMap::new();
if let Ok(env_table) = table.get::<Table>("env") {
// 获取当前系统的路径分隔符Windows 为 ";"Unix 为 ":"
#[cfg(windows)]
let sep = ";";
#[cfg(not(windows))]
let sep = ":";
for pair in env_table.pairs::<String, Value>() {
let (k, v) = pair?;
match v {
// 情况 1: 普通字符串,如 HOME = "C:/path" -> 直接覆盖
Value::String(s) => {
envs_map.insert(k, s.to_str()?.to_string());
let table = match value {
Value::Table(t) => t,
_ => {
return Err(conversion_error(format!(
"期望得到一个 Lua Table 配置对象,实际是 {}",
value.type_name()
)));
}
// 情况 2: 数组 Table如 PATH = { bin_dir, get_env("PATH") }
Value::Table(arr) => {
let paths: Vec<String> = arr
// 将 arr 作为序列(数组)处理,每个元素转为 String
.sequence_values::<String>()
.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)
};
// 必填字段: target严格限定为字符串避免数字被 mlua 宽松转为字符串后掩盖错误)
let target = match table.get::<Value>("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::<Option<Value>>("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<Table> 转换报错,不再静默忽略)
let env_table: Option<Table> = 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::<String, Value>() {
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);
}
}
println!("环境变量结果:{:?}", env);
Ok(ShimConfig {
target_path: PathBuf::from(path_str),
target: PathBuf::from(target),
args,
envs,
env,
})
}
_ => Err(mlua::Error::FromLuaConversionError {
from: value.type_name(),
to: "ShimConfig".into(),
message: Some("Expected a Lua table".to_string()),
}),
}
#[cfg(test)]
mod tests {
use super::*;
fn parse(src: &str) -> mlua::Result<ShimConfig> {
let lua = Lua::new();
let value = lua.load(src).eval::<Value>()?;
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());
}
}

View File

@@ -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<Self, ShimError> {
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,
})
}
}

View File

@@ -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,

45
src/layout.rs Normal file
View File

@@ -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<Path>) -> Result<Self, ShimError> {
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,
})
}
}

View File

@@ -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;

39
src/logger.rs Normal file
View File

@@ -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<Path>) -> Option<WorkerGuard> {
// 允许通过环境变量动态控制日志级别,如 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<dyn std::error::Error>> {
// // 假设日志存放在安装根目录下的 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(())
// }

View File

@@ -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<String, String>>,
env_vars: Option<&std::collections::HashMap<String, String>>,
) -> 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(&params[..]).unwrap();
let arguments = CString::new(&arguments[..]).unwrap();
let mut info = SHELLEXECUTEINFOA::default();
info.cbSize = size_of::<SHELLEXECUTEINFOA>() as DWORD;
info.cbSize = size_of::<SHELLEXECUTEINFOA>() 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::<u8>();
info.lpFile = program.as_ptr().cast::<u8>();
info.lpParameters = arguments.as_ptr().cast::<u8>();
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 {

View File

@@ -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<Self, ShimError> {
pub fn new(layout: &ShimLayout) -> Result<Self, ShimError> {
// 只加载安全的标准库,剥离 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<String> {
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::<Table>("package") {
let _ = package.set("cpath", "");
let _ = package.set("loadlib", Value::Nil);
if let Ok(path) = package.get::<String>("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<Value> {
let original: mlua::Function = lua.globals().get("_rshim_original_require")?;
match original.call::<Value>(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<Table, ShimError> {
let code = fs::read_to_string(path)?;
pub fn eval_script<T: FromLua>(&self, path: impl AsRef<Path>) -> Result<T, ShimError> {
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::<Table>()
.map_err(|e| ShimError::LuaExecutionError {
.set_name(&chunk_name)
.eval::<T>()
.map_err(|e| ShimError::LuaExecution {
file: path.display().to_string(),
source: e,
})
}
/// 将 Lua Value 解析转化为 ShimConfig 数据对象
pub fn parse_config(&self, value: Value) -> Result<ShimConfig, ShimError> {
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, ShimError> {
// 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)));
}
}

View File

@@ -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<ShimConfig, ShimError> {
pub fn load() -> Result<ShimConfig, ShimError> {
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<ShimConfig, ShimError> {
fn resolve_config(
runtime: &LuaRuntime,
paths: &ShimLayout,
target_name: &str,
) -> Result<ShimConfig, ShimError> {
// 策略 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::<Value>(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::<ShimConfig>(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
)))
}
}