1use std::collections::BTreeMap;
2
3use axum::extract::State;
4use futures::StreamExt;
5use ruma::{
6 OwnedUserId, UInt,
7 api::Direction::{self, Backward, Forward},
8};
9use synapse_admin_api::statistics::user_media_statistics::v1::{
10 Request, Response, UserMediaSortOrder, UserMediaStat,
11};
12use tuwunel_core::{
13 Err, Result,
14 utils::{IterStream, ReadyExt, math::ruma_from_usize, stream::BroadbandExt},
15};
16
17use super::{SortKey, usize_from};
18use crate::{Ruma, client::admin::require_admin};
19
20pub(crate) async fn admin_user_media_statistics_route(
26 State(services): State<crate::State>,
27 body: Ruma<Request>,
28) -> Result<Response> {
29 require_admin(&services, body.sender_user()).await?;
30
31 let order_by = body
32 .order_by
33 .as_ref()
34 .unwrap_or(&UserMediaSortOrder::UserId);
35
36 if !matches!(
37 order_by,
38 UserMediaSortOrder::MediaLength
39 | UserMediaSortOrder::MediaCount
40 | UserMediaSortOrder::UserId
41 | UserMediaSortOrder::Displayname
42 ) {
43 return Err!(Request(InvalidParam(
44 "Query parameter order_by must be one of media_length, media_count, user_id, \
45 displayname."
46 )));
47 }
48
49 let from_ts = body.from_ts.map_or(0, u64::from);
50 let until_ts = body.until_ts.map(u64::from);
51 if until_ts.is_some_and(|until_ts| until_ts <= from_ts) {
52 return Err!(Request(InvalidParam(
53 "Query parameter until_ts must be greater than from_ts."
54 )));
55 }
56
57 let search_term = body.search_term.as_deref();
58 if search_term.is_some_and(str::is_empty) {
59 return Err!(Request(InvalidParam(
60 "Query parameter search_term cannot be an empty string."
61 )));
62 }
63
64 let usage = services
65 .media
66 .upload_stats()
67 .ready_filter(|stat| in_window(stat.created_ts, from_ts, until_ts))
68 .ready_fold(BTreeMap::new(), |mut usage: BTreeMap<OwnedUserId, (u64, u64)>, stat| {
69 let (count, length) = usage.entry(stat.user_id).or_default();
70 *count = count.saturating_add(1);
71 *length = length.saturating_add(stat.media_length);
72 usage
73 })
74 .await;
75
76 let rows: Vec<UserMediaStat> = usage
77 .into_iter()
78 .stream()
79 .broad_then(async |(user_id, (count, length))| {
80 let displayname = services.profile.displayname(&user_id).await.ok();
81
82 UserMediaStat {
83 displayname,
84 ..UserMediaStat::new(user_id, uint_from_u64(count), uint_from_u64(length))
85 }
86 })
87 .ready_filter(|row| search_term.is_none_or(|term| row_matches(row, term)))
88 .collect()
89 .await;
90
91 let from = body.from.map_or(0, usize_from);
92 let limit = body.limit.map_or(100, usize_from);
93 let dir = body.dir.unwrap_or(Forward);
94
95 Ok(into_response(rows, order_by, dir, from, limit))
96}
97
98fn into_response(
99 rows: Vec<UserMediaStat>,
100 order_by: &UserMediaSortOrder,
101 dir: Direction,
102 from: usize,
103 limit: usize,
104) -> Response {
105 let total = rows.len();
106 let users = paginate(rows, order_by, dir, from, limit);
107 let end = from.saturating_add(users.len());
108 let next_token = (end < total).then(|| ruma_from_usize(end));
109
110 Response {
111 users,
112 next_token,
113 total: ruma_from_usize(total),
114 }
115}
116
117fn uint_from_u64(value: u64) -> UInt { UInt::try_from(value).unwrap_or(UInt::MAX) }
118
119fn in_window(created_ts: u64, from_ts: u64, until_ts: Option<u64>) -> bool {
120 created_ts >= from_ts && until_ts.is_none_or(|until_ts| created_ts <= until_ts)
121}
122
123fn row_matches(row: &UserMediaStat, term: &str) -> bool {
124 row.user_id.localpart().contains(term)
125 || row
126 .displayname
127 .as_deref()
128 .is_some_and(|name| name.contains(term))
129}
130
131fn paginate(
132 mut rows: Vec<UserMediaStat>,
133 order_by: &UserMediaSortOrder,
134 dir: Direction,
135 from: usize,
136 limit: usize,
137) -> Vec<UserMediaStat> {
138 rows.sort_unstable_by(|a, b| {
139 sort_key(order_by, a)
140 .cmp(&sort_key(order_by, b))
141 .then_with(|| a.user_id.cmp(&b.user_id))
142 });
143
144 if matches!(dir, Backward) {
145 rows.reverse();
146 }
147
148 rows.into_iter().skip(from).take(limit).collect()
149}
150
151fn sort_key<'a>(order_by: &UserMediaSortOrder, row: &'a UserMediaStat) -> SortKey<'a> {
152 match order_by {
153 | UserMediaSortOrder::MediaLength => SortKey::Num(Some(row.media_length.into())),
154 | UserMediaSortOrder::MediaCount => SortKey::Num(Some(row.media_count.into())),
155 | UserMediaSortOrder::Displayname => SortKey::OptStr(row.displayname.as_deref()),
156 | _ => SortKey::Str(row.user_id.as_str()),
157 }
158}
159
160#[cfg(test)]
161mod tests {
162 use ruma::{
163 UInt,
164 api::Direction::{Backward, Forward},
165 };
166
167 use super::{
168 UserMediaSortOrder, UserMediaStat, in_window, into_response, paginate, row_matches,
169 uint_from_u64,
170 };
171
172 fn row(local: &str, displayname: Option<&str>, count: u32, length: u32) -> UserMediaStat {
173 let user_id = format!("@{local}:example.org")
174 .try_into()
175 .unwrap();
176
177 UserMediaStat {
178 displayname: displayname.map(ToOwned::to_owned),
179 ..UserMediaStat::new(user_id, count.into(), length.into())
180 }
181 }
182
183 fn ids(rows: &[UserMediaStat]) -> Vec<&str> {
184 rows.iter()
185 .map(|row| row.user_id.localpart())
186 .collect()
187 }
188 #[test]
189 fn user_id_forward_orders_lexically() {
190 let rows = vec![row("b", None, 1, 1), row("a", None, 1, 1), row("c", None, 1, 1)];
191
192 let page = paginate(rows, &UserMediaSortOrder::UserId, Forward, 0, 100);
193
194 assert_eq!(ids(&page), ["a", "b", "c"]);
195 }
196
197 #[test]
198 fn backward_reverses_order() {
199 let rows = vec![row("b", None, 1, 1), row("a", None, 1, 1), row("c", None, 1, 1)];
200
201 let page = paginate(rows, &UserMediaSortOrder::UserId, Backward, 0, 100);
202
203 assert_eq!(ids(&page), ["c", "b", "a"]);
204 }
205
206 #[test]
207 fn media_length_orders_numerically() {
208 let rows = vec![row("a", None, 1, 5), row("b", None, 1, 30), row("c", None, 1, 20)];
209
210 let forward = paginate(rows.clone(), &UserMediaSortOrder::MediaLength, Forward, 0, 100);
211
212 assert_eq!(ids(&forward), ["a", "c", "b"]);
213
214 let backward = paginate(rows, &UserMediaSortOrder::MediaLength, Backward, 0, 100);
215
216 assert_eq!(ids(&backward), ["b", "c", "a"]);
217 }
218
219 #[test]
220 fn equal_counts_tiebreak_by_user_id() {
221 let rows = vec![row("c", None, 5, 1), row("a", None, 5, 1), row("b", None, 5, 1)];
222
223 let page = paginate(rows, &UserMediaSortOrder::MediaCount, Forward, 0, 100);
224
225 assert_eq!(ids(&page), ["a", "b", "c"]);
226 }
227
228 #[test]
229 fn displayname_none_sorts_first_forward() {
230 let rows = vec![row("a", Some("zed"), 1, 1), row("b", None, 1, 1)];
231
232 let page = paginate(rows, &UserMediaSortOrder::Displayname, Forward, 0, 100);
233
234 assert_eq!(ids(&page), ["b", "a"]);
235 }
236
237 #[test]
238 fn from_and_limit_window_the_page() {
239 let rows = vec![row("a", None, 1, 1), row("b", None, 1, 1), row("c", None, 1, 1)];
240
241 let page = paginate(rows, &UserMediaSortOrder::UserId, Forward, 1, 1);
242
243 assert_eq!(ids(&page), ["b"]);
244 }
245
246 #[test]
247 fn oversized_aggregate_saturates_on_wire() {
248 assert_eq!(uint_from_u64(42), UInt::from(42_u32));
249 assert_eq!(uint_from_u64(u64::MAX), UInt::MAX);
250 }
251
252 #[test]
253 fn response_reports_total_and_next_token() {
254 let rows = vec![row("a", None, 1, 1), row("b", None, 1, 1), row("c", None, 1, 1)];
255
256 let response = into_response(rows, &UserMediaSortOrder::UserId, Forward, 0, 1);
257
258 assert_eq!(ids(&response.users), ["a"]);
259 assert_eq!(response.total, UInt::from(3_u32));
260 assert_eq!(response.next_token, Some(UInt::from(1_u32)));
261 }
262
263 #[test]
264 fn final_page_omits_next_token() {
265 let rows = vec![row("a", None, 1, 1), row("b", None, 1, 1), row("c", None, 1, 1)];
266
267 let response = into_response(rows, &UserMediaSortOrder::UserId, Forward, 1, 2);
268
269 assert_eq!(ids(&response.users), ["b", "c"]);
270 assert_eq!(response.total, UInt::from(3_u32));
271 assert_eq!(response.next_token, None);
272 }
273
274 #[test]
275 fn zero_limit_repeats_from_as_next_token() {
276 let rows = vec![row("a", None, 1, 1), row("b", None, 1, 1), row("c", None, 1, 1)];
277
278 let response = into_response(rows, &UserMediaSortOrder::UserId, Forward, 1, 0);
279
280 assert!(response.users.is_empty());
281 assert_eq!(response.total, UInt::from(3_u32));
282 assert_eq!(response.next_token, Some(UInt::from(1_u32)));
283 }
284
285 #[test]
286 fn window_bounds_are_inclusive() {
287 assert!(in_window(5, 5, Some(5)));
288 assert!(!in_window(4, 5, None));
289 assert!(!in_window(6, 0, Some(5)));
290 assert!(in_window(0, 0, None));
291 }
292
293 #[test]
294 fn search_matches_localpart_or_displayname() {
295 assert!(row_matches(&row("alice", None, 1, 1), "lic"));
296 assert!(row_matches(&row("bob", Some("Wonderland"), 1, 1), "onder"));
297 assert!(!row_matches(&row("alice", Some("alice"), 1, 1), "example"));
298 assert!(!row_matches(&row("alice", Some("alice"), 1, 1), "Alice"));
299 }
300}