refactor(spec,loader): 优化项目结构,
- 提取公共方法 - 优化 Mirror FromLua trait 处理逻辑 - 新增 logger.rs 统一日志处理 - 其他优化
This commit is contained in:
3
Cargo.lock
generated
3
Cargo.lock
generated
@@ -269,11 +269,9 @@ version = "0.1.0"
|
|||||||
dependencies = [
|
dependencies = [
|
||||||
"anyhow",
|
"anyhow",
|
||||||
"clap",
|
"clap",
|
||||||
"clap_derive",
|
|
||||||
"mirror-core",
|
"mirror-core",
|
||||||
"mlua",
|
"mlua",
|
||||||
"tracing",
|
"tracing",
|
||||||
"tracing-subscriber",
|
|
||||||
]
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
@@ -297,6 +295,7 @@ version = "0.1.0"
|
|||||||
dependencies = [
|
dependencies = [
|
||||||
"anyhow",
|
"anyhow",
|
||||||
"mirror-core",
|
"mirror-core",
|
||||||
|
"mlua",
|
||||||
"tracing",
|
"tracing",
|
||||||
"tracing-appender",
|
"tracing-appender",
|
||||||
"tracing-subscriber",
|
"tracing-subscriber",
|
||||||
|
|||||||
@@ -11,8 +11,6 @@ path = "src/main.rs"
|
|||||||
[dependencies]
|
[dependencies]
|
||||||
mirror-core = { path = "../mirror-core" }
|
mirror-core = { path = "../mirror-core" }
|
||||||
clap = { version = "4.6.6", features = ["cargo", "color", "derive","string"] }
|
clap = { version = "4.6.6", features = ["cargo", "color", "derive","string"] }
|
||||||
clap_derive = { version = "4.6.4" }
|
|
||||||
mlua = { workspace = true }
|
mlua = { workspace = true }
|
||||||
anyhow = { workspace = true }
|
anyhow = { workspace = true }
|
||||||
tracing = { workspace = true }
|
tracing = { workspace = true }
|
||||||
tracing-subscriber = { workspace = true }
|
|
||||||
@@ -1,11 +1,16 @@
|
|||||||
use crate::cli;
|
use crate::cli;
|
||||||
use crate::commands::{builtin, dynamic::MrCommand};
|
use crate::commands::{builtin, dynamic::MrCommand};
|
||||||
use anyhow::{Context, Result, bail};
|
use anyhow::{Context, Result, bail};
|
||||||
use mirror_core::{Layout, LuaRuntime};
|
use tracing::debug;
|
||||||
|
use mirror_core::{init_logging_from, Layout, LuaRuntime};
|
||||||
|
|
||||||
pub fn run() -> Result<()> {
|
pub fn run() -> Result<()> {
|
||||||
let exe = std::env::current_exe().context("获取代理程序路径失败")?;
|
let exe = std::env::current_exe().context("获取代理程序路径失败")?;
|
||||||
let layout = Layout::from(&exe)?;
|
let layout = Layout::from(&exe)?;
|
||||||
|
let _guard = init_logging_from(&layout);
|
||||||
|
|
||||||
|
debug!("=== 管理程序启动 ===");
|
||||||
|
|
||||||
let registry_file = layout.base_dir.join("commands").join("commands.lua");
|
let registry_file = layout.base_dir.join("commands").join("commands.lua");
|
||||||
|
|
||||||
let runtime = LuaRuntime::new(&layout)?;
|
let runtime = LuaRuntime::new(&layout)?;
|
||||||
|
|||||||
@@ -1,13 +1,13 @@
|
|||||||
use anyhow::{Context, Result};
|
use anyhow::{Context, Result};
|
||||||
use mirror_core::Layout;
|
use mirror_core::{Layout, LogLevel};
|
||||||
use std::fs;
|
use std::fs;
|
||||||
|
|
||||||
// Rust 原生处理 log 命令逻辑
|
// Rust 原生处理 log 命令逻辑
|
||||||
pub(crate) fn log_handle(layout: &Layout, matches: &clap::ArgMatches) -> Result<()> {
|
pub(crate) fn log_handle(layout: &Layout, matches: &clap::ArgMatches) -> Result<()> {
|
||||||
let log_ini_path = layout.base_dir.join("mirror-log.ini");
|
let log_ini_path = layout.base_dir.join("../../../mirror-core/mirror.ini");
|
||||||
|
|
||||||
// 1. 处理设置日志等级
|
// 1. 处理设置日志等级
|
||||||
if let Some(level) = matches.get_one::<String>("level") {
|
if let Some(&level) = matches.get_one::<LogLevel>("level") {
|
||||||
let content = format!("level = \"{}\"\nlog_dir = \"logs\"\n", level.as_str());
|
let content = format!("level = \"{}\"\nlog_dir = \"logs\"\n", level.as_str());
|
||||||
fs::write(&log_ini_path, content)
|
fs::write(&log_ini_path, content)
|
||||||
.with_context(|| format!("写入日志配置文件失败: {}", log_ini_path.display()))?;
|
.with_context(|| format!("写入日志配置文件失败: {}", log_ini_path.display()))?;
|
||||||
@@ -26,9 +26,9 @@ pub(crate) fn log_handle(layout: &Layout, matches: &clap::ArgMatches) -> Result<
|
|||||||
if log_ini_path.is_file() {
|
if log_ini_path.is_file() {
|
||||||
let current_ini = fs::read_to_string(&log_ini_path)
|
let current_ini = fs::read_to_string(&log_ini_path)
|
||||||
.with_context(|| format!("读取配置文件失败: {}", log_ini_path.display()))?;
|
.with_context(|| format!("读取配置文件失败: {}", log_ini_path.display()))?;
|
||||||
println!("当前 mirror-log.ini 配置:\n{current_ini}");
|
println!("当前 mirror.ini 配置:\n{current_ini}");
|
||||||
} else {
|
} else {
|
||||||
println!("未找到 mirror-log.ini,当前使用默认全局级别: info");
|
println!("未找到 mirror.ini,当前使用默认全局级别: warn");
|
||||||
}
|
}
|
||||||
|
|
||||||
Ok(())
|
Ok(())
|
||||||
|
|||||||
@@ -1,14 +1,14 @@
|
|||||||
pub mod error;
|
pub mod error;
|
||||||
mod layout;
|
mod layout;
|
||||||
mod logger;
|
mod logger;
|
||||||
mod mirror;
|
|
||||||
mod runtime;
|
mod runtime;
|
||||||
mod utils;
|
mod utils;
|
||||||
pub mod validators;
|
mod validators;
|
||||||
|
|
||||||
pub use layout::Layout;
|
pub use layout::Layout;
|
||||||
pub use mirror::Mirror;
|
// pub use mirror_shim::mirror::Mirror;
|
||||||
pub use runtime::LuaRuntime;
|
pub use runtime::LuaRuntime;
|
||||||
|
|
||||||
pub use logger::{init_logging_from, Logger,LogLevel};
|
pub use logger::{init_logging_from, LogLevel, Logger};
|
||||||
pub use utils::{lua_string_2_os_string, normalize_path_for_lua, parse_tokens};
|
// pub use mirror_shim::utils::{lua_string_2_os_string, normalize_path_for_lua, parse_tokens};
|
||||||
|
pub use validators::validate_command_name;
|
||||||
@@ -111,7 +111,7 @@ impl Logger {
|
|||||||
pub fn init_logging_from(layout: &Layout) -> Option<WorkerGuard> {
|
pub fn init_logging_from(layout: &Layout) -> Option<WorkerGuard> {
|
||||||
// 1. 在 layout.base_dir 目录下寻找 log.toml
|
// 1. 在 layout.base_dir 目录下寻找 log.toml
|
||||||
// let config_path = layout.base_dir.join("mirror.toml");
|
// let config_path = layout.base_dir.join("mirror.toml");
|
||||||
let config_path = layout.base_dir.join("../../mirror.ini");
|
let config_path = layout.base_dir.join("mirror.ini");
|
||||||
|
|
||||||
// 读取配置文件(如不存在或解析失败,降级回退到默认设置)
|
// 读取配置文件(如不存在或解析失败,降级回退到默认设置)
|
||||||
let config = if config_path.exists() {
|
let config = if config_path.exists() {
|
||||||
|
|||||||
@@ -204,7 +204,7 @@ impl LuaRuntime {
|
|||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
mod tests {
|
mod tests {
|
||||||
use super::*;
|
use super::*;
|
||||||
use crate::{Layout, Mirror};
|
use crate::{Layout, };
|
||||||
|
|
||||||
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");
|
||||||
@@ -278,56 +278,56 @@ mod tests {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
// #[test]
|
||||||
fn path_with_get_env_joins_without_quote_error() {
|
// fn path_with_get_env_joins_without_quote_error() {
|
||||||
// 复现用户场景:PATH = { base_dir .. "/tools/numa", get_env("PATH") }
|
// // 复现用户场景:PATH = { base_dir .. "/tools/numa", get_env("PATH") }
|
||||||
// 宿主 PATH 即使含双引号,也应拆分后正常拼接,而不是报错
|
// // 宿主 PATH 即使含双引号,也应拆分后正常拼接,而不是报错
|
||||||
let runtime = LuaRuntime::new(&test_layout()).unwrap();
|
// let runtime = LuaRuntime::new(&test_layout()).unwrap();
|
||||||
let cfg: Mirror = runtime
|
// let cfg: Mirror = runtime
|
||||||
.lua
|
// .lua
|
||||||
.load(
|
// .load(
|
||||||
r#"
|
// r#"
|
||||||
return {
|
// return {
|
||||||
target = __SHIM_DIR__ .. "/tools/numa/numa.exe",
|
// target = __SHIM_DIR__ .. "/tools/numa/numa.exe",
|
||||||
args = { "--help" },
|
// args = { "--help" },
|
||||||
env = {
|
// env = {
|
||||||
PATH = { __SHIM_DIR__ .. "/tools/numa", get_env("PATH") }
|
// PATH = { __SHIM_DIR__ .. "/tools/numa", get_env("PATH") }
|
||||||
}
|
// }
|
||||||
}
|
// }
|
||||||
"#,
|
// "#,
|
||||||
)
|
// )
|
||||||
.eval()
|
// .eval()
|
||||||
.unwrap();
|
// .unwrap();
|
||||||
|
//
|
||||||
let path = cfg.env.get("PATH").unwrap().to_str().unwrap();
|
// let path = cfg.env.get("PATH").unwrap().to_str().unwrap();
|
||||||
let prefix = std::env::temp_dir()
|
// let prefix = std::env::temp_dir()
|
||||||
.join("rshim-test-layout")
|
// .join("rshim-test-layout")
|
||||||
.to_string_lossy()
|
// .to_string_lossy()
|
||||||
.replace('\\', "/")
|
// .replace('\\', "/")
|
||||||
+ "/tools/numa;";
|
// + "/tools/numa;";
|
||||||
assert!(path.starts_with(&prefix), "unexpected PATH: {path}");
|
// assert!(path.starts_with(&prefix), "unexpected PATH: {path}");
|
||||||
|
//
|
||||||
// 宿主 PATH 的段应被附加在配置前缀之后
|
// // 宿主 PATH 的段应被附加在配置前缀之后
|
||||||
let host = std::env::var("PATH").unwrap_or_default();
|
// let host = std::env::var("PATH").unwrap_or_default();
|
||||||
if !host.is_empty() {
|
// if !host.is_empty() {
|
||||||
let host_first = std::env::split_paths(&host)
|
// let host_first = std::env::split_paths(&host)
|
||||||
.next()
|
// .next()
|
||||||
.unwrap()
|
// .unwrap()
|
||||||
.to_string_lossy()
|
// .to_string_lossy()
|
||||||
.into_owned();
|
// .into_owned();
|
||||||
assert!(
|
// assert!(
|
||||||
path.contains(&host_first),
|
// path.contains(&host_first),
|
||||||
"missing host PATH segment: {host_first}"
|
// "missing host PATH segment: {host_first}"
|
||||||
);
|
// );
|
||||||
}
|
// }
|
||||||
|
//
|
||||||
// 宿主 PATH 里引号包裹的畸形段(如 "D:\\...\\bin;")应被原样保留,
|
// // 宿主 PATH 里引号包裹的畸形段(如 "D:\\...\\bin;")应被原样保留,
|
||||||
// 而不是让整个配置加载失败
|
// // 而不是让整个配置加载失败
|
||||||
if host.contains('"') {
|
// if host.contains('"') {
|
||||||
assert!(
|
// assert!(
|
||||||
path.contains('"'),
|
// path.contains('"'),
|
||||||
"quoted host segments should be preserved"
|
// "quoted host segments should be preserved"
|
||||||
);
|
// );
|
||||||
}
|
// }
|
||||||
}
|
// }
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -9,225 +9,3 @@ pub fn normalize_path_for_lua(path: &Path) -> String {
|
|||||||
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"]);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -1,706 +1,41 @@
|
|||||||
use crate::error::validation_error;
|
use crate::error::validation_error;
|
||||||
use crate::{lua_string_2_os_string, parse_tokens};
|
use mlua::Value;
|
||||||
use anyhow::{Context, Result, anyhow};
|
|
||||||
use mlua::{LuaString, Table, Value};
|
|
||||||
use std::collections::{HashMap, HashSet};
|
|
||||||
use std::ffi::OsString;
|
|
||||||
use std::fmt;
|
|
||||||
use std::path::PathBuf;
|
|
||||||
use std::str::FromStr;
|
|
||||||
use tinyjson::JsonValue;
|
|
||||||
use tracing::debug;
|
|
||||||
|
|
||||||
/// Lua 值校验器:针对不同上下文定义校验规则
|
pub fn validate_command_name(name: &Value) -> mlua::Result<String> {
|
||||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
|
||||||
pub enum LuaValidator {
|
|
||||||
/// 目标程序路径:必须是字符串路径
|
|
||||||
Target,
|
|
||||||
/// 命令行参数元素:仅支持一维基础标量,禁止嵌套 Table
|
|
||||||
Args,
|
|
||||||
/// 环境变量值:支持基础标量及多维嵌套 Table(递归展平)
|
|
||||||
Env,
|
|
||||||
Aliases,
|
|
||||||
}
|
|
||||||
|
|
||||||
impl fmt::Display for LuaValidator {
|
|
||||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
|
||||||
match self {
|
|
||||||
Self::Target => write!(f, "目标路径 (target)"),
|
|
||||||
Self::Args => write!(f, "命令行参数 (args)"),
|
|
||||||
Self::Env => write!(f, "环境变量 (env)"),
|
|
||||||
Self::Aliases => write!(f, "命令行别名 (aliases)"),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
impl LuaValidator {
|
|
||||||
fn parse_sequence(&self, index: i64, item: &Value) -> mlua::Result<()> {
|
|
||||||
match item {
|
|
||||||
Value::String(_) | Value::Integer(_) | Value::Number(_) | Value::Boolean(_) => Ok(()),
|
|
||||||
Value::Table(tbl) => match self {
|
|
||||||
Self::Env | Self::Aliases => self.validate_sequence_table(tbl),
|
|
||||||
Self::Args => Err(validation_error(format!(
|
|
||||||
"{} 仅支持一维数组,不能包含嵌套 Table,索引位置: {index}",
|
|
||||||
self
|
|
||||||
))),
|
|
||||||
Self::Target => Err(validation_error(format!("{} 仅支持字符串", self))),
|
|
||||||
},
|
|
||||||
other => Err(validation_error(format!(
|
|
||||||
"{} 第 {} 个元素类型无效: {}",
|
|
||||||
self,
|
|
||||||
index,
|
|
||||||
other.type_name()
|
|
||||||
))),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
/// 校验 Table 是否为严格连续的纯数组,并递归校验其内部元素
|
|
||||||
fn validate_sequence_table(&self, tbl: &Table) -> mlua::Result<()> {
|
|
||||||
let mut index = 1i64;
|
|
||||||
|
|
||||||
// 1. 顺序遍历连续整数索引 1..N
|
|
||||||
loop {
|
|
||||||
let item: Value = tbl.raw_get(index)?;
|
|
||||||
if matches!(&item, Value::Nil) {
|
|
||||||
break;
|
|
||||||
}
|
|
||||||
self.parse_sequence(index, &item)?;
|
|
||||||
|
|
||||||
index += 1;
|
|
||||||
}
|
|
||||||
|
|
||||||
// 2. 查漏:校验是否存在空洞索引或字典键 (Key-Value 键值对)
|
|
||||||
for pair in tbl.pairs::<Value, Value>() {
|
|
||||||
let (key, _) = pair?;
|
|
||||||
match key {
|
|
||||||
Value::Integer(i) if i >= 1 && i < index => {} // 已遍历放行
|
|
||||||
Value::Integer(_) => {
|
|
||||||
return Err(validation_error(format!(
|
|
||||||
"{} 第 {} 个元素为 nil 或不存在(请检查漏写引号或变量未定义)",
|
|
||||||
self, index
|
|
||||||
)));
|
|
||||||
}
|
|
||||||
_ => {
|
|
||||||
return Err(validation_error(format!(
|
|
||||||
"{} 必须是纯列表,不能包含键值对/字典结构",
|
|
||||||
self
|
|
||||||
)));
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
Ok(())
|
|
||||||
}
|
|
||||||
|
|
||||||
/// 底层 NUL 字符跨平台安全检查
|
|
||||||
fn ensure_no_nul(context: &impl fmt::Display, os_str: &std::ffi::OsStr) -> mlua::Result<()> {
|
|
||||||
#[cfg(unix)]
|
|
||||||
{
|
|
||||||
use std::os::unix::ffi::OsStrExt;
|
|
||||||
if os_str.as_bytes().contains(&0) {
|
|
||||||
return Err(conversion_error(format!(
|
|
||||||
"{} 的值不能包含 NUL 字符",
|
|
||||||
context
|
|
||||||
)));
|
|
||||||
}
|
|
||||||
}
|
|
||||||
#[cfg(windows)]
|
|
||||||
{
|
|
||||||
use std::os::windows::ffi::OsStrExt;
|
|
||||||
if os_str.encode_wide().any(|c| c == 0) {
|
|
||||||
return Err(validation_error(format!(
|
|
||||||
"{} 的值不能包含 NUL 字符",
|
|
||||||
context
|
|
||||||
)));
|
|
||||||
}
|
|
||||||
}
|
|
||||||
Ok(())
|
|
||||||
}
|
|
||||||
|
|
||||||
/// 【纯粹数据转换与展开】将已通过校验的 Value 递归解析为 OsString 动态数组
|
|
||||||
fn collect_value_into(&self, value: &Value, out: &mut Vec<OsString>) -> mlua::Result<()> {
|
|
||||||
match value {
|
|
||||||
Value::Nil => {}
|
|
||||||
Value::String(s) => {
|
|
||||||
let os_str = lua_string_2_os_string(s)?;
|
|
||||||
match self {
|
|
||||||
Self::Args | Self::Aliases => {
|
|
||||||
if let Some(str_ref) = os_str.to_str() {
|
|
||||||
// 使用 Tokenizer 切分空格与引号
|
|
||||||
out.extend(parse_tokens(str_ref));
|
|
||||||
} else {
|
|
||||||
// 对于无法转为 UTF-8 的特殊二进制数据,作为整体追加
|
|
||||||
out.push(os_str);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
Self::Env => {
|
|
||||||
out.push(os_str);
|
|
||||||
}
|
|
||||||
Self::Target => {}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
Value::Integer(i) => out.push(OsString::from(i.to_string())),
|
|
||||||
Value::Number(n) => {
|
|
||||||
tracing::warn!(%self, value = %n, "浮点数将按十进制格式转换为字符串");
|
|
||||||
out.push(OsString::from(n.to_string()));
|
|
||||||
}
|
|
||||||
Value::Boolean(b) => {
|
|
||||||
tracing::warn!(%self, value = %b, "布尔值将转换为字符串");
|
|
||||||
out.push(OsString::from(b.to_string()));
|
|
||||||
}
|
|
||||||
Value::Table(tbl) => {
|
|
||||||
let mut index = 1i64;
|
|
||||||
loop {
|
|
||||||
let item: Value = tbl.raw_get(index)?;
|
|
||||||
if matches!(item, Value::Nil) {
|
|
||||||
break;
|
|
||||||
}
|
|
||||||
|
|
||||||
Self::collect_value_into(self, &item, out)?;
|
|
||||||
index += 1;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
_ => unreachable!("传入收集器的 Value 应已通过 validate 校验"),
|
|
||||||
}
|
|
||||||
Ok(())
|
|
||||||
}
|
|
||||||
fn parse_name(name: &Value) -> mlua::Result<String> {
|
|
||||||
let name_str = match name {
|
let name_str = match name {
|
||||||
Value::String(s) => match s.to_str() {
|
Value::String(s) => s
|
||||||
Ok(str_ref) => str_ref.to_string(),
|
.to_str()
|
||||||
Err(_) => return Err(validation_error("环境变量名必须是合法的 UTF-8 字符串")),
|
.map_err(|_| validation_error("子命令名称必须是合法的 UTF-8 字符串"))?
|
||||||
},
|
.to_string(),
|
||||||
other => {
|
other => {
|
||||||
return Err(validation_error(format!(
|
return Err(validation_error(format!(
|
||||||
"环境变量键名类型错误:期望 string,实际是 {}",
|
"命令行键名类型错误:期望 string,实际是 {}",
|
||||||
other.type_name()
|
other.type_name()
|
||||||
)));
|
)));
|
||||||
}
|
}
|
||||||
};
|
};
|
||||||
if name_str.is_empty() {
|
let trimmed = name_str.trim();
|
||||||
return Err(validation_error("环境变量名不能为空"));
|
if trimmed.is_empty() {
|
||||||
|
return Err(validation_error("子命令名称不能为空字符串"));
|
||||||
}
|
}
|
||||||
|
// 2. 禁止包含空格与不可见控制符(否则 shell 与 clap 无法正确定位)
|
||||||
|
if name_str.contains(|c: char| c.is_whitespace()) {
|
||||||
|
return Err(validation_error(format!(
|
||||||
|
"子命令名称 [{name_str}] 非法:命令名不能包含空格或空白字符"
|
||||||
|
)));
|
||||||
|
}
|
||||||
|
|
||||||
if name_str.contains('=') {
|
if name_str.contains('=') {
|
||||||
return Err(validation_error(format!(
|
return Err(validation_error(format!(
|
||||||
"环境变量名 [{}] 不能包含 '='",
|
"子命令名称 [{}] 不能包含 '='",
|
||||||
name_str
|
name_str
|
||||||
)));
|
)));
|
||||||
}
|
}
|
||||||
if name_str.contains('\0') {
|
if name_str.contains('\0') {
|
||||||
return Err(validation_error(format!(
|
return Err(validation_error(format!(
|
||||||
"环境变量名 [{}] 不能包含 NUL 字符",
|
"子命令名称 [{}] 不能包含 NUL 字符",
|
||||||
name_str
|
name_str
|
||||||
)));
|
)));
|
||||||
}
|
}
|
||||||
Ok(name_str)
|
Ok(name_str)
|
||||||
}
|
}
|
||||||
fn parse_val(&self, name: &str, value: &Value) -> mlua::Result<()> {
|
|
||||||
match value {
|
|
||||||
Value::Nil
|
|
||||||
| Value::String(_)
|
|
||||||
| Value::Integer(_)
|
|
||||||
| Value::Number(_)
|
|
||||||
| Value::Boolean(_) => Ok(()),
|
|
||||||
Value::Table(tbl) => self.validate_sequence_table(tbl),
|
|
||||||
other => Err(validation_error(format!(
|
|
||||||
"{} [{}]不支持 {} 类型, 仅支持 string/number/boolean 或嵌套数组",
|
|
||||||
self,
|
|
||||||
name,
|
|
||||||
other.type_name()
|
|
||||||
))),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
impl LuaValidator {
|
|
||||||
/// 解析并校验 `target`
|
|
||||||
pub fn parse_target(value: &Value) -> mlua::Result<PathBuf> {
|
|
||||||
let ctx = Self::Target;
|
|
||||||
match value {
|
|
||||||
Value::String(s) => {
|
|
||||||
let os_str = lua_string_2_os_string(&s)?;
|
|
||||||
Self::ensure_no_nul(&ctx, &os_str)?;
|
|
||||||
Ok(PathBuf::from(os_str))
|
|
||||||
}
|
|
||||||
Value::Nil => Err(validation_error("缺少必填字段 target(应为字符串路径)")),
|
|
||||||
other => Err(validation_error(format!(
|
|
||||||
"{} 需为有效的路径且类型必须是字符串,实际类型是 {}",
|
|
||||||
ctx,
|
|
||||||
other.type_name()
|
|
||||||
))),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
/// 解析并校验 `args`
|
|
||||||
pub fn parse_args(value: &Value) -> mlua::Result<Vec<OsString>> {
|
|
||||||
{
|
|
||||||
debug!("args 配置已关闭 ");
|
|
||||||
#[cfg(not(feature = "args"))]
|
|
||||||
Ok(Vec::new())
|
|
||||||
}
|
|
||||||
#[cfg(feature = "args")]
|
|
||||||
{
|
|
||||||
debug!("args 配置已开启");
|
|
||||||
let ctx = Self::Args;
|
|
||||||
match value {
|
|
||||||
Value::Table(tbl) => {
|
|
||||||
ctx.validate_sequence_table(tbl)?;
|
|
||||||
|
|
||||||
let capacity = tbl.raw_len().min(128);
|
|
||||||
|
|
||||||
let mut raw_parts = Vec::with_capacity(capacity);
|
|
||||||
|
|
||||||
Self::collect_value_into(&ctx, value, &mut raw_parts)?;
|
|
||||||
|
|
||||||
for part in &raw_parts {
|
|
||||||
Self::ensure_no_nul(&ctx, &part)?;
|
|
||||||
}
|
|
||||||
Ok(raw_parts)
|
|
||||||
}
|
|
||||||
other => Err(validation_error(format!(
|
|
||||||
"{} 必须是数组列表,实际类型是 {}",
|
|
||||||
ctx,
|
|
||||||
other.type_name()
|
|
||||||
)))?,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
/// 校验环境变量名的合法性(Windows 约束:非空、不含 '='、不含 NUL)
|
|
||||||
fn parse_env_name(name: &Value) -> mlua::Result<String> {
|
|
||||||
Self::parse_name(name)
|
|
||||||
}
|
|
||||||
|
|
||||||
/// 解析、校验并拼接单个环境变量键值对 (`key`, `value`) -> (`OsString`)
|
|
||||||
fn parse_env_val(
|
|
||||||
context: &LuaValidator,
|
|
||||||
name: &str,
|
|
||||||
raw_val: &Value,
|
|
||||||
) -> mlua::Result<OsString> {
|
|
||||||
context.parse_val(name, raw_val)?;
|
|
||||||
|
|
||||||
// let ctx = format!("{} [{}]", context, name);
|
|
||||||
let capacity = match raw_val {
|
|
||||||
Value::Table(t) => t.raw_len().min(128),
|
|
||||||
Value::Nil => 0,
|
|
||||||
_ => 1,
|
|
||||||
};
|
|
||||||
|
|
||||||
let mut parts = Vec::with_capacity(capacity);
|
|
||||||
|
|
||||||
Self::collect_value_into(&context, raw_val, &mut parts)?;
|
|
||||||
// context.collect_value_into(raw_val, &mut parts)?;
|
|
||||||
|
|
||||||
// 校验每个展开元素的 NUL 字符
|
|
||||||
for part in &parts {
|
|
||||||
Self::ensure_no_nul(&name, part)?;
|
|
||||||
}
|
|
||||||
// 3. 使用系统路径分隔符拼接数组列表
|
|
||||||
let joined_os_str = std::env::join_paths(parts).map_err(|e| {
|
|
||||||
validation_error(format!("{} 的值无法用系统路径分隔符拼接: {}", context, e))
|
|
||||||
})?;
|
|
||||||
Ok(joined_os_str)
|
|
||||||
}
|
|
||||||
|
|
||||||
/// 解析整个 `env` Table,直接返回安全的环境变量 Map
|
|
||||||
pub fn parse_env(value: &Value) -> mlua::Result<HashMap<String, OsString>> {
|
|
||||||
let ctx = Self::Env;
|
|
||||||
|
|
||||||
let Value::Table(tbl) = value else {
|
|
||||||
// 这里作为防御性校验,外部虽然 match 过了,但内层仍保持强类型安全
|
|
||||||
return Err(validation_error(format!(
|
|
||||||
"{} 必须是键值表 (table),实际类型是 {}",
|
|
||||||
ctx,
|
|
||||||
value.type_name()
|
|
||||||
)));
|
|
||||||
};
|
|
||||||
|
|
||||||
let mut env_map = HashMap::new();
|
|
||||||
let mut n = 0i32;
|
|
||||||
|
|
||||||
for pair in tbl.pairs::<Value, Value>() {
|
|
||||||
n += 1;
|
|
||||||
|
|
||||||
let (raw_key, raw_val) = pair?;
|
|
||||||
debug!("env {} {}", n, raw_val.type_name());
|
|
||||||
|
|
||||||
// 1. 解析并校验 Key,拿到安全的 String
|
|
||||||
let key = Self::parse_env_name(&raw_key)?;
|
|
||||||
|
|
||||||
// 2. 借用 &name 传递给 Value 解析器作为上下文
|
|
||||||
let val = Self::parse_env_val(&ctx, &key, &raw_val)?;
|
|
||||||
|
|
||||||
env_map.insert(key, val);
|
|
||||||
}
|
|
||||||
|
|
||||||
Ok(env_map)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
impl LuaValidator {
|
|
||||||
/// 校验环境变量名的合法性(Windows 约束:非空、不含 '='、不含 NUL)
|
|
||||||
fn parse_aliases_name(name: &Value) -> mlua::Result<String> {
|
|
||||||
Self::parse_name(name)
|
|
||||||
}
|
|
||||||
fn parse_aliases_val(
|
|
||||||
context: &LuaValidator,
|
|
||||||
name: &str,
|
|
||||||
raw_val: &Value,
|
|
||||||
) -> mlua::Result<Vec<OsString>> {
|
|
||||||
context.parse_val(name, raw_val)?;
|
|
||||||
|
|
||||||
// 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::with_capacity(capacity);
|
|
||||||
Self::collect_value_into(context, raw_val, &mut parts)?;
|
|
||||||
Ok(parts)
|
|
||||||
}
|
|
||||||
/// 解析并打平别名表 (aliases)
|
|
||||||
/// - 支持输入为 Nil / None / Table
|
|
||||||
/// - 别名的值支持:String, Number, Boolean, Nil, " " 空白串, Table(连续数组)
|
|
||||||
/// - 字符串作为整体参数保存,仅做 trim() 清理首尾空格,不按空格拆分
|
|
||||||
/// - 包含拓扑展开与死环检测
|
|
||||||
pub fn parse_aliases(value: &Value) -> mlua::Result<HashMap<String, Vec<OsString>>> {
|
|
||||||
let ctx = Self::Aliases;
|
|
||||||
// 1. 处理 nil / None 的情况,直接返回空 HashMap
|
|
||||||
let Value::Table(table) = value else {
|
|
||||||
return Err(validation_error(format!(
|
|
||||||
"{} 必须是键值表 (table) ,当前类型: {}",
|
|
||||||
ctx,
|
|
||||||
value.type_name()
|
|
||||||
)));
|
|
||||||
};
|
|
||||||
|
|
||||||
// 阶段一:提取原始别名映射 (Raw Extraction)
|
|
||||||
let mut raw_aliases: HashMap<String, Vec<OsString>> = HashMap::new();
|
|
||||||
|
|
||||||
for pair in table.pairs::<Value, Value>() {
|
|
||||||
let (raw_key, raw_val) = pair?;
|
|
||||||
let key = Self::parse_aliases_name(&raw_key)?;
|
|
||||||
let val = Self::parse_aliases_val(&ctx, &key, &raw_val)?;
|
|
||||||
|
|
||||||
raw_aliases.insert(key, val);
|
|
||||||
}
|
|
||||||
// 阶段二:递归拓扑打平与循环引用检测 (Flattening & Cycle Detection)
|
|
||||||
let mut visited_stack = HashSet::new();
|
|
||||||
|
|
||||||
for key in raw_aliases.keys() {
|
|
||||||
visited_stack.clear();
|
|
||||||
Self::detect_alias_cycle(key, &raw_aliases, &mut visited_stack)?;
|
|
||||||
}
|
|
||||||
|
|
||||||
// 阶段三:无环前提下的高效展开
|
|
||||||
let mut flattened_aliases: HashMap<String, Vec<OsString>> =
|
|
||||||
HashMap::with_capacity(raw_aliases.len());
|
|
||||||
|
|
||||||
let mut resolved_args = Vec::new();
|
|
||||||
|
|
||||||
for key in raw_aliases.keys() {
|
|
||||||
resolved_args.clear();
|
|
||||||
Self::expand_alias_dfs(key, &raw_aliases, &mut resolved_args);
|
|
||||||
flattened_aliases.insert(key.clone(), resolved_args.clone());
|
|
||||||
}
|
|
||||||
|
|
||||||
Ok(flattened_aliases)
|
|
||||||
}
|
|
||||||
|
|
||||||
/// 仅用于校验别名依赖图中是否存在死循环(不消耗额外的参数拼接内存)
|
|
||||||
fn detect_alias_cycle(
|
|
||||||
current_key: &str,
|
|
||||||
raw_aliases: &HashMap<String, Vec<OsString>>,
|
|
||||||
visited_stack: &mut HashSet<String>,
|
|
||||||
) -> mlua::Result<()> {
|
|
||||||
// 递归栈中再次遇到相同的 Key,说明存在死循环
|
|
||||||
if visited_stack.contains(current_key) {
|
|
||||||
return Err(validation_error(format!(
|
|
||||||
"配置加载失败: 检测到别名循环嵌套依赖 'mr:{}'",
|
|
||||||
current_key
|
|
||||||
)));
|
|
||||||
}
|
|
||||||
|
|
||||||
if let Some(args) = raw_aliases.get(current_key) {
|
|
||||||
// 压栈
|
|
||||||
visited_stack.insert(current_key.to_string());
|
|
||||||
|
|
||||||
for arg in args {
|
|
||||||
let arg_str = arg.to_string_lossy();
|
|
||||||
if let Some(sub_key) = arg_str.strip_prefix("mr:") {
|
|
||||||
// 如果引用的子别名在映射表中存在,则深度优先校验
|
|
||||||
if raw_aliases.contains_key(sub_key) {
|
|
||||||
Self::detect_alias_cycle(sub_key, raw_aliases, visited_stack)?;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// 出栈(回溯)
|
|
||||||
visited_stack.remove(current_key);
|
|
||||||
}
|
|
||||||
|
|
||||||
Ok(())
|
|
||||||
}
|
|
||||||
/// 安全拓扑展开:在保证绝对无环的前提下递归展开 mr: 前缀参数
|
|
||||||
fn expand_alias_dfs(
|
|
||||||
current_key: &str,
|
|
||||||
raw_aliases: &HashMap<String, Vec<OsString>>,
|
|
||||||
out: &mut Vec<OsString>,
|
|
||||||
) {
|
|
||||||
if let Some(args) = raw_aliases.get(current_key) {
|
|
||||||
for arg in args {
|
|
||||||
let arg_str = arg.to_string_lossy();
|
|
||||||
if let Some(sub_key) = arg_str.strip_prefix("mr:") {
|
|
||||||
if raw_aliases.contains_key(sub_key) {
|
|
||||||
// 安全地直接递归展开,无需再检查死循环
|
|
||||||
Self::expand_alias_dfs(sub_key, raw_aliases, out);
|
|
||||||
} else {
|
|
||||||
// 找不到对应的别名,按原样参数输出
|
|
||||||
out.push(arg.clone());
|
|
||||||
}
|
|
||||||
} else {
|
|
||||||
// 普通参数,直接输出
|
|
||||||
out.push(arg.clone());
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
#[cfg(test)]
|
|
||||||
mod tests {
|
|
||||||
use super::*;
|
|
||||||
use crate::Mirror;
|
|
||||||
use mlua::{FromLua, Lua};
|
|
||||||
|
|
||||||
fn parse(src: &str) -> mlua::Result<Mirror> {
|
|
||||||
let lua = Lua::new();
|
|
||||||
// 模拟 runtime 注入的 get_env:返回按平台分隔符拆分的段数组(空变量返回空表)
|
|
||||||
let get_env = lua
|
|
||||||
.create_function(|lua, key: String| -> mlua::Result<Table> {
|
|
||||||
let value = std::env::var(key).unwrap_or_default();
|
|
||||||
let segments: Vec<String> = if value.is_empty() {
|
|
||||||
Vec::new()
|
|
||||||
} else {
|
|
||||||
std::env::split_paths(std::ffi::OsStr::new(&value))
|
|
||||||
.map(|p| p.to_string_lossy().into_owned())
|
|
||||||
.collect()
|
|
||||||
};
|
|
||||||
lua.create_sequence_from(segments)
|
|
||||||
})
|
|
||||||
.unwrap();
|
|
||||||
lua.globals().set("get_env", get_env).unwrap();
|
|
||||||
|
|
||||||
let value = lua.load(src).eval::<Value>()?;
|
|
||||||
let t = Mirror::from_lua(value, &lua);
|
|
||||||
println!("读取出的数据:{:?}", t.clone()?);
|
|
||||||
t
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn parses_basic_config() {
|
|
||||||
let cfg = parse(
|
|
||||||
r#"
|
|
||||||
return {
|
|
||||||
target = "C:/tools/git.exe",
|
|
||||||
args = { "--no-pager",2},
|
|
||||||
env = {
|
|
||||||
PATH = { "C:/tools/git/bin", "C:/Windows", get_env("PATH")},
|
|
||||||
HOME = "C:/tools/home",
|
|
||||||
CONST = 3,
|
|
||||||
BOOL = true,
|
|
||||||
},
|
|
||||||
aliases = {
|
|
||||||
-- 1. 整体字符串(会自动 trim 首尾空格,保留为单个参数)
|
|
||||||
st = "status -s",
|
|
||||||
|
|
||||||
-- 2. 连续数组表 (Sequence Table)
|
|
||||||
lg = { "log", "--oneline", "-n", 10 },
|
|
||||||
|
|
||||||
-- 3. 标量数字与布尔值支持
|
|
||||||
v = 1,
|
|
||||||
quiet = true,
|
|
||||||
|
|
||||||
-- 4. 嵌套别名组合(加载期会自动展开并进行死环检测)
|
|
||||||
base_log = "log --graph",
|
|
||||||
all_log = { "mr:base_log", "--all","D:\\CNWei\\CNW\\Rust","D:/CNWei/CNW/Rust/" ,[[D:\CNWei\CNW\Rust\]]},
|
|
||||||
|
|
||||||
-- 5. nil 或空字符串(解析为空参数列表)
|
|
||||||
empty_alias = nil,
|
|
||||||
blank = " "
|
|
||||||
}
|
|
||||||
}
|
|
||||||
"#,
|
|
||||||
)
|
|
||||||
.unwrap();
|
|
||||||
assert_eq!(cfg.target, PathBuf::from("C:/tools/git.exe"));
|
|
||||||
// assert_eq!(cfg.args, vec!["--no-pager", "2"]);
|
|
||||||
assert_eq!(cfg.env.get("HOME").unwrap().to_str(), Some("C:/tools/home"));
|
|
||||||
// PATH 前缀来自配置,随后附加 get_env("PATH") 拆出的宿主 PATH 段
|
|
||||||
let path = cfg.env.get("PATH").unwrap().to_str().unwrap();
|
|
||||||
assert!(
|
|
||||||
path == "C:/tools/git/bin;C:/Windows"
|
|
||||||
|| path.starts_with("C:/tools/git/bin;C:/Windows;"),
|
|
||||||
"unexpected PATH: {path}"
|
|
||||||
);
|
|
||||||
// ==================== aliases 断言校验 ====================
|
|
||||||
|
|
||||||
// 1. 整体字符串:保留原样(trim 后),不切分空格
|
|
||||||
assert_eq!(
|
|
||||||
cfg.aliases.get("st").unwrap(),
|
|
||||||
&vec![OsString::from("status"), OsString::from("-s")]
|
|
||||||
);
|
|
||||||
|
|
||||||
// 2. 连续数组表:按顺序转为 OsString 列表
|
|
||||||
assert_eq!(
|
|
||||||
cfg.aliases.get("lg").unwrap(),
|
|
||||||
&vec![
|
|
||||||
OsString::from("log"),
|
|
||||||
OsString::from("--oneline"),
|
|
||||||
OsString::from("-n"),
|
|
||||||
OsString::from("10")
|
|
||||||
]
|
|
||||||
);
|
|
||||||
|
|
||||||
// 3. 标量数字与布尔值支持
|
|
||||||
assert_eq!(cfg.aliases.get("v").unwrap(), &vec![OsString::from("1")]);
|
|
||||||
assert_eq!(
|
|
||||||
cfg.aliases.get("quiet").unwrap(),
|
|
||||||
&vec![OsString::from("true")]
|
|
||||||
);
|
|
||||||
|
|
||||||
// 4. 嵌套别名组合:加载期递归拓扑打平
|
|
||||||
// base_log 本身为 "log --graph"
|
|
||||||
assert_eq!(
|
|
||||||
cfg.aliases.get("base_log").unwrap(),
|
|
||||||
&vec![OsString::from("log"), OsString::from("--graph")]
|
|
||||||
);
|
|
||||||
// all_log 展开 mr:base_log 替换为 "log --graph",追加 "--all"
|
|
||||||
assert_eq!(
|
|
||||||
cfg.aliases.get("all_log").unwrap(),
|
|
||||||
&vec![
|
|
||||||
OsString::from("log"),
|
|
||||||
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
|
|
||||||
assert_eq!(cfg.aliases.get("blank").unwrap(), &Vec::<OsString>::new());
|
|
||||||
// nil 键在遍历表时会被当作空或不存在,不产生 key 或值为空 Vec
|
|
||||||
assert!(
|
|
||||||
cfg.aliases
|
|
||||||
.get("empty_alias")
|
|
||||||
.map_or(true, |v| v.is_empty())
|
|
||||||
);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[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").unwrap().to_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").unwrap().to_str(), Some(""));
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn expands_nested_env_array() {
|
|
||||||
// get_env("PATH") 现在返回拆分后的段数组,嵌套表应被递归展开
|
|
||||||
let cfg = parse(
|
|
||||||
r#"return { target = "t.exe", env = { PATH = { "C:/a", { "C:/b", "C:/c" } } } }"#,
|
|
||||||
)
|
|
||||||
.unwrap();
|
|
||||||
assert_eq!(
|
|
||||||
cfg.env.get("PATH").unwrap().to_str(),
|
|
||||||
Some("C:/a;C:/b;C:/c")
|
|
||||||
);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[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() {
|
|
||||||
// ValueShunt::Args 允许基础标量(数字/布尔)转字符串,仅禁止嵌套表
|
|
||||||
let cfg = parse(r#"return { target = "t.exe", args = { 1, true } }"#).unwrap();
|
|
||||||
assert_eq!(cfg.args, vec!["1", "true"]);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[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());
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -12,6 +12,7 @@ windows-sys = { workspace = true }
|
|||||||
tracing = { workspace = true }
|
tracing = { workspace = true }
|
||||||
tracing-subscriber = { workspace = true }
|
tracing-subscriber = { workspace = true }
|
||||||
tracing-appender = { workspace = true }
|
tracing-appender = { workspace = true }
|
||||||
|
mlua ={workspace = true}
|
||||||
|
|
||||||
[features]
|
[features]
|
||||||
default = []
|
default = []
|
||||||
|
|||||||
@@ -1,4 +1,5 @@
|
|||||||
use mirror_core::{init_logging_from, Layout, Mirror};
|
use mirror_core::{init_logging_from, Layout, };
|
||||||
|
use crate::mirror::Mirror;
|
||||||
use std::ffi::OsString;
|
use std::ffi::OsString;
|
||||||
use std::{env, process::exit};
|
use std::{env, process::exit};
|
||||||
use anyhow::Context;
|
use anyhow::Context;
|
||||||
@@ -6,6 +7,9 @@ use tracing::{debug, error};
|
|||||||
use crate::sys::*;
|
use crate::sys::*;
|
||||||
|
|
||||||
pub mod sys;
|
pub mod sys;
|
||||||
|
pub mod mirror;
|
||||||
|
pub mod validators;
|
||||||
|
pub mod utils;
|
||||||
|
|
||||||
fn main() {
|
fn main() {
|
||||||
let current_exe = match env::current_exe().context("获取代理程序路径失败"){
|
let current_exe = match env::current_exe().context("获取代理程序路径失败"){
|
||||||
|
|||||||
@@ -1,6 +1,5 @@
|
|||||||
use crate::error::validation_error;
|
use mirror_core::error::validation_error;
|
||||||
use crate::validators::LuaValidator;
|
use crate::validators::LuaValidator;
|
||||||
use crate::{Layout, LuaRuntime};
|
|
||||||
use anyhow::{Context, Result, bail};
|
use anyhow::{Context, Result, bail};
|
||||||
use mlua::{FromLua, Lua, Table, Value};
|
use mlua::{FromLua, Lua, Table, Value};
|
||||||
use std::collections::HashMap;
|
use std::collections::HashMap;
|
||||||
@@ -9,6 +8,47 @@ use std::ffi::OsString;
|
|||||||
use std::path::PathBuf;
|
use std::path::PathBuf;
|
||||||
use std::process::Command;
|
use std::process::Command;
|
||||||
use tracing::{debug, trace};
|
use tracing::{debug, trace};
|
||||||
|
use mirror_core::{Layout, LuaRuntime};
|
||||||
|
use mirror_core::validate_command_name;
|
||||||
|
|
||||||
|
struct MasterMirror{
|
||||||
|
master:HashMap<String,Mirror>,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl FromLua for MasterMirror {
|
||||||
|
fn from_lua(value: Value, lua: &Lua) -> mlua::Result<Self> {
|
||||||
|
let root_tbl = match value {
|
||||||
|
Value::Table(t) => t,
|
||||||
|
other => {
|
||||||
|
return Err(validation_error(format!(
|
||||||
|
"mirror 注册表的顶层配置必须是 Table,实际检测到: {}",
|
||||||
|
other.type_name()
|
||||||
|
)));
|
||||||
|
}
|
||||||
|
};
|
||||||
|
let mut master = HashMap::new();
|
||||||
|
for pair in root_tbl.pairs::<Value, Value>() {
|
||||||
|
let (cmd_name, cmd_entry) = pair?;
|
||||||
|
|
||||||
|
let cmd_name = validate_command_name(&cmd_name)?;
|
||||||
|
// 2. 校验 entry 是否是 Table
|
||||||
|
if !cmd_entry.is_table() {
|
||||||
|
return Err(validation_error(format!(
|
||||||
|
"子命令 '{cmd_name}' 的配置必须是 Table,实际检测到: {}",
|
||||||
|
cmd_entry.type_name()
|
||||||
|
)));
|
||||||
|
}
|
||||||
|
// 3. 直接交由子命令解析,外层负责补充错误上下文
|
||||||
|
let cmd_def = Mirror::from_lua(cmd_entry, lua).map_err(|err| {
|
||||||
|
validation_error(format!("子命令 '{cmd_name}' 配置解析失败:\n{err}"))
|
||||||
|
})?;
|
||||||
|
|
||||||
|
master.insert(cmd_name, cmd_def);
|
||||||
|
}
|
||||||
|
Ok(Self { master })
|
||||||
|
}}
|
||||||
|
|
||||||
|
|
||||||
#[derive(Debug, Clone, Default)]
|
#[derive(Debug, Clone, Default)]
|
||||||
pub struct Mirror {
|
pub struct Mirror {
|
||||||
pub target: PathBuf,
|
pub target: PathBuf,
|
||||||
@@ -20,7 +60,7 @@ pub struct Mirror {
|
|||||||
impl Mirror {
|
impl Mirror {
|
||||||
pub fn load(layout: &Layout) -> Result<Self> {
|
pub fn load(layout: &Layout) -> Result<Self> {
|
||||||
// 策略 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() {
|
if !global_config.is_file() {
|
||||||
bail!("主配置文件不存在: {}", global_config.display());
|
bail!("主配置文件不存在: {}", global_config.display());
|
||||||
225
mirror-shim/src/utils.rs
Normal file
225
mirror-shim/src/utils.rs
Normal file
@@ -0,0 +1,225 @@
|
|||||||
|
use mlua::LuaString;
|
||||||
|
use std::ffi::OsString;
|
||||||
|
|
||||||
|
/// 将 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"]);
|
||||||
|
}
|
||||||
|
}
|
||||||
706
mirror-shim/src/validators.rs
Normal file
706
mirror-shim/src/validators.rs
Normal file
@@ -0,0 +1,706 @@
|
|||||||
|
use mirror_core::error::validation_error;
|
||||||
|
use anyhow::{Context, Result, anyhow};
|
||||||
|
use mlua::{LuaString, Table, Value};
|
||||||
|
use std::collections::{HashMap, HashSet};
|
||||||
|
use std::ffi::OsString;
|
||||||
|
use std::fmt;
|
||||||
|
use std::path::PathBuf;
|
||||||
|
use std::str::FromStr;
|
||||||
|
// use tinyjson::JsonValue;
|
||||||
|
use tracing::debug;
|
||||||
|
use crate::utils::{lua_string_2_os_string, parse_tokens};
|
||||||
|
|
||||||
|
/// Lua 值校验器:针对不同上下文定义校验规则
|
||||||
|
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||||
|
pub enum LuaValidator {
|
||||||
|
/// 目标程序路径:必须是字符串路径
|
||||||
|
Target,
|
||||||
|
/// 命令行参数元素:仅支持一维基础标量,禁止嵌套 Table
|
||||||
|
Args,
|
||||||
|
/// 环境变量值:支持基础标量及多维嵌套 Table(递归展平)
|
||||||
|
Env,
|
||||||
|
Aliases,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl fmt::Display for LuaValidator {
|
||||||
|
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||||
|
match self {
|
||||||
|
Self::Target => write!(f, "目标路径 (target)"),
|
||||||
|
Self::Args => write!(f, "命令行参数 (args)"),
|
||||||
|
Self::Env => write!(f, "环境变量 (env)"),
|
||||||
|
Self::Aliases => write!(f, "命令行别名 (aliases)"),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
impl LuaValidator {
|
||||||
|
fn parse_sequence(&self, index: i64, item: &Value) -> mlua::Result<()> {
|
||||||
|
match item {
|
||||||
|
Value::String(_) | Value::Integer(_) | Value::Number(_) | Value::Boolean(_) => Ok(()),
|
||||||
|
Value::Table(tbl) => match self {
|
||||||
|
Self::Env | Self::Aliases => self.validate_sequence_table(tbl),
|
||||||
|
Self::Args => Err(validation_error(format!(
|
||||||
|
"{} 仅支持一维数组,不能包含嵌套 Table,索引位置: {index}",
|
||||||
|
self
|
||||||
|
))),
|
||||||
|
Self::Target => Err(validation_error(format!("{} 仅支持字符串", self))),
|
||||||
|
},
|
||||||
|
other => Err(validation_error(format!(
|
||||||
|
"{} 第 {} 个元素类型无效: {}",
|
||||||
|
self,
|
||||||
|
index,
|
||||||
|
other.type_name()
|
||||||
|
))),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 校验 Table 是否为严格连续的纯数组,并递归校验其内部元素
|
||||||
|
fn validate_sequence_table(&self, tbl: &Table) -> mlua::Result<()> {
|
||||||
|
let mut index = 1i64;
|
||||||
|
|
||||||
|
// 1. 顺序遍历连续整数索引 1..N
|
||||||
|
loop {
|
||||||
|
let item: Value = tbl.raw_get(index)?;
|
||||||
|
if matches!(&item, Value::Nil) {
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
self.parse_sequence(index, &item)?;
|
||||||
|
|
||||||
|
index += 1;
|
||||||
|
}
|
||||||
|
|
||||||
|
// 2. 查漏:校验是否存在空洞索引或字典键 (Key-Value 键值对)
|
||||||
|
for pair in tbl.pairs::<Value, Value>() {
|
||||||
|
let (key, _) = pair?;
|
||||||
|
match key {
|
||||||
|
Value::Integer(i) if i >= 1 && i < index => {} // 已遍历放行
|
||||||
|
Value::Integer(_) => {
|
||||||
|
return Err(validation_error(format!(
|
||||||
|
"{} 第 {} 个元素为 nil 或不存在(请检查漏写引号或变量未定义)",
|
||||||
|
self, index
|
||||||
|
)));
|
||||||
|
}
|
||||||
|
_ => {
|
||||||
|
return Err(validation_error(format!(
|
||||||
|
"{} 必须是纯列表,不能包含键值对/字典结构",
|
||||||
|
self
|
||||||
|
)));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 底层 NUL 字符跨平台安全检查
|
||||||
|
fn ensure_no_nul(context: &impl fmt::Display, os_str: &std::ffi::OsStr) -> mlua::Result<()> {
|
||||||
|
#[cfg(unix)]
|
||||||
|
{
|
||||||
|
use std::os::unix::ffi::OsStrExt;
|
||||||
|
if os_str.as_bytes().contains(&0) {
|
||||||
|
return Err(conversion_error(format!(
|
||||||
|
"{} 的值不能包含 NUL 字符",
|
||||||
|
context
|
||||||
|
)));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
#[cfg(windows)]
|
||||||
|
{
|
||||||
|
use std::os::windows::ffi::OsStrExt;
|
||||||
|
if os_str.encode_wide().any(|c| c == 0) {
|
||||||
|
return Err(validation_error(format!(
|
||||||
|
"{} 的值不能包含 NUL 字符",
|
||||||
|
context
|
||||||
|
)));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 【纯粹数据转换与展开】将已通过校验的 Value 递归解析为 OsString 动态数组
|
||||||
|
fn collect_value_into(&self, value: &Value, out: &mut Vec<OsString>) -> mlua::Result<()> {
|
||||||
|
match value {
|
||||||
|
Value::Nil => {}
|
||||||
|
Value::String(s) => {
|
||||||
|
let os_str = lua_string_2_os_string(s)?;
|
||||||
|
match self {
|
||||||
|
Self::Args | Self::Aliases => {
|
||||||
|
if let Some(str_ref) = os_str.to_str() {
|
||||||
|
// 使用 Tokenizer 切分空格与引号
|
||||||
|
out.extend(parse_tokens(str_ref));
|
||||||
|
} else {
|
||||||
|
// 对于无法转为 UTF-8 的特殊二进制数据,作为整体追加
|
||||||
|
out.push(os_str);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
Self::Env => {
|
||||||
|
out.push(os_str);
|
||||||
|
}
|
||||||
|
Self::Target => {}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
Value::Integer(i) => out.push(OsString::from(i.to_string())),
|
||||||
|
Value::Number(n) => {
|
||||||
|
tracing::warn!(%self, value = %n, "浮点数将按十进制格式转换为字符串");
|
||||||
|
out.push(OsString::from(n.to_string()));
|
||||||
|
}
|
||||||
|
Value::Boolean(b) => {
|
||||||
|
tracing::warn!(%self, value = %b, "布尔值将转换为字符串");
|
||||||
|
out.push(OsString::from(b.to_string()));
|
||||||
|
}
|
||||||
|
Value::Table(tbl) => {
|
||||||
|
let mut index = 1i64;
|
||||||
|
loop {
|
||||||
|
let item: Value = tbl.raw_get(index)?;
|
||||||
|
if matches!(item, Value::Nil) {
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
|
||||||
|
Self::collect_value_into(self, &item, out)?;
|
||||||
|
index += 1;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
_ => unreachable!("传入收集器的 Value 应已通过 validate 校验"),
|
||||||
|
}
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
fn parse_name(name: &Value) -> mlua::Result<String> {
|
||||||
|
let name_str = match name {
|
||||||
|
Value::String(s) => match s.to_str() {
|
||||||
|
Ok(str_ref) => str_ref.to_string(),
|
||||||
|
Err(_) => return Err(validation_error("环境变量名必须是合法的 UTF-8 字符串")),
|
||||||
|
},
|
||||||
|
other => {
|
||||||
|
return Err(validation_error(format!(
|
||||||
|
"环境变量键名类型错误:期望 string,实际是 {}",
|
||||||
|
other.type_name()
|
||||||
|
)));
|
||||||
|
}
|
||||||
|
};
|
||||||
|
if name_str.is_empty() {
|
||||||
|
return Err(validation_error("环境变量名不能为空"));
|
||||||
|
}
|
||||||
|
if name_str.contains('=') {
|
||||||
|
return Err(validation_error(format!(
|
||||||
|
"环境变量名 [{}] 不能包含 '='",
|
||||||
|
name_str
|
||||||
|
)));
|
||||||
|
}
|
||||||
|
if name_str.contains('\0') {
|
||||||
|
return Err(validation_error(format!(
|
||||||
|
"环境变量名 [{}] 不能包含 NUL 字符",
|
||||||
|
name_str
|
||||||
|
)));
|
||||||
|
}
|
||||||
|
Ok(name_str)
|
||||||
|
}
|
||||||
|
fn parse_val(&self, name: &str, value: &Value) -> mlua::Result<()> {
|
||||||
|
match value {
|
||||||
|
Value::Nil
|
||||||
|
| Value::String(_)
|
||||||
|
| Value::Integer(_)
|
||||||
|
| Value::Number(_)
|
||||||
|
| Value::Boolean(_) => Ok(()),
|
||||||
|
Value::Table(tbl) => self.validate_sequence_table(tbl),
|
||||||
|
other => Err(validation_error(format!(
|
||||||
|
"{} [{}]不支持 {} 类型, 仅支持 string/number/boolean 或嵌套数组",
|
||||||
|
self,
|
||||||
|
name,
|
||||||
|
other.type_name()
|
||||||
|
))),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
impl LuaValidator {
|
||||||
|
/// 解析并校验 `target`
|
||||||
|
pub fn parse_target(value: &Value) -> mlua::Result<PathBuf> {
|
||||||
|
let ctx = Self::Target;
|
||||||
|
match value {
|
||||||
|
Value::String(s) => {
|
||||||
|
let os_str = lua_string_2_os_string(&s)?;
|
||||||
|
Self::ensure_no_nul(&ctx, &os_str)?;
|
||||||
|
Ok(PathBuf::from(os_str))
|
||||||
|
}
|
||||||
|
Value::Nil => Err(validation_error("缺少必填字段 target(应为字符串路径)")),
|
||||||
|
other => Err(validation_error(format!(
|
||||||
|
"{} 需为有效的路径且类型必须是字符串,实际类型是 {}",
|
||||||
|
ctx,
|
||||||
|
other.type_name()
|
||||||
|
))),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 解析并校验 `args`
|
||||||
|
pub fn parse_args(value: &Value) -> mlua::Result<Vec<OsString>> {
|
||||||
|
{
|
||||||
|
debug!("args 配置已关闭 ");
|
||||||
|
#[cfg(not(feature = "args"))]
|
||||||
|
Ok(Vec::new())
|
||||||
|
}
|
||||||
|
#[cfg(feature = "args")]
|
||||||
|
{
|
||||||
|
debug!("args 配置已开启");
|
||||||
|
let ctx = Self::Args;
|
||||||
|
match value {
|
||||||
|
Value::Table(tbl) => {
|
||||||
|
ctx.validate_sequence_table(tbl)?;
|
||||||
|
|
||||||
|
let capacity = tbl.raw_len().min(128);
|
||||||
|
|
||||||
|
let mut raw_parts = Vec::with_capacity(capacity);
|
||||||
|
|
||||||
|
Self::collect_value_into(&ctx, value, &mut raw_parts)?;
|
||||||
|
|
||||||
|
for part in &raw_parts {
|
||||||
|
Self::ensure_no_nul(&ctx, &part)?;
|
||||||
|
}
|
||||||
|
Ok(raw_parts)
|
||||||
|
}
|
||||||
|
other => Err(validation_error(format!(
|
||||||
|
"{} 必须是数组列表,实际类型是 {}",
|
||||||
|
ctx,
|
||||||
|
other.type_name()
|
||||||
|
)))?,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 校验环境变量名的合法性(Windows 约束:非空、不含 '='、不含 NUL)
|
||||||
|
fn parse_env_name(name: &Value) -> mlua::Result<String> {
|
||||||
|
Self::parse_name(name)
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 解析、校验并拼接单个环境变量键值对 (`key`, `value`) -> (`OsString`)
|
||||||
|
fn parse_env_val(
|
||||||
|
context: &LuaValidator,
|
||||||
|
name: &str,
|
||||||
|
raw_val: &Value,
|
||||||
|
) -> mlua::Result<OsString> {
|
||||||
|
context.parse_val(name, raw_val)?;
|
||||||
|
|
||||||
|
// let ctx = format!("{} [{}]", context, name);
|
||||||
|
let capacity = match raw_val {
|
||||||
|
Value::Table(t) => t.raw_len().min(128),
|
||||||
|
Value::Nil => 0,
|
||||||
|
_ => 1,
|
||||||
|
};
|
||||||
|
|
||||||
|
let mut parts = Vec::with_capacity(capacity);
|
||||||
|
|
||||||
|
Self::collect_value_into(&context, raw_val, &mut parts)?;
|
||||||
|
// context.collect_value_into(raw_val, &mut parts)?;
|
||||||
|
|
||||||
|
// 校验每个展开元素的 NUL 字符
|
||||||
|
for part in &parts {
|
||||||
|
Self::ensure_no_nul(&name, part)?;
|
||||||
|
}
|
||||||
|
// 3. 使用系统路径分隔符拼接数组列表
|
||||||
|
let joined_os_str = std::env::join_paths(parts).map_err(|e| {
|
||||||
|
validation_error(format!("{} 的值无法用系统路径分隔符拼接: {}", context, e))
|
||||||
|
})?;
|
||||||
|
Ok(joined_os_str)
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 解析整个 `env` Table,直接返回安全的环境变量 Map
|
||||||
|
pub fn parse_env(value: &Value) -> mlua::Result<HashMap<String, OsString>> {
|
||||||
|
let ctx = Self::Env;
|
||||||
|
|
||||||
|
let Value::Table(tbl) = value else {
|
||||||
|
// 这里作为防御性校验,外部虽然 match 过了,但内层仍保持强类型安全
|
||||||
|
return Err(validation_error(format!(
|
||||||
|
"{} 必须是键值表 (table),实际类型是 {}",
|
||||||
|
ctx,
|
||||||
|
value.type_name()
|
||||||
|
)));
|
||||||
|
};
|
||||||
|
|
||||||
|
let mut env_map = HashMap::new();
|
||||||
|
let mut n = 0i32;
|
||||||
|
|
||||||
|
for pair in tbl.pairs::<Value, Value>() {
|
||||||
|
n += 1;
|
||||||
|
|
||||||
|
let (raw_key, raw_val) = pair?;
|
||||||
|
debug!("env {} {}", n, raw_val.type_name());
|
||||||
|
|
||||||
|
// 1. 解析并校验 Key,拿到安全的 String
|
||||||
|
let key = Self::parse_env_name(&raw_key)?;
|
||||||
|
|
||||||
|
// 2. 借用 &name 传递给 Value 解析器作为上下文
|
||||||
|
let val = Self::parse_env_val(&ctx, &key, &raw_val)?;
|
||||||
|
|
||||||
|
env_map.insert(key, val);
|
||||||
|
}
|
||||||
|
|
||||||
|
Ok(env_map)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
impl LuaValidator {
|
||||||
|
/// 校验环境变量名的合法性(Windows 约束:非空、不含 '='、不含 NUL)
|
||||||
|
fn parse_aliases_name(name: &Value) -> mlua::Result<String> {
|
||||||
|
Self::parse_name(name)
|
||||||
|
}
|
||||||
|
fn parse_aliases_val(
|
||||||
|
context: &LuaValidator,
|
||||||
|
name: &str,
|
||||||
|
raw_val: &Value,
|
||||||
|
) -> mlua::Result<Vec<OsString>> {
|
||||||
|
context.parse_val(name, raw_val)?;
|
||||||
|
|
||||||
|
// 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::with_capacity(capacity);
|
||||||
|
Self::collect_value_into(context, raw_val, &mut parts)?;
|
||||||
|
Ok(parts)
|
||||||
|
}
|
||||||
|
/// 解析并打平别名表 (aliases)
|
||||||
|
/// - 支持输入为 Nil / None / Table
|
||||||
|
/// - 别名的值支持:String, Number, Boolean, Nil, " " 空白串, Table(连续数组)
|
||||||
|
/// - 字符串作为整体参数保存,仅做 trim() 清理首尾空格,不按空格拆分
|
||||||
|
/// - 包含拓扑展开与死环检测
|
||||||
|
pub fn parse_aliases(value: &Value) -> mlua::Result<HashMap<String, Vec<OsString>>> {
|
||||||
|
let ctx = Self::Aliases;
|
||||||
|
// 1. 处理 nil / None 的情况,直接返回空 HashMap
|
||||||
|
let Value::Table(table) = value else {
|
||||||
|
return Err(validation_error(format!(
|
||||||
|
"{} 必须是键值表 (table) ,当前类型: {}",
|
||||||
|
ctx,
|
||||||
|
value.type_name()
|
||||||
|
)));
|
||||||
|
};
|
||||||
|
|
||||||
|
// 阶段一:提取原始别名映射 (Raw Extraction)
|
||||||
|
let mut raw_aliases: HashMap<String, Vec<OsString>> = HashMap::new();
|
||||||
|
|
||||||
|
for pair in table.pairs::<Value, Value>() {
|
||||||
|
let (raw_key, raw_val) = pair?;
|
||||||
|
let key = Self::parse_aliases_name(&raw_key)?;
|
||||||
|
let val = Self::parse_aliases_val(&ctx, &key, &raw_val)?;
|
||||||
|
|
||||||
|
raw_aliases.insert(key, val);
|
||||||
|
}
|
||||||
|
// 阶段二:递归拓扑打平与循环引用检测 (Flattening & Cycle Detection)
|
||||||
|
let mut visited_stack = HashSet::new();
|
||||||
|
|
||||||
|
for key in raw_aliases.keys() {
|
||||||
|
visited_stack.clear();
|
||||||
|
Self::detect_alias_cycle(key, &raw_aliases, &mut visited_stack)?;
|
||||||
|
}
|
||||||
|
|
||||||
|
// 阶段三:无环前提下的高效展开
|
||||||
|
let mut flattened_aliases: HashMap<String, Vec<OsString>> =
|
||||||
|
HashMap::with_capacity(raw_aliases.len());
|
||||||
|
|
||||||
|
let mut resolved_args = Vec::new();
|
||||||
|
|
||||||
|
for key in raw_aliases.keys() {
|
||||||
|
resolved_args.clear();
|
||||||
|
Self::expand_alias_dfs(key, &raw_aliases, &mut resolved_args);
|
||||||
|
flattened_aliases.insert(key.clone(), resolved_args.clone());
|
||||||
|
}
|
||||||
|
|
||||||
|
Ok(flattened_aliases)
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 仅用于校验别名依赖图中是否存在死循环(不消耗额外的参数拼接内存)
|
||||||
|
fn detect_alias_cycle(
|
||||||
|
current_key: &str,
|
||||||
|
raw_aliases: &HashMap<String, Vec<OsString>>,
|
||||||
|
visited_stack: &mut HashSet<String>,
|
||||||
|
) -> mlua::Result<()> {
|
||||||
|
// 递归栈中再次遇到相同的 Key,说明存在死循环
|
||||||
|
if visited_stack.contains(current_key) {
|
||||||
|
return Err(validation_error(format!(
|
||||||
|
"配置加载失败: 检测到别名循环嵌套依赖 'mr:{}'",
|
||||||
|
current_key
|
||||||
|
)));
|
||||||
|
}
|
||||||
|
|
||||||
|
if let Some(args) = raw_aliases.get(current_key) {
|
||||||
|
// 压栈
|
||||||
|
visited_stack.insert(current_key.to_string());
|
||||||
|
|
||||||
|
for arg in args {
|
||||||
|
let arg_str = arg.to_string_lossy();
|
||||||
|
if let Some(sub_key) = arg_str.strip_prefix("mr:") {
|
||||||
|
// 如果引用的子别名在映射表中存在,则深度优先校验
|
||||||
|
if raw_aliases.contains_key(sub_key) {
|
||||||
|
Self::detect_alias_cycle(sub_key, raw_aliases, visited_stack)?;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// 出栈(回溯)
|
||||||
|
visited_stack.remove(current_key);
|
||||||
|
}
|
||||||
|
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
/// 安全拓扑展开:在保证绝对无环的前提下递归展开 mr: 前缀参数
|
||||||
|
fn expand_alias_dfs(
|
||||||
|
current_key: &str,
|
||||||
|
raw_aliases: &HashMap<String, Vec<OsString>>,
|
||||||
|
out: &mut Vec<OsString>,
|
||||||
|
) {
|
||||||
|
if let Some(args) = raw_aliases.get(current_key) {
|
||||||
|
for arg in args {
|
||||||
|
let arg_str = arg.to_string_lossy();
|
||||||
|
if let Some(sub_key) = arg_str.strip_prefix("mr:") {
|
||||||
|
if raw_aliases.contains_key(sub_key) {
|
||||||
|
// 安全地直接递归展开,无需再检查死循环
|
||||||
|
Self::expand_alias_dfs(sub_key, raw_aliases, out);
|
||||||
|
} else {
|
||||||
|
// 找不到对应的别名,按原样参数输出
|
||||||
|
out.push(arg.clone());
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
// 普通参数,直接输出
|
||||||
|
out.push(arg.clone());
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod tests {
|
||||||
|
use super::*;
|
||||||
|
use crate::Mirror;
|
||||||
|
use mlua::{FromLua, Lua};
|
||||||
|
|
||||||
|
fn parse(src: &str) -> mlua::Result<Mirror> {
|
||||||
|
let lua = Lua::new();
|
||||||
|
// 模拟 runtime 注入的 get_env:返回按平台分隔符拆分的段数组(空变量返回空表)
|
||||||
|
let get_env = lua
|
||||||
|
.create_function(|lua, key: String| -> mlua::Result<Table> {
|
||||||
|
let value = std::env::var(key).unwrap_or_default();
|
||||||
|
let segments: Vec<String> = if value.is_empty() {
|
||||||
|
Vec::new()
|
||||||
|
} else {
|
||||||
|
std::env::split_paths(std::ffi::OsStr::new(&value))
|
||||||
|
.map(|p| p.to_string_lossy().into_owned())
|
||||||
|
.collect()
|
||||||
|
};
|
||||||
|
lua.create_sequence_from(segments)
|
||||||
|
})
|
||||||
|
.unwrap();
|
||||||
|
lua.globals().set("get_env", get_env).unwrap();
|
||||||
|
|
||||||
|
let value = lua.load(src).eval::<Value>()?;
|
||||||
|
let t = Mirror::from_lua(value, &lua);
|
||||||
|
println!("读取出的数据:{:?}", t.clone()?);
|
||||||
|
t
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn parses_basic_config() {
|
||||||
|
let cfg = parse(
|
||||||
|
r#"
|
||||||
|
return {
|
||||||
|
target = "C:/tools/git.exe",
|
||||||
|
args = { "--no-pager",2},
|
||||||
|
env = {
|
||||||
|
PATH = { "C:/tools/git/bin", "C:/Windows", get_env("PATH")},
|
||||||
|
HOME = "C:/tools/home",
|
||||||
|
CONST = 3,
|
||||||
|
BOOL = true,
|
||||||
|
},
|
||||||
|
aliases = {
|
||||||
|
-- 1. 整体字符串(会自动 trim 首尾空格,保留为单个参数)
|
||||||
|
st = "status -s",
|
||||||
|
|
||||||
|
-- 2. 连续数组表 (Sequence Table)
|
||||||
|
lg = { "log", "--oneline", "-n", 10 },
|
||||||
|
|
||||||
|
-- 3. 标量数字与布尔值支持
|
||||||
|
v = 1,
|
||||||
|
quiet = true,
|
||||||
|
|
||||||
|
-- 4. 嵌套别名组合(加载期会自动展开并进行死环检测)
|
||||||
|
base_log = "log --graph",
|
||||||
|
all_log = { "mr:base_log", "--all","D:\\CNWei\\CNW\\Rust","D:/CNWei/CNW/Rust/" ,[[D:\CNWei\CNW\Rust\]]},
|
||||||
|
|
||||||
|
-- 5. nil 或空字符串(解析为空参数列表)
|
||||||
|
empty_alias = nil,
|
||||||
|
blank = " "
|
||||||
|
}
|
||||||
|
}
|
||||||
|
"#,
|
||||||
|
)
|
||||||
|
.unwrap();
|
||||||
|
assert_eq!(cfg.target, PathBuf::from("C:/tools/git.exe"));
|
||||||
|
// assert_eq!(cfg.args, vec!["--no-pager", "2"]);
|
||||||
|
assert_eq!(cfg.env.get("HOME").unwrap().to_str(), Some("C:/tools/home"));
|
||||||
|
// PATH 前缀来自配置,随后附加 get_env("PATH") 拆出的宿主 PATH 段
|
||||||
|
let path = cfg.env.get("PATH").unwrap().to_str().unwrap();
|
||||||
|
assert!(
|
||||||
|
path == "C:/tools/git/bin;C:/Windows"
|
||||||
|
|| path.starts_with("C:/tools/git/bin;C:/Windows;"),
|
||||||
|
"unexpected PATH: {path}"
|
||||||
|
);
|
||||||
|
// ==================== aliases 断言校验 ====================
|
||||||
|
|
||||||
|
// 1. 整体字符串:保留原样(trim 后),不切分空格
|
||||||
|
assert_eq!(
|
||||||
|
cfg.aliases.get("st").unwrap(),
|
||||||
|
&vec![OsString::from("status"), OsString::from("-s")]
|
||||||
|
);
|
||||||
|
|
||||||
|
// 2. 连续数组表:按顺序转为 OsString 列表
|
||||||
|
assert_eq!(
|
||||||
|
cfg.aliases.get("lg").unwrap(),
|
||||||
|
&vec![
|
||||||
|
OsString::from("log"),
|
||||||
|
OsString::from("--oneline"),
|
||||||
|
OsString::from("-n"),
|
||||||
|
OsString::from("10")
|
||||||
|
]
|
||||||
|
);
|
||||||
|
|
||||||
|
// 3. 标量数字与布尔值支持
|
||||||
|
assert_eq!(cfg.aliases.get("v").unwrap(), &vec![OsString::from("1")]);
|
||||||
|
assert_eq!(
|
||||||
|
cfg.aliases.get("quiet").unwrap(),
|
||||||
|
&vec![OsString::from("true")]
|
||||||
|
);
|
||||||
|
|
||||||
|
// 4. 嵌套别名组合:加载期递归拓扑打平
|
||||||
|
// base_log 本身为 "log --graph"
|
||||||
|
assert_eq!(
|
||||||
|
cfg.aliases.get("base_log").unwrap(),
|
||||||
|
&vec![OsString::from("log"), OsString::from("--graph")]
|
||||||
|
);
|
||||||
|
// all_log 展开 mr:base_log 替换为 "log --graph",追加 "--all"
|
||||||
|
assert_eq!(
|
||||||
|
cfg.aliases.get("all_log").unwrap(),
|
||||||
|
&vec![
|
||||||
|
OsString::from("log"),
|
||||||
|
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
|
||||||
|
assert_eq!(cfg.aliases.get("blank").unwrap(), &Vec::<OsString>::new());
|
||||||
|
// nil 键在遍历表时会被当作空或不存在,不产生 key 或值为空 Vec
|
||||||
|
assert!(
|
||||||
|
cfg.aliases
|
||||||
|
.get("empty_alias")
|
||||||
|
.map_or(true, |v| v.is_empty())
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[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").unwrap().to_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").unwrap().to_str(), Some(""));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn expands_nested_env_array() {
|
||||||
|
// get_env("PATH") 现在返回拆分后的段数组,嵌套表应被递归展开
|
||||||
|
let cfg = parse(
|
||||||
|
r#"return { target = "t.exe", env = { PATH = { "C:/a", { "C:/b", "C:/c" } } } }"#,
|
||||||
|
)
|
||||||
|
.unwrap();
|
||||||
|
assert_eq!(
|
||||||
|
cfg.env.get("PATH").unwrap().to_str(),
|
||||||
|
Some("C:/a;C:/b;C:/c")
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[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() {
|
||||||
|
// ValueShunt::Args 允许基础标量(数字/布尔)转字符串,仅禁止嵌套表
|
||||||
|
let cfg = parse(r#"return { target = "t.exe", args = { 1, true } }"#).unwrap();
|
||||||
|
assert_eq!(cfg.args, vec!["1", "true"]);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[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());
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user