use multi_codec::Codec;
use multi_hash::{Multihash, mh};
use multi_key::{
AttrId, AttrView, Builder, ConvView, DataView, Error, FingerprintView, Multikey, SignView,
VerifyView, ViewBuilder,
};
use multi_sig::Multisig;
use multi_trait::Null;
use multi_util::CodecInfo;
use ssh_key::{PrivateKey, PublicKey};
use zeroize::Zeroizing;
fn key_type_byte(mk: &Multikey) -> u8 {
mk.attributes
.get(&AttrId::KeyType)
.and_then(|t| t.first().copied())
.unwrap_or(0)
}
fn algorithm_name(mk: &Multikey) -> Result<&str, Error> {
let bytes = mk
.attributes
.get(&AttrId::AlgorithmName)
.ok_or_else(|| Error::UnsupportedAlgorithm("AlgorithmName missing".into()))?;
std::str::from_utf8(bytes)
.map_err(|_| Error::UnsupportedAlgorithm("AlgorithmName not UTF-8".into()))
}
struct CustomAttrs {
key_type: u8,
}
impl AttrView for CustomAttrs {
fn is_encrypted(&self) -> bool {
false
}
fn is_public_key(&self) -> bool {
self.key_type == 0
}
fn is_secret_key(&self) -> bool {
self.key_type == 1
}
fn is_secret_key_share(&self) -> bool {
false
}
}
struct CustomData {
key_bytes: Vec<u8>,
}
impl DataView for CustomData {
fn key_bytes(&self) -> Result<Zeroizing<Vec<u8>>, Error> {
Ok(Zeroizing::new(self.key_bytes.clone()))
}
fn secret_bytes(&self) -> Result<Zeroizing<Vec<u8>>, Error> {
Ok(Zeroizing::new(self.key_bytes.clone()))
}
}
struct CustomConv {
key_type: u8,
key: Multikey,
}
impl ConvView for CustomConv {
fn to_public_key(&self) -> Result<Multikey, Error> {
if self.key_type != 1 {
return Err(Error::UnsupportedAlgorithm(
"custom public key has no public form".into(),
));
}
let mut pk = self.key.clone();
pk.attributes
.insert(AttrId::KeyType, Zeroizing::new(vec![0]));
Ok(pk)
}
fn to_ssh_public_key(&self) -> Result<PublicKey, Error> {
Err(Error::UnsupportedAlgorithm(
"custom keys have no ssh form".into(),
))
}
fn to_ssh_private_key(&self) -> Result<PrivateKey, Error> {
Err(Error::UnsupportedAlgorithm(
"custom keys have no ssh form".into(),
))
}
}
struct CustomFingerprint {
key_bytes: Vec<u8>,
}
impl FingerprintView for CustomFingerprint {
fn fingerprint(&self, hash: Codec) -> Result<Multihash, Error> {
Ok(mh::Builder::new_from_bytes(hash, &self.key_bytes)?.try_build()?)
}
}
struct NamedSign {
name: &'static str,
}
impl SignView for NamedSign {
fn sign(&self, _: &[u8], _: bool, _: Option<u8>) -> Result<Multisig, Error> {
Err(Error::UnsupportedAlgorithm(format!("{} sign", self.name)))
}
}
struct NamedVerify {
name: String,
}
impl VerifyView for NamedVerify {
fn verify(&self, _: &Multisig, _: Option<&[u8]>) -> Result<(), Error> {
Err(Error::UnsupportedAlgorithm(format!("{} verify", self.name)))
}
}
fn attr_factory<'a>(mk: &'a Multikey) -> Result<Box<dyn AttrView + 'a>, Error> {
algorithm_name(mk)?;
Ok(Box::new(CustomAttrs {
key_type: key_type_byte(mk),
}))
}
fn data_factory<'a>(mk: &'a Multikey) -> Result<Box<dyn DataView + 'a>, Error> {
let key_bytes = mk
.attributes
.get(&AttrId::KeyData)
.ok_or_else(|| Error::UnsupportedAlgorithm("KeyData missing".into()))?
.to_vec();
Ok(Box::new(CustomData { key_bytes }))
}
fn conv_factory<'a>(mk: &'a Multikey) -> Result<Box<dyn ConvView + 'a>, Error> {
algorithm_name(mk)?;
Ok(Box::new(CustomConv {
key_type: key_type_byte(mk),
key: mk.clone(),
}))
}
fn fingerprint_factory<'a>(mk: &'a Multikey) -> Result<Box<dyn FingerprintView + 'a>, Error> {
let key_bytes = mk
.attributes
.get(&AttrId::KeyData)
.ok_or_else(|| Error::UnsupportedAlgorithm("KeyData missing".into()))?
.to_vec();
Ok(Box::new(CustomFingerprint { key_bytes }))
}
fn verify_factory<'a>(mk: &'a Multikey) -> Result<Box<dyn VerifyView + 'a>, Error> {
Ok(Box::new(NamedVerify {
name: algorithm_name(mk)?.to_string(),
}))
}
fn sign_factory(
serves: &'static str,
) -> impl for<'a> Fn(&'a Multikey) -> Result<Box<dyn SignView + 'a>, Error> + Send + Sync + 'static
{
move |mk| {
if algorithm_name(mk)? == serves {
Ok(Box::new(NamedSign { name: serves }))
} else {
Err(Error::UnsupportedAlgorithm(format!(
"this factory serves {serves} only"
)))
}
}
}
fn decoded_custom_key(name: &str, key_type: u8) -> Multikey {
let built = Builder::new(Codec::Identity)
.with_comment("custom protocol key")
.with_key_bytes(b"custom-key-seed".as_slice())
.with_algorithm_name(name)
.with_key_type(key_type)
.try_build()
.unwrap();
let bytes: Vec<u8> = built.into();
Multikey::try_from(bytes.as_ref()).unwrap()
}
#[test]
fn test_custom_key_builder_stamps_attributes() {
let built = Builder::new(Codec::Identity)
.with_key_bytes(b"custom-key-seed".as_slice())
.with_algorithm_name("my-protocol")
.with_key_type(1)
.try_build()
.unwrap();
assert_eq!(built.codec(), Codec::Identity);
assert_eq!(
built
.attributes
.get(&AttrId::AlgorithmName)
.unwrap()
.as_slice(),
b"my-protocol".as_slice()
);
assert_eq!(
built.attributes.get(&AttrId::KeyType).unwrap().as_slice(),
b"\x01".as_slice()
);
}
#[test]
fn test_custom_key_wire_roundtrip() {
let built = Builder::new(Codec::Identity)
.with_comment("custom protocol key")
.with_key_bytes(b"custom-key-seed".as_slice())
.with_algorithm_name("my-protocol")
.with_key_type(1)
.try_build()
.unwrap();
let bytes: Vec<u8> = built.clone().into();
let decoded = Multikey::try_from(bytes.as_ref()).unwrap();
assert_eq!(decoded, built);
assert_eq!(decoded.codec(), Codec::Identity);
assert_eq!(
decoded
.attributes
.get(&AttrId::AlgorithmName)
.unwrap()
.as_slice(),
b"my-protocol".as_slice()
);
assert_eq!(
decoded.attributes.get(&AttrId::KeyType).unwrap().as_slice(),
b"\x01".as_slice()
);
}
#[test]
fn test_empty_algorithm_name_is_accepted() {
let built = Builder::new(Codec::Identity)
.with_key_bytes(b"custom-key-seed".as_slice())
.with_algorithm_name("")
.with_key_type(0)
.try_build()
.unwrap();
assert_eq!(
built
.attributes
.get(&AttrId::AlgorithmName)
.unwrap()
.as_slice(),
b"".as_slice()
);
let bytes: Vec<u8> = built.into();
let decoded = Multikey::try_from(bytes.as_ref()).unwrap();
assert_eq!(
decoded
.attributes
.get(&AttrId::AlgorithmName)
.unwrap()
.as_slice(),
b"".as_slice()
);
}
#[test]
fn test_custom_key_factory_dispatch() {
let mk = decoded_custom_key("my-protocol", 1);
let attrs = ViewBuilder::new(&mk)
.attr()
.with_local_codec(Codec::Identity, attr_factory)
.build()
.unwrap();
assert!(attrs.is_secret_key());
assert!(!attrs.is_public_key());
assert!(!attrs.is_encrypted());
assert!(!attrs.is_secret_key_share());
let public = decoded_custom_key("my-protocol", 0);
let attrs = ViewBuilder::new(&public)
.attr()
.with_local_codec(Codec::Identity, attr_factory)
.build()
.unwrap();
assert!(attrs.is_public_key());
assert!(!attrs.is_secret_key());
let data = ViewBuilder::new(&mk)
.data()
.with_local_codec(Codec::Identity, data_factory)
.build()
.unwrap();
assert_eq!(
data.key_bytes().unwrap().as_slice(),
b"custom-key-seed".as_slice()
);
let conv = ViewBuilder::new(&mk)
.conv()
.with_local_codec(Codec::Identity, conv_factory)
.build()
.unwrap();
let public_key = conv.to_public_key().unwrap();
assert_eq!(
public_key
.attributes
.get(&AttrId::KeyType)
.unwrap()
.as_slice(),
b"\x00".as_slice()
);
assert_eq!(
conv.to_ssh_public_key().err().unwrap().to_string(),
"Unsupported key algorithm: custom keys have no ssh form"
);
let fp = ViewBuilder::new(&mk)
.fingerprint()
.with_local_codec(Codec::Identity, fingerprint_factory)
.build()
.unwrap();
let expected = mh::Builder::new_from_bytes(Codec::Blake2S256, b"custom-key-seed".as_slice())
.unwrap()
.try_build()
.unwrap();
assert_eq!(fp.fingerprint(Codec::Blake2S256).unwrap(), expected);
}
#[test]
fn test_algorithm_name_branching() {
let alpha = decoded_custom_key("alpha-protocol", 1);
let beta = decoded_custom_key("beta-protocol", 1);
let sign = ViewBuilder::new(&alpha)
.sign()
.with_local_codec(Codec::Identity, sign_factory("alpha-protocol"))
.build()
.unwrap();
assert_eq!(
sign.sign(b"message", false, None)
.err()
.unwrap()
.to_string(),
"Unsupported key algorithm: alpha-protocol sign"
);
let sign = ViewBuilder::new(&beta)
.sign()
.with_local_codec(Codec::Identity, sign_factory("beta-protocol"))
.build()
.unwrap();
assert_eq!(
sign.sign(b"message", false, None)
.err()
.unwrap()
.to_string(),
"Unsupported key algorithm: beta-protocol sign"
);
let err = ViewBuilder::new(&beta)
.sign()
.with_local_codec(Codec::Identity, sign_factory("alpha-protocol"))
.build()
.err()
.unwrap();
assert_eq!(
err.to_string(),
"Unsupported key algorithm: this factory serves alpha-protocol only"
);
let verify = ViewBuilder::new(&beta)
.verify()
.with_local_codec(Codec::Identity, verify_factory)
.build()
.unwrap();
let sig = Multisig::null();
assert_eq!(
verify.verify(&sig, None).err().unwrap().to_string(),
"Unsupported key algorithm: beta-protocol verify"
);
}
#[cfg(feature = "serde")]
#[test]
fn test_custom_key_serde_roundtrip() {
let mk = decoded_custom_key("my-protocol", 1);
let json = serde_json::to_string(&mk).unwrap();
let back: Multikey = serde_json::from_str(&json).unwrap();
assert_eq!(back, mk);
let mut cbor = Vec::new();
ciborium::into_writer(&mk, &mut cbor).unwrap();
let back: Multikey = ciborium::from_reader(cbor.as_slice()).unwrap();
assert_eq!(back, mk);
}