use std::collections::{BTreeSet, HashMap};
use std::net::{IpAddr, SocketAddr};
use std::sync::{Arc, Mutex};
use std::time::{Duration, Instant};
use axum::body::Body;
use axum::extract::{ConnectInfo, Request, State};
use axum::http::{HeaderMap, StatusCode};
use axum::middleware::Next;
use axum::response::{IntoResponse, Response};
use axum::routing::get;
use crate::context::SelfIdentity;
pub const DEFAULT_BIND: &str = "127.0.0.1:8449";
pub const MAX_BODY_BYTES: usize = 4 << 20;
pub const RATE_BURST: u32 = 120;
pub const RATE_WINDOW: Duration = Duration::from_secs(60);
pub const HEALTH_PATH: &str = "/health";
pub const MCP_PATH: &str = "/mcp";
#[derive(Debug, Clone)]
pub struct Guard {
token: Option<Arc<tailscale_rest::Secret>>,
hosts: Arc<BTreeSet<String>>,
origins: Arc<BTreeSet<String>>,
limiter: Arc<Mutex<RateLimiter>>,
peers: Arc<HashMap<IpAddr, String>>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Refusal {
UnknownHost,
ForbiddenOrigin,
RateLimited,
BadToken,
}
impl Refusal {
pub const fn status(self) -> StatusCode {
match self {
Self::UnknownHost | Self::ForbiddenOrigin => StatusCode::FORBIDDEN,
Self::RateLimited => StatusCode::TOO_MANY_REQUESTS,
Self::BadToken => StatusCode::UNAUTHORIZED,
}
}
pub const fn message(self) -> &'static str {
match self {
Self::UnknownHost => {
"this server does not answer for that `Host`; add it with `--http-allow-host`"
}
Self::ForbiddenOrigin => {
"this server does not answer requests from a browser page; add the origin with \
`--http-allow-origin` if that is what you meant"
}
Self::RateLimited => "too many requests from this address; wait and try again",
Self::BadToken => "a bearer token is required and did not match",
}
}
}
impl IntoResponse for Refusal {
fn into_response(self) -> Response {
(self.status(), format!("{}\n", self.message())).into_response()
}
}
#[derive(Debug, Clone)]
pub struct Caller {
pub address: IpAddr,
pub name: Option<String>,
}
impl Caller {
pub fn describe(&self) -> String {
match &self.name {
Some(name) => format!("{name} ({})", self.address),
None => self.address.to_string(),
}
}
}
impl Guard {
pub fn for_session(
settings: &crate::config::HttpConfig,
identity: &SelfIdentity,
peers: HashMap<IpAddr, String>,
) -> Self {
Self::new(
settings.token.clone(),
&settings.allow_hosts,
&settings.allow_origins,
identity,
peers,
)
}
pub fn new(
token: Option<tailscale_rest::Secret>,
extra_hosts: &[String],
origins: &[String],
identity: &SelfIdentity,
peers: HashMap<IpAddr, String>,
) -> Self {
let mut hosts: BTreeSet<String> = ["localhost", "127.0.0.1", "[::1]", "::1"]
.iter()
.map(|host| (*host).to_owned())
.collect();
if let Some(dns_name) = &identity.dns_name {
let full = dns_name.trim_end_matches('.').to_ascii_lowercase();
if let Some(short) = full.split('.').next() {
hosts.insert(short.to_owned());
}
hosts.insert(full);
}
hosts.extend(identity.addresses.iter().map(|a| bracketed(a)));
hosts.extend(extra_hosts.iter().map(|host| normalise_host(host)));
Self {
token: token.map(Arc::new),
hosts: Arc::new(hosts),
origins: Arc::new(origins.iter().map(|o| normalise_origin(o)).collect()),
limiter: Arc::new(Mutex::new(RateLimiter::default())),
peers: Arc::new(peers),
}
}
pub fn hosts(&self) -> Vec<&str> {
self.hosts.iter().map(String::as_str).collect()
}
pub fn admit(&self, headers: &HeaderMap, address: IpAddr, now: Instant) -> Result<(), Refusal> {
let host = headers
.get(axum::http::header::HOST)
.and_then(|value| value.to_str().ok())
.unwrap_or_default();
if !self.hosts.contains(&normalise_host(host)) {
return Err(Refusal::UnknownHost);
}
if let Some(origin) = headers.get(axum::http::header::ORIGIN) {
let origin = normalise_origin(origin.to_str().unwrap_or_default());
if !self.origins.contains(&origin) {
return Err(Refusal::ForbiddenOrigin);
}
}
if !self
.limiter
.lock()
.map(|mut limiter| limiter.allow(address, now))
.unwrap_or(true)
{
return Err(Refusal::RateLimited);
}
let Some(expected) = &self.token else {
return Ok(());
};
let given = headers
.get(axum::http::header::AUTHORIZATION)
.and_then(|value| value.to_str().ok())
.and_then(|value| {
value
.strip_prefix("Bearer ")
.or_else(|| value.strip_prefix("bearer "))
})
.unwrap_or_default();
if same_secret(given.as_bytes(), expected.expose().as_bytes()) {
Ok(())
} else {
Err(Refusal::BadToken)
}
}
pub fn caller(&self, address: IpAddr) -> Caller {
Caller {
address,
name: self.peers.get(&address).cloned(),
}
}
}
fn normalise_host(host: &str) -> String {
let host = host.trim().to_ascii_lowercase();
if let Some(rest) = host.strip_prefix('[') {
let closed = rest.split_once(']').map(|(inside, _)| inside);
return closed.map_or(host.clone(), |inside| format!("[{inside}]"));
}
host.split_once(':')
.map_or(host.clone(), |(name, _)| name.to_owned())
}
fn normalise_origin(origin: &str) -> String {
let origin = origin.trim();
let Some((scheme, rest)) = origin.split_once("://") else {
return origin.to_ascii_lowercase();
};
let scheme = scheme.to_ascii_lowercase();
let authority = normalise_host(rest.split('/').next().unwrap_or_default());
let port = rest
.split('/')
.next()
.and_then(|a| a.rsplit_once(':'))
.filter(|(before, _)| !before.ends_with(':') && !before.is_empty())
.and_then(|(_, port)| port.parse::<u16>().ok())
.filter(|port| !matches!((scheme.as_str(), port), ("http", 80) | ("https", 443)));
match port {
Some(port) => format!("{scheme}://{authority}:{port}"),
None => format!("{scheme}://{authority}"),
}
}
fn bracketed(address: &str) -> String {
if address.contains(':') {
format!("[{}]", address.to_ascii_lowercase())
} else {
address.to_ascii_lowercase()
}
}
fn same_secret(given: &[u8], expected: &[u8]) -> bool {
let mut difference = (given.len() ^ expected.len()) as u32;
let longest = given.len().max(expected.len());
for i in 0..longest {
let a = given.get(i).copied().unwrap_or(0);
let b = expected.get(i).copied().unwrap_or(0);
difference |= u32::from(a ^ b);
}
difference == 0
}
#[derive(Debug, Default)]
struct RateLimiter {
seen: HashMap<IpAddr, Bucket>,
}
#[derive(Debug, Clone, Copy)]
struct Bucket {
left: f64,
at: Instant,
}
impl RateLimiter {
fn allow(&mut self, address: IpAddr, now: Instant) -> bool {
let rate = f64::from(RATE_BURST) / RATE_WINDOW.as_secs_f64();
self.seen.retain(|_, bucket| {
let refill = now.saturating_duration_since(bucket.at).as_secs_f64() * rate;
refill < f64::from(RATE_BURST) - bucket.left
});
let bucket = self.seen.entry(address).or_insert(Bucket {
left: f64::from(RATE_BURST),
at: now,
});
let refill = now.saturating_duration_since(bucket.at).as_secs_f64() * rate;
bucket.left = (bucket.left + refill).min(f64::from(RATE_BURST));
bucket.at = now;
if bucket.left < 1.0 {
return false;
}
bucket.left -= 1.0;
true
}
}
async fn admission(
State(guard): State<Guard>,
ConnectInfo(peer): ConnectInfo<SocketAddr>,
request: Request,
next: Next,
) -> Response {
let address = peer.ip();
if let Err(refusal) = guard.admit(request.headers(), address, Instant::now()) {
tracing::warn!(
caller = guard.caller(address).describe(),
refusal = ?refusal,
"refused an HTTP request"
);
return refusal.into_response();
}
let caller = guard.caller(address);
tracing::info!(caller = caller.describe(), path = %request.uri().path(), "http request");
let mut request = request;
request.extensions_mut().insert(caller);
next.run(request).await
}
pub fn router<S>(guard: Guard, mcp: S) -> axum::Router
where
S: tower::Service<Request<Body>, Response = Response, Error = std::convert::Infallible>
+ Clone
+ Send
+ Sync
+ 'static,
S::Future: Send + 'static,
{
axum::Router::new()
.route_service(MCP_PATH, mcp)
.layer(axum::middleware::from_fn_with_state(
guard.clone(),
admission,
))
.route(HEALTH_PATH, get(health))
.with_state(guard)
}
pub async fn serve(
settings: &crate::config::HttpConfig,
guard: Guard,
server: crate::server::TailscaleMcpServer,
) -> std::io::Result<()> {
use rmcp::transport::streamable_http_server::{
StreamableHttpService, session::local::LocalSessionManager,
};
let transport = StreamableHttpService::new(
move || Ok(server.clone()),
Arc::new(LocalSessionManager::default()),
rmcp::transport::streamable_http_server::StreamableHttpServerConfig::default()
.with_allowed_hosts(guard.hosts())
.with_max_request_body_bytes(MAX_BODY_BYTES)
.with_legacy_session_mode(settings.stateful),
);
let transport = tower::util::ServiceExt::<Request<Body>>::map_response(transport, |response| {
axum::http::Response::map(response, Body::new)
});
let listener = tokio::net::TcpListener::bind(settings.bind).await?;
tracing::info!(
address = %settings.bind,
authenticated = guard.token.is_some(),
sessions = settings.stateful,
"serving MCP over HTTP"
);
axum::serve(
listener,
router(guard, transport).into_make_service_with_connect_info::<SocketAddr>(),
)
.await
}
async fn health() -> impl IntoResponse {
(
StatusCode::OK,
[("content-type", "application/json")],
concat!(
"{\"status\":\"ok\",\"server\":\"",
env!("CARGO_PKG_NAME"),
"\"}\n"
),
)
}
#[cfg(test)]
mod tests {
use super::*;
fn identity() -> SelfIdentity {
SelfIdentity {
node_id: Some("n1111111CNTRL".to_owned()),
numeric_id: None,
addresses: vec!["100.64.0.1".to_owned(), "fd7a:115c:a1e0::1".to_owned()],
dns_name: Some("workstation.example-tailnet.ts.net.".to_owned()),
}
}
fn guard(token: Option<&str>, origins: &[&str]) -> Guard {
Guard::new(
token.map(tailscale_rest::Secret::new),
&[],
&origins.iter().map(|o| (*o).to_owned()).collect::<Vec<_>>(),
&identity(),
HashMap::new(),
)
}
fn headers(pairs: &[(&str, &str)]) -> HeaderMap {
let mut headers = HeaderMap::new();
for (name, value) in pairs {
headers.insert(
axum::http::HeaderName::from_bytes(name.as_bytes()).expect("a header name"),
value.parse().expect("a header value"),
);
}
headers
}
fn here() -> IpAddr {
IpAddr::from([127, 0, 0, 1])
}
#[test]
fn this_nodes_own_names_are_allowed_without_configuration() {
let guard = guard(None, &[]);
for host in [
"localhost",
"127.0.0.1:8449",
"workstation",
"workstation.example-tailnet.ts.net",
"WORKSTATION.EXAMPLE-TAILNET.TS.NET:8449",
"100.64.0.1:8449",
"[fd7a:115c:a1e0::1]:8449",
] {
assert_eq!(
guard.admit(&headers(&[("host", host)]), here(), Instant::now()),
Ok(()),
"`{host}` is one of this node's own names"
);
}
}
#[test]
fn a_host_this_server_does_not_answer_for_is_refused_before_anything_else() {
let guard = guard(Some("s3cret-token-value"), &[]);
assert_eq!(
guard.admit(
&headers(&[
("host", "evil.example"),
("authorization", "Bearer s3cret-token-value")
]),
here(),
Instant::now()
),
Err(Refusal::UnknownHost),
"and refused for the host, not for the token, which was right"
);
}
#[test]
fn a_browser_origin_is_refused_unless_it_was_listed() {
let closed = guard(None, &[]);
assert_eq!(
closed.admit(
&headers(&[("host", "localhost"), ("origin", "https://app.example")]),
here(),
Instant::now()
),
Err(Refusal::ForbiddenOrigin)
);
let opened = guard(None, &["https://app.example"]);
assert_eq!(
opened.admit(
&headers(&[("host", "localhost"), ("origin", "https://app.example")]),
here(),
Instant::now()
),
Ok(())
);
assert_eq!(
opened.admit(
&headers(&[("host", "localhost"), ("origin", "https://other.example")]),
here(),
Instant::now()
),
Err(Refusal::ForbiddenOrigin)
);
}
#[test]
fn the_token_has_to_match_and_a_missing_one_is_the_same_answer_as_a_wrong_one() {
let guard = guard(Some("s3cret-token-value"), &[]);
let host = ("host", "localhost");
assert_eq!(
guard.admit(
&headers(&[host, ("authorization", "Bearer s3cret-token-value")]),
here(),
Instant::now()
),
Ok(())
);
assert_eq!(
guard.admit(&headers(&[host]), here(), Instant::now()),
Err(Refusal::BadToken)
);
assert_eq!(
guard.admit(
&headers(&[host, ("authorization", "Bearer wrong")]),
here(),
Instant::now()
),
Err(Refusal::BadToken)
);
assert_eq!(
guard.admit(
&headers(&[host, ("authorization", "Bearer s3cret-token-valu")]),
here(),
Instant::now()
),
Err(Refusal::BadToken)
);
}
#[test]
fn comparing_a_secret_folds_the_length_in_rather_than_checking_it_first() {
assert!(same_secret(b"abc", b"abc"));
assert!(!same_secret(b"abc", b"abd"));
assert!(!same_secret(b"ab", b"abc"));
assert!(!same_secret(b"abcd", b"abc"));
assert!(same_secret(b"", b""));
assert!(!same_secret(b"abc", b"abc\0"));
}
#[test]
fn the_rate_limit_triggers_and_then_recovers() {
let mut limiter = RateLimiter::default();
let start = Instant::now();
for i in 0..RATE_BURST {
assert!(
limiter.allow(here(), start),
"request {i} is inside the burst"
);
}
assert!(
!limiter.allow(here(), start),
"and one more is not, at the same instant"
);
let later = start + RATE_WINDOW / RATE_BURST + Duration::from_millis(1);
assert!(limiter.allow(here(), later), "a bucket refills with time");
assert!(!limiter.allow(here(), later), "one at a time, though");
assert!(limiter.allow(IpAddr::from([127, 0, 0, 2]), start));
}
#[test]
fn an_address_that_has_gone_quiet_is_forgotten() {
let mut limiter = RateLimiter::default();
let start = Instant::now();
assert!(limiter.allow(here(), start));
assert_eq!(limiter.seen.len(), 1, "while it is still spending");
assert!(limiter.allow(IpAddr::from([127, 0, 0, 2]), start + RATE_WINDOW * 2));
assert_eq!(
limiter.seen.keys().collect::<Vec<_>>(),
vec![&IpAddr::from([127, 0, 0, 2])],
"a server up for a year should not hold one entry per address that ever reached it"
);
}
#[test]
fn neither_the_guard_nor_the_settings_print_the_token() {
let guard = guard(Some("s3cret-token-value"), &[]);
assert!(
!format!("{guard:?}").contains("s3cret-token-value"),
"the guard printed its token: {guard:?}"
);
}
#[test]
fn an_origin_is_compared_as_an_origin_and_not_as_a_string() {
assert_eq!(
normalise_origin("https://App.Example"),
"https://app.example"
);
assert_eq!(
normalise_origin("https://app.example:443"),
"https://app.example",
"the default port for the scheme is not part of the origin"
);
assert_eq!(normalise_origin("http://localhost:80"), "http://localhost");
assert_eq!(
normalise_origin("http://localhost:3000/some/page"),
"http://localhost:3000",
"a path is not part of an origin either"
);
assert_eq!(normalise_origin("null"), "null");
let guard = guard(None, &["https://app.example"]);
assert_eq!(
guard.admit(
&headers(&[("host", "localhost"), ("origin", "https://App.Example:443")]),
here(),
Instant::now()
),
Ok(()),
"listing an origin lists it however a browser spells it"
);
}
#[test]
fn a_host_header_is_matched_without_its_port_or_its_case() {
assert_eq!(normalise_host("LocalHost:8449"), "localhost");
assert_eq!(normalise_host("127.0.0.1"), "127.0.0.1");
assert_eq!(normalise_host("[::1]:8449"), "[::1]");
assert_eq!(normalise_host("[FD7A:115C:A1E0::1]"), "[fd7a:115c:a1e0::1]");
assert_eq!(
normalise_host(" example-tailnet.ts.net "),
"example-tailnet.ts.net"
);
}
}