Skip to content

Commit 5d59c8a

Browse files
committed
fix: learn IPs from DNS responses to handle anycast rotation in AllowList
Signed-off-by: danbugs <danilochiarlone@gmail.com>
1 parent a4831d7 commit 5d59c8a

1 file changed

Lines changed: 155 additions & 0 deletions

File tree

host/src/lib.rs

Lines changed: 155 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -185,6 +185,7 @@ pub enum NetworkPolicy {
185185
pub struct AllowList {
186186
allowed_ips: HashSet<IpAddr>,
187187
hostnames: Vec<String>,
188+
learned_ips: Arc<Mutex<HashSet<IpAddr>>>,
188189
}
189190

190191
impl AllowList {
@@ -219,13 +220,19 @@ impl AllowList {
219220
Ok(Self {
220221
allowed_ips,
221222
hostnames,
223+
learned_ips: Arc::new(Mutex::new(HashSet::new())),
222224
})
223225
}
224226

225227
fn is_allowed(&self, ip: &IpAddr) -> bool {
226228
if self.allowed_ips.contains(ip) {
227229
return true;
228230
}
231+
if let Ok(learned) = self.learned_ips.lock() {
232+
if learned.contains(ip) {
233+
return true;
234+
}
235+
}
229236
// Re-resolve hostnames to catch CDN/anycast IP rotation.
230237
use std::net::ToSocketAddrs;
231238
for host in &self.hostnames {
@@ -239,6 +246,12 @@ impl AllowList {
239246
}
240247
false
241248
}
249+
250+
fn learn_ip(&self, ip: IpAddr) {
251+
if let Ok(mut learned) = self.learned_ips.lock() {
252+
learned.insert(ip);
253+
}
254+
}
242255
}
243256

244257
/// A set of blocked network destinations.
@@ -990,6 +1003,121 @@ fn sockaddr_to_json(addr: SocketAddr) -> serde_json::Value {
9901003
})
9911004
}
9921005

1006+
/// Extract IPs from a DNS response for hostnames that match the allow list.
1007+
/// Minimal parser — handles standard A (type 1) and AAAA (type 28) answers.
1008+
fn learn_ips_from_dns_response(data: &[u8], al: &AllowList) {
1009+
if data.len() < 12 {
1010+
return;
1011+
}
1012+
let flags = u16::from_be_bytes([data[2], data[3]]);
1013+
let is_response = (flags & 0x8000) != 0;
1014+
if !is_response {
1015+
return;
1016+
}
1017+
let qdcount = u16::from_be_bytes([data[4], data[5]]) as usize;
1018+
let ancount = u16::from_be_bytes([data[6], data[7]]) as usize;
1019+
if qdcount == 0 || ancount == 0 {
1020+
return;
1021+
}
1022+
1023+
// Parse question section to extract the queried name.
1024+
let mut pos = 12;
1025+
let qname = match dns_read_name(data, &mut pos) {
1026+
Some(n) => n,
1027+
None => return,
1028+
};
1029+
// Skip QTYPE (2) + QCLASS (2)
1030+
pos += 4;
1031+
if pos > data.len() {
1032+
return;
1033+
}
1034+
1035+
// Check if the queried name matches any allowed hostname.
1036+
let qname_lower = qname.to_lowercase();
1037+
let is_allowed_host = al.hostnames.iter().any(|h| h.to_lowercase() == qname_lower);
1038+
if !is_allowed_host {
1039+
return;
1040+
}
1041+
1042+
// Parse answer records and learn IPs.
1043+
for _ in 0..ancount {
1044+
// Skip name (may be a pointer)
1045+
if dns_read_name(data, &mut pos).is_none() {
1046+
return;
1047+
}
1048+
if pos + 10 > data.len() {
1049+
return;
1050+
}
1051+
let rtype = u16::from_be_bytes([data[pos], data[pos + 1]]);
1052+
let rdlen = u16::from_be_bytes([data[pos + 8], data[pos + 9]]) as usize;
1053+
pos += 10;
1054+
if pos + rdlen > data.len() {
1055+
return;
1056+
}
1057+
match rtype {
1058+
1 if rdlen == 4 => {
1059+
let ip = IpAddr::V4(std::net::Ipv4Addr::new(
1060+
data[pos], data[pos + 1], data[pos + 2], data[pos + 3],
1061+
));
1062+
al.learn_ip(ip);
1063+
}
1064+
28 if rdlen == 16 => {
1065+
let mut octets = [0u8; 16];
1066+
octets.copy_from_slice(&data[pos..pos + 16]);
1067+
al.learn_ip(IpAddr::V6(std::net::Ipv6Addr::from(octets)));
1068+
}
1069+
_ => {}
1070+
}
1071+
pos += rdlen;
1072+
}
1073+
}
1074+
1075+
/// Read a DNS name at `pos`, advancing pos past it. Returns the decoded name.
1076+
fn dns_read_name(data: &[u8], pos: &mut usize) -> Option<String> {
1077+
let mut name = String::new();
1078+
let mut p = *pos;
1079+
let mut jumped = false;
1080+
let mut jump_save = 0;
1081+
loop {
1082+
if p >= data.len() {
1083+
return None;
1084+
}
1085+
let len = data[p] as usize;
1086+
if len == 0 {
1087+
p += 1;
1088+
break;
1089+
}
1090+
if (len & 0xC0) == 0xC0 {
1091+
// Pointer
1092+
if p + 1 >= data.len() {
1093+
return None;
1094+
}
1095+
let offset = ((len & 0x3F) << 8) | data[p + 1] as usize;
1096+
if !jumped {
1097+
jump_save = p + 2;
1098+
jumped = true;
1099+
}
1100+
p = offset;
1101+
continue;
1102+
}
1103+
p += 1;
1104+
if p + len > data.len() {
1105+
return None;
1106+
}
1107+
if !name.is_empty() {
1108+
name.push('.');
1109+
}
1110+
name.push_str(&String::from_utf8_lossy(&data[p..p + len]));
1111+
p += len;
1112+
}
1113+
if jumped {
1114+
*pos = jump_save;
1115+
} else {
1116+
*pos = p;
1117+
}
1118+
Some(name)
1119+
}
1120+
9931121
fn register_net_tools(
9941122
tools: &mut ToolRegistry,
9951123
policy: &NetworkPolicy,
@@ -1152,6 +1280,7 @@ fn register_net_tools(
11521280

11531281
// net_recvfrom
11541282
let t = table.clone();
1283+
let pol_recv = policy.clone();
11551284
tools.register("net_recvfrom", move |args| {
11561285
let fd = args["fd"].as_u64().ok_or_else(|| anyhow!("missing 'fd'"))?;
11571286
let len = args["len"].as_u64().unwrap_or(4096) as usize;
@@ -1166,6 +1295,18 @@ fn register_net_tools(
11661295
sock.recv_from(buf_init)?
11671296
};
11681297
buf.truncate(n);
1298+
1299+
// Learn IPs from DNS responses so AllowList stays current with
1300+
// anycast/CDN rotation (guest may resolve via a different DNS
1301+
// server than the host, getting different IPs for the same name).
1302+
if let Some(pa) = peer.as_socket() {
1303+
if pa.port() == 53 {
1304+
if let NetworkPolicy::AllowList(al) = &*pol_recv {
1305+
learn_ips_from_dns_response(&buf, al);
1306+
}
1307+
}
1308+
}
1309+
11691310
let encoded = base64::engine::general_purpose::STANDARD.encode(&buf);
11701311
let mut resp = json!({ "data": encoded, "len": n });
11711312
if let Some(pa) = peer.as_socket() {
@@ -1377,10 +1518,24 @@ impl FsRouter {
13771518
let target = fs.resolve(rel)?;
13781519
let md =
13791520
std::fs::metadata(&target).map_err(|e| anyhow!("fs_stat {:?}: {}", path, e))?;
1521+
let mtime_ns = md
1522+
.modified()
1523+
.ok()
1524+
.and_then(|t| t.duration_since(std::time::UNIX_EPOCH).ok())
1525+
.map(|d| d.as_nanos() as u64)
1526+
.unwrap_or(0);
1527+
let atime_ns = md
1528+
.accessed()
1529+
.ok()
1530+
.and_then(|t| t.duration_since(std::time::UNIX_EPOCH).ok())
1531+
.map(|d| d.as_nanos() as u64)
1532+
.unwrap_or(0);
13801533
Ok(json!({
13811534
"size": md.len(),
13821535
"is_dir": md.is_dir(),
13831536
"is_file": md.is_file(),
1537+
"mtime_ns": mtime_ns,
1538+
"atime_ns": atime_ns,
13841539
}))
13851540
});
13861541

0 commit comments

Comments
 (0)