// DNS domain names and RFC 1035 compression (section 4.1.4).
///|
fn name_to_lowercase(name : String) -> String {
let result = Ref("")
for c in name {
result.val = result.val + c.to_ascii_lowercase().to_string()
}
result.val
}
///|
fn is_root_domain(name : String) -> Bool {
name == "" || name == "."
}
///|
fn is_ascii_decimal_digit(char : Char) -> Bool {
char >= '0' && char <= '9'
}
///|
fn raw_label_is_valid(label : String) -> Result[Unit, String] {
if label.length() == 0 {
return Err("DNS name contains an empty label")
}
if label.length() > max_label_len {
return Err("DNS label exceeds 63 octets")
}
for char in label {
if char.to_int() > 0xFF {
return Err("DNS label contains a non-octet character")
}
}
Ok(())
}
// DNS uses octets rather than Unicode strings. Public names use the standard
// presentation convention: an unescaped dot separates labels; `\\DDD` is an
// octet written in decimal, and `\\X` makes a following non-digit octet
// literal (digits use their unambiguous three-digit form). This keeps decoded
// labels containing dots, backslashes, or control bytes reversible through
// the String API.
///|
fn split_labels_checked(name : String) -> Result[Array[String], String] {
if is_root_domain(name) {
return Ok(Array::new(capacity=0))
}
let labels : Array[String] = Array::new(capacity=6)
let current = Ref("")
let pos = Ref(0)
while pos.val < name.length() {
let char = name[pos.val].unsafe_to_char()
if char == '.' {
match raw_label_is_valid(current.val) {
Ok(_) => labels.push(current.val)
Err(error) => return Err(error)
}
current.val = ""
pos.val = pos.val + 1
} else if char == '\\' {
if pos.val + 1 >= name.length() {
return Err("DNS name ends with an incomplete escape")
}
let next = name[pos.val + 1].unsafe_to_char()
if is_ascii_decimal_digit(next) {
if pos.val + 3 >= name.length() ||
!is_ascii_decimal_digit(name[pos.val + 2].unsafe_to_char()) ||
!is_ascii_decimal_digit(name[pos.val + 3].unsafe_to_char()) {
return Err("DNS decimal escape must contain exactly three digits")
}
let value = (next.to_int() - '0'.to_int()) * 100 +
(name[pos.val + 2].unsafe_to_char().to_int() - '0'.to_int()) * 10 +
(name[pos.val + 3].unsafe_to_char().to_int() - '0'.to_int())
if value > 255 {
return Err("DNS decimal escape exceeds one octet")
}
current.val = current.val + value.unsafe_to_char().to_string()
pos.val = pos.val + 4
} else {
if next.to_int() > 0xFF {
return Err("DNS escape contains a non-octet character")
}
current.val = current.val + next.to_string()
pos.val = pos.val + 2
}
} else {
if char.to_int() > 0xFF {
return Err("DNS label contains a non-octet character")
}
current.val = current.val + char.to_string()
pos.val = pos.val + 1
}
}
if current.val != "" {
match raw_label_is_valid(current.val) {
Ok(_) => labels.push(current.val)
Err(error) => return Err(error)
}
} else if !name.has_suffix(".") {
return Err("DNS name contains an empty label")
}
let wire_len = Ref(1)
for label in labels {
wire_len.val = wire_len.val + 1 + label.length()
}
if wire_len.val > max_name_len {
Err("DNS name exceeds 255 octets")
} else {
Ok(labels)
}
}
///|
fn decimal_escape(byte : Int) -> String {
"\\" +
(byte / 100 + '0'.to_int()).unsafe_to_char().to_string() +
(byte / 10 % 10 + '0'.to_int()).unsafe_to_char().to_string() +
(byte % 10 + '0'.to_int()).unsafe_to_char().to_string()
}
///|
fn presentation_label(label : String) -> String {
let out = StringBuilder()
for char in label {
let byte = char.to_int()
if byte >= 0x21 && byte <= 0x7E && char != '.' && char != '\\' {
out.write_char(char)
} else {
out.write_string(decimal_escape(byte))
}
}
out.to_string()
}
///|
fn presentation_name(labels : Array[String]) -> String {
let out = StringBuilder()
for index, label in labels {
if index > 0 {
out.write_char('.')
}
out.write_string(presentation_label(label))
}
out.to_string()
}
// Length-prefix each raw label before it enters a key. A literal dot, colon,
// or backslash is then data rather than an ambiguous separator. DNS comparison
// folds ASCII letters only (RFC 4343), never arbitrary bytes.
///|
fn canonical_label_sequence(labels : Array[String], start : Int) -> String {
let out = StringBuilder()
for index in start.. Result[Array[Byte], String] {
let labels = match split_labels_checked(name) {
Ok(value) => value
Err(err) => return Err(err)
}
let length = Ref(1)
for label in labels {
length.val = length.val + 1 + label.length()
}
let out = Array::make(length.val, (0).to_byte())
let pos = Ref(0)
for label in labels {
out[pos.val] = label.length().to_byte()
pos.val = pos.val + 1
for i in 0.. Array[Byte] {
match encode_name_checked(name) {
Ok(bytes) => bytes
Err(error) => abort(error)
}
}
///|
fn suffix_from(labels : Array[String], start : Int) -> String {
canonical_label_sequence(labels, start)
}
///|
fn write_name_compressed(
out : Array[Byte],
offsets : Map[String, Int],
name : String,
) -> Result[Unit, String] {
let labels = match split_labels_checked(name) {
Ok(value) => value
Err(err) => return Err(err)
}
if labels.length() == 0 {
out.push(0)
return Ok(())
}
let pointer_at = Ref(-1)
let pointer_target = Ref(0)
for i in 0..= 0 && target < 0x4000 {
pointer_at.val = i
pointer_target.val = target
}
}
}
}
let stop = if pointer_at.val == -1 { labels.length() } else { pointer_at.val }
for i in 0..> 8) & 0x3F)).to_byte())
out.push((pointer_target.val & 0xFF).to_byte())
}
Ok(())
}
///|
fn decode_name_limited(
bytes : Array[Byte],
offset : Int,
msg_start : Int,
encoded_end : Int,
) -> Result[(String, Int), String] {
if offset < msg_start || offset >= encoded_end || encoded_end > bytes.length() {
return Err("DNS name starts outside its containing field")
}
let labels : Array[String] = Array::new(capacity=8)
let current = Ref(offset)
let jumped = Ref(false)
let next_offset = Ref(offset)
let hops = Ref(0)
let wire_length = Ref(1)
let terminated = Ref(false)
while !terminated.val {
let limit = if jumped.val { bytes.length() } else { encoded_end }
if current.val < msg_start || current.val >= limit {
return Err("truncated DNS name")
}
let first = bytes[current.val].to_int()
if first == 0 {
if !jumped.val {
next_offset.val = current.val + 1
}
terminated.val = true
} else if (first & 0xC0) == 0xC0 {
if current.val + 1 >= limit {
return Err("truncated DNS compression pointer")
}
let target = msg_start +
((first & 0x3F) << 8) +
bytes[current.val + 1].to_int()
// RFC 1035 pointers point to an earlier occurrence. Requiring this also
// rejects self loops and all forward/cyclic pointer constructions.
if target < msg_start + dns_header_size ||
target >= current.val ||
target >= bytes.length() {
return Err("invalid DNS compression pointer")
}
hops.val = hops.val + 1
if hops.val > max_pointer_hops {
return Err("too many DNS compression pointers")
}
if !jumped.val {
next_offset.val = current.val + 2
}
jumped.val = true
current.val = target
} else if (first & 0xC0) != 0 {
return Err("invalid DNS label tag")
} else if first > max_label_len {
return Err("DNS label exceeds 63 octets")
} else {
if current.val + 1 + first > limit {
return Err("truncated DNS label")
}
if labels.length() >= max_dns_labels {
return Err("too many DNS labels")
}
wire_length.val = wire_length.val + 1 + first
if wire_length.val > max_name_len {
return Err("DNS name exceeds 255 octets")
}
let chars = Array::make(first, ' ')
for i in 0.. Result[(String, Int), String] {
decode_name_limited(bytes, offset, msg_start, bytes.length())
}
///|
fn decode_name_in_rdata(
bytes : Array[Byte],
offset : Int,
msg_start : Int,
rdata_end : Int,
) -> Result[(String, Int), String] {
decode_name_limited(bytes, offset, msg_start, rdata_end)
}
///|
fn name_wire_size(name : String) -> Int {
match encode_name_checked(name) {
Ok(bytes) => bytes.length()
Err(_) => 0
}
}