use std::collections::HashMap;
use std::sync::Mutex;
use smallvec::SmallVec;
use tokio::sync::broadcast;
use super::protocol_hid::get_key_name;
use crate::kernel::event::{BoardEvent, ComboKeyEvent, KeyPressEvent, KeySource};
#[derive(Debug, Clone)]
pub struct PressedKeyMeta {
pub key_index: usize,
pub key_name: String,
pub key_value: u16,
pub source: KeySource,
}
struct AggState {
state: HashMap<KeySource, HashMap<usize, PressedKeyMeta>>,
prev_all_pressed: SmallVec<[usize; 12]>,
last_config_mask: Option<u8>,
}
pub struct KeyStateAggregator {
inner: Mutex<AggState>,
event_tx: broadcast::Sender<BoardEvent>,
}
impl KeyStateAggregator {
pub fn new(event_tx: broadcast::Sender<BoardEvent>) -> Self {
Self {
inner: Mutex::new(AggState {
state: HashMap::new(),
prev_all_pressed: SmallVec::new(),
last_config_mask: None,
}),
event_tx,
}
}
pub fn report_change(
&self,
source: KeySource,
pressed_keys: Vec<PressedKeyMeta>,
config_mask: Option<u8>,
) {
let (released_events, pressed_events, combo_event) = {
let mut inner = self.inner.lock().unwrap_or_else(|e| e.into_inner());
if source == KeySource::Config {
inner.last_config_mask = config_mask;
}
let new_source_map: HashMap<usize, PressedKeyMeta> =
pressed_keys.into_iter().map(|m| (m.key_index, m)).collect();
inner.state.insert(source, new_source_map);
let mut all_pressed: SmallVec<[usize; 12]> = inner
.state
.values()
.flat_map(|m| m.keys())
.copied()
.collect();
all_pressed.sort();
all_pressed.dedup();
let mut newly_pressed: SmallVec<[usize; 12]> = SmallVec::new();
for &k in &all_pressed {
if !inner.prev_all_pressed.contains(&k) {
newly_pressed.push(k);
}
}
let mut released: SmallVec<[usize; 12]> = SmallVec::new();
for &k in &inner.prev_all_pressed {
if !all_pressed.contains(&k) {
released.push(k);
}
}
let all_sorted: Vec<usize> = all_pressed.iter().copied().collect();
let mut rel_events: Vec<KeyPressEvent> = Vec::new();
for &key_index in &released {
rel_events.push(KeyPressEvent {
key_index,
key_name: get_key_name(key_index).to_string(),
key_value: 0x0000,
pressed: false,
source,
});
}
let mut newly_sorted: Vec<usize> = newly_pressed.iter().copied().collect();
newly_sorted.sort();
let mut press_events: Vec<KeyPressEvent> = Vec::new();
for &key_index in &newly_sorted {
if let Some(meta) = Self::find_meta(&inner.state, key_index) {
press_events.push(KeyPressEvent {
key_index: meta.key_index,
key_name: meta.key_name.clone(),
key_value: meta.key_value,
pressed: true,
source: meta.source,
});
}
}
let combo = if all_pressed.len() > 1 && !newly_sorted.is_empty() {
Some(ComboKeyEvent {
keys: all_sorted.clone(),
key_names: all_sorted
.iter()
.map(|&i| get_key_name(i).to_string())
.collect(),
config_mask: inner.last_config_mask,
})
} else {
None
};
inner.prev_all_pressed = all_pressed;
(rel_events, press_events, combo)
};
for evt in released_events {
let _ = self.event_tx.send(BoardEvent::KeyPress(evt));
}
for evt in pressed_events {
log::debug!(target: "hid", "键按下: key_index={}", evt.key_index);
let _ = self.event_tx.send(BoardEvent::KeyPress(evt));
}
if let Some(combo) = combo_event {
let _ = self.event_tx.send(BoardEvent::ComboKey(combo));
}
}
fn find_meta(
state: &HashMap<KeySource, HashMap<usize, PressedKeyMeta>>,
key_index: usize,
) -> Option<&PressedKeyMeta> {
state.values().find_map(|m| m.get(&key_index))
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::kernel::protocol_hid::get_key_name;
fn meta(idx: usize, val: u16, source: KeySource) -> PressedKeyMeta {
PressedKeyMeta {
key_index: idx,
key_name: get_key_name(idx).to_string(),
key_value: val,
source,
}
}
fn drain(rx: &mut broadcast::Receiver<BoardEvent>) -> Vec<BoardEvent> {
let mut v = Vec::new();
while let Ok(e) = rx.try_recv() {
v.push(e);
}
v
}
#[test]
fn press_and_release_single_key() {
let (tx, mut rx) = broadcast::channel(64);
let agg = KeyStateAggregator::new(tx);
agg.report_change(
KeySource::Config,
vec![meta(6, 0x01, KeySource::Config)],
Some(0x01),
);
let evts = drain(&mut rx);
assert_eq!(evts.len(), 1);
match &evts[0] {
BoardEvent::KeyPress(k) => {
assert_eq!(k.key_index, 6);
assert!(k.pressed);
assert_eq!(k.source, KeySource::Config);
assert_eq!(k.key_value, 0x01);
}
_ => panic!("expected KeyPress"),
}
agg.report_change(KeySource::Config, vec![], None);
let evts = drain(&mut rx);
assert_eq!(evts.len(), 1);
if let BoardEvent::KeyPress(k) = &evts[0] {
assert_eq!(k.key_index, 6);
assert!(!k.pressed);
assert_eq!(k.key_value, 0x0000);
} else {
panic!("expected KeyPress release");
}
}
#[test]
fn combo_when_two_keys_pressed() {
let (tx, mut rx) = broadcast::channel(64);
let agg = KeyStateAggregator::new(tx);
agg.report_change(
KeySource::Config,
vec![
meta(6, 0x01, KeySource::Config),
meta(7, 0x02, KeySource::Config),
],
Some(0x03),
);
let evts = drain(&mut rx);
let presses: usize = evts
.iter()
.filter(|e| matches!(e, BoardEvent::KeyPress(k) if k.pressed))
.count();
let combos: Vec<&ComboKeyEvent> = evts
.iter()
.filter_map(|e| match e {
BoardEvent::ComboKey(c) => Some(c),
_ => None,
})
.collect();
assert_eq!(presses, 2, "应有 2 个按下事件");
assert_eq!(combos.len(), 1, "应有 1 个组合键事件");
assert_eq!(combos[0].keys, vec![6, 7]);
assert_eq!(combos[0].config_mask, Some(0x03));
}
#[test]
fn cross_source_merge() {
let (tx, mut rx) = broadcast::channel(64);
let agg = KeyStateAggregator::new(tx);
agg.report_change(
KeySource::Config,
vec![meta(6, 0x01, KeySource::Config)],
Some(0x01),
);
agg.report_change(
KeySource::Consumer,
vec![meta(3, 0x0F01, KeySource::Consumer)],
None,
);
let evts = drain(&mut rx);
let pressed_indices: Vec<usize> = evts
.iter()
.filter_map(|e| match e {
BoardEvent::KeyPress(k) if k.pressed => Some(k.key_index),
_ => None,
})
.collect();
assert_eq!(pressed_indices, vec![6, 3]);
}
}