use dashmap::DashMap;
use std::collections::HashMap;
use std::sync::atomic::{AtomicU64, AtomicU8, AtomicUsize, Ordering};
use std::sync::{Arc, Mutex as StdMutex};
use tokio::sync::{broadcast, watch, Mutex, RwLock};
use crate::ilink::{QrLoginUiEvent, UpstreamSink};
use crate::store::Store;
use super::*;
pub const MAX_CONCURRENT_POLLS_PER_VTOKEN: usize = 3;
pub const MAX_HUB_POLLS_DEFAULT: usize = 8192;
pub const HISTOGRAM_BUCKETS_MS: &[u64] = &[1, 5, 25, 100, 500, 2_500, 10_000];
pub const HISTOGRAM_BUCKET_COUNT: usize = HISTOGRAM_BUCKETS_MS.len() + 1;
#[derive(Debug)]
pub struct LatencyHistogram {
pub count: AtomicU64,
pub sum_us: AtomicU64,
pub buckets: [AtomicU64; HISTOGRAM_BUCKET_COUNT],
}
impl Default for LatencyHistogram {
fn default() -> Self {
Self::new()
}
}
impl LatencyHistogram {
pub fn new() -> Self {
Self {
count: AtomicU64::new(0),
sum_us: AtomicU64::new(0),
buckets: std::array::from_fn(|_| AtomicU64::new(0)),
}
}
pub fn observe(&self, elapsed: std::time::Duration) {
let us = u64::try_from(elapsed.as_micros()).unwrap_or(u64::MAX);
let ms = elapsed.as_millis() as u64;
self.count.fetch_add(1, Ordering::Relaxed);
self.sum_us.fetch_add(us, Ordering::Relaxed);
for (i, boundary) in HISTOGRAM_BUCKETS_MS.iter().enumerate() {
if ms <= *boundary {
self.buckets[i].fetch_add(1, Ordering::Relaxed);
return;
}
}
self.buckets[HISTOGRAM_BUCKETS_MS.len()].fetch_add(1, Ordering::Relaxed);
}
}
pub struct Metrics {
pub messages_dispatched: AtomicU64,
pub messages_dropped: AtomicU64,
pub upstream_user_messages: AtomicU64,
pub sendmessage_total: AtomicU64,
pub sendmessage_errors: AtomicU64,
pub relogin_attempts: AtomicU64,
pub messages_persist_dropped: AtomicU64,
pub getupdates_latency_ms: LatencyHistogram,
pub sendmessage_upstream_latency_ms: LatencyHistogram,
pub dispatch_latency_ms: LatencyHistogram,
pub process_start_unix_secs: f64,
pub dispatcher_lagged: AtomicU64,
pub persist_fire_and_forget_failures_forward: AtomicU64,
pub persist_fire_and_forget_failures_broadcast: AtomicU64,
pub quote_resolve_miss_total: AtomicU64,
}
impl Metrics {
pub fn new() -> Self {
let process_start_unix_secs = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_secs_f64())
.unwrap_or(0.0);
Self {
messages_dispatched: AtomicU64::new(0),
messages_dropped: AtomicU64::new(0),
upstream_user_messages: AtomicU64::new(0),
sendmessage_total: AtomicU64::new(0),
sendmessage_errors: AtomicU64::new(0),
relogin_attempts: AtomicU64::new(0),
messages_persist_dropped: AtomicU64::new(0),
getupdates_latency_ms: LatencyHistogram::new(),
sendmessage_upstream_latency_ms: LatencyHistogram::new(),
dispatch_latency_ms: LatencyHistogram::new(),
process_start_unix_secs,
dispatcher_lagged: AtomicU64::new(0),
persist_fire_and_forget_failures_forward: AtomicU64::new(0),
persist_fire_and_forget_failures_broadcast: AtomicU64::new(0),
quote_resolve_miss_total: AtomicU64::new(0),
}
}
}
impl Default for Metrics {
fn default() -> Self {
Self::new()
}
}
pub struct LatencyGuard<'a> {
start: std::time::Instant,
histogram: &'a LatencyHistogram,
}
impl<'a> LatencyGuard<'a> {
pub fn new(histogram: &'a LatencyHistogram) -> Self {
Self {
start: std::time::Instant::now(),
histogram,
}
}
}
impl Drop for LatencyGuard<'_> {
fn drop(&mut self) {
self.histogram.observe(self.start.elapsed());
}
}
#[derive(Debug, Default)]
pub struct PollTracker {
pub counts: StdMutex<HashMap<String, usize>>,
total: AtomicUsize,
hub_cap: AtomicUsize,
}
impl PollTracker {
pub fn set_hub_cap(&self, cap: usize) {
self.hub_cap.store(cap, Ordering::Relaxed);
}
pub fn total_polls(&self) -> usize {
self.total.load(Ordering::Relaxed)
}
pub fn enter(self: &Arc<Self>, vtoken: &str) -> EnterOutcome {
let cap = self.hub_cap.load(Ordering::Relaxed);
let prev_total = self.total.fetch_add(1, Ordering::AcqRel);
if prev_total >= cap {
self.total.fetch_sub(1, Ordering::AcqRel);
return EnterOutcome::HubLimitReached {
total: prev_total,
cap,
};
}
let count = {
let Ok(mut counts) = self.counts.lock() else {
return EnterOutcome::Poisoned {
guard: PollGuard {
tracker: Arc::clone(self),
vtoken: vtoken.to_string(),
},
};
};
let c = counts.entry(vtoken.to_string()).or_insert(0);
*c += 1;
*c
};
EnterOutcome::Ok {
per_vtoken: count,
guard: PollGuard {
tracker: Arc::clone(self),
vtoken: vtoken.to_string(),
},
}
}
}
#[derive(Debug)]
pub enum EnterOutcome {
Ok { per_vtoken: usize, guard: PollGuard },
HubLimitReached { total: usize, cap: usize },
Poisoned { guard: PollGuard },
}
impl EnterOutcome {
#[allow(dead_code)]
pub fn guard(self) -> Option<PollGuard> {
match self {
EnterOutcome::Ok { guard, .. } | EnterOutcome::Poisoned { guard } => Some(guard),
EnterOutcome::HubLimitReached { .. } => None,
}
}
}
#[derive(Debug)]
pub struct PollGuard {
tracker: Arc<PollTracker>,
vtoken: String,
}
impl Drop for PollGuard {
fn drop(&mut self) {
self.tracker.total.fetch_sub(1, Ordering::AcqRel);
let Ok(mut counts) = self.tracker.counts.lock() else {
return;
};
if let Some(c) = counts.get_mut(&self.vtoken) {
*c = c.saturating_sub(1);
if *c == 0 {
counts.remove(&self.vtoken);
}
}
}
}
pub struct IlinkConnState {
pub upstream: Arc<dyn UpstreamSink>,
pub shutdown: watch::Receiver<bool>,
pub ilink_status: Arc<AtomicU8>,
pub qr_tx: broadcast::Sender<QrLoginUiEvent>,
pub qr_last_ready: Arc<std::sync::Mutex<Option<QrLoginUiEvent>>>,
pub relogin_tx: broadcast::Sender<()>,
pub qr_ticket: crate::server::sse_ticket::SseTicketStore,
}
impl IlinkConnState {
pub(crate) fn new(upstream: Arc<dyn UpstreamSink>, shutdown: watch::Receiver<bool>) -> Self {
let (qr_tx, _) = broadcast::channel(16);
let (relogin_tx, _) = broadcast::channel(4);
Self {
upstream,
shutdown,
ilink_status: Arc::new(AtomicU8::new(ilink_status::UNKNOWN)),
qr_tx,
qr_last_ready: Arc::new(std::sync::Mutex::new(None)),
relogin_tx,
qr_ticket: crate::server::sse_ticket::SseTicketStore::new(),
}
}
}
pub struct RoutingState {
pub router: Mutex<Router>,
}
impl RoutingState {
pub(crate) fn new() -> Self {
Self {
router: Mutex::new(Router::new(None)),
}
}
}
pub struct ClientState {
pub registry: RwLock<ClientRegistry>,
pub pairing: RwLock<PairingRegistry>,
pub pairing_notify: Arc<tokio::sync::Notify>,
pub queue: Arc<dyn MessageQueue>,
pub poll_tracker: Arc<PollTracker>,
pub last_seen: Arc<DashMap<String, AtomicU64>>,
}
impl ClientState {
pub(crate) fn new(queue: Arc<dyn MessageQueue>) -> Self {
let poll_tracker = Arc::new(PollTracker::default());
poll_tracker.set_hub_cap(MAX_HUB_POLLS_DEFAULT);
Self {
registry: RwLock::new(ClientRegistry::new()),
pairing: RwLock::new(PairingRegistry::new()),
pairing_notify: Arc::new(tokio::sync::Notify::new()),
queue,
poll_tracker,
last_seen: Arc::new(DashMap::new()),
}
}
}
const MAX_CONCURRENT_PERSIST_TASKS: usize = 32;
#[derive(Debug, Clone)]
pub struct AdminConfig {
pub token: Option<String>,
pub insecure_no_auth: bool,
pub outbound_origin_label: Option<String>,
}
impl AdminConfig {
pub fn from_env() -> Self {
let token = std::env::var("ILINK_ADMIN_TOKEN")
.ok()
.filter(|s| !s.is_empty());
let insecure_no_auth = token.is_none()
&& std::env::var("ILINK_ADMIN_INSECURE_NO_AUTH")
.ok()
.map(|v| matches!(v.trim().to_ascii_lowercase().as_str(), "1" | "true" | "yes"))
.unwrap_or(false);
let outbound_origin_label = std::env::var("ILINKHUB_OUTBOUND_ORIGIN_LABEL").ok();
Self {
token,
insecure_no_auth,
outbound_origin_label,
}
}
}
pub struct HubState {
pub ilink: IlinkConnState,
pub routing: RoutingState,
pub clients: ClientState,
pub store: Arc<Store>,
pub metrics: Arc<Metrics>,
pub persist_sem: Arc<tokio::sync::Semaphore>,
pub relay_secret: String,
pub admin: AdminConfig,
pub a2a_waiter: Arc<crate::mcp::A2aWaiter>,
}
impl HubState {
pub fn new(
upstream: Arc<dyn UpstreamSink>,
store: Arc<Store>,
queue: Arc<dyn MessageQueue>,
shutdown: watch::Receiver<bool>,
relay_secret: String,
admin: AdminConfig,
) -> Arc<Self> {
Arc::new(Self {
ilink: IlinkConnState::new(upstream, shutdown),
routing: RoutingState::new(),
clients: ClientState::new(queue),
store,
metrics: Arc::new(Metrics::new()),
persist_sem: Arc::new(tokio::sync::Semaphore::new(MAX_CONCURRENT_PERSIST_TASKS)),
relay_secret,
admin,
a2a_waiter: Arc::new(crate::mcp::A2aWaiter::new()),
})
}
pub async fn with_router_and_registry<R>(
self: &Arc<Self>,
f: impl FnOnce(&mut Router, &ClientRegistry) -> R,
) -> Option<R> {
let mut router_guard = self.routing.router.lock().await;
let registry_guard = self.clients.registry.read().await;
let result = f(&mut router_guard, ®istry_guard);
drop(registry_guard);
drop(router_guard);
Some(result)
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::Ordering;
use std::time::Duration;
#[test]
fn observe_increments_count_and_sum_us() {
let h = LatencyHistogram::new();
h.observe(Duration::from_millis(10));
assert_eq!(h.count.load(Ordering::Relaxed), 1);
assert_eq!(h.sum_us.load(Ordering::Relaxed), 10_000);
}
#[test]
fn observe_sub_ms_contributes_non_zero_to_sum_us() {
let h = LatencyHistogram::new();
h.observe(Duration::from_micros(500)); assert_eq!(
h.sum_us.load(Ordering::Relaxed),
500,
"500µs observation must add 500 to sum_us, not 0"
);
}
#[test]
fn observe_exactly_at_boundary_lands_in_lower_bucket() {
let h = LatencyHistogram::new();
h.observe(Duration::from_millis(1));
assert_eq!(
h.buckets[0].load(Ordering::Relaxed),
1,
"1ms (= boundary) must land in bucket[0] (≤ 1ms)"
);
assert_eq!(
h.buckets[1].load(Ordering::Relaxed),
0,
"bucket[1] must be empty"
);
}
#[test]
fn observe_just_above_boundary_lands_in_next_bucket() {
let h = LatencyHistogram::new();
h.observe(Duration::from_millis(2));
assert_eq!(
h.buckets[0].load(Ordering::Relaxed),
0,
"2ms must NOT land in bucket[0] (> 1ms boundary)"
);
assert_eq!(
h.buckets[1].load(Ordering::Relaxed),
1,
"2ms must land in bucket[1] (≤ 5ms)"
);
}
#[test]
fn observe_above_all_boundaries_lands_in_overflow_bucket() {
let h = LatencyHistogram::new();
h.observe(Duration::from_millis(11_000)); let overflow_idx = HISTOGRAM_BUCKETS_MS.len();
assert_eq!(
h.buckets[overflow_idx].load(Ordering::Relaxed),
1,
"11s must land in overflow bucket"
);
for i in 0..HISTOGRAM_BUCKETS_MS.len() {
assert_eq!(
h.buckets[i].load(Ordering::Relaxed),
0,
"bucket[{i}] must be empty"
);
}
}
fn make_tracker(cap: usize) -> Arc<PollTracker> {
let t = Arc::new(PollTracker::default());
t.set_hub_cap(cap);
t
}
#[test]
fn enter_first_poll_returns_ok_with_per_vtoken_count_one() {
let t = make_tracker(10);
let outcome = t.enter("vtok-A");
assert!(
matches!(outcome, EnterOutcome::Ok { per_vtoken: 1, .. }),
"first enter must be Ok with per_vtoken=1"
);
assert_eq!(t.total_polls(), 1);
}
#[test]
fn enter_second_poll_same_vtoken_returns_count_two() {
let t = make_tracker(10);
let _g1 = t.enter("vtok-A");
let outcome = t.enter("vtok-A");
assert!(
matches!(outcome, EnterOutcome::Ok { per_vtoken: 2, .. }),
"second enter on same vtoken must be Ok with per_vtoken=2"
);
}
#[test]
fn enter_hub_cap_exceeded_returns_hub_limit_reached() {
let t = make_tracker(1);
let _g1 = t.enter("vtok-A"); let outcome = t.enter("vtok-B");
assert!(
matches!(outcome, EnterOutcome::HubLimitReached { .. }),
"cap=1 with 1 active poll must reject next enter"
);
}
#[test]
fn enter_hub_limit_reached_rolls_back_total_counter() {
let t = make_tracker(1);
let _g1 = t.enter("vtok-A"); let _rejected = t.enter("vtok-B"); assert_eq!(
t.total_polls(),
1,
"rejected enter must roll back Hub-wide total"
);
}
#[test]
fn poll_guard_drop_decrements_total_and_removes_vtoken_entry() {
let t = make_tracker(10);
{
let outcome = t.enter("vtok-A");
assert_eq!(t.total_polls(), 1);
drop(outcome); }
assert_eq!(
t.total_polls(),
0,
"PollGuard drop must decrement Hub-wide total to zero"
);
let counts = t.counts.lock().unwrap();
assert!(
!counts.contains_key("vtok-A"),
"zero-count vtoken entry must be removed from map"
);
}
#[test]
fn enter_outcome_guard_some_for_ok_none_for_hub_limit_reached() {
let t = make_tracker(10);
let ok = t.enter("vtok-A");
assert!(ok.guard().is_some(), "Ok variant must return Some(guard)");
let t2 = make_tracker(0);
let rejected = t2.enter("vtok-X");
assert!(
rejected.guard().is_none(),
"HubLimitReached must return None from guard()"
);
}
}