1use std::collections::HashMap;
7
8use axum::extract::{FromRef, Query, State};
9use axum::routing::get;
10use axum::{Json, Router};
11use sea_orm::{ColumnTrait, EntityTrait, PaginatorTrait, QueryFilter};
12use serde::{Deserialize, Serialize};
13
14use crate::auth::{AuthPrincipal, MiryadAuthState};
15use crate::query::{PagedResult, Pagination};
16use crate::rest::error::RestError;
17use crate::users::{group, is_admin, membership, resolve_user, user};
18
19#[derive(Debug, Clone, PartialEq, Serialize)]
20pub struct UserSummary {
21 pub id: i32,
22 pub subject: String,
23 pub email: Option<String>,
24 pub groups: Vec<String>,
25}
26
27#[derive(Deserialize)]
28struct ListParams {
29 page: Option<u64>,
30 per_page: Option<u64>,
31}
32
33pub fn users_router<S>() -> Router<S>
38where
39 S: Clone + Send + Sync + 'static,
40 MiryadAuthState: FromRef<S>,
41{
42 Router::new().nest("/api/v1", Router::new().route("/users", get(list_users_handler)))
43}
44
45async fn list_users_handler(
46 State(auth): State<MiryadAuthState>,
47 principal: AuthPrincipal,
48 Query(params): Query<ListParams>,
49) -> Result<Json<PagedResult<UserSummary>>, RestError> {
50 let caller = resolve_user(&auth.db, &principal.subject, principal.email.as_deref()).await?;
51 if !is_admin(&auth.db, caller.id).await? {
52 return Err(RestError::Forbidden);
53 }
54
55 let pagination = Pagination::from_raw(params.page, params.per_page);
56 let paginator = user::Entity::find().paginate(&auth.db, pagination.per_page);
57 let totals = paginator.num_items_and_pages().await?;
58 let users = paginator.fetch_page(pagination.page - 1).await?;
59
60 let mut groups_by_user = groups_by_user(&auth.db, users.iter().map(|u| u.id)).await?;
61
62 let items = users
63 .into_iter()
64 .map(|u| UserSummary {
65 groups: groups_by_user.remove(&u.id).unwrap_or_default(),
66 id: u.id,
67 subject: u.subject,
68 email: u.email,
69 })
70 .collect();
71
72 Ok(Json(PagedResult {
73 items,
74 page: pagination.page,
75 per_page: pagination.per_page,
76 total_items: totals.number_of_items,
77 total_pages: totals.number_of_pages,
78 }))
79}
80
81async fn groups_by_user(
85 db: &sea_orm::DatabaseConnection,
86 user_ids: impl Iterator<Item = i32>,
87) -> Result<HashMap<i32, Vec<String>>, sea_orm::DbErr> {
88 let user_ids: Vec<i32> = user_ids.collect();
89 if user_ids.is_empty() {
90 return Ok(HashMap::new());
91 }
92
93 let memberships = membership::Entity::find()
94 .filter(membership::Column::UserId.is_in(user_ids))
95 .all(db)
96 .await?;
97 if memberships.is_empty() {
98 return Ok(HashMap::new());
99 }
100
101 let group_ids: Vec<i32> = memberships.iter().map(|m| m.group_id).collect();
102 let group_names: HashMap<i32, String> = group::Entity::find()
103 .filter(group::Column::Id.is_in(group_ids))
104 .all(db)
105 .await?
106 .into_iter()
107 .map(|g| (g.id, g.name))
108 .collect();
109
110 let mut result: HashMap<i32, Vec<String>> = HashMap::new();
111 for m in memberships {
112 if let Some(name) = group_names.get(&m.group_id) {
113 result.entry(m.user_id).or_default().push(name.clone());
114 }
115 }
116 Ok(result)
117}
118
119#[cfg(test)]
120mod tests {
121 use super::*;
122 use crate::auth::issue_token;
123 use crate::auth::oidc::MockOidcClient;
124 use crate::migration::Migrator;
125 use crate::users::sync_group_memberships;
126 use axum::body::Body;
127 use axum::http::{Request, StatusCode};
128 use sea_orm::{Database, DatabaseConnection};
129 use sea_orm_migration::MigratorTrait;
130 use tower::ServiceExt;
131
132 async fn test_db() -> DatabaseConnection {
133 let db = Database::connect("sqlite::memory:")
134 .await
135 .expect("in-memory sqlite connects");
136 Migrator::up(&db, None).await.expect("migrations apply cleanly");
137 db
138 }
139
140 fn test_state(db: DatabaseConnection) -> MiryadAuthState {
141 MiryadAuthState {
142 oidc_client: std::sync::Arc::new(MockOidcClient),
143 cookie_key: ::cookie::Key::from(&[0u8; 64]),
144 post_login_redirect: "/".to_string(),
145 post_logout_redirect: "/".to_string(),
146 db,
147 }
148 }
149
150 fn app(state: MiryadAuthState) -> Router {
151 Router::new()
152 .merge(users_router::<MiryadAuthState>())
153 .with_state(state)
154 }
155
156 async fn json_body(resp: axum::response::Response) -> serde_json::Value {
157 let bytes = axum::body::to_bytes(resp.into_body(), usize::MAX)
158 .await
159 .expect("readable body");
160 serde_json::from_slice(&bytes).expect("valid JSON body")
161 }
162
163 fn get_request(uri: &str, token: &str) -> Request<Body> {
164 Request::builder()
165 .method("GET")
166 .uri(uri)
167 .header("Authorization", format!("Bearer {token}"))
168 .body(Body::empty())
169 .expect("valid request")
170 }
171
172 #[tokio::test]
173 async fn admin_sees_paginated_users_with_their_groups() {
174 let db = test_db().await;
175 let admin = resolve_user(&db, "admin-sub", None)
176 .await
177 .expect("resolve succeeds");
178 sync_group_memberships(&db, admin.id, &["admin".to_string()])
179 .await
180 .expect("sync succeeds");
181 let alice = resolve_user(&db, "alice-sub", Some("alice@example.com"))
182 .await
183 .expect("resolve succeeds");
184 sync_group_memberships(&db, alice.id, &["editors".to_string(), "viewers".to_string()])
185 .await
186 .expect("sync succeeds");
187 resolve_user(&db, "bob-sub", None)
189 .await
190 .expect("resolve succeeds");
191
192 let token = issue_token(&db, "admin-sub", "test", None)
193 .await
194 .expect("issuing succeeds")
195 .token;
196 let app = app(test_state(db));
197
198 let resp = app
199 .oneshot(get_request("/api/v1/users", &token))
200 .await
201 .expect("router does not fail");
202 assert_eq!(resp.status(), StatusCode::OK);
203
204 let body = json_body(resp).await;
205 assert_eq!(body["total_items"], 3);
206 let items = body["items"].as_array().expect("items array");
207 assert_eq!(items.len(), 3);
208
209 let alice_entry = items
210 .iter()
211 .find(|u| u["subject"] == "alice-sub")
212 .expect("alice present");
213 assert_eq!(alice_entry["email"], "alice@example.com");
214 let mut groups: Vec<&str> = alice_entry["groups"]
215 .as_array()
216 .expect("groups array")
217 .iter()
218 .map(|g| g.as_str().expect("group is a string"))
219 .collect();
220 groups.sort_unstable();
221 assert_eq!(groups, vec!["editors", "viewers"]);
222
223 let bob_entry = items
224 .iter()
225 .find(|u| u["subject"] == "bob-sub")
226 .expect("bob present");
227 assert_eq!(bob_entry["groups"].as_array().expect("groups array").len(), 0);
228 }
229
230 #[tokio::test]
231 async fn non_admin_is_forbidden() {
232 let db = test_db().await;
233 resolve_user(&db, "alice-sub", None)
234 .await
235 .expect("resolve succeeds");
236 let token = issue_token(&db, "alice-sub", "test", None)
237 .await
238 .expect("issuing succeeds")
239 .token;
240 let app = app(test_state(db));
241
242 let resp = app
243 .oneshot(get_request("/api/v1/users", &token))
244 .await
245 .expect("router does not fail");
246 assert_eq!(resp.status(), StatusCode::FORBIDDEN);
247 }
248
249 #[tokio::test]
250 async fn pagination_params_are_respected() {
251 let db = test_db().await;
252 let admin = resolve_user(&db, "admin-sub", None)
253 .await
254 .expect("resolve succeeds");
255 sync_group_memberships(&db, admin.id, &["admin".to_string()])
256 .await
257 .expect("sync succeeds");
258 for n in 0..3 {
259 resolve_user(&db, &format!("user-{n}"), None)
260 .await
261 .expect("resolve succeeds");
262 }
263 let token = issue_token(&db, "admin-sub", "test", None)
266 .await
267 .expect("issuing succeeds")
268 .token;
269 let app = app(test_state(db));
270
271 let resp = app
272 .oneshot(get_request("/api/v1/users?page=2&per_page=3", &token))
273 .await
274 .expect("router does not fail");
275 assert_eq!(resp.status(), StatusCode::OK);
276
277 let body = json_body(resp).await;
278 assert_eq!(body["page"], 2);
279 assert_eq!(body["per_page"], 3);
280 assert_eq!(body["total_items"], 4);
281 assert_eq!(body["total_pages"], 2);
282 assert_eq!(body["items"].as_array().expect("items array").len(), 1);
283 }
284}