1use std::str::FromStr;
9
10use aide::{OperationIo, transform::TransformOperation};
11use axum::{Json, response::IntoResponse};
12use axum_extra::extract::{Query, QueryRejection};
13use axum_macros::FromRequestParts;
14use chrono::{DateTime, Utc};
15use hyper::StatusCode;
16use mas_axum_utils::record_error;
17use mas_storage::{Page, oauth2::OAuth2SessionFilter};
18use oauth2_types::scope::{Scope, ScopeToken};
19use schemars::JsonSchema;
20use serde::Deserialize;
21use ulid::Ulid;
22
23use crate::{
24 admin::{
25 call_context::CallContext,
26 model::{OAuth2Session, Resource},
27 params::{IncludeCount, Pagination},
28 response::{ErrorResponse, PaginatedResponse},
29 },
30 impl_from_error_for_route,
31};
32
33#[derive(Deserialize, JsonSchema, Clone, Copy)]
34#[serde(rename_all = "snake_case")]
35enum OAuth2SessionStatus {
36 Active,
37 Finished,
38}
39
40impl std::fmt::Display for OAuth2SessionStatus {
41 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
42 match self {
43 Self::Active => write!(f, "active"),
44 Self::Finished => write!(f, "finished"),
45 }
46 }
47}
48
49#[derive(Deserialize, JsonSchema, Clone, Copy)]
50#[serde(rename_all = "snake_case")]
51enum OAuth2ClientKind {
52 Dynamic,
53 Static,
54}
55
56impl std::fmt::Display for OAuth2ClientKind {
57 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
58 match self {
59 Self::Dynamic => write!(f, "dynamic"),
60 Self::Static => write!(f, "static"),
61 }
62 }
63}
64
65#[derive(FromRequestParts, Deserialize, JsonSchema, OperationIo)]
66#[serde(rename = "OAuth2SessionFilter")]
67#[aide(input_with = "Query<FilterParams>")]
68#[from_request(via(Query), rejection(RouteError))]
69pub struct FilterParams {
70 #[serde(rename = "filter[user]")]
72 #[schemars(with = "Option<crate::admin::schema::Ulid>")]
73 user: Option<Ulid>,
74
75 #[serde(default, rename = "filter[client]")]
80 #[schemars(with = "Vec<crate::admin::schema::Ulid>")]
81 client: Vec<Ulid>,
82
83 #[serde(rename = "filter[client-kind]")]
85 client_kind: Option<OAuth2ClientKind>,
86
87 #[serde(rename = "filter[user-session]")]
89 #[schemars(with = "Option<crate::admin::schema::Ulid>")]
90 user_session: Option<Ulid>,
91
92 #[serde(default, rename = "filter[scope]")]
94 scope: Vec<String>,
95
96 #[serde(rename = "filter[status]")]
104 status: Option<OAuth2SessionStatus>,
105
106 #[serde(rename = "filter[created-before]")]
108 created_before: Option<DateTime<Utc>>,
109
110 #[serde(rename = "filter[created-after]")]
112 created_after: Option<DateTime<Utc>>,
113
114 #[serde(rename = "filter[last-active-before]")]
116 last_active_before: Option<DateTime<Utc>>,
117
118 #[serde(rename = "filter[last-active-after]")]
120 last_active_after: Option<DateTime<Utc>>,
121}
122
123impl std::fmt::Display for FilterParams {
124 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
125 let mut sep = '?';
126
127 if let Some(user) = self.user {
128 write!(f, "{sep}filter[user]={user}")?;
129 sep = '&';
130 }
131
132 for client in &self.client {
133 write!(f, "{sep}filter[client]={client}")?;
134 sep = '&';
135 }
136
137 if let Some(client_kind) = self.client_kind {
138 write!(f, "{sep}filter[client-kind]={client_kind}")?;
139 sep = '&';
140 }
141
142 if let Some(user_session) = self.user_session {
143 write!(f, "{sep}filter[user-session]={user_session}")?;
144 sep = '&';
145 }
146
147 for scope in &self.scope {
148 write!(f, "{sep}filter[scope]={scope}")?;
149 sep = '&';
150 }
151
152 if let Some(status) = self.status {
153 write!(f, "{sep}filter[status]={status}")?;
154 sep = '&';
155 }
156
157 if let Some(created_before) = self.created_before {
158 write!(
159 f,
160 "{sep}filter[created-before]={}",
161 created_before.format("%Y-%m-%dT%H:%M:%SZ")
162 )?;
163 sep = '&';
164 }
165
166 if let Some(created_after) = self.created_after {
167 write!(
168 f,
169 "{sep}filter[created-after]={}",
170 created_after.format("%Y-%m-%dT%H:%M:%SZ")
171 )?;
172 sep = '&';
173 }
174
175 if let Some(last_active_before) = self.last_active_before {
176 write!(
177 f,
178 "{sep}filter[last-active-before]={}",
179 last_active_before.format("%Y-%m-%dT%H:%M:%SZ")
180 )?;
181 sep = '&';
182 }
183
184 if let Some(last_active_after) = self.last_active_after {
185 write!(
186 f,
187 "{sep}filter[last-active-after]={}",
188 last_active_after.format("%Y-%m-%dT%H:%M:%SZ")
189 )?;
190 sep = '&';
191 }
192
193 let _ = sep;
194 Ok(())
195 }
196}
197
198#[derive(Debug, thiserror::Error, OperationIo)]
199#[aide(output_with = "Json<ErrorResponse>")]
200pub enum RouteError {
201 #[error(transparent)]
202 Internal(Box<dyn std::error::Error + Send + Sync + 'static>),
203
204 #[error("User ID {0} not found")]
205 UserNotFound(Ulid),
206
207 #[error("Client ID {0} not found")]
208 ClientNotFound(Ulid),
209
210 #[error("User session ID {0} not found")]
211 UserSessionNotFound(Ulid),
212
213 #[error("Invalid filter parameters")]
214 InvalidFilter(#[from] QueryRejection),
215
216 #[error("Invalid scope {0:?} in filter parameters")]
217 InvalidScope(String),
218}
219
220impl_from_error_for_route!(mas_storage::RepositoryError);
221
222impl IntoResponse for RouteError {
223 fn into_response(self) -> axum::response::Response {
224 let error = ErrorResponse::from_error(&self);
225 let sentry_event_id = record_error!(self, RouteError::Internal(_));
226 let status = match self {
227 Self::Internal(_) => StatusCode::INTERNAL_SERVER_ERROR,
228 Self::UserNotFound(_) | Self::ClientNotFound(_) | Self::UserSessionNotFound(_) => {
229 StatusCode::NOT_FOUND
230 }
231 Self::InvalidScope(_) | Self::InvalidFilter(_) => StatusCode::BAD_REQUEST,
232 };
233 (status, sentry_event_id, Json(error)).into_response()
234 }
235}
236
237pub fn doc(operation: TransformOperation) -> TransformOperation {
238 operation
239 .id("listOAuth2Sessions")
240 .summary("List OAuth 2.0 sessions")
241 .description("Retrieve a list of OAuth 2.0 sessions.
242Note that by default, all sessions, including finished ones are returned, with the oldest first.
243Use the `filter[status]` parameter to filter the sessions by their status and `page[last]` parameter to retrieve the last N sessions.
244The `filter[client]` parameter may be repeated to filter on multiple clients at once.")
245 .tag("oauth2-session")
246 .response_with::<200, Json<PaginatedResponse<OAuth2Session>>, _>(|t| {
247 let sessions = OAuth2Session::samples();
248 let pagination = mas_storage::Pagination::first(sessions.len());
249 let page = Page {
250 edges: sessions
251 .into_iter()
252 .map(|node| mas_storage::pagination::Edge {
253 cursor: node.id(),
254 node,
255 })
256 .collect(),
257 has_next_page: true,
258 has_previous_page: false,
259 };
260
261 t.description("Paginated response of OAuth 2.0 sessions")
262 .example(PaginatedResponse::for_page(
263 page,
264 pagination,
265 Some(42),
266 OAuth2Session::PATH,
267 ))
268 })
269 .response_with::<404, RouteError, _>(|t| {
270 let response = ErrorResponse::from_error(&RouteError::UserNotFound(Ulid::nil()));
271 t.description("User was not found").example(response)
272 })
273 .response_with::<400, RouteError, _>(|t| {
274 let response = ErrorResponse::from_error(&RouteError::InvalidScope("not a valid scope".to_owned()));
275 t.description("Invalid scope").example(response)
276 })
277}
278
279#[tracing::instrument(name = "handler.admin.v1.oauth2_sessions.list", skip_all)]
280pub async fn handler(
281 CallContext { mut repo, .. }: CallContext,
282 Pagination(pagination, include_count): Pagination,
283 params: FilterParams,
284) -> Result<Json<PaginatedResponse<OAuth2Session>>, RouteError> {
285 let base = format!("{path}{params}", path = OAuth2Session::PATH);
286 let base = include_count.add_to_base(&base);
287 let filter = OAuth2SessionFilter::default();
288
289 let user = if let Some(user_id) = params.user {
291 let user = repo
292 .user()
293 .lookup(user_id)
294 .await?
295 .ok_or(RouteError::UserNotFound(user_id))?;
296
297 Some(user)
298 } else {
299 None
300 };
301
302 let filter = match &user {
303 Some(user) => filter.for_user(user),
304 None => filter,
305 };
306
307 let mut clients = Vec::with_capacity(params.client.len());
308 for client_id in params.client {
309 let client = repo
310 .oauth2_client()
311 .lookup(client_id)
312 .await?
313 .ok_or(RouteError::ClientNotFound(client_id))?;
314 clients.push(client);
315 }
316 let client_refs: Vec<&_> = clients.iter().collect();
317
318 let filter = if client_refs.is_empty() {
319 filter
320 } else {
321 filter.for_clients(&client_refs)
322 };
323
324 let filter = match params.client_kind {
325 Some(OAuth2ClientKind::Dynamic) => filter.only_dynamic_clients(),
326 Some(OAuth2ClientKind::Static) => filter.only_static_clients(),
327 None => filter,
328 };
329
330 let user_session = if let Some(user_session_id) = params.user_session {
331 let user_session = repo
332 .browser_session()
333 .lookup(user_session_id)
334 .await?
335 .ok_or(RouteError::UserSessionNotFound(user_session_id))?;
336
337 Some(user_session)
338 } else {
339 None
340 };
341
342 let filter = match &user_session {
343 Some(user_session) => filter.for_browser_session(user_session),
344 None => filter,
345 };
346
347 let scope: Scope = params
348 .scope
349 .into_iter()
350 .map(|s| ScopeToken::from_str(&s).map_err(|_| RouteError::InvalidScope(s)))
351 .collect::<Result<_, _>>()?;
352
353 let filter = if scope.is_empty() {
354 filter
355 } else {
356 filter.with_scope(&scope)
357 };
358
359 let filter = match params.status {
360 Some(OAuth2SessionStatus::Active) => filter.active_only(),
361 Some(OAuth2SessionStatus::Finished) => filter.finished_only(),
362 None => filter,
363 };
364
365 let filter = if let Some(created_before) = params.created_before {
366 filter.with_created_before(created_before)
367 } else {
368 filter
369 };
370
371 let filter = if let Some(created_after) = params.created_after {
372 filter.with_created_after(created_after)
373 } else {
374 filter
375 };
376
377 let filter = if let Some(last_active_before) = params.last_active_before {
378 filter.with_last_active_before(last_active_before)
379 } else {
380 filter
381 };
382
383 let filter = if let Some(last_active_after) = params.last_active_after {
384 filter.with_last_active_after(last_active_after)
385 } else {
386 filter
387 };
388
389 let response = match include_count {
390 IncludeCount::True => {
391 let page = repo
392 .oauth2_session()
393 .list(filter, pagination)
394 .await?
395 .map(OAuth2Session::from);
396 let count = repo.oauth2_session().count(filter).await?;
397 PaginatedResponse::for_page(page, pagination, Some(count), &base)
398 }
399 IncludeCount::False => {
400 let page = repo
401 .oauth2_session()
402 .list(filter, pagination)
403 .await?
404 .map(OAuth2Session::from);
405 PaginatedResponse::for_page(page, pagination, None, &base)
406 }
407 IncludeCount::Only => {
408 let count = repo.oauth2_session().count(filter).await?;
409 PaginatedResponse::for_count_only(count, &base)
410 }
411 };
412
413 Ok(Json(response))
414}
415
416#[cfg(test)]
417mod tests {
418 use chrono::Duration;
419 use hyper::{Request, StatusCode};
420 use mas_data_model::Clock;
421 use oauth2_types::{
422 requests::GrantType,
423 scope::{OPENID, Scope},
424 };
425 use sqlx::PgPool;
426
427 use crate::test_utils::{RequestBuilderExt, ResponseExt, TestState, setup};
428
429 #[sqlx::test(migrator = "mas_storage_pg::MIGRATOR")]
430 async fn test_oauth2_simple_session_list(pool: PgPool) {
431 setup();
432 let mut state = TestState::from_pool(pool).await.unwrap();
433 let token = state.token_with_scope("urn:mas:admin").await;
434
435 let request = Request::get("/api/admin/v1/oauth2-sessions")
437 .bearer(&token)
438 .empty();
439 let response = state.request(request).await;
440 response.assert_status(StatusCode::OK);
441 let body: serde_json::Value = response.json();
442 insta::assert_json_snapshot!(body, @r#"
443 {
444 "meta": {
445 "count": 1
446 },
447 "data": [
448 {
449 "type": "oauth2-session",
450 "id": "01FSHN9AG0MKGTBNZ16RDR3PVY",
451 "attributes": {
452 "created_at": "2022-01-16T14:40:00Z",
453 "finished_at": null,
454 "user_id": null,
455 "user_session_id": null,
456 "client_id": "01FSHN9AG0FAQ50MT1E9FFRPZR",
457 "scope": "urn:mas:admin",
458 "user_agent": null,
459 "last_active_at": null,
460 "last_active_ip": null,
461 "human_name": null
462 },
463 "links": {
464 "self": "/api/admin/v1/oauth2-sessions/01FSHN9AG0MKGTBNZ16RDR3PVY"
465 },
466 "meta": {
467 "page": {
468 "cursor": "01FSHN9AG0MKGTBNZ16RDR3PVY"
469 }
470 }
471 }
472 ],
473 "links": {
474 "self": "/api/admin/v1/oauth2-sessions?page[first]=10",
475 "first": "/api/admin/v1/oauth2-sessions?page[first]=10",
476 "last": "/api/admin/v1/oauth2-sessions?page[last]=10"
477 }
478 }
479 "#);
480
481 let request = Request::get("/api/admin/v1/oauth2-sessions?count=false")
483 .bearer(&token)
484 .empty();
485 let response = state.request(request).await;
486 response.assert_status(StatusCode::OK);
487 let body: serde_json::Value = response.json();
488 insta::assert_json_snapshot!(body, @r#"
489 {
490 "data": [
491 {
492 "type": "oauth2-session",
493 "id": "01FSHN9AG0MKGTBNZ16RDR3PVY",
494 "attributes": {
495 "created_at": "2022-01-16T14:40:00Z",
496 "finished_at": null,
497 "user_id": null,
498 "user_session_id": null,
499 "client_id": "01FSHN9AG0FAQ50MT1E9FFRPZR",
500 "scope": "urn:mas:admin",
501 "user_agent": null,
502 "last_active_at": null,
503 "last_active_ip": null,
504 "human_name": null
505 },
506 "links": {
507 "self": "/api/admin/v1/oauth2-sessions/01FSHN9AG0MKGTBNZ16RDR3PVY"
508 },
509 "meta": {
510 "page": {
511 "cursor": "01FSHN9AG0MKGTBNZ16RDR3PVY"
512 }
513 }
514 }
515 ],
516 "links": {
517 "self": "/api/admin/v1/oauth2-sessions?count=false&page[first]=10",
518 "first": "/api/admin/v1/oauth2-sessions?count=false&page[first]=10",
519 "last": "/api/admin/v1/oauth2-sessions?count=false&page[last]=10"
520 }
521 }
522 "#);
523
524 let request = Request::get("/api/admin/v1/oauth2-sessions?count=only")
526 .bearer(&token)
527 .empty();
528 let response = state.request(request).await;
529 response.assert_status(StatusCode::OK);
530 let body: serde_json::Value = response.json();
531 insta::assert_json_snapshot!(body, @r#"
532 {
533 "meta": {
534 "count": 1
535 },
536 "links": {
537 "self": "/api/admin/v1/oauth2-sessions?count=only"
538 }
539 }
540 "#);
541 }
542
543 #[sqlx::test(migrator = "mas_storage_pg::MIGRATOR")]
544 async fn test_oauth2_session_list_by_last_active(pool: PgPool) {
545 setup();
546 let mut state = TestState::from_pool(pool).await.unwrap();
547 let token = state.token_with_scope("urn:mas:admin").await;
548 let mut rng = state.rng();
549
550 let mut repo = state.repository().await.unwrap();
551 let client = repo
552 .oauth2_client()
553 .add(
554 &mut rng,
555 &state.clock,
556 vec!["https://example.com/redirect".parse().unwrap()],
557 None,
558 None,
559 None,
560 vec![GrantType::ClientCredentials],
561 Some("Test client".to_owned()),
562 Some("https://example.com/logo.png".parse().unwrap()),
563 Some("https://example.com/".parse().unwrap()),
564 Some("https://example.com/policy".parse().unwrap()),
565 Some("https://example.com/tos".parse().unwrap()),
566 Some("https://example.com/jwks.json".parse().unwrap()),
567 None,
568 None,
569 None,
570 None,
571 None,
572 Some("https://example.com/login".parse().unwrap()),
573 )
574 .await
575 .unwrap();
576
577 let session = repo
578 .oauth2_session()
579 .add_from_client_credentials(
580 &mut rng,
581 &state.clock,
582 &client,
583 Scope::from_iter([OPENID]),
584 )
585 .await
586 .unwrap();
587
588 state.clock.advance(Duration::minutes(5));
589 let activity_at = state.clock.now();
590 repo.oauth2_session()
591 .record_batch_activity(vec![(session.id, activity_at, None)])
592 .await
593 .unwrap();
594 repo.save().await.unwrap();
595
596 let threshold = activity_at - Duration::minutes(1);
599 let request = Request::get(format!(
600 "/api/admin/v1/oauth2-sessions?filter[last-active-after]={}",
601 threshold.format("%Y-%m-%dT%H:%M:%SZ")
602 ))
603 .bearer(&token)
604 .empty();
605 let response = state.request(request).await;
606 response.assert_status(StatusCode::OK);
607 let body: serde_json::Value = response.json();
608 let ids: Vec<&str> = body["data"]
609 .as_array()
610 .unwrap()
611 .iter()
612 .map(|v| v["id"].as_str().unwrap())
613 .collect();
614 assert!(ids.contains(&session.id.to_string().as_str()));
615 }
616
617 #[sqlx::test(migrator = "mas_storage_pg::MIGRATOR")]
618 async fn test_oauth2_session_list_by_created_at(pool: PgPool) {
619 setup();
620 let mut state = TestState::from_pool(pool).await.unwrap();
621 let token = state.token_with_scope("urn:mas:admin").await;
622 let mut rng = state.rng();
623
624 let mut repo = state.repository().await.unwrap();
625 let client = repo
626 .oauth2_client()
627 .add(
628 &mut rng,
629 &state.clock,
630 vec!["https://example.com/redirect".parse().unwrap()],
631 None,
632 None,
633 None,
634 vec![GrantType::ClientCredentials],
635 Some("Test client".to_owned()),
636 Some("https://example.com/logo.png".parse().unwrap()),
637 Some("https://example.com/".parse().unwrap()),
638 Some("https://example.com/policy".parse().unwrap()),
639 Some("https://example.com/tos".parse().unwrap()),
640 Some("https://example.com/jwks.json".parse().unwrap()),
641 None,
642 None,
643 None,
644 None,
645 None,
646 Some("https://example.com/login".parse().unwrap()),
647 )
648 .await
649 .unwrap();
650
651 repo.oauth2_session()
654 .add_from_client_credentials(
655 &mut rng,
656 &state.clock,
657 &client,
658 Scope::from_iter([OPENID]),
659 )
660 .await
661 .unwrap();
662 state.clock.advance(Duration::minutes(1));
663 let cutoff = state.clock.now();
664 state.clock.advance(Duration::minutes(1));
665 let new_session = repo
666 .oauth2_session()
667 .add_from_client_credentials(
668 &mut rng,
669 &state.clock,
670 &client,
671 Scope::from_iter([OPENID]),
672 )
673 .await
674 .unwrap();
675 repo.save().await.unwrap();
676
677 let request = Request::get(format!(
679 "/api/admin/v1/oauth2-sessions?filter[created-after]={}",
680 cutoff.format("%Y-%m-%dT%H:%M:%SZ")
681 ))
682 .bearer(&token)
683 .empty();
684 let response = state.request(request).await;
685 response.assert_status(StatusCode::OK);
686 let body: serde_json::Value = response.json();
687 assert_eq!(body["meta"]["count"], 1);
688 let ids: Vec<&str> = body["data"]
689 .as_array()
690 .unwrap()
691 .iter()
692 .map(|v| v["id"].as_str().unwrap())
693 .collect();
694 assert_eq!(ids, vec![new_session.id.to_string()]);
695 }
696
697 #[sqlx::test(migrator = "mas_storage_pg::MIGRATOR")]
702 async fn test_oauth2_session_list_multiple_clients(pool: PgPool) {
703 setup();
704 let mut state = TestState::from_pool(pool).await.unwrap();
705 let token = state.token_with_scope("urn:mas:admin").await;
706 let mut rng = state.rng();
707
708 let mut repo = state.repository().await.unwrap();
711 let client_a = repo
712 .oauth2_client()
713 .add(
714 &mut rng,
715 &state.clock,
716 vec!["https://a.example.com/redirect".parse().unwrap()],
717 None,
718 None,
719 None,
720 vec![GrantType::ClientCredentials],
721 Some("client a".to_owned()),
722 None,
723 None,
724 None,
725 None,
726 None,
727 None,
728 None,
729 None,
730 None,
731 None,
732 None,
733 )
734 .await
735 .unwrap();
736 let client_b = repo
737 .oauth2_client()
738 .add(
739 &mut rng,
740 &state.clock,
741 vec!["https://b.example.com/redirect".parse().unwrap()],
742 None,
743 None,
744 None,
745 vec![GrantType::ClientCredentials],
746 Some("client b".to_owned()),
747 None,
748 None,
749 None,
750 None,
751 None,
752 None,
753 None,
754 None,
755 None,
756 None,
757 None,
758 )
759 .await
760 .unwrap();
761
762 let scope: Scope = "urn:mas:admin".parse().unwrap();
763 let session_a = repo
764 .oauth2_session()
765 .add_from_client_credentials(&mut rng, &state.clock, &client_a, scope.clone())
766 .await
767 .unwrap();
768 let session_b = repo
769 .oauth2_session()
770 .add_from_client_credentials(&mut rng, &state.clock, &client_b, scope.clone())
771 .await
772 .unwrap();
773 repo.save().await.unwrap();
774
775 let url = format!(
778 "/api/admin/v1/oauth2-sessions?filter[client]={}&filter[client]={}",
779 client_a.id, client_b.id,
780 );
781 let request = Request::get(&url).bearer(&token).empty();
782 let response = state.request(request).await;
783 response.assert_status(StatusCode::OK);
784 let body: serde_json::Value = response.json();
785
786 assert_eq!(body["meta"]["count"], 2);
787 let ids: Vec<&str> = body["data"]
788 .as_array()
789 .unwrap()
790 .iter()
791 .map(|v| v["id"].as_str().unwrap())
792 .collect();
793 let session_a_id = session_a.id.to_string();
794 let session_b_id = session_b.id.to_string();
795 assert!(ids.contains(&session_a_id.as_str()));
796 assert!(ids.contains(&session_b_id.as_str()));
797 assert_eq!(ids.len(), 2);
798
799 let self_link = body["links"]["self"].as_str().unwrap();
801 assert!(self_link.contains(&format!("filter[client]={}", client_a.id)));
802 assert!(self_link.contains(&format!("filter[client]={}", client_b.id)));
803 }
804}