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::{io, sync::Arc};
11use uuid::Uuid;
12
13const IDENTITY_ROLES_MIGRATION: &str =
14    include_str!("../migrations/90000000000000_identity_roles/up.sql");
15
16#[derive(Clone)]
17pub struct DbIdentityStore {
18    database_url: String,
19}
20impl DbIdentityStore {
21    pub fn new(database_url: impl Into<String>) -> Self {
22        Self {
23            database_url: database_url.into(),
24        }
25    }
26    fn connect(&self) -> Result<SqliteConnection, AuthError> {
27        SqliteConnection::establish(&self.database_url).map_err(|e| AuthError::Store(e.to_string()))
28    }
29}
30
31#[derive(QueryableByName)]
32struct UserRow {
33    #[diesel(sql_type = diesel::sql_types::Text)]
34    id: String,
35    #[diesel(sql_type = diesel::sql_types::Text)]
36    name: String,
37    #[diesel(sql_type = diesel::sql_types::Text)]
38    email: String,
39    #[diesel(sql_type = diesel::sql_types::Text)]
40    password_hash: String,
41    #[diesel(sql_type = diesel::sql_types::Integer)]
42    active: i32,
43}
44#[derive(QueryableByName)]
45struct RoleRow {
46    #[diesel(sql_type = diesel::sql_types::Text)]
47    name: String,
48}
49
50fn now_us() -> i64 {
51    std::time::SystemTime::now()
52        .duration_since(std::time::UNIX_EPOCH)
53        .unwrap_or_default()
54        .as_micros() as i64
55}
56fn validate(name: &str, email: &str, password: Option<&str>) -> Result<String, AuthError> {
57    if name.trim().is_empty() {
58        return Err(AuthError::InvalidInput("name is required".into()));
59    }
60    let email = email.trim().to_ascii_lowercase();
61    if !email
62        .split_once('@')
63        .is_some_and(|(l, r)| !l.is_empty() && r.contains('.'))
64    {
65        return Err(AuthError::InvalidInput("valid email is required".into()));
66    }
67    if password.is_some_and(|p| p.len() < 12) {
68        return Err(AuthError::InvalidInput(
69            "password must contain at least 12 characters".into(),
70        ));
71    }
72    Ok(email)
73}
74fn hash(password: &str) -> Result<String, AuthError> {
75    Argon2::default()
76        .hash_password(password.as_bytes(), &SaltString::generate(&mut OsRng))
77        .map(|h| h.to_string())
78        .map_err(|e| AuthError::Store(e.to_string()))
79}
80fn roles(connection: &mut SqliteConnection, id: &str) -> Result<Vec<String>, AuthError> {
81    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()))
82}
83fn identity(connection: &mut SqliteConnection, row: UserRow) -> Result<Identity, AuthError> {
84    let assigned = roles(connection, &row.id)?;
85    Ok(Identity {
86        id: row.id,
87        name: row.name,
88        email: row.email,
89        active: row.active != 0,
90        roles: assigned,
91    })
92}
93fn get_row(connection: &mut SqliteConnection, id: &str) -> Result<Option<UserRow>, AuthError> {
94    sql_query("SELECT id,name,email,password_hash,active FROM users WHERE id = ?")
95        .bind::<diesel::sql_types::Text, _>(id)
96        .get_result::<UserRow>(connection)
97        .optional()
98        .map_err(|e| AuthError::Store(e.to_string()))
99}
100
101#[async_trait]
102impl IdentityStore for DbIdentityStore {
103    async fn authenticate(&self, email: &str, password: &str) -> Result<Identity, AuthError> {
104        let mut c = self.connect()?;
105        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)?;
106        let parsed =
107            PasswordHash::new(&row.password_hash).map_err(|_| AuthError::InvalidCredentials)?;
108        Argon2::default()
109            .verify_password(password.as_bytes(), &parsed)
110            .map_err(|_| AuthError::InvalidCredentials)?;
111        identity(&mut c, row)
112    }
113    async fn get(&self, id: &str) -> Result<Option<Identity>, AuthError> {
114        let mut c = self.connect()?;
115        get_row(&mut c, id)?
116            .map(|row| identity(&mut c, row))
117            .transpose()
118    }
119    async fn list(&self) -> Result<Vec<Identity>, AuthError> {
120        let mut c = self.connect()?;
121        let rows = sql_query("SELECT id,name,email,password_hash,active FROM users ORDER BY email")
122            .load::<UserRow>(&mut c)
123            .map_err(|e| AuthError::Store(e.to_string()))?;
124        rows.into_iter().map(|row| identity(&mut c, row)).collect()
125    }
126    async fn has_users(&self) -> Result<bool, AuthError> {
127        let mut c = self.connect()?;
128        #[derive(QueryableByName)]
129        struct Count {
130            #[diesel(sql_type=diesel::sql_types::BigInt)]
131            count: i64,
132        }
133        Ok(sql_query("SELECT COUNT(*) AS count FROM users")
134            .get_result::<Count>(&mut c)
135            .map_err(|e| AuthError::Store(e.to_string()))?
136            .count
137            > 0)
138    }
139    async fn create_user(
140        &self,
141        name: &str,
142        email: &str,
143        password: &str,
144        assigned: &[String],
145    ) -> Result<Identity, AuthError> {
146        let email = validate(name, email, Some(password))?;
147        let mut c = self.connect()?;
148        let id = Uuid::new_v4().to_string();
149        let password_hash = hash(password)?;
150        let now = now_us();
151        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())})?;
152        self.get(&id).await?.ok_or(AuthError::NotFound)
153    }
154    async fn update_profile(
155        &self,
156        id: &str,
157        name: &str,
158        email: &str,
159    ) -> Result<Identity, AuthError> {
160        let email = validate(name, email, None)?;
161        let mut c = self.connect()?;
162        sql_query("UPDATE users SET name=?,email=?,updated_at=? WHERE id=?")
163            .bind::<diesel::sql_types::Text, _>(name.trim())
164            .bind::<diesel::sql_types::Text, _>(&email)
165            .bind::<diesel::sql_types::BigInt, _>(now_us())
166            .bind::<diesel::sql_types::Text, _>(id)
167            .execute(&mut c)
168            .map_err(|e| AuthError::Store(e.to_string()))?;
169        self.get(id).await?.ok_or(AuthError::NotFound)
170    }
171    async fn change_password(&self, id: &str, password: &str) -> Result<(), AuthError> {
172        validate("valid", "v@e.co", Some(password))?;
173        let mut c = self.connect()?;
174        sql_query("UPDATE users SET password_hash=?,updated_at=? WHERE id=?")
175            .bind::<diesel::sql_types::Text, _>(hash(password)?)
176            .bind::<diesel::sql_types::BigInt, _>(now_us())
177            .bind::<diesel::sql_types::Text, _>(id)
178            .execute(&mut c)
179            .map_err(|e| AuthError::Store(e.to_string()))?;
180        Ok(())
181    }
182    async fn set_roles(&self, id: &str, assigned: &[String]) -> Result<Identity, AuthError> {
183        let mut c = self.connect()?;
184        c.transaction::<_, diesel::result::Error, _>(|c| {
185            sql_query("DELETE FROM user_roles WHERE user_id=?")
186                .bind::<diesel::sql_types::Text, _>(id)
187                .execute(c)?;
188            for role in assigned {
189                sql_query(
190                    "INSERT INTO user_roles (user_id,role_id) SELECT ?,id FROM roles WHERE name=?",
191                )
192                .bind::<diesel::sql_types::Text, _>(id)
193                .bind::<diesel::sql_types::Text, _>(role)
194                .execute(c)?;
195            }
196            Ok(())
197        })
198        .map_err(|e| AuthError::Store(e.to_string()))?;
199        self.get(id).await?.ok_or(AuthError::NotFound)
200    }
201}
202
203pub struct DbIdentityPlugin {
204    store: Arc<dyn IdentityStore>,
205}
206impl DbIdentityPlugin {
207    pub fn new(database_url: impl Into<String>) -> Self {
208        Self {
209            store: Arc::new(DbIdentityStore::new(database_url)),
210        }
211    }
212}
213#[async_trait]
214impl ArcPlugin for DbIdentityPlugin {
215    fn name(&self) -> &'static str {
216        "auth-db"
217    }
218    fn register(&self, builder: ArcAppBuilder) -> ArcAppBuilder {
219        builder.register_data(self.store.clone())
220    }
221    async fn setup(&self, context: &PluginSetupContext<'_>) -> io::Result<()> {
222        let driver = std::env::var("DATABASE_DRIVER").unwrap_or_else(|_| "sqlite".into());
223        if driver != "sqlite" {
224            return Err(io::Error::new(
225                io::ErrorKind::Unsupported,
226                format!("arc-auth-db currently supports DATABASE_DRIVER=sqlite, not `{driver}`"),
227            ));
228        }
229        let mut c = SqliteConnection::establish(context.database_url).map_err(io::Error::other)?;
230        c.batch_execute(IDENTITY_ROLES_MIGRATION)
231            .map_err(io::Error::other)?;
232        if !self.store.has_users().await.map_err(io::Error::other)? {
233            let name = std::env::var("ARC_SETUP_ADMIN_NAME").map_err(|_| io::Error::other("no users exist; set ARC_SETUP_ADMIN_NAME, ARC_SETUP_ADMIN_EMAIL, and ARC_SETUP_ADMIN_PASSWORD for setup"))?;
234            let email = std::env::var("ARC_SETUP_ADMIN_EMAIL").map_err(|_| {
235                io::Error::other("ARC_SETUP_ADMIN_EMAIL is required for first-admin setup")
236            })?;
237            let password = std::env::var("ARC_SETUP_ADMIN_PASSWORD").map_err(|_| {
238                io::Error::other("ARC_SETUP_ADMIN_PASSWORD is required for first-admin setup")
239            })?;
240            self.store
241                .create_user(&name, &email, &password, &["admin".into()])
242                .await
243                .map_err(io::Error::other)?;
244        }
245        Ok(())
246    }
247}