use rusqlite::{Connection, params};
use crate::auth::{self, Role};
#[derive(Debug, Clone)]
pub struct UserRow {
pub id: i64,
pub username: String,
pub password_hash: String,
pub role: Role,
pub created_at: Option<String>,
}
#[derive(Debug, Clone)]
pub struct RefreshTokenRow {
pub id: String,
pub user_id: i64,
pub expires_at: i64,
pub revoked: bool,
pub created_at: Option<String>,
}
pub const LOCAL_USER: i64 = 0;
pub fn first_admin(conn: &Connection) -> Result<Option<i64>, rusqlite::Error> {
conn.query_row("SELECT MIN(id) FROM users WHERE role = 'admin'", [], |r| {
r.get(0)
})
}
pub fn resolve_user(conn: &Connection, user: i64) -> Result<i64, rusqlite::Error> {
if user != LOCAL_USER {
return Ok(user);
}
Ok(first_admin(conn)?.unwrap_or(LOCAL_USER))
}
pub fn is_local_user(conn: &Connection, user: i64) -> Result<bool, rusqlite::Error> {
Ok(resolve_user(conn, user)? == resolve_user(conn, LOCAL_USER)?)
}
pub fn adopt_local_rows(conn: &Connection) -> Result<(), rusqlite::Error> {
let Some(admin) = first_admin(conn)? else {
return Ok(());
};
for table in [
"favourites",
"favourite_albums",
"favourite_artists",
"play_history",
"playlists",
"shares",
] {
let pending: bool = conn.query_row(
&format!("SELECT EXISTS(SELECT 1 FROM {table} WHERE user_id = ?1)"),
params![LOCAL_USER],
|r| r.get(0),
)?;
if !pending {
continue;
}
conn.execute(
&format!("UPDATE OR IGNORE {table} SET user_id = ?1 WHERE user_id = ?2"),
params![admin, LOCAL_USER],
)?;
conn.execute(
&format!("DELETE FROM {table} WHERE user_id = ?1"),
params![LOCAL_USER],
)?;
}
Ok(())
}
pub fn set_sealed_password(
conn: &Connection,
username: &str,
sealed: &[u8],
) -> Result<(), rusqlite::Error> {
conn.execute(
"UPDATE users SET sealed_password = ?2 WHERE username = ?1",
params![username, sealed],
)?;
Ok(())
}
pub fn sealed_password(
conn: &Connection,
username: &str,
) -> Result<Option<Vec<u8>>, rusqlite::Error> {
use rusqlite::OptionalExtension;
Ok(conn
.query_row(
"SELECT sealed_password FROM users WHERE username = ?1",
params![username],
|r| r.get::<_, Option<Vec<u8>>>(0),
)
.optional()?
.flatten())
}
pub fn remember_password(
conn: &Connection,
username: &str,
password: &str,
) -> Result<(), Box<dyn std::error::Error>> {
let key = auth::subsonic_key()?;
set_sealed_password(
conn,
username,
&auth::seal_password(&key, username, password)?,
)?;
Ok(())
}
pub fn create_user(
conn: &Connection,
username: &str,
password: &str,
role: Role,
) -> Result<i64, rusqlite::Error> {
let hash = auth::hash_password(password)
.map_err(|e| rusqlite::Error::ToSqlConversionFailure(e.into()))?;
conn.execute(
"INSERT INTO users (username, password_hash, role) VALUES (?1, ?2, ?3)",
params![username, hash, role.as_str()],
)?;
let id = conn.last_insert_rowid();
adopt_local_rows(conn)?;
Ok(id)
}
pub fn get_user_by_username(
conn: &Connection,
username: &str,
) -> Result<Option<UserRow>, rusqlite::Error> {
let mut stmt = conn.prepare(
"SELECT id, username, password_hash, role, created_at FROM users WHERE username = ?1",
)?;
let mut rows = stmt.query_map(params![username], |row| {
let role_str: String = row.get(3)?;
Ok(UserRow {
id: row.get(0)?,
username: row.get(1)?,
password_hash: row.get(2)?,
role: role_str.parse().unwrap_or(Role::Readonly),
created_at: row.get(4)?,
})
})?;
match rows.next() {
Some(Ok(user)) => Ok(Some(user)),
Some(Err(e)) => Err(e),
None => Ok(None),
}
}
pub fn get_user_by_id(conn: &Connection, user_id: i64) -> Result<Option<UserRow>, rusqlite::Error> {
let mut stmt = conn
.prepare("SELECT id, username, password_hash, role, created_at FROM users WHERE id = ?1")?;
let mut rows = stmt.query_map(params![user_id], |row| {
let role_str: String = row.get(3)?;
Ok(UserRow {
id: row.get(0)?,
username: row.get(1)?,
password_hash: row.get(2)?,
role: role_str.parse().unwrap_or(Role::Readonly),
created_at: row.get(4)?,
})
})?;
match rows.next() {
Some(Ok(user)) => Ok(Some(user)),
Some(Err(e)) => Err(e),
None => Ok(None),
}
}
pub fn list_users(conn: &Connection) -> Result<Vec<UserRow>, rusqlite::Error> {
let mut stmt = conn
.prepare("SELECT id, username, password_hash, role, created_at FROM users ORDER BY id")?;
let rows = stmt.query_map([], |row| {
let role_str: String = row.get(3)?;
Ok(UserRow {
id: row.get(0)?,
username: row.get(1)?,
password_hash: row.get(2)?,
role: role_str.parse().unwrap_or(Role::Readonly),
created_at: row.get(4)?,
})
})?;
rows.collect()
}
pub fn delete_user(conn: &Connection, user_id: i64) -> Result<bool, rusqlite::Error> {
let count = conn.execute("DELETE FROM users WHERE id = ?1", params![user_id])?;
Ok(count > 0)
}
pub fn update_password(
conn: &Connection,
username: &str,
new_password: &str,
) -> Result<bool, Box<dyn std::error::Error>> {
let hash = crate::auth::hash_password(new_password)?;
let updated = conn.execute(
"UPDATE users SET password_hash = ?1 WHERE username = ?2",
params![hash, username],
)?;
if updated > 0 {
if let Some(user) = get_user_by_username(conn, username)? {
revoke_all_user_tokens(conn, user.id)?;
}
}
Ok(updated > 0)
}
pub fn update_role(
conn: &Connection,
username: &str,
role: crate::auth::Role,
) -> Result<bool, rusqlite::Error> {
let updated = conn.execute(
"UPDATE users SET role = ?1 WHERE username = ?2",
params![role.as_str(), username],
)?;
adopt_local_rows(conn)?;
Ok(updated > 0)
}
pub fn has_users(conn: &Connection) -> Result<bool, rusqlite::Error> {
let count: i64 = conn.query_row("SELECT COUNT(*) FROM users", [], |row| row.get(0))?;
Ok(count > 0)
}
pub fn admin_count(conn: &Connection) -> Result<i64, rusqlite::Error> {
conn.query_row(
"SELECT COUNT(*) FROM users WHERE role = 'admin'",
[],
|row| row.get(0),
)
}
pub fn store_refresh_token(
conn: &Connection,
token_id: &str,
user_id: i64,
expires_at: i64,
) -> Result<(), rusqlite::Error> {
conn.execute(
"INSERT INTO refresh_tokens (id, user_id, expires_at) VALUES (?1, ?2, ?3)",
params![auth::sha256_hex(token_id), user_id, expires_at],
)?;
Ok(())
}
pub fn get_valid_refresh_token(
conn: &Connection,
token_id: &str,
) -> Result<Option<RefreshTokenRow>, rusqlite::Error> {
let now = auth::now_unix() as i64;
let mut stmt = conn.prepare(
"SELECT id, user_id, expires_at, revoked, created_at
FROM refresh_tokens
WHERE id = ?1 AND revoked = 0 AND expires_at > ?2",
)?;
let mut rows = stmt.query_map(params![auth::sha256_hex(token_id), now], |row| {
Ok(RefreshTokenRow {
id: row.get(0)?,
user_id: row.get(1)?,
expires_at: row.get(2)?,
revoked: row.get::<_, i32>(3)? != 0,
created_at: row.get(4)?,
})
})?;
match rows.next() {
Some(Ok(token)) => Ok(Some(token)),
Some(Err(e)) => Err(e),
None => Ok(None),
}
}
pub fn consume_refresh_token(
conn: &Connection,
token_id: &str,
) -> Result<Option<RefreshTokenRow>, rusqlite::Error> {
let now = auth::now_unix() as i64;
let mut stmt = conn.prepare(
"UPDATE refresh_tokens SET revoked = 1
WHERE id = ?1 AND revoked = 0 AND expires_at > ?2
RETURNING id, user_id, expires_at, revoked, created_at",
)?;
let mut rows = stmt.query_map(params![auth::sha256_hex(token_id), now], |row| {
Ok(RefreshTokenRow {
id: row.get(0)?,
user_id: row.get(1)?,
expires_at: row.get(2)?,
revoked: row.get::<_, i32>(3)? != 0,
created_at: row.get(4)?,
})
})?;
match rows.next() {
Some(Ok(token)) => Ok(Some(token)),
Some(Err(e)) => Err(e),
None => Ok(None),
}
}
pub fn revoke_refresh_token(conn: &Connection, token_id: &str) -> Result<bool, rusqlite::Error> {
let count = conn.execute(
"UPDATE refresh_tokens SET revoked = 1 WHERE id = ?1",
params![auth::sha256_hex(token_id)],
)?;
Ok(count > 0)
}
pub fn revoke_all_user_tokens(conn: &Connection, user_id: i64) -> Result<usize, rusqlite::Error> {
let count = conn.execute(
"UPDATE refresh_tokens SET revoked = 1 WHERE user_id = ?1 AND revoked = 0",
params![user_id],
)?;
Ok(count)
}
pub fn cleanup_expired_tokens(conn: &Connection) -> Result<usize, rusqlite::Error> {
let now = auth::now_unix() as i64;
let count = conn.execute(
"DELETE FROM refresh_tokens WHERE revoked = 1 OR expires_at <= ?1",
params![now],
)?;
Ok(count)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::db::connection::Database;
use tempfile::TempDir;
fn test_db() -> (Database, TempDir) {
let tmp = TempDir::new().unwrap();
let db_path = tmp.path().join("test.db");
let db = Database::open(&db_path).unwrap();
(db, tmp)
}
#[test]
fn create_and_get_user() {
let (db, _tmp) = test_db();
let id = create_user(&db.conn, "alice", "password123", Role::Admin).unwrap();
assert!(id > 0);
let user = get_user_by_username(&db.conn, "alice").unwrap().unwrap();
assert_eq!(user.username, "alice");
assert_eq!(user.role, Role::Admin);
assert!(user.password_hash.starts_with("$argon2"));
}
#[test]
fn duplicate_username_rejected() {
let (db, _tmp) = test_db();
create_user(&db.conn, "bob", "pass1", Role::User).unwrap();
let result = create_user(&db.conn, "bob", "pass2", Role::User);
assert!(result.is_err());
}
#[test]
fn list_and_delete_users() {
let (db, _tmp) = test_db();
let id1 = create_user(&db.conn, "user1", "pass", Role::Admin).unwrap();
create_user(&db.conn, "user2", "pass", Role::User).unwrap();
let users = list_users(&db.conn).unwrap();
assert_eq!(users.len(), 2);
assert!(delete_user(&db.conn, id1).unwrap());
let users = list_users(&db.conn).unwrap();
assert_eq!(users.len(), 1);
assert_eq!(users[0].username, "user2");
}
#[test]
fn has_users_empty_and_populated() {
let (db, _tmp) = test_db();
assert!(!has_users(&db.conn).unwrap());
create_user(&db.conn, "first", "pass", Role::Admin).unwrap();
assert!(has_users(&db.conn).unwrap());
}
#[test]
fn refresh_token_lifecycle() {
let (db, _tmp) = test_db();
let uid = create_user(&db.conn, "user", "pass", Role::User).unwrap();
let future_ts = auth::now_unix() as i64 + 86400;
store_refresh_token(&db.conn, "tok-123", uid, future_ts).unwrap();
let tok = get_valid_refresh_token(&db.conn, "tok-123")
.unwrap()
.unwrap();
assert_eq!(tok.user_id, uid);
assert!(revoke_refresh_token(&db.conn, "tok-123").unwrap());
assert!(
get_valid_refresh_token(&db.conn, "tok-123")
.unwrap()
.is_none()
);
}
#[test]
fn expired_token_not_returned() {
let (db, _tmp) = test_db();
let uid = create_user(&db.conn, "user", "pass", Role::User).unwrap();
store_refresh_token(&db.conn, "tok-old", uid, 0).unwrap();
assert!(
get_valid_refresh_token(&db.conn, "tok-old")
.unwrap()
.is_none()
);
}
#[test]
fn cleanup_removes_expired_and_revoked() {
let (db, _tmp) = test_db();
let uid = create_user(&db.conn, "user", "pass", Role::User).unwrap();
let future = auth::now_unix() as i64 + 86400;
store_refresh_token(&db.conn, "active", uid, future).unwrap();
store_refresh_token(&db.conn, "expired", uid, 0).unwrap();
store_refresh_token(&db.conn, "revoked", uid, future).unwrap();
revoke_refresh_token(&db.conn, "revoked").unwrap();
let cleaned = cleanup_expired_tokens(&db.conn).unwrap();
assert_eq!(cleaned, 2);
assert!(
get_valid_refresh_token(&db.conn, "active")
.unwrap()
.is_some()
);
}
use crate::db::queries::{self, sample_meta, upsert_track};
use std::path::Path;
fn count(db: &Database, sql: &str) -> i64 {
db.conn.query_row(sql, [], |r| r.get(0)).unwrap()
}
#[test]
fn two_users_star_the_same_track_independently() {
let (db, _tmp) = test_db();
let admin = create_user(&db.conn, "owner", "pw", Role::Admin).unwrap();
let mate = create_user(&db.conn, "mate", "pw", Role::User).unwrap();
let path = Path::new("/music/a.flac");
queries::add_favourite(&db.conn, admin, path).unwrap();
queries::add_favourite(&db.conn, mate, path).unwrap();
queries::remove_favourite(&db.conn, admin, path).unwrap();
assert!(
queries::load_favourites(&db.conn, admin)
.unwrap()
.is_empty()
);
assert!(
queries::load_favourites(&db.conn, mate)
.unwrap()
.contains(path)
);
assert!(queries::toggle_favourite_album(&db.conn, mate, "Coil", "Scatology").unwrap());
assert!(queries::toggle_favourite_album(&db.conn, admin, "Coil", "Scatology").unwrap());
assert_eq!(count(&db, "SELECT COUNT(*) FROM favourite_albums"), 2);
}
#[test]
fn the_local_user_is_the_first_admin_once_there_is_one() {
let (db, _tmp) = test_db();
let path = Path::new("/music/a.flac");
let track = upsert_track(&db.conn, &sample_meta("T", "A", "B")).unwrap();
queries::add_favourite(&db.conn, LOCAL_USER, path).unwrap();
queries::record_play(&db.conn, LOCAL_USER, track, None).unwrap();
let list = queries::create_playlist(&db.conn, LOCAL_USER, "Mine", None).unwrap();
assert_eq!(resolve_user(&db.conn, LOCAL_USER).unwrap(), LOCAL_USER);
create_user(&db.conn, "mate", "pw", Role::User).unwrap();
assert_eq!(resolve_user(&db.conn, LOCAL_USER).unwrap(), LOCAL_USER);
let admin = create_user(&db.conn, "owner", "pw", Role::Admin).unwrap();
assert_eq!(resolve_user(&db.conn, LOCAL_USER).unwrap(), admin);
assert!(
queries::load_favourites(&db.conn, admin)
.unwrap()
.contains(path)
);
assert_eq!(queries::play_count(&db.conn, admin, track).unwrap(), 1);
assert_eq!(
queries::get_playlist(&db.conn, list)
.unwrap()
.unwrap()
.user_id,
admin
);
assert_eq!(
count(&db, "SELECT COUNT(*) FROM favourites WHERE user_id = 0"),
0
);
}
#[test]
fn playlists_are_the_owners_plus_everyones_public_ones() {
let (db, _tmp) = test_db();
let admin = create_user(&db.conn, "owner", "pw", Role::Admin).unwrap();
let mate = create_user(&db.conn, "mate", "pw", Role::User).unwrap();
let private = queries::create_playlist(&db.conn, admin, "Private", None).unwrap();
let public = queries::create_playlist(&db.conn, admin, "Public", None).unwrap();
db.conn
.execute("UPDATE playlists SET public = 1 WHERE id = ?1", [public])
.unwrap();
let own = queries::create_playlist(&db.conn, mate, "Mate's", None).unwrap();
let ids = |user| -> Vec<i64> {
let mut ids: Vec<i64> = queries::list_playlists(&db.conn, user)
.unwrap()
.into_iter()
.map(|p| p.id)
.collect();
ids.sort_unstable();
ids
};
assert_eq!(ids(mate), vec![public, own]);
assert_eq!(ids(admin), vec![private, public]);
assert_eq!(ids(LOCAL_USER), vec![private, public]);
let row = queries::get_playlist(&db.conn, public).unwrap().unwrap();
assert!(row.readable_by(mate) && !row.editable_by(mate));
assert_eq!(row.owner.as_deref(), Some("owner"));
let row = queries::get_playlist(&db.conn, private).unwrap().unwrap();
assert!(!row.readable_by(mate));
}
#[test]
fn deleting_an_account_takes_its_data_with_it() {
let (db, _tmp) = test_db();
let admin = create_user(&db.conn, "owner", "pw", Role::Admin).unwrap();
let mate = create_user(&db.conn, "mate", "pw", Role::User).unwrap();
let track = upsert_track(&db.conn, &sample_meta("T", "A", "B")).unwrap();
for user in [admin, mate] {
queries::add_favourite(&db.conn, user, Path::new("/music/a.flac")).unwrap();
queries::set_favourite_album(&db.conn, user, "A", "B", true).unwrap();
queries::set_favourite_artist(&db.conn, user, "A", true).unwrap();
queries::record_play(&db.conn, user, track, None).unwrap();
queries::create_playlist(&db.conn, user, "List", None).unwrap();
queries::shares::create_share(
&db.conn,
user,
queries::shares::Slice::TRACKS,
&[track],
None,
0,
None,
)
.unwrap();
}
assert!(delete_user(&db.conn, mate).unwrap());
for table in [
"favourites",
"favourite_albums",
"favourite_artists",
"play_history",
"playlists",
"shares",
] {
assert_eq!(
count(
&db,
&format!("SELECT COUNT(*) FROM {table} WHERE user_id = {mate}")
),
0,
"{table} kept the deleted account's rows"
);
assert_eq!(
count(
&db,
&format!("SELECT COUNT(*) FROM {table} WHERE user_id = {admin}")
),
1,
"{table} lost another account's rows"
);
}
}
}