mod self_origin;
mod whois;
use std::collections::HashMap;
use std::future::Future;
use std::net::{IpAddr, SocketAddr};
use std::sync::{Arc, Mutex, PoisonError};
use std::time::Duration;
use anyhow::Context;
use axum::Router;
use axum::extract::{ConnectInfo, Request, State};
use axum::http::StatusCode;
use axum::middleware::Next;
use axum::response::{IntoResponse, Response};
use futures_util::future::BoxFuture;
use tokio::sync::OnceCell;
use tokio::time::Instant;
pub use self_origin::{SelfAuthorities, check_target};
pub use whois::{PeerIdentity, PeerResolver, TailscaleCliResolver, WhoisError};
pub const WHOIS_TIMEOUT: Duration = Duration::from_secs(2);
pub const IDENTITY_TTL: Duration = Duration::from_secs(30);
pub const FAILURE_TTL: Duration = Duration::from_secs(5);
const CACHE_SWEEP_THRESHOLD: usize = 256;
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum PeerVerdict {
Allow,
ForeignLogin(String),
Tagged,
Unresolved(String),
}
struct Resolved {
at: Instant,
result: Result<PeerIdentity, String>,
}
impl Resolved {
fn expired(&self, now: Instant) -> bool {
let ttl = if self.result.is_ok() {
IDENTITY_TTL
} else {
FAILURE_TTL
};
now.duration_since(self.at) >= ttl
}
}
type Slot = Arc<OnceCell<Resolved>>;
pub struct TailnetPeerGate {
resolver: Arc<dyn PeerResolver>,
host_ip: IpAddr,
timeout: Duration,
cache: Mutex<HashMap<IpAddr, Slot>>,
}
impl TailnetPeerGate {
pub fn new(resolver: Arc<dyn PeerResolver>, host_ip: IpAddr) -> Self {
Self::with_timeout(resolver, host_ip, WHOIS_TIMEOUT)
}
pub fn with_timeout(
resolver: Arc<dyn PeerResolver>,
host_ip: IpAddr,
timeout: Duration,
) -> Self {
Self {
resolver,
host_ip,
timeout,
cache: Mutex::new(HashMap::new()),
}
}
pub async fn authorize(&self, peer: IpAddr) -> PeerVerdict {
if self.host_ip.is_unspecified() {
return PeerVerdict::Unresolved(format!(
"listener is bound to the wildcard address {}",
self.host_ip
));
}
let host = match self.identity(self.host_ip).await {
Ok(id) => id,
Err(e) => return PeerVerdict::Unresolved(format!("host identity: {e}")),
};
let peer_id = match self.identity(peer).await {
Ok(id) => id,
Err(e) => return PeerVerdict::Unresolved(e),
};
if host.tagged || peer_id.tagged {
return PeerVerdict::Tagged;
}
if peer_id.login == host.login {
PeerVerdict::Allow
} else {
PeerVerdict::ForeignLogin(peer_id.login)
}
}
pub async fn host_node_name(&self) -> Option<String> {
self.identity(self.host_ip).await.ok()?.node_name
}
async fn identity(&self, ip: IpAddr) -> Result<PeerIdentity, String> {
let slot = self.slot_for(ip);
let resolved = slot
.get_or_init(|| async {
let result = match tokio::time::timeout(self.timeout, self.resolver.whois(ip)).await
{
Ok(Ok(id)) => Ok(id),
Ok(Err(e)) => Err(e.to_string()),
Err(_) => Err(WhoisError::Timeout(self.timeout).to_string()),
};
Resolved {
at: Instant::now(),
result,
}
})
.await;
resolved.result.clone()
}
fn slot_for(&self, ip: IpAddr) -> Slot {
let now = Instant::now();
let mut cache = self.cache.lock().unwrap_or_else(PoisonError::into_inner);
if let Some(slot) = cache.get(&ip)
&& slot.get().is_none_or(|r| !r.expired(now))
{
return Arc::clone(slot);
}
if cache.len() >= CACHE_SWEEP_THRESHOLD {
cache.retain(|_, s| s.get().is_none_or(|r| !r.expired(now)));
}
let slot = Slot::default();
cache.insert(ip, Arc::clone(&slot));
slot
}
}
fn forbidden() -> Response {
(
StatusCode::FORBIDDEN,
axum::Json(serde_json::json!({
"error": "tailnet peer is not authorized for this console",
})),
)
.into_response()
}
#[derive(Clone)]
pub struct TailnetGuard {
pub gate: Arc<TailnetPeerGate>,
pub listen: SocketAddr,
}
pub async fn guard_tailnet_peer(
State(guard): State<TailnetGuard>,
req: Request,
next: Next,
) -> Response {
let Some(ConnectInfo(peer)) = req.extensions().get::<ConnectInfo<SocketAddr>>().copied() else {
tracing::warn!("tailnet listener refused a request with no peer address");
return forbidden();
};
match guard.gate.authorize(peer.ip()).await {
PeerVerdict::Allow => {}
PeerVerdict::ForeignLogin(login) => {
tracing::warn!(peer = %peer, login = %login, "tailnet listener refused a peer owned by another login");
return forbidden();
}
PeerVerdict::Tagged => {
tracing::warn!(peer = %peer, "tailnet listener refused a tagged node (or the host is tagged)");
return forbidden();
}
PeerVerdict::Unresolved(reason) => {
tracing::warn!(peer = %peer, reason = %reason, "tailnet listener refused a peer whose identity could not be determined");
return forbidden();
}
}
let name = guard.gate.host_node_name().await;
let allowed = SelfAuthorities::new(guard.listen, name.as_deref());
if let Err(reason) = check_target(req.headers(), req.uri(), &allowed) {
tracing::warn!(peer = %peer, reason, "tailnet listener refused a request not addressed to itself");
return forbidden();
}
next.run(req).await
}
pub async fn serve_tailnet(
listener: tokio::net::TcpListener,
router: Router,
gate: Arc<TailnetPeerGate>,
shutdown: impl Future<Output = ()> + Send + 'static,
) -> std::io::Result<()> {
let listen = listener.local_addr()?;
let app = router.layer(axum::middleware::from_fn_with_state(
TailnetGuard { gate, listen },
guard_tailnet_peer,
));
axum::serve(
listener,
app.into_make_service_with_connect_info::<SocketAddr>(),
)
.with_graceful_shutdown(shutdown)
.await
}
pub struct ConsoleListeners {
pub bound: Vec<SocketAddr>,
primary: tokio::task::JoinHandle<std::io::Result<()>>,
}
impl ConsoleListeners {
pub async fn wait_primary(self) -> anyhow::Result<()> {
self.primary
.await
.context("primary listener task failed")?
.context("server error")
}
}
pub async fn serve_listeners<S, F>(
addrs: &[SocketAddr],
router: &Router,
resolver: Arc<dyn PeerResolver>,
shutdown: S,
) -> anyhow::Result<ConsoleListeners>
where
S: Fn() -> F,
F: Future<Output = ()> + Send + 'static,
{
let mut bound = Vec::with_capacity(addrs.len());
let mut primary = None;
for &addr in addrs {
let listener = crate::bind::bind_listener(addr).await?;
let local = listener.local_addr().context("get local addr")?;
let serve: BoxFuture<'static, std::io::Result<()>> = if local.ip().is_loopback() {
let (app, stop) = (router.clone(), shutdown());
Box::pin(async move {
axum::serve(listener, app)
.with_graceful_shutdown(stop)
.await
})
} else {
if local.ip().is_unspecified() {
tracing::warn!(
"trusty-console listener {local} is a wildcard bind; the tailnet peer \
gate refuses every request on it (bind the tailnet address, or use --tailscale)"
);
}
eprintln!("trusty-console (tailnet peer gate): http://{local}");
let gate = Arc::new(TailnetPeerGate::new(Arc::clone(&resolver), local.ip()));
Box::pin(serve_tailnet(listener, router.clone(), gate, shutdown()))
};
if primary.is_none() {
tracing::info!("trusty-console listening on http://{local}");
primary = Some(tokio::spawn(serve));
} else {
tracing::info!("trusty-console also listening on http://{local}");
tokio::spawn(async move {
if let Err(e) = serve.await {
tracing::warn!("extra listener {local} exited: {e}");
}
});
}
bound.push(local);
}
let primary = primary.context("bind address list is empty")?;
Ok(ConsoleListeners { bound, primary })
}
#[cfg(test)]
mod tests;