refactor: Fix errors in api/client/state.rs

This commit is contained in:
Ginger
2026-04-28 09:16:52 -04:00
parent 360e0dada8
commit 1d39210a0c
2 changed files with 125 additions and 118 deletions
+50 -43
View File
@@ -5,13 +5,15 @@ use axum_client_ip::ClientIp;
use conduwuit::{ use conduwuit::{
Err, Result, err, Err, Result, err,
matrix::{Event, pdu::PartialPdu}, matrix::{Event, pdu::PartialPdu},
utils::BoolExt,
}; };
use conduwuit_service::Services; use conduwuit_service::Services;
use futures::{FutureExt, TryStreamExt}; use futures::{FutureExt, TryStreamExt};
use ruma::{ use ruma::{
MilliSecondsSinceUnixEpoch, OwnedEventId, RoomId, UserId, MilliSecondsSinceUnixEpoch, OwnedEventId, RoomId, UserId,
api::client::state::{get_state_events, get_state_events_for_key, send_state_event}, api::client::state::{
get_state_event_for_key::{self, v3::StateEventFormat},
get_state_events, send_state_event,
},
events::{ events::{
AnyStateEventContent, StateEventType, AnyStateEventContent, StateEventType,
room::{ room::{
@@ -24,7 +26,7 @@ use ruma::{
}, },
serde::Raw, serde::Raw,
}; };
use serde_json::json; use serde_json::{json, value::to_raw_value};
use crate::{Ruma, RumaResponse}; use crate::{Ruma, RumaResponse};
@@ -46,8 +48,7 @@ pub(crate) async fn send_state_event_for_key_route(
return Err!(Request(UserSuspended("You cannot perform this action while suspended."))); return Err!(Request(UserSuspended("You cannot perform this action while suspended.")));
} }
Ok(send_state_event::v3::Response { let event_id = send_state_event_for_key_helper(
event_id: send_state_event_for_key_helper(
&services, &services,
sender_user, sender_user,
&body.room_id, &body.room_id,
@@ -61,8 +62,9 @@ pub(crate) async fn send_state_event_for_key_route(
}, },
) )
.boxed() .boxed()
.await?, .await?;
})
Ok(send_state_event::v3::Response::new(event_id))
} }
/// # `PUT /_matrix/client/*/rooms/{roomId}/state/{eventType}` /// # `PUT /_matrix/client/*/rooms/{roomId}/state/{eventType}`
@@ -100,15 +102,15 @@ pub(crate) async fn get_state_events_route(
return Err!(Request(Forbidden("You don't have permission to view the room state."))); return Err!(Request(Forbidden("You don't have permission to view the room state.")));
} }
Ok(get_state_events::v3::Response { let room_state = services
room_state: services
.rooms .rooms
.state_accessor .state_accessor
.room_state_full_pdus(&body.room_id) .room_state_full_pdus(&body.room_id)
.map_ok(Event::into_format) .map_ok(Event::into_format)
.try_collect() .try_collect()
.await?, .await?;
})
Ok(get_state_events::v3::Response::new(room_state))
} }
/// # `GET /_matrix/client/v3/rooms/{roomid}/state/{eventType}/{stateKey}` /// # `GET /_matrix/client/v3/rooms/{roomid}/state/{eventType}/{stateKey}`
@@ -119,10 +121,10 @@ pub(crate) async fn get_state_events_route(
/// ///
/// - If not joined: Only works if current room history visibility is world /// - If not joined: Only works if current room history visibility is world
/// readable /// readable
pub(crate) async fn get_state_events_for_key_route( pub(crate) async fn get_state_event_for_key_route(
State(services): State<crate::State>, State(services): State<crate::State>,
body: Ruma<get_state_events_for_key::v3::Request>, body: Ruma<get_state_event_for_key::v3::Request>,
) -> Result<get_state_events_for_key::v3::Response> { ) -> Result<get_state_event_for_key::v3::Response> {
let sender_user = body.sender_user(); let sender_user = body.sender_user();
if !services if !services
@@ -149,15 +151,9 @@ pub(crate) async fn get_state_events_for_key_route(
)))) ))))
})?; })?;
let event_format = body let content = match body.format {
.format | StateEventFormat::Content => event.content,
.as_ref() | StateEventFormat::Event => to_raw_value(&json!({
.is_some_and(|f| f.to_lowercase().eq("event"));
Ok(get_state_events_for_key::v3::Response {
content: event_format.or(|| event.get_content_as_value()),
event: event_format.then(|| {
json!({
"content": event.content(), "content": event.content(),
"event_id": event.event_id(), "event_id": event.event_id(),
"origin_server_ts": event.origin_server_ts(), "origin_server_ts": event.origin_server_ts(),
@@ -166,9 +162,12 @@ pub(crate) async fn get_state_events_for_key_route(
"state_key": event.state_key(), "state_key": event.state_key(),
"type": event.kind(), "type": event.kind(),
"unsigned": event.unsigned(), "unsigned": event.unsigned(),
}) }))
}), .unwrap(),
}) | _ => return Err!(Request(InvalidParam("Unknown response format"))),
};
Ok(get_state_event_for_key::v3::Response::new(content))
} }
/// # `GET /_matrix/client/v3/rooms/{roomid}/state/{eventType}` /// # `GET /_matrix/client/v3/rooms/{roomid}/state/{eventType}`
@@ -181,9 +180,9 @@ pub(crate) async fn get_state_events_for_key_route(
/// readable /// readable
pub(crate) async fn get_state_events_for_empty_key_route( pub(crate) async fn get_state_events_for_empty_key_route(
State(services): State<crate::State>, State(services): State<crate::State>,
body: Ruma<get_state_events_for_key::v3::Request>, body: Ruma<get_state_event_for_key::v3::Request>,
) -> Result<RumaResponse<get_state_events_for_key::v3::Response>> { ) -> Result<RumaResponse<get_state_event_for_key::v3::Response>> {
get_state_events_for_key_route(State(services), body) get_state_event_for_key_route(State(services), body)
.await .await
.map(RumaResponse) .map(RumaResponse)
} }
@@ -237,9 +236,16 @@ async fn allowed_to_send_state_event(
| StateEventType::RoomServerAcl => { | StateEventType::RoomServerAcl => {
// prevents common ACL paw-guns as ACL management is difficult and prone to // prevents common ACL paw-guns as ACL management is difficult and prone to
// irreversible mistakes // irreversible mistakes
match json.deserialize_as::<RoomServerAclEventContent>() { match json.deserialize_as_unchecked::<RoomServerAclEventContent>() {
| Ok(acl_content) => { | Ok(acl_content) => {
if acl_content.allow_is_empty() { let allow_has_wildcard = acl_content.allow.iter().any(|entry| entry == "*");
let deny_has_wildcard = acl_content.deny.iter().any(|entry| entry == "*");
let allow_has_server = acl_content
.allow
.iter()
.any(|entry| entry == services.globals.server_name().as_str());
if acl_content.allow.is_empty() {
return Err!(Request(BadJson(debug_warn!( return Err!(Request(BadJson(debug_warn!(
%room_id, %room_id,
"Sending an ACL event with an empty allow key will permanently \ "Sending an ACL event with an empty allow key will permanently \
@@ -248,7 +254,7 @@ async fn allowed_to_send_state_event(
)))); ))));
} }
if acl_content.deny_contains("*") && acl_content.allow_contains("*") { if allow_has_wildcard && deny_has_wildcard {
return Err!(Request(BadJson(debug_warn!( return Err!(Request(BadJson(debug_warn!(
%room_id, %room_id,
"Sending an ACL event with a deny and allow key value of \"*\" will \ "Sending an ACL event with a deny and allow key value of \"*\" will \
@@ -257,9 +263,9 @@ async fn allowed_to_send_state_event(
)))); ))));
} }
if acl_content.deny_contains("*") if deny_has_wildcard
&& !acl_content.is_allowed(services.globals.server_name()) && !acl_content.is_allowed(services.globals.server_name())
&& !acl_content.allow_contains(services.globals.server_name().as_str()) && !allow_has_server
{ {
return Err!(Request(BadJson(debug_warn!( return Err!(Request(BadJson(debug_warn!(
%room_id, %room_id,
@@ -269,9 +275,9 @@ async fn allowed_to_send_state_event(
)))); ))));
} }
if !acl_content.allow_contains("*") if !allow_has_wildcard
&& !acl_content.is_allowed(services.globals.server_name()) && !acl_content.is_allowed(services.globals.server_name())
&& !acl_content.allow_contains(services.globals.server_name().as_str()) && !allow_has_server
{ {
return Err!(Request(BadJson(debug_warn!( return Err!(Request(BadJson(debug_warn!(
%room_id, %room_id,
@@ -297,7 +303,7 @@ async fn allowed_to_send_state_event(
// admin room is a sensitive room, it should not ever be made public // admin room is a sensitive room, it should not ever be made public
if let Ok(admin_room_id) = services.admin.get_admin_room().await { if let Ok(admin_room_id) = services.admin.get_admin_room().await {
if admin_room_id == room_id { if admin_room_id == room_id {
match json.deserialize_as::<RoomJoinRulesEventContent>() { match json.deserialize_as_unchecked::<RoomJoinRulesEventContent>() {
| Ok(join_rule) => | Ok(join_rule) =>
if join_rule.join_rule == JoinRule::Public { if join_rule.join_rule == JoinRule::Public {
return Err!(Request(Forbidden( return Err!(Request(Forbidden(
@@ -316,7 +322,7 @@ async fn allowed_to_send_state_event(
| StateEventType::RoomHistoryVisibility => { | StateEventType::RoomHistoryVisibility => {
// admin room is a sensitive room, it should not ever be made world readable // admin room is a sensitive room, it should not ever be made world readable
if let Ok(admin_room_id) = services.admin.get_admin_room().await { if let Ok(admin_room_id) = services.admin.get_admin_room().await {
match json.deserialize_as::<RoomHistoryVisibilityEventContent>() { match json.deserialize_as_unchecked::<RoomHistoryVisibilityEventContent>() {
| Ok(visibility_content) => { | Ok(visibility_content) => {
if admin_room_id == room_id if admin_room_id == room_id
&& visibility_content.history_visibility && visibility_content.history_visibility
@@ -337,7 +343,7 @@ async fn allowed_to_send_state_event(
} }
}, },
| StateEventType::RoomCanonicalAlias => { | StateEventType::RoomCanonicalAlias => {
match json.deserialize_as::<RoomCanonicalAliasEventContent>() { match json.deserialize_as_unchecked::<RoomCanonicalAliasEventContent>() {
| Ok(canonical_alias_content) => { | Ok(canonical_alias_content) => {
let mut aliases = canonical_alias_content.alt_aliases.clone(); let mut aliases = canonical_alias_content.alt_aliases.clone();
@@ -369,7 +375,8 @@ async fn allowed_to_send_state_event(
}, },
} }
}, },
| StateEventType::RoomMember => match json.deserialize_as::<RoomMemberEventContent>() { | StateEventType::RoomMember =>
match json.deserialize_as_unchecked::<RoomMemberEventContent>() {
| Ok(mut membership_content) => { | Ok(mut membership_content) => {
let Ok(state_key) = UserId::parse(state_key) else { let Ok(state_key) = UserId::parse(state_key) else {
return Err!(Request(BadJson( return Err!(Request(BadJson(
@@ -380,12 +387,12 @@ async fn allowed_to_send_state_event(
if let Some(authorising_user) = if let Some(authorising_user) =
membership_content.join_authorized_via_users_server membership_content.join_authorized_via_users_server
{ {
// join_authorized_via_users_server must be thrown away, if user is already a // join_authorized_via_users_server must be thrown away, if user is
// member of the room. // already a member of the room.
if services if services
.rooms .rooms
.state_cache .state_cache
.is_joined(state_key, room_id) .is_joined(&state_key, room_id)
.await .await
{ {
membership_content.join_authorized_via_users_server = None; membership_content.join_authorized_via_users_server = None;
+1 -1
View File
@@ -118,7 +118,7 @@ pub fn build(router: Router<State>, server: &Server) -> Router<State> {
.ruma_route(&client::send_message_event_route) .ruma_route(&client::send_message_event_route)
.ruma_route(&client::send_state_event_for_key_route) .ruma_route(&client::send_state_event_for_key_route)
.ruma_route(&client::get_state_events_route) .ruma_route(&client::get_state_events_route)
.ruma_route(&client::get_state_events_for_key_route) .ruma_route(&client::get_state_event_for_key_route)
// Ruma doesn't have support for multiple paths for a single endpoint yet, and these routes // Ruma doesn't have support for multiple paths for a single endpoint yet, and these routes
// share one Ruma request / response type pair with {get,send}_state_event_for_key_route // share one Ruma request / response type pair with {get,send}_state_event_for_key_route
.route( .route(