use super::*;
use std::collections::HashMap;
use std::time::Instant;
#[derive(Debug, Clone)]
pub struct SecurityContext {
config: SecurityConfig,
request_history: Arc<Mutex<HashMap<String, Vec<Instant>>>>,
}
impl SecurityContext {
pub fn new(config: SecurityConfig) -> Self {
Self {
config,
request_history: Arc::new(Mutex::new(HashMap::new())),
}
}
pub fn validate_server(&self, server_addr: &IpAddr) -> Result<()> {
if !self.config.validate_server {
return Ok(());
}
if self.config.allowed_servers.is_empty() {
return Ok(());
}
if self.config.allowed_servers.contains(server_addr) {
Ok(())
} else {
Err(DhcpClientError::SecurityViolation(
format!("Server {} is not in allowed list", server_addr)
))
}
}
pub async fn check_rate_limit(&self, interface: &str) -> Result<()> {
if self.config.max_request_rate == 0 {
return Ok(());
}
let mut history = self.request_history.lock().await;
let now = Instant::now();
let requests = history.entry(interface.to_string()).or_insert_with(Vec::new);
requests.retain(|&time| now.duration_since(time).as_secs() < 1);
if requests.len() >= self.config.max_request_rate as usize {
return Err(DhcpClientError::SecurityViolation(
format!("Rate limit exceeded for interface {}", interface)
));
}
requests.push(now);
Ok(())
}
pub fn validate_lease_time(&self, lease_time: u32) -> Result<()> {
if lease_time < self.config.min_lease_time {
return Err(DhcpClientError::SecurityViolation(
format!("Lease time {} is below minimum {}", lease_time, self.config.min_lease_time)
));
}
if lease_time > self.config.max_lease_time {
return Err(DhcpClientError::SecurityViolation(
format!("Lease time {} exceeds maximum {}", lease_time, self.config.max_lease_time)
));
}
Ok(())
}
pub fn validate_ipv4_address(&self, addr: &Ipv4Addr) -> Result<()> {
if addr.is_unspecified() {
return Err(DhcpClientError::SecurityViolation(
"Unspecified IP address (0.0.0.0)".to_string()
));
}
if addr.is_broadcast() {
return Err(DhcpClientError::SecurityViolation(
"Broadcast IP address (255.255.255.255)".to_string()
));
}
if addr.is_multicast() {
return Err(DhcpClientError::SecurityViolation(
format!("Multicast IP address {}", addr)
));
}
Ok(())
}
pub fn validate_ipv6_address(&self, addr: &Ipv6Addr) -> Result<()> {
if addr.is_unspecified() {
return Err(DhcpClientError::SecurityViolation(
"Unspecified IPv6 address (::)".to_string()
));
}
if addr.is_multicast() {
return Err(DhcpClientError::SecurityViolation(
format!("Multicast IPv6 address {}", addr)
));
}
Ok(())
}
pub fn validate_mac_address(&self, mac: &[u8; 6]) -> Result<()> {
if mac.iter().all(|&b| b == 0) {
return Err(DhcpClientError::SecurityViolation(
"Invalid MAC address (all zeros)".to_string()
));
}
if mac.iter().all(|&b| b == 0xFF) {
return Err(DhcpClientError::SecurityViolation(
"Invalid MAC address (broadcast)".to_string()
));
}
Ok(())
}
pub fn generate_xid() -> u32 {
use rand::Rng;
let mut rng = rand::rng();
rng.random()
}
pub fn config(&self) -> &SecurityConfig {
&self.config
}
}
pub fn sanitize_domain_name(domain: &str) -> Option<String> {
if domain.is_empty() || domain.len() > 253 {
return None;
}
let labels: Vec<&str> = domain.split('.').collect();
if labels.is_empty() {
return None;
}
for label in &labels {
if label.is_empty() || label.len() > 63 {
return None;
}
let first = label.chars().next()?;
if !first.is_ascii_alphanumeric() {
return None;
}
let last = label.chars().last()?;
if !last.is_ascii_alphanumeric() {
return None;
}
if !label.chars().all(|c| c.is_ascii_alphanumeric() || c == '-') {
return None;
}
}
Some(domain.to_ascii_lowercase())
}
pub fn sanitize_hostname(hostname: &str) -> Option<String> {
if hostname.is_empty() || hostname.len() > 63 {
return None;
}
let first = hostname.chars().next()?;
if !first.is_ascii_alphanumeric() {
return None;
}
let last = hostname.chars().last()?;
if !last.is_ascii_alphanumeric() {
return None;
}
if !hostname.chars().all(|c| c.is_ascii_alphanumeric() || c == '-') {
return None;
}
Some(hostname.to_ascii_lowercase())
}
pub fn sanitize_resolv_conf_value(value: &str) -> Option<String> {
let sanitized: String = value
.chars()
.filter(|c| c.is_ascii_alphanumeric() || *c == '.' || *c == '-' || *c == '_')
.collect();
if sanitized.is_empty() || sanitized.len() > 253 {
return None;
}
Some(sanitized)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_validate_ipv4() {
let sec = SecurityContext::new(SecurityConfig::default());
assert!(sec.validate_ipv4_address(&Ipv4Addr::new(192, 168, 1, 1)).is_ok());
assert!(sec.validate_ipv4_address(&Ipv4Addr::new(10, 0, 0, 1)).is_ok());
assert!(sec.validate_ipv4_address(&Ipv4Addr::new(0, 0, 0, 0)).is_err());
assert!(sec.validate_ipv4_address(&Ipv4Addr::new(255, 255, 255, 255)).is_err());
assert!(sec.validate_ipv4_address(&Ipv4Addr::new(224, 0, 0, 1)).is_err()); }
#[test]
fn test_validate_mac() {
let sec = SecurityContext::new(SecurityConfig::default());
assert!(sec.validate_mac_address(&[0x00, 0x11, 0x22, 0x33, 0x44, 0x55]).is_ok());
assert!(sec.validate_mac_address(&[0x00, 0x00, 0x00, 0x00, 0x00, 0x00]).is_err());
assert!(sec.validate_mac_address(&[0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF]).is_err());
}
#[test]
fn test_validate_lease_time() {
let sec = SecurityContext::new(SecurityConfig::default());
assert!(sec.validate_lease_time(3600).is_ok());
assert!(sec.validate_lease_time(86400).is_ok());
assert!(sec.validate_lease_time(60).is_err()); assert!(sec.validate_lease_time(86400 * 30).is_err()); }
#[tokio::test]
async fn test_rate_limiting() {
let mut config = SecurityConfig::default();
config.max_request_rate = 2;
let sec = SecurityContext::new(config);
assert!(sec.check_rate_limit("eth0").await.is_ok());
assert!(sec.check_rate_limit("eth0").await.is_ok());
assert!(sec.check_rate_limit("eth0").await.is_err());
assert!(sec.check_rate_limit("eth1").await.is_ok());
}
#[test]
fn test_server_validation() {
let mut config = SecurityConfig::default();
config.validate_server = true; config.allowed_servers = vec![IpAddr::V4(Ipv4Addr::new(192, 168, 1, 1))];
let sec = SecurityContext::new(config);
assert!(sec.validate_server(&IpAddr::V4(Ipv4Addr::new(192, 168, 1, 1))).is_ok());
assert!(sec.validate_server(&IpAddr::V4(Ipv4Addr::new(192, 168, 1, 2))).is_err());
}
#[test]
fn test_sanitize_domain_name() {
assert_eq!(sanitize_domain_name("example.com"), Some("example.com".to_string()));
assert_eq!(sanitize_domain_name("sub.example.com"), Some("sub.example.com".to_string()));
assert_eq!(sanitize_domain_name("my-domain.org"), Some("my-domain.org".to_string()));
assert_eq!(sanitize_domain_name("Example.COM"), Some("example.com".to_string()));
assert_eq!(sanitize_domain_name(""), None); assert_eq!(sanitize_domain_name("-example.com"), None); assert_eq!(sanitize_domain_name("example-.com"), None); assert_eq!(sanitize_domain_name("exam ple.com"), None); assert_eq!(sanitize_domain_name("example.com\nmalicious"), None); assert_eq!(sanitize_domain_name("example.com;rm -rf /"), None); assert_eq!(sanitize_domain_name(&"a".repeat(254)), None); }
#[test]
fn test_sanitize_hostname() {
assert_eq!(sanitize_hostname("myhost"), Some("myhost".to_string()));
assert_eq!(sanitize_hostname("my-host"), Some("my-host".to_string()));
assert_eq!(sanitize_hostname("host1"), Some("host1".to_string()));
assert_eq!(sanitize_hostname("MyHost"), Some("myhost".to_string()));
assert_eq!(sanitize_hostname(""), None); assert_eq!(sanitize_hostname("-myhost"), None); assert_eq!(sanitize_hostname("myhost-"), None); assert_eq!(sanitize_hostname("my host"), None); assert_eq!(sanitize_hostname("host\nmalicious"), None); assert_eq!(sanitize_hostname(&"a".repeat(64)), None); }
#[test]
fn test_sanitize_resolv_conf_value() {
assert_eq!(sanitize_resolv_conf_value("example.com"), Some("example.com".to_string()));
assert_eq!(sanitize_resolv_conf_value("my-domain_test.org"), Some("my-domain_test.org".to_string()));
assert_eq!(sanitize_resolv_conf_value("example.com\nmalicious"), Some("example.commalicious".to_string()));
assert_eq!(sanitize_resolv_conf_value("example.com; rm -rf /"), Some("example.comrm-rf".to_string()));
assert_eq!(sanitize_resolv_conf_value("$(whoami)"), Some("whoami".to_string()));
assert_eq!(sanitize_resolv_conf_value(""), None);
assert_eq!(sanitize_resolv_conf_value(" "), None); }
}