Skip to main content

haste_server/auth_n/certificates/providers/
local.rs

1use base64::{Engine as _, engine::general_purpose::URL_SAFE_NO_PAD};
2use haste_fhir_model::r4::generated::terminology::IssueType;
3use haste_fhir_operation_error::OperationOutcomeError;
4use rand::rngs::OsRng;
5use rsa::{
6    RsaPrivateKey,
7    pkcs1::{DecodeRsaPrivateKey, EncodeRsaPrivateKey, EncodeRsaPublicKey},
8    pkcs8::LineEnding,
9    traits::PublicKeyParts,
10};
11use sha1::{Digest, Sha1};
12use std::{path::Path, sync::Arc};
13use walkdir::{DirEntry, WalkDir};
14
15use crate::{
16    auth_n::certificates::{
17        JSONWebKey, JSONWebKeyAlgorithm, JSONWebKeySet, JSONWebKeyType,
18        traits::{CertificationProvider, DecodingKey, EncodingKey},
19    },
20    config::ServerConfig,
21};
22
23fn derive_kid(cert_path: &Path) -> String {
24    let file_name = Path::file_stem(cert_path)
25        .unwrap()
26        .to_str()
27        .unwrap()
28        .to_string();
29    let chunks = file_name.split("_").collect::<Vec<&str>>();
30    chunks.first().unwrap().to_string()
31}
32
33fn get_sorted_private_cert_paths(config: &ServerConfig) -> Vec<DirEntry> {
34    let certificate_dir = &config.certification_dir;
35
36    let cert_dir: &Path = Path::new(&certificate_dir);
37    let walker = WalkDir::new(cert_dir).into_iter();
38    let mut entries = walker
39        .filter_map(|e| e.ok())
40        .filter(|e| e.metadata().unwrap().is_file())
41        .filter(|e| e.file_name().to_str().unwrap().ends_with(".pem"))
42        .collect::<Vec<DirEntry>>();
43
44    entries.sort_by(|a, b| {
45        let a_chunks = Path::file_stem(a.path())
46            .unwrap()
47            .to_str()
48            .unwrap()
49            .split("_")
50            .collect::<Vec<&str>>();
51        let b_chunks = Path::file_stem(b.path())
52            .unwrap()
53            .to_str()
54            .unwrap()
55            .split("_")
56            .collect::<Vec<&str>>();
57
58        let date_a =
59            chrono::NaiveDate::parse_from_str(a_chunks.get(1).unwrap(), "%Y-%m-%d").unwrap();
60        let date_b =
61            chrono::NaiveDate::parse_from_str(b_chunks.get(1).unwrap(), "%Y-%m-%d").unwrap();
62
63        // latest first.
64        date_b.cmp(&date_a)
65    });
66
67    entries
68}
69
70fn create_jwk_set(
71    certificate_entries: &Vec<DirEntry>,
72) -> Result<JSONWebKeySet, OperationOutcomeError> {
73    let mut jsonweb_key_set = JSONWebKeySet { keys: vec![] };
74
75    for certification_entry in certificate_entries.iter() {
76        let cert_path = certification_entry.path();
77        let rsa_private =
78            RsaPrivateKey::from_pkcs1_pem(&std::fs::read_to_string(cert_path).unwrap()).unwrap();
79        let rsa_public_key = rsa_private.to_public_key();
80
81        let mut hasher = Sha1::new();
82        hasher.update(rsa_public_key.to_pkcs1_der().unwrap().as_bytes());
83        let x5t = hasher.finalize();
84
85        jsonweb_key_set.keys.push(JSONWebKey {
86            kid: derive_kid(cert_path),
87            alg: JSONWebKeyAlgorithm::RS256,
88            kty: JSONWebKeyType::RSA,
89            e: URL_SAFE_NO_PAD.encode(rsa_public_key.e().clone().to_bytes_be()),
90            n: URL_SAFE_NO_PAD.encode(rsa_public_key.n().clone().to_bytes_be()),
91            x5t: Some(URL_SAFE_NO_PAD.encode(x5t)),
92        });
93    }
94
95    Ok(jsonweb_key_set)
96}
97
98fn create_decoding_keys(
99    certificate_entries: &Vec<DirEntry>,
100) -> Result<Vec<DecodingKey>, OperationOutcomeError> {
101    let mut decoding_keys = vec![];
102
103    for certification_entry in certificate_entries.iter() {
104        let cert_path = certification_entry.path();
105        let rsa_private =
106            RsaPrivateKey::from_pkcs1_pem(&std::fs::read_to_string(cert_path).unwrap()).unwrap();
107
108        let rsa_public_key = rsa_private.to_public_key();
109
110        let decoding_key = jsonwebtoken::DecodingKey::from_rsa_pem(
111            rsa_public_key
112                .to_pkcs1_pem(LineEnding::default())
113                .unwrap()
114                .as_bytes(),
115        )
116        .unwrap();
117
118        decoding_keys.push(DecodingKey {
119            kid: derive_kid(cert_path),
120            decoding_key,
121        });
122    }
123
124    Ok(decoding_keys)
125}
126
127/// Latest key is first. this is set by date_b.cmp(&date_a) in get_sorted_private_cert_paths
128fn get_encoding_keys(
129    certificate_entries: &Vec<DirEntry>,
130) -> Result<Vec<EncodingKey>, OperationOutcomeError> {
131    let mut encoding_keys = vec![];
132
133    for certification_entry in certificate_entries.iter() {
134        let cert_path = certification_entry.path();
135        let encoding_key =
136            jsonwebtoken::EncodingKey::from_rsa_pem(&std::fs::read(cert_path).unwrap()).unwrap();
137
138        encoding_keys.push(EncodingKey {
139            kid: derive_kid(cert_path),
140            encoding_key,
141        });
142    }
143
144    Ok(encoding_keys)
145}
146
147fn create_certifications_if_needed(config: &ServerConfig) -> Result<(), OperationOutcomeError> {
148    let certificate_dir = &config.certification_dir;
149    let cert_dir: &Path = Path::new(&certificate_dir);
150
151    let private_key_files = get_sorted_private_cert_paths(config);
152
153    // If no private key than write.
154    if private_key_files.is_empty() {
155        let mut rng = OsRng;
156        let bits: usize = 2048;
157
158        // Use rfc 3339 format for date. Same as time_rotating.id.
159        let date = chrono::Utc::now();
160        let date2 = date + chrono::Days::new(5);
161
162        let private_key_file_name1 = format!("k1_{}.pem", date.format("%Y-%m-%d"));
163        let private_key_file_name2 = format!("k2_{}.pem", date2.format("%Y-%m-%d"));
164
165        let priv_key1 = RsaPrivateKey::new(&mut rng, bits).expect("failed to generate a key");
166        let priv_key2 = RsaPrivateKey::new(&mut rng, bits).expect("failed to generate a key");
167
168        std::fs::create_dir_all(cert_dir).unwrap();
169        std::fs::write(
170            cert_dir.join(private_key_file_name1),
171            priv_key1.to_pkcs1_pem(LineEnding::default()).unwrap(),
172        )
173        .map_err(|e| OperationOutcomeError::fatal(IssueType::exception(), e.to_string()))?;
174        std::fs::write(
175            cert_dir.join(private_key_file_name2),
176            priv_key2.to_pkcs1_pem(LineEnding::default()).unwrap(),
177        )
178        .map_err(|e| OperationOutcomeError::fatal(IssueType::exception(), e.to_string()))?;
179    }
180
181    Ok(())
182}
183
184pub struct LocalCertifications {
185    decoding_key: Arc<Vec<DecodingKey>>,
186    encoding_keys: Arc<Vec<EncodingKey>>,
187    jwk_set: Arc<JSONWebKeySet>,
188}
189
190impl LocalCertifications {
191    pub fn new(config: &ServerConfig) -> Result<Self, OperationOutcomeError> {
192        create_certifications_if_needed(config)?;
193
194        let private_certificate_entries = get_sorted_private_cert_paths(config);
195
196        Ok(LocalCertifications {
197            decoding_key: Arc::new(create_decoding_keys(&private_certificate_entries)?),
198            encoding_keys: Arc::new(get_encoding_keys(&private_certificate_entries)?),
199            jwk_set: Arc::new(create_jwk_set(&private_certificate_entries)?),
200        })
201    }
202}
203
204impl CertificationProvider for LocalCertifications {
205    fn decoding_key<'a>(&'a self, kid: &str) -> Result<&'a DecodingKey, OperationOutcomeError> {
206        self.decoding_key
207            .iter()
208            .find(|d| d.kid == kid)
209            .ok_or_else(|| {
210                OperationOutcomeError::error(
211                    IssueType::exception(),
212                    format!("No decoding key found for kid: '{}'", kid),
213                )
214            })
215    }
216
217    fn encoding_key(&self) -> Result<&EncodingKey, OperationOutcomeError> {
218        self.encoding_keys.first().ok_or_else(|| {
219            OperationOutcomeError::error(
220                IssueType::exception(),
221                "No encoding key available".to_string(),
222            )
223        })
224    }
225
226    fn jwk_set(&self) -> Arc<JSONWebKeySet> {
227        self.jwk_set.clone()
228    }
229}