use std::collections::HashMap;
use std::time::Duration;
use anyhow::{Context, Result, bail};
use bytes::Bytes;
use http_body_util::BodyExt;
use hyper_util::client::legacy::Client;
use hyper_util::rt::TokioExecutor;
use kanade_shared::nats_client::{CredentialKind, CredentialProbe, NatsRole, parse_client_name};
use serde::Deserialize;
use sqlx::{Row, SqlitePool};
use tracing::{debug, info, warn};
const POLL_INTERVAL: Duration = Duration::from_secs(60);
const PAGE_SIZE: usize = 1024;
const MAX_PAGES: usize = 64;
const POLL_TIMEOUT: Duration = Duration::from_secs(20);
const MAX_BODY_BYTES: usize = 16 * 1024 * 1024;
const REDACTED_BY_BROKER: &str = "[REDACTED]";
pub const LABEL_SHARED_TOKEN: &str = "shared-token";
pub const LABEL_NO_AUTH: &str = "no-auth";
pub const LABEL_UNKNOWN: &str = "unknown";
#[derive(Debug, Deserialize)]
struct Connz {
#[serde(default)]
total: usize,
#[serde(default)]
connections: Vec<ConnInfo>,
}
#[derive(Debug, Deserialize)]
struct ConnInfo {
#[serde(default)]
cid: u64,
#[serde(default)]
name: Option<String>,
#[serde(default)]
authorized_user: Option<String>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum Evidence {
UsersMode,
TokenMode,
Unproven,
}
fn evidence(probe: &CredentialProbe, connected: bool) -> Evidence {
if !connected {
return Evidence::Unproven;
}
match probe.kind() {
CredentialKind::User => Evidence::UsersMode,
CredentialKind::Token => Evidence::TokenMode,
CredentialKind::None => Evidence::Unproven,
}
}
fn classify(authorized_user: Option<&str>, probe: &CredentialProbe, ev: Evidence) -> String {
match authorized_user.map(str::trim).filter(|s| !s.is_empty()) {
None => LABEL_NO_AUTH.to_string(),
Some(u) if u == REDACTED_BY_BROKER => LABEL_SHARED_TOKEN.to_string(),
Some(u) if probe.is_ours(u) => LABEL_SHARED_TOKEN.to_string(),
Some(u) if ev == Evidence::UsersMode => u.to_string(),
Some(_) => LABEL_UNKNOWN.to_string(),
}
}
#[derive(Debug, Default, PartialEq, Eq)]
struct Anomalies {
non_agent: Vec<String>,
conflicting: Vec<String>,
}
impl Anomalies {
fn is_empty(&self) -> bool {
self.non_agent.is_empty() && self.conflicting.is_empty()
}
}
fn correlate(
conns: &[ConnInfo],
probe: &CredentialProbe,
ev: Evidence,
) -> (HashMap<String, String>, Anomalies) {
let mut best: HashMap<String, (u64, String)> = HashMap::new();
let mut anomalies = Anomalies::default();
for c in conns {
let Some(name) = c.name.as_deref() else {
continue;
};
let Some(parsed) = parse_client_name(name) else {
continue;
};
let Some(pc_id) = parsed.identity else {
continue;
};
if parsed.role != NatsRole::Agent.as_str() {
anomalies.non_agent.push(name.to_string());
continue;
}
let label = classify(c.authorized_user.as_deref(), probe, ev);
match best.get(pc_id) {
Some((cid, existing)) if *cid >= c.cid => {
if existing != &label {
anomalies
.conflicting
.push(format!("{pc_id} (kept {existing}, ignored {label})"));
}
}
_ => {
best.insert(pc_id.to_string(), (c.cid, label));
}
}
}
anomalies.non_agent.sort();
anomalies.conflicting.sort();
let labels = best
.into_iter()
.map(|(pc, (_cid, label))| (pc, label))
.collect();
(labels, anomalies)
}
async fn apply(pool: &SqlitePool, labels: &HashMap<String, String>) -> Result<u64> {
if labels.is_empty() {
return Ok(0);
}
let current: HashMap<String, Option<String>> =
sqlx::query("SELECT pc_id, nats_user FROM agents")
.fetch_all(pool)
.await
.context("read current nats_user labels")?
.into_iter()
.map(|r| {
(
r.try_get::<String, _>("pc_id").unwrap_or_default(),
r.try_get::<Option<String>, _>("nats_user").ok().flatten(),
)
})
.collect();
let pending: Vec<(&String, &String)> = labels
.iter()
.filter(|(pc_id, label)| match current.get(pc_id.as_str()) {
None => false,
Some(existing) => existing.as_deref() != Some(label.as_str()),
})
.collect();
if pending.is_empty() {
return Ok(0);
}
let mut tx = pool.begin().await.context("begin nats_user tx")?;
let mut changed = 0;
for (pc_id, label) in pending {
let res = sqlx::query(
"UPDATE agents
SET nats_user = ?, nats_user_since = CURRENT_TIMESTAMP
WHERE pc_id = ?
AND (nats_user IS NULL OR nats_user <> ?)",
)
.bind(label)
.bind(pc_id)
.bind(label)
.execute(&mut *tx)
.await
.with_context(|| format!("update nats_user for {pc_id}"))?;
changed += res.rows_affected();
}
tx.commit().await.context("commit nats_user tx")?;
Ok(changed)
}
type MonitorClient =
Client<hyper_util::client::legacy::connect::HttpConnector, http_body_util::Empty<Bytes>>;
async fn fetch_page(client: &MonitorClient, base: &str, offset: usize) -> Result<Connz> {
let url = format!("{base}/connz?auth=1&subs=0&limit={PAGE_SIZE}&offset={offset}");
let uri: hyper::Uri = url.parse().with_context(|| format!("parse {url}"))?;
let res = client
.get(uri)
.await
.with_context(|| format!("GET {url}"))?;
let status = res.status();
if !status.is_success() {
bail!("GET {url} returned HTTP {status}");
}
let body = http_body_util::Limited::new(res.into_body(), MAX_BODY_BYTES)
.collect()
.await
.map_err(|e| anyhow::anyhow!("{e}"))
.with_context(|| format!("read body of {url} (cap {MAX_BODY_BYTES} bytes)"))?
.to_bytes();
serde_json::from_slice(&body).with_context(|| format!("decode /connz from {url}"))
}
async fn fetch_all(client: &MonitorClient, base: &str) -> Result<Vec<ConnInfo>> {
let mut out: Vec<ConnInfo> = Vec::new();
for page in 0..MAX_PAGES {
let z = fetch_page(client, base, out.len()).await?;
let got = z.connections.len();
out.extend(z.connections);
if got == 0 || out.len() >= z.total {
return Ok(out);
}
if page + 1 == MAX_PAGES {
warn!(
pages = MAX_PAGES,
fetched = out.len(),
total = z.total,
"/connz paging hit its cap; the tail of the connection list was not read",
);
}
}
Ok(out)
}
pub async fn run(pool: SqlitePool, monitor_url: String, nats: async_nats::Client) -> Result<()> {
let probe = CredentialProbe::for_role(NatsRole::Backend);
let client: MonitorClient = Client::builder(TokioExecutor::new()).build_http();
info!(
monitor_url = %monitor_url,
poll_secs = POLL_INTERVAL.as_secs(),
"nats connections projector started",
);
let mut healthy: Option<bool> = None;
let mut reported = Anomalies::default();
let mut tick = tokio::time::interval(POLL_INTERVAL);
tick.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
loop {
tick.tick().await;
let polled = match tokio::time::timeout(POLL_TIMEOUT, fetch_all(&client, &monitor_url))
.await
{
Ok(result) => result,
Err(_) => Err(anyhow::anyhow!(
"poll timed out after {POLL_TIMEOUT:?} (endpoint accepted the connection but did \
not finish answering)"
)),
};
match polled {
Ok(conns) => {
let seen = conns.len();
let ev = evidence(
&probe,
nats.connection_state() == async_nats::connection::State::Connected,
);
let (labels, anomalies) = correlate(&conns, &probe, ev);
let correlated = labels.len();
if anomalies != reported {
if !anomalies.non_agent.is_empty() {
warn!(
connections = ?anomalies.non_agent,
"NATS connections name a pc_id under a non-agent role; ignored",
);
}
if !anomalies.conflicting.is_empty() {
warn!(
hosts = ?anomalies.conflicting,
"two live NATS connections claim one pc_id with different credentials",
);
}
if anomalies.is_empty() {
info!("previously reported NATS connection anomalies have cleared");
}
reported = anomalies;
}
match apply(&pool, &labels).await {
Ok(changed) => {
if healthy != Some(true) {
info!(
monitor_url = %monitor_url,
connections = seen,
correlated,
"reading NATS connection credentials",
);
healthy = Some(true);
}
if changed > 0 {
info!(changed, correlated, "agent NATS credentials updated");
} else {
debug!(correlated, "agent NATS credentials unchanged");
}
}
Err(e) => warn!(error = %format!("{e:#}"), "nats_user projection write failed"),
}
}
Err(e) => {
if healthy != Some(false) {
warn!(
error = %format!("{e:#}"),
monitor_url = %monitor_url,
"NATS monitoring endpoint unreadable; agent NATS credentials will not be \
updated until it recovers",
);
healthy = Some(false);
}
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use sqlx::sqlite::SqlitePoolOptions;
fn parse(json: &str) -> Connz {
serde_json::from_str(json).expect("decode /connz")
}
fn probe_holding(token: Option<&str>) -> CredentialProbe {
CredentialProbe::from_token(token.map(str::to_string))
}
#[test]
fn a_token_authenticated_connection_is_the_shared_token() {
let probe = probe_holding(Some("fleet-secret"));
let ev = Evidence::TokenMode;
assert_eq!(classify(Some("[REDACTED]"), &probe, ev), LABEL_SHARED_TOKEN);
assert_eq!(
classify(Some("fleet-secret"), &probe, ev),
LABEL_SHARED_TOKEN,
);
assert_eq!(
classify(Some("[REDACTED]"), &probe_holding(None), Evidence::Unproven),
LABEL_SHARED_TOKEN,
);
}
#[test]
fn a_credential_we_cannot_vouch_for_is_never_stored_verbatim() {
let probe = probe_holding(Some("fleet-secret"));
let out = classify(Some("some-other-secret"), &probe, Evidence::TokenMode);
assert_eq!(out, LABEL_UNKNOWN);
assert!(!out.contains("secret"), "{out}");
}
#[test]
fn a_username_is_kept_only_against_positive_proof_of_users_mode() {
let probe = CredentialProbe::from_user("kanade-backend");
assert_eq!(
classify(Some("kanade-agent"), &probe, Evidence::UsersMode),
"kanade-agent",
);
assert_eq!(
classify(
Some("fleet-secret"),
&probe_holding(None),
Evidence::Unproven
),
LABEL_UNKNOWN,
);
assert_eq!(
classify(Some("fleet-secret"), &probe, Evidence::Unproven),
LABEL_UNKNOWN,
);
}
#[test]
fn evidence_comes_from_the_live_connection_not_the_resolved_credential() {
let token = probe_holding(Some("fleet-secret"));
let user = CredentialProbe::from_user("kanade-backend");
let none = probe_holding(None);
assert_eq!(evidence(&token, true), Evidence::TokenMode);
assert_eq!(evidence(&user, true), Evidence::UsersMode);
assert_eq!(evidence(&none, true), Evidence::Unproven);
for p in [&token, &user, &none] {
assert_eq!(
evidence(p, false),
Evidence::Unproven,
"a credential proves nothing while the link is down",
);
}
}
#[test]
fn no_credential_at_all_is_its_own_state() {
let probe = probe_holding(None);
for reported in [None, Some(""), Some(" ")] {
assert_eq!(
classify(reported, &probe, Evidence::Unproven),
LABEL_NO_AUTH,
"{reported:?} means the broker authenticated nobody",
);
}
let probe = probe_holding(Some("fleet-secret"));
assert_eq!(classify(None, &probe, Evidence::TokenMode), LABEL_NO_AUTH);
}
#[test]
fn only_agent_connections_carrying_a_pc_id_are_correlated() {
let probe = probe_holding(Some("fleet-secret"));
let z = parse(
r#"{"total":5,"connections":[
{"cid":1,"name":"kanade-agent/PC001","authorized_user":"[REDACTED]"},
{"cid":2,"name":"kanade-backend","authorized_user":"[REDACTED]"},
{"cid":3,"name":"kanade-agent","authorized_user":"[REDACTED]"},
{"cid":4,"name":"NATS CLI Version 0.1.5","authorized_user":"[REDACTED]"},
{"cid":5,"authorized_user":"[REDACTED]"}
]}"#,
);
let (out, anomalies) = correlate(&z.connections, &probe, Evidence::TokenMode);
assert_eq!(out.len(), 1, "only the named agent connection: {out:?}");
assert!(anomalies.is_empty(), "{anomalies:?}");
assert_eq!(
out.get("PC001").map(String::as_str),
Some(LABEL_SHARED_TOKEN)
);
assert!(!out.contains_key("kanade-agent"));
}
#[test]
fn the_newest_connection_wins_when_two_claim_one_host() {
let probe = CredentialProbe::from_user("kanade-backend");
let z = parse(
r#"{"total":2,"connections":[
{"cid":7,"name":"kanade-agent/PC001","authorized_user":"legacy"},
{"cid":9,"name":"kanade-agent/PC001","authorized_user":"kanade-agent"}
]}"#,
);
let (out, _) = correlate(&z.connections, &probe, Evidence::UsersMode);
assert_eq!(out.get("PC001").map(String::as_str), Some("kanade-agent"));
let z = parse(
r#"{"total":2,"connections":[
{"cid":9,"name":"kanade-agent/PC001","authorized_user":"kanade-agent"},
{"cid":7,"name":"kanade-agent/PC001","authorized_user":"legacy"}
]}"#,
);
let (out, _) = correlate(&z.connections, &probe, Evidence::UsersMode);
assert_eq!(out.get("PC001").map(String::as_str), Some("kanade-agent"));
}
async fn canned_connz(
pages: Vec<String>,
) -> (String, std::sync::Arc<std::sync::Mutex<Vec<String>>>) {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
let queries = std::sync::Arc::new(std::sync::Mutex::new(Vec::new()));
let seen = queries.clone();
tokio::spawn(async move {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
for body in pages {
let Ok((mut sock, _)) = listener.accept().await else {
return;
};
let mut buf = [0u8; 2048];
let n = sock.read(&mut buf).await.unwrap_or(0);
let req = String::from_utf8_lossy(&buf[..n]).to_string();
if let Some(line) = req.lines().next() {
seen.lock().unwrap().push(line.to_string());
}
let res = format!(
"HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{body}",
body.len(),
);
let _ = sock.write_all(res.as_bytes()).await;
let _ = sock.shutdown().await;
}
});
(format!("http://127.0.0.1:{port}"), queries)
}
#[tokio::test]
async fn a_fleet_larger_than_one_page_is_walked_to_the_end() {
let page = |cids: &[u64]| {
let conns: Vec<String> = cids
.iter()
.map(|c| {
format!(
r#"{{"cid":{c},"name":"kanade-agent/PC{c:03}","authorized_user":"[REDACTED]"}}"#
)
})
.collect();
format!(r#"{{"total":3,"connections":[{}]}}"#, conns.join(","))
};
let (base, queries) = canned_connz(vec![page(&[1, 2]), page(&[3])]).await;
let client: MonitorClient = Client::builder(TokioExecutor::new()).build_http();
let conns = fetch_all(&client, &base).await.expect("walk both pages");
assert_eq!(conns.len(), 3, "the tail of the fleet must not be dropped");
let queries = queries.lock().unwrap().clone();
assert_eq!(queries.len(), 2, "{queries:?}");
assert!(queries[0].contains("offset=0"), "{}", queries[0]);
assert!(queries[1].contains("offset=2"), "{}", queries[1]);
assert!(queries.iter().all(|q| q.contains("auth=1")), "{queries:?}");
}
#[tokio::test]
async fn a_total_the_broker_will_not_serve_still_terminates() {
let (base, queries) = canned_connz(vec![
r#"{"total":9999,"connections":[{"cid":1,"name":"kanade-agent/PC001"}]}"#.to_string(),
r#"{"total":9999,"connections":[]}"#.to_string(),
])
.await;
let client: MonitorClient = Client::builder(TokioExecutor::new()).build_http();
let conns = fetch_all(&client, &base).await.expect("walk terminates");
assert_eq!(conns.len(), 1);
assert_eq!(
queries.lock().unwrap().len(),
2,
"an empty page ends the walk"
);
}
#[tokio::test]
async fn a_hung_endpoint_loses_one_poll_rather_than_the_projector() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
let _accepting = tokio::spawn(async move {
let mut held = Vec::new();
while let Ok((sock, _)) = listener.accept().await {
held.push(sock);
}
});
let client: MonitorClient = Client::builder(TokioExecutor::new()).build_http();
let base = format!("http://127.0.0.1:{port}");
let out = tokio::time::timeout(Duration::from_millis(300), fetch_all(&client, &base)).await;
assert!(
out.is_err(),
"a hung endpoint must be interruptible by the caller's deadline, not answer",
);
}
async fn pool_with(pcs: &[&str]) -> SqlitePool {
let pool = SqlitePoolOptions::new()
.max_connections(1)
.connect("sqlite::memory:")
.await
.unwrap();
sqlx::migrate!("./migrations").run(&pool).await.unwrap();
for pc in pcs {
sqlx::query("INSERT INTO agents (pc_id) VALUES (?)")
.bind(pc)
.execute(&pool)
.await
.unwrap();
}
pool
}
async fn read(pool: &SqlitePool, pc: &str) -> (Option<String>, Option<String>) {
let r = sqlx::query("SELECT nats_user, nats_user_since FROM agents WHERE pc_id = ?")
.bind(pc)
.fetch_one(pool)
.await
.unwrap();
(r.try_get(0).unwrap(), r.try_get(1).unwrap())
}
#[tokio::test]
async fn a_host_that_was_never_seen_stays_null() {
let pool = pool_with(&["PC001", "PC002"]).await;
let labels = HashMap::from([("PC001".to_string(), LABEL_SHARED_TOKEN.to_string())]);
assert_eq!(apply(&pool, &labels).await.unwrap(), 1);
assert_eq!(
read(&pool, "PC001").await.0.as_deref(),
Some(LABEL_SHARED_TOKEN)
);
assert_eq!(read(&pool, "PC002").await.0, None);
}
#[tokio::test]
async fn an_unchanged_label_is_not_rewritten() {
let pool = pool_with(&["PC001"]).await;
let labels = HashMap::from([("PC001".to_string(), LABEL_SHARED_TOKEN.to_string())]);
apply(&pool, &labels).await.unwrap();
let (_, first_since) = read(&pool, "PC001").await;
assert!(first_since.is_some());
assert_eq!(apply(&pool, &labels).await.unwrap(), 0);
assert_eq!(read(&pool, "PC001").await.1, first_since);
}
#[tokio::test]
async fn moving_to_a_new_credential_restamps_since() {
let pool = pool_with(&["PC001"]).await;
apply(
&pool,
&HashMap::from([("PC001".to_string(), LABEL_SHARED_TOKEN.to_string())]),
)
.await
.unwrap();
sqlx::query(
"UPDATE agents SET nats_user_since = '2000-01-01 00:00:00' WHERE pc_id='PC001'",
)
.execute(&pool)
.await
.unwrap();
assert_eq!(
apply(
&pool,
&HashMap::from([("PC001".to_string(), "kanade-agent".to_string())]),
)
.await
.unwrap(),
1,
);
let (user, after) = read(&pool, "PC001").await;
assert_eq!(user.as_deref(), Some("kanade-agent"));
assert_ne!(
after.as_deref(),
Some("2000-01-01 00:00:00"),
"a credential change must restamp `since`",
);
}
#[tokio::test]
#[ignore = "requires nats-server in PATH; cargo test -- --ignored"]
async fn live_connz_reports_the_token_as_the_user_and_echoes_our_name() {
use std::io::Write as _;
const TOKEN: &str = "live-test-fleet-token";
let client_port = portpicker::pick_unused_port().expect("pick client port");
let http_port = portpicker::pick_unused_port().expect("pick monitor port");
let dir = tempfile::TempDir::new().expect("tempdir");
let conf = dir.path().join("nats.conf");
let mut f = std::fs::File::create(&conf).expect("write conf");
write!(
f,
"port: {client_port}\nhttp_port: {http_port}\nauthorization {{ token: \"{TOKEN}\" }}\n"
)
.expect("write conf");
drop(f);
let _server = tokio::process::Command::new("nats-server")
.arg("-c")
.arg(&conf)
.stdout(std::process::Stdio::null())
.stderr(std::process::Stdio::null())
.kill_on_drop(true)
.spawn()
.expect("spawn nats-server (is it in PATH?)");
let name = kanade_shared::nats_client::client_name(NatsRole::Agent, Some("PC-LIVE"));
let url = format!("nats://127.0.0.1:{client_port}");
let mut conn = None;
for _ in 0..50 {
match async_nats::ConnectOptions::new()
.token(TOKEN.to_string())
.name(name.clone())
.connect(&url)
.await
{
Ok(c) => {
conn = Some(c);
break;
}
Err(_) => tokio::time::sleep(Duration::from_millis(100)).await,
}
}
let _conn = conn.expect("nats-server did not come up in 5s");
let http: MonitorClient = Client::builder(TokioExecutor::new()).build_http();
let base = format!("http://127.0.0.1:{http_port}");
let conns = fetch_all(&http, &base).await.expect("read live /connz");
let ours = conns
.iter()
.find(|c| c.name.as_deref() == Some(name.as_str()))
.expect("our connection is in /connz under the name we announced");
let reported = ours
.authorized_user
.as_deref()
.expect("a token-authenticated connection reports an authorized_user");
assert!(
reported == REDACTED_BY_BROKER || reported == TOKEN,
"unexpected authorized_user shape under token auth; revisit classify()",
);
assert_eq!(
reported, REDACTED_BY_BROKER,
"nats-server {} redacts it; if a build stops doing that, the classifier still \
holds (it compares against our own credential) but this assertion documents when \
the behaviour changed",
"2.14.3",
);
let probe = probe_holding(Some(TOKEN));
let (labels, _) = correlate(&conns, &probe, Evidence::TokenMode);
assert_eq!(
labels.get("PC-LIVE").map(String::as_str),
Some(LABEL_SHARED_TOKEN),
);
assert!(
!labels.values().any(|v| v.contains(TOKEN)),
"the token must never reach a stored value: {labels:?}",
);
}
#[tokio::test]
async fn a_connection_for_an_unknown_pc_id_creates_no_row() {
let pool = pool_with(&["PC001"]).await;
let labels = HashMap::from([("GHOST".to_string(), LABEL_SHARED_TOKEN.to_string())]);
assert_eq!(apply(&pool, &labels).await.unwrap(), 0);
let n: i64 = sqlx::query("SELECT COUNT(*) FROM agents")
.fetch_one(&pool)
.await
.unwrap()
.try_get(0)
.unwrap();
assert_eq!(n, 1);
}
}