use std::sync::atomic::{AtomicUsize, Ordering};
use windows::Win32::Foundation::{
CLASS_E_CLASSNOTAVAILABLE, E_POINTER, HMODULE, MAX_PATH, S_FALSE, S_OK,
};
use windows::Win32::System::LibraryLoader::{
GET_MODULE_HANDLE_EX_FLAG_FROM_ADDRESS, GET_MODULE_HANDLE_EX_FLAG_UNCHANGED_REFCOUNT,
GetModuleFileNameW, GetModuleHandleExW,
};
use windows::Win32::System::Registry::{
HKEY, HKEY_CLASSES_ROOT, KEY_WRITE, REG_OPTION_NON_VOLATILE, REG_SZ, RegCloseKey,
RegCreateKeyExW, RegDeleteKeyW, RegSetValueExW,
};
use windows_core::{GUID, HRESULT, Interface, PCWSTR};
use crate::control::DenisePanel;
use crate::factory::PanelFactory;
use crate::registry;
use crate::typelib;
pub const CLSID_DENISE_PANEL: GUID = GUID::from_u128(0x7F1B_483A_5853_4348_9081_D5BD_502B_51E8);
static OUTSTANDING: AtomicUsize = AtomicUsize::new(0);
pub(crate) fn lock_server() {
OUTSTANDING.fetch_add(1, Ordering::Relaxed);
}
pub(crate) fn unlock_server() {
OUTSTANDING.fetch_sub(1, Ordering::Relaxed);
}
#[unsafe(no_mangle)]
pub unsafe extern "system" fn DllGetClassObject(
rclsid: *const GUID,
riid: *const GUID,
ppv: *mut *mut core::ffi::c_void,
) -> HRESULT {
if ppv.is_null() || rclsid.is_null() || riid.is_null() {
return E_POINTER;
}
unsafe {
ppv.write(core::ptr::null_mut());
if *rclsid != CLSID_DENISE_PANEL {
return CLASS_E_CLASSNOTAVAILABLE;
}
let factory: windows_core::IUnknown = PanelFactory.into();
factory.query(riid, ppv)
}
}
#[unsafe(no_mangle)]
pub extern "system" fn DllCanUnloadNow() -> HRESULT {
if OUTSTANDING.load(Ordering::Relaxed) == 0 {
S_OK
} else {
S_FALSE
}
}
#[unsafe(no_mangle)]
pub extern "system" fn DllRegisterServer() -> HRESULT {
match register() {
Ok(()) => S_OK,
Err(e) => e.into(),
}
}
#[unsafe(no_mangle)]
pub extern "system" fn DllUnregisterServer() -> HRESULT {
match unregister() {
Ok(()) => S_OK,
Err(e) => e.into(),
}
}
fn register() -> windows_core::Result<()> {
let path = server_path()?;
for entry in registry::entries(&path) {
write_value(&entry.key, &entry.name, &entry.value)?;
}
let tlb = typelib::path_beside(&path);
typelib::build(&tlb)?;
typelib::register(&tlb)
}
fn unregister() -> windows_core::Result<()> {
let _ = typelib::unregister();
for key in registry::keys_to_remove() {
let wide = wide(&key);
unsafe {
let _ = RegDeleteKeyW(HKEY_CLASSES_ROOT, PCWSTR(wide.as_ptr()));
}
}
Ok(())
}
fn write_value(key: &str, name: &str, value: &str) -> windows_core::Result<()> {
let key_wide = wide(key);
let name_wide = wide(name);
let value_wide = wide(value);
let mut handle = HKEY::default();
unsafe {
RegCreateKeyExW(
HKEY_CLASSES_ROOT,
PCWSTR(key_wide.as_ptr()),
None,
None,
REG_OPTION_NON_VOLATILE,
KEY_WRITE,
None,
&mut handle,
None,
)
.ok()?;
let bytes = core::slice::from_raw_parts(
value_wide.as_ptr().cast::<u8>(),
value_wide.len() * core::mem::size_of::<u16>(),
);
let result = RegSetValueExW(
handle,
if name.is_empty() {
PCWSTR::null()
} else {
PCWSTR(name_wide.as_ptr())
},
None,
REG_SZ,
Some(bytes),
);
let _ = RegCloseKey(handle);
result.ok()?;
}
Ok(())
}
fn server_path() -> windows_core::Result<String> {
let mut module = HMODULE::default();
unsafe {
GetModuleHandleExW(
GET_MODULE_HANDLE_EX_FLAG_FROM_ADDRESS | GET_MODULE_HANDLE_EX_FLAG_UNCHANGED_REFCOUNT,
PCWSTR(server_path as *const () as *const u16),
&mut module,
)?;
}
let mut buffer = [0u16; MAX_PATH as usize];
let written = unsafe { GetModuleFileNameW(Some(module), &mut buffer) };
if written == 0 {
return Err(windows_core::Error::from_thread());
}
Ok(String::from_utf16_lossy(&buffer[..written as usize]))
}
pub(crate) fn library_path() -> Option<String> {
server_path().ok().map(|dll| typelib::path_beside(&dll))
}
fn wide(text: &str) -> Vec<u16> {
text.encode_utf16().chain(core::iter::once(0)).collect()
}
const _: fn() -> DenisePanel = DenisePanel::new;
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn the_binary_and_text_class_ids_are_the_same_identity() {
let text = format!("{{{:?}}}", CLSID_DENISE_PANEL).to_uppercase();
assert_eq!(text, registry::CLSID_TEXT.to_uppercase());
}
#[test]
fn a_wide_string_is_terminated() {
let w = wide("AB");
assert_eq!(w, vec![0x41, 0x42, 0x00]);
assert_eq!(*wide("").last().expect("terminator"), 0);
}
#[test]
fn the_dll_only_unloads_with_nothing_outstanding() {
assert_eq!(DllCanUnloadNow(), S_OK);
lock_server();
assert_eq!(DllCanUnloadNow(), S_FALSE);
unlock_server();
assert_eq!(DllCanUnloadNow(), S_OK);
}
}