crdhcpc 0.1.1

Standalone DHCP Client for Linux with DHCPv4, DHCPv6, PXE, and Dynamic DNS support
Documentation
//! Security features for DHCP client
//!
//! Implements:
//! - Message validation
//! - Server authentication
//! - Rate limiting
//! - DHCP snooping
//! - Anti-spoofing

use super::*;
use std::collections::HashMap;
use std::time::Instant;

/// Security context for DHCP client
#[derive(Debug, Clone)]
pub struct SecurityContext {
    config: SecurityConfig,
    request_history: Arc<Mutex<HashMap<String, Vec<Instant>>>>,
}

impl SecurityContext {
    /// Create new security context
    pub fn new(config: SecurityConfig) -> Self {
        Self {
            config,
            request_history: Arc::new(Mutex::new(HashMap::new())),
        }
    }

    /// Validate a DHCP server address
    pub fn validate_server(&self, server_addr: &IpAddr) -> Result<()> {
        if !self.config.validate_server {
            return Ok(());
        }

        // If allowlist is empty, accept all servers
        if self.config.allowed_servers.is_empty() {
            return Ok(());
        }

        // Check if server is in allowlist
        if self.config.allowed_servers.contains(server_addr) {
            Ok(())
        } else {
            Err(DhcpClientError::SecurityViolation(
                format!("Server {} is not in allowed list", server_addr)
            ))
        }
    }

    /// Check rate limiting for interface
    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();

        // Get or create history for this interface
        let requests = history.entry(interface.to_string()).or_insert_with(Vec::new);

        // Remove requests older than 1 second
        requests.retain(|&time| now.duration_since(time).as_secs() < 1);

        // Check if rate limit exceeded
        if requests.len() >= self.config.max_request_rate as usize {
            return Err(DhcpClientError::SecurityViolation(
                format!("Rate limit exceeded for interface {}", interface)
            ));
        }

        // Add current request
        requests.push(now);

        Ok(())
    }

    /// Validate lease time
    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(())
    }

    /// Validate IPv4 address
    pub fn validate_ipv4_address(&self, addr: &Ipv4Addr) -> Result<()> {
        // Reject invalid addresses
        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)
            ));
        }

        // Note: We allow link-local and private addresses as they're valid for DHCP

        Ok(())
    }

    /// Validate IPv6 address
    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(())
    }

    /// Validate MAC address
    pub fn validate_mac_address(&self, mac: &[u8; 6]) -> Result<()> {
        // Check for all-zeros
        if mac.iter().all(|&b| b == 0) {
            return Err(DhcpClientError::SecurityViolation(
                "Invalid MAC address (all zeros)".to_string()
            ));
        }

        // Check for all-ones (broadcast)
        if mac.iter().all(|&b| b == 0xFF) {
            return Err(DhcpClientError::SecurityViolation(
                "Invalid MAC address (broadcast)".to_string()
            ));
        }

        Ok(())
    }

    /// Generate secure transaction ID
    pub fn generate_xid() -> u32 {
        use rand::Rng;
        let mut rng = rand::rng();
        rng.random()
    }

    /// Get security config
    pub fn config(&self) -> &SecurityConfig {
        &self.config
    }
}

/// Validate and sanitize a domain name per RFC 1123
/// Returns sanitized domain name or None if invalid
pub fn sanitize_domain_name(domain: &str) -> Option<String> {
    // RFC 1123: max 253 characters total
    if domain.is_empty() || domain.len() > 253 {
        return None;
    }

    // Split into labels and validate each
    let labels: Vec<&str> = domain.split('.').collect();

    // Must have at least one label
    if labels.is_empty() {
        return None;
    }

    for label in &labels {
        // Each label: 1-63 characters
        if label.is_empty() || label.len() > 63 {
            return None;
        }

        // Must start with alphanumeric
        let first = label.chars().next()?;
        if !first.is_ascii_alphanumeric() {
            return None;
        }

        // Must end with alphanumeric
        let last = label.chars().last()?;
        if !last.is_ascii_alphanumeric() {
            return None;
        }

        // All characters must be alphanumeric or hyphen
        if !label.chars().all(|c| c.is_ascii_alphanumeric() || c == '-') {
            return None;
        }
    }

    // Return lowercase sanitized domain
    Some(domain.to_ascii_lowercase())
}

/// Validate and sanitize a hostname per RFC 952/1123
/// Returns sanitized hostname or None if invalid
pub fn sanitize_hostname(hostname: &str) -> Option<String> {
    // RFC 1123: max 63 characters for a single hostname label
    if hostname.is_empty() || hostname.len() > 63 {
        return None;
    }

    // Must start with alphanumeric
    let first = hostname.chars().next()?;
    if !first.is_ascii_alphanumeric() {
        return None;
    }

    // Must end with alphanumeric
    let last = hostname.chars().last()?;
    if !last.is_ascii_alphanumeric() {
        return None;
    }

    // All characters must be alphanumeric or hyphen
    if !hostname.chars().all(|c| c.is_ascii_alphanumeric() || c == '-') {
        return None;
    }

    // Return lowercase sanitized hostname
    Some(hostname.to_ascii_lowercase())
}

/// Sanitize a string for safe inclusion in resolv.conf
/// Removes any characters that could be used for injection
pub fn sanitize_resolv_conf_value(value: &str) -> Option<String> {
    // Remove any control characters, newlines, or shell metacharacters
    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());

        // Valid addresses
        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());

        // Invalid addresses
        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()); // Multicast
    }

    #[test]
    fn test_validate_mac() {
        let sec = SecurityContext::new(SecurityConfig::default());

        // Valid MAC
        assert!(sec.validate_mac_address(&[0x00, 0x11, 0x22, 0x33, 0x44, 0x55]).is_ok());

        // Invalid MACs
        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());

        // Valid lease times
        assert!(sec.validate_lease_time(3600).is_ok());
        assert!(sec.validate_lease_time(86400).is_ok());

        // Invalid lease times
        assert!(sec.validate_lease_time(60).is_err()); // Too short
        assert!(sec.validate_lease_time(86400 * 30).is_err()); // Too long
    }

    #[tokio::test]
    async fn test_rate_limiting() {
        let mut config = SecurityConfig::default();
        config.max_request_rate = 2;
        let sec = SecurityContext::new(config);

        // First two requests should succeed
        assert!(sec.check_rate_limit("eth0").await.is_ok());
        assert!(sec.check_rate_limit("eth0").await.is_ok());

        // Third request should fail
        assert!(sec.check_rate_limit("eth0").await.is_err());

        // Different interface should succeed
        assert!(sec.check_rate_limit("eth1").await.is_ok());
    }

    #[test]
    fn test_server_validation() {
        let mut config = SecurityConfig::default();
        config.validate_server = true;  // Enable server validation
        config.allowed_servers = vec![IpAddr::V4(Ipv4Addr::new(192, 168, 1, 1))];
        let sec = SecurityContext::new(config);

        // Allowed server
        assert!(sec.validate_server(&IpAddr::V4(Ipv4Addr::new(192, 168, 1, 1))).is_ok());

        // Not allowed server
        assert!(sec.validate_server(&IpAddr::V4(Ipv4Addr::new(192, 168, 1, 2))).is_err());
    }

    #[test]
    fn test_sanitize_domain_name() {
        // Valid domains
        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())); // Lowercased

        // Invalid domains
        assert_eq!(sanitize_domain_name(""), None); // Empty
        assert_eq!(sanitize_domain_name("-example.com"), None); // Starts with hyphen
        assert_eq!(sanitize_domain_name("example-.com"), None); // Ends with hyphen
        assert_eq!(sanitize_domain_name("exam ple.com"), None); // Contains space
        assert_eq!(sanitize_domain_name("example.com\nmalicious"), None); // Newline injection
        assert_eq!(sanitize_domain_name("example.com;rm -rf /"), None); // Command injection
        assert_eq!(sanitize_domain_name(&"a".repeat(254)), None); // Too long
    }

    #[test]
    fn test_sanitize_hostname() {
        // Valid hostnames
        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())); // Lowercased

        // Invalid hostnames
        assert_eq!(sanitize_hostname(""), None); // Empty
        assert_eq!(sanitize_hostname("-myhost"), None); // Starts with hyphen
        assert_eq!(sanitize_hostname("myhost-"), None); // Ends with hyphen
        assert_eq!(sanitize_hostname("my host"), None); // Contains space
        assert_eq!(sanitize_hostname("host\nmalicious"), None); // Newline injection
        assert_eq!(sanitize_hostname(&"a".repeat(64)), None); // Too long
    }

    #[test]
    fn test_sanitize_resolv_conf_value() {
        // Valid values
        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()));

        // Injection attempts - should be sanitized
        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()));

        // Empty after sanitization
        assert_eq!(sanitize_resolv_conf_value(""), None);
        assert_eq!(sanitize_resolv_conf_value("   "), None); // Only spaces
    }
}