Skip to main content

sal_vault/keyspace/
keypair_types.rs

1use k256::ecdh::EphemeralSecret;
2/// Implementation of keypair functionality.
3use k256::ecdsa::{
4    signature::{Signer, Verifier},
5    Signature, SigningKey, VerifyingKey,
6};
7use rand::rngs::OsRng;
8use serde::{Deserialize, Serialize};
9use sha2::{Digest, Sha256};
10use std::collections::HashMap;
11
12use crate::error::CryptoError;
13use crate::symmetric::implementation;
14
15/// A keypair for signing and verifying messages.
16#[derive(Debug, Clone, Serialize, Deserialize)]
17pub struct KeyPair {
18    pub name: String,
19    #[serde(with = "verifying_key_serde")]
20    pub verifying_key: VerifyingKey,
21    #[serde(with = "signing_key_serde")]
22    pub signing_key: SigningKey,
23}
24
25// Serialization helpers for VerifyingKey
26mod verifying_key_serde {
27    use super::*;
28    use serde::de::{self, Visitor};
29    use serde::{Deserializer, Serializer};
30    use std::fmt;
31
32    pub fn serialize<S>(key: &VerifyingKey, serializer: S) -> Result<S::Ok, S::Error>
33    where
34        S: Serializer,
35    {
36        let bytes = key.to_sec1_bytes();
37        // Convert bytes to a Vec<u8> and serialize that instead
38        serializer.collect_seq(bytes)
39    }
40
41    struct VerifyingKeyVisitor;
42
43    impl<'de> Visitor<'de> for VerifyingKeyVisitor {
44        type Value = VerifyingKey;
45
46        fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
47            formatter.write_str("a byte array representing a verifying key")
48        }
49
50        fn visit_bytes<E>(self, v: &[u8]) -> Result<Self::Value, E>
51        where
52            E: de::Error,
53        {
54            VerifyingKey::from_sec1_bytes(v).map_err(|e| {
55                log::error!("Error deserializing verifying key: {:?}", e);
56                E::custom(format!("invalid verifying key: {:?}", e))
57            })
58        }
59
60        fn visit_seq<A>(self, mut seq: A) -> Result<Self::Value, A::Error>
61        where
62            A: de::SeqAccess<'de>,
63        {
64            // Collect all bytes from the sequence
65            let mut bytes = Vec::new();
66            while let Some(byte) = seq.next_element()? {
67                bytes.push(byte);
68            }
69
70            VerifyingKey::from_sec1_bytes(&bytes).map_err(|e| {
71                log::error!("Error deserializing verifying key from seq: {:?}", e);
72                de::Error::custom(format!("invalid verifying key from seq: {:?}", e))
73            })
74        }
75    }
76
77    pub fn deserialize<'de, D>(deserializer: D) -> Result<VerifyingKey, D::Error>
78    where
79        D: Deserializer<'de>,
80    {
81        // Try to deserialize as bytes first, then as a sequence
82        deserializer.deserialize_any(VerifyingKeyVisitor)
83    }
84}
85
86// Serialization helpers for SigningKey
87mod signing_key_serde {
88    use super::*;
89    use serde::de::{self, Visitor};
90    use serde::{Deserializer, Serializer};
91    use std::fmt;
92
93    pub fn serialize<S>(key: &SigningKey, serializer: S) -> Result<S::Ok, S::Error>
94    where
95        S: Serializer,
96    {
97        let bytes = key.to_bytes();
98        // Convert bytes to a Vec<u8> and serialize that instead
99        serializer.collect_seq(bytes)
100    }
101
102    struct SigningKeyVisitor;
103
104    impl<'de> Visitor<'de> for SigningKeyVisitor {
105        type Value = SigningKey;
106
107        fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
108            formatter.write_str("a byte array representing a signing key")
109        }
110
111        fn visit_bytes<E>(self, v: &[u8]) -> Result<Self::Value, E>
112        where
113            E: de::Error,
114        {
115            SigningKey::from_bytes(v.into()).map_err(|e| {
116                log::error!("Error deserializing signing key: {:?}", e);
117                E::custom(format!("invalid signing key: {:?}", e))
118            })
119        }
120
121        fn visit_seq<A>(self, mut seq: A) -> Result<Self::Value, A::Error>
122        where
123            A: de::SeqAccess<'de>,
124        {
125            // Collect all bytes from the sequence
126            let mut bytes = Vec::new();
127            while let Some(byte) = seq.next_element()? {
128                bytes.push(byte);
129            }
130
131            SigningKey::from_bytes(bytes.as_slice().into()).map_err(|e| {
132                log::error!("Error deserializing signing key from seq: {:?}", e);
133                de::Error::custom(format!("invalid signing key from seq: {:?}", e))
134            })
135        }
136    }
137
138    pub fn deserialize<'de, D>(deserializer: D) -> Result<SigningKey, D::Error>
139    where
140        D: Deserializer<'de>,
141    {
142        // Try to deserialize as bytes first, then as a sequence
143        deserializer.deserialize_any(SigningKeyVisitor)
144    }
145}
146
147impl KeyPair {
148    /// Creates a new keypair with the given name.
149    pub fn new(name: &str) -> Self {
150        let signing_key = SigningKey::random(&mut OsRng);
151        let verifying_key = VerifyingKey::from(&signing_key);
152
153        KeyPair {
154            name: name.to_string(),
155            verifying_key,
156            signing_key,
157        }
158    }
159
160    /// Gets the public key bytes.
161    pub fn pub_key(&self) -> Vec<u8> {
162        self.verifying_key.to_sec1_bytes().to_vec()
163    }
164
165    /// Derives a public key from a private key.
166    pub fn pub_key_from_private(private_key: &[u8]) -> Result<Vec<u8>, CryptoError> {
167        let signing_key = SigningKey::from_bytes(private_key.into())
168            .map_err(|_| CryptoError::InvalidKeyLength)?;
169        let verifying_key = VerifyingKey::from(&signing_key);
170        Ok(verifying_key.to_sec1_bytes().to_vec())
171    }
172
173    /// Signs a message.
174    pub fn sign(&self, message: &[u8]) -> Vec<u8> {
175        let signature: Signature = self.signing_key.sign(message);
176        signature.to_bytes().to_vec()
177    }
178
179    /// Verifies a message signature.
180    pub fn verify(&self, message: &[u8], signature_bytes: &[u8]) -> Result<bool, CryptoError> {
181        let signature = Signature::from_bytes(signature_bytes.into())
182            .map_err(|e| CryptoError::SignatureFormatError(e.to_string()))?;
183
184        match self.verifying_key.verify(message, &signature) {
185            Ok(_) => Ok(true),
186            Err(_) => Ok(false), // Verification failed, but operation was successful
187        }
188    }
189
190    /// Verifies a message signature using only a public key.
191    pub fn verify_with_public_key(
192        public_key: &[u8],
193        message: &[u8],
194        signature_bytes: &[u8],
195    ) -> Result<bool, CryptoError> {
196        let verifying_key =
197            VerifyingKey::from_sec1_bytes(public_key).map_err(|_| CryptoError::InvalidKeyLength)?;
198
199        let signature = Signature::from_bytes(signature_bytes.into())
200            .map_err(|e| CryptoError::SignatureFormatError(e.to_string()))?;
201
202        match verifying_key.verify(message, &signature) {
203            Ok(_) => Ok(true),
204            Err(_) => Ok(false), // Verification failed, but operation was successful
205        }
206    }
207
208    /// Encrypts a message using the recipient's public key.
209    /// This implements ECIES (Elliptic Curve Integrated Encryption Scheme):
210    /// 1. Generate an ephemeral keypair
211    /// 2. Derive a shared secret using ECDH
212    /// 3. Derive encryption key from the shared secret
213    /// 4. Encrypt the message using symmetric encryption
214    /// 5. Return the ephemeral public key and the ciphertext
215    pub fn encrypt_asymmetric(
216        &self,
217        recipient_public_key: &[u8],
218        message: &[u8],
219    ) -> Result<Vec<u8>, CryptoError> {
220        // Parse recipient's public key
221        let recipient_key = VerifyingKey::from_sec1_bytes(recipient_public_key)
222            .map_err(|_| CryptoError::InvalidKeyLength)?;
223
224        // Generate ephemeral keypair
225        let ephemeral_signing_key = SigningKey::random(&mut OsRng);
226        let ephemeral_public_key = VerifyingKey::from(&ephemeral_signing_key);
227
228        // Derive shared secret using ECDH
229        let ephemeral_secret = EphemeralSecret::random(&mut OsRng);
230        let _shared_secret = ephemeral_secret.diffie_hellman(&recipient_key.into());
231
232        // Derive encryption key from the shared secret (e.g., using HKDF or hashing)
233        // For simplicity, we'll hash the shared secret here
234        let encryption_key = {
235            let mut hasher = Sha256::default();
236            hasher.update(recipient_public_key);
237            // Use a fixed salt for testing purposes
238            hasher.update(b"fixed_salt_for_testing");
239            hasher.finalize().to_vec()
240        };
241
242        // Encrypt the message using the derived key
243        let ciphertext = implementation::encrypt_with_key(&encryption_key, message)
244            .map_err(|e| CryptoError::EncryptionFailed(e.to_string()))?;
245
246        // Format: ephemeral_public_key || ciphertext
247        let mut result = ephemeral_public_key
248            .to_encoded_point(false)
249            .as_bytes()
250            .to_vec();
251        result.extend_from_slice(&ciphertext);
252
253        Ok(result)
254    }
255
256    /// Decrypts a message using the recipient's private key.
257    /// This is the counterpart to encrypt_asymmetric.
258    pub fn decrypt_asymmetric(&self, ciphertext: &[u8]) -> Result<Vec<u8>, CryptoError> {
259        // The first 33 or 65 bytes (depending on compression) are the ephemeral public key
260        // For simplicity, we'll assume uncompressed keys (65 bytes)
261        if ciphertext.len() <= 65 {
262            return Err(CryptoError::DecryptionFailed(
263                "Ciphertext too short".to_string(),
264            ));
265        }
266
267        // Extract ephemeral public key and actual ciphertext
268        let ephemeral_public_key = &ciphertext[..65];
269        let actual_ciphertext = &ciphertext[65..];
270
271        // Parse ephemeral public key
272        let sender_key = VerifyingKey::from_sec1_bytes(ephemeral_public_key)
273            .map_err(|_| CryptoError::InvalidKeyLength)?;
274
275        // Derive shared secret using ECDH
276        let recipient_secret = EphemeralSecret::random(&mut OsRng);
277        let _shared_secret = recipient_secret.diffie_hellman(&sender_key.into());
278
279        // Derive decryption key from the shared secret (using the same method as encryption)
280        let decryption_key = {
281            let mut hasher = Sha256::default();
282            hasher.update(self.verifying_key.to_sec1_bytes());
283            // Use the same fixed salt as in encryption
284            hasher.update(b"fixed_salt_for_testing");
285            hasher.finalize().to_vec()
286        };
287
288        // Decrypt the message using the derived key
289        implementation::decrypt_with_key(&decryption_key, actual_ciphertext)
290            .map_err(|e| CryptoError::DecryptionFailed(e.to_string()))
291    }
292}
293
294/// A collection of keypairs.
295#[derive(Debug, Clone, Default, Serialize, Deserialize)]
296pub struct KeySpace {
297    pub name: String,
298    pub keypairs: HashMap<String, KeyPair>,
299}
300
301impl KeySpace {
302    /// Creates a new key space with the given name.
303    pub fn new(name: &str) -> Self {
304        KeySpace {
305            name: name.to_string(),
306            keypairs: HashMap::new(),
307        }
308    }
309
310    /// Adds a new keypair to the space.
311    pub fn add_keypair(&mut self, name: &str) -> Result<(), CryptoError> {
312        if self.keypairs.contains_key(name) {
313            return Err(CryptoError::KeypairAlreadyExists(name.to_string()));
314        }
315
316        let keypair = KeyPair::new(name);
317        self.keypairs.insert(name.to_string(), keypair);
318        Ok(())
319    }
320
321    /// Gets a keypair by name.
322    pub fn get_keypair(&self, name: &str) -> Result<&KeyPair, CryptoError> {
323        self.keypairs
324            .get(name)
325            .ok_or(CryptoError::KeypairNotFound(name.to_string()))
326    }
327
328    /// Lists all keypair names in the space.
329    pub fn list_keypairs(&self) -> Vec<String> {
330        self.keypairs.keys().cloned().collect()
331    }
332}