use std::io::{self, Read, Write};
use std::sync::Arc;
use rustls::ServerConnection;
use super::consts::{KEY_EXPANSION_ID, KEY_METHOD_MASK};
use super::data;
use super::keys::PeerKeys;
use super::options::Options;
use super::prf::prf10;
use super::reliable::Reliable;
use super::window::Window;
use super::Opcode;
fn invalid(msg: impl Into<String>) -> io::Error {
io::Error::new(io::ErrorKind::InvalidData, msg.into())
}
#[derive(Debug, Clone)]
pub struct AuthInfo {
pub username: String,
pub password: String,
pub peer_info: std::collections::HashMap<String, String>,
pub dev_type: String,
}
#[derive(Debug, Clone)]
pub struct PeerConfig {
pub ip: std::net::IpAddr,
pub gateway: std::net::IpAddr,
pub mask: std::net::IpAddr,
pub prefix_len: u8,
}
pub type OnAuth = Arc<dyn Fn(&AuthInfo) -> io::Result<PeerConfig> + Send + Sync>;
#[derive(Default, Debug)]
pub struct PeerOutput {
pub send: Vec<Vec<u8>>,
pub deliver: Option<Vec<u8>>,
pub authenticated: bool,
pub close: bool,
}
#[derive(Debug, PartialEq, Eq)]
enum Phase {
Init,
Handshaking,
Established,
}
pub struct Peer {
tls: ServerConnection,
reliable: Reliable,
phase: Phase,
on_auth: OnAuth,
opts: Option<Options>,
keys: Option<PeerKeys>,
replay: Window,
layer: u8,
out_pid: u32,
ctrl_buf: Vec<u8>,
kx_done: bool,
peer_cfg: Option<PeerConfig>,
peer_info: std::collections::HashMap<String, String>,
server_random: [u8; 64],
}
impl std::fmt::Debug for Peer {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Peer")
.field("phase", &self.phase)
.field("authenticated", &self.kx_done)
.finish()
}
}
impl Peer {
pub fn new(
config: Arc<rustls::ServerConfig>,
local_id: [u8; 8],
on_auth: OnAuth,
) -> io::Result<Peer> {
let tls = ServerConnection::new(config)
.map_err(|e| invalid(format!("rustls server connection: {e}")))?;
Ok(Peer {
tls,
reliable: Reliable::new(local_id),
phase: Phase::Init,
on_auth,
opts: None,
keys: None,
replay: Window::new(),
layer: 3,
out_pid: 0,
ctrl_buf: Vec::new(),
kx_done: false,
peer_cfg: None,
peer_info: std::collections::HashMap::new(),
server_random: [0u8; 64],
})
}
pub fn peer_config(&self) -> Option<&PeerConfig> {
self.peer_cfg.as_ref()
}
pub fn layer(&self) -> u8 {
self.layer
}
pub fn tick(&mut self, now: std::time::Instant) -> io::Result<PeerOutput> {
let mut out = PeerOutput::default();
let tick = self.reliable.tick(now);
out.send = tick.resend;
out.close = tick.timed_out;
Ok(out)
}
pub fn handle_packet(&mut self, data: &[u8]) -> io::Result<PeerOutput> {
if data.is_empty() {
return Ok(PeerOutput::default());
}
let (opcode, _kid) = Opcode::from_byte(data[0]);
if opcode.is_control() {
self.handle_control(data)
} else if opcode == Opcode::DATA_V1 || opcode == Opcode::DATA_V2 {
self.handle_data(data)
} else {
Ok(PeerOutput::default())
}
}
fn handle_control(&mut self, data: &[u8]) -> io::Result<PeerOutput> {
let mut out = PeerOutput::default();
let recv = self.reliable.recv(data)?;
if recv.got_client_reset {
if self.phase == Phase::Init {
self.phase = Phase::Handshaking;
}
let reset = self.reliable.build_hard_reset();
let acks = self.reliable.take_pending_acks();
out.send.push(reset.to_bytes(&acks));
self.pump_tls(&mut out)?;
return Ok(out);
}
if !recv.tls_bytes.is_empty() {
let mut cursor = io::Cursor::new(&recv.tls_bytes);
while (cursor.position() as usize) < recv.tls_bytes.len() {
let n = self
.tls
.read_tls(&mut cursor)
.map_err(|e| invalid(format!("read_tls: {e}")))?;
if n == 0 {
break;
}
self.tls
.process_new_packets()
.map_err(|e| invalid(format!("tls process: {e}")))?;
}
let mut plain = Vec::new();
let _ = self.tls.reader().read_to_end(&mut plain);
if !plain.is_empty() {
self.ctrl_buf.extend_from_slice(&plain);
}
self.advance_control()?;
}
self.pump_tls(&mut out)?;
if self.kx_done {
out.authenticated = true;
}
Ok(out)
}
fn pump_tls(&mut self, out: &mut PeerOutput) -> io::Result<()> {
let mut tls_out = Vec::new();
while self.tls.wants_write() {
let n = self
.tls
.write_tls(&mut tls_out)
.map_err(|e| invalid(format!("write_tls: {e}")))?;
if n == 0 {
break;
}
}
if !tls_out.is_empty() {
let chunks = self.reliable.chunk_tls_stream(&tls_out);
for (i, pkt) in chunks.iter().enumerate() {
let acks = if i == 0 {
self.reliable.take_pending_acks()
} else {
Vec::new()
};
out.send.push(pkt.to_bytes(&acks));
}
}
if self.reliable.has_pending_acks() {
let acks = self.reliable.take_pending_acks();
let ack = self.reliable.build_ack();
out.send.push(ack.to_bytes(&acks));
}
Ok(())
}
fn advance_control(&mut self) -> io::Result<()> {
if self.kx_done {
self.handle_post_auth_control()?;
return Ok(());
}
let parsed = match self.try_parse_key_exchange()? {
Some(p) => p,
None => return Ok(()), };
fill_random(&mut self.server_random)?;
let reply = self.build_kx_reply(&parsed)?;
self.tls
.writer()
.write_all(&reply)
.map_err(|e| invalid(format!("tls write reply: {e}")))?;
self.derive_keys(&parsed)?;
let auth = AuthInfo {
username: parsed.username.clone(),
password: parsed.password.clone(),
peer_info: parsed.peer_info.clone(),
dev_type: parsed.opts.dev_type.clone(),
};
let cfg = (self.on_auth)(&auth)?;
self.peer_cfg = Some(cfg);
self.layer = match parsed.opts.dev_type.as_str() {
"tap" => 2,
_ => 3,
};
self.peer_info = parsed.peer_info;
self.opts = Some(parsed.opts);
self.kx_done = true;
self.phase = Phase::Established;
Ok(())
}
fn handle_post_auth_control(&mut self) -> io::Result<()> {
while let Some(nul) = self.ctrl_buf.iter().position(|&b| b == 0) {
let line: Vec<u8> = self.ctrl_buf.drain(..=nul).collect();
let body = &line[..line.len() - 1];
if body.is_empty() {
continue;
}
let s = String::from_utf8_lossy(body).into_owned();
let verb = s.split(',').next().unwrap_or("");
if verb == "PUSH_REQUEST" {
let reply = self.build_push_reply();
self.tls
.writer()
.write_all(reply.as_bytes())
.map_err(|e| invalid(format!("tls push reply: {e}")))?;
}
}
Ok(())
}
pub fn peer_info(&self) -> &std::collections::HashMap<String, String> {
&self.peer_info
}
fn build_push_reply(&self) -> String {
let cfg = self.peer_cfg.as_ref();
let (ip, gw_or_mask) = match cfg {
Some(c) if self.layer == 2 => (c.ip.to_string(), c.mask.to_string()),
Some(c) => (c.ip.to_string(), c.gateway.to_string()),
None => ("0.0.0.0".to_string(), "0.0.0.0".to_string()),
};
format!(
"PUSH_REPLY,ping 10,comp-lzo no,topology net30,ifconfig {} {}\0",
ip, gw_or_mask
)
}
fn try_parse_key_exchange(&self) -> io::Result<Option<KeyExchange>> {
let buf = &self.ctrl_buf;
let fixed = 4 + 1 + 48 + 32 + 32;
if buf.len() < fixed {
return Ok(None);
}
let zero = u32::from_be_bytes([buf[0], buf[1], buf[2], buf[3]]);
if zero != 0 {
return Err(invalid("control channel: expected 4 zero bytes"));
}
let key_method = buf[4];
if key_method & KEY_METHOD_MASK != 2 {
return Err(invalid("invalid key method, expected method 2"));
}
let mut pos = 5;
let mut pre_master = [0u8; 48];
pre_master.copy_from_slice(&buf[pos..pos + 48]);
pos += 48;
let mut random1 = [0u8; 32];
random1.copy_from_slice(&buf[pos..pos + 32]);
pos += 32;
let mut random2 = [0u8; 32];
random2.copy_from_slice(&buf[pos..pos + 32]);
pos += 32;
let (options_string, p1) = match read_control_string(buf, pos)? {
Some(v) => v,
None => return Ok(None),
};
let (username, p2) = match read_control_string(buf, p1)? {
Some(v) => v,
None => return Ok(None),
};
let (password, p3) = match read_control_string(buf, p2)? {
Some(v) => v,
None => return Ok(None),
};
let (peer_info_raw, _p4) = match read_control_string(buf, p3)? {
Some(v) => v,
None => return Ok(None),
};
let mut opts = Options::parse(&options_string).map_err(invalid)?;
opts.is_server = false;
if opts.to_string() != options_string {
return Err(invalid("invalid options provided"));
}
let peer_info = parse_peer_info(&peer_info_raw)?;
let options_server = {
let mut o = opts.clone();
o.is_server = true;
o.to_string()
};
let _ = options_string;
Ok(Some(KeyExchange {
pre_master,
random1,
random2,
options_server,
opts,
username,
password,
peer_info,
peer_info_raw,
}))
}
fn build_kx_reply(&self, kx: &KeyExchange) -> io::Result<Vec<u8>> {
let mut buf = Vec::new();
buf.extend_from_slice(&0u32.to_be_bytes());
buf.push(2u8);
buf.extend_from_slice(&self.server_random);
write_control_string(&mut buf, &kx.options_server);
write_control_string(&mut buf, ""); write_control_string(&mut buf, ""); write_control_string(&mut buf, &kx.peer_info_raw);
Ok(buf)
}
fn derive_keys(&mut self, kx: &KeyExchange) -> io::Result<()> {
let (sr1, sr2) = self.server_random.split_at(32);
let mut master = [0u8; 48];
let mut seed = Vec::with_capacity(64);
seed.extend_from_slice(&kx.random1);
seed.extend_from_slice(sr1);
let label = format!("{} master secret", KEY_EXPANSION_ID);
prf10(&mut master, &kx.pre_master, label.as_bytes(), &seed);
let mut expansion = [0u8; 256];
let mut seed2 = Vec::with_capacity(32 + 32 + 8 + 8);
seed2.extend_from_slice(&kx.random2);
seed2.extend_from_slice(sr2);
seed2.extend_from_slice(&self.reliable.peer_id);
seed2.extend_from_slice(&self.reliable.local_id);
let label2 = format!("{} key expansion", KEY_EXPANSION_ID);
prf10(&mut expansion, &master, label2.as_bytes(), &seed2);
self.keys = Some(PeerKeys::from_expansion(&expansion));
Ok(())
}
fn handle_data(&mut self, data: &[u8]) -> io::Result<PeerOutput> {
let mut out = PeerOutput::default();
let (Some(opts), Some(keys)) = (self.opts.as_ref(), self.keys.as_ref()) else {
return Err(invalid("stream not ready for data transmission"));
};
let mut buf = data.to_vec();
if let Some(dec) = data::decrypt(opts, keys, &mut buf)? {
if !self.replay.check(dec.pid) {
return Ok(out); }
if dec.is_ping {
return Ok(out);
}
out.deliver = Some(dec.payload.to_vec());
}
Ok(out)
}
pub fn send_data(&mut self, payload: &[u8]) -> io::Result<Vec<u8>> {
let (Some(opts), Some(keys)) = (self.opts.as_ref(), self.keys.as_ref()) else {
return Err(invalid("stream not ready for data transmission"));
};
self.out_pid = self.out_pid.wrapping_add(1);
data::encrypt(opts, keys, self.out_pid, payload, fill_random)
}
}
struct KeyExchange {
pre_master: [u8; 48],
random1: [u8; 32],
random2: [u8; 32],
options_server: String,
opts: Options,
username: String,
password: String,
peer_info: std::collections::HashMap<String, String>,
peer_info_raw: String,
}
fn read_control_string(buf: &[u8], pos: usize) -> io::Result<Option<(String, usize)>> {
if pos + 2 > buf.len() {
return Ok(None);
}
let len = u16::from_be_bytes([buf[pos], buf[pos + 1]]) as usize;
if len == 0 {
return Err(invalid("empty control string"));
}
let start = pos + 2;
if start + len > buf.len() {
return Ok(None);
}
let raw = &buf[start..start + len];
if raw[len - 1] != 0 {
return Err(invalid("control string not NUL-terminated"));
}
let s = String::from_utf8_lossy(&raw[..len - 1]).into_owned();
Ok(Some((s, start + len)))
}
fn write_control_string(buf: &mut Vec<u8>, s: &str) {
let len = s.len() + 1;
buf.extend_from_slice(&(len as u16).to_be_bytes());
buf.extend_from_slice(s.as_bytes());
buf.push(0);
}
fn parse_peer_info(raw: &str) -> io::Result<std::collections::HashMap<String, String>> {
let mut peer_info = std::collections::HashMap::new();
for line in raw.split('\n') {
let line = line.strip_suffix('\r').unwrap_or(line);
if line.is_empty() {
continue;
}
match line.find('=') {
Some(i) => {
peer_info.insert(line[..i].to_string(), line[i + 1..].to_string());
}
None => return Err(invalid("invalid string in peer_info")),
}
}
Ok(peer_info)
}
pub(crate) fn fill_random(buf: &mut [u8]) -> io::Result<()> {
getrandom::getrandom(buf).map_err(|e| io::Error::other(format!("getrandom: {e}")))
}
#[cfg(test)]
mod tests {
use super::parse_peer_info;
#[test]
fn peer_info_parses_iv_keys() {
let raw = "IV_VER=2.6.0\nIV_PLAT=linux\nIV_PROTO=6\nIV_CIPHERS=AES-256-GCM:AES-128-GCM\n";
let pi = parse_peer_info(raw).unwrap();
assert_eq!(pi.get("IV_VER").map(String::as_str), Some("2.6.0"));
assert_eq!(pi.get("IV_PLAT").map(String::as_str), Some("linux"));
assert_eq!(pi.get("IV_PROTO").map(String::as_str), Some("6"));
assert_eq!(
pi.get("IV_CIPHERS").map(String::as_str),
Some("AES-256-GCM:AES-128-GCM")
);
assert_eq!(pi.len(), 4);
}
#[test]
fn peer_info_value_may_contain_equals() {
let pi = parse_peer_info("UV_OPT=a=b=c\n").unwrap();
assert_eq!(pi.get("UV_OPT").map(String::as_str), Some("a=b=c"));
}
#[test]
fn peer_info_skips_blank_and_crlf_lines() {
let pi = parse_peer_info("\nIV_VER=2.6\r\n\n").unwrap();
assert_eq!(pi.len(), 1);
assert_eq!(pi.get("IV_VER").map(String::as_str), Some("2.6"));
}
#[test]
fn peer_info_rejects_line_without_equals() {
assert!(parse_peer_info("IV_VER=2.6\nGARBAGE\n").is_err());
}
#[test]
fn peer_info_empty_is_ok() {
let pi = parse_peer_info("").unwrap();
assert!(pi.is_empty());
}
}