use core::marker::PhantomData;
use embassy_hal_internal::{Peri, PeripheralType};
use embassy_sync::waitqueue::AtomicWaker;
pub use crate::aes::{
AesCbc, AesCcm, AesCtr, AesEcb, AesGcm, Cipher, CipherAuthenticated, CipherSized, Context, Direction, Error,
IVSized, KeySize,
};
use crate::dma::ChannelAndRequest;
use crate::interrupt::typelevel::Interrupt;
use crate::mode::{Async, Blocking, Mode};
use crate::{interrupt, pac, peripherals, rcc};
static SAES_WAKER: AtomicWaker = AtomicWaker::new();
pub struct InterruptHandler<T: Instance> {
_phantom: PhantomData<T>,
}
impl<T: Instance> interrupt::typelevel::Handler<T::Interrupt> for InterruptHandler<T> {
unsafe fn on_interrupt() {
let isr = T::regs().isr().read();
if isr.ccf() {
T::regs().icr().write(|w| w.0 = 0xFFFF_FFFF);
SAES_WAKER.wake();
}
if isr.rweif() {
T::regs().icr().write(|w| w.0 = 0xFFFF_FFFF);
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[cfg_attr(feature = "defmt", derive(defmt::Format))]
pub enum HardwareKeySource {
DHUK = 1,
BHK = 2,
XorDhukBhk = 3,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[cfg_attr(feature = "defmt", derive(defmt::Format))]
pub enum KeyMode {
Normal = 0,
WrappedKey = 1,
SharedKey = 2,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[cfg_attr(feature = "defmt", derive(defmt::Format))]
pub enum KeyShareTarget {
AES = 0,
}
pub struct Saes<'d, T: Instance, M: Mode> {
_peripheral: Peri<'d, T>,
_phantom: PhantomData<M>,
#[allow(dead_code)] dma_in: Option<ChannelAndRequest<'d>>,
#[allow(dead_code)] dma_out: Option<ChannelAndRequest<'d>>,
}
impl<'d, T: Instance> Saes<'d, T, Blocking> {
pub fn new_blocking(
peripheral: Peri<'d, T>,
_irq: impl interrupt::typelevel::Binding<T::Interrupt, InterruptHandler<T>> + 'd,
) -> Self {
#[cfg(rng_wba6)]
{
let rcc = pac::RCC;
if !rcc.ahb2enr().read().rngen() {
rcc.ccipr2().modify(|w| w.set_rngsel(pac::rcc::vals::Rngsel::HSI));
rcc.ahb2enr().modify(|w| w.set_rngen(true));
pac::RNG.cr().modify(|w| w.set_rngen(true));
cortex_m::asm::delay(10_000);
}
}
rcc::enable_and_reset::<T>();
let p = T::regs();
while p.sr().read().busy() {}
assert!(!p.isr().read().rngeif(), "SAES: RNG error during initialization");
let instance = Self {
_peripheral: peripheral,
_phantom: PhantomData,
dma_in: None,
dma_out: None,
};
T::Interrupt::unpend();
unsafe { T::Interrupt::enable() };
instance
}
}
impl<'d, T: Instance> Saes<'d, T, Async> {
pub fn new<D1: DmaIn<T>, D2: DmaOut<T>>(
peripheral: Peri<'d, T>,
dma_in: Peri<'d, D1>,
dma_out: Peri<'d, D2>,
_irq: impl interrupt::typelevel::Binding<T::Interrupt, InterruptHandler<T>>
+ interrupt::typelevel::Binding<D1::Interrupt, crate::dma::InterruptHandler<D1>>
+ interrupt::typelevel::Binding<D2::Interrupt, crate::dma::InterruptHandler<D2>>
+ 'd,
) -> Self {
#[cfg(rng_wba6)]
{
let rcc = pac::RCC;
if !rcc.ahb2enr().read().rngen() {
rcc.ccipr2().modify(|w| w.set_rngsel(pac::rcc::vals::Rngsel::HSI));
rcc.ahb2enr().modify(|w| w.set_rngen(true));
pac::RNG.cr().modify(|w| w.set_rngen(true));
cortex_m::asm::delay(10_000);
}
}
rcc::enable_and_reset::<T>();
let p = T::regs();
while p.sr().read().busy() {}
assert!(!p.isr().read().rngeif(), "SAES: RNG error during initialization");
let instance = Self {
_peripheral: peripheral,
_phantom: PhantomData,
dma_in: new_dma!(dma_in, _irq),
dma_out: new_dma!(dma_out, _irq),
};
T::Interrupt::unpend();
unsafe { T::Interrupt::enable() };
instance
}
}
impl<'d, T: Instance, M: Mode> Saes<'d, T, M> {
pub fn start<'c, C>(&mut self, cipher: &'c C, dir: Direction) -> Context<'c, C>
where
C: Cipher<'c> + CipherSized + IVSized,
{
self.start_with_key_mode(cipher, dir, KeyMode::Normal, None)
}
pub fn start_with_hw_key<'c, C>(
&mut self,
key_source: HardwareKeySource,
cipher: &'c C,
dir: Direction,
) -> Context<'c, C>
where
C: Cipher<'c> + CipherSized + IVSized,
{
self.start_with_key_mode(cipher, dir, KeyMode::Normal, Some(key_source))
}
fn start_with_key_mode<'c, C>(
&mut self,
cipher: &'c C,
dir: Direction,
key_mode: KeyMode,
hw_key: Option<HardwareKeySource>,
) -> Context<'c, C>
where
C: Cipher<'c> + CipherSized + IVSized,
{
let p = T::regs();
p.cr().modify(|w| w.set_en(false));
while p.sr().read().busy() {}
p.cr().modify(|w| w.set_iprst(true));
p.cr().modify(|w| w.set_iprst(false));
p.icr().write(|w| w.0 = 0xFFFF_FFFF);
p.cr()
.modify(|w| w.set_datatype(pac::saes::vals::Datatype::from_bits(cipher.datatype())));
let keysize = cipher.key_size();
let keysize_val = match keysize {
KeySize::Bits128 => pac::saes::vals::Keysize::BITS128,
KeySize::Bits256 => pac::saes::vals::Keysize::BITS256,
};
p.cr().modify(|w| w.set_keysize(keysize_val));
while p.sr().read().busy() {}
self.set_cipher_mode(p, cipher);
let is_gcm_ccm = cipher.uses_gcm_phases();
let mode_val = match dir {
Direction::Encrypt => pac::saes::vals::Mode::ENCRYPTION,
Direction::Decrypt => pac::saes::vals::Mode::DECRYPTION,
};
p.cr().modify(|w| w.set_mode(mode_val));
let kmod_val = pac::saes::vals::Kmod::from_bits(key_mode as u8);
p.cr().modify(|w| w.set_kmod(kmod_val));
if is_gcm_ccm {
p.cr().modify(|w| w.set_gcmph(pac::saes::vals::Gcmph::from_bits(0)));
}
if let Some(hw_key_src) = hw_key {
let keysel_val = pac::saes::vals::Keysel::from_bits(hw_key_src as u8);
p.cr().modify(|w| w.set_keysel(keysel_val));
p.cr().modify(|w| w.set_keyprot(true));
while !p.sr().read().keyvalid() {}
} else {
self.load_key(cipher.key());
while !p.sr().read().keyvalid() {}
}
let needs_key_derivation = dir == Direction::Decrypt && matches!(cipher.chmod_bits(), 0 | 1);
if needs_key_derivation {
p.cr().modify(|w| w.set_mode(pac::saes::vals::Mode::KEY_DERIVATION));
p.cr().modify(|w| w.set_en(true));
while !p.isr().read().ccf() {}
p.icr().write(|w| w.0 = 0xFFFF_FFFF);
p.cr().modify(|w| w.set_mode(mode_val));
}
self.load_iv(cipher.iv());
if is_gcm_ccm {
p.cr().modify(|w| w.set_en(true));
while !p.isr().read().ccf() {}
p.icr().write(|w| w.0 = 0xFFFF_FFFF);
} else {
p.cr().modify(|w| w.set_en(true));
while p.sr().read().busy() {}
}
Context {
cipher,
dir,
last_block_processed: false,
is_gcm_ccm,
header_processed: false,
header_len: 0,
payload_len: 0,
aad_buffer: [0; 16],
aad_buffer_len: 0,
cr: p.cr().read().0,
iv: [p.ivr(0).read(), p.ivr(1).read(), p.ivr(2).read(), p.ivr(3).read()],
suspr: [0; 8],
}
}
pub fn share_key_with(&mut self, target: KeyShareTarget) {
let p = T::regs();
let kshareid_val = match target {
KeyShareTarget::AES => pac::saes::vals::Kshareid::AES,
};
p.cr().modify(|w| w.set_kshareid(kshareid_val));
}
fn set_cipher_mode<'c, C>(&mut self, p: pac::saes::Saes, cipher: &C)
where
C: Cipher<'c>,
{
p.cr()
.modify(|w| w.set_chmod(pac::saes::vals::Chmod::from_bits(cipher.chmod_bits())));
}
pub fn aad_blocking<'c, C>(&mut self, ctx: &mut Context<'c, C>, aad: &[u8], last: bool) -> Result<(), Error>
where
C: Cipher<'c> + CipherAuthenticated<16>,
{
let p = T::regs();
if ctx.header_processed && last {
return Ok(());
}
p.cr().modify(|w| w.set_gcmph(pac::saes::vals::Gcmph::from_bits(1)));
p.cr().modify(|w| w.set_en(true));
let mut aad_remaining = aad.len();
let mut aad_index = 0;
if ctx.aad_buffer_len > 0 {
let space_available = 16 - ctx.aad_buffer_len;
let to_copy = core::cmp::min(space_available, aad_remaining);
ctx.aad_buffer[ctx.aad_buffer_len..ctx.aad_buffer_len + to_copy].copy_from_slice(&aad[..to_copy]);
ctx.aad_buffer_len += to_copy;
aad_index += to_copy;
aad_remaining -= to_copy;
if ctx.aad_buffer_len == 16 {
self.write_block_blocking(&ctx.aad_buffer)?;
while !p.isr().read().ccf() {}
p.icr().write(|w| w.0 = 0xFFFF_FFFF);
ctx.header_len += 16;
ctx.aad_buffer_len = 0;
}
}
while aad_remaining >= 16 {
self.write_block_blocking(&aad[aad_index..aad_index + 16])?;
while !p.isr().read().ccf() {}
p.icr().write(|w| w.0 = 0xFFFF_FFFF);
ctx.header_len += 16;
aad_index += 16;
aad_remaining -= 16;
}
if aad_remaining > 0 {
ctx.aad_buffer[..aad_remaining].copy_from_slice(&aad[aad_index..aad_index + aad_remaining]);
ctx.aad_buffer_len = aad_remaining;
}
if last {
if ctx.aad_buffer_len > 0 {
for i in ctx.aad_buffer_len..16 {
ctx.aad_buffer[i] = 0;
}
self.write_block_blocking(&ctx.aad_buffer)?;
while !p.isr().read().ccf() {}
p.icr().write(|w| w.0 = 0xFFFF_FFFF);
ctx.header_len += ctx.aad_buffer_len as u64;
ctx.aad_buffer_len = 0;
}
ctx.header_processed = true;
}
Ok(())
}
pub fn payload_blocking<'c, C>(
&mut self,
ctx: &mut Context<'c, C>,
input: &[u8],
output: &mut [u8],
last: bool,
) -> Result<(), Error>
where
C: Cipher<'c>,
{
let p = T::regs();
if output.len() < input.len() {
return Err(Error::ConfigError);
}
if ctx.is_gcm_ccm {
if !ctx.header_processed {
ctx.header_processed = true;
}
p.cr().modify(|w| w.set_gcmph(pac::saes::vals::Gcmph::from_bits(2)));
p.cr().modify(|w| w.set_npblb(0));
p.cr().modify(|w| w.set_en(true));
}
let block_size = C::BLOCK_SIZE;
let mut processed = 0;
if C::REQUIRES_PADDING && !last && input.len() % block_size != 0 {
return Err(Error::ConfigError);
}
let complete_blocks = if last {
input.len() / block_size
} else {
input.len() / block_size
};
for _ in 0..complete_blocks {
let block = &input[processed..processed + block_size];
let out_block = &mut output[processed..processed + block_size];
self.write_block_blocking(block)?;
self.read_block_blocking(out_block)?;
processed += block_size;
ctx.payload_len += block_size as u64;
}
if last && processed < input.len() {
if C::REQUIRES_PADDING {
return Err(Error::ConfigError);
}
let remaining = input.len() - processed;
let mut partial_block = [0u8; 16];
partial_block[..remaining].copy_from_slice(&input[processed..]);
let padding_bytes = (16 - remaining) as u8;
p.cr().modify(|w| w.set_npblb(padding_bytes));
self.write_block_blocking(&partial_block)?;
self.read_block_blocking(&mut partial_block)?;
output[processed..processed + remaining].copy_from_slice(&partial_block[..remaining]);
ctx.payload_len += remaining as u64;
}
if last {
ctx.last_block_processed = true;
}
Ok(())
}
pub fn finish_blocking<'c, C>(&mut self, ctx: Context<'c, C>) -> Result<Option<[u8; 16]>, Error>
where
C: Cipher<'c>,
{
let p = T::regs();
if ctx.is_gcm_ccm {
while p.sr().read().busy() {}
p.cr().modify(|w| w.set_gcmph(pac::saes::vals::Gcmph::from_bits(3)));
let header_bits = (ctx.header_len * 8) as u64;
let payload_bits = (ctx.payload_len * 8) as u64;
let mut length_block = [0u8; 16];
length_block[0..8].copy_from_slice(&header_bits.to_be_bytes());
length_block[8..16].copy_from_slice(&payload_bits.to_be_bytes());
self.write_block_blocking(&length_block)?;
let mut tag = [0u8; 16];
self.read_block_blocking(&mut tag)?;
p.cr().modify(|w| w.set_en(false));
Ok(Some(tag))
} else {
p.cr().modify(|w| w.set_en(false));
Ok(None)
}
}
fn load_key(&mut self, key: &[u8]) {
let p = T::regs();
let key_words = key.len() / 4;
for i in 0..key_words {
let word = u32::from_be_bytes([key[i * 4], key[i * 4 + 1], key[i * 4 + 2], key[i * 4 + 3]]);
p.keyr(key_words - 1 - i).write_value(word); }
}
fn load_iv(&mut self, iv: &[u8]) {
if iv.is_empty() {
return;
}
let p = T::regs();
let iv_words = core::cmp::min(iv.len(), 16) / 4;
for i in 0..iv_words {
let word = u32::from_be_bytes([iv[i * 4], iv[i * 4 + 1], iv[i * 4 + 2], iv[i * 4 + 3]]);
p.ivr(i).write_value(word);
}
let remaining = core::cmp::min(iv.len(), 16) % 4;
if remaining > 0 {
let i = iv_words * 4;
let mut bytes = [0u8; 4];
bytes[..remaining].copy_from_slice(&iv[i..i + remaining]);
let word = u32::from_be_bytes(bytes);
p.ivr(iv_words).write_value(word);
}
}
fn write_block_blocking(&mut self, block: &[u8]) -> Result<(), Error> {
let p = T::regs();
if p.sr().read().wrerr() {
return Err(Error::WriteError);
}
for i in 0..4 {
let word = u32::from_be_bytes([block[i * 4], block[i * 4 + 1], block[i * 4 + 2], block[i * 4 + 3]]);
p.dinr().write_value(word);
}
Ok(())
}
fn read_block_blocking(&mut self, block: &mut [u8]) -> Result<(), Error> {
let p = T::regs();
while !p.isr().read().ccf() {}
if p.isr().read().rweif() {
p.icr().write(|w| w.0 = 0xFFFF_FFFF);
return Err(Error::ReadError);
}
for i in 0..4 {
let word = p.doutr().read();
let bytes = word.to_be_bytes();
block[i * 4..i * 4 + 4].copy_from_slice(&bytes);
}
p.icr().write(|w| w.0 = 0xFFFF_FFFF);
Ok(())
}
}
trait SealedInstance {
fn regs() -> pac::saes::Saes;
}
#[allow(private_bounds)]
pub trait Instance: SealedInstance + PeripheralType + crate::rcc::RccPeripheral + 'static + Send {
type Interrupt: interrupt::typelevel::Interrupt;
}
foreach_interrupt!(
($inst:ident, saes, SAES, GLOBAL, $irq:ident) => {
impl Instance for peripherals::$inst {
type Interrupt = crate::interrupt::typelevel::$irq;
}
impl SealedInstance for peripherals::$inst {
fn regs() -> crate::pac::saes::Saes {
crate::pac::$inst
}
}
};
);
dma_trait!(DmaIn, Instance);
dma_trait!(DmaOut, Instance);