///|
fn v2_codes(header : ProxyHeader) -> (Byte, Byte) {
let command = if header.command == Local { 0 } else { 1 }
let family = match header.family {
Unspec => 0
Inet => 1
Inet6 => 2
Unix => 3
}
let transport = match header.transport {
Unspec => 0
Stream => 1
Datagram => 2
}
((0x20 | command).to_byte(), ((family << 4) | transport).to_byte())
}
///|
fn encode_v2_address(header : ProxyHeader) -> Result[Bytes, ProxyError] {
if header.command == Local {
return Ok(Bytes::new(0))
}
let buffer = @buffer.Buffer(size_hint=216)
match header.address {
NoAddress => ()
RawUnsupported(raw) => buffer.write_bytes(raw)
Ipv4(endpoints) => {
buffer.write_bytes(Bytes::from_array(endpoints.source_address.octets))
buffer.write_bytes(
Bytes::from_array(endpoints.destination_address.octets),
)
buffer.write_bytes(write_u16_be(endpoints.source_port))
buffer.write_bytes(write_u16_be(endpoints.destination_port))
}
Ipv6(endpoints) => {
buffer.write_bytes(Bytes::from_array(endpoints.source_address.octets))
buffer.write_bytes(
Bytes::from_array(endpoints.destination_address.octets),
)
buffer.write_bytes(write_u16_be(endpoints.source_port))
buffer.write_bytes(write_u16_be(endpoints.destination_port))
}
UnixAddress(endpoints) => {
buffer.write_bytes(endpoints.source)
buffer.write_bytes(endpoints.destination)
}
}
Ok(buffer.to_bytes())
}
///|
pub fn encode_v2(header : ProxyHeader) -> Result[Bytes, ProxyError] {
if header.version != V2 {
return Err(
proxy_error(PolicyViolation, 0, "v2 encoder requires a V2 header"),
)
}
let address = match encode_v2_address(header) {
Err(err) => return Err(err)
Ok(v) => v
}
let tlvs = match encode_tlvs(header.tlvs) {
Err(err) => return Err(err)
Ok(v) => v
}
let body_length = address.length() + tlvs.length()
if body_length > 65535 {
return Err(
proxy_error(HeaderTooLarge, 14, "v2 body is longer than 65535 bytes"),
)
}
let codes = v2_codes(header)
let buffer = @buffer.Buffer(size_hint=16 + body_length)
buffer.write_bytes(v2_signature())
buffer.write_byte(codes.0)
buffer.write_byte(codes.1)
buffer.write_bytes(write_u16_be(body_length))
buffer.write_bytes(address)
buffer.write_bytes(tlvs)
Ok(buffer.to_bytes())
}
///|
pub fn insert_proxy_crc32c(
header_without_checksum : ProxyHeader,
) -> Result[Bytes, ProxyError] {
if header_without_checksum.version != V2 {
return Err(proxy_error(PolicyViolation, 0, "CRC32C exists only in v2"))
}
if find_tlv(header_without_checksum.tlvs, TLV_CRC32C) is Some(_) {
return Err(
proxy_error(DuplicateCrc32c, 0, "header already has a CRC32C TLV"),
)
}
let tlvs = header_without_checksum.tlvs
tlvs.push({
type_code: TLV_CRC32C,
value: Bytes::from_array([b'\x00', b'\x00', b'\x00', b'\x00']),
})
let with_zero = { ..header_without_checksum, tlvs, }
match encode_v2(with_zero) {
Err(err) => Err(err)
Ok(encoded) => {
let checksum = crc32c(encoded)
let result = encoded.to_array()
let mut at = 16
let family = encoded[13].to_int() >> 4
let address_length = if family == 1 {
12
} else if family == 2 {
36
} else if family == 3 {
216
} else {
0
}
at = at + address_length
while at < result.length() {
let length = (result[at + 1].to_int() << 8) | result[at + 2].to_int()
if result[at] == TLV_CRC32C {
result[at + 3] = (checksum >> 24).to_byte()
result[at + 4] = (checksum >> 16).to_byte()
result[at + 5] = (checksum >> 8).to_byte()
result[at + 6] = checksum.to_byte()
break
}
at = at + 3 + length
}
Ok(Bytes::from_array(result))
}
}
}