use std::time::{Duration, SystemTime, UNIX_EPOCH};
use jsonwebtoken::{Algorithm, EncodingKey, Header};
use rsa::pkcs1::{DecodeRsaPrivateKey, EncodeRsaPrivateKey};
use rsa::pkcs8::{DecodePrivateKey, EncodePublicKey, LineEnding};
use rsa::{RsaPrivateKey, RsaPublicKey};
use serde::Serialize;
use uuid::Uuid;
use zeroize::Zeroizing;
use super::BackendError;
const RSA_BITS: usize = 2048;
#[derive(Serialize)]
struct SvidClaims {
sub: String,
aud: String,
iat: u64,
exp: u64,
}
pub struct SvidMinter {
encoding_key: EncodingKey,
public_key_pem: String,
header: Header,
spiffe_id: String,
audience: String,
ttl: Duration,
}
impl SvidMinter {
pub fn generate(
spiffe_id: impl Into<String>,
audience: impl Into<String>,
ttl: Duration,
) -> Result<Self, BackendError> {
let mut rng = rand::thread_rng();
let private = RsaPrivateKey::new(&mut rng, RSA_BITS)
.map_err(|e| BackendError::Backend(format!("rsa keygen: {e}")))?;
Self::from_rsa_private(&private, spiffe_id, audience, ttl)
}
pub fn from_pem(
key_pem: &str,
spiffe_id: impl Into<String>,
audience: impl Into<String>,
ttl: Duration,
) -> Result<Self, BackendError> {
let private = RsaPrivateKey::from_pkcs8_pem(key_pem)
.or_else(|_| RsaPrivateKey::from_pkcs1_pem(key_pem))
.map_err(|e| BackendError::Backend(format!("decode signer private key: {e}")))?;
Self::from_rsa_private(&private, spiffe_id, audience, ttl)
}
fn from_rsa_private(
private: &RsaPrivateKey,
spiffe_id: impl Into<String>,
audience: impl Into<String>,
ttl: Duration,
) -> Result<Self, BackendError> {
let public = RsaPublicKey::from(private);
let private_pem: Zeroizing<String> = private
.to_pkcs1_pem(LineEnding::LF)
.map_err(|e| BackendError::Backend(format!("encode private key: {e}")))?;
let public_key_pem = public
.to_public_key_pem(LineEnding::LF)
.map_err(|e| BackendError::Backend(format!("encode public key: {e}")))?;
let encoding_key = EncodingKey::from_rsa_pem(private_pem.as_bytes())
.map_err(|e| BackendError::Backend(format!("load signing key: {e}")))?;
drop(private_pem);
let mut header = Header::new(Algorithm::RS256);
header.kid = Some(Uuid::new_v4().simple().to_string());
Ok(Self {
encoding_key,
public_key_pem,
header,
spiffe_id: spiffe_id.into(),
audience: audience.into(),
ttl,
})
}
pub fn public_key_pem(&self) -> &str {
&self.public_key_pem
}
pub fn spiffe_id(&self) -> &str {
&self.spiffe_id
}
pub fn mint(&self) -> Result<String, BackendError> {
let now = SystemTime::now()
.duration_since(UNIX_EPOCH)
.map_err(|e| BackendError::Backend(e.to_string()))?
.as_secs();
let claims = SvidClaims {
sub: self.spiffe_id.clone(),
aud: self.audience.clone(),
iat: now,
exp: now + self.ttl.as_secs(),
};
jsonwebtoken::encode(&self.header, &claims, &self.encoding_key)
.map_err(|e| BackendError::Backend(format!("mint svid: {e}")))
}
}
#[cfg(test)]
mod tests {
use super::{Duration, RsaPrivateKey, SvidMinter};
use rsa::pkcs1::EncodeRsaPrivateKey;
use rsa::pkcs8::{EncodePrivateKey, LineEnding};
fn test_key() -> RsaPrivateKey {
let mut rng = rand::thread_rng();
RsaPrivateKey::new(&mut rng, 2048).expect("rsa keygen")
}
#[test]
fn from_pem_loads_pkcs8_and_mints() {
let key = test_key();
let pem = key.to_pkcs8_pem(LineEnding::LF).expect("pkcs8 pem");
let minter = SvidMinter::from_pem(
&pem,
"spiffe://example.test/basil",
"openbao",
Duration::from_mins(1),
)
.expect("from_pem pkcs8");
assert_eq!(minter.spiffe_id(), "spiffe://example.test/basil");
assert!(minter.public_key_pem().contains("BEGIN PUBLIC KEY"));
assert!(!minter.mint().expect("mint").is_empty());
}
#[test]
fn from_pem_loads_pkcs1() {
let key = test_key();
let pem = key.to_pkcs1_pem(LineEnding::LF).expect("pkcs1 pem");
let minter = SvidMinter::from_pem(
&pem,
"spiffe://example.test/basil",
"openbao",
Duration::from_mins(1),
)
.expect("from_pem pkcs1");
assert_eq!(minter.spiffe_id(), "spiffe://example.test/basil");
}
#[test]
fn private_pem_intermediate_is_zeroizing() {
use zeroize::Zeroizing;
let key = test_key();
let private_pem: Zeroizing<String> = key.to_pkcs1_pem(LineEnding::LF).expect("pkcs1 pem");
assert!(private_pem.contains("PRIVATE KEY"));
let minter = SvidMinter::from_pem(
&private_pem,
"spiffe://example.test/basil",
"openbao",
Duration::from_mins(1),
)
.expect("from_pem after zeroizing-pem round-trip");
assert!(!minter.mint().expect("mint").is_empty());
}
#[test]
fn from_pem_rejects_garbage() {
match SvidMinter::from_pem(
"not a pem",
"spiffe://example.test/basil",
"openbao",
Duration::from_mins(1),
) {
Err(e) => assert!(e.to_string().contains("decode signer private key")),
Ok(_) => panic!("garbage pem must be rejected"),
}
}
}