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 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 {
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 async fn spawn_tailnet_listeners<S, F>(
addrs: &[SocketAddr],
router: &Router,
resolver: Arc<dyn PeerResolver>,
shutdown: S,
) -> anyhow::Result<Vec<SocketAddr>>
where
S: Fn() -> F,
F: Future<Output = ()> + Send + 'static,
{
let mut bound = Vec::with_capacity(addrs.len());
for &addr in addrs {
let listener = crate::bind::bind_listener(addr).await?;
let local = listener.local_addr().context("get extra local addr")?;
tracing::info!("trusty-console also listening on http://{local}");
eprintln!("trusty-console (tailnet): http://{local}");
let gate = Arc::new(TailnetPeerGate::new(Arc::clone(&resolver), local.ip()));
let serve = serve_tailnet(listener, router.clone(), gate, shutdown());
tokio::spawn(async move {
if let Err(e) = serve.await {
tracing::warn!("extra listener {local} exited: {e}");
}
});
bound.push(local);
}
Ok(bound)
}
#[cfg(test)]
mod tests;