diff --git a/kernel/include/kernel/USB/Device.h b/kernel/include/kernel/USB/Device.h index 03ede93a..3f07a936 100644 --- a/kernel/include/kernel/USB/Device.h +++ b/kernel/include/kernel/USB/Device.h @@ -23,6 +23,8 @@ namespace Kernel virtual void handle_stall(uint8_t endpoint_id) = 0; virtual void handle_input_data(size_t byte_count, uint8_t endpoint_id) = 0; + + virtual bool is_hid_driver() const { return false; } }; class USBDevice @@ -62,6 +64,12 @@ namespace Kernel uint8_t tt_think_time; }; + struct HIDInfo + { + BAN::Atomic led_mask { 0 }; + BAN::Vector led_controls; + }; + public: USBDevice(USBController& controller, USB::SpeedClass speed_class, uint8_t depth) : m_controller(controller) @@ -92,6 +100,8 @@ namespace Kernel void register_hub_to_init() { m_controller.register_hub_to_init(m_depth + 1); }; void mark_hub_init_done() { m_controller.mark_hub_init_done(m_depth + 1); }; + void update_led_mask(uint32_t led_mask); + protected: void handle_stall(uint8_t endpoint_id); void handle_input_data(size_t byte_count, uint8_t endpoint_id); @@ -112,6 +122,7 @@ namespace Kernel BAN::UniqPtr m_dma_buffer; BAN::Vector> m_class_drivers; + HIDInfo m_hid_info; }; } diff --git a/kernel/include/kernel/USB/HID/HIDDriver.h b/kernel/include/kernel/USB/HID/HIDDriver.h index 0a581d0c..97c860e5 100644 --- a/kernel/include/kernel/USB/HID/HIDDriver.h +++ b/kernel/include/kernel/USB/HID/HIDDriver.h @@ -74,6 +74,7 @@ namespace Kernel struct DeviceReport { BAN::Vector inputs; + BAN::Vector outputs; BAN::RefPtr device; }; @@ -81,27 +82,43 @@ namespace Kernel void handle_stall(uint8_t endpoint_id) override; void handle_input_data(size_t byte_count, uint8_t endpoint_id) override; + bool is_hid_driver() const override { return true; } + USBDevice& device() { return m_device; } const USBDevice::InterfaceDescriptor& interface() const { return m_interface; } + bool has_led_control() const { return !m_led_controls.empty(); } + void set_leds(uint32_t led_mask); + private: USBHIDDriver(USBDevice&, const USBDevice::InterfaceDescriptor&); ~USBHIDDriver(); BAN::ErrorOr initialize() override; - BAN::ErrorOr> initializes_device_reports(const BAN::Vector&); + BAN::ErrorOr initializes_device_reports(const BAN::Vector&); + + private: + struct LEDControl + { + DeviceReport* report; + uint32_t report_id; + uint32_t report_bits; + }; private: USBDevice& m_device; USBDevice::InterfaceDescriptor m_interface; bool m_uses_report_id { false }; - BAN::Vector m_device_inputs; + BAN::Vector m_device_reports; uint8_t m_data_endpoint_id = 0; BAN::UniqPtr m_data_buffer; + BAN::Vector m_led_controls; + BAN::UniqPtr m_led_region; + friend class BAN::UniqPtr; }; diff --git a/kernel/include/kernel/USB/HID/Keyboard.h b/kernel/include/kernel/USB/HID/Keyboard.h index 4359dd34..d3d1a5ea 100644 --- a/kernel/include/kernel/USB/HID/Keyboard.h +++ b/kernel/include/kernel/USB/HID/Keyboard.h @@ -11,8 +11,6 @@ namespace Kernel BAN_NON_MOVABLE(USBKeyboard); public: - BAN::ErrorOr initialize() override; - void start_report() override; void stop_report() override; @@ -23,12 +21,9 @@ namespace Kernel void update() override; private: - USBKeyboard(USBHIDDriver& driver, BAN::Vector&& outputs); + USBKeyboard(USBHIDDriver& driver); ~USBKeyboard() = default; - void set_leds(uint16_t mask); - void set_leds(uint8_t report_id, uint16_t mask); - private: USBHIDDriver& m_driver; @@ -38,11 +33,7 @@ namespace Kernel BAN::Array m_keyboard_state { false }; BAN::Array m_keyboard_state_temp { false }; uint16_t m_toggle_mask { 0 }; - uint16_t m_led_mask { 0 }; - BAN::UniqPtr m_led_region; - - BAN::Vector m_outputs; BAN::Optional m_repeat_scancode; uint8_t m_repeat_modifier { 0 }; diff --git a/kernel/kernel/USB/Device.cpp b/kernel/kernel/USB/Device.cpp index 0f219db6..f6770937 100644 --- a/kernel/kernel/USB/Device.cpp +++ b/kernel/kernel/USB/Device.cpp @@ -252,14 +252,19 @@ namespace Kernel dwarnln("Could not initialize USB interface {}", ret.error()); m_class_drivers.remove(i--); } + else if (m_class_drivers[i]->is_hid_driver() && static_cast(m_class_drivers[i].ptr())->has_led_control()) + TRY(m_hid_info.led_controls.push_back(m_class_drivers[i].ptr())); } if (!m_class_drivers.empty()) { + update_led_mask(m_hid_info.led_mask); + dprintln("Successfully initialized USB device with {}/{} interfaces", m_class_drivers.size(), configuration.interfaces.size() ); + return {}; } } @@ -363,4 +368,13 @@ namespace Kernel driver->handle_input_data(byte_count, endpoint_id); } + void USBDevice::update_led_mask(uint32_t led_mask) + { + const uint32_t old_mask = m_hid_info.led_mask.exchange(led_mask); + if (old_mask == led_mask) + return; + for (auto* led_control : m_hid_info.led_controls) + static_cast(led_control)->set_leds(led_mask); + } + } diff --git a/kernel/kernel/USB/HID/HIDDriver.cpp b/kernel/kernel/USB/HID/HIDDriver.cpp index a2b47413..833d9838 100644 --- a/kernel/kernel/USB/HID/HIDDriver.cpp +++ b/kernel/kernel/USB/HID/HIDDriver.cpp @@ -77,9 +77,9 @@ namespace Kernel USBHIDDriver::~USBHIDDriver() { - for (auto& device_input : m_device_inputs) - if (device_input.device) - DevFileSystem::get().remove_device(device_input.device); + for (auto& device_report : m_device_reports) + if (device_report.device) + DevFileSystem::get().remove_device(device_report.device); } BAN::ErrorOr USBHIDDriver::initialize() @@ -184,7 +184,10 @@ namespace Kernel return BAN::Error::from_errno(EFAULT); } - m_device_inputs = TRY(initializes_device_reports(collections)); + TRY(initializes_device_reports(collections)); + + if (has_led_control()) + m_led_region = TRY(DMARegion::create(PAGE_SIZE, PageTable::MemoryType::Normal)); for (const auto& endpoint : m_interface.endpoints) { @@ -196,7 +199,7 @@ namespace Kernel continue; TRY(m_device.configure_endpoint(desc)); - m_data_buffer = TRY(DMARegion::create(desc.wMaxPacketSize & 0x07FF)); + m_data_buffer = TRY(DMARegion::create(desc.wMaxPacketSize & 0x07FF, PageTable::MemoryType::Normal)); m_data_endpoint_id = (desc.bEndpointAddress & 0x0F) * 2 + !!(desc.bEndpointAddress & 0x80); @@ -209,7 +212,7 @@ namespace Kernel return BAN::Error::from_errno(EINVAL); } - for (auto& report : m_device_inputs) + for (auto& report : m_device_reports) if (report.device && report.device->initialize().is_error()) report.device.clear(); @@ -238,10 +241,9 @@ namespace Kernel return {}; } - BAN::ErrorOr> USBHIDDriver::initializes_device_reports(const BAN::Vector& collection_list) + BAN::ErrorOr USBHIDDriver::initializes_device_reports(const BAN::Vector& collection_list) { - BAN::Vector result; - TRY(result.reserve(collection_list.size())); + TRY(m_device_reports.reserve(collection_list.size())); for (size_t i = 0; i < collection_list.size(); i++) { @@ -249,9 +251,7 @@ namespace Kernel USBHIDDriver::DeviceReport report; TRY(gather_collection_reports(collection, report.inputs, USBHID::Report::Type::Input)); - - BAN::Vector outputs; - TRY(gather_collection_reports(collection, outputs, USBHID::Report::Type::Output)); + TRY(gather_collection_reports(collection, report.outputs, USBHID::Report::Type::Output)); switch (collection.usage_page) { @@ -267,7 +267,7 @@ namespace Kernel dprintln("Initialized an USB Joystick"); break; case 0x06: - report.device = TRY(BAN::RefPtr::create(*this, BAN::move(outputs))); + report.device = TRY(BAN::RefPtr::create(*this)); dprintln("Initialized an USB Keyboard"); break; default: @@ -279,19 +279,47 @@ namespace Kernel switch (collection.usage_id) { case 0x01: - report.device = TRY(BAN::RefPtr::create(*this, BAN::move(outputs))); + report.device = TRY(BAN::RefPtr::create(*this)); dprintln("Initialized an USB Consumer Control"); break; } } - TRY(result.push_back(BAN::move(report))); + TRY(m_device_reports.push_back(BAN::move(report))); } - for (auto& report : result) - if (report.device) - DevFileSystem::get().add_device(report.device); - return BAN::move(result); + for (auto& report : m_device_reports) + { + if (!report.device) + continue; + DevFileSystem::get().add_device(report.device); + + uint8_t led_report_ids[0x100 / 8] {}; + for (const auto& output : report.outputs) + { + if (output.usage_page != 0x08) + continue; + + const auto byte = output.report_id / 8; + const auto bit = output.report_id % 8; + if (led_report_ids[byte] & (1u << bit)) + continue; + led_report_ids[byte] |= (1u << bit); + + uint32_t report_bits = 0; + for (const auto& temp : report.outputs) + if (temp.report_id == output.report_id) + report_bits += temp.report_size * temp.report_count; + + TRY(m_led_controls.push_back({ + .report = &m_device_reports.back(), + .report_id = output.report_id, + .report_bits = report_bits, + })); + } + } + + return {}; } void USBHIDDriver::handle_stall(uint8_t endpoint_id) @@ -368,19 +396,19 @@ namespace Kernel } size_t bit_offset = 0; - for (auto& device_input : m_device_inputs) + for (auto& device_report : m_device_reports) { - if (device_input.device) - device_input.device->start_report(); + if (device_report.device) + device_report.device->start_report(); - for (const auto& input : device_input.inputs) + for (const auto& input : device_report.inputs) { if (report_id.value_or(input.report_id) != input.report_id) continue; ASSERT(input.report_size <= 32); - if (!device_input.device || (input.usage_id == 0 && input.usage_minimum == 0 && input.usage_maximum == 0)) + if (!device_report.device || (input.usage_id == 0 && input.usage_minimum == 0 && input.usage_maximum == 0)) { bit_offset += input.report_size * input.report_count; continue; @@ -401,7 +429,7 @@ namespace Kernel const auto usage = input.usage_id ? input.usage_id : (input.usage_minimum + (variable ? i : logical)); if (!variable) - device_input.device->handle_array(input.usage_page, usage); + device_report.device->handle_array(input.usage_page, usage); else { const int64_t physical = @@ -411,17 +439,51 @@ namespace Kernel input.physical_minimum; if (relative) - device_input.device->handle_variable(input.usage_page, usage, physical); + device_report.device->handle_variable(input.usage_page, usage, physical); else - device_input.device->handle_variable_absolute(input.usage_page, usage, physical, input.physical_minimum, input.physical_maximum); + device_report.device->handle_variable_absolute(input.usage_page, usage, physical, input.physical_minimum, input.physical_maximum); } bit_offset += input.report_size; } } - if (device_input.device) - device_input.device->stop_report(); + if (device_report.device) + device_report.device->stop_report(); + } + } + + void USBHIDDriver::set_leds(uint32_t led_mask) + { + for (const auto& led_control : m_led_controls) + { + const size_t report_bytes = BAN::Math::div_round_up(led_control.report_bits, 8); + + auto led_data = BAN::ByteSpan(reinterpret_cast(m_led_region->vaddr()), report_bytes); + memset(led_data.data(), 0, report_bytes); + + size_t bit_offset = 0; + for (const auto& output : led_control.report->outputs) + { + if (output.report_id != led_control.report_id) + continue; + + const size_t usage_base = output.usage_id ? output.usage_id : output.usage_minimum; + for (size_t i = 0; output.report_size == 1 && i < output.report_count; i++, bit_offset++) + if (led_mask & (1u << (usage_base + bit_offset))) + led_data[bit_offset / 8] |= 1u << (bit_offset % 8); + + bit_offset += output.report_size * output.report_count; + } + + USBDeviceRequest request; + request.bmRequestType = USB::RequestType::HostToDevice | USB::RequestType::Class | USB::RequestType::Interface; + request.bRequest = 0x09; + request.wValue = 0x0200 | led_control.report_id; + request.wIndex = m_interface.descriptor.bInterfaceNumber; + request.wLength = report_bytes; + if (auto ret = m_device.send_request(request, m_led_region->paddr()); ret.is_error()) + dprintln_if(DEBUG_USB_HID, "Failed to update LEDs: {}", ret.error()); } } diff --git a/kernel/kernel/USB/HID/Keyboard.cpp b/kernel/kernel/USB/HID/Keyboard.cpp index 5b955e63..f74ef47f 100644 --- a/kernel/kernel/USB/HID/Keyboard.cpp +++ b/kernel/kernel/USB/HID/Keyboard.cpp @@ -14,19 +14,10 @@ namespace Kernel static void initialize_scancode_to_keycode(); static constexpr bool is_repeatable_scancode(uint8_t scancode); - USBKeyboard::USBKeyboard(USBHIDDriver& driver, BAN::Vector&& outputs) + USBKeyboard::USBKeyboard(USBHIDDriver& driver) : USBHIDDevice(InputDevice::Type::Keyboard) , m_driver(driver) - , m_outputs(BAN::move(outputs)) - { - set_leds(0); - } - - BAN::ErrorOr USBKeyboard::initialize() - { - m_led_region = TRY(DMARegion::create(PAGE_SIZE, PageTable::MemoryType::Normal)); - return {}; - } + { } void USBKeyboard::start_report() { @@ -173,8 +164,19 @@ namespace Kernel const auto toggle_mask = ({ SpinLockGuard _(m_keyboard_lock); m_toggle_mask; }); if (m_led_mask != toggle_mask) - set_leds(toggle_mask); - m_led_mask = toggle_mask; + { + uint32_t usb_led_mask { 0 }; + if (toggle_mask & KeyModifier::NumLock) + usb_led_mask |= 1 << 1; + if (toggle_mask & KeyModifier::CapsLock) + usb_led_mask |= 1 << 2; + if (toggle_mask & KeyModifier::ScrollLock) + usb_led_mask |= 1 << 3; + + m_driver.device().update_led_mask(usb_led_mask); + + m_led_mask = toggle_mask; + } SpinLockGuard _(m_keyboard_lock); @@ -193,88 +195,6 @@ namespace Kernel m_next_repeat_event_ms += s_repeat_interval_ms; } - void USBKeyboard::set_leds(uint16_t mask) - { - uint8_t report_ids_done[0x100 / 8] {}; - - for (const auto& report : m_outputs) - { - if (report.usage_page != 0x08) - continue; - - const auto byte = report.report_id / 8; - const auto bit = report.report_id % 8; - if (report_ids_done[byte] & (1u << bit)) - continue; - - set_leds(report.report_id, mask); - report_ids_done[byte] |= (1u << bit); - } - } - - void USBKeyboard::set_leds(uint8_t report_id, uint16_t mask) - { - using KeyModifier = LibInput::KeyEvent::Modifier; - - if (!m_led_region) - return; - - size_t report_bits = 0; - for (const auto& report : m_outputs) - { - if (report.report_id != report_id) - continue; - report_bits += report.report_size * report.report_count; - } - - const size_t report_bytes = (report_bits + 7) / 8; - ASSERT(report_bytes <= PAGE_SIZE); - - PageTable::with_fast_page(m_led_region->paddr(), [&] { - uint8_t* data = &PageTable::fast_page_as(); - - memset(data, 0, report_bytes); - - size_t bit_offset = 0; - for (const auto& report : m_outputs) - { - if (report.report_id != report_id) - continue; - - for (size_t i = 0; report.report_size == 1 && i < report.report_count; i++, bit_offset++) - { - const size_t usage = (report.usage_id ? report.usage_id : report.usage_minimum) + bit_offset; - switch (usage) - { - case 0x01: - if (mask & KeyModifier::NumLock) - data[bit_offset / 8] |= 1u << (bit_offset % 8); - break; - case 0x02: - if (mask & KeyModifier::CapsLock) - data[bit_offset / 8] |= 1u << (bit_offset % 8); - break; - case 0x03: - if (mask & KeyModifier::ScrollLock) - data[bit_offset / 8] |= 1u << (bit_offset % 8); - break; - } - } - - bit_offset += report.report_size * report.report_count; - } - }); - - USBDeviceRequest request; - request.bmRequestType = USB::RequestType::HostToDevice | USB::RequestType::Class | USB::RequestType::Interface; - request.bRequest = 0x09; - request.wValue = 0x0200 | report_id; - request.wIndex = m_driver.interface().descriptor.bInterfaceNumber; - request.wLength = report_bytes; - if (auto ret = m_driver.device().send_request(request, m_led_region->paddr()); ret.is_error()) - dprintln_if(DEBUG_USB_KEYBOARD, "Failed to update LEDs: {}", ret.error()); - } - void initialize_scancode_to_keycode() { using LibInput::keycode_function;