switch-pico/switch2_driver.cpp
Joey Yakimowich-Payne 95b0576d1d Add Switch2 driver
Ultraworked with [Sisyphus](https://github.com/code-yeongyu/oh-my-openagent)

Co-authored-by: Sisyphus <clio-agent@sisyphuslabs.ai>
2026-08-11 12:23:07 +09:00

237 lines
6.9 KiB
C++

#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;
}