use std::collections::HashMap;
use std::io;
use std::net::{SocketAddr, ToSocketAddrs};
use std::time::Duration;
use crate::entropy::Entropy;
use crate::esp::ChildSa;
use crate::ikev1::isakmp::{self, exchange, payload, IsakmpHeader};
use crate::ikev1::phase1::{respond_aggressive, Phase1Config, Phase1State};
use crate::ikev1::quick::{respond_quick, QuickResponder};
use crate::transport::{DriverError, UdpTransport};
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum ServerEvent {
Phase1SaInit,
Phase1Established,
QuickSaInit,
ChildSaEstablished { cky_i: [u8; 8] },
Ignored,
}
struct Session {
phase1: Option<Phase1State>,
quick: Option<QuickResponder>,
}
pub struct Server<E> {
transport: UdpTransport,
entropy: E,
cfg: Phase1Config,
sessions: HashMap<[u8; 8], Session>,
children: HashMap<[u8; 8], ChildSa>,
}
impl<E: Entropy> Server<E> {
pub fn bind(addr: impl ToSocketAddrs, entropy: E, cfg: Phase1Config) -> io::Result<Self> {
Ok(Self {
transport: UdpTransport::bind(addr)?,
entropy,
cfg,
sessions: HashMap::new(),
children: HashMap::new(),
})
}
pub fn local_addr(&self) -> io::Result<SocketAddr> {
self.transport.local_addr()
}
pub fn set_read_timeout(&self, dur: Option<Duration>) -> io::Result<()> {
self.transport.set_read_timeout(dur)
}
pub fn child(&self, cky_i: [u8; 8]) -> Option<&ChildSa> {
self.children.get(&cky_i)
}
pub fn take_child(&mut self, cky_i: [u8; 8]) -> Option<ChildSa> {
self.children.remove(&cky_i)
}
pub fn handle_one(&mut self) -> Result<ServerEvent, DriverError> {
let (data, from) = self.transport.recv_from()?;
let hdr = IsakmpHeader::parse(&data)?;
let cky_i = hdr.init_cookie;
match hdr.exchange_type {
exchange::AGGRESSIVE => {
let ps = isakmp::parse_payloads(hdr.next_payload, &data[IsakmpHeader::LEN..])?;
let has = |t: u8| ps.iter().any(|p| p.payload_type == t);
if has(payload::SA) {
let (msg2, st) = respond_aggressive(&self.cfg, &data, &mut self.entropy)?;
self.transport.send_to(&msg2, from)?;
self.sessions.insert(cky_i, Session { phase1: Some(st), quick: None });
Ok(ServerEvent::Phase1SaInit)
} else if has(payload::HASH) {
let Some(st) = self.sessions.get(&cky_i).and_then(|s| s.phase1.as_ref()) else {
return Ok(ServerEvent::Ignored);
};
st.verify_hash_i(&data)?;
Ok(ServerEvent::Phase1Established)
} else {
Ok(ServerEvent::Ignored)
}
}
exchange::QUICK => {
let (has_phase1, quick_started) = match self.sessions.get(&cky_i) {
Some(s) => (s.phase1.is_some(), s.quick.is_some()),
None => return Ok(ServerEvent::Ignored),
};
if !has_phase1 {
return Ok(ServerEvent::Ignored);
}
if !quick_started {
let st = self.sessions.get(&cky_i).unwrap().phase1.clone().unwrap();
let (msg2, qr) = respond_quick(&st, &data, &mut self.entropy)?;
self.transport.send_to(&msg2, from)?;
self.sessions.get_mut(&cky_i).unwrap().quick = Some(qr);
Ok(ServerEvent::QuickSaInit)
} else {
let qr = self.sessions.get_mut(&cky_i).unwrap().quick.take().unwrap();
let child = qr.complete(&data)?;
self.children.insert(cky_i, child);
Ok(ServerEvent::ChildSaEstablished { cky_i })
}
}
_ => Ok(ServerEvent::Ignored),
}
}
}