use crate::{
ascii::{checksum::LrcWriter, id},
error::{AsciiCommandSplitError, AsciiError, AsciiReservedCharacterError},
};
use std::borrow::Cow;
use std::io;
mod private {
use super::Target;
pub trait Sealed {}
impl Sealed for str {}
impl Sealed for String {}
impl Sealed for [u8] {}
impl<const N: usize> Sealed for [u8; N] {}
impl Sealed for Vec<u8> {}
impl<T, D> Sealed for (T, D)
where
T: Into<Target> + Copy,
D: AsRef<[u8]>,
{
}
impl<D: AsRef<[u8]>> Sealed for (u8, u8, D) {}
impl<T> Sealed for &T where T: Sealed + ?Sized {}
impl<T> Sealed for &mut T where T: Sealed + ?Sized {}
impl<T> Sealed for Box<T> where T: Sealed + ?Sized {}
}
pub trait Command: private::Sealed {
fn target(&self) -> Target;
fn data(&self) -> Cow<'_, [u8]>;
}
impl Command for str {
fn target(&self) -> Target {
Target::for_all()
}
fn data(&self) -> Cow<'_, [u8]> {
self.as_bytes().into()
}
}
impl Command for [u8] {
fn target(&self) -> Target {
Target::for_all()
}
fn data(&self) -> Cow<'_, [u8]> {
self.into()
}
}
impl<const N: usize> Command for [u8; N] {
fn target(&self) -> Target {
Target::for_all()
}
fn data(&self) -> Cow<'_, [u8]> {
self.as_slice().into()
}
}
impl Command for String {
fn target(&self) -> Target {
Target::for_all()
}
fn data(&self) -> Cow<'_, [u8]> {
self.as_bytes().into()
}
}
impl Command for Vec<u8> {
fn target(&self) -> Target {
Target::for_all()
}
fn data(&self) -> Cow<'_, [u8]> {
self.as_slice().into()
}
}
impl<T, D> Command for (T, D)
where
T: Into<Target> + Copy,
D: AsRef<[u8]>,
{
fn target(&self) -> Target {
self.0.into()
}
fn data(&self) -> Cow<'_, [u8]> {
self.1.as_ref().into()
}
}
impl<D> Command for (u8, u8, D)
where
D: AsRef<[u8]>,
{
fn target(&self) -> Target {
Target::for_device(self.0).with_axis(self.1)
}
fn data(&self) -> Cow<'_, [u8]> {
self.2.as_ref().into()
}
}
impl<T> Command for &T
where
T: Command + ?Sized,
{
fn target(&self) -> Target {
(**self).target()
}
fn data(&self) -> Cow<'_, [u8]> {
(**self).data()
}
}
impl<T> Command for &mut T
where
T: Command + ?Sized,
{
fn target(&self) -> Target {
(**self).target()
}
fn data(&self) -> Cow<'_, [u8]> {
(**self).data()
}
}
impl<T> Command for Box<T>
where
T: Command + ?Sized,
{
fn target(&self) -> Target {
(**self).target()
}
fn data(&self) -> Cow<'_, [u8]> {
(**self).data()
}
}
#[derive(Debug, Copy, Clone, PartialEq, Eq, Hash)]
pub struct MaxPacketSize(usize);
impl MaxPacketSize {
const MIN_PACKET_SIZE: usize = 80;
pub const fn new(value: usize) -> Option<Self> {
if value >= Self::MIN_PACKET_SIZE {
Some(MaxPacketSize(value))
} else {
None
}
}
pub const fn default() -> MaxPacketSize {
MaxPacketSize(Self::MIN_PACKET_SIZE)
}
pub const fn as_usize(self) -> usize {
self.0
}
}
impl Default for MaxPacketSize {
fn default() -> Self {
MaxPacketSize::default()
}
}
pub(crate) struct CommandWriter<'a> {
pub target: Target,
pub id: Option<u8>,
pub data: Cow<'a, [u8]>,
pub offset: usize,
pub checksum: bool,
pub max_packet_size: MaxPacketSize,
pub packet_index: usize,
}
impl<'a> CommandWriter<'a> {
pub fn new<C, G>(
command: &'a C,
mut generator: G,
generate_id: bool,
generate_checksum: bool,
max_packet_size: MaxPacketSize,
) -> Result<CommandWriter<'a>, AsciiReservedCharacterError>
where
C: Command,
G: id::Generator,
{
let data = command.data();
if let Some(reserved) = data
.iter()
.find(|byte| **byte > 127 || b"/@#!\r\n:\\".contains(byte))
{
return Err(AsciiReservedCharacterError::new(command, *reserved));
}
Ok(CommandWriter {
target: command.target(),
id: if generate_id {
Some(generator.next_id())
} else {
None
},
data,
offset: 0,
checksum: generate_checksum,
max_packet_size,
packet_index: 0,
})
}
fn is_complete(&self) -> bool {
self.packet_index > 0
&& (self.offset >= self.data.len() || self.data.iter().all(u8::is_ascii_whitespace))
}
fn write_packet_header<W: io::Write>(
&mut self,
writer: &mut LrcWriter<W>,
) -> io::Result<usize> {
use std::io::Write as _;
let device_char_count = ascii_char_count(self.target.device() as usize);
let axis_char_count = ascii_char_count(self.target.axis() as usize);
write!(writer, "/")?;
let mut bytes_written = 1;
writer.reset_hash();
match self.id {
Some(id) => {
write!(
writer,
"{} {} {}",
self.target.device(),
self.target.axis(),
id
)?;
bytes_written +=
device_char_count + axis_char_count + ascii_char_count(id as usize) + 2;
}
None => {
if self.target.axis() != 0 {
write!(writer, "{} {}", self.target.device(), self.target.axis())?;
bytes_written += device_char_count + axis_char_count + 1; } else if self.target.device() != 0 {
write!(writer, "{}", self.target.device())?;
bytes_written += device_char_count;
}
}
};
Ok(bytes_written)
}
pub fn write_packet<W: io::Write + ?Sized>(
&mut self,
writer: &mut W,
) -> Result<bool, AsciiError> {
use std::io::Write as _;
if self.is_complete() {
return Ok(false);
}
let writer = &mut LrcWriter::new(writer);
let mut bytes_written = self.write_packet_header(writer)?;
let data = &self.data[self.offset..];
let mut words = data
.split(u8::is_ascii_whitespace)
.filter(|word| !word.is_empty()) .enumerate()
.peekable();
if words.peek().is_some() {
if bytes_written > 1 {
write!(writer, " ")?;
bytes_written += 1;
}
if self.packet_index != 0 {
write!(writer, "cont {} ", self.packet_index)?;
bytes_written += 6 + ascii_char_count(self.packet_index);
}
let mut remaining = self.max_packet_size.as_usize()
- bytes_written - if self.checksum { 3 } else { 0 } - 1; let mut wrote_word = false;
while let Some((index, word)) = words.next() {
let mut needed_bytes = word.len();
if index != 0 {
needed_bytes += 1; }
if needed_bytes <= remaining {
if words.peek().is_some() && needed_bytes == remaining {
self.offset = word.as_ptr() as usize - self.data.as_ptr() as usize;
break;
}
if index != 0 {
writer.write_all(b" ")?;
}
writer.write_all(word)?;
wrote_word = true;
remaining -= needed_bytes;
self.offset = words.peek().map_or_else(
|| self.data.len(), |(_, word)| word.as_ptr() as usize - self.data.as_ptr() as usize,
);
} else {
self.offset = word.as_ptr() as usize - self.data.as_ptr() as usize;
break;
}
}
if !wrote_word {
return Err(AsciiCommandSplitError::new((self.target, self.data.to_vec())).into());
}
if self.offset != self.data.len() {
writer.write_all(b"\\")?;
}
}
if self.checksum {
let checksum = writer.finish_hash();
write!(writer, ":{checksum:02X}")?;
}
writer.write_all(b"\n")?;
self.packet_index += 1;
Ok(!self.is_complete())
}
}
fn ascii_char_count(mut num: usize) -> usize {
let mut count = 1;
while num >= 10 {
count += 1;
num /= 10;
}
count
}
#[derive(Debug, Copy, Clone, PartialEq, Eq, Hash)]
pub struct Target(u8, u8);
impl Target {
pub const fn new(device: u8, axis: u8) -> Target {
Target(device, axis)
}
#[must_use]
pub const fn for_all() -> Target {
Target(0, 0)
}
#[must_use]
pub const fn for_device(address: u8) -> Target {
Target(address, 0)
}
#[must_use]
pub const fn with_all_axes(self) -> Target {
Target(self.0, 0)
}
#[must_use]
pub const fn with_axis(self, axis: u8) -> Target {
Target(self.0, axis)
}
pub const fn device(self) -> u8 {
self.0
}
pub const fn axis(self) -> u8 {
self.1
}
pub(crate) const fn elicited_by_command_to(self, target: Target) -> bool {
if target.0 == 0 || self.0 == target.0 {
target.1 == 0 || self.1 == target.1
} else {
false
}
}
}
impl Default for Target {
fn default() -> Target {
Target::for_all()
}
}
impl From<u8> for Target {
fn from(other: u8) -> Target {
Target(other, 0)
}
}
impl From<(u8, u8)> for Target {
fn from(other: (u8, u8)) -> Target {
Target(other.0, other.1)
}
}
#[cfg(test)]
mod test {
use super::*;
struct ConstId {}
impl id::Generator for ConstId {
fn next_id(&mut self) -> u8 {
5
}
}
#[test]
fn target_default_is_all() {
assert_eq!(Target::default(), Target::for_all());
}
#[test]
fn test_target() {
assert_eq!(Target::default(), Target::for_all());
assert_eq!(Target::for_device(1), Target(1, 0));
assert_eq!(Target::for_device(1).with_axis(1), Target(1, 1));
assert_eq!(Target::new(5, 9).with_all_axes(), Target(5, 0));
}
#[test]
fn test_command_writer() {
let mut buf = Vec::with_capacity(500);
struct Case {
command: &'static (u8, u8, &'static str),
generate_id: bool,
generate_checksum: bool,
expected: &'static [u8],
}
let cases = [
Case {
command: &(0, 0, ""),
generate_id: false,
generate_checksum: false,
expected: b"/\n",
},
Case {
command: &(0, 0, " \t"),
generate_id: false,
generate_checksum: false,
expected: b"/\n",
},
Case {
command: &(0, 0, ""),
generate_id: false,
generate_checksum: true,
expected: b"/:00\n",
},
Case {
command: &(0, 0, " \t"),
generate_id: false,
generate_checksum: true,
expected: b"/:00\n",
},
Case {
command: &(0, 0, ""),
generate_id: true,
generate_checksum: true,
expected: b"/0 0 5:2B\n",
},
Case {
command: &(0, 0, " \t"),
generate_id: true,
generate_checksum: true,
expected: b"/0 0 5:2B\n",
},
Case {
command: &(1, 0, ""),
generate_id: false,
generate_checksum: true,
expected: b"/1:CF\n",
},
Case {
command: &(0, 1, ""),
generate_id: false,
generate_checksum: true,
expected: b"/0 1:7F\n",
},
Case {
command: &(2, 1, ""),
generate_id: false,
generate_checksum: true,
expected: b"/2 1:7D\n",
},
Case {
command: &(1, 0, "tools echo"),
generate_id: false,
generate_checksum: true,
expected: b"/1 tools echo:BF\n",
},
Case {
command: &(0, 0, "get maxspeed"),
generate_id: false,
generate_checksum: true,
expected: b"/get maxspeed:49\n",
},
Case {
command: &(2, 0, "get maxspeed"),
generate_id: true,
generate_checksum: true,
expected: b"/2 0 5 get maxspeed:52\n",
},
Case {
command: &(1, 0, "tools echo aaaaaaaaaa bbbbbbbbbb ccccccccc dddddddddd eeeeeeeeee ffffffffff gggggggggg hhhhhhhhhh iiiiiiiiii jjjjjjjjj"),
generate_id: false,
generate_checksum: true,
expected: b"/1 tools echo aaaaaaaaaa bbbbbbbbbb ccccccccc dddddddddd eeeeeeeeee\\:D0\n/1 cont 1 ffffffffff gggggggggg hhhhhhhhhh iiiiiiiiii jjjjjjjjj:24\n",
},
Case {
command: &(0, 0, "aaaaaaaaaa bbbbbbbbbb ccccccccc dddddddddd eeeeeeeeee ffffffffff gggggggggg"), generate_id: false,
generate_checksum: true,
expected: b"/aaaaaaaaaa bbbbbbbbbb ccccccccc dddddddddd eeeeeeeeee ffffffffff gggggggggg:4B\n",
},
Case {
command: &(0, 0, "aaaaaaaaaa bbbbbbbbbb ccccccccc dddddddddd eeeeeeeeee ffffffffff gggggggggg h"), generate_id: false,
generate_checksum: true,
expected: b"/aaaaaaaaaa bbbbbbbbbb ccccccccc dddddddddd eeeeeeeeee ffffffffff\\:15\n/cont 1 gggggggggg h:4D\n",
},
Case {
command: &(0, 0, "aaaaaaaaaa bbbbbbbbbb ccccccccc dddddddddd eeeeeeeeee ffffffffff ggggggggg h"), generate_id: false,
generate_checksum: true,
expected: b"/aaaaaaaaaa bbbbbbbbbb ccccccccc dddddddddd eeeeeeeeee ffffffffff ggggggggg\\:56\n/cont 1 h:73\n",
},
Case {
command: &(1, 0, "aaaaaaaaaa bbbbbbbbbb ccccccccc dddddddddd eeeeeeeeee ffffffffff gggggggggg"), generate_id: false,
generate_checksum: true,
expected: b"/1 aaaaaaaaaa bbbbbbbbbb ccccccccc dddddddddd eeeeeeeeee ffffffffff\\:C4\n/1 cont 1 gggggggggg:84\n",
},
Case {
command: &(0, 0, "aaaaaaaaaa bbbbbbbbbb ccccccccc dddddddddd eeeeeeeeee ffffffffff gggggggggg"), generate_id: true,
generate_checksum: true,
expected: b"/0 0 5 aaaaaaaaaa bbbbbbbbbb ccccccccc dddddddddd eeeeeeeeee ffffffffff\\:20\n/0 0 5 cont 1 gggggggggg:E0\n",
},
Case {
command: &(0, 0, "aaaaaaaaaa bbbbbbbbbb ccccccccc dddddddddd \teeeeeeeeee ffffffffff gggggggggg h"), generate_id: false,
generate_checksum: true,
expected: b"/aaaaaaaaaa bbbbbbbbbb ccccccccc dddddddddd eeeeeeeeee ffffffffff\\:15\n/cont 1 gggggggggg h:4D\n",
},
];
for (case_index, case) in cases.into_iter().enumerate() {
eprintln!("cases[{}] = {:?}", case_index, case.command);
buf.clear();
let mut writer = CommandWriter::new(
&case.command,
ConstId {},
case.generate_id,
case.generate_checksum,
MaxPacketSize::default(),
)
.unwrap();
let num_expected_packets = case.expected.iter().filter(|byte| **byte == b'\n').count();
for index in 0..num_expected_packets {
let more = writer.write_packet(&mut buf).unwrap();
assert_eq!(
more,
index + 1 != num_expected_packets,
"packet {index}: unexpected write_packet result ({more}): {}",
std::str::from_utf8(&buf).unwrap()
);
}
assert_eq!(
buf,
case.expected,
"unexpected output: {}",
String::from_utf8_lossy(&buf)
);
}
}
#[test]
fn test_command_writer_custom_packet_size() {
let mut buf = vec![];
let _79_bytes =
"1234567891123456789212345678931234567894123456789512345678961234567897123456789";
{
let mut writer = CommandWriter::new(
&_79_bytes,
ConstId {},
false,
false,
MaxPacketSize::default(),
)
.unwrap();
writer.write_packet(&mut buf).unwrap_err();
}
{
let mut writer = CommandWriter::new(
&_79_bytes,
ConstId {},
false,
false,
MaxPacketSize::new(81).unwrap(),
)
.unwrap();
assert!(!writer.write_packet(&mut buf).unwrap());
}
}
#[test]
fn test_command_writer_cannot_split() {
let mut buf = vec![];
let mut writer = CommandWriter::new(&(1, "tools echo aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa"), ConstId {}, false, true, MaxPacketSize::default()).unwrap();
assert!(writer.write_packet(&mut buf).unwrap());
let _: AsciiCommandSplitError = writer
.write_packet(&mut buf)
.unwrap_err()
.try_into()
.unwrap();
}
#[test]
fn test_max_packet_size() {
assert_eq!(MaxPacketSize::default().as_usize(), 80);
assert!(MaxPacketSize::new(79).is_none());
assert!(MaxPacketSize::new(80).is_some());
}
#[test]
fn ascii_char_count() {
let cases = [
(0, 1),
(1, 1),
(2, 1),
(3, 1),
(4, 1),
(5, 1),
(6, 1),
(7, 1),
(8, 1),
(9, 1),
(10, 2),
(99, 2),
(100, 3),
(999, 3),
(1000, 4),
(9999, 4),
(10000, 5),
(100000, 6),
(1000000, 7),
(10000000, 8),
(100000000, 9),
(usize::MAX, 20),
];
for (input, expected_count) in cases {
eprintln!("case {input}");
let actual_count = super::ascii_char_count(input);
assert_eq!(actual_count, expected_count);
}
}
}