Add versioned persistent configuration protocol

This commit is contained in:
Joey Yakimowich-Payne 2026-09-02 15:24:35 -06:00
commit 7afc9981fe
36 changed files with 2784 additions and 756 deletions

View file

@ -285,6 +285,15 @@ uint32_t btstack_run_loop_get_time_ms() {
#include "../bluepad32_input_backend.cpp"
void configuration_service_prepare() {}
void configuration_service_initialize_on_storage_core() {}
void configuration_service_task_on_storage_core(uint32_t) {}
void configuration_service_snapshot(ConfigurationServiceSnapshot* output) {
*output = {};
output->state = ConfigurationServiceState::kReady;
output->configuration.pairing_window_seconds =
ADAPTER_PAIRING_WINDOW_SECONDS_DEFAULT;
}
#ifdef SWITCH_PICO_ADAPTER_FEASIBILITY
AdapterUsbMode test_adapter_mode = AdapterUsbMode::kXInput;
AdapterUsbMode adapter_host_probe_mode() {

View file

@ -0,0 +1,232 @@
#include "adapter_configuration.h"
#include "configuration_storage.h"
#include "configuration_transaction.h"
#include <cstdlib>
#include <cstring>
#include <iostream>
namespace {
constexpr size_t kSectorSize = 4096;
constexpr size_t kPageSize = 256;
struct FakeFlash {
uint8_t bytes[CONFIGURATION_STORAGE_COPY_COUNT][kSectorSize];
int successful_programs = 0;
int fail_after_programs = -1;
FakeFlash() { memset(bytes, 0xff, sizeof(bytes)); }
};
void require(bool condition, const char* message) {
if (!condition) {
std::cerr << message << '\n';
std::exit(1);
}
}
bool fake_read(void* context, uint8_t copy, size_t offset,
uint8_t* output, size_t size) {
auto* flash = static_cast<FakeFlash*>(context);
if (copy >= CONFIGURATION_STORAGE_COPY_COUNT ||
offset + size > kSectorSize) {
return false;
}
memcpy(output, &flash->bytes[copy][offset], size);
return true;
}
bool fake_erase(void* context, uint8_t copy) {
auto* flash = static_cast<FakeFlash*>(context);
if (copy >= CONFIGURATION_STORAGE_COPY_COUNT) {
return false;
}
memset(flash->bytes[copy], 0xff, kSectorSize);
return true;
}
bool fake_program(void* context, uint8_t copy, size_t offset,
const uint8_t* data, size_t size) {
auto* flash = static_cast<FakeFlash*>(context);
if (copy >= CONFIGURATION_STORAGE_COPY_COUNT || size != kPageSize ||
offset + size > kSectorSize) {
return false;
}
if (flash->fail_after_programs >= 0 &&
flash->successful_programs >= flash->fail_after_programs) {
return false;
}
for (size_t index = 0; index < size; ++index) {
flash->bytes[copy][offset + index] &= data[index];
}
++flash->successful_programs;
return true;
}
ConfigurationStorageIo fake_io(FakeFlash* flash) {
return {
flash,
kSectorSize,
kPageSize,
fake_read,
fake_erase,
fake_program,
};
}
void test_schema_encoding() {
AdapterConfiguration configuration{};
configuration.pairing_window_seconds = 90;
uint8_t payload[ADAPTER_CONFIGURATION_ENCODED_SIZE]{};
require(adapter_configuration_encode(configuration, payload,
sizeof(payload)),
"valid configuration did not encode");
AdapterConfiguration decoded{};
require(adapter_configuration_decode(payload, sizeof(payload),
&decoded) &&
decoded.pairing_window_seconds == 90,
"configuration did not round trip");
payload[2] = 1;
require(!adapter_configuration_decode(payload, sizeof(payload),
&decoded),
"nonzero reserved configuration byte was accepted");
configuration.pairing_window_seconds = 9;
require(!adapter_configuration_encode(configuration, payload,
sizeof(payload)),
"out-of-range pairing window was accepted");
}
void test_two_copy_recovery() {
FakeFlash flash;
ConfigurationStorage store;
require(store.initialize(fake_io(&flash)),
"storage did not initialize");
require(!store.snapshot().valid,
"erased storage appeared valid");
const uint8_t first[] = {1, 2, 3, 4};
require(store.commit(1, first, sizeof(first)) ==
ConfigurationStorageResult::kOk &&
store.snapshot().generation == 1,
"first generation did not commit");
const int programs_after_first = flash.successful_programs;
require(store.commit(1, first, sizeof(first)) ==
ConfigurationStorageResult::kUnchanged &&
flash.successful_programs == programs_after_first,
"identical configuration consumed a flash write");
const uint8_t second[] = {5, 6, 7, 8};
require(store.commit(1, second, sizeof(second)) ==
ConfigurationStorageResult::kOk &&
store.snapshot().generation == 2,
"second generation did not commit");
ConfigurationStorage after_power_cycle;
require(after_power_cycle.initialize(fake_io(&flash)) &&
after_power_cycle.snapshot().generation == 2 &&
memcmp(after_power_cycle.snapshot().payload, second,
sizeof(second)) == 0,
"latest generation did not survive reinitialization");
flash.bytes[1][CONFIGURATION_STORAGE_RECORD_HEADER_SIZE] ^= 0x01;
ConfigurationStorage after_corruption;
require(after_corruption.initialize(fake_io(&flash)) &&
after_corruption.snapshot().generation == 1 &&
memcmp(after_corruption.snapshot().payload, first,
sizeof(first)) == 0,
"corrupt newest generation did not roll back");
}
void test_interrupted_write_retains_previous_generation() {
FakeFlash flash;
ConfigurationStorage store;
require(store.initialize(fake_io(&flash)),
"storage did not initialize for interruption test");
const uint8_t first[] = {9, 8, 7, 6};
require(store.commit(1, first, sizeof(first)) ==
ConfigurationStorageResult::kOk,
"baseline generation did not commit");
uint8_t large[300];
memset(large, 0x5a, sizeof(large));
flash.fail_after_programs = flash.successful_programs + 1;
require(store.commit(1, large, sizeof(large)) ==
ConfigurationStorageResult::kIoError,
"interrupted multi-page write reported success");
flash.fail_after_programs = -1;
ConfigurationStorage recovered;
require(recovered.initialize(fake_io(&flash)) &&
recovered.snapshot().generation == 1 &&
recovered.snapshot().payload_size == sizeof(first) &&
memcmp(recovered.snapshot().payload, first,
sizeof(first)) == 0,
"interrupted write replaced the previous generation");
}
void test_transaction_validation() {
AdapterConfiguration configuration{};
configuration.pairing_window_seconds = 120;
uint8_t payload[ADAPTER_CONFIGURATION_ENCODED_SIZE]{};
require(adapter_configuration_encode(configuration, payload,
sizeof(payload)),
"transaction payload did not encode");
const uint32_t crc = configuration_crc32(payload, sizeof(payload));
ConfigurationTransaction transaction;
require(transaction.begin(7, ADAPTER_CONFIGURATION_SCHEMA_VERSION,
sizeof(payload), crc) ==
ConfigurationTransactionStatus::kReceiving,
"valid transaction did not begin");
require(transaction.append(7, 1, payload, 1) ==
ConfigurationTransactionStatus::kOutOfOrder,
"out-of-order chunk was accepted");
require(transaction.begin(8, ADAPTER_CONFIGURATION_SCHEMA_VERSION,
sizeof(payload), crc ^ 1u) ==
ConfigurationTransactionStatus::kReceiving,
"replacement transaction did not begin");
require(transaction.append(8, 0, payload, 2) ==
ConfigurationTransactionStatus::kReceiving &&
transaction.append(8, 2, &payload[2], 2) ==
ConfigurationTransactionStatus::kReceiving,
"ordered chunks were rejected");
require(transaction.finish(8) ==
ConfigurationTransactionStatus::kBadCrc,
"bad transaction CRC was accepted");
require(transaction.begin(9, ADAPTER_CONFIGURATION_SCHEMA_VERSION,
sizeof(payload), crc) ==
ConfigurationTransactionStatus::kReceiving &&
transaction.append(9, 0, payload, sizeof(payload)) ==
ConfigurationTransactionStatus::kReceiving &&
transaction.finish(9) ==
ConfigurationTransactionStatus::kPending,
"valid transaction did not reach pending commit");
require(transaction.begin(10, ADAPTER_CONFIGURATION_SCHEMA_VERSION,
sizeof(payload), crc) ==
ConfigurationTransactionStatus::kBusy,
"pending transaction was replaced");
transaction.set_result(ConfigurationTransactionStatus::kCommitted,
4, crc);
require(transaction.begin(11, 99, sizeof(payload), crc) ==
ConfigurationTransactionStatus::kUnsupportedSchema,
"unsupported schema was accepted");
require(transaction.begin(
12, ADAPTER_CONFIGURATION_SCHEMA_VERSION,
CONFIGURATION_STORAGE_MAX_PAYLOAD_SIZE + 1, crc) ==
ConfigurationTransactionStatus::kTooLarge,
"oversized transaction was accepted");
}
} // namespace
int main() {
test_schema_encoding();
test_two_copy_recovery();
test_interrupted_write_retains_previous_generation();
test_transaction_validation();
return 0;
}

View file

@ -0,0 +1,244 @@
from __future__ import annotations
import struct
import zlib
import pytest
import switch_pico_bridge.config_manager as config_manager
def make_response(
operation: int,
payload: bytes = b"",
*,
status: int = config_manager.STATUS_OK,
flags: int = 0,
schema: int = 0,
generation: int = 0,
) -> bytes:
return struct.pack(
"<4sBBBBHHII",
b"SPMG",
config_manager.PROTOCOL_VERSION,
operation,
status,
flags,
len(payload),
schema,
generation,
zlib.crc32(payload) & 0xFFFFFFFF,
) + payload
class FakeDevice:
bus = 1
address = 7
def __init__(self) -> None:
self.configuration = struct.pack("<Hxx", 60)
self.configuration_generation = 3
self.transaction_id = 0
self.transaction_payload = bytearray()
self.transaction_expected_size = 0
self.transaction_expected_crc = 0
self.transaction_status = config_manager.STATUS_OK
self.records = [
(
config_manager.TRANSPORT_CLASSIC,
0xFE,
bytes.fromhex("010203040506"),
),
(
config_manager.TRANSPORT_BLE,
2,
bytes.fromhex("A1A2A3A4A5A6"),
),
]
self.pairing_generation = 4
self.requests: list[int] = []
def _pairing_payload(self) -> bytes:
payload = bytearray([len(self.records), 0, 0, 0])
for transport, address_type, address in self.records:
payload.extend([transport, address_type])
payload.extend(address)
return bytes(payload)
def _transaction_payload(self) -> bytes:
stored_crc = zlib.crc32(self.configuration) & 0xFFFFFFFF
return struct.pack(
"<IHHIII",
self.transaction_id,
len(self.transaction_payload),
self.transaction_expected_size,
self.transaction_expected_crc,
self.configuration_generation,
stored_crc,
)
def ctrl_transfer(
self,
bm_request_type: int,
request: int,
value: int,
index: int,
data_or_w_length: object,
timeout: int,
) -> bytes | int:
assert value == config_manager.REQUEST_VALUE
assert index == config_manager.REQUEST_INDEX
assert timeout == config_manager.USB_TIMEOUT_MS
self.requests.append(request)
if bm_request_type == 0xC0:
if request == config_manager.OP_INFO:
return make_response(request, bytes([0, 2, 0, 2, 0, 0, 0, 2]))
if request == config_manager.OP_CONFIGURATION_READ:
return make_response(
request,
self.configuration,
schema=config_manager.CONFIGURATION_SCHEMA_VERSION,
generation=self.configuration_generation,
)
if request == config_manager.OP_TRANSACTION_STATUS:
return make_response(
request,
self._transaction_payload(),
status=self.transaction_status,
schema=config_manager.CONFIGURATION_SCHEMA_VERSION,
generation=self.configuration_generation,
)
if request == config_manager.OP_PAIRING_READ:
return make_response(
request,
self._pairing_payload(),
generation=self.pairing_generation,
)
raise AssertionError(f"unexpected IN request {request}")
assert bm_request_type == 0x40
encoded = bytes(data_or_w_length)
assert encoded[:4] == b"SPMG"
payload_size = struct.unpack_from("<H", encoded, 8)[0]
payload = encoded[config_manager.REQUEST_HEADER_SIZE :]
assert payload_size == len(payload)
assert struct.unpack_from("<I", encoded, 12)[0] == (
zlib.crc32(payload) & 0xFFFFFFFF
)
if request == config_manager.OP_CONFIGURATION_BEGIN:
(
self.transaction_id,
_schema,
self.transaction_expected_size,
self.transaction_expected_crc,
) = struct.unpack("<IHHI", payload)
self.transaction_payload = bytearray()
self.transaction_status = config_manager.STATUS_PENDING
elif request == config_manager.OP_CONFIGURATION_CHUNK:
transaction_id, offset, chunk_size = struct.unpack_from(
"<IHH", payload
)
assert transaction_id == self.transaction_id
assert offset == len(self.transaction_payload)
self.transaction_payload.extend(payload[8 : 8 + chunk_size])
elif request == config_manager.OP_CONFIGURATION_COMMIT:
assert struct.unpack("<I", payload)[0] == self.transaction_id
assert len(self.transaction_payload) == self.transaction_expected_size
assert (
zlib.crc32(self.transaction_payload) & 0xFFFFFFFF
) == self.transaction_expected_crc
self.configuration = bytes(self.transaction_payload)
self.configuration_generation += 1
self.transaction_status = config_manager.STATUS_OK
elif request == config_manager.OP_CONFIGURATION_RESET:
self.transaction_id = struct.unpack("<I", payload)[0]
self.configuration = struct.pack("<Hxx", 60)
self.configuration_generation += 1
self.transaction_payload = bytearray(self.configuration)
self.transaction_expected_size = len(self.configuration)
self.transaction_expected_crc = (
zlib.crc32(self.configuration) & 0xFFFFFFFF
)
self.transaction_status = config_manager.STATUS_OK
elif request == config_manager.OP_PAIRING_REFRESH:
self.pairing_generation += 1
elif request == config_manager.OP_PAIRING_CLEAR:
self.records = []
self.pairing_generation += 1
else:
raise AssertionError(f"unexpected OUT request {request}")
return len(encoded)
def test_response_validation() -> None:
payload = make_response(config_manager.OP_INFO, b"12345678")
envelope = config_manager.parse_response(payload, config_manager.OP_INFO)
assert envelope.payload == b"12345678"
malformed = [
b"",
b"NOPE" + bytes(config_manager.RESPONSE_HEADER_SIZE - 4),
make_response(config_manager.OP_INFO, b"12345678")[:-1],
make_response(config_manager.OP_INFO, b"12345678") + b"x",
]
bad_crc = bytearray(make_response(config_manager.OP_INFO, b"12345678"))
bad_crc[-1] ^= 1
malformed.append(bytes(bad_crc))
for response in malformed:
with pytest.raises(config_manager.ConfigManagerError):
config_manager.parse_response(response, config_manager.OP_INFO)
def test_configuration_transaction_and_reset() -> None:
device = FakeDevice()
before = config_manager.read_configuration(device)
assert before.pairing_window_seconds == 60
status = config_manager.write_configuration(
device,
config_manager.AdapterConfiguration(90, before.generation, before.crc),
1.0,
)
assert status.stored_generation == 4
assert config_manager.read_configuration(device).pairing_window_seconds == 90
reset = config_manager.reset_configuration(device, 1.0)
assert reset.stored_generation == 5
assert config_manager.read_configuration(device).pairing_window_seconds == 60
def test_status_and_pairing_commands(
monkeypatch: pytest.MonkeyPatch,
capsys: pytest.CaptureFixture[str],
) -> None:
device = FakeDevice()
monkeypatch.setattr(config_manager, "_candidate_devices", lambda: [device])
assert config_manager.main(["status"]) == 0
output = capsys.readouterr().out
assert "Firmware: 0.2.0" in output
assert "Pairing window: 60 seconds" in output
assert config_manager.main(["pairings", "list"]) == 0
output = capsys.readouterr().out
assert "Classic 01:02:03:04:05:06" in output
assert "BLE (public identity) A1:A2:A3:A4:A5:A6" in output
assert config_manager.main(["pairings", "clear"]) == 2
assert "requires --yes" in capsys.readouterr().err
assert config_manager.main(["pairings", "clear", "--yes"]) == 0
assert capsys.readouterr().out == "Cleared 2 stored pairing(s).\n"
def test_find_requires_selector_for_multiple_picos(
monkeypatch: pytest.MonkeyPatch,
) -> None:
first = FakeDevice()
second = FakeDevice()
second.address = 8
monkeypatch.setattr(
config_manager, "_candidate_devices", lambda: [first, second]
)
with pytest.raises(
config_manager.ConfigManagerError,
match="multiple switch-pico devices",
):
config_manager.find_pico(None, None)
assert config_manager.find_pico(1, 8) is second

View file

@ -0,0 +1,31 @@
import shutil
import subprocess
from pathlib import Path
def test_configuration_storage_native(tmp_path: Path) -> None:
root = Path(__file__).resolve().parents[1]
compiler = shutil.which("c++") or shutil.which("g++")
assert compiler is not None, "a host C++ compiler is required"
executable = tmp_path / "configuration_storage_test"
subprocess.run(
[
compiler,
"-std=c++17",
"-Wall",
"-Wextra",
"-Werror",
"-pedantic",
f"-I{root}",
str(root / "tests" / "configuration_storage_test.cpp"),
str(root / "adapter_configuration.cpp"),
str(root / "configuration_storage.cpp"),
str(root / "configuration_transaction.cpp"),
"-o",
str(executable),
],
check=True,
cwd=root,
)
subprocess.run([str(executable)], check=True, cwd=root)

View file

@ -1,156 +0,0 @@
from __future__ import annotations
import struct
import pytest
import switch_pico_bridge.pairing_manager as pairing_manager
def make_payload(
generation: int,
records: list[tuple[int, int, bytes]],
*,
status: int = pairing_manager.STATUS_READY,
overflow: bool = False,
) -> bytes:
payload = bytearray(b"SPPM")
payload.extend(
[
pairing_manager.PROTOCOL_VERSION,
status,
len(records),
int(overflow),
]
)
payload.extend(struct.pack("<I", generation))
for transport, address_type, address in records:
payload.extend([transport, address_type])
payload.extend(address)
return bytes(payload)
class FakeDevice:
bus = 1
address = 7
def __init__(self) -> None:
self.generation = 3
self.records = [
(
pairing_manager.TRANSPORT_CLASSIC,
0xFE,
bytes.fromhex("010203040506"),
),
(
pairing_manager.TRANSPORT_BLE,
2,
bytes.fromhex("A1A2A3A4A5A6"),
),
]
self.requests: list[int] = []
def ctrl_transfer(
self,
bm_request_type: int,
request: int,
value: int,
index: int,
data_or_w_length: object,
timeout: int,
) -> bytes | int:
assert value == pairing_manager.REQUEST_VALUE
assert index == pairing_manager.REQUEST_INDEX
assert timeout == pairing_manager.USB_TIMEOUT_MS
self.requests.append(request)
if bm_request_type == 0xC0:
assert request == pairing_manager.REQUEST_GET
return make_payload(self.generation, self.records)
assert bm_request_type == 0x40
if request == pairing_manager.REQUEST_REFRESH:
self.generation += 1
elif request == pairing_manager.REQUEST_CLEAR:
self.records = []
self.generation += 1
else:
raise AssertionError(f"unexpected request {request}")
return 0
def test_parse_snapshot() -> None:
snapshot = pairing_manager.parse_snapshot(
make_payload(
0x78563412,
[
(
pairing_manager.TRANSPORT_CLASSIC,
0xFE,
bytes.fromhex("010203040506"),
),
(
pairing_manager.TRANSPORT_BLE,
3,
bytes.fromhex("A1A2A3A4A5A6"),
),
],
overflow=True,
)
)
assert snapshot.generation == 0x78563412
assert snapshot.overflow
assert snapshot.records[0].transport_text == "Classic"
assert snapshot.records[0].address_text == "01:02:03:04:05:06"
assert snapshot.records[1].transport_text == "BLE (random identity)"
@pytest.mark.parametrize(
"payload",
[
b"",
b"NOPE" + bytes(8),
b"SPPM\x02" + bytes(7),
b"SPPM\x01\x00\x11\x00" + bytes(4),
],
)
def test_parse_rejects_invalid_payload(payload: bytes) -> None:
with pytest.raises(pairing_manager.PairingManagerError):
pairing_manager.parse_snapshot(payload)
def test_list_and_clear_commands(
monkeypatch: pytest.MonkeyPatch,
capsys: pytest.CaptureFixture[str],
) -> None:
device = FakeDevice()
monkeypatch.setattr(pairing_manager, "_candidate_devices", lambda: [device])
assert pairing_manager.main(["list"]) == 0
output = capsys.readouterr().out
assert "Classic 01:02:03:04:05:06" in output
assert "BLE (public identity) A1:A2:A3:A4:A5:A6" in output
assert pairing_manager.main(["clear"]) == 2
assert "requires --yes" in capsys.readouterr().err
assert pairing_manager.main(["clear", "--yes"]) == 0
assert capsys.readouterr().out == "Cleared 2 stored pairing(s).\n"
assert device.records == []
assert pairing_manager.REQUEST_REFRESH in device.requests
assert pairing_manager.REQUEST_CLEAR in device.requests
def test_find_requires_selector_for_multiple_picos(
monkeypatch: pytest.MonkeyPatch,
) -> None:
first = FakeDevice()
second = FakeDevice()
second.address = 8
monkeypatch.setattr(
pairing_manager, "_candidate_devices", lambda: [first, second]
)
with pytest.raises(
pairing_manager.PairingManagerError,
match="multiple switch-pico devices",
):
pairing_manager.find_pico(None, None)
assert pairing_manager.find_pico(1, 8) is second

View file

@ -4,12 +4,12 @@ from pathlib import Path
def test_usb_pairing_management_native(tmp_path: Path) -> None:
def test_usb_configuration_management_native(tmp_path: Path) -> None:
root = Path(__file__).resolve().parents[1]
compiler = shutil.which("c++") or shutil.which("g++")
assert compiler is not None, "a host C++ compiler is required"
executable = tmp_path / "usb_pairing_management_test"
executable = tmp_path / "usb_configuration_management_test"
subprocess.run(
[
compiler,
@ -20,7 +20,7 @@ def test_usb_pairing_management_native(tmp_path: Path) -> None:
"-pedantic",
f"-I{root / 'tests' / 'usb_management_native_stubs'}",
f"-I{root}",
str(root / "tests" / "usb_pairing_management_test.cpp"),
str(root / "tests" / "usb_configuration_management_test.cpp"),
"-o",
str(executable),
],

View file

@ -0,0 +1,301 @@
#include "usb_configuration_management.h"
#include <cstdlib>
#include <cstring>
#include <iostream>
#include <vector>
#include <tusb.h>
namespace {
Bluepad32PairingSnapshot current_pairings{};
ConfigurationServiceSnapshot current_configuration{};
bool refresh_requested = false;
bool clear_requested = false;
std::vector<uint8_t> control_payload;
std::vector<uint8_t> next_out_payload;
uint32_t begin_transaction_id = 0;
uint32_t append_transaction_id = 0;
uint32_t commit_transaction_id = 0;
size_t append_offset = 0;
std::vector<uint8_t> appended_bytes;
void require(bool condition, const char* message) {
if (!condition) {
std::cerr << message << '\n';
std::exit(1);
}
}
void write_u16(std::vector<uint8_t>* output, size_t offset,
uint16_t value) {
(*output)[offset] = static_cast<uint8_t>(value);
(*output)[offset + 1] = static_cast<uint8_t>(value >> 8);
}
void write_u32(std::vector<uint8_t>* output, size_t offset,
uint32_t value) {
(*output)[offset] = static_cast<uint8_t>(value);
(*output)[offset + 1] = static_cast<uint8_t>(value >> 8);
(*output)[offset + 2] = static_cast<uint8_t>(value >> 16);
(*output)[offset + 3] = static_cast<uint8_t>(value >> 24);
}
std::vector<uint8_t> make_request(
UsbConfigurationManagement::Operation operation,
const std::vector<uint8_t>& payload) {
using namespace UsbConfigurationManagement;
std::vector<uint8_t> request(kRequestHeaderSize + payload.size());
memcpy(request.data(), "SPMG", 4);
request[4] = kProtocolVersion;
request[5] = static_cast<uint8_t>(operation);
write_u16(&request, 8, static_cast<uint16_t>(payload.size()));
write_u32(&request, 12,
configuration_crc32(payload.data(), payload.size()));
memcpy(request.data() + kRequestHeaderSize, payload.data(),
payload.size());
return request;
}
tusb_control_request_t setup_request(
UsbConfigurationManagement::Operation operation, uint8_t direction,
uint16_t length) {
tusb_control_request_t request{};
request.bmRequestType_bit.recipient = TUSB_REQ_RCPT_DEVICE;
request.bmRequestType_bit.type = TUSB_REQ_TYPE_VENDOR;
request.bmRequestType_bit.direction = direction;
request.bRequest = static_cast<uint8_t>(operation);
request.wValue = UsbConfigurationManagement::kRequestValue;
request.wIndex = UsbConfigurationManagement::kRequestIndex;
request.wLength = length;
return request;
}
void test_envelope_encoding() {
using namespace UsbConfigurationManagement;
const uint8_t payload[] = {1, 2, 3};
uint8_t encoded[32]{};
const size_t size = encode_response(
Operation::kConfigurationRead, Status::kOk, 5, 1,
0x78563412, payload, sizeof(payload), encoded, sizeof(encoded));
require(size == kResponseHeaderSize + sizeof(payload) &&
memcmp(encoded, "SPMG", 4) == 0 &&
encoded[4] == kProtocolVersion &&
encoded[5] ==
static_cast<uint8_t>(Operation::kConfigurationRead) &&
encoded[6] == static_cast<uint8_t>(Status::kOk) &&
encoded[7] == 5 && encoded[8] == 3 &&
encoded[10] == 1 && encoded[12] == 0x12 &&
encoded[15] == 0x78 &&
memcmp(&encoded[kResponseHeaderSize], payload,
sizeof(payload)) == 0,
"versioned response envelope encoded incorrectly");
require(encode_response(
Operation::kConfigurationRead, Status::kOk, 0, 1, 0,
payload, sizeof(payload), encoded, size - 1) == 0,
"response encoder accepted a short destination");
}
void test_pairing_encoding() {
using namespace UsbConfigurationManagement;
Bluepad32PairingSnapshot snapshot{};
snapshot.generation = 0x78563412;
snapshot.status = Bluepad32PairingSnapshotStatus::kReady;
snapshot.record_count = 2;
snapshot.overflow = true;
snapshot.records[0].transport =
Bluepad32PairingTransport::kClassic;
snapshot.records[0].address_type = 0xfe;
const uint8_t classic_address[6] = {1, 2, 3, 4, 5, 6};
memcpy(snapshot.records[0].address, classic_address, 6);
snapshot.records[1].transport = Bluepad32PairingTransport::kBle;
snapshot.records[1].address_type = 2;
const uint8_t ble_address[6] = {6, 5, 4, 3, 2, 1};
memcpy(snapshot.records[1].address, ble_address, 6);
uint8_t payload[kMaximumResponseSize]{};
const size_t size =
encode_pairing_snapshot(snapshot, payload, sizeof(payload));
require(size == kResponseHeaderSize + kPairingPayloadHeaderSize +
2 * kPairingRecordSize &&
payload[5] ==
static_cast<uint8_t>(Operation::kPairingRead) &&
payload[7] == 1 && payload[12] == 0x12 &&
payload[kResponseHeaderSize] == 2 &&
payload[kResponseHeaderSize + 1] == 1 &&
payload[kResponseHeaderSize + 4] == 1 &&
payload[kResponseHeaderSize + 5] == 0xfe &&
memcmp(&payload[kResponseHeaderSize + 6],
classic_address, 6) == 0,
"pairings were not migrated into the versioned envelope");
}
void perform_out(UsbConfigurationManagement::Operation operation,
const std::vector<uint8_t>& payload,
bool expected_ack = true) {
next_out_payload = make_request(operation, payload);
tusb_control_request_t request = setup_request(
operation, TUSB_DIR_OUT,
static_cast<uint16_t>(next_out_payload.size()));
require(tud_vendor_control_xfer_cb(
0, CONTROL_STAGE_SETUP, &request),
"valid OUT setup was rejected");
require(tud_vendor_control_xfer_cb(
0, CONTROL_STAGE_ACK, &request) == expected_ack,
"OUT acknowledgement result was incorrect");
}
void test_vendor_requests() {
using namespace UsbConfigurationManagement;
current_pairings = {};
current_pairings.generation = 7;
current_pairings.status = Bluepad32PairingSnapshotStatus::kReady;
current_pairings.record_count = 1;
current_pairings.records[0].transport =
Bluepad32PairingTransport::kClassic;
tusb_control_request_t request = setup_request(
Operation::kPairingRead, TUSB_DIR_IN, kMaximumResponseSize);
require(tud_vendor_control_xfer_cb(
0, CONTROL_STAGE_SETUP, &request) &&
control_payload[5] ==
static_cast<uint8_t>(Operation::kPairingRead) &&
control_payload[12] == 7,
"pairing read did not use the versioned envelope");
perform_out(Operation::kPairingRefresh, {});
require(refresh_requested,
"pairing refresh was not dispatched");
perform_out(Operation::kPairingClear, {});
require(clear_requested, "pairing clear was not dispatched");
std::vector<uint8_t> begin(12);
write_u32(&begin, 0, 0x11223344);
write_u16(&begin, 4, ADAPTER_CONFIGURATION_SCHEMA_VERSION);
write_u16(&begin, 6, ADAPTER_CONFIGURATION_ENCODED_SIZE);
write_u32(&begin, 8, 0xaabbccdd);
perform_out(Operation::kConfigurationBegin, begin);
require(begin_transaction_id == 0x11223344,
"configuration begin was not dispatched");
std::vector<uint8_t> chunk(12);
write_u32(&chunk, 0, 0x11223344);
write_u16(&chunk, 4, 0);
write_u16(&chunk, 6, 4);
chunk[8] = 60;
perform_out(Operation::kConfigurationChunk, chunk);
require(append_transaction_id == 0x11223344 &&
append_offset == 0 && appended_bytes.size() == 4,
"configuration chunk was not dispatched");
std::vector<uint8_t> commit(4);
write_u32(&commit, 0, 0x11223344);
perform_out(Operation::kConfigurationCommit, commit);
require(commit_transaction_id == 0x11223344,
"configuration commit was not dispatched");
next_out_payload =
make_request(Operation::kPairingRefresh, {});
next_out_payload[12] ^= 1;
request = setup_request(
Operation::kPairingRefresh, TUSB_DIR_OUT,
static_cast<uint16_t>(next_out_payload.size()));
require(tud_vendor_control_xfer_cb(
0, CONTROL_STAGE_SETUP, &request) &&
!tud_vendor_control_xfer_cb(
0, CONTROL_STAGE_ACK, &request),
"bad request CRC was accepted");
request.wValue = 0;
require(!tud_vendor_control_xfer_cb(
0, CONTROL_STAGE_SETUP, &request),
"request with invalid magic was accepted");
}
} // namespace
uint32_t configuration_crc32(const uint8_t* data, size_t size) {
uint32_t crc = 0xffffffffu;
for (size_t index = 0; index < size; ++index) {
crc ^= data[index];
for (uint8_t bit = 0; bit < 8; ++bit) {
const uint32_t mask = 0u - (crc & 1u);
crc = (crc >> 1) ^ (0xedb88320u & mask);
}
}
return ~crc;
}
void configuration_service_snapshot(ConfigurationServiceSnapshot* output) {
*output = current_configuration;
}
ConfigurationTransactionStatus configuration_service_begin(
uint32_t transaction_id, uint16_t, size_t, uint32_t) {
begin_transaction_id = transaction_id;
return ConfigurationTransactionStatus::kReceiving;
}
ConfigurationTransactionStatus configuration_service_append(
uint32_t transaction_id, size_t offset, const uint8_t* data,
size_t size) {
append_transaction_id = transaction_id;
append_offset = offset;
appended_bytes.assign(data, data + size);
return ConfigurationTransactionStatus::kReceiving;
}
ConfigurationTransactionStatus configuration_service_commit(
uint32_t transaction_id) {
commit_transaction_id = transaction_id;
return ConfigurationTransactionStatus::kPending;
}
ConfigurationTransactionStatus configuration_service_reset(uint32_t) {
return ConfigurationTransactionStatus::kPending;
}
void bluepad32_input_backend_request_pairing_snapshot() {
refresh_requested = true;
}
void bluepad32_input_backend_clear_pairings() {
clear_requested = true;
}
void bluepad32_input_backend_pairing_snapshot(
Bluepad32PairingSnapshot* out) {
*out = current_pairings;
}
bool tud_control_xfer(uint8_t, const tusb_control_request_t* request,
void* buffer, uint16_t length) {
if (request->bmRequestType_bit.direction == TUSB_DIR_OUT) {
if (next_out_payload.size() != length) {
return false;
}
memcpy(buffer, next_out_payload.data(), length);
} else {
const auto* bytes = static_cast<const uint8_t*>(buffer);
control_payload.assign(bytes, bytes + length);
}
return true;
}
bool tud_control_status(uint8_t, const tusb_control_request_t*) {
return true;
}
#include "../adapter_configuration.cpp"
#include "../usb_configuration_management.cpp"
int main() {
current_configuration.state = ConfigurationServiceState::kReady;
current_configuration.configuration =
adapter_configuration_default();
test_envelope_encoding();
test_pairing_encoding();
test_vendor_requests();
return 0;
}

View file

@ -9,6 +9,7 @@ enum {
CONTROL_STAGE_ACK = 2,
TUSB_REQ_RCPT_DEVICE = 0,
TUSB_DIR_OUT = 0,
TUSB_REQ_TYPE_VENDOR = 2,
TUSB_DIR_IN = 1,
};

View file

@ -1,144 +0,0 @@
#include "usb_pairing_management.h"
#include <cstdlib>
#include <cstring>
#include <iostream>
#include <vector>
#include <tusb.h>
namespace {
Bluepad32PairingSnapshot current_snapshot{};
bool refresh_requested = false;
bool clear_requested = false;
bool control_status_sent = false;
std::vector<uint8_t> control_payload;
void require(bool condition, const char* message) {
if (!condition) {
std::cerr << message << '\n';
std::exit(1);
}
}
void test_encoding() {
Bluepad32PairingSnapshot snapshot{};
snapshot.generation = 0x78563412;
snapshot.status = Bluepad32PairingSnapshotStatus::kReady;
snapshot.record_count = 2;
snapshot.overflow = true;
snapshot.records[0].transport =
Bluepad32PairingTransport::kClassic;
snapshot.records[0].address_type = 0xfe;
const uint8_t classic_address[6] = {1, 2, 3, 4, 5, 6};
memcpy(snapshot.records[0].address, classic_address, 6);
snapshot.records[1].transport = Bluepad32PairingTransport::kBle;
snapshot.records[1].address_type = 2;
const uint8_t ble_address[6] = {6, 5, 4, 3, 2, 1};
memcpy(snapshot.records[1].address, ble_address, 6);
uint8_t payload[UsbPairingManagement::kMaximumResponseSize]{};
const size_t size = UsbPairingManagement::encode_snapshot(
snapshot, payload, sizeof(payload));
require(size == UsbPairingManagement::kResponseHeaderSize +
2 * UsbPairingManagement::kRecordSize,
"snapshot encoded with the wrong size");
require(memcmp(payload, "SPPM", 4) == 0 &&
payload[4] == UsbPairingManagement::kProtocolVersion &&
payload[5] == 0 && payload[6] == 2 && payload[7] == 1,
"snapshot header encoding is invalid");
require(payload[8] == 0x12 && payload[9] == 0x34 &&
payload[10] == 0x56 && payload[11] == 0x78,
"snapshot generation is not little endian");
require(payload[12] == 1 && payload[13] == 0xfe &&
memcmp(&payload[14], classic_address, 6) == 0 &&
payload[20] == 2 && payload[21] == 2 &&
memcmp(&payload[22], ble_address, 6) == 0,
"pairing records are encoded incorrectly");
require(UsbPairingManagement::encode_snapshot(
snapshot, payload, size - 1) == 0,
"encoder accepted a short destination buffer");
}
void test_vendor_requests() {
current_snapshot = {};
current_snapshot.generation = 7;
current_snapshot.status = Bluepad32PairingSnapshotStatus::kReady;
current_snapshot.record_count = 1;
current_snapshot.records[0].transport =
Bluepad32PairingTransport::kClassic;
tusb_control_request_t request{};
request.bmRequestType_bit.recipient = TUSB_REQ_RCPT_DEVICE;
request.bmRequestType_bit.direction = TUSB_DIR_IN;
request.bRequest = UsbPairingManagement::kRequestGet;
request.wValue = UsbPairingManagement::kRequestValue;
request.wIndex = UsbPairingManagement::kRequestIndex;
request.wLength = UsbPairingManagement::kMaximumResponseSize;
require(tud_vendor_control_xfer_cb(
0, CONTROL_STAGE_SETUP, &request) &&
control_payload.size() ==
UsbPairingManagement::kResponseHeaderSize +
UsbPairingManagement::kRecordSize &&
control_payload[8] == 7,
"GET request did not return the current pairing snapshot");
request.bmRequestType_bit.direction = TUSB_DIR_OUT;
request.wLength = 0;
request.bRequest = UsbPairingManagement::kRequestRefresh;
require(tud_vendor_control_xfer_cb(
0, CONTROL_STAGE_SETUP, &request) &&
refresh_requested && control_status_sent,
"REFRESH request was not acknowledged and queued");
control_status_sent = false;
request.bRequest = UsbPairingManagement::kRequestClear;
require(tud_vendor_control_xfer_cb(
0, CONTROL_STAGE_SETUP, &request) &&
clear_requested && control_status_sent,
"CLEAR request was not acknowledged and queued");
request.wValue = 0;
require(!tud_vendor_control_xfer_cb(
0, CONTROL_STAGE_SETUP, &request),
"request with invalid magic was accepted");
require(tud_vendor_control_xfer_cb(
0, CONTROL_STAGE_ACK, &request),
"non-setup control stage was rejected");
}
} // namespace
void bluepad32_input_backend_request_pairing_snapshot() {
refresh_requested = true;
}
void bluepad32_input_backend_clear_pairings() {
clear_requested = true;
}
void bluepad32_input_backend_pairing_snapshot(
Bluepad32PairingSnapshot* out) {
*out = current_snapshot;
}
bool tud_control_xfer(uint8_t, const tusb_control_request_t*,
void* buffer, uint16_t length) {
const auto* bytes = static_cast<const uint8_t*>(buffer);
control_payload.assign(bytes, bytes + length);
return true;
}
bool tud_control_status(uint8_t, const tusb_control_request_t*) {
control_status_sent = true;
return true;
}
#include "../usb_pairing_management.cpp"
int main() {
test_encoding();
test_vendor_requests();
return 0;
}