use std::fmt;
use sha2::{Digest, Sha384};
use crate::cbor::{self, Value};
use crate::node_key::{carried_key_well_formed, verify, KeyError, NodeKey};
use crate::profile::Profile;
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum ObjectError {
Malformed,
SignatureInvalid,
AlgMismatch,
FieldKeyNotText,
DuplicateField(String),
Key(KeyError),
}
impl fmt::Display for ObjectError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
ObjectError::Malformed => f.write_str("malformed signed object"),
ObjectError::SignatureInvalid => {
f.write_str("the signed object's signature does not verify")
}
ObjectError::AlgMismatch => {
f.write_str("the signed object names another profile's algorithm")
}
ObjectError::FieldKeyNotText => f.write_str("a field key that is not text"),
ObjectError::DuplicateField(name) => write!(f, "two fields named {name:?}"),
ObjectError::Key(e) => write!(f, "{e}"),
}
}
}
impl std::error::Error for ObjectError {}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Object {
pub key: Vec<u8>,
pub tbs: Vec<u8>,
pub signature: Vec<u8>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct HeldObject {
pub tbs: Vec<u8>,
pub signature: Vec<u8>,
}
#[derive(Debug, Clone, PartialEq)]
pub struct VerifiedObject {
pub key: Vec<u8>,
pub tbs: Vec<u8>,
pub fields: Value,
}
impl Object {
pub fn to_value(&self) -> Value {
Value::Map(vec![
(Value::text("key"), Value::Bytes(self.key.clone())),
(Value::text("tbs"), Value::Bytes(self.tbs.clone())),
(
Value::text("signature"),
Value::Bytes(self.signature.clone()),
),
])
}
pub fn from_value(value: &Value) -> Result<Object, ObjectError> {
let [key, tbs, signature] = exact_byte_fields(value, ["key", "tbs", "signature"])?;
Ok(Object {
key,
tbs,
signature,
})
}
}
impl HeldObject {
pub fn to_value(&self) -> Value {
Value::Map(vec![
(Value::text("tbs"), Value::Bytes(self.tbs.clone())),
(
Value::text("signature"),
Value::Bytes(self.signature.clone()),
),
])
}
pub fn from_value(value: &Value) -> Result<HeldObject, ObjectError> {
let [tbs, signature] = exact_byte_fields(value, ["tbs", "signature"])?;
Ok(HeldObject { tbs, signature })
}
}
pub fn sign_object(
label: &str,
fields: &[(Value, Value)],
key: &NodeKey,
) -> Result<Object, ObjectError> {
let carried = key.public_key();
let tbs = object_tbs(fields, key.profile())?;
let signature = key
.sign(&object_signed_bytes(label, &carried, &tbs))
.map_err(ObjectError::Key)?;
Ok(Object {
key: carried,
tbs,
signature,
})
}
pub fn sign_held_object(
label: &str,
fields: &[(Value, Value)],
key: &NodeKey,
) -> Result<HeldObject, ObjectError> {
let object = sign_object(label, fields, key)?;
Ok(HeldObject {
tbs: object.tbs,
signature: object.signature,
})
}
pub fn verify_object(
label: &str,
value: &Value,
profile: Profile,
) -> Result<VerifiedObject, ObjectError> {
let object = Object::from_value(value)?;
if !carried_key_well_formed(&object.key, profile) {
return Err(ObjectError::Malformed);
}
verified(label, object.key, object.tbs, &object.signature, profile)
}
pub fn verify_held_object(
label: &str,
value: &Value,
key: &[u8],
profile: Profile,
) -> Result<VerifiedObject, ObjectError> {
let held = HeldObject::from_value(value)?;
verified(label, key.to_vec(), held.tbs, &held.signature, profile)
}
fn verified(
label: &str,
key: Vec<u8>,
tbs: Vec<u8>,
signature: &[u8],
profile: Profile,
) -> Result<VerifiedObject, ObjectError> {
if !verify(
&object_signed_bytes(label, &key, &tbs),
signature,
&key,
profile,
) {
return Err(ObjectError::SignatureInvalid);
}
let fields = cbor::decode(&tbs).map_err(|_| ObjectError::Malformed)?;
let alg = match (&fields, fields.get("alg")) {
(Value::Map(_), Some(Value::Text(alg))) => alg.clone(),
_ => return Err(ObjectError::Malformed),
};
if alg != profile.sig_alg() {
return Err(ObjectError::AlgMismatch);
}
Ok(VerifiedObject { key, tbs, fields })
}
fn object_tbs(fields: &[(Value, Value)], profile: Profile) -> Result<Vec<u8>, ObjectError> {
let mut seen = std::collections::HashSet::with_capacity(fields.len());
let mut with_alg = Vec::with_capacity(fields.len() + 1);
for (key, value) in fields {
let Value::Text(name) = key else {
return Err(ObjectError::FieldKeyNotText);
};
if seen.contains(name.as_str()) {
return Err(ObjectError::DuplicateField(name.clone()));
}
if name == "alg" {
continue;
}
seen.insert(name.as_str());
with_alg.push((key.clone(), value.clone()));
}
with_alg.push((Value::text("alg"), Value::text(profile.sig_alg())));
cbor::encode(&Value::Map(with_alg)).map_err(|_| ObjectError::Malformed)
}
fn object_signed_bytes(label: &str, key: &[u8], tbs: &[u8]) -> Vec<u8> {
let mut out = Vec::with_capacity(label.len() + 1 + 48 + tbs.len());
out.extend_from_slice(label.as_bytes());
out.push(0);
out.extend_from_slice(&Sha384::digest(key));
out.extend_from_slice(tbs);
out
}
fn exact_byte_fields<const N: usize>(
value: &Value,
names: [&str; N],
) -> Result<[Vec<u8>; N], ObjectError> {
let Value::Map(pairs) = value else {
return Err(ObjectError::Malformed);
};
if pairs.len() != N {
return Err(ObjectError::Malformed);
}
let mut out: [Vec<u8>; N] = std::array::from_fn(|_| Vec::new());
for (slot, name) in out.iter_mut().zip(names) {
match value.get(name) {
Some(Value::Bytes(b)) => *slot = b.clone(),
_ => return Err(ObjectError::Malformed),
}
}
Ok(out)
}