pub const ENVELOPE_MAGIC: &[u8; 4] = b"RKST";
pub const ENVELOPE_VERSION: u16 = 1;
#[repr(u8)]
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum FormatId {
Unknown = 0,
Clap = 1,
Vst3 = 2,
AuV2 = 3,
AuV3 = 4,
Vst2 = 5,
Lv2 = 6,
Aax = 7,
}
#[derive(Debug, thiserror::Error)]
pub enum StateLoadError {
#[error("state blob too short: expected at least {expected} bytes, got {actual}")]
Truncated {
expected: usize,
actual: usize,
},
#[error("state magic mismatch")]
BadMagic,
#[error("state envelope version {found} > supported {supported}")]
UnsupportedVersion {
found: u16,
supported: u16,
},
#[error("state payload length {declared} != trailing bytes {actual}")]
LengthMismatch {
declared: u32,
actual: usize,
},
#[error("state format mismatch: payload is {found:?}, host expected {expected:?}")]
FormatMismatch {
found: FormatId,
expected: FormatId,
},
}
#[derive(Debug, Clone)]
pub struct StateEnvelope<'a> {
pub format: FormatId,
pub payload: &'a [u8],
}
const HEADER_LEN: usize = 4 + 2 + 1 + 1 + 4;
impl StateEnvelope<'_> {
#[must_use]
pub fn encode(&self) -> Vec<u8> {
let mut out = Vec::with_capacity(HEADER_LEN + self.payload.len());
out.extend_from_slice(ENVELOPE_MAGIC);
out.extend_from_slice(&ENVELOPE_VERSION.to_le_bytes());
out.push(self.format as u8);
out.push(0);
#[allow(clippy::cast_possible_truncation)]
let payload_len = self.payload.len() as u32;
out.extend_from_slice(&payload_len.to_le_bytes());
out.extend_from_slice(self.payload);
out
}
pub fn decode(bytes: &[u8]) -> Result<StateEnvelope<'_>, StateLoadError> {
if bytes.len() < HEADER_LEN {
return Err(StateLoadError::Truncated {
expected: HEADER_LEN,
actual: bytes.len(),
});
}
if &bytes[0..4] != ENVELOPE_MAGIC {
return Err(StateLoadError::BadMagic);
}
let version = u16::from_le_bytes([bytes[4], bytes[5]]);
if version > ENVELOPE_VERSION {
return Err(StateLoadError::UnsupportedVersion {
found: version,
supported: ENVELOPE_VERSION,
});
}
let format = match bytes[6] {
1 => FormatId::Clap,
2 => FormatId::Vst3,
3 => FormatId::AuV2,
4 => FormatId::AuV3,
5 => FormatId::Vst2,
6 => FormatId::Lv2,
7 => FormatId::Aax,
_ => FormatId::Unknown,
};
let declared_len = u32::from_le_bytes([bytes[8], bytes[9], bytes[10], bytes[11]]);
let payload = &bytes[HEADER_LEN..];
if declared_len as usize != payload.len() {
return Err(StateLoadError::LengthMismatch {
declared: declared_len,
actual: payload.len(),
});
}
Ok(StateEnvelope { format, payload })
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn round_trip() {
let payload = b"opaque plugin state bytes";
let env = StateEnvelope {
format: FormatId::Clap,
payload,
};
let encoded = env.encode();
let decoded = StateEnvelope::decode(&encoded).expect("decode");
assert_eq!(decoded.format, FormatId::Clap);
assert_eq!(decoded.payload, payload);
}
#[test]
fn truncated_header_rejected() {
let short = b"RK";
let err = StateEnvelope::decode(short).unwrap_err();
assert!(matches!(err, StateLoadError::Truncated { .. }));
}
#[test]
fn bad_magic_rejected() {
let mut buf = b"XXXX".to_vec();
buf.extend(std::iter::repeat_n(0u8, HEADER_LEN));
let err = StateEnvelope::decode(&buf).unwrap_err();
assert!(matches!(err, StateLoadError::BadMagic));
}
#[test]
fn length_mismatch_rejected() {
let env = StateEnvelope {
format: FormatId::Vst3,
payload: b"abcd",
};
let mut buf = env.encode();
buf.push(0xFF); let err = StateEnvelope::decode(&buf).unwrap_err();
assert!(matches!(err, StateLoadError::LengthMismatch { .. }));
}
}