Initial commit: all files from local directory
This commit is contained in:
commit
9491dafcca
33 files changed
+6027
No files matched your search
+102
@@ -0,0 +1,102 @@
|
||||
//! DNS cache implementation.
|
||||
//!
|
||||
//! This module provides a simple in-memory cache for DNS records to improve
|
||||
//! performance by avoiding repeated database lookups for frequently accessed domains.
|
||||
#![allow(dead_code)]
|
||||
#[allow(unused_variables)]
|
||||
|
||||
use std::{
|
||||
collections::HashMap,
|
||||
sync::{Arc, Mutex, OnceLock},
|
||||
time::{Duration, SystemTime},
|
||||
};
|
||||
use log::debug;
|
||||
|
||||
/// Interval for cleaning up expired cache entries (in seconds).
|
||||
pub const CACHE_CLEANUP_INTERVAL: Duration = Duration::from_secs(300);
|
||||
|
||||
/// Global cache instance.
|
||||
pub static CACHE: OnceLock<DnsCache> = OnceLock::new();
|
||||
|
||||
/// An entry in the DNS cache.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct CacheEntry {
|
||||
/// The IP address for the domain.
|
||||
pub ip: String,
|
||||
|
||||
/// When this entry was added to the cache.
|
||||
pub inserted: SystemTime,
|
||||
|
||||
/// Time-to-live in seconds.
|
||||
pub ttl: u64,
|
||||
}
|
||||
|
||||
/// Cache for DNS records to improve performance.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct DnsCache {
|
||||
/// Map of domain names to cache entries.
|
||||
pub entries: Arc<Mutex<HashMap<String, CacheEntry>>>,
|
||||
|
||||
/// List of NS records for zones this server is authoritative for.
|
||||
pub ns_records: Vec<String>,
|
||||
}
|
||||
|
||||
impl DnsCache {
|
||||
/// Create a new DNS cache.
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `ns_records` - List of NS records for zones this server is authoritative for.
|
||||
///
|
||||
/// # Returns
|
||||
/// A new `DnsCache` instance.
|
||||
pub fn new(ns_records: Vec<String>) -> Self {
|
||||
Self {
|
||||
entries: Arc::new(Mutex::new(HashMap::new())),
|
||||
ns_records,
|
||||
}
|
||||
}
|
||||
|
||||
/// Get a cached IP address for a domain.
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `domain` - The domain name to look up.
|
||||
///
|
||||
/// # Returns
|
||||
/// An `Option` containing the IP address and TTL if found and not expired.
|
||||
pub fn get(&self, domain: &str) -> Option<(String, u64)> {
|
||||
let cache = self.entries.lock().unwrap();
|
||||
if let Some(entry) = cache.get(domain) {
|
||||
if entry.inserted.elapsed().map(|d| d.as_secs() <= entry.ttl).unwrap_or(true) {
|
||||
return Some((entry.ip.clone(), entry.ttl));
|
||||
}
|
||||
}
|
||||
None
|
||||
}
|
||||
|
||||
/// Add or update a domain in the cache.
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `domain` - The domain name to cache.
|
||||
/// * `ip` - The IP address for the domain.
|
||||
/// * `ttl` - Time-to-live in seconds.
|
||||
pub fn set(&self, domain: String, ip: String, ttl: u64) {
|
||||
let mut cache = self.entries.lock().unwrap();
|
||||
cache.insert(
|
||||
domain,
|
||||
CacheEntry {
|
||||
ip,
|
||||
inserted: SystemTime::now(),
|
||||
ttl,
|
||||
},
|
||||
);
|
||||
}
|
||||
|
||||
/// Remove expired entries from the cache.
|
||||
pub fn cleanup(&self) {
|
||||
let mut cache = self.entries.lock().unwrap();
|
||||
cache.retain(|_, entry| {
|
||||
entry.inserted.elapsed().map(|d| d.as_secs() <= entry.ttl).unwrap_or(true)
|
||||
});
|
||||
debug!("Cache cleanup completed");
|
||||
}
|
||||
}
|
||||
+117
@@ -0,0 +1,117 @@
|
||||
//! Configuration for the DNS server.
|
||||
//!
|
||||
//! This module defines the configuration structure and methods to load
|
||||
//! configuration from environment variables.
|
||||
#![allow(dead_code)]
|
||||
#[allow(unused_variables)]
|
||||
|
||||
use std::{env, fs, net::SocketAddr};
|
||||
use log::{error, info};
|
||||
|
||||
use crate::errors::DnsError;
|
||||
|
||||
/// Default TTL for DNS records in seconds.
|
||||
pub const DEFAULT_TTL: u64 = 600;
|
||||
|
||||
/// Maximum size of DNS packets in bytes.
|
||||
pub const MAX_PACKET_SIZE: usize = 4096;
|
||||
|
||||
/// Server configuration loaded from environment variables.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct ServerConfig {
|
||||
/// Address to bind the DNS server to.
|
||||
pub bind_addr: SocketAddr,
|
||||
|
||||
/// Path to the SQLite database file.
|
||||
pub db_path: String,
|
||||
|
||||
/// Time-to-live for cached DNS records.
|
||||
pub cache_ttl: u64,
|
||||
|
||||
/// Whether to enable IPv6 support.
|
||||
pub enable_ipv6: bool,
|
||||
|
||||
/// Maximum size of DNS packets.
|
||||
pub max_packet_size: usize,
|
||||
|
||||
/// Whether this server is authoritative for its zones.
|
||||
pub authoritative: bool,
|
||||
|
||||
/// List of NS records for zones this server is authoritative for.
|
||||
pub ns_records: Vec<String>,
|
||||
|
||||
/// Default domain for the server.
|
||||
pub default_domain: String,
|
||||
|
||||
/// Default IP address for the server.
|
||||
pub default_ip: String,
|
||||
|
||||
/// List of upstream DNS servers to forward queries to.
|
||||
pub forwarders: Vec<SocketAddr>,
|
||||
|
||||
/// List of DS records for DNSSEC.
|
||||
pub ds_records: Vec<String>,
|
||||
|
||||
/// List of DNSKEY records for DNSSEC.
|
||||
pub dnskey_records: Vec<String>,
|
||||
}
|
||||
|
||||
impl ServerConfig {
|
||||
/// Load server configuration from environment variables.
|
||||
///
|
||||
/// # Returns
|
||||
/// A `Result` containing either the loaded `ServerConfig` or a `DnsError`.
|
||||
pub fn from_env() -> Result<Self, DnsError> {
|
||||
let bind_addr = env::var("DNS_BIND")
|
||||
.unwrap_or_else(|_| "0.0.0.0:53".into())
|
||||
.parse()
|
||||
.map_err(|_| DnsError::Config("Invalid DNS_BIND address".into()))?;
|
||||
|
||||
let forwarders = env::var("DNS_FORWARDERS")
|
||||
.unwrap_or_else(|_| "8.8.8.8:53,1.1.1.1:53,9.9.9.9:53".into())
|
||||
.split(',')
|
||||
.filter_map(|s| s.trim().parse().ok())
|
||||
.collect();
|
||||
|
||||
let key_path = env::var("DNSSEC_KEY_FILE").unwrap_or_else(|_| "Kbzo.in.+008+24550.key".to_string());
|
||||
let dnskey_records = match fs::read_to_string(&key_path) {
|
||||
Ok(content) => {
|
||||
info!("Loaded DNSSEC key from {}", key_path);
|
||||
vec![content.trim().to_string()]
|
||||
},
|
||||
Err(e) => {
|
||||
error!("Failed to load DNSSEC key from {}: {}", key_path, e);
|
||||
vec![]
|
||||
}
|
||||
};
|
||||
|
||||
Ok(Self {
|
||||
bind_addr,
|
||||
db_path: env::var("DNS_DB_PATH").unwrap_or_else(|_| "dns.db".into()),
|
||||
cache_ttl: env::var("DNS_CACHE_TTL")
|
||||
.ok()
|
||||
.and_then(|v| v.parse().ok())
|
||||
.unwrap_or(DEFAULT_TTL),
|
||||
enable_ipv6: env::var("DNS_ENABLE_IPV6")
|
||||
.map(|v| v == "1" || v.eq_ignore_ascii_case("true"))
|
||||
.unwrap_or(false),
|
||||
max_packet_size: env::var("DNS_MAX_PACKET_SIZE")
|
||||
.ok()
|
||||
.and_then(|v| v.parse().ok())
|
||||
.unwrap_or(MAX_PACKET_SIZE),
|
||||
authoritative: env::var("DNS_AUTHORITATIVE")
|
||||
.map(|v| v == "1" || v.eq_ignore_ascii_case("true"))
|
||||
.unwrap_or(false),
|
||||
ns_records: env::var("DNS_NS_RECORDS")
|
||||
.map(|v| v.split(',').map(|s| s.trim().to_string()).collect())
|
||||
.unwrap_or_else(|_| vec!["ns1.yourdomain.tld.".into(), "ns2.yourdomain.tld.".into()]),
|
||||
default_domain: env::var("DNS_DEFAULT_DOMAIN").unwrap_or_else(|_| "bzo.in".into()),
|
||||
default_ip: env::var("DNS_DEFAULT_IP").unwrap_or_else(|_| "<your-public-ip4-here>".into()),
|
||||
ds_records: vec![
|
||||
"yourdomain.tld. IN DS 24550 8 2 1F21CA282945434EE0662805430599CB2A6C479D9F934087150901CE2DA580A0".to_string()
|
||||
],
|
||||
dnskey_records,
|
||||
forwarders,
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,224 @@
|
||||
//! Database operations for the DNS server.
|
||||
//!
|
||||
//! This module provides functions for interacting with the SQLite database
|
||||
//! that stores DNS records and zone information.
|
||||
|
||||
use rusqlite::{params, Connection};
|
||||
|
||||
use crate::errors::DnsError;
|
||||
use crate::config::ServerConfig;
|
||||
|
||||
/// Information about a DNS zone.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct ZoneInfo {
|
||||
/// The domain name of the zone.
|
||||
pub name: String,
|
||||
|
||||
/// List of NS records for the zone.
|
||||
pub ns_records: Vec<String>,
|
||||
|
||||
/// SOA record for the zone, if available.
|
||||
pub soa_record: Option<String>,
|
||||
}
|
||||
|
||||
/// Initialize the DNS database.
|
||||
///
|
||||
/// Creates the database schema if it doesn't exist and populates it with default records
|
||||
/// if the database is empty.
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `db_path` - Path to the SQLite database file.
|
||||
/// * `default_domain` - Default domain name to use for initial records.
|
||||
/// * `default_ip` - Default IP address to use for initial records.
|
||||
///
|
||||
/// # Returns
|
||||
/// A `Result` indicating success or failure.
|
||||
pub fn init_db(db_path: &str, default_domain: &str, default_ip: &str) -> Result<(), DnsError> {
|
||||
let conn = Connection::open(db_path)?;
|
||||
|
||||
// Updated schema to allow multiple NS records
|
||||
conn.execute(
|
||||
"CREATE TABLE IF NOT EXISTS dns_records (
|
||||
domain TEXT NOT NULL,
|
||||
record_type TEXT NOT NULL CHECK(record_type IN (
|
||||
'A','AAAA','MX','TXT','NS','CNAME','PTR','SOA',
|
||||
'SRV','CAA','NAPTR','DS','DNSKEY','RRSIG','NSEC',
|
||||
'TLSA','SSHFP'
|
||||
)),
|
||||
value TEXT NOT NULL,
|
||||
ttl INTEGER DEFAULT 3600,
|
||||
PRIMARY KEY (domain, record_type, value) -- Now allows multiple NS records
|
||||
) WITHOUT ROWID",
|
||||
[],
|
||||
)?;
|
||||
|
||||
let count: i64 = conn.query_row("SELECT COUNT(*) FROM dns_records", [], |row| row.get(0))?;
|
||||
|
||||
if count == 0 && !default_ip.is_empty() {
|
||||
let mail_domain = format!("mail.{}", default_domain);
|
||||
let ns1 = format!("ns1.{}", default_domain);
|
||||
let ns2 = format!("ns2.{}", default_domain);
|
||||
let soa_record = format!("{} hostmaster.{} 1 10800 3600 604800 86400", ns1, default_domain);
|
||||
|
||||
conn.execute_batch(
|
||||
&format!(
|
||||
r#"
|
||||
INSERT OR IGNORE INTO dns_records VALUES('{0}', 'A', ?, 3600);
|
||||
INSERT OR IGNORE INTO dns_records VALUES('www.{0}', 'A', ?, 3600);
|
||||
INSERT OR IGNORE INTO dns_records VALUES('api.{0}', 'A', ?, 3600);
|
||||
INSERT OR IGNORE INTO dns_records VALUES('mail.{0}', 'A', ?, 3600);
|
||||
INSERT OR IGNORE INTO dns_records VALUES('ns1.{0}', 'A', ?, 3600);
|
||||
INSERT OR IGNORE INTO dns_records VALUES('ns2.{0}', 'A', ?, 3600);
|
||||
INSERT OR IGNORE INTO dns_records VALUES('{0}', 'MX', '10 {1}', 3600);
|
||||
INSERT OR IGNORE INTO dns_records VALUES('{0}', 'TXT', '\"v=spf1 a mx ~all\"', 3600);
|
||||
INSERT OR IGNORE INTO dns_records VALUES('{0}', 'NS', '{2}', 3600);
|
||||
INSERT OR IGNORE INTO dns_records VALUES('{0}', 'NS', '{3}', 3600);
|
||||
INSERT OR IGNORE INTO dns_records VALUES('{0}', 'SOA', '{4}', 3600);
|
||||
"#,
|
||||
default_domain, mail_domain, ns1, ns2, soa_record
|
||||
),
|
||||
)?;
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Look up DNS records for a domain.
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `db_path` - Path to the SQLite database file.
|
||||
/// * `domain` - Domain name to look up.
|
||||
///
|
||||
/// # Returns
|
||||
/// A vector of tuples containing (value, ttl, record_type) for each record found.
|
||||
pub fn lookup_records(db_path: &str, domain: &str) -> Vec<(String, u64, String)> {
|
||||
let conn = Connection::open(db_path);
|
||||
match conn {
|
||||
Ok(conn) => {
|
||||
match conn.prepare(
|
||||
"SELECT value, ttl, record_type FROM dns_records WHERE domain = ?"
|
||||
) {
|
||||
Ok(mut stmt) => {
|
||||
match stmt.query_map(params![domain], |row| {
|
||||
Ok((
|
||||
row.get(0).unwrap_or_default(),
|
||||
row.get(1).unwrap_or_default(),
|
||||
row.get(2).unwrap_or_default(),
|
||||
))
|
||||
}) {
|
||||
Ok(rows) => rows.filter_map(Result::ok).collect(),
|
||||
Err(_) => Vec::new(),
|
||||
}
|
||||
}
|
||||
Err(_) => Vec::new(),
|
||||
}
|
||||
}
|
||||
Err(_) => Vec::new(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Get information about all zones for which this server is authoritative.
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `db_path` - Path to the SQLite database file.
|
||||
///
|
||||
/// # Returns
|
||||
/// A vector of `ZoneInfo` structs containing information about each zone.
|
||||
pub fn get_authoritative_zones(db_path: &str) -> Vec<ZoneInfo> {
|
||||
let mut zones = Vec::new();
|
||||
|
||||
if let Ok(conn) = Connection::open(db_path) {
|
||||
// Find all domains with NS records (these are zones)
|
||||
if let Ok(mut stmt) = conn.prepare(
|
||||
"SELECT DISTINCT domain FROM dns_records WHERE record_type = 'NS'"
|
||||
) {
|
||||
if let Ok(rows) = stmt.query_map([], |row| {
|
||||
Ok(row.get::<_, String>(0)?)
|
||||
}) {
|
||||
for domain_result in rows {
|
||||
if let Ok(domain) = domain_result {
|
||||
let mut zone_info = ZoneInfo {
|
||||
name: domain.clone(),
|
||||
ns_records: Vec::new(),
|
||||
soa_record: None,
|
||||
};
|
||||
|
||||
// Get NS records for this zone
|
||||
if let Ok(mut ns_stmt) = conn.prepare(
|
||||
"SELECT value FROM dns_records WHERE domain = ? AND record_type = 'NS'"
|
||||
) {
|
||||
if let Ok(ns_rows) = ns_stmt.query_map([&domain], |row| {
|
||||
Ok(row.get::<_, String>(0)?)
|
||||
}) {
|
||||
zone_info.ns_records = ns_rows.filter_map(Result::ok).collect();
|
||||
}
|
||||
}
|
||||
|
||||
// Get SOA record if exists
|
||||
if let Ok(mut soa_stmt) = conn.prepare(
|
||||
"SELECT value FROM dns_records WHERE domain = ? AND record_type = 'SOA' LIMIT 1"
|
||||
) {
|
||||
if let Ok(mut soa_rows) = soa_stmt.query_map([&domain], |row| {
|
||||
Ok(row.get::<_, String>(0)?)
|
||||
}) {
|
||||
zone_info.soa_record = soa_rows.next().and_then(|r| r.ok());
|
||||
}
|
||||
}
|
||||
|
||||
zones.push(zone_info);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Add default zone information from config
|
||||
if zones.is_empty() {
|
||||
if let Ok(config) = ServerConfig::from_env() {
|
||||
let default_zone = ZoneInfo {
|
||||
name: config.default_domain.clone(),
|
||||
ns_records: config.ns_records.clone(),
|
||||
soa_record: Some(format!(
|
||||
"{} hostmaster.{} 1 10800 3600 604800 86400",
|
||||
config.ns_records.first().unwrap_or(&String::from("ns1.example.com.")),
|
||||
config.default_domain
|
||||
)),
|
||||
};
|
||||
zones.push(default_zone);
|
||||
}
|
||||
}
|
||||
|
||||
zones
|
||||
}
|
||||
|
||||
/// Find the closest parent zone for a given domain.
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `domain` - Domain name to find the parent zone for.
|
||||
/// * `zones` - List of zones to search in.
|
||||
///
|
||||
/// # Returns
|
||||
/// An `Option` containing the closest parent zone, if found.
|
||||
pub fn find_closest_parent_zone(domain: &str, zones: &[ZoneInfo]) -> Option<ZoneInfo> {
|
||||
let domain_parts: Vec<&str> = domain.split('.').collect();
|
||||
|
||||
// Try progressively shorter parent domains
|
||||
for i in 0..domain_parts.len() {
|
||||
let candidate = domain_parts[i..].join(".");
|
||||
|
||||
// Exact match
|
||||
if let Some(zone) = zones.iter().find(|z| z.name == candidate) {
|
||||
return Some(zone.clone());
|
||||
}
|
||||
}
|
||||
|
||||
// Check if domain is a subdomain of any zone we're authoritative for
|
||||
for zone in zones {
|
||||
if domain.ends_with(&format!(".{}", zone.name)) {
|
||||
return Some(zone.clone());
|
||||
}
|
||||
}
|
||||
|
||||
// If no match found, return None - we are not authoritative for this domain
|
||||
None
|
||||
}
|
||||
+1229
File diff suppressed because it is too large.
Load diff
@@ -0,0 +1,39 @@
|
||||
//! Error types for the DNS server
|
||||
#![allow(dead_code)]
|
||||
#[allow(unused_variables)]
|
||||
|
||||
use std::io;
|
||||
use rusqlite;
|
||||
use thiserror::Error;
|
||||
|
||||
/// Errors that can occur in the DNS server
|
||||
#[derive(Debug, Error)]
|
||||
pub enum DnsError {
|
||||
/// I/O errors from the underlying system
|
||||
#[error("I/O error: {0}")]
|
||||
Io(#[from] io::Error),
|
||||
|
||||
/// Database errors from SQLite operations
|
||||
#[error("Database error: {0}")]
|
||||
Db(#[from] rusqlite::Error),
|
||||
|
||||
/// Protocol errors related to DNS message format or content
|
||||
#[error("Protocol error: {0}")]
|
||||
Protocol(String),
|
||||
|
||||
/// Configuration errors from invalid settings
|
||||
#[error("Configuration error: {0}")]
|
||||
Config(String),
|
||||
|
||||
/// Parse errors from string to number conversions
|
||||
#[error("Parse error: {0}")]
|
||||
Parse(#[from] std::num::ParseIntError),
|
||||
|
||||
/// Base64 decoding errors
|
||||
#[error("Base64 error: {0}")]
|
||||
Base64(String),
|
||||
|
||||
/// Shutdown signal received
|
||||
#[error("Shutdown signal received")]
|
||||
Shutdown,
|
||||
}
|
||||
@@ -0,0 +1,39 @@
|
||||
//! Error types for the DNS server.
|
||||
//!
|
||||
//! This module defines the error types used throughout the DNS server implementation.
|
||||
#![allow(dead_code)]
|
||||
#[allow(unused_variables)]
|
||||
|
||||
use thiserror::Error;
|
||||
|
||||
/// Represents errors that can occur in the DNS server.
|
||||
#[derive(Error, Debug)]
|
||||
pub enum DnsError {
|
||||
/// I/O errors from the standard library.
|
||||
#[error("I/O error: {0}")]
|
||||
Io(#[from] std::io::Error),
|
||||
|
||||
/// Database errors from rusqlite.
|
||||
#[error("Database error: {0}")]
|
||||
Db(#[from] rusqlite::Error),
|
||||
|
||||
/// Errors related to DNS protocol parsing or formatting.
|
||||
#[error("Invalid DNS packet: {0}")]
|
||||
Protocol(String),
|
||||
|
||||
/// Configuration errors.
|
||||
#[error("Configuration error: {0}")]
|
||||
Config(String),
|
||||
|
||||
/// Integer parsing errors.
|
||||
#[error("Parse error: {0}")]
|
||||
Parse(#[from] std::num::ParseIntError),
|
||||
|
||||
/// Base64 decoding errors.
|
||||
#[error("Base64 error: {0}")]
|
||||
Base64(String),
|
||||
|
||||
/// Shutdown signal received.
|
||||
#[error("Shutdown signal received")]
|
||||
Shutdown,
|
||||
}
|
||||
+190
@@ -0,0 +1,190 @@
|
||||
//! Request handlers for the DNS server.
|
||||
//!
|
||||
//! This module provides functions for handling DNS requests over UDP and TCP.
|
||||
#![allow(dead_code)]
|
||||
#[allow(unused_variables)]
|
||||
|
||||
use std::net::SocketAddr;
|
||||
use std::sync::Arc;
|
||||
use log::{debug, error, info, warn};
|
||||
use tokio::{
|
||||
io::AsyncReadExt,
|
||||
net::{TcpListener, TcpStream, UdpSocket},
|
||||
task,
|
||||
};
|
||||
|
||||
use crate::errors::DnsError;
|
||||
use crate::config::ServerConfig;
|
||||
use crate::utils::extract_domain;
|
||||
use crate::dns::{
|
||||
build_not_implemented_response, build_nxdomain_response, generate_dns_response,
|
||||
send_tcp_response,
|
||||
};
|
||||
|
||||
/// Run the UDP DNS server.
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `config` - The server configuration.
|
||||
///
|
||||
/// # Returns
|
||||
/// A `Result` indicating success or failure.
|
||||
pub async fn run_udp_server(config: ServerConfig) -> Result<(), DnsError> {
|
||||
let socket = UdpSocket::bind(config.bind_addr).await?;
|
||||
info!("UDP DNS server listening on {}", config.bind_addr);
|
||||
let socket = Arc::new(socket);
|
||||
let mut buf = vec![0u8; config.max_packet_size];
|
||||
|
||||
loop {
|
||||
match socket.recv_from(&mut buf).await {
|
||||
Ok((amt, src)) => {
|
||||
let query = buf[..amt].to_vec();
|
||||
let socket = socket.clone();
|
||||
let config = config.clone();
|
||||
task::spawn(async move {
|
||||
if let Err(e) = handle_udp_query(query, src, socket, config).await {
|
||||
warn!("UDP query error: {}", e);
|
||||
}
|
||||
});
|
||||
}
|
||||
Err(e) => error!("UDP receive error: {}", e),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Handle a UDP DNS query.
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `query` - The DNS query.
|
||||
/// * `src` - The source address of the query.
|
||||
/// * `socket` - The UDP socket to send the response on.
|
||||
/// * `config` - The server configuration.
|
||||
///
|
||||
/// # Returns
|
||||
/// A `Result` indicating success or failure.
|
||||
pub async fn handle_udp_query(
|
||||
query: Vec<u8>,
|
||||
src: SocketAddr,
|
||||
socket: Arc<UdpSocket>,
|
||||
config: ServerConfig,
|
||||
) -> Result<(), DnsError> {
|
||||
if query.len() < 12 {
|
||||
debug!("Received malformed query from {}", src);
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let opcode = (query[2] & 0x78) >> 3;
|
||||
if opcode != 0 {
|
||||
if let Some(response) = build_not_implemented_response(&query, config.authoritative) {
|
||||
socket.send_to(&response, src).await?;
|
||||
}
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let domain = match extract_domain(&query) {
|
||||
Some(d) => d,
|
||||
None => {
|
||||
info!("Failed to extract domain from query");
|
||||
return Ok(());
|
||||
}
|
||||
};
|
||||
|
||||
debug!("UDP query for {} from {}", domain, src);
|
||||
info!("Processing query for domain: {}", domain);
|
||||
|
||||
let response = match generate_dns_response(&query, domain.clone(), &config).await {
|
||||
Ok(resp) => resp,
|
||||
Err(_) => {
|
||||
build_nxdomain_response(&query, config.authoritative)
|
||||
.ok_or(DnsError::Protocol("NXDOMAIN".into()))?
|
||||
}
|
||||
};
|
||||
|
||||
socket.send_to(&response, src).await?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Run the TCP DNS server.
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `config` - The server configuration.
|
||||
///
|
||||
/// # Returns
|
||||
/// A `Result` indicating success or failure.
|
||||
pub async fn run_tcp_server(config: ServerConfig) -> Result<(), DnsError> {
|
||||
let listener = TcpListener::bind(config.bind_addr).await?;
|
||||
info!("TCP DNS server listening on {}", config.bind_addr);
|
||||
|
||||
loop {
|
||||
match listener.accept().await {
|
||||
Ok((stream, addr)) => {
|
||||
let config = config.clone();
|
||||
task::spawn(async move {
|
||||
if let Err(e) = handle_tcp_connection(stream, addr, config).await {
|
||||
warn!("TCP connection error: {}", e);
|
||||
}
|
||||
});
|
||||
}
|
||||
Err(e) => error!("TCP accept error: {}", e),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Handle a TCP DNS connection.
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `stream` - The TCP stream.
|
||||
/// * `addr` - The client address.
|
||||
/// * `config` - The server configuration.
|
||||
///
|
||||
/// # Returns
|
||||
/// A `Result` indicating success or failure.
|
||||
pub async fn handle_tcp_connection(
|
||||
mut stream: TcpStream,
|
||||
addr: SocketAddr,
|
||||
config: ServerConfig,
|
||||
) -> Result<(), DnsError> {
|
||||
// Read the 2-byte length prefix
|
||||
let mut len_buf = [0u8; 2];
|
||||
stream.read_exact(&mut len_buf).await?;
|
||||
let len = u16::from_be_bytes(len_buf) as usize;
|
||||
|
||||
// Read the DNS query
|
||||
let mut query = vec![0u8; len];
|
||||
stream.read_exact(&mut query).await?;
|
||||
|
||||
if query.len() < 12 {
|
||||
debug!("Received malformed TCP query from {}", addr);
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let opcode = (query[2] & 0x78) >> 3;
|
||||
if opcode != 0 {
|
||||
if let Some(response) = build_not_implemented_response(&query, config.authoritative) {
|
||||
send_tcp_response(&mut stream, &response).await?;
|
||||
}
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let domain = match extract_domain(&query) {
|
||||
Some(d) => d,
|
||||
None => {
|
||||
info!("Failed to extract domain from TCP query");
|
||||
return Ok(());
|
||||
}
|
||||
};
|
||||
|
||||
debug!("TCP query for {} from {}", domain, addr);
|
||||
info!("Processing TCP query for domain: {}", domain);
|
||||
|
||||
let response = match generate_dns_response(&query, domain.clone(), &config).await {
|
||||
Ok(resp) => resp,
|
||||
Err(_) => {
|
||||
build_nxdomain_response(&query, config.authoritative)
|
||||
.ok_or(DnsError::Protocol("NXDOMAIN".into()))?
|
||||
}
|
||||
};
|
||||
|
||||
// Send the response (local/cache answer)
|
||||
send_tcp_response(&mut stream, &response).await?;
|
||||
Ok(())
|
||||
}
|
||||
+23
@@ -0,0 +1,23 @@
|
||||
//! NX9 DNS Server Library
|
||||
//!
|
||||
//! This library provides functionality for a DNS server implementation.
|
||||
//! It handles DNS queries over UDP and TCP, supports various record types,
|
||||
//! and can forward queries to upstream DNS servers.
|
||||
|
||||
#![allow(dead_code)]
|
||||
#[allow(unused_variables)]
|
||||
|
||||
// Define modules
|
||||
pub mod errors;
|
||||
pub mod config;
|
||||
pub mod cache;
|
||||
pub mod db;
|
||||
pub mod dns;
|
||||
pub mod handlers;
|
||||
pub mod utils;
|
||||
mod error;
|
||||
|
||||
// Re-export commonly used items
|
||||
pub use errors::DnsError;
|
||||
pub use config::ServerConfig;
|
||||
pub use cache::DnsCache;
|
||||
+69
@@ -0,0 +1,69 @@
|
||||
//! NX9 DNS Server
|
||||
//!
|
||||
//! A DNS server implementation that supports various record types and can forward
|
||||
//! queries to upstream DNS servers.
|
||||
//!
|
||||
//! Author: Sunil Purushottam Thakare
|
||||
#![allow(dead_code)]
|
||||
#[allow(unused_variables)]
|
||||
|
||||
use log::info;
|
||||
use tokio::{signal, task};
|
||||
|
||||
use nx9_dns_server::{
|
||||
cache::{CACHE, CACHE_CLEANUP_INTERVAL},
|
||||
config::ServerConfig,
|
||||
db::init_db,
|
||||
errors::DnsError,
|
||||
handlers::{run_tcp_server, run_udp_server},
|
||||
};
|
||||
|
||||
#[tokio::main]
|
||||
async fn main() -> Result<(), DnsError> {
|
||||
// Initialize the logger
|
||||
env_logger::Builder::from_env(env_logger::Env::default().default_filter_or("info"))
|
||||
.format_timestamp_micros()
|
||||
.init();
|
||||
|
||||
// Load configuration from environment variables
|
||||
let config = ServerConfig::from_env()?;
|
||||
|
||||
// Initialize cache with NS records from config
|
||||
let cache = CACHE.get_or_init(|| nx9_dns_server::cache::DnsCache::new(config.ns_records.clone()));
|
||||
|
||||
// Initialize the database
|
||||
init_db(&config.db_path, &config.default_domain, &config.default_ip)?;
|
||||
|
||||
// Set up cache cleanup task
|
||||
let cache_cleanup = task::spawn({
|
||||
let cache = cache.clone();
|
||||
async move {
|
||||
let mut interval = tokio::time::interval(CACHE_CLEANUP_INTERVAL);
|
||||
loop {
|
||||
interval.tick().await;
|
||||
cache.cleanup();
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
// Set up shutdown signal handler
|
||||
let shutdown_signal = async {
|
||||
signal::ctrl_c().await.expect("Failed to listen for shutdown signal");
|
||||
info!("Shutdown signal received");
|
||||
};
|
||||
|
||||
// Start UDP and TCP servers
|
||||
let udp_server = run_udp_server(config.clone());
|
||||
let tcp_server = run_tcp_server(config.clone());
|
||||
|
||||
// Wait for either a shutdown signal or server error
|
||||
tokio::select! {
|
||||
_ = shutdown_signal => {
|
||||
info!("Initiating graceful shutdown...");
|
||||
cache_cleanup.abort();
|
||||
Ok(())
|
||||
},
|
||||
res = udp_server => res,
|
||||
res = tcp_server => res,
|
||||
}
|
||||
}
|
||||
+427
@@ -0,0 +1,427 @@
|
||||
//! Utility functions for DNS operations.
|
||||
//!
|
||||
//! This module provides helper functions for parsing and encoding DNS data.
|
||||
#![allow(dead_code)]
|
||||
#[allow(unused_variables)]
|
||||
|
||||
use std::str;
|
||||
use chrono::{NaiveDateTime, TimeZone, Utc};
|
||||
|
||||
use crate::errors::DnsError;
|
||||
|
||||
/// Extract the domain name from a DNS query packet.
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `query` - The DNS query packet.
|
||||
///
|
||||
/// # Returns
|
||||
/// An `Option` containing the domain name if successfully extracted.
|
||||
pub fn extract_domain(query: &[u8]) -> Option<String> {
|
||||
if query.len() < 12 {
|
||||
return None; // DNS header is 12 bytes
|
||||
}
|
||||
|
||||
let mut pos = 12; // Start after header
|
||||
let mut domain = String::new();
|
||||
|
||||
// Extract QNAME (domain)
|
||||
loop {
|
||||
if pos >= query.len() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let len = query[pos] as usize;
|
||||
if len == 0 {
|
||||
break; // End of QNAME
|
||||
}
|
||||
pos += 1;
|
||||
|
||||
if pos + len > query.len() {
|
||||
return None; // Invalid length
|
||||
}
|
||||
|
||||
if !domain.is_empty() {
|
||||
domain.push('.');
|
||||
}
|
||||
|
||||
let label = match str::from_utf8(&query[pos..pos + len]) {
|
||||
Ok(l) => l,
|
||||
Err(_) => return None, // Invalid UTF-8
|
||||
};
|
||||
domain.push_str(label);
|
||||
pos += len;
|
||||
}
|
||||
|
||||
// Skip QTYPE and QCLASS (4 bytes)
|
||||
pos += 4;
|
||||
|
||||
// Verify we have enough data for at least QTYPE/QCLASS
|
||||
if pos > query.len() {
|
||||
return None;
|
||||
}
|
||||
|
||||
Some(domain)
|
||||
}
|
||||
|
||||
/// Extract the query type from a DNS query packet.
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `query` - The DNS query packet.
|
||||
///
|
||||
/// # Returns
|
||||
/// An `Option` containing the query type as a u16 if successfully extracted.
|
||||
pub fn extract_query_type(query: &[u8]) -> Option<u16> {
|
||||
if query.len() < 12 {
|
||||
return None; // DNS header is 12 bytes
|
||||
}
|
||||
|
||||
let mut pos = 12; // Start after header
|
||||
|
||||
// Skip QNAME
|
||||
loop {
|
||||
if pos >= query.len() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let len = query[pos] as usize;
|
||||
if len == 0 {
|
||||
pos += 1;
|
||||
break; // End of QNAME
|
||||
}
|
||||
|
||||
pos += len + 1;
|
||||
}
|
||||
|
||||
// Get QTYPE (2 bytes after QNAME)
|
||||
if pos + 1 < query.len() {
|
||||
Some(((query[pos] as u16) << 8) | query[pos + 1] as u16)
|
||||
} else {
|
||||
None
|
||||
}
|
||||
}
|
||||
|
||||
/// Encode a domain name in DNS wire format.
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `name` - The domain name to encode.
|
||||
///
|
||||
/// # Returns
|
||||
/// A vector of bytes containing the encoded domain name.
|
||||
pub fn encode_dns_name(name: &str) -> Vec<u8> {
|
||||
let mut out = Vec::new();
|
||||
for part in name.trim_end_matches('.').split('.') {
|
||||
if part.len() > 63 {
|
||||
continue; // Skip invalid labels
|
||||
}
|
||||
out.push(part.len() as u8);
|
||||
out.extend_from_slice(part.as_bytes());
|
||||
}
|
||||
out.push(0); // Null terminator
|
||||
out
|
||||
}
|
||||
|
||||
/// Parse a signature time in YYYYMMDDHHMMSS format to seconds since epoch.
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `s` - The signature time string.
|
||||
///
|
||||
/// # Returns
|
||||
/// A `Result` containing the parsed time as a u32 or an error.
|
||||
pub fn parse_sig_time(s: &str) -> Result<u32, DnsError> {
|
||||
let dt = NaiveDateTime::parse_from_str(s, "%Y%m%d%H%M%S")
|
||||
.map_err(|e| DnsError::Config(format!("Invalid sigtime: {e}")))?;
|
||||
Ok(Utc.from_utc_datetime(&dt).timestamp() as u32)
|
||||
}
|
||||
|
||||
/// Check if a DNS query packet has an OPT record (EDNS).
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `query` - The DNS query packet.
|
||||
///
|
||||
/// # Returns
|
||||
/// A boolean indicating whether the query has an OPT record.
|
||||
pub fn has_opt_record(query: &[u8]) -> bool {
|
||||
if query.len() < 12 {
|
||||
return false;
|
||||
}
|
||||
|
||||
// Get ARCOUNT (number of additional records)
|
||||
let arcount = ((query[10] as u16) << 8) | query[11] as u16;
|
||||
if arcount == 0 {
|
||||
return false;
|
||||
}
|
||||
|
||||
// Skip header
|
||||
let mut pos = 12;
|
||||
|
||||
// Skip question section
|
||||
// First skip QNAME
|
||||
loop {
|
||||
if pos >= query.len() {
|
||||
return false;
|
||||
}
|
||||
let len = query[pos] as usize;
|
||||
if len == 0 {
|
||||
pos += 1;
|
||||
break;
|
||||
}
|
||||
pos += len + 1;
|
||||
}
|
||||
|
||||
// Skip QTYPE and QCLASS
|
||||
pos += 4;
|
||||
|
||||
// Skip answer and authority sections
|
||||
let ancount = ((query[6] as u16) << 8) | query[7] as u16;
|
||||
let nscount = ((query[8] as u16) << 8) | query[9] as u16;
|
||||
|
||||
for _ in 0..(ancount + nscount) {
|
||||
// Skip name
|
||||
if pos >= query.len() {
|
||||
return false;
|
||||
}
|
||||
|
||||
// Handle compression pointers
|
||||
if (query[pos] & 0xC0) == 0xC0 {
|
||||
pos += 2; // Skip compression pointer
|
||||
} else {
|
||||
// Skip labels
|
||||
loop {
|
||||
if pos >= query.len() {
|
||||
return false;
|
||||
}
|
||||
let len = query[pos] as usize;
|
||||
if len == 0 {
|
||||
pos += 1;
|
||||
break;
|
||||
}
|
||||
pos += len + 1;
|
||||
}
|
||||
}
|
||||
|
||||
// Skip TYPE, CLASS, TTL, RDLENGTH, RDATA
|
||||
if pos + 10 > query.len() {
|
||||
return false;
|
||||
}
|
||||
let rdlength = ((query[pos + 8] as usize) << 8) | query[pos + 9] as usize;
|
||||
pos += 10 + rdlength;
|
||||
}
|
||||
|
||||
// Check additional records for OPT
|
||||
for _ in 0..arcount {
|
||||
if pos >= query.len() {
|
||||
return false;
|
||||
}
|
||||
|
||||
// OPT record has empty (root) name
|
||||
if query[pos] == 0 {
|
||||
// Check if TYPE is OPT (41)
|
||||
if pos + 2 < query.len() && query[pos + 1] == 0 && query[pos + 2] == 41 {
|
||||
return true;
|
||||
}
|
||||
}
|
||||
|
||||
// Skip this record
|
||||
if pos + 10 >= query.len() { break; }
|
||||
let rdlength = ((query[pos + 8] as usize) << 8) | query[pos + 9] as usize;
|
||||
pos += 10 + rdlength;
|
||||
}
|
||||
|
||||
false // No OPT record found
|
||||
}
|
||||
|
||||
/// Extract the EDNS payload size from a DNS query packet.
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `query` - The DNS query packet.
|
||||
///
|
||||
/// # Returns
|
||||
/// An `Option` containing the EDNS payload size if found.
|
||||
pub fn extract_edns_payload_size(query: &[u8]) -> Option<u16> {
|
||||
if query.len() < 12 {
|
||||
return None;
|
||||
}
|
||||
|
||||
// Get ARCOUNT (number of additional records)
|
||||
let arcount = ((query[10] as u16) << 8) | query[11] as u16;
|
||||
if arcount == 0 {
|
||||
return None;
|
||||
}
|
||||
|
||||
// Skip header
|
||||
let mut pos = 12;
|
||||
|
||||
// Skip question section
|
||||
// First skip QNAME
|
||||
loop {
|
||||
if pos >= query.len() {
|
||||
return None;
|
||||
}
|
||||
let len = query[pos] as usize;
|
||||
if len == 0 {
|
||||
pos += 1;
|
||||
break;
|
||||
}
|
||||
pos += len + 1;
|
||||
}
|
||||
|
||||
// Skip QTYPE and QCLASS
|
||||
pos += 4;
|
||||
|
||||
// Skip answer and authority sections
|
||||
let ancount = ((query[6] as u16) << 8) | query[7] as u16;
|
||||
let nscount = ((query[8] as u16) << 8) | query[9] as u16;
|
||||
|
||||
for _ in 0..(ancount + nscount) {
|
||||
// Skip name
|
||||
if pos >= query.len() {
|
||||
return None;
|
||||
}
|
||||
|
||||
// Handle compression pointers
|
||||
if (query[pos] & 0xC0) == 0xC0 {
|
||||
pos += 2; // Skip compression pointer
|
||||
} else {
|
||||
// Skip labels
|
||||
loop {
|
||||
if pos >= query.len() {
|
||||
return None;
|
||||
}
|
||||
let len = query[pos] as usize;
|
||||
if len == 0 {
|
||||
pos += 1;
|
||||
break;
|
||||
}
|
||||
pos += len + 1;
|
||||
}
|
||||
}
|
||||
|
||||
// Skip TYPE, CLASS, TTL, RDLENGTH, RDATA
|
||||
if pos + 10 > query.len() {
|
||||
return None;
|
||||
}
|
||||
let rdlength = ((query[pos + 8] as usize) << 8) | query[pos + 9] as usize;
|
||||
pos += 10 + rdlength;
|
||||
}
|
||||
|
||||
// Check additional records for OPT
|
||||
for _ in 0..arcount {
|
||||
if pos >= query.len() {
|
||||
return None;
|
||||
}
|
||||
|
||||
// OPT record has empty (root) name
|
||||
if query[pos] == 0 {
|
||||
// Check if TYPE is OPT (41)
|
||||
if pos + 5 < query.len() && query[pos + 1] == 0 && query[pos + 2] == 41 {
|
||||
// Extract UDP payload size (CLASS field in OPT record)
|
||||
return Some(((query[pos + 3] as u16) << 8) | query[pos + 4] as u16);
|
||||
}
|
||||
}
|
||||
|
||||
// Skip this record
|
||||
if pos + 10 >= query.len() { break; }
|
||||
let rdlength = ((query[pos + 8] as usize) << 8) | query[pos + 9] as usize;
|
||||
pos += 10 + rdlength;
|
||||
}
|
||||
|
||||
None // No OPT record found
|
||||
}
|
||||
|
||||
/// Extract the DO (DNSSEC OK) bit from a DNS query packet.
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `query` - The DNS query packet.
|
||||
///
|
||||
/// # Returns
|
||||
/// A boolean indicating whether the DO bit is set.
|
||||
pub fn extract_do_bit(query: &[u8]) -> bool {
|
||||
if query.len() < 12 {
|
||||
return false;
|
||||
}
|
||||
|
||||
// Get ARCOUNT (number of additional records)
|
||||
let arcount = ((query[10] as u16) << 8) | query[11] as u16;
|
||||
if arcount == 0 {
|
||||
return false;
|
||||
}
|
||||
|
||||
// Skip header
|
||||
let mut pos = 12;
|
||||
|
||||
// Skip question section
|
||||
// First skip QNAME
|
||||
loop {
|
||||
if pos >= query.len() {
|
||||
return false;
|
||||
}
|
||||
let len = query[pos] as usize;
|
||||
if len == 0 {
|
||||
pos += 1;
|
||||
break;
|
||||
}
|
||||
pos += len + 1;
|
||||
}
|
||||
|
||||
// Skip QTYPE and QCLASS
|
||||
pos += 4;
|
||||
|
||||
// Skip answer and authority sections
|
||||
let ancount = ((query[6] as u16) << 8) | query[7] as u16;
|
||||
let nscount = ((query[8] as u16) << 8) | query[9] as u16;
|
||||
|
||||
for _ in 0..(ancount + nscount) {
|
||||
// Skip name
|
||||
if pos >= query.len() {
|
||||
return false;
|
||||
}
|
||||
|
||||
// Handle compression pointers
|
||||
if (query[pos] & 0xC0) == 0xC0 {
|
||||
pos += 2; // Skip compression pointer
|
||||
} else {
|
||||
// Skip labels
|
||||
loop {
|
||||
if pos >= query.len() {
|
||||
return false;
|
||||
}
|
||||
let len = query[pos] as usize;
|
||||
if len == 0 {
|
||||
pos += 1;
|
||||
break;
|
||||
}
|
||||
pos += len + 1;
|
||||
}
|
||||
}
|
||||
|
||||
// Skip TYPE, CLASS, TTL, RDLENGTH, RDATA
|
||||
if pos + 10 > query.len() {
|
||||
return false;
|
||||
}
|
||||
let rdlength = ((query[pos + 8] as usize) << 8) | query[pos + 9] as usize;
|
||||
pos += 10 + rdlength;
|
||||
}
|
||||
|
||||
// Check additional records for OPT
|
||||
for _ in 0..arcount {
|
||||
if pos >= query.len() {
|
||||
return false;
|
||||
}
|
||||
|
||||
// OPT record has empty (root) name
|
||||
if query[pos] == 0 {
|
||||
// Check if TYPE is OPT (41)
|
||||
if pos + 7 < query.len() && query[pos + 1] == 0 && query[pos + 2] == 41 {
|
||||
// Check DO bit (bit 15 of TTL field, which is used for flags in OPT)
|
||||
return (query[pos + 6] & 0x80) != 0;
|
||||
}
|
||||
}
|
||||
|
||||
// Skip this record
|
||||
if pos + 10 >= query.len() { break; }
|
||||
let rdlength = ((query[pos + 8] as usize) << 8) | query[pos + 9] as usize;
|
||||
pos += 10 + rdlength;
|
||||
}
|
||||
|
||||
false // No OPT record found or no DO bit set
|
||||
}
|
||||
Reference in new issue
Block a user