Skip to main content

mas_handlers/graphql/
mod.rs

1// Copyright 2025, 2026 Element Creations Ltd.
2// Copyright 2024, 2025 New Vector Ltd.
3// Copyright 2022-2024 The Matrix.org Foundation C.I.C.
4//
5// SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-Element-Commercial
6// Please see LICENSE files in the repository root for full details.
7
8#![allow(
9    clippy::module_name_repetitions,
10    // async-graphql's `#[Object]` requires resolvers to be `async fn`, and its
11    // `Enum` derive emits async trait impls without awaits
12    clippy::unused_async_trait_impl,
13)]
14
15use std::{net::IpAddr, ops::Deref, sync::Arc};
16
17use async_graphql::{
18    EmptySubscription, InputObject,
19    extensions::Tracing,
20    http::MultipartOptions,
21    parser::types::{DocumentOperations, OperationType},
22};
23use axum::{
24    Extension, Json,
25    body::Body,
26    extract::{RawQuery, State as AxumState},
27    http::StatusCode,
28    response::{IntoResponse, Response},
29};
30use axum_extra::typed_header::TypedHeader;
31use chrono::{DateTime, Utc};
32use futures_util::TryStreamExt;
33use headers::{Authorization, ContentType, HeaderValue, authorization::Bearer};
34use hyper::header::CACHE_CONTROL;
35use mas_axum_utils::{
36    InternalError, RecordAsRequester, SessionInfo, SessionInfoExt, cookies::CookieJar,
37    sentry::SentryEventID,
38};
39use mas_data_model::{
40    BoxClock, BoxRng, BrowserSession, Clock, Session, SiteConfig, SystemClock, User,
41};
42use mas_matrix::HomeserverConnection;
43use mas_policy::{InstantiateError, Policy, PolicyFactory};
44use mas_router::UrlBuilder;
45use mas_storage::{BoxRepository, BoxRepositoryFactory, RepositoryError};
46use opentelemetry_semantic_conventions::trace::{
47    GRAPHQL_DOCUMENT, GRAPHQL_OPERATION_NAME, GRAPHQL_OPERATION_TYPE,
48};
49use rand::{SeedableRng, thread_rng};
50use rand_chacha::ChaChaRng;
51use state::has_session_ended;
52use tracing::{Instrument, info_span};
53use ulid::Ulid;
54
55mod model;
56mod mutations;
57mod query;
58mod state;
59
60pub use self::state::{BoxState, State};
61use self::{
62    model::{CreationEvent, Node},
63    mutations::Mutation,
64    query::Query,
65};
66use crate::{
67    BoundActivityTracker, Limiter, RequesterFingerprint, impl_from_error_for_route,
68    passwords::PasswordManager,
69};
70
71#[cfg(test)]
72mod tests;
73
74/// Extra parameters we get from the listener configuration, because they are
75/// per-listener options. We pass them through request extensions.
76#[derive(Debug, Clone)]
77pub struct ExtraRouterParameters {
78    pub undocumented_oauth2_access: bool,
79}
80
81struct GraphQLState {
82    repository_factory: BoxRepositoryFactory,
83    homeserver_connection: Arc<dyn HomeserverConnection>,
84    policy_factory: Arc<PolicyFactory>,
85    site_config: SiteConfig,
86    password_manager: PasswordManager,
87    url_builder: UrlBuilder,
88    limiter: Limiter,
89}
90
91#[async_trait::async_trait]
92impl state::State for GraphQLState {
93    async fn repository(&self) -> Result<BoxRepository, RepositoryError> {
94        self.repository_factory.create().await
95    }
96
97    async fn policy(&self) -> Result<Policy, InstantiateError> {
98        self.policy_factory.instantiate().await
99    }
100
101    fn password_manager(&self) -> PasswordManager {
102        self.password_manager.clone()
103    }
104
105    fn site_config(&self) -> &SiteConfig {
106        &self.site_config
107    }
108
109    fn homeserver_connection(&self) -> &dyn HomeserverConnection {
110        self.homeserver_connection.as_ref()
111    }
112
113    fn url_builder(&self) -> &UrlBuilder {
114        &self.url_builder
115    }
116
117    fn limiter(&self) -> &Limiter {
118        &self.limiter
119    }
120
121    fn clock(&self) -> BoxClock {
122        let clock = SystemClock::default();
123        Box::new(clock)
124    }
125
126    fn rng(&self) -> BoxRng {
127        #[expect(clippy::disallowed_methods)]
128        let rng = thread_rng();
129
130        let rng = ChaChaRng::from_rng(rng).expect("Failed to seed rng");
131        Box::new(rng)
132    }
133}
134
135#[must_use]
136pub fn schema(
137    repository_factory: BoxRepositoryFactory,
138    policy_factory: &Arc<PolicyFactory>,
139    homeserver_connection: impl HomeserverConnection + 'static,
140    site_config: SiteConfig,
141    password_manager: PasswordManager,
142    url_builder: UrlBuilder,
143    limiter: Limiter,
144) -> Schema {
145    let state = GraphQLState {
146        repository_factory,
147        policy_factory: Arc::clone(policy_factory),
148        homeserver_connection: Arc::new(homeserver_connection),
149        site_config,
150        password_manager,
151        url_builder,
152        limiter,
153    };
154    let state: BoxState = Box::new(state);
155
156    schema_builder().extension(Tracing).data(state).finish()
157}
158
159fn span_and_operation_for_graphql_request(
160    request: &mut async_graphql::Request,
161) -> (tracing::Span, GraphQLOperation) {
162    let span = info_span!(
163        "GraphQL operation",
164        "otel.name" = tracing::field::Empty,
165        "otel.kind" = "server",
166        { GRAPHQL_DOCUMENT } = request.query,
167        { GRAPHQL_OPERATION_NAME } = tracing::field::Empty,
168        { GRAPHQL_OPERATION_TYPE } = tracing::field::Empty,
169    );
170
171    let mut graphql_operation = GraphQLOperation {
172        operation_type: None,
173        operation_name: None,
174    };
175
176    // We need to clone the operation_name before parsing the query, else we're
177    // going to have a borrow conflict between request.parsed_query() and
178    // request.operation_name
179    let operation_name = request.operation_name.clone();
180    if let Ok(document) = request.parsed_query() {
181        match (&document.operations, operation_name) {
182            // A single anonymous operation, with no name requested: the
183            // document defines no name for it, so we only record the type.
184            (DocumentOperations::Single(operation), None) => {
185                span.record("otel.name", format!("GraphQL {}", operation.node.ty));
186                span.record(
187                    GRAPHQL_OPERATION_TYPE,
188                    tracing::field::display(operation.node.ty),
189                );
190                graphql_operation.operation_type = Some(operation.node.ty);
191            }
192
193            (DocumentOperations::Multiple(operations), Some(name)) => {
194                if let Some((name, operation)) = operations.get_key_value(name.as_str()) {
195                    span.record(
196                        "otel.name",
197                        format!("GraphQL {} {}", operation.node.ty, name),
198                    );
199                    span.record(
200                        GRAPHQL_OPERATION_TYPE,
201                        tracing::field::display(operation.node.ty),
202                    );
203                    span.record(GRAPHQL_OPERATION_NAME, tracing::field::display(name));
204                    graphql_operation.operation_type = Some(operation.node.ty);
205                    graphql_operation.operation_name = Some(name.to_string());
206                }
207            }
208
209            (DocumentOperations::Multiple(operations), None) if operations.len() == 1 => {
210                let mut iter = operations.iter();
211                let (name, operation) = iter.next().unwrap();
212                span.record(
213                    "otel.name",
214                    format!("GraphQL {} {}", operation.node.ty, name),
215                );
216                span.record(
217                    GRAPHQL_OPERATION_TYPE,
218                    tracing::field::display(operation.node.ty),
219                );
220                span.record(GRAPHQL_OPERATION_NAME, name.as_ref());
221                graphql_operation.operation_type = Some(operation.node.ty);
222                graphql_operation.operation_name = Some(name.to_string());
223            }
224
225            // Cases the executor rejects, so we don't record a misleading
226            // operation: a single anonymous operation with a name requested, or
227            // several named operations with no requested name (ambiguous).
228            (DocumentOperations::Single(_), Some(_)) | (DocumentOperations::Multiple(_), None) => {}
229        }
230    }
231
232    (span, graphql_operation)
233}
234
235/// The GraphQL operation being executed, attached to the response extensions so
236/// the HTTP logging middleware can record it on the request log line.
237#[derive(Clone, Debug)]
238pub struct GraphQLOperation {
239    /// The type of the operation: query, mutation or subscription.
240    pub operation_type: Option<OperationType>,
241    /// The name of the operation, as defined in the query document.
242    pub operation_name: Option<String>,
243}
244
245#[derive(thiserror::Error, Debug)]
246pub enum RouteError {
247    #[error(transparent)]
248    Internal(Box<dyn std::error::Error + Send + Sync + 'static>),
249
250    #[error("Loading of some database objects failed")]
251    LoadFailed,
252
253    #[error("Invalid access token")]
254    InvalidToken,
255
256    #[error("Missing scope")]
257    MissingScope,
258
259    #[error(transparent)]
260    ParseRequest(#[from] async_graphql::ParseRequestError),
261}
262
263impl_from_error_for_route!(mas_storage::RepositoryError);
264
265impl IntoResponse for RouteError {
266    fn into_response(self) -> Response {
267        let event_id = sentry::capture_error(&self);
268
269        let response = match self {
270            e @ (Self::Internal(_) | Self::LoadFailed) => {
271                let error = async_graphql::Error::new_with_source(e);
272                (
273                    StatusCode::INTERNAL_SERVER_ERROR,
274                    Json(serde_json::json!({"errors": [error]})),
275                )
276                    .into_response()
277            }
278
279            Self::InvalidToken => {
280                let error = async_graphql::Error::new("Invalid token");
281                (
282                    StatusCode::UNAUTHORIZED,
283                    Json(serde_json::json!({"errors": [error]})),
284                )
285                    .into_response()
286            }
287
288            Self::MissingScope => {
289                let error = async_graphql::Error::new("Missing urn:mas:graphql:* scope");
290                (
291                    StatusCode::UNAUTHORIZED,
292                    Json(serde_json::json!({"errors": [error]})),
293                )
294                    .into_response()
295            }
296
297            Self::ParseRequest(e) => {
298                let error = async_graphql::Error::new_with_source(e);
299                (
300                    StatusCode::BAD_REQUEST,
301                    Json(serde_json::json!({"errors": [error]})),
302                )
303                    .into_response()
304            }
305        };
306
307        (SentryEventID::from(event_id), response).into_response()
308    }
309}
310
311async fn get_requester(
312    undocumented_oauth2_access: bool,
313    clock: &impl Clock,
314    activity_tracker: &BoundActivityTracker,
315    mut repo: BoxRepository,
316    session_info: &SessionInfo,
317    user_agent: Option<String>,
318    token: Option<&str>,
319) -> Result<Requester, RouteError> {
320    let entity = if let Some(token) = token {
321        // If we haven't enabled undocumented_oauth2_access on the listener, we bail out
322        if !undocumented_oauth2_access {
323            return Err(RouteError::InvalidToken);
324        }
325
326        let token = repo
327            .oauth2_access_token()
328            .find_by_token(token)
329            .await?
330            .ok_or(RouteError::InvalidToken)?;
331
332        let session = repo
333            .oauth2_session()
334            .lookup(token.session_id)
335            .await?
336            .ok_or(RouteError::LoadFailed)?;
337
338        activity_tracker
339            .record_oauth2_session(clock, &session)
340            .await;
341
342        // Load the user if there is one
343        let user = if let Some(user_id) = session.user_id {
344            let user = repo
345                .user()
346                .lookup(user_id)
347                .await?
348                .ok_or(RouteError::LoadFailed)?;
349            Some(user)
350        } else {
351            None
352        };
353
354        // If there is a user for this session, check that it is not locked
355        let user_valid = user.as_ref().is_none_or(User::is_valid);
356
357        if !token.is_valid(clock.now()) || !session.is_valid() || !user_valid {
358            return Err(RouteError::InvalidToken);
359        }
360
361        if !session.scope.contains("urn:mas:graphql:*") {
362            return Err(RouteError::MissingScope);
363        }
364
365        if let Some(user) = &user {
366            user.maybe_record_as_requester();
367        }
368
369        RequestingEntity::OAuth2Session(Box::new((session, user)))
370    } else {
371        let maybe_session = session_info.load_active_session(&mut repo).await?;
372
373        if let Some(session) = maybe_session.as_ref() {
374            activity_tracker
375                .record_browser_session(clock, session)
376                .await;
377        }
378
379        RequestingEntity::from(maybe_session)
380    };
381
382    let requester = Requester {
383        entity,
384        ip_address: activity_tracker.ip(),
385        user_agent,
386    };
387
388    repo.cancel().await?;
389    Ok(requester)
390}
391
392pub async fn post(
393    AxumState(schema): AxumState<Schema>,
394    Extension(ExtraRouterParameters {
395        undocumented_oauth2_access,
396    }): Extension<ExtraRouterParameters>,
397    clock: BoxClock,
398    repo: BoxRepository,
399    activity_tracker: BoundActivityTracker,
400    cookie_jar: CookieJar,
401    content_type: Option<TypedHeader<ContentType>>,
402    authorization: Option<TypedHeader<Authorization<Bearer>>>,
403    user_agent: Option<TypedHeader<headers::UserAgent>>,
404    body: Body,
405) -> Result<impl IntoResponse, RouteError> {
406    let body = body.into_data_stream();
407    let token = authorization
408        .as_ref()
409        .map(|TypedHeader(Authorization(bearer))| bearer.token());
410    let user_agent = user_agent.map(|TypedHeader(h)| h.to_string());
411    let (session_info, mut cookie_jar) = cookie_jar.session_info();
412    let requester = get_requester(
413        undocumented_oauth2_access,
414        &clock,
415        &activity_tracker,
416        repo,
417        &session_info,
418        user_agent,
419        token,
420    )
421    .await?;
422
423    let content_type = content_type.map(|TypedHeader(h)| h.to_string());
424
425    let mut request = async_graphql::http::receive_body(
426        content_type,
427        body.map_err(std::io::Error::other).into_async_read(),
428        MultipartOptions::default(),
429    )
430    .await?
431    .data(requester); // XXX: this should probably return another error response?
432
433    let (span, operation) = span_and_operation_for_graphql_request(&mut request);
434    let mut response = schema.execute(request).instrument(span).await;
435
436    if has_session_ended(&mut response) {
437        let session_info = session_info.mark_session_ended(clock.now());
438        cookie_jar = cookie_jar.update_session_info(&session_info);
439    }
440
441    let cache_control = response
442        .cache_control
443        .value()
444        .and_then(|v| HeaderValue::from_str(&v).ok())
445        .map(|h| [(CACHE_CONTROL, h)]);
446
447    let headers = response.http_headers.clone();
448
449    Ok((
450        headers,
451        cache_control,
452        cookie_jar,
453        Extension(operation),
454        Json(response),
455    ))
456}
457
458pub async fn get(
459    AxumState(schema): AxumState<Schema>,
460    Extension(ExtraRouterParameters {
461        undocumented_oauth2_access,
462    }): Extension<ExtraRouterParameters>,
463    clock: BoxClock,
464    repo: BoxRepository,
465    activity_tracker: BoundActivityTracker,
466    cookie_jar: CookieJar,
467    authorization: Option<TypedHeader<Authorization<Bearer>>>,
468    user_agent: Option<TypedHeader<headers::UserAgent>>,
469    RawQuery(query): RawQuery,
470) -> Result<impl IntoResponse, InternalError> {
471    let token = authorization
472        .as_ref()
473        .map(|TypedHeader(Authorization(bearer))| bearer.token());
474    let user_agent = user_agent.map(|TypedHeader(h)| h.to_string());
475    let (session_info, mut cookie_jar) = cookie_jar.session_info();
476    let requester = get_requester(
477        undocumented_oauth2_access,
478        &clock,
479        &activity_tracker,
480        repo,
481        &session_info,
482        user_agent,
483        token,
484    )
485    .await?;
486
487    let mut request =
488        async_graphql::http::parse_query_string(&query.unwrap_or_default())?.data(requester);
489
490    let (span, operation) = span_and_operation_for_graphql_request(&mut request);
491    let mut response = schema.execute(request).instrument(span).await;
492
493    if has_session_ended(&mut response) {
494        let session_info = session_info.mark_session_ended(clock.now());
495        cookie_jar = cookie_jar.update_session_info(&session_info);
496    }
497
498    let cache_control = response
499        .cache_control
500        .value()
501        .and_then(|v| HeaderValue::from_str(&v).ok())
502        .map(|h| [(CACHE_CONTROL, h)]);
503
504    let headers = response.http_headers.clone();
505
506    Ok((
507        headers,
508        cache_control,
509        cookie_jar,
510        Extension(operation),
511        Json(response),
512    ))
513}
514
515pub type Schema = async_graphql::Schema<Query, Mutation, EmptySubscription>;
516pub type SchemaBuilder = async_graphql::SchemaBuilder<Query, Mutation, EmptySubscription>;
517
518#[must_use]
519pub fn schema_builder() -> SchemaBuilder {
520    async_graphql::Schema::build(Query::new(), Mutation::new(), EmptySubscription)
521        .register_output_type::<Node>()
522        .register_output_type::<CreationEvent>()
523}
524
525pub struct Requester {
526    entity: RequestingEntity,
527    ip_address: Option<IpAddr>,
528    user_agent: Option<String>,
529}
530
531impl Requester {
532    pub fn fingerprint(&self) -> RequesterFingerprint {
533        if let Some(ip) = self.ip_address {
534            RequesterFingerprint::new(ip)
535        } else {
536            RequesterFingerprint::EMPTY
537        }
538    }
539
540    pub fn for_policy(&self) -> mas_policy::Requester {
541        mas_policy::Requester {
542            ip_address: self.ip_address,
543            user_agent: self.user_agent.clone(),
544        }
545    }
546}
547
548impl Deref for Requester {
549    type Target = RequestingEntity;
550
551    fn deref(&self) -> &Self::Target {
552        &self.entity
553    }
554}
555
556/// The identity of the requester.
557#[derive(Debug, Clone, Default, PartialEq, Eq)]
558pub enum RequestingEntity {
559    /// The requester presented no authentication information.
560    #[default]
561    Anonymous,
562
563    /// The requester is a browser session, stored in a cookie.
564    BrowserSession(Box<BrowserSession>),
565
566    /// The requester is a `OAuth2` session, with an access token.
567    OAuth2Session(Box<(Session, Option<User>)>),
568}
569
570trait OwnerId {
571    fn owner_id(&self) -> Option<Ulid>;
572}
573
574impl OwnerId for User {
575    fn owner_id(&self) -> Option<Ulid> {
576        Some(self.id)
577    }
578}
579
580impl OwnerId for BrowserSession {
581    fn owner_id(&self) -> Option<Ulid> {
582        Some(self.user.id)
583    }
584}
585
586impl OwnerId for mas_data_model::UserEmail {
587    fn owner_id(&self) -> Option<Ulid> {
588        Some(self.user_id)
589    }
590}
591
592impl OwnerId for Session {
593    fn owner_id(&self) -> Option<Ulid> {
594        self.user_id
595    }
596}
597
598impl OwnerId for mas_data_model::CompatSession {
599    fn owner_id(&self) -> Option<Ulid> {
600        Some(self.user_id)
601    }
602}
603
604impl OwnerId for mas_data_model::UpstreamOAuthLink {
605    fn owner_id(&self) -> Option<Ulid> {
606        self.user_id
607    }
608}
609
610/// A dumb wrapper around a `Ulid` to implement `OwnerId` for it.
611pub struct UserId(Ulid);
612
613impl OwnerId for UserId {
614    fn owner_id(&self) -> Option<Ulid> {
615        Some(self.0)
616    }
617}
618
619impl RequestingEntity {
620    fn browser_session(&self) -> Option<&BrowserSession> {
621        match self {
622            Self::BrowserSession(session) => Some(session),
623            Self::OAuth2Session(_) | Self::Anonymous => None,
624        }
625    }
626
627    fn user(&self) -> Option<&User> {
628        match self {
629            Self::BrowserSession(session) => Some(&session.user),
630            Self::OAuth2Session(tuple) => tuple.1.as_ref(),
631            Self::Anonymous => None,
632        }
633    }
634
635    fn oauth2_session(&self) -> Option<&Session> {
636        match self {
637            Self::OAuth2Session(tuple) => Some(&tuple.0),
638            Self::BrowserSession(_) | Self::Anonymous => None,
639        }
640    }
641
642    /// Returns true if the requester can access the resource.
643    fn is_owner_or_admin(&self, resource: &impl OwnerId) -> bool {
644        // If the requester is an admin, they can do anything.
645        if self.is_admin() {
646            return true;
647        }
648
649        // Otherwise, they must be the owner of the resource.
650        let Some(owner_id) = resource.owner_id() else {
651            return false;
652        };
653
654        let Some(user) = self.user() else {
655            return false;
656        };
657
658        user.id == owner_id
659    }
660
661    fn is_admin(&self) -> bool {
662        match self {
663            Self::OAuth2Session(tuple) => {
664                // TODO: is this the right scope?
665                // This has to be in sync with the policy
666                tuple.0.scope.contains("urn:mas:admin")
667            }
668            Self::BrowserSession(_) | Self::Anonymous => false,
669        }
670    }
671}
672
673impl From<BrowserSession> for RequestingEntity {
674    fn from(session: BrowserSession) -> Self {
675        Self::BrowserSession(Box::new(session))
676    }
677}
678
679impl<T> From<Option<T>> for RequestingEntity
680where
681    T: Into<RequestingEntity>,
682{
683    fn from(session: Option<T>) -> Self {
684        session.map(Into::into).unwrap_or_default()
685    }
686}
687
688/// A filter for dates, with a lower bound and an upper bound
689#[derive(InputObject, Default, Clone, Copy)]
690pub struct DateFilter {
691    /// The lower bound of the date range
692    after: Option<DateTime<Utc>>,
693
694    /// The upper bound of the date range
695    before: Option<DateTime<Utc>>,
696}