use super::*;
use std::net::{SocketAddrV6, UdpSocket};
use tokio::time::sleep;
#[cfg(feature = "dhcp-debug")]
macro_rules! dhcpv6_debug {
($($arg:tt)*) => {
debug!("[{}:{}] {}", file!(), line!(), format!($($arg)*))
};
}
#[cfg(not(feature = "dhcp-debug"))]
macro_rules! dhcpv6_debug {
($($arg:tt)*) => {};
}
#[cfg(feature = "dhcp-debug")]
macro_rules! dhcpv6_info {
($($arg:tt)*) => {
info!("[{}:{}] {}", file!(), line!(), format!($($arg)*))
};
}
#[cfg(not(feature = "dhcp-debug"))]
macro_rules! dhcpv6_info {
($($arg:tt)*) => {};
}
const DHCPV6_CLIENT_PORT: u16 = 546;
const DHCPV6_SERVER_PORT: u16 = 547;
const ALL_DHCP_RELAY_AGENTS_AND_SERVERS: Ipv6Addr = Ipv6Addr::new(0xff02, 0, 0, 0, 0, 0, 1, 2);
const DHCPV6_SOLICIT: u8 = 1;
const DHCPV6_ADVERTISE: u8 = 2;
const DHCPV6_REQUEST: u8 = 3;
const DHCPV6_CONFIRM: u8 = 4;
const DHCPV6_RENEW: u8 = 5;
const DHCPV6_REBIND: u8 = 6;
const DHCPV6_REPLY: u8 = 7;
const DHCPV6_RELEASE: u8 = 8;
const DHCPV6_DECLINE: u8 = 9;
const DHCPV6_INFORMATION_REQUEST: u8 = 11;
const OPTION_CLIENTID: u16 = 1;
const OPTION_SERVERID: u16 = 2;
const OPTION_IA_NA: u16 = 3;
const OPTION_IA_TA: u16 = 4;
const OPTION_IAADDR: u16 = 5;
const OPTION_ORO: u16 = 6; const OPTION_PREFERENCE: u16 = 7;
const OPTION_ELAPSED_TIME: u16 = 8;
const OPTION_STATUS_CODE: u16 = 13;
const OPTION_RAPID_COMMIT: u16 = 14;
const OPTION_DNS_SERVERS: u16 = 23;
const OPTION_DOMAIN_LIST: u16 = 24;
const OPTION_IA_PD: u16 = 25;
const OPTION_IAPREFIX: u16 = 26;
const STATUS_SUCCESS: u16 = 0;
const STATUS_UNSPEC_FAIL: u16 = 1;
const STATUS_NO_ADDRS_AVAIL: u16 = 2;
const STATUS_NO_BINDING: u16 = 3;
const STATUS_NOT_ON_LINK: u16 = 4;
const STATUS_USE_MULTICAST: u16 = 5;
const STATUS_NO_PREFIX_AVAIL: u16 = 6;
const OPTION_AUTH: u16 = 11;
const DHCPV6_STARTTLS: u8 = 23;
const AUTH_PROTOCOL_DELAYED: u8 = 2; const AUTH_PROTOCOL_RECONFIGURE_KEY: u8 = 3;
const AUTH_ALGORITHM_HMAC_MD5: u8 = 1; const AUTH_ALGORITHM_HMAC_SHA1: u8 = 2; const AUTH_ALGORITHM_HMAC_SHA256: u8 = 3; const AUTH_ALGORITHM_HMAC_SHA384: u8 = 4; const AUTH_ALGORITHM_HMAC_SHA512: u8 = 5;
const RDM_MONOTONIC_COUNTER: u8 = 0;
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum Dhcpv6State {
Init,
Solicit,
Request,
Confirm,
Renew,
Rebind,
Bound,
}
impl std::fmt::Display for Dhcpv6State {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Init => write!(f, "INIT"),
Self::Solicit => write!(f, "SOLICIT"),
Self::Request => write!(f, "REQUEST"),
Self::Confirm => write!(f, "CONFIRM"),
Self::Renew => write!(f, "RENEW"),
Self::Rebind => write!(f, "REBIND"),
Self::Bound => write!(f, "BOUND"),
}
}
}
#[derive(Debug, Clone)]
pub struct Dhcpv6IaNa {
pub iaid: u32,
pub addresses: Vec<Ipv6Addr>,
pub t1: Duration, pub t2: Duration, pub preferred_lifetime: Duration,
pub valid_lifetime: Duration,
pub acquired_at: SystemTime,
}
impl Dhcpv6IaNa {
pub fn to_info(&self) -> Dhcpv6IaNaInfo {
Dhcpv6IaNaInfo {
addresses: self.addresses.clone(),
preferred_lifetime: self.preferred_lifetime,
valid_lifetime: self.valid_lifetime,
acquired_at: self.acquired_at,
}
}
pub fn renewal_at(&self) -> SystemTime {
self.acquired_at + self.t1
}
pub fn rebinding_at(&self) -> SystemTime {
self.acquired_at + self.t2
}
pub fn expires_at(&self) -> SystemTime {
self.acquired_at + self.valid_lifetime
}
pub fn is_expired(&self) -> bool {
SystemTime::now() > self.expires_at()
}
}
#[derive(Debug, Clone)]
pub struct Dhcpv6IaPd {
pub iaid: u32,
pub prefix: Ipv6Addr,
pub prefix_len: u8,
pub t1: Duration,
pub t2: Duration,
pub preferred_lifetime: Duration,
pub valid_lifetime: Duration,
pub acquired_at: SystemTime,
}
impl Dhcpv6IaPd {
pub fn to_info(&self) -> Dhcpv6IaPdInfo {
Dhcpv6IaPdInfo {
prefix: self.prefix,
prefix_len: self.prefix_len,
preferred_lifetime: self.preferred_lifetime,
valid_lifetime: self.valid_lifetime,
acquired_at: self.acquired_at,
}
}
}
#[derive(Debug, Clone)]
struct Dhcpv6ServerInfo {
server_id: Vec<u8>,
address: Ipv6Addr,
preference: u8,
ia_na: Option<Dhcpv6IaNa>,
ia_pd: Option<Dhcpv6IaPd>,
dns_servers: Vec<Ipv6Addr>,
domain_list: Vec<String>,
}
#[derive(Debug, Clone)]
pub struct Dhcpv6AuthContext {
protocol: u8,
algorithm: u8,
key: Option<Vec<u8>>,
key_id: Option<u32>,
realm: Option<String>,
replay_counter: u64,
server_replay_counter: Option<u64>,
enabled: bool,
require_server_auth: bool,
accept_unauthenticated: bool,
}
impl Dhcpv6AuthContext {
pub fn from_config(config: &Dhcpv6AuthConfig) -> Self {
let protocol = match config.protocol.as_str() {
"reconfigure-key" => AUTH_PROTOCOL_RECONFIGURE_KEY,
_ => AUTH_PROTOCOL_DELAYED,
};
let algorithm = match config.algorithm.as_str() {
"hmac-md5" => AUTH_ALGORITHM_HMAC_MD5,
"hmac-sha1" => AUTH_ALGORITHM_HMAC_SHA1,
"hmac-sha384" => AUTH_ALGORITHM_HMAC_SHA384,
"hmac-sha512" => AUTH_ALGORITHM_HMAC_SHA512,
_ => AUTH_ALGORITHM_HMAC_SHA256,
};
let key = config.key.as_ref().and_then(|k| {
if let Ok(decoded) = hex::decode(k) {
return Some(decoded);
}
use base64::Engine;
if let Ok(decoded) = base64::engine::general_purpose::STANDARD.decode(k) {
return Some(decoded);
}
if !k.is_empty() {
return Some(k.as_bytes().to_vec());
}
None
});
let replay_counter = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_secs())
.unwrap_or(0);
Self {
protocol,
algorithm,
key,
key_id: config.key_id,
realm: config.realm.clone(),
replay_counter,
server_replay_counter: None,
enabled: config.enabled,
require_server_auth: config.require_server_auth,
accept_unauthenticated: config.accept_unauthenticated,
}
}
fn next_replay_counter(&mut self) -> u64 {
self.replay_counter += 1;
self.replay_counter
}
fn hmac_length(&self) -> usize {
match self.algorithm {
AUTH_ALGORITHM_HMAC_MD5 => 16, AUTH_ALGORITHM_HMAC_SHA1 => 20, AUTH_ALGORITHM_HMAC_SHA256 => 32, AUTH_ALGORITHM_HMAC_SHA384 => 48, AUTH_ALGORITHM_HMAC_SHA512 => 64, _ => 32,
}
}
fn compute_hmac(&self, message: &[u8]) -> Option<Vec<u8>> {
use ring::hmac;
let key = self.key.as_ref()?;
let algorithm = match self.algorithm {
AUTH_ALGORITHM_HMAC_SHA256 => hmac::HMAC_SHA256,
AUTH_ALGORITHM_HMAC_SHA384 => hmac::HMAC_SHA384,
AUTH_ALGORITHM_HMAC_SHA512 => hmac::HMAC_SHA512,
AUTH_ALGORITHM_HMAC_SHA1 => {
warn!("HMAC-SHA1 not directly supported by ring, using SHA256");
hmac::HMAC_SHA256
}
AUTH_ALGORITHM_HMAC_MD5 => {
warn!("HMAC-MD5 is deprecated and not supported, using SHA256");
hmac::HMAC_SHA256
}
_ => hmac::HMAC_SHA256,
};
let signing_key = hmac::Key::new(algorithm, key);
let tag = hmac::sign(&signing_key, message);
Some(tag.as_ref().to_vec())
}
fn verify_hmac(&self, message: &[u8], received_hmac: &[u8]) -> bool {
use ring::hmac;
let key = match self.key.as_ref() {
Some(k) => k,
None => return false,
};
let algorithm = match self.algorithm {
AUTH_ALGORITHM_HMAC_SHA256 => hmac::HMAC_SHA256,
AUTH_ALGORITHM_HMAC_SHA384 => hmac::HMAC_SHA384,
AUTH_ALGORITHM_HMAC_SHA512 => hmac::HMAC_SHA512,
_ => hmac::HMAC_SHA256,
};
let verification_key = hmac::Key::new(algorithm, key);
hmac::verify(&verification_key, message, received_hmac).is_ok()
}
fn build_auth_option(&mut self, message_without_auth: &[u8]) -> Option<Vec<u8>> {
if !self.enabled || self.key.is_none() {
return None;
}
let mut auth_data = Vec::new();
auth_data.push(self.protocol);
auth_data.push(self.algorithm);
auth_data.push(RDM_MONOTONIC_COUNTER);
let replay = self.next_replay_counter();
auth_data.extend_from_slice(&replay.to_be_bytes());
if self.protocol == AUTH_PROTOCOL_DELAYED {
if let Some(ref realm) = self.realm {
auth_data.extend_from_slice(realm.as_bytes());
auth_data.push(0); } else {
auth_data.push(0); }
let key_id = self.key_id.unwrap_or(0);
auth_data.extend_from_slice(&key_id.to_be_bytes());
let hmac_len = self.hmac_length();
let zeros = vec![0u8; hmac_len];
let mut temp_auth = auth_data.clone();
temp_auth.extend_from_slice(&zeros);
let mut full_msg = message_without_auth.to_vec();
full_msg.extend_from_slice(&OPTION_AUTH.to_be_bytes());
full_msg.extend_from_slice(&(temp_auth.len() as u16).to_be_bytes());
full_msg.extend_from_slice(&temp_auth);
if let Some(hmac) = self.compute_hmac(&full_msg) {
let hmac_to_use = if hmac.len() > hmac_len {
&hmac[..hmac_len]
} else {
&hmac
};
auth_data.extend_from_slice(hmac_to_use);
} else {
return None;
}
} else if self.protocol == AUTH_PROTOCOL_RECONFIGURE_KEY {
let key_id = self.key_id.unwrap_or(0);
auth_data.extend_from_slice(&key_id.to_be_bytes());
if let Some(hmac) = self.compute_hmac(message_without_auth) {
auth_data.extend_from_slice(&hmac);
} else {
return None;
}
}
Some(auth_data)
}
fn parse_and_verify_auth(&mut self, message: &[u8], auth_option_data: &[u8]) -> std::result::Result<bool, String> {
if auth_option_data.len() < 11 {
return Err("Authentication option too short".to_string());
}
let protocol = auth_option_data[0];
let algorithm = auth_option_data[1];
let rdm = auth_option_data[2];
let replay_bytes: [u8; 8] = match auth_option_data[3..11].try_into() {
Ok(b) => b,
Err(_) => return Err("Invalid replay detection field".to_string()),
};
let replay_counter = u64::from_be_bytes(replay_bytes);
if protocol != self.protocol {
return Err(format!("Protocol mismatch: expected {}, got {}", self.protocol, protocol));
}
if algorithm != self.algorithm {
return Err(format!("Algorithm mismatch: expected {}, got {}", self.algorithm, algorithm));
}
if rdm != RDM_MONOTONIC_COUNTER {
return Err(format!("Unsupported RDM: {}", rdm));
}
if let Some(last_counter) = self.server_replay_counter {
if replay_counter <= last_counter {
return Err(format!("Replay attack detected: {} <= {}", replay_counter, last_counter));
}
}
let auth_info = &auth_option_data[11..];
if self.protocol == AUTH_PROTOCOL_DELAYED {
let realm_end = match auth_info.iter().position(|&b| b == 0) {
Some(pos) => pos,
None => return Err("No realm terminator found".to_string()),
};
let _realm = &auth_info[..realm_end];
let after_realm = &auth_info[realm_end + 1..];
if after_realm.len() < 4 {
return Err("Missing key ID".to_string());
}
let key_id_bytes: [u8; 4] = match after_realm[..4].try_into() {
Ok(b) => b,
Err(_) => return Err("Invalid key ID".to_string()),
};
let _key_id = u32::from_be_bytes(key_id_bytes);
let received_hmac = &after_realm[4..];
let hmac_len = received_hmac.len();
let auth_opt_pos = match self.find_option_position(message, OPTION_AUTH) {
Some(pos) => pos,
None => return Err("Cannot find auth option in message".to_string()),
};
let mut verify_msg = message.to_vec();
let hmac_start = auth_opt_pos + 4 + 11 + realm_end + 1 + 4; if hmac_start + hmac_len <= verify_msg.len() {
for i in 0..hmac_len {
verify_msg[hmac_start + i] = 0;
}
}
if !self.verify_hmac(&verify_msg, received_hmac) {
return Err("HMAC verification failed".to_string());
}
} else {
let received_hmac = auth_info;
if !self.verify_hmac(message, received_hmac) {
return Err("HMAC verification failed".to_string());
}
}
self.server_replay_counter = Some(replay_counter);
Ok(true)
}
fn find_option_position(&self, message: &[u8], option_code: u16) -> Option<usize> {
if message.len() < 4 {
return None;
}
let mut pos = 4;
while pos + 4 <= message.len() {
let opt_code = u16::from_be_bytes([message[pos], message[pos + 1]]);
let opt_len = u16::from_be_bytes([message[pos + 2], message[pos + 3]]) as usize;
if opt_code == option_code {
return Some(pos);
}
pos += 4 + opt_len;
}
None
}
}
pub use config::Dhcpv6AuthConfig;
pub use config::Dhcpv6TlsConfig;
pub struct Dhcpv6TlsContext {
config: Dhcpv6TlsConfig,
tls_config: Option<Arc<rustls::ClientConfig>>,
tls_established: bool,
}
impl std::fmt::Debug for Dhcpv6TlsContext {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Dhcpv6TlsContext")
.field("enabled", &self.config.enabled)
.field("tls_established", &self.tls_established)
.finish()
}
}
impl Dhcpv6TlsContext {
pub fn from_config(config: &Dhcpv6TlsConfig) -> Self {
Self {
config: config.clone(),
tls_config: None,
tls_established: false,
}
}
pub fn is_enabled(&self) -> bool {
self.config.enabled
}
pub fn use_tcp(&self) -> bool {
self.config.use_tcp || self.config.enabled
}
pub fn is_tls_established(&self) -> bool {
self.tls_established
}
pub fn init_tls_config(&mut self) -> Result<()> {
use rustls::ClientConfig;
use std::fs::File;
use std::io::BufReader;
if self.tls_config.is_some() {
return Ok(());
}
let mut root_store = rustls::RootCertStore::empty();
if let Some(ref ca_path) = self.config.ca_cert {
let ca_file = File::open(ca_path)
.map_err(|e| DhcpClientError::InvalidConfig(
format!("Failed to open CA cert file: {}", e)
))?;
let mut ca_reader = BufReader::new(ca_file);
let certs = rustls_pemfile::certs(&mut ca_reader)
.collect::<std::result::Result<Vec<_>, _>>()
.map_err(|e| DhcpClientError::InvalidConfig(
format!("Failed to parse CA certs: {}", e)
))?;
for cert in certs {
root_store.add(cert)
.map_err(|e| DhcpClientError::InvalidConfig(
format!("Failed to add CA cert: {}", e)
))?;
}
} else {
root_store.extend(webpki_roots::TLS_SERVER_ROOTS.iter().cloned());
}
let builder = ClientConfig::builder()
.with_root_certificates(root_store);
let tls_config = if let (Some(ref cert_path), Some(ref key_path)) =
(&self.config.client_cert, &self.config.client_key)
{
let cert_file = File::open(cert_path)
.map_err(|e| DhcpClientError::InvalidConfig(
format!("Failed to open client cert: {}", e)
))?;
let mut cert_reader = BufReader::new(cert_file);
let certs = rustls_pemfile::certs(&mut cert_reader)
.collect::<std::result::Result<Vec<_>, _>>()
.map_err(|e| DhcpClientError::InvalidConfig(
format!("Failed to parse client certs: {}", e)
))?;
let key_file = File::open(key_path)
.map_err(|e| DhcpClientError::InvalidConfig(
format!("Failed to open client key: {}", e)
))?;
let mut key_reader = BufReader::new(key_file);
let key = rustls_pemfile::private_key(&mut key_reader)
.map_err(|e| DhcpClientError::InvalidConfig(
format!("Failed to parse client key: {}", e)
))?
.ok_or_else(|| DhcpClientError::InvalidConfig(
"No private key found in key file".to_string()
))?;
builder.with_client_auth_cert(certs, key)
.map_err(|e| DhcpClientError::InvalidConfig(
format!("Failed to configure client auth: {}", e)
))?
} else {
builder.with_no_client_auth()
};
self.tls_config = Some(Arc::new(tls_config));
Ok(())
}
pub fn get_tls_config(&self) -> Option<Arc<rustls::ClientConfig>> {
self.tls_config.clone()
}
pub fn set_tls_established(&mut self, established: bool) {
self.tls_established = established;
}
pub fn build_starttls_message(xid: &[u8; 3]) -> Vec<u8> {
let mut msg = Vec::with_capacity(4);
msg.push(DHCPV6_STARTTLS);
msg.extend_from_slice(xid);
msg
}
pub fn parse_starttls_response(data: &[u8], expected_xid: &[u8; 3]) -> bool {
if data.len() < 4 {
return false;
}
if data[0] != DHCPV6_STARTTLS {
return false;
}
&data[1..4] == expected_xid
}
}
pub struct Dhcpv6TcpConnection {
stream: Option<tokio::net::TcpStream>,
tls_stream: Option<tokio_rustls::client::TlsStream<tokio::net::TcpStream>>,
server_addr: std::net::SocketAddrV6,
}
impl Dhcpv6TcpConnection {
pub async fn connect(server_addr: std::net::SocketAddrV6) -> Result<Self> {
let stream = tokio::net::TcpStream::connect(server_addr).await
.map_err(|e| DhcpClientError::Network(e))?;
Ok(Self {
stream: Some(stream),
tls_stream: None,
server_addr,
})
}
pub async fn starttls(
&mut self,
tls_config: Arc<rustls::ClientConfig>,
server_name: &str,
) -> Result<()> {
use tokio_rustls::TlsConnector;
let stream = self.stream.take()
.ok_or_else(|| DhcpClientError::InvalidConfig(
"No TCP stream available for TLS upgrade".to_string()
))?;
let server_name = rustls::pki_types::ServerName::try_from(server_name.to_string())
.map_err(|_| DhcpClientError::InvalidConfig(
format!("Invalid server name for TLS: {}", server_name)
))?;
let connector = TlsConnector::from(tls_config);
let tls_stream = connector.connect(server_name, stream).await
.map_err(|e| DhcpClientError::Network(std::io::Error::new(
std::io::ErrorKind::ConnectionRefused,
format!("TLS handshake failed: {}", e)
)))?;
self.tls_stream = Some(tls_stream);
info!("DHCPv6 TLS connection established");
Ok(())
}
pub async fn send(&mut self, message: &[u8]) -> Result<()> {
use tokio::io::AsyncWriteExt;
let len = message.len() as u16;
let mut buf = Vec::with_capacity(2 + message.len());
buf.extend_from_slice(&len.to_be_bytes());
buf.extend_from_slice(message);
if let Some(ref mut tls_stream) = self.tls_stream {
tls_stream.write_all(&buf).await
.map_err(|e| DhcpClientError::Network(e))?;
tls_stream.flush().await
.map_err(|e| DhcpClientError::Network(e))?;
} else if let Some(ref mut stream) = self.stream {
stream.write_all(&buf).await
.map_err(|e| DhcpClientError::Network(e))?;
stream.flush().await
.map_err(|e| DhcpClientError::Network(e))?;
} else {
return Err(DhcpClientError::Network(std::io::Error::new(
std::io::ErrorKind::NotConnected,
"No connection available"
)));
}
Ok(())
}
pub async fn recv(&mut self, timeout: Duration) -> Result<Vec<u8>> {
use tokio::io::AsyncReadExt;
use tokio::time::timeout as tokio_timeout;
let mut len_buf = [0u8; 2];
let read_result = if let Some(ref mut tls_stream) = self.tls_stream {
tokio_timeout(timeout, tls_stream.read_exact(&mut len_buf)).await
} else if let Some(ref mut stream) = self.stream {
tokio_timeout(timeout, stream.read_exact(&mut len_buf)).await
} else {
return Err(DhcpClientError::Network(std::io::Error::new(
std::io::ErrorKind::NotConnected,
"No connection available"
)));
};
match read_result {
Ok(Ok(_)) => {}
Ok(Err(e)) => return Err(DhcpClientError::Network(e)),
Err(_) => return Err(DhcpClientError::Timeout("TCP receive timeout".to_string())),
}
let msg_len = u16::from_be_bytes(len_buf) as usize;
if msg_len > 65535 {
return Err(DhcpClientError::InvalidMessage(
format!("Message too large: {} bytes", msg_len)
));
}
let mut msg_buf = vec![0u8; msg_len];
let read_result = if let Some(ref mut tls_stream) = self.tls_stream {
tokio_timeout(timeout, tls_stream.read_exact(&mut msg_buf)).await
} else if let Some(ref mut stream) = self.stream {
tokio_timeout(timeout, stream.read_exact(&mut msg_buf)).await
} else {
unreachable!()
};
match read_result {
Ok(Ok(_)) => Ok(msg_buf),
Ok(Err(e)) => Err(DhcpClientError::Network(e)),
Err(_) => Err(DhcpClientError::Timeout("TCP receive timeout".to_string())),
}
}
pub fn is_tls(&self) -> bool {
self.tls_stream.is_some()
}
pub async fn close(&mut self) {
use tokio::io::AsyncWriteExt;
if let Some(ref mut tls_stream) = self.tls_stream {
let _ = tls_stream.shutdown().await;
}
if let Some(ref mut stream) = self.stream {
let _ = stream.shutdown().await;
}
self.tls_stream = None;
self.stream = None;
}
}
pub struct Dhcpv6Client {
interface: String,
mac_address: [u8; 6],
duid: Vec<u8>,
state: Arc<RwLock<Dhcpv6State>>,
config: Dhcpv6Config,
ia_na: Arc<RwLock<Option<Dhcpv6IaNa>>>,
ia_pd: Arc<RwLock<Option<Dhcpv6IaPd>>>,
servers: Arc<RwLock<Vec<DhcpServer>>>,
security: SecurityContext,
shutdown: Arc<Mutex<bool>>,
transaction_id: Arc<Mutex<[u8; 3]>>,
server_duid: Arc<RwLock<Option<Vec<u8>>>>,
dns_servers: Arc<RwLock<Vec<Ipv6Addr>>>,
domain_list: Arc<RwLock<Vec<String>>>,
iaid: u32,
auth_context: Arc<Mutex<Dhcpv6AuthContext>>,
tls_context: Arc<Mutex<Dhcpv6TlsContext>>,
}
impl std::fmt::Debug for Dhcpv6Client {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Dhcpv6Client")
.field("interface", &self.interface)
.field("mac_address", &format!("{:02x}:{:02x}:{:02x}:{:02x}:{:02x}:{:02x}",
self.mac_address[0], self.mac_address[1], self.mac_address[2],
self.mac_address[3], self.mac_address[4], self.mac_address[5]))
.finish()
}
}
impl Dhcpv6Client {
pub fn new(interface: &str, config: Dhcpv6Config, security: SecurityConfig) -> Result<Self> {
validate_interface_name(interface)?;
let mac_address = Self::get_mac_address(interface)?;
let duid = Self::generate_duid(&mac_address);
let iaid = Self::generate_iaid(interface);
let auth_context = Dhcpv6AuthContext::from_config(&config.authentication);
let tls_context = Dhcpv6TlsContext::from_config(&config.tls);
Ok(Self {
interface: interface.to_string(),
mac_address,
duid,
state: Arc::new(RwLock::new(Dhcpv6State::Init)),
config,
ia_na: Arc::new(RwLock::new(None)),
ia_pd: Arc::new(RwLock::new(None)),
servers: Arc::new(RwLock::new(Vec::new())),
security: SecurityContext::new(security),
shutdown: Arc::new(Mutex::new(false)),
transaction_id: Arc::new(Mutex::new([0u8; 3])),
server_duid: Arc::new(RwLock::new(None)),
dns_servers: Arc::new(RwLock::new(Vec::new())),
domain_list: Arc::new(RwLock::new(Vec::new())),
iaid,
auth_context: Arc::new(Mutex::new(auth_context)),
tls_context: Arc::new(Mutex::new(tls_context)),
})
}
fn get_mac_address(interface: &str) -> Result<[u8; 6]> {
use std::fs;
let mac_path = format!("/sys/class/net/{}/address", interface);
let mac_str = fs::read_to_string(&mac_path)
.map_err(|e| DhcpClientError::InterfaceNotFound(
format!("Failed to read MAC address for {}: {}", interface, e)
))?;
let mac_str = mac_str.trim();
let parts: Vec<&str> = mac_str.split(':').collect();
if parts.len() != 6 {
return Err(DhcpClientError::InterfaceNotFound(
format!("Invalid MAC address format for interface {}", interface)
));
}
let mut mac = [0u8; 6];
for (i, part) in parts.iter().enumerate() {
mac[i] = u8::from_str_radix(part, 16)
.map_err(|_| DhcpClientError::InterfaceNotFound(
format!("Invalid MAC address for interface {}", interface)
))?;
}
Ok(mac)
}
fn generate_duid(mac: &[u8; 6]) -> Vec<u8> {
use std::time::{SystemTime, UNIX_EPOCH};
let mut duid = Vec::new();
duid.extend_from_slice(&1u16.to_be_bytes());
duid.extend_from_slice(&1u16.to_be_bytes());
let epoch_2000 = UNIX_EPOCH + Duration::from_secs(946684800);
let now = SystemTime::now();
let time_since = now.duration_since(epoch_2000)
.unwrap_or(Duration::from_secs(0))
.as_secs() as u32;
duid.extend_from_slice(&time_since.to_be_bytes());
duid.extend_from_slice(mac);
duid
}
fn generate_iaid(interface: &str) -> u32 {
use std::collections::hash_map::DefaultHasher;
use std::hash::{Hash, Hasher};
let mut hasher = DefaultHasher::new();
interface.hash(&mut hasher);
hasher.finish() as u32
}
pub fn interface(&self) -> &str {
&self.interface
}
pub fn get_status(&self) -> Dhcpv6Status {
let state = self.state.blocking_read();
let ia_na = self.ia_na.blocking_read();
let ia_pd = self.ia_pd.blocking_read();
let servers = self.servers.blocking_read();
Dhcpv6Status {
interface: self.interface.clone(),
state: state.to_string(),
ia_na: ia_na.as_ref().map(|i| i.to_info()),
ia_pd: ia_pd.as_ref().map(|i| i.to_info()),
servers: servers.clone(),
}
}
fn create_socket(&self) -> Result<UdpSocket> {
let bind_addr = SocketAddrV6::new(
Ipv6Addr::UNSPECIFIED,
DHCPV6_CLIENT_PORT,
0,
self.get_interface_index()?,
);
let socket = UdpSocket::bind(bind_addr)
.map_err(|e| DhcpClientError::Network(std::io::Error::new(
std::io::ErrorKind::Other,
format!("Failed to bind DHCPv6 socket: {}", e)
)))?;
socket.set_read_timeout(Some(Duration::from_secs(self.config.timeout)))
.map_err(|e| DhcpClientError::Network(e))?;
socket.join_multicast_v6(&ALL_DHCP_RELAY_AGENTS_AND_SERVERS, self.get_interface_index()?)
.map_err(|e| DhcpClientError::Network(std::io::Error::new(
std::io::ErrorKind::Other,
format!("Failed to join DHCPv6 multicast group: {}", e)
)))?;
dhcpv6_debug!("DHCPv6 socket created on {}", self.interface);
Ok(socket)
}
fn get_interface_index(&self) -> Result<u32> {
use std::ffi::CString;
let iface_cstr = CString::new(self.interface.as_str())
.map_err(|_| DhcpClientError::InvalidInterfaceName(
"Interface name contains null byte".to_string()
))?;
let index = unsafe { libc::if_nametoindex(iface_cstr.as_ptr()) };
if index == 0 {
return Err(DhcpClientError::InterfaceNotFound(
format!("Interface {} not found", self.interface)
));
}
Ok(index)
}
pub async fn run(&self) -> Result<()> {
info!("========================================");
info!("Starting DHCPv6 client on {}", self.interface);
info!(" DUID: {}", hex::encode(&self.duid));
info!(" IAID: 0x{:08x}", self.iaid);
info!("========================================");
loop {
if *self.shutdown.lock().await {
info!("DHCPv6 client shutting down on {}", self.interface);
break;
}
let state = self.state.read().await.clone();
dhcpv6_debug!("Current state: {}", state);
match state {
Dhcpv6State::Init => {
dhcpv6_info!("STATE: INIT -> Sending SOLICIT");
match self.do_solicit().await {
Ok(Some(server_info)) => {
dhcpv6_info!("Received ADVERTISE from server");
*self.server_duid.write().await = Some(server_info.server_id.clone());
match self.do_request(&server_info).await {
Ok(true) => {
*self.state.write().await = Dhcpv6State::Bound;
info!("DHCPv6 lease acquired on {}", self.interface);
if let Err(e) = self.apply_lease().await {
error!("Failed to apply DHCPv6 lease: {}", e);
}
}
Ok(false) => {
warn!("REQUEST failed, restarting");
*self.state.write().await = Dhcpv6State::Init;
sleep(Duration::from_secs(self.config.timeout)).await;
}
Err(e) => {
warn!("REQUEST error: {}", e);
*self.state.write().await = Dhcpv6State::Init;
sleep(Duration::from_secs(self.config.timeout)).await;
}
}
}
Ok(None) => {
dhcpv6_info!("No ADVERTISE received, retrying");
sleep(Duration::from_secs(self.config.timeout)).await;
}
Err(e) => {
warn!("SOLICIT failed: {}", e);
sleep(Duration::from_secs(self.config.timeout)).await;
}
}
}
Dhcpv6State::Solicit | Dhcpv6State::Request => {
*self.state.write().await = Dhcpv6State::Init;
}
Dhcpv6State::Bound => {
if let Some(ia_na) = self.ia_na.read().await.as_ref() {
let now = SystemTime::now();
let renewal_at = ia_na.renewal_at();
if now >= renewal_at {
dhcpv6_info!("Renewal time reached, transitioning to RENEW");
*self.state.write().await = Dhcpv6State::Renew;
} else if ia_na.is_expired() {
warn!("Lease expired, returning to INIT");
*self.ia_na.write().await = None;
*self.state.write().await = Dhcpv6State::Init;
} else {
let wait_time = renewal_at.duration_since(now)
.unwrap_or(Duration::from_secs(60));
sleep(wait_time.min(Duration::from_secs(60))).await;
}
} else {
*self.state.write().await = Dhcpv6State::Init;
}
}
Dhcpv6State::Renew => {
dhcpv6_info!("STATE: RENEW");
match self.do_renew().await {
Ok(true) => {
dhcpv6_info!("Renewal successful");
*self.state.write().await = Dhcpv6State::Bound;
if let Err(e) = self.apply_lease().await {
warn!("Failed to apply renewed lease: {}", e);
}
}
Ok(false) | Err(_) => {
if let Some(ia_na) = self.ia_na.read().await.as_ref() {
let now = SystemTime::now();
if now >= ia_na.rebinding_at() {
dhcpv6_info!("Rebinding time reached, transitioning to REBIND");
*self.state.write().await = Dhcpv6State::Rebind;
} else {
sleep(Duration::from_secs(self.config.timeout)).await;
}
} else {
*self.state.write().await = Dhcpv6State::Init;
}
}
}
}
Dhcpv6State::Rebind => {
dhcpv6_info!("STATE: REBIND");
match self.do_rebind().await {
Ok(true) => {
dhcpv6_info!("Rebind successful");
*self.state.write().await = Dhcpv6State::Bound;
if let Err(e) = self.apply_lease().await {
warn!("Failed to apply rebound lease: {}", e);
}
}
Ok(false) | Err(_) => {
if let Some(ia_na) = self.ia_na.read().await.as_ref() {
if ia_na.is_expired() {
warn!("Lease expired during rebind, returning to INIT");
*self.ia_na.write().await = None;
*self.state.write().await = Dhcpv6State::Init;
} else {
sleep(Duration::from_secs(self.config.timeout)).await;
}
} else {
*self.state.write().await = Dhcpv6State::Init;
}
}
}
}
Dhcpv6State::Confirm => {
*self.state.write().await = Dhcpv6State::Bound;
}
}
sleep(Duration::from_millis(100)).await;
}
Ok(())
}
async fn do_solicit(&self) -> Result<Option<Dhcpv6ServerInfo>> {
self.security.check_rate_limit(&self.interface).await?;
let xid: [u8; 3] = rand::random();
*self.transaction_id.lock().await = xid;
dhcpv6_info!("Sending DHCPv6 SOLICIT (xid: {:02x}{:02x}{:02x})", xid[0], xid[1], xid[2]);
let solicit = self.build_solicit_message(&xid);
let socket = self.create_socket()?;
let dest = SocketAddrV6::new(
ALL_DHCP_RELAY_AGENTS_AND_SERVERS,
DHCPV6_SERVER_PORT,
0,
self.get_interface_index()?,
);
socket.send_to(&solicit, dest)
.map_err(|e| DhcpClientError::Network(e))?;
info!("DHCPv6 SOLICIT sent on {}", self.interface);
let mut buf = [0u8; 1500];
let timeout = Duration::from_secs(self.config.timeout);
let start = std::time::Instant::now();
while start.elapsed() < timeout {
match socket.recv_from(&mut buf) {
Ok((size, _src)) => {
dhcpv6_debug!("Received {} bytes from {}", size, _src);
if let Some(server_info) = self.parse_advertise(&buf[..size], &xid) {
dhcpv6_info!("Received valid ADVERTISE");
return Ok(Some(server_info));
}
}
Err(e) if e.kind() == std::io::ErrorKind::WouldBlock => {
continue;
}
Err(e) if e.kind() == std::io::ErrorKind::TimedOut => {
break;
}
Err(_e) => {
dhcpv6_debug!("Recv error: {}", _e);
}
}
}
Ok(None)
}
async fn do_request(&self, server_info: &Dhcpv6ServerInfo) -> Result<bool> {
self.security.check_rate_limit(&self.interface).await?;
let xid: [u8; 3] = rand::random();
*self.transaction_id.lock().await = xid;
dhcpv6_info!("Sending DHCPv6 REQUEST");
let request = self.build_request_message(&xid, server_info);
let socket = self.create_socket()?;
let dest = SocketAddrV6::new(
ALL_DHCP_RELAY_AGENTS_AND_SERVERS,
DHCPV6_SERVER_PORT,
0,
self.get_interface_index()?,
);
socket.send_to(&request, dest)
.map_err(|e| DhcpClientError::Network(e))?;
info!("DHCPv6 REQUEST sent on {}", self.interface);
let mut buf = [0u8; 1500];
let timeout = Duration::from_secs(self.config.timeout);
let start = std::time::Instant::now();
while start.elapsed() < timeout {
match socket.recv_from(&mut buf) {
Ok((size, _src)) => {
if let Some((ia_na, ia_pd, dns, domains)) = self.parse_reply(&buf[..size], &xid) {
dhcpv6_info!("Received valid REPLY");
if let Some(ia_na) = ia_na {
*self.ia_na.write().await = Some(ia_na);
}
if let Some(ia_pd) = ia_pd {
*self.ia_pd.write().await = Some(ia_pd);
}
*self.dns_servers.write().await = dns;
*self.domain_list.write().await = domains;
return Ok(true);
}
}
Err(e) if e.kind() == std::io::ErrorKind::WouldBlock => continue,
Err(e) if e.kind() == std::io::ErrorKind::TimedOut => break,
Err(_) => {}
}
}
Ok(false)
}
async fn do_renew(&self) -> Result<bool> {
self.security.check_rate_limit(&self.interface).await?;
let xid: [u8; 3] = rand::random();
*self.transaction_id.lock().await = xid;
let server_duid = self.server_duid.read().await.clone()
.ok_or_else(|| DhcpClientError::NoLease)?;
dhcpv6_info!("Sending DHCPv6 RENEW");
let renew = self.build_renew_message(&xid, &server_duid);
let socket = self.create_socket()?;
let dest = SocketAddrV6::new(
ALL_DHCP_RELAY_AGENTS_AND_SERVERS,
DHCPV6_SERVER_PORT,
0,
self.get_interface_index()?,
);
socket.send_to(&renew, dest)
.map_err(|e| DhcpClientError::Network(e))?;
let mut buf = [0u8; 1500];
let timeout = Duration::from_secs(self.config.timeout);
let start = std::time::Instant::now();
while start.elapsed() < timeout {
match socket.recv_from(&mut buf) {
Ok((size, _)) => {
if let Some((ia_na, ia_pd, dns, domains)) = self.parse_reply(&buf[..size], &xid) {
if let Some(ia_na) = ia_na {
*self.ia_na.write().await = Some(ia_na);
}
if let Some(ia_pd) = ia_pd {
*self.ia_pd.write().await = Some(ia_pd);
}
*self.dns_servers.write().await = dns;
*self.domain_list.write().await = domains;
return Ok(true);
}
}
Err(e) if e.kind() == std::io::ErrorKind::WouldBlock => continue,
Err(e) if e.kind() == std::io::ErrorKind::TimedOut => break,
Err(_) => {}
}
}
Ok(false)
}
async fn do_rebind(&self) -> Result<bool> {
self.security.check_rate_limit(&self.interface).await?;
let xid: [u8; 3] = rand::random();
*self.transaction_id.lock().await = xid;
dhcpv6_info!("Sending DHCPv6 REBIND");
let rebind = self.build_rebind_message(&xid);
let socket = self.create_socket()?;
let dest = SocketAddrV6::new(
ALL_DHCP_RELAY_AGENTS_AND_SERVERS,
DHCPV6_SERVER_PORT,
0,
self.get_interface_index()?,
);
socket.send_to(&rebind, dest)
.map_err(|e| DhcpClientError::Network(e))?;
let mut buf = [0u8; 1500];
let timeout = Duration::from_secs(self.config.timeout);
let start = std::time::Instant::now();
while start.elapsed() < timeout {
match socket.recv_from(&mut buf) {
Ok((size, _)) => {
if let Some((ia_na, ia_pd, dns, domains)) = self.parse_reply(&buf[..size], &xid) {
if let Some(ia_na) = ia_na {
*self.ia_na.write().await = Some(ia_na);
}
if let Some(ia_pd) = ia_pd {
*self.ia_pd.write().await = Some(ia_pd);
}
*self.dns_servers.write().await = dns;
*self.domain_list.write().await = domains;
return Ok(true);
}
}
Err(e) if e.kind() == std::io::ErrorKind::WouldBlock => continue,
Err(e) if e.kind() == std::io::ErrorKind::TimedOut => break,
Err(_) => {}
}
}
Ok(false)
}
fn build_solicit_message(&self, xid: &[u8; 3]) -> Vec<u8> {
let mut msg = Vec::new();
msg.push(DHCPV6_SOLICIT);
msg.extend_from_slice(xid);
self.add_option(&mut msg, OPTION_CLIENTID, &self.duid);
self.add_option(&mut msg, OPTION_ELAPSED_TIME, &[0, 0]);
let ia_na_data = self.build_ia_na_option(None);
self.add_option(&mut msg, OPTION_IA_NA, &ia_na_data);
if self.config.prefix_delegation {
let ia_pd_data = self.build_ia_pd_option(None);
self.add_option(&mut msg, OPTION_IA_PD, &ia_pd_data);
}
let mut oro_data = Vec::new();
for opt in &self.config.request_options {
oro_data.extend_from_slice(&opt.to_be_bytes());
}
if !self.config.request_options.contains(&OPTION_DNS_SERVERS) {
oro_data.extend_from_slice(&OPTION_DNS_SERVERS.to_be_bytes());
}
if !self.config.request_options.contains(&OPTION_DOMAIN_LIST) {
oro_data.extend_from_slice(&OPTION_DOMAIN_LIST.to_be_bytes());
}
self.add_option(&mut msg, OPTION_ORO, &oro_data);
if self.config.rapid_commit {
self.add_option(&mut msg, OPTION_RAPID_COMMIT, &[]);
}
self.add_auth_option_blocking(&mut msg);
msg
}
fn build_request_message(&self, xid: &[u8; 3], server_info: &Dhcpv6ServerInfo) -> Vec<u8> {
let mut msg = Vec::new();
msg.push(DHCPV6_REQUEST);
msg.extend_from_slice(xid);
self.add_option(&mut msg, OPTION_CLIENTID, &self.duid);
self.add_option(&mut msg, OPTION_SERVERID, &server_info.server_id);
self.add_option(&mut msg, OPTION_ELAPSED_TIME, &[0, 0]);
if let Some(ref ia_na) = server_info.ia_na {
let ia_na_data = self.build_ia_na_option(Some(ia_na));
self.add_option(&mut msg, OPTION_IA_NA, &ia_na_data);
} else {
let ia_na_data = self.build_ia_na_option(None);
self.add_option(&mut msg, OPTION_IA_NA, &ia_na_data);
}
if self.config.prefix_delegation {
if let Some(ref ia_pd) = server_info.ia_pd {
let ia_pd_data = self.build_ia_pd_option(Some(ia_pd));
self.add_option(&mut msg, OPTION_IA_PD, &ia_pd_data);
} else {
let ia_pd_data = self.build_ia_pd_option(None);
self.add_option(&mut msg, OPTION_IA_PD, &ia_pd_data);
}
}
let mut oro_data = Vec::new();
oro_data.extend_from_slice(&OPTION_DNS_SERVERS.to_be_bytes());
oro_data.extend_from_slice(&OPTION_DOMAIN_LIST.to_be_bytes());
self.add_option(&mut msg, OPTION_ORO, &oro_data);
self.add_auth_option_blocking(&mut msg);
msg
}
fn build_renew_message(&self, xid: &[u8; 3], server_duid: &[u8]) -> Vec<u8> {
let mut msg = Vec::new();
msg.push(DHCPV6_RENEW);
msg.extend_from_slice(xid);
self.add_option(&mut msg, OPTION_CLIENTID, &self.duid);
self.add_option(&mut msg, OPTION_SERVERID, server_duid);
self.add_option(&mut msg, OPTION_ELAPSED_TIME, &[0, 0]);
if let Some(ref ia_na) = *self.ia_na.blocking_read() {
let ia_na_data = self.build_ia_na_option(Some(ia_na));
self.add_option(&mut msg, OPTION_IA_NA, &ia_na_data);
}
if self.config.prefix_delegation {
if let Some(ref ia_pd) = *self.ia_pd.blocking_read() {
let ia_pd_data = self.build_ia_pd_option(Some(ia_pd));
self.add_option(&mut msg, OPTION_IA_PD, &ia_pd_data);
}
}
self.add_auth_option_blocking(&mut msg);
msg
}
fn build_rebind_message(&self, xid: &[u8; 3]) -> Vec<u8> {
let mut msg = Vec::new();
msg.push(DHCPV6_REBIND);
msg.extend_from_slice(xid);
self.add_option(&mut msg, OPTION_CLIENTID, &self.duid);
self.add_option(&mut msg, OPTION_ELAPSED_TIME, &[0, 0]);
if let Some(ref ia_na) = *self.ia_na.blocking_read() {
let ia_na_data = self.build_ia_na_option(Some(ia_na));
self.add_option(&mut msg, OPTION_IA_NA, &ia_na_data);
}
if self.config.prefix_delegation {
if let Some(ref ia_pd) = *self.ia_pd.blocking_read() {
let ia_pd_data = self.build_ia_pd_option(Some(ia_pd));
self.add_option(&mut msg, OPTION_IA_PD, &ia_pd_data);
}
}
self.add_auth_option_blocking(&mut msg);
msg
}
fn build_release_message(&self, xid: &[u8; 3], server_duid: &[u8]) -> Vec<u8> {
let mut msg = Vec::new();
msg.push(DHCPV6_RELEASE);
msg.extend_from_slice(xid);
self.add_option(&mut msg, OPTION_CLIENTID, &self.duid);
self.add_option(&mut msg, OPTION_SERVERID, server_duid);
self.add_option(&mut msg, OPTION_ELAPSED_TIME, &[0, 0]);
if let Some(ref ia_na) = *self.ia_na.blocking_read() {
let ia_na_data = self.build_ia_na_option(Some(ia_na));
self.add_option(&mut msg, OPTION_IA_NA, &ia_na_data);
}
if let Some(ref ia_pd) = *self.ia_pd.blocking_read() {
let ia_pd_data = self.build_ia_pd_option(Some(ia_pd));
self.add_option(&mut msg, OPTION_IA_PD, &ia_pd_data);
}
self.add_auth_option_blocking(&mut msg);
msg
}
fn build_decline_message(&self, xid: &[u8; 3], server_duid: &[u8], addresses: &[Ipv6Addr]) -> Vec<u8> {
let mut msg = Vec::new();
msg.push(DHCPV6_DECLINE);
msg.extend_from_slice(xid);
self.add_option(&mut msg, OPTION_CLIENTID, &self.duid);
self.add_option(&mut msg, OPTION_SERVERID, server_duid);
self.add_option(&mut msg, OPTION_ELAPSED_TIME, &[0, 0]);
let mut ia_na_data = Vec::new();
ia_na_data.extend_from_slice(&self.iaid.to_be_bytes()); ia_na_data.extend_from_slice(&0u32.to_be_bytes()); ia_na_data.extend_from_slice(&0u32.to_be_bytes());
for addr in addresses {
let mut iaaddr = Vec::new();
iaaddr.extend_from_slice(&addr.octets()); iaaddr.extend_from_slice(&0u32.to_be_bytes()); iaaddr.extend_from_slice(&0u32.to_be_bytes());
ia_na_data.extend_from_slice(&OPTION_IAADDR.to_be_bytes());
ia_na_data.extend_from_slice(&(iaaddr.len() as u16).to_be_bytes());
ia_na_data.extend_from_slice(&iaaddr);
}
self.add_option(&mut msg, OPTION_IA_NA, &ia_na_data);
self.add_auth_option_blocking(&mut msg);
msg
}
fn build_ia_na_option(&self, existing: Option<&Dhcpv6IaNa>) -> Vec<u8> {
let mut data = Vec::new();
let iaid = existing.map(|i| i.iaid).unwrap_or(self.iaid);
data.extend_from_slice(&iaid.to_be_bytes());
let t1 = existing.map(|i| i.t1.as_secs() as u32).unwrap_or(0);
data.extend_from_slice(&t1.to_be_bytes());
let t2 = existing.map(|i| i.t2.as_secs() as u32).unwrap_or(0);
data.extend_from_slice(&t2.to_be_bytes());
if let Some(ia_na) = existing {
for addr in &ia_na.addresses {
let mut iaaddr = Vec::new();
iaaddr.extend_from_slice(&addr.octets());
iaaddr.extend_from_slice(&(ia_na.preferred_lifetime.as_secs() as u32).to_be_bytes());
iaaddr.extend_from_slice(&(ia_na.valid_lifetime.as_secs() as u32).to_be_bytes());
data.extend_from_slice(&OPTION_IAADDR.to_be_bytes());
data.extend_from_slice(&(iaaddr.len() as u16).to_be_bytes());
data.extend_from_slice(&iaaddr);
}
}
data
}
fn build_ia_pd_option(&self, existing: Option<&Dhcpv6IaPd>) -> Vec<u8> {
let mut data = Vec::new();
let iaid = existing.map(|i| i.iaid).unwrap_or(self.iaid.wrapping_add(1));
data.extend_from_slice(&iaid.to_be_bytes());
let t1 = existing.map(|i| i.t1.as_secs() as u32).unwrap_or(0);
data.extend_from_slice(&t1.to_be_bytes());
let t2 = existing.map(|i| i.t2.as_secs() as u32).unwrap_or(0);
data.extend_from_slice(&t2.to_be_bytes());
if let Some(ia_pd) = existing {
let mut iaprefix = Vec::new();
iaprefix.extend_from_slice(&(ia_pd.preferred_lifetime.as_secs() as u32).to_be_bytes());
iaprefix.extend_from_slice(&(ia_pd.valid_lifetime.as_secs() as u32).to_be_bytes());
iaprefix.push(ia_pd.prefix_len);
iaprefix.extend_from_slice(&ia_pd.prefix.octets());
data.extend_from_slice(&OPTION_IAPREFIX.to_be_bytes());
data.extend_from_slice(&(iaprefix.len() as u16).to_be_bytes());
data.extend_from_slice(&iaprefix);
}
data
}
fn add_option(&self, msg: &mut Vec<u8>, code: u16, data: &[u8]) {
msg.extend_from_slice(&code.to_be_bytes());
msg.extend_from_slice(&(data.len() as u16).to_be_bytes());
msg.extend_from_slice(data);
}
async fn add_auth_option(&self, msg: &mut Vec<u8>) {
let mut auth_ctx = self.auth_context.lock().await;
if !auth_ctx.enabled {
return;
}
if let Some(auth_data) = auth_ctx.build_auth_option(msg) {
self.add_option(msg, OPTION_AUTH, &auth_data);
dhcpv6_debug!("Added authentication option ({} bytes)", auth_data.len());
}
}
fn add_auth_option_blocking(&self, msg: &mut Vec<u8>) {
let mut auth_ctx = self.auth_context.blocking_lock();
if !auth_ctx.enabled {
return;
}
if let Some(auth_data) = auth_ctx.build_auth_option(msg) {
self.add_option(msg, OPTION_AUTH, &auth_data);
dhcpv6_debug!("Added authentication option ({} bytes)", auth_data.len());
}
}
async fn verify_message_auth(&self, message: &[u8]) -> Result<bool> {
let mut auth_ctx = self.auth_context.lock().await;
if !auth_ctx.enabled {
return Ok(false);
}
if let Some(auth_opt_data) = self.extract_auth_option(message) {
match auth_ctx.parse_and_verify_auth(message, &auth_opt_data) {
Ok(true) => {
dhcpv6_debug!("Authentication verified successfully");
Ok(true)
}
Ok(false) => {
warn!("Authentication verification returned false");
if auth_ctx.require_server_auth {
Err(DhcpClientError::SecurityViolation(
"Server authentication failed".to_string()
))
} else {
Ok(false)
}
}
Err(e) => {
warn!("Authentication verification error: {}", e);
if auth_ctx.require_server_auth {
Err(DhcpClientError::SecurityViolation(
format!("Server authentication error: {}", e)
))
} else {
Ok(false)
}
}
}
} else {
if auth_ctx.require_server_auth && !auth_ctx.accept_unauthenticated {
Err(DhcpClientError::SecurityViolation(
"Server message has no authentication".to_string()
))
} else {
dhcpv6_debug!("No authentication option in server message");
Ok(false)
}
}
}
fn extract_auth_option(&self, message: &[u8]) -> Option<Vec<u8>> {
if message.len() < 4 {
return None;
}
let mut pos = 4;
while pos + 4 <= message.len() {
let opt_code = u16::from_be_bytes([message[pos], message[pos + 1]]);
let opt_len = u16::from_be_bytes([message[pos + 2], message[pos + 3]]) as usize;
if pos + 4 + opt_len > message.len() {
break;
}
if opt_code == OPTION_AUTH {
return Some(message[pos + 4..pos + 4 + opt_len].to_vec());
}
pos += 4 + opt_len;
}
None
}
fn parse_advertise(&self, data: &[u8], expected_xid: &[u8; 3]) -> Option<Dhcpv6ServerInfo> {
if data.len() < 4 {
return None;
}
if data[0] != DHCPV6_ADVERTISE {
return None;
}
if &data[1..4] != expected_xid {
dhcpv6_debug!("XID mismatch in ADVERTISE");
return None;
}
let mut server_id: Option<Vec<u8>> = None;
let mut preference: u8 = 0;
let mut ia_na: Option<Dhcpv6IaNa> = None;
let mut ia_pd: Option<Dhcpv6IaPd> = None;
let mut dns_servers: Vec<Ipv6Addr> = Vec::new();
let mut domain_list: Vec<String> = Vec::new();
let mut i = 4;
while i + 4 <= data.len() {
let opt_code = u16::from_be_bytes([data[i], data[i + 1]]);
let opt_len = u16::from_be_bytes([data[i + 2], data[i + 3]]) as usize;
i += 4;
if i + opt_len > data.len() {
break;
}
let opt_data = &data[i..i + opt_len];
match opt_code {
OPTION_SERVERID => {
server_id = Some(opt_data.to_vec());
}
OPTION_PREFERENCE => {
if opt_len >= 1 {
preference = opt_data[0];
}
}
OPTION_IA_NA => {
ia_na = self.parse_ia_na(opt_data);
}
OPTION_IA_PD => {
ia_pd = self.parse_ia_pd(opt_data);
}
OPTION_DNS_SERVERS => {
dns_servers = self.parse_dns_servers(opt_data);
}
OPTION_DOMAIN_LIST => {
domain_list = self.parse_domain_list(opt_data);
}
OPTION_STATUS_CODE => {
if opt_len >= 2 {
let status = u16::from_be_bytes([opt_data[0], opt_data[1]]);
if status != STATUS_SUCCESS {
dhcpv6_debug!("ADVERTISE contains error status: {}", status);
return None;
}
}
}
_ => {
dhcpv6_debug!("Unknown option {} (len {})", opt_code, opt_len);
}
}
i += opt_len;
}
let server_id = server_id?;
Some(Dhcpv6ServerInfo {
server_id,
address: Ipv6Addr::UNSPECIFIED, preference,
ia_na,
ia_pd,
dns_servers,
domain_list,
})
}
fn parse_reply(&self, data: &[u8], expected_xid: &[u8; 3]) -> Option<(Option<Dhcpv6IaNa>, Option<Dhcpv6IaPd>, Vec<Ipv6Addr>, Vec<String>)> {
if data.len() < 4 {
return None;
}
if data[0] != DHCPV6_REPLY {
return None;
}
if &data[1..4] != expected_xid {
dhcpv6_debug!("XID mismatch in REPLY");
return None;
}
let mut ia_na: Option<Dhcpv6IaNa> = None;
let mut ia_pd: Option<Dhcpv6IaPd> = None;
let mut dns_servers: Vec<Ipv6Addr> = Vec::new();
let mut domain_list: Vec<String> = Vec::new();
let mut status_ok = true;
let mut i = 4;
while i + 4 <= data.len() {
let opt_code = u16::from_be_bytes([data[i], data[i + 1]]);
let opt_len = u16::from_be_bytes([data[i + 2], data[i + 3]]) as usize;
i += 4;
if i + opt_len > data.len() {
break;
}
let opt_data = &data[i..i + opt_len];
match opt_code {
OPTION_IA_NA => {
ia_na = self.parse_ia_na(opt_data);
}
OPTION_IA_PD => {
ia_pd = self.parse_ia_pd(opt_data);
}
OPTION_DNS_SERVERS => {
dns_servers = self.parse_dns_servers(opt_data);
}
OPTION_DOMAIN_LIST => {
domain_list = self.parse_domain_list(opt_data);
}
OPTION_STATUS_CODE => {
if opt_len >= 2 {
let status = u16::from_be_bytes([opt_data[0], opt_data[1]]);
if status != STATUS_SUCCESS {
warn!("REPLY contains error status: {}", status);
status_ok = false;
}
}
}
_ => {}
}
i += opt_len;
}
if !status_ok {
return None;
}
Some((ia_na, ia_pd, dns_servers, domain_list))
}
fn parse_ia_na(&self, data: &[u8]) -> Option<Dhcpv6IaNa> {
if data.len() < 12 {
return None;
}
let iaid = u32::from_be_bytes([data[0], data[1], data[2], data[3]]);
let t1 = u32::from_be_bytes([data[4], data[5], data[6], data[7]]);
let t2 = u32::from_be_bytes([data[8], data[9], data[10], data[11]]);
let mut addresses = Vec::new();
let mut preferred_lifetime = Duration::from_secs(0);
let mut valid_lifetime = Duration::from_secs(0);
let mut i = 12;
while i + 4 <= data.len() {
let opt_code = u16::from_be_bytes([data[i], data[i + 1]]);
let opt_len = u16::from_be_bytes([data[i + 2], data[i + 3]]) as usize;
i += 4;
if i + opt_len > data.len() {
break;
}
if opt_code == OPTION_IAADDR && opt_len >= 24 {
let addr_bytes: [u8; 16] = data[i..i + 16].try_into().ok()?;
let addr = Ipv6Addr::from(addr_bytes);
preferred_lifetime = Duration::from_secs(
u32::from_be_bytes([data[i + 16], data[i + 17], data[i + 18], data[i + 19]]) as u64
);
valid_lifetime = Duration::from_secs(
u32::from_be_bytes([data[i + 20], data[i + 21], data[i + 22], data[i + 23]]) as u64
);
addresses.push(addr);
}
i += opt_len;
}
if addresses.is_empty() {
return None;
}
Some(Dhcpv6IaNa {
iaid,
addresses,
t1: Duration::from_secs(t1 as u64),
t2: Duration::from_secs(t2 as u64),
preferred_lifetime,
valid_lifetime,
acquired_at: SystemTime::now(),
})
}
fn parse_ia_pd(&self, data: &[u8]) -> Option<Dhcpv6IaPd> {
if data.len() < 12 {
return None;
}
let iaid = u32::from_be_bytes([data[0], data[1], data[2], data[3]]);
let t1 = u32::from_be_bytes([data[4], data[5], data[6], data[7]]);
let t2 = u32::from_be_bytes([data[8], data[9], data[10], data[11]]);
let mut i = 12;
while i + 4 <= data.len() {
let opt_code = u16::from_be_bytes([data[i], data[i + 1]]);
let opt_len = u16::from_be_bytes([data[i + 2], data[i + 3]]) as usize;
i += 4;
if i + opt_len > data.len() {
break;
}
if opt_code == OPTION_IAPREFIX && opt_len >= 25 {
let preferred_lifetime = Duration::from_secs(
u32::from_be_bytes([data[i], data[i + 1], data[i + 2], data[i + 3]]) as u64
);
let valid_lifetime = Duration::from_secs(
u32::from_be_bytes([data[i + 4], data[i + 5], data[i + 6], data[i + 7]]) as u64
);
let prefix_len = data[i + 8];
let prefix_bytes: [u8; 16] = data[i + 9..i + 25].try_into().ok()?;
let prefix = Ipv6Addr::from(prefix_bytes);
return Some(Dhcpv6IaPd {
iaid,
prefix,
prefix_len,
t1: Duration::from_secs(t1 as u64),
t2: Duration::from_secs(t2 as u64),
preferred_lifetime,
valid_lifetime,
acquired_at: SystemTime::now(),
});
}
i += opt_len;
}
None
}
fn parse_dns_servers(&self, data: &[u8]) -> Vec<Ipv6Addr> {
let mut servers = Vec::new();
let mut i = 0;
while i + 16 <= data.len() {
let addr_bytes: [u8; 16] = data[i..i + 16].try_into().unwrap();
servers.push(Ipv6Addr::from(addr_bytes));
i += 16;
}
servers
}
fn parse_domain_list(&self, data: &[u8]) -> Vec<String> {
let mut domains = Vec::new();
let mut i = 0;
while i < data.len() {
let mut labels = Vec::new();
loop {
if i >= data.len() {
break;
}
let len = data[i] as usize;
i += 1;
if len == 0 {
break;
}
if i + len > data.len() {
break;
}
if let Ok(label) = std::str::from_utf8(&data[i..i + len]) {
labels.push(label.to_string());
}
i += len;
}
if !labels.is_empty() {
domains.push(labels.join("."));
}
}
domains
}
async fn apply_lease(&self) -> Result<()> {
use tokio::process::Command;
let ia_na = self.ia_na.read().await;
let dns_servers = self.dns_servers.read().await;
if let Some(ref ia_na) = *ia_na {
for addr in &ia_na.addresses {
info!("Configuring IPv6 address {} on {}", addr, self.interface);
let output = Command::new("ip")
.args(&["-6", "addr", "add", &format!("{}/128", addr), "dev", &self.interface])
.output()
.await
.map_err(|e| DhcpClientError::Network(e))?;
if !output.status.success() {
let stderr = String::from_utf8_lossy(&output.stderr);
if !stderr.contains("File exists") {
warn!("Failed to add IPv6 address: {}", stderr);
}
}
}
}
if !dns_servers.is_empty() {
if let Err(e) = self.update_resolv_conf(&dns_servers).await {
warn!("Failed to update resolv.conf: {}", e);
}
}
info!("✓ DHCPv6 lease applied on {}", self.interface);
Ok(())
}
async fn update_resolv_conf(&self, dns_servers: &[Ipv6Addr]) -> Result<()> {
use tokio::fs;
let existing = fs::read_to_string("/etc/resolv.conf").await.unwrap_or_default();
let lines: Vec<&str> = existing.lines()
.filter(|line| {
if line.starts_with("nameserver ") {
let addr = line.trim_start_matches("nameserver ").trim();
!addr.contains(':')
} else {
true
}
})
.collect();
let mut new_content = String::new();
new_content.push_str(&format!("# DHCPv6 DNS servers for {}\n", self.interface));
for dns in dns_servers {
new_content.push_str(&format!("nameserver {}\n", dns));
}
for line in &lines {
new_content.push_str(line);
new_content.push('\n');
}
let temp_path = "/etc/resolv.conf.dhcpv6.tmp";
fs::write(temp_path, &new_content).await
.map_err(|e| DhcpClientError::Network(e))?;
fs::rename(temp_path, "/etc/resolv.conf").await
.map_err(|e| DhcpClientError::Network(e))?;
info!("✓ Updated /etc/resolv.conf with {} IPv6 DNS servers", dns_servers.len());
Ok(())
}
pub async fn renew(&self) -> Result<()> {
info!("Renewing DHCPv6 lease on {}", self.interface);
*self.state.write().await = Dhcpv6State::Renew;
Ok(())
}
pub async fn release(&self) -> Result<()> {
info!("Releasing DHCPv6 lease on {}", self.interface);
if let Some(ref server_duid) = *self.server_duid.read().await {
let xid: [u8; 3] = rand::random();
let release = self.build_release_message(&xid, server_duid);
if let Ok(socket) = self.create_socket() {
let dest = SocketAddrV6::new(
ALL_DHCP_RELAY_AGENTS_AND_SERVERS,
DHCPV6_SERVER_PORT,
0,
self.get_interface_index().unwrap_or(0),
);
let _ = socket.send_to(&release, dest);
info!("DHCPv6 RELEASE sent on {}", self.interface);
}
}
*self.ia_na.write().await = None;
*self.ia_pd.write().await = None;
*self.server_duid.write().await = None;
*self.dns_servers.write().await = Vec::new();
*self.domain_list.write().await = Vec::new();
*self.state.write().await = Dhcpv6State::Init;
Ok(())
}
pub async fn decline(&self, addresses: &[Ipv6Addr]) -> Result<()> {
info!("Declining DHCPv6 addresses on {}", self.interface);
if let Some(ref server_duid) = *self.server_duid.read().await {
let xid: [u8; 3] = rand::random();
let decline = self.build_decline_message(&xid, server_duid, addresses);
if let Ok(socket) = self.create_socket() {
let dest = SocketAddrV6::new(
ALL_DHCP_RELAY_AGENTS_AND_SERVERS,
DHCPV6_SERVER_PORT,
0,
self.get_interface_index().unwrap_or(0),
);
let _ = socket.send_to(&decline, dest);
info!("DHCPv6 DECLINE sent on {}", self.interface);
}
}
*self.state.write().await = Dhcpv6State::Init;
Ok(())
}
pub async fn shutdown(&self) {
*self.shutdown.lock().await = true;
}
pub async fn starttls(&self, server_addr: std::net::SocketAddrV6) -> Result<Dhcpv6TcpConnection> {
let mut tls_ctx = self.tls_context.lock().await;
if !tls_ctx.is_enabled() {
return Err(DhcpClientError::InvalidConfig(
"TLS is not enabled in configuration".to_string()
));
}
tls_ctx.init_tls_config()?;
info!("Initiating DHCPv6 STARTTLS to {}", server_addr);
let mut conn = Dhcpv6TcpConnection::connect(server_addr).await?;
let xid: [u8; 3] = rand::random();
let starttls_msg = Dhcpv6TlsContext::build_starttls_message(&xid);
conn.send(&starttls_msg).await?;
dhcpv6_debug!("Sent STARTTLS message with XID {:02x}{:02x}{:02x}", xid[0], xid[1], xid[2]);
let timeout = Duration::from_secs(self.config.timeout);
let response = conn.recv(timeout).await?;
if !Dhcpv6TlsContext::parse_starttls_response(&response, &xid) {
conn.close().await;
return Err(DhcpClientError::ServerError(
"Server rejected STARTTLS or invalid response".to_string()
));
}
dhcpv6_debug!("Server accepted STARTTLS, upgrading to TLS");
let server_name = tls_ctx.config.server_name
.clone()
.unwrap_or_else(|| server_addr.ip().to_string());
let tls_config = tls_ctx.get_tls_config()
.ok_or_else(|| DhcpClientError::InvalidConfig(
"TLS configuration not initialized".to_string()
))?;
conn.starttls(tls_config, &server_name).await?;
tls_ctx.set_tls_established(true);
info!("DHCPv6 STARTTLS completed successfully");
Ok(conn)
}
pub async fn is_tls_enabled(&self) -> bool {
self.tls_context.lock().await.is_enabled()
}
pub async fn is_tls_established(&self) -> bool {
self.tls_context.lock().await.is_tls_established()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_dhcpv6_state_display() {
assert_eq!(Dhcpv6State::Init.to_string(), "INIT");
assert_eq!(Dhcpv6State::Solicit.to_string(), "SOLICIT");
assert_eq!(Dhcpv6State::Bound.to_string(), "BOUND");
}
#[test]
fn test_generate_duid() {
let mac = [0x00, 0x11, 0x22, 0x33, 0x44, 0x55];
let duid = Dhcpv6Client::generate_duid(&mac);
assert_eq!(duid.len(), 14);
assert_eq!(u16::from_be_bytes([duid[0], duid[1]]), 1);
assert_eq!(u16::from_be_bytes([duid[2], duid[3]]), 1);
assert_eq!(&duid[8..14], &mac);
}
#[test]
fn test_generate_iaid() {
let iaid1 = Dhcpv6Client::generate_iaid("eth0");
let iaid2 = Dhcpv6Client::generate_iaid("eth1");
assert_ne!(iaid1, iaid2);
let iaid1_again = Dhcpv6Client::generate_iaid("eth0");
assert_eq!(iaid1, iaid1_again);
}
#[test]
fn test_auth_context_creation() {
use config::Dhcpv6AuthConfig;
let config = Dhcpv6AuthConfig::default();
let auth_ctx = Dhcpv6AuthContext::from_config(&config);
assert!(!auth_ctx.enabled);
assert!(auth_ctx.key.is_none());
let config = Dhcpv6AuthConfig {
enabled: true,
protocol: "delayed".to_string(),
algorithm: "hmac-sha256".to_string(),
key: Some("0102030405060708090a0b0c0d0e0f10".to_string()),
key_id: Some(1),
realm: Some("example.com".to_string()),
require_server_auth: false,
accept_unauthenticated: true,
};
let auth_ctx = Dhcpv6AuthContext::from_config(&config);
assert!(auth_ctx.enabled);
assert!(auth_ctx.key.is_some());
assert_eq!(auth_ctx.key.as_ref().unwrap().len(), 16);
assert_eq!(auth_ctx.protocol, AUTH_PROTOCOL_DELAYED);
assert_eq!(auth_ctx.algorithm, AUTH_ALGORITHM_HMAC_SHA256);
}
#[test]
fn test_auth_context_hmac() {
use config::Dhcpv6AuthConfig;
let config = Dhcpv6AuthConfig {
enabled: true,
protocol: "delayed".to_string(),
algorithm: "hmac-sha256".to_string(),
key: Some("secret_key_for_testing_12345678".to_string()),
key_id: Some(1),
realm: None,
require_server_auth: false,
accept_unauthenticated: true,
};
let auth_ctx = Dhcpv6AuthContext::from_config(&config);
let test_message = b"Test DHCPv6 message data";
let hmac = auth_ctx.compute_hmac(test_message);
assert!(hmac.is_some());
let hmac = hmac.unwrap();
assert_eq!(hmac.len(), 32);
let hmac2 = auth_ctx.compute_hmac(test_message).unwrap();
assert_eq!(hmac, hmac2);
let hmac3 = auth_ctx.compute_hmac(b"Different message").unwrap();
assert_ne!(hmac, hmac3);
}
#[test]
fn test_auth_option_building() {
use config::Dhcpv6AuthConfig;
let config = Dhcpv6AuthConfig {
enabled: true,
protocol: "delayed".to_string(),
algorithm: "hmac-sha256".to_string(),
key: Some("test_key_123456789012345678901234".to_string()),
key_id: Some(42),
realm: Some("test.local".to_string()),
require_server_auth: false,
accept_unauthenticated: true,
};
let mut auth_ctx = Dhcpv6AuthContext::from_config(&config);
let test_msg = vec![
DHCPV6_SOLICIT, 0x01, 0x02, 0x03, 0x00, 0x01, 0x00, 0x04, 0xde, 0xad, 0xbe, 0xef, ];
let auth_option = auth_ctx.build_auth_option(&test_msg);
assert!(auth_option.is_some());
let auth_data = auth_option.unwrap();
assert_eq!(auth_data[0], AUTH_PROTOCOL_DELAYED);
assert_eq!(auth_data[1], AUTH_ALGORITHM_HMAC_SHA256);
assert_eq!(auth_data[2], RDM_MONOTONIC_COUNTER);
}
#[test]
fn test_tls_context_creation() {
use config::Dhcpv6TlsConfig;
let config = Dhcpv6TlsConfig::default();
let tls_ctx = Dhcpv6TlsContext::from_config(&config);
assert!(!tls_ctx.is_enabled());
assert!(!tls_ctx.use_tcp());
assert!(!tls_ctx.is_tls_established());
let config = Dhcpv6TlsConfig {
enabled: true,
use_tcp: true,
require_tls: true,
verify_server: true,
server_name: Some("dhcp.example.com".to_string()),
ca_cert: None,
client_cert: None,
client_key: None,
min_version: "1.3".to_string(),
tcp_port: 547,
};
let tls_ctx = Dhcpv6TlsContext::from_config(&config);
assert!(tls_ctx.is_enabled());
assert!(tls_ctx.use_tcp());
}
#[test]
fn test_starttls_message_building() {
let xid = [0x12, 0x34, 0x56];
let msg = Dhcpv6TlsContext::build_starttls_message(&xid);
assert_eq!(msg.len(), 4);
assert_eq!(msg[0], DHCPV6_STARTTLS);
assert_eq!(&msg[1..4], &xid);
}
#[test]
fn test_starttls_response_parsing() {
let xid = [0xAB, 0xCD, 0xEF];
let valid_response = vec![DHCPV6_STARTTLS, 0xAB, 0xCD, 0xEF];
assert!(Dhcpv6TlsContext::parse_starttls_response(&valid_response, &xid));
let wrong_type = vec![DHCPV6_SOLICIT, 0xAB, 0xCD, 0xEF];
assert!(!Dhcpv6TlsContext::parse_starttls_response(&wrong_type, &xid));
let wrong_xid = vec![DHCPV6_STARTTLS, 0x00, 0x00, 0x00];
assert!(!Dhcpv6TlsContext::parse_starttls_response(&wrong_xid, &xid));
let too_short = vec![DHCPV6_STARTTLS, 0xAB, 0xCD];
assert!(!Dhcpv6TlsContext::parse_starttls_response(&too_short, &xid));
let empty: Vec<u8> = vec![];
assert!(!Dhcpv6TlsContext::parse_starttls_response(&empty, &xid));
}
}