use embassy_futures::join::join5;
use embassy_futures::select::{Either, select};
use embassy_sync::signal::Signal;
#[cfg(feature = "usb_log")]
use embassy_usb::class::cdc_acm::CdcAcmClass;
use embassy_usb::class::hid::{HidReader, HidReaderWriter, HidWriter, ReportId, RequestHandler};
use embassy_usb::control::OutResponse;
use embassy_usb::driver::{Driver, EndpointError};
use embassy_usb::{Builder, Handler, UsbDevice};
use rmk_types::connection::{ConnectionType, UsbState};
use static_cell::StaticCell;
use usbd_hid::descriptor::AsInputReport;
use crate::RawMutex;
use crate::channel::USB_REPORT_CHANNEL;
use crate::config::DeviceConfig;
use crate::core_traits::Runnable;
#[cfg(feature = "steno")]
use crate::hid::StenoReport;
use crate::hid::{
CompositeReport, CompositeReportType, HidError, HidWriterTrait, KeyboardReport, Report, run_led_reader,
};
use crate::light::UsbLedReader;
use crate::state::{current_usb_state, set_usb_state};
#[cfg(any(feature = "rynk", all(feature = "dongle", not(feature = "vial"))))]
pub(crate) mod rynk;
#[cfg(feature = "vial")]
pub(crate) mod vial;
#[cfg(any(feature = "rynk", all(feature = "dongle", not(feature = "vial"))))]
use rynk as host_usb;
#[cfg(feature = "vial")]
use vial as host_usb;
pub(crate) static USB_REMOTE_WAKEUP: Signal<RawMutex, ()> = Signal::new();
pub(crate) trait HostSession {
async fn serve<R: embedded_io_async::Read, W: embedded_io_async::Write>(&self, rx: &mut R, tx: &mut W);
}
impl HostSession for () {
async fn serve<R: embedded_io_async::Read, W: embedded_io_async::Write>(&self, _rx: &mut R, _tx: &mut W) {
core::future::pending().await
}
}
pub(crate) struct UsbKeyboardWriter<'a, 'd, D: Driver<'d>> {
pub(crate) keyboard_writer: &'a mut HidWriter<'d, D, 8>,
pub(crate) other_writer: &'a mut HidWriter<'d, D, 9>,
#[cfg(feature = "steno")]
pub(crate) steno_writer: &'a mut HidWriter<'d, D, 9>,
}
impl<'d, D: Driver<'d>> UsbKeyboardWriter<'_, 'd, D> {
pub(crate) async fn run_writer(&mut self) -> ! {
loop {
let report = USB_REPORT_CHANNEL.receive().await;
if current_usb_state() == UsbState::Suspended {
USB_REMOTE_WAKEUP.signal(());
continue;
}
if let Err(e) = self.write_report(&report).await {
error!("Failed to send report: {:?}", e);
if let HidError::UsbEndpointError(EndpointError::Disabled) = e {
USB_REMOTE_WAKEUP.signal(());
embassy_time::Timer::after_millis(500).await;
if let Err(e) = self.write_report(&report).await {
error!("Failed to send report after wakeup: {:?}", e);
}
}
}
}
}
async fn write_composite<R: AsInputReport>(
&mut self,
kind: CompositeReportType,
report: &R,
) -> Result<usize, HidError> {
let mut buf = [0u8; 9];
buf[0] = kind as u8;
let n = report
.serialize(&mut buf[1..])
.map_err(|_| HidError::ReportSerializeError)?;
self.other_writer
.write(&buf[0..n + 1])
.await
.map_err(HidError::UsbEndpointError)?;
Ok(n)
}
}
impl<'d, D: Driver<'d>> HidWriterTrait for UsbKeyboardWriter<'_, 'd, D> {
type ReportType = Report;
async fn write_report(&mut self, report: &Self::ReportType) -> Result<usize, HidError> {
match report {
Report::KeyboardReport(keyboard_report) => {
let mut buf: [u8; 8] = [0; 8];
let n: usize = keyboard_report
.serialize(&mut buf)
.map_err(|_| HidError::ReportSerializeError)?;
self.keyboard_writer
.write(&buf[0..n])
.await
.map_err(HidError::UsbEndpointError)?;
Ok(n)
}
Report::MouseReport(r) => self.write_composite(CompositeReportType::Mouse, r).await,
Report::MediaKeyboardReport(r) => self.write_composite(CompositeReportType::Media, r).await,
Report::SystemControlReport(r) => self.write_composite(CompositeReportType::System, r).await,
#[cfg(feature = "steno")]
Report::StenoReport(steno_report) => {
let mut buf: [u8; 9] = [0; 9];
let n = steno_report
.serialize(&mut buf)
.map_err(|_| HidError::ReportSerializeError)?;
match embassy_time::with_timeout(
embassy_time::Duration::from_millis(5),
self.steno_writer.write(&buf[0..n]),
)
.await
{
Ok(Ok(())) => {}
Ok(Err(e)) => return Err(HidError::UsbEndpointError(e)),
Err(_) => {} }
Ok(n)
}
}
}
}
const DEFAULT_CONFIG_DESC_SIZE: usize = if cfg!(any(
feature = "usb_log",
feature = "steno",
feature = "dfu",
feature = "rynk",
all(feature = "dongle", not(feature = "vial"))
)) {
256
} else {
128
};
fn default_config_descriptor() -> &'static mut [u8] {
static CONFIG_DESC: StaticCell<[u8; DEFAULT_CONFIG_DESC_SIZE]> = StaticCell::new();
&mut CONFIG_DESC.init([0; DEFAULT_CONFIG_DESC_SIZE])[..]
}
pub(crate) fn new_usb_builder<'d, D: Driver<'d>>(
driver: D,
keyboard_config: DeviceConfig<'d>,
config_descriptor: &'d mut [u8],
) -> Builder<'d, D> {
let mut usb_config = embassy_usb::Config::new(keyboard_config.vid, keyboard_config.pid);
usb_config.manufacturer = Some(keyboard_config.manufacturer);
usb_config.product = Some(keyboard_config.product_name);
#[cfg(any(feature = "rynk", all(feature = "dongle", not(feature = "vial"))))]
let serial_number = {
static SERIAL: StaticCell<heapless::String<64>> = StaticCell::new();
let s = SERIAL.init(heapless::String::new());
let _ = s.push_str(rmk_types::protocol::rynk::RYNK_MAGIC);
let _ = s.push_str(keyboard_config.serial_number);
s.as_str()
};
#[cfg(not(any(feature = "rynk", all(feature = "dongle", not(feature = "vial")))))]
let serial_number = keyboard_config.serial_number;
usb_config.serial_number = Some(serial_number);
usb_config.max_power = 450;
usb_config.supports_remote_wakeup = true;
usb_config.max_packet_size_0 = 64;
usb_config.device_class = 0xEF;
usb_config.device_sub_class = 0x02;
usb_config.device_protocol = 0x01;
usb_config.composite_with_iads = true;
#[cfg(feature = "dfu")]
const CONTROL_BUF_SIZE: usize = crate::dfu::BLOCK_SIZE_DFU;
#[cfg(not(feature = "dfu"))]
const CONTROL_BUF_SIZE: usize = DEFAULT_CONFIG_DESC_SIZE;
const RYNK_INTERFACE: bool = cfg!(any(feature = "rynk", all(feature = "dongle", not(feature = "vial"))));
const BOS_BUF_SIZE: usize = if RYNK_INTERFACE { 64 } else { 16 };
const MSOS_BUF_SIZE: usize = if RYNK_INTERFACE { 256 } else { 16 };
static BOS_DESC: StaticCell<[u8; BOS_BUF_SIZE]> = StaticCell::new();
static MSOS_DESC: StaticCell<[u8; MSOS_BUF_SIZE]> = StaticCell::new();
static CONTROL_BUF: StaticCell<[u8; CONTROL_BUF_SIZE]> = StaticCell::new();
let mut builder = Builder::new(
driver,
usb_config,
config_descriptor,
&mut BOS_DESC.init([0; BOS_BUF_SIZE])[..],
&mut MSOS_DESC.init([0; MSOS_BUF_SIZE])[..],
&mut CONTROL_BUF.init([0; CONTROL_BUF_SIZE])[..],
);
static device_handler: StaticCell<UsbDeviceHandler> = StaticCell::new();
builder.handler(device_handler.init(UsbDeviceHandler::new()));
builder
}
pub struct UsbTransport<'a, D: Driver<'static>, S = ()> {
device: UsbDevice<'static, D>,
keyboard_reader: HidReader<'static, D, 1>,
keyboard_writer: HidWriter<'static, D, 8>,
other_writer: HidWriter<'static, D, 9>,
#[cfg(feature = "steno")]
steno_writer: HidWriter<'static, D, 9>,
#[cfg(feature = "usb_log")]
logger: Option<embassy_usb::class::cdc_acm::CdcAcmClass<'static, D>>,
#[cfg(any(feature = "host", feature = "dongle"))]
host_reader: host_usb::HostUsbReader<D>,
#[cfg(any(feature = "host", feature = "dongle"))]
host_writer: host_usb::HostUsbWriter<D>,
session: &'a S,
}
impl<'a, D: Driver<'static>> UsbTransport<'a, D> {
pub fn new(driver: D, device_config: DeviceConfig<'static>) -> Self {
UsbTransportBuilder::new(driver, device_config, default_config_descriptor()).build()
}
pub fn builder(driver: D, device_config: DeviceConfig<'static>) -> UsbTransportBuilder<D> {
const SIZE: usize = DEFAULT_CONFIG_DESC_SIZE + 256;
static CONFIG_DESC: StaticCell<[u8; SIZE]> = StaticCell::new();
UsbTransportBuilder::new(driver, device_config, &mut CONFIG_DESC.init([0; SIZE])[..])
}
}
pub struct UsbTransportBuilder<D: Driver<'static>> {
builder: Builder<'static, D>,
keyboard_rw: HidReaderWriter<'static, D, 1, 8>,
other_writer: HidWriter<'static, D, 9>,
#[cfg(feature = "steno")]
steno_writer: HidWriter<'static, D, 9>,
#[cfg(feature = "usb_log")]
logger: CdcAcmClass<'static, D>,
#[cfg(any(feature = "host", feature = "dongle"))]
host_reader: host_usb::HostUsbReader<D>,
#[cfg(any(feature = "host", feature = "dongle"))]
host_writer: host_usb::HostUsbWriter<D>,
}
impl<D: Driver<'static>> UsbTransportBuilder<D> {
#[inline(always)]
fn new(driver: D, device_config: DeviceConfig<'static>, config_descriptor: &'static mut [u8]) -> Self {
#[cfg(feature = "_nrf_ble")]
let device_config = {
let mut device_config = device_config;
device_config.serial_number = crate::ble::nrf::get_serial_number();
device_config
};
let mut builder: Builder<'static, D> = new_usb_builder(driver, device_config, config_descriptor);
let keyboard_rw = add_usb_reader_writer!(
&mut builder,
KeyboardReport,
1,
8,
8,
::embassy_usb::class::hid::HidSubclass::Boot,
::embassy_usb::class::hid::HidBootProtocol::Keyboard
);
let other_writer = add_usb_writer!(&mut builder, CompositeReport, 9, 16);
#[cfg(feature = "steno")]
let steno_writer = add_usb_writer!(&mut builder, StenoReport, 9, 16);
#[cfg(feature = "usb_log")]
let logger = add_usb_logger!(&mut builder);
#[cfg(any(feature = "dfu_rp", feature = "dfu_nrf"))]
if let Some(mgr) = crate::dfu::get_manager() {
crate::dfu::register_dfu_interface(
&mut builder,
mgr,
device_config.product_name,
#[cfg(feature = "dfu_split")]
crate::SPLIT_PERIPHERALS_NUM,
);
}
#[cfg(any(feature = "host", feature = "dongle"))]
let (host_reader, host_writer) = host_usb::build_host_usb(&mut builder);
Self {
builder,
keyboard_rw,
other_writer,
#[cfg(feature = "steno")]
steno_writer,
#[cfg(feature = "usb_log")]
logger,
#[cfg(any(feature = "host", feature = "dongle"))]
host_reader,
#[cfg(any(feature = "host", feature = "dongle"))]
host_writer,
}
}
pub fn usb_builder(&mut self) -> &mut Builder<'static, D> {
&mut self.builder
}
#[inline(always)]
pub fn build<'a>(self) -> UsbTransport<'a, D> {
let (keyboard_reader, keyboard_writer) = self.keyboard_rw.split();
UsbTransport {
device: self.builder.build(),
keyboard_reader,
keyboard_writer,
other_writer: self.other_writer,
#[cfg(feature = "steno")]
steno_writer: self.steno_writer,
#[cfg(feature = "usb_log")]
logger: Some(self.logger),
#[cfg(any(feature = "host", feature = "dongle"))]
host_reader: self.host_reader,
#[cfg(any(feature = "host", feature = "dongle"))]
host_writer: self.host_writer,
session: &(),
}
}
}
impl<'a, D: Driver<'static>, S> UsbTransport<'a, D, S> {
#[cfg(feature = "host")]
pub fn with_host_service(
self,
service: &'a crate::host::HostService<'a>,
) -> UsbTransport<'a, D, crate::host::HostService<'a>> {
self.serving(service)
}
#[cfg(feature = "dongle")]
pub fn with_dongle_router(
self,
router: &'a crate::dongle::DongleRouter,
) -> UsbTransport<'a, D, crate::dongle::DongleRouter> {
self.serving(router)
}
fn serving<T>(self, session: &'a T) -> UsbTransport<'a, D, T> {
UsbTransport {
device: self.device,
keyboard_reader: self.keyboard_reader,
keyboard_writer: self.keyboard_writer,
other_writer: self.other_writer,
#[cfg(feature = "steno")]
steno_writer: self.steno_writer,
#[cfg(feature = "usb_log")]
logger: self.logger,
#[cfg(any(feature = "host", feature = "dongle"))]
host_reader: self.host_reader,
#[cfg(any(feature = "host", feature = "dongle"))]
host_writer: self.host_writer,
session,
}
}
}
impl<D: Driver<'static>, S: HostSession> Runnable for UsbTransport<'_, D, S> {
async fn run(&mut self) -> ! {
let Self {
device,
keyboard_reader,
keyboard_writer,
other_writer,
#[cfg(feature = "steno")]
steno_writer,
#[cfg(feature = "usb_log")]
logger,
#[cfg(any(feature = "host", feature = "dongle"))]
host_reader,
#[cfg(any(feature = "host", feature = "dongle"))]
host_writer,
session,
} = self;
let usb_device_task = async {
loop {
device.run_until_suspend().await;
match select(device.wait_resume(), USB_REMOTE_WAKEUP.wait()).await {
Either::First(_) => continue,
Either::Second(_) => {
info!("USB remote wakeup requested");
if let Err(e) = device.remote_wakeup().await {
warn!("Remote wakeup failed: {:?}", e);
}
}
}
}
};
let mut writer = UsbKeyboardWriter {
keyboard_writer,
other_writer,
#[cfg(feature = "steno")]
steno_writer,
};
let writer_task = writer.run_writer();
let mut led_reader = UsbLedReader::new(keyboard_reader);
let led_task = run_led_reader(&mut led_reader, ConnectionType::Usb);
#[cfg(any(feature = "host", feature = "dongle"))]
let host_task = host_usb::run_host_usb(host_reader, host_writer, *session);
#[cfg(not(any(feature = "host", feature = "dongle")))]
let host_task = {
let _ = session;
core::future::pending::<()>()
};
#[cfg(feature = "usb_log")]
let logger_task = run_usb_logger(logger.take().expect("UsbTransport::run called twice"));
#[cfg(not(feature = "usb_log"))]
let logger_task = core::future::pending::<()>();
join5(usb_device_task, writer_task, led_task, host_task, logger_task).await;
unreachable!("UsbTransport sub-tasks must run forever");
}
}
#[cfg(feature = "usb_log")]
async fn run_usb_logger<D: Driver<'static>>(logger_class: CdcAcmClass<'static, D>) {
let logger_fut =
::embassy_usb_logger::with_custom_style!(1024, log::LevelFilter::Trace, logger_class, |record, writer| {
use core::fmt::Write;
let ms = embassy_time::Instant::now().as_millis();
let _ = write!(writer, "[{:>8}ms {:5}] {}\r\n", ms, record.level(), record.args());
});
logger_fut.await;
}
#[cfg(any(feature = "usb_log", feature = "dfu_nrf", feature = "dfu_rp"))]
pub async fn run_peripheral_usb<D: Driver<'static>>(driver: D, config: DeviceConfig<'static>) {
let mut builder = new_usb_builder(driver, config, default_config_descriptor());
#[cfg(feature = "usb_log")]
let logger_fut = run_usb_logger(add_usb_logger!(&mut builder));
#[cfg(not(feature = "usb_log"))]
let logger_fut = ::core::future::pending::<()>();
#[cfg(any(feature = "dfu_rp", feature = "dfu_nrf"))]
if let Some(mgr) = crate::dfu::get_manager() {
crate::dfu::register_dfu_interface(
&mut builder,
mgr,
config.product_name,
#[cfg(feature = "dfu_split")]
0,
);
}
let mut usb_device = builder.build();
::embassy_futures::join::join(usb_device.run(), logger_fut).await;
}
#[cfg(feature = "usb_log")]
macro_rules! add_usb_logger {
($usb_builder:expr) => {{
use embassy_usb::class::cdc_acm::{CdcAcmClass, State};
use static_cell::StaticCell;
static LOGGER_STATE: StaticCell<State> = StaticCell::new();
let state = LOGGER_STATE.init(State::new());
CdcAcmClass::new($usb_builder, state, embassy_usb_logger::MAX_PACKET_SIZE as u16)
}};
}
macro_rules! usb_hid_state_and_config {
($descriptor:ty, $max_packet:expr, $subclass:expr, $protocol:expr) => {{
use usbd_hid::descriptor::SerializedDescriptor;
paste::paste! {
static [<$descriptor:snake:upper _STATE>]: ::static_cell::StaticCell<::embassy_usb::class::hid::State> = ::static_cell::StaticCell::new();
static [<$descriptor:snake:upper _HANDLER>]: ::static_cell::StaticCell<$crate::usb::UsbRequestHandler> = ::static_cell::StaticCell::new();
}
let state = paste::paste! { [<$descriptor:snake:upper _STATE>].init(::embassy_usb::class::hid::State::new()) };
let request_handler = paste::paste! { [<$descriptor:snake:upper _HANDLER>].init($crate::usb::UsbRequestHandler {}) };
let hid_config = ::embassy_usb::class::hid::Config {
report_descriptor: <$descriptor>::desc(),
request_handler: Some(request_handler),
poll_ms: 1,
max_packet_size: $max_packet,
hid_subclass: $subclass,
hid_boot_protocol: $protocol,
};
(state, hid_config)
}};
}
macro_rules! add_usb_writer {
($usb_builder:expr, $descriptor:ty, $n:expr, $max_packet:expr) => {{
let (state, hid_config) = $crate::usb::usb_hid_state_and_config!(
$descriptor,
$max_packet,
::embassy_usb::class::hid::HidSubclass::No,
::embassy_usb::class::hid::HidBootProtocol::None
);
::embassy_usb::class::hid::HidWriter::<_, $n>::new($usb_builder, state, hid_config)
}};
}
macro_rules! add_usb_reader_writer {
($usb_builder:expr, $descriptor:ty, $read_n:expr, $write_n:expr, $max_packet:expr) => {
$crate::usb::add_usb_reader_writer!(
$usb_builder,
$descriptor,
$read_n,
$write_n,
$max_packet,
::embassy_usb::class::hid::HidSubclass::No,
::embassy_usb::class::hid::HidBootProtocol::None
)
};
($usb_builder:expr, $descriptor:ty, $read_n:expr, $write_n:expr, $max_packet:expr, $subclass:expr, $protocol:expr) => {{
let (state, hid_config) =
$crate::usb::usb_hid_state_and_config!($descriptor, $max_packet, $subclass, $protocol);
::embassy_usb::class::hid::HidReaderWriter::<_, $read_n, $write_n>::new($usb_builder, state, hid_config)
}};
}
#[cfg(feature = "usb_log")]
pub(crate) use add_usb_logger;
pub(crate) use add_usb_reader_writer;
pub(crate) use add_usb_writer;
pub(crate) use usb_hid_state_and_config;
pub(crate) struct UsbRequestHandler {}
impl RequestHandler for UsbRequestHandler {
fn set_report(&mut self, id: ReportId, data: &[u8]) -> OutResponse {
info!("Set report for {:?}: {:?}", id, data);
OutResponse::Accepted
}
}
pub(crate) struct UsbDeviceHandler {
pre_suspend: UsbState,
}
impl UsbDeviceHandler {
fn new() -> Self {
UsbDeviceHandler {
pre_suspend: UsbState::Disabled,
}
}
}
impl Handler for UsbDeviceHandler {
fn enabled(&mut self, enabled: bool) {
if enabled {
info!("Device enabled");
set_usb_state(UsbState::Enabled);
} else {
info!("Device disabled");
set_usb_state(UsbState::Disabled);
}
}
fn reset(&mut self) {
info!("Bus reset, the Vbus current limit is 100mA");
}
fn addressed(&mut self, addr: u8) {
info!("USB address set to: {}", addr);
}
fn configured(&mut self, configured: bool) {
if configured {
set_usb_state(UsbState::Configured);
info!("Device configured, it may now draw up to the configured current from Vbus.")
} else {
set_usb_state(UsbState::Enabled);
info!("Device is no longer configured, the Vbus current limit is 100mA.");
}
}
fn suspended(&mut self, suspended: bool) {
if suspended {
let live = current_usb_state();
if live == UsbState::Configured {
self.pre_suspend = live;
set_usb_state(UsbState::Suspended);
info!(
"Device suspended, the Vbus current limit is 500µA (or 2.5mA for high-power devices with remote wakeup enabled)."
);
} else if live != UsbState::Suspended {
info!("Bus suspended before enumeration (charger or charge-only cable?), USB stays inactive");
}
} else {
if current_usb_state() == UsbState::Suspended {
set_usb_state(self.pre_suspend);
}
info!(
"Device resumed, the Vbus current limit is 500µA (or 2.5mA for high-power devices with remote wakeup enabled)."
);
}
}
fn remote_wakeup_enabled(&mut self, enabled: bool) {
info!("Remote wakeup enabled state: {}", enabled);
}
}
#[cfg(test)]
mod tests {
use embassy_usb::Handler;
use rmk_types::connection::UsbState;
use super::UsbDeviceHandler;
use crate::state::{current_usb_state, set_usb_state};
#[test]
fn suspend_without_enumeration_stays_enabled() {
let mut handler = UsbDeviceHandler::new();
handler.enabled(true);
assert_eq!(current_usb_state(), UsbState::Enabled);
handler.suspended(true);
assert_eq!(current_usb_state(), UsbState::Enabled);
handler.suspended(false);
assert_eq!(current_usb_state(), UsbState::Enabled);
handler.configured(true);
assert_eq!(current_usb_state(), UsbState::Configured);
}
#[test]
fn suspend_after_configured_publishes_suspended_and_resume_restores() {
let mut handler = UsbDeviceHandler::new();
handler.enabled(true);
handler.configured(true);
handler.suspended(true);
assert_eq!(current_usb_state(), UsbState::Suspended);
handler.suspended(false);
assert_eq!(current_usb_state(), UsbState::Configured);
}
#[test]
fn duplicate_suspend_preserves_pre_suspend_state() {
let mut handler = UsbDeviceHandler::new();
handler.enabled(true);
handler.configured(true);
handler.suspended(true);
handler.suspended(true);
assert_eq!(current_usb_state(), UsbState::Suspended);
handler.suspended(false);
assert_eq!(current_usb_state(), UsbState::Configured);
}
#[test]
fn resume_without_suspend_is_a_no_op() {
let mut handler = UsbDeviceHandler::new();
set_usb_state(UsbState::Configured);
handler.suspended(false);
assert_eq!(current_usb_state(), UsbState::Configured);
}
}