Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
45 changes: 26 additions & 19 deletions credentialsd/src/credential_service/usb.rs
Original file line number Diff line number Diff line change
@@ -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};
Expand All @@ -16,6 +16,7 @@ use libwebauthn::{
use tokio::sync::{
broadcast,
mpsc::{self, Receiver, Sender, WeakSender},
Mutex as AsyncMutex,
};
use tracing::{debug, warn};

Expand Down Expand Up @@ -70,17 +71,18 @@ impl InProcessUsbHandler {
}

async fn process_selecting_device(
hid_devices: Vec<HidDevice>,
hid_devices: &[HidDevice],
) -> Result<UsbStateInternal, Error> {
let expected_answers = hid_devices.len();
let (blinking_tx, mut blinking_rx) =
tokio::sync::mpsc::channel::<Option<usize>>(expected_answers);
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();

Expand Down Expand Up @@ -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 => {
Expand All @@ -144,7 +146,7 @@ impl InProcessUsbHandler {
}

async fn process_select_credential(
response: GetAssertionResponse,
response: &GetAssertionResponse,
cred_rx: &mut Receiver<String>,
) -> Result<UsbStateInternal, Error> {
match cred_rx.recv().await {
Expand Down Expand Up @@ -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
}
Expand All @@ -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<AsyncMutex<HidDevice>>,
signal_tx: &Sender<Result<UsbUvMessage, Error>>,
) {
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);
}
Expand Down Expand Up @@ -403,7 +410,7 @@ pub(super) enum UsbStateInternal {
SelectingDevice(Vec<HidDevice>),

/// USB device connected, prompt user to tap
Connected(HidDevice),
Connected(Arc<AsyncMutex<HidDevice>>),

/// The device needs the PIN to be entered.
NeedsPin {
Expand Down