use crate::error::HandshakeResult;
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());
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
}
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) -> crate::cipherstate::CipherStates<C> {
self.inner.get_ciphers()
}
fn get_hash(&self) -> <H as Hash>::Output {
self.inner.get_hash()
}
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,
{
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");
}
}