Skip to main content

auth/
repositories.rs

1use crate::models::{AuthSession, AuthUser, AuthUserId};
2use crate::resolver::session_token_hash;
3use chrono::{DateTime, Utc};
4use platform_core::{AppError, AppResult, DbPool, ErrorCode};
5
6#[async_trait::async_trait]
7pub trait AuthUserRepository: std::fmt::Debug + Send + Sync {
8    async fn insert(&self, user: &AuthUser) -> AppResult<()>;
9    async fn find_by_id(&self, user_id: &AuthUserId) -> AppResult<Option<AuthUser>>;
10    async fn list(&self, limit: i64, cursor: Option<&str>) -> AppResult<Vec<AuthUser>>;
11}
12
13#[derive(Debug, Clone)]
14pub struct PostgresAuthUserRepository {
15    pool: DbPool,
16}
17
18impl PostgresAuthUserRepository {
19    #[must_use]
20    pub fn new(pool: DbPool) -> Self {
21        Self { pool }
22    }
23
24    pub async fn create_dev_session(
25        &self,
26        user_id: AuthUserId,
27        session_id: String,
28        token: String,
29        created_at: DateTime<Utc>,
30        expires_at: DateTime<Utc>,
31    ) -> AppResult<AuthSession> {
32        let mut tx = self.pool.begin().await.map_err(map_sql_error)?;
33
34        sqlx::query(
35            r#"
36            insert into auth.users (id, created_at, disabled_at)
37            values ($1, $2, null)
38            on conflict (id) do nothing
39            "#,
40        )
41        .bind(&user_id.0)
42        .bind(created_at)
43        .execute(&mut *tx)
44        .await
45        .map_err(map_sql_error)?;
46
47        let disabled_at = sqlx::query_scalar::<_, Option<DateTime<Utc>>>(
48            "select disabled_at from auth.users where id = $1",
49        )
50        .bind(&user_id.0)
51        .fetch_one(&mut *tx)
52        .await
53        .map_err(map_sql_error)?;
54
55        if disabled_at.is_some() {
56            return Err(AppError::new(ErrorCode::Forbidden, "Auth user is disabled"));
57        }
58
59        sqlx::query(
60            r#"
61            insert into auth.sessions (id, user_id, token_hash, created_at, expires_at, revoked_at)
62            values ($1, $2, $3, $4, $5, null)
63            "#,
64        )
65        .bind(&session_id)
66        .bind(&user_id.0)
67        .bind(session_token_hash(&token))
68        .bind(created_at)
69        .bind(expires_at)
70        .execute(&mut *tx)
71        .await
72        .map_err(map_sql_error)?;
73
74        tx.commit().await.map_err(map_sql_error)?;
75
76        Ok(AuthSession {
77            id: session_id,
78            user_id,
79            token,
80            expires_at,
81        })
82    }
83
84    pub async fn revoke_session_token(
85        &self,
86        token: &str,
87        revoked_at: DateTime<Utc>,
88    ) -> AppResult<bool> {
89        let result = sqlx::query(
90            r#"
91            update auth.sessions
92            set revoked_at = $2
93            where token_hash = $1
94              and revoked_at is null
95            "#,
96        )
97        .bind(session_token_hash(token))
98        .bind(revoked_at)
99        .execute(&self.pool)
100        .await
101        .map_err(map_sql_error)?;
102
103        Ok(result.rows_affected() > 0)
104    }
105}
106
107#[async_trait::async_trait]
108impl AuthUserRepository for PostgresAuthUserRepository {
109    async fn insert(&self, user: &AuthUser) -> AppResult<()> {
110        sqlx::query(
111            r#"
112            insert into auth.users (id, created_at, disabled_at)
113            values ($1, $2, $3)
114            "#,
115        )
116        .bind(&user.id.0)
117        .bind(user.created_at)
118        .bind(user.disabled_at)
119        .execute(&self.pool)
120        .await
121        .map(|_| ())
122        .map_err(map_sql_error)
123    }
124
125    async fn find_by_id(&self, user_id: &AuthUserId) -> AppResult<Option<AuthUser>> {
126        sqlx::query_as::<_, UserRow>(
127            r#"
128            select id, created_at, disabled_at
129            from auth.users
130            where id = $1
131            "#,
132        )
133        .bind(&user_id.0)
134        .fetch_optional(&self.pool)
135        .await
136        .map(|row| row.map(user_from_row))
137        .map_err(map_sql_error)
138    }
139
140    async fn list(&self, limit: i64, cursor: Option<&str>) -> AppResult<Vec<AuthUser>> {
141        let rows = match cursor {
142            Some(after) => {
143                sqlx::query_as::<_, UserRow>(
144                    r#"
145                    select id, created_at, disabled_at
146                    from auth.users
147                    where id > $1
148                    order by id asc
149                    limit $2
150                    "#,
151                )
152                .bind(after)
153                .bind(limit)
154                .fetch_all(&self.pool)
155                .await
156            }
157            None => {
158                sqlx::query_as::<_, UserRow>(
159                    r#"
160                    select id, created_at, disabled_at
161                    from auth.users
162                    order by id asc
163                    limit $1
164                    "#,
165                )
166                .bind(limit)
167                .fetch_all(&self.pool)
168                .await
169            }
170        }
171        .map_err(map_sql_error)?;
172
173        Ok(rows.into_iter().map(user_from_row).collect())
174    }
175}
176
177type UserRow = (String, DateTime<Utc>, Option<DateTime<Utc>>);
178
179fn user_from_row(row: UserRow) -> AuthUser {
180    let (id, created_at, disabled_at) = row;
181    AuthUser {
182        id: AuthUserId(id),
183        created_at,
184        disabled_at,
185    }
186}
187
188fn map_sql_error(source: sqlx::Error) -> AppError {
189    AppError::new(ErrorCode::Internal, "Internal server error").with_source(source)
190}