Skip to main content

mentedb_server/
cluster.rs

1//! Self-organizing sharding for mentedb-server.
2//!
3//! Running N nodes shards the fleet with no external coordinator: nodes gossip to
4//! converge on the live set (engine `mentedb::sharding::gossip`), rendezvous
5//! placement picks one owner per agent (every node computes the same owner), and a
6//! request that lands on the wrong node is forwarded to the owner. Ownership is the
7//! single-writer lock, so an agent's database is only ever written by one node.
8//!
9//! Off unless `MENTEDB_SHARDING` is set. Routing needs JWT auth (the agent id comes
10//! from the token); without a `--jwt-secret` the middleware is a pass-through.
11
12use std::sync::Arc;
13use std::time::Duration;
14
15use axum::{
16    Json,
17    body::Body,
18    extract::{Request, State},
19    http::{StatusCode, header},
20    middleware::Next,
21    response::{IntoResponse, Response},
22};
23use mentedb::sharding::gossip::{GossipMembership, GossipTransport, GossipView};
24use mentedb::sharding::{NodeRegistry, placement};
25
26use crate::state::AppState;
27
28/// Path peers POST their gossip view to (and the reply carries ours back).
29pub const GOSSIP_PATH: &str = "/v1/cluster/gossip";
30/// Marks a request already forwarded once, so a placement disagreement cannot loop.
31const FORWARDED: &str = "x-mentedb-forwarded";
32
33/// reqwest-backed gossip transport: POST our view to a peer, get theirs back.
34#[derive(Clone)]
35pub struct HttpGossip {
36    client: reqwest::Client,
37}
38
39impl GossipTransport for HttpGossip {
40    async fn exchange(&self, peer_addr: &str, ours: GossipView) -> Result<GossipView, String> {
41        let url = format!("{}{}", peer_addr.trim_end_matches('/'), GOSSIP_PATH);
42        let resp = self
43            .client
44            .post(&url)
45            .json(&ours)
46            .send()
47            .await
48            .map_err(|e| e.to_string())?;
49        resp.json::<GossipView>().await.map_err(|e| e.to_string())
50    }
51}
52
53type Membership = GossipMembership<HttpGossip>;
54
55/// Cluster handle held in `AppState`; present only when sharding is enabled.
56#[derive(Clone)]
57pub struct Cluster {
58    node_id: String,
59    membership: Arc<Membership>,
60    http: reqwest::Client,
61}
62
63impl Cluster {
64    /// Build from the environment, or return `None` when sharding is off or
65    /// misconfigured:
66    /// - `MENTEDB_SHARDING` = `1`/`true` to enable.
67    /// - `MENTEDB_NODE_ID` (defaults to `$HOSTNAME`, else `node-<pid>`).
68    /// - `MENTEDB_NODE_ADDR` this node's base URL peers reach it at (required).
69    /// - `MENTEDB_SEEDS` comma-separated peer base URLs to bootstrap from.
70    pub fn from_env() -> Option<Self> {
71        let enabled = std::env::var("MENTEDB_SHARDING")
72            .map(|v| v == "1" || v.eq_ignore_ascii_case("true"))
73            .unwrap_or(false);
74        if !enabled {
75            return None;
76        }
77        let node_id = std::env::var("MENTEDB_NODE_ID")
78            .ok()
79            .or_else(|| std::env::var("HOSTNAME").ok())
80            .unwrap_or_else(|| format!("node-{}", std::process::id()));
81        let node_addr = std::env::var("MENTEDB_NODE_ADDR").unwrap_or_default();
82        if node_addr.is_empty() {
83            tracing::warn!(
84                "MENTEDB_SHARDING is set but MENTEDB_NODE_ADDR is empty; sharding disabled"
85            );
86            return None;
87        }
88        let seeds: Vec<String> = std::env::var("MENTEDB_SEEDS")
89            .unwrap_or_default()
90            .split(',')
91            .map(|s| s.trim().to_string())
92            .filter(|s| !s.is_empty())
93            .collect();
94        let http = reqwest::Client::new();
95        let membership = Arc::new(GossipMembership::new(
96            node_id.clone(),
97            node_addr,
98            seeds,
99            Duration::from_secs(15),
100            HttpGossip {
101                client: http.clone(),
102            },
103        ));
104        tracing::info!(node = %node_id, "sharding enabled: self-organizing gossip fleet");
105        Some(Self {
106            node_id,
107            membership,
108            http,
109        })
110    }
111
112    /// How many nodes the gossip layer currently considers live.
113    pub async fn live_node_count(&self) -> usize {
114        self.membership
115            .live_nodes()
116            .await
117            .map(|n| n.len())
118            .unwrap_or(0)
119    }
120
121    /// Run the gossip loop in the background: one anti-entropy round per interval.
122    pub fn spawn_gossip(&self) {
123        let membership = self.membership.clone();
124        tokio::spawn(async move {
125            let mut ticker = tokio::time::interval(Duration::from_secs(3));
126            loop {
127                ticker.tick().await;
128                membership.round().await;
129            }
130        });
131    }
132}
133
134/// Handler for `GOSSIP_PATH`: merge a peer's view and answer with ours. 404 when
135/// sharding is off.
136pub async fn gossip_handler(
137    State(state): State<Arc<AppState>>,
138    Json(incoming): Json<GossipView>,
139) -> Response {
140    match &state.cluster {
141        Some(cluster) => {
142            cluster.membership.merge(incoming);
143            Json(cluster.membership.snapshot()).into_response()
144        }
145        None => StatusCode::NOT_FOUND.into_response(),
146    }
147}
148
149/// Middleware: forward a request to the node that owns its agent, or serve it here.
150/// A pass-through when sharding is off, the request carries no resolvable agent, or
151/// we already own it.
152pub async fn route(State(state): State<Arc<AppState>>, req: Request, next: Next) -> Response {
153    let Some(cluster) = state.cluster.clone() else {
154        return next.run(req).await;
155    };
156    // The gossip endpoint and already-forwarded requests are always served locally.
157    if req.uri().path() == GOSSIP_PATH || req.headers().contains_key(FORWARDED) {
158        return next.run(req).await;
159    }
160    // Identify the agent from the JWT as an owned value before any await, so no
161    // reference to the request is held across one. No token, no per-agent affinity.
162    let Some(agent) = agent_from_request(&req, &state) else {
163        return next.run(req).await;
164    };
165
166    let nodes = cluster.membership.live_nodes().await.unwrap_or_default();
167    let ids: Vec<String> = nodes.iter().map(|n| n.id.clone()).collect();
168    let owner = placement::owner(&agent, &ids)
169        .unwrap_or(cluster.node_id.as_str())
170        .to_string();
171    if owner == cluster.node_id {
172        return next.run(req).await;
173    }
174    match nodes.iter().find(|n| n.id == owner).map(|n| n.addr.clone()) {
175        Some(addr) => forward(&cluster.http, &addr, req).await,
176        // Owner has no known address (mid-convergence): serve locally rather than fail.
177        None => next.run(req).await,
178    }
179}
180
181/// The agent id a request belongs to, from its Bearer JWT. `None` without auth (no
182/// secret, no token, or a token that does not validate).
183fn agent_from_request(req: &Request, state: &AppState) -> Option<String> {
184    let secret = state.jwt_secret.as_deref()?;
185    let token = req
186        .headers()
187        .get(header::AUTHORIZATION)
188        .and_then(|v| v.to_str().ok())
189        .and_then(|v| v.strip_prefix("Bearer "))?;
190    crate::auth::validate_token(secret, token)
191        .ok()
192        .map(|c| c.agent_id)
193}
194
195/// Reverse-proxy the request to the owning node and relay its response.
196async fn forward(http: &reqwest::Client, addr: &str, req: Request) -> Response {
197    let (parts, body) = req.into_parts();
198    let path = parts
199        .uri
200        .path_and_query()
201        .map(|p| p.as_str())
202        .unwrap_or("/");
203    let url = format!("{}{}", addr.trim_end_matches('/'), path);
204
205    let bytes = match axum::body::to_bytes(body, 16 * 1024 * 1024).await {
206        Ok(b) => b,
207        Err(_) => {
208            return error_response(StatusCode::BAD_REQUEST, "request body too large to forward");
209        }
210    };
211
212    let mut builder = http.request(parts.method, &url).body(bytes.to_vec());
213    for (name, value) in parts.headers.iter() {
214        if name != header::HOST {
215            builder = builder.header(name, value);
216        }
217    }
218    builder = builder.header(FORWARDED, "1");
219
220    match builder.send().await {
221        Ok(resp) => {
222            let status = resp.status();
223            let headers = resp.headers().clone();
224            let body = resp.bytes().await.unwrap_or_default();
225            let mut out = Response::builder().status(status);
226            for (name, value) in headers.iter() {
227                if name != header::TRANSFER_ENCODING && name != header::CONNECTION {
228                    out = out.header(name, value);
229                }
230            }
231            out.body(Body::from(body)).unwrap_or_else(|_| {
232                error_response(StatusCode::BAD_GATEWAY, "bad upstream response")
233            })
234        }
235        Err(e) => {
236            tracing::warn!(error = %e, url = %url, "sharding: forward to owner failed");
237            error_response(StatusCode::BAD_GATEWAY, "owner node unreachable")
238        }
239    }
240}
241
242fn error_response(status: StatusCode, msg: &str) -> Response {
243    (
244        status,
245        [(header::CONTENT_TYPE, "application/json")],
246        serde_json::json!({ "error": msg }).to_string(),
247    )
248        .into_response()
249}