#[cfg(feature = "kms")]
use aes_gcm::aead::{Aead, KeyInit, OsRng};
#[cfg(feature = "kms")]
use aes_gcm::{AeadCore, Aes256Gcm, Key, Nonce};
use bytes::Bytes;
#[cfg(feature = "kms")]
use std::collections::HashMap;
#[cfg(feature = "kms")]
use zeroize::Zeroizing;
#[cfg(feature = "kms")]
use crate::errors::Error;
use crate::errors::Result;
#[cfg(not(feature = "kms"))]
#[derive(Debug)]
pub enum SessionCrypto {}
#[cfg(not(feature = "kms"))]
impl SessionCrypto {
pub fn encrypt(&self, _plaintext: &[u8]) -> Result<Bytes> {
match *self {}
}
pub fn decrypt(&self, _frame: &[u8]) -> Result<Bytes> {
match *self {}
}
}
#[cfg(feature = "kms")]
const DATA_KEY_BYTES: i32 = 64;
#[cfg(feature = "kms")]
const NONCE_LEN: usize = 12;
#[cfg(feature = "kms")]
const TAG_LEN: usize = 16;
#[cfg(feature = "kms")]
pub struct SessionCrypto {
cipher_text_blob: Bytes,
encrypt: Aes256Gcm,
decrypt: Aes256Gcm,
}
#[cfg(feature = "kms")]
impl std::fmt::Debug for SessionCrypto {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("SessionCrypto(AES-256-GCM, keys redacted)")
}
}
#[cfg(feature = "kms")]
impl SessionCrypto {
pub async fn negotiate(
kms: &aws_sdk_kms::Client,
kms_key_id: &str,
session_id: &str,
target_id: &str,
) -> Result<Self> {
let context = HashMap::from([
("aws:ssm:SessionId".to_owned(), session_id.to_owned()),
("aws:ssm:TargetId".to_owned(), target_id.to_owned()),
]);
let output = kms
.generate_data_key()
.key_id(kms_key_id)
.number_of_bytes(DATA_KEY_BYTES)
.set_encryption_context(Some(context))
.send()
.await
.map_err(Error::from)?;
let plaintext = output
.plaintext()
.ok_or_else(|| Error::Crypto("KMS returned no plaintext data key".into()))?;
let blob = output
.ciphertext_blob()
.ok_or_else(|| Error::Crypto("KMS returned no ciphertext blob".into()))?;
let material = Zeroizing::new(plaintext.as_ref().to_vec());
Self::from_key_material(&material, Bytes::copy_from_slice(blob.as_ref()))
}
fn from_key_material(material: &[u8], cipher_text_blob: Bytes) -> Result<Self> {
if material.len() != DATA_KEY_BYTES as usize {
return Err(Error::Crypto(format!(
"expected a {DATA_KEY_BYTES}-byte data key, got {}",
material.len()
)));
}
let (decrypt_half, encrypt_half) = material.split_at(material.len() / 2);
Ok(Self {
cipher_text_blob,
encrypt: Aes256Gcm::new(Key::<Aes256Gcm>::from_slice(encrypt_half)),
decrypt: Aes256Gcm::new(Key::<Aes256Gcm>::from_slice(decrypt_half)),
})
}
pub fn cipher_text_blob(&self) -> &Bytes {
&self.cipher_text_blob
}
pub fn encrypt(&self, plaintext: &[u8]) -> Result<Bytes> {
let nonce = Aes256Gcm::generate_nonce(&mut OsRng);
let ciphertext = self
.encrypt
.encrypt(&nonce, plaintext)
.map_err(|_| Error::Crypto("AES-GCM encryption failed".into()))?;
let mut out = Vec::with_capacity(NONCE_LEN + ciphertext.len());
out.extend_from_slice(&nonce);
out.extend_from_slice(&ciphertext);
Ok(Bytes::from(out))
}
pub fn decrypt(&self, frame: &[u8]) -> Result<Bytes> {
if frame.len() < NONCE_LEN + TAG_LEN {
return Err(Error::Crypto(format!(
"encrypted frame is {} bytes, shorter than the {}-byte minimum",
frame.len(),
NONCE_LEN + TAG_LEN
)));
}
let (nonce, ciphertext) = frame.split_at(NONCE_LEN);
let plaintext = self
.decrypt
.decrypt(Nonce::from_slice(nonce), ciphertext)
.map_err(|_| {
Error::Crypto("AES-GCM authentication failed: frame corrupt or key mismatch".into())
})?;
Ok(Bytes::from(plaintext))
}
}
#[cfg(all(test, feature = "kms"))]
mod tests {
use super::*;
fn peers() -> (SessionCrypto, SessionCrypto) {
let material: Vec<u8> = (0..64u8).collect();
let client = SessionCrypto::from_key_material(&material, Bytes::from_static(b"blob"))
.expect("valid key material");
let mut swapped = material[32..].to_vec();
swapped.extend_from_slice(&material[..32]);
let agent = SessionCrypto::from_key_material(&swapped, Bytes::from_static(b"blob"))
.expect("valid key material");
(client, agent)
}
#[test]
fn client_ciphertext_decrypts_on_the_agent_side() {
let (client, agent) = peers();
let plaintext = b"echo hello from the client\n";
let frame = client.encrypt(plaintext).unwrap();
assert_eq!(agent.decrypt(&frame).unwrap().as_ref(), plaintext);
}
#[test]
fn agent_ciphertext_decrypts_on_the_client_side() {
let (client, agent) = peers();
let plaintext = b"total 0\r\n";
let frame = agent.encrypt(plaintext).unwrap();
assert_eq!(client.decrypt(&frame).unwrap().as_ref(), plaintext);
}
#[test]
fn a_peer_cannot_decrypt_its_own_ciphertext() {
let (client, _agent) = peers();
let frame = client.encrypt(b"secret").unwrap();
assert!(client.decrypt(&frame).is_err());
}
#[test]
fn frame_layout_is_nonce_then_ciphertext_and_tag() {
let (client, _) = peers();
let plaintext = b"1234567890";
let frame = client.encrypt(plaintext).unwrap();
assert_eq!(frame.len(), NONCE_LEN + plaintext.len() + TAG_LEN);
}
#[test]
fn nonces_are_unique_per_message() {
let (client, _) = peers();
let a = client.encrypt(b"same plaintext").unwrap();
let b = client.encrypt(b"same plaintext").unwrap();
assert_ne!(a[..NONCE_LEN], b[..NONCE_LEN], "nonce must not repeat");
assert_ne!(a, b);
}
#[test]
fn tampering_is_detected() {
let (client, agent) = peers();
let mut frame = client.encrypt(b"transfer $10").unwrap().to_vec();
let last = frame.len() - 1;
frame[last] ^= 0x01;
assert!(
agent.decrypt(&frame).is_err(),
"GCM tag must reject tampering"
);
}
#[test]
fn short_frames_are_rejected_without_panicking() {
let (_, agent) = peers();
for len in 0..(NONCE_LEN + TAG_LEN) {
assert!(agent.decrypt(&vec![0u8; len]).is_err(), "len {len}");
}
}
#[test]
fn empty_payload_round_trips() {
let (client, agent) = peers();
let frame = client.encrypt(b"").unwrap();
assert!(agent.decrypt(&frame).unwrap().is_empty());
}
#[test]
fn wrong_key_length_is_rejected() {
let Err(err) = SessionCrypto::from_key_material(&[0u8; 32], Bytes::new()) else {
panic!("a 32-byte key must be rejected");
};
assert!(err.to_string().contains("64-byte data key"), "{err}");
}
}