1use std::path::Path;
4
5use crate::error::{Result, SignError};
6use crate::ml_dsa_bridge::{MlDsaSignerVariant, load_ml_dsa_signer_from_pkcs8_der};
7
8pub(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
60pub 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
79pub 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
95pub 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}