use std::net::SocketAddr;
use std::path::PathBuf;
use std::sync::{Arc, Weak};
use std::time::Duration;
use base64::{Engine as _, engine::general_purpose::STANDARD};
use tokio::io::{AsyncReadExt as _, AsyncWriteExt as _};
use tokio::net::{TcpListener, TcpStream};
use tokio::sync::OnceCell;
use tokio::task::AbortHandle;
use ts_keys::PersistState;
use crate::config::Config;
use crate::netstack;
use crate::{Device, RegistrationError, ServiceError, ServiceMode, Status, StatusNode};
pub const STATE_FILE: &str = "tailscale-rs.state";
pub const STATE_KEY: &str = "_tailscale-rs/persist";
pub trait StateStore: Send + Sync {
fn read_state(&self, id: &str) -> std::io::Result<Option<Vec<u8>>>;
fn write_state(&self, id: &str, value: &[u8]) -> std::io::Result<()>;
}
pub struct FileStore {
path: PathBuf,
}
impl FileStore {
pub fn new(dir: impl Into<PathBuf>) -> Self {
Self {
path: dir.into().join(STATE_FILE),
}
}
pub fn at(path: impl Into<PathBuf>) -> Self {
Self { path: path.into() }
}
}
impl StateStore for FileStore {
fn read_state(&self, _id: &str) -> std::io::Result<Option<Vec<u8>>> {
match std::fs::read(&self.path) {
Ok(bytes) => Ok(Some(bytes)),
Err(e) if e.kind() == std::io::ErrorKind::NotFound => Ok(None),
Err(e) => Err(e),
}
}
fn write_state(&self, _id: &str, value: &[u8]) -> std::io::Result<()> {
if let Some(parent) = self.path.parent() {
std::fs::create_dir_all(parent)?;
}
std::fs::write(&self.path, value)
}
}
#[derive(Default)]
pub struct MemStore {
inner: std::sync::Mutex<std::collections::HashMap<String, Vec<u8>>>,
}
impl StateStore for MemStore {
fn read_state(&self, id: &str) -> std::io::Result<Option<Vec<u8>>> {
Ok(self
.inner
.lock()
.unwrap_or_else(|e| e.into_inner())
.get(id)
.cloned())
}
fn write_state(&self, id: &str, value: &[u8]) -> std::io::Result<()> {
self.inner
.lock()
.unwrap_or_else(|e| e.into_inner())
.insert(id.to_string(), value.to_vec());
Ok(())
}
}
#[derive(Debug, thiserror::Error)]
#[non_exhaustive]
pub enum Error {
#[error("device error: {0}")]
Device(#[from] crate::Error),
#[error("registration error: {0}")]
Registration(#[from] RegistrationError),
#[error("invalid control URL: {0}")]
InvalidControlUrl(url::ParseError),
#[error("invalid address {addr:?}: {source}")]
InvalidAddr {
addr: String,
source: std::net::AddrParseError,
},
#[error(
"address {addr:?} is not an {want} address (required by the family-pinned listen network)"
)]
AddrFamilyMismatch {
addr: String,
want: &'static str,
},
#[error("unsupported network {network:?} (want tcp, tcp4, tcp6, udp, udp4, or udp6)")]
UnsupportedNetwork {
network: String,
},
#[error("invalid or unsupported network {network:?}")]
InvalidNetwork {
network: String,
},
#[error("state store I/O error: {0}")]
Store(std::io::Error),
#[error("node identity (de)serialization error: {0}")]
State(serde_json::Error),
#[error("logout error: {0}")]
Logout(#[from] crate::LogoutError),
#[error("loopback I/O error: {0}")]
Loopback(std::io::Error),
}
#[derive(Default)]
pub struct FunnelOptions {
pub funnel_only: bool,
pub tls: Option<crate::TlsAcceptor>,
}
impl FunnelOptions {
pub fn funnel_only() -> Self {
Self {
funnel_only: true,
..Default::default()
}
}
pub fn with_tls(mut self, acceptor: crate::TlsAcceptor) -> Self {
self.tls = Some(acceptor);
self
}
}
impl From<FunnelOptions> for ts_control::FunnelOptions {
fn from(o: FunnelOptions) -> Self {
if o.tls.is_some() {
tracing::warn!(
"tsnet::FunnelOptions::tls (FunnelTLSConfig) is set but the engine terminates Funnel \
with the node's own certificate; the supplied acceptor is ignored (see design doc)"
);
}
ts_control::FunnelOptions {
funnel_only: o.funnel_only,
}
}
}
#[derive(Debug, thiserror::Error)]
#[non_exhaustive]
pub enum ListenFunnelError {
#[error("failed to start the node before listening on funnel: {0}")]
Start(#[from] Error),
#[error(transparent)]
Funnel(#[from] ts_control::FunnelError),
}
#[derive(Debug, thiserror::Error)]
#[non_exhaustive]
pub enum ListenServiceError {
#[error("failed to start the node before listening on service: {0}")]
Start(#[from] Error),
#[error(transparent)]
Service(#[from] ServiceError),
}
pub struct ServiceListener {
inner: netstack::TcpListener,
fqdn: String,
}
impl ServiceListener {
pub fn fqdn(&self) -> &str {
&self.fqdn
}
pub fn into_inner(self) -> netstack::TcpListener {
self.inner
}
}
impl std::ops::Deref for ServiceListener {
type Target = netstack::TcpListener;
fn deref(&self) -> &Self::Target {
&self.inner
}
}
#[derive(Clone)]
#[non_exhaustive]
pub struct Loopback {
pub address: SocketAddr,
pub proxy_cred: String,
pub local_api_address: SocketAddr,
pub local_api_cred: String,
}
impl std::fmt::Debug for Loopback {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Loopback")
.field("address", &self.address)
.field("proxy_cred", &"<redacted>")
.field("local_api_address", &self.local_api_address)
.field("local_api_cred", &"<redacted>")
.finish()
}
}
#[derive(Clone)]
pub struct LocalClient {
address: SocketAddr,
cred: String,
}
impl std::fmt::Debug for LocalClient {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("LocalClient")
.field("address", &self.address)
.field("cred", &"<redacted>")
.finish()
}
}
impl LocalClient {
pub fn address(&self) -> SocketAddr {
self.address
}
pub fn credential(&self) -> &str {
&self.cred
}
pub async fn status(&self) -> Result<Vec<u8>, Error> {
match self.get("/localapi/v0/status").await? {
(200, body) => Ok(body),
(code, _) => Err(Error::Loopback(std::io::Error::other(format!(
"localapi /status returned HTTP {code}"
)))),
}
}
pub async fn get(&self, path: &str) -> Result<(u16, Vec<u8>), Error> {
localapi_client_get(self.address, &self.cred, path)
.await
.map_err(Error::Loopback)
}
}
#[derive(Clone, Debug, Default)]
pub struct TunSpec {
pub name: Option<String>,
pub mtu: Option<u16>,
}
pub type ConfigureHook = Box<dyn Fn(&mut Config) + Send + Sync>;
#[derive(Default)]
pub struct Server {
pub hostname: Option<String>,
pub auth_key: Option<String>,
pub control_url: Option<String>,
pub ephemeral: bool,
pub advertise_tags: Vec<String>,
pub port: Option<u16>,
pub run_web_client: bool,
pub client_id: Option<String>,
pub client_secret: Option<String>,
pub id_token: Option<String>,
pub audience: Option<String>,
pub dir: Option<PathBuf>,
pub store: Option<Arc<dyn StateStore>>,
pub tun: Option<TunSpec>,
configure: Option<ConfigureHook>,
device: OnceCell<Arc<Device>>,
loopback_rt: OnceCell<LoopbackRt>,
}
impl Server {
pub fn new() -> Self {
Self::default()
}
pub fn configure<F>(&mut self, f: F) -> &mut Self
where
F: Fn(&mut Config) + Send + Sync + 'static,
{
self.configure = Some(Box::new(f));
self
}
async fn resolve_key_state(&self) -> Result<Option<PersistState>, Error> {
if let Some(store) = &self.store {
return match store.read_state(STATE_KEY).map_err(Error::Store)? {
Some(bytes) => Ok(Some(serde_json::from_slice(&bytes).map_err(Error::State)?)),
None => {
let fresh = PersistState::default();
let bytes = serde_json::to_vec(&fresh).map_err(Error::State)?;
store.write_state(STATE_KEY, &bytes).map_err(Error::Store)?;
Ok(Some(fresh))
}
};
}
Ok(None)
}
async fn build_config(&self) -> Result<Config, Error> {
let mut config = match (&self.store, &self.dir) {
(None, Some(dir)) => Config::default_with_key_file(dir.join(STATE_FILE)).await?,
_ => {
let mut c = Config::default();
if let Some(ks) = self.resolve_key_state().await? {
c.key_state = ks;
}
c
}
};
config.ephemeral = self.ephemeral;
config.requested_hostname = self.hostname.clone();
config.requested_tags = self.advertise_tags.clone();
config.wireguard_listen_port = self.port;
config.run_web_client = self.run_web_client;
config.auth_key = self.auth_key.clone();
config.client_id = self.client_id.clone();
config.client_secret = self.client_secret.clone();
config.id_token = self.id_token.clone();
config.audience = self.audience.clone();
if let Some(raw) = &self.control_url {
config.control_server_url = raw.parse().map_err(Error::InvalidControlUrl)?;
}
if let Some(tun) = &self.tun {
config = config.use_tun(tun.name.clone(), tun.mtu);
}
if let Some(hook) = &self.configure {
hook(&mut config);
}
Ok(config)
}
async fn build_and_start(&self) -> Result<Arc<Device>, Error> {
let config = self.build_config().await?;
Ok(Arc::new(Device::new(&config, self.auth_key.clone()).await?))
}
async fn started(&self) -> Result<&Arc<Device>, Error> {
self.device.get_or_try_init(|| self.build_and_start()).await
}
pub async fn start(&self) -> Result<(), Error> {
self.started().await.map(|_| ())
}
pub async fn up(&self, timeout: Option<Duration>) -> Result<Status, Error> {
let dev = self.started().await?;
dev.wait_until_running(timeout).await?;
Ok(dev.status().await?)
}
pub async fn device(&self) -> Result<&Device, Error> {
Ok(&**self.started().await?)
}
pub async fn dial(&self, network: &str, addr: &str) -> Result<crate::DialConn, Error> {
let net = parse_network(network).map_err(|_| Error::UnsupportedNetwork {
network: network.to_string(),
})?;
let dev = self.started().await?;
Ok(match (net.transport, net.family) {
(Transport::Tcp, Family::Any) => crate::DialConn::Tcp(dev.dial_tcp(addr).await?),
(Transport::Udp, Family::Any) => crate::DialConn::Udp(dev.dial_udp(addr).await?),
_ => dev.dial(network, addr).await?,
})
}
pub async fn dial_tcp(&self, addr: &str) -> Result<netstack::TcpStream, Error> {
Ok(self.started().await?.dial_tcp(addr).await?)
}
pub async fn dial_udp(&self, addr: &str) -> Result<crate::ConnectedUdpSocket, Error> {
Ok(self.started().await?.dial_udp(addr).await?)
}
pub async fn listen(&self, network: &str, addr: &str) -> Result<netstack::TcpListener, Error> {
let net = parse_network(network)?;
if net.transport != Transport::Tcp {
return Err(Error::InvalidNetwork {
network: network.to_string(),
});
}
let sa = parse_listen_addr(addr, net.family)?;
Ok(self.started().await?.tcp_listen(sa).await?)
}
pub async fn listen_packet(
&self,
network: &str,
addr: &str,
) -> Result<netstack::UdpSocket, Error> {
let net = parse_network(network)?;
if net.transport != Transport::Udp {
return Err(Error::InvalidNetwork {
network: network.to_string(),
});
}
let addr = normalize_listen_addr(addr, net.family);
Ok(self.started().await?.listen_packet(network, &addr).await?)
}
pub async fn listen_funnel(
&self,
cfg: &crate::ServeConfig,
opts: FunnelOptions,
) -> Result<ts_runtime::funnel::FunnelAcceptedReceiver, ListenFunnelError> {
let dev = self.started().await?;
Ok(dev.listen_funnel(cfg, opts.into()).await?)
}
pub async fn listen_service(
&self,
name: &str,
mode: ServiceMode,
) -> Result<ServiceListener, ListenServiceError> {
let dev = self.started().await?;
let inner = dev.listen_service(name, mode).await?;
let fqdn = dev
.self_node()
.await
.map(|n| n.fqdn(false))
.unwrap_or_default();
Ok(ServiceListener { inner, fqdn })
}
pub async fn loopback(&self) -> Result<Loopback, Error> {
let rt = self.ensure_loopback().await?;
Ok(Loopback {
address: rt.socks_address,
proxy_cred: rt.proxy_cred.clone(),
local_api_address: rt.local_api_address,
local_api_cred: rt.local_api_cred.clone(),
})
}
pub async fn local_client(&self) -> Result<LocalClient, Error> {
let rt = self.ensure_loopback().await?;
Ok(LocalClient {
address: rt.local_api_address,
cred: rt.local_api_cred.clone(),
})
}
async fn ensure_loopback(&self) -> Result<&LoopbackRt, Error> {
let device = self.started().await?.clone();
self.loopback_rt
.get_or_try_init(|| build_loopback_rt(device))
.await
}
#[cfg(feature = "hyper")]
pub async fn http_client<B>(
&self,
) -> Result<hyper_util::client::legacy::Client<crate::http::TailnetConnector, B>, Error>
where
B: hyper::body::Body + Send + 'static,
B::Data: Send,
{
let connector = self.started().await?.http_connector().await?;
Ok(
hyper_util::client::legacy::Client::builder(hyper_util::rt::TokioExecutor::new())
.build(connector),
)
}
pub async fn tailscale_ips(
&self,
) -> Result<(std::net::Ipv4Addr, Option<std::net::Ipv6Addr>), Error> {
Ok(self.started().await?.tailscale_ips().await?)
}
pub async fn listen_tls(
&self,
cfg: &crate::ServeConfig,
) -> Result<crate::TlsAcceptor, ts_control::CertError> {
cfg.validate()?;
self.started()
.await
.map_err(start_failed_cert)?
.listen_tls(cfg)
.await
}
#[cfg(feature = "acme")]
pub async fn cert_pair(
&self,
name: &str,
min_validity: Option<Duration>,
) -> Result<(String, String), ts_control::CertError> {
if !ts_control::is_tailnet_name(name) {
return Err(ts_control::CertError::NotTailnetName(name.to_string()));
}
self.started()
.await
.map_err(start_failed_cert)?
.cert_pair(name, min_validity)
.await
}
pub async fn cert_domains(&self) -> Result<Vec<String>, Error> {
Ok(self.started().await?.cert_domains().await?)
}
pub async fn status(&self) -> Result<Status, Error> {
Ok(self.started().await?.status().await?)
}
pub async fn logout(&self) -> Result<(), Error> {
self.started().await?.logout().await?;
Ok(())
}
pub async fn close(self, timeout: Option<Duration>) -> bool {
drop(self.loopback_rt);
match self.device.into_inner() {
None => true,
Some(arc) => match Arc::into_inner(arc) {
Some(dev) => dev.shutdown(timeout).await,
None => false,
},
}
}
}
fn start_failed_cert(e: Error) -> ts_control::CertError {
ts_control::CertError::Io(std::io::Error::other(format!(
"server failed to start: {e}"
)))
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum Transport {
Tcp,
Udp,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum Family {
Any,
V4,
V6,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
struct Network {
transport: Transport,
family: Family,
}
fn parse_network(network: &str) -> Result<Network, Error> {
let (transport, family) = match network {
"tcp" => (Transport::Tcp, Family::Any),
"tcp4" => (Transport::Tcp, Family::V4),
"tcp6" => (Transport::Tcp, Family::V6),
"udp" => (Transport::Udp, Family::Any),
"udp4" => (Transport::Udp, Family::V4),
"udp6" => (Transport::Udp, Family::V6),
_ => {
return Err(Error::InvalidNetwork {
network: network.to_string(),
});
}
};
Ok(Network { transport, family })
}
fn normalize_listen_addr(addr: &str, family: Family) -> String {
if addr.starts_with(':') {
match family {
Family::V6 => format!("[::]{addr}"),
Family::Any | Family::V4 => format!("0.0.0.0{addr}"),
}
} else {
addr.to_string()
}
}
fn parse_listen_addr(addr: &str, family: Family) -> Result<SocketAddr, Error> {
let sa: SocketAddr = normalize_listen_addr(addr, family)
.parse()
.map_err(|source| Error::InvalidAddr {
addr: addr.to_string(),
source,
})?;
let mismatch = match family {
Family::Any => None,
Family::V4 => (!sa.is_ipv4()).then_some("IPv4"),
Family::V6 => (!sa.is_ipv6()).then_some("IPv6"),
};
match mismatch {
Some(want) => Err(Error::AddrFamilyMismatch {
addr: addr.to_string(),
want,
}),
None => Ok(sa),
}
}
struct LoopbackRt {
socks_address: SocketAddr,
proxy_cred: String,
local_api_address: SocketAddr,
local_api_cred: String,
_socks_handle: crate::LoopbackHandle,
localapi_task: AbortHandle,
}
impl Drop for LoopbackRt {
fn drop(&mut self) {
self.localapi_task.abort();
}
}
async fn build_loopback_rt(device: Arc<Device>) -> Result<LoopbackRt, Error> {
let (socks_address, proxy_cred, socks_handle) = device.loopback().await?;
let listener = TcpListener::bind((std::net::Ipv4Addr::LOCALHOST, 0))
.await
.map_err(Error::Loopback)?;
let local_api_address = listener.local_addr().map_err(Error::Loopback)?;
let local_api_cred = gen_cred();
let weak: Weak<Device> = Arc::downgrade(&device);
let status: localapi::StatusFn = Arc::new(move || {
let weak = weak.clone();
Box::pin(async move {
match weak.upgrade() {
Some(dev) => dev
.status()
.await
.map(|s| status_json(&s))
.map_err(|e| e.to_string()),
None => Err("device has shut down".to_string()),
}
})
});
let task = tokio::spawn(localapi::serve(listener, local_api_cred.clone(), status));
Ok(LoopbackRt {
socks_address,
proxy_cred,
local_api_address,
local_api_cred,
_socks_handle: socks_handle,
localapi_task: task.abort_handle(),
})
}
fn gen_cred() -> String {
let bytes: [u8; 16] = rand::random();
bytes.iter().map(|b| format!("{b:02x}")).collect()
}
fn status_node_json(n: &StatusNode) -> serde_json::Value {
serde_json::json!({
"stable_id": n.stable_id.0,
"display_name": n.display_name,
"ipv4": n.ipv4.to_string(),
"ipv6": n.ipv6.to_string(),
"online": n.online,
"last_seen": n.last_seen.map(|t| t.timestamp()),
"allowed_routes": n.allowed_routes.iter().map(|r| r.to_string()).collect::<Vec<_>>(),
"is_exit_node": n.is_exit_node,
"cur_addr": n.cur_addr.map(|a| a.to_string()),
"relay": n.relay,
"ssh_host_keys": n.ssh_host_keys,
})
}
fn status_json(s: &Status) -> Vec<u8> {
let value = serde_json::json!({
"self": s.self_node.as_ref().map(status_node_json),
"peers": s.peers.iter().map(status_node_json).collect::<Vec<_>>(),
"active_exit_node": s.active_exit_node.as_ref().map(|id| id.0.clone()),
"magic_dns_suffix": s.magic_dns_suffix,
});
serde_json::to_vec(&value).unwrap_or_else(|_| b"{}".to_vec())
}
fn find_subslice(hay: &[u8], needle: &[u8]) -> Option<usize> {
if needle.is_empty() || hay.len() < needle.len() {
return None;
}
hay.windows(needle.len()).position(|w| w == needle)
}
fn parse_response(resp: &[u8]) -> Option<(u16, Vec<u8>)> {
let head_end = find_subslice(resp, b"\r\n\r\n")?;
let head = std::str::from_utf8(&resp[..head_end]).ok()?;
let status_line = head.split("\r\n").next()?;
let code: u16 = status_line.split_whitespace().nth(1)?.parse().ok()?;
Some((code, resp[head_end + 4..].to_vec()))
}
async fn localapi_client_get(
addr: SocketAddr,
cred: &str,
path: &str,
) -> std::io::Result<(u16, Vec<u8>)> {
let mut sock = TcpStream::connect(addr).await?;
let auth = STANDARD.encode(format!(":{cred}"));
let req = format!(
"GET {path} HTTP/1.1\r\nHost: 127.0.0.1\r\nSec-Tailscale: localapi\r\nAuthorization: Basic {auth}\r\nConnection: close\r\n\r\n"
);
sock.write_all(req.as_bytes()).await?;
let mut resp = Vec::new();
sock.read_to_end(&mut resp).await?;
parse_response(&resp).ok_or_else(|| {
std::io::Error::new(std::io::ErrorKind::InvalidData, "malformed HTTP response")
})
}
mod localapi {
use super::{Duration, STANDARD, TcpListener, TcpStream, find_subslice};
use base64::Engine as _;
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use tokio::io::{AsyncReadExt as _, AsyncWriteExt as _};
use tokio::sync::Semaphore;
const MAX_HEAD: usize = 8 * 1024;
const REQUEST_TIMEOUT: Duration = Duration::from_secs(30);
const MAX_CONCURRENT: usize = 64;
const SEC_TAILSCALE_HEADER: &str = "Sec-Tailscale";
const SEC_TAILSCALE_VALUE: &str = "localapi";
pub(super) type StatusFn = Arc<
dyn Fn() -> Pin<Box<dyn Future<Output = Result<Vec<u8>, String>> + Send>> + Send + Sync,
>;
pub(super) async fn serve(listener: TcpListener, cred: String, status: StatusFn) {
let sem = Arc::new(Semaphore::new(MAX_CONCURRENT));
loop {
let permit = match sem.clone().acquire_owned().await {
Ok(p) => p,
Err(_) => return,
};
let (sock, _peer) = match listener.accept().await {
Ok(pair) => pair,
Err(e) => {
tracing::warn!(error = %e, "loopback LocalAPI accept failed; stopping accept loop");
return;
}
};
let cred = cred.clone();
let status = status.clone();
tokio::spawn(async move {
let _permit = permit;
match tokio::time::timeout(REQUEST_TIMEOUT, handle(sock, &cred, &status)).await {
Ok(Ok(())) => {}
Ok(Err(e)) => tracing::debug!(error = %e, "loopback LocalAPI connection ended"),
Err(_) => tracing::debug!("loopback LocalAPI request timed out"),
}
});
}
}
async fn handle(mut sock: TcpStream, cred: &str, status: &StatusFn) -> std::io::Result<()> {
let mut buf = Vec::with_capacity(1024);
let mut chunk = [0u8; 1024];
let head_len = loop {
if let Some(pos) = find_subslice(&buf, b"\r\n\r\n") {
break pos;
}
if buf.len() > MAX_HEAD {
let r = response(
431,
"Request Header Fields Too Large",
"text/plain",
b"header too large",
&[],
);
sock.write_all(&r).await?;
return Ok(());
}
let n = sock.read(&mut chunk).await?;
if n == 0 {
return Ok(()); }
buf.extend_from_slice(&chunk[..n]);
};
let Some((method, target, password, sec_tailscale)) = parse_head(&buf[..head_len]) else {
let r = response(400, "Bad Request", "text/plain", b"bad request", &[]);
sock.write_all(&r).await?;
return Ok(());
};
if sec_tailscale.as_deref() != Some(SEC_TAILSCALE_VALUE) {
let r = response(
403,
"Forbidden",
"text/plain",
b"missing 'Sec-Tailscale: localapi' header",
&[],
);
sock.write_all(&r).await?;
return Ok(());
}
if !password.as_deref().is_some_and(|p| cred_ok(p, cred)) {
let r = response(
401,
"Unauthorized",
"text/plain",
b"unauthorized",
&[("WWW-Authenticate", "Basic realm=\"tailscale localapi\"")],
);
sock.write_all(&r).await?;
return Ok(());
}
let path = target.split('?').next().unwrap_or(&target);
let resp = match (method.as_str(), path) {
("GET", "/localapi/v0/status") => match status().await {
Ok(body) => response(200, "OK", "application/json", &body, &[]),
Err(_) => response(
500,
"Internal Server Error",
"text/plain",
b"status error",
&[],
),
},
_ => response(404, "Not Found", "text/plain", b"not found", &[]),
};
sock.write_all(&resp).await?;
Ok(())
}
pub(super) fn parse_head(
head: &[u8],
) -> Option<(String, String, Option<String>, Option<String>)> {
let text = std::str::from_utf8(head).ok()?;
let mut lines = text.split("\r\n");
let mut request_line = lines.next()?.split(' ');
let method = request_line.next()?.to_string();
let target = request_line.next()?.to_string();
request_line.next()?; let mut password = None;
let mut sec_tailscale = None;
for line in lines {
let Some((name, value)) = line.split_once(':') else {
continue;
};
let name = name.trim();
if name.eq_ignore_ascii_case("authorization") {
password = basic_auth_password(value.trim());
} else if name.eq_ignore_ascii_case(SEC_TAILSCALE_HEADER) {
sec_tailscale = Some(value.trim().to_string());
}
}
Some((method, target, password, sec_tailscale))
}
pub(super) fn basic_auth_password(value: &str) -> Option<String> {
let (scheme, b64) = value.split_once(' ')?;
if !scheme.eq_ignore_ascii_case("basic") {
return None;
}
let decoded = STANDARD.decode(b64.trim()).ok()?;
let decoded = String::from_utf8(decoded).ok()?;
decoded
.split_once(':')
.map(|(_user, pass)| pass.to_string())
}
pub(super) fn cred_ok(provided: &str, expected: &str) -> bool {
let (a, b) = (provided.as_bytes(), expected.as_bytes());
if a.len() != b.len() {
return false;
}
let mut diff = 0u8;
for (x, y) in a.iter().zip(b.iter()) {
diff |= x ^ y;
}
diff == 0
}
pub(super) fn response(
code: u16,
reason: &str,
content_type: &str,
body: &[u8],
extra_headers: &[(&str, &str)],
) -> Vec<u8> {
let mut head = format!(
"HTTP/1.1 {code} {reason}\r\nContent-Type: {content_type}\r\nContent-Length: {}\r\nConnection: close\r\n",
body.len()
);
for (name, value) in extra_headers {
head.push_str(name);
head.push_str(": ");
head.push_str(value);
head.push_str("\r\n");
}
head.push_str("\r\n");
let mut out = head.into_bytes();
out.extend_from_slice(body);
out
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn parse_network_accepts_the_tsnet_set() {
assert_eq!(
parse_network("tcp").unwrap(),
Network {
transport: Transport::Tcp,
family: Family::Any
}
);
assert_eq!(parse_network("tcp4").unwrap().family, Family::V4);
assert_eq!(parse_network("tcp6").unwrap().family, Family::V6);
assert_eq!(parse_network("udp").unwrap().transport, Transport::Udp);
assert_eq!(parse_network("udp4").unwrap().family, Family::V4);
assert_eq!(parse_network("udp6").unwrap().family, Family::V6);
}
#[test]
fn parse_network_rejects_unknown_strings() {
for n in ["", "tcp5", "sctp", "unix", "TCP", "udp7", "ip", "tcp ", "0"] {
assert!(
matches!(parse_network(n), Err(Error::InvalidNetwork { network }) if network == n),
"network {n:?} must be rejected as InvalidNetwork carrying the offending value"
);
}
}
#[test]
fn parse_colon_port_is_family_aware_wildcard() {
let v4 = parse_listen_addr(":8080", Family::Any).unwrap();
assert!(v4.ip().is_unspecified() && v4.is_ipv4());
assert_eq!(v4.port(), 8080);
assert!(parse_listen_addr(":80", Family::V4).unwrap().is_ipv4());
let v6 = parse_listen_addr(":80", Family::V6).unwrap();
assert!(v6.ip().is_unspecified() && v6.is_ipv6());
assert_eq!(v6.port(), 80);
}
#[test]
fn parse_full_addr_is_used_verbatim() {
let sa = parse_listen_addr("100.64.0.1:443", Family::Any).unwrap();
assert_eq!(sa.port(), 443);
assert_eq!(sa.ip().to_string(), "100.64.0.1");
}
#[test]
fn parse_bad_addr_is_typed_error() {
let err = parse_listen_addr("not-an-addr", Family::Any).unwrap_err();
assert!(matches!(err, Error::InvalidAddr { .. }));
}
#[test]
fn parse_listen_addr_enforces_the_pinned_family_on_explicit_hosts() {
assert!(
matches!(
parse_listen_addr("[::1]:80", Family::V4),
Err(Error::AddrFamilyMismatch { want: "IPv4", ref addr }) if addr == "[::1]:80"
),
"a v6 literal under tcp4 must be AddrFamilyMismatch(IPv4)"
);
assert!(
matches!(
parse_listen_addr("127.0.0.1:80", Family::V6),
Err(Error::AddrFamilyMismatch { want: "IPv6", .. })
),
"a v4 literal under tcp6 must be AddrFamilyMismatch(IPv6)"
);
assert!(matches!(
parse_listen_addr("[::1]:80", Family::V6),
Ok(SocketAddr::V6(_))
));
assert!(matches!(
parse_listen_addr("127.0.0.1:80", Family::V4),
Ok(SocketAddr::V4(_))
));
assert!(
parse_listen_addr("[::1]:80", Family::Any)
.unwrap()
.is_ipv6()
);
assert!(
parse_listen_addr("127.0.0.1:80", Family::Any)
.unwrap()
.is_ipv4()
);
}
#[test]
fn normalize_listen_addr_only_fills_a_bare_port() {
assert_eq!(normalize_listen_addr(":0", Family::Any), "0.0.0.0:0");
assert_eq!(normalize_listen_addr(":0", Family::V4), "0.0.0.0:0");
assert_eq!(normalize_listen_addr(":0", Family::V6), "[::]:0");
assert_eq!(normalize_listen_addr("0.0.0.0:0", Family::V4), "0.0.0.0:0");
assert_eq!(normalize_listen_addr("[::]:53", Family::V6), "[::]:53");
assert_eq!(normalize_listen_addr("host:53", Family::Any), "host:53");
}
#[tokio::test]
async fn listen_rejects_a_non_tcp_network_before_starting() {
let s = Server::new();
assert!(matches!(
s.listen("udp", ":80").await,
Err(Error::InvalidNetwork { network }) if network == "udp"
));
assert!(matches!(
s.listen("sctp", ":80").await,
Err(Error::InvalidNetwork { .. })
));
}
#[tokio::test]
async fn listen_reports_a_bad_addr_before_starting() {
let s = Server::new();
assert!(matches!(
s.listen("tcp", "not-an-addr").await,
Err(Error::InvalidAddr { .. })
));
}
#[tokio::test]
async fn listen_rejects_a_family_mismatched_explicit_host_before_starting() {
let s = Server::new();
assert!(matches!(
s.listen("tcp4", "[::1]:80").await,
Err(Error::AddrFamilyMismatch { want: "IPv4", .. })
));
assert!(matches!(
s.listen("tcp6", "127.0.0.1:80").await,
Err(Error::AddrFamilyMismatch { want: "IPv6", .. })
));
}
#[tokio::test]
async fn listen_packet_rejects_a_non_udp_network_before_starting() {
let s = Server::new();
assert!(matches!(
s.listen_packet("tcp", "0.0.0.0:0").await,
Err(Error::InvalidNetwork { network }) if network == "tcp"
));
assert!(matches!(
s.listen_packet("nope", "0.0.0.0:0").await,
Err(Error::InvalidNetwork { .. })
));
}
#[tokio::test]
async fn listen_tls_rejects_a_non_tailnet_name_before_starting() {
let s = Server::new();
let cfg = ts_control::ServeConfig {
name: "example.com".into(), port: 443,
target: ts_control::ServeTarget::Accept,
};
assert!(matches!(
s.listen_tls(&cfg).await,
Err(ts_control::CertError::NotTailnetName(n)) if n == "example.com"
));
}
#[tokio::test]
async fn listen_tls_rejects_a_zero_port_before_starting() {
let s = Server::new();
let cfg = ts_control::ServeConfig {
name: "host.tailnet.ts.net".into(), port: 0, target: ts_control::ServeTarget::Accept,
};
assert!(matches!(
s.listen_tls(&cfg).await,
Err(ts_control::CertError::Acme(_))
));
}
#[test]
fn start_failure_maps_to_a_typed_cert_io_error() {
let e = start_failed_cert(Error::Store(std::io::Error::other("boom")));
assert!(matches!(e, ts_control::CertError::Io(_)));
let msg = e.to_string();
assert!(msg.contains("server failed to start"), "got {msg:?}");
assert!(
msg.contains("boom"),
"underlying reason must be preserved, got {msg:?}"
);
}
#[cfg(feature = "acme")]
#[tokio::test]
async fn cert_pair_rejects_a_non_tailnet_name_before_starting() {
let s = Server::new();
assert!(matches!(
s.cert_pair("example.com", None).await,
Err(ts_control::CertError::NotTailnetName(n)) if n == "example.com"
));
}
#[test]
fn server_default_is_go_shaped() {
let s = Server::new();
assert!(!s.ephemeral);
assert!(s.dir.is_none());
assert!(s.store.is_none());
assert!(s.hostname.is_none());
}
#[test]
fn mem_store_round_trips() {
let store = MemStore::default();
assert!(store.read_state(STATE_KEY).unwrap().is_none());
store.write_state(STATE_KEY, b"blob").unwrap();
assert_eq!(
store.read_state(STATE_KEY).unwrap().as_deref(),
Some(&b"blob"[..])
);
}
#[test]
fn funnel_options_map_to_engine() {
let engine: ts_control::FunnelOptions = FunnelOptions::funnel_only().into();
assert!(engine.funnel_only);
}
#[test]
fn funnel_options_default_maps_to_non_funnel_only() {
let engine: ts_control::FunnelOptions = FunnelOptions::default().into();
assert!(!engine.funnel_only);
}
fn dummy_serve_config() -> crate::ServeConfig {
crate::ServeConfig {
name: "node.example.ts.net".into(),
port: 443,
target: crate::ServeTarget::Accept,
}
}
#[tokio::test]
async fn listen_funnel_reports_a_start_failure_not_a_funnel_denial() {
let mut s = Server::new();
s.control_url = Some("not a url".into());
let cfg = dummy_serve_config();
let res = s.listen_funnel(&cfg, FunnelOptions::default()).await;
match &res {
Err(e @ ListenFunnelError::Start(Error::InvalidControlUrl(_))) => {
assert!(
e.to_string().contains("invalid control URL"),
"start error dropped its underlying cause from Display: {e}"
);
assert!(
std::error::Error::source(e).is_some(),
"start error must expose the underlying Error as its source"
);
}
Err(ListenFunnelError::Start(e)) => panic!("start failed with the wrong cause: {e:?}"),
Err(ListenFunnelError::Funnel(f)) => {
panic!("a startup failure was misdiagnosed as a Funnel error: {f:?}")
}
Ok(_) => panic!("a bad control_url must not yield a live funnel listener"),
}
}
#[tokio::test]
async fn listen_service_reports_a_start_failure_not_a_bind_error() {
let mut s = Server::new();
s.control_url = Some("not a url".into());
let res = s
.listen_service("svc:web", ServiceMode::Tcp { port: 80 })
.await;
match &res {
Err(e @ ListenServiceError::Start(Error::InvalidControlUrl(_))) => {
assert!(
e.to_string().contains("invalid control URL"),
"start error dropped its underlying cause from Display: {e}"
);
assert!(
std::error::Error::source(e).is_some(),
"start error must expose the underlying Error as its source"
);
}
Err(ListenServiceError::Start(e)) => panic!("start failed with the wrong cause: {e:?}"),
Err(ListenServiceError::Service(se)) => {
panic!("a startup failure was misdiagnosed as a ServiceError: {se:?}")
}
Ok(_) => panic!("a bad control_url must not yield a live service listener"),
}
}
#[test]
fn listen_funnel_error_carries_the_engine_funnel_error_unchanged() {
let e: ListenFunnelError = ts_control::FunnelError::PortNotAllowed(8443).into();
assert!(
matches!(
e,
ListenFunnelError::Funnel(ts_control::FunnelError::PortNotAllowed(8443))
),
"engine FunnelError must pass through as ListenFunnelError::Funnel, unchanged"
);
}
#[test]
fn listen_service_error_carries_the_engine_service_error_unchanged() {
let e: ListenServiceError = ServiceError::UntaggedHost.into();
assert!(
matches!(e, ListenServiceError::Service(ServiceError::UntaggedHost)),
"engine ServiceError must pass through as ListenServiceError::Service, unchanged"
);
}
#[test]
fn documented_construction_idiom_compiles() {
let mut srv = Server::new();
srv.hostname = Some("web".into());
srv.auth_key = Some("tskey-xxxx".into());
srv.dir = Some("/var/lib/web".into());
srv.ephemeral = false;
srv.advertise_tags = vec!["tag:web".into()];
srv.port = Some(41641);
srv.configure(|c| c.accept_routes = true);
srv.store = Some(Arc::new(MemStore::default()));
assert_eq!(srv.hostname.as_deref(), Some("web"));
assert!(!srv.ephemeral);
assert!(srv.store.is_some());
assert!(srv.configure.is_some());
}
fn _assert_send_sync<T: Send + Sync>() {}
#[allow(dead_code)]
fn _server_is_send_sync() {
_assert_send_sync::<Server>();
}
fn scratch_dir(label: &str) -> PathBuf {
let pid = std::process::id();
let dir = std::env::temp_dir().join(format!("tsnet-rs-test-{pid}-{label}"));
std::fs::remove_dir_all(&dir).ok();
dir
}
#[tokio::test]
async fn build_config_maps_every_go_field_onto_config() {
let mut srv = Server::new();
srv.hostname = Some("web".into());
srv.auth_key = Some("tskey-auth-xxxx".into());
srv.control_url = Some("https://control.example.com".into());
srv.ephemeral = false;
srv.advertise_tags = vec!["tag:web".into(), "tag:prod".into()];
srv.port = Some(41641);
srv.run_web_client = true;
srv.client_id = Some("cid".into());
srv.client_secret = Some("csecret".into());
srv.id_token = Some("idtok".into());
srv.audience = Some("aud".into());
let cfg = srv.build_config().await.unwrap();
assert_eq!(cfg.requested_hostname.as_deref(), Some("web"));
assert_eq!(cfg.auth_key.as_deref(), Some("tskey-auth-xxxx"));
assert_eq!(cfg.control_server_url.scheme(), "https");
assert_eq!(
cfg.control_server_url.host_str(),
Some("control.example.com")
);
assert_eq!(
cfg.requested_tags,
vec!["tag:web".to_string(), "tag:prod".to_string()]
);
assert_eq!(cfg.wireguard_listen_port, Some(41641));
assert!(cfg.run_web_client);
assert_eq!(cfg.client_id.as_deref(), Some("cid"));
assert_eq!(cfg.client_secret.as_deref(), Some("csecret"));
assert_eq!(cfg.id_token.as_deref(), Some("idtok"));
assert_eq!(cfg.audience.as_deref(), Some("aud"));
assert_eq!(cfg.transport_mode, crate::TransportMode::Netstack);
}
#[tokio::test]
async fn build_config_maps_go_zero_value_defaults() {
let cfg = Server::new().build_config().await.unwrap();
assert!(cfg.requested_hostname.is_none(), "unset hostname ⇒ None");
assert!(cfg.requested_tags.is_empty(), "no advertise_tags ⇒ empty");
assert!(cfg.wireguard_listen_port.is_none(), "unset port ⇒ None");
assert!(!cfg.run_web_client, "run_web_client defaults off");
assert!(cfg.auth_key.is_none(), "unset auth_key ⇒ None");
assert!(
cfg.client_id.is_none() && cfg.client_secret.is_none(),
"unset OAuth/WIF client fields ⇒ None"
);
assert!(
cfg.id_token.is_none() && cfg.audience.is_none(),
"unset id_token/audience ⇒ None"
);
assert_eq!(cfg.transport_mode, crate::TransportMode::Netstack);
}
#[tokio::test]
async fn build_config_forces_go_default_ephemeral() {
assert!(
Config::default().ephemeral,
"precondition: a bare Config defaults to ephemeral=true"
);
let cfg = Server::new().build_config().await.unwrap();
assert!(
!cfg.ephemeral,
"a default tsnet::Server maps to a non-ephemeral Config (Go parity)"
);
let mut srv = Server::new();
srv.ephemeral = true;
assert!(srv.build_config().await.unwrap().ephemeral);
}
#[tokio::test]
async fn build_config_none_control_url_keeps_engine_default() {
let cfg = Server::new().build_config().await.unwrap();
assert_eq!(cfg.control_server_url, Config::default().control_server_url);
}
#[tokio::test]
async fn build_config_rejects_a_bad_control_url() {
let mut srv = Server::new();
srv.control_url = Some("not a url".into());
assert!(matches!(
srv.build_config().await,
Err(Error::InvalidControlUrl(_))
));
}
#[tokio::test]
async fn build_config_tun_selects_kernel_tun_transport() {
let mut srv = Server::new();
srv.tun = Some(TunSpec {
name: Some("tailscale0".into()),
mtu: Some(1280),
});
let cfg = srv.build_config().await.unwrap();
assert_eq!(
cfg.transport_mode,
crate::TransportMode::Tun(crate::TunConfig {
name: Some("tailscale0".into()),
mtu: Some(1280),
})
);
}
#[tokio::test]
async fn build_config_runs_configure_hook_after_mapping() {
let mut srv = Server::new();
srv.hostname = Some("exit".into());
srv.configure(|c| {
c.advertise_exit_node = true;
c.accept_routes = true;
});
let cfg = srv.build_config().await.unwrap();
assert!(cfg.advertise_exit_node);
assert!(cfg.accept_routes);
assert_eq!(cfg.requested_hostname.as_deref(), Some("exit"));
}
#[test]
fn file_store_round_trips_on_disk() {
let dir = scratch_dir("filestore");
let store = FileStore::new(dir.clone());
assert!(
store.read_state(STATE_KEY).unwrap().is_none(),
"never-written ⇒ None (a missing file is not an error)"
);
store.write_state(STATE_KEY, b"identity-blob").unwrap();
assert!(
dir.join(STATE_FILE).exists(),
"FileStore::new persists under dir/STATE_FILE"
);
assert_eq!(
store.read_state(STATE_KEY).unwrap().as_deref(),
Some(&b"identity-blob"[..])
);
assert_eq!(
FileStore::new(dir.clone())
.read_state(STATE_KEY)
.unwrap()
.as_deref(),
Some(&b"identity-blob"[..]),
"a fresh FileStore over the same dir reloads the persisted value"
);
std::fs::remove_dir_all(&dir).ok();
}
#[test]
fn file_store_at_writes_the_exact_path() {
let dir = scratch_dir("filestore-at");
let path = dir.join("custom.state");
FileStore::at(path.clone())
.write_state(STATE_KEY, b"x")
.unwrap();
assert!(
path.exists(),
"FileStore::at writes to the exact path given"
);
std::fs::remove_dir_all(&dir).ok();
}
#[tokio::test]
async fn dir_persists_node_identity_across_builds() {
let dir = scratch_dir("dir-state-root");
let mut srv = Server::new();
srv.dir = Some(dir.clone());
let cfg1 = srv.build_config().await.unwrap();
assert!(
dir.join(STATE_FILE).exists(),
"Dir persists identity to dir/STATE_FILE via the engine key-file format"
);
let mut srv2 = Server::new();
srv2.dir = Some(dir.clone());
let cfg2 = srv2.build_config().await.unwrap();
assert_eq!(
serde_json::to_vec(&cfg1.key_state).unwrap(),
serde_json::to_vec(&cfg2.key_state).unwrap(),
"a Dir-rooted node reloads a stable identity rather than re-minting each boot"
);
std::fs::remove_dir_all(&dir).ok();
}
#[tokio::test]
async fn custom_store_round_trips_identity_through_build_config() {
let store: Arc<dyn StateStore> = Arc::new(MemStore::default());
let mut srv = Server::new();
srv.store = Some(store.clone());
let cfg1 = srv.build_config().await.unwrap();
assert!(
store.read_state(STATE_KEY).unwrap().is_some(),
"the store now holds the minted identity blob"
);
let mut srv2 = Server::new();
srv2.store = Some(store.clone());
let cfg2 = srv2.build_config().await.unwrap();
assert_eq!(
serde_json::to_vec(&cfg1.key_state).unwrap(),
serde_json::to_vec(&cfg2.key_state).unwrap(),
"a shared store yields a stable identity across servers"
);
}
#[tokio::test]
async fn store_takes_precedence_over_dir() {
let dir = scratch_dir("store-precedence");
let store: Arc<dyn StateStore> = Arc::new(MemStore::default());
let mut srv = Server::new();
srv.dir = Some(dir.clone());
srv.store = Some(store.clone());
srv.build_config().await.unwrap();
assert!(store.read_state(STATE_KEY).unwrap().is_some());
assert!(
!dir.join(STATE_FILE).exists(),
"store must take precedence over dir (design §8)"
);
std::fs::remove_dir_all(&dir).ok();
}
#[tokio::test]
async fn close_on_never_started_server_is_clean() {
assert!(Server::new().close(None).await);
assert!(Server::new().close(Some(Duration::from_millis(1))).await);
}
#[tokio::test]
async fn start_fails_fast_on_bad_config_before_touching_the_network() {
let mut s = Server::new();
s.control_url = Some("not a url".into());
assert!(matches!(s.start().await, Err(Error::InvalidControlUrl(_))));
}
#[tokio::test]
async fn up_fails_fast_on_bad_config() {
let mut s = Server::new();
s.control_url = Some("not a url".into());
assert!(matches!(s.up(None).await, Err(Error::InvalidControlUrl(_))));
}
#[tokio::test]
async fn close_is_clean_after_a_failed_start() {
let mut s = Server::new();
s.control_url = Some("not a url".into());
assert!(matches!(s.start().await, Err(Error::InvalidControlUrl(_))));
assert!(s.close(None).await);
}
#[tokio::test]
async fn dial_rejects_unsupported_network_fail_fast() {
for n in [
"", "TCP", "tcp5", "sctp", "unix", "ip", "udplite", "tcp ", " udp",
] {
match Server::new().dial(n, "host:80").await {
Err(Error::UnsupportedNetwork { network }) => assert_eq!(network, n),
Err(e) => panic!("dial({n:?}) should be UnsupportedNetwork, got Err({e:?})"),
Ok(_) => panic!("dial({n:?}) should be UnsupportedNetwork, got Ok(conn)"),
}
}
}
#[tokio::test]
async fn dial_accepts_every_tsnet_network_then_reaches_lazy_start() {
for (n, addr) in [
("tcp", "host:80"),
("tcp4", "1.2.3.4:80"),
("tcp6", "[2001:db8::1]:80"),
("udp", "host:53"),
("udp4", "1.2.3.4:53"),
("udp6", "[2001:db8::1]:53"),
] {
let mut s = Server::new();
s.control_url = Some("not a url".into());
assert!(
matches!(s.dial(n, addr).await, Err(Error::InvalidControlUrl(_))),
"dial({n:?}, {addr:?}) with a bad control_url should fail fast at config build"
);
}
}
#[tokio::test]
async fn dial_parses_network_before_touching_config() {
let mut s = Server::new();
s.control_url = Some("not a url".into());
match s.dial("sctp", "host:80").await {
Err(Error::UnsupportedNetwork { network }) => assert_eq!(network, "sctp"),
Err(e) => panic!("network parse must precede config build, got Err({e:?})"),
Ok(_) => panic!("network parse must precede config build, got Ok(conn)"),
}
}
#[tokio::test]
async fn dial_tcp_and_dial_udp_fail_fast_on_bad_config() {
let mut s = Server::new();
s.control_url = Some("not a url".into());
assert!(matches!(
s.dial_tcp("host:80").await,
Err(Error::InvalidControlUrl(_))
));
let mut s = Server::new();
s.control_url = Some("not a url".into());
assert!(matches!(
s.dial_udp("host:80").await,
Err(Error::InvalidControlUrl(_))
));
}
fn mock_status(body: &'static [u8]) -> localapi::StatusFn {
Arc::new(move || Box::pin(async move { Ok(body.to_vec()) }))
}
#[test]
fn gen_cred_is_32_lowercase_hex() {
let cred = gen_cred();
assert_eq!(cred.len(), 32, "16 random bytes → 32 hex chars (Go parity)");
assert!(
cred.chars()
.all(|c| c.is_ascii_hexdigit() && !c.is_ascii_uppercase())
);
assert_ne!(gen_cred(), gen_cred(), "credentials are random per call");
}
#[test]
fn find_subslice_locates_header_terminator() {
assert_eq!(find_subslice(b"ab\r\n\r\ncd", b"\r\n\r\n"), Some(2));
assert_eq!(find_subslice(b"no terminator", b"\r\n\r\n"), None);
assert_eq!(find_subslice(b"", b"\r\n\r\n"), None);
}
#[test]
fn parse_response_extracts_code_and_body() {
let (code, body) =
parse_response(b"HTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\nhi").unwrap();
assert_eq!(code, 200);
assert_eq!(body, b"hi");
assert_eq!(parse_response(b"garbage without terminator"), None);
}
#[test]
fn parse_head_reads_method_target_auth_and_sec_tailscale() {
let head = b"GET /localapi/v0/status HTTP/1.1\r\nHost: x\r\nSec-Tailscale: localapi\r\nAuthorization: Basic dXNlcjpwYXNz\r\nAccept: */*";
let (method, target, password, sec_tailscale) = localapi::parse_head(head).unwrap();
assert_eq!(method, "GET");
assert_eq!(target, "/localapi/v0/status");
assert_eq!(password.as_deref(), Some("pass"));
assert_eq!(
sec_tailscale.as_deref(),
Some("localapi"),
"captures the anti-rebinding header"
);
let no_hdr =
b"GET /localapi/v0/status HTTP/1.1\r\nHost: x\r\nAuthorization: Basic dXNlcjpwYXNz";
let (_, _, _, sec_tailscale) = localapi::parse_head(no_hdr).unwrap();
assert_eq!(sec_tailscale, None);
}
#[test]
fn parse_head_rejects_malformed_request_line() {
assert!(localapi::parse_head(b"GET-only-one-token").is_none());
assert!(
localapi::parse_head(b"GET /x").is_none(),
"needs a version token"
);
}
#[test]
fn basic_auth_password_ignores_username() {
let with_user = STANDARD.encode("anyuser:the-cred");
let no_user = STANDARD.encode(":the-cred");
assert_eq!(
localapi::basic_auth_password(&format!("Basic {with_user}")).as_deref(),
Some("the-cred")
);
assert_eq!(
localapi::basic_auth_password(&format!("Basic {no_user}")).as_deref(),
Some("the-cred")
);
assert_eq!(
localapi::basic_auth_password(&format!("bAsIc {with_user}")).as_deref(),
Some("the-cred")
);
assert!(localapi::basic_auth_password("Bearer xyz").is_none());
assert!(localapi::basic_auth_password("Basic !!!not-base64").is_none());
assert!(localapi::basic_auth_password("no-space-token").is_none());
}
#[test]
fn cred_ok_matches_only_exact_credentials() {
assert!(localapi::cred_ok("abc123", "abc123"));
assert!(!localapi::cred_ok("abc123", "abc124"));
assert!(
!localapi::cred_ok("abc", "abc123"),
"length mismatch is a mismatch"
);
assert!(!localapi::cred_ok("", "x"));
}
#[test]
fn status_json_serializes_status_snapshot() {
use crate::StableNodeId;
let node = StatusNode {
stable_id: StableNodeId("nabc123".to_string()),
display_name: "web.tail0.ts.net".to_string(),
ipv4: "100.64.0.1".parse().unwrap(),
ipv6: "fd7a:115c:a1e0::1".parse().unwrap(),
online: Some(true),
last_seen: None,
allowed_routes: vec![],
is_exit_node: false,
cur_addr: None,
relay: Some("nyc".to_string()),
ssh_host_keys: vec![],
};
let status = Status {
self_node: Some(node),
peers: vec![],
active_exit_node: None,
magic_dns_suffix: Some("tail0.ts.net".to_string()),
};
let bytes = status_json(&status);
let v: serde_json::Value = serde_json::from_slice(&bytes).expect("valid JSON");
assert_eq!(v["self"]["stable_id"], "nabc123");
assert_eq!(v["self"]["display_name"], "web.tail0.ts.net");
assert_eq!(v["self"]["ipv4"], "100.64.0.1");
assert_eq!(v["self"]["online"], true);
assert_eq!(v["self"]["relay"], "nyc");
assert_eq!(v["magic_dns_suffix"], "tail0.ts.net");
assert!(v["peers"].as_array().unwrap().is_empty());
assert!(v["active_exit_node"].is_null());
}
#[tokio::test]
async fn localapi_server_authenticates_and_routes_over_real_socket() {
let listener = TcpListener::bind((std::net::Ipv4Addr::LOCALHOST, 0))
.await
.unwrap();
let addr = listener.local_addr().unwrap();
let cred = "s3cr3t-cred".to_string();
let task = tokio::spawn(localapi::serve(
listener,
cred.clone(),
mock_status(br#"{"ok":true}"#),
));
let (code, body) = localapi_client_get(addr, &cred, "/localapi/v0/status")
.await
.unwrap();
assert_eq!(code, 200);
assert_eq!(body, br#"{"ok":true}"#);
let (code, _) = localapi_client_get(addr, &cred, "/localapi/v0/nope")
.await
.unwrap();
assert_eq!(code, 404);
let (code, _) = localapi_client_get(addr, "wrong-cred", "/localapi/v0/status")
.await
.unwrap();
assert_eq!(code, 401);
{
use tokio::io::{AsyncReadExt as _, AsyncWriteExt as _};
let mut sock = TcpStream::connect(addr).await.unwrap();
sock.write_all(
b"GET /localapi/v0/status HTTP/1.1\r\nHost: x\r\nSec-Tailscale: localapi\r\nConnection: close\r\n\r\n",
)
.await
.unwrap();
let mut resp = Vec::new();
sock.read_to_end(&mut resp).await.unwrap();
assert_eq!(parse_response(&resp).unwrap().0, 401);
}
{
use tokio::io::{AsyncReadExt as _, AsyncWriteExt as _};
let auth = STANDARD.encode(format!(":{cred}"));
let mut sock = TcpStream::connect(addr).await.unwrap();
sock.write_all(
format!(
"GET /localapi/v0/status HTTP/1.1\r\nHost: x\r\nAuthorization: Basic {auth}\r\nConnection: close\r\n\r\n"
)
.as_bytes(),
)
.await
.unwrap();
let mut resp = Vec::new();
sock.read_to_end(&mut resp).await.unwrap();
assert_eq!(
parse_response(&resp).unwrap().0,
403,
"no Sec-Tailscale header → 403 even with a valid credential"
);
}
task.abort();
}
#[tokio::test]
async fn local_client_round_trips_through_the_localapi_server() {
let listener = TcpListener::bind((std::net::Ipv4Addr::LOCALHOST, 0))
.await
.unwrap();
let addr = listener.local_addr().unwrap();
let cred = "local-api-cred".to_string();
let task = tokio::spawn(localapi::serve(
listener,
cred.clone(),
mock_status(br#"{"self":null,"peers":[]}"#),
));
let client = LocalClient {
address: addr,
cred: cred.clone(),
};
assert_eq!(client.address(), addr);
assert_eq!(client.credential(), cred);
let body = client.status().await.unwrap();
assert_eq!(body, br#"{"self":null,"peers":[]}"#);
let (code, _) = client.get("/localapi/v0/status").await.unwrap();
assert_eq!(code, 200);
let bad = LocalClient {
address: addr,
cred: "nope".to_string(),
};
assert!(matches!(bad.status().await, Err(Error::Loopback(_))));
task.abort();
}
#[test]
fn loopback_result_carries_both_distinct_credentials() {
let lb = Loopback {
address: "127.0.0.1:1080".parse().unwrap(),
proxy_cred: "proxy".to_string(),
local_api_address: "127.0.0.1:1081".parse().unwrap(),
local_api_cred: "localapi".to_string(),
};
let cloned = lb.clone();
assert_ne!(cloned.proxy_cred, cloned.local_api_cred);
assert_ne!(cloned.address, cloned.local_api_address);
assert!(format!("{cloned:?}").contains("local_api_cred"));
}
}