use std::{collections::HashMap, sync::Arc};
use hidpp::{
channel::{HidppChannel, HidppMessage},
receiver::{self, Receiver},
};
use serde::{Deserialize, Serialize};
use thiserror::Error;
use tokio::sync::mpsc;
use tracing::{debug, trace};
pub use hidpp::receiver::bolt::DeviceKind as BoltDeviceKind;
use crate::transport::{enumerate_hidpp_devices, open_hidpp_channel};
mod notification;
mod registers;
use notification::{Notification, decode, parse_notification, subscribe};
use registers::{
BOLT_DISCOVERY, BOLT_PAIRING, NOTIFICATION_FLAGS, NOTIFICATIONS, UNIFYING_PAIRING,
write_long_register, write_register,
};
const RECEIVER_INDEX: u8 = 0xff;
#[derive(Clone, Copy, PartialEq, Eq, Debug)]
pub enum ReceiverFamily {
Bolt,
Unifying,
}
#[derive(Clone, Copy, PartialEq, Eq, Debug)]
enum PairingPhase {
BoltDiscovery,
BoltPairing,
UnifyingPairing,
}
impl From<ReceiverFamily> for PairingPhase {
fn from(family: ReceiverFamily) -> Self {
match family {
ReceiverFamily::Bolt => Self::BoltDiscovery,
ReceiverFamily::Unifying => Self::UnifyingPairing,
}
}
}
fn family_for(product_id: u16) -> Option<ReceiverFamily> {
if crate::BOLT_PIDS.contains(&product_id) {
Some(ReceiverFamily::Bolt)
} else if crate::speaks_unifying_protocol(product_id) {
Some(ReceiverFamily::Unifying)
} else {
None
}
}
#[derive(Clone, Debug)]
pub struct PairingReceiver {
pub uid: Option<String>,
pub family: ReceiverFamily,
pub product_id: u16,
}
#[derive(Clone, Debug, Serialize, Deserialize)]
pub enum ReceiverSelector {
First,
BoltUid(String),
}
#[derive(Clone, Debug)]
pub struct DiscoveredDevice {
pub address: [u8; 6],
pub authentication: u8,
pub kind: BoltDeviceKind,
pub name: String,
}
impl DiscoveredDevice {
#[must_use]
pub fn passkey_on_keyboard(&self) -> bool {
self.authentication & 0x01 != 0
}
fn entropy(&self) -> u8 {
if self.kind == BoltDeviceKind::Keyboard {
20
} else {
10
}
}
}
#[derive(Clone, Copy, PartialEq, Eq, Debug, Serialize, Deserialize)]
pub enum Click {
Left,
Right,
}
#[derive(Clone, Debug, Serialize, Deserialize)]
pub enum PasskeyMethod {
Keyboard(String),
Pointer {
passkey: String,
clicks: Vec<Click>,
},
}
fn passkey_to_clicks(value: u32) -> Vec<Click> {
(0..10)
.rev()
.map(|bit| {
if value & (1 << bit) != 0 {
Click::Right
} else {
Click::Left
}
})
.collect()
}
#[derive(Clone, Debug)]
pub enum PairingEvent {
Searching,
DeviceFound(DiscoveredDevice),
Passkey(PasskeyMethod),
Paired {
slot: u8,
},
Failed(PairingError),
}
#[derive(Clone, Debug)]
pub enum PairingCommand {
Pair(DiscoveredDevice),
Cancel,
}
#[derive(Clone, Debug, Error)]
pub enum PairingError {
#[error("HID transport error: {0}")]
Hid(String),
#[error("no supported pairing-capable receiver found")]
ReceiverNotFound,
#[error("receiver register access failed: {0}")]
Register(String),
#[error("pairing timed out")]
Timeout,
#[error("receiver reported pairing error {0:#04x}")]
Device(u8),
#[error("pairing was cancelled")]
Cancelled,
#[error("malformed pairing notification ({0})")]
MalformedNotification(&'static str),
}
impl From<async_hid::HidError> for PairingError {
fn from(e: async_hid::HidError) -> Self {
PairingError::Hid(e.to_string())
}
}
pub async fn list_pairing_receivers() -> Result<Vec<PairingReceiver>, PairingError> {
let mut out = Vec::new();
for dev in enumerate_hidpp_devices().await? {
let Some((_, channel)) = open_hidpp_channel(dev).await? else {
continue;
};
let Some(family) = family_for(channel.product_id) else {
continue;
};
let uid = match family {
ReceiverFamily::Bolt => read_bolt_uid(&channel).await,
ReceiverFamily::Unifying => None,
};
out.push(PairingReceiver {
uid,
family,
product_id: channel.product_id,
});
}
Ok(out)
}
async fn read_bolt_uid(channel: &Arc<HidppChannel>) -> Option<String> {
let Some(Receiver::Bolt(bolt)) = receiver::detect(Arc::clone(channel)) else {
return None;
};
bolt.get_unique_id().await.ok()
}
async fn open_receiver(
target: &ReceiverSelector,
) -> Result<(Arc<HidppChannel>, ReceiverFamily), PairingError> {
for dev in enumerate_hidpp_devices().await? {
let Some((_, channel)) = open_hidpp_channel(dev).await? else {
continue;
};
let Some(family) = family_for(channel.product_id) else {
continue;
};
match target {
ReceiverSelector::First => return Ok((channel, family)),
ReceiverSelector::BoltUid(want) => {
if family == ReceiverFamily::Bolt
&& read_bolt_uid(&channel)
.await
.is_some_and(|uid| uid.eq_ignore_ascii_case(want))
{
return Ok((channel, family));
}
}
}
}
Err(PairingError::ReceiverNotFound)
}
const SESSION_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(90);
const DISCOVERY_TIMEOUT: u8 = 30;
pub async fn run_pairing(
target: ReceiverSelector,
mut commands: mpsc::UnboundedReceiver<PairingCommand>,
events: mpsc::UnboundedSender<PairingEvent>,
) -> Result<(), PairingError> {
let (channel, family) = match open_receiver(&target).await {
Ok(receiver) => receiver,
Err(e) => {
let _ = events.send(PairingEvent::Failed(e.clone()));
return Err(e);
}
};
let (listener, mut notifications) = subscribe(&channel);
let result = run_session(&channel, family, &mut commands, &mut notifications, &events).await;
drop(listener);
let _ = channel
.write_register(RECEIVER_INDEX, NOTIFICATIONS, [0, 0, 0])
.await;
if let Err(ref e) = result {
let _ = events.send(PairingEvent::Failed(e.clone()));
}
result
}
async fn run_session(
channel: &HidppChannel,
family: ReceiverFamily,
commands: &mut mpsc::UnboundedReceiver<PairingCommand>,
notifications: &mut mpsc::UnboundedReceiver<HidppMessage>,
events: &mpsc::UnboundedSender<PairingEvent>,
) -> Result<(), PairingError> {
let mut phase = PairingPhase::from(family);
let result = drive(channel, family, &mut phase, commands, notifications, events).await;
if result.is_err() {
cancel(channel, phase).await;
}
result
}
async fn drive(
channel: &HidppChannel,
family: ReceiverFamily,
phase: &mut PairingPhase,
commands: &mut mpsc::UnboundedReceiver<PairingCommand>,
notifications: &mut mpsc::UnboundedReceiver<HidppMessage>,
events: &mpsc::UnboundedSender<PairingEvent>,
) -> Result<(), PairingError> {
write_register(channel, NOTIFICATIONS, NOTIFICATION_FLAGS).await?;
match family {
ReceiverFamily::Bolt => {
write_register(channel, BOLT_DISCOVERY, [DISCOVERY_TIMEOUT, 0x01, 0x00]).await?;
}
ReceiverFamily::Unifying => {
write_register(channel, UNIFYING_PAIRING, [0x01, 0x00, DISCOVERY_TIMEOUT]).await?;
}
}
let _ = events.send(PairingEvent::Searching);
let mut partial: HashMap<u16, PartialDevice> = HashMap::new();
let mut pairing_auth: Option<u8> = None;
let deadline = tokio::time::sleep(SESSION_TIMEOUT);
tokio::pin!(deadline);
loop {
tokio::select! {
() = &mut deadline => return Err(PairingError::Timeout),
cmd = commands.recv() => match cmd {
Some(PairingCommand::Pair(device)) => {
pairing_auth = Some(device.authentication);
if *phase == PairingPhase::BoltDiscovery {
*phase = PairingPhase::BoltPairing;
}
pair_bolt_device(channel, &device).await?;
}
Some(PairingCommand::Cancel) | None => {
return Err(PairingError::Cancelled);
}
},
msg = notifications.recv() => {
let Some(msg) = msg else {
return Err(PairingError::Hid("receiver channel closed".into()));
};
let (device_index, sub_id, payload) = decode(&msg);
trace!(sub_id = format_args!("{sub_id:#04x}"), ?payload, "pairing notification");
let Some(note) = parse_notification(sub_id, device_index, payload) else {
continue;
};
match note {
Notification::DiscoveryInfo { counter, kind, address, authentication } => {
let entry = partial.entry(counter).or_default();
entry.kind = Some(kind);
entry.address = Some(address);
entry.authentication = Some(authentication);
if let Some(device) = entry.build() {
let _ = events.send(PairingEvent::DeviceFound(device));
}
}
Notification::DiscoveryName { counter, name } => {
let entry = partial.entry(counter).or_default();
entry.name = Some(name);
if let Some(device) = entry.build() {
let _ = events.send(PairingEvent::DeviceFound(device));
}
}
Notification::Passkey { digits, value } => {
let method = match pairing_auth {
Some(auth) if auth & 0x01 != 0 => PasskeyMethod::Keyboard(digits),
_ => PasskeyMethod::Pointer {
clicks: passkey_to_clicks(value),
passkey: digits,
},
};
let _ = events.send(PairingEvent::Passkey(method));
}
Notification::MalformedPasskey => {
return Err(PairingError::MalformedNotification("passkey digits"));
}
Notification::PairingSucceeded { slot } => {
let _ = events.send(PairingEvent::Paired { slot });
return Ok(());
}
Notification::PairingError(code) => return Err(PairingError::Device(code)),
Notification::Connected { slot, established } if family == ReceiverFamily::Unifying => {
if established {
let _ = events.send(PairingEvent::Paired { slot });
return Ok(());
}
}
Notification::Connected { .. } => {}
Notification::UnifyingLock { open, error } => {
if error != 0 {
return Err(PairingError::Device(error));
}
if !open {
return Err(PairingError::Timeout);
}
}
}
}
}
}
}
#[derive(Default)]
struct PartialDevice {
kind: Option<u8>,
address: Option<[u8; 6]>,
authentication: Option<u8>,
name: Option<String>,
emitted: bool,
}
impl PartialDevice {
fn build(&mut self) -> Option<DiscoveredDevice> {
if self.emitted {
return None;
}
let (kind, address, authentication, name) = (
self.kind?,
self.address?,
self.authentication?,
self.name.clone()?,
);
self.emitted = true;
Some(DiscoveredDevice {
address,
authentication,
kind: BoltDeviceKind::from(kind & 0x0f),
name,
})
}
}
async fn pair_bolt_device(
channel: &HidppChannel,
device: &DiscoveredDevice,
) -> Result<(), PairingError> {
let mut payload = [0u8; 16];
payload[0] = 0x01; payload[1] = 0x00; payload[2..8].copy_from_slice(&device.address);
payload[8] = device.authentication;
payload[9] = device.entropy();
write_long_register(channel, BOLT_PAIRING, payload).await
}
async fn cancel(channel: &HidppChannel, phase: PairingPhase) {
let res = match phase {
PairingPhase::BoltDiscovery => {
write_register(channel, BOLT_DISCOVERY, [DISCOVERY_TIMEOUT, 0x02, 0x00]).await
}
PairingPhase::BoltPairing => {
let mut payload = [0u8; 16];
payload[0] = 0x02;
write_long_register(channel, BOLT_PAIRING, payload).await
}
PairingPhase::UnifyingPairing => {
write_register(channel, UNIFYING_PAIRING, [0x02, 0x00, 0x00]).await
}
};
if let Err(e) = res {
debug!(?phase, ?e, "cancel write failed");
}
}
pub async fn unpair(target: ReceiverSelector, slot: u8) -> Result<(), PairingError> {
let (channel, family) = open_receiver(&target).await?;
match family {
ReceiverFamily::Bolt => {
let mut payload = [0u8; 16];
payload[0] = 0x03; payload[1] = slot;
write_long_register(&channel, BOLT_PAIRING, payload).await
}
ReceiverFamily::Unifying => {
write_register(&channel, UNIFYING_PAIRING, [0x03, slot, 0x00]).await
}
}
}
#[cfg(test)]
mod tests;