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 fn spawn_gossip(&self) {
114 let membership = self.membership.clone();
115 tokio::spawn(async move {
116 let mut ticker = tokio::time::interval(Duration::from_secs(3));
117 loop {
118 ticker.tick().await;
119 membership.round().await;
120 }
121 });
122 }
123}
124
125pub async fn gossip_handler(
128 State(state): State<Arc<AppState>>,
129 Json(incoming): Json<GossipView>,
130) -> Response {
131 match &state.cluster {
132 Some(cluster) => {
133 cluster.membership.merge(incoming);
134 Json(cluster.membership.snapshot()).into_response()
135 }
136 None => StatusCode::NOT_FOUND.into_response(),
137 }
138}
139
140pub async fn route(State(state): State<Arc<AppState>>, req: Request, next: Next) -> Response {
144 let Some(cluster) = state.cluster.clone() else {
145 return next.run(req).await;
146 };
147 if req.uri().path() == GOSSIP_PATH || req.headers().contains_key(FORWARDED) {
149 return next.run(req).await;
150 }
151 let Some(agent) = agent_from_request(&req, &state) else {
154 return next.run(req).await;
155 };
156
157 let nodes = cluster.membership.live_nodes().await.unwrap_or_default();
158 let ids: Vec<String> = nodes.iter().map(|n| n.id.clone()).collect();
159 let owner = placement::owner(&agent, &ids)
160 .unwrap_or(cluster.node_id.as_str())
161 .to_string();
162 if owner == cluster.node_id {
163 return next.run(req).await;
164 }
165 match nodes.iter().find(|n| n.id == owner).map(|n| n.addr.clone()) {
166 Some(addr) => forward(&cluster.http, &addr, req).await,
167 None => next.run(req).await,
169 }
170}
171
172fn agent_from_request(req: &Request, state: &AppState) -> Option<String> {
175 let secret = state.jwt_secret.as_deref()?;
176 let token = req
177 .headers()
178 .get(header::AUTHORIZATION)
179 .and_then(|v| v.to_str().ok())
180 .and_then(|v| v.strip_prefix("Bearer "))?;
181 crate::auth::validate_token(secret, token)
182 .ok()
183 .map(|c| c.agent_id)
184}
185
186async fn forward(http: &reqwest::Client, addr: &str, req: Request) -> Response {
188 let (parts, body) = req.into_parts();
189 let path = parts
190 .uri
191 .path_and_query()
192 .map(|p| p.as_str())
193 .unwrap_or("/");
194 let url = format!("{}{}", addr.trim_end_matches('/'), path);
195
196 let bytes = match axum::body::to_bytes(body, 16 * 1024 * 1024).await {
197 Ok(b) => b,
198 Err(_) => {
199 return error_response(StatusCode::BAD_REQUEST, "request body too large to forward");
200 }
201 };
202
203 let mut builder = http.request(parts.method, &url).body(bytes.to_vec());
204 for (name, value) in parts.headers.iter() {
205 if name != header::HOST {
206 builder = builder.header(name, value);
207 }
208 }
209 builder = builder.header(FORWARDED, "1");
210
211 match builder.send().await {
212 Ok(resp) => {
213 let status = resp.status();
214 let headers = resp.headers().clone();
215 let body = resp.bytes().await.unwrap_or_default();
216 let mut out = Response::builder().status(status);
217 for (name, value) in headers.iter() {
218 if name != header::TRANSFER_ENCODING && name != header::CONNECTION {
219 out = out.header(name, value);
220 }
221 }
222 out.body(Body::from(body)).unwrap_or_else(|_| {
223 error_response(StatusCode::BAD_GATEWAY, "bad upstream response")
224 })
225 }
226 Err(e) => {
227 tracing::warn!(error = %e, url = %url, "sharding: forward to owner failed");
228 error_response(StatusCode::BAD_GATEWAY, "owner node unreachable")
229 }
230 }
231}
232
233fn error_response(status: StatusCode, msg: &str) -> Response {
234 (
235 status,
236 [(header::CONTENT_TYPE, "application/json")],
237 serde_json::json!({ "error": msg }).to_string(),
238 )
239 .into_response()
240}