use bnb::{BitEnum, BitReader, Sink, Source, bin, bitfield, u3, u4};
use std::net::UdpSocket;
use std::thread;
use std::time::Duration;
use tracing::info;
#[derive(BitEnum, Clone, Copy, Debug, PartialEq, Eq)]
#[bit_enum(u4)]
enum OpCode {
Query,
IQuery,
Status,
#[catch_all]
Other(u4),
}
#[derive(BitEnum, Clone, Copy, Debug, PartialEq, Eq)]
#[bit_enum(u4)]
enum RCode {
NoError,
FormErr,
ServFail,
NxDomain,
#[catch_all]
Other(u4),
}
#[bitfield(u16, bits = msb, bytes = big)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
struct Flags {
qr: bool, opcode: OpCode, aa: bool, tc: bool, rd: bool, ra: bool, z: u3, rcode: RCode, }
fn read_name<S: Source>(r: &mut S) -> Result<Vec<String>, bnb::BitError> {
let mut labels = Vec::new();
let mut resume_at: Option<usize> = None; let mut hops = 0;
loop {
let n: u8 = r.read()?;
if n == 0 {
break; }
if n & 0xC0 == 0xC0 {
let lo: u8 = r.read()?;
let offset = (((n & 0x3F) as usize) << 8) | lo as usize;
resume_at.get_or_insert_with(|| r.bit_pos());
hops += 1;
if hops > 128 {
break; }
r.seek_to_bit(offset * 8)?; continue;
}
let mut bytes = Vec::with_capacity(n as usize);
for _ in 0..n {
bytes.push(r.read::<u8>()?);
}
labels.push(String::from_utf8_lossy(&bytes).into_owned());
}
if let Some(pos) = resume_at {
r.seek_to_bit(pos)?; }
Ok(labels)
}
fn write_name<K: Sink>(labels: &[String], w: &mut K) -> Result<(), bnb::BitError> {
for label in labels {
if label.len() > 63 {
return Err(bnb::BitError::convert(
format!(
"DNS label '{label}' is {} bytes; the maximum is 63",
label.len()
),
w.bit_pos(),
));
}
w.write(label.len() as u8)?;
w.write_bytes(label.as_bytes())?;
}
w.write(0u8) }
#[bin(codec(parse = read_name, write = write_name))]
#[derive(Debug, Clone, PartialEq, Eq)]
struct DnsName(Vec<String>);
fn dotted(labels: &[String]) -> String {
labels.join(".")
}
#[bin(big)]
#[derive(Debug, Clone, PartialEq, Eq)]
struct Question {
#[brw(variable)]
name: DnsName,
qtype: u16,
qclass: u16,
}
#[bin(big)]
#[derive(Debug, Clone, PartialEq, Eq)]
struct Record {
name: DnsName,
rtype: u16,
rclass: u16,
ttl: u32,
#[br(temp)]
#[bw(calc = self.rdata.len() as u16)]
rdlength: u16,
#[br(count = rdlength)]
rdata: Vec<u8>,
}
#[bin(big)]
#[derive(Debug, Clone, PartialEq, Eq)]
struct Message {
id: u16,
flags: Flags,
#[br(temp)]
#[bw(calc = self.questions.len() as u16)]
qdcount: u16,
#[br(temp)]
#[bw(calc = self.answers.len() as u16)]
ancount: u16,
#[br(temp)]
#[bw(calc = self.authority.len() as u16)]
nscount: u16,
#[br(temp)]
#[bw(calc = self.additional.len() as u16)]
arcount: u16,
#[br(count = qdcount)]
questions: Vec<Question>,
#[builder(default)]
#[br(count = ancount)]
answers: Vec<Record>,
#[builder(default)]
#[br(count = nscount)]
authority: Vec<Record>,
#[builder(default)]
#[br(count = arcount)]
additional: Vec<Record>,
}
fn hex(bytes: &[u8]) -> String {
bytes
.iter()
.map(|b| format!("{b:02x}"))
.collect::<Vec<_>>()
.join(" ")
}
fn serve_one(sock: &UdpSocket) -> std::io::Result<()> {
let mut buf = [0u8; 512]; let (n, client) = sock.recv_from(&mut buf)?;
let query = Message::decode_exact(&buf[..n]).expect("decode query");
let q = query.questions[0].clone();
let name = q.name.clone();
let response = Message::builder()
.id(query.id) .flags(Flags::new().with_qr(true).with_rd(true).with_ra(true)) .questions(vec![q])
.answers(vec![Record {
name,
rtype: 1, rclass: 1, ttl: 60,
rdata: vec![93, 184, 216, 34],
}])
.build()
.expect("build response");
sock.send_to(&response.to_bytes().expect("encode response"), client)?;
Ok(())
}
fn main() -> Result<(), Box<dyn std::error::Error>> {
tracing_subscriber::fmt()
.with_max_level(tracing::Level::INFO)
.with_target(false)
.without_time()
.init();
let query = Message::builder()
.id(0x1234)
.flags(Flags::new().with_rd(true)) .questions(vec![Question {
name: DnsName(vec!["example".into(), "com".into()]),
qtype: 1, qclass: 1, }])
.build()?;
let q_bytes = query.to_bytes()?;
info!(
question = %dotted(&query.questions[0].name.0),
bytes = %hex(&q_bytes),
"built query",
);
assert_eq!(Message::decode_exact(&q_bytes)?, query);
let oversized = Question {
name: DnsName(vec!["x".repeat(64), "com".into()]),
qtype: 1,
qclass: 1,
};
let err = oversized.to_bytes().unwrap_err();
info!(%err, "a 64-byte label is refused at encode (checked write)");
let wire: &[u8] = &[
0x12, 0x34, 0x81, 0x80, 0x00, 0x01, 0x00, 0x01, 0x00, 0x00, 0x00, 0x00, 0x07, b'e', b'x', b'a', b'm', b'p', b'l', b'e', 0x03, b'c', b'o', b'm',
0x00, 0x00, 0x01, 0x00, 0x01, 0xc0, 0x0c, 0x00, 0x01, 0x00, 0x01, 0x00, 0x00, 0x00, 0x3c, 0x00, 0x04, 0x5d, 0xb8, 0xd8, 0x22, ];
info!(
len = wire.len(),
"decoding a response with a compressed answer name"
);
let resp = Message::decode_exact(wire)?;
info!("the full decoded structure:\n{resp:#?}"); let answer = &resp.answers[0];
let ip = &answer.rdata;
info!(
qr = resp.flags.qr(),
opcode = ?resp.flags.opcode(),
ra = resp.flags.ra(),
rcode = ?resp.flags.rcode(),
question = %dotted(&resp.questions[0].name.0),
answer_name = %dotted(&answer.name.0),
ttl = answer.ttl,
address = %format!("{}.{}.{}.{}", ip[0], ip[1], ip[2], ip[3]),
"decoded response (compression pointer followed)",
);
assert_eq!(dotted(&answer.name.0), "example.com"); assert_eq!(answer.rdata, vec![93, 184, 216, 34]);
let reencoded = resp.to_bytes()?;
assert_eq!(Message::decode_exact(&reencoded)?, resp);
info!(
on_wire = wire.len(),
re_encoded = reencoded.len(),
"round-trips by value; re-encode is uncompressed (larger), as expected",
);
let mut r = BitReader::new(wire);
r.seek_to_bit(29 * 8)?;
assert_eq!(read_name(&mut r)?, vec!["example", "com"]);
let server = UdpSocket::bind("127.0.0.1:0")?; let server_addr = server.local_addr()?;
let server_thread = thread::spawn(move || serve_one(&server));
let client = UdpSocket::bind("127.0.0.1:0")?;
client.set_read_timeout(Some(Duration::from_secs(2)))?; client.send_to(&q_bytes, server_addr)?;
info!(to = %server_addr, bytes = q_bytes.len(), "client → query over UDP loopback");
let mut inbox = [0u8; 512];
let (n, from) = client.recv_from(&mut inbox)?;
let reply = Message::decode_exact(&inbox[..n])?;
let a = &reply.answers[0];
info!(
from = %from,
id = %format!("0x{:04x}", reply.id),
name = %dotted(&a.name.0),
address = %format!("{}.{}.{}.{}", a.rdata[0], a.rdata[1], a.rdata[2], a.rdata[3]),
"client ← response decoded",
);
assert_eq!(a.name, query.questions[0].name); server_thread.join().expect("server thread panicked")?;
info!("all checks passed");
Ok(())
}