mod logfmt;
mod mcp;
mod proxy;
mod secret;
mod ws;
use std::collections::{BTreeMap, HashMap, HashSet};
use std::net::{IpAddr, Ipv4Addr};
use std::sync::{Arc, Mutex};
use std::time::{Duration, Instant};
use http_body_util::{BodyExt, Full};
use hyper::body::Bytes;
use hyper::service::service_fn;
use hyper::{Method, Request, Response, StatusCode};
use hyper_util::client::legacy::connect::HttpConnector;
use hyper_util::client::legacy::Client;
use hyper_util::rt::{TokioExecutor, TokioIo};
use serde_json::{json, Value};
use logfmt::log;
pub type BoxedBody =
http_body_util::combinators::BoxBody<Bytes, Box<dyn std::error::Error + Send + Sync>>;
pub type HttpClient = Client<HttpConnector, Full<Bytes>>;
#[derive(Clone)]
pub struct HttpClients {
pooled: HttpClient,
fresh: HttpClient,
}
impl HttpClients {
fn new() -> Self {
HttpClients {
pooled: Client::builder(TokioExecutor::new()).build_http(),
fresh: Client::builder(TokioExecutor::new())
.pool_max_idle_per_host(0)
.build_http(),
}
}
async fn send_once(
&self,
req: Request<Full<Bytes>>,
t: Duration,
) -> Result<hyper::Response<hyper::body::Incoming>, String> {
tokio::time::timeout(t, self.pooled.request(req))
.await
.map_err(|_| format!("timeout after {:.1}s", t.as_secs_f64()))?
.map_err(|e| e.to_string())
}
async fn send(
&self,
build: impl Fn() -> Result<Request<Full<Bytes>>, String>,
t: Duration,
) -> Result<hyper::Response<hyper::body::Incoming>, String> {
let first = tokio::time::timeout(t, self.pooled.request(build()?))
.await
.map_err(|_| format!("timeout after {:.1}s", t.as_secs_f64()))?;
let e = match first {
Ok(r) => return Ok(r),
Err(e) if e.is_connect() => return Err(e.to_string()),
Err(e) => e,
};
match tokio::time::timeout(t, self.fresh.request(build()?)).await {
Err(_) => Err(format!("timeout after {:.1}s", t.as_secs_f64())),
Ok(Ok(r)) => Ok(r),
Ok(Err(retry)) => Err(format!("{retry} (first attempt: {e})")),
}
}
}
pub struct Config {
pub daemons: Vec<String>,
pub port: u16,
pub announce_port: u16,
pub poll_interval: Duration,
pub probe_interval: Duration,
pub probe_timeout: Duration,
pub probe_fresh: Duration,
pub probe_promote: Duration,
pub serving_timeout: Duration,
pub mcp_timeout: Duration,
pub tools_ttl: Duration,
pub allowed_sources: Vec<String>,
pub discover_peers: bool,
}
fn env_str(name: &str, default: &str) -> String {
std::env::var(name)
.ok()
.filter(|s| !s.trim().is_empty())
.unwrap_or_else(|| default.to_string())
}
fn env_secs(name: &str, default: f64) -> Duration {
let v = std::env::var(name)
.ok()
.and_then(|s| s.trim().parse::<f64>().ok())
.filter(|v| *v > 0.0)
.unwrap_or(default);
Duration::from_secs_f64(v)
}
impl Config {
fn from_env() -> Config {
let probe_interval = env_secs("PROBE_INTERVAL_S", 5.0);
let probe_timeout = env_secs("PROBE_TIMEOUT_S", 3.0);
Config {
daemons: std::env::var("MENTAT_DAEMONS")
.unwrap_or_else(|_| "127.0.0.1:6380".to_string())
.split(',')
.map(|s| s.trim().to_string())
.filter(|s| !s.is_empty())
.collect(),
port: env_str("SERVE_PORT", "6381").parse().unwrap_or(6381),
announce_port: env_str("MENTAT_ANNOUNCE_PORT", "6382")
.parse()
.unwrap_or(6382),
poll_interval: env_secs("POLL_INTERVAL_S", 10.0),
probe_interval,
probe_timeout,
probe_fresh: env_secs(
"PROBE_FRESH_S",
(probe_interval * 3 + probe_timeout).as_secs_f64(),
),
probe_promote: env_secs("PROBE_PROMOTE_S", (probe_interval * 6).as_secs_f64()),
serving_timeout: env_secs("SERVING_TIMEOUT_S", 1800.0),
mcp_timeout: env_secs("MCP_TIMEOUT_S", 180.0),
tools_ttl: env_secs("TOOLS_TTL_S", 60.0),
allowed_sources: env_str("ALLOWED_SOURCES", "10.100.0.,192.168.1.,127.0.0.1,::1,172.")
.split(',')
.map(|s| s.trim().to_string())
.filter(|s| !s.is_empty())
.collect(),
discover_peers: env_str("DISCOVER_PEERS", "1") == "1",
}
}
}
pub struct DaemonView {
pub status: Option<Value>,
pub seen: Option<Instant>,
pub error: Option<String>,
}
pub struct ProbeResult {
pub ok: bool,
pub models: Vec<String>,
pub seen: Instant,
pub error: Option<String>,
pub selected: Option<String>,
pub promoted_at: Instant,
}
pub struct Shared {
pub cfg: Config,
pub started: Instant,
pub client: HttpClients,
pub daemons: Mutex<HashMap<String, DaemonView>>,
pub watched: Mutex<HashSet<String>>,
pub probes: Mutex<HashMap<String, ProbeResult>>,
pub tools: Mutex<HashMap<String, (Instant, Vec<Value>)>>,
pub refresh: tokio::sync::Notify,
}
#[derive(Clone)]
pub struct Endpoint {
pub candidates: Vec<String>,
pub announced: String,
pub note: Option<String>,
}
impl Endpoint {
pub fn best(&self) -> Option<&str> {
self.candidates.first().map(String::as_str)
}
}
#[derive(Clone)]
pub struct GroupEntry {
pub group: String,
pub daemon: String,
pub agents_alive: usize,
pub running: usize,
pub openai: Option<Endpoint>,
pub mcp: Option<Endpoint>,
}
fn collect_node_addrs(snap: &Value, out: &mut HashMap<String, Vec<String>>) {
let mut record = |ids: Vec<&str>, addrs: Vec<String>| {
if addrs.is_empty() {
return;
}
for id in ids.into_iter().filter(|s| !s.is_empty()) {
match out.get(id) {
Some(prev) if prev.len() >= addrs.len() => {}
_ => {
out.insert(id.to_string(), addrs.clone());
}
}
}
};
let listed = |v: &Value| -> Vec<String> {
v.as_array()
.into_iter()
.flatten()
.filter_map(|a| a.as_str())
.map(str::to_string)
.collect()
};
let own = listed(&snap["addrs"]);
let own_ip = snap["node_ip"].as_str().unwrap_or_default();
let mut own_ids: Vec<&str> = vec![own_ip];
own_ids.extend(own.iter().map(String::as_str));
record(
own_ids,
if own.is_empty() {
vec![own_ip.to_string()]
} else {
own.clone()
},
);
for (_, p) in snap["peers"].as_object().into_iter().flatten() {
let addrs = listed(&p["addrs"]);
let ip = p["node_ip"].as_str().unwrap_or_default();
let mut ids: Vec<&str> = vec![ip, p["link_ip"].as_str().unwrap_or_default()];
ids.extend(addrs.iter().map(String::as_str));
record(
ids,
if addrs.is_empty() {
vec![ip.to_string()]
} else {
addrs.clone()
},
);
}
}
fn endpoint_of(
agent: &Value,
svc: &str,
nodes: &HashMap<String, Vec<String>>,
allowed: &[String],
subnets: &[(u32, u32)],
) -> Option<Endpoint> {
let note = agent["service_notes"][svc]
.as_str()
.filter(|n| !n.is_empty())
.map(str::to_string);
if let Some(url) = agent["services"][svc].as_str().filter(|u| !u.is_empty()) {
return Some(Endpoint {
candidates: vec![url.to_string()],
announced: url.to_string(),
note,
});
}
let sp = &agent["services_ports"][svc];
let port = sp["port"].as_u64()?;
let path = sp["path"].as_str().unwrap_or_default();
let node_ip = agent["node_ip"].as_str().unwrap_or_default();
let mut hosts: Vec<String> = nodes
.get(node_ip)
.cloned()
.unwrap_or_else(|| vec![node_ip.to_string()]);
hosts.retain(|h| !h.is_empty() && prefix_allowed(allowed, h));
let mut seen = HashSet::new();
hosts.retain(|h| seen.insert(h.clone()));
hosts.sort_by_key(|h| !on_local_subnet(h, subnets));
Some(Endpoint {
candidates: hosts
.iter()
.map(|h| format!("http://{h}:{port}{path}"))
.collect(),
announced: format!("port {port}{path} on node {node_ip}"),
note,
})
}
pub fn group_table(shared: &Shared) -> BTreeMap<String, GroupEntry> {
let stale = shared.cfg.poll_interval * 3;
let subnets = local_subnets();
let daemons = shared.daemons.lock().unwrap();
let mut nodes: HashMap<String, Vec<String>> = HashMap::new();
for view in daemons.values() {
if let Some(snap) = view.status.as_ref() {
collect_node_addrs(snap, &mut nodes);
}
}
let mut out: BTreeMap<String, GroupEntry> = BTreeMap::new();
for (addr, view) in daemons.iter() {
let fresh = view.seen.map(|s| s.elapsed() <= stale).unwrap_or(false);
let Some(snap) = view.status.as_ref().filter(|_| fresh) else {
continue;
};
for (name, g) in snap["groups"].as_object().into_iter().flatten() {
let agents: Vec<&Value> = g["agents"]
.as_array()
.into_iter()
.flatten()
.filter(|a| a["alive"].as_bool().unwrap_or(false))
.collect();
let running = g["actors"]
.as_array()
.into_iter()
.flatten()
.filter(|a| a["state"].as_str() == Some("running"))
.count();
let resolve = |a: &Value, svc: &str| {
endpoint_of(a, svc, &nodes, &shared.cfg.allowed_sources, &subnets)
};
let openai = agents
.iter()
.filter_map(|a| resolve(a, "openai"))
.min_by(|x, y| x.best().cmp(&y.best()));
let mcp = agents
.iter()
.filter_map(|a| {
resolve(a, "mcp").map(|m| {
(
a["services"]["openai"].is_null()
&& a["services_ports"]["openai"].is_null(),
m,
)
})
})
.min_by(|x, y| (x.0, x.1.best()).cmp(&(y.0, y.1.best())))
.map(|(_, m)| m);
let entry = GroupEntry {
group: name.clone(),
daemon: addr.clone(),
agents_alive: agents.len(),
running,
openai,
mcp,
};
let replace = match out.get(name) {
None => true,
Some(prev) => {
(entry.running, entry.agents_alive) > (prev.running, prev.agents_alive)
|| ((entry.running, entry.agents_alive)
== (prev.running, prev.agents_alive)
&& entry.daemon < prev.daemon)
}
};
if replace {
out.insert(name.clone(), entry);
}
}
}
out
}
pub fn health_of(shared: &Shared, e: &GroupEntry) -> Result<Vec<String>, String> {
let Some(ep) = e.openai.as_ref() else {
return Err("no announced OpenAI endpoint".into());
};
if ep.candidates.is_empty() {
return Err(format!(
"announced {}, and no address of that node passes ALLOWED_SOURCES",
ep.announced
));
}
if e.running == 0 {
return Err("no running actors".into());
}
let probes = shared.probes.lock().unwrap();
match probes.get(&e.group) {
None => Err("not probed yet".into()),
Some(p) if !p.ok => Err(format!(
"endpoint probe failed: {}{}",
p.error.as_deref().unwrap_or("unknown"),
match &ep.note {
Some(n) => format!(" (agent reports: {n})"),
None => String::new(),
}
)),
Some(p) if p.seen.elapsed() > shared.cfg.probe_fresh => Err("endpoint probe stale".into()),
Some(p) => Ok(p.models.clone()),
}
}
pub fn endpoint_url(shared: &Shared, e: &GroupEntry) -> Option<String> {
let ep = e.openai.as_ref()?;
shared
.probes
.lock()
.unwrap()
.get(&e.group)
.and_then(|p| p.selected.clone())
.or_else(|| ep.best().map(str::to_string))
}
pub fn model_table(shared: &Shared) -> BTreeMap<String, (String, String)> {
let mut out = BTreeMap::new();
for e in group_table(shared).values() {
if let Ok(models) = health_of(shared, e) {
let url = endpoint_url(shared, e).unwrap_or_default();
for m in models {
out.entry(m)
.or_insert_with(|| (e.group.clone(), url.clone()));
}
}
}
out
}
pub fn not_ready(shared: &Shared) -> BTreeMap<String, String> {
let mut out = BTreeMap::new();
for e in group_table(shared).values() {
if let Err(why) = health_of(shared, e) {
out.insert(e.group.clone(), why);
}
}
out
}
pub fn status_view(shared: &Shared) -> Value {
let daemons: BTreeMap<String, Value> = shared
.daemons
.lock()
.unwrap()
.iter()
.map(|(addr, v)| {
(
addr.clone(),
json!({
"connected": v.status.is_some(),
"age_s": v.seen.map(|s| s.elapsed().as_secs()),
"error": v.error,
}),
)
})
.collect();
let groups: BTreeMap<String, Value> = group_table(shared)
.values()
.map(|e| {
let health = health_of(shared, e);
(
e.group.clone(),
json!({
"daemon": e.daemon,
"agents_alive": e.agents_alive,
"actors_running": e.running,
"openai": endpoint_url(shared, e),
"openai_candidates": e.openai.as_ref().map(|x| x.candidates.clone()),
"openai_note": e.openai.as_ref().and_then(|x| x.note.clone()),
"mcp": e.mcp.as_ref().and_then(|x| x.best()),
"healthy": health.is_ok(),
"models": health.as_ref().ok(),
"why_not": health.as_ref().err(),
}),
)
})
.collect();
let models: BTreeMap<String, Value> = model_table(shared)
.into_iter()
.map(|(m, (g, url))| (m, json!({ "group": g, "url": url })))
.collect();
json!({
"uptime_s": shared.started.elapsed().as_secs(),
"daemons": daemons,
"groups": groups,
"models": models,
})
}
pub fn ensure_watched(shared: &Arc<Shared>, addr: String) {
{
let mut w = shared.watched.lock().unwrap();
if !w.insert(addr.clone()) {
return;
}
}
log("daemon_watch", &[("daemon", addr.clone())]);
let shared = shared.clone();
tokio::spawn(async move { watch_daemon(shared, addr).await });
}
async fn watch_daemon(shared: Arc<Shared>, addr: String) {
loop {
poll_status(&shared, &addr).await;
match ws::EventStream::connect(&addr).await {
Ok(mut es) => loop {
match es.next(shared.cfg.poll_interval).await {
Ok(Some(_event)) => {
while let Ok(Some(_)) = es.next(Duration::from_millis(200)).await {}
poll_status(&shared, &addr).await;
}
Ok(None) => poll_status(&shared, &addr).await,
Err(e) => {
log(
"daemon_events_lost",
&[("daemon", addr.clone()), ("error", e.to_string())],
);
break;
}
}
},
Err(e) => {
let mut d = shared.daemons.lock().unwrap();
let v = d.entry(addr.clone()).or_insert(DaemonView {
status: None,
seen: None,
error: None,
});
v.error = Some(format!("events: {e}"));
}
}
tokio::time::sleep(Duration::from_secs(2)).await;
}
}
async fn udp_listener(shared: Arc<Shared>) {
let port = shared.cfg.announce_port;
if port == 0 {
return;
}
let sock = match tokio::net::UdpSocket::bind(("0.0.0.0", port)).await {
Ok(s) => s,
Err(e) => {
log(
"announce_listen_failed",
&[("port", port.to_string()), ("error", e.to_string())],
);
return;
}
};
let key = secret::load();
log(
"announce_listen",
&[
("port", port.to_string()),
(
"verify",
match key {
Some(_) => "required".to_string(),
None => "off (no MENTAT_SECRET)".to_string(),
},
),
],
);
let mut seen: HashMap<String, (String, u64)> = HashMap::new();
let mut warned: HashSet<String> = HashSet::new();
let mut noted: HashSet<String> = HashSet::new();
let mut chosen: HashMap<String, String> = HashMap::new();
let universe = secret::universe();
let mut buf = [0u8; 2048];
loop {
let Ok((n, src)) = sock.recv_from(&mut buf).await else {
continue;
};
match secret::peek_universe(&buf[..n]) {
Some(u) if u != universe => continue,
_ => {}
}
let v = match &key {
Some(k) => {
let Some(p) = secret::verify(&buf[..n], k) else {
if warned.insert(src.ip().to_string()) {
log(
"announce_rejected",
&[
("src", src.ip().to_string()),
("why", "bad signature or unsigned".to_string()),
],
);
}
continue;
};
if p["mentat_announce"].as_u64() != Some(secret::SIGNED_VERSION) {
continue;
}
let Some(t) = p["t"].as_f64() else { continue };
if !secret::fresh(t, secret::now_s()) {
continue;
}
let node = p["node_id"].as_str().unwrap_or_default().to_string();
let boot = p["boot_id"].as_str().unwrap_or_default().to_string();
let seq = p["seq"].as_u64().unwrap_or(0);
if node.is_empty() || boot.is_empty() {
continue;
}
match seen.get(&node) {
Some((b, last)) if *b == boot && seq <= *last => continue,
_ => seen.insert(node, (boot, seq)),
};
p
}
None => {
let Ok(p) = serde_json::from_slice::<Value>(&buf[..n]) else {
continue;
};
if p["mentat_announce"].as_u64() != Some(1) {
if p.get("sig").is_some() && warned.insert(src.ip().to_string()) {
log(
"announce_unverifiable",
&[
("src", src.ip().to_string()),
(
"why",
"signed announcement, no MENTAT_SECRET here".to_string(),
),
],
);
}
continue;
}
p
}
};
let Some(http) = v["http"].as_str() else {
continue;
};
let Some((http_ip, http_port)) = http.rsplit_once(':') else {
continue;
};
if http_port.parse::<u16>().map(|p| p == 0).unwrap_or(true) {
continue;
}
let src_ip = src.ip().to_string();
if !source_allowed(&shared.cfg, &src_ip) {
if warned.insert(src_ip.clone()) {
log(
"announce_source_not_allowed",
&[
("src", src_ip.clone()),
("allowed_sources", shared.cfg.allowed_sources.join(",")),
],
);
}
continue;
}
let node = v["node_id"].as_str().unwrap_or_default().to_string();
if !node.is_empty() && chosen.contains_key(&node) {
continue;
}
let ranked: Vec<String> = v["addrs"]
.as_array()
.into_iter()
.flatten()
.filter_map(|a| a.as_str())
.map(str::to_string)
.collect();
let pick = announce_address(
&ranked,
&src_ip,
&shared.cfg.allowed_sources,
&local_subnets(),
);
if pick != src_ip && noted.insert(src_ip.clone()) {
log(
"announce_preferred_addr",
&[
("src", src_ip.clone()),
("advertised", http.to_string()),
("watching", format!("{pick}:{http_port}")),
],
);
} else if http_ip != src_ip && noted.insert(src_ip.clone()) {
log(
"announce_addr_mismatch",
&[
("src", src_ip.clone()),
("advertised", http.to_string()),
("watching", format!("{pick}:{http_port}")),
],
);
}
if !node.is_empty() {
chosen.insert(node, pick.clone());
}
ensure_watched(&shared, format!("{pick}:{http_port}"));
}
}
fn local_subnets() -> Vec<(u32, u32)> {
let Ok(ifaces) = getifaddrs::InterfaceFilter::new().v4().get() else {
return Vec::new();
};
ifaces
.filter_map(|i| match (i.address.ip_addr(), i.address.netmask()) {
(Some(IpAddr::V4(a)), Some(IpAddr::V4(m))) => {
let (a, m) = (u32::from(a), u32::from(m));
Some((a & m, m))
}
_ => None,
})
.collect()
}
fn on_local_subnet(ip: &str, subnets: &[(u32, u32)]) -> bool {
let Ok(v4) = ip.parse::<Ipv4Addr>() else {
return false;
};
let a = u32::from(v4);
subnets.iter().any(|(net, mask)| a & mask == *net)
}
fn announce_address(
ranked: &[String],
src_ip: &str,
allowed: &[String],
subnets: &[(u32, u32)],
) -> String {
ranked
.iter()
.filter(|a| prefix_allowed(allowed, a))
.find(|a| on_local_subnet(a, subnets))
.cloned()
.unwrap_or_else(|| src_ip.to_string())
}
fn peer_address(p: &Value, subnets: &[(u32, u32)]) -> Option<String> {
let mut cands: Vec<String> = Vec::new();
let mut seen = HashSet::new();
let mut push = |v: Option<&str>| {
if let Some(s) = v.filter(|s| !s.is_empty()) {
if seen.insert(s.to_string()) {
cands.push(s.to_string());
}
}
};
push(p["link_ip"].as_str());
for a in p["addrs"].as_array().into_iter().flatten() {
push(a.as_str());
}
push(p["node_ip"].as_str());
cands
.iter()
.find(|c| on_local_subnet(c, subnets))
.or_else(|| cands.first())
.cloned()
}
async fn poll_status(shared: &Arc<Shared>, addr: &str) {
let url = format!("http://{addr}/status");
match http_get_json(&shared.client, &url, Duration::from_secs(5)).await {
Ok(snap) => {
if shared.cfg.discover_peers {
let subnets = local_subnets();
for (_, p) in snap["peers"].as_object().into_iter().flatten() {
let Some(port) = p["http_port"].as_u64() else {
continue;
};
if port == 0 {
continue;
}
if let Some(ip) = peer_address(p, &subnets) {
ensure_watched(shared, format!("{ip}:{port}"));
}
}
}
shared.daemons.lock().unwrap().insert(
addr.to_string(),
DaemonView {
status: Some(snap),
seen: Some(Instant::now()),
error: None,
},
);
shared.refresh.notify_one();
}
Err(e) => {
let mut d = shared.daemons.lock().unwrap();
let v = d.entry(addr.to_string()).or_insert(DaemonView {
status: None,
seen: None,
error: None,
});
v.error = Some(e);
}
}
}
async fn prober(shared: Arc<Shared>) {
loop {
let table = group_table(&shared);
{
let mut probes = shared.probes.lock().unwrap();
probes.retain(|k, _| {
table
.get(k)
.map(|e| e.openai.is_some() && e.running > 0)
.unwrap_or(false)
});
}
let mut set = tokio::task::JoinSet::new();
for e in table.values() {
if e.running == 0 {
continue;
}
let Some(ep) = e.openai.clone().filter(|x| !x.candidates.is_empty()) else {
continue;
};
let group = e.group.clone();
let client = shared.client.clone();
let t = shared.cfg.probe_timeout;
let (sticky, promote) = {
let probes = shared.probes.lock().unwrap();
match probes.get(&group) {
Some(p) => (
p.selected.clone(),
p.promoted_at.elapsed() >= shared.cfg.probe_promote,
),
None => (None, true),
}
};
set.spawn(async move {
let r = probe_candidates(&client, &ep.candidates, sticky, promote, t).await;
(group, r)
});
}
while let Some(Ok((group, (tried_top, res)))) = set.join_next().await {
let now = Instant::now();
let prev = {
let probes = shared.probes.lock().unwrap();
probes
.get(&group)
.map(|p| (p.ok, p.selected.clone(), p.promoted_at))
};
let (was_ok, was_sel, was_promoted) = match prev {
Some((a, b, c)) => (Some(a), b, c),
None => (None, None, now),
};
let pr = match res {
Ok((url, mut models)) => {
if models.is_empty() {
models.push(group.clone());
}
ProbeResult {
ok: true,
models,
seen: now,
error: None,
selected: Some(url),
promoted_at: if tried_top { now } else { was_promoted },
}
}
Err(e) => ProbeResult {
ok: false,
models: Vec::new(),
seen: now,
error: Some(e),
selected: None,
promoted_at: if tried_top { now } else { was_promoted },
},
};
if was_ok != Some(pr.ok) {
log(
"group_probe",
&[
("group", group.clone()),
("ok", pr.ok.to_string()),
("models", format!("{:?}", pr.models)),
("error", pr.error.clone().unwrap_or_default()),
],
);
}
if pr.ok && pr.selected != was_sel {
log(
"group_endpoint",
&[
("group", group.clone()),
("url", pr.selected.clone().unwrap_or_default()),
("previous", was_sel.unwrap_or_default()),
],
);
}
shared.probes.lock().unwrap().insert(group, pr);
}
tokio::select! {
_ = tokio::time::sleep(shared.cfg.probe_interval) => {}
_ = shared.refresh.notified() => {}
}
}
}
async fn probe_candidates(
client: &HttpClients,
candidates: &[String],
sticky: Option<String>,
promote: bool,
t: Duration,
) -> (bool, Result<(String, Vec<String>), String>) {
let at = sticky
.as_ref()
.and_then(|s| candidates.iter().position(|c| c == s));
let order: Vec<&String> = match at.filter(|_| !promote) {
Some(i) => std::iter::once(&candidates[i])
.chain(
candidates
.iter()
.enumerate()
.filter_map(|(j, c)| (j != i).then_some(c)),
)
.collect(),
None => candidates.iter().collect(),
};
let tried_top = order.first().copied() == candidates.first();
let mut errors: Vec<String> = Vec::new();
for base in order {
let url = format!("{}/models", base.trim_end_matches('/'));
match http_get_json(client, &url, t).await {
Ok(v) => {
let models: Vec<String> = v["data"]
.as_array()
.into_iter()
.flatten()
.filter_map(|m| m["id"].as_str().map(String::from))
.collect();
return (tried_top, Ok((base.clone(), models)));
}
Err(e) => errors.push(format!("{base}: {e}")),
}
}
(tried_top, Err(errors.join("; ")))
}
pub fn full_body(bytes: impl Into<Bytes>) -> BoxedBody {
Full::new(bytes.into()).map_err(|e| match e {}).boxed()
}
pub fn json_response(status: StatusCode, v: &Value) -> Response<BoxedBody> {
Response::builder()
.status(status)
.header(hyper::header::CONTENT_TYPE, "application/json")
.body(full_body(v.to_string()))
.expect("static response")
}
pub async fn http_get_json(client: &HttpClients, url: &str, t: Duration) -> Result<Value, String> {
http_json(
client,
|| {
Request::builder()
.method(Method::GET)
.uri(url)
.body(Full::new(Bytes::new()))
.map_err(|e| e.to_string())
},
t,
)
.await
}
pub async fn http_post_json(
client: &HttpClients,
url: &str,
body: &Value,
t: Duration,
) -> Result<Value, String> {
let req = Request::builder()
.method(Method::POST)
.uri(url)
.header(hyper::header::CONTENT_TYPE, "application/json")
.body(Full::new(Bytes::from(body.to_string())))
.map_err(|e| e.to_string())?;
read_json(client.send_once(req, t).await?, t).await
}
async fn http_json(
client: &HttpClients,
build: impl Fn() -> Result<Request<Full<Bytes>>, String>,
t: Duration,
) -> Result<Value, String> {
read_json(client.send(build, t).await?, t).await
}
async fn read_json(
resp: hyper::Response<hyper::body::Incoming>,
t: Duration,
) -> Result<Value, String> {
let status = resp.status();
let body = tokio::time::timeout(t, resp.into_body().collect())
.await
.map_err(|_| "timeout reading body".to_string())?
.map_err(|e| e.to_string())?
.to_bytes();
if !status.is_success() {
return Err(format!("HTTP {status}"));
}
serde_json::from_slice(&body).map_err(|e| e.to_string())
}
fn source_allowed(cfg: &Config, addr: &str) -> bool {
prefix_allowed(&cfg.allowed_sources, addr)
}
fn prefix_allowed(allowed: &[String], addr: &str) -> bool {
allowed.iter().any(|p| addr.starts_with(p))
}
async fn handle(
shared: Arc<Shared>,
peer_ip: String,
req: Request<hyper::body::Incoming>,
) -> Response<BoxedBody> {
if !source_allowed(&shared.cfg, &peer_ip) {
return json_response(
StatusCode::FORBIDDEN,
&json!({"error": "source not permitted"}),
);
}
let path = {
let p = req.uri().path().trim_end_matches('/');
if p.is_empty() { "/" } else { p }.to_string()
};
match (req.method().clone(), path.as_str()) {
(Method::GET, "/" | "/healthz" | "/status.json") => {
json_response(StatusCode::OK, &status_view(&shared))
}
(Method::GET, "/v1" | "/v1/models") => {
let data: Vec<Value> = model_table(&shared)
.iter()
.map(|(m, (g, _))| json!({"id": m, "object": "model", "owned_by": g}))
.collect();
json_response(StatusCode::OK, &json!({"object": "list", "data": data}))
}
(Method::POST, "/mcp") => mcp::handle(&shared, req).await,
(Method::POST, _) => proxy::forward(&shared, req).await,
_ => json_response(StatusCode::NOT_FOUND, &json!({"error": "not found"})),
}
}
#[tokio::main]
async fn main() {
if std::env::args().nth(1).as_deref() == Some("--version") {
println!("mentatd-serve {}", env!("CARGO_PKG_VERSION"));
return;
}
let cfg = Config::from_env();
let shared = Arc::new(Shared {
started: Instant::now(),
client: HttpClients::new(),
daemons: Mutex::new(HashMap::new()),
watched: Mutex::new(HashSet::new()),
probes: Mutex::new(HashMap::new()),
tools: Mutex::new(HashMap::new()),
refresh: tokio::sync::Notify::new(),
cfg,
});
log(
"serve_up",
&[
("port", shared.cfg.port.to_string()),
("daemons", shared.cfg.daemons.join(",")),
],
);
for d in shared.cfg.daemons.clone() {
ensure_watched(&shared, d);
}
{
let shared = shared.clone();
tokio::spawn(async move { prober(shared).await });
}
{
let shared = shared.clone();
tokio::spawn(async move { udp_listener(shared).await });
}
let listener = tokio::net::TcpListener::bind(("0.0.0.0", shared.cfg.port))
.await
.unwrap_or_else(|e| panic!("bind 0.0.0.0:{}: {e}", shared.cfg.port));
loop {
let Ok((stream, peer)) = listener.accept().await else {
continue;
};
let shared = shared.clone();
tokio::spawn(async move {
let peer_ip = peer.ip().to_string();
let svc = service_fn(move |req| {
let shared = shared.clone();
let peer_ip = peer_ip.clone();
async move { Ok::<_, std::convert::Infallible>(handle(shared, peer_ip, req).await) }
});
let _ = hyper::server::conn::http1::Builder::new()
.serve_connection(TokioIo::new(stream), svc)
.await;
});
}
}
#[cfg(test)]
mod tests {
use super::*;
fn subnets() -> Vec<(u32, u32)> {
let mask = u32::from(Ipv4Addr::new(255, 255, 255, 0));
vec![
(u32::from(Ipv4Addr::new(10, 0, 0, 0)), mask),
(u32::from(Ipv4Addr::new(192, 168, 1, 0)), mask),
]
}
async fn serves_once_per_connection() -> std::net::SocketAddr {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
loop {
let Ok((mut sock, _)) = listener.accept().await else {
return;
};
tokio::spawn(async move {
let mut buf = [0u8; 2048];
if sock.read(&mut buf).await.unwrap_or(0) == 0 {
return;
}
let body = b"{\"data\":[]}";
let head = format!(
"HTTP/1.1 200 OK\r\nContent-Type: application/json\r\n\
Content-Length: {}\r\nConnection: keep-alive\r\n\r\n",
body.len()
);
let _ = sock.write_all(head.as_bytes()).await;
let _ = sock.flush().await;
let _ = sock.write_all(body).await;
let _ = sock.flush().await;
let _ = sock.read(&mut buf).await;
});
}
});
addr
}
fn get(url: &str) -> Request<Full<Bytes>> {
Request::builder()
.method(Method::GET)
.uri(url)
.body(Full::new(Bytes::new()))
.unwrap()
}
#[tokio::test]
async fn without_a_retry_a_reused_connection_fails() {
let addr = serves_once_per_connection().await;
let url = format!("http://{addr}/v1/models");
let pooled: HttpClient = Client::builder(TokioExecutor::new()).build_http();
let first = pooled.request(get(&url)).await;
assert!(first.is_ok(), "first request: {first:?}");
let _ = first.unwrap().into_body().collect().await;
tokio::time::sleep(Duration::from_millis(150)).await;
let second = pooled.request(get(&url)).await;
let e = second.expect_err("reusing the dead connection should fail");
assert!(
!e.is_connect(),
"the endpoint is up, so this must not read as a connect failure: {e}"
);
}
#[tokio::test]
async fn a_stale_pooled_connection_retries_instead_of_failing() {
let addr = serves_once_per_connection().await;
let url = format!("http://{addr}/v1/models");
let t = Duration::from_secs(5);
let clients = HttpClients::new();
assert!(
http_get_json(&clients, &url, t).await.is_ok(),
"first probe"
);
tokio::time::sleep(Duration::from_millis(150)).await;
let second = http_get_json(&clients, &url, t).await;
assert!(
second.is_ok(),
"probe over a stale pooled connection: {second:?}"
);
}
#[test]
fn an_unlisted_identity_subnet_does_not_block_discovery() {
let allowed = vec!["192.168.1.".to_string()];
let subnets = subnets();
assert_eq!(
announce_address(&["10.100.0.2".into()], "192.168.1.77", &allowed, &subnets),
"192.168.1.77"
);
}
#[test]
fn an_advertised_candidate_is_still_gated() {
let subnets = subnets();
let allowed = vec!["10.0.0.".to_string()];
assert_eq!(
announce_address(&["10.0.0.7".into()], "10.0.0.1", &allowed, &subnets),
"10.0.0.7",
"allowed and local, so it is preferred"
);
assert_eq!(
announce_address(&["192.168.1.13".into()], "10.0.0.1", &allowed, &subnets),
"10.0.0.1",
"local but not allowed, so the source stands"
);
}
#[tokio::test]
async fn a_post_is_not_retried() {
let addr = serves_once_per_connection().await;
let url = format!("http://{addr}/mcp");
let t = Duration::from_secs(5);
let clients = HttpClients::new();
let body = serde_json::json!({"jsonrpc": "2.0", "method": "tools/list"});
assert!(
http_post_json(&clients, &url, &body, t).await.is_ok(),
"first post"
);
tokio::time::sleep(Duration::from_millis(150)).await;
let second = http_post_json(&clients, &url, &body, t).await;
assert!(
second.is_err(),
"a POST over a stale connection fails rather than replaying: {second:?}"
);
}
#[tokio::test]
async fn a_dead_endpoint_still_fails() {
let l = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = l.local_addr().unwrap();
drop(l);
let clients = HttpClients::new();
let r = http_get_json(
&clients,
&format!("http://{addr}/v1/models"),
Duration::from_secs(2),
)
.await;
assert!(r.is_err(), "a closed port must still fail");
}
#[test]
fn subnet_membership() {
let s = subnets();
assert!(on_local_subnet("10.0.0.7", &s));
assert!(on_local_subnet("192.168.1.13", &s));
assert!(!on_local_subnet("10.100.0.1", &s));
assert!(!on_local_subnet("not-an-ip", &s));
}
#[test]
fn a_reachable_addr_beats_an_unroutable_identity() {
let p = serde_json::json!({
"node_ip": "10.100.0.1",
"link_ip": "10.100.0.1",
"addrs": ["10.100.0.1", "192.168.1.11"],
});
assert_eq!(
peer_address(&p, &subnets()).as_deref(),
Some("192.168.1.11")
);
}
#[test]
fn link_ip_leads_when_nothing_is_local() {
let p = serde_json::json!({
"node_ip": "10.100.0.1",
"link_ip": "172.16.4.4",
"addrs": ["172.16.9.9"],
});
assert_eq!(peer_address(&p, &subnets()).as_deref(), Some("172.16.4.4"));
}
#[test]
fn an_advertised_address_is_still_a_claim() {
let allowed = vec!["127.".to_string()];
assert!(prefix_allowed(&allowed, "127.0.0.1"));
assert!(!prefix_allowed(&allowed, "192.168.1.109"));
assert!(!prefix_allowed(&allowed, "203.0.113.7"));
}
#[test]
fn a_port_announcement_resolves_to_every_address_of_its_node() {
let snap = serde_json::json!({
"node_ip": "10.100.0.1",
"addrs": ["10.100.0.1", "192.168.1.11"],
"peers": {},
});
let mut nodes = HashMap::new();
collect_node_addrs(&snap, &mut nodes);
let agent = serde_json::json!({
"node_ip": "10.100.0.1",
"services": {},
"services_ports": {"openai": {"port": 8000, "path": "/v1"}},
});
let allowed = vec!["10.".to_string(), "192.168.1.".to_string()];
let ep = endpoint_of(&agent, "openai", &nodes, &allowed, &subnets()).unwrap();
assert_eq!(
ep.candidates,
vec![
"http://192.168.1.11:8000/v1",
"http://10.100.0.1:8000/v1",
]
);
}
#[test]
fn a_verbatim_url_is_neither_re_derived_nor_gated() {
let agent = serde_json::json!({
"node_ip": "10.100.0.1",
"services": {"openai": "http://203.0.113.7:8000/v1"},
"services_ports": {},
});
let ep = endpoint_of(
&agent,
"openai",
&HashMap::new(),
&["10.".to_string()],
&subnets(),
)
.unwrap();
assert_eq!(ep.candidates, vec!["http://203.0.113.7:8000/v1"]);
}
#[test]
fn derived_addresses_are_gated_and_may_leave_nothing() {
let mut nodes = HashMap::new();
nodes.insert(
"10.100.0.1".to_string(),
vec!["10.100.0.1".to_string(), "192.168.1.11".to_string()],
);
let agent = serde_json::json!({
"node_ip": "10.100.0.1",
"services": {},
"services_ports": {"openai": {"port": 8000, "path": "/v1"}},
});
let ep = endpoint_of(&agent, "openai", &nodes, &["172.".to_string()], &subnets()).unwrap();
assert!(ep.candidates.is_empty(), "{:?}", ep.candidates);
assert!(ep.announced.contains("port 8000/v1"), "{}", ep.announced);
}
#[test]
fn an_agent_joins_its_node_by_any_of_its_addresses() {
let snap = serde_json::json!({
"node_ip": "192.168.1.13",
"addrs": ["192.168.1.13"],
"peers": {"n1": {
"node_ip": "10.100.0.1",
"link_ip": "192.168.1.11",
"addrs": ["192.168.1.11", "10.100.0.1"],
}},
});
let mut nodes = HashMap::new();
collect_node_addrs(&snap, &mut nodes);
for key in ["10.100.0.1", "192.168.1.11"] {
assert_eq!(
nodes.get(key).map(Vec::len),
Some(2),
"{key} must resolve to the peer's whole address list"
);
}
}
#[test]
fn the_fuller_description_of_a_node_wins() {
let mut nodes = HashMap::new();
collect_node_addrs(
&serde_json::json!({
"node_ip": "10.100.0.1", "addrs": [], "peers": {}
}),
&mut nodes,
);
collect_node_addrs(
&serde_json::json!({
"node_ip": "10.100.0.1",
"addrs": ["10.100.0.1", "192.168.1.11"],
"peers": {},
}),
&mut nodes,
);
assert_eq!(nodes["10.100.0.1"].len(), 2);
}
#[tokio::test]
async fn a_working_selection_is_probed_first_and_kept() {
let (top, low) = (
models_endpoint("model-x").await,
models_endpoint("model-x").await,
);
let c = vec![format!("http://{top}/v1"), format!("http://{low}/v1")];
let clients = HttpClients::new();
let (tried_top, r) = probe_candidates(
&clients,
&c,
Some(c[1].clone()),
false,
Duration::from_secs(2),
)
.await;
assert!(!tried_top, "the sticky candidate is not the top-ranked one");
assert_eq!(r.unwrap().0, c[1], "the sticky candidate must be kept");
}
#[tokio::test]
async fn a_dead_top_candidate_falls_through() {
let up = models_endpoint("model-x").await;
let dead = dead_addr().await;
let c = vec![format!("http://{dead}/v1"), format!("http://{up}/v1")];
let clients = HttpClients::new();
let (_, r) = probe_candidates(&clients, &c, None, false, Duration::from_secs(2)).await;
let (url, models) = r.expect("the second candidate answers");
assert_eq!(url, c[1]);
assert_eq!(models, vec!["model-x"]);
}
#[tokio::test]
async fn a_promotion_round_takes_the_preferred_address_back() {
let top = models_endpoint("model-x").await;
let low = models_endpoint("model-x").await;
let c = vec![format!("http://{top}/v1"), format!("http://{low}/v1")];
let clients = HttpClients::new();
let (tried_top, r) = probe_candidates(
&clients,
&c,
Some(c[1].clone()),
true,
Duration::from_secs(2),
)
.await;
assert!(tried_top);
assert_eq!(r.unwrap().0, c[0], "the preferred address answers again");
}
#[tokio::test]
async fn every_candidate_down_reports_every_candidate() {
let (d1, d2) = (dead_addr().await, dead_addr().await);
let c = vec![format!("http://{d1}/v1"), format!("http://{d2}/v1")];
let clients = HttpClients::new();
let (_, r) = probe_candidates(&clients, &c, None, false, Duration::from_secs(2)).await;
let e = r.expect_err("nothing answers");
assert!(
e.contains(&d1.to_string()) && e.contains(&d2.to_string()),
"{e}"
);
}
async fn models_endpoint(model: &str) -> std::net::SocketAddr {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let body = format!("{{\"data\":[{{\"id\":\"{model}\"}}]}}");
tokio::spawn(async move {
loop {
let Ok((mut sock, _)) = listener.accept().await else {
return;
};
let body = body.clone();
tokio::spawn(async move {
let mut buf = [0u8; 2048];
while sock.read(&mut buf).await.unwrap_or(0) > 0 {
let head = format!(
"HTTP/1.1 200 OK\r\nContent-Type: application/json\r\n\
Content-Length: {}\r\n\r\n",
body.len()
);
let _ = sock.write_all(head.as_bytes()).await;
let _ = sock.write_all(body.as_bytes()).await;
let _ = sock.flush().await;
}
});
}
});
addr
}
async fn dead_addr() -> std::net::SocketAddr {
let l = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = l.local_addr().unwrap();
drop(l);
addr
}
#[test]
fn an_old_daemon_reports_node_ip_alone() {
let p = serde_json::json!({"node_ip": "192.168.1.11"});
assert_eq!(
peer_address(&p, &subnets()).as_deref(),
Some("192.168.1.11")
);
assert_eq!(peer_address(&serde_json::json!({}), &subnets()), None);
}
}