use jsonwebtoken::jwk::{AlgorithmParameters, EllipticCurve, Jwk, ThumbprintHash};
use jsonwebtoken::{crypto, Algorithm, AlgorithmFamily, DecodingKey};
use serde::Deserialize;
use serde_json::Value;
use sha2::{Digest, Sha256};
use crate::config::VerifierConfig;
use crate::error::VerifyError;
use crate::jws_util::{
constant_time_eq, decode_json_segment, parse_and_check_alg, reject_private_jwk_members,
split_compact, RawJws,
};
use crate::request::SignedRequest;
#[derive(Debug, Deserialize)]
struct RawSignatureClaims {
iat: i64,
exp: i64,
jti: String,
mth: String,
pth: String,
#[serde(default)]
qsh: Option<String>,
#[serde(default)]
bdh: Option<String>,
aud: String,
}
pub struct ParsedSignature<'a> {
raw: RawJws<'a>,
alg: Algorithm,
jwk_json: Value,
}
pub fn parse<'a>(
token: &'a str,
allowed_algs: &[Algorithm],
) -> Result<ParsedSignature<'a>, VerifyError> {
let raw = split_compact(token)?;
let header_json = decode_json_segment(raw.header_b64)?;
let alg = parse_and_check_alg(&header_json, allowed_algs)?;
let jwk_json = header_json
.get("jwk")
.cloned()
.ok_or_else(|| VerifyError::BadJwk("signature header has no embedded jwk".to_string()))?;
Ok(ParsedSignature { raw, alg, jwk_json })
}
pub struct VerifiedSignature {
pub jti: String,
pub exp: i64,
pub jwk_thumbprint: String,
}
pub fn verify_bound_and_signed(
parsed: &ParsedSignature<'_>,
request: &SignedRequest<'_>,
config: &VerifierConfig,
cnf_jkt: &str,
now: i64,
) -> Result<VerifiedSignature, VerifyError> {
reject_private_jwk_members(&parsed.jwk_json)?;
let jwk: Jwk = serde_json::from_value(parsed.jwk_json.clone())
.map_err(|e| VerifyError::BadJwk(format!("embedded jwk does not parse: {e}")))?;
let type_matches_alg = matches!(
(&jwk.algorithm, parsed.alg.family()),
(AlgorithmParameters::EllipticCurve(_), AlgorithmFamily::Ec)
| (AlgorithmParameters::RSA(_), AlgorithmFamily::Rsa)
| (AlgorithmParameters::OctetKeyPair(_), AlgorithmFamily::Ed)
);
if !type_matches_alg {
return Err(VerifyError::BadJwk(format!(
"embedded jwk type does not match header alg {:?}",
parsed.alg
)));
}
if !curve_is_consistent(&jwk.algorithm) {
return Err(VerifyError::BadJwk(
"embedded jwk has an inconsistent kty/crv pairing (e.g. an EC key claiming the \
Ed25519 curve, or an OKP key claiming a Weierstrass curve)"
.to_string(),
));
}
let thumbprint = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
jwk.thumbprint(ThumbprintHash::SHA256)
}))
.map_err(|_| {
tracing::error!(
target: "authkestra_devsig",
"embedded jwk panicked inside jsonwebtoken::jwk::Jwk::thumbprint() -- rejecting as \
bad_jwk instead of propagating the panic"
);
VerifyError::BadJwk("embedded jwk is not a supported/consistent key shape".to_string())
})?;
if !constant_time_eq(&thumbprint, cnf_jkt) {
tracing::warn!(
target: "authkestra_devsig",
"device-signature rejected: embedded jwk thumbprint does not match attestation cnf.jkt \
(key_not_bound) — this is the forgery the binding check exists to catch"
);
return Err(VerifyError::KeyNotBound);
}
let decoding_key = DecodingKey::from_jwk(&jwk).map_err(|e| {
VerifyError::BadJwk(format!("embedded jwk unusable as a decoding key: {e}"))
})?;
let signing_input = format!("{}.{}", parsed.raw.header_b64, parsed.raw.payload_b64);
let sig_ok = crypto::verify(
parsed.raw.signature_b64,
signing_input.as_bytes(),
&decoding_key,
parsed.alg,
)
.map_err(|e| VerifyError::BadSignature(format!("verification error: {e}")))?;
if !sig_ok {
return Err(VerifyError::BadSignature(
"signature does not verify".to_string(),
));
}
let claims_json = decode_json_segment(parsed.raw.payload_b64)?;
let claims: RawSignatureClaims = serde_json::from_value(claims_json)
.map_err(|e| VerifyError::Malformed(format!("invalid signature claims: {e}")))?;
let skew = config.max_clock_skew.as_secs() as i64;
if now < claims.iat - skew || now > claims.exp + skew {
return Err(VerifyError::SignatureExpired);
}
let lifetime = claims.exp - claims.iat;
if lifetime < 0 || lifetime as u64 > config.max_signature_lifetime.as_secs() {
return Err(VerifyError::LifetimeTooLong);
}
if claims.mth != request.method {
return Err(VerifyError::MethodMismatch);
}
if claims.pth != request.path {
return Err(VerifyError::PathMismatch);
}
if claims.aud != config.audience {
return Err(VerifyError::AudienceMismatch);
}
if let Some(query) = request.query {
let computed = hex_sha256(query.as_bytes());
match &claims.qsh {
Some(qsh) if constant_time_eq(qsh, &computed) => {}
_ => return Err(VerifyError::QueryMismatch),
}
}
if let Some(body) = request.body {
let computed = hex_sha256(body);
match &claims.bdh {
Some(bdh) if constant_time_eq(bdh, &computed) => {}
_ => return Err(VerifyError::BodyMismatch),
}
}
Ok(VerifiedSignature {
jti: claims.jti,
exp: claims.exp,
jwk_thumbprint: thumbprint,
})
}
fn curve_is_consistent(algorithm: &AlgorithmParameters) -> bool {
match algorithm {
AlgorithmParameters::EllipticCurve(p) => !matches!(p.curve, EllipticCurve::Ed25519),
AlgorithmParameters::OctetKeyPair(p) => matches!(p.curve, EllipticCurve::Ed25519),
AlgorithmParameters::RSA(_) | AlgorithmParameters::OctetKey(_) => true,
}
}
fn hex_sha256(data: &[u8]) -> String {
let digest = Sha256::digest(data);
digest.iter().map(|byte| format!("{byte:02x}")).collect()
}