use std::{
net::{Ipv4Addr, SocketAddr},
sync::Arc,
time::Duration,
};
use base64::Engine;
use netstack::{CreateSocket, netcore::Channel, netsock::TcpStream};
use tokio::{
io::{AsyncReadExt, AsyncWriteExt},
sync::watch,
time::timeout,
};
use ts_dns_wire::{Rcode, decode_query, encode_response};
use ts_packetfilter::FilterExt;
use crate::{
magic_dns::{
ClientTransport, Decision, DnsView, check_response_size_and_set_tc, decide, forward_query,
},
packetfilter::LiveFilterRx,
};
const MAX_REQUEST: usize = 8 * 1024;
const CLIENT_TIMEOUT: Duration = Duration::from_secs(5);
const PEER_CLIENT_TRANSPORT: ClientTransport = ClientTransport::Tcp;
const MAX_CLIENT_RESPONSE: usize = u16::MAX as usize;
pub(crate) async fn forward_doh(
channel: &Channel,
doh_addr: SocketAddr,
query: &[u8],
fallback: Vec<u8>,
client: ClientTransport,
) -> Vec<u8> {
match timeout(CLIENT_TIMEOUT, doh_round_trip(channel, doh_addr, query)).await {
Ok(Ok(resp)) if !resp.is_empty() => check_response_size_and_set_tc(query, resp, client),
Ok(Ok(_)) => {
tracing::warn!(%doh_addr, "peerapi doh client: empty response from exit node");
fallback
}
Ok(Err(e)) => {
tracing::warn!(error = %e, %doh_addr, "peerapi doh client: delegation failed");
fallback
}
Err(_) => {
tracing::warn!(%doh_addr, "peerapi doh client: delegation timed out");
fallback
}
}
}
async fn doh_round_trip(
channel: &Channel,
doh_addr: SocketAddr,
query: &[u8],
) -> std::io::Result<Vec<u8>> {
let local = SocketAddr::new(Ipv4Addr::UNSPECIFIED.into(), 0);
let mut stream = channel
.tcp_connect(local, doh_addr)
.await
.map_err(|e| std::io::Error::other(e.to_string()))?;
let request = format!(
"POST /dns-query HTTP/1.1\r\nHost: {doh_addr}\r\nContent-Type: application/dns-message\r\nAccept: application/dns-message\r\nContent-Length: {}\r\nConnection: close\r\n\r\n",
query.len()
);
stream.write_all(request.as_bytes()).await?;
stream.write_all(query).await?;
stream.flush().await?;
read_doh_response(&mut stream).await
}
async fn read_doh_response(stream: &mut TcpStream) -> std::io::Result<Vec<u8>> {
let mut buf = Vec::with_capacity(1024);
let mut tmp = [0u8; 1024];
let header_end = loop {
if let Some(pos) = find_header_end(&buf) {
break pos;
}
if buf.len() > MAX_CLIENT_RESPONSE {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
"doh response headers too large",
));
}
let n = stream.read(&mut tmp).await?;
if n == 0 {
return Err(std::io::Error::new(
std::io::ErrorKind::UnexpectedEof,
"eof before doh response headers",
));
}
buf.extend_from_slice(&tmp[..n]);
};
let content_length = parse_response_head(&buf)?;
let mut body = buf[header_end..].to_vec();
while body.len() < content_length {
let n = stream.read(&mut tmp).await?;
if n == 0 {
return Err(std::io::Error::new(
std::io::ErrorKind::UnexpectedEof,
"eof before doh response body complete",
));
}
body.extend_from_slice(&tmp[..n]);
if body.len() > MAX_CLIENT_RESPONSE {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
"doh response body too large",
));
}
}
body.truncate(content_length);
Ok(body)
}
fn parse_response_head(buf: &[u8]) -> std::io::Result<usize> {
let mut headers = [httparse::EMPTY_HEADER; 32];
let mut resp = httparse::Response::new(&mut headers);
match resp.parse(buf) {
Ok(httparse::Status::Complete(_)) => {}
Ok(httparse::Status::Partial) | Err(_) => {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
"malformed doh response headers",
));
}
}
if resp.code != Some(200) {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!("doh response status {:?}", resp.code),
));
}
let content_length = resp
.headers
.iter()
.find(|h| h.name.eq_ignore_ascii_case("content-length"))
.and_then(|h| std::str::from_utf8(h.value).ok())
.and_then(|v| v.trim().parse::<usize>().ok())
.ok_or_else(|| {
std::io::Error::new(
std::io::ErrorKind::InvalidData,
"doh response missing length",
)
})?;
if content_length > MAX_CLIENT_RESPONSE {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
"doh response body too large",
));
}
Ok(content_length)
}
pub(crate) async fn handle_conn(
mut stream: TcpStream,
seed: Vec<u8>,
header_end: usize,
channel: &Channel,
view_rx: &watch::Receiver<Arc<DnsView>>,
filter_rx: &LiveFilterRx,
forward_exit_egress: bool,
) -> std::io::Result<()> {
let request = match read_request(&mut stream, seed, header_end).await? {
Some(r) => r,
None => return Ok(()),
};
let query = match request {
DohRequest::TooLarge => {
return write_status(&mut stream, "413 Payload Too Large").await;
}
DohRequest::BadRequest => {
return write_status(&mut stream, "400 Bad Request").await;
}
DohRequest::NotFound => {
return write_status(&mut stream, "404 Not Found").await;
}
DohRequest::Query(bytes) => bytes,
};
let view = view_rx.borrow().clone();
let filter = filter_rx.borrow().clone();
let src = stream.remote_addr().ip();
if !dns_source_allowed(&view, filter.as_ref().map(|f| &*f.0), src) {
tracing::debug!(
%src,
"peerapi doh: 403, the ACL does not grant this peer internet access through us"
);
return write_status(&mut stream, "403 Forbidden").await;
}
let response = resolve(&view, &query, channel, forward_exit_egress).await;
write_dns_response(&mut stream, &response).await
}
const TCP_SYN_FLAG: u8 = 0x02;
fn internet_probe_dst(src: std::net::IpAddr) -> std::net::IpAddr {
match src {
std::net::IpAddr::V4(_) => std::net::Ipv4Addr::UNSPECIFIED.into(),
std::net::IpAddr::V6(_) => std::net::Ipv6Addr::new(0x2000, 0, 0, 0, 0, 0, 0, 0).into(),
}
}
pub(crate) fn dns_source_allowed(
view: &DnsView,
filter: Option<&(dyn ts_packetfilter::Filter + Send + Sync)>,
src: std::net::IpAddr,
) -> bool {
if view.peers.as_ref().and_then(|p| p.get(&src)).is_none() {
tracing::debug!(%src, "peerapi doh: source is not a known tailnet peer");
return false;
}
let Some(filter) = filter else {
tracing::debug!(%src, "peerapi doh: no packet filter compiled yet; refusing");
return false;
};
let info = ts_packetfilter::PacketInfo {
src,
dst: internet_probe_dst(src),
ip_proto: ts_packetfilter::IpProto::TCP,
port: 53,
l4: ts_packetfilter::L4Header::Tcp {
flags: TCP_SYN_FLAG,
},
};
let caps = [];
filter.can_access(&info, caps)
}
async fn resolve(
view: &DnsView,
query: &[u8],
channel: &Channel,
forward_exit_egress: bool,
) -> Vec<u8> {
match server_decide(view, query, forward_exit_egress) {
ServerDecision::Reply(resp) => resp,
ServerDecision::Forward {
upstreams,
query,
servfail,
} => forward_query(channel, &upstreams, &query, servfail, PEER_CLIENT_TRANSPORT).await,
}
}
enum ServerDecision {
Reply(Vec<u8>),
Forward {
upstreams: Vec<SocketAddr>,
query: Vec<u8>,
servfail: Vec<u8>,
},
}
fn server_decide(view: &DnsView, query: &[u8], forward_exit_egress: bool) -> ServerDecision {
let Ok(decoded) = decode_query(query) else {
let id = if query.len() >= 2 {
u16::from_be_bytes([query[0], query[1]])
} else {
0
};
return ServerDecision::Reply(encode_formerr(id));
};
let canon = decoded.question.name.to_canon();
if view.cfg.exit_node_filters(&canon) {
return ServerDecision::Reply(encode_response(
decoded.id,
&decoded.question,
decoded.recursion_desired,
Rcode::Refused,
&[],
None,
));
}
match decide(view, query) {
None => ServerDecision::Reply(encode_formerr(decoded.id)),
Some(Decision::Reply(resp)) => ServerDecision::Reply(resp),
Some(Decision::Forward {
upstreams,
query,
servfail,
recursive: _,
}) => {
if !forward_exit_egress {
return ServerDecision::Reply(encode_response(
decoded.id,
&decoded.question,
decoded.recursion_desired,
Rcode::Refused,
&[],
None,
));
}
ServerDecision::Forward {
upstreams,
query,
servfail,
}
}
}
}
fn encode_formerr(id: u16) -> Vec<u8> {
let mut msg = vec![0u8; 12];
msg[0..2].copy_from_slice(&id.to_be_bytes());
msg[2] = 0x80; msg[3] = 0x01; msg
}
enum DohRequest {
Query(Vec<u8>),
TooLarge,
BadRequest,
NotFound,
}
async fn read_request(
stream: &mut TcpStream,
buf: Vec<u8>,
header_end: usize,
) -> std::io::Result<Option<DohRequest>> {
let mut tmp = [0u8; 1024];
let mut headers = [httparse::EMPTY_HEADER; 32];
let mut req = httparse::Request::new(&mut headers);
let parsed = match req.parse(&buf) {
Ok(httparse::Status::Complete(n)) => n,
Ok(httparse::Status::Partial) => return Ok(Some(DohRequest::BadRequest)),
Err(_) => return Ok(Some(DohRequest::BadRequest)),
};
debug_assert_eq!(parsed, header_end);
let method = req.method.unwrap_or("");
let path = req.path.unwrap_or("");
let (raw_path, query_str) = match path.split_once('?') {
Some((p, q)) => (p, Some(q)),
None => (path, None),
};
if raw_path != "/dns-query" {
return Ok(Some(DohRequest::NotFound));
}
match method {
"GET" => Ok(Some(parse_get(query_str))),
"POST" => {
let content_length =
header_value(&req, "content-length").and_then(|v| v.trim().parse::<usize>().ok());
let Some(len) = content_length else {
return Ok(Some(DohRequest::BadRequest));
};
if len > MAX_REQUEST {
return Ok(Some(DohRequest::TooLarge));
}
if !header_value(&req, "content-type")
.is_some_and(|v| v.trim().eq_ignore_ascii_case("application/dns-message"))
{
return Ok(Some(DohRequest::BadRequest));
}
let mut body = buf[header_end..].to_vec();
while body.len() < len {
if buf.len() + tmp.len() > MAX_REQUEST + 1024 {
return Ok(Some(DohRequest::TooLarge));
}
let n = stream.read(&mut tmp).await?;
if n == 0 {
return Ok(Some(DohRequest::BadRequest));
}
body.extend_from_slice(&tmp[..n]);
}
body.truncate(len);
Ok(Some(DohRequest::Query(body)))
}
_ => Ok(Some(DohRequest::BadRequest)),
}
}
fn parse_get(query_str: Option<&str>) -> DohRequest {
let Some(qs) = query_str else {
return DohRequest::BadRequest;
};
let Some(dns_param) = qs.split('&').find_map(|kv| kv.strip_prefix("dns=")) else {
return DohRequest::BadRequest;
};
match base64::engine::general_purpose::URL_SAFE_NO_PAD.decode(dns_param) {
Ok(bytes) if bytes.len() <= MAX_REQUEST => DohRequest::Query(bytes),
Ok(_) => DohRequest::TooLarge,
Err(_) => DohRequest::BadRequest,
}
}
fn header_value<'a>(req: &'a httparse::Request<'_, '_>, name: &str) -> Option<&'a str> {
req.headers
.iter()
.find(|h| h.name.eq_ignore_ascii_case(name))
.and_then(|h| std::str::from_utf8(h.value).ok())
}
pub(crate) fn find_header_end(buf: &[u8]) -> Option<usize> {
buf.windows(4).position(|w| w == b"\r\n\r\n").map(|p| p + 4)
}
async fn write_dns_response(stream: &mut TcpStream, dns_msg: &[u8]) -> std::io::Result<()> {
let head = format!(
"HTTP/1.1 200 OK\r\nContent-Type: application/dns-message\r\nContent-Length: {}\r\nConnection: close\r\n\r\n",
dns_msg.len()
);
stream.write_all(head.as_bytes()).await?;
stream.write_all(dns_msg).await?;
stream.flush().await
}
pub(crate) async fn write_status(stream: &mut TcpStream, status: &str) -> std::io::Result<()> {
let head = format!("HTTP/1.1 {status}\r\nContent-Length: 0\r\nConnection: close\r\n\r\n");
stream.write_all(head.as_bytes()).await?;
stream.flush().await
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn find_header_end_locates_terminator() {
assert_eq!(find_header_end(b"GET / HTTP/1.1\r\n\r\n"), Some(18));
assert_eq!(
find_header_end(b"GET / HTTP/1.1\r\nX: 1\r\n\r\nBODY"),
Some(24)
);
assert_eq!(find_header_end(b"GET / HTTP/1.1\r\n"), None);
}
#[test]
fn parse_get_decodes_base64url_dns_param() {
let raw = [0xab, 0xcd, 0x01, 0x00];
let encoded = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(raw);
match parse_get(Some(&format!("dns={encoded}"))) {
DohRequest::Query(b) => assert_eq!(b, raw),
_ => panic!("expected Query"),
}
}
#[test]
fn parse_get_rejects_missing_or_bad_param() {
assert!(matches!(parse_get(None), DohRequest::BadRequest));
assert!(matches!(parse_get(Some("foo=bar")), DohRequest::BadRequest));
assert!(matches!(
parse_get(Some("dns=!!!notbase64!!!")),
DohRequest::BadRequest
));
}
#[test]
fn parse_response_head_returns_content_length_on_200() {
let head = b"HTTP/1.1 200 OK\r\nContent-Type: application/dns-message\r\nContent-Length: 42\r\nConnection: close\r\n\r\n";
assert_eq!(parse_response_head(head).unwrap(), 42);
}
#[test]
fn parse_response_head_rejects_non_200() {
let head = b"HTTP/1.1 404 Not Found\r\nContent-Length: 0\r\n\r\n";
assert!(parse_response_head(head).is_err());
}
#[test]
fn parse_response_head_rejects_missing_length() {
let head = b"HTTP/1.1 200 OK\r\nContent-Type: application/dns-message\r\n\r\n";
assert!(parse_response_head(head).is_err());
}
#[test]
fn parse_response_head_rejects_oversized_body() {
let head = format!(
"HTTP/1.1 200 OK\r\nContent-Length: {}\r\n\r\n",
MAX_CLIENT_RESPONSE + 1
);
assert!(parse_response_head(head.as_bytes()).is_err());
}
#[test]
fn parse_response_head_rejects_body_too_large_to_frame_over_tcp() {
let head = b"HTTP/1.1 200 OK\r\nContent-Length: 65536\r\n\r\n";
assert!(
parse_response_head(head).is_err(),
"a 65,536-byte body exceeds the DNS-over-TCP framing limit"
);
let head = b"HTTP/1.1 200 OK\r\nContent-Length: 65535\r\n\r\n";
let len = parse_response_head(head).expect("65,535 bytes is a legal DNS message");
assert_eq!(len, 65_535);
assert!(
u16::try_from(len).is_ok(),
"every accepted body must fit a two-byte length prefix"
);
}
#[test]
fn encode_formerr_sets_response_and_rcode() {
let msg = encode_formerr(0x1234);
assert_eq!(&msg[0..2], &[0x12, 0x34]);
assert_eq!(msg[2] & 0x80, 0x80, "QR response bit set");
assert_eq!(msg[3] & 0x0F, 0x01, "FORMERR rcode");
}
fn peer_node(hostname: &str, v4: &str, v6: &str) -> ts_control::Node {
use ts_control::{Node, NodeCapMap, StableNodeId, TailnetAddress};
Node {
id: 1,
stable_id: StableNodeId("n1".to_string()),
hostname: hostname.to_string(),
user_id: 0,
tailnet: Some("user.ts.net".to_string()),
tags: vec![],
addresses: vec![v4.parse().unwrap(), v6.parse().unwrap()],
tailnet_address: TailnetAddress {
ipv4: v4.parse().unwrap(),
ipv6: v6.parse().unwrap(),
},
node_key: [0u8; 32].into(),
node_key_expiry: None,
expired: false,
online: None,
last_seen: None,
key_signature: vec![],
machine_key: None,
disco_key: None,
accepted_routes: vec![],
underlay_addresses: vec![],
derp_region: None,
cap: Default::default(),
cap_map: NodeCapMap::new(),
peerapi_port: None,
peerapi_dns_proxy: false,
is_wireguard_only: false,
exit_node_dns_resolvers: vec![],
peer_relay: false,
ssh_host_keys: vec![],
service_vips: Default::default(),
unsigned_peer_api_only: false,
}
}
use ts_control::DnsConfig;
fn query_for(id: u16, labels: &[&str]) -> Vec<u8> {
let mut buf: Vec<u8> = Vec::new();
buf.extend_from_slice(&id.to_be_bytes());
buf.extend_from_slice(&0u16.to_be_bytes()); buf.extend_from_slice(&1u16.to_be_bytes()); buf.extend_from_slice(&0u16.to_be_bytes()); buf.extend_from_slice(&0u16.to_be_bytes()); buf.extend_from_slice(&0u16.to_be_bytes()); for label in labels {
buf.push(label.len() as u8);
buf.extend_from_slice(label.as_bytes());
}
buf.push(0); buf.extend_from_slice(&1u16.to_be_bytes()); buf.extend_from_slice(&1u16.to_be_bytes()); buf
}
fn rcode(resp: &[u8]) -> u8 {
resp[3] & 0x0F
}
fn view(filtered: &[&str]) -> DnsView {
DnsView {
cfg: DnsConfig {
magic_dns: true,
search_domains: vec!["user.ts.net".to_string()],
fallback_resolvers: vec![ts_control::DnsResolver {
transport: ts_control::ResolverTransport::Udp("9.9.9.9:53".parse().unwrap()),
use_with_exit_node: false,
}],
exit_node_filtered_set: filtered.iter().map(|s| s.to_string()).collect(),
..Default::default()
},
accept_dns: true,
..Default::default()
}
}
#[test]
fn filtered_name_is_refused() {
let v = view(&["blocked.example.com"]);
let q = query_for(0x1, &["blocked", "example", "com"]);
match server_decide(&v, &q, true) {
ServerDecision::Reply(resp) => assert_eq!(rcode(&resp), 5, "REFUSED"),
ServerDecision::Forward { .. } => panic!("filtered name must not forward"),
}
}
#[test]
fn recursive_query_refused_when_egress_disabled() {
let v = view(&[]);
let q = query_for(0x2, &["example", "com"]);
match server_decide(&v, &q, false) {
ServerDecision::Reply(resp) => assert_eq!(rcode(&resp), 5, "REFUSED"),
ServerDecision::Forward { .. } => panic!("must not forward when egress disabled"),
}
}
#[test]
fn recursive_query_forwards_when_egress_enabled() {
let v = view(&[]);
let q = query_for(0x3, &["example", "com"]);
match server_decide(&v, &q, true) {
ServerDecision::Forward { upstreams, .. } => {
assert_eq!(upstreams, vec!["9.9.9.9:53".parse().unwrap()]);
}
ServerDecision::Reply(_) => panic!("expected forward when egress enabled"),
}
}
#[test]
fn authoritative_answer_is_not_gated() {
let v = view(&[]);
let q = query_for(0x4, &["host", "user", "ts", "net"]);
match server_decide(&v, &q, false) {
ServerDecision::Reply(resp) => assert_eq!(rcode(&resp), 3, "NXDOMAIN, not REFUSED"),
ServerDecision::Forward { .. } => panic!("tailnet name must not forward"),
}
}
#[test]
fn a_subdomain_host_answers_here_like_any_other_peer_name() {
use std::sync::Arc;
use crate::peer_tracker::PeerDb;
let mut node = peer_node("host", "100.64.0.1/32", "fd7a::1/128");
node.cap_map
.insert("dns-subdomain-resolve".to_string(), vec![]);
let mut db = PeerDb::default();
db.upsert(&node);
let mut v = view(&[]);
v.peers = Some(Arc::new(db));
let q = query_for(0x9, &["my", "host", "user", "ts", "net"]);
match server_decide(&v, &q, false) {
ServerDecision::Reply(resp) => {
assert_eq!(rcode(&resp), 0, "NoError from the subdomain host");
assert_eq!(
u16::from_be_bytes([resp[6], resp[7]]),
1,
"one A record, the peer's own address"
);
assert_eq!(&resp[resp.len() - 4..], &[100, 64, 0, 1]);
}
ServerDecision::Forward { .. } => {
panic!("an authoritative subdomain answer must not forward")
}
}
}
#[test]
fn a_peers_doh_answer_is_never_marked_truncated_for_size() {
let query = query_for(0x5, &["example", "com"]);
let mut answer = query.clone();
answer[2] |= 0x80; answer.resize(900, 0xAB);
let out = check_response_size_and_set_tc(&query, answer.clone(), PEER_CLIENT_TRANSPORT);
assert_eq!(
out, answer,
"the DoH server relays the peer's answer byte-for-byte"
);
assert_eq!(
out[2] & 0x02,
0,
"TC is the requesting peer's call to make, not ours"
);
}
#[test]
fn unparseable_body_is_formerr() {
match server_decide(&view(&[]), &[0xAB, 0xCD, 0xFF], true) {
ServerDecision::Reply(resp) => {
assert_eq!(&resp[0..2], &[0xAB, 0xCD]);
assert_eq!(rcode(&resp), 1, "FORMERR");
}
ServerDecision::Forward { .. } => panic!("garbage must not forward"),
}
}
fn tcp_rule(
src_pfx: &str,
dst_pfx: &str,
ports: std::ops::RangeInclusive<u16>,
) -> ts_packetfilter::Rule {
ts_packetfilter::Rule {
src: ts_packetfilter::SrcMatch {
pfxs: vec![src_pfx.parse().unwrap()],
caps: vec![],
},
protos: vec![ts_packetfilter::IpProto::TCP],
dst: vec![ts_packetfilter::DstMatch {
ports,
ips: vec![dst_pfx.parse().unwrap()],
}],
}
}
fn filter_of(
rules: Vec<ts_packetfilter::Rule>,
) -> Arc<dyn ts_packetfilter::Filter + Send + Sync> {
let mut f = ts_packetfilter::HashbrownFilter::new();
f.insert("acl".to_string(), rules);
Arc::new(f)
}
fn view_with_peer() -> DnsView {
let mut db = crate::peer_tracker::PeerDb::default();
db.upsert(&peer_node("host", "100.64.0.1/32", "fd7a::1/128"));
let mut v = view(&[]);
v.peers = Some(Arc::new(db));
v
}
enum Served {
Forbidden,
Dns(ServerDecision),
}
fn doh_request_path(
view: &DnsView,
filter: Option<&(dyn ts_packetfilter::Filter + Send + Sync)>,
src: std::net::IpAddr,
query: &[u8],
forward_exit_egress: bool,
) -> Served {
if !dns_source_allowed(view, filter, src) {
return Served::Forbidden;
}
Served::Dns(server_decide(view, query, forward_exit_egress))
}
#[test]
fn probe_destination_is_off_tailnet_per_family() {
assert_eq!(
internet_probe_dst("100.64.0.1".parse().unwrap()),
"0.0.0.0".parse::<std::net::IpAddr>().unwrap()
);
assert_eq!(
internet_probe_dst("fd7a::1".parse().unwrap()),
"2000::".parse::<std::net::IpAddr>().unwrap()
);
}
#[test]
fn a_peer_the_acl_grants_internet_is_answered() {
let v = view_with_peer();
let filter = filter_of(vec![tcp_rule("100.64.0.1/32", "0.0.0.0/0", 53..=53)]);
let src: std::net::IpAddr = "100.64.0.1".parse().unwrap();
assert!(dns_source_allowed(&v, Some(&*filter), src));
let q = query_for(0x10, &["example", "com"]);
match doh_request_path(&v, Some(&*filter), src, &q, true) {
Served::Dns(ServerDecision::Forward { .. }) => {}
Served::Dns(ServerDecision::Reply(_)) => panic!("expected the forward it always got"),
Served::Forbidden => panic!("an ACL-granted peer must still get its answer"),
}
}
#[test]
fn a_peer_the_acl_denies_gets_403_and_no_resolution() {
let v = view_with_peer();
let filter = filter_of(vec![tcp_rule(
"100.64.0.1/32",
"100.64.0.0/10",
0..=u16::MAX,
)]);
let src: std::net::IpAddr = "100.64.0.1".parse().unwrap();
assert!(!dns_source_allowed(&v, Some(&*filter), src));
let q = query_for(0x11, &["example", "com"]);
assert!(
matches!(server_decide(&v, &q, true), ServerDecision::Forward { .. }),
"the name itself is one this node would have resolved"
);
assert!(matches!(
doh_request_path(&v, Some(&*filter), src, &q, true),
Served::Forbidden
));
let tailnet_q = query_for(0x12, &["host", "user", "ts", "net"]);
assert!(matches!(
doh_request_path(&v, Some(&*filter), src, &tailnet_q, true),
Served::Forbidden
));
}
#[test]
fn a_source_that_is_no_known_peer_is_refused() {
let v = view_with_peer();
let filter = filter_of(vec![tcp_rule("0.0.0.0/0", "0.0.0.0/0", 0..=u16::MAX)]);
assert!(!dns_source_allowed(
&v,
Some(&*filter),
"198.51.100.7".parse().unwrap()
));
}
#[test]
fn no_compiled_filter_yet_refuses() {
let v = view_with_peer();
assert!(!dns_source_allowed(&v, None, "100.64.0.1".parse().unwrap()));
}
#[test]
fn an_ipv6_peer_is_checked_against_the_global_unicast_probe() {
let v = view_with_peer();
let src: std::net::IpAddr = "fd7a::1".parse().unwrap();
let internet = filter_of(vec![tcp_rule("fd7a::1/128", "2000::/3", 53..=53)]);
assert!(dns_source_allowed(&v, Some(&*internet), src));
let tailnet_only = filter_of(vec![tcp_rule("fd7a::1/128", "fd7a::/48", 0..=u16::MAX)]);
assert!(!dns_source_allowed(&v, Some(&*tailnet_only), src));
}
#[test]
fn a_rule_on_another_port_does_not_open_dns() {
let v = view_with_peer();
let filter = filter_of(vec![tcp_rule("100.64.0.1/32", "0.0.0.0/0", 443..=443)]);
assert!(!dns_source_allowed(
&v,
Some(&*filter),
"100.64.0.1".parse().unwrap()
));
}
use kameo::actor::Spawn;
fn forwarder_cfg() -> crate::env::ForwarderConfig {
crate::env::ForwarderConfig {
accept_routes: false,
accept_dns: true,
exit_node: None,
forward_routes: vec![],
forward_tcp_ports: vec![],
forward_udp_ports: vec![],
forward_all_ports: false,
forward_exit_egress: false,
block_incoming: false,
exit_proxy: None,
peerapi_port: None,
taildrop_dir: None,
enable_ipv6: false,
wireguard_listen_port: None,
network_monitor: false,
persistent_keepalive_interval: None,
ingress_active: Arc::new(std::sync::atomic::AtomicBool::new(false)),
}
}
fn netmap_with(rules: Vec<ts_packetfilter::Rule>) -> Arc<ts_control::StateUpdate> {
Arc::new(ts_control::StateUpdate {
session_handle: None,
seq: 0,
keep_alive: false,
derp: None,
node: None,
peer_update: None,
peer_patches: Vec::new(),
user_profiles: Vec::new(),
ping: None,
packetfilter: Some((Some(rules), Default::default())),
cap_grants: None,
pop_browser_url: None,
dial_plan: None,
dns_config: None,
ssh_policy: None,
tka: None,
online_change: Default::default(),
peer_seen_change: Default::default(),
control_time: None,
})
}
fn updater_with_cell() -> (
kameo::actor::ActorRef<crate::packetfilter::PacketfilterUpdater>,
crate::packetfilter::LiveFilterRx,
watch::Sender<Option<crate::packetfilter::PacketFilterState>>,
crate::env::Env,
) {
let (_shutdown_tx, shutdown_rx) = watch::channel(false);
let env =
crate::env::Env::new(ts_keys::NodeState::generate(), shutdown_rx, forwarder_cfg());
let (cap_grants_tx, _cap_grants_rx) = watch::channel(Default::default());
let (filter_tx, filter_rx) = watch::channel(None);
let updater = crate::packetfilter::PacketfilterUpdater::spawn((
env.clone(),
cap_grants_tx,
filter_tx.clone(),
));
(updater, filter_rx, filter_tx, env)
}
fn gate_says_yes(
view: &DnsView,
filter_rx: &crate::packetfilter::LiveFilterRx,
src: &str,
) -> bool {
let filter = filter_rx.borrow().clone();
dns_source_allowed(view, filter.as_ref().map(|f| &*f.0), src.parse().unwrap())
}
#[tokio::test]
async fn a_gate_that_attaches_after_the_first_netmap_still_reads_the_filter() {
let (updater, filter_rx, filter_tx, _env) = updater_with_cell();
let v = view_with_peer();
assert!(
!gate_says_yes(&v, &filter_rx, "100.64.0.1"),
"no compiled filter yet must refuse"
);
updater
.tell(netmap_with(vec![tcp_rule(
"100.64.0.1/32",
"0.0.0.0/0",
53..=53,
)]))
.await
.expect("netmap delivered to the packet-filter updater");
wait_for_filter(&filter_rx).await;
assert!(
gate_says_yes(&v, &filter_rx, "100.64.0.1"),
"the receiver held since before the updater spawned must see the compiled filter"
);
let late_rx = filter_tx.subscribe();
assert!(
gate_says_yes(&v, &late_rx, "100.64.0.1"),
"a gate attaching after the compile must still see the filter"
);
}
#[tokio::test]
async fn a_revoked_grant_reaches_the_gate_too() {
let (updater, mut filter_rx, _filter_tx, _env) = updater_with_cell();
let v = view_with_peer();
updater
.tell(netmap_with(vec![tcp_rule(
"100.64.0.1/32",
"0.0.0.0/0",
53..=53,
)]))
.await
.expect("first netmap delivered");
wait_for_filter(&filter_rx).await;
assert!(gate_says_yes(&v, &filter_rx, "100.64.0.1"));
filter_rx.mark_unchanged();
updater
.tell(netmap_with(vec![tcp_rule(
"100.64.0.1/32",
"100.64.0.0/10",
0..=u16::MAX,
)]))
.await
.expect("second netmap delivered");
filter_rx
.changed()
.await
.expect("the revocation reaches the cell");
assert!(
!gate_says_yes(&v, &filter_rx, "100.64.0.1"),
"a revoked grant must reach the gate, not leave it admitting on a stale filter"
);
}
#[tokio::test]
async fn a_parked_bus_subscriber_goes_stale_while_the_cell_does_not() {
let (updater, mut filter_rx, _filter_tx, env) = updater_with_cell();
let v = view_with_peer();
let seen = Arc::new(std::sync::atomic::AtomicUsize::new(0));
let last_seen: TapState = Arc::new(std::sync::Mutex::new(None));
let (release_tx, release_rx) = watch::channel(false);
let tap = ParkedFilterTap::spawn_with_mailbox(
(seen.clone(), last_seen.clone(), release_rx),
kameo::mailbox::bounded(1),
);
env.subscribe::<crate::packetfilter::PacketFilterState>(&tap)
.await
.expect("tap subscribed to the filter bus");
for rules in [
vec![tcp_rule("100.64.0.1/32", "100.64.0.0/10", 0..=u16::MAX)],
vec![tcp_rule("100.64.0.1/32", "0.0.0.0/0", 443..=443)],
vec![tcp_rule("100.64.0.1/32", "0.0.0.0/0", 53..=53)],
] {
filter_rx.mark_unchanged();
updater
.tell(netmap_with(rules))
.await
.expect("netmap delivered");
filter_rx
.changed()
.await
.expect("each compiled filter reaches the cell");
}
wait_for_count(&seen, 1).await;
let delivered = seen.load(std::sync::atomic::Ordering::SeqCst);
assert_eq!(
delivered, 1,
"the tap must be parked on the first policy for the contrast below to mean anything"
);
assert!(
gate_says_yes(&v, &filter_rx, "100.64.0.1"),
"the gate answers from the newest compiled filter, not from what the bus managed to deliver"
);
let stale = last_seen.lock().unwrap().clone();
assert!(
!dns_source_allowed(
&v,
stale.as_ref().map(|f| &*f.0),
"100.64.0.1".parse().unwrap()
),
"the parked subscriber is supposed to be stuck on an older policy here"
);
release_tx.send_replace(true);
}
async fn wait_for_count(counter: &std::sync::atomic::AtomicUsize, want: usize) {
for _ in 0..200 {
if counter.load(std::sync::atomic::Ordering::SeqCst) >= want {
return;
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
panic!(
"timed out waiting for the bus tap to receive {want}: got {}",
counter.load(std::sync::atomic::Ordering::SeqCst)
);
}
async fn wait_for_filter(filter_rx: &crate::packetfilter::LiveFilterRx) {
for _ in 0..200 {
if filter_rx.borrow().is_some() {
return;
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
panic!("timed out waiting for the packet-filter updater to write the live filter cell");
}
type TapState = Arc<std::sync::Mutex<Option<crate::packetfilter::PacketFilterState>>>;
struct ParkedFilterTap {
seen: Arc<std::sync::atomic::AtomicUsize>,
last_seen: TapState,
release: watch::Receiver<bool>,
}
impl kameo::Actor for ParkedFilterTap {
type Args = (
Arc<std::sync::atomic::AtomicUsize>,
TapState,
watch::Receiver<bool>,
);
type Error = crate::Error;
async fn on_start(
(seen, last_seen, release): Self::Args,
_slf: kameo::actor::ActorRef<Self>,
) -> Result<Self, Self::Error> {
Ok(Self {
seen,
last_seen,
release,
})
}
}
impl kameo::message::Message<crate::packetfilter::PacketFilterState> for ParkedFilterTap {
type Reply = ();
async fn handle(
&mut self,
msg: crate::packetfilter::PacketFilterState,
_ctx: &mut kameo::message::Context<Self, Self::Reply>,
) {
{
*self.last_seen.lock().unwrap() = Some(msg);
}
self.seen.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
loop {
let released = *self.release.borrow_and_update();
if released || self.release.changed().await.is_err() {
return;
}
}
}
}
}