Skip to main content

strat9_kernel/hardware/usb/
hid.rs

1// USB HID (Human Interface Device) Driver
2// Supports boot protocol keyboards and mice
3//
4// Features:
5// - Boot protocol keyboard support
6// - Boot protocol mouse support
7// - Event queue for key presses and mouse movements
8// - PS/2 to USB keycode translation
9// - Interrupt transfer polling via xHCI
10// - Unification with PS/2: events feed into the same keyboard/mouse buffers
11//
12// Inspired by Redox usbhid, Asterinas input subsystem, Maestro device manager.
13
14#![allow(dead_code)]
15
16use crate::arch::{keyboard, mouse};
17use alloc::{sync::Arc, vec::Vec};
18use core::sync::atomic::{AtomicBool, Ordering};
19use spin::Mutex;
20
21pub const HID_BOOT_KEYBOARD: u8 = 0x01;
22pub const HID_BOOT_MOUSE: u8 = 0x02;
23
24const KBD_REPORT_SIZE: usize = 8;
25const MOUSE_REPORT_SIZE: usize = 4;
26
27#[derive(Clone, Copy, Debug)]
28pub struct KeyEvent {
29    pub keycode: u8,
30    pub pressed: bool,
31    pub modifiers: u8,
32}
33
34#[derive(Clone, Copy, Debug)]
35pub struct MouseEvent {
36    pub dx: i8,
37    pub dy: i8,
38    pub dz: i8,
39    pub buttons: u8,
40}
41
42const USB_TO_PS2: [u8; 128] = [
43    0x00, 0x00, 0x00, 0x00, 0x1C, 0x32, 0x21, 0x23, 0x1D, 0x24, 0x2B, 0x34, 0x33, 0x43, 0x35, 0x0E,
44    0x15, 0x16, 0x17, 0x1C, 0x18, 0x19, 0x14, 0x1A, 0x1B, 0x1D, 0x1E, 0x21, 0x22, 0x23, 0x24, 0x2B,
45    0x29, 0x2F, 0x2E, 0x30, 0x20, 0x31, 0x32, 0x33, 0x2C, 0x2D, 0x11, 0x12, 0x13, 0x3F, 0x3E, 0x46,
46    0x45, 0x5D, 0x4C, 0x36, 0x4A, 0x55, 0x37, 0x4E, 0x57, 0x5E, 0x5C, 0x41, 0x52, 0x4D, 0x4B, 0x5B,
47    0x5A, 0x69, 0x6A, 0x6B, 0x6C, 0x6D, 0x6E, 0x6F, 0x70, 0x71, 0x72, 0x73, 0x74, 0x75, 0x76, 0x77,
48    0x78, 0x79, 0x7A, 0x7B, 0x7C, 0x7D, 0x7E, 0x7F, 0x80, 0x81, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00,
49    0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00,
50    0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00,
51];
52
53fn usb_to_ps2(keycode: u8) -> u8 {
54    if keycode < USB_TO_PS2.len() as u8 {
55        USB_TO_PS2[keycode as usize]
56    } else {
57        0x00
58    }
59}
60
61pub struct HidKeyboard {
62    port: usize,
63    slot_id: u8,
64    interface: u8,
65    endpoint: u8,
66    max_packet: u16,
67    interval: u8,
68    event_queue: Vec<KeyEvent>,
69    last_report: [u8; KBD_REPORT_SIZE],
70    report_buf: *mut u8,
71}
72
73unsafe impl Send for HidKeyboard {}
74unsafe impl Sync for HidKeyboard {}
75
76impl HidKeyboard {
77    pub fn new(
78        port: usize,
79        slot_id: u8,
80        interface: u8,
81        endpoint: u8,
82        max_packet: u16,
83        interval: u8,
84    ) -> Self {
85        Self {
86            port,
87            slot_id,
88            interface,
89            endpoint,
90            max_packet,
91            interval,
92            event_queue: Vec::new(),
93            last_report: [0; KBD_REPORT_SIZE],
94            report_buf: core::ptr::null_mut(),
95        }
96    }
97
98    pub fn read_event(&mut self) -> Option<KeyEvent> {
99        if self.event_queue.is_empty() {
100            None
101        } else {
102            Some(self.event_queue.remove(0))
103        }
104    }
105
106    pub fn process_report(&mut self, report: &[u8]) {
107        if report.len() < KBD_REPORT_SIZE {
108            return;
109        }
110        let modifiers = report[0];
111
112        for i in 2..8 {
113            let keycode = report[i];
114            if keycode == 0 {
115                continue;
116            }
117            let was_pressed = self.last_report[2..8].contains(&keycode);
118            if !was_pressed {
119                self.event_queue.push(KeyEvent {
120                    keycode: usb_to_ps2(keycode),
121                    pressed: true,
122                    modifiers,
123                });
124            }
125        }
126
127        for i in 2..8 {
128            let keycode = self.last_report[i];
129            if keycode != 0 && !report[2..8].contains(&keycode) {
130                self.event_queue.push(KeyEvent {
131                    keycode: usb_to_ps2(keycode),
132                    pressed: false,
133                    modifiers,
134                });
135            }
136        }
137
138        for i in 0..8 {
139            self.last_report[i] = report[i];
140        }
141    }
142
143    pub fn is_modifier_pressed(&self, modifier: u8) -> bool {
144        self.last_report[0] & modifier != 0
145    }
146
147    pub fn drain_into_unified(&mut self) {
148        for ev in self.event_queue.drain(..) {
149            keyboard::inject_hid_scancode(ev.keycode, ev.pressed);
150        }
151    }
152}
153
154pub struct HidMouse {
155    port: usize,
156    slot_id: u8,
157    interface: u8,
158    endpoint: u8,
159    max_packet: u16,
160    interval: u8,
161    event_queue: Vec<MouseEvent>,
162    last_buttons: u8,
163    report_buf: *mut u8,
164}
165
166unsafe impl Send for HidMouse {}
167unsafe impl Sync for HidMouse {}
168
169impl HidMouse {
170    pub fn new(
171        port: usize,
172        slot_id: u8,
173        interface: u8,
174        endpoint: u8,
175        max_packet: u16,
176        interval: u8,
177    ) -> Self {
178        Self {
179            port,
180            slot_id,
181            interface,
182            endpoint,
183            max_packet,
184            interval,
185            event_queue: Vec::new(),
186            last_buttons: 0,
187            report_buf: core::ptr::null_mut(),
188        }
189    }
190
191    pub fn read_event(&mut self) -> Option<MouseEvent> {
192        if self.event_queue.is_empty() {
193            None
194        } else {
195            Some(self.event_queue.remove(0))
196        }
197    }
198
199    pub fn process_report(&mut self, report: &[u8]) {
200        if report.len() < 3 {
201            return;
202        }
203
204        let buttons = report[0];
205        let dx = report[1] as i8;
206        let dy = report[2] as i8;
207        let dz = if report.len() > 3 { report[3] as i8 } else { 0 };
208
209        for i in 0..5 {
210            let mask = 1 << i;
211            let was_pressed = self.last_buttons & mask != 0;
212            let is_pressed = buttons & mask != 0;
213
214            if was_pressed != is_pressed {
215                self.event_queue.push(MouseEvent {
216                    dx: 0,
217                    dy: 0,
218                    dz: 0,
219                    buttons: if is_pressed { mask } else { 0 },
220                });
221            }
222        }
223
224        if dx != 0 || dy != 0 || dz != 0 {
225            self.event_queue.push(MouseEvent {
226                dx,
227                dy,
228                dz,
229                buttons,
230            });
231        }
232
233        self.last_buttons = buttons;
234    }
235
236    pub fn is_button_pressed(&self, button: u8) -> bool {
237        self.last_buttons & (1 << button) != 0
238    }
239
240    pub fn drain_into_unified(&mut self) {
241        for ev in self.event_queue.drain(..) {
242            let left = ev.buttons & 0x01 != 0;
243            let right = ev.buttons & 0x02 != 0;
244            let middle = ev.buttons & 0x04 != 0;
245            mouse::push_event_from_hid(ev.dx as i16, ev.dy as i16, ev.dz, left, right, middle);
246        }
247    }
248}
249
250static KEYBOARDS: Mutex<Vec<Arc<Mutex<HidKeyboard>>>> = Mutex::new(Vec::new());
251static MICE: Mutex<Vec<Arc<Mutex<HidMouse>>>> = Mutex::new(Vec::new());
252static HID_INITIALIZED: AtomicBool = AtomicBool::new(false);
253
254pub fn init() {
255    log::info!("[USB-HID] Initializing HID drivers...");
256    HID_INITIALIZED.store(true, Ordering::SeqCst);
257    log::info!(
258        "[USB-HID] Initialized: {} keyboard(s), {} mouse/mice",
259        KEYBOARDS.lock().len(),
260        MICE.lock().len()
261    );
262}
263
264pub fn enumerate_device(port: usize, slot_id: u8, dev_desc: &[u8; 18]) {
265    let dev_class = dev_desc[4];
266    unsafe {
267        core::arch::asm!("out 0xe9, al", in("al") b'h', options(nomem, nostack));
268        core::arch::asm!("out 0xe9, al", in("al") dev_class, options(nomem, nostack));
269        core::arch::asm!("out 0xe9, al", in("al") dev_desc[6], options(nomem, nostack));
270        core::arch::asm!("out 0xe9, al", in("al") b'\n', options(nomem, nostack));
271    }
272
273    if dev_class == 0x03 {
274        let protocol = dev_desc[6];
275        log::info!(
276            "[USB-HID] HID device: port={} slot={} class=03 protocol={:02x}",
277            port,
278            slot_id,
279            protocol
280        );
281
282        if let Some(controller_arc) = crate::hardware::usb::xhci::get_controller(0) {
283            let mut controller = controller_arc.lock();
284
285            if protocol == 1 {
286                controller.set_protocol(port as u8, 0, 0).ok();
287            } else if protocol == 2 {
288                controller.set_protocol(port as u8, 0, 1).ok();
289            }
290
291            let mut config_desc = [0u8; 256];
292            if controller
293                .get_configuration_descriptor(slot_id, 0, &mut config_desc, 9)
294                .is_ok()
295            {
296                let total_len = u16::from_le_bytes([config_desc[2], config_desc[3]]) as usize;
297                if total_len > 9 && total_len <= 256 {
298                    controller
299                        .get_configuration_descriptor(slot_id, 0, &mut config_desc, total_len)
300                        .ok();
301                }
302
303                let mut offset = 9;
304                while offset + 9 <= total_len {
305                    let b_length = config_desc[offset];
306                    let b_descriptor_type = config_desc[offset + 1];
307                    if b_length < 9 || offset + b_length as usize > total_len {
308                        break;
309                    }
310                    if b_descriptor_type == 4 {
311                        let b_interface_class = config_desc[offset + 5];
312                        let b_interface_protocol = config_desc[offset + 7];
313
314                        if b_interface_class == 0x03 {
315                            let mut ep_offset = offset + 9;
316                            while ep_offset + 7 <= offset + b_length as usize {
317                                let ep_b_length = config_desc[ep_offset];
318                                let ep_b_descriptor_type = config_desc[ep_offset + 1];
319                                if ep_b_length < 7 || ep_b_descriptor_type != 5 {
320                                    break;
321                                }
322                                let ep_addr = config_desc[ep_offset + 2];
323                                let ep_max_packet = u16::from_le_bytes([
324                                    config_desc[ep_offset + 4],
325                                    config_desc[ep_offset + 5],
326                                ]);
327                                let ep_interval = config_desc[ep_offset + 6];
328
329                                if (ep_addr & 0x80) != 0 {
330                                    let ep_num = ep_addr & 0x0F;
331                                    let ep_type = 7;
332                                    unsafe {
333                                        core::arch::asm!("out 0xe9, al", in("al") b'K', options(nomem, nostack));
334                                        core::arch::asm!("out 0xe9, al", in("al") b'b', options(nomem, nostack));
335                                        core::arch::asm!("out 0xe9, al", in("al") b'0'+b_interface_protocol, options(nomem, nostack));
336                                        core::arch::asm!("out 0xe9, al", in("al") b'\n', options(nomem, nostack));
337                                    }
338
339                                    let setup_ok = controller
340                                        .setup_endpoint(
341                                            slot_id,
342                                            ep_num,
343                                            ep_max_packet as u32,
344                                            ep_type,
345                                            ep_interval as u32,
346                                            0,
347                                        )
348                                        .is_ok();
349                                    unsafe {
350                                        core::arch::asm!("out 0xe9, al", in("al") b'K', options(nomem, nostack));
351                                        core::arch::asm!("out 0xe9, al", in("al") if setup_ok { b's' } else { b'S' }, options(nomem, nostack));
352                                        core::arch::asm!("out 0xe9, al", in("al") b'\n', options(nomem, nostack));
353                                    }
354
355                                    let buf_size = ep_max_packet as usize;
356                                    let alloc_res = controller
357                                        .alloc_interrupt_buffer(slot_id, ep_num, buf_size);
358                                    unsafe {
359                                        core::arch::asm!("out 0xe9, al", in("al") b'K', options(nomem, nostack));
360                                        core::arch::asm!("out 0xe9, al", in("al") if alloc_res.is_ok() { b'a' } else { b'A' }, options(nomem, nostack));
361                                        core::arch::asm!("out 0xe9, al", in("al") b'\n', options(nomem, nostack));
362                                    }
363                                    if let Ok((_buf_virt, _buf_phys)) = alloc_res {
364                                        if b_interface_protocol == 1 {
365                                            let mut keyboard = HidKeyboard::new(
366                                                port,
367                                                slot_id,
368                                                config_desc[offset + 2],
369                                                ep_addr,
370                                                ep_max_packet,
371                                                ep_interval,
372                                            );
373                                            keyboard.report_buf = _buf_virt;
374                                            log::info!(
375                                                "[USB-HID] Keyboard: port={} slot={} ep={:02x} max_pkt={} interval={}",
376                                                port,
377                                                slot_id,
378                                                ep_addr,
379                                                ep_max_packet,
380                                                ep_interval
381                                            );
382                                            KEYBOARDS.lock().push(Arc::new(Mutex::new(keyboard)));
383                                            unsafe {
384                                                core::arch::asm!("out 0xe9, al", in("al") b'B', options(nomem, nostack));
385                                                core::arch::asm!("out 0xe9, al", in("al") b'!', options(nomem, nostack));
386                                            }
387
388                                            controller
389                                                .submit_interrupt_transfer(slot_id, ep_num)
390                                                .ok();
391                                        } else if b_interface_protocol == 2 {
392                                            let mut mouse_dev = HidMouse::new(
393                                                port,
394                                                slot_id,
395                                                config_desc[offset + 2],
396                                                ep_addr,
397                                                ep_max_packet,
398                                                ep_interval,
399                                            );
400                                            mouse_dev.report_buf = _buf_virt;
401                                            log::info!(
402                                                "[USB-HID] Mouse: port={} slot={} ep={:02x} max_pkt={} interval={}",
403                                                port,
404                                                slot_id,
405                                                ep_addr,
406                                                ep_max_packet,
407                                                ep_interval
408                                            );
409                                            MICE.lock().push(Arc::new(Mutex::new(mouse_dev)));
410                                            unsafe {
411                                                core::arch::asm!("out 0xe9, al", in("al") b'M', options(nomem, nostack));
412                                                core::arch::asm!("out 0xe9, al", in("al") b'!', options(nomem, nostack));
413                                            }
414
415                                            controller
416                                                .submit_interrupt_transfer(slot_id, ep_num)
417                                                .ok();
418                                        }
419                                    }
420                                }
421                                ep_offset += ep_b_length as usize;
422                            }
423                        }
424                    }
425                    offset += b_length as usize;
426                }
427            }
428
429            controller.set_configuration(slot_id, 1).ok();
430        }
431    } else {
432        log::info!(
433            "[USB-HID] Non-HID device: port={} slot={} class={:02x}",
434            port,
435            slot_id,
436            dev_class
437        );
438    }
439}
440
441/// # Safety
442///
443/// - `buf` must point to `len` readable bytes of a valid HID report.
444pub unsafe fn receive_interrupt_report(slot_id: u8, ep_id: u8, buf: *const u8, len: usize) {
445    if buf.is_null() || len == 0 {
446        return;
447    }
448
449    let report = unsafe { core::slice::from_raw_parts(buf, len) };
450
451    for kbd in KEYBOARDS.lock().iter() {
452        let mut k = kbd.lock();
453        if k.slot_id == slot_id && (k.endpoint & 0x0F) == ep_id {
454            k.process_report(report);
455            k.drain_into_unified();
456            return;
457        }
458    }
459
460    for m in MICE.lock().iter() {
461        let mut dev = m.lock();
462        if dev.slot_id == slot_id && (dev.endpoint & 0x0F) == ep_id {
463            dev.process_report(report);
464            dev.drain_into_unified();
465            return;
466        }
467    }
468}
469
470pub fn get_keyboard(index: usize) -> Option<Arc<Mutex<HidKeyboard>>> {
471    KEYBOARDS.lock().get(index).cloned()
472}
473
474pub fn get_mouse(index: usize) -> Option<Arc<Mutex<HidMouse>>> {
475    MICE.lock().get(index).cloned()
476}
477
478pub fn keyboard_count() -> usize {
479    KEYBOARDS.lock().len()
480}
481
482pub fn mouse_count() -> usize {
483    MICE.lock().len()
484}
485
486pub fn is_available() -> bool {
487    HID_INITIALIZED.load(Ordering::Relaxed)
488}
489
490pub fn poll_all() {
491    for kbd in KEYBOARDS.lock().iter() {
492        let mut k = kbd.lock();
493        k.drain_into_unified();
494    }
495    for m in MICE.lock().iter() {
496        let mut dev = m.lock();
497        dev.drain_into_unified();
498    }
499}
500
501pub fn notify_transfer_complete(_slot_id: u8, _ep_id: u8) {
502    poll_all();
503}