1use 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#[derive(Clone, Debug, PartialEq, Eq)]
21pub struct PushedRoomEvent {
22 pub room: String,
24 pub sender: String,
26 pub body: String,
30 pub event_id: String,
32 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 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)); }
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 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 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}