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 pub grant: Option<OAuthGrant>,
30}
31
32#[derive(Debug, Clone, PartialEq, Eq)]
33pub struct OAuthGrant {
34 pub id: String,
35 pub client_name: String,
36}
37
38const TOKEN_COLUMNS: &str = "id, user_id, expires_at, revoked, created_at, grant_id, client_name";
39
40fn token_row(row: &rusqlite::Row) -> rusqlite::Result<RefreshTokenRow> {
41 let grant_id: Option<String> = row.get(5)?;
42 let client_name: Option<String> = row.get(6)?;
43 Ok(RefreshTokenRow {
44 id: row.get(0)?,
45 user_id: row.get(1)?,
46 expires_at: row.get(2)?,
47 revoked: row.get::<_, i32>(3)? != 0,
48 created_at: row.get(4)?,
49 grant: grant_id.map(|id| OAuthGrant {
50 id,
51 client_name: client_name.unwrap_or_default(),
52 }),
53 })
54}
55
56pub const LOCAL_USER: i64 = 0;
68
69pub fn first_admin(conn: &Connection) -> Result<Option<i64>, rusqlite::Error> {
71 conn.prepare_cached("SELECT MIN(id) FROM users WHERE role = 'admin'")?
72 .query_row([], |r| r.get(0))
73}
74
75pub fn resolve_user(conn: &Connection, user: i64) -> Result<i64, rusqlite::Error> {
78 if user != LOCAL_USER {
79 return Ok(user);
80 }
81 Ok(first_admin(conn)?.unwrap_or(LOCAL_USER))
82}
83
84pub fn is_local_user(conn: &Connection, user: i64) -> Result<bool, rusqlite::Error> {
87 Ok(resolve_user(conn, user)? == resolve_user(conn, LOCAL_USER)?)
88}
89
90pub fn adopt_local_rows(conn: &Connection) -> Result<(), rusqlite::Error> {
94 let Some(admin) = first_admin(conn)? else {
95 return Ok(());
96 };
97 for table in [
98 "favourites",
99 "favourite_albums",
100 "favourite_artists",
101 "play_history",
102 "playlists",
103 "shares",
104 ] {
105 let pending: bool = conn.query_row(
108 &format!("SELECT EXISTS(SELECT 1 FROM {table} WHERE user_id = ?1)"),
109 params![LOCAL_USER],
110 |r| r.get(0),
111 )?;
112 if !pending {
113 continue;
114 }
115 conn.execute(
116 &format!("UPDATE OR IGNORE {table} SET user_id = ?1 WHERE user_id = ?2"),
117 params![admin, LOCAL_USER],
118 )?;
119 conn.execute(
120 &format!("DELETE FROM {table} WHERE user_id = ?1"),
121 params![LOCAL_USER],
122 )?;
123 }
124 Ok(())
125}
126
127pub fn create_user(
133 conn: &Connection,
134 username: &str,
135 password: &str,
136 role: Role,
137) -> Result<i64, rusqlite::Error> {
138 let hash = auth::hash_password(password)
139 .map_err(|e| rusqlite::Error::ToSqlConversionFailure(e.into()))?;
140 conn.execute(
141 "INSERT INTO users (username, password_hash, role) VALUES (?1, ?2, ?3)",
142 params![username, hash, role.as_str()],
143 )?;
144 let id = conn.last_insert_rowid();
145 adopt_local_rows(conn)?;
146 Ok(id)
147}
148
149pub fn get_user_by_username(
151 conn: &Connection,
152 username: &str,
153) -> Result<Option<UserRow>, rusqlite::Error> {
154 let mut stmt = conn.prepare_cached(
155 "SELECT id, username, password_hash, role, created_at FROM users WHERE username = ?1",
156 )?;
157 let mut rows = stmt.query_map(params![username], |row| {
158 let role_str: String = row.get(3)?;
159 Ok(UserRow {
160 id: row.get(0)?,
161 username: row.get(1)?,
162 password_hash: row.get(2)?,
163 role: role_str.parse().unwrap_or(Role::Readonly),
164 created_at: row.get(4)?,
165 })
166 })?;
167 match rows.next() {
168 Some(Ok(user)) => Ok(Some(user)),
169 Some(Err(e)) => Err(e),
170 None => Ok(None),
171 }
172}
173
174pub fn get_user_by_id(conn: &Connection, user_id: i64) -> Result<Option<UserRow>, rusqlite::Error> {
176 let mut stmt = conn
177 .prepare("SELECT id, username, password_hash, role, created_at FROM users WHERE id = ?1")?;
178 let mut rows = stmt.query_map(params![user_id], |row| {
179 let role_str: String = row.get(3)?;
180 Ok(UserRow {
181 id: row.get(0)?,
182 username: row.get(1)?,
183 password_hash: row.get(2)?,
184 role: role_str.parse().unwrap_or(Role::Readonly),
185 created_at: row.get(4)?,
186 })
187 })?;
188 match rows.next() {
189 Some(Ok(user)) => Ok(Some(user)),
190 Some(Err(e)) => Err(e),
191 None => Ok(None),
192 }
193}
194
195pub fn list_users(conn: &Connection) -> Result<Vec<UserRow>, rusqlite::Error> {
197 let mut stmt = conn
198 .prepare("SELECT id, username, password_hash, role, created_at FROM users ORDER BY id")?;
199 let rows = stmt.query_map([], |row| {
200 let role_str: String = row.get(3)?;
201 Ok(UserRow {
202 id: row.get(0)?,
203 username: row.get(1)?,
204 password_hash: row.get(2)?,
205 role: role_str.parse().unwrap_or(Role::Readonly),
206 created_at: row.get(4)?,
207 })
208 })?;
209 rows.collect()
210}
211
212pub fn delete_user(conn: &Connection, user_id: i64) -> Result<bool, rusqlite::Error> {
214 let count = conn.execute("DELETE FROM users WHERE id = ?1", params![user_id])?;
215 Ok(count > 0)
216}
217
218pub fn update_password(
222 conn: &Connection,
223 username: &str,
224 new_password: &str,
225) -> Result<bool, Box<dyn std::error::Error>> {
226 let hash = crate::auth::hash_password(new_password)?;
227 let updated = conn.execute(
228 "UPDATE users SET password_hash = ?1 WHERE username = ?2",
229 params![hash, username],
230 )?;
231 if updated > 0 {
232 if let Some(user) = get_user_by_username(conn, username)? {
234 revoke_all_user_tokens(conn, user.id)?;
235 super::api_keys::revoke_user_api_keys(conn, user.id)?;
236 }
237 }
238 Ok(updated > 0)
239}
240
241pub fn update_role(
243 conn: &Connection,
244 username: &str,
245 role: crate::auth::Role,
246) -> Result<bool, rusqlite::Error> {
247 let updated = conn.execute(
248 "UPDATE users SET role = ?1 WHERE username = ?2",
249 params![role.as_str(), username],
250 )?;
251 adopt_local_rows(conn)?;
252 Ok(updated > 0)
253}
254
255pub fn has_users(conn: &Connection) -> Result<bool, rusqlite::Error> {
257 let count: i64 = conn.query_row("SELECT COUNT(*) FROM users", [], |row| row.get(0))?;
258 Ok(count > 0)
259}
260
261pub fn admin_count(conn: &Connection) -> Result<i64, rusqlite::Error> {
263 conn.query_row(
264 "SELECT COUNT(*) FROM users WHERE role = 'admin'",
265 [],
266 |row| row.get(0),
267 )
268}
269
270pub fn store_refresh_token(
277 conn: &Connection,
278 token_id: &str,
279 user_id: i64,
280 expires_at: i64,
281) -> Result<(), rusqlite::Error> {
282 store_grant_token(conn, token_id, user_id, expires_at, None)
283}
284
285pub fn store_grant_token(
287 conn: &Connection,
288 token_id: &str,
289 user_id: i64,
290 expires_at: i64,
291 grant: Option<&OAuthGrant>,
292) -> Result<(), rusqlite::Error> {
293 conn.execute(
294 "INSERT INTO refresh_tokens (id, user_id, expires_at, grant_id, client_name)
295 VALUES (?1, ?2, ?3, ?4, ?5)",
296 params![
297 auth::sha256_hex(token_id),
298 user_id,
299 expires_at,
300 grant.map(|g| &g.id),
301 grant.map(|g| &g.client_name),
302 ],
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_cached(&format!(
314 "SELECT {TOKEN_COLUMNS} 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], token_row)?;
318 match rows.next() {
319 Some(Ok(token)) => Ok(Some(token)),
320 Some(Err(e)) => Err(e),
321 None => Ok(None),
322 }
323}
324
325pub fn consume_refresh_token(
329 conn: &Connection,
330 token_id: &str,
331) -> Result<Option<RefreshTokenRow>, rusqlite::Error> {
332 let now = auth::now_unix() as i64;
333 let mut stmt = conn.prepare(&format!(
334 "UPDATE refresh_tokens SET revoked = 1, used_at = ?2
335 WHERE id = ?1 AND revoked = 0 AND expires_at > ?2
336 RETURNING {TOKEN_COLUMNS}"
337 ))?;
338 let mut rows = stmt.query_map(params![auth::sha256_hex(token_id), now], token_row)?;
339 match rows.next() {
340 Some(Ok(token)) => Ok(Some(token)),
341 Some(Err(e)) => Err(e),
342 None => Ok(None),
343 }
344}
345
346pub fn revoke_refresh_token(conn: &Connection, token_id: &str) -> Result<bool, rusqlite::Error> {
348 let count = conn.execute(
349 "UPDATE refresh_tokens SET revoked = 1 WHERE id = ?1",
350 params![auth::sha256_hex(token_id)],
351 )?;
352 Ok(count > 0)
353}
354
355pub fn revoke_all_user_tokens(conn: &Connection, user_id: i64) -> Result<usize, rusqlite::Error> {
357 let count = conn.execute(
358 "UPDATE refresh_tokens SET revoked = 1 WHERE user_id = ?1 AND revoked = 0",
359 params![user_id],
360 )?;
361 Ok(count)
362}
363
364pub fn revoke_replayed_grant(
372 conn: &Connection,
373 token_id: &str,
374 grace_secs: i64,
375) -> Result<usize, rusqlite::Error> {
376 let cutoff = auth::now_unix() as i64 - grace_secs;
377 conn.execute(
378 "UPDATE refresh_tokens SET revoked = 1
379 WHERE revoked = 0 AND grant_id = (
380 SELECT grant_id FROM refresh_tokens
381 WHERE id = ?1 AND revoked = 1 AND grant_id IS NOT NULL AND used_at < ?2)",
382 params![auth::sha256_hex(token_id), cutoff],
383 )
384}
385
386pub fn revoke_grant(conn: &Connection, grant_id: &str) -> Result<usize, rusqlite::Error> {
388 conn.execute(
389 "UPDATE refresh_tokens SET revoked = 1 WHERE grant_id = ?1 AND revoked = 0",
390 params![grant_id],
391 )
392}
393
394pub fn cleanup_expired_tokens(conn: &Connection) -> Result<usize, rusqlite::Error> {
398 let now = auth::now_unix() as i64;
399 let count = conn.execute(
400 "DELETE FROM refresh_tokens
401 WHERE expires_at <= ?1 OR (revoked = 1 AND (grant_id IS NULL OR used_at IS NULL))",
402 params![now],
403 )?;
404 Ok(count)
405}
406
407#[cfg(test)]
412mod tests {
413 use super::*;
414 use crate::db::connection::Database;
415 use tempfile::TempDir;
416
417 fn test_db() -> (Database, TempDir) {
418 let tmp = TempDir::new().unwrap();
419 let db_path = tmp.path().join("test.db");
420 let db = Database::open(&db_path).unwrap();
421 (db, tmp)
422 }
423
424 #[test]
425 fn create_and_get_user() {
426 let (db, _tmp) = test_db();
427 let id = create_user(&db.conn, "alice", "password123", Role::Admin).unwrap();
428 assert!(id > 0);
429
430 let user = get_user_by_username(&db.conn, "alice").unwrap().unwrap();
431 assert_eq!(user.username, "alice");
432 assert_eq!(user.role, Role::Admin);
433 assert!(user.password_hash.starts_with("$argon2"));
434 }
435
436 #[test]
437 fn duplicate_username_rejected() {
438 let (db, _tmp) = test_db();
439 create_user(&db.conn, "bob", "pass1", Role::User).unwrap();
440 let result = create_user(&db.conn, "bob", "pass2", Role::User);
441 assert!(result.is_err());
442 }
443
444 #[test]
445 fn list_and_delete_users() {
446 let (db, _tmp) = test_db();
447 let id1 = create_user(&db.conn, "user1", "pass", Role::Admin).unwrap();
448 create_user(&db.conn, "user2", "pass", Role::User).unwrap();
449
450 let users = list_users(&db.conn).unwrap();
451 assert_eq!(users.len(), 2);
452
453 assert!(delete_user(&db.conn, id1).unwrap());
454 let users = list_users(&db.conn).unwrap();
455 assert_eq!(users.len(), 1);
456 assert_eq!(users[0].username, "user2");
457 }
458
459 #[test]
460 fn has_users_empty_and_populated() {
461 let (db, _tmp) = test_db();
462 assert!(!has_users(&db.conn).unwrap());
463 create_user(&db.conn, "first", "pass", Role::Admin).unwrap();
464 assert!(has_users(&db.conn).unwrap());
465 }
466
467 #[test]
468 fn refresh_token_lifecycle() {
469 let (db, _tmp) = test_db();
470 let uid = create_user(&db.conn, "user", "pass", Role::User).unwrap();
471
472 let future_ts = auth::now_unix() as i64 + 86400;
473 store_refresh_token(&db.conn, "tok-123", uid, future_ts).unwrap();
474
475 let tok = get_valid_refresh_token(&db.conn, "tok-123")
477 .unwrap()
478 .unwrap();
479 assert_eq!(tok.user_id, uid);
480
481 assert!(revoke_refresh_token(&db.conn, "tok-123").unwrap());
483 assert!(
484 get_valid_refresh_token(&db.conn, "tok-123")
485 .unwrap()
486 .is_none()
487 );
488 }
489
490 #[test]
491 fn expired_token_not_returned() {
492 let (db, _tmp) = test_db();
493 let uid = create_user(&db.conn, "user", "pass", Role::User).unwrap();
494
495 store_refresh_token(&db.conn, "tok-old", uid, 0).unwrap();
497 assert!(
498 get_valid_refresh_token(&db.conn, "tok-old")
499 .unwrap()
500 .is_none()
501 );
502 }
503
504 #[test]
505 fn a_spent_grant_token_coming_back_revokes_the_grant() {
506 let (db, _tmp) = test_db();
507 let uid = create_user(&db.conn, "user", "pass", Role::User).unwrap();
508 let future = auth::now_unix() as i64 + 86400;
509 let grant = OAuthGrant {
510 id: "g1".into(),
511 client_name: "Claude".into(),
512 };
513 store_grant_token(&db.conn, "first", uid, future, Some(&grant)).unwrap();
514 let spent = consume_refresh_token(&db.conn, "first").unwrap().unwrap();
515 assert_eq!(spent.grant.as_ref(), Some(&grant));
516 store_grant_token(&db.conn, "second", uid, future, Some(&grant)).unwrap();
517 store_refresh_token(&db.conn, "app", uid, future).unwrap();
519 consume_refresh_token(&db.conn, "app").unwrap().unwrap();
520
521 assert_eq!(revoke_replayed_grant(&db.conn, "first", 30).unwrap(), 0);
523 assert_eq!(revoke_replayed_grant(&db.conn, "app", -1).unwrap(), 0);
524 cleanup_expired_tokens(&db.conn).unwrap();
526 assert_eq!(revoke_replayed_grant(&db.conn, "first", -1).unwrap(), 1);
527 assert!(
528 get_valid_refresh_token(&db.conn, "second")
529 .unwrap()
530 .is_none()
531 );
532 }
533
534 #[test]
535 fn cleanup_removes_expired_and_revoked() {
536 let (db, _tmp) = test_db();
537 let uid = create_user(&db.conn, "user", "pass", Role::User).unwrap();
538
539 let future = auth::now_unix() as i64 + 86400;
540 store_refresh_token(&db.conn, "active", uid, future).unwrap();
541 store_refresh_token(&db.conn, "expired", uid, 0).unwrap();
542 store_refresh_token(&db.conn, "revoked", uid, future).unwrap();
543 revoke_refresh_token(&db.conn, "revoked").unwrap();
544
545 let cleaned = cleanup_expired_tokens(&db.conn).unwrap();
546 assert_eq!(cleaned, 2);
547
548 assert!(
550 get_valid_refresh_token(&db.conn, "active")
551 .unwrap()
552 .is_some()
553 );
554 }
555
556 use crate::db::queries::{self, sample_meta, upsert_track};
559
560 fn count(db: &Database, sql: &str) -> i64 {
561 db.conn.query_row(sql, [], |r| r.get(0)).unwrap()
562 }
563
564 #[test]
565 fn two_users_star_the_same_track_independently() {
566 let (db, _tmp) = test_db();
567 let admin = create_user(&db.conn, "owner", "pw", Role::Admin).unwrap();
568 let mate = create_user(&db.conn, "mate", "pw", Role::User).unwrap();
569 let track = upsert_track(&db.conn, &sample_meta("Scatology", "Coil", "Scatology")).unwrap();
570 let album = queries::get_track_row(&db.conn, track)
571 .unwrap()
572 .unwrap()
573 .album_id
574 .unwrap();
575
576 queries::add_favourite(&db.conn, admin, track).unwrap();
577 queries::add_favourite(&db.conn, mate, track).unwrap();
578 queries::remove_favourite(&db.conn, admin, track).unwrap();
579
580 assert!(
581 queries::load_favourites(&db.conn, admin)
582 .unwrap()
583 .is_empty()
584 );
585 assert!(
586 queries::load_favourites(&db.conn, mate)
587 .unwrap()
588 .contains(&track)
589 );
590 assert!(queries::toggle_favourite_album(&db.conn, mate, album).unwrap());
591 assert!(queries::toggle_favourite_album(&db.conn, admin, album).unwrap());
592 assert_eq!(count(&db, "SELECT COUNT(*) FROM favourite_albums"), 2);
593 }
594
595 #[test]
596 fn the_local_user_is_the_first_admin_once_there_is_one() {
597 let (db, _tmp) = test_db();
598 let track = upsert_track(&db.conn, &sample_meta("T", "A", "B")).unwrap();
599 queries::add_favourite(&db.conn, LOCAL_USER, track).unwrap();
600 queries::record_play(&db.conn, LOCAL_USER, track, None).unwrap();
601 let list = queries::create_playlist(&db.conn, LOCAL_USER, "Mine", None).unwrap();
602 assert_eq!(resolve_user(&db.conn, LOCAL_USER).unwrap(), LOCAL_USER);
603
604 create_user(&db.conn, "mate", "pw", Role::User).unwrap();
605 assert_eq!(resolve_user(&db.conn, LOCAL_USER).unwrap(), LOCAL_USER);
606 let admin = create_user(&db.conn, "owner", "pw", Role::Admin).unwrap();
607
608 assert_eq!(resolve_user(&db.conn, LOCAL_USER).unwrap(), admin);
609 assert!(
610 queries::load_favourites(&db.conn, admin)
611 .unwrap()
612 .contains(&track)
613 );
614 assert_eq!(queries::play_count(&db.conn, admin, track).unwrap(), 1);
615 assert_eq!(
616 queries::get_playlist(&db.conn, list)
617 .unwrap()
618 .unwrap()
619 .user_id,
620 admin
621 );
622 assert_eq!(
623 count(&db, "SELECT COUNT(*) FROM favourites WHERE user_id = 0"),
624 0
625 );
626 }
627
628 #[test]
629 fn playlists_are_the_owners_plus_everyones_public_ones() {
630 let (db, _tmp) = test_db();
631 let admin = create_user(&db.conn, "owner", "pw", Role::Admin).unwrap();
632 let mate = create_user(&db.conn, "mate", "pw", Role::User).unwrap();
633 let private = queries::create_playlist(&db.conn, admin, "Private", None).unwrap();
634 let public = queries::create_playlist(&db.conn, admin, "Public", None).unwrap();
635 db.conn
636 .execute("UPDATE playlists SET public = 1 WHERE id = ?1", [public])
637 .unwrap();
638 let own = queries::create_playlist(&db.conn, mate, "Mate's", None).unwrap();
639
640 let ids = |user| -> Vec<i64> {
641 let mut ids: Vec<i64> = queries::list_playlists(&db.conn, user)
642 .unwrap()
643 .into_iter()
644 .map(|p| p.id)
645 .collect();
646 ids.sort_unstable();
647 ids
648 };
649 assert_eq!(ids(mate), vec![public, own]);
650 assert_eq!(ids(admin), vec![private, public]);
651 assert_eq!(ids(LOCAL_USER), vec![private, public]);
653
654 let row = queries::get_playlist(&db.conn, public).unwrap().unwrap();
655 assert!(row.readable_by(mate) && !row.editable_by(mate));
656 assert_eq!(row.owner.as_deref(), Some("owner"));
657 let row = queries::get_playlist(&db.conn, private).unwrap().unwrap();
658 assert!(!row.readable_by(mate));
659 }
660
661 #[test]
662 fn deleting_an_account_takes_its_data_with_it() {
663 let (db, _tmp) = test_db();
664 let admin = create_user(&db.conn, "owner", "pw", Role::Admin).unwrap();
665 let mate = create_user(&db.conn, "mate", "pw", Role::User).unwrap();
666 let track = upsert_track(&db.conn, &sample_meta("T", "A", "B")).unwrap();
667 let row = queries::get_track_row(&db.conn, track).unwrap().unwrap();
668 for user in [admin, mate] {
669 queries::add_favourite(&db.conn, user, track).unwrap();
670 queries::set_favourite_album(&db.conn, user, row.album_id.unwrap(), true).unwrap();
671 queries::set_favourite_artist(&db.conn, user, row.artist_id.unwrap(), true).unwrap();
672 queries::record_play(&db.conn, user, track, None).unwrap();
673 queries::create_playlist(&db.conn, user, "List", None).unwrap();
674 queries::shares::create_share(
675 &db.conn,
676 user,
677 queries::shares::Slice::TRACKS,
678 &[track],
679 None,
680 0,
681 None,
682 )
683 .unwrap();
684 }
685
686 assert!(delete_user(&db.conn, mate).unwrap());
687
688 for table in [
689 "favourites",
690 "favourite_albums",
691 "favourite_artists",
692 "play_history",
693 "playlists",
694 "shares",
695 ] {
696 assert_eq!(
697 count(
698 &db,
699 &format!("SELECT COUNT(*) FROM {table} WHERE user_id = {mate}")
700 ),
701 0,
702 "{table} kept the deleted account's rows"
703 );
704 assert_eq!(
705 count(
706 &db,
707 &format!("SELECT COUNT(*) FROM {table} WHERE user_id = {admin}")
708 ),
709 1,
710 "{table} lost another account's rows"
711 );
712 }
713 }
714}