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