diff --git a/credentialsd/src/credential_service/usb.rs b/credentialsd/src/credential_service/usb.rs index 88069c7..0744cf6 100644 --- a/credentialsd/src/credential_service/usb.rs +++ b/credentialsd/src/credential_service/usb.rs @@ -1,4 +1,4 @@ -use std::{collections::HashMap, time::Duration}; +use std::{collections::HashMap, sync::Arc, time::Duration}; use async_stream::stream; use base64::{self, engine::general_purpose::URL_SAFE_NO_PAD, Engine}; @@ -16,6 +16,7 @@ use libwebauthn::{ use tokio::sync::{ broadcast, mpsc::{self, Receiver, Sender, WeakSender}, + Mutex as AsyncMutex, }; use tracing::{debug, warn}; @@ -70,7 +71,7 @@ impl InProcessUsbHandler { } async fn process_selecting_device( - hid_devices: Vec, + hid_devices: &[HidDevice], ) -> Result { let expected_answers = hid_devices.len(); let (blinking_tx, mut blinking_rx) = @@ -78,9 +79,10 @@ impl InProcessUsbHandler { let mut channel_map = HashMap::new(); let (setup_tx, mut setup_rx) = tokio::sync::mpsc::channel::<(usize, HidDevice, HidChannelHandle)>(expected_answers); - for (idx, mut device) in hid_devices.into_iter().enumerate() { + for (idx, device) in hid_devices.iter().enumerate() { let stx = setup_tx.clone(); let tx = blinking_tx.clone(); + let mut device = device.clone(); tokio::spawn(async move { let dev = device.clone(); @@ -132,7 +134,7 @@ impl InProcessUsbHandler { tracing::info!("Cancelling device {device:?}."); handle.cancel_ongoing_operation().await; } - state = UsbStateInternal::Connected(device); + state = UsbStateInternal::Connected(Arc::new(AsyncMutex::new(device))); break; } None => { @@ -144,7 +146,7 @@ impl InProcessUsbHandler { } async fn process_select_credential( - response: GetAssertionResponse, + response: &GetAssertionResponse, cred_rx: &mut Receiver, ) -> Result { match cred_rx.recv().await { @@ -246,20 +248,21 @@ impl InProcessUsbHandler { // act on current USB USB state, send state changes to the stream, and // loop until a credential or error is returned. loop { - tracing::debug!("current usb state: {:?}", state); + tracing::trace!("current usb state: {:?}", state); let prev_usb_state = state; let next_usb_state = match prev_usb_state { UsbStateInternal::Idle | UsbStateInternal::Waiting => { Self::process_idle_waiting(&mut failures, &prev_usb_state).await } - UsbStateInternal::SelectingDevice(hid_devices) => { - Self::process_selecting_device(hid_devices).await + UsbStateInternal::SelectingDevice(ref hid_devices) => { + Self::process_selecting_device(hid_devices.as_slice()).await } - UsbStateInternal::Connected(device) => { + UsbStateInternal::Connected(ref device) => { + let device = std::sync::Arc::clone(device); let signal_tx2 = signal_tx.clone(); let cred_request = cred_request.clone(); tokio::spawn(async move { - handle_events(&cred_request, device, &signal_tx2).await; + handle_events(&cred_request, device.clone(), &signal_tx2).await; }); Self::process_user_interaction(&mut signal_rx, &cred_tx).await } @@ -269,32 +272,36 @@ impl InProcessUsbHandler { Self::process_user_interaction(&mut signal_rx, &cred_tx).await } UsbStateInternal::SelectCredential { - response, + ref response, cred_tx: _, } => Self::process_select_credential(response, &mut cred_rx).await, UsbStateInternal::Completed(_) => break Ok(()), UsbStateInternal::Failed(err) => break Err(err), }; state = next_usb_state.unwrap_or_else(UsbStateInternal::Failed); - tx.send(state.clone()).await.map_err(|_| { - Error::Internal("USB state channel receiver closed prematurely".to_string()) - })?; + if std::mem::discriminant(&state) != std::mem::discriminant(&prev_usb_state) { + tracing::debug!("USB current state: {state:?}"); + tx.send(state.clone()).await.map_err(|_| { + Error::Internal("USB state channel receiver closed prematurely".to_string()) + })?; + } } } } async fn handle_events( cred_request: &CredentialRequest, - mut device: HidDevice, + device: Arc>, signal_tx: &Sender>, ) { + let mut device = device.lock().await; let device_debug = device.to_string(); - match device + let channel = device .channel(ChannelSettings { persistent_token_store: Some(super::persistent_token_store()), }) - .await - { + .await; + match channel { Err(err) => { tracing::error!("Failed to open channel to USB authenticator, cannot receive user verification events: {:?}", err); } @@ -403,7 +410,7 @@ pub(super) enum UsbStateInternal { SelectingDevice(Vec), /// USB device connected, prompt user to tap - Connected(HidDevice), + Connected(Arc>), /// The device needs the PIN to be entered. NeedsPin {