use std::{
collections::{HashMap, VecDeque},
sync::{
Arc, Mutex, MutexGuard, Weak,
atomic::{AtomicBool, AtomicU8, AtomicU64, Ordering},
},
thread::{self, JoinHandle},
time::Duration,
};
use futures::{FutureExt, channel::oneshot, select};
use rand::Rng;
use tracing::trace;
use crate::nibble::U4;
mod error;
mod message;
mod raw;
#[cfg(test)]
#[allow(
clippy::unwrap_used,
clippy::expect_used,
reason = "expect/unwrap are idiomatic in tests"
)]
pub(crate) mod tests;
pub use error::ChannelError;
pub use message::{
HidppMessage, LONG_REPORT_ID, LONG_REPORT_LENGTH, SHORT_REPORT_ID, SHORT_REPORT_LENGTH,
};
pub use raw::RawHidChannel;
use raw::supports_short_long_hidpp;
const MAX_REPORT_LENGTH: usize = LONG_REPORT_LENGTH;
const MAX_RAW_REPORT_LENGTH: usize = 64;
pub const SEND_RESPONSE_TIMEOUT: Duration = Duration::from_secs(5);
type MessageListener = Arc<dyn Fn(HidppMessage, bool) + Send + Sync + 'static>;
#[expect(
clippy::expect_used,
reason = "mutex poisoning is unrecoverable here — see doc comment"
)]
fn lock<T>(mutex: &Mutex<T>) -> MutexGuard<'_, T> {
mutex.lock().expect("mutex poisoned")
}
pub struct MessageListenerGuard {
message_listeners: Weak<Mutex<HashMap<u32, MessageListener>>>,
hdl: u32,
}
impl Drop for MessageListenerGuard {
fn drop(&mut self) {
if let Some(message_listeners) = self.message_listeners.upgrade() {
lock(&message_listeners).remove(&self.hdl);
}
}
}
pub struct HidppChannel {
pub supports_short: bool,
pub supports_long: bool,
pub vendor_id: u16,
pub product_id: u16,
raw_channel: Arc<dyn RawHidChannel>,
rotate_software_id: AtomicBool,
software_id: AtomicU8,
pending_messages: Arc<Mutex<VecDeque<PendingMessage>>>,
pending_message_id: AtomicU64,
message_listeners: Arc<Mutex<HashMap<u32, MessageListener>>>,
read_thread_close: Option<oneshot::Sender<()>>,
read_thread_hdl: Option<JoinHandle<()>>,
sw_id_lease: Option<(u8, fn(u8))>,
}
impl Drop for HidppChannel {
fn drop(&mut self) {
if let Some((id, free)) = self.sw_id_lease.take() {
free(id);
}
if let Some(read_thread_close) = self.read_thread_close.take() {
let _ = read_thread_close.send(());
}
if let Some(read_thread_hdl) = self.read_thread_hdl.take() {
#[expect(
clippy::unwrap_used,
reason = "propagate a read-thread panic instead of ignoring a crashed background worker"
)]
read_thread_hdl.join().unwrap();
}
}
}
struct PendingMessage {
id: u64,
response_predicate: Box<dyn Fn(&HidppMessage) -> bool + Send>,
sender: oneshot::Sender<HidppMessage>,
}
impl HidppChannel {
pub async fn from_raw_channel(raw: impl RawHidChannel) -> Result<Self, ChannelError> {
let (supports_short, supports_long) = supports_short_long_hidpp(&raw).await?;
if !supports_short && !supports_long {
return Err(ChannelError::HidppNotSupported);
}
let raw_channel_rc = Arc::new(raw);
let pending_messages_rc = Arc::new(Mutex::new(VecDeque::<PendingMessage>::new()));
let message_listeners_rc = Arc::new(Mutex::new(HashMap::<u32, MessageListener>::new()));
let (close_sender, close_receiver) = oneshot::channel::<()>();
let read_thread_hdl = thread::spawn({
let raw_channel = Arc::clone(&raw_channel_rc);
let pending_messages = Arc::clone(&pending_messages_rc);
let message_listeners = Arc::clone(&message_listeners_rc);
move || {
futures::executor::block_on(read_loop(
&*raw_channel,
&pending_messages,
&message_listeners,
close_receiver,
));
}
});
Ok(Self {
supports_short,
supports_long,
vendor_id: raw_channel_rc.vendor_id(),
product_id: raw_channel_rc.product_id(),
raw_channel: raw_channel_rc,
rotate_software_id: AtomicBool::new(false),
software_id: AtomicU8::new(0x01),
pending_messages: pending_messages_rc,
pending_message_id: AtomicU64::new(1),
message_listeners: message_listeners_rc,
read_thread_close: Some(close_sender),
read_thread_hdl: Some(read_thread_hdl),
sw_id_lease: None,
})
}
pub fn is_connected(&self) -> bool {
self.raw_channel.is_connected()
}
pub fn set_sw_id(&self, sw_id: U4) {
self.software_id.store(sw_id.to_lo(), Ordering::SeqCst);
}
pub fn set_rotating_sw_id(&self, enable: bool) {
self.rotate_software_id.store(enable, Ordering::SeqCst);
}
pub fn set_sw_id_lease(&mut self, id: u8, free: fn(u8)) {
self.sw_id_lease = Some((id, free));
}
pub fn get_sw_id(&self) -> U4 {
if self.rotate_software_id.load(Ordering::SeqCst) {
let previous =
match self
.software_id
.fetch_update(Ordering::SeqCst, Ordering::SeqCst, |old| {
Some(if old & 0x0f == 0x0f {
0x01
} else {
old.wrapping_add(1)
})
}) {
Ok(previous) | Err(previous) => previous,
};
U4::from_lo(previous)
} else {
U4::from_lo(self.software_id.load(Ordering::SeqCst))
}
}
pub fn supports_msg(&self, msg: &HidppMessage) -> bool {
match msg {
HidppMessage::Short(_) => self.supports_short,
HidppMessage::Long(_) => self.supports_long,
}
}
fn normalize_outgoing(&self, msg: HidppMessage) -> HidppMessage {
match msg {
HidppMessage::Short(_) if !self.supports_short && self.supports_long => msg.widened(),
other => other,
}
}
pub async fn send(
&self,
msg: HidppMessage,
response_predicate: impl Fn(&HidppMessage) -> bool + Send + 'static,
) -> Result<HidppMessage, ChannelError> {
self.send_with_timeout(msg, response_predicate, SEND_RESPONSE_TIMEOUT)
.await
}
pub async fn send_with_timeout(
&self,
msg: HidppMessage,
response_predicate: impl Fn(&HidppMessage) -> bool + Send + 'static,
timeout: Duration,
) -> Result<HidppMessage, ChannelError> {
let msg = self.normalize_outgoing(msg);
if !self.supports_msg(&msg) {
return Err(ChannelError::MessageTypeNotSupported);
}
let (dev, feat, func) = msg.header();
trace!(dev, feat, func, "hidpp request");
let (sender, receiver) = oneshot::channel::<HidppMessage>();
let pending_id = self.pending_message_id.fetch_add(1, Ordering::SeqCst);
{
let mut pending = lock(&self.pending_messages);
pending.retain(|m| !m.sender.is_canceled());
pending.push_back(PendingMessage {
id: pending_id,
response_predicate: Box::new(response_predicate),
sender,
});
}
let mut request = std::pin::pin!(
async {
self.send_and_forget(msg).await?;
receiver.await.map_err(|_| ChannelError::NoResponse)
}
.fuse()
);
let result = select! {
result = request => result,
() = futures_timer::Delay::new(timeout).fuse() => Err(ChannelError::Timeout),
};
match &result {
Ok(_) => trace!(dev, feat, "hidpp response"),
Err(e) => trace!(dev, feat, error = ?e, "hidpp no response"),
}
if result.is_err() {
self.remove_pending_message(pending_id);
}
result
}
fn remove_pending_message(&self, id: u64) {
let mut pending = lock(&self.pending_messages);
if let Some(pos) = pending.iter().position(|msg| msg.id == id) {
pending.remove(pos);
}
}
pub async fn send_and_forget(&self, msg: HidppMessage) -> Result<(), ChannelError> {
let msg = self.normalize_outgoing(msg);
if !self.supports_msg(&msg) {
return Err(ChannelError::MessageTypeNotSupported);
}
let mut buf = [0u8; LONG_REPORT_LENGTH];
let len = msg.write_raw(&mut buf);
self.raw_channel
.write_report(&buf[..len])
.await
.map(|_| ())
.map_err(ChannelError::Implementation)
}
pub async fn write_raw_report(&self, report: &[u8]) -> Result<usize, ChannelError> {
self.write_raw_report_with_timeout(report, SEND_RESPONSE_TIMEOUT)
.await
}
async fn write_raw_report_with_timeout(
&self,
report: &[u8],
timeout: Duration,
) -> Result<usize, ChannelError> {
if !(1..=MAX_RAW_REPORT_LENGTH).contains(&report.len()) {
return Err(ChannelError::InvalidRawReportLength(report.len()));
}
let mut write = std::pin::pin!(self.raw_channel.write_report(report).fuse());
select! {
result = write => result.map_err(ChannelError::Implementation),
() = futures_timer::Delay::new(timeout).fuse() => Err(ChannelError::Timeout),
}
}
pub fn add_msg_listener(
&self,
listener: impl Fn(HidppMessage, bool) + Send + Sync + 'static,
) -> u32 {
let mut listeners = lock(&self.message_listeners);
let mut rng = rand::rng();
let mut hdl = rng.random::<u32>();
while listeners.contains_key(&hdl) {
hdl = rng.random::<u32>();
}
listeners.insert(hdl, Arc::new(listener));
hdl
}
pub fn add_msg_listener_guarded(
&self,
listener: impl Fn(HidppMessage, bool) + Send + Sync + 'static,
) -> MessageListenerGuard {
let hdl = self.add_msg_listener(listener);
MessageListenerGuard {
message_listeners: Arc::downgrade(&self.message_listeners),
hdl,
}
}
pub fn remove_msg_listener(&self, hdl: u32) -> bool {
lock(&self.message_listeners).remove(&hdl).is_some()
}
}
async fn read_loop(
raw_channel: &dyn RawHidChannel,
pending_messages: &Mutex<VecDeque<PendingMessage>>,
message_listeners: &Mutex<HashMap<u32, MessageListener>>,
mut close: oneshot::Receiver<()>,
) {
let mut buf = [0u8; MAX_REPORT_LENGTH];
loop {
let res = select! {
_ = close => break,
res = raw_channel.read_report(&mut buf).fuse() => res,
};
let len = match res {
Ok(len) => len,
Err(error) => {
trace!(?error, "read_report error");
continue;
}
};
let Some(msg) = HidppMessage::read_raw(&buf[..len]) else {
trace!(len, "report not HID++ — dropped");
continue;
};
let mut matched = false;
let pending_count;
{
let mut msgs = lock(pending_messages);
pending_count = msgs.len();
if let Some(pos) = msgs.iter().position(|elem| (elem.response_predicate)(&msg))
&& let Some(waiting) = msgs.remove(pos)
{
let _ = waiting.sender.send(msg);
matched = true;
}
}
trace!(
len,
matched,
pending_count,
payload = format!("{:02x?}", &buf[..len.min(16)]),
"raw report received"
);
let listeners: Vec<_> = lock(message_listeners).values().cloned().collect();
for listener in listeners {
listener(msg, matched);
}
}
}