systemprompt_security/keys/
mod.rs1use std::fs;
13use std::path::Path;
14
15use base64::Engine;
16use base64::engine::general_purpose::URL_SAFE_NO_PAD;
17use rsa::pkcs8::{DecodePrivateKey, EncodePrivateKey, EncodePublicKey, LineEnding};
18use rsa::rand_core::OsRng;
19use rsa::{RsaPrivateKey, RsaPublicKey};
20use sha2::{Digest, Sha256};
21
22pub mod authority;
23pub mod jwks;
24pub mod jwks_client;
25
26pub use authority::{TokenAuthorityError, TokenAuthorityResult};
27pub use jwks::{Jwk, Jwks};
28pub use jwks_client::{JwksClient, JwksClientError};
29
30pub const DEFAULT_RSA_BITS: usize = 2048;
31
32#[derive(Debug, thiserror::Error)]
33pub enum KeyError {
34 #[error("RSA key generation failed: {0}")]
35 Generation(#[source] rsa::Error),
36 #[error("PKCS#8 encoding failed: {0}")]
37 Encode(#[source] rsa::pkcs8::Error),
38 #[error("SPKI encoding failed: {0}")]
39 EncodeSpki(#[source] rsa::pkcs8::spki::Error),
40 #[error("PKCS#8 decoding failed: {0}")]
41 Decode(#[source] rsa::pkcs8::Error),
42 #[error("I/O error for {path}: {source}")]
43 Io {
44 path: String,
45 #[source]
46 source: std::io::Error,
47 },
48}
49
50#[derive(Clone)]
51pub struct RsaSigningKey {
52 private_key: RsaPrivateKey,
53 public_key: RsaPublicKey,
54 kid: String,
55}
56
57impl std::fmt::Debug for RsaSigningKey {
58 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
59 f.debug_struct("RsaSigningKey")
60 .field("kid", &self.kid)
61 .finish_non_exhaustive()
62 }
63}
64
65impl RsaSigningKey {
66 pub fn generate() -> Result<Self, KeyError> {
67 Self::generate_bits(DEFAULT_RSA_BITS)
68 }
69
70 pub fn generate_bits(bits: usize) -> Result<Self, KeyError> {
71 let mut rng = OsRng;
72 let private_key = RsaPrivateKey::new(&mut rng, bits).map_err(KeyError::Generation)?;
73 Self::from_private(private_key)
74 }
75
76 pub fn from_pkcs8_pem(pem: &str) -> Result<Self, KeyError> {
77 let private_key = RsaPrivateKey::from_pkcs8_pem(pem).map_err(KeyError::Decode)?;
78 Self::from_private(private_key)
79 }
80
81 pub fn load_from_pem_file(path: &Path) -> Result<Self, KeyError> {
82 let pem = fs::read_to_string(path).map_err(|source| KeyError::Io {
83 path: path.display().to_string(),
84 source,
85 })?;
86 Self::from_pkcs8_pem(&pem)
87 }
88
89 pub fn to_pkcs8_pem(&self) -> Result<String, KeyError> {
90 self.private_key
91 .to_pkcs8_pem(LineEnding::LF)
92 .map(|s| s.to_string())
93 .map_err(KeyError::Encode)
94 }
95
96 pub fn write_pem_file(&self, path: &Path) -> Result<(), KeyError> {
97 let pem = self.to_pkcs8_pem()?;
98 fs::write(path, pem).map_err(|source| KeyError::Io {
99 path: path.display().to_string(),
100 source,
101 })
102 }
103
104 pub const fn public_key(&self) -> &RsaPublicKey {
105 &self.public_key
106 }
107
108 pub const fn private_key(&self) -> &RsaPrivateKey {
109 &self.private_key
110 }
111
112 pub fn kid(&self) -> &str {
113 &self.kid
114 }
115
116 pub fn jwk(&self) -> Jwk {
117 Jwk::from_rsa_public_key(&self.public_key, self.kid.clone())
118 }
119
120 pub fn jwks(&self) -> Jwks {
121 Jwks {
122 keys: vec![self.jwk()],
123 }
124 }
125
126 fn from_private(private_key: RsaPrivateKey) -> Result<Self, KeyError> {
127 let public_key = RsaPublicKey::from(&private_key);
128 let kid = compute_kid(&public_key)?;
129 Ok(Self {
130 private_key,
131 public_key,
132 kid,
133 })
134 }
135}
136
137pub fn compute_kid(public_key: &RsaPublicKey) -> Result<String, KeyError> {
138 let der = public_key
139 .to_public_key_der()
140 .map_err(KeyError::EncodeSpki)?;
141 let digest = Sha256::digest(der.as_bytes());
142 Ok(URL_SAFE_NO_PAD.encode(digest))
143}