Skip to main content

mas_storage_pg/upstream_oauth2/
mod.rs

1// Copyright 2024, 2025 New Vector Ltd.
2// Copyright 2022-2024 The Matrix.org Foundation C.I.C.
3//
4// SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-Element-Commercial
5// Please see LICENSE files in the repository root for full details.
6
7//! A module containing the PostgreSQL implementation of the repositories
8//! related to the upstream OAuth 2.0 providers
9
10mod link;
11mod provider;
12mod session;
13
14pub use self::{
15    link::PgUpstreamOAuthLinkRepository, provider::PgUpstreamOAuthProviderRepository,
16    session::PgUpstreamOAuthSessionRepository,
17};
18
19#[cfg(test)]
20mod tests {
21    use chrono::Duration;
22    use mas_data_model::{
23        UpstreamOAuthProviderClaimsImports, UpstreamOAuthProviderOnBackchannelLogout,
24        UpstreamOAuthProviderTokenAuthMethod, clock::MockClock,
25    };
26    use mas_iana::jose::JsonWebSignatureAlg;
27    use mas_storage::{
28        Pagination, RepositoryAccess,
29        upstream_oauth2::{
30            UpstreamOAuthLinkFilter, UpstreamOAuthLinkRepository, UpstreamOAuthProviderFilter,
31            UpstreamOAuthProviderParams, UpstreamOAuthProviderRepository,
32            UpstreamOAuthSessionFilter, UpstreamOAuthSessionRepository,
33        },
34        user::UserRepository,
35    };
36    use oauth2_types::scope::{OPENID, Scope};
37    use rand::SeedableRng;
38    use sqlx::PgPool;
39
40    use crate::PgRepository;
41
42    #[sqlx::test(migrator = "crate::MIGRATOR")]
43    async fn test_repository(pool: PgPool) {
44        let mut rng = rand_chacha::ChaChaRng::seed_from_u64(42);
45        let clock = MockClock::default();
46        let mut repo = PgRepository::from_pool(&pool).await.unwrap();
47
48        // The provider list should be empty at the start
49        let all_providers = repo.upstream_oauth_provider().all_enabled().await.unwrap();
50        assert!(all_providers.is_empty());
51
52        // Let's add a provider
53        let provider = repo
54            .upstream_oauth_provider()
55            .add(
56                &mut rng,
57                &clock,
58                UpstreamOAuthProviderParams {
59                    issuer: Some("https://example.com/".to_owned()),
60                    human_name: None,
61                    brand_name: None,
62                    scope: Scope::from_iter([OPENID]),
63                    token_endpoint_auth_method: UpstreamOAuthProviderTokenAuthMethod::None,
64                    id_token_signed_response_alg: JsonWebSignatureAlg::Rs256,
65                    fetch_userinfo: false,
66                    userinfo_signed_response_alg: None,
67                    token_endpoint_signing_alg: None,
68                    client_id: "client-id".to_owned(),
69                    encrypted_client_secret: None,
70                    claims_imports: UpstreamOAuthProviderClaimsImports::default(),
71                    token_endpoint_override: None,
72                    authorization_endpoint_override: None,
73                    userinfo_endpoint_override: None,
74                    jwks_uri_override: None,
75                    discovery_mode: mas_data_model::UpstreamOAuthProviderDiscoveryMode::Oidc,
76                    pkce_mode: mas_data_model::UpstreamOAuthProviderPkceMode::Auto,
77                    response_mode: None,
78                    additional_authorization_parameters: Vec::new(),
79                    forward_login_hint: false,
80                    ui_order: 0,
81                    on_backchannel_logout: UpstreamOAuthProviderOnBackchannelLogout::DoNothing,
82                    registration_token_required: false,
83                },
84            )
85            .await
86            .unwrap();
87
88        // Look it up in the database
89        let provider = repo
90            .upstream_oauth_provider()
91            .lookup(provider.id)
92            .await
93            .unwrap()
94            .expect("provider to be found in the database");
95        assert_eq!(provider.issuer.as_deref(), Some("https://example.com/"));
96        assert_eq!(provider.client_id, "client-id");
97
98        // It should be in the list of all providers
99        let providers = repo.upstream_oauth_provider().all_enabled().await.unwrap();
100        assert_eq!(providers.len(), 1);
101        assert_eq!(providers[0].issuer.as_deref(), Some("https://example.com/"));
102        assert_eq!(providers[0].client_id, "client-id");
103
104        // Start a session
105        let session = repo
106            .upstream_oauth_session()
107            .add(
108                &mut rng,
109                &clock,
110                &provider,
111                "some-state".to_owned(),
112                None,
113                Some("some-nonce".to_owned()),
114            )
115            .await
116            .unwrap();
117
118        // Look it up in the database
119        let session = repo
120            .upstream_oauth_session()
121            .lookup(session.id)
122            .await
123            .unwrap()
124            .expect("session to be found in the database");
125        assert_eq!(session.provider_id, provider.id);
126        assert_eq!(session.link_id(), None);
127        assert!(session.is_pending());
128        assert!(!session.is_completed());
129        assert!(!session.is_consumed());
130
131        // Create a link
132        let link = repo
133            .upstream_oauth_link()
134            .add(
135                &mut rng,
136                &clock,
137                &provider,
138                "a-subject".to_owned(),
139                Some("alice@example.com".to_owned()),
140            )
141            .await
142            .unwrap();
143
144        // We can look it up by its ID
145        repo.upstream_oauth_link()
146            .lookup(link.id)
147            .await
148            .unwrap()
149            .expect("link to be found in database");
150
151        // or by its subject
152        let link = repo
153            .upstream_oauth_link()
154            .find_by_subject(&provider, "a-subject")
155            .await
156            .unwrap()
157            .expect("link to be found in database");
158        assert_eq!(link.subject, "a-subject");
159        assert_eq!(link.provider_id, provider.id);
160
161        let session = repo
162            .upstream_oauth_session()
163            .complete_with_link(&clock, session, &link, None, None, None, None)
164            .await
165            .unwrap();
166        // Reload the session
167        let session = repo
168            .upstream_oauth_session()
169            .lookup(session.id)
170            .await
171            .unwrap()
172            .expect("session to be found in the database");
173        assert!(session.is_completed());
174        assert!(!session.is_consumed());
175        assert_eq!(session.link_id(), Some(link.id));
176
177        // We need to create a user and start a browser session to consume the session
178        let user = repo
179            .user()
180            .add(&mut rng, &clock, "john".to_owned())
181            .await
182            .unwrap();
183        let browser_session = repo
184            .browser_session()
185            .add(&mut rng, &clock, &user, None)
186            .await
187            .unwrap();
188
189        let session = repo
190            .upstream_oauth_session()
191            .consume(&clock, session, &browser_session)
192            .await
193            .unwrap();
194
195        // Reload the session
196        let session = repo
197            .upstream_oauth_session()
198            .lookup(session.id)
199            .await
200            .unwrap()
201            .expect("session to be found in the database");
202        assert!(session.is_consumed());
203
204        repo.upstream_oauth_link()
205            .associate_to_user(&link, &user)
206            .await
207            .unwrap();
208
209        // XXX: we should also try other combinations of the filter
210        let filter = UpstreamOAuthLinkFilter::new()
211            .for_user(&user)
212            .for_provider(&provider)
213            .for_subject("a-subject")
214            .enabled_providers_only();
215
216        let links = repo
217            .upstream_oauth_link()
218            .list(filter, Pagination::first(10))
219            .await
220            .unwrap();
221        assert!(!links.has_previous_page);
222        assert!(!links.has_next_page);
223        assert_eq!(links.edges.len(), 1);
224        assert_eq!(links.edges[0].node.id, link.id);
225        assert_eq!(links.edges[0].node.user_id, Some(user.id));
226
227        assert_eq!(repo.upstream_oauth_link().count(filter).await.unwrap(), 1);
228
229        // Filtering on a case-insensitive partial human_account_name should match
230        let matching_filter = UpstreamOAuthLinkFilter::new().matching_human_account_name("LICE");
231        let links = repo
232            .upstream_oauth_link()
233            .list(matching_filter, Pagination::first(10))
234            .await
235            .unwrap();
236        assert_eq!(links.edges.len(), 1);
237        assert_eq!(links.edges[0].node.id, link.id);
238        assert_eq!(
239            repo.upstream_oauth_link()
240                .count(matching_filter)
241                .await
242                .unwrap(),
243            1
244        );
245
246        // A non-matching human_account_name should return nothing
247        let non_matching_filter =
248            UpstreamOAuthLinkFilter::new().matching_human_account_name("nope");
249        let links = repo
250            .upstream_oauth_link()
251            .list(non_matching_filter, Pagination::first(10))
252            .await
253            .unwrap();
254        assert_eq!(links.edges.len(), 0);
255        assert_eq!(
256            repo.upstream_oauth_link()
257                .count(non_matching_filter)
258                .await
259                .unwrap(),
260            0
261        );
262
263        // There should be exactly one enabled provider
264        assert_eq!(
265            repo.upstream_oauth_provider()
266                .count(UpstreamOAuthProviderFilter::new())
267                .await
268                .unwrap(),
269            1
270        );
271        assert_eq!(
272            repo.upstream_oauth_provider()
273                .count(UpstreamOAuthProviderFilter::new().enabled_only())
274                .await
275                .unwrap(),
276            1
277        );
278        assert_eq!(
279            repo.upstream_oauth_provider()
280                .count(UpstreamOAuthProviderFilter::new().disabled_only())
281                .await
282                .unwrap(),
283            0
284        );
285
286        // Disable the provider
287        repo.upstream_oauth_provider()
288            .disable(&clock, provider.clone())
289            .await
290            .unwrap();
291
292        // There should be exactly one disabled provider
293        assert_eq!(
294            repo.upstream_oauth_provider()
295                .count(UpstreamOAuthProviderFilter::new())
296                .await
297                .unwrap(),
298            1
299        );
300        assert_eq!(
301            repo.upstream_oauth_provider()
302                .count(UpstreamOAuthProviderFilter::new().enabled_only())
303                .await
304                .unwrap(),
305            0
306        );
307        assert_eq!(
308            repo.upstream_oauth_provider()
309                .count(UpstreamOAuthProviderFilter::new().disabled_only())
310                .await
311                .unwrap(),
312            1
313        );
314
315        // Test listing and counting sessions
316        let session_filter = UpstreamOAuthSessionFilter::new().for_provider(&provider);
317
318        // Count the sessions for the provider
319        let session_count = repo
320            .upstream_oauth_session()
321            .count(session_filter)
322            .await
323            .unwrap();
324        assert_eq!(session_count, 1);
325
326        // List the sessions for the provider
327        let session_page = repo
328            .upstream_oauth_session()
329            .list(session_filter, Pagination::first(10))
330            .await
331            .unwrap();
332
333        assert_eq!(session_page.edges.len(), 1);
334        assert_eq!(session_page.edges[0].node.id, session.id);
335        assert!(!session_page.has_next_page);
336        assert!(!session_page.has_previous_page);
337
338        // Try deleting the provider
339        repo.upstream_oauth_provider()
340            .delete(provider)
341            .await
342            .unwrap();
343        assert_eq!(
344            repo.upstream_oauth_provider()
345                .count(UpstreamOAuthProviderFilter::new())
346                .await
347                .unwrap(),
348            0
349        );
350    }
351
352    /// Test that the pagination works as expected in the upstream OAuth
353    /// provider repository
354    #[sqlx::test(migrator = "crate::MIGRATOR")]
355    async fn test_provider_repository_pagination(pool: PgPool) {
356        let scope = Scope::from_iter([OPENID]);
357
358        let mut rng = rand_chacha::ChaChaRng::seed_from_u64(42);
359        let clock = MockClock::default();
360        let mut repo = PgRepository::from_pool(&pool).await.unwrap();
361
362        let filter = UpstreamOAuthProviderFilter::new();
363
364        // Count the number of providers before we start
365        assert_eq!(
366            repo.upstream_oauth_provider().count(filter).await.unwrap(),
367            0
368        );
369
370        let mut ids = Vec::with_capacity(20);
371        // Create 20 providers
372        for idx in 0..20 {
373            let client_id = format!("client-{idx}");
374            let provider = repo
375                .upstream_oauth_provider()
376                .add(
377                    &mut rng,
378                    &clock,
379                    UpstreamOAuthProviderParams {
380                        issuer: None,
381                        human_name: None,
382                        brand_name: None,
383                        scope: scope.clone(),
384                        token_endpoint_auth_method: UpstreamOAuthProviderTokenAuthMethod::None,
385                        fetch_userinfo: false,
386                        userinfo_signed_response_alg: None,
387                        token_endpoint_signing_alg: None,
388                        id_token_signed_response_alg: JsonWebSignatureAlg::Rs256,
389                        client_id,
390                        encrypted_client_secret: None,
391                        claims_imports: UpstreamOAuthProviderClaimsImports::default(),
392                        token_endpoint_override: None,
393                        authorization_endpoint_override: None,
394                        userinfo_endpoint_override: None,
395                        jwks_uri_override: None,
396                        discovery_mode: mas_data_model::UpstreamOAuthProviderDiscoveryMode::Oidc,
397                        pkce_mode: mas_data_model::UpstreamOAuthProviderPkceMode::Auto,
398                        response_mode: None,
399                        additional_authorization_parameters: Vec::new(),
400                        forward_login_hint: false,
401                        ui_order: 0,
402                        on_backchannel_logout: UpstreamOAuthProviderOnBackchannelLogout::DoNothing,
403                        registration_token_required: false,
404                    },
405                )
406                .await
407                .unwrap();
408            ids.push(provider.id);
409            clock.advance(Duration::microseconds(10 * 1000 * 1000));
410        }
411
412        // Now we have 20 providers
413        assert_eq!(
414            repo.upstream_oauth_provider().count(filter).await.unwrap(),
415            20
416        );
417
418        // Lookup the first 10 items
419        let page = repo
420            .upstream_oauth_provider()
421            .list(filter, Pagination::first(10))
422            .await
423            .unwrap();
424
425        // It returned the first 10 items
426        assert!(page.has_next_page);
427        let edge_ids: Vec<_> = page.edges.iter().map(|p| p.node.id).collect();
428        assert_eq!(&edge_ids, &ids[..10]);
429
430        // Getting the same page with the "enabled only" filter should return the same
431        // results
432        let other_page = repo
433            .upstream_oauth_provider()
434            .list(filter.enabled_only(), Pagination::first(10))
435            .await
436            .unwrap();
437
438        assert_eq!(page, other_page);
439
440        // Lookup the next 10 items
441        let page = repo
442            .upstream_oauth_provider()
443            .list(filter, Pagination::first(10).after(ids[9]))
444            .await
445            .unwrap();
446
447        // It returned the next 10 items
448        assert!(!page.has_next_page);
449        let edge_ids: Vec<_> = page.edges.iter().map(|p| p.node.id).collect();
450        assert_eq!(&edge_ids, &ids[10..]);
451
452        // Lookup the last 10 items
453        let page = repo
454            .upstream_oauth_provider()
455            .list(filter, Pagination::last(10))
456            .await
457            .unwrap();
458
459        // It returned the last 10 items
460        assert!(page.has_previous_page);
461        let edge_ids: Vec<_> = page.edges.iter().map(|p| p.node.id).collect();
462        assert_eq!(&edge_ids, &ids[10..]);
463
464        // Lookup the previous 10 items
465        let page = repo
466            .upstream_oauth_provider()
467            .list(filter, Pagination::last(10).before(ids[10]))
468            .await
469            .unwrap();
470
471        // It returned the previous 10 items
472        assert!(!page.has_previous_page);
473        let edge_ids: Vec<_> = page.edges.iter().map(|p| p.node.id).collect();
474        assert_eq!(&edge_ids, &ids[..10]);
475
476        // Lookup 10 items between two IDs
477        let page = repo
478            .upstream_oauth_provider()
479            .list(filter, Pagination::first(10).after(ids[5]).before(ids[8]))
480            .await
481            .unwrap();
482
483        // It returned the items in between
484        assert!(!page.has_next_page);
485        let edge_ids: Vec<_> = page.edges.iter().map(|p| p.node.id).collect();
486        assert_eq!(&edge_ids, &ids[6..8]);
487
488        // There should not be any disabled providers
489        assert!(
490            repo.upstream_oauth_provider()
491                .list(
492                    UpstreamOAuthProviderFilter::new().disabled_only(),
493                    Pagination::first(1)
494                )
495                .await
496                .unwrap()
497                .edges
498                .is_empty()
499        );
500    }
501
502    /// Test that the pagination works as expected in the upstream OAuth
503    /// session repository
504    #[sqlx::test(migrator = "crate::MIGRATOR")]
505    async fn test_session_repository_pagination(pool: PgPool) {
506        let scope = Scope::from_iter([OPENID]);
507
508        let mut rng = rand_chacha::ChaChaRng::seed_from_u64(42);
509        let clock = MockClock::default();
510        let mut repo = PgRepository::from_pool(&pool).await.unwrap();
511
512        // Create a provider
513        let provider = repo
514            .upstream_oauth_provider()
515            .add(
516                &mut rng,
517                &clock,
518                UpstreamOAuthProviderParams {
519                    issuer: Some("https://example.com/".to_owned()),
520                    human_name: None,
521                    brand_name: None,
522                    scope,
523                    token_endpoint_auth_method: UpstreamOAuthProviderTokenAuthMethod::None,
524                    id_token_signed_response_alg: JsonWebSignatureAlg::Rs256,
525                    fetch_userinfo: false,
526                    userinfo_signed_response_alg: None,
527                    token_endpoint_signing_alg: None,
528                    client_id: "client-id".to_owned(),
529                    encrypted_client_secret: None,
530                    claims_imports: UpstreamOAuthProviderClaimsImports::default(),
531                    token_endpoint_override: None,
532                    authorization_endpoint_override: None,
533                    userinfo_endpoint_override: None,
534                    jwks_uri_override: None,
535                    discovery_mode: mas_data_model::UpstreamOAuthProviderDiscoveryMode::Oidc,
536                    pkce_mode: mas_data_model::UpstreamOAuthProviderPkceMode::Auto,
537                    response_mode: None,
538                    additional_authorization_parameters: Vec::new(),
539                    forward_login_hint: false,
540                    ui_order: 0,
541                    on_backchannel_logout: UpstreamOAuthProviderOnBackchannelLogout::DoNothing,
542                    registration_token_required: false,
543                },
544            )
545            .await
546            .unwrap();
547
548        let filter = UpstreamOAuthSessionFilter::new().for_provider(&provider);
549
550        // Count the number of sessions before we start
551        assert_eq!(
552            repo.upstream_oauth_session().count(filter).await.unwrap(),
553            0
554        );
555
556        let mut links = Vec::with_capacity(3);
557        for subject in ["alice", "bob", "charlie"] {
558            let link = repo
559                .upstream_oauth_link()
560                .add(&mut rng, &clock, &provider, subject.to_owned(), None)
561                .await
562                .unwrap();
563            links.push(link);
564        }
565
566        let mut ids = Vec::with_capacity(20);
567        let sids = ["one", "two"].into_iter().cycle();
568        // Create 20 sessions
569        for (idx, (link, sid)) in links.iter().cycle().zip(sids).enumerate().take(20) {
570            let state = format!("state-{idx}");
571            let session = repo
572                .upstream_oauth_session()
573                .add(&mut rng, &clock, &provider, state, None, None)
574                .await
575                .unwrap();
576            let id_token_claims = serde_json::json!({
577                "sub": link.subject,
578                "sid": sid,
579                "aud": provider.client_id,
580                "iss": "https://example.com/",
581            });
582            let session = repo
583                .upstream_oauth_session()
584                .complete_with_link(
585                    &clock,
586                    session,
587                    link,
588                    None,
589                    Some(id_token_claims),
590                    None,
591                    None,
592                )
593                .await
594                .unwrap();
595            ids.push(session.id);
596            clock.advance(Duration::microseconds(10 * 1000 * 1000));
597        }
598
599        // Now we have 20 sessions
600        assert_eq!(
601            repo.upstream_oauth_session().count(filter).await.unwrap(),
602            20
603        );
604
605        // Lookup the first 10 items
606        let page = repo
607            .upstream_oauth_session()
608            .list(filter, Pagination::first(10))
609            .await
610            .unwrap();
611
612        // It returned the first 10 items
613        assert!(page.has_next_page);
614        let edge_ids: Vec<_> = page.edges.iter().map(|s| s.node.id).collect();
615        assert_eq!(&edge_ids, &ids[..10]);
616
617        // Lookup the next 10 items
618        let page = repo
619            .upstream_oauth_session()
620            .list(filter, Pagination::first(10).after(ids[9]))
621            .await
622            .unwrap();
623
624        // It returned the next 10 items
625        assert!(!page.has_next_page);
626        let edge_ids: Vec<_> = page.edges.iter().map(|s| s.node.id).collect();
627        assert_eq!(&edge_ids, &ids[10..]);
628
629        // Lookup the last 10 items
630        let page = repo
631            .upstream_oauth_session()
632            .list(filter, Pagination::last(10))
633            .await
634            .unwrap();
635
636        // It returned the last 10 items
637        assert!(page.has_previous_page);
638        let edge_ids: Vec<_> = page.edges.iter().map(|s| s.node.id).collect();
639        assert_eq!(&edge_ids, &ids[10..]);
640
641        // Lookup the previous 10 items
642        let page = repo
643            .upstream_oauth_session()
644            .list(filter, Pagination::last(10).before(ids[10]))
645            .await
646            .unwrap();
647
648        // It returned the previous 10 items
649        assert!(!page.has_previous_page);
650        let edge_ids: Vec<_> = page.edges.iter().map(|s| s.node.id).collect();
651        assert_eq!(&edge_ids, &ids[..10]);
652
653        // Lookup 5 items between two IDs
654        let page = repo
655            .upstream_oauth_session()
656            .list(filter, Pagination::first(10).after(ids[5]).before(ids[11]))
657            .await
658            .unwrap();
659
660        // It returned the items in between
661        assert!(!page.has_next_page);
662        let edge_ids: Vec<_> = page.edges.iter().map(|s| s.node.id).collect();
663        assert_eq!(&edge_ids, &ids[6..11]);
664
665        // Check the sub/sid filters
666        assert_eq!(
667            repo.upstream_oauth_session()
668                .count(filter.with_sub_claim("alice").with_sid_claim("one"))
669                .await
670                .unwrap(),
671            4
672        );
673        assert_eq!(
674            repo.upstream_oauth_session()
675                .count(filter.with_sub_claim("bob").with_sid_claim("two"))
676                .await
677                .unwrap(),
678            4
679        );
680
681        let page = repo
682            .upstream_oauth_session()
683            .list(
684                filter.with_sub_claim("alice").with_sid_claim("one"),
685                Pagination::first(10),
686            )
687            .await
688            .unwrap();
689        assert_eq!(page.edges.len(), 4);
690        for edge in page.edges {
691            assert_eq!(
692                edge.node
693                    .id_token_claims()
694                    .unwrap()
695                    .get("sub")
696                    .unwrap()
697                    .as_str(),
698                Some("alice")
699            );
700            assert_eq!(
701                edge.node
702                    .id_token_claims()
703                    .unwrap()
704                    .get("sid")
705                    .unwrap()
706                    .as_str(),
707                Some("one")
708            );
709        }
710    }
711}