use cbor2::Cbor;
use crate::{
header::{decode_protected, encode_protected, validate_header_buckets},
iana, tag, util, Error, Header, Label, Signer, Verifier,
};
#[derive(Clone, Debug, PartialEq, Cbor)]
#[cbor(array)]
struct SignatureWire {
#[serde(with = "crate::strict::bytes")]
protected: Vec<u8>,
unprotected: Header,
#[serde(with = "crate::strict::bytes")]
signature: Vec<u8>,
}
#[derive(Clone, Debug, PartialEq, Cbor)]
#[cbor(tag = 98, array)]
struct SignWire {
#[serde(with = "crate::strict::bytes")]
protected: Vec<u8>,
unprotected: Header,
#[serde(with = "crate::strict::optional_bytes")]
payload: Option<Vec<u8>>,
signatures: Vec<SignatureWire>,
}
struct SignatureRef<'a>(&'a Signature);
impl serde::Serialize for SignatureRef<'_> {
fn serialize<S: serde::Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
(
serde_bytes::Bytes::new(&self.0.protected_raw),
&self.0.unprotected,
serde_bytes::Bytes::new(&self.0.signature),
)
.serialize(serializer)
}
}
struct SignaturesRef<'a>(&'a [Signature]);
impl serde::Serialize for SignaturesRef<'_> {
fn serialize<S: serde::Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
use serde::ser::SerializeSeq;
let mut sequence = serializer.serialize_seq(Some(self.0.len()))?;
for signature in self.0 {
sequence.serialize_element(&SignatureRef(signature))?;
}
sequence.end()
}
}
#[derive(Clone, Debug, Default, PartialEq)]
pub struct Signature {
pub protected: Header,
pub unprotected: Header,
signature: Vec<u8>,
protected_raw: Vec<u8>,
state: util::OperationState,
}
impl Signature {
pub fn new() -> Self {
Self::default()
}
pub fn with_alg_kid(alg: Option<Label>, kid: Option<&[u8]>) -> Self {
let mut signature = Self::new();
if let Some(alg) = alg {
signature.protected.set_alg(alg);
}
if let Some(kid) = kid {
signature.unprotected.set_kid(kid.to_vec());
}
signature
}
fn kid(&self) -> Result<Option<&[u8]>, Error> {
match self.protected.kid()? {
Some(kid) => Ok(Some(kid)),
None => self.unprotected.kid(),
}
}
pub fn signature(&self) -> &[u8] {
&self.signature
}
pub fn protected_raw(&self) -> &[u8] {
&self.protected_raw
}
pub fn set_signature(&mut self, signature: impl Into<Vec<u8>>) -> Result<(), Error> {
validate_header_buckets(&self.protected, &self.unprotected)?;
if !self.state.initialized() {
self.protected_raw = encode_protected(&self.protected)?;
}
crate::header::validate_protected_state(&self.protected, &self.protected_raw)?;
self.signature = signature.into();
self.state = util::OperationState::Complete;
Ok(())
}
}
#[derive(Clone, Debug, Default, PartialEq)]
pub struct SignMessage {
pub protected: Header,
pub unprotected: Header,
pub payload: Option<Vec<u8>>,
pub signatures: Vec<Signature>,
protected_raw: Vec<u8>,
state: util::OperationState,
}
impl SignMessage {
pub fn new(payload: Option<Vec<u8>>) -> Self {
SignMessage {
payload,
..Default::default()
}
}
pub fn to_be_signed(
body_protected: &[u8],
sign_protected: &[u8],
external_aad: &[u8],
payload: &[u8],
) -> Result<Vec<u8>, Error> {
util::encode_structure(&(
"Signature",
serde_bytes::Bytes::new(body_protected),
serde_bytes::Bytes::new(sign_protected),
serde_bytes::Bytes::new(external_aad),
serde_bytes::Bytes::new(payload),
))
}
pub fn prepare_signatures(
&mut self,
signatures: Vec<Signature>,
external_aad: Option<&[u8]>,
) -> Result<Vec<Vec<u8>>, Error> {
self.prepare_signature_headers(signatures)?;
let payload =
util::require_embedded_payload(&self.payload, "SignMessage::prepare_signatures")?;
self.signature_inputs(payload, external_aad.unwrap_or(&[]))
}
pub fn prepare_detached_signatures(
&mut self,
signatures: Vec<Signature>,
detached_payload: &[u8],
external_aad: Option<&[u8]>,
) -> Result<Vec<Vec<u8>>, Error> {
self.prepare_signature_headers(signatures)?;
let to_be_signed = self.signature_inputs(detached_payload, external_aad.unwrap_or(&[]))?;
self.payload = None;
Ok(to_be_signed)
}
fn prepare_signature_headers(&mut self, mut signatures: Vec<Signature>) -> Result<(), Error> {
if signatures.is_empty() {
return Err(Error::Custom(
"SignMessage requires at least one signature".into(),
));
}
validate_header_buckets(&self.protected, &self.unprotected)?;
let protected_raw = encode_protected(&self.protected)?;
for signature in &mut signatures {
validate_header_buckets(&signature.protected, &signature.unprotected)?;
let sign_protected_raw = encode_protected(&signature.protected)?;
signature.protected_raw = sign_protected_raw;
signature.state = util::OperationState::Prepared;
signature.signature.clear();
}
self.protected_raw = protected_raw;
self.state = util::OperationState::Prepared;
self.signatures = signatures;
Ok(())
}
fn signature_inputs(&self, payload: &[u8], external_aad: &[u8]) -> Result<Vec<Vec<u8>>, Error> {
self.signatures
.iter()
.map(|signature| {
Self::to_be_signed(
&self.protected_raw,
&signature.protected_raw,
external_aad,
payload,
)
})
.collect()
}
pub fn set_signatures<I, S>(&mut self, signatures: I) -> Result<(), Error>
where
I: IntoIterator<Item = S>,
S: Into<Vec<u8>>,
{
let signatures = signatures.into_iter().map(Into::into).collect::<Vec<_>>();
if signatures.is_empty() {
return Err(Error::Custom(
"SignMessage requires at least one signature".into(),
));
}
if signatures.len() != self.signatures.len() {
return Err(Error::Custom(format!(
"signature count mismatch, message has {}, got {}",
self.signatures.len(),
signatures.len()
)));
}
validate_header_buckets(&self.protected, &self.unprotected)?;
if !self.state.initialized() {
self.protected_raw = encode_protected(&self.protected)?;
}
crate::header::validate_protected_state(&self.protected, &self.protected_raw)?;
for (slot, signature) in self.signatures.iter_mut().zip(signatures) {
slot.set_signature(signature)?;
}
self.state = util::OperationState::Complete;
Ok(())
}
pub fn sign(
&mut self,
signers: &[&dyn Signer],
external_aad: Option<&[u8]>,
) -> Result<(), Error> {
self.sign_with_payload(signers, None, external_aad.unwrap_or(&[]))
}
pub fn sign_detached(
&mut self,
signers: &[&dyn Signer],
detached_payload: &[u8],
external_aad: Option<&[u8]>,
) -> Result<(), Error> {
self.sign_with_payload(signers, Some(detached_payload), external_aad.unwrap_or(&[]))?;
self.payload = None;
Ok(())
}
fn sign_with_payload(
&mut self,
signers: &[&dyn Signer],
detached_payload: Option<&[u8]>,
external_aad: &[u8],
) -> Result<(), Error> {
if signers.is_empty() {
return Err(Error::Custom(
"SignMessage requires at least one signer".into(),
));
}
let signature_headers = signers
.iter()
.map(|signer| Signature::with_alg_kid(signer.alg(), signer.kid()))
.collect::<Vec<_>>();
self.prepare_signature_headers(signature_headers)?;
let payload = match detached_payload {
Some(payload) => payload,
None => util::require_embedded_payload(&self.payload, "SignMessage::sign")?,
};
let signatures = signers
.iter()
.zip(&self.signatures)
.map(|(signer, signature)| {
let tbs = Self::to_be_signed(
&self.protected_raw,
&signature.protected_raw,
external_aad,
payload,
)?;
signer.sign(&tbs)
})
.collect::<Result<Vec<_>, _>>()?;
self.set_signatures(signatures)
}
pub fn sign_and_encode(
&mut self,
signers: &[&dyn Signer],
external_aad: Option<&[u8]>,
) -> Result<Vec<u8>, Error> {
self.sign(signers, external_aad)?;
self.to_vec()
}
pub fn sign_detached_and_encode(
&mut self,
signers: &[&dyn Signer],
detached_payload: &[u8],
external_aad: Option<&[u8]>,
) -> Result<Vec<u8>, Error> {
self.sign_detached(signers, detached_payload, external_aad)?;
self.to_vec()
}
pub fn to_vec(&self) -> Result<Vec<u8>, Error> {
self.encode(tag::SIGN_PREFIX)
}
pub fn to_cwt_vec(&self) -> Result<Vec<u8>, Error> {
self.encode(tag::CWT_SIGN_PREFIX)
}
pub fn to_untagged_vec(&self) -> Result<Vec<u8>, Error> {
self.encode(&[])
}
fn encode(&self, prefix: &[u8]) -> Result<Vec<u8>, Error> {
if !self.state.complete() {
return Err(Error::InvalidState(
"SignMessage must be signed before encoding".into(),
));
}
if self.signatures.is_empty() {
return Err(Error::Custom("SignMessage has no signatures".into()));
}
validate_header_buckets(&self.protected, &self.unprotected)?;
crate::header::validate_protected_state(&self.protected, &self.protected_raw)?;
for sig in &self.signatures {
validate_header_buckets(&sig.protected, &sig.unprotected)?;
crate::header::validate_protected_state(&sig.protected, &sig.protected_raw)?;
}
let unprotected = util::canonical_raw(&self.unprotected)?;
let signatures = util::canonical_raw(&SignaturesRef(&self.signatures))?;
util::encode_prefixed(
prefix,
&(
serde_bytes::Bytes::new(&self.protected_raw),
&unprotected,
self.payload.as_deref().map(serde_bytes::Bytes::new),
&signatures,
),
)
}
pub fn from_slice(data: &[u8]) -> Result<Self, Error> {
let body = tag::message_body(data, Self::TAG)?;
let wire: SignWire = cbor2::from_slice(body)?;
if wire.signatures.is_empty() {
return Err(Error::Custom("SignMessage has no signatures".into()));
}
let protected = decode_protected(&wire.protected)?;
validate_header_buckets(&protected, &wire.unprotected)?;
let mut signatures = Vec::with_capacity(wire.signatures.len());
for sw in wire.signatures {
let sig_protected = decode_protected(&sw.protected)?;
validate_header_buckets(&sig_protected, &sw.unprotected)?;
signatures.push(Signature {
protected: sig_protected,
unprotected: sw.unprotected,
signature: sw.signature,
protected_raw: sw.protected,
state: util::OperationState::Complete,
});
}
Ok(SignMessage {
protected,
unprotected: wire.unprotected,
payload: wire.payload,
signatures,
protected_raw: wire.protected,
state: util::OperationState::Complete,
})
}
pub fn verify(
&self,
verifiers: &[&dyn Verifier],
external_aad: Option<&[u8]>,
) -> Result<(), Error> {
let payload = util::require_embedded_payload(&self.payload, "SignMessage::verify")?;
self.verify_payload(verifiers, payload, external_aad.unwrap_or(&[]))
}
pub fn verify_detached(
&self,
verifiers: &[&dyn Verifier],
detached_payload: &[u8],
external_aad: Option<&[u8]>,
) -> Result<(), Error> {
if self.payload.is_some() {
return Err(Error::Custom(
"SignMessage carries an embedded payload; use verify".into(),
));
}
self.verify_payload(verifiers, detached_payload, external_aad.unwrap_or(&[]))
}
fn verify_payload(
&self,
verifiers: &[&dyn Verifier],
payload: &[u8],
external_aad: &[u8],
) -> Result<(), Error> {
if !self.state.complete() {
return Err(Error::InvalidState(
"SignMessage must be decoded before verifying".into(),
));
}
if verifiers.is_empty() {
return Err(Error::Custom(
"SignMessage requires at least one verifier".into(),
));
}
if self.signatures.is_empty() {
return Err(Error::Custom("SignMessage has no signatures".into()));
}
crate::header::validate_protected_state(&self.protected, &self.protected_raw)?;
for sig in &self.signatures {
crate::header::validate_protected_state(&sig.protected, &sig.protected_raw)?;
let kid = sig.kid()?;
let tbs = Self::to_be_signed(
&self.protected_raw,
&sig.protected_raw,
external_aad,
payload,
)?;
let mut matched_kid = false;
let mut last_error = None;
let mut verified = false;
'candidates: for rank in 0..=1 {
for verifier in verifiers
.iter()
.filter(|verifier| util::kid_match_rank(kid, verifier.kid()) == Some(rank))
{
matched_kid = true;
if let Err(err) = self
.protected
.ensure_crit_understood(verifier.understood_critical_headers())
.and_then(|_| {
sig.protected
.ensure_crit_understood(verifier.understood_critical_headers())
})
{
last_error = Some(err);
continue;
}
if let Err(err) =
util::check_protected_alg(&sig.protected, &sig.unprotected, verifier.alg())
{
last_error = Some(err);
continue;
}
match verifier.verify(&tbs, &sig.signature) {
Ok(()) => {
verified = true;
break 'candidates;
}
Err(err) => last_error = Some(err),
}
}
}
if !matched_kid {
return Err(Error::verify("no verifier for signature kid"));
}
if !verified {
return Err(last_error.unwrap_or_else(|| Error::verify("signature mismatch")));
}
}
Ok(())
}
pub fn verify_and_decode(
verifiers: &[&dyn Verifier],
data: &[u8],
external_aad: Option<&[u8]>,
) -> Result<Self, Error> {
let msg = Self::from_slice(data)?;
msg.verify(verifiers, external_aad)?;
Ok(msg)
}
pub fn verify_detached_and_decode(
verifiers: &[&dyn Verifier],
data: &[u8],
detached_payload: &[u8],
external_aad: Option<&[u8]>,
) -> Result<Self, Error> {
let msg = Self::from_slice(data)?;
msg.verify_detached(verifiers, detached_payload, external_aad)?;
Ok(msg)
}
pub fn protected_raw(&self) -> &[u8] {
&self.protected_raw
}
pub const TAG: u64 = iana::CBORTagCOSESign;
}
#[cfg(test)]
mod tests {
use super::*;
fn assert_cbor_shape<T: cbor2::Cbor>(tag: Option<u64>, array: bool) {
assert_eq!(T::TAG, tag);
assert_eq!(T::ARRAY, array);
}
#[test]
fn wire_metadata_declares_tagged_array_shape() {
assert_cbor_shape::<SignatureWire>(None, true);
assert_cbor_shape::<SignWire>(Some(iana::CBORTagCOSESign), true);
}
}