use std::fmt;
use jsonwebtoken::{Algorithm, DecodingKey, Validation};
#[cfg(test)]
use jsonwebtoken::{EncodingKey, Header};
#[cfg(test)]
use pubky_common::crypto::Keypair;
use pubky_common::crypto::PublicKey;
use serde::{Deserialize, Deserializer};
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct JwsCompact(String);
impl JwsCompact {
pub fn parse(s: &str) -> Result<Self, JwsCompactError> {
if s.splitn(4, '.').count() != 3 {
return Err(JwsCompactError);
}
Ok(Self(s.to_string()))
}
pub fn as_str(&self) -> &str {
&self.0
}
}
impl fmt::Display for JwsCompact {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(&self.0)
}
}
impl<'de> Deserialize<'de> for JwsCompact {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
let s = String::deserialize(deserializer)?;
Self::parse(&s).map_err(serde::de::Error::custom)
}
}
#[derive(Debug)]
pub struct JwsCompactError;
impl fmt::Display for JwsCompactError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("JWS Compact Serialization must have exactly 3 dot-separated parts")
}
}
impl std::error::Error for JwsCompactError {}
const ED25519_SPKI_PREFIX: [u8; 12] = [
0x30, 0x2a, 0x30, 0x05, 0x06, 0x03, 0x2b, 0x65, 0x70, 0x03, 0x21, 0x00,
];
#[cfg(test)]
pub fn encoding_key(keypair: &Keypair) -> EncodingKey {
let pem = ed25519_keypair_to_pem(keypair);
EncodingKey::from_ed_pem(pem.as_bytes())
.expect("invariant: PEM is constructed from valid Ed25519 key bytes")
}
pub fn decoding_key(pubkey: &PublicKey) -> DecodingKey {
let pem = ed25519_pubkey_to_pem(pubkey.as_bytes());
DecodingKey::from_ed_pem(pem.as_bytes())
.expect("invariant: PEM is constructed from valid Ed25519 key bytes")
}
#[cfg(test)]
pub fn eddsa_header(typ: &str) -> Header {
let mut header = Header::new(Algorithm::EdDSA);
header.typ = Some(typ.to_string());
header
}
pub fn eddsa_validation() -> Validation {
let mut validation = Validation::new(Algorithm::EdDSA);
validation.validate_exp = false;
validation.validate_aud = false;
validation.required_spec_claims.clear();
validation
}
#[cfg(test)]
fn ed25519_keypair_to_pem(keypair: &Keypair) -> String {
use base64::{engine::general_purpose::STANDARD, Engine};
let seed = keypair.secret();
let pubkey = keypair.public_key();
let mut der = Vec::with_capacity(85);
der.extend_from_slice(&[0x30, 0x53]);
der.extend_from_slice(&[0x02, 0x01, 0x01]);
der.extend_from_slice(&[0x30, 0x05, 0x06, 0x03, 0x2b, 0x65, 0x70]);
der.extend_from_slice(&[0x04, 0x22, 0x04, 0x20]);
der.extend_from_slice(&seed);
der.extend_from_slice(&[0xa1, 0x23, 0x03, 0x21, 0x00]);
der.extend_from_slice(pubkey.as_bytes());
debug_assert_eq!(der.len(), 85);
let b64 = STANDARD.encode(&der);
format!(
"-----BEGIN PRIVATE KEY-----\n{}\n-----END PRIVATE KEY-----\n",
b64
)
}
fn ed25519_pubkey_to_pem(pubkey: &[u8; 32]) -> String {
use base64::{engine::general_purpose::STANDARD, Engine};
let mut der = Vec::with_capacity(ED25519_SPKI_PREFIX.len() + 32);
der.extend_from_slice(&ED25519_SPKI_PREFIX);
der.extend_from_slice(pubkey);
let b64 = STANDARD.encode(&der);
format!(
"-----BEGIN PUBLIC KEY-----\n{}\n-----END PUBLIC KEY-----\n",
b64
)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn encoding_decoding_key_roundtrip() {
let keypair = Keypair::random();
let enc = encoding_key(&keypair);
let dec = decoding_key(&keypair.public_key());
let header = eddsa_header("test-jws");
let claims = serde_json::json!({"sub": "test", "exp": 9999999999u64});
let token = jsonwebtoken::encode(&header, &claims, &enc).unwrap();
let validation = eddsa_validation();
let decoded: jsonwebtoken::TokenData<serde_json::Value> =
jsonwebtoken::decode(&token, &dec, &validation).unwrap();
assert_eq!(decoded.claims["sub"], "test");
}
#[test]
fn pubky_common_sign_jws_round_trips_through_jsonwebtoken_decode() {
let kp = Keypair::random();
let claims = serde_json::json!({"sub": "interop", "exp": 9_999_999_999u64});
let compact = pubky_common::auth::jws::sign_jws(&kp, "test-jws", &claims);
let dec = decoding_key(&kp.public_key());
let validation = eddsa_validation();
let decoded: jsonwebtoken::TokenData<serde_json::Value> =
jsonwebtoken::decode(&compact, &dec, &validation).unwrap();
assert_eq!(decoded.claims["sub"], "interop");
assert_eq!(decoded.header.typ.as_deref(), Some("test-jws"));
}
#[test]
fn wrong_key_fails_verification() {
let keypair = Keypair::random();
let wrong_keypair = Keypair::random();
let enc = encoding_key(&keypair);
let wrong_dec = decoding_key(&wrong_keypair.public_key());
let header = eddsa_header("test-jws");
let claims = serde_json::json!({"sub": "test"});
let token = jsonwebtoken::encode(&header, &claims, &enc).unwrap();
let validation = eddsa_validation();
let result = jsonwebtoken::decode::<serde_json::Value>(&token, &wrong_dec, &validation);
assert!(result.is_err());
}
}