use crate::ndr::{NdrDecoder, NdrEncoder};
use crate::transport::SmbPipe;
use crate::{Result, Syntax};
use smb2_client::SmbClient;
pub fn wkssvc_syntax() -> Syntax {
Syntax::new("6bffd098-a112-3610-9833-46c3f87e345a", 1, 0)
}
pub mod opnum {
pub const NETR_WKSTA_USER_ENUM: u16 = 2;
}
const WKSTA_USER_LEVEL_1: u32 = 1;
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct WkstaUser {
pub username: String,
pub logon_domain: String,
pub other_domains: String,
pub logon_server: String,
}
pub fn encode_wksta_user_enum() -> Vec<u8> {
let mut e = NdrEncoder::new();
e.null_ptr();
e.u32(WKSTA_USER_LEVEL_1); e.u32(WKSTA_USER_LEVEL_1); e.referent(); e.u32(0); e.null_ptr();
e.u32(0xFFFF_FFFF); e.null_ptr(); e.into_bytes()
}
pub fn decode_wksta_user_enum(stub: &[u8]) -> Result<(Vec<WkstaUser>, u32, u32)> {
let mut d = NdrDecoder::new(stub);
let _level = d.u32()?; let _tag = d.u32()?; let container_ref = d.u32()?; let mut users = Vec::new();
if container_ref != 0 {
let entries_read = d.u32()? as usize;
let buffer_ref = d.u32()?;
if buffer_ref != 0 {
let _max = d.u32()?; if entries_read
.checked_mul(16)
.map_or(true, |need| need > d.remaining())
{
return Err(crate::RpcError::Protocol(format!(
"NetrWkstaUserEnum: EntriesRead={entries_read} exceeds remaining stub"
)));
}
let mut refs = Vec::with_capacity(entries_read);
for _ in 0..entries_read {
let user_ref = d.u32()?;
let logon_domain_ref = d.u32()?;
let oth_domains_ref = d.u32()?;
let logon_server_ref = d.u32()?;
refs.push((
user_ref,
logon_domain_ref,
oth_domains_ref,
logon_server_ref,
));
}
for (u_ref, ld_ref, od_ref, ls_ref) in refs {
let username = if u_ref != 0 {
d.conformant_varying_wstr()?
} else {
String::new()
};
let logon_domain = if ld_ref != 0 {
d.conformant_varying_wstr()?
} else {
String::new()
};
let other_domains = if od_ref != 0 {
d.conformant_varying_wstr()?
} else {
String::new()
};
let logon_server = if ls_ref != 0 {
d.conformant_varying_wstr()?
} else {
String::new()
};
users.push(WkstaUser {
username,
logon_domain,
other_domains,
logon_server,
});
}
}
}
let total_entries = d.u32()?;
let resume_ref = d.u32()?;
if resume_ref != 0 {
let _resume = d.u32();
}
let ret = d.u32()?;
Ok((users, total_entries, ret))
}
pub struct WkstaUserClient<'a> {
pipe: SmbPipe<'a>,
}
impl<'a> WkstaUserClient<'a> {
pub async fn bind(client: &'a mut SmbClient, file_id: [u8; 16]) -> Result<Self> {
let mut pipe = SmbPipe::new(client, file_id);
pipe.bind(wkssvc_syntax()).await?;
Ok(WkstaUserClient { pipe })
}
pub async fn enum_users(&mut self) -> Result<(Vec<WkstaUser>, u32)> {
let resp = self
.pipe
.call(opnum::NETR_WKSTA_USER_ENUM, &encode_wksta_user_enum())
.await?;
let (users, _total, ret) = decode_wksta_user_enum(&resp)?;
Ok((users, ret))
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn request_selects_level_1_with_null_server() {
let stub = encode_wksta_user_enum();
assert_eq!(&stub[0..4], &0u32.to_le_bytes()); assert_eq!(u32::from_le_bytes(stub[4..8].try_into().unwrap()), 1); assert_eq!(u32::from_le_bytes(stub[8..12].try_into().unwrap()), 1); }
#[test]
fn decodes_one_user_from_a_handbuilt_reply() {
fn wstr(out: &mut Vec<u8>, s: &str) {
let units: Vec<u16> = s.encode_utf16().chain(std::iter::once(0)).collect();
out.extend_from_slice(&(units.len() as u32).to_le_bytes()); out.extend_from_slice(&0u32.to_le_bytes()); out.extend_from_slice(&(units.len() as u32).to_le_bytes()); for u in units {
out.extend_from_slice(&u.to_le_bytes());
}
while out.len() % 4 != 0 {
out.push(0);
}
}
let mut r = Vec::new();
r.extend_from_slice(&1u32.to_le_bytes()); r.extend_from_slice(&1u32.to_le_bytes()); r.extend_from_slice(&0x2_0000u32.to_le_bytes()); r.extend_from_slice(&1u32.to_le_bytes()); r.extend_from_slice(&0x2_0004u32.to_le_bytes()); r.extend_from_slice(&1u32.to_le_bytes()); r.extend_from_slice(&0x2_0008u32.to_le_bytes()); r.extend_from_slice(&0x2_000cu32.to_le_bytes()); r.extend_from_slice(&0u32.to_le_bytes()); r.extend_from_slice(&0x2_0010u32.to_le_bytes()); wstr(&mut r, "alice"); wstr(&mut r, "CORP"); wstr(&mut r, "DC01"); r.extend_from_slice(&1u32.to_le_bytes()); r.extend_from_slice(&0u32.to_le_bytes()); r.extend_from_slice(&0u32.to_le_bytes());
let (users, total, ret) = decode_wksta_user_enum(&r).unwrap();
assert_eq!(ret, 0);
assert_eq!(total, 1);
assert_eq!(
users,
vec![WkstaUser {
username: "alice".into(),
logon_domain: "CORP".into(),
other_domains: String::new(),
logon_server: "DC01".into(),
}]
);
}
#[test]
fn empty_reply_is_not_a_panic() {
for cut in 0..24 {
let _ = decode_wksta_user_enum(&vec![0u8; cut]);
}
}
#[test]
fn entries_read_is_bounded_against_stub() {
let mut r = Vec::new();
r.extend_from_slice(&1u32.to_le_bytes()); r.extend_from_slice(&1u32.to_le_bytes()); r.extend_from_slice(&0x2_0000u32.to_le_bytes()); r.extend_from_slice(&0xFFFF_FFFFu32.to_le_bytes()); r.extend_from_slice(&0x2_0004u32.to_le_bytes()); r.extend_from_slice(&0xFFFF_FFFFu32.to_le_bytes()); let err = decode_wksta_user_enum(&r).unwrap_err();
assert!(
matches!(err, crate::RpcError::Protocol(ref s) if s.contains("EntriesRead")),
"expected Protocol(EntriesRead …), got {err:?}"
);
}
}