Skip to main content

m4a_agent/engine/
push.rs

1//! The client's push socket. The client opens it against the
2//! homeserver it already has. Push contract **v1**: metadata only
3//! (`room`, `sender`, `event_id`, `recipient`, `wire_type`) — never a
4//! plaintext `body`. The ack goes out before the event is handed to that
5//! session. This path does not POST. The client drives `/sync` and
6//! decrypts after the push.
7
8use std::io::ErrorKind;
9use std::sync::atomic::{AtomicBool, Ordering};
10use std::sync::{Arc, Mutex};
11use std::thread::{self, JoinHandle};
12use std::time::{Duration, Instant};
13
14use tungstenite::{stream::MaybeTlsStream, Message};
15
16use super::http::clip_public;
17use crate::error::AgentError;
18
19/// One room event the homeserver pushed for a single session.
20#[derive(Clone, Debug, PartialEq, Eq)]
21pub struct PushedRoomEvent {
22    /// Room id.
23    pub room: String,
24    /// Sender mxid.
25    pub sender: String,
26    /// Always empty on push v1 (server never sends plaintext). Kept so
27    /// call sites compiling against this struct stay stable; wake uses
28    /// `/sync` + decrypt, not this field.
29    pub body: String,
30    /// Matrix event id.
31    pub event_id: String,
32    /// `m.room.message` or `m.room.encrypted`.
33    pub wire_type: String,
34}
35
36struct Incoming {
37    recipient: String,
38    event: PushedRoomEvent,
39}
40
41pub struct PushLink {
42    inbox: Arc<Mutex<Vec<Incoming>>>,
43    stop: Arc<AtomicBool>,
44    workers: Vec<JoinHandle<()>>,
45}
46
47impl Drop for PushLink {
48    fn drop(&mut self) {
49        self.stop.store(true, Ordering::Release);
50        for worker in self.workers.drain(..) {
51            let _ = worker.join();
52        }
53    }
54}
55
56impl PushLink {
57    /// One socket for all `tokens`, or with `separate` (product mode: a handshake
58    /// is vouched for one identity) one socket per token sharing one inbox.
59    pub fn open(base_url: &str, keep_prefix: bool, tokens: Vec<String>, separate: bool) -> Result<Self, AgentError> {
60        if tokens.is_empty() {
61            return Err(AgentError::Transport("push socket has no session".into()));
62        }
63        let url = push_ws_url(base_url, keep_prefix)?;
64        let inbox = Arc::new(Mutex::new(Vec::new()));
65        let stop = Arc::new(AtomicBool::new(false));
66        let groups: Vec<Vec<String>> = if separate { tokens.into_iter().map(|t| vec![t]).collect() } else { vec![tokens] };
67        let mut link = Self { inbox: Arc::clone(&inbox), stop: Arc::clone(&stop), workers: Vec::new() };
68        let mut readies = Vec::new();
69        let mut faileds = Vec::new();
70        for group in groups {
71            let ready = Arc::new(AtomicBool::new(false));
72            let failed = Arc::new(Mutex::new(None));
73            let worker = thread::spawn({
74                let (url, inbox, ready, failed, stop) = (url.clone(), Arc::clone(&inbox), Arc::clone(&ready), Arc::clone(&failed), Arc::clone(&stop));
75                move || worker_main(url, group, inbox, ready, failed, stop)
76            });
77            link.workers.push(worker);
78            readies.push(ready);
79            faileds.push(failed);
80        }
81        let start = Instant::now();
82        loop {
83            if readies.iter().all(|r| r.load(Ordering::Acquire)) {
84                return Ok(link);
85            }
86            for failed in &faileds {
87                if let Some(err) = failed.lock().unwrap_or_else(|err| err.into_inner()).clone() {
88                    return Err(AgentError::Transport(err)); // Drop stops and joins the workers
89                }
90            }
91            if start.elapsed() > Duration::from_secs(5) {
92                return Err(AgentError::Transport("push socket did not register".into()));
93            }
94            thread::sleep(Duration::from_millis(20));
95        }
96    }
97
98    pub fn drain(&self) -> Vec<(String, PushedRoomEvent)> {
99        self.inbox
100            .lock()
101            .unwrap_or_else(|err| err.into_inner())
102            .drain(..)
103            .map(|incoming| (incoming.recipient, incoming.event))
104            .collect()
105    }
106}
107
108fn push_ws_url(base: &str, keep_prefix: bool) -> Result<String, AgentError> {
109    let url = reqwest::Url::parse(base).map_err(|_| AgentError::Protocol("not a valid server URL".into()))?;
110    let scheme = match url.scheme() {
111        "http" => "ws",
112        "https" => "wss",
113        _ => return Err(AgentError::Protocol("not a valid server URL".into())),
114    };
115    let host = url
116        .host_str()
117        .filter(|host| !host.is_empty())
118        .ok_or(AgentError::Protocol("not a valid server URL".into()))?;
119    let mut raw = format!("{scheme}://{host}");
120    if let Some(port) = url.port() {
121        raw.push(':');
122        raw.push_str(&port.to_string());
123    }
124    raw.push_str(if keep_prefix { "/_matrix/client/v3/push" } else { "/client/v3/push" });
125    Ok(raw)
126}
127
128fn worker_main(
129    url: String,
130    mut tokens: Vec<String>,
131    inbox: Arc<Mutex<Vec<Incoming>>>,
132    ready: Arc<AtomicBool>,
133    failed: Arc<Mutex<Option<String>>>,
134    stop: Arc<AtomicBool>,
135) {
136    let result = run_socket(&url, &mut tokens, &inbox, &ready, &stop);
137    for token in &mut tokens {
138        zeroize::Zeroize::zeroize(token);
139    }
140    if let Err(err) = result {
141        if !stop.load(Ordering::Acquire) {
142            *failed.lock().unwrap_or_else(|err| err.into_inner()) = Some(err);
143        }
144    }
145}
146
147fn run_socket(
148    url: &str,
149    tokens: &mut [String],
150    inbox: &Mutex<Vec<Incoming>>,
151    ready: &AtomicBool,
152    stop: &AtomicBool,
153) -> Result<(), String> {
154    // The first token also rides the handshake as a bearer: a product proxy
155    // authenticates the handshake with it and signs it; a plain server ignores it.
156    let mut request = tungstenite::client::IntoClientRequest::into_client_request(url)
157        .map_err(|err| clip_public(err.to_string()))?;
158    if let Some(first) = tokens.first() {
159        if let Ok(value) = tungstenite::http::HeaderValue::from_str(&format!("Bearer {first}")) {
160            request.headers_mut().insert(tungstenite::http::header::AUTHORIZATION, value);
161        }
162    }
163    let (mut socket, _response) =
164        tungstenite::connect(request).map_err(|err| clip_public(err.to_string()))?;
165    match socket.get_mut() {
166        MaybeTlsStream::Plain(tcp) => {
167            let _ = tcp.set_read_timeout(Some(Duration::from_millis(200)));
168        }
169        MaybeTlsStream::Rustls(tls) => {
170            let _ = tls.sock.set_read_timeout(Some(Duration::from_millis(200)));
171        }
172        _ => {}
173    }
174    let register = serde_json::json!({ "type": "register", "tokens": tokens }).to_string();
175    for token in tokens.iter_mut() {
176        zeroize::Zeroize::zeroize(token);
177    }
178    socket
179        .send(Message::text(register))
180        .map_err(|err| clip_public(err.to_string()))?;
181    loop {
182        if stop.load(Ordering::Acquire) {
183            return Ok(());
184        }
185        let message = match socket.read() {
186            Ok(message) => message,
187            Err(tungstenite::Error::Io(err))
188                if matches!(err.kind(), ErrorKind::TimedOut | ErrorKind::WouldBlock) =>
189            {
190                continue;
191            }
192            Err(err) => return Err(clip_public(err.to_string())),
193        };
194        let text = match message {
195            Message::Text(text) => text,
196            Message::Close(_) => return Err("push socket closed".into()),
197            _ => continue,
198        };
199        let Ok(value) = serde_json::from_str::<serde_json::Value>(text.as_str()) else {
200            continue;
201        };
202        match value.get("type").and_then(|item| item.as_str()) {
203            Some("registered") => ready.store(true, Ordering::Release),
204            Some("event") => {
205                if let Some(envelope_id) = value.get("envelope_id").and_then(|item| item.as_str()) {
206                    let ack = serde_json::json!({ "type": "ack", "envelope_id": envelope_id })
207                        .to_string();
208                    socket
209                        .send(Message::text(ack))
210                        .map_err(|err| clip_public(err.to_string()))?;
211                }
212                if let Some(incoming) = parse_event(&value) {
213                    inbox
214                        .lock()
215                        .unwrap_or_else(|err| err.into_inner())
216                        .push(incoming);
217                }
218            }
219            _ => {}
220        }
221    }
222}
223
224fn parse_event(value: &serde_json::Value) -> Option<Incoming> {
225    if value.get("access_token").is_some() || value.get("bearer").is_some() {
226        return None;
227    }
228    let event = value.get("event")?;
229    if event.get("access_token").is_some() || event.get("bearer").is_some() {
230        return None;
231    }
232    let room = event.get("room")?.as_str()?.to_string();
233    let sender = event.get("sender")?.as_str()?.to_string();
234    let event_id = event.get("event_id")?.as_str()?.to_string();
235    let recipient = event.get("recipient")?.as_str()?.to_string();
236    let wire_type = event
237        .get("wire_type")
238        .and_then(|item| item.as_str())
239        .unwrap_or("m.room.message")
240        .to_string();
241    // Push v1: never trust/require body. Always empty locally.
242    let body = String::new();
243    if room.is_empty() || sender.is_empty() || event_id.is_empty() || recipient.is_empty() || wire_type.is_empty()
244    {
245        return None;
246    }
247    Some(Incoming {
248        recipient,
249        event: PushedRoomEvent {
250            room,
251            sender,
252            body,
253            event_id,
254            wire_type,
255        },
256    })
257}
258
259#[cfg(test)]
260mod tests {
261    use super::*;
262
263    #[test]
264    fn https_homeserver_opens_the_push_socket_as_wss() {
265        let url = push_ws_url("https://example.test/ignored", false).expect("url");
266        assert_eq!(url, "wss://example.test/client/v3/push");
267        let local = push_ws_url("http://127.0.0.1:9", false).expect("local");
268        assert_eq!(local, "ws://127.0.0.1:9/client/v3/push");
269    }
270
271    #[test]
272    fn encrypted_push_keeps_the_event_id_and_drops_any_body() {
273        let value = serde_json::json!({
274            "type": "event",
275            "v": 1,
276            "envelope_id": "p1",
277            "event": {
278                "room": "!room:example",
279                "sender": "@a:example",
280                "event_id": "$evt",
281                "recipient": "@b:example",
282                "wire_type": "m.room.encrypted",
283                "body": "not-plaintext",
284            }
285        });
286        let incoming = parse_event(&value).expect("parsed");
287        assert_eq!(incoming.recipient, "@b:example");
288        assert_eq!(incoming.event.event_id, "$evt");
289        assert_eq!(incoming.event.wire_type, "m.room.encrypted");
290        assert!(incoming.event.body.is_empty());
291        assert_eq!(incoming.event.room, "!room:example");
292        assert_eq!(incoming.event.sender, "@a:example");
293    }
294
295    #[test]
296    fn v1_plaintext_wire_type_parses_without_body() {
297        let value = serde_json::json!({
298            "type": "event",
299            "v": 1,
300            "envelope_id": "p2",
301            "event": {
302                "room": "!room:example",
303                "sender": "@a:example",
304                "event_id": "$evt2",
305                "recipient": "@b:example",
306                "wire_type": "m.room.message",
307            }
308        });
309        let incoming = parse_event(&value).expect("v1 without body must parse");
310        assert!(incoming.event.body.is_empty());
311        assert_eq!(incoming.event.wire_type, "m.room.message");
312        assert_eq!(incoming.event.event_id, "$evt2");
313    }
314}