use std::{io, iter::once, time::Duration};
use neli::{
consts::{nl::NlmF, socket::NlFamily},
genl::{AttrTypeBuilder, Genlmsghdr, GenlmsghdrBuilder, NlattrBuilder},
nl::NlPayload,
router::{asynchronous, synchronous},
types::GenlBuffer,
utils::Groups,
};
use tokio::{pin, select, time::sleep};
use super::{
Backend, Error, Scan, ScanCompleted, ScanInternal, ScanTriggered, network_manager, nl80211,
};
use crate::{
Bss, Interface,
nl80211::{Attr, Cmd, NL80211_FAMILY_NAME},
};
const SCAN_TIMEOUT: Duration = Duration::from_secs(25);
const IDLE_TIMEOUT: Duration = Duration::from_secs(2);
const SCAN_MULTICAST_NAME: &str = "scan";
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum ScanState {
ScanInProgress, WaitingForNextScan, }
#[tracing::instrument(skip(interface), fields(interface = %interface.name(), ifindex = interface.index()))]
pub(crate) async fn scan(interface: &Interface, backend: Backend) -> Result<Scan, Error> {
let (socket, mut multicast) =
asynchronous::NlRouter::connect(NlFamily::Generic, None, Groups::empty()).await?;
let id = socket
.resolve_nl_mcast_group(NL80211_FAMILY_NAME, SCAN_MULTICAST_NAME)
.await?;
socket.add_mcast_membership(Groups::new_groups(&[id]))?;
match backend {
Backend::Nl80211 => nl80211::trigger_scan(&socket, interface).await?,
Backend::NetworkManager => network_manager::trigger_scan(interface).await?,
};
let mut state = ScanState::ScanInProgress;
let timeout_fut = sleep(SCAN_TIMEOUT);
pin!(timeout_fut);
let mut last_scan_triggered = None;
let mut scans = Vec::new();
let mut first_scan_results_error = None;
loop {
select! {
Some(msg) = multicast.next::<u16, Genlmsghdr<Cmd, Attr>>() => {
let Ok(msg) = msg else {
continue;
};
let Some(payload) = msg.get_payload() else {
continue;
};
match *payload.cmd() {
Cmd::TriggerScan => {
let Ok(scan_triggered) = ScanTriggered::try_from(payload) else {
continue;
};
if scan_triggered.ifindex != interface.index() {
continue;
};
last_scan_triggered = Some(scan_triggered);
timeout_fut.set(sleep(SCAN_TIMEOUT));
state = ScanState::ScanInProgress;
}
Cmd::NewScanResults => {
let Ok(scan_completed) = ScanCompleted::try_from(payload) else {
continue;
};
if scan_completed.ifindex != interface.index() {
continue;
}
let bss_list = match scan_results(interface, &socket).await {
Ok(bss_list) => bss_list,
Err(error) => {
tracing::warn!(error = %error, "Failed to fetch scan results");
first_scan_results_error.get_or_insert(error);
continue;
}
};
if let Some(scan_triggered) = last_scan_triggered.take() {
scans.push(ScanInternal {
bss_list,
scan_triggered,
scan_completed,
});
}
timeout_fut.set(sleep(IDLE_TIMEOUT));
state = ScanState::WaitingForNextScan;
}
_ => {}
}
},
_ = &mut timeout_fut => {
match state {
ScanState::ScanInProgress => {
return if scans.is_empty() {
Err(first_scan_results_error.unwrap_or_else(|| {
io::Error::new(io::ErrorKind::TimedOut, "Scanning timed out").into()
}))
} else {
Scan::new(scans)
}
}
ScanState::WaitingForNextScan => {
break;
}
}
}
}
}
Scan::new(scans)
}
#[tracing::instrument(skip(interface, socket), fields(interface = %interface.name(), ifindex = interface.index()))]
pub(crate) async fn scan_results(
interface: &Interface,
socket: &asynchronous::NlRouter,
) -> Result<Vec<Bss>, Error> {
tracing::debug!("Fetching nl80211 scan results");
let family_id = socket.resolve_genl_family(NL80211_FAMILY_NAME).await?;
let ifindex_attr = NlattrBuilder::default()
.nla_type(AttrTypeBuilder::default().nla_type(Attr::Ifindex).build()?)
.nla_payload(interface.index())
.build()?;
let attrs = once(ifindex_attr).collect::<GenlBuffer<_, _>>();
let genlmsghdr = GenlmsghdrBuilder::default()
.cmd(Cmd::GetScan)
.attrs(attrs)
.version(1)
.build()?;
let mut responses = socket
.send::<_, _, u16, Genlmsghdr<Cmd, Attr>>(
family_id,
NlmF::DUMP | NlmF::ACK,
NlPayload::Payload(genlmsghdr),
)
.await?;
let mut scan_results = Vec::new();
while let Some(response) = responses.next::<u16, Genlmsghdr<Cmd, Attr>>().await {
match response {
Ok(response) => {
let Some(payload) = response.get_payload() else {
continue;
};
match Bss::try_from(payload) {
Ok(bss) => scan_results.push(bss),
Err(error) => {
tracing::warn!(error = %error, "Skipping malformed BSS scan result");
}
}
}
Err(error) => {
tracing::warn!(error = %error, "Skipping failed scan response");
}
}
}
tracing::debug!(bss_count = scan_results.len(), "Retrieved scan results");
Ok(scan_results)
}
#[tracing::instrument(skip(interface, socket), fields(interface = %interface.name(), ifindex = interface.index()))]
pub(crate) fn scan_results_blocking(
interface: &Interface,
socket: &synchronous::NlRouter,
) -> Result<Vec<Bss>, Error> {
tracing::debug!("Fetching nl80211 scan results (blocking)");
let family_id = socket.resolve_genl_family(NL80211_FAMILY_NAME)?;
let ifindex_attr = NlattrBuilder::default()
.nla_type(AttrTypeBuilder::default().nla_type(Attr::Ifindex).build()?)
.nla_payload(interface.index())
.build()?;
let attrs = once(ifindex_attr).collect::<GenlBuffer<_, _>>();
let genlmsghdr = GenlmsghdrBuilder::default()
.cmd(Cmd::GetScan)
.attrs(attrs)
.version(1)
.build()?;
let responses = socket.send::<_, _, u16, Genlmsghdr<Cmd, Attr>>(
family_id,
NlmF::DUMP | NlmF::ACK,
NlPayload::Payload(genlmsghdr),
)?;
let mut scan_results = Vec::new();
for response in responses {
match response {
Ok(response) => {
let Some(payload) = response.get_payload() else {
continue;
};
match Bss::try_from(payload) {
Ok(bss) => scan_results.push(bss),
Err(error) => {
tracing::warn!(error = %error, "Skipping malformed BSS scan result");
}
}
}
Err(error) => {
tracing::warn!(error = %error, "Skipping failed scan response");
}
}
}
tracing::debug!(bss_count = scan_results.len(), "Retrieved scan results");
Ok(scan_results)
}