use core::sync::atomic::{AtomicU8, Ordering};
use embedded_io_async::{ErrorType, Write};
use heapless::Vec;
use rmk_types::protocol::rynk::{RYNK_BLE_CHUNK_SIZE, RYNK_HID_REPORT_SIZE};
use trouble_host::prelude::*;
use super::HostWriteOutcome;
use crate::ble::ble_server::Server;
use crate::channel::RYNK_BLE_RX_PIPE;
use crate::host::rynk::RynkService;
use crate::host::transport::HostTransportError;
pub(crate) struct HostGattHandler {
custom_output_handle: u16,
custom_input_cccd_handle: u16,
hid_output_handle: u16,
hid_input_cccd_handle: u16,
hid_control_point_handle: u16,
}
impl HostGattHandler {
pub(crate) fn new(server: &Server<'_>) -> Self {
Self {
custom_output_handle: server.rynk_service.output_data.handle,
custom_input_cccd_handle: server
.rynk_service
.input_data
.cccd_handle
.expect("No CCCD for Rynk input"),
hid_output_handle: server.rynk_hid_service.output_data.handle,
hid_input_cccd_handle: server
.rynk_hid_service
.input_data
.cccd_handle
.expect("No CCCD for Rynk HID input"),
hid_control_point_handle: server.rynk_hid_service.hid_control_point.handle,
}
}
pub(crate) async fn handle_write(&mut self, handle: u16, data: &[u8], encrypted: bool) -> HostWriteOutcome {
if handle == self.custom_output_handle {
if !data.is_empty() {
if encrypted {
debug!("Got Rynk packet ({} bytes)", data.len());
RYNK_BLE_RX_PIPE.write_all(data).await;
RynkBleSource::Custom.activate();
} else {
warn!("Rynk: dropping {}-byte write on unencrypted link", data.len());
}
}
HostWriteOutcome::Handled
} else if handle == self.custom_input_cccd_handle {
HostWriteOutcome::CccdUpdated
} else if handle == self.hid_output_handle {
if encrypted {
if data.len() == RYNK_HID_REPORT_SIZE {
RYNK_BLE_RX_PIPE.write_all(data).await;
RynkBleSource::Hid.activate();
} else {
warn!("Wrong Rynk HID report size: {}", data.len());
}
} else {
warn!("Rynk HID: dropping {}-byte write on unencrypted link", data.len());
}
HostWriteOutcome::Handled
} else if handle == self.hid_input_cccd_handle {
HostWriteOutcome::CccdUpdated
} else if handle == self.hid_control_point_handle {
HostWriteOutcome::ControlPoint
} else {
HostWriteOutcome::Unhandled
}
}
pub(crate) async fn run<'stack, 'server, P: PacketPool>(
server: &'server Server<'_>,
conn: &GattConnection<'stack, 'server, P>,
service: &RynkService<'_>,
) {
RYNK_BLE_RX_PIPE.clear();
RynkBleSource::None.activate();
let mut rx = &RYNK_BLE_RX_PIPE;
let mut tx = RynkBleTx {
custom_input: server.rynk_service.input_data.clone(),
hid_input: server.rynk_hid_service.input_data,
conn,
};
service.run_session(&mut rx, &mut tx).await;
}
}
#[derive(Clone, Copy, PartialEq, Eq)]
#[repr(u8)]
pub(crate) enum RynkBleSource {
None,
Custom,
Hid,
}
static ACTIVE_SOURCE: AtomicU8 = AtomicU8::new(RynkBleSource::None as u8);
impl RynkBleSource {
pub(crate) fn activate(self) {
ACTIVE_SOURCE.store(self as u8, Ordering::Relaxed);
}
fn active() -> Self {
match ACTIVE_SOURCE.load(Ordering::Relaxed) {
v if v == Self::Custom as u8 => Self::Custom,
v if v == Self::Hid as u8 => Self::Hid,
_ => Self::None,
}
}
}
struct RynkBleTx<'a, 'b, 'c, P: PacketPool> {
custom_input: Characteristic<Vec<u8, RYNK_BLE_CHUNK_SIZE>>,
hid_input: Characteristic<[u8; RYNK_HID_REPORT_SIZE]>,
conn: &'a GattConnection<'b, 'c, P>,
}
impl<P: PacketPool> ErrorType for RynkBleTx<'_, '_, '_, P> {
type Error = HostTransportError;
}
impl<P: PacketPool> Write for RynkBleTx<'_, '_, '_, P> {
async fn write(&mut self, buf: &[u8]) -> Result<usize, Self::Error> {
if buf.is_empty() {
return Ok(0);
}
match RynkBleSource::active() {
RynkBleSource::Hid => {
for chunk in buf.chunks(RYNK_HID_REPORT_SIZE) {
let mut report = [0u8; RYNK_HID_REPORT_SIZE];
report[..chunk.len()].copy_from_slice(chunk);
if let Err(e) = self.hid_input.notify(self.conn, &report, true).await {
error!("Failed to notify Rynk HID reply: {:?}", e);
return Err(HostTransportError);
}
}
}
RynkBleSource::Custom => {
let max_notify = (self.conn.raw().att_mtu() as usize).saturating_sub(3);
let chunk_size = RYNK_BLE_CHUNK_SIZE.min(max_notify).max(1);
for chunk in buf.chunks(chunk_size) {
let payload =
Vec::<u8, RYNK_BLE_CHUNK_SIZE>::from_slice(chunk).expect("chunk size <= RYNK_BLE_CHUNK_SIZE");
if let Err(e) = self.custom_input.notify(self.conn, &payload, true).await {
error!("Failed to notify Rynk reply: {:?}", e);
return Err(HostTransportError);
}
}
}
RynkBleSource::None => {}
}
Ok(buf.len())
}
async fn flush(&mut self) -> Result<(), Self::Error> {
Ok(())
}
}
#[cfg(test)]
mod tests {
use rmk_types::protocol::rynk::{Cmd, Deframer, RYNK_HEADER_SIZE, RynkHeader, encode_frame};
use super::*;
#[test]
fn classifies_rynk_gatt_handles() {
use crate::test_support::test_block_on as block_on;
let server = Server::new_default("rmk").unwrap();
let mut handler = HostGattHandler::new(&server);
assert_eq!(
block_on(handler.handle_write(server.rynk_service.output_data.handle, &[], true)),
HostWriteOutcome::Handled
);
assert_eq!(
block_on(handler.handle_write(server.rynk_hid_service.output_data.handle, &[], true,)),
HostWriteOutcome::Handled
);
assert_eq!(
block_on(handler.handle_write(server.rynk_service.input_data.cccd_handle.unwrap(), &[], true,)),
HostWriteOutcome::CccdUpdated
);
assert_eq!(
block_on(handler.handle_write(server.rynk_hid_service.input_data.cccd_handle.unwrap(), &[], true,)),
HostWriteOutcome::CccdUpdated
);
assert_eq!(
block_on(handler.handle_write(server.rynk_hid_service.hid_control_point.handle, &[0], true,)),
HostWriteOutcome::ControlPoint
);
assert_eq!(
block_on(handler.handle_write(u16::MAX, &[], true)),
HostWriteOutcome::Unhandled
);
}
fn cobs_frame(buf: &mut [u8], seq: u8, payload: &[u8]) -> usize {
encode_frame(
buf,
RynkHeader {
cmd: Cmd::GetMacro,
seq,
},
&payload,
)
.unwrap()
}
fn read_one_frame() -> (u8, Vec<u8, 64>) {
use crate::test_support::test_block_on as block_on;
let rx = &RYNK_BLE_RX_PIPE;
let mut buf = [0u8; 256];
let mut df = Deframer::new();
loop {
let n = block_on(rx.read(df.tail(&mut buf)));
df.commit(n);
if let Some(frame_len) = df.next(&mut buf) {
let frame = &buf[..frame_len];
return (frame[2], postcard::from_bytes(&frame[RYNK_HEADER_SIZE..]).unwrap());
}
}
}
#[test]
fn fragments_reassemble_through_pipe_for_session() {
RYNK_BLE_RX_PIPE.clear();
let payload: [u8; 60] = core::array::from_fn(|i| i as u8);
let mut buf = [0u8; 128];
let n = cobs_frame(&mut buf, 9, &payload);
assert!(n > RYNK_HID_REPORT_SIZE, "frame must span multiple reports");
for chunk in buf[..n].chunks(RYNK_HID_REPORT_SIZE) {
let mut report = [0u8; RYNK_HID_REPORT_SIZE];
report[..chunk.len()].copy_from_slice(chunk);
assert_eq!(RYNK_BLE_RX_PIPE.try_write(&report).unwrap(), report.len());
}
let (seq, decoded) = read_one_frame();
assert_eq!(seq, 9);
assert_eq!(&decoded[..], &payload[..]);
}
#[test]
fn small_frame_drops_padding() {
RYNK_BLE_RX_PIPE.clear();
let mut buf = [0u8; 64];
let n = cobs_frame(&mut buf, 3, &[0xAA, 0xBB]);
let mut report = [0u8; RYNK_HID_REPORT_SIZE];
report[..n].copy_from_slice(&buf[..n]);
assert_eq!(RYNK_BLE_RX_PIPE.try_write(&report).unwrap(), report.len());
let (seq, decoded) = read_one_frame();
assert_eq!(seq, 3);
assert_eq!(&decoded[..], &[0xAA, 0xBB]);
}
}