refactor: Adjust auth check logic structure

This commit is contained in:
Ginger
2026-10-05 14:22:53 +00:00
committed by Ellis Git
parent 4d4693494b
commit 3a20ad7ff2
4 changed files with 209 additions and 224 deletions
Generated
+24 -23
View File
@@ -4946,8 +4946,8 @@ dependencies = [
[[package]]
name = "ruma"
version = "0.17.0"
source = "git+https://github.com/ruma/ruma.git?rev=e7384f01eacd454c593683de50491201e2160fd1#e7384f01eacd454c593683de50491201e2160fd1"
version = "0.16.0"
source = "git+https://github.com/ruma/ruma.git?rev=e927699d374114843f9f41ee133a5ced08449229#e927699d374114843f9f41ee133a5ced08449229"
dependencies = [
"assign",
"js_int",
@@ -4965,8 +4965,8 @@ dependencies = [
[[package]]
name = "ruma-appservice-api"
version = "0.17.0"
source = "git+https://github.com/ruma/ruma.git?rev=e7384f01eacd454c593683de50491201e2160fd1#e7384f01eacd454c593683de50491201e2160fd1"
version = "0.16.0"
source = "git+https://github.com/ruma/ruma.git?rev=e927699d374114843f9f41ee133a5ced08449229#e927699d374114843f9f41ee133a5ced08449229"
dependencies = [
"http",
"js_int",
@@ -4978,8 +4978,8 @@ dependencies = [
[[package]]
name = "ruma-client-api"
version = "0.25.0"
source = "git+https://github.com/ruma/ruma.git?rev=e7384f01eacd454c593683de50491201e2160fd1#e7384f01eacd454c593683de50491201e2160fd1"
version = "0.24.0"
source = "git+https://github.com/ruma/ruma.git?rev=e927699d374114843f9f41ee133a5ced08449229#e927699d374114843f9f41ee133a5ced08449229"
dependencies = [
"as_variant",
"assign",
@@ -5000,11 +5000,11 @@ dependencies = [
[[package]]
name = "ruma-common"
version = "0.20.0"
source = "git+https://github.com/ruma/ruma.git?rev=e7384f01eacd454c593683de50491201e2160fd1#e7384f01eacd454c593683de50491201e2160fd1"
version = "0.19.0"
source = "git+https://github.com/ruma/ruma.git?rev=e927699d374114843f9f41ee133a5ced08449229#e927699d374114843f9f41ee133a5ced08449229"
dependencies = [
"as_variant",
"base64 0.23.1",
"base64 0.22.1",
"bytes",
"date_header",
"form_urlencoded",
@@ -5033,8 +5033,8 @@ dependencies = [
[[package]]
name = "ruma-events"
version = "0.35.0"
source = "git+https://github.com/ruma/ruma.git?rev=e7384f01eacd454c593683de50491201e2160fd1#e7384f01eacd454c593683de50491201e2160fd1"
version = "0.34.0"
source = "git+https://github.com/ruma/ruma.git?rev=e927699d374114843f9f41ee133a5ced08449229#e927699d374114843f9f41ee133a5ced08449229"
dependencies = [
"as_variant",
"indexmap 2.14.2",
@@ -5054,10 +5054,11 @@ dependencies = [
[[package]]
name = "ruma-federation-api"
version = "0.16.0"
source = "git+https://github.com/ruma/ruma.git?rev=e7384f01eacd454c593683de50491201e2160fd1#e7384f01eacd454c593683de50491201e2160fd1"
version = "0.15.0"
source = "git+https://github.com/ruma/ruma.git?rev=e927699d374114843f9f41ee133a5ced08449229#e927699d374114843f9f41ee133a5ced08449229"
dependencies = [
"bytes",
"headers",
"http",
"http-auth",
"httparse",
@@ -5077,7 +5078,7 @@ dependencies = [
[[package]]
name = "ruma-identifiers-validation"
version = "0.12.1"
source = "git+https://github.com/ruma/ruma.git?rev=e7384f01eacd454c593683de50491201e2160fd1#e7384f01eacd454c593683de50491201e2160fd1"
source = "git+https://github.com/ruma/ruma.git?rev=e927699d374114843f9f41ee133a5ced08449229#e927699d374114843f9f41ee133a5ced08449229"
dependencies = [
"js_int",
"thiserror 2.0.21",
@@ -5085,8 +5086,8 @@ dependencies = [
[[package]]
name = "ruma-macros"
version = "0.20.0"
source = "git+https://github.com/ruma/ruma.git?rev=e7384f01eacd454c593683de50491201e2160fd1#e7384f01eacd454c593683de50491201e2160fd1"
version = "0.19.0"
source = "git+https://github.com/ruma/ruma.git?rev=e927699d374114843f9f41ee133a5ced08449229#e927699d374114843f9f41ee133a5ced08449229"
dependencies = [
"as_variant",
"cfg-if",
@@ -5101,8 +5102,8 @@ dependencies = [
[[package]]
name = "ruma-push-gateway-api"
version = "0.16.0"
source = "git+https://github.com/ruma/ruma.git?rev=e7384f01eacd454c593683de50491201e2160fd1#e7384f01eacd454c593683de50491201e2160fd1"
version = "0.15.0"
source = "git+https://github.com/ruma/ruma.git?rev=e927699d374114843f9f41ee133a5ced08449229#e927699d374114843f9f41ee133a5ced08449229"
dependencies = [
"js_int",
"ruma-common",
@@ -5113,10 +5114,10 @@ dependencies = [
[[package]]
name = "ruma-signatures"
version = "0.22.0"
source = "git+https://github.com/ruma/ruma.git?rev=e7384f01eacd454c593683de50491201e2160fd1#e7384f01eacd454c593683de50491201e2160fd1"
version = "0.21.0"
source = "git+https://github.com/ruma/ruma.git?rev=e927699d374114843f9f41ee133a5ced08449229#e927699d374114843f9f41ee133a5ced08449229"
dependencies = [
"base64 0.23.1",
"base64 0.22.1",
"ed25519-dalek 3.0.0",
"memchr",
"pkcs8 0.11.0",
@@ -5130,8 +5131,8 @@ dependencies = [
[[package]]
name = "ruma-state-res"
version = "0.18.0"
source = "git+https://github.com/ruma/ruma.git?rev=e7384f01eacd454c593683de50491201e2160fd1#e7384f01eacd454c593683de50491201e2160fd1"
version = "0.17.0"
source = "git+https://github.com/ruma/ruma.git?rev=e927699d374114843f9f41ee133a5ced08449229#e927699d374114843f9f41ee133a5ced08449229"
dependencies = [
"js_int",
"ruma-common",
+2 -1
View File
@@ -348,7 +348,7 @@ version = "1.1.1"
# Used for matrix spec type definitions and helpers
[workspace.dependencies.ruma]
git = "https://github.com/ruma/ruma.git"
rev = "e7384f01eacd454c593683de50491201e2160fd1"
rev = "e927699d374114843f9f41ee133a5ced08449229"
features = [
"appservice-api-c",
"client-api",
@@ -385,6 +385,7 @@ features = [
"unstable-msc4480",
"unstable-msc4466",
"unstable-msc4494",
"unstable-msc4484",
"unstable-extensible-events",
]
+13 -4
View File
@@ -8,7 +8,7 @@
use conduwuit::{Error, Result, err};
use ruma::{
CanonicalJsonObject,
api::{IncomingRequest, IncomingRequestExt},
api::{IncomingRequest, IncomingRequestExt, auth_scheme::AuthScheme},
};
use serde::Deserialize;
@@ -92,13 +92,22 @@ async fn from_request(
let request = hyper::Request::from_parts(parts, body.as_ref());
// Check authentication
let auth =
R::Authentication::authenticate::<R, &[u8]>(services, &request, auth_query).await?;
let authentication =
R::Authentication::extract_authentication(&request).map_err(|err| {
err!(Request(Unauthorized(warn!(
path = request.uri().path(),
err = err.into(),
"Failed to extract authentication"
))))
})?;
let identity =
R::Authentication::check::<R>(services, &request, authentication, auth_query).await?;
// Deserialize the body
let body = R::try_from_http_request(request, &borrowed_path)
.map_err(|e| err!(Request(BadJson(debug_warn!("{e}")))))?;
Ok(Self { body, json_body, identity: auth })
Ok(Self { body, json_body, identity })
}
}
+170 -196
View File
@@ -84,75 +84,50 @@ pub(crate) fn is_appservice(&self) -> bool { matches!(self, Self::Appservice { .
pub(crate) trait CheckAuth: AuthScheme {
type Identity: Send;
fn authenticate<R: IncomingRequest + Any, B: AsRef<[u8]> + Sync>(
fn check<R: IncomingRequest + Any>(
services: &Services,
incoming_request: &hyper::Request<B>,
incoming_request: &hyper::Request<&[u8]>,
authentication: Self::Output,
query: AuthQueryParams,
) -> impl Future<Output = Result<Self::Identity>> + Send {
async move {
let route = TypeId::of::<R>();
let output = Self::extract_authentication(incoming_request).map_err(|err| {
err!(Request(Unauthorized(warn!(
"Failed to extract authorization: {}",
err.into()
))))
})?;
Self::verify(services, output, incoming_request, query, route).await
}
}
fn verify<B: AsRef<[u8]> + Sync>(
services: &Services,
output: Self::Output,
request: &hyper::Request<B>,
query: AuthQueryParams,
route: TypeId,
) -> impl Future<Output = Result<Self::Identity>> + Send;
}
impl CheckAuth for ServerSignatures {
type Identity = OwnedServerName;
async fn verify<B: AsRef<[u8]> + Sync>(
async fn check<R: IncomingRequest + Any>(
services: &Services,
output: Self::Output,
request: &hyper::Request<B>,
incoming_request: &hyper::Request<&[u8]>,
authentication: Self::Output,
_query: AuthQueryParams,
_route: TypeId,
) -> Result<Self::Identity> {
let destination = services.globals.server_name();
if output
.destination
.as_ref()
.is_some_and(|supplied_destination| supplied_destination != destination)
{
return Err!(Request(Unauthorized("Destination mismatch.")));
}
let min_valid_ts =
MilliSecondsSinceUnixEpoch(UInt::new_saturating(millis_since_unix_epoch()));
let key = services
.server_keys
.get_single_verify_key(&output.origin, &output.key, min_valid_ts)
.get_single_verify_key(&authentication.origin, &authentication.key, min_valid_ts)
.await
.ok_or_else(|| {
err!(Request(Unauthorized(debug_warn!(
origin=%output.origin,
origin=%authentication.origin,
"Unable to acquire your signing key (ID: {})",
output.key
authentication.key
))))
})?;
let keys: PubKeys = [(output.key.to_string(), key.key)].into();
let keys: PubKeyMap = [(output.origin.as_str().into(), keys)].into();
let keys: PubKeys = [(authentication.key.to_string(), key.key)].into();
let keys: PubKeyMap = [(authentication.origin.as_str().into(), keys)].into();
match output.verify_http_request(request, destination, &keys) {
match authentication.verify_http_request(
incoming_request,
services.globals.server_name(),
&keys,
) {
| Ok(()) => {
if services
.moderation
.is_remote_server_forbidden(&output.origin)
.is_remote_server_forbidden(&authentication.origin)
{
return Err!(Request(Forbidden(
"You are blocked from federating with this server."
@@ -160,15 +135,15 @@ async fn verify<B: AsRef<[u8]> + Sync>(
}
// Ping the server as healthy
if services.federation.mark_healthy(&output.origin) {
if services.federation.mark_healthy(&authentication.origin) {
services
.sending
.flush_servers([output.origin.clone()].stream())
.flush_servers([authentication.origin.clone()].stream())
.await
.ok();
}
Ok(output.origin)
Ok(authentication.origin)
},
| Err(err) => Err!(Request(Unauthorized(debug_warn!(
"Failed to verify X-Matrix header: {err}"
@@ -180,120 +155,31 @@ async fn verify<B: AsRef<[u8]> + Sync>(
impl CheckAuth for AccessToken {
type Identity = ClientIdentity;
async fn verify<B: AsRef<[u8]> + Sync>(
async fn check<R: IncomingRequest + Any>(
services: &Services,
output: Self::Output,
_request: &hyper::Request<B>,
_incoming_request: &hyper::Request<&[u8]>,
authentication: Self::Output,
query: AuthQueryParams,
route: TypeId,
) -> Result<Self::Identity> {
if output.is_empty() {
return Err!(Request(Unauthorized("Missing access token.")));
}
if let Some((sender_user, sender_device, status)) =
services.users.find_from_token(&output).await
{
// If the token is expired we return a soft logout
if matches!(status, AccessTokenStatus::Expired) {
return Err(Error::Request(
ErrorKind::UnknownToken(
assign!(UnknownTokenErrorData::new(), { soft_logout: true }),
),
"This token has expired".into(),
StatusCode::UNAUTHORIZED,
));
}
// Locked users can only use /logout and /logout/all
if services
.users
.is_locked(&sender_user)
.await
.is_ok_and(std::convert::identity)
{
if !(route == TypeId::of::<client::session::logout::v3::Request>()
|| route == TypeId::of::<client::session::logout_all::v3::Request>())
{
return Err!(Request(UserLocked("Your account is locked.")));
}
}
Ok(ClientIdentity::User { sender_user, sender_device })
} else if let Ok(appservice_info) = services.appservice.find_from_token(&output).await {
let Ok(sender_user) = query.user_id.clone().map_or_else(
|| {
UserId::parse_with_server_name(
appservice_info.registration.sender_localpart.as_str(),
services.globals.server_name(),
)
},
UserId::parse,
) else {
return Err!(Request(InvalidUsername("Username is invalid.")));
};
if !appservice_info.is_user_match(&sender_user) {
return Err!(Request(Exclusive("User is not in namespace.")));
}
// MSC3202/MSC4190: Handle device_id masquerading for appservices.
// The device_id can be provided via `device_id` or
// `org.matrix.msc3202.device_id` query parameter.
let sender_device = if let Some(device_id) = query
.device_id
.or(query.legacy_device_id)
.as_deref()
.map(Into::into)
{
// Verify the device exists for this user
if services
.users
.get_device_metadata(&sender_user, device_id)
.await
.is_err()
{
return Err!(Request(Forbidden(
"Device does not exist for user or appservice cannot masquerade as this \
device."
)));
}
Some(device_id.to_owned())
} else {
None
};
Ok(ClientIdentity::Appservice {
sender_user,
sender_device,
appservice_info: Box::new(appservice_info),
})
} else {
Err(Error::Request(
ErrorKind::UnknownToken(UnknownTokenErrorData::new()),
"Invalid token".into(),
StatusCode::UNAUTHORIZED,
))
}
check_access_token(services, &authentication, query, TypeId::of::<R>()).await
}
}
impl CheckAuth for AccessTokenOptional {
type Identity = Option<ClientIdentity>;
async fn verify<B: AsRef<[u8]> + Sync>(
async fn check<R: IncomingRequest + Any>(
services: &Services,
output: Self::Output,
request: &hyper::Request<B>,
_incoming_request: &hyper::Request<&[u8]>,
authentication: Self::Output,
query: AuthQueryParams,
route: TypeId,
) -> Result<Self::Identity> {
match output {
| Some(token) =>
<AccessToken as CheckAuth>::verify(services, token, request, query, route)
.await
.map(Some),
| None => Ok(None),
if let Some(authentication) = authentication {
check_access_token(services, &authentication, query, TypeId::of::<R>())
.await
.map(Some)
} else {
Ok(None)
}
}
}
@@ -301,40 +187,31 @@ async fn verify<B: AsRef<[u8]> + Sync>(
impl CheckAuth for AppserviceToken {
type Identity = RegistrationInfo;
async fn verify<B: AsRef<[u8]> + Sync>(
async fn check<R: IncomingRequest + Any>(
services: &Services,
output: Self::Output,
_request: &hyper::Request<B>,
_incoming_request: &hyper::Request<&[u8]>,
authentication: Self::Output,
_query: AuthQueryParams,
_route: TypeId,
) -> Result<Self::Identity> {
if output.is_empty() {
return Err!(Request(Unauthorized("Missing access token.")));
}
let Ok(appservice_info) = services.appservice.find_from_token(&output).await else {
return Err!(Request(Unauthorized("Invalid appservice token.")));
};
Ok(appservice_info)
check_appservice_token(services, &authentication).await
}
}
impl CheckAuth for AppserviceTokenOptional {
type Identity = Option<RegistrationInfo>;
async fn verify<B: AsRef<[u8]> + Sync>(
async fn check<R: IncomingRequest + Any>(
services: &Services,
output: Self::Output,
request: &hyper::Request<B>,
query: AuthQueryParams,
route: TypeId,
_incoming_request: &hyper::Request<&[u8]>,
authentication: Self::Output,
_query: AuthQueryParams,
) -> Result<Self::Identity> {
match output {
| Some(token) =>
<AppserviceToken as CheckAuth>::verify(services, token, request, query, route)
.await
.map(Some),
| None => Ok(None),
if let Some(authentication) = authentication {
check_appservice_token(services, &authentication)
.await
.map(Some)
} else {
Ok(None)
}
}
}
@@ -342,13 +219,12 @@ async fn verify<B: AsRef<[u8]> + Sync>(
impl CheckAuth for NoAuthentication {
type Identity = ();
fn verify<B: AsRef<[u8]> + Sync>(
fn check<R: IncomingRequest + Any>(
_services: &Services,
_output: Self::Output,
_request: &hyper::Request<B>,
_incoming_request: &hyper::Request<&[u8]>,
_authentication: Self::Output,
_query: AuthQueryParams,
_route: TypeId,
) -> impl Future<Output = Result<Self::Identity>> {
) -> impl Future<Output = Result<Self::Identity>> + Send {
std::future::ready(Ok(()))
}
}
@@ -356,31 +232,129 @@ fn verify<B: AsRef<[u8]> + Sync>(
impl CheckAuth for NoAccessToken {
type Identity = Option<ClientIdentity>;
async fn verify<B: AsRef<[u8]> + Sync>(
async fn check<R: IncomingRequest + Any>(
services: &Services,
_output: Self::Output,
request: &hyper::Request<B>,
incoming_request: &hyper::Request<&[u8]>,
_authentication: Self::Output,
query: AuthQueryParams,
route: TypeId,
) -> Result<Self::Identity> {
// We handle these the same as AccessTokenOptional
let token = AccessTokenOptional::extract_authentication(request).map_err(|err| {
err!(Request(Unauthorized(warn!("Failed to extract authorization: {}", err))))
})?;
let authentication = AccessTokenOptional::extract_authentication(incoming_request)
.map_err(|err| {
err!(Request(Unauthorized(warn!("Failed to extract authorization: {}", err))))
})?;
// Check special access restrictions
if (route == TypeId::of::<client::profile::get_avatar_url::v3::Request>()
|| route == TypeId::of::<client::profile::get_display_name::v3::Request>()
|| route == TypeId::of::<client::profile::get_profile_field::v3::Request>()
|| route == TypeId::of::<client::profile::get_profile::v3::Request>())
&& services.config.require_auth_for_profile_requests
&& token.is_none()
{
return Err!(Request(Unauthorized(
"This server requires authentication to access user profiles."
)));
if let Some(authentication) = authentication {
check_access_token(services, &authentication, query, TypeId::of::<R>())
.await
.map(Some)
} else {
Ok(None)
}
<AccessTokenOptional as CheckAuth>::verify(services, token, request, query, route).await
}
}
async fn check_access_token(
services: &Services,
token: &str,
query: AuthQueryParams,
route: TypeId,
) -> Result<ClientIdentity> {
if token.is_empty() {
return Err!(Request(Unauthorized("Empty access token.")));
}
if let Some((sender_user, sender_device, status)) =
services.users.find_from_token(token).await
{
// If the token is expired we return a soft logout
if matches!(status, AccessTokenStatus::Expired) {
return Err(Error::Request(
ErrorKind::UnknownToken(
assign!(UnknownTokenErrorData::new(), { soft_logout: true }),
),
"This token has expired".into(),
StatusCode::UNAUTHORIZED,
));
}
// Locked users can only use /logout and /logout/all
if services
.users
.is_locked(&sender_user)
.await
.is_ok_and(std::convert::identity)
{
if !(route == TypeId::of::<client::session::logout::v3::Request>()
|| route == TypeId::of::<client::session::logout_all::v3::Request>())
{
return Err!(Request(UserLocked("Your account is locked.")));
}
}
Ok(ClientIdentity::User { sender_user, sender_device })
} else if let Ok(appservice_info) = services.appservice.find_from_token(token).await {
let Ok(sender_user) = query.user_id.clone().map_or_else(
|| {
UserId::parse_with_server_name(
appservice_info.registration.sender_localpart.as_str(),
services.globals.server_name(),
)
},
UserId::parse,
) else {
return Err!(Request(InvalidUsername("Username is invalid.")));
};
if !appservice_info.is_user_match(&sender_user) {
return Err!(Request(Exclusive("User is not in namespace.")));
}
// MSC3202/MSC4190: Handle device_id masquerading for appservices.
// The device_id can be provided via `device_id` or
// `org.matrix.msc3202.device_id` query parameter.
let sender_device = if let Some(device_id) = query
.device_id
.or(query.legacy_device_id)
.as_deref()
.map(Into::into)
{
// Verify the device exists for this user
if services
.users
.get_device_metadata(&sender_user, device_id)
.await
.is_err()
{
return Err!(Request(Forbidden(
"Device does not exist for user or appservice cannot masquerade as this \
device."
)));
}
Some(device_id.to_owned())
} else {
None
};
Ok(ClientIdentity::Appservice {
sender_user,
sender_device,
appservice_info: Box::new(appservice_info),
})
} else {
Err(Error::Request(
ErrorKind::UnknownToken(UnknownTokenErrorData::new()),
"Invalid token".into(),
StatusCode::UNAUTHORIZED,
))
}
}
async fn check_appservice_token(services: &Services, token: &str) -> Result<RegistrationInfo> {
let Ok(appservice_info) = services.appservice.find_from_token(token).await else {
return Err!(Request(Unauthorized("Invalid appservice token.")));
};
Ok(appservice_info)
}