use aws_lc_rs::{
digest::{Digest, SHA256, digest},
pkcs8::Document,
rand::SystemRandom,
signature::{
self, ECDSA_P256_SHA256_FIXED_SIGNING, EcdsaKeyPair, EcdsaSigningAlgorithm,
EcdsaVerificationAlgorithm, KeyPair, Signature,
},
};
use base64::{Engine as _, prelude::BASE64_URL_SAFE_NO_PAD};
use rama_core::error::{ErrorContext, OpaqueError};
use serde::{Deserialize, Serialize, Serializer, ser::SerializeStruct};
use crate::jose::{JWA, Signer};
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq)]
pub struct JWK {
pub alg: JWA,
#[serde(flatten)]
pub key_type: JWKType,
#[serde(skip_serializing_if = "Option::is_none")]
pub r#use: Option<JWKUse>,
#[serde(skip_serializing_if = "Option::is_none")]
pub key_ops: Option<Vec<String>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub x5c: Option<Vec<String>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub x5t: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
#[serde(rename = "x5t#S256")]
pub x5t_sha256: Option<String>,
}
#[derive(Clone, Debug, Deserialize, PartialEq, Eq)]
#[serde(tag = "kty")]
pub enum JWKType {
RSA {
n: String,
e: String,
},
EC {
crv: JWKEllipticCurves,
x: String,
y: String,
},
OCT {
k: String,
},
}
impl Serialize for JWKType {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: Serializer,
{
match &self {
Self::EC { crv, x, y } => {
let mut state = serializer.serialize_struct("JWKType", 4)?;
state.serialize_field("crv", crv)?;
state.serialize_field("kty", "EC")?;
state.serialize_field("x", x)?;
state.serialize_field("y", y)?;
state.end()
}
Self::RSA { n, e } => {
let mut state = serializer.serialize_struct("JWKType", 3)?;
state.serialize_field("e", e)?;
state.serialize_field("kty", "RSA")?;
state.serialize_field("n", n)?;
state.end()
}
Self::OCT { k } => {
let mut state = serializer.serialize_struct("JWKType", 2)?;
state.serialize_field("k", k)?;
state.serialize_field("kty", "OCT")?;
state.end()
}
}
}
}
#[derive(Clone, Copy, Debug, Serialize, Deserialize, PartialEq, Eq)]
pub enum JWKEllipticCurves {
#[serde(rename = "P-256")]
P256,
#[serde(rename = "P-384")]
P384,
#[serde(rename = "P-521")]
P521,
}
#[derive(Clone, Copy, Debug, Serialize, Deserialize, PartialEq, Eq)]
pub enum JWKUse {
#[serde(rename = "sig")]
Signature,
#[serde(rename = "enc")]
Encryption,
}
impl JWK {
fn new_from_escdsa_keypair(key: &EcdsaKeyPair, alg: JWA) -> Result<Self, OpaqueError> {
let curve = alg.try_into()?;
let pub_key = key.public_key().as_ref();
let middle = (pub_key.len() - 1) / 2;
let (x, y) = key.public_key().as_ref()[1..].split_at(middle);
Ok(Self {
alg,
key_type: JWKType::EC {
crv: curve,
x: BASE64_URL_SAFE_NO_PAD.encode(x),
y: BASE64_URL_SAFE_NO_PAD.encode(y),
},
r#use: Some(JWKUse::Signature),
key_ops: None,
x5c: None,
x5t: None,
x5t_sha256: None,
})
}
pub fn thumb_sha256(&self) -> Result<Digest, OpaqueError> {
Ok(digest(
&SHA256,
&serde_json::to_vec(&self.key_type).context("failed to serialise JWK")?,
))
}
pub fn unparsed_public_key(
&self,
) -> Result<signature::UnparsedPublicKey<Vec<u8>>, OpaqueError> {
match &self.key_type {
JWKType::RSA { .. } => Err(OpaqueError::from_display("currently not supported")),
JWKType::OCT { .. } => Err(OpaqueError::from_display(
"Symmetric key cannot be converted to public key",
)),
JWKType::EC { crv, x, y } => {
let alg: &'static EcdsaVerificationAlgorithm =
JWA::from(crv.to_owned()).try_into()?;
let x_bytes = BASE64_URL_SAFE_NO_PAD
.decode(x)
.context("decode ec curve x point")?;
let y_bytes = BASE64_URL_SAFE_NO_PAD
.decode(y)
.context("decode ec curve y point")?;
let mut point_bytes = Vec::with_capacity(1 + x_bytes.len() + y_bytes.len());
point_bytes.push(0x04);
point_bytes.extend_from_slice(&x_bytes);
point_bytes.extend_from_slice(&y_bytes);
Ok(signature::UnparsedPublicKey::new(alg, point_bytes))
}
}
}
}
pub struct EcdsaKey {
rng: SystemRandom,
alg: JWA,
inner: EcdsaKeyPair,
}
impl EcdsaKey {
pub fn new(key_pair: EcdsaKeyPair, alg: JWA, rng: SystemRandom) -> Result<Self, OpaqueError> {
let _curve = JWKEllipticCurves::try_from(alg)?;
Ok(Self {
rng,
alg,
inner: key_pair,
})
}
pub fn generate() -> Result<Self, OpaqueError> {
let key_pair = EcdsaKeyPair::generate(&ECDSA_P256_SHA256_FIXED_SIGNING)
.context("generate EcdsaKeyPair")?;
Self::new(key_pair, JWA::ES256, SystemRandom::new())
}
pub fn from_pkcs8_der(
pkcs8_der: &[u8],
alg: JWA,
rng: SystemRandom,
) -> Result<Self, OpaqueError> {
let ec_alg: &'static EcdsaSigningAlgorithm = alg.try_into()?;
let key_pair = EcdsaKeyPair::from_pkcs8(ec_alg, pkcs8_der)
.context("create EcdsaKeyPair from pkcs8")?;
Self::new(key_pair, alg, rng)
}
pub fn pkcs8_der(&self) -> Result<(JWA, Document), OpaqueError> {
let doc = self
.inner
.to_pkcs8v1()
.context("create pkcs8 der from keypair")?;
Ok((self.alg, doc))
}
#[must_use]
pub fn create_jwk(&self) -> JWK {
JWK::new_from_escdsa_keypair(&self.inner, self.alg).expect("create JWK from escdsa keypair")
}
#[must_use]
pub fn rng(&self) -> &SystemRandom {
&self.rng
}
#[must_use]
pub fn alg(&self) -> JWA {
self.alg
}
}
#[derive(Serialize)]
struct EcdsaKeySigningHeaders<'a> {
alg: JWA,
jwk: &'a JWK,
}
impl Signer for EcdsaKey {
type Signature = Signature;
type Error = OpaqueError;
fn set_headers(
&self,
protected_headers: &mut super::jws::Headers,
_unprotected_headers: &mut super::jws::Headers,
) -> Result<(), Self::Error> {
let jwk = self.create_jwk();
protected_headers.try_set_headers(EcdsaKeySigningHeaders {
alg: jwk.alg,
jwk: &jwk,
})?;
Ok(())
}
fn sign(&self, data: &str) -> Result<Self::Signature, Self::Error> {
let sig = self
.inner
.sign(self.rng(), data.as_bytes())
.context("sign protected data")?;
Ok(sig)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn jwk_thumb_order_is_correct() {
let jwk_type = JWKType::EC {
crv: JWKEllipticCurves::P256,
x: "x".into(),
y: "y".into(),
};
let output = serde_json::to_string(&jwk_type).unwrap();
let expected_output = r##"{"crv":"P-256","kty":"EC","x":"x","y":"y"}"##;
assert_eq!(&output, expected_output);
let jwk_type = JWKType::RSA {
n: "n".into(),
e: "e".into(),
};
let output = serde_json::to_string(&jwk_type).unwrap();
let expected_output = r##"{"e":"e","kty":"RSA","n":"n"}"##;
assert_eq!(&output, expected_output);
let jwk_type = JWKType::OCT { k: "k".into() };
let output = serde_json::to_string(&jwk_type).unwrap();
let expected_output = r##"{"k":"k","kty":"OCT"}"##;
assert_eq!(&output, expected_output);
}
#[test]
fn can_generate_and_reuse_keys() {
let key = EcdsaKey::generate().unwrap();
let stored = key.pkcs8_der().unwrap();
let recreated_key =
EcdsaKey::from_pkcs8_der(stored.1.as_ref(), stored.0, SystemRandom::new()).unwrap();
assert_eq!(key.create_jwk(), recreated_key.create_jwk())
}
}