use anyhow::{bail, Context, Result};
use std::convert::TryFrom;
use std::net::IpAddr;
use std::str;
use bytes::{BufMut, BytesMut};
use packed_struct::prelude::*;
use crate::codec::{domain_name, rdata};
use crate::specs::enums_generated::{self, ResourceClass, ResourceType};
use crate::specs::message::*;
#[derive(Clone, Debug)]
pub struct RequestInfo {
pub name: String,
pub resource_type: ResourceType,
pub dnssec_ok: bool,
pub received_request_id: u16,
pub requested_udp_size: u16,
}
pub const OPT_RESOURCE_TYPE: u16 = 41;
#[derive(PackedStruct)]
#[packed_struct(endian = "msb", bit_numbering = "msb0")]
pub struct HeaderBits {
#[packed_field(bits = "0:15")]
pub id: u16,
#[packed_field(bits = "16:16")]
pub is_response: bool,
#[packed_field(bits = "17:20")]
pub op_code: Integer<u8, packed_bits::Bits<4>>,
#[packed_field(bits = "21:21")]
pub authoritative: bool,
#[packed_field(bits = "22:22")]
pub truncated: bool,
#[packed_field(bits = "23:23")]
pub recursion_desired: bool,
#[packed_field(bits = "24:24")]
pub recursion_available: bool,
#[packed_field(bits = "25:25")]
pub reserved_9: bool,
#[packed_field(bits = "26:26")]
pub authentic_data: bool,
#[packed_field(bits = "27:27")]
pub checking_disabled: bool,
#[packed_field(bits = "28:31")]
pub response_code: Integer<u8, packed_bits::Bits<4>>,
#[packed_field(bits = "32:47")]
pub question_count: u16,
#[packed_field(bits = "48:63")]
pub answer_count: u16,
#[packed_field(bits = "64:79")]
pub authority_count: u16,
#[packed_field(bits = "80:95")]
pub additional_count: u16,
}
#[derive(Clone, Debug, PartialEq)]
pub struct RecordCounts {
pub question: u16,
pub answer: u16,
pub authority: u16,
pub additional: u16,
}
impl RecordCounts {
pub fn new() -> RecordCounts {
RecordCounts {
question: 0,
answer: 0,
authority: 0,
additional: 0,
}
}
}
pub fn write_header_id(message: &Message, id_override: u16, buf: &mut BytesMut) -> Result<()> {
let opt_len = match &message.opt {
Some(_opt) => 1,
None => 0,
};
write_header_bits(
HeaderBits {
id: id_override,
is_response: message.header.is_response,
op_code: match message.header.op_code {
IntEnum::Enum(e) => Integer::from(e as u8),
IntEnum::Unknown(i) => Integer::from(i),
},
authoritative: message.header.authoritative,
truncated: message.header.truncated,
recursion_desired: message.header.recursion_desired,
recursion_available: message.header.recursion_available,
reserved_9: message.header.reserved_9,
authentic_data: message.header.authentic_data,
checking_disabled: message.header.checking_disabled,
response_code: match message.header.response_code {
IntEnum::Enum(e) => Integer::from(e as u8),
IntEnum::Unknown(i) => Integer::from(i),
},
question_count: u16::try_from(message.question.len())
.with_context(|| "question length doesn't fit")?,
answer_count: u16::try_from(message.answer.len())
.with_context(|| "answer length doesn't fit")?,
authority_count: u16::try_from(message.authority.len())
.with_context(|| "authority length doesn't fit")?,
additional_count: u16::try_from(message.additional.len() + opt_len)
.with_context(|| "additional length doesn't fit")?,
},
buf,
)
}
pub fn write_header_bits(bits: HeaderBits, buf: &mut BytesMut) -> Result<()> {
let bits_packed = bits.pack()?;
buf.reserve(bits_packed.len());
buf.put_slice(&bits_packed);
Ok(())
}
pub fn read_header(buf: &[u8], offset: &mut usize) -> Result<Option<(Header, RecordCounts, bool)>> {
let headerbits_size = 12; if buf.len() < *offset + headerbits_size {
return Ok(None);
}
let bits = HeaderBits::unpack_from_slice(&buf[*offset..*offset + headerbits_size])
.with_context(|| "couldn't unpack header bits")?;
let header = Header {
id: bits.id,
is_response: bits.is_response,
op_code: match enums_generated::opcode_int(*bits.op_code as usize) {
Some(e) => IntEnum::Enum(e),
None => IntEnum::Unknown(*bits.op_code),
},
authoritative: bits.authoritative,
truncated: bits.truncated,
recursion_desired: bits.recursion_desired,
recursion_available: bits.recursion_available,
reserved_9: bits.reserved_9,
authentic_data: bits.authentic_data,
checking_disabled: bits.checking_disabled,
response_code: match enums_generated::responsecode_int(*bits.response_code as usize) {
Some(e) => IntEnum::Enum(e),
None => IntEnum::Unknown(*bits.response_code),
},
};
let record_counts = RecordCounts {
question: bits.question_count,
answer: bits.answer_count,
authority: bits.authority_count,
additional: bits.additional_count,
};
*offset += headerbits_size;
Ok(Some((header, record_counts, bits.truncated)))
}
#[derive(PackedStruct)]
#[packed_struct(endian = "msb", bit_numbering = "msb0")]
pub struct HeaderIDBits {
#[packed_field(bits = "0:15")]
id: u16,
}
pub fn update_message_id(id: u16, buf: &mut BytesMut, message_offset: usize) -> Result<()> {
let bits = HeaderIDBits { id }.pack()?;
if buf.len() < message_offset + bits.len() {
bail!(
"Buffer is too small to update request ID of length {} at message_offset={}: buf=0x{:X}",
buf.len(),
message_offset,
buf
);
}
for i in 0..bits.len() {
buf[message_offset + i] = bits[i];
}
Ok(())
}
pub fn get_min_ttl_secs(message: &Message) -> Option<u32> {
let mut min_ttl: Option<u32> = None;
for a in &message.answer {
match min_ttl {
Some(min_val) => {
if a.ttl < min_val {
min_ttl = Some(a.ttl);
}
}
None => {
min_ttl = Some(a.ttl);
}
}
}
for a in &message.authority {
match min_ttl {
Some(min_val) => {
if a.ttl < min_val {
min_ttl = Some(a.ttl);
}
}
None => {
min_ttl = Some(a.ttl);
}
}
}
for a in &message.additional {
match min_ttl {
Some(min_val) => {
if a.ttl < min_val {
min_ttl = Some(a.ttl);
}
}
None => {
min_ttl = Some(a.ttl);
}
}
}
return min_ttl;
}
pub fn update_cached_response(
response: &mut Message,
request_info: &RequestInfo,
cache_remaining_ttl_secs: u32,
) -> Result<()> {
let min_ttl_secs = get_min_ttl_secs(&response).with_context(|| {
format!(
"Missing resources in cached {:?} response for {}",
request_info.resource_type, request_info.name
)
})?;
if min_ttl_secs < cache_remaining_ttl_secs {
bail!(
"Redis had invalid TTL {}s with {:?} result for {}: {}",
cache_remaining_ttl_secs,
request_info.resource_type,
request_info.name,
response
);
}
let ttl_subtract_secs = min_ttl_secs - cache_remaining_ttl_secs;
response.header.id = request_info.received_request_id;
for r in &mut response.answer {
r.ttl -= ttl_subtract_secs;
}
for r in &mut response.authority {
r.ttl -= ttl_subtract_secs;
}
if let Some(opt) = &mut response.opt {
opt.udp_size = request_info.requested_udp_size;
}
for r in &mut response.additional {
r.ttl -= ttl_subtract_secs;
}
Ok(())
}
#[derive(PackedStruct)]
#[packed_struct(endian = "msb", bit_numbering = "msb0")]
pub struct QuestionBits {
#[packed_field(bits = "0:15")]
resource_type: u16,
#[packed_field(bits = "16:31")]
resource_class: u16,
}
pub fn write_question(
question: &Question,
buf: &mut BytesMut,
ptr_offsets: &mut domain_name::LabelOffsets,
) -> Result<()> {
domain_name::write(&question.name, buf, ptr_offsets, "question.name")?;
let bits = QuestionBits {
resource_type: match question.resource_type {
IntEnum::Enum(e) => e as u16,
IntEnum::Unknown(i) => i,
},
resource_class: match question.resource_class {
IntEnum::Enum(e) => e as u16,
IntEnum::Unknown(i) => i,
},
}
.pack()?;
buf.reserve(bits.len());
buf.put_slice(&bits);
Ok(())
}
pub fn read_question(buf: &[u8], offset: &mut usize) -> Result<Option<Question>> {
let (name_bytes_consumed, name_str) = domain_name::read(buf, *offset, "question.name")?;
let bits_size = 4; if buf.len() < *offset + name_bytes_consumed + bits_size {
return Ok(None);
}
let bits = QuestionBits::unpack_from_slice(
&buf[*offset + name_bytes_consumed..*offset + name_bytes_consumed + bits_size],
)
.with_context(|| "couldn't unpack question bits")?;
*offset += name_bytes_consumed + bits_size;
Ok(Some(Question {
name: name_str,
resource_type: match enums_generated::resourcetype_int(bits.resource_type as usize) {
Some(e) => IntEnum::Enum(e),
None => IntEnum::Unknown(bits.resource_type),
},
resource_class: match enums_generated::resourceclass_int(bits.resource_class as usize) {
Some(e) => IntEnum::Enum(e),
None => IntEnum::Unknown(bits.resource_class),
},
}))
}
#[derive(PackedStruct)]
#[packed_struct(endian = "msb", bit_numbering = "msb0")]
pub struct ResourceTypeBits {
#[packed_field(bits = "0:15")]
resource_type: u16,
}
#[derive(PackedStruct)]
#[packed_struct(endian = "msb", bit_numbering = "msb0")]
pub struct ResourceClassTTLBits {
#[packed_field(bits = "0:15")]
class: u16,
#[packed_field(bits = "16:47")]
ttl: u32,
}
#[derive(Default, PackedStruct)]
#[packed_struct(endian = "msb", bit_numbering = "msb0")]
pub struct ResourceClassTTLOPTBits {
#[packed_field(bits = "0:15")]
pub udp_size: u16,
#[packed_field(bits = "16:23")]
pub response_code: u8,
#[packed_field(bits = "24:31")]
pub version: u8,
#[packed_field(bits = "32:32")]
pub dnssec_ok: bool,
#[packed_field(bits = "33:47")]
_reserved: ReservedZero<packed_bits::Bits<15>>,
}
#[derive(PackedStruct)]
#[packed_struct(endian = "msb", bit_numbering = "msb0")]
pub struct ResourceRdataLengthBits {
#[packed_field(bits = "0:15")]
rdata_len: u16,
}
pub enum RDataFields<'a> {
RDATA(&'a ResourceData),
IP(IpAddr),
SOA(rdata::SOAFields<'a>),
TXT(&'a String),
}
pub struct ResourceFields<'a> {
pub name: &'a str,
pub resource_type: IntEnum<u16, ResourceType>,
pub resource_class: IntEnum<u16, ResourceClass>,
pub ttl: u32,
pub rdata: RDataFields<'a>,
}
pub fn write_resource(
resource: &Resource,
buf: &mut BytesMut,
ptr_offsets: &mut domain_name::LabelOffsets,
) -> Result<()> {
write_resource_fields(
&ResourceFields {
name: resource.name.as_str(),
resource_type: resource.resource_type,
resource_class: resource.resource_class,
ttl: resource.ttl,
rdata: RDataFields::RDATA(&resource.rdata),
},
buf,
ptr_offsets,
)
}
pub fn write_resource_fields(
resource: &ResourceFields,
buf: &mut BytesMut,
ptr_offsets: &mut domain_name::LabelOffsets,
) -> Result<()> {
domain_name::write(resource.name, buf, ptr_offsets, "resource.name")?;
let type_bits = ResourceTypeBits {
resource_type: match resource.resource_type {
IntEnum::Enum(e) => e as u16,
IntEnum::Unknown(i) => i,
},
}
.pack()?;
let classttl_bits = ResourceClassTTLBits {
class: match resource.resource_class {
IntEnum::Enum(e) => e as u16,
IntEnum::Unknown(i) => i,
},
ttl: resource.ttl,
}
.pack()?;
let rdlen_zero_bits = ResourceRdataLengthBits {
rdata_len: 0 as u16, }
.pack()?;
buf.reserve(type_bits.len() + classttl_bits.len() + rdlen_zero_bits.len());
buf.put_slice(&type_bits);
buf.put_slice(&classttl_bits);
let rdata_len_offset = buf.len();
buf.put_slice(&rdlen_zero_bits);
let rdata_offset = buf.len();
match &resource.rdata {
RDataFields::RDATA(rdata) => {
rdata::write_rdata(&resource.resource_type, &rdata, buf, ptr_offsets)?
}
RDataFields::IP(IpAddr::V4(ip)) => rdata::write_a_ip(&ip, buf)?,
RDataFields::IP(IpAddr::V6(ip)) => rdata::write_aaaa_ip(&ip, buf)?,
RDataFields::SOA(soa) => rdata::write_soa_fields(&soa, buf, ptr_offsets)?,
RDataFields::TXT(txt) => rdata::write_txt_entry(txt.as_bytes(), buf)?,
}
if buf.len() > rdata_offset {
let rdlen_updated_bits = ResourceRdataLengthBits {
rdata_len: u16::try_from(buf.len() - rdata_offset)
.with_context(|| "rdata.length doesn't fit")?,
}
.pack()?;
let buf_len = buf.len();
let bits = buf
.get_mut(rdata_len_offset..rdata_offset)
.with_context(|| {
format!(
"failed to get bytes {}..{} from buffer with length {}",
rdata_len_offset, rdata_offset, buf_len
)
})?;
for i in 0..(rdata_offset - rdata_len_offset) {
bits[i] = rdlen_updated_bits[i];
}
}
Ok(())
}
pub fn read_resource_non_opt<'a>(buf: &[u8], offset: &mut usize) -> Result<Option<Resource>> {
match read_resource_name_type(buf, offset).context("Failed to read non-OPT resource")? {
Some((_name, OPT_RESOURCE_TYPE)) => bail!("Got OPT resource in unexpected part of message"),
Some((name, resource_type)) => {
match read_resource_remainder(buf, offset, name, resource_type)? {
Some(resource) => Ok(Some(resource)),
None => Ok(None),
}
}
None => Ok(None),
}
}
pub fn read_resource_name_type(buf: &[u8], offset: &mut usize) -> Result<Option<(String, u16)>> {
let (mut total_bytes_consumed, name_str) = domain_name::read(buf, *offset, "resource.name")?;
let type_bits_size = 2; if buf.len() < *offset + total_bytes_consumed + type_bits_size {
return Ok(None);
}
let resource_type = ResourceTypeBits::unpack_from_slice(
&buf[*offset + total_bytes_consumed..*offset + total_bytes_consumed + type_bits_size],
)
.with_context(|| "couldn't unpack resource prelude bits")?
.resource_type;
total_bytes_consumed += type_bits_size;
*offset += total_bytes_consumed;
Ok(Some((name_str, resource_type)))
}
pub fn read_resource_remainder_opt(buf: &[u8], offset: &mut usize) -> Result<Option<OPT>> {
let opt_class_ttl_bits_size = 6; if buf.len() < *offset + opt_class_ttl_bits_size {
return Ok(None);
}
let opt_class_ttl_bits = ResourceClassTTLOPTBits::unpack_from_slice(
&buf[*offset..*offset + opt_class_ttl_bits_size],
)
.with_context(|| "couldn't unpack OPT resource custom class/ttl bits")?;
let mut total_bytes_consumed = opt_class_ttl_bits_size;
let rdata_location = read_rdata_location(
buf,
*offset + total_bytes_consumed,
&mut total_bytes_consumed,
)
.context("Failed to read OPT rdata location")?;
match rdata_location {
Some((rdata_offset, rdata_len)) => {
let opt = rdata::read_opt(buf, rdata_offset, rdata_len, opt_class_ttl_bits)?;
total_bytes_consumed += rdata_len;
*offset += total_bytes_consumed;
Ok(Some(opt))
}
None => {
Ok(None)
}
}
}
pub fn read_resource_remainder(
buf: &[u8],
offset: &mut usize,
name: String,
resource_type: u16,
) -> Result<Option<Resource>> {
let class_ttl_bits_size = 6; if buf.len() < *offset + class_ttl_bits_size {
return Ok(None);
}
let class_ttl_bits =
ResourceClassTTLBits::unpack_from_slice(&buf[*offset..*offset + class_ttl_bits_size])
.with_context(|| "couldn't unpack resource class/ttl bits")?;
let mut total_bytes_consumed = class_ttl_bits_size;
let rdata_location = read_rdata_location(
buf,
*offset + total_bytes_consumed,
&mut total_bytes_consumed,
)
.with_context(|| {
format!(
"Failed to read {:?} resource rdata location for {}",
resource_type, name
)
})?;
match rdata_location {
Some((rdata_offset, rdata_len)) => {
let rdata = rdata::read_rdata(buf, resource_type, rdata_offset, rdata_len)
.with_context(|| {
format!("Failed to read {:?} resource for {}", resource_type, name)
})?;
total_bytes_consumed += rdata_len;
*offset += total_bytes_consumed;
Ok(Some(Resource {
name,
resource_type: match enums_generated::resourcetype_int(resource_type as usize) {
Some(r) => IntEnum::Enum(r),
None => IntEnum::Unknown(resource_type),
},
resource_class: match enums_generated::resourceclass_int(
class_ttl_bits.class as usize,
) {
Some(r) => IntEnum::Enum(r),
None => IntEnum::Unknown(resource_type),
},
ttl: class_ttl_bits.ttl,
rdata,
}))
}
None => {
Ok(None)
}
}
}
fn read_rdata_location(
buf: &[u8],
offset: usize,
total_bytes_consumed: &mut usize,
) -> Result<Option<(usize, usize)>> {
let rdata_length_bits_size = 2; let rdata_offset = offset + rdata_length_bits_size;
if buf.len() < rdata_offset {
return Ok(None);
}
let rdata_length_bits =
ResourceRdataLengthBits::unpack_from_slice(&buf[offset..offset + rdata_length_bits_size])
.with_context(|| "couldn't unpack rdata length bits")?;
let rdata_len = rdata_length_bits.rdata_len as usize;
if buf.len() < rdata_offset + rdata_len {
return Ok(None);
}
*total_bytes_consumed += rdata_length_bits_size;
Ok(Some((rdata_offset, rdata_len)))
}
pub fn write_opt(opt: &OPT, udp_size_override: Option<u16>, buf: &mut BytesMut) -> Result<()> {
domain_name::write_nopointer(".", buf, "opt.name")?;
let type_bits = ResourceTypeBits {
resource_type: ResourceType::OPT as u16,
}
.pack()?;
let classttl_bits = ResourceClassTTLOPTBits {
udp_size: match udp_size_override {
Some(udp_size) => udp_size,
None => opt.udp_size,
},
response_code: opt.response_code,
version: opt.response_code,
dnssec_ok: opt.dnssec_ok,
..ResourceClassTTLOPTBits::default() }
.pack()?;
let rdlen_zero_bits = ResourceRdataLengthBits {
rdata_len: 0 as u16, }
.pack()?;
buf.reserve(type_bits.len() + classttl_bits.len() + rdlen_zero_bits.len());
buf.put_slice(&type_bits);
buf.put_slice(&classttl_bits);
let rdata_len_offset = buf.len();
buf.put_slice(&rdlen_zero_bits);
let rdata_offset = buf.len();
rdata::write_opt(opt, buf)?;
if buf.len() > rdata_offset {
let rdlen_updated_bits = ResourceRdataLengthBits {
rdata_len: u16::try_from(buf.len() - rdata_offset)
.with_context(|| "rdata.length doesn't fit")?,
}
.pack()?;
let buf_len = buf.len();
let bits = buf
.get_mut(rdata_len_offset..rdata_offset)
.with_context(|| {
format!(
"failed to get bytes {}..{} from buffer with length {}",
rdata_len_offset, rdata_offset, buf_len
)
})?;
for i in 0..(rdata_offset - rdata_len_offset) {
bits[i] = rdlen_updated_bits[i];
}
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_update_message_id() {
let mut buf = BytesMut::with_capacity(5);
buf.resize(5, 0);
assert_eq!(5, buf.len(), "{:?}", buf);
let id = 511;
let mut res = update_message_id(id, &mut buf, 0);
assert!(res.is_ok(), "{:?}", res);
assert_eq!(1, buf[0]);
assert_eq!(255, buf[1]);
assert_eq!(0, buf[2]);
assert_eq!(0, buf[3]);
assert_eq!(0, buf[4]);
buf.clear();
buf.resize(5, 0);
res = update_message_id(id, &mut buf, 1);
assert!(res.is_ok(), "{:?}", res);
assert_eq!(0, buf[0]);
assert_eq!(1, buf[1]);
assert_eq!(255, buf[2]);
assert_eq!(0, buf[3]);
assert_eq!(0, buf[4]);
buf.clear();
buf.resize(5, 0);
res = update_message_id(id, &mut buf, 2);
assert!(res.is_ok(), "{:?}", res);
assert_eq!(0, buf[0]);
assert_eq!(0, buf[1]);
assert_eq!(1, buf[2]);
assert_eq!(255, buf[3]);
assert_eq!(0, buf[4]);
buf.clear();
buf.resize(5, 0);
res = update_message_id(id, &mut buf, 3);
assert!(res.is_ok(), "{:?}", res);
assert_eq!(0, buf[0]);
assert_eq!(0, buf[1]);
assert_eq!(0, buf[2]);
assert_eq!(1, buf[3]);
assert_eq!(255, buf[4]);
for offset in 4..10 {
buf.clear();
buf.resize(5, 0);
res = update_message_id(id, &mut buf, offset);
assert!(res.is_err(), "offset={} {:?}", offset, res);
assert_eq!(0, buf[0]);
assert_eq!(0, buf[1]);
assert_eq!(0, buf[2]);
assert_eq!(0, buf[3]);
assert_eq!(0, buf[4]);
}
}
}