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
21pub struct WifiStation {
23 socket_path: std::path::PathBuf,
25 request_receiver: mpsc::Receiver<Request>,
27 broadcast_sender: broadcast::Sender<Broadcast>,
29 self_sender: mpsc::Sender<Request>,
31 select_timeout: Duration,
33 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 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 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 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 let field = match ¶m {
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 ¶m {
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 => (), }
308 Ok(())
309 }
310}
311
312fn 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 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 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 let hex_psk = "8dbbe42cb44f21088fbb9cfbf24dc9b39787d6026d436b01b3ac7d34afb4416d";
377 assert_eq!(hex_psk.parse::<Psk>().unwrap().to_field(), hex_psk);
378 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 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 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 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}