use parking_lot::RwLock;
use std::collections::HashMap;
use std::fs;
use std::path::Path;
use std::sync::Arc;
use thiserror::Error;
#[derive(Debug, Error)]
pub enum EnvError {
#[error(".env 文件读取失败: {path} — {source}")]
FileRead {
path: String,
#[source]
source: std::io::Error,
},
#[error(".env 文件解析失败: {path} — 行 {line}: {message}")]
Parse {
path: String,
line: usize,
message: String,
},
}
#[derive(Debug, Clone, Default)]
pub struct Env {
data: Arc<RwLock<HashMap<String, String>>>,
}
impl Env {
pub fn new() -> Self {
Self::default()
}
pub fn load_from_file(&self, path: impl AsRef<Path>) -> Result<(), EnvError> {
let path_ref = path.as_ref();
let content = fs::read_to_string(path_ref).map_err(|e| EnvError::FileRead {
path: path_ref.display().to_string(),
source: e,
})?;
self.parse_ini_content(&content, &path_ref.display().to_string())
}
fn parse_ini_content(&self, content: &str, path: &str) -> Result<(), EnvError> {
let mut data = self.data.write();
let mut current_section: String = String::new();
for (line_idx, raw_line) in content.lines().enumerate() {
let line_no = line_idx + 1;
let line = raw_line.trim();
if line.is_empty() {
continue;
}
if line.starts_with('#') || line.starts_with(';') {
continue;
}
if line.starts_with('[') {
if let Some(end) = line.find(']') {
current_section = line[1..end].trim().to_string();
} else {
return Err(EnvError::Parse {
path: path.to_string(),
line: line_no,
message: "section 头缺少闭合的 ']'".to_string(),
});
}
continue;
}
if let Some(eq_pos) = line.find('=') {
let key = line[..eq_pos].trim().to_string();
let mut value = line[eq_pos + 1..].trim().to_string();
if key.is_empty() {
return Err(EnvError::Parse {
path: path.to_string(),
line: line_no,
message: "键为空".to_string(),
});
}
if value.len() >= 2 {
let first = value.chars().next().expect("已检查 value.len() >= 2");
let last = value.chars().last().expect("已检查 value.len() >= 2");
if (first == '"' && last == '"') || (first == '\'' && last == '\'') {
value = value[1..value.len() - 1].to_string();
}
}
let full_key = if current_section.is_empty() {
key
} else {
format!("{}.{}", current_section, key)
};
data.insert(full_key, value);
} else {
return Err(EnvError::Parse {
path: path.to_string(),
line: line_no,
message: "缺少 '=' 分隔符".to_string(),
});
}
}
Ok(())
}
pub fn get(&self, name: &str) -> Option<String> {
if let Ok(value) = std::env::var(name) {
if !value.is_empty() {
return Some(value);
}
}
let data = self.data.read();
data.get(name).cloned()
}
pub fn get_with_default(&self, name: &str, default: &str) -> String {
self.get(name).unwrap_or_else(|| default.to_string())
}
pub fn has(&self, name: &str) -> bool {
if let Ok(value) = std::env::var(name) {
if !value.is_empty() {
return true;
}
}
let data = self.data.read();
data.contains_key(name)
}
pub fn set(&self, name: &str, value: &str) {
let mut data = self.data.write();
data.insert(name.to_string(), value.to_string());
}
pub fn remove(&self, name: &str) -> bool {
let mut data = self.data.write();
data.remove(name).is_some()
}
pub fn all(&self) -> HashMap<String, String> {
let data = self.data.read();
data.clone()
}
pub fn clear(&self) {
let mut data = self.data.write();
data.clear();
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::io::Write;
#[test]
fn test_new_env_is_empty() {
let env = Env::new();
assert!(env.all().is_empty());
assert!(!env.has("NON_EXISTENT_KEY"));
assert_eq!(env.get("NON_EXISTENT_KEY"), None);
}
#[test]
fn test_set_get_remove() {
let env = Env::new();
env.set("APP_KEY", "base64:xxxxxx");
assert!(env.has("APP_KEY"));
assert_eq!(env.get("APP_KEY"), Some("base64:xxxxxx".to_string()));
assert!(env.remove("APP_KEY"));
assert!(!env.has("APP_KEY"));
assert_eq!(env.get("APP_KEY"), None);
}
#[test]
fn test_get_with_default() {
let env = Env::new();
assert_eq!(env.get_with_default("MISSING", "fallback"), "fallback");
env.set("EXISTING", "actual");
assert_eq!(env.get_with_default("EXISTING", "fallback"), "actual");
}
#[test]
fn test_load_from_ini_content_with_section() {
let env = Env::new();
let content = r#"
# 顶层配置
APP_DEBUG = true
APP_KEY = "base64:secret"
[database]
hostname = localhost
port = 3306
[redis]
host = "127.0.0.1"
"#;
env.parse_ini_content(content, "<test>").unwrap();
assert_eq!(env.get("APP_DEBUG"), Some("true".to_string()));
assert_eq!(env.get("APP_KEY"), Some("base64:secret".to_string()));
assert_eq!(env.get("database.hostname"), Some("localhost".to_string()));
assert_eq!(env.get("database.port"), Some("3306".to_string()));
assert_eq!(env.get("redis.host"), Some("127.0.0.1".to_string()));
}
#[test]
fn test_quote_stripping() {
let env = Env::new();
let content = r#"
DOUBLE = "value with spaces"
SINGLE = 'another value'
NO_QUOTE = plain
EMPTY = ""
"#;
env.parse_ini_content(content, "<test>").unwrap();
assert_eq!(env.get("DOUBLE"), Some("value with spaces".to_string()));
assert_eq!(env.get("SINGLE"), Some("another value".to_string()));
assert_eq!(env.get("NO_QUOTE"), Some("plain".to_string()));
assert_eq!(env.get("EMPTY"), Some("".to_string()));
}
#[test]
fn test_comment_lines_skipped() {
let env = Env::new();
let content = r#"
# 这是注释
APP_KEY = value1
; 这也是注释
APP_DEBUG = value2
"#;
env.parse_ini_content(content, "<test>").unwrap();
assert_eq!(env.get("APP_KEY"), Some("value1".to_string()));
assert_eq!(env.get("APP_DEBUG"), Some("value2".to_string()));
}
#[test]
fn test_load_from_file() {
let temp_dir = std::env::temp_dir().join("sz_rust_env_test");
let _ = std::fs::create_dir_all(&temp_dir);
let env_file = temp_dir.join(".env");
let mut file = std::fs::File::create(&env_file).unwrap();
writeln!(file, "TEST_KEY = test_value").unwrap();
writeln!(file).unwrap();
writeln!(file, "[section]").unwrap();
writeln!(file, "inner = inner_value").unwrap();
drop(file);
let env = Env::new();
env.load_from_file(&env_file).unwrap();
assert_eq!(env.get("TEST_KEY"), Some("test_value".to_string()));
assert_eq!(env.get("section.inner"), Some("inner_value".to_string()));
let _ = std::fs::remove_dir_all(&temp_dir);
}
#[test]
fn test_load_nonexistent_file_errors() {
let env = Env::new();
let result = env.load_from_file("/nonexistent/path/.env");
assert!(result.is_err());
match result {
Err(EnvError::FileRead { .. }) => {}
_ => panic!("期望 FileRead 错误"),
}
}
#[test]
fn test_parse_unclosed_section_errors() {
let env = Env::new();
let content = "[unclosed_section\nkey = value";
let result = env.parse_ini_content(content, "<test>");
assert!(result.is_err());
match result {
Err(EnvError::Parse { line, .. }) => {
assert_eq!(line, 1);
}
_ => panic!("期望 Parse 错误"),
}
}
#[test]
fn test_parse_missing_equals_errors() {
let env = Env::new();
let content = "this_is_not_a_key_value_pair";
let result = env.parse_ini_content(content, "<test>");
assert!(result.is_err());
match result {
Err(EnvError::Parse { line, .. }) => {
assert_eq!(line, 1);
}
_ => panic!("期望 Parse 错误"),
}
}
#[test]
fn test_parse_empty_key_errors() {
let env = Env::new();
let content = " = value";
let result = env.parse_ini_content(content, "<test>");
assert!(result.is_err());
match result {
Err(EnvError::Parse { line, .. }) => {
assert_eq!(line, 1);
}
_ => panic!("期望 Parse 错误"),
}
}
#[test]
fn test_process_env_takes_priority() {
let env = Env::new();
env.set("SZ_RUST_TEST_ENV_PRIORITY", "internal_value");
std::env::set_var("SZ_RUST_TEST_ENV_PRIORITY", "process_value");
assert_eq!(
env.get("SZ_RUST_TEST_ENV_PRIORITY"),
Some("process_value".to_string())
);
std::env::remove_var("SZ_RUST_TEST_ENV_PRIORITY");
}
#[test]
fn test_empty_process_env_falls_back_to_internal() {
let env = Env::new();
env.set("SZ_RUST_TEST_EMPTY_FALLBACK", "internal_value");
std::env::set_var("SZ_RUST_TEST_EMPTY_FALLBACK", "");
assert_eq!(
env.get("SZ_RUST_TEST_EMPTY_FALLBACK"),
Some("internal_value".to_string())
);
std::env::remove_var("SZ_RUST_TEST_EMPTY_FALLBACK");
}
#[test]
fn test_clear() {
let env = Env::new();
env.set("KEY1", "value1");
env.set("KEY2", "value2");
assert_eq!(env.all().len(), 2);
env.clear();
assert!(env.all().is_empty());
}
#[test]
fn test_all_returns_snapshot() {
let env = Env::new();
env.set("KEY1", "value1");
env.set("KEY2", "value2");
let snapshot = env.all();
assert_eq!(snapshot.len(), 2);
assert_eq!(snapshot.get("KEY1"), Some(&"value1".to_string()));
assert_eq!(snapshot.get("KEY2"), Some(&"value2".to_string()));
env.set("KEY3", "value3");
assert_eq!(snapshot.len(), 2);
}
#[test]
fn test_remove_nonexistent_returns_false() {
let env = Env::new();
assert!(!env.remove("NON_EXISTENT"));
}
#[test]
fn test_section_isolation() {
let env = Env::new();
let content = r#"
[section1]
key = value1
[section2]
key = value2
"#;
env.parse_ini_content(content, "<test>").unwrap();
assert_eq!(env.get("section1.key"), Some("value1".to_string()));
assert_eq!(env.get("section2.key"), Some("value2".to_string()));
}
#[test]
fn test_multiple_load_accumulates() {
let env = Env::new();
let content1 = "KEY1 = value1";
let content2 = "KEY2 = value2";
env.parse_ini_content(content1, "<test1>").unwrap();
env.parse_ini_content(content2, "<test2>").unwrap();
assert_eq!(env.get("KEY1"), Some("value1".to_string()));
assert_eq!(env.get("KEY2"), Some("value2".to_string()));
}
}