Skip to main content

wifi_ctrl/sta/
mod.rs

1use crate::error::ClientError;
2
3use super::*;
4
5use tokio::time::Duration;
6
7mod types;
8pub use types::*;
9
10mod client;
11pub use client::*;
12
13mod setup;
14pub use setup::*;
15
16mod event_socket;
17use event_socket::*;
18
19const PATH_DEFAULT_SERVER: &str = "/var/run/wpa_supplicant/wlan2";
20
21/// Instance that runs the Wifi process
22pub struct WifiStation {
23    /// Path to the socket
24    socket_path: std::path::PathBuf,
25    /// Channel for receiving requests
26    request_receiver: mpsc::Receiver<Request>,
27    /// Channel for broadcasting alerts
28    broadcast_sender: broadcast::Sender<Broadcast>,
29    /// Channel for sending requests to itself
30    self_sender: mpsc::Sender<Request>,
31    /// Timeout duration in case no valid select response is received
32    select_timeout: Duration,
33    /// How long to wait for a reply to a control command/request
34    command_timeout: Duration,
35}
36
37impl WifiStation {
38    pub async fn run(&mut self) -> SocketResult {
39        info!("Starting Wifi Station process");
40        let (socket_handle, mut deferred_requests) = SocketHandle::open(
41            &self.socket_path,
42            "mapper_wpa_ctrl_sync.sock",
43            &mut self.request_receiver,
44            self.command_timeout,
45        )
46        .await?;
47        // We start up a separate socket for receiving the "unexpected" events that
48        // gets forwarded to us via the unsolicited_receiver
49        let (next_deferred_requests, unsolicited) = EventSocket::new(
50            &self.socket_path,
51            &mut self.request_receiver,
52            self.command_timeout,
53        )
54        .await?;
55        deferred_requests.extend(next_deferred_requests);
56        for request in deferred_requests {
57            self.self_sender
58                .send(request)
59                .await
60                .expect("self_sender should never close as same struct owns both ends");
61        }
62        self.broadcast(Broadcast::Ready);
63        self.run_internal(unsolicited, socket_handle).await
64    }
65
66    fn broadcast(&self, event: Broadcast) {
67        if self.broadcast_sender.send(event).is_err() {
68            debug!("broadcast listener closed")
69        }
70    }
71
72    async fn run_internal(
73        &mut self,
74        mut unsolicited: EventSocket,
75        mut socket_handle: SocketHandle<10240>,
76    ) -> SocketResult {
77        // We will collect scan requests and batch respond to them when results are ready
78        let mut scan_requests = Vec::new();
79        let mut select_request: Option<SelectRequest> = None;
80        loop {
81            enum EventOrRequest {
82                Event(Event),
83                Request(Option<Request>),
84                SelectTimeout,
85            }
86
87            let event_or_request = tokio::select!(
88                unsolicited_msg = unsolicited.recv() => {
89                    EventOrRequest::Event(unsolicited_msg?)
90                },
91                request = self.request_receiver.recv() => {
92                    EventOrRequest::Request(request)
93                },
94                _ = async {
95                    match select_request.as_mut() {
96                        Some(select_request) => select_request.timeout.as_mut().await,
97                        None => std::future::pending().await,
98                    }
99                } => EventOrRequest::SelectTimeout,
100            );
101
102            match event_or_request {
103                EventOrRequest::Event(unsolicited_msg) => {
104                    debug!("Unsolicited event: {unsolicited_msg:?}");
105                    self.handle_event(
106                        &mut socket_handle,
107                        unsolicited_msg,
108                        &mut scan_requests,
109                        &mut select_request,
110                    )
111                    .await?
112                }
113                EventOrRequest::Request(request) => match request {
114                    Some(Request::Shutdown) => return Ok(()),
115                    Some(request) => {
116                        self.handle_request(
117                            &mut socket_handle,
118                            request,
119                            &mut scan_requests,
120                            &mut select_request,
121                        )
122                        .await?;
123                    }
124                    None => return Err(error::SocketError::ClientChannelClosed),
125                },
126                EventOrRequest::SelectTimeout => {
127                    if let Some(sender) = select_request.take() {
128                        sender.send(Err(ClientError::Timeout));
129                    };
130                }
131            }
132        }
133    }
134
135    async fn handle_event<const N: usize>(
136        &mut self,
137        socket_handle: &mut SocketHandle<N>,
138        event: Event,
139        scan_requests: &mut Vec<oneshot::Sender<Result<Arc<Vec<ScanResult>>>>>,
140        select_request: &mut Option<SelectRequest>,
141    ) -> SocketResult {
142        match event {
143            Event::ScanComplete => {
144                let scan_results = socket_handle
145                    .request("SCAN_RESULTS", ScanResult::vec_from_str)
146                    .await?;
147                while let Some(scan_request) = scan_requests.pop() {
148                    let _ = scan_request.send(scan_results.clone());
149                }
150            }
151            Event::ScanFailed => {
152                while let Some(scan_request) = scan_requests.pop() {
153                    let _ = scan_request.send(Err(ClientError::Failed));
154                }
155            }
156            Event::Connected => {
157                self.broadcast(Broadcast::Connected);
158                if let Some(sender) = select_request.take() {
159                    sender.send(Ok(SelectResult::Success));
160                }
161            }
162            Event::Disconnected => {
163                self.broadcast(Broadcast::Disconnected);
164            }
165            Event::NetworkNotFound => {
166                self.broadcast(Broadcast::NetworkNotFound);
167                if let Some(sender) = select_request.take() {
168                    sender.send(Ok(SelectResult::NotFound));
169                }
170            }
171            Event::WrongPsk => {
172                self.broadcast(Broadcast::WrongPsk);
173                if let Some(sender) = select_request.take() {
174                    sender.send(Ok(SelectResult::WrongPsk));
175                }
176            }
177            Event::Unknown(msg) => {
178                self.broadcast(Broadcast::Unknown(msg));
179            }
180        }
181        Ok(())
182    }
183
184    async fn get_status<const N: usize>(
185        socket_handle: &mut SocketHandle<N>,
186    ) -> SocketResult<Result<Status>> {
187        socket_handle.request("STATUS", parse_status).await
188    }
189
190    async fn handle_request<const N: usize>(
191        &self,
192        socket_handle: &mut SocketHandle<N>,
193        request: Request,
194        scan_requests: &mut Vec<oneshot::Sender<Result<Arc<Vec<ScanResult>>>>>,
195        select_request: &mut Option<SelectRequest>,
196    ) -> SocketResult {
197        debug!("Handling request: {request:?}");
198        match request {
199            Request::Custom(custom, response_channel) => {
200                let data_str = socket_handle.request(&custom, TryInto::try_into).await?;
201                debug!("Custom request response: {data_str:?}");
202                let _ = response_channel.send(data_str);
203            }
204            Request::Scan(response_channel) => {
205                // wpa_supplicant replies FAIL-BUSY when a scan is already in
206                // progress; the pending CTRL-EVENT-SCAN-RESULTS will answer
207                // this request too, so treat it as accepted
208                match socket_handle
209                    .command_matching(b"SCAN", |data| data == "OK" || data == "FAIL-BUSY")
210                    .await?
211                {
212                    Ok(_) => {
213                        scan_requests.push(response_channel);
214                    }
215                    Err(e) => {
216                        let _ = response_channel.send(Err(e));
217                    }
218                };
219            }
220            Request::Networks(response_channel) => {
221                let network_list = NetworkResult::request_results(socket_handle).await?;
222                let _ = response_channel.send(network_list);
223            }
224            Request::Status(response_channel) => {
225                let status = Self::get_status(socket_handle).await?;
226                let _ = response_channel.send(status);
227            }
228            Request::AddNetwork(response_channel) => {
229                let network_id = socket_handle
230                    .request("ADD_NETWORK", usize::from_str)
231                    .await?;
232                debug!("wpa_ctrl created network {network_id:?}");
233                let _ = response_channel.send(network_id);
234            }
235            Request::SetNetwork(id, param, response) => {
236                // Psk and Bssid are validated at construction, so every
237                // variant formats infallibly; Psk's Debug impl redacts the
238                // key wherever the request is logged.
239                let field = match &param {
240                    SetNetwork::Ssid(ssid) => format!("ssid {}", conf_escape(ssid)),
241                    SetNetwork::Bssid(bssid) => format!("bssid {bssid}"),
242                    SetNetwork::Psk(psk) => format!("psk {}", psk.to_field()),
243                    SetNetwork::KeyMgmt(mgmt) => format!("key_mgmt {mgmt}"),
244                };
245                let cmd = format!("SET_NETWORK {id} {field}");
246                match &param {
247                    SetNetwork::Psk(_) => debug!("wpa_ctrl SET_NETWORK {id} psk <redacted>"),
248                    _ => debug!("wpa_ctrl {cmd:?}"),
249                }
250                let _ = response.send(socket_handle.command(cmd.as_bytes()).await?);
251            }
252            Request::SaveConfig(response) => {
253                debug!("wpa_ctrl config saved");
254                let _ = response.send(socket_handle.command(b"SAVE_CONFIG").await?);
255            }
256            Request::ReloadConfig(response) => {
257                debug!("wpa_ctrl config reloaded");
258                let _ = response.send(socket_handle.command(b"RECONFIGURE").await?);
259            }
260            Request::RemoveNetwork(remove_network, response) => {
261                let str = match remove_network {
262                    RemoveNetwork::All => "all".to_string(),
263                    RemoveNetwork::Id(id) => id.to_string(),
264                };
265                let cmd = format!("REMOVE_NETWORK {str}");
266                let bytes = cmd.into_bytes();
267                debug!("wpa_ctrl removed network {str}");
268                let _ = response.send(socket_handle.command(&bytes).await?);
269            }
270            Request::SelectNetwork(id, response_sender) => {
271                match select_request {
272                    None => {
273                        let cmd = format!("SELECT_NETWORK {id}");
274                        let bytes = cmd.into_bytes();
275                        if let Err(e) = socket_handle.command(&bytes).await? {
276                            warn!("Error while selecting network {id}: {e}");
277                            let _ = response_sender.send(Err(e));
278                        } else {
279                            debug!("wpa_ctrl selected network {id}");
280                            match Self::get_status(socket_handle).await? {
281                                Err(e) => {
282                                    let _ = response_sender.send(Err(e));
283                                }
284                                Ok(status) if status.id == Some(id) => {
285                                    let _ =
286                                        response_sender.send(Ok(SelectResult::AlreadyConnected));
287                                }
288                                Ok(_) => {
289                                    *select_request = Some(SelectRequest::new(
290                                        response_sender,
291                                        self.select_timeout,
292                                    ));
293                                }
294                            }
295                        }
296                    }
297                    Some(_) => {
298                        warn!("Select request already pending! Dropping this one.");
299                        let _ = response_sender.send(Err(ClientError::PendingSelect));
300                        debug!(
301                            "wpa_ctrl rejected select of network {id}: a select is already pending"
302                        );
303                    }
304                };
305            }
306            Request::Shutdown => (), //shutdown is handled at the scope above
307        }
308        Ok(())
309    }
310}
311
312/// convert to wpa config format, ideally a "quoted string"
313/// in case of new-lines, quotes or emoji fall back to hex encoding the whole thing
314fn conf_escape(raw: &str) -> String {
315    if raw.bytes().all(|b| b.is_ascii_graphic() && b != b'"') {
316        format!("\"{raw}\"")
317    } else {
318        hex::encode(raw)
319    }
320}
321
322struct SelectRequest {
323    response: oneshot::Sender<Result<SelectResult>>,
324    /// Polled as a branch of the main event loop; expiry resolves the request
325    /// with a timeout error
326    timeout: std::pin::Pin<Box<tokio::time::Sleep>>,
327}
328
329impl SelectRequest {
330    fn new(response: oneshot::Sender<Result<SelectResult>>, timeout: Duration) -> Self {
331        Self {
332            response,
333            timeout: Box::pin(tokio::time::sleep(timeout)),
334        }
335    }
336
337    fn send(self, result: Result<SelectResult>) {
338        let _ = self.response.send(result);
339    }
340}
341
342#[cfg(test)]
343mod tests {
344    use super::*;
345
346    #[test]
347    fn psk_passphrase_is_quoted() {
348        assert_eq!(
349            Psk::passphrase("password123").unwrap().to_field(),
350            "\"password123\""
351        );
352    }
353
354    #[test]
355    fn psk_passphrase_with_spaces_is_quoted_not_hex() {
356        // A passphrase may contain spaces; it must stay a quoted string, since
357        // an unquoted/hex value would be read as a raw pre-shared key.
358        assert_eq!(
359            Psk::passphrase("correct horse battery").unwrap().to_field(),
360            "\"correct horse battery\""
361        );
362    }
363
364    #[test]
365    fn raw_psk_is_bare_hex() {
366        let hex_psk = "8dbbe42cb44f21088fbb9cfbf24dc9b39787d6026d436b01b3ac7d34afb4416d";
367        let mut key = [0u8; 32];
368        hex::decode_to_slice(hex_psk, &mut key).unwrap();
369        assert_eq!(Psk::raw(key).to_field(), hex_psk);
370    }
371
372    #[test]
373    fn psk_from_str_uses_conf_semantics() {
374        // Exactly 64 hex digits parses as a raw key, anything else as a
375        // passphrase; unambiguous since a passphrase is at most 63 chars.
376        let hex_psk = "8dbbe42cb44f21088fbb9cfbf24dc9b39787d6026d436b01b3ac7d34afb4416d";
377        assert_eq!(hex_psk.parse::<Psk>().unwrap().to_field(), hex_psk);
378        // 63 hex digits is a passphrase
379        assert_eq!(
380            hex_psk[..63].parse::<Psk>().unwrap().to_field(),
381            format!("\"{}\"", &hex_psk[..63])
382        );
383    }
384
385    #[test]
386    fn psk_with_quote_is_rejected() {
387        // A literal quote could break out of the quoted value, so it is rejected
388        // rather than emitted into the SET_NETWORK command.
389        assert!(matches!(
390            Psk::passphrase("pass\"; extra"),
391            Err(ClientError::InvalidPsk)
392        ));
393    }
394
395    #[test]
396    fn psk_with_control_char_is_rejected() {
397        assert!(matches!(
398            Psk::passphrase("pass\nword"),
399            Err(ClientError::InvalidPsk)
400        ));
401    }
402
403    #[test]
404    fn psk_outside_wpa_length_is_rejected() {
405        assert!(matches!(
406            Psk::passphrase("short"),
407            Err(ClientError::InvalidPsk)
408        ));
409        assert!(matches!(
410            Psk::passphrase("x".repeat(64)),
411            Err(ClientError::InvalidPsk)
412        ));
413    }
414
415    #[test]
416    fn psk_debug_is_redacted() {
417        let psk = Psk::passphrase("password123").unwrap();
418        assert_eq!(format!("{psk:?}"), "Psk(<redacted>)");
419        // The whole request is debug-logged by handle_request; the key must
420        // not appear through that path either.
421        let request = SetNetwork::Psk(psk);
422        assert!(!format!("{request:?}").contains("password123"));
423    }
424
425    #[test]
426    fn bssid_roundtrips_raw_and_unquoted() {
427        let bssid: Bssid = "cc:7b:5c:1a:d2:21".parse().unwrap();
428        assert_eq!(bssid.to_string(), "cc:7b:5c:1a:d2:21");
429        assert_eq!(Bssid::from([0xcc, 0x7b, 0x5c, 0x1a, 0xd2, 0x21]), bssid);
430    }
431
432    #[test]
433    fn bssid_is_canonicalized() {
434        // Mixed case parses but is always emitted lowercase
435        let bssid: Bssid = "CC:7B:5C:1A:D2:21".parse().unwrap();
436        assert_eq!(bssid.to_string(), "cc:7b:5c:1a:d2:21");
437    }
438
439    #[test]
440    fn malformed_bssid_is_rejected() {
441        for bad in [
442            "cc:7b:5c:1a:d2",
443            "cc:7b:5c:1a:d2:21:33",
444            "cc:7b:5c:1a:d2:21 x",
445            "cc:7b:5c:1a:d2:+1",
446            "not-a-mac",
447            "",
448        ] {
449            assert!(
450                matches!(bad.parse::<Bssid>(), Err(ClientError::InvalidBssid)),
451                "expected {bad:?} to be rejected"
452            );
453        }
454    }
455}