Skip to main content

mail4agent_server/http/
identity.rs

1//! The messenger side of the seam: the signed-assertion middleware and the
2//! lifecycle events endpoint. Verification and wire formats live in
3//! [`m4a_seam`]; this file only applies them. Inert until the deployment sets
4//! [`Homeserver::seam`].
5//!
6//! The middleware verifies the assertion header against the method and the
7//! path as this router receives it, resolves the identity (creating it, which
8//! is first contact), and hands the result to [`super::resolve_caller`]
9//! through an internal header that is stripped from every incoming request
10//! first.
11
12use 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
30/// Internal hand-over header; never trusted from the wire.
31pub(super) const RESOLVED_HEADER: &str = "x-m4a-resolved";
32/// Pseudo user id of an anonymous read-only caller (never a row).
33pub(super) const ANON_USER: i64 = -1;
34
35/// Seam configuration of this server.
36pub struct Seam {
37    /// Shared secrets, current first (a second one is accepted during rotation).
38    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    /// `M4A_ASSERTION_SECRET` (>= 16 chars; `M4A_ASSERTION_SECRET_PREV` accepted too),
57    /// `M4A_ASSERTION_SKEW_S` (30), `M4A_ASSERTION_HEADER`, `M4A_EVENT_SIG_HEADER`.
58    /// `Ok(None)` when no secret is configured.
59    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
84/// Parsed internal header: `(user_id, device_id, claims)`.
85pub(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
107/// Strip the internal header, then honour a valid signed assertion.
108pub(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
154/// Re-stamp the identity's current nick into its member events; returns users to wake.
155fn 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
161/// Product -> messenger lifecycle events. Body is a [`m4a_seam::Event`]; the
162/// signature header carries the hex HMAC of the raw body.
163/// Product's startup snapshot of live credentials; see [`m4a_seam::Reconcile`].
164async 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![])); // replay of an applied event: success, nothing to do
197        }
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            // Not applied: allow a corrected retry with the same id.
217            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        // no assertion at all never creates anything
284        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}