Skip to main content

arc_auth_db/
lib.rs

1use arc_auth_core::{AuthError, Identity, IdentityStore};
2use arc_web::{ArcAppBuilder, ArcPlugin, PluginSetupContext};
3use argon2::{
4    password_hash::{PasswordHash, PasswordHasher, PasswordVerifier, SaltString},
5    Argon2,
6};
7use async_trait::async_trait;
8use diesel::{connection::SimpleConnection, prelude::*, sql_query};
9use rand::rngs::OsRng;
10use std::{
11    io::{self, IsTerminal, Write},
12    sync::Arc,
13};
14use uuid::Uuid;
15
16const IDENTITY_ROLES_MIGRATION: &str =
17    include_str!("../migrations/90000000000000_identity_roles/up.sql");
18
19#[derive(Clone)]
20pub struct DbIdentityStore {
21    database_url: String,
22}
23impl DbIdentityStore {
24    pub fn new(database_url: impl Into<String>) -> Self {
25        Self {
26            database_url: database_url.into(),
27        }
28    }
29    fn connect(&self) -> Result<SqliteConnection, AuthError> {
30        SqliteConnection::establish(&self.database_url).map_err(|e| AuthError::Store(e.to_string()))
31    }
32}
33
34#[derive(QueryableByName)]
35struct UserRow {
36    #[diesel(sql_type = diesel::sql_types::Text)]
37    id: String,
38    #[diesel(sql_type = diesel::sql_types::Text)]
39    name: String,
40    #[diesel(sql_type = diesel::sql_types::Text)]
41    email: String,
42    #[diesel(sql_type = diesel::sql_types::Text)]
43    password_hash: String,
44    #[diesel(sql_type = diesel::sql_types::Integer)]
45    active: i32,
46}
47#[derive(QueryableByName)]
48struct RoleRow {
49    #[diesel(sql_type = diesel::sql_types::Text)]
50    name: String,
51}
52
53fn now_us() -> i64 {
54    std::time::SystemTime::now()
55        .duration_since(std::time::UNIX_EPOCH)
56        .unwrap_or_default()
57        .as_micros() as i64
58}
59fn validate(name: &str, email: &str, password: Option<&str>) -> Result<String, AuthError> {
60    if name.trim().is_empty() {
61        return Err(AuthError::InvalidInput("name is required".into()));
62    }
63    let email = email.trim().to_ascii_lowercase();
64    if !email
65        .split_once('@')
66        .is_some_and(|(l, r)| !l.is_empty() && r.contains('.'))
67    {
68        return Err(AuthError::InvalidInput("valid email is required".into()));
69    }
70    if password.is_some_and(|p| p.len() < 12) {
71        return Err(AuthError::InvalidInput(
72            "password must contain at least 12 characters".into(),
73        ));
74    }
75    Ok(email)
76}
77fn hash(password: &str) -> Result<String, AuthError> {
78    Argon2::default()
79        .hash_password(password.as_bytes(), &SaltString::generate(&mut OsRng))
80        .map(|h| h.to_string())
81        .map_err(|e| AuthError::Store(e.to_string()))
82}
83fn roles(connection: &mut SqliteConnection, id: &str) -> Result<Vec<String>, AuthError> {
84    sql_query("SELECT roles.name AS name FROM roles JOIN user_roles ON user_roles.role_id = roles.id WHERE user_roles.user_id = ? ORDER BY roles.name").bind::<diesel::sql_types::Text,_>(id).load::<RoleRow>(connection).map(|rows|rows.into_iter().map(|r|r.name).collect()).map_err(|e|AuthError::Store(e.to_string()))
85}
86fn identity(connection: &mut SqliteConnection, row: UserRow) -> Result<Identity, AuthError> {
87    let assigned = roles(connection, &row.id)?;
88    Ok(Identity {
89        id: row.id,
90        name: row.name,
91        email: row.email,
92        active: row.active != 0,
93        roles: assigned,
94    })
95}
96fn get_row(connection: &mut SqliteConnection, id: &str) -> Result<Option<UserRow>, AuthError> {
97    sql_query("SELECT id,name,email,password_hash,active FROM users WHERE id = ?")
98        .bind::<diesel::sql_types::Text, _>(id)
99        .get_result::<UserRow>(connection)
100        .optional()
101        .map_err(|e| AuthError::Store(e.to_string()))
102}
103
104#[async_trait]
105impl IdentityStore for DbIdentityStore {
106    async fn authenticate(&self, email: &str, password: &str) -> Result<Identity, AuthError> {
107        let mut c = self.connect()?;
108        let row=sql_query("SELECT id,name,email,password_hash,active FROM users WHERE email = ? COLLATE NOCASE AND active = 1").bind::<diesel::sql_types::Text,_>(email.trim()).get_result::<UserRow>(&mut c).optional().map_err(|e|AuthError::Store(e.to_string()))?.ok_or(AuthError::InvalidCredentials)?;
109        let parsed =
110            PasswordHash::new(&row.password_hash).map_err(|_| AuthError::InvalidCredentials)?;
111        Argon2::default()
112            .verify_password(password.as_bytes(), &parsed)
113            .map_err(|_| AuthError::InvalidCredentials)?;
114        identity(&mut c, row)
115    }
116    async fn get(&self, id: &str) -> Result<Option<Identity>, AuthError> {
117        let mut c = self.connect()?;
118        get_row(&mut c, id)?
119            .map(|row| identity(&mut c, row))
120            .transpose()
121    }
122    async fn list(&self) -> Result<Vec<Identity>, AuthError> {
123        let mut c = self.connect()?;
124        let rows = sql_query("SELECT id,name,email,password_hash,active FROM users ORDER BY email")
125            .load::<UserRow>(&mut c)
126            .map_err(|e| AuthError::Store(e.to_string()))?;
127        rows.into_iter().map(|row| identity(&mut c, row)).collect()
128    }
129    async fn has_users(&self) -> Result<bool, AuthError> {
130        let mut c = self.connect()?;
131        #[derive(QueryableByName)]
132        struct Count {
133            #[diesel(sql_type=diesel::sql_types::BigInt)]
134            count: i64,
135        }
136        Ok(sql_query("SELECT COUNT(*) AS count FROM users")
137            .get_result::<Count>(&mut c)
138            .map_err(|e| AuthError::Store(e.to_string()))?
139            .count
140            > 0)
141    }
142    async fn create_user(
143        &self,
144        name: &str,
145        email: &str,
146        password: &str,
147        assigned: &[String],
148    ) -> Result<Identity, AuthError> {
149        let email = validate(name, email, Some(password))?;
150        let mut c = self.connect()?;
151        let id = Uuid::new_v4().to_string();
152        let password_hash = hash(password)?;
153        let now = now_us();
154        c.transaction::<_,diesel::result::Error,_>(|c|{sql_query("INSERT INTO users (id,name,email,password_hash,active,created_at,updated_at) VALUES (?,?,?,?,1,?,?)").bind::<diesel::sql_types::Text,_>(&id).bind::<diesel::sql_types::Text,_>(name.trim()).bind::<diesel::sql_types::Text,_>(&email).bind::<diesel::sql_types::Text,_>(&password_hash).bind::<diesel::sql_types::BigInt,_>(now).bind::<diesel::sql_types::BigInt,_>(now).execute(c)?;for role in assigned{sql_query("INSERT INTO user_roles (user_id,role_id) SELECT ?,id FROM roles WHERE name = ?").bind::<diesel::sql_types::Text,_>(&id).bind::<diesel::sql_types::Text,_>(role).execute(c)?;}Ok(())}).map_err(|e|if matches!(e,diesel::result::Error::DatabaseError(diesel::result::DatabaseErrorKind::UniqueViolation,_)){AuthError::DuplicateEmail}else{AuthError::Store(e.to_string())})?;
155        self.get(&id).await?.ok_or(AuthError::NotFound)
156    }
157    async fn update_profile(
158        &self,
159        id: &str,
160        name: &str,
161        email: &str,
162    ) -> Result<Identity, AuthError> {
163        let email = validate(name, email, None)?;
164        let mut c = self.connect()?;
165        sql_query("UPDATE users SET name=?,email=?,updated_at=? WHERE id=?")
166            .bind::<diesel::sql_types::Text, _>(name.trim())
167            .bind::<diesel::sql_types::Text, _>(&email)
168            .bind::<diesel::sql_types::BigInt, _>(now_us())
169            .bind::<diesel::sql_types::Text, _>(id)
170            .execute(&mut c)
171            .map_err(|e| AuthError::Store(e.to_string()))?;
172        self.get(id).await?.ok_or(AuthError::NotFound)
173    }
174    async fn change_password(&self, id: &str, password: &str) -> Result<(), AuthError> {
175        validate("valid", "v@e.co", Some(password))?;
176        let mut c = self.connect()?;
177        sql_query("UPDATE users SET password_hash=?,updated_at=? WHERE id=?")
178            .bind::<diesel::sql_types::Text, _>(hash(password)?)
179            .bind::<diesel::sql_types::BigInt, _>(now_us())
180            .bind::<diesel::sql_types::Text, _>(id)
181            .execute(&mut c)
182            .map_err(|e| AuthError::Store(e.to_string()))?;
183        Ok(())
184    }
185    async fn set_roles(&self, id: &str, assigned: &[String]) -> Result<Identity, AuthError> {
186        let mut c = self.connect()?;
187        if !assigned.iter().any(|role| role == "admin") {
188            ensure_not_final_admin(&mut c, id)?;
189        }
190        c.transaction::<_, diesel::result::Error, _>(|c| {
191            sql_query("DELETE FROM user_roles WHERE user_id=?")
192                .bind::<diesel::sql_types::Text, _>(id)
193                .execute(c)?;
194            for role in assigned {
195                sql_query(
196                    "INSERT INTO user_roles (user_id,role_id) SELECT ?,id FROM roles WHERE name=?",
197                )
198                .bind::<diesel::sql_types::Text, _>(id)
199                .bind::<diesel::sql_types::Text, _>(role)
200                .execute(c)?;
201            }
202            Ok(())
203        })
204        .map_err(|e| AuthError::Store(e.to_string()))?;
205        self.get(id).await?.ok_or(AuthError::NotFound)
206    }
207    async fn set_active(&self, id: &str, active: bool) -> Result<Identity, AuthError> {
208        let mut c = self.connect()?;
209        if !active {
210            ensure_not_final_admin(&mut c, id)?;
211        }
212        let changed = sql_query("UPDATE users SET active=?,updated_at=? WHERE id=?")
213            .bind::<diesel::sql_types::Integer, _>(i32::from(active))
214            .bind::<diesel::sql_types::BigInt, _>(now_us())
215            .bind::<diesel::sql_types::Text, _>(id)
216            .execute(&mut c)
217            .map_err(|e| AuthError::Store(e.to_string()))?;
218        if changed == 0 {
219            return Err(AuthError::NotFound);
220        }
221        self.get(id).await?.ok_or(AuthError::NotFound)
222    }
223}
224
225fn ensure_not_final_admin(c: &mut SqliteConnection, id: &str) -> Result<(), AuthError> {
226    #[derive(QueryableByName)]
227    struct Count {
228        #[diesel(sql_type=diesel::sql_types::BigInt)]
229        count: i64,
230    }
231    let target = sql_query("SELECT COUNT(*) AS count FROM users JOIN user_roles ON users.id=user_roles.user_id JOIN roles ON roles.id=user_roles.role_id WHERE users.id=? AND users.active=1 AND roles.name='admin'")
232        .bind::<diesel::sql_types::Text,_>(id).get_result::<Count>(c).map_err(|e|AuthError::Store(e.to_string()))?.count;
233    if target == 0 {
234        return Ok(());
235    }
236    let admins = sql_query("SELECT COUNT(DISTINCT users.id) AS count FROM users JOIN user_roles ON users.id=user_roles.user_id JOIN roles ON roles.id=user_roles.role_id WHERE users.active=1 AND roles.name='admin'")
237        .get_result::<Count>(c).map_err(|e|AuthError::Store(e.to_string()))?.count;
238    if admins <= 1 {
239        Err(AuthError::InvalidInput(
240            "the final active administrator cannot be deactivated or lose the admin role".into(),
241        ))
242    } else {
243        Ok(())
244    }
245}
246
247pub struct DbIdentityPlugin {
248    store: Arc<dyn IdentityStore>,
249}
250impl DbIdentityPlugin {
251    pub fn new(database_url: impl Into<String>) -> Self {
252        Self {
253            store: Arc::new(DbIdentityStore::new(database_url)),
254        }
255    }
256}
257#[async_trait]
258impl ArcPlugin for DbIdentityPlugin {
259    fn name(&self) -> &'static str {
260        "auth-db"
261    }
262    fn register(&self, builder: ArcAppBuilder) -> ArcAppBuilder {
263        builder.register_data(self.store.clone())
264    }
265    async fn setup(&self, context: &PluginSetupContext<'_>) -> io::Result<()> {
266        let driver = std::env::var("DATABASE_DRIVER").unwrap_or_else(|_| "sqlite".into());
267        if driver != "sqlite" {
268            return Err(io::Error::new(
269                io::ErrorKind::Unsupported,
270                format!("arc-auth-db currently supports DATABASE_DRIVER=sqlite, not `{driver}`"),
271            ));
272        }
273        let mut c = SqliteConnection::establish(context.database_url).map_err(io::Error::other)?;
274        c.batch_execute(IDENTITY_ROLES_MIGRATION)
275            .map_err(io::Error::other)?;
276        if !self.store.has_users().await.map_err(io::Error::other)? {
277            let name = bootstrap_value("ARC_SETUP_ADMIN_NAME", "Administrator name", false)?;
278            let email = bootstrap_value("ARC_SETUP_ADMIN_EMAIL", "Administrator email", false)?;
279            let password =
280                bootstrap_value("ARC_SETUP_ADMIN_PASSWORD", "Administrator password", true)?;
281            self.store
282                .create_user(&name, &email, &password, &["admin".into()])
283                .await
284                .map_err(io::Error::other)?;
285        }
286        Ok(())
287    }
288}
289
290fn bootstrap_value(key: &str, label: &str, secret: bool) -> io::Result<String> {
291    if let Ok(value) = std::env::var(key) {
292        if !value.trim().is_empty() {
293            return Ok(value);
294        }
295    }
296    if !io::stdin().is_terminal() {
297        return Err(io::Error::other(format!(
298            "{key} is required for noninteractive first-administrator setup"
299        )));
300    }
301    let value = if secret {
302        rpassword::prompt_password(format!("{label}: "))?
303    } else {
304        print!("{label}: ");
305        io::stdout().flush()?;
306        let mut value = String::new();
307        io::stdin().read_line(&mut value)?;
308        value.trim().to_owned()
309    };
310    if value.trim().is_empty() {
311        Err(io::Error::other(format!("{label} cannot be empty")))
312    } else {
313        Ok(value)
314    }
315}