1use rusqlite::{Connection, params};
4
5use crate::auth::{self, Role};
6
7#[derive(Debug, Clone)]
12pub struct UserRow {
13 pub id: i64,
14 pub username: String,
15 pub password_hash: String,
16 pub role: Role,
17 pub created_at: Option<String>,
18}
19
20#[derive(Debug, Clone)]
21pub struct RefreshTokenRow {
22 pub id: String,
23 pub user_id: i64,
24 pub expires_at: i64,
25 pub revoked: bool,
26 pub created_at: Option<String>,
27}
28
29pub const LOCAL_USER: i64 = 0;
41
42pub fn first_admin(conn: &Connection) -> Result<Option<i64>, rusqlite::Error> {
44 conn.prepare_cached("SELECT MIN(id) FROM users WHERE role = 'admin'")?
45 .query_row([], |r| r.get(0))
46}
47
48pub fn resolve_user(conn: &Connection, user: i64) -> Result<i64, rusqlite::Error> {
51 if user != LOCAL_USER {
52 return Ok(user);
53 }
54 Ok(first_admin(conn)?.unwrap_or(LOCAL_USER))
55}
56
57pub fn is_local_user(conn: &Connection, user: i64) -> Result<bool, rusqlite::Error> {
60 Ok(resolve_user(conn, user)? == resolve_user(conn, LOCAL_USER)?)
61}
62
63pub fn adopt_local_rows(conn: &Connection) -> Result<(), rusqlite::Error> {
67 let Some(admin) = first_admin(conn)? else {
68 return Ok(());
69 };
70 for table in [
71 "favourites",
72 "favourite_albums",
73 "favourite_artists",
74 "play_history",
75 "playlists",
76 "shares",
77 ] {
78 let pending: bool = conn.query_row(
81 &format!("SELECT EXISTS(SELECT 1 FROM {table} WHERE user_id = ?1)"),
82 params![LOCAL_USER],
83 |r| r.get(0),
84 )?;
85 if !pending {
86 continue;
87 }
88 conn.execute(
89 &format!("UPDATE OR IGNORE {table} SET user_id = ?1 WHERE user_id = ?2"),
90 params![admin, LOCAL_USER],
91 )?;
92 conn.execute(
93 &format!("DELETE FROM {table} WHERE user_id = ?1"),
94 params![LOCAL_USER],
95 )?;
96 }
97 Ok(())
98}
99
100pub fn set_sealed_password(
106 conn: &Connection,
107 username: &str,
108 sealed: &[u8],
109) -> Result<(), rusqlite::Error> {
110 conn.execute(
111 "UPDATE users SET sealed_password = ?2 WHERE username = ?1",
112 params![username, sealed],
113 )?;
114 Ok(())
115}
116
117pub fn sealed_password(
118 conn: &Connection,
119 username: &str,
120) -> Result<Option<Vec<u8>>, rusqlite::Error> {
121 use rusqlite::OptionalExtension;
122 Ok(conn
123 .query_row(
124 "SELECT sealed_password FROM users WHERE username = ?1",
125 params![username],
126 |r| r.get::<_, Option<Vec<u8>>>(0),
127 )
128 .optional()?
129 .flatten())
130}
131
132pub fn remember_password(
135 conn: &Connection,
136 username: &str,
137 password: &str,
138) -> Result<(), Box<dyn std::error::Error>> {
139 let key = auth::subsonic_key()?;
140 set_sealed_password(
141 conn,
142 username,
143 &auth::seal_password(&key, username, password)?,
144 )?;
145 Ok(())
146}
147
148pub fn create_user(
150 conn: &Connection,
151 username: &str,
152 password: &str,
153 role: Role,
154) -> Result<i64, rusqlite::Error> {
155 let hash = auth::hash_password(password)
156 .map_err(|e| rusqlite::Error::ToSqlConversionFailure(e.into()))?;
157 conn.execute(
158 "INSERT INTO users (username, password_hash, role) VALUES (?1, ?2, ?3)",
159 params![username, hash, role.as_str()],
160 )?;
161 let id = conn.last_insert_rowid();
162 adopt_local_rows(conn)?;
163 Ok(id)
164}
165
166pub fn get_user_by_username(
168 conn: &Connection,
169 username: &str,
170) -> Result<Option<UserRow>, rusqlite::Error> {
171 let mut stmt = conn.prepare_cached(
172 "SELECT id, username, password_hash, role, created_at FROM users WHERE username = ?1",
173 )?;
174 let mut rows = stmt.query_map(params![username], |row| {
175 let role_str: String = row.get(3)?;
176 Ok(UserRow {
177 id: row.get(0)?,
178 username: row.get(1)?,
179 password_hash: row.get(2)?,
180 role: role_str.parse().unwrap_or(Role::Readonly),
181 created_at: row.get(4)?,
182 })
183 })?;
184 match rows.next() {
185 Some(Ok(user)) => Ok(Some(user)),
186 Some(Err(e)) => Err(e),
187 None => Ok(None),
188 }
189}
190
191pub fn get_user_by_id(conn: &Connection, user_id: i64) -> Result<Option<UserRow>, rusqlite::Error> {
193 let mut stmt = conn
194 .prepare("SELECT id, username, password_hash, role, created_at FROM users WHERE id = ?1")?;
195 let mut rows = stmt.query_map(params![user_id], |row| {
196 let role_str: String = row.get(3)?;
197 Ok(UserRow {
198 id: row.get(0)?,
199 username: row.get(1)?,
200 password_hash: row.get(2)?,
201 role: role_str.parse().unwrap_or(Role::Readonly),
202 created_at: row.get(4)?,
203 })
204 })?;
205 match rows.next() {
206 Some(Ok(user)) => Ok(Some(user)),
207 Some(Err(e)) => Err(e),
208 None => Ok(None),
209 }
210}
211
212pub fn list_users(conn: &Connection) -> Result<Vec<UserRow>, rusqlite::Error> {
214 let mut stmt = conn
215 .prepare("SELECT id, username, password_hash, role, created_at FROM users ORDER BY id")?;
216 let rows = stmt.query_map([], |row| {
217 let role_str: String = row.get(3)?;
218 Ok(UserRow {
219 id: row.get(0)?,
220 username: row.get(1)?,
221 password_hash: row.get(2)?,
222 role: role_str.parse().unwrap_or(Role::Readonly),
223 created_at: row.get(4)?,
224 })
225 })?;
226 rows.collect()
227}
228
229pub fn delete_user(conn: &Connection, user_id: i64) -> Result<bool, rusqlite::Error> {
231 let count = conn.execute("DELETE FROM users WHERE id = ?1", params![user_id])?;
232 Ok(count > 0)
233}
234
235pub fn update_password(
239 conn: &Connection,
240 username: &str,
241 new_password: &str,
242) -> Result<bool, Box<dyn std::error::Error>> {
243 let hash = crate::auth::hash_password(new_password)?;
244 let updated = conn.execute(
245 "UPDATE users SET password_hash = ?1 WHERE username = ?2",
246 params![hash, username],
247 )?;
248 if updated > 0 {
249 if let Some(user) = get_user_by_username(conn, username)? {
251 revoke_all_user_tokens(conn, user.id)?;
252 super::api_keys::revoke_user_api_keys(conn, user.id)?;
253 }
254 }
255 Ok(updated > 0)
256}
257
258pub fn update_role(
260 conn: &Connection,
261 username: &str,
262 role: crate::auth::Role,
263) -> Result<bool, rusqlite::Error> {
264 let updated = conn.execute(
265 "UPDATE users SET role = ?1 WHERE username = ?2",
266 params![role.as_str(), username],
267 )?;
268 adopt_local_rows(conn)?;
269 Ok(updated > 0)
270}
271
272pub fn has_users(conn: &Connection) -> Result<bool, rusqlite::Error> {
274 let count: i64 = conn.query_row("SELECT COUNT(*) FROM users", [], |row| row.get(0))?;
275 Ok(count > 0)
276}
277
278pub fn admin_count(conn: &Connection) -> Result<i64, rusqlite::Error> {
280 conn.query_row(
281 "SELECT COUNT(*) FROM users WHERE role = 'admin'",
282 [],
283 |row| row.get(0),
284 )
285}
286
287pub fn store_refresh_token(
294 conn: &Connection,
295 token_id: &str,
296 user_id: i64,
297 expires_at: i64,
298) -> Result<(), rusqlite::Error> {
299 conn.execute(
300 "INSERT INTO refresh_tokens (id, user_id, expires_at) VALUES (?1, ?2, ?3)",
301 params![auth::sha256_hex(token_id), user_id, expires_at],
302 )?;
303 Ok(())
304}
305
306pub fn get_valid_refresh_token(
308 conn: &Connection,
309 token_id: &str,
310) -> Result<Option<RefreshTokenRow>, rusqlite::Error> {
311 let now = auth::now_unix() as i64;
312 let mut stmt = conn.prepare_cached(
313 "SELECT id, user_id, expires_at, revoked, created_at
314 FROM refresh_tokens
315 WHERE id = ?1 AND revoked = 0 AND expires_at > ?2",
316 )?;
317 let mut rows = stmt.query_map(params![auth::sha256_hex(token_id), now], |row| {
318 Ok(RefreshTokenRow {
319 id: row.get(0)?,
320 user_id: row.get(1)?,
321 expires_at: row.get(2)?,
322 revoked: row.get::<_, i32>(3)? != 0,
323 created_at: row.get(4)?,
324 })
325 })?;
326 match rows.next() {
327 Some(Ok(token)) => Ok(Some(token)),
328 Some(Err(e)) => Err(e),
329 None => Ok(None),
330 }
331}
332
333pub fn consume_refresh_token(
337 conn: &Connection,
338 token_id: &str,
339) -> Result<Option<RefreshTokenRow>, rusqlite::Error> {
340 let now = auth::now_unix() as i64;
341 let mut stmt = conn.prepare(
342 "UPDATE refresh_tokens SET revoked = 1
343 WHERE id = ?1 AND revoked = 0 AND expires_at > ?2
344 RETURNING id, user_id, expires_at, revoked, created_at",
345 )?;
346 let mut rows = stmt.query_map(params![auth::sha256_hex(token_id), now], |row| {
347 Ok(RefreshTokenRow {
348 id: row.get(0)?,
349 user_id: row.get(1)?,
350 expires_at: row.get(2)?,
351 revoked: row.get::<_, i32>(3)? != 0,
352 created_at: row.get(4)?,
353 })
354 })?;
355 match rows.next() {
356 Some(Ok(token)) => Ok(Some(token)),
357 Some(Err(e)) => Err(e),
358 None => Ok(None),
359 }
360}
361
362pub fn revoke_refresh_token(conn: &Connection, token_id: &str) -> Result<bool, rusqlite::Error> {
364 let count = conn.execute(
365 "UPDATE refresh_tokens SET revoked = 1 WHERE id = ?1",
366 params![auth::sha256_hex(token_id)],
367 )?;
368 Ok(count > 0)
369}
370
371pub fn revoke_all_user_tokens(conn: &Connection, user_id: i64) -> Result<usize, rusqlite::Error> {
373 let count = conn.execute(
374 "UPDATE refresh_tokens SET revoked = 1 WHERE user_id = ?1 AND revoked = 0",
375 params![user_id],
376 )?;
377 Ok(count)
378}
379
380pub fn cleanup_expired_tokens(conn: &Connection) -> Result<usize, rusqlite::Error> {
382 let now = auth::now_unix() as i64;
383 let count = conn.execute(
384 "DELETE FROM refresh_tokens WHERE revoked = 1 OR expires_at <= ?1",
385 params![now],
386 )?;
387 Ok(count)
388}
389
390#[cfg(test)]
395mod tests {
396 use super::*;
397 use crate::db::connection::Database;
398 use tempfile::TempDir;
399
400 fn test_db() -> (Database, TempDir) {
401 let tmp = TempDir::new().unwrap();
402 let db_path = tmp.path().join("test.db");
403 let db = Database::open(&db_path).unwrap();
404 (db, tmp)
405 }
406
407 #[test]
408 fn create_and_get_user() {
409 let (db, _tmp) = test_db();
410 let id = create_user(&db.conn, "alice", "password123", Role::Admin).unwrap();
411 assert!(id > 0);
412
413 let user = get_user_by_username(&db.conn, "alice").unwrap().unwrap();
414 assert_eq!(user.username, "alice");
415 assert_eq!(user.role, Role::Admin);
416 assert!(user.password_hash.starts_with("$argon2"));
417 }
418
419 #[test]
420 fn duplicate_username_rejected() {
421 let (db, _tmp) = test_db();
422 create_user(&db.conn, "bob", "pass1", Role::User).unwrap();
423 let result = create_user(&db.conn, "bob", "pass2", Role::User);
424 assert!(result.is_err());
425 }
426
427 #[test]
428 fn list_and_delete_users() {
429 let (db, _tmp) = test_db();
430 let id1 = create_user(&db.conn, "user1", "pass", Role::Admin).unwrap();
431 create_user(&db.conn, "user2", "pass", Role::User).unwrap();
432
433 let users = list_users(&db.conn).unwrap();
434 assert_eq!(users.len(), 2);
435
436 assert!(delete_user(&db.conn, id1).unwrap());
437 let users = list_users(&db.conn).unwrap();
438 assert_eq!(users.len(), 1);
439 assert_eq!(users[0].username, "user2");
440 }
441
442 #[test]
443 fn has_users_empty_and_populated() {
444 let (db, _tmp) = test_db();
445 assert!(!has_users(&db.conn).unwrap());
446 create_user(&db.conn, "first", "pass", Role::Admin).unwrap();
447 assert!(has_users(&db.conn).unwrap());
448 }
449
450 #[test]
451 fn refresh_token_lifecycle() {
452 let (db, _tmp) = test_db();
453 let uid = create_user(&db.conn, "user", "pass", Role::User).unwrap();
454
455 let future_ts = auth::now_unix() as i64 + 86400;
456 store_refresh_token(&db.conn, "tok-123", uid, future_ts).unwrap();
457
458 let tok = get_valid_refresh_token(&db.conn, "tok-123")
460 .unwrap()
461 .unwrap();
462 assert_eq!(tok.user_id, uid);
463
464 assert!(revoke_refresh_token(&db.conn, "tok-123").unwrap());
466 assert!(
467 get_valid_refresh_token(&db.conn, "tok-123")
468 .unwrap()
469 .is_none()
470 );
471 }
472
473 #[test]
474 fn expired_token_not_returned() {
475 let (db, _tmp) = test_db();
476 let uid = create_user(&db.conn, "user", "pass", Role::User).unwrap();
477
478 store_refresh_token(&db.conn, "tok-old", uid, 0).unwrap();
480 assert!(
481 get_valid_refresh_token(&db.conn, "tok-old")
482 .unwrap()
483 .is_none()
484 );
485 }
486
487 #[test]
488 fn cleanup_removes_expired_and_revoked() {
489 let (db, _tmp) = test_db();
490 let uid = create_user(&db.conn, "user", "pass", Role::User).unwrap();
491
492 let future = auth::now_unix() as i64 + 86400;
493 store_refresh_token(&db.conn, "active", uid, future).unwrap();
494 store_refresh_token(&db.conn, "expired", uid, 0).unwrap();
495 store_refresh_token(&db.conn, "revoked", uid, future).unwrap();
496 revoke_refresh_token(&db.conn, "revoked").unwrap();
497
498 let cleaned = cleanup_expired_tokens(&db.conn).unwrap();
499 assert_eq!(cleaned, 2);
500
501 assert!(
503 get_valid_refresh_token(&db.conn, "active")
504 .unwrap()
505 .is_some()
506 );
507 }
508
509 use crate::db::queries::{self, sample_meta, upsert_track};
512 use std::path::Path;
513
514 fn count(db: &Database, sql: &str) -> i64 {
515 db.conn.query_row(sql, [], |r| r.get(0)).unwrap()
516 }
517
518 #[test]
519 fn two_users_star_the_same_track_independently() {
520 let (db, _tmp) = test_db();
521 let admin = create_user(&db.conn, "owner", "pw", Role::Admin).unwrap();
522 let mate = create_user(&db.conn, "mate", "pw", Role::User).unwrap();
523 let path = Path::new("/music/a.flac");
524
525 queries::add_favourite(&db.conn, admin, path).unwrap();
526 queries::add_favourite(&db.conn, mate, path).unwrap();
527 queries::remove_favourite(&db.conn, admin, path).unwrap();
528
529 assert!(
530 queries::load_favourites(&db.conn, admin)
531 .unwrap()
532 .is_empty()
533 );
534 assert!(
535 queries::load_favourites(&db.conn, mate)
536 .unwrap()
537 .contains(path)
538 );
539 assert!(queries::toggle_favourite_album(&db.conn, mate, "Coil", "Scatology").unwrap());
540 assert!(queries::toggle_favourite_album(&db.conn, admin, "Coil", "Scatology").unwrap());
541 assert_eq!(count(&db, "SELECT COUNT(*) FROM favourite_albums"), 2);
542 }
543
544 #[test]
545 fn the_local_user_is_the_first_admin_once_there_is_one() {
546 let (db, _tmp) = test_db();
547 let path = Path::new("/music/a.flac");
548 let track = upsert_track(&db.conn, &sample_meta("T", "A", "B")).unwrap();
549 queries::add_favourite(&db.conn, LOCAL_USER, path).unwrap();
550 queries::record_play(&db.conn, LOCAL_USER, track, None).unwrap();
551 let list = queries::create_playlist(&db.conn, LOCAL_USER, "Mine", None).unwrap();
552 assert_eq!(resolve_user(&db.conn, LOCAL_USER).unwrap(), LOCAL_USER);
553
554 create_user(&db.conn, "mate", "pw", Role::User).unwrap();
555 assert_eq!(resolve_user(&db.conn, LOCAL_USER).unwrap(), LOCAL_USER);
556 let admin = create_user(&db.conn, "owner", "pw", Role::Admin).unwrap();
557
558 assert_eq!(resolve_user(&db.conn, LOCAL_USER).unwrap(), admin);
559 assert!(
560 queries::load_favourites(&db.conn, admin)
561 .unwrap()
562 .contains(path)
563 );
564 assert_eq!(queries::play_count(&db.conn, admin, track).unwrap(), 1);
565 assert_eq!(
566 queries::get_playlist(&db.conn, list)
567 .unwrap()
568 .unwrap()
569 .user_id,
570 admin
571 );
572 assert_eq!(
573 count(&db, "SELECT COUNT(*) FROM favourites WHERE user_id = 0"),
574 0
575 );
576 }
577
578 #[test]
579 fn playlists_are_the_owners_plus_everyones_public_ones() {
580 let (db, _tmp) = test_db();
581 let admin = create_user(&db.conn, "owner", "pw", Role::Admin).unwrap();
582 let mate = create_user(&db.conn, "mate", "pw", Role::User).unwrap();
583 let private = queries::create_playlist(&db.conn, admin, "Private", None).unwrap();
584 let public = queries::create_playlist(&db.conn, admin, "Public", None).unwrap();
585 db.conn
586 .execute("UPDATE playlists SET public = 1 WHERE id = ?1", [public])
587 .unwrap();
588 let own = queries::create_playlist(&db.conn, mate, "Mate's", None).unwrap();
589
590 let ids = |user| -> Vec<i64> {
591 let mut ids: Vec<i64> = queries::list_playlists(&db.conn, user)
592 .unwrap()
593 .into_iter()
594 .map(|p| p.id)
595 .collect();
596 ids.sort_unstable();
597 ids
598 };
599 assert_eq!(ids(mate), vec![public, own]);
600 assert_eq!(ids(admin), vec![private, public]);
601 assert_eq!(ids(LOCAL_USER), vec![private, public]);
603
604 let row = queries::get_playlist(&db.conn, public).unwrap().unwrap();
605 assert!(row.readable_by(mate) && !row.editable_by(mate));
606 assert_eq!(row.owner.as_deref(), Some("owner"));
607 let row = queries::get_playlist(&db.conn, private).unwrap().unwrap();
608 assert!(!row.readable_by(mate));
609 }
610
611 #[test]
612 fn deleting_an_account_takes_its_data_with_it() {
613 let (db, _tmp) = test_db();
614 let admin = create_user(&db.conn, "owner", "pw", Role::Admin).unwrap();
615 let mate = create_user(&db.conn, "mate", "pw", Role::User).unwrap();
616 let track = upsert_track(&db.conn, &sample_meta("T", "A", "B")).unwrap();
617 for user in [admin, mate] {
618 queries::add_favourite(&db.conn, user, Path::new("/music/a.flac")).unwrap();
619 queries::set_favourite_album(&db.conn, user, "A", "B", true).unwrap();
620 queries::set_favourite_artist(&db.conn, user, "A", true).unwrap();
621 queries::record_play(&db.conn, user, track, None).unwrap();
622 queries::create_playlist(&db.conn, user, "List", None).unwrap();
623 queries::shares::create_share(
624 &db.conn,
625 user,
626 queries::shares::Slice::TRACKS,
627 &[track],
628 None,
629 0,
630 None,
631 )
632 .unwrap();
633 }
634
635 assert!(delete_user(&db.conn, mate).unwrap());
636
637 for table in [
638 "favourites",
639 "favourite_albums",
640 "favourite_artists",
641 "play_history",
642 "playlists",
643 "shares",
644 ] {
645 assert_eq!(
646 count(
647 &db,
648 &format!("SELECT COUNT(*) FROM {table} WHERE user_id = {mate}")
649 ),
650 0,
651 "{table} kept the deleted account's rows"
652 );
653 assert_eq!(
654 count(
655 &db,
656 &format!("SELECT COUNT(*) FROM {table} WHERE user_id = {admin}")
657 ),
658 1,
659 "{table} lost another account's rows"
660 );
661 }
662 }
663}