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}