Skip to main content

hoike_sign/
keyfile.rs

1//! Software key loading from PKCS#8 PEM/DER files.
2
3use std::path::Path;
4
5use crate::error::{Result, SignError};
6use crate::ml_dsa_bridge::{MlDsaSignerVariant, load_ml_dsa_signer_from_pkcs8_der};
7
8/// Decode a PEM block into DER bytes.
9///
10/// Validates the label contains `expected_label` (e.g. "PRIVATE KEY"),
11/// handles whitespace trimming, and rejects truncated or missing PEM structure.
12pub(crate) fn pem_to_der(pem_data: &[u8], expected_label: &str) -> Result<Vec<u8>> {
13    let pem_str = std::str::from_utf8(pem_data)
14        .map_err(|e| SignError::KeyLoad(format!("key file is not valid UTF-8: {e}")))?;
15
16    use base64::Engine;
17    let mut collecting = false;
18    let mut found_end = false;
19    let mut b64 = String::new();
20    for line in pem_str.lines() {
21        if line.starts_with("-----BEGIN") {
22            if !line.contains(expected_label) {
23                return Err(SignError::KeyLoad(format!(
24                    "expected PEM {} but found: {}",
25                    expected_label,
26                    line.trim()
27                )));
28            }
29            collecting = true;
30            continue;
31        }
32        if line.starts_with("-----END") {
33            if !line.contains(expected_label) {
34                return Err(SignError::KeyLoad(format!(
35                    "PEM label mismatch: BEGIN {} but END says: {}",
36                    expected_label,
37                    line.trim()
38                )));
39            }
40            found_end = true;
41            break;
42        }
43        if collecting {
44            b64.push_str(line.trim());
45        }
46    }
47    if !collecting {
48        return Err(SignError::KeyLoad("no PEM header found in key file".into()));
49    }
50    if !found_end {
51        return Err(SignError::KeyLoad(
52            "PEM is truncated: found BEGIN but no END marker".into(),
53        ));
54    }
55    base64::engine::general_purpose::STANDARD
56        .decode(&b64)
57        .map_err(|e| SignError::KeyLoad(format!("invalid base64 in PEM: {e}")))
58}
59
60/// Load an ECDSA P-256 signing key from a PKCS#8 PEM or DER file.
61pub fn load_ecdsa_p256_key(path: &Path) -> Result<p256::ecdsa::SigningKey> {
62    let data = std::fs::read(path).map_err(|e| {
63        SignError::KeyLoad(format!("failed to read key file {}: {e}", path.display()))
64    })?;
65
66    if data.starts_with(b"-----BEGIN") {
67        let pem_str = std::str::from_utf8(&data)
68            .map_err(|e| SignError::KeyLoad(format!("key file is not valid UTF-8: {e}")))?;
69        use p256::pkcs8::DecodePrivateKey;
70        p256::ecdsa::SigningKey::from_pkcs8_pem(pem_str)
71            .map_err(|e| SignError::KeyLoad(format!("PKCS#8 PEM decode: {e}")))
72    } else {
73        use p256::pkcs8::DecodePrivateKey;
74        p256::ecdsa::SigningKey::from_pkcs8_der(&data)
75            .map_err(|e| SignError::KeyLoad(format!("PKCS#8 DER decode: {e}")))
76    }
77}
78
79/// Load an ML-DSA signer from a PKCS#8 PEM or DER file, auto-detecting the
80/// parameter set (44/65/87) from the AlgorithmIdentifier OID.
81pub fn load_ml_dsa_key(path: &Path) -> Result<MlDsaSignerVariant> {
82    let data = std::fs::read(path).map_err(|e| {
83        SignError::KeyLoad(format!("failed to read key file {}: {e}", path.display()))
84    })?;
85
86    let der_bytes = if data.starts_with(b"-----BEGIN") {
87        pem_to_der(&data, "PRIVATE KEY")?
88    } else {
89        data
90    };
91
92    load_ml_dsa_signer_from_pkcs8_der(&der_bytes).map_err(SignError::KeyLoad)
93}
94
95/// Generate an ephemeral ECDSA P-256 signing key for demo/testing use only.
96pub fn demo_ecdsa_p256_key() -> p256::ecdsa::SigningKey {
97    let secret = [42u8; 32];
98    p256::ecdsa::SigningKey::from_bytes((&secret).into()).expect("demo key generation failed")
99}
100
101#[cfg(test)]
102mod tests {
103    use super::*;
104    use p256::pkcs8::EncodePrivateKey;
105
106    #[test]
107    fn load_pkcs8_der_round_trip() {
108        let dir = tempfile::tempdir().unwrap();
109        let key_path = dir.path().join("test.der");
110
111        let original = demo_ecdsa_p256_key();
112        let der_bytes = original.to_pkcs8_der().expect("PKCS#8 DER encode failed");
113        std::fs::write(&key_path, der_bytes.as_bytes()).unwrap();
114
115        let loaded = load_ecdsa_p256_key(&key_path).unwrap();
116
117        use signature::Signer;
118        let msg = b"test message for signing";
119        let sig_orig: p256::ecdsa::DerSignature = original.sign(msg);
120        let sig_loaded: p256::ecdsa::DerSignature = loaded.sign(msg);
121        assert_eq!(sig_orig.to_bytes(), sig_loaded.to_bytes());
122    }
123
124    #[test]
125    fn load_pkcs8_pem_round_trip() {
126        let dir = tempfile::tempdir().unwrap();
127        let key_path = dir.path().join("test.pem");
128
129        let original = demo_ecdsa_p256_key();
130        let pem_str = original
131            .to_pkcs8_pem(p256::pkcs8::LineEnding::LF)
132            .expect("PKCS#8 PEM encode failed");
133        std::fs::write(&key_path, pem_str.as_bytes()).unwrap();
134
135        let loaded = load_ecdsa_p256_key(&key_path).unwrap();
136
137        use signature::Signer;
138        let sig_orig: p256::ecdsa::DerSignature = original.sign(b"msg");
139        let sig_loaded: p256::ecdsa::DerSignature = loaded.sign(b"msg");
140        assert_eq!(sig_orig.to_bytes(), sig_loaded.to_bytes());
141    }
142
143    #[test]
144    fn load_nonexistent_file_errors() {
145        let result = load_ecdsa_p256_key(Path::new("/nonexistent/key.pem"));
146        assert!(result.is_err());
147        let err = result.unwrap_err().to_string();
148        assert!(err.contains("failed to read key file"));
149    }
150
151    #[test]
152    fn load_invalid_data_errors() {
153        let dir = tempfile::tempdir().unwrap();
154        let key_path = dir.path().join("garbage.der");
155        std::fs::write(&key_path, b"not a valid key").unwrap();
156
157        let result = load_ecdsa_p256_key(&key_path);
158        assert!(result.is_err());
159    }
160
161    #[test]
162    fn load_ml_dsa_pkcs8_der_file() {
163        use ml_dsa::pkcs8::EncodePrivateKey;
164        let dir = tempfile::tempdir().unwrap();
165        let key_path = dir.path().join("ml-dsa-87.der");
166
167        let sk = ml_dsa::SigningKey::<ml_dsa::MlDsa87>::from_seed((&[42u8; 32]).into());
168        let der_doc = sk.to_pkcs8_der().expect("encode PKCS#8");
169        std::fs::write(&key_path, der_doc.as_bytes()).unwrap();
170
171        let variant = load_ml_dsa_key(&key_path).unwrap();
172        assert_eq!(variant.algorithm_name(), "ml-dsa-87");
173    }
174
175    #[test]
176    fn load_ml_dsa_pkcs8_pem_file() {
177        use ml_dsa::pkcs8::EncodePrivateKey;
178        let dir = tempfile::tempdir().unwrap();
179        let key_path = dir.path().join("ml-dsa-65.pem");
180
181        let sk = ml_dsa::SigningKey::<ml_dsa::MlDsa65>::from_seed((&[99u8; 32]).into());
182        let der_doc = sk.to_pkcs8_der().expect("encode PKCS#8");
183
184        use base64::Engine;
185        let b64 = base64::engine::general_purpose::STANDARD.encode(der_doc.as_bytes());
186        let mut pem = String::from("-----BEGIN PRIVATE KEY-----\n");
187        for chunk in b64.as_bytes().chunks(64) {
188            pem.push_str(std::str::from_utf8(chunk).unwrap());
189            pem.push('\n');
190        }
191        pem.push_str("-----END PRIVATE KEY-----\n");
192        std::fs::write(&key_path, pem.as_bytes()).unwrap();
193
194        let variant = load_ml_dsa_key(&key_path).unwrap();
195        assert_eq!(variant.algorithm_name(), "ml-dsa-65");
196    }
197
198    #[test]
199    fn load_ml_dsa_nonexistent_file_errors() {
200        let result = load_ml_dsa_key(Path::new("/nonexistent/ml-dsa.pem"));
201        assert!(result.is_err());
202    }
203
204    #[test]
205    fn load_ml_dsa_ecdsa_key_errors() {
206        let dir = tempfile::tempdir().unwrap();
207        let key_path = dir.path().join("ecdsa.der");
208
209        let ecdsa_key = demo_ecdsa_p256_key();
210        let der_bytes = ecdsa_key.to_pkcs8_der().expect("PKCS#8 DER encode failed");
211        std::fs::write(&key_path, der_bytes.as_bytes()).unwrap();
212
213        let result = load_ml_dsa_key(&key_path);
214        assert!(result.is_err());
215    }
216}