mentedb_server/
cluster.rs1use 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
28pub const GOSSIP_PATH: &str = "/v1/cluster/gossip";
30const FORWARDED: &str = "x-mentedb-forwarded";
32
33#[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#[derive(Clone)]
57pub struct Cluster {
58 node_id: String,
59 membership: Arc<Membership>,
60 http: reqwest::Client,
61}
62
63impl Cluster {
64 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 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 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
134pub 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
149pub 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 if req.uri().path() == GOSSIP_PATH || req.headers().contains_key(FORWARDED) {
158 return next.run(req).await;
159 }
160 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 None => next.run(req).await,
178 }
179}
180
181fn 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
195async 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}