use crate::sys::mib_table::MibTable;
use crate::sys::InterfaceHandle;
use crate::{Error, Interface};
use ipnet::IpNet;
use std::collections::HashSet;
use std::io::{self, ErrorKind};
use std::net::{Ipv4Addr, Ipv6Addr, SocketAddr};
use widestring::U16CString;
use windows::core::{Error as WinError, GUID, HRESULT, HSTRING};
use windows::Win32::Foundation;
use windows::Win32::NetworkManagement::IpHelper::{
ConvertInterfaceAliasToLuid, ConvertInterfaceGuidToLuid, ConvertInterfaceIndexToLuid,
ConvertInterfaceLuidToAlias, ConvertInterfaceLuidToGuid, ConvertInterfaceLuidToIndex,
ConvertInterfaceLuidToNameW, ConvertInterfaceNameToLuidW, CreateUnicastIpAddressEntry,
DeleteUnicastIpAddressEntry, GetIfEntry2, GetIpInterfaceEntry, InitializeIpInterfaceEntry,
InitializeUnicastIpAddressEntry, SetIpInterfaceEntry, MIB_IF_ROW2, MIB_IPINTERFACE_ROW,
MIB_UNICASTIPADDRESS_ROW,
};
use windows::Win32::NetworkManagement::Ndis::{IF_MAX_STRING_SIZE, NET_LUID_LH};
use windows::Win32::Networking::WinSock::{
ADDRESS_FAMILY, AF_INET, AF_INET6, AF_UNSPEC, SOCKADDR_INET,
};
const ERROR_ACCESS_DENIED: HRESULT = Foundation::ERROR_ACCESS_DENIED.to_hresult();
const ERROR_FILE_NOT_FOUND: HRESULT = Foundation::ERROR_FILE_NOT_FOUND.to_hresult();
const ERROR_INVALID_NAME: HRESULT = Foundation::ERROR_INVALID_NAME.to_hresult();
fn convert_sockaddr(sa: SOCKADDR_INET) -> SocketAddr {
unsafe {
match sa.si_family {
AF_INET => SocketAddr::new(
Ipv4Addr::from(sa.Ipv4.sin_addr).into(),
u16::from_be(sa.Ipv4.sin_port),
),
AF_INET6 => SocketAddr::new(
Ipv6Addr::from(sa.Ipv6.sin6_addr).into(),
u16::from_be(sa.Ipv6.sin6_port),
),
_ => panic!("Invalid address family"),
}
}
}
impl InterfaceHandle {
fn mib_if_row2(&self) -> Result<MIB_IF_ROW2, Error> {
let mut row = MIB_IF_ROW2 {
InterfaceIndex: self.index,
..Default::default()
};
let code = unsafe { GetIfEntry2(&mut row) };
match code.ok().map_err(HRESULT::from) {
Ok(_) => Ok(row),
Err(ERROR_FILE_NOT_FOUND) => Err(Error::InterfaceNotFound),
Err(e) => Err(WinError::from(e).into()),
}
}
fn mib_unicastipaddress_row(&self, network: IpNet) -> MIB_UNICASTIPADDRESS_ROW {
let mut row = MIB_UNICASTIPADDRESS_ROW::default();
unsafe { InitializeUnicastIpAddressEntry(&mut row as _) };
row.InterfaceIndex = self.index;
row.Address = SocketAddr::new(network.addr(), 0).into();
row.OnLinkPrefixLength = network.prefix_len();
row
}
fn net_luid_lh(&self) -> Result<NET_LUID_LH, Error> {
Ok(NET_LUID_LH {
Value: self.interface().luid()?,
})
}
}
pub trait InterfaceExt {
fn try_from_luid(luid: u64) -> Result<Interface, Error>;
fn try_from_guid(guid: u128) -> Result<Interface, Error>;
fn try_from_alias(alias: &str) -> Result<Interface, Error>;
fn luid(&self) -> Result<u64, Error>;
fn guid(&self) -> Result<u128, Error>;
fn alias(&self) -> Result<String, Error>;
fn description(&self) -> Result<String, Error>;
}
impl InterfaceExt for Interface {
fn try_from_luid(luid: u64) -> Result<Interface, Error> {
let luid = NET_LUID_LH { Value: luid };
let mut index = 0;
unsafe { ConvertInterfaceLuidToIndex(&luid, &mut index).ok()? };
Ok(unsafe { Self::from_index_unchecked(index) })
}
fn try_from_guid(guid: u128) -> Result<Interface, Error> {
let mut luid = NET_LUID_LH::default();
unsafe { ConvertInterfaceGuidToLuid(&GUID::from_u128(guid), &mut luid).ok()? };
Self::try_from_luid(unsafe { luid.Value })
}
fn try_from_alias(alias: &str) -> Result<Interface, Error> {
let mut luid = NET_LUID_LH::default();
let alias = HSTRING::from(alias);
let code = unsafe { ConvertInterfaceAliasToLuid(&alias, &mut luid) }
.ok()
.map_err(HRESULT::from);
match code {
Ok(_) => Self::try_from_luid(unsafe { luid.Value }),
Err(ERROR_INVALID_NAME) => Err(Error::InterfaceNotFound),
Err(e) => Err(WinError::from(e).into()),
}
}
fn luid(&self) -> Result<u64, Error> {
let mut luid = NET_LUID_LH::default();
let code = unsafe { ConvertInterfaceIndexToLuid(self.index()?, &mut luid) };
match code.ok().map_err(HRESULT::from) {
Ok(_) => Ok(unsafe { luid.Value }),
Err(ERROR_FILE_NOT_FOUND) => Err(Error::InterfaceNotFound),
Err(e) => Err(WinError::from(e).into()),
}
}
fn guid(&self) -> Result<u128, Error> {
let mut guid = GUID::zeroed();
let code = unsafe { ConvertInterfaceLuidToGuid(&self.0.net_luid_lh()?, &mut guid) };
match code.ok().map_err(HRESULT::from) {
Ok(_) => Ok(guid.into()),
Err(ERROR_FILE_NOT_FOUND) => Err(Error::InterfaceNotFound),
Err(e) => Err(WinError::from(e).into()),
}
}
fn alias(&self) -> Result<String, Error> {
let mut alias_buf = vec![0u16; (IF_MAX_STRING_SIZE + 1) as _];
let code = unsafe { ConvertInterfaceLuidToAlias(&self.0.net_luid_lh()?, &mut alias_buf) };
match code.ok().map_err(HRESULT::from) {
Ok(_) => Ok(U16CString::from_vec_truncate(alias_buf).to_string()?),
Err(ERROR_FILE_NOT_FOUND) => Err(Error::InterfaceNotFound),
Err(e) => Err(WinError::from(e).into()),
}
}
fn description(&self) -> Result<String, Error> {
Ok(
U16CString::from_vec_truncate(self.0.mib_if_row2()?.Description.to_vec())
.to_string()?,
)
}
}
impl InterfaceHandle {
pub fn addresses(&self) -> Result<Vec<IpNet>, Error> {
let address_set: Result<HashSet<IpNet>, Error> =
MibTable::GetUnicastIpAddressTable(&AF_UNSPEC)?
.as_slice()
.iter()
.filter(|row| row.InterfaceIndex == self.index)
.map(|row| {
IpNet::new(convert_sockaddr(row.Address).ip(), row.OnLinkPrefixLength)
.map_err(|_| Error::UnexpectedMetadata)
})
.collect();
Ok(address_set?.into_iter().collect())
}
pub fn add_address(&self, network: IpNet) -> Result<(), Error> {
let entry = self.mib_unicastipaddress_row(network);
unsafe { Ok(CreateUnicastIpAddressEntry(&entry).ok()?) }
}
pub fn remove_address(&self, network: IpNet) -> Result<(), Error> {
let entry = self.mib_unicastipaddress_row(network);
unsafe { Ok(DeleteUnicastIpAddressEntry(&entry).ok()?) }
}
pub fn mtu(&self) -> Result<u32, Error> {
self.get_mtu_for_family(AF_INET)
}
pub fn mtu_v4(&self) -> Result<u32, Error> {
self.get_mtu_for_family(AF_INET)
}
pub fn mtu_v6(&self) -> Result<u32, Error> {
self.get_mtu_for_family(AF_INET6)
}
pub fn set_mtu_v4(&self, mtu: u32) -> Result<(), Error> {
self.set_mtu_impl(mtu, AF_INET)
}
pub fn set_mtu_v6(&self, mtu: u32) -> Result<(), Error> {
self.set_mtu_impl(mtu, AF_INET6)
}
pub fn set_mtu(&self, mtu: u32) -> Result<(), Error> {
for family in [AF_INET, AF_INET6] {
self.set_mtu_impl(mtu, family)?;
}
Ok(())
}
fn set_mtu_impl(&self, mtu: u32, family: ADDRESS_FAMILY) -> Result<(), Error> {
let mut row: MIB_IPINTERFACE_ROW = unsafe { std::mem::zeroed() };
unsafe { InitializeIpInterfaceEntry(&mut row) };
row.Family = family;
row.InterfaceIndex = self.index;
row.NlMtu = mtu;
let code = unsafe { SetIpInterfaceEntry(&mut row) };
match code.ok().map_err(HRESULT::from) {
Ok(_) => Ok(()),
Err(ERROR_FILE_NOT_FOUND) => Err(Error::InterfaceNotFound),
Err(ERROR_ACCESS_DENIED) => Err(io::Error::from(ErrorKind::PermissionDenied).into()),
Err(e) => Err(WinError::from(e).into()),
}
}
pub fn get_mtu_for_family(&self, family: ADDRESS_FAMILY) -> Result<u32, Error> {
let mut row: MIB_IPINTERFACE_ROW = unsafe { std::mem::zeroed() };
unsafe { InitializeIpInterfaceEntry(&mut row) };
row.Family = family;
row.InterfaceIndex = self.index;
let code = unsafe { GetIpInterfaceEntry(&mut row) };
match code.ok().map_err(HRESULT::from) {
Ok(_) => {
if row.NlMtu < 65535 {
return Ok(row.NlMtu);
}
let link_mtu = self.mib_if_row2()?.Mtu;
Ok(link_mtu)
}
Err(ERROR_FILE_NOT_FOUND) => Err(Error::InterfaceNotFound),
Err(e) => Err(WinError::from(e).into()),
}
}
pub fn name(&self) -> Result<String, Error> {
let mut name_buf = vec![0u16; (IF_MAX_STRING_SIZE + 1) as _];
let code = unsafe { ConvertInterfaceLuidToNameW(&self.net_luid_lh()?, &mut name_buf) };
match code.ok().map_err(HRESULT::from) {
Ok(_) => Ok(U16CString::from_vec_truncate(name_buf).to_string()?),
Err(ERROR_FILE_NOT_FOUND) => Err(Error::InterfaceNotFound),
Err(e) => Err(WinError::from(e).into()),
}
}
pub fn try_from_name(name: &str) -> Result<Interface, Error> {
let mut luid = NET_LUID_LH::default();
let name = HSTRING::from(name);
let code = unsafe { ConvertInterfaceNameToLuidW(&name, &mut luid) };
match code.ok().map_err(HRESULT::from) {
Ok(_) => Interface::try_from_luid(unsafe { luid.Value }),
Err(ERROR_FILE_NOT_FOUND) => Err(Error::InterfaceNotFound),
Err(e) => Err(WinError::from(e).into()),
}
}
pub fn index(&self) -> Result<u32, Error> {
Ok(self.index)
}
pub fn try_from_index(index: u32) -> Result<Interface, Error> {
let mut luid = NET_LUID_LH::default();
let code = unsafe { ConvertInterfaceIndexToLuid(index, &mut luid) };
match code.ok().map_err(HRESULT::from) {
Ok(_) => Ok(unsafe { Interface::from_index_unchecked(index) }),
Err(ERROR_FILE_NOT_FOUND) => Err(Error::InterfaceNotFound),
Err(e) => Err(WinError::from(e).into()),
}
}
pub fn hwaddress(&self) -> Result<[u8; 6], Error> {
self.mib_if_row2()?.PhysicalAddress[..6]
.try_into()
.map_err(|_| Error::UnexpectedMetadata)
}
}