use embedded_io::Error as _;
use heapless::Vec as HVec;
use crate::protocol::{self, MessageId, sd};
use crate::traits::{PayloadWireFormat, WireFormat};
pub const ENTRY_CAP: usize = 8;
pub const OPT_CAP: usize = 8;
pub const PAYLOAD_CAP: usize = 256;
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct HeaplessSdHeader {
pub flags: sd::Flags,
pub entries: HVec<sd::Entry, ENTRY_CAP>,
pub options: HVec<sd::Options, OPT_CAP>,
}
impl WireFormat for HeaplessSdHeader {
fn required_size(&self) -> usize {
sd::Header::new(self.flags, &self.entries, &self.options).required_size()
}
fn encode<T: embedded_io::Write>(&self, writer: &mut T) -> Result<usize, protocol::Error> {
sd::Header::new(self.flags, &self.entries, &self.options).encode(writer)
}
}
#[allow(clippy::large_enum_variant)]
#[derive(Clone, Debug, Eq, PartialEq)]
enum HeaplessPayloadKind {
Sd(HeaplessSdHeader),
Raw(HVec<u8, PAYLOAD_CAP>),
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct HeaplessPayload {
message_id: MessageId,
kind: HeaplessPayloadKind,
}
impl HeaplessPayload {
#[must_use]
pub fn raw_bytes(&self) -> Option<&[u8]> {
match &self.kind {
HeaplessPayloadKind::Raw(bytes) => Some(bytes),
HeaplessPayloadKind::Sd(_) => None,
}
}
}
impl PayloadWireFormat for HeaplessPayload {
type SdHeader = HeaplessSdHeader;
fn message_id(&self) -> MessageId {
self.message_id
}
fn as_sd_header(&self) -> Option<&HeaplessSdHeader> {
match &self.kind {
HeaplessPayloadKind::Sd(header) => Some(header),
HeaplessPayloadKind::Raw(_) => None,
}
}
fn from_payload_bytes(message_id: MessageId, payload: &[u8]) -> Result<Self, protocol::Error> {
if message_id == MessageId::SD {
let view = sd::SdHeaderView::parse(payload)?;
let mut entries: HVec<sd::Entry, ENTRY_CAP> = HVec::new();
for ev in view.entries() {
let entry = ev.to_owned()?;
entries
.push(entry)
.map_err(|_| protocol::Error::Io(embedded_io::ErrorKind::OutOfMemory))?;
}
let mut options: HVec<sd::Options, OPT_CAP> = HVec::new();
for ov in view.options() {
let opt = ov.to_owned()?;
options
.push(opt)
.map_err(|_| protocol::Error::Io(embedded_io::ErrorKind::OutOfMemory))?;
}
Ok(Self {
message_id,
kind: HeaplessPayloadKind::Sd(HeaplessSdHeader {
flags: view.flags(),
entries,
options,
}),
})
} else {
let mut bytes: HVec<u8, PAYLOAD_CAP> = HVec::new();
bytes
.extend_from_slice(payload)
.map_err(|_| protocol::Error::Io(embedded_io::ErrorKind::OutOfMemory))?;
Ok(Self {
message_id,
kind: HeaplessPayloadKind::Raw(bytes),
})
}
}
fn new_sd_payload(header: &HeaplessSdHeader) -> Self {
Self {
message_id: MessageId::SD,
kind: HeaplessPayloadKind::Sd(header.clone()),
}
}
fn sd_flags(&self) -> Option<sd::Flags> {
match &self.kind {
HeaplessPayloadKind::Sd(header) => Some(header.flags),
HeaplessPayloadKind::Raw(_) => None,
}
}
fn required_size(&self) -> usize {
match &self.kind {
HeaplessPayloadKind::Sd(header) => header.required_size(),
HeaplessPayloadKind::Raw(bytes) => bytes.len(),
}
}
fn encode<T: embedded_io::Write>(&self, writer: &mut T) -> Result<usize, protocol::Error> {
match &self.kind {
HeaplessPayloadKind::Sd(header) => header.encode(writer),
HeaplessPayloadKind::Raw(bytes) => {
writer
.write_all(bytes)
.map_err(|e| protocol::Error::Io(e.kind()))?;
Ok(bytes.len())
}
}
}
fn new_subscription_sd_header(
service_id: u16,
instance_id: u16,
major_version: u8,
ttl: u32,
event_group_id: u16,
client_ip: core::net::Ipv4Addr,
protocol: sd::TransportProtocol,
client_port: u16,
reboot_flag: sd::RebootFlag,
) -> HeaplessSdHeader {
let entry = sd::Entry::SubscribeEventGroup(sd::EventGroupEntry::new(
service_id,
instance_id,
major_version,
ttl,
event_group_id,
));
let endpoint = sd::Options::IpV4Endpoint {
ip: client_ip,
protocol,
port: client_port,
};
let mut entries: HVec<sd::Entry, ENTRY_CAP> = HVec::new();
let _ = entries.push(entry); let mut options: HVec<sd::Options, OPT_CAP> = HVec::new();
let _ = options.push(endpoint);
HeaplessSdHeader {
flags: sd::Flags::new_sd(reboot_flag),
entries,
options,
}
}
fn set_reboot_flag(header: &mut HeaplessSdHeader, reboot: sd::RebootFlag) {
header.flags = sd::Flags::new(bool::from(reboot), header.flags.unicast());
}
fn for_each_offered_endpoint<F>(&self, mut f: F)
where
F: FnMut(crate::OfferedEndpoint),
{
let header = match &self.kind {
HeaplessPayloadKind::Sd(header) => header,
HeaplessPayloadKind::Raw(_) => return,
};
for entry in &header.entries {
if let sd::Entry::OfferService(svc) | sd::Entry::StopOfferService(svc) = entry {
let is_offer = matches!(entry, sd::Entry::OfferService(_));
let endpoint =
sd::extract_ipv4_endpoint(&header.options).map(|(addr, protocol)| {
crate::NetEndpoint::new(core::net::SocketAddr::V4(addr), protocol)
});
f(crate::OfferedEndpoint {
service_id: svc.service_id,
instance_id: svc.instance_id,
major_version: svc.major_version,
minor_version: svc.minor_version,
endpoint,
is_offer,
});
}
}
}
fn for_each_service_instance<F>(&self, mut f: F)
where
F: FnMut(u16, u16),
{
let header = match &self.kind {
HeaplessPayloadKind::Sd(header) => header,
HeaplessPayloadKind::Raw(_) => return,
};
for entry in &header.entries {
let (svc, inst) = match entry {
sd::Entry::FindService(svc)
| sd::Entry::OfferService(svc)
| sd::Entry::StopOfferService(svc) => (svc.service_id, svc.instance_id),
sd::Entry::SubscribeEventGroup(eg) | sd::Entry::SubscribeAckEventGroup(eg) => {
(eg.service_id, eg.instance_id)
}
};
f(svc, inst);
}
}
}