use std::sync::Arc;
use netstack::{netcore::Channel, netsock::TcpListener};
use tokio::{
io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt},
sync::{Semaphore, watch},
time::{Duration, timeout},
};
use crate::magic_dns::{
ClientTransport, Decision, DnsView, RecursivePlan, check_response_size_and_set_tc, decide,
forward_plan, forward_query,
};
async fn answer_query(view: &DnsView, channel: &Channel, query: &[u8]) -> Option<Vec<u8>> {
match decide(view, query)? {
Decision::Reply(response) => Some(check_response_size_and_set_tc(
query,
response,
ClientTransport::Tcp,
)),
Decision::Forward {
upstreams,
query,
servfail,
recursive,
} => Some(match forward_plan(view, upstreams, recursive) {
RecursivePlan::Udp(upstreams) => {
forward_query(channel, &upstreams, &query, servfail, ClientTransport::Tcp).await
}
RecursivePlan::Doh(doh_addr) => {
crate::peerapi_doh::forward_doh(
channel,
doh_addr,
&query,
servfail,
ClientTransport::Tcp,
)
.await
}
}),
}
}
const IDLE_TIMEOUT: Duration = Duration::from_secs(30);
const MESSAGE_TIMEOUT: Duration = Duration::from_secs(5);
const MAX_INFLIGHT_CONNS: usize = 64;
pub(crate) async fn serve(
listener: TcpListener,
view_rx: watch::Receiver<Arc<DnsView>>,
forward_channel: Channel,
) {
tracing::debug!(addr = %listener.local_addr(), "magic dns tcp accepting");
let inflight = Arc::new(Semaphore::new(MAX_INFLIGHT_CONNS));
loop {
let stream = match listener.accept().await {
Ok(s) => s,
Err(e) => {
tracing::warn!(error = %e, "magic dns tcp accept failed, stopping server");
return;
}
};
let Ok(permit) = inflight.clone().try_acquire_owned() else {
tracing::warn!(
client = %stream.remote_addr(),
"magic dns tcp drop: at max in-flight connections ({MAX_INFLIGHT_CONNS})"
);
continue;
};
let view_rx = view_rx.clone();
let forward_channel = forward_channel.clone();
tokio::spawn(async move {
let _permit = permit;
serve_conn(stream, view_rx, forward_channel).await;
});
}
}
async fn serve_conn<S>(mut stream: S, view_rx: watch::Receiver<Arc<DnsView>>, channel: Channel)
where
S: AsyncRead + AsyncWrite + Unpin,
{
loop {
let mut len_buf = [0u8; 2];
match timeout(IDLE_TIMEOUT, stream.read_exact(&mut len_buf)).await {
Ok(Ok(_)) => {}
Ok(Err(_)) => return,
Err(_) => {
tracing::debug!("magic dns tcp: idle connection closed");
return;
}
}
let len = usize::from(u16::from_be_bytes(len_buf));
if len == 0 {
tracing::debug!("magic dns tcp: zero-length message; closing");
return;
}
let mut query = vec![0u8; len];
match timeout(MESSAGE_TIMEOUT, stream.read_exact(&mut query)).await {
Ok(Ok(_)) => {}
Ok(Err(_)) => return,
Err(_) => {
tracing::debug!(len, "magic dns tcp: message body stalled; closing");
return;
}
}
let view = view_rx.borrow().clone();
let Some(response) = answer_query(&view, &channel, &query).await else {
tracing::debug!("magic dns tcp: malformed query; closing");
return;
};
let Ok(response_len) = u16::try_from(response.len()) else {
tracing::warn!(
len = response.len(),
"magic dns tcp: response too long to frame"
);
return;
};
let mut framed = Vec::with_capacity(2 + response.len());
framed.extend_from_slice(&response_len.to_be_bytes());
framed.extend_from_slice(&response);
if stream.write_all(&framed).await.is_err() || stream.flush().await.is_err() {
return;
}
}
}
#[cfg(test)]
mod tests {
use netstack::HasChannel;
use ts_control::{DnsConfig, Node, StableNodeId, TailnetAddress};
use super::*;
use crate::peer_tracker::PeerDb;
fn peer_node() -> Node {
Node {
id: 1,
user_id: 0,
stable_id: StableNodeId("n1".to_string()),
hostname: "host".to_string(),
tailnet: Some("user.ts.net".to_string()),
tags: vec![],
addresses: vec![
"100.64.0.1/32".parse().unwrap(),
"fd7a::1/128".parse().unwrap(),
],
tailnet_address: TailnetAddress {
ipv4: "100.64.0.1/32".parse().unwrap(),
ipv6: "fd7a::1/128".parse().unwrap(),
},
node_key: [1u8; 32].into(),
node_key_expiry: None,
key_signature: vec![],
machine_key: None,
disco_key: None,
accepted_routes: vec![],
underlay_addresses: vec![],
derp_region: None,
cap: Default::default(),
cap_map: Default::default(),
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,
online: None,
last_seen: None,
}
}
fn forwardless_view() -> DnsView {
let mut db = PeerDb::default();
db.upsert(&peer_node());
DnsView {
cfg: DnsConfig {
magic_dns: true,
search_domains: vec!["user.ts.net".to_string()],
..Default::default()
},
peers: Some(Arc::new(db)),
self_node: None,
exit_doh: None,
enable_ipv6: false,
accept_dns: true,
}
}
fn a_query(id: u16, name: &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 name.split('.') {
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 response_rcode(resp: &[u8]) -> u8 {
resp[3] & 0x0F
}
fn framed(msg: &[u8]) -> Vec<u8> {
let len = u16::try_from(msg.len()).expect("test message fits a length prefix");
let mut out = len.to_be_bytes().to_vec();
out.extend_from_slice(msg);
out
}
async fn framed_exchange(view: DnsView, queries: &[Vec<u8>]) -> Vec<Vec<u8>> {
let (netstack, _pipe) = netstack::piped(netstack::netcore::Config::default());
let channel = netstack.command_channel();
let (_view_tx, view_rx) = watch::channel(Arc::new(view));
let (client, server) = tokio::io::duplex(64 * 1024);
let served = tokio::spawn(serve_conn(server, view_rx, channel));
let (mut read_half, mut write_half) = tokio::io::split(client);
for q in queries {
write_half.write_all(q).await.expect("write query");
}
write_half.shutdown().await.expect("half-close");
let mut out = Vec::new();
loop {
let mut len = [0u8; 2];
if read_half.read_exact(&mut len).await.is_err() {
break;
}
let mut body = vec![0u8; usize::from(u16::from_be_bytes(len))];
if read_half.read_exact(&mut body).await.is_err() {
break;
}
out.push(body);
}
served.await.expect("serve_conn must not panic");
out
}
#[tokio::test]
async fn length_prefixed_query_is_answered() {
let query = a_query(0x1234, "host.user.ts.net");
let replies = framed_exchange(forwardless_view(), &[framed(&query)]).await;
assert_eq!(replies.len(), 1, "exactly one framed answer");
let reply = &replies[0];
assert_eq!(
reply[0..2],
query[0..2],
"the answer echoes the query's transaction id"
);
assert_eq!(
response_rcode(reply),
0,
"an in-tailnet name resolves NOERROR over TCP just as it does over UDP"
);
assert_eq!(
reply[2] & 0x02,
0,
"a TCP answer is never marked truncated: the client already did the TCP retry"
);
assert_eq!(
u16::from_be_bytes([reply[6], reply[7]]),
1,
"the answer section carries the peer's address"
);
assert_eq!(
&reply[reply.len() - 4..],
&[100, 64, 0, 1],
"and it is the peer's tailnet IPv4"
);
}
#[tokio::test]
async fn connection_serves_more_than_one_query() {
let first = a_query(0x0001, "host.user.ts.net");
let second = a_query(0x0002, "host.user.ts.net");
let replies = framed_exchange(forwardless_view(), &[framed(&first), framed(&second)]).await;
assert_eq!(replies.len(), 2, "both queries answered on one connection");
assert_eq!(replies[0][0..2], first[0..2], "first answer, first");
assert_eq!(replies[1][0..2], second[0..2], "second answer, second");
}
#[tokio::test]
async fn malformed_query_closes_without_answering() {
let replies = framed_exchange(forwardless_view(), &[framed(&[0u8; 3])]).await;
assert!(
replies.is_empty(),
"a malformed query is answered with nothing at all"
);
}
#[tokio::test]
async fn zero_length_message_closes_the_connection() {
let good = a_query(0x0003, "host.user.ts.net");
let replies = framed_exchange(forwardless_view(), &[vec![0, 0], framed(&good)]).await;
assert!(
replies.is_empty(),
"the zero-length prefix ends the connection before the following query is read"
);
}
#[tokio::test]
async fn accept_dns_off_refuses_over_tcp_too() {
let mut view = forwardless_view();
view.accept_dns = false;
let replies = framed_exchange(view, &[framed(&a_query(0x0004, "host.user.ts.net"))]).await;
assert_eq!(replies.len(), 1, "the refusal is still an answer");
assert_eq!(response_rcode(&replies[0]), 5, "REFUSED");
}
}