Skip to main content

auth/
public.rs

1use crate::resolver::session_token_hash;
2use crate::session_policy::{AllowSessionPolicy, AuthSessionPolicy, SessionCreateInput};
3use chrono::{DateTime, Utc};
4use platform_core::{AppError, AppResult, DbPool, ErrorCode};
5use sqlx::{Postgres, Transaction};
6
7pub use crate::models::{AuthSession, AuthUserId};
8pub use crate::session_policy::SessionCreateOptions;
9
10#[derive(Debug, Clone, PartialEq, Eq)]
11pub struct AuthIdentity {
12    pub id: String,
13    pub user_id: AuthUserId,
14}
15
16pub async fn create_user_identity_in_tx(
17    tx: &mut Transaction<'_, Postgres>,
18    user_id: AuthUserId,
19    identity_id: String,
20    provider: &str,
21    provider_subject: &str,
22    created_at: DateTime<Utc>,
23) -> AppResult<AuthIdentity> {
24    sqlx::query(
25        r#"
26        insert into auth.users (id, created_at, disabled_at, disabled_reason, disabled_until)
27        values ($1, $2, null, null, null)
28        "#,
29    )
30    .bind(&user_id.0)
31    .bind(created_at)
32    .execute(&mut **tx)
33    .await
34    .map_err(map_sql_error)?;
35
36    sqlx::query(
37        r#"
38        insert into auth.identities (id, user_id, provider, provider_subject, created_at, updated_at)
39        values ($1, $2, $3, $4, $5, $5)
40        "#,
41    )
42    .bind(&identity_id)
43    .bind(&user_id.0)
44    .bind(provider)
45    .bind(provider_subject)
46    .bind(created_at)
47    .execute(&mut **tx)
48    .await
49    .map_err(map_sql_error)?;
50
51    Ok(AuthIdentity {
52        id: identity_id,
53        user_id,
54    })
55}
56
57pub async fn find_active_identity(
58    pool: &DbPool,
59    provider: &str,
60    provider_subject: &str,
61) -> AppResult<Option<AuthIdentity>> {
62    sqlx::query_as::<_, IdentityRow>(
63        r#"
64        select identities.id, identities.user_id
65        from auth.identities identities
66        join auth.users users on users.id = identities.user_id
67        where identities.provider = $1
68          and identities.provider_subject = $2
69          and (users.disabled_at is null or users.disabled_until <= now())
70        limit 1
71        "#,
72    )
73    .bind(provider)
74    .bind(provider_subject)
75    .fetch_optional(pool)
76    .await
77    .map(|row| row.map(identity_from_row))
78    .map_err(map_sql_error)
79}
80
81pub async fn create_session(
82    pool: &DbPool,
83    user_id: &AuthUserId,
84    session_id: String,
85    token: String,
86    created_at: DateTime<Utc>,
87    expires_at: DateTime<Utc>,
88) -> AppResult<AuthSession> {
89    create_session_with_policy(
90        pool,
91        user_id,
92        session_id,
93        token,
94        created_at,
95        expires_at,
96        SessionCreateOptions::default(),
97        &AllowSessionPolicy,
98    )
99    .await
100}
101
102pub async fn create_session_with_policy(
103    pool: &DbPool,
104    user_id: &AuthUserId,
105    session_id: String,
106    token: String,
107    created_at: DateTime<Utc>,
108    expires_at: DateTime<Utc>,
109    options: SessionCreateOptions,
110    policy: &dyn AuthSessionPolicy,
111) -> AppResult<AuthSession> {
112    let mut tx = pool.begin().await.map_err(map_sql_error)?;
113    let session = create_session_in_tx_with_policy(
114        &mut tx, user_id, session_id, token, created_at, expires_at, options, policy,
115    )
116    .await?;
117    tx.commit().await.map_err(map_sql_error)?;
118    Ok(session)
119}
120
121pub async fn create_session_in_tx(
122    tx: &mut Transaction<'_, Postgres>,
123    user_id: &AuthUserId,
124    session_id: String,
125    token: String,
126    created_at: DateTime<Utc>,
127    expires_at: DateTime<Utc>,
128) -> AppResult<AuthSession> {
129    create_session_in_tx_with_policy(
130        tx,
131        user_id,
132        session_id,
133        token,
134        created_at,
135        expires_at,
136        SessionCreateOptions::default(),
137        &AllowSessionPolicy,
138    )
139    .await
140}
141
142pub async fn create_session_in_tx_with_policy(
143    tx: &mut Transaction<'_, Postgres>,
144    user_id: &AuthUserId,
145    session_id: String,
146    token: String,
147    created_at: DateTime<Utc>,
148    expires_at: DateTime<Utc>,
149    options: SessionCreateOptions,
150    policy: &dyn AuthSessionPolicy,
151) -> AppResult<AuthSession> {
152    let active_user_exists = sqlx::query_scalar::<_, bool>(
153        r#"
154        select exists(
155            select 1
156            from auth.users
157            where id = $1
158              and (disabled_at is null or disabled_until <= now())
159        )
160        "#,
161    )
162    .bind(&user_id.0)
163    .fetch_one(&mut **tx)
164    .await
165    .map_err(map_sql_error)?;
166
167    if !active_user_exists {
168        return Err(AppError::new(ErrorCode::Forbidden, "Auth user is disabled"));
169    }
170
171    let decision = policy
172        .before_session_create(&SessionCreateInput {
173            user_id: user_id.clone(),
174            session_id: session_id.clone(),
175            proposed_device_id: options.device_id,
176            created_at,
177            expires_at,
178        })
179        .await?;
180
181    sqlx::query(
182        r#"
183        insert into auth.sessions (id, user_id, token_hash, device_id, created_at, expires_at, revoked_at)
184        values ($1, $2, $3, $4, $5, $6, null)
185        "#,
186    )
187    .bind(&session_id)
188    .bind(&user_id.0)
189    .bind(session_token_hash(&token))
190    .bind(decision.device_id.as_deref())
191    .bind(created_at)
192    .bind(expires_at)
193    .execute(&mut **tx)
194    .await
195    .map_err(map_sql_error)?;
196
197    Ok(AuthSession {
198        id: session_id,
199        user_id: user_id.clone(),
200        token,
201        device_id: decision.device_id,
202        expires_at,
203    })
204}
205
206type IdentityRow = (String, String);
207
208fn identity_from_row(row: IdentityRow) -> AuthIdentity {
209    let (id, user_id) = row;
210    AuthIdentity {
211        id,
212        user_id: AuthUserId(user_id),
213    }
214}
215
216fn map_sql_error(source: sqlx::Error) -> AppError {
217    if let sqlx::Error::Database(database_error) = &source {
218        if database_error.constraint() == Some("identities_provider_subject_key") {
219            return AppError::new(ErrorCode::Conflict, "An auth identity already exists")
220                .with_source(source);
221        }
222    }
223
224    AppError::new(ErrorCode::Internal, "Internal server error").with_source(source)
225}