Skip to main content

koan_core/db/queries/
auth.rs

1//! Auth queries: user CRUD, refresh token management.
2
3use rusqlite::{Connection, params};
4
5use crate::auth::{self, Role};
6
7// ---------------------------------------------------------------------------
8// Row types
9// ---------------------------------------------------------------------------
10
11#[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
29// ---------------------------------------------------------------------------
30// User CRUD
31// ---------------------------------------------------------------------------
32
33/// Create a new user. Returns the user ID.
34pub fn create_user(
35    conn: &Connection,
36    username: &str,
37    password: &str,
38    role: Role,
39) -> Result<i64, rusqlite::Error> {
40    let hash = auth::hash_password(password)
41        .map_err(|e| rusqlite::Error::ToSqlConversionFailure(e.into()))?;
42    conn.execute(
43        "INSERT INTO users (username, password_hash, role) VALUES (?1, ?2, ?3)",
44        params![username, hash, role.as_str()],
45    )?;
46    Ok(conn.last_insert_rowid())
47}
48
49/// Get a user by username.
50pub fn get_user_by_username(
51    conn: &Connection,
52    username: &str,
53) -> Result<Option<UserRow>, rusqlite::Error> {
54    let mut stmt = conn.prepare(
55        "SELECT id, username, password_hash, role, created_at FROM users WHERE username = ?1",
56    )?;
57    let mut rows = stmt.query_map(params![username], |row| {
58        let role_str: String = row.get(3)?;
59        Ok(UserRow {
60            id: row.get(0)?,
61            username: row.get(1)?,
62            password_hash: row.get(2)?,
63            role: role_str.parse().unwrap_or(Role::Readonly),
64            created_at: row.get(4)?,
65        })
66    })?;
67    match rows.next() {
68        Some(Ok(user)) => Ok(Some(user)),
69        Some(Err(e)) => Err(e),
70        None => Ok(None),
71    }
72}
73
74/// Get a user by ID.
75pub fn get_user_by_id(conn: &Connection, user_id: i64) -> Result<Option<UserRow>, rusqlite::Error> {
76    let mut stmt = conn
77        .prepare("SELECT id, username, password_hash, role, created_at FROM users WHERE id = ?1")?;
78    let mut rows = stmt.query_map(params![user_id], |row| {
79        let role_str: String = row.get(3)?;
80        Ok(UserRow {
81            id: row.get(0)?,
82            username: row.get(1)?,
83            password_hash: row.get(2)?,
84            role: role_str.parse().unwrap_or(Role::Readonly),
85            created_at: row.get(4)?,
86        })
87    })?;
88    match rows.next() {
89        Some(Ok(user)) => Ok(Some(user)),
90        Some(Err(e)) => Err(e),
91        None => Ok(None),
92    }
93}
94
95/// List all users (no password hashes).
96pub fn list_users(conn: &Connection) -> Result<Vec<UserRow>, rusqlite::Error> {
97    let mut stmt = conn
98        .prepare("SELECT id, username, password_hash, role, created_at FROM users ORDER BY id")?;
99    let rows = stmt.query_map([], |row| {
100        let role_str: String = row.get(3)?;
101        Ok(UserRow {
102            id: row.get(0)?,
103            username: row.get(1)?,
104            password_hash: row.get(2)?,
105            role: role_str.parse().unwrap_or(Role::Readonly),
106            created_at: row.get(4)?,
107        })
108    })?;
109    rows.collect()
110}
111
112/// Delete a user by ID. Returns true if a row was deleted.
113pub fn delete_user(conn: &Connection, user_id: i64) -> Result<bool, rusqlite::Error> {
114    let count = conn.execute("DELETE FROM users WHERE id = ?1", params![user_id])?;
115    Ok(count > 0)
116}
117
118/// Update a user's password. Revokes all their refresh tokens.
119pub fn update_password(
120    conn: &Connection,
121    username: &str,
122    new_password: &str,
123) -> Result<bool, Box<dyn std::error::Error>> {
124    let hash = crate::auth::hash_password(new_password)?;
125    let updated = conn.execute(
126        "UPDATE users SET password_hash = ?1 WHERE username = ?2",
127        params![hash, username],
128    )?;
129    if updated > 0 {
130        // Revoke all existing tokens for this user.
131        if let Some(user) = get_user_by_username(conn, username)? {
132            revoke_all_user_tokens(conn, user.id)?;
133        }
134    }
135    Ok(updated > 0)
136}
137
138/// Update a user's role.
139pub fn update_role(
140    conn: &Connection,
141    username: &str,
142    role: crate::auth::Role,
143) -> Result<bool, rusqlite::Error> {
144    let updated = conn.execute(
145        "UPDATE users SET role = ?1 WHERE username = ?2",
146        params![role.as_str(), username],
147    )?;
148    Ok(updated > 0)
149}
150
151/// Check if any users exist (for first-run detection).
152pub fn has_users(conn: &Connection) -> Result<bool, rusqlite::Error> {
153    let count: i64 = conn.query_row("SELECT COUNT(*) FROM users", [], |row| row.get(0))?;
154    Ok(count > 0)
155}
156
157/// Count users with admin role.
158pub fn admin_count(conn: &Connection) -> Result<i64, rusqlite::Error> {
159    conn.query_row(
160        "SELECT COUNT(*) FROM users WHERE role = 'admin'",
161        [],
162        |row| row.get(0),
163    )
164}
165
166// ---------------------------------------------------------------------------
167// Refresh tokens
168// ---------------------------------------------------------------------------
169
170/// Store a refresh token. Only `sha256(token)` is persisted — the raw token is
171/// a bearer credential and read access to the database must not yield one.
172pub fn store_refresh_token(
173    conn: &Connection,
174    token_id: &str,
175    user_id: i64,
176    expires_at: i64,
177) -> Result<(), rusqlite::Error> {
178    conn.execute(
179        "INSERT INTO refresh_tokens (id, user_id, expires_at) VALUES (?1, ?2, ?3)",
180        params![auth::sha256_hex(token_id), user_id, expires_at],
181    )?;
182    Ok(())
183}
184
185/// Look up a refresh token. Returns None if not found, expired, or revoked.
186pub fn get_valid_refresh_token(
187    conn: &Connection,
188    token_id: &str,
189) -> Result<Option<RefreshTokenRow>, rusqlite::Error> {
190    let now = auth::now_unix() as i64;
191    let mut stmt = conn.prepare(
192        "SELECT id, user_id, expires_at, revoked, created_at
193         FROM refresh_tokens
194         WHERE id = ?1 AND revoked = 0 AND expires_at > ?2",
195    )?;
196    let mut rows = stmt.query_map(params![auth::sha256_hex(token_id), now], |row| {
197        Ok(RefreshTokenRow {
198            id: row.get(0)?,
199            user_id: row.get(1)?,
200            expires_at: row.get(2)?,
201            revoked: row.get::<_, i32>(3)? != 0,
202            created_at: row.get(4)?,
203        })
204    })?;
205    match rows.next() {
206        Some(Ok(token)) => Ok(Some(token)),
207        Some(Err(e)) => Err(e),
208        None => Ok(None),
209    }
210}
211
212/// Atomically consume a valid refresh token: revoke it and return the row in one
213/// statement. Returns `None` if the token doesn't exist, is already revoked, or
214/// has expired. This prevents TOCTOU races in refresh-token rotation.
215pub fn consume_refresh_token(
216    conn: &Connection,
217    token_id: &str,
218) -> Result<Option<RefreshTokenRow>, rusqlite::Error> {
219    let now = auth::now_unix() as i64;
220    let mut stmt = conn.prepare(
221        "UPDATE refresh_tokens SET revoked = 1
222         WHERE id = ?1 AND revoked = 0 AND expires_at > ?2
223         RETURNING id, user_id, expires_at, revoked, created_at",
224    )?;
225    let mut rows = stmt.query_map(params![auth::sha256_hex(token_id), now], |row| {
226        Ok(RefreshTokenRow {
227            id: row.get(0)?,
228            user_id: row.get(1)?,
229            expires_at: row.get(2)?,
230            revoked: row.get::<_, i32>(3)? != 0,
231            created_at: row.get(4)?,
232        })
233    })?;
234    match rows.next() {
235        Some(Ok(token)) => Ok(Some(token)),
236        Some(Err(e)) => Err(e),
237        None => Ok(None),
238    }
239}
240
241/// Revoke a single refresh token (logout).
242pub fn revoke_refresh_token(conn: &Connection, token_id: &str) -> Result<bool, rusqlite::Error> {
243    let count = conn.execute(
244        "UPDATE refresh_tokens SET revoked = 1 WHERE id = ?1",
245        params![auth::sha256_hex(token_id)],
246    )?;
247    Ok(count > 0)
248}
249
250/// Revoke all refresh tokens for a user (password change, account delete).
251pub fn revoke_all_user_tokens(conn: &Connection, user_id: i64) -> Result<usize, rusqlite::Error> {
252    let count = conn.execute(
253        "UPDATE refresh_tokens SET revoked = 1 WHERE user_id = ?1 AND revoked = 0",
254        params![user_id],
255    )?;
256    Ok(count)
257}
258
259/// Clean up expired/revoked refresh tokens (housekeeping).
260pub fn cleanup_expired_tokens(conn: &Connection) -> Result<usize, rusqlite::Error> {
261    let now = auth::now_unix() as i64;
262    let count = conn.execute(
263        "DELETE FROM refresh_tokens WHERE revoked = 1 OR expires_at <= ?1",
264        params![now],
265    )?;
266    Ok(count)
267}
268
269// ---------------------------------------------------------------------------
270// Tests
271// ---------------------------------------------------------------------------
272
273#[cfg(test)]
274mod tests {
275    use super::*;
276    use crate::db::connection::Database;
277    use tempfile::TempDir;
278
279    fn test_db() -> (Database, TempDir) {
280        let tmp = TempDir::new().unwrap();
281        let db_path = tmp.path().join("test.db");
282        let db = Database::open(&db_path).unwrap();
283        (db, tmp)
284    }
285
286    #[test]
287    fn create_and_get_user() {
288        let (db, _tmp) = test_db();
289        let id = create_user(&db.conn, "alice", "password123", Role::Admin).unwrap();
290        assert!(id > 0);
291
292        let user = get_user_by_username(&db.conn, "alice").unwrap().unwrap();
293        assert_eq!(user.username, "alice");
294        assert_eq!(user.role, Role::Admin);
295        assert!(user.password_hash.starts_with("$argon2"));
296    }
297
298    #[test]
299    fn duplicate_username_rejected() {
300        let (db, _tmp) = test_db();
301        create_user(&db.conn, "bob", "pass1", Role::User).unwrap();
302        let result = create_user(&db.conn, "bob", "pass2", Role::User);
303        assert!(result.is_err());
304    }
305
306    #[test]
307    fn list_and_delete_users() {
308        let (db, _tmp) = test_db();
309        let id1 = create_user(&db.conn, "user1", "pass", Role::Admin).unwrap();
310        create_user(&db.conn, "user2", "pass", Role::User).unwrap();
311
312        let users = list_users(&db.conn).unwrap();
313        assert_eq!(users.len(), 2);
314
315        assert!(delete_user(&db.conn, id1).unwrap());
316        let users = list_users(&db.conn).unwrap();
317        assert_eq!(users.len(), 1);
318        assert_eq!(users[0].username, "user2");
319    }
320
321    #[test]
322    fn has_users_empty_and_populated() {
323        let (db, _tmp) = test_db();
324        assert!(!has_users(&db.conn).unwrap());
325        create_user(&db.conn, "first", "pass", Role::Admin).unwrap();
326        assert!(has_users(&db.conn).unwrap());
327    }
328
329    #[test]
330    fn refresh_token_lifecycle() {
331        let (db, _tmp) = test_db();
332        let uid = create_user(&db.conn, "user", "pass", Role::User).unwrap();
333
334        let future_ts = auth::now_unix() as i64 + 86400;
335        store_refresh_token(&db.conn, "tok-123", uid, future_ts).unwrap();
336
337        // Valid lookup.
338        let tok = get_valid_refresh_token(&db.conn, "tok-123")
339            .unwrap()
340            .unwrap();
341        assert_eq!(tok.user_id, uid);
342
343        // Revoke.
344        assert!(revoke_refresh_token(&db.conn, "tok-123").unwrap());
345        assert!(
346            get_valid_refresh_token(&db.conn, "tok-123")
347                .unwrap()
348                .is_none()
349        );
350    }
351
352    #[test]
353    fn expired_token_not_returned() {
354        let (db, _tmp) = test_db();
355        let uid = create_user(&db.conn, "user", "pass", Role::User).unwrap();
356
357        // Already expired.
358        store_refresh_token(&db.conn, "tok-old", uid, 0).unwrap();
359        assert!(
360            get_valid_refresh_token(&db.conn, "tok-old")
361                .unwrap()
362                .is_none()
363        );
364    }
365
366    #[test]
367    fn cleanup_removes_expired_and_revoked() {
368        let (db, _tmp) = test_db();
369        let uid = create_user(&db.conn, "user", "pass", Role::User).unwrap();
370
371        let future = auth::now_unix() as i64 + 86400;
372        store_refresh_token(&db.conn, "active", uid, future).unwrap();
373        store_refresh_token(&db.conn, "expired", uid, 0).unwrap();
374        store_refresh_token(&db.conn, "revoked", uid, future).unwrap();
375        revoke_refresh_token(&db.conn, "revoked").unwrap();
376
377        let cleaned = cleanup_expired_tokens(&db.conn).unwrap();
378        assert_eq!(cleaned, 2);
379
380        // Active token still there.
381        assert!(
382            get_valid_refresh_token(&db.conn, "active")
383                .unwrap()
384                .is_some()
385        );
386    }
387}