use super::did_resolver::{DidError, DidResolver, ResolvedKey};
use base64::Engine;
use jsonwebtoken::{Algorithm, Validation, decode};
use serde::Deserialize;
#[derive(Debug)]
pub enum AuthError {
InvalidToken(String),
Transient(String),
}
impl std::fmt::Display for AuthError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
AuthError::InvalidToken(s) | AuthError::Transient(s) => s.fmt(f),
}
}
}
#[derive(Debug, Deserialize)]
#[serde(untagged)]
pub enum AudClaim {
Single(String),
Multiple(Vec<String>),
}
impl AudClaim {
pub fn contains(&self, value: &str) -> bool {
match self {
AudClaim::Single(s) => s == value,
AudClaim::Multiple(v) => v.iter().any(|s| s == value),
}
}
}
#[derive(Debug, Deserialize)]
#[allow(dead_code)] pub struct AtprotoClaims {
pub iss: String,
pub exp: u64,
pub nbf: Option<u64>,
pub aud: Option<AudClaim>,
}
#[derive(Debug)]
pub struct ValidatedIdentity {
pub did: String,
}
pub async fn validate_atproto_jwt(
token: &str,
resolver: &DidResolver,
expected_aud: Option<&str>,
) -> Result<ValidatedIdentity, AuthError> {
let claims = decode_claims(token).map_err(AuthError::InvalidToken)?;
let did = &claims.iss;
if !did.starts_with("did:") {
return Err(AuthError::InvalidToken(format!(
"invalid DID in JWT issuer: {did}"
)));
}
if !did.starts_with("did:plc:") && !did.starts_with("did:web:") {
return Err(AuthError::InvalidToken(format!(
"unsupported DID method in JWT issuer: {did}"
)));
}
if let Some(expected) = expected_aud {
match &claims.aud {
Some(aud) if aud.contains(expected) => {}
Some(aud) => {
let displayed = match aud {
AudClaim::Single(s) => format!("'{s}'"),
AudClaim::Multiple(v) => format!("{v:?}"),
};
return Err(AuthError::InvalidToken(format!(
"JWT audience mismatch: token is for {displayed}, expected '{expected}'"
)));
}
None => {
return Err(AuthError::InvalidToken(format!(
"JWT missing aud claim, expected '{expected}'"
)));
}
}
}
let resolved = resolver.resolve_key(did).await.map_err(|e| match e {
DidError::Authoritative(msg) => AuthError::InvalidToken(msg),
DidError::Transient(msg) => AuthError::Transient(msg),
})?;
verify_signature(token, &resolved)
.await
.map_err(AuthError::InvalidToken)?;
Ok(ValidatedIdentity { did: did.clone() })
}
fn decode_claims(token: &str) -> Result<AtprotoClaims, String> {
let parts: Vec<&str> = token.split('.').collect();
if parts.len() != 3 {
return Err(format!(
"malformed JWT: expected 3 parts, got {}",
parts.len()
));
}
let payload_b64 = parts[1];
let payload_bytes = base64::engine::general_purpose::URL_SAFE_NO_PAD
.decode(payload_b64)
.map_err(|e| format!("JWT payload base64 decode failed: {e}"))?;
let claims: AtprotoClaims = serde_json::from_slice(&payload_bytes)
.map_err(|e| format!("JWT claims parse failed: {e}"))?;
Ok(claims)
}
const LEEWAY_SECS: u64 = 60;
async fn verify_signature(token: &str, resolved: &ResolvedKey) -> Result<(), String> {
let token_owned = token.to_string();
let resolved_owned = resolved.clone();
tokio::task::spawn_blocking(move || match resolved_owned {
ResolvedKey::P256(key) => {
let mut validation = Validation::new(Algorithm::ES256);
validation.validate_exp = true;
validation.validate_nbf = true;
validation.validate_aud = false;
validation.leeway = LEEWAY_SECS;
decode::<AtprotoClaims>(&token_owned, &key, &validation)
.map_err(|e| format!("JWT signature verification failed: {e}"))?;
Ok(())
}
ResolvedKey::K256(public_key) => verify_es256k(&token_owned, &public_key),
})
.await
.map_err(|e| format!("ECDSA verification task panicked: {e}"))
.and_then(|r| r)
}
fn verify_es256k(token: &str, public_key: &k256::PublicKey) -> Result<(), String> {
use k256::ecdsa::{Signature, VerifyingKey, signature::Verifier};
let parts: Vec<&str> = token.rsplitn(2, '.').collect();
if parts.len() != 2 {
return Err("malformed JWT: expected header.payload.signature".to_string());
}
let sig_b64 = parts[0];
let signing_input = parts[1];
let header_b64 = signing_input
.split('.')
.next()
.ok_or("malformed JWT: missing header")?;
let header_bytes = base64::engine::general_purpose::URL_SAFE_NO_PAD
.decode(header_b64)
.map_err(|e| format!("JWT header base64 decode failed: {e}"))?;
let header: serde_json::Value = serde_json::from_slice(&header_bytes)
.map_err(|e| format!("JWT header parse failed: {e}"))?;
match header.get("alg").and_then(|v| v.as_str()) {
Some("ES256K") => {}
Some(other) => {
return Err(format!(
"JWT alg mismatch: token header says '{other}', but key requires ES256K"
));
}
None => {
return Err("JWT header missing 'alg' field".to_string());
}
}
use base64::Engine;
let sig_bytes = base64::engine::general_purpose::URL_SAFE_NO_PAD
.decode(sig_b64)
.map_err(|e| format!("JWT signature base64 decode failed: {e}"))?;
let signature =
Signature::from_slice(&sig_bytes).map_err(|e| format!("invalid ES256K signature: {e}"))?;
if signature.normalize_s().is_some() {
return Err("ES256K signature rejected: high-S (non-canonical) form".to_string());
}
let verifying_key = VerifyingKey::from(public_key);
verifying_key
.verify(signing_input.as_bytes(), &signature)
.map_err(|e| format!("ES256K signature verification failed: {e}"))?;
let claims = decode_claims(token)?;
let now = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map_err(|e| format!("system time error: {e}"))?
.as_secs();
if claims.exp.saturating_add(LEEWAY_SECS) < now {
return Err("JWT has expired".to_string());
}
if let Some(nbf) = claims.nbf
&& now + LEEWAY_SECS < nbf
{
return Err("JWT not yet valid (nbf)".to_string());
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use jsonwebtoken::{DecodingKey, EncodingKey, Header, encode};
use p256::pkcs8::EncodePrivateKey;
use serde::Serialize;
fn sign_test_jwt(did: &str, secret: &p256::SecretKey) -> String {
#[derive(Serialize)]
struct Claims {
iss: String,
exp: u64,
}
let private_pem = secret
.to_pkcs8_pem(p256::pkcs8::LineEnding::LF)
.expect("PEM encode private key");
let encoding_key = EncodingKey::from_ec_pem(private_pem.as_bytes()).expect("parse EC PEM");
let claims = Claims {
iss: did.to_string(),
exp: 9_999_999_999, };
encode(&Header::new(Algorithm::ES256), &claims, &encoding_key).expect("sign JWT")
}
fn test_key_pair() -> p256::SecretKey {
p256::SecretKey::from_slice(&[
0x9f, 0x86, 0xd0, 0x81, 0x88, 0x4c, 0x7d, 0x65, 0x9a, 0x2f, 0xea, 0xa0, 0xc5, 0x5a,
0xd0, 0x15, 0xa3, 0xbf, 0x4f, 0x1b, 0x2b, 0x0b, 0x82, 0x2c, 0xd1, 0x5d, 0x6c, 0x15,
0xb0, 0xf0, 0x0a, 0x08,
])
.expect("valid test key")
}
fn dummy_resolver() -> DidResolver {
DidResolver::with_plc_directory("http://127.0.0.1:1".to_string())
}
#[tokio::test]
async fn rejects_empty_token() {
assert!(
validate_atproto_jwt("", &dummy_resolver(), None)
.await
.is_err()
);
}
#[tokio::test]
async fn rejects_garbage_token() {
assert!(
validate_atproto_jwt("not.a.jwt", &dummy_resolver(), None)
.await
.is_err()
);
}
fn decoding_key_from_p256(public: &p256::PublicKey) -> DecodingKey {
use p256::pkcs8::EncodePublicKey;
let pem = public
.to_public_key_pem(p256::pkcs8::LineEnding::LF)
.expect("PEM encode public key");
DecodingKey::from_ec_pem(pem.as_bytes()).expect("parse EC PEM")
}
#[tokio::test]
async fn verify_signature_accepts_valid_jwt() {
let secret = test_key_pair();
let token = sign_test_jwt("did:plc:testuser", &secret);
let decoding_key = decoding_key_from_p256(&secret.public_key());
let result = verify_signature(&token, &ResolvedKey::P256(decoding_key)).await;
assert!(
result.is_ok(),
"verify_signature failed: {:?}",
result.err()
);
}
#[tokio::test]
async fn verify_signature_rejects_wrong_key() {
let secret = test_key_pair();
let token = sign_test_jwt("did:plc:testuser", &secret);
let wrong_secret = p256::SecretKey::from_slice(&[
0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, 0x08, 0x09, 0x0a, 0x0b, 0x0c, 0x0d, 0x0e,
0x0f, 0x10, 0x11, 0x12, 0x13, 0x14, 0x15, 0x16, 0x17, 0x18, 0x19, 0x1a, 0x1b, 0x1c,
0x1d, 0x1e, 0x1f, 0x20,
])
.expect("valid key");
let wrong_key = decoding_key_from_p256(&wrong_secret.public_key());
let result = verify_signature(&token, &ResolvedKey::P256(wrong_key)).await;
assert!(result.is_err());
assert!(
result
.unwrap_err()
.contains("signature verification failed"),
"should report signature failure"
);
}
#[tokio::test]
async fn verify_es256k_rejects_high_s_malleated_signature() {
use base64::Engine;
use k256::ecdsa::{Signature as K256Sig, SigningKey, signature::Signer};
let sk_bytes = [0x01u8; 32];
let secret = k256::SecretKey::from_slice(&sk_bytes).expect("valid test key");
let public = secret.public_key();
let signing_key = SigningKey::from(&secret);
let header_b64 = base64::engine::general_purpose::URL_SAFE_NO_PAD
.encode(br#"{"alg":"ES256K","typ":"JWT"}"#);
let payload_b64 = base64::engine::general_purpose::URL_SAFE_NO_PAD
.encode(br#"{"iss":"did:plc:test","exp":9999999999}"#);
let signing_input = format!("{header_b64}.{payload_b64}");
let sig: K256Sig = signing_key.sign(signing_input.as_bytes());
let low_sig_b64 =
base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(sig.to_bytes().as_slice());
let low_token = format!("{signing_input}.{low_sig_b64}");
assert!(
verify_es256k(&low_token, &public).is_ok(),
"canonical low-S signature must verify"
);
let (r, s) = sig.split_scalars();
let high_s = -*s;
let high_sig = K256Sig::from_scalars(r.to_bytes(), high_s.to_bytes())
.expect("negated s is a valid scalar");
let high_sig_b64 =
base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(high_sig.to_bytes().as_slice());
let high_token = format!("{signing_input}.{high_sig_b64}");
let err = verify_es256k(&high_token, &public).expect_err("high-S must be rejected");
assert!(
err.contains("high-S"),
"expected high-S rejection, got: {err}"
);
}
#[tokio::test]
async fn rejects_valid_jwt_with_no_matching_key() {
let secret = test_key_pair();
let token = sign_test_jwt("did:plc:testuser", &secret);
let resolver = DidResolver::with_plc_directory("http://127.0.0.1:1".to_string());
let result = validate_atproto_jwt(&token, &resolver, None).await;
assert!(result.is_err(), "should reject when key resolution fails");
}
}