Skip to main content

systemprompt_security/keys/
mod.rs

1//! RSA signing-key infrastructure for systemprompt.io's federated JWT plane.
2//!
3//! Provides an [`RsaSigningKey`] wrapper around an `rsa::RsaPrivateKey` that
4//! can be generated, loaded from PKCS#8 PEM, persisted to PEM, and exposes a
5//! deterministic `kid` (SHA-256 of the DER-encoded `SubjectPublicKeyInfo`,
6//! base64 URL-encoded, no padding). The accompanying [`jwks`] module turns the
7//! public half into a JWKS document.
8//!
9//! Copyright (c) systemprompt.io — Business Source License 1.1.
10//! See <https://systemprompt.io> for licensing details.
11
12use 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}