use crate::error::{Error, Result, check_wnet};
use crate::options::{ConnectOptions, DisconnectOptions, DriveLetter};
use crate::strings::{WideSecret, from_wide_buf, len_u32, opt_ptr, to_wide};
use tracing::{debug, trace};
use windows_sys::Win32::NetworkManagement::WNet;
struct ConnectArgs {
remote: Vec<u16>,
local: Option<Vec<u16>>,
username: Option<Vec<u16>>,
password: Option<WideSecret>,
}
impl ConnectArgs {
fn new(target: &SmbTarget) -> Result<Self> {
let password = target
.password
.as_deref()
.map(|password| to_wide(password))
.transpose()?
.map(WideSecret::from);
Ok(Self {
remote: to_wide(&target.remote)?,
local: target.local.as_deref().map(to_wide).transpose()?,
username: target.username.as_deref().map(to_wide).transpose()?,
password,
})
}
fn resource(&self) -> WNet::NETRESOURCEW {
WNet::NETRESOURCEW {
dwType: WNet::RESOURCETYPE_DISK,
lpLocalName: opt_ptr(self.local.as_deref()).cast_mut(),
lpRemoteName: self.remote.as_ptr().cast_mut(),
..Default::default()
}
}
}
pub struct SmbTarget {
remote: String,
username: Option<String>,
password: Option<zeroize::Zeroizing<String>>,
local: Option<String>,
}
impl std::fmt::Debug for SmbTarget {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("SmbTarget")
.field("remote", &self.remote)
.field("username", &self.username)
.field("password", &self.password.as_ref().map(|_| "<redacted>"))
.field("local", &self.local)
.finish()
}
}
impl SmbTarget {
pub fn new(remote: impl Into<String>) -> Self {
Self {
remote: remote.into(),
username: None,
password: None,
local: None,
}
}
#[must_use]
pub fn credentials(self, username: impl Into<String>, password: impl Into<String>) -> Self {
self.username(username).password(password)
}
#[must_use]
pub fn username(mut self, username: impl Into<String>) -> Self {
self.username = Some(username.into());
self
}
#[must_use]
pub fn password(mut self, password: impl Into<String>) -> Self {
self.password = Some(zeroize::Zeroizing::new(password.into()));
self
}
#[must_use]
pub fn mount_on(mut self, letter: DriveLetter) -> Self {
self.local = Some(letter.to_string());
self
}
pub fn connect(&self) -> Result<()> {
self.connect_with(ConnectOptions::new())
}
pub fn connect_with(&self, options: ConnectOptions) -> Result<()> {
if options.is_persistent() && self.local.is_none() {
return Err(Error::PersistenceRequiresDrive);
}
let flags = options.flags;
let args = ConnectArgs::new(self)?;
let resource = args.resource();
trace!(
"connecting to {} as {} with flags {flags:#x}",
self.remote,
self.username.as_deref().unwrap_or("<logged-on user>")
);
let status = unsafe {
WNet::WNetAddConnection2W(
&raw const resource,
opt_ptr(args.password.as_ref().map(|password| password.as_slice())),
opt_ptr(args.username.as_deref()),
flags,
)
};
debug!("WNetAddConnection2W returned {status}");
check_wnet(status)
}
pub fn connect_auto(&self, options: ConnectOptions) -> Result<String> {
let flags = options.flags | WNet::CONNECT_REDIRECT;
let args = ConnectArgs::new(self)?;
let resource = args.resource();
trace!(
"auto-connecting to {} as {} with flags {flags:#x}",
self.remote,
self.username.as_deref().unwrap_or("<logged-on user>")
);
let mut access_name = vec![0u16; 1024 + args.remote.len()];
let mut size = len_u32(access_name.len());
let mut result = 0u32;
let status = unsafe {
WNet::WNetUseConnectionW(
std::ptr::null_mut(), &raw const resource,
opt_ptr(args.password.as_ref().map(|password| password.as_slice())),
opt_ptr(args.username.as_deref()),
flags,
access_name.as_mut_ptr(),
&raw mut size,
&raw mut result,
)
};
debug!("WNetUseConnectionW returned {status}, result {result:#x}");
check_wnet(status)?;
Ok(from_wide_buf(&access_name))
}
pub fn connect_guarded(&self, options: ConnectOptions) -> Result<Connection> {
if options.is_persistent() {
return Err(Error::PersistentGuard);
}
let device = self.local.clone().ok_or(Error::GuardRequiresDrive)?;
self.connect_with(options)?;
Ok(Connection {
device,
armed: true,
})
}
pub fn connect_auto_guarded(&self, options: ConnectOptions) -> Result<Connection> {
if options.is_persistent() {
return Err(Error::PersistentGuard);
}
let device = self.connect_auto(options)?;
Ok(Connection {
device,
armed: true,
})
}
pub fn disconnect(&self) -> Result<()> {
self.disconnect_with(DisconnectOptions::default())
}
pub fn disconnect_with(&self, options: DisconnectOptions) -> Result<()> {
let name = self.local.as_deref().unwrap_or(&self.remote);
cancel_connection(name, options)
}
}
pub fn cancel_connection(name: &str, options: DisconnectOptions) -> Result<()> {
if options.forget && !is_drive(name) {
return Err(Error::ForgetRequiresDrive);
}
let wide = to_wide(name)?;
trace!(
"disconnecting {name} (force={}, forget={})",
options.force, options.forget
);
let flags = if options.forget {
WNet::CONNECT_UPDATE_PROFILE
} else {
0
};
let status =
unsafe { WNet::WNetCancelConnection2W(wide.as_ptr(), flags, i32::from(options.force)) };
debug!("WNetCancelConnection2W returned {status}");
check_wnet(status)
}
fn is_drive(name: &str) -> bool {
matches!(name.as_bytes(), [letter, b':'] if letter.is_ascii_alphabetic())
}
#[derive(Debug)]
#[must_use = "dropping the guard disconnects the connection immediately"]
pub struct Connection {
device: String,
armed: bool,
}
impl Connection {
#[must_use]
pub fn device(&self) -> &str {
&self.device
}
pub fn disconnect(mut self, options: DisconnectOptions) -> Result<()> {
self.armed = false;
cancel_connection(&self.device, options)
}
}
impl Drop for Connection {
fn drop(&mut self) {
if self.armed
&& let Err(e) = cancel_connection(&self.device, DisconnectOptions::default())
{
debug!("failed to disconnect {} on guard drop: {e}", self.device);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn debug_redacts_password() {
let target = SmbTarget::new(r"\\server\share").credentials("user", "secret-value");
let debug = format!("{target:?}");
assert!(debug.contains("user"));
assert!(!debug.contains("secret-value"));
}
#[test]
fn credential_forms_keep_missing_fields_absent() {
let default = SmbTarget::new(r"\\server\share");
let args = ConnectArgs::new(&default).unwrap();
assert!(args.username.is_none());
assert!(args.password.is_none());
let username_only = SmbTarget::new(r"\\server\share").username("user");
let args = ConnectArgs::new(&username_only).unwrap();
assert!(args.username.is_some());
assert!(args.password.is_none());
let password_only = SmbTarget::new(r"\\server\share").password("secret-value");
let args = ConnectArgs::new(&password_only).unwrap();
assert!(args.username.is_none());
assert!(args.password.is_some());
}
#[test]
fn invalid_lifetime_combinations_are_rejected_before_windows_calls() {
let deviceless = SmbTarget::new(r"\\server\share");
let persistent = ConnectOptions::new().persist(true);
assert_eq!(
deviceless.connect_with(persistent),
Err(Error::PersistenceRequiresDrive)
);
assert_eq!(
deviceless
.connect_guarded(ConnectOptions::new())
.unwrap_err(),
Error::GuardRequiresDrive
);
assert_eq!(
deviceless.connect_auto_guarded(persistent).unwrap_err(),
Error::PersistentGuard
);
let mapped = SmbTarget::new(r"\\server\share").mount_on(DriveLetter::D);
assert_eq!(
mapped.connect_guarded(persistent).unwrap_err(),
Error::PersistentGuard
);
assert_eq!(
deviceless.disconnect_with(DisconnectOptions::default().forget(true)),
Err(Error::ForgetRequiresDrive)
);
}
#[test]
fn drive_names_are_recognized_for_forgetting() {
assert!(is_drive("D:"));
assert!(is_drive("z:"));
assert!(!is_drive(r"\\server\share"));
assert!(!is_drive("D:\\"));
}
}