use alloc::collections::BTreeMap;
use alloc::string::String;
use alloc::vec::Vec;
use core::fmt;
use crate::LeanEvent;
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum AuthError {
NotMember { sender: String },
InsufficientPowerLevel {
required: i64,
actual: i64,
event_type: String,
},
BannedUser { sender: String },
InvalidStateKey { expected: String, actual: String },
CreateWithPrevEvents,
MissingAuthEvent(String),
}
impl fmt::Display for AuthError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
AuthError::NotMember { sender } => {
write!(f, "sender {} is not a joined member", sender)
}
AuthError::InsufficientPowerLevel {
required,
actual,
event_type,
} => write!(
f,
"power level {} < {} required for {}",
actual, required, event_type
),
AuthError::BannedUser { sender } => {
write!(f, "sender {} is banned", sender)
}
AuthError::InvalidStateKey { expected, actual } => {
write!(
f,
"invalid state_key: expected {}, got {}",
expected, actual
)
}
AuthError::CreateWithPrevEvents => {
write!(f, "m.room.create must not have prev_events")
}
AuthError::MissingAuthEvent(id) => {
write!(f, "missing auth event: {}", id)
}
}
}
}
pub type RoomState = BTreeMap<(String, String), LeanEvent>;
pub fn check_auth(event: &LeanEvent, state: &RoomState) -> Result<(), AuthError> {
if event.event_type == "m.room.create" {
if !event.prev_events.is_empty() {
return Err(AuthError::CreateWithPrevEvents);
}
return Ok(());
}
let member_key = ("m.room.member".into(), event.sender.clone());
if let Some(member_event) = state.get(&member_key) {
if let Some(membership) = member_event
.content
.get("membership")
.and_then(|m| m.as_str())
{
if membership == "ban" {
return Err(AuthError::BannedUser {
sender: event.sender.clone(),
});
}
if event.event_type != "m.room.member" && membership != "join" {
return Err(AuthError::NotMember {
sender: event.sender.clone(),
});
}
}
} else if event.event_type != "m.room.member" {
return Err(AuthError::NotMember {
sender: event.sender.clone(),
});
}
let sender_pl = get_sender_power_level(&event.sender, state);
let required_pl = get_required_power_level(&event.event_type, state);
if sender_pl < required_pl {
return Err(AuthError::InsufficientPowerLevel {
required: required_pl,
actual: sender_pl,
event_type: event.event_type.clone(),
});
}
if event.event_type == "m.room.member" {
check_membership_rules(event, state)?;
}
Ok(())
}
fn get_sender_power_level(sender: &str, state: &RoomState) -> i64 {
let pl_key = ("m.room.power_levels".into(), String::new());
if let Some(pl_event) = state.get(&pl_key) {
if let Some(users) = pl_event.content.get("users").and_then(|u| u.as_object()) {
if let Some(pl) = users.get(sender).and_then(|v| v.as_i64()) {
return pl;
}
}
if let Some(default) = pl_event
.content
.get("users_default")
.and_then(|v| v.as_i64())
{
return default;
}
}
0 }
fn get_required_power_level(event_type: &str, state: &RoomState) -> i64 {
let pl_key = ("m.room.power_levels".into(), String::new());
if let Some(pl_event) = state.get(&pl_key) {
if let Some(events) = pl_event.content.get("events").and_then(|e| e.as_object()) {
if let Some(pl) = events.get(event_type).and_then(|v| v.as_i64()) {
return pl;
}
}
if event_type.starts_with("m.room.") {
if let Some(default) = pl_event
.content
.get("state_default")
.and_then(|v| v.as_i64())
{
return default;
}
}
if let Some(default) = pl_event
.content
.get("events_default")
.and_then(|v| v.as_i64())
{
return default;
}
}
0 }
fn check_membership_rules(event: &LeanEvent, state: &RoomState) -> Result<(), AuthError> {
let target_user = &event.state_key;
let _sender = &event.sender;
let new_membership = event
.content
.get("membership")
.and_then(|m| m.as_str())
.unwrap_or("");
match new_membership {
"join" if event.state_key != event.sender => {
return Err(AuthError::InvalidStateKey {
expected: event.sender.clone(),
actual: event.state_key.clone(),
});
}
"leave" if event.state_key != event.sender => {
let sender_pl = get_sender_power_level(&event.sender, state);
let kick_pl = get_kick_power_level(state);
if sender_pl < kick_pl {
return Err(AuthError::InsufficientPowerLevel {
required: kick_pl,
actual: sender_pl,
event_type: "kick".into(),
});
}
}
"ban" => {
let sender_pl = get_sender_power_level(&event.sender, state);
let ban_pl = get_ban_power_level(state);
if sender_pl < ban_pl {
return Err(AuthError::InsufficientPowerLevel {
required: ban_pl,
actual: sender_pl,
event_type: "ban".into(),
});
}
}
"invite" => {
if event.state_key == event.sender {
return Err(AuthError::InvalidStateKey {
expected: alloc::format!("!= {}", event.sender),
actual: event.state_key.clone(),
});
}
let target_key = ("m.room.member".into(), target_user.clone());
if let Some(target_member) = state.get(&target_key) {
if target_member
.content
.get("membership")
.and_then(|m| m.as_str())
== Some("ban")
{
return Err(AuthError::BannedUser {
sender: target_user.clone(),
});
}
}
}
_ => {}
}
Ok(())
}
fn get_kick_power_level(state: &RoomState) -> i64 {
let pl_key = ("m.room.power_levels".into(), String::new());
if let Some(pl_event) = state.get(&pl_key) {
if let Some(kick) = pl_event.content.get("kick").and_then(|v| v.as_i64()) {
return kick;
}
}
50 }
fn get_ban_power_level(state: &RoomState) -> i64 {
let pl_key = ("m.room.power_levels".into(), String::new());
if let Some(pl_event) = state.get(&pl_key) {
if let Some(ban) = pl_event.content.get("ban").and_then(|v| v.as_i64()) {
return ban;
}
}
50 }
pub fn check_auth_chain(
sorted_events: &[LeanEvent],
initial_state: &RoomState,
) -> (Vec<String>, Vec<(String, AuthError)>) {
let mut state = initial_state.clone();
let mut accepted = Vec::new();
let mut rejected = Vec::new();
for event in sorted_events {
match check_auth(event, &state) {
Ok(()) => {
if !event.state_key.is_empty() || event.event_type == "m.room.create" {
state.insert(
(event.event_type.clone(), event.state_key.clone()),
event.clone(),
);
}
accepted.push(event.event_id.clone());
}
Err(e) => {
rejected.push((event.event_id.clone(), e));
}
}
}
(accepted, rejected)
}
#[cfg(test)]
mod tests {
use super::*;
use alloc::vec;
use serde_json::json;
fn make_event(
id: &str,
event_type: &str,
state_key: &str,
sender: &str,
content: serde_json::Value,
) -> LeanEvent {
LeanEvent {
event_id: id.into(),
event_type: event_type.into(),
state_key: state_key.into(),
sender: sender.into(),
content,
..Default::default()
}
}
#[test]
fn test_create_event_no_prev_events() {
let create = make_event(
"$create",
"m.room.create",
"",
"@alice:example.com",
json!({}),
);
let state = RoomState::new();
assert!(check_auth(&create, &state).is_ok());
}
#[test]
fn test_create_event_with_prev_events() {
let mut create = make_event(
"$create",
"m.room.create",
"",
"@alice:example.com",
json!({}),
);
create.prev_events = vec!["$other".into()];
let state = RoomState::new();
assert_eq!(
check_auth(&create, &state),
Err(AuthError::CreateWithPrevEvents)
);
}
#[test]
fn test_non_member_rejection() {
let msg = make_event("$msg", "m.room.message", "", "@bob:example.com", json!({}));
let state = RoomState::new();
assert!(matches!(
check_auth(&msg, &state),
Err(AuthError::NotMember { .. })
));
}
#[test]
fn test_joined_member_can_send() {
let msg = make_event(
"$msg",
"m.room.message",
"",
"@alice:example.com",
json!({}),
);
let mut state = RoomState::new();
state.insert(
("m.room.member".into(), "@alice:example.com".into()),
make_event(
"$join",
"m.room.member",
"@alice:example.com",
"@alice:example.com",
json!({"membership": "join"}),
),
);
assert!(check_auth(&msg, &state).is_ok());
}
#[test]
fn test_banned_user_rejected() {
let msg = make_event(
"$msg",
"m.room.message",
"",
"@alice:example.com",
json!({}),
);
let mut state = RoomState::new();
state.insert(
("m.room.member".into(), "@alice:example.com".into()),
make_event(
"$ban",
"m.room.member",
"@alice:example.com",
"@admin:example.com",
json!({"membership": "ban"}),
),
);
assert!(matches!(
check_auth(&msg, &state),
Err(AuthError::BannedUser { .. })
));
}
#[test]
fn test_insufficient_power_level() {
let msg = make_event(
"$msg",
"m.room.power_levels",
"",
"@alice:example.com",
json!({}),
);
let mut state = RoomState::new();
state.insert(
("m.room.member".into(), "@alice:example.com".into()),
make_event(
"$join",
"m.room.member",
"@alice:example.com",
"@alice:example.com",
json!({"membership": "join"}),
),
);
state.insert(
("m.room.power_levels".into(), String::new()),
make_event(
"$pl",
"m.room.power_levels",
"",
"@admin:example.com",
json!({"state_default": 50, "users": {"@admin:example.com": 100}}),
),
);
assert!(matches!(
check_auth(&msg, &state),
Err(AuthError::InsufficientPowerLevel { .. })
));
}
#[test]
fn test_join_self_only() {
let join = make_event(
"$join",
"m.room.member",
"@bob:example.com",
"@alice:example.com",
json!({"membership": "join"}),
);
let state = RoomState::new();
assert!(matches!(
check_auth(&join, &state),
Err(AuthError::InvalidStateKey { .. })
));
}
#[test]
fn test_iterative_auth_chain() {
let create = make_event(
"$create",
"m.room.create",
"",
"@alice:example.com",
json!({}),
);
let join = make_event(
"$join",
"m.room.member",
"@alice:example.com",
"@alice:example.com",
json!({"membership": "join"}),
);
let msg = make_event(
"$msg",
"m.room.message",
"",
"@alice:example.com",
json!({"body": "hello"}),
);
let (accepted, rejected) = check_auth_chain(&[create, join, msg], &RoomState::new());
assert_eq!(accepted, vec!["$create", "$join", "$msg"]);
assert!(rejected.is_empty());
}
#[test]
fn test_auth_error_display() {
let err = AuthError::NotMember {
sender: "@bob:example.com".into(),
};
let msg = alloc::format!("{}", err);
assert!(msg.contains("bob"));
}
}