use std::ffi::{OsStr, c_void};
use std::io;
use std::os::windows::ffi::OsStrExt as _;
use std::path::Path;
use windows_sys::Win32::Foundation::{CloseHandle, ERROR_SUCCESS, FALSE, HANDLE, LocalFree};
use windows_sys::Win32::Security::Authorization::{
ConvertSidToStringSidW, ConvertStringSecurityDescriptorToSecurityDescriptorW,
GetNamedSecurityInfoW, SDDL_REVISION_1, SE_FILE_OBJECT, SetNamedSecurityInfoW,
};
use windows_sys::Win32::Security::{
ACCESS_ALLOWED_ACE, ACL, ACL_SIZE_INFORMATION, AclSizeInformation, DACL_SECURITY_INFORMATION,
EqualSid, GetAce, GetAclInformation, GetSecurityDescriptorDacl,
PROTECTED_DACL_SECURITY_INFORMATION, PSECURITY_DESCRIPTOR, PSID, TOKEN_QUERY, TOKEN_USER,
TokenUser,
};
use windows_sys::Win32::System::Threading::{GetCurrentProcess, OpenProcessToken};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Inheritance {
ObjectOnly,
Propagating,
}
impl Inheritance {
fn ace_flags(self) -> &'static str {
match self {
Self::ObjectOnly => "",
Self::Propagating => "OICI",
}
}
}
#[derive(Debug)]
struct LocalPtr(*mut c_void);
impl Drop for LocalPtr {
fn drop(&mut self) {
if !self.0.is_null() {
unsafe {
LocalFree(self.0);
}
}
}
}
#[derive(Debug)]
struct TokenHandle(HANDLE);
impl Drop for TokenHandle {
fn drop(&mut self) {
if !self.0.is_null() {
unsafe {
CloseHandle(self.0);
}
}
}
}
fn wide(path: &OsStr) -> Vec<u16> {
path.encode_wide().chain(std::iter::once(0)).collect()
}
fn current_user_sid_buffer() -> io::Result<Vec<u64>> {
let mut raw_token: HANDLE = std::ptr::null_mut();
let ok = unsafe { OpenProcessToken(GetCurrentProcess(), TOKEN_QUERY, &raw mut raw_token) };
if ok == FALSE {
return Err(io::Error::last_os_error());
}
let token = TokenHandle(raw_token);
let mut needed: u32 = 0;
unsafe {
windows_sys::Win32::Security::GetTokenInformation(
token.0,
TokenUser,
std::ptr::null_mut(),
0,
&raw mut needed,
);
}
if needed == 0 {
return Err(io::Error::last_os_error());
}
let words = (needed as usize).div_ceil(size_of::<u64>());
let mut buffer = vec![0u64; words];
let ok = unsafe {
windows_sys::Win32::Security::GetTokenInformation(
token.0,
TokenUser,
buffer.as_mut_ptr().cast(),
needed,
&raw mut needed,
)
};
if ok == FALSE {
return Err(io::Error::last_os_error());
}
Ok(buffer)
}
unsafe fn sid_from_buffer(buffer: &[u64]) -> PSID {
let token_user = unsafe { &*buffer.as_ptr().cast::<TOKEN_USER>() };
token_user.User.Sid
}
fn current_user_sid_string() -> io::Result<String> {
let buffer = current_user_sid_buffer()?;
let sid = unsafe { sid_from_buffer(&buffer) };
let mut raw: *mut u16 = std::ptr::null_mut();
let ok = unsafe { ConvertSidToStringSidW(sid, &raw mut raw) };
if ok == FALSE {
return Err(io::Error::last_os_error());
}
let owned = LocalPtr(raw.cast());
let text = unsafe {
let mut len = 0usize;
while *raw.add(len) != 0 {
len += 1;
}
String::from_utf16_lossy(std::slice::from_raw_parts(raw, len))
};
drop(owned);
Ok(text)
}
pub fn restrict_to_current_user(path: &Path, inheritance: Inheritance) -> io::Result<()> {
let sid = current_user_sid_string()?;
let sddl = format!("D:P(A;{};FA;;;{sid})", inheritance.ace_flags());
let sddl_wide = wide(OsStr::new(&sddl));
let mut descriptor: PSECURITY_DESCRIPTOR = std::ptr::null_mut();
let ok = unsafe {
ConvertStringSecurityDescriptorToSecurityDescriptorW(
sddl_wide.as_ptr(),
SDDL_REVISION_1,
&raw mut descriptor,
std::ptr::null_mut(),
)
};
if ok == FALSE {
return Err(io::Error::last_os_error());
}
let owned = LocalPtr(descriptor.cast());
let mut present: i32 = 0;
let mut dacl: *mut ACL = std::ptr::null_mut();
let mut defaulted: i32 = 0;
let ok = unsafe {
GetSecurityDescriptorDacl(descriptor, &raw mut present, &raw mut dacl, &raw mut defaulted)
};
if ok == FALSE {
return Err(io::Error::last_os_error());
}
if present == FALSE || dacl.is_null() {
return Err(io::Error::other(
"the generated security descriptor has no access control list",
));
}
let mut object = wide(path.as_os_str());
let status = unsafe {
SetNamedSecurityInfoW(
object.as_mut_ptr(),
SE_FILE_OBJECT,
DACL_SECURITY_INFORMATION | PROTECTED_DACL_SECURITY_INFORMATION,
std::ptr::null_mut(),
std::ptr::null_mut(),
dacl,
std::ptr::null_mut(),
)
};
drop(owned);
if status != ERROR_SUCCESS {
return Err(io::Error::from_raw_os_error(i32::try_from(status).unwrap_or(i32::MAX)));
}
Ok(())
}
pub fn is_restricted_to_current_user(path: &Path) -> io::Result<bool> {
let sid_buffer = current_user_sid_buffer()?;
let expected = unsafe { sid_from_buffer(&sid_buffer) };
let mut dacl: *mut ACL = std::ptr::null_mut();
let mut descriptor: PSECURITY_DESCRIPTOR = std::ptr::null_mut();
let mut object = wide(path.as_os_str());
let status = unsafe {
GetNamedSecurityInfoW(
object.as_mut_ptr(),
SE_FILE_OBJECT,
DACL_SECURITY_INFORMATION,
std::ptr::null_mut(),
std::ptr::null_mut(),
&raw mut dacl,
std::ptr::null_mut(),
&raw mut descriptor,
)
};
if status != ERROR_SUCCESS {
return Err(io::Error::from_raw_os_error(i32::try_from(status).unwrap_or(i32::MAX)));
}
let _owned = LocalPtr(descriptor.cast());
if dacl.is_null() {
return Ok(false);
}
let mut info: ACL_SIZE_INFORMATION = unsafe { std::mem::zeroed() };
let size = u32::try_from(size_of::<ACL_SIZE_INFORMATION>())
.map_err(|_| io::Error::other("access control list size structure is implausibly large"))?;
let ok = unsafe { GetAclInformation(dacl, (&raw mut info).cast(), size, AclSizeInformation) };
if ok == FALSE {
return Err(io::Error::last_os_error());
}
if info.AceCount == 0 {
return Ok(false);
}
for index in 0..info.AceCount {
let mut ace: *mut c_void = std::ptr::null_mut();
let ok = unsafe { GetAce(dacl, index, &raw mut ace) };
if ok == FALSE {
return Err(io::Error::last_os_error());
}
let sid = unsafe {
let ace = &*ace.cast::<ACCESS_ALLOWED_ACE>();
(&raw const ace.SidStart).cast::<c_void>().cast_mut()
};
let same = unsafe { EqualSid(sid, expected) };
if same == FALSE {
return Ok(false);
}
}
Ok(true)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn the_current_user_has_a_readable_sid() {
let sid = current_user_sid_string().expect("the process has a user");
assert!(sid.starts_with("S-1-"), "not a SID: {sid}");
}
#[test]
fn the_sddl_says_protected_and_full_access() {
let sid = current_user_sid_string().expect("the process has a user");
let sddl = format!("D:P(A;{};FA;;;{sid})", Inheritance::ObjectOnly.ace_flags());
assert!(sddl.starts_with("D:P("), "the DACL must be protected: {sddl}");
assert!(sddl.contains(";FA;"));
assert!(sddl.contains(&sid));
}
#[test]
fn a_directory_entry_propagates_and_a_file_entry_does_not() {
assert_eq!(Inheritance::ObjectOnly.ace_flags(), "");
assert_eq!(Inheritance::Propagating.ace_flags(), "OICI");
}
#[test]
fn a_restricted_file_verifies_and_an_untouched_one_does_not() {
let dir = tempfile::tempdir().expect("a temporary directory");
let loose = dir.path().join("loose.json");
std::fs::write(&loose, b"{}").expect("write");
assert!(
!is_restricted_to_current_user(&loose).expect("reading the DACL must succeed"),
"an inherited DACL was wrongly reported as owner-only"
);
let tight = dir.path().join("tight.json");
std::fs::write(&tight, b"{}").expect("write");
restrict_to_current_user(&tight, Inheritance::ObjectOnly)
.expect("restricting must succeed");
assert!(
is_restricted_to_current_user(&tight).expect("reading the DACL must succeed"),
"a restricted file was not recognised as owner-only"
);
}
#[test]
fn restricting_a_missing_path_is_an_error_rather_than_a_silent_success() {
let dir = tempfile::tempdir().expect("a temporary directory");
let missing = dir.path().join("does-not-exist");
assert!(restrict_to_current_user(&missing, Inheritance::ObjectOnly).is_err());
}
#[test]
fn a_restricted_directory_still_accepts_new_files() {
let dir = tempfile::tempdir().expect("a temporary directory");
let inner = dir.path().join("store");
std::fs::create_dir(&inner).expect("create");
restrict_to_current_user(&inner, Inheritance::Propagating).expect("restrict");
let file = inner.join("tokens.json");
std::fs::write(&file, b"{}").expect("the owner must still be able to write");
assert!(is_restricted_to_current_user(&inner).expect("read"));
}
}