use std::io::Cursor;
use std::io::Read;
use std::sync::Arc;
use bitvec::field::BitField;
use bitvec::order::Msb0;
use bitvec::view::BitView;
#[cfg(feature = "search")]
use memmem::Searcher;
#[cfg(feature = "search")]
use memmem::TwoWaySearcher;
#[cfg(feature = "tracing")]
use tracing::debug;
#[cfg(feature = "tracing")]
use tracing::info;
#[cfg(feature = "tracing")]
use tracing::trace;
use crate::ErrorKind;
use crate::klv::Klv;
use crate::klv_value::KlvValue;
use crate::tag::Tag;
pub const UAS_LOCAL_SET_UNIVERSAL_LABEL: [u8; 16] = [
0x06, 0x0E, 0x2B, 0x34, 0x02, 0x0B, 0x01, 0x01, 0x0E, 0x01, 0x03, 0x01,
0x01, 0x00, 0x00, 0x00,
];
#[derive(Clone, Debug)]
pub struct KlvPacket {
fields: Vec<Klv>,
}
impl KlvPacket {
fn get_tag(buf: &mut Cursor<&[u8]>) -> usize {
Self::get_ber_value(buf)
}
fn get_length(buf: &mut Cursor<&[u8]>) -> usize {
Self::get_ber_value(buf)
}
fn get_ber_value(buf: &mut Cursor<&[u8]>) -> usize {
let mut new_byte: [u8; 1] = [0];
buf.read_exact(&mut new_byte).expect("Can't read from bytes");
let bits = new_byte.view_bits::<Msb0>();
let msb =
bits.get(0).expect("Cannot get the first bit from the byte array");
if msb == true {
let Some(remainder) = bits.get(1..bits.len()) else {
panic!("Cannot get bits after first for BER byte")
};
let long_length = remainder.load_be();
let mut len_bytes = vec![0; long_length];
buf.read_exact(&mut len_bytes).expect("Can't read from bytes");
return len_bytes.view_bits::<Msb0>().load_be();
}
bits.load_be()
}
fn get_value(
buf: &mut Cursor<&[u8]>,
tag: usize,
length: usize,
) -> Result<Klv, ErrorKind> {
let mut value_buf = vec![0; length];
buf.read_exact(&mut value_buf)
.expect("Couldn't read all value data necessary from buffer");
Klv::new(tag, value_buf.into())
}
fn calculate_checksum(buf: &[u8]) -> u16 {
let mut bcc: u16 = 0;
buf[..buf.len()].iter().enumerate().for_each(|(idx, byte)| {
bcc =
bcc.wrapping_add((*byte as u16) << (8 * ((idx + 1) % 2)) as u16)
});
bcc
}
pub fn from_bytes(bytes: &[u8]) -> Result<Option<KlvPacket>, ErrorKind> {
let start_index: usize;
#[cfg(feature = "search")]
{
let search = TwoWaySearcher::new(&UAS_LOCAL_SET_UNIVERSAL_LABEL);
start_index = match search.search_in(&bytes) {
Some(idx) => idx,
None => return Ok(None),
};
}
#[cfg(not(feature = "search"))]
{
let test_bytes = &bytes[0..UAS_LOCAL_SET_UNIVERSAL_LABEL.len()];
if !test_bytes.iter().eq(UAS_LOCAL_SET_UNIVERSAL_LABEL.iter()) {
return Ok(None);
}
start_index = 0;
}
#[cfg(feature = "tracing")]
{
trace!("Parsing KLV packet: {:02X?}", bytes);
trace!("Start index [{}]", start_index);
}
let length_position = start_index + UAS_LOCAL_SET_UNIVERSAL_LABEL.len();
#[cfg(feature = "tracing")]
trace!("Length position [{}]", length_position);
let mut buffer = Cursor::new(bytes);
buffer.set_position(length_position as u64);
let klv_length = Self::get_length(&mut buffer);
#[cfg(feature = "tracing")]
trace!("Length of packet [{}]", klv_length);
let length_field_length = buffer.position() as usize - length_position;
#[cfg(feature = "tracing")]
trace!("Length field length [{}]", length_field_length);
let klv_packet_end = length_position + length_field_length + klv_length;
#[cfg(feature = "tracing")]
trace!("KLV packet end [{}]", klv_packet_end);
let max_tag_id = Tag::COUNT;
let mut fields = Vec::new();
while buffer.position() < klv_packet_end as u64 {
let tag = Self::get_tag(&mut buffer);
if tag > max_tag_id {
return Err(ErrorKind::UnsupportedTag(tag));
}
let length = Self::get_length(&mut buffer);
if length == 0 {
#[cfg(feature = "tracing")]
debug!("Length of tag [{}] is 0", tag);
continue;
}
let value = Self::get_value(&mut buffer, tag, length)?;
#[cfg(feature = "tracing")]
trace!(
"Added tag to KLV packet: [{}]",
Into::<&'static str>::into(value.tag())
);
fields.push(value);
}
let packet = KlvPacket { fields };
let packet_checksum = packet.checksum();
let checksum_bytes_length = klv_packet_end - 2;
let calculated_checksum = Self::calculate_checksum(
bytes.get(start_index..checksum_bytes_length).unwrap(),
);
if packet_checksum != calculated_checksum {
#[cfg(feature = "tracing")]
debug!(
"Checksum for packet [{}] vs calculated checksum [{}]",
packet_checksum, calculated_checksum
);
return Err(ErrorKind::InvalidChecksum);
}
#[cfg(feature = "tracing")]
debug!(
"Found valid KLV packet with timestamp {}",
packet.precision_time_stamp()
);
Ok(Some(packet))
}
pub fn get_id(&self, tag: usize) -> Option<Klv> {
self.fields
.iter()
.find(|field_tag| tag == field_tag.tag().into())
.cloned()
}
pub fn get(&self, tag: Tag) -> Option<Klv> {
self.get_id(tag.into())
}
pub fn checksum(&self) -> u16 {
match self
.get(Tag::Checksum)
.expect("KLV packets must have a checksum")
.value()
{
KlvValue::Uint16(value) => *value,
_ => panic!(
"This packet does not have a checksum and that error was not caught. This should be unreachable"
),
}
}
pub fn precision_time_stamp(&self) -> u64 {
match self
.get(Tag::PrecisionTimeStamp)
.expect("KLV packets must have a precision time stamp")
.value()
{
KlvValue::Uint64(value) => *value,
_ => panic!(
"This packet does not have a checksum and that error was not caught. This should be unreachable"
),
}
}
pub fn mission_id(&self) -> Option<Arc<str>> {
match self.get(Tag::MissionID)?.value() {
KlvValue::Utf8(value) => Some(value.clone()),
_ => panic!(
"This packet does not have a mission ID and that error was not caught. This should be unreachable"
),
}
}
}
#[cfg(test)]
mod tests {
use std::io::Cursor;
use std::sync::Arc;
use std::vec;
use itertools::chain;
use test_case::test_case;
use super::KlvPacket;
use super::UAS_LOCAL_SET_UNIVERSAL_LABEL;
fn packet_from_value(test_value: Vec<u8>) -> Vec<u8> {
let precision_timestamp_bytes =
vec![0x02, 0x08, 0x00, 0x11, 0x22, 0x33, 0x44, 0x55, 0x66, 0x77];
let checksum_header = vec![0x01, 0x02];
let length: Vec<u8> = vec![
(precision_timestamp_bytes.len() + test_value.len() + 4)
.try_into()
.expect("Length does not fit in 1 byte"),
];
let packet_minus_checksum: Box<[u8]> = chain!(
Vec::from(UAS_LOCAL_SET_UNIVERSAL_LABEL),
length,
precision_timestamp_bytes,
test_value,
checksum_header
)
.collect();
let checksum = KlvPacket::calculate_checksum(&packet_minus_checksum);
chain!(
Vec::from(packet_minus_checksum),
Vec::from(checksum.to_be_bytes())
)
.collect()
}
fn packet_1() -> Vec<u8> {
packet_from_value(vec![0x03, 0x02, b'I', b'D']) }
#[test_case(packet_1, 47467, Some("ID".into()))]
fn from_bytes(
packet: fn() -> Vec<u8>,
checksum: u16,
mission_id: Option<Arc<str>>,
) {
let bytes = packet();
let packet = KlvPacket::from_bytes(bytes.into()).unwrap().unwrap();
assert_eq!(packet.checksum(), checksum, "Checksum is incorrect");
assert_eq!(
packet.precision_time_stamp(),
4822678189205111,
"Precision Time Stamp is incorrect"
);
assert_eq!(packet.mission_id(), mission_id)
}
#[test_case(&[0x71, 0xF1, 0x00], 113; "Short form")]
#[test_case(&[0x81, 0xF1, 0x00], 241; "Long-Form: One byte")]
#[test_case(&[0x83, 0xF1, 0xFF, 0xF1], 15859697; "Long-Form Three bytes")]
fn get_ber_value(bytes: &[u8], correct_length: usize) {
let mut test_bytes = Cursor::new(bytes.clone());
let length = KlvPacket::get_ber_value(&mut test_bytes);
assert_eq!(length, correct_length, "Failed to BER")
}
}