use crate::cipherstate::CipherStates;
use crate::error::{CipherResult, HandshakeResult};
use crate::symmetricstate::SymmetricState;
use crate::traits::{Cipher, Handshaker, HandshakerInternal, Hash};
use crate::transportstate::TransportState;
pub struct DualLayerHandshake<Outer, Inner, C, H, const BUF: usize>
where
Inner: Handshaker<C, H>,
Outer: Handshaker<C, H>,
C: Cipher,
H: Hash,
{
outer: Option<Outer>,
inner: Inner,
outer_transport: Option<TransportState<C, H>>,
outer_is_finished: bool,
outer_receive_buf: [u8; BUF],
}
impl<Outer, Inner, C, H, const BUF: usize> DualLayerHandshake<Outer, Inner, C, H, BUF>
where
Inner: Handshaker<C, H>,
Outer: Handshaker<C, H>,
C: Cipher,
H: Hash,
{
pub fn new(outer: Outer, inner: Inner) -> Self {
assert!(outer.is_initiator() == inner.is_initiator());
assert!(!outer.get_pattern().is_one_way());
Self {
outer: Some(outer),
inner,
outer_transport: None,
outer_is_finished: false,
outer_receive_buf: [0u8; BUF],
}
}
pub fn outer_completed(&self) -> bool {
self.outer_is_finished
}
pub fn inner(&self) -> &Inner {
&self.inner
}
pub fn inner_mut(&mut self) -> &mut Inner {
&mut self.inner
}
pub fn outer(&self) -> Option<&Outer> {
self.outer.as_ref()
}
pub fn outer_mut(&mut self) -> Option<&mut Outer> {
self.outer.as_mut()
}
fn update_outer_state(&mut self) -> HandshakeResult<()> {
if self.outer.as_ref().unwrap().is_finished() {
self.outer_transport = Some(self.outer.take().unwrap().finalize()?);
self.outer_is_finished = true;
}
Ok(())
}
}
impl<Outer, Inner, C, H, const BUF: usize> HandshakerInternal<C, H>
for DualLayerHandshake<Outer, Inner, C, H, BUF>
where
Inner: Handshaker<C, H>,
Outer: Handshaker<C, H>,
C: Cipher,
H: Hash,
{
fn status(&self) -> super::HandshakeStatus {
if self.outer_completed() {
self.inner.status()
} else {
self.outer.as_ref().unwrap().status()
}
}
fn set_error(&mut self) {
self.inner.set_error();
if self.outer.is_some() {
self.outer.as_mut().unwrap().set_error();
}
}
fn write_message_impl(
&mut self,
payload: &[u8],
out: &mut [u8],
) -> crate::error::HandshakeResult<usize> {
if self.outer_completed() {
let n = self.inner.write_message_impl(payload, out)?;
let n = self
.outer_transport
.as_mut()
.unwrap()
.send_in_place(out, n)?;
Ok(n)
} else {
let r = self
.outer
.as_mut()
.unwrap()
.write_message_impl(payload, out)?;
self.update_outer_state()?;
Ok(r)
}
}
fn read_message_impl(
&mut self,
message: &[u8],
out: &mut [u8],
) -> crate::error::HandshakeResult<usize> {
if self.outer_completed() {
let n = self
.outer_transport
.as_mut()
.unwrap()
.receive(message, &mut self.outer_receive_buf)?;
self.inner
.read_message_impl(&self.outer_receive_buf[..n], out)
} else {
let r = self
.outer
.as_mut()
.unwrap()
.read_message_impl(message, out)?;
self.update_outer_state()?;
Ok(r)
}
}
fn get_ciphers(&self) -> CipherResult<CipherStates<C>> {
self.inner.get_ciphers()
}
fn get_hash(&self) -> <H as Hash>::Output {
self.inner.get_hash()
}
fn mix_hash(&mut self, data: &[u8]) {
if self.outer_is_finished {
self.inner.mix_hash(data);
} else {
self.outer.as_mut().unwrap().mix_hash(data);
}
}
fn mix_key_and_hash(&mut self, data: &[u8]) {
if self.outer_is_finished {
self.inner.mix_key_and_hash(data);
} else {
self.outer.as_mut().unwrap().mix_key_and_hash(data);
}
}
fn get_pattern(&self) -> crate::handshakepattern::HandshakePattern {
self.inner.get_pattern()
}
}
impl<Outer, Inner, C, H, const BUF: usize> Handshaker<C, H>
for DualLayerHandshake<Outer, Inner, C, H, BUF>
where
Inner: Handshaker<C, H>,
Outer: Handshaker<C, H>,
C: Cipher,
H: Hash,
{
type E = Inner::E;
type S = Inner::S;
fn push_psk(&mut self, _psk: &[u8]) {
panic!("Not applicable for dual-layer handshakes");
}
fn is_write_turn(&self) -> bool {
if self.outer_completed() {
self.inner.is_write_turn()
} else {
self.outer.as_ref().unwrap().is_write_turn()
}
}
fn is_initiator(&self) -> bool {
self.inner.is_initiator()
}
fn get_next_message_overhead(&self) -> HandshakeResult<usize> {
if self.outer_completed() {
self.inner.get_next_message_overhead()
} else {
self.outer.as_ref().unwrap().get_next_message_overhead()
}
}
fn build_name(_: &crate::handshakepattern::HandshakePattern) -> arrayvec::ArrayString<128> {
panic!("Not applicable for dual-layer handshakes");
}
fn get_remote_static(&self) -> Option<Self::S> {
self.inner.get_remote_static()
}
fn get_remote_ephemeral(&self) -> Option<Self::E> {
self.inner.get_remote_ephemeral()
}
fn get_state(&self) -> SymmetricState<C, H> {
self.inner.get_state()
}
fn get_state_mut(&mut self) -> &mut SymmetricState<C, H> {
self.inner.get_state_mut()
}
}