use crate::{ShimConfig, Layout}; use anyhow::{Context, Result, anyhow, bail}; use mlua::{FromLua, Lua, StdLib, Table, Value}; use std::ffi::OsStr; use std::path::Path; use std::{env, fs}; /// 将 Path 转换为适合 Lua 使用的安全字符串路径 fn normalize_path_for_lua(path: &Path) -> String { // 自动将 Windows UNC 规范路径转回传统路径 let simplified = dunce::simplified(path); simplified.to_string_lossy().replace('\\', "/") } pub struct LuaRuntime { lua: Lua, } impl LuaRuntime { /// 初始化限定权限的 Lua 沙箱环境 pub fn new(layout: &Layout) -> Result { // 只加载安全的标准库,剥离 os / io 等风险模块 let lua = Lua::new_with( StdLib::TABLE | StdLib::STRING | StdLib::MATH | StdLib::PACKAGE, mlua::LuaOptions::default(), ) .context("初始化 Lua 失败")?; let globals = lua.globals(); // 统一使用 POSIX 风格路径规范化路径字符串 let base_dir = normalize_path_for_lua(&layout.base_dir); let tools_dir = normalize_path_for_lua(&layout.tools_dir); // 1. 注入锚点变量 __SHIM_DIR__(shim 安装根目录) globals .set("__SHIM_DIR__", base_dir.clone()) .context("设置 __SHIM_DIR__ 环境变量失败")?; // 2. 安全暴露 get_env 供配置读取环境变量 // 返回按平台路径分隔符拆分后的段数组(自动剥离引号包裹), // 便于 PATH 等列表变量直接嵌入数组:PATH = { prefix, get_env("PATH") } let get_env = lua .create_function(|lua, key: String| -> mlua::Result { // 缺失变量视为空字符串,拆分后得到空表(不贡献任何路径段) let value = env::var_os(key).unwrap_or_default(); // 空输入返回空表;否则按平台分隔符拆分(split_paths 会剥离引号包裹) if value.is_empty() { return lua.create_table(); } let table = lua.create_table()?; for (i, p) in env::split_paths(&value).enumerate() { // 3. 跨平台提取原始字节并转为 LuaString,保证 100% 无损 #[cfg(unix)] let lua_str = { use std::os::unix::ffi::OsStrExt; lua.create_string(p.as_os_str().as_bytes())? }; #[cfg(windows)] let lua_str = { // Windows 路径是 UTF-16,转成字符串或保持其字节表达 let s = p.to_string_lossy(); lua.create_string(s.as_bytes())? }; table.set(i + 1, lua_str)?; } Ok(table) }) .context("注册 get_env 函数失败")?; globals .set("get_env", get_env) .context("挂载 get_env 全局函数失败")?; // 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, base_dir, base_dir, tools_dir, tools_dir ); let _ = package.set("path", new_path); } } // 4. 包装 require:配置模块缺失/加载失败时记录日志并跳过该条目, // 而不是让整个 mimic.lua 解析失败(排查问题时日志可见) let original_require: mlua::Function = globals .get("require") .context("获取内置 require 函数失败")?; globals .set("_rshim_original_require", &original_require) .context("备份原始 require 函数失败")?; let wrapped_require = lua .create_function(|lua, module: String| -> mlua::Result { let original: mlua::Function = lua.globals().get("_rshim_original_require")?; match original.call::(module.clone()) { Ok(value) => Ok(value), Err(e) => { tracing::warn!( module = %module, error = %e, "配置模块加载失败,已跳过该条目(可在独立配置文件中定义)" ); Ok(Value::Nil) } } }) .context("创建包装版 require 函数失败")?; globals .set("require", wrapped_require) .context("重载 require 函数失败")?; Ok(Self { lua }) } /// 执行指定脚本文件,直接返回完整的 Lua Table pub fn eval_script(&self, path: impl AsRef) -> Result { let path = path.as_ref(); // println!("path {:?}", path); let bytes = fs::read(path).with_context(|| format!("无法读取配置文件: {}", path.display()))?; let code = String::from_utf8(bytes).with_context(|| { format!( "{} 不是有效的 UTF-8 文件(请将 Lua 配置文件保存为 UTF-8 编码)", path.display() ) })?; println!("code {:?}", code); let chunk_name = format!("@{}", path.display()); self.lua .load(&code) .set_name(&chunk_name) .eval::() // .map_err(|e| anyhow!(e.to_string())) .with_context(|| format!("执行 Lua 配置文件失败: {}", path.display())) } } #[cfg(test)] mod tests { use super::*; use crate::Layout; fn test_layout() -> Layout { let root = std::env::temp_dir().join("rshim-test-layout"); Layout { base_dir: root.clone(), bin_dir: root.join("bin"), tools_dir: root.join("tools"), lua_file: Default::default(), } } #[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))); } #[test] fn get_env_returns_split_table() { let runtime = LuaRuntime::new(&test_layout()).unwrap(); let value: Value = runtime .lua .load(r#"return get_env("PATH")"#) .eval() .unwrap(); let table = match value { Value::Table(t) => t, other => panic!("expected table, got {}", other.type_name()), }; assert!( table.raw_len() >= 1, "PATH should have at least one segment" ); // split_paths 会剥离引号包裹,拆分段不应再含双引号 for i in 1..=table.raw_len() { let seg: String = table.raw_get(i).unwrap(); assert!( !seg.contains('"'), "segment should not contain quote: {seg:?}" ); } } #[test] fn get_env_missing_returns_empty_table() { let runtime = LuaRuntime::new(&test_layout()).unwrap(); let value: Value = runtime .lua .load(r#"return get_env("RSHIM_TEST_NO_SUCH_VAR_12345")"#) .eval() .unwrap(); match value { Value::Table(t) => assert_eq!(t.raw_len(), 0), other => panic!("expected table, got {}", other.type_name()), } } #[test] fn path_with_get_env_joins_without_quote_error() { // 复现用户场景:PATH = { base_dir .. "/tools/numa", get_env("PATH") } // 宿主 PATH 即使含双引号,也应拆分后正常拼接,而不是报错 let runtime = LuaRuntime::new(&test_layout()).unwrap(); let cfg: ShimConfig = runtime .lua .load( r#" return { target = __SHIM_DIR__ .. "/tools/numa/numa.exe", args = { "--help" }, env = { PATH = { __SHIM_DIR__ .. "/tools/numa", get_env("PATH") } } } "#, ) .eval() .unwrap(); let path = cfg.env.get("PATH").unwrap().to_str().unwrap(); let prefix = std::env::temp_dir() .join("rshim-test-layout") .to_string_lossy() .replace('\\', "/") + "/tools/numa;"; assert!(path.starts_with(&prefix), "unexpected PATH: {path}"); // 宿主 PATH 的段应被附加在配置前缀之后 let host = std::env::var("PATH").unwrap_or_default(); if !host.is_empty() { let host_first = std::env::split_paths(&host) .next() .unwrap() .to_string_lossy() .into_owned(); assert!( path.contains(&host_first), "missing host PATH segment: {host_first}" ); } // 宿主 PATH 里引号包裹的畸形段(如 "D:\\...\\bin;")应被原样保留, // 而不是让整个配置加载失败 if host.contains('"') { assert!( path.contains('"'), "quoted host segments should be preserved" ); } } }