1mod 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 let all_providers = repo.upstream_oauth_provider().all_enabled().await.unwrap();
50 assert!(all_providers.is_empty());
51
52 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 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 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 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 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 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 repo.upstream_oauth_link()
146 .lookup(link.id)
147 .await
148 .unwrap()
149 .expect("link to be found in database");
150
151 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 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 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 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 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 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 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 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 repo.upstream_oauth_provider()
288 .disable(&clock, provider.clone())
289 .await
290 .unwrap();
291
292 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 let session_filter = UpstreamOAuthSessionFilter::new().for_provider(&provider);
317
318 let session_count = repo
320 .upstream_oauth_session()
321 .count(session_filter)
322 .await
323 .unwrap();
324 assert_eq!(session_count, 1);
325
326 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 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 #[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 assert_eq!(
366 repo.upstream_oauth_provider().count(filter).await.unwrap(),
367 0
368 );
369
370 let mut ids = Vec::with_capacity(20);
371 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 assert_eq!(
414 repo.upstream_oauth_provider().count(filter).await.unwrap(),
415 20
416 );
417
418 let page = repo
420 .upstream_oauth_provider()
421 .list(filter, Pagination::first(10))
422 .await
423 .unwrap();
424
425 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 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 let page = repo
442 .upstream_oauth_provider()
443 .list(filter, Pagination::first(10).after(ids[9]))
444 .await
445 .unwrap();
446
447 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 let page = repo
454 .upstream_oauth_provider()
455 .list(filter, Pagination::last(10))
456 .await
457 .unwrap();
458
459 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 let page = repo
466 .upstream_oauth_provider()
467 .list(filter, Pagination::last(10).before(ids[10]))
468 .await
469 .unwrap();
470
471 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 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 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 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 #[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 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 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 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 assert_eq!(
601 repo.upstream_oauth_session().count(filter).await.unwrap(),
602 20
603 );
604
605 let page = repo
607 .upstream_oauth_session()
608 .list(filter, Pagination::first(10))
609 .await
610 .unwrap();
611
612 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 let page = repo
619 .upstream_oauth_session()
620 .list(filter, Pagination::first(10).after(ids[9]))
621 .await
622 .unwrap();
623
624 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 let page = repo
631 .upstream_oauth_session()
632 .list(filter, Pagination::last(10))
633 .await
634 .unwrap();
635
636 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 let page = repo
643 .upstream_oauth_session()
644 .list(filter, Pagination::last(10).before(ids[10]))
645 .await
646 .unwrap();
647
648 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 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 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 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}