use std::{io, path::PathBuf, slice};
use windows::{
Win32::{
Foundation::ERROR_CANCELLED,
Globalization::lstrlenW,
System::Com::{
CLSCTX_INPROC_SERVER, COINIT_APARTMENTTHREADED, CoCreateInstance, CoInitializeEx,
CoTaskMemFree, CoUninitialize,
},
UI::Shell::{
FOS_FORCEFILESYSTEM, FOS_PICKFOLDERS, FileOpenDialog, IFileOpenDialog,
SIGDN_FILESYSPATH,
},
},
core::{HRESULT, PCWSTR, PWSTR},
};
pub const MAX_DIALOG_PATH_UNITS: usize = 32_767;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum DialogSelection {
File,
Folder,
}
pub fn pick(selection: DialogSelection) -> io::Result<Option<PathBuf>> {
let _apartment = ComApartment::initialize()?;
let dialog: IFileOpenDialog =
unsafe { CoCreateInstance(&FileOpenDialog, None, CLSCTX_INPROC_SERVER) }
.map_err(windows_error)?;
let mut options = unsafe { dialog.GetOptions() }.map_err(windows_error)?;
options |= FOS_FORCEFILESYSTEM;
if selection == DialogSelection::Folder {
options |= FOS_PICKFOLDERS;
}
unsafe { dialog.SetOptions(options) }.map_err(windows_error)?;
let shown = unsafe { dialog.Show(None) };
if let Err(error) = shown {
if error.code() == HRESULT::from_win32(ERROR_CANCELLED.0) {
return Ok(None);
}
return Err(windows_error(error));
}
let item = unsafe { dialog.GetResult() }.map_err(windows_error)?;
let raw_path = unsafe { item.GetDisplayName(SIGDN_FILESYSPATH) }.map_err(windows_error)?;
TaskMemoryPath::new(raw_path).to_path().map(Some)
}
struct ComApartment {
initialized: bool,
}
impl ComApartment {
fn initialize() -> io::Result<Self> {
let result = unsafe { CoInitializeEx(None, COINIT_APARTMENTTHREADED) };
if result.is_ok() {
Ok(Self { initialized: true })
} else {
Err(windows_error(result.into()))
}
}
}
impl Drop for ComApartment {
fn drop(&mut self) {
if self.initialized {
unsafe { CoUninitialize() };
}
}
}
struct TaskMemoryPath(PWSTR);
impl TaskMemoryPath {
fn new(value: PWSTR) -> Self {
Self(value)
}
fn to_path(&self) -> io::Result<PathBuf> {
if self.0.0.is_null() {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"Windows common dialog returned a null path",
));
}
let length = unsafe { lstrlenW(PCWSTR(self.0.0)) };
if length < 0 {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"Windows common dialog returned a negative path length",
));
}
let length = usize::try_from(length).map_err(|_| {
io::Error::new(
io::ErrorKind::InvalidData,
"Windows common dialog path length overflowed usize",
)
})?;
if length > MAX_DIALOG_PATH_UNITS {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"Windows common dialog path exceeds the bounded UTF-16 limit",
));
}
let units = unsafe { slice::from_raw_parts(self.0.0, length) };
let mut value = String::new();
value
.try_reserve_exact(length)
.map_err(|_| io::Error::new(io::ErrorKind::OutOfMemory, "path allocation failed"))?;
for character in char::decode_utf16(units.iter().copied()) {
value.push(character.map_err(|_| {
io::Error::new(
io::ErrorKind::InvalidData,
"Windows common dialog path contains invalid UTF-16",
)
})?);
}
Ok(PathBuf::from(value))
}
}
impl Drop for TaskMemoryPath {
fn drop(&mut self) {
if !self.0.0.is_null() {
unsafe { CoTaskMemFree(Some(self.0.0.cast())) };
}
}
}
fn windows_error(error: windows::core::Error) -> io::Error {
io::Error::from_raw_os_error(error.code().0)
}
#[cfg(test)]
mod tests {
use super::*;
use windows::Win32::System::Com::CoTaskMemAlloc;
#[test]
fn selection_kinds_are_explicit() {
assert_ne!(DialogSelection::File, DialogSelection::Folder);
}
#[test]
fn a_task_memory_path_decodes_to_the_buffer_it_owns() {
let expected = PathBuf::from("C:\\Moirai\\\u{e9}t\u{e9}.txt");
let units: Vec<u16> = expected
.to_str()
.expect("the fixture is UTF-8")
.encode_utf16()
.chain(std::iter::once(0))
.collect();
let buffer = unsafe { CoTaskMemAlloc(units.len() * size_of::<u16>()) }.cast::<u16>();
assert!(!buffer.is_null());
unsafe { std::ptr::copy_nonoverlapping(units.as_ptr(), buffer, units.len()) };
let owner = TaskMemoryPath::new(PWSTR(buffer));
assert_eq!(owner.to_path().expect("valid path"), expected);
assert_eq!(
owner
.to_path()
.expect("decoding leaves the buffer readable"),
expected
);
}
#[test]
fn cancellation_uses_the_windows_hresult() {
assert_eq!(HRESULT::from_win32(ERROR_CANCELLED.0).0 as u32, 0x8007_04c7);
}
}