Add Switch2 driver
Ultraworked with [Sisyphus](https://github.com/code-yeongyu/oh-my-openagent) Co-authored-by: Sisyphus <clio-agent@sisyphuslabs.ai>
This commit is contained in:
parent
3251e49191
commit
95b0576d1d
1 changed files with 237 additions and 0 deletions
237
switch2_driver.cpp
Normal file
237
switch2_driver.cpp
Normal file
|
|
@ -0,0 +1,237 @@
|
|||
#include "switch_protocol.h"
|
||||
|
||||
#include <array>
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <cstring>
|
||||
|
||||
#include "pico/time.h"
|
||||
#include "switch2_commands.h"
|
||||
#include "switch2_descriptors.h"
|
||||
#include "switch2_reports.h"
|
||||
#include "tusb.h"
|
||||
|
||||
namespace {
|
||||
|
||||
constexpr uint32_t kReportIntervalMs = 4;
|
||||
constexpr std::size_t kMaxVendorRequestLength = 64;
|
||||
constexpr std::size_t kMaxVendorResponseLength = 12;
|
||||
|
||||
SwitchInputState input_state{};
|
||||
Switch2InputReportId selected_report = Switch2InputReportId::Pro;
|
||||
uint32_t report_counter = 0;
|
||||
uint32_t last_report_time = 0;
|
||||
bool mounted = false;
|
||||
bool initialized = false;
|
||||
std::array<uint8_t, kMaxVendorResponseLength> pending_response{};
|
||||
std::size_t pending_response_length = 0;
|
||||
SwitchRumbleCallback rumble_callback = nullptr;
|
||||
|
||||
void reset_connection_state(bool is_mounted) {
|
||||
input_state = {};
|
||||
selected_report = Switch2InputReportId::Pro;
|
||||
report_counter = 0;
|
||||
last_report_time = 0;
|
||||
mounted = is_mounted;
|
||||
initialized = false;
|
||||
pending_response_length = 0;
|
||||
}
|
||||
|
||||
void flush_vendor_response() {
|
||||
if (pending_response_length == 0) return;
|
||||
|
||||
const uint32_t written = tud_vendor_write(
|
||||
pending_response.data(), static_cast<uint32_t>(pending_response_length));
|
||||
if (written == 0) return;
|
||||
|
||||
const std::size_t consumed = written;
|
||||
pending_response_length -= consumed;
|
||||
if (pending_response_length > 0) {
|
||||
std::memmove(
|
||||
pending_response.data(),
|
||||
pending_response.data() + consumed,
|
||||
pending_response_length);
|
||||
}
|
||||
tud_vendor_write_flush();
|
||||
}
|
||||
|
||||
} // namespace
|
||||
|
||||
void switch_protocol_init() {
|
||||
rumble_callback = nullptr;
|
||||
reset_connection_state(false);
|
||||
}
|
||||
|
||||
void switch_protocol_set_input(const SwitchInputState& state) {
|
||||
input_state = state;
|
||||
}
|
||||
|
||||
void switch_protocol_task() {
|
||||
if (!mounted) return;
|
||||
|
||||
const bool had_pending_response = pending_response_length > 0;
|
||||
flush_vendor_response();
|
||||
if (had_pending_response || !initialized || pending_response_length > 0) return;
|
||||
|
||||
const uint32_t now = static_cast<uint32_t>(
|
||||
to_ms_since_boot(get_absolute_time()));
|
||||
if ((now - last_report_time) < kReportIntervalMs || !tud_hid_ready()) return;
|
||||
|
||||
const Switch2InputReport report =
|
||||
switch2_build_input_report(selected_report, input_state, report_counter);
|
||||
if (tud_hid_report(
|
||||
static_cast<uint8_t>(report.id),
|
||||
report.payload.data(),
|
||||
static_cast<uint16_t>(report.payload.size()))) {
|
||||
++report_counter;
|
||||
last_report_time = now;
|
||||
}
|
||||
}
|
||||
|
||||
bool switch_protocol_is_ready() {
|
||||
return mounted && initialized;
|
||||
}
|
||||
|
||||
void switch_protocol_set_rumble_callback(SwitchRumbleCallback callback) {
|
||||
rumble_callback = callback;
|
||||
}
|
||||
|
||||
uint8_t const* tud_descriptor_device_cb() {
|
||||
return switch2_device_descriptor;
|
||||
}
|
||||
|
||||
uint8_t const* tud_descriptor_configuration_cb(uint8_t index) {
|
||||
return index == 0 ? switch2_configuration_descriptor : nullptr;
|
||||
}
|
||||
|
||||
uint8_t const* tud_hid_descriptor_report_cb(uint8_t instance) {
|
||||
return instance == 0 ? switch2_hid_report_descriptor : nullptr;
|
||||
}
|
||||
|
||||
uint16_t tud_hid_get_report_cb(
|
||||
uint8_t instance,
|
||||
uint8_t report_id,
|
||||
hid_report_type_t report_type,
|
||||
uint8_t* buffer,
|
||||
uint16_t reqlen) {
|
||||
if (instance != 0 || report_type != HID_REPORT_TYPE_INPUT || buffer == nullptr ||
|
||||
!mounted || !initialized) {
|
||||
return 0;
|
||||
}
|
||||
|
||||
Switch2InputReportId id;
|
||||
if (report_id == static_cast<uint8_t>(Switch2InputReportId::Common)) {
|
||||
id = Switch2InputReportId::Common;
|
||||
} else if (report_id == static_cast<uint8_t>(Switch2InputReportId::Pro)) {
|
||||
id = Switch2InputReportId::Pro;
|
||||
} else {
|
||||
return 0;
|
||||
}
|
||||
|
||||
const Switch2InputReport report =
|
||||
switch2_build_input_report(id, input_state, report_counter);
|
||||
const uint16_t length = reqlen < report.payload.size()
|
||||
? reqlen
|
||||
: static_cast<uint16_t>(report.payload.size());
|
||||
std::memcpy(buffer, report.payload.data(), length);
|
||||
return length;
|
||||
}
|
||||
|
||||
void tud_hid_set_report_cb(
|
||||
uint8_t instance,
|
||||
uint8_t report_id,
|
||||
hid_report_type_t report_type,
|
||||
uint8_t const* buffer,
|
||||
uint16_t bufsize) {
|
||||
if (instance != 0 || report_type != HID_REPORT_TYPE_OUTPUT || buffer == nullptr) return;
|
||||
|
||||
uint8_t const* payload = nullptr;
|
||||
if (report_id == 0 && bufsize == 64 && buffer[0] == 0x02) {
|
||||
payload = buffer + 1;
|
||||
} else if (report_id == 0x02 && bufsize == 63) {
|
||||
payload = buffer;
|
||||
} else {
|
||||
return;
|
||||
}
|
||||
|
||||
// Report 0x02 carries two native HD-rumble payloads. No legacy UART mapping exists.
|
||||
(void)payload;
|
||||
}
|
||||
|
||||
void tud_vendor_rx_cb(uint8_t itf, uint8_t const* buffer, uint16_t bufsize) {
|
||||
if (itf != 0) return;
|
||||
|
||||
std::array<uint8_t, kMaxVendorRequestLength> request{};
|
||||
const std::size_t request_length = bufsize < request.size() ? bufsize : request.size();
|
||||
if (buffer != nullptr) {
|
||||
std::memcpy(request.data(), buffer, request_length);
|
||||
}
|
||||
tud_vendor_read_flush();
|
||||
if (!mounted || buffer == nullptr) return;
|
||||
|
||||
const Switch2VendorCommand command =
|
||||
switch2_classify_vendor_request(request.data(), request_length);
|
||||
|
||||
switch (command) {
|
||||
case Switch2VendorCommand::Unsupported:
|
||||
return;
|
||||
case Switch2VendorCommand::SelectReport05:
|
||||
selected_report = Switch2InputReportId::Common;
|
||||
break;
|
||||
case Switch2VendorCommand::SelectReport09:
|
||||
selected_report = Switch2InputReportId::Pro;
|
||||
break;
|
||||
case Switch2VendorCommand::InitializeUsb:
|
||||
initialized = true;
|
||||
break;
|
||||
}
|
||||
|
||||
pending_response_length = switch2_build_vendor_response(
|
||||
command, pending_response.data(), pending_response.size());
|
||||
}
|
||||
|
||||
void tud_mount_cb() {
|
||||
reset_connection_state(true);
|
||||
}
|
||||
|
||||
void tud_umount_cb() {
|
||||
reset_connection_state(false);
|
||||
}
|
||||
|
||||
uint16_t const* tud_descriptor_string_cb(uint8_t index, uint16_t langid) {
|
||||
(void)langid;
|
||||
static uint16_t descriptor[32];
|
||||
|
||||
if (index == 0) {
|
||||
descriptor[1] = static_cast<uint16_t>(switch2_string_language[0]) |
|
||||
(static_cast<uint16_t>(switch2_string_language[1]) << 8);
|
||||
descriptor[0] = 0x0304;
|
||||
return descriptor;
|
||||
}
|
||||
|
||||
const uint8_t* string = nullptr;
|
||||
std::size_t length = 0;
|
||||
switch (index) {
|
||||
case 1:
|
||||
string = switch2_string_manufacturer;
|
||||
length = switch2_string_manufacturer_length;
|
||||
break;
|
||||
case 2:
|
||||
string = switch2_string_product;
|
||||
length = switch2_string_product_length;
|
||||
break;
|
||||
case 3:
|
||||
string = switch2_string_serial;
|
||||
length = switch2_string_serial_length;
|
||||
break;
|
||||
default:
|
||||
return nullptr;
|
||||
}
|
||||
|
||||
if (length > 31) length = 31;
|
||||
for (std::size_t position = 0; position < length; ++position) {
|
||||
descriptor[1 + position] = string[position];
|
||||
}
|
||||
descriptor[0] = static_cast<uint16_t>(0x0300 | (2 * length + 2));
|
||||
return descriptor;
|
||||
}
|
||||
Loading…
Add table
Add a link
Reference in a new issue