Skip to content

Commit 86d2609

Browse files
authored
Merge pull request #56 from hyperlight-dev/fix/allowlist-dns-rotation
2 parents c6128d7 + 4c374d9 commit 86d2609

1 file changed

Lines changed: 158 additions & 0 deletions

File tree

host/src/lib.rs

Lines changed: 158 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.
@@ -1016,6 +1029,124 @@ fn sockaddr_to_json(addr: SocketAddr) -> serde_json::Value {
10161029
})
10171030
}
10181031

1032+
/// Extract IPs from a DNS response for hostnames that match the allow list.
1033+
/// Minimal parser — handles standard A (type 1) and AAAA (type 28) answers.
1034+
fn learn_ips_from_dns_response(data: &[u8], al: &AllowList) {
1035+
if data.len() < 12 {
1036+
return;
1037+
}
1038+
let flags = u16::from_be_bytes([data[2], data[3]]);
1039+
let is_response = (flags & 0x8000) != 0;
1040+
if !is_response {
1041+
return;
1042+
}
1043+
let qdcount = u16::from_be_bytes([data[4], data[5]]) as usize;
1044+
let ancount = u16::from_be_bytes([data[6], data[7]]) as usize;
1045+
if qdcount == 0 || ancount == 0 {
1046+
return;
1047+
}
1048+
1049+
// Parse question section to extract the queried name.
1050+
let mut pos = 12;
1051+
let qname = match dns_read_name(data, &mut pos) {
1052+
Some(n) => n,
1053+
None => return,
1054+
};
1055+
// Skip QTYPE (2) + QCLASS (2)
1056+
pos += 4;
1057+
if pos > data.len() {
1058+
return;
1059+
}
1060+
1061+
// Check if the queried name matches any allowed hostname.
1062+
let qname_lower = qname.to_lowercase();
1063+
let is_allowed_host = al.hostnames.iter().any(|h| h.to_lowercase() == qname_lower);
1064+
if !is_allowed_host {
1065+
return;
1066+
}
1067+
1068+
// Parse answer records and learn IPs.
1069+
for _ in 0..ancount {
1070+
// Skip name (may be a pointer)
1071+
if dns_read_name(data, &mut pos).is_none() {
1072+
return;
1073+
}
1074+
if pos + 10 > data.len() {
1075+
return;
1076+
}
1077+
let rtype = u16::from_be_bytes([data[pos], data[pos + 1]]);
1078+
let rdlen = u16::from_be_bytes([data[pos + 8], data[pos + 9]]) as usize;
1079+
pos += 10;
1080+
if pos + rdlen > data.len() {
1081+
return;
1082+
}
1083+
match rtype {
1084+
1 if rdlen == 4 => {
1085+
let ip = IpAddr::V4(std::net::Ipv4Addr::new(
1086+
data[pos],
1087+
data[pos + 1],
1088+
data[pos + 2],
1089+
data[pos + 3],
1090+
));
1091+
al.learn_ip(ip);
1092+
}
1093+
28 if rdlen == 16 => {
1094+
let mut octets = [0u8; 16];
1095+
octets.copy_from_slice(&data[pos..pos + 16]);
1096+
al.learn_ip(IpAddr::V6(std::net::Ipv6Addr::from(octets)));
1097+
}
1098+
_ => {}
1099+
}
1100+
pos += rdlen;
1101+
}
1102+
}
1103+
1104+
/// Read a DNS name at `pos`, advancing pos past it. Returns the decoded name.
1105+
fn dns_read_name(data: &[u8], pos: &mut usize) -> Option<String> {
1106+
let mut name = String::new();
1107+
let mut p = *pos;
1108+
let mut jumped = false;
1109+
let mut jump_save = 0;
1110+
loop {
1111+
if p >= data.len() {
1112+
return None;
1113+
}
1114+
let len = data[p] as usize;
1115+
if len == 0 {
1116+
p += 1;
1117+
break;
1118+
}
1119+
if (len & 0xC0) == 0xC0 {
1120+
// Pointer
1121+
if p + 1 >= data.len() {
1122+
return None;
1123+
}
1124+
let offset = ((len & 0x3F) << 8) | data[p + 1] as usize;
1125+
if !jumped {
1126+
jump_save = p + 2;
1127+
jumped = true;
1128+
}
1129+
p = offset;
1130+
continue;
1131+
}
1132+
p += 1;
1133+
if p + len > data.len() {
1134+
return None;
1135+
}
1136+
if !name.is_empty() {
1137+
name.push('.');
1138+
}
1139+
name.push_str(&String::from_utf8_lossy(&data[p..p + len]));
1140+
p += len;
1141+
}
1142+
if jumped {
1143+
*pos = jump_save;
1144+
} else {
1145+
*pos = p;
1146+
}
1147+
Some(name)
1148+
}
1149+
10191150
fn register_net_tools(
10201151
tools: &mut ToolRegistry,
10211152
policy: &NetworkPolicy,
@@ -1178,6 +1309,7 @@ fn register_net_tools(
11781309

11791310
// net_recvfrom
11801311
let t = table.clone();
1312+
let pol_recv = policy.clone();
11811313
tools.register("net_recvfrom", move |args| {
11821314
let fd = args["fd"].as_u64().ok_or_else(|| anyhow!("missing 'fd'"))?;
11831315
let len = args["len"].as_u64().unwrap_or(4096) as usize;
@@ -1192,6 +1324,18 @@ fn register_net_tools(
11921324
sock.recv_from(buf_init)?
11931325
};
11941326
buf.truncate(n);
1327+
1328+
// Learn IPs from DNS responses so AllowList stays current with
1329+
// anycast/CDN rotation (guest may resolve via a different DNS
1330+
// server than the host, getting different IPs for the same name).
1331+
if let Some(pa) = peer.as_socket() {
1332+
if pa.port() == 53 {
1333+
if let NetworkPolicy::AllowList(al) = &*pol_recv {
1334+
learn_ips_from_dns_response(&buf, al);
1335+
}
1336+
}
1337+
}
1338+
11951339
let encoded = base64::engine::general_purpose::STANDARD.encode(&buf);
11961340
let mut resp = json!({ "data": encoded, "len": n });
11971341
if let Some(pa) = peer.as_socket() {
@@ -1405,10 +1549,24 @@ impl FsRouter {
14051549
let target = fs.resolve(rel)?;
14061550
let md =
14071551
std::fs::metadata(&target).map_err(|e| anyhow!("fs_stat {:?}: {}", path, e))?;
1552+
let mtime_ns = md
1553+
.modified()
1554+
.ok()
1555+
.and_then(|t| t.duration_since(std::time::UNIX_EPOCH).ok())
1556+
.map(|d| d.as_nanos() as u64)
1557+
.unwrap_or(0);
1558+
let atime_ns = md
1559+
.accessed()
1560+
.ok()
1561+
.and_then(|t| t.duration_since(std::time::UNIX_EPOCH).ok())
1562+
.map(|d| d.as_nanos() as u64)
1563+
.unwrap_or(0);
14081564
Ok(json!({
14091565
"size": md.len(),
14101566
"is_dir": md.is_dir(),
14111567
"is_file": md.is_file(),
1568+
"mtime_ns": mtime_ns,
1569+
"atime_ns": atime_ns,
14121570
}))
14131571
});
14141572

0 commit comments

Comments
 (0)