use base64::engine::general_purpose::NO_PAD;
use base64::engine::{GeneralPurpose, GeneralPurposeConfig};
use base64::{alphabet, Engine as _};
use jwtk::jwk::WithKid;
use jwtk::{decode_without_verify, ecdsa, sign, verify, HeaderAndClaims};
use reqwest;
use serde_derive::{Deserialize, Serialize};
use serde_json::Value;
use std::error::Error;
const NO_PAD_TRAILING_BITS: GeneralPurposeConfig = NO_PAD.with_decode_allow_trailing_bits(true);
const URL_SAFE_NO_PAD: GeneralPurpose =
GeneralPurpose::new(&alphabet::URL_SAFE, NO_PAD_TRAILING_BITS);
#[derive(Debug, Serialize, Deserialize, PartialEq)]
pub struct XQR {
pub token: String,
}
impl XQR {
pub fn get_kid(&self) -> Option<String> {
match decode_without_verify::<XQRClaims>(&self.token) {
Ok(header) => header.header().kid.clone().map(|s| s.to_string()),
Err(_) => None,
}
}
pub fn get_iss(&self) -> Option<String> {
match decode_without_verify::<XQRClaims>(&self.token) {
Ok(header) => header.claims().iss.clone().map(|s| s.to_string()),
Err(_) => None,
}
}
}
impl ToString for XQR {
fn to_string(&self) -> String {
self.token.clone()
}
}
impl From<String> for XQR {
fn from(s: String) -> Self {
XQR { token: s }
}
}
#[derive(Debug, Serialize, Deserialize)]
pub struct XQRClaims {
value: String,
}
pub fn encode(
private_key_pem: &str,
value: &str,
iss: &str,
valid_for: Option<std::time::Duration>,
) -> jwtk::Result<XQR> {
let private_key = ecdsa::EcdsaPrivateKey::from_pem(private_key_pem.as_ref())?;
let private_key = WithKid::new_with_thumbprint_id(private_key)?;
let mut claims = HeaderAndClaims::new_dynamic();
let claims = claims
.insert("value", value)
.set_iss(iss)
.set_iat_now()
.set_nbf_from_now(std::time::Duration::from_secs(0));
if valid_for.is_some() {
claims.set_exp_from_now(valid_for.unwrap());
}
let token = sign(claims, &private_key)?;
Ok(XQR { token })
}
pub fn decode(public_key_pem: &str, xqr: &XQR) -> jwtk::Result<String> {
let public_key = ecdsa::EcdsaPublicKey::from_pem(public_key_pem.as_ref())?;
let verified = verify::<XQRClaims>(&xqr.token, &public_key)?;
Ok(verified.claims().extra.value.clone())
}
pub fn fetch_public_key(issuer: &str, key_id: &str) -> Result<String, Box<dyn Error>> {
let domain = url::Url::parse(issuer)?;
let domain = domain.host_str().unwrap();
let url = format!("https://{}/.well-known/jwks.json", domain);
let response = reqwest::blocking::get(&url)?;
let jwks: Value = response.json()?;
if let Some(keys) = jwks["keys"].as_array() {
for key in keys {
if key["kid"].as_str() == Some(key_id) {
let pub_key = ecdsa::EcdsaPublicKey::from_coordinates(
&URL_SAFE_NO_PAD.decode(key["x"].as_str().unwrap())?,
&URL_SAFE_NO_PAD.decode(key["y"].as_str().unwrap())?,
ecdsa::EcdsaAlgorithm::ES256,
)?;
return Ok(pub_key.to_pem()?);
}
}
}
Err(Box::new(std::io::Error::new(
std::io::ErrorKind::NotFound,
"Key ID not found",
)))
}
pub fn generate_key() -> jwtk::Result<ecdsa::EcdsaPrivateKey> {
ecdsa::EcdsaPrivateKey::generate(ecdsa::EcdsaAlgorithm::ES256)
}
#[cfg(test)]
mod tests {
use super::*;
use std::time::Duration;
#[test]
fn encode_decode_test() {
let key = generate_key().unwrap();
let private_key = key.private_key_to_pem_pkcs8().unwrap();
let public_key = key.public_key_to_pem().unwrap();
let encoded_xqr = encode(&private_key, "value", "https://example.com", None).unwrap();
let decoded_value = decode(&public_key, &encoded_xqr).unwrap();
assert_eq!(decoded_value, "value");
}
#[test]
fn decode_with_wrong_pub_key_fails() {
let key = generate_key().unwrap();
let private_key = key.private_key_to_pem_pkcs8().unwrap();
let public_key = generate_key().unwrap().public_key_to_pem().unwrap();
let encoded_xqr = encode(&private_key, "value", "https://example.com", None).unwrap();
let decoded_value = decode(&public_key, &encoded_xqr);
assert!(decoded_value.is_err());
}
#[test]
fn get_kid_test() {
let key = generate_key().unwrap();
let private_key = key.private_key_to_pem_pkcs8().unwrap();
let encoded_xqr = encode(&private_key, "value", "https://example.com", None).unwrap();
assert!(encoded_xqr.get_kid().is_some());
}
#[test]
fn get_iss_test() {
let key = generate_key().unwrap();
let private_key = key.private_key_to_pem_pkcs8().unwrap();
let encoded_xqr = encode(&private_key, "value", "https://example.com", None).unwrap();
assert_eq!(encoded_xqr.get_iss().unwrap(), "https://example.com");
}
#[test]
fn pem_serialization_test() {
let key = generate_key().unwrap();
let private_pem = key.private_key_to_pem_pkcs8().unwrap();
let public_pem = key.public_key_to_pem().unwrap();
assert!(private_pem.contains("-----BEGIN PRIVATE KEY-----"));
assert!(private_pem.contains("-----END PRIVATE KEY-----"));
assert!(public_pem.contains("-----BEGIN PUBLIC KEY-----"));
assert!(public_pem.contains("-----END PUBLIC KEY-----"));
}
#[test]
fn xqr_to_string_ergonomics() {
let key = generate_key().unwrap();
let private_key = key.private_key_to_pem_pkcs8().unwrap();
let encoded_xqr = encode(&private_key, "value", "https://example.com", None).unwrap();
assert_eq!(encoded_xqr.to_string(), encoded_xqr.token);
}
#[test]
fn xqr_from_string_ergonomics() {
let key = generate_key().unwrap();
let private_key = key.private_key_to_pem_pkcs8().unwrap();
let encoded_xqr = encode(&private_key, "value", "https://example.com", None).unwrap();
let encoded_xqr_string = encoded_xqr.to_string();
assert_eq!(XQR::from(encoded_xqr_string), encoded_xqr);
}
#[test]
fn expiration_is_not_set_when_valid_for_is_none() {
let key = generate_key().unwrap();
let private_key = key.private_key_to_pem_pkcs8().unwrap();
let public_key = key.public_key_to_pem().unwrap();
let public_key = ecdsa::EcdsaPublicKey::from_pem(public_key.as_ref()).unwrap();
let encoded_xqr = encode(&private_key, "value", "https://example.com", None).unwrap();
let claims = verify::<XQRClaims>(&encoded_xqr.token, &public_key).unwrap();
assert!(claims.claims().exp.is_none());
}
#[test]
fn expiration_is_set_when_valid_for_is_not_none() {
let key = generate_key().unwrap();
let private_key = key.private_key_to_pem_pkcs8().unwrap();
let public_key = key.public_key_to_pem().unwrap();
let public_key = ecdsa::EcdsaPublicKey::from_pem(public_key.as_ref()).unwrap();
let encoded_xqr = encode(
&private_key,
"value",
"https://example.com",
Some(Duration::from_secs(60)),
)
.unwrap();
let claims = verify::<XQRClaims>(&encoded_xqr.token, &public_key).unwrap();
assert!(claims.claims().exp.is_some());
}
}