1use std::sync::Arc;
13
14use axum::body::Bytes;
15use axum::extract::{Request, State};
16use axum::http::{HeaderMap, HeaderValue};
17use axum::middleware::Next;
18use axum::response::{IntoResponse, Response};
19use axum::routing::post;
20use axum::{Json, Router};
21use m4a_seam::{Event, EventKind, NonceCache, SeamError};
22use serde_json::{json, Value};
23
24use crate::error::MatrixError;
25use crate::identities;
26use crate::policy::{claims_from, Claims};
27
28use super::{wake_users, Homeserver};
29
30pub(super) const RESOLVED_HEADER: &str = "x-m4a-resolved";
32pub(super) const ANON_USER: i64 = -1;
34
35pub struct Seam {
37 pub secrets: Vec<Vec<u8>>,
39 pub skew_ms: i64,
40 pub assertion_header: String,
41 pub event_sig_header: String,
42 nonces: NonceCache,
43}
44
45impl Seam {
46 pub fn new(secrets: Vec<Vec<u8>>, skew_s: i64, assertion_header: Option<String>, event_sig_header: Option<String>) -> Self {
47 Self {
48 secrets,
49 skew_ms: skew_s * 1000,
50 assertion_header: assertion_header.unwrap_or_else(|| m4a_seam::DEFAULT_ASSERTION_HEADER.to_string()).to_ascii_lowercase(),
51 event_sig_header: event_sig_header.unwrap_or_else(|| m4a_seam::DEFAULT_EVENT_SIG_HEADER.to_string()).to_ascii_lowercase(),
52 nonces: NonceCache::new(),
53 }
54 }
55
56 pub fn from_env() -> Result<Option<Self>, String> {
60 let Ok(secret) = std::env::var("M4A_ASSERTION_SECRET") else { return Ok(None) };
61 if secret.len() < 16 {
62 return Err("M4A_ASSERTION_SECRET must be at least 16 characters".into());
63 }
64 let mut secrets = vec![secret.into_bytes()];
65 if let Ok(prev) = std::env::var("M4A_ASSERTION_SECRET_PREV") {
66 if !prev.is_empty() {
67 secrets.push(prev.into_bytes());
68 }
69 }
70 let skew: i64 = std::env::var("M4A_ASSERTION_SKEW_S").ok().and_then(|v| v.parse().ok()).unwrap_or(30);
71 let name = |k: &str| std::env::var(k).ok().filter(|v| !v.is_empty());
72 Ok(Some(Self::new(secrets, skew, name("M4A_ASSERTION_HEADER"), name("M4A_EVENT_SIG_HEADER"))))
73 }
74}
75
76pub(super) fn routes() -> Router<Arc<Homeserver>> {
77 Router::new().route("/account-source/v1/events", post(lifecycle_event)).route("/account-source/v1/reconcile", post(reconcile_snapshot))
78}
79
80fn now_ms() -> i64 {
81 chrono::Utc::now().timestamp_millis()
82}
83
84pub(super) fn read_resolved(headers: &HeaderMap) -> Option<(i64, String, Claims)> {
86 let v = headers.get(RESOLVED_HEADER)?.to_str().ok()?;
87 let mut it = v.splitn(3, '|');
88 let uid = it.next()?.parse().ok()?;
89 let device = it.next()?.to_string();
90 let flag: u8 = it.next()?.parse().ok()?;
91 let mut claims = claims_from(flag);
92 if uid == ANON_USER {
93 claims.insert("anon".into(), "1".into());
94 }
95 Some((uid, device, claims))
96}
97
98fn seam_error(e: SeamError) -> MatrixError {
99 match e {
100 SeamError::Malformed(m) => MatrixError::unauthorized(format!("assertion rejected: {m}")),
101 SeamError::BadSignature => MatrixError::unauthorized("assertion rejected: bad signature"),
102 SeamError::Expired => MatrixError::unauthorized("assertion expired"),
103 SeamError::Replay => MatrixError::unauthorized("assertion replayed"),
104 }
105}
106
107pub(super) async fn assertion_layer(State(state): State<Arc<Homeserver>>, mut req: Request, next: Next) -> Response {
109 req.headers_mut().remove(RESOLVED_HEADER);
110 let anon_marked = req.headers_mut().remove(m4a_seam::ANON_READ_HEADER).is_some_and(|v| v == "1");
111 let Some(seam) = state.seam.get().cloned() else { return next.run(req).await };
112 let Some(value) = req.headers().get(seam.assertion_header.as_str()).and_then(|v| v.to_str().ok()).map(str::to_string) else {
113 if anon_marked && state.anon_read.get().is_some() {
114 let path = req.uri().path_and_query().map(|p| p.as_str()).unwrap_or("");
115 if !m4a_seam::anon_read_path_ok(req.method().as_str(), path) {
116 return MatrixError::unauthorized("anonymous read is limited to public read-only paths").into_response();
117 }
118 req.headers_mut().insert(RESOLVED_HEADER, HeaderValue::from_static("-1|anon|0"));
119 }
120 return next.run(req).await;
121 };
122 let method = req.method().as_str().to_string();
123 let path = req.uri().path_and_query().map(|p| p.as_str().to_string()).unwrap_or_default();
124 let now = now_ms();
125 let a = match m4a_seam::verify_assertion(&seam.secrets, &value, &method, &path, now, seam.skew_ms) {
126 Ok(a) => a,
127 Err(e) => return seam_error(e).into_response(),
128 };
129 if method != "GET" && method != "HEAD" {
130 if let Err(e) = seam.nonces.check_and_insert(&a, now, seam.skew_ms) {
131 return seam_error(e).into_response();
132 }
133 }
134 let st = Arc::clone(&state);
135 let done = tokio::task::spawn_blocking(move || -> Result<String, MatrixError> {
136 st.conn_scope(|conn: &mut rusqlite::Connection| {
137 let r = identities::resolve_assertion(&mut *conn, &a.nick, &a.cred_ref, now)?;
138 Ok(format!("{}|{}|{}", r.identity.id, r.device_id, a.paid))
139 })
140 })
141 .await;
142 match done {
143 Ok(Ok(v)) => {
144 if let Ok(hv) = HeaderValue::from_str(&v) {
145 req.headers_mut().insert(RESOLVED_HEADER, hv);
146 }
147 next.run(req).await
148 }
149 Ok(Err(e)) => e.into_response(),
150 Err(_) => MatrixError::internal().into_response(),
151 }
152}
153
154fn restamp(conn: &mut rusqlite::Connection, id: i64) -> Result<Vec<i64>, MatrixError> {
156 let Some(idn) = identities::identity_by_id(conn, id)? else { return Ok(vec![]) };
157 let r = crate::store::refresh_member_displayname(conn, id, &idn.nick, &chrono::Utc::now().to_rfc3339(), now_ms())?;
158 Ok(r.affected_user_ids.into_iter().collect())
159}
160
161async fn reconcile_snapshot(State(state): State<Arc<Homeserver>>, headers: HeaderMap, body: Bytes) -> Result<Json<Value>, MatrixError> {
165 let seam = state.seam.get().cloned().ok_or_else(MatrixError::unrecognized)?;
166 let sig = headers.get(seam.event_sig_header.as_str()).and_then(|v| v.to_str().ok()).unwrap_or("");
167 if !m4a_seam::verify_body(&seam.secrets, &body, sig) {
168 return Err(MatrixError::unauthorized("bad event signature"));
169 }
170 let snap: m4a_seam::Reconcile = serde_json::from_slice(&body).map_err(|_| MatrixError::bad_json("malformed reconcile body"))?;
171 let out = super::with_conn_pub(&state, move |conn| identities::reconcile(conn, &snap, now_ms())).await?;
172 wake_users(&state, out.wake.clone());
173 for u in &out.closed {
174 state.push.close_user(*u);
175 }
176 Ok(Json(json!({ "ok": true, "devices_removed": out.devices_removed, "identities_retired": out.identities_retired })))
177}
178
179async fn lifecycle_event(State(state): State<Arc<Homeserver>>, headers: HeaderMap, body: Bytes) -> Result<Json<Value>, MatrixError> {
180 let seam = state.seam.get().cloned().ok_or_else(MatrixError::unrecognized)?;
181 let sig = headers.get(seam.event_sig_header.as_str()).and_then(|v| v.to_str().ok()).unwrap_or("");
182 if !m4a_seam::verify_body(&seam.secrets, &body, sig) {
183 return Err(MatrixError::unauthorized("bad event signature"));
184 }
185 let ev: Event = serde_json::from_slice(&body).map_err(|_| MatrixError::bad_json("unknown event type or missing field"))?;
186 if ev.id.is_empty() {
187 return Err(MatrixError::bad_json("missing id"));
188 }
189 let closed = Arc::new(std::sync::Mutex::new(Vec::<i64>::new()));
190 let closed_in = Arc::clone(&closed);
191 let (applied, woke) = super::with_conn_pub(&state, move |conn| {
192 let closed = closed_in;
193 let now = now_ms();
194 let fresh = conn.execute("INSERT OR IGNORE INTO account_events_seen (event_id, at_ms) VALUES (?1, ?2)", rusqlite::params![ev.id, now])?;
195 if fresh == 0 {
196 return Ok((false, vec![])); }
198 let res: Result<(bool, Vec<i64>), MatrixError> = match &ev.kind {
199 EventKind::CredentialRevoked { cred_ref } => identities::apply_credential_revoked(conn, cred_ref).map(|u| {
200 closed.lock().unwrap().extend(u);
201 (u.is_some(), u.into_iter().collect())
202 }),
203 EventKind::AccountDeleted { nick } => {
204 let id = identities::identity_by_nick(conn, nick)?.map(|i| i.id);
205 identities::apply_account_deleted(conn, nick, now).map(|b| {
206 closed.lock().unwrap().extend(id);
207 (b, vec![])
208 })
209 }
210 EventKind::NickChanged { old, new } => identities::apply_nick_changed(conn, old, new).and_then(|id| match id {
211 Some(id) => Ok((true, restamp(conn, id)?)),
212 None => Ok((false, vec![])),
213 }),
214 };
215 if res.is_err() {
216 conn.execute("DELETE FROM account_events_seen WHERE event_id = ?1", rusqlite::params![ev.id])?;
218 }
219 res
220 })
221 .await?;
222 wake_users(&state, woke);
223 for u in closed.lock().unwrap().iter() {
224 state.push.close_user(*u);
225 }
226 Ok(Json(json!({ "ok": true, "applied": applied })))
227}
228
229#[cfg(test)]
230mod tests {
231 use super::*;
232 use axum::body::Body;
233 use axum::http::{Request as Req, StatusCode};
234 use m4a_seam::{sign_assertion, sign_body, Assertion};
235 use tower::ServiceExt;
236
237 const SECRET: &[u8] = b"0123456789abcdef0123";
238
239 fn state() -> Arc<Homeserver> {
240 let c = rusqlite::Connection::open_in_memory().unwrap();
241 crate::store::create_matrix_schema(&c).unwrap();
242 crate::keys::create_matrix_keys_schema(&c).unwrap();
243 let hs = Arc::new(Homeserver::new(c));
244 let _ = hs.seam.set(Arc::new(Seam::new(vec![SECRET.to_vec()], 30, None, None)));
245 hs
246 }
247
248 async fn call(hs: &Arc<Homeserver>, req: Req<Body>) -> (StatusCode, Value) {
249 let resp = crate::http::router(hs.clone()).oneshot(req).await.unwrap();
250 let st = resp.status();
251 let b = axum::body::to_bytes(resp.into_body(), 1 << 20).await.unwrap();
252 (st, serde_json::from_slice(&b).unwrap_or(Value::Null))
253 }
254
255 fn assertion(nick: &str, cred: &str, nonce: &str, paid: u8) -> Assertion {
256 let now = chrono::Utc::now().timestamp();
257 Assertion { nick: nick.into(), cred_ref: cred.into(), authenticated: 1, paid, iat: now, exp: now + 60, nonce: nonce.into() }
258 }
259
260 fn signed_req(method: &str, path: &str, a: &Assertion, body: Body) -> Req<Body> {
261 Req::builder().method(method).uri(path).header(m4a_seam::DEFAULT_ASSERTION_HEADER, sign_assertion(SECRET, method, path, a))
262 .header("content-type", "application/json").body(body).unwrap()
263 }
264
265 #[tokio::test]
266 async fn assertion_creates_identity_contact_and_device_and_forged_header_is_ignored() {
267 let hs = state();
268 let p = "/client/v3/profile/@ivy:example.org/displayname";
269 let (st, body) = call(&hs, signed_req("GET", p, &assertion("ivy", "c1", "n1", 0), Body::empty())).await;
270 assert_eq!(st, StatusCode::OK, "{body}");
271 assert_eq!(body["displayname"], "ivy");
272 hs.conn_async(|c| {
273 let a = identities::identity_by_nick(c, "ivy").unwrap().unwrap();
274 assert_eq!(crate::store::mxid_of(c, a.id).unwrap().as_deref(), Some("@ivy:example.org"));
275 assert_eq!(crate::keys::list_devices(c, a.id).unwrap().len(), 1);
276 })
277 .await;
278 let forged = Req::builder().uri(p).header(RESOLVED_HEADER, "1|X|1").body(Body::empty()).unwrap();
279 assert_eq!(call(&hs, forged).await.0, StatusCode::UNAUTHORIZED, "no token, forged hand-over is stripped");
280 let mut r = signed_req("GET", p, &assertion("ivy", "c1", "n2", 0), Body::empty());
281 *r.uri_mut() = "/client/v3/profile/@other:example.org/displayname".parse().unwrap();
282 assert_eq!(call(&hs, r).await.0, StatusCode::UNAUTHORIZED);
283 let none = Req::builder().uri("/client/v3/capabilities").body(Body::empty()).unwrap();
285 assert_eq!(call(&hs, none).await.0, StatusCode::UNAUTHORIZED);
286 let n: i64 = hs.conn_async(|c| c.query_row("SELECT COUNT(*) FROM identities", [], |r| r.get(0)).unwrap()).await;
287 assert_eq!(n, 1);
288 }
289
290 #[tokio::test]
291 async fn replayed_nonce_on_a_write_is_refused_and_nick_conflict_has_a_stable_shape() {
292 let hs = state();
293 let a = assertion("jay", "c2", "n-write", 0);
294 let (st, _) = call(&hs, signed_req("PUT", "/client/v3/profile/@jay:example.org/displayname", &a, Body::from("{}"))).await;
295 assert_eq!(st, StatusCode::FORBIDDEN, "displayname is the product's, the core never sets it");
296 let (st, _) = call(&hs, signed_req("PUT", "/client/v3/profile/@jay:example.org/displayname", &a, Body::from("{}"))).await;
297 assert_eq!(st, StatusCode::UNAUTHORIZED, "same nonce again");
298 hs.conn_async(|c| c.execute_batch("INSERT INTO matrix_users (user_id, mxid, created_at) VALUES (77, '@taken:example.org', 't')").unwrap()).await;
299 let (st, body) = call(&hs, signed_req("GET", "/client/v3/capabilities", &assertion("taken", "c3", "n9", 0), Body::empty())).await;
300 assert_eq!((st, body["errcode"].as_str()), (StatusCode::CONFLICT, Some("M4A_NICK_CONFLICT")));
301 }
302
303 #[tokio::test]
304 async fn login_is_closed_and_flows_are_empty() {
305 let hs = state();
306 let r = Req::builder().method("POST").uri("/client/v3/login").header("content-type", "application/json").body(Body::from("{}")).unwrap();
307 assert_eq!(call(&hs, r).await.0, StatusCode::FORBIDDEN);
308 let (_, flows) = call(&hs, Req::builder().uri("/client/v3/login").body(Body::empty()).unwrap()).await;
309 assert_eq!(flows["flows"], json!([]));
310 }
311
312 #[tokio::test]
313 async fn lifecycle_events_are_signed_idempotent_and_revoke_wakes() {
314 let hs = state();
315 let (st, _) = call(&hs, signed_req("GET", "/client/v3/capabilities", &assertion("lee", "c9", "n7", 0), Body::empty())).await;
316 assert_eq!(st, StatusCode::OK);
317 let post = |body: &str, sig: Option<String>| {
318 let sig = sig.unwrap_or_else(|| sign_body(SECRET, body.as_bytes()));
319 Req::builder().method("POST").uri("/account-source/v1/events").header(m4a_seam::DEFAULT_EVENT_SIG_HEADER, sig).body(Body::from(body.to_string())).unwrap()
320 };
321 let rev = r#"{"id":"e1","type":"credential.revoked","cred_ref":"c9"}"#;
322 assert_eq!(call(&hs, post(rev, Some("00".into()))).await.0, StatusCode::UNAUTHORIZED);
323 let (st, b) = call(&hs, post(rev, None)).await;
324 assert_eq!((st, b["applied"].clone()), (StatusCode::OK, json!(true)));
325 assert_eq!(call(&hs, post(rev, None)).await.1["applied"], json!(false), "replayed id is a no-op success");
326 let ren = r#"{"id":"e2","type":"nick.changed","old":"lee","new":"leigh"}"#;
327 assert_eq!(call(&hs, post(ren, None)).await.1["applied"], json!(true));
328 let del = r#"{"id":"e3","type":"account.deleted","nick":"leigh"}"#;
329 assert_eq!(call(&hs, post(del, None)).await.1["applied"], json!(true));
330 assert_eq!(call(&hs, post(r#"{"id":"e4","type":"bogus"}"#, None)).await.0, StatusCode::BAD_REQUEST);
331 }
332
333 #[tokio::test]
334 async fn policy_hook_is_consulted_with_claims_for_every_action() {
335 use crate::policy::{Action, Decision, PolicyContext, PolicyHook};
336 struct Rec(std::sync::Mutex<Vec<(Action, String)>>);
337 impl PolicyHook for Rec {
338 fn decide(&self, c: &PolicyContext<'_>) -> Decision {
339 self.0.lock().unwrap().push((c.action, c.claims.get("flag").cloned().unwrap_or_default()));
340 Decision::Deny("closed".into())
341 }
342 }
343 let hs = state();
344 let rec = Arc::new(Rec(Default::default()));
345 let _ = hs.policy.set(rec.clone());
346 let routes: [(&str, &str, &str); 5] = [
347 ("POST", "/client/v3/createRoom", "{}"),
348 ("POST", "/client/v3/rooms/!r:example.org/join", "{}"),
349 ("POST", "/client/v3/rooms/!r:example.org/invite", "{\"user_id\":\"@a:example.org\"}"),
350 ("PUT", "/client/v3/rooms/!r:example.org/send/m.room.message/t1", "{}"),
351 ("POST", "/media/v3/upload", "x"),
352 ];
353 for (i, (m, p, b)) in routes.iter().enumerate() {
354 let (st, body) = call(&hs, signed_req(m, p, &assertion("pol", "cp", &format!("np{i}"), 1), Body::from(*b))).await;
355 assert_eq!((st, body["errcode"].as_str()), (StatusCode::FORBIDDEN, Some("M4A_POLICY_DENIED")), "{m} {p}: {body}");
356 }
357 let seen = rec.0.lock().unwrap().clone();
358 for a in [Action::CreateRoom, Action::JoinRoom, Action::Invite, Action::SendEvent, Action::UploadMedia] {
359 assert!(seen.iter().any(|(x, f)| *x == a && f == "1"), "{a:?} not consulted with the flag claim: {seen:?}");
360 }
361 }
362}