use std::fmt;
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr};
use bincode::{Encode, Decode};
#[derive(Debug, Clone, PartialEq, Eq, Encode, Decode)]
pub struct Request {
pub id: u16,
pub flags: Flags,
pub query: Query,
pub client_address: Option<ClientAddress>,
}
#[derive(Debug, Clone, PartialEq, Encode, Decode)]
pub struct Response {
pub id: u16,
pub flags: Flags,
pub queries: Vec<Query>,
pub answers: Vec<Record>,
pub authorities: Vec<Record>,
pub additionals: Vec<Record>,
}
#[derive(Debug, Clone, PartialEq, Eq, Hash, Encode, Decode)]
pub struct Query {
pub name: String,
pub qtype: RecordType,
pub qclass: QClass,
}
#[derive(Debug, Clone, PartialEq, Encode, Decode)]
pub struct Record {
pub name: String,
pub rtype: RecordType,
pub class: QClass,
pub ttl: u32,
pub data: RecordData,
}
#[derive(Debug, Clone, PartialEq, Encode, Decode)]
pub enum RecordData {
A(Ipv4Addr),
AAAA(Ipv6Addr),
CNAME(String),
MX {
priority: u16,
exchange: String
},
NS(String),
PTR(String),
SOA {
mname: String,
rname: String,
serial: u32,
refresh: u32,
retry: u32,
expire: u32,
minimum: u32,
},
TXT(Vec<String>),
SRV {
priority: u16,
weight: u16,
port: u16,
target: String,
},
Unknown(Vec<u8>),
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Encode, Decode)]
#[repr(u16)]
pub enum RecordType {
A = 1,
NS = 2,
CNAME = 5,
SOA = 6,
PTR = 12,
MX = 15,
TXT = 16,
AAAA = 28,
SRV = 33,
Unknown(u16),
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Encode, Decode)]
#[repr(u16)]
pub enum QClass {
IN = 1,
CH = 3,
HS = 4,
ANY = 255,
Unknown(u16),
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Encode, Decode)]
pub struct Flags {
pub qr: bool,
pub opcode: u8,
pub aa: bool,
pub tc: bool,
pub rd: bool,
pub ra: bool,
pub z: u8,
pub rcode: u8,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[repr(u8)]
pub enum ResponseCode {
NoError = 0,
FormatError = 1,
ServerFailure = 2,
NxDomain = 3,
NotImplemented = 4,
Refused = 5,
Unknown(u8),
}
#[derive(Debug, Clone, PartialEq, Eq, Encode, Decode)]
pub struct ClientAddress {
pub address: IpAddr,
pub source_prefix_length: u8,
pub scope_prefix_length: u8,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct EdnsOption {
pub code: u16,
pub data: Vec<u8>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct EdnsRecord {
pub udp_payload_size: u16,
pub extended_rcode: u8,
pub version: u8,
pub dnssec_ok: bool,
pub options: Vec<EdnsOption>,
}
pub mod edns_option_codes {
pub const CLIENT_ADDRESS: u16 = 8;
pub const COOKIE: u16 = 10;
pub const KEEPALIVE: u16 = 11;
pub const PADDING: u16 = 12;
}
impl ClientAddress {
pub fn new(address: IpAddr, source_prefix_length: u8) -> Self {
Self {
address,
source_prefix_length,
scope_prefix_length: 0, }
}
pub fn from_ipv4(address: Ipv4Addr, prefix_length: u8) -> Self {
Self::new(IpAddr::V4(address), prefix_length)
}
pub fn from_ipv6(address: Ipv6Addr, prefix_length: u8) -> Self {
Self::new(IpAddr::V6(address), prefix_length)
}
pub fn family(&self) -> u16 {
match self.address {
IpAddr::V4(_) => 1,
IpAddr::V6(_) => 2,
}
}
pub fn encode(&self) -> Vec<u8> {
let mut data = Vec::new();
data.extend_from_slice(&self.family().to_be_bytes());
data.push(self.source_prefix_length);
data.push(self.scope_prefix_length);
match self.address {
IpAddr::V4(addr) => {
let bytes = addr.octets();
let byte_count = (self.source_prefix_length + 7) / 8;
data.extend_from_slice(&bytes[..byte_count as usize]);
}
IpAddr::V6(addr) => {
let bytes = addr.octets();
let byte_count = (self.source_prefix_length + 7) / 8;
data.extend_from_slice(&bytes[..byte_count as usize]);
}
}
data
}
pub fn decode(data: &[u8]) -> Result<Self, &'static str> {
if data.len() < 4 {
return Err("Client subnet data too short");
}
let family = u16::from_be_bytes([data[0], data[1]]);
let source_prefix_length = data[2];
let scope_prefix_length = data[3];
let address = match family {
1 => {
if data.len() < 4 {
return Err("IPv4 client subnet data too short");
}
let mut addr_bytes = [0u8; 4];
let available_bytes = data.len() - 4;
let copy_bytes = std::cmp::min(4, available_bytes);
addr_bytes[..copy_bytes].copy_from_slice(&data[4..4 + copy_bytes]);
IpAddr::V4(Ipv4Addr::from(addr_bytes))
}
2 => {
if data.len() < 4 {
return Err("IPv6 client subnet data too short");
}
let mut addr_bytes = [0u8; 16];
let available_bytes = data.len() - 4;
let copy_bytes = std::cmp::min(16, available_bytes);
addr_bytes[..copy_bytes].copy_from_slice(&data[4..4 + copy_bytes]);
IpAddr::V6(Ipv6Addr::from(addr_bytes))
}
_ => return Err("Unsupported address family"),
};
Ok(Self {
address,
source_prefix_length,
scope_prefix_length,
})
}
}
impl From<u16> for RecordType {
fn from(value: u16) -> Self {
match value {
1 => RecordType::A,
2 => RecordType::NS,
5 => RecordType::CNAME,
6 => RecordType::SOA,
12 => RecordType::PTR,
15 => RecordType::MX,
16 => RecordType::TXT,
28 => RecordType::AAAA,
33 => RecordType::SRV,
_ => RecordType::Unknown(value),
}
}
}
impl From<RecordType> for u16 {
fn from(rtype: RecordType) -> Self {
match rtype {
RecordType::A => 1,
RecordType::NS => 2,
RecordType::CNAME => 5,
RecordType::SOA => 6,
RecordType::PTR => 12,
RecordType::MX => 15,
RecordType::TXT => 16,
RecordType::AAAA => 28,
RecordType::SRV => 33,
RecordType::Unknown(value) => value,
}
}
}
impl From<u16> for QClass {
fn from(value: u16) -> Self {
match value {
1 => QClass::IN,
3 => QClass::CH,
4 => QClass::HS,
255 => QClass::ANY,
_ => QClass::Unknown(value),
}
}
}
impl From<QClass> for u16 {
fn from(qclass: QClass) -> Self {
match qclass {
QClass::IN => 1,
QClass::CH => 3,
QClass::HS => 4,
QClass::ANY => 255,
QClass::Unknown(value) => value,
}
}
}
impl From<u8> for ResponseCode {
fn from(value: u8) -> Self {
match value {
0 => ResponseCode::NoError,
1 => ResponseCode::FormatError,
2 => ResponseCode::ServerFailure,
3 => ResponseCode::NxDomain,
4 => ResponseCode::NotImplemented,
5 => ResponseCode::Refused,
_ => ResponseCode::Unknown(value),
}
}
}
impl From<ResponseCode> for u8 {
fn from(rcode: ResponseCode) -> Self {
match rcode {
ResponseCode::NoError => 0,
ResponseCode::FormatError => 1,
ResponseCode::ServerFailure => 2,
ResponseCode::NxDomain => 3,
ResponseCode::NotImplemented => 4,
ResponseCode::Refused => 5,
ResponseCode::Unknown(value) => value,
}
}
}
impl Default for Flags {
fn default() -> Self {
Self {
qr: false,
opcode: 0,
aa: false,
tc: false,
rd: true,
ra: false,
z: 0,
rcode: 0,
}
}
}
impl fmt::Display for RecordType {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
RecordType::A => write!(f, "A"),
RecordType::NS => write!(f, "NS"),
RecordType::CNAME => write!(f, "CNAME"),
RecordType::SOA => write!(f, "SOA"),
RecordType::PTR => write!(f, "PTR"),
RecordType::MX => write!(f, "MX"),
RecordType::TXT => write!(f, "TXT"),
RecordType::AAAA => write!(f, "AAAA"),
RecordType::SRV => write!(f, "SRV"),
RecordType::Unknown(value) => write!(f, "TYPE{}", value),
}
}
}