use std::{io, sync::RwLock};
use windows::{
core::{HSTRING, PCWSTR},
Win32::{
Foundation::{LPARAM, WPARAM},
UI::WindowsAndMessaging::{
SendMessageTimeoutW, HWND_BROADCAST, SMTO_ABORTIFHUNG, WM_SETTINGCHANGE,
},
},
};
use winreg::{
enums::{HKEY_CURRENT_USER, KEY_READ, KEY_WRITE},
RegKey,
};
static LOCK: RwLock<()> = RwLock::new(());
fn regkey() -> io::Result<RegKey> {
let hkcu = RegKey::predef(HKEY_CURRENT_USER);
hkcu.open_subkey_with_flags("Environment", KEY_READ | KEY_WRITE)
}
fn check_separator(value: &str) -> io::Result<()> {
if value.contains(';') {
Err(io::Error::new(
io::ErrorKind::InvalidInput,
"value contains `;`",
))
} else {
Ok(())
}
}
pub fn append<T1, T2>(var: T1, value: T2) -> io::Result<()>
where
T1: AsRef<str>,
T2: AsRef<str>,
{
add_inner(var.as_ref(), value.as_ref(), false)
}
pub fn prepend<T1, T2>(var: T1, value: T2) -> io::Result<()>
where
T1: AsRef<str>,
T2: AsRef<str>,
{
add_inner(var.as_ref(), value.as_ref(), true)
}
fn add_inner(var: &str, value: &str, front: bool) -> io::Result<()> {
check_separator(value)?;
let _lock = LOCK.write().unwrap();
let env = regkey()?;
let get_res = env.get_value(var);
let env_var: String = match get_res {
Ok(s) => s,
Err(err) if err.kind() == io::ErrorKind::NotFound => String::default(),
Err(err) => return Err(err),
};
let mut values = env_var
.split(';')
.filter(|x| !x.is_empty())
.collect::<Vec<&str>>();
if !values.contains(&value) {
if front {
values.insert(0, value);
} else {
values.push(value);
}
let new_env_var = values.join(";");
env.set_value(var, &new_env_var)?;
unsafe { std::env::set_var(var, &new_env_var) };
notify_system();
}
Ok(())
}
pub fn remove_from_list(var: &str, value: &str) -> io::Result<bool> {
check_separator(value)?;
let _lock = LOCK.write().unwrap();
let env = regkey()?;
let get_res = env.get_value(var);
let env_var: String = match get_res {
Ok(s) => s,
Err(err) if err.kind() == io::ErrorKind::NotFound => return Ok(false),
Err(err) => return Err(err),
};
let mut values = env_var.split(';').collect::<Vec<&str>>();
let len = values.len();
values.retain(|p| p != &value);
let found = len != values.len();
let new_env_var = values.join(";");
env.set_value(var, &new_env_var)?;
unsafe { std::env::set_var(var, &new_env_var) };
notify_system();
Ok(found)
}
pub fn exists_in_list(var: &str, value: &str) -> io::Result<bool> {
check_separator(value)?;
let env_var = get(var)?;
match env_var {
Some(s) => Ok(s.split(';').any(|p| p == value)),
None => Ok(false),
}
}
pub fn set<T1: AsRef<str>, T2: AsRef<str>>(var: T1, value: T2) -> io::Result<()> {
let _lock = LOCK.write().unwrap();
let env = regkey()?;
env.set_value(var.as_ref(), &value.as_ref())?;
unsafe { std::env::set_var(var.as_ref(), value.as_ref()) };
notify_system();
Ok(())
}
pub fn get<T: AsRef<str>>(var: T) -> io::Result<Option<String>> {
let _lock = LOCK.read().unwrap();
let env = regkey()?;
let res = env.get_value(var.as_ref());
match res {
Ok(s) => Ok(Some(s)),
Err(err) if err.kind() == io::ErrorKind::NotFound => Ok(None),
Err(err) => Err(err),
}
}
pub fn remove<T: AsRef<str>>(var: T) -> io::Result<()> {
let _lock = LOCK.write().unwrap();
let env = regkey()?;
if let Err(err) = env.delete_value(var.as_ref()) {
if err.kind() != io::ErrorKind::NotFound {
return Err(err);
}
};
unsafe { std::env::remove_var(var.as_ref()) };
notify_system();
Ok(())
}
fn w<T: Into<HSTRING>>(x: T) -> PCWSTR {
PCWSTR::from_raw(x.into().as_ptr())
}
fn notify_system() {
let msg = w("Environment");
unsafe {
SendMessageTimeoutW(
HWND_BROADCAST,
WM_SETTINGCHANGE,
WPARAM(0),
LPARAM(msg.as_ptr() as isize),
SMTO_ABORTIFHUNG,
500,
None,
);
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_get_set() -> Result<(), Box<dyn std::error::Error>> {
const ENV_VAR: &str = "TEST-GET-SET";
set(ENV_VAR, "test")?;
assert_eq!(get(ENV_VAR)?.unwrap(), "test");
remove(ENV_VAR)?;
assert!(get(ENV_VAR)?.is_none());
Ok(())
}
#[test]
fn test_list_operations() -> Result<(), Box<dyn std::error::Error>> {
const ENV_VAR: &str = "TEST-LIST-OPERATIONS";
set(ENV_VAR, "test1;test2;te")?;
assert!(exists_in_list(ENV_VAR, "test1")?);
assert!(exists_in_list(ENV_VAR, "test2")?);
assert!(exists_in_list(ENV_VAR, "te")?);
append(ENV_VAR, "st3")?;
assert_eq!(get(ENV_VAR)?.unwrap(), "test1;test2;te;st3");
prepend(ENV_VAR, "st4")?;
assert_eq!(get(ENV_VAR)?.unwrap(), "st4;test1;test2;te;st3");
remove_from_list(ENV_VAR, "test1")?;
assert_eq!(get(ENV_VAR)?.unwrap(), "st4;test2;te;st3");
assert!(!exists_in_list(ENV_VAR, "test1")?);
remove(ENV_VAR)?;
Ok(())
}
#[test]
fn test_reset_one_var() -> Result<(), Box<dyn std::error::Error>> {
const ENV_VAR: &str = "TEST-RESET-ONE-VAR";
set(ENV_VAR, "test")?;
assert_eq!(get(ENV_VAR)?.unwrap(), "test");
set(ENV_VAR, "new_test")?;
assert_eq!(get(ENV_VAR)?.unwrap(), "new_test");
remove(ENV_VAR)?;
Ok(())
}
#[test]
fn test_operate_with_not_exist_var() -> Result<(), Box<dyn std::error::Error>> {
const NOT_EXIST: &str = "A_VAR_DOES_NOT_EXIST";
remove(NOT_EXIST)?;
assert!(get(NOT_EXIST)?.is_none());
assert!(!exists_in_list(NOT_EXIST, "test")?);
assert!(!remove_from_list(NOT_EXIST, "test")?);
append(NOT_EXIST, "test")?;
assert_eq!(get(NOT_EXIST)?.unwrap(), "test");
remove(NOT_EXIST)?;
prepend(NOT_EXIST, "test")?;
assert_eq!(get(NOT_EXIST)?.unwrap(), "test");
remove(NOT_EXIST)?;
Ok(())
}
#[test]
fn test_operation_will_affect_current_process() -> Result<(), Box<dyn std::error::Error>> {
let env_var = "TEST-OPERATION-WILL-AFFECT-CURRENT-PROCESS";
set(env_var, "test")?;
assert_eq!(std::env::var(env_var)?, "test");
remove(env_var)?;
assert_eq!(std::env::var(env_var), Err(std::env::VarError::NotPresent));
Ok(())
}
#[test]
fn test_invalid_value() -> Result<(), Box<dyn std::error::Error>> {
let env_var = "TEST-INVALID-VALUE";
assert!(append(env_var, "123;456").is_err());
assert!(prepend(env_var, "123;456").is_err());
assert!(remove_from_list(env_var, "123;456").is_err());
Ok(())
}
}