use anyhow::{bail, Result};
use base64::engine::general_purpose::URL_SAFE_NO_PAD;
use base64::Engine;
use serde_json::json;
use super::jwt;
use super::keys::SigningKey;
const JWT_BEARER: &str = "urn:ietf:params:oauth:client-assertion-type:jwt-bearer";
const ASSERTION_LIFETIME_SECS: i64 = 60;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum AuthMethod {
None,
PrivateKeyJwt,
}
impl AuthMethod {
pub fn negotiate(dev: bool) -> Self {
if dev {
AuthMethod::None
} else {
AuthMethod::PrivateKeyJwt
}
}
pub fn as_str(self) -> &'static str {
match self {
AuthMethod::None => "none",
AuthMethod::PrivateKeyJwt => "private_key_jwt",
}
}
}
impl std::str::FromStr for AuthMethod {
type Err = anyhow::Error;
fn from_str(value: &str) -> Result<Self> {
match value {
"none" => Ok(AuthMethod::None),
"private_key_jwt" => Ok(AuthMethod::PrivateKeyJwt),
other => bail!("unsupported client-auth method {other:?}"),
}
}
}
pub fn credential_params(
method: AuthMethod,
client_id: &str,
assertion: Option<&str>,
) -> Result<Vec<(&'static str, String)>> {
let mut params = vec![("client_id", client_id.to_string())];
if method == AuthMethod::PrivateKeyJwt {
let Some(assertion) = assertion else {
bail!("private_key_jwt requires a client assertion, but none was supplied");
};
params.push(("client_assertion_type", JWT_BEARER.to_string()));
params.push(("client_assertion", assertion.to_string()));
}
Ok(params)
}
pub fn client_assertion(
key: &SigningKey,
client_id: &str,
issuer: &str,
now: i64,
) -> Result<String> {
let mut jti = [0u8; 16];
getrandom::fill(&mut jti).expect("OS CSPRNG unavailable; refusing to mint a client assertion");
jwt::sign(
key,
&json!({ "kid": key.kid() }),
&json!({
"iss": client_id,
"sub": client_id,
"aud": issuer,
"jti": URL_SAFE_NO_PAD.encode(jti),
"iat": now,
"exp": now + ASSERTION_LIFETIME_SECS,
}),
)
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::Value;
const CLIENT_ID: &str = "https://feather-reader.com/oauth/client-metadata.json";
const ISSUER: &str = "https://auth.example.com";
const KID: &str = "featherreader-oauth-1";
const NOW: i64 = 1_700_000_000;
fn part(jws: &str, index: usize) -> Value {
let raw = URL_SAFE_NO_PAD
.decode(jws.split('.').nth(index).unwrap())
.unwrap();
serde_json::from_slice(&raw).unwrap()
}
#[test]
fn the_method_follows_the_client_shape() {
assert_eq!(AuthMethod::negotiate(true), AuthMethod::None);
assert_eq!(AuthMethod::negotiate(false), AuthMethod::PrivateKeyJwt);
assert_eq!(AuthMethod::None.as_str(), "none");
assert_eq!(AuthMethod::PrivateKeyJwt.as_str(), "private_key_jwt");
}
#[test]
fn the_method_round_trips_through_its_wire_name() {
for method in [AuthMethod::None, AuthMethod::PrivateKeyJwt] {
assert_eq!(method.as_str().parse::<AuthMethod>().unwrap(), method);
}
assert!("client_secret_basic".parse::<AuthMethod>().is_err());
assert!("".parse::<AuthMethod>().is_err());
}
#[test]
fn a_public_client_sends_only_its_client_id() {
let params = credential_params(AuthMethod::None, CLIENT_ID, None).unwrap();
assert_eq!(params, vec![("client_id", CLIENT_ID.to_string())]);
let params = credential_params(AuthMethod::None, CLIENT_ID, Some("jws")).unwrap();
assert_eq!(params.len(), 1);
assert!(!params
.iter()
.any(|(k, _)| k.starts_with("client_assertion")));
}
#[test]
fn a_confidential_client_sends_the_assertion_and_its_type() {
let params =
credential_params(AuthMethod::PrivateKeyJwt, CLIENT_ID, Some("the-jws")).unwrap();
assert_eq!(
params,
vec![
("client_id", CLIENT_ID.to_string()),
(
"client_assertion_type",
"urn:ietf:params:oauth:client-assertion-type:jwt-bearer".to_string()
),
("client_assertion", "the-jws".to_string()),
]
);
}
#[test]
fn a_confidential_client_without_an_assertion_is_an_error() {
assert!(credential_params(AuthMethod::PrivateKeyJwt, CLIENT_ID, None).is_err());
}
#[test]
fn the_assertion_carries_the_required_claims() {
let key = SigningKey::generate(KID);
let jws = client_assertion(&key, CLIENT_ID, ISSUER, NOW).unwrap();
let claims = part(&jws, 1);
assert_eq!(claims["iss"], CLIENT_ID);
assert_eq!(claims["sub"], CLIENT_ID);
assert_eq!(claims["iat"], NOW);
assert_eq!(claims["exp"], NOW + 60);
assert!(claims["jti"].as_str().is_some_and(|j| j.len() >= 22));
}
#[test]
fn the_audience_is_the_issuer() {
let key = SigningKey::generate(KID);
let claims = part(&client_assertion(&key, CLIENT_ID, ISSUER, NOW).unwrap(), 1);
assert_eq!(claims["aud"], ISSUER);
}
#[test]
fn the_header_carries_kid_and_alg_and_nothing_else() {
let key = SigningKey::generate(KID);
let jws = client_assertion(&key, CLIENT_ID, ISSUER, NOW).unwrap();
let header = part(&jws, 0);
assert_eq!(header["alg"], "ES256");
assert_eq!(header["kid"], KID);
assert!(header.get("typ").is_none(), "header should carry no typ");
assert_eq!(header.as_object().unwrap().len(), 2);
assert!(part(&jws, 1).get("nbf").is_none());
}
#[test]
fn every_assertion_gets_a_fresh_jti() {
let key = SigningKey::generate(KID);
let mut seen = std::collections::HashSet::new();
for _ in 0..32 {
let claims = part(&client_assertion(&key, CLIENT_ID, ISSUER, NOW).unwrap(), 1);
let jti = claims["jti"].as_str().unwrap().to_string();
assert!(seen.insert(jti), "jti repeated across assertions");
}
}
#[test]
fn the_assertion_verifies_under_the_clients_key() {
let key = SigningKey::generate(KID);
let jws = client_assertion(&key, CLIENT_ID, ISSUER, NOW).unwrap();
assert!(jwt::verify(&key, &jws).is_ok());
assert!(jwt::verify(&SigningKey::generate(KID), &jws).is_err());
}
#[test]
fn the_header_kid_matches_the_signing_key() {
let key = SigningKey::generate("some-other-kid");
let jws = client_assertion(&key, CLIENT_ID, ISSUER, NOW).unwrap();
assert_eq!(part(&jws, 0)["kid"], "some-other-kid");
assert_eq!(part(&jws, 0)["kid"], key.kid());
}
}