1use alloc::collections::BTreeMap;
22use alloc::string::String;
23use alloc::vec::Vec;
24use core::fmt;
25
26use crate::LeanEvent;
27
28#[derive(Debug, Clone, PartialEq, Eq)]
30pub enum AuthError {
31 NotMember { sender: String },
33 InsufficientPowerLevel {
35 required: i64,
36 actual: i64,
37 event_type: String,
38 },
39 BannedUser { sender: String },
41 InvalidStateKey { expected: String, actual: String },
44 CreateWithPrevEvents,
46 MissingAuthEvent(String),
48}
49
50impl fmt::Display for AuthError {
51 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
52 match self {
53 AuthError::NotMember { sender } => {
54 write!(f, "sender {} is not a joined member", sender)
55 }
56 AuthError::InsufficientPowerLevel {
57 required,
58 actual,
59 event_type,
60 } => write!(
61 f,
62 "power level {} < {} required for {}",
63 actual, required, event_type
64 ),
65 AuthError::BannedUser { sender } => {
66 write!(f, "sender {} is banned", sender)
67 }
68 AuthError::InvalidStateKey { expected, actual } => {
69 write!(
70 f,
71 "invalid state_key: expected {}, got {}",
72 expected, actual
73 )
74 }
75 AuthError::CreateWithPrevEvents => {
76 write!(f, "m.room.create must not have prev_events")
77 }
78 AuthError::MissingAuthEvent(id) => {
79 write!(f, "missing auth event: {}", id)
80 }
81 }
82 }
83}
84
85pub type RoomState = BTreeMap<(String, String), LeanEvent>;
87
88pub fn check_auth(event: &LeanEvent, state: &RoomState) -> Result<(), AuthError> {
97 if event.event_type == "m.room.create" {
99 if !event.prev_events.is_empty() {
100 return Err(AuthError::CreateWithPrevEvents);
101 }
102 return Ok(());
104 }
105
106 let member_key = ("m.room.member".into(), event.sender.clone());
108 if let Some(member_event) = state.get(&member_key) {
109 if let Some(membership) = member_event
110 .content
111 .get("membership")
112 .and_then(|m| m.as_str())
113 {
114 if membership == "ban" {
115 return Err(AuthError::BannedUser {
116 sender: event.sender.clone(),
117 });
118 }
119
120 if event.event_type != "m.room.member" && membership != "join" {
122 return Err(AuthError::NotMember {
123 sender: event.sender.clone(),
124 });
125 }
126 }
127 } else if event.event_type != "m.room.member" {
128 return Err(AuthError::NotMember {
130 sender: event.sender.clone(),
131 });
132 }
133
134 let sender_pl = get_sender_power_level(&event.sender, state);
136 let required_pl = get_required_power_level(&event.event_type, state);
137
138 if sender_pl < required_pl {
139 return Err(AuthError::InsufficientPowerLevel {
140 required: required_pl,
141 actual: sender_pl,
142 event_type: event.event_type.clone(),
143 });
144 }
145
146 if event.event_type == "m.room.member" {
148 check_membership_rules(event, state)?;
149 }
150
151 Ok(())
152}
153
154fn get_sender_power_level(sender: &str, state: &RoomState) -> i64 {
156 let pl_key = ("m.room.power_levels".into(), String::new());
157 if let Some(pl_event) = state.get(&pl_key) {
158 if let Some(users) = pl_event.content.get("users").and_then(|u| u.as_object()) {
159 if let Some(pl) = users.get(sender).and_then(|v| v.as_i64()) {
160 return pl;
161 }
162 }
163 if let Some(default) = pl_event
165 .content
166 .get("users_default")
167 .and_then(|v| v.as_i64())
168 {
169 return default;
170 }
171 }
172 0 }
174
175fn get_required_power_level(event_type: &str, state: &RoomState) -> i64 {
177 let pl_key = ("m.room.power_levels".into(), String::new());
178 if let Some(pl_event) = state.get(&pl_key) {
179 if let Some(events) = pl_event.content.get("events").and_then(|e| e.as_object()) {
181 if let Some(pl) = events.get(event_type).and_then(|v| v.as_i64()) {
182 return pl;
183 }
184 }
185 if event_type.starts_with("m.room.") {
187 if let Some(default) = pl_event
188 .content
189 .get("state_default")
190 .and_then(|v| v.as_i64())
191 {
192 return default;
193 }
194 }
195 if let Some(default) = pl_event
196 .content
197 .get("events_default")
198 .and_then(|v| v.as_i64())
199 {
200 return default;
201 }
202 }
203 0 }
205
206fn check_membership_rules(event: &LeanEvent, state: &RoomState) -> Result<(), AuthError> {
208 let target_user = &event.state_key;
209 let _sender = &event.sender;
210
211 let new_membership = event
212 .content
213 .get("membership")
214 .and_then(|m| m.as_str())
215 .unwrap_or("");
216
217 match new_membership {
218 "join" if event.state_key != event.sender => {
220 return Err(AuthError::InvalidStateKey {
221 expected: event.sender.clone(),
222 actual: event.state_key.clone(),
223 });
224 }
225 "leave" if event.state_key != event.sender => {
227 let sender_pl = get_sender_power_level(&event.sender, state);
228 let kick_pl = get_kick_power_level(state);
229 if sender_pl < kick_pl {
230 return Err(AuthError::InsufficientPowerLevel {
231 required: kick_pl,
232 actual: sender_pl,
233 event_type: "kick".into(),
234 });
235 }
236 }
237 "ban" => {
238 let sender_pl = get_sender_power_level(&event.sender, state);
240 let ban_pl = get_ban_power_level(state);
241 if sender_pl < ban_pl {
242 return Err(AuthError::InsufficientPowerLevel {
243 required: ban_pl,
244 actual: sender_pl,
245 event_type: "ban".into(),
246 });
247 }
248 }
249 "invite" => {
250 if event.state_key == event.sender {
252 return Err(AuthError::InvalidStateKey {
253 expected: alloc::format!("!= {}", event.sender),
254 actual: event.state_key.clone(),
255 });
256 }
257 let target_key = ("m.room.member".into(), target_user.clone());
259 if let Some(target_member) = state.get(&target_key) {
260 if target_member
261 .content
262 .get("membership")
263 .and_then(|m| m.as_str())
264 == Some("ban")
265 {
266 return Err(AuthError::BannedUser {
267 sender: target_user.clone(),
268 });
269 }
270 }
271 }
272 _ => {}
273 }
274
275 Ok(())
276}
277
278fn get_kick_power_level(state: &RoomState) -> i64 {
280 let pl_key = ("m.room.power_levels".into(), String::new());
281 if let Some(pl_event) = state.get(&pl_key) {
282 if let Some(kick) = pl_event.content.get("kick").and_then(|v| v.as_i64()) {
283 return kick;
284 }
285 }
286 50 }
288
289fn get_ban_power_level(state: &RoomState) -> i64 {
291 let pl_key = ("m.room.power_levels".into(), String::new());
292 if let Some(pl_event) = state.get(&pl_key) {
293 if let Some(ban) = pl_event.content.get("ban").and_then(|v| v.as_i64()) {
294 return ban;
295 }
296 }
297 50 }
299
300pub fn check_auth_chain(
304 sorted_events: &[LeanEvent],
305 initial_state: &RoomState,
306) -> (Vec<String>, Vec<(String, AuthError)>) {
307 let mut state = initial_state.clone();
308 let mut accepted = Vec::new();
309 let mut rejected = Vec::new();
310
311 for event in sorted_events {
312 match check_auth(event, &state) {
313 Ok(()) => {
314 if !event.state_key.is_empty() || event.event_type == "m.room.create" {
316 state.insert(
317 (event.event_type.clone(), event.state_key.clone()),
318 event.clone(),
319 );
320 }
321 accepted.push(event.event_id.clone());
322 }
323 Err(e) => {
324 rejected.push((event.event_id.clone(), e));
325 }
326 }
327 }
328
329 (accepted, rejected)
330}
331
332#[cfg(test)]
333mod tests {
334 use super::*;
335 use alloc::vec;
336 use serde_json::json;
337
338 fn make_event(
339 id: &str,
340 event_type: &str,
341 state_key: &str,
342 sender: &str,
343 content: serde_json::Value,
344 ) -> LeanEvent {
345 LeanEvent {
346 event_id: id.into(),
347 event_type: event_type.into(),
348 state_key: state_key.into(),
349 sender: sender.into(),
350 content,
351 ..Default::default()
352 }
353 }
354
355 #[test]
356 fn test_create_event_no_prev_events() {
357 let create = make_event(
358 "$create",
359 "m.room.create",
360 "",
361 "@alice:example.com",
362 json!({}),
363 );
364 let state = RoomState::new();
365 assert!(check_auth(&create, &state).is_ok());
366 }
367
368 #[test]
369 fn test_create_event_with_prev_events() {
370 let mut create = make_event(
371 "$create",
372 "m.room.create",
373 "",
374 "@alice:example.com",
375 json!({}),
376 );
377 create.prev_events = vec!["$other".into()];
378 let state = RoomState::new();
379 assert_eq!(
380 check_auth(&create, &state),
381 Err(AuthError::CreateWithPrevEvents)
382 );
383 }
384
385 #[test]
386 fn test_non_member_rejection() {
387 let msg = make_event("$msg", "m.room.message", "", "@bob:example.com", json!({}));
388 let state = RoomState::new();
389 assert!(matches!(
390 check_auth(&msg, &state),
391 Err(AuthError::NotMember { .. })
392 ));
393 }
394
395 #[test]
396 fn test_joined_member_can_send() {
397 let msg = make_event(
398 "$msg",
399 "m.room.message",
400 "",
401 "@alice:example.com",
402 json!({}),
403 );
404 let mut state = RoomState::new();
405 state.insert(
406 ("m.room.member".into(), "@alice:example.com".into()),
407 make_event(
408 "$join",
409 "m.room.member",
410 "@alice:example.com",
411 "@alice:example.com",
412 json!({"membership": "join"}),
413 ),
414 );
415 assert!(check_auth(&msg, &state).is_ok());
416 }
417
418 #[test]
419 fn test_banned_user_rejected() {
420 let msg = make_event(
421 "$msg",
422 "m.room.message",
423 "",
424 "@alice:example.com",
425 json!({}),
426 );
427 let mut state = RoomState::new();
428 state.insert(
429 ("m.room.member".into(), "@alice:example.com".into()),
430 make_event(
431 "$ban",
432 "m.room.member",
433 "@alice:example.com",
434 "@admin:example.com",
435 json!({"membership": "ban"}),
436 ),
437 );
438 assert!(matches!(
439 check_auth(&msg, &state),
440 Err(AuthError::BannedUser { .. })
441 ));
442 }
443
444 #[test]
445 fn test_insufficient_power_level() {
446 let msg = make_event(
447 "$msg",
448 "m.room.power_levels",
449 "",
450 "@alice:example.com",
451 json!({}),
452 );
453 let mut state = RoomState::new();
454 state.insert(
455 ("m.room.member".into(), "@alice:example.com".into()),
456 make_event(
457 "$join",
458 "m.room.member",
459 "@alice:example.com",
460 "@alice:example.com",
461 json!({"membership": "join"}),
462 ),
463 );
464 state.insert(
465 ("m.room.power_levels".into(), String::new()),
466 make_event(
467 "$pl",
468 "m.room.power_levels",
469 "",
470 "@admin:example.com",
471 json!({"state_default": 50, "users": {"@admin:example.com": 100}}),
472 ),
473 );
474 assert!(matches!(
475 check_auth(&msg, &state),
476 Err(AuthError::InsufficientPowerLevel { .. })
477 ));
478 }
479
480 #[test]
481 fn test_join_self_only() {
482 let join = make_event(
483 "$join",
484 "m.room.member",
485 "@bob:example.com",
486 "@alice:example.com",
487 json!({"membership": "join"}),
488 );
489 let state = RoomState::new();
490 assert!(matches!(
491 check_auth(&join, &state),
492 Err(AuthError::InvalidStateKey { .. })
493 ));
494 }
495
496 #[test]
497 fn test_iterative_auth_chain() {
498 let create = make_event(
499 "$create",
500 "m.room.create",
501 "",
502 "@alice:example.com",
503 json!({}),
504 );
505 let join = make_event(
506 "$join",
507 "m.room.member",
508 "@alice:example.com",
509 "@alice:example.com",
510 json!({"membership": "join"}),
511 );
512 let msg = make_event(
513 "$msg",
514 "m.room.message",
515 "",
516 "@alice:example.com",
517 json!({"body": "hello"}),
518 );
519 let (accepted, rejected) = check_auth_chain(&[create, join, msg], &RoomState::new());
520 assert_eq!(accepted, vec!["$create", "$join", "$msg"]);
521 assert!(rejected.is_empty());
522 }
523
524 #[test]
525 fn test_auth_error_display() {
526 let err = AuthError::NotMember {
527 sender: "@bob:example.com".into(),
528 };
529 let msg = alloc::format!("{}", err);
530 assert!(msg.contains("bob"));
531 }
532}