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}