use axum::{ extract::State, http::{header::SET_COOKIE, HeaderMap, HeaderValue, StatusCode}, response::{IntoResponse, Response}, Json, }; use axum_extra::{ headers::{authorization::Bearer, Authorization, Cookie}, typed_header::TypedHeader, }; use chrono::{Duration as ChronoDuration, Utc}; use diesel::{pg::PgConnection, prelude::*}; use rand::{rngs::OsRng, RngCore}; use serde::{Deserialize, Serialize}; use sha2::{Digest, Sha256}; use uuid::Uuid; use crate::{ auth::{password, AuthenticatedUser}, error::{AppError, AppResult}, models::{NewRefreshToken, NewUser, RefreshToken, Tenant, TenantStatus, User, UserMembership}, schema::{ refresh_tokens, tenants::dsl as tenant_dsl, user_memberships::dsl as memberships_dsl, users::dsl, }, state::AppState, }; use crate::schema::refresh_tokens::dsl as refresh_dsl; const REFRESH_COOKIE_NAME: &str = "refresh_token"; #[derive(Deserialize)] pub struct LoginRequest { pub username: String, pub password: String, #[serde(default)] pub preferred_tenant_slug: Option, } #[derive(Deserialize)] pub struct SignupRequest { pub username: String, pub password: String, } #[derive(Serialize)] pub struct LoginResponse { pub access_token: String, pub token_type: String, pub expires_in: i64, pub tenant: TenantSnippet, } #[derive(Serialize)] pub struct TenantSnippet { pub id: Uuid, pub slug: String, } #[derive(Serialize)] pub struct TenantSelectionResponse { pub access_token: String, pub tenants: Vec, } #[derive(Serialize)] pub struct TenantListResponse { pub tenants: Vec, } #[derive(Deserialize)] pub struct TenantSelectionRequest { pub tenant_id: Uuid, } pub async fn signup( State(state): State, Json(payload): Json, ) -> AppResult { let username = payload.username.trim(); let password = payload.password.trim(); if username.is_empty() || password.is_empty() { return Err(AppError::bad_request( "username and password must not be empty", )); } let mut conn = state.db_unscoped()?; let exists: bool = dsl::users .filter(dsl::username.eq(username)) .first::(&mut conn) .optional()? .is_some(); if exists { return Err(AppError::conflict("username already exists")); } let user_id = Uuid::new_v4(); let password_hash = password::hash_password(password)?; insert_user(&mut conn, user_id, username, &password_hash)?; let tenant_slug = username.to_lowercase(); let tenant = state.tenants.create_tenant( &tenant_slug, None, None, TenantStatus::Creating, &[user_id], Some(user_id), )?; let user: User = dsl::users.find(user_id).first(&mut conn)?; issue_session(&state, &mut conn, &user, tenant.id) } pub async fn login( State(state): State, Json(payload): Json, ) -> AppResult { let mut conn = state.db_unscoped()?; let user: Option = dsl::users .filter(dsl::username.eq(&payload.username)) .first(&mut conn) .optional()?; let user = match user { Some(user) => user, None => return Err(AppError::unauthorized()), }; let valid = password::verify_password(&payload.password, &user.password_hash) .map_err(|_| AppError::unauthorized())?; if !valid { return Err(AppError::unauthorized()); } let memberships: Vec<(UserMembership, Tenant)> = memberships_dsl::user_memberships .inner_join(tenant_dsl::tenants) .filter(memberships_dsl::user_id.eq(user.id)) .load(&mut conn)?; if memberships.is_empty() { return Err(AppError::unauthorized()); } let preferred_slug = payload .preferred_tenant_slug .as_ref() .map(|slug| slug.trim().to_string()) .filter(|slug| !slug.is_empty()); if let Some(tenant) = preferred_slug.as_ref().and_then(|slug| { memberships .iter() .find(|(_, tenant)| tenant.slug.eq_ignore_ascii_case(slug)) }) { return issue_session(&state, &mut conn, &user, tenant.1.id); } if memberships.len() == 1 { let tenant_id = memberships[0].1.id; return issue_session(&state, &mut conn, &user, tenant_id); } let selection_token = state .jwt .generate_tenant_selector_token(user.id) .map_err(AppError::from)?; let tenants = memberships .into_iter() .map(|(_, tenant)| TenantSnippet { id: tenant.id, slug: tenant.slug, }) .collect(); let response = Json(TenantSelectionResponse { access_token: selection_token, tenants, }) .into_response(); Ok(response) } pub async fn refresh( State(state): State, jar: Option>, ) -> AppResult { let cookies = jar.ok_or_else(AppError::unauthorized)?; let refresh_value = cookies .get(REFRESH_COOKIE_NAME) .ok_or_else(AppError::unauthorized)?; let hashed = hash_refresh_token(refresh_value); let mut conn = state.db_unscoped()?; let now = Utc::now(); let now_naive = now.naive_utc(); let token = match refresh_dsl::refresh_tokens .filter(refresh_dsl::token_hash.eq(&hashed)) .filter(refresh_dsl::revoked_at.is_null()) .filter(refresh_dsl::expires_at.gt(now_naive)) .first::(&mut conn) { Ok(token) => token, Err(diesel::result::Error::NotFound) => return Err(AppError::unauthorized()), Err(err) => return Err(AppError::from(err)), }; diesel::update(refresh_dsl::refresh_tokens.filter(refresh_dsl::id.eq(token.id))) .set(( refresh_dsl::revoked_at.eq(now_naive), refresh_dsl::updated_at.eq(now_naive), )) .execute(&mut conn)?; let user: User = dsl::users .find(token.user_id) .first(&mut conn) .map_err(AppError::from)?; issue_session(&state, &mut conn, &user, token.tenant_id) } fn insert_user( conn: &mut PgConnection, id: Uuid, username: &str, password_hash: &str, ) -> AppResult<()> { let new_user = NewUser { id, username: username.to_string(), password_hash: password_hash.to_string(), }; diesel::insert_into(dsl::users) .values(&new_user) .execute(conn) .map(|_| ()) .map_err(AppError::from) } pub async fn select_tenant( State(state): State, TypedHeader(Authorization(bearer)): TypedHeader>, Json(payload): Json, ) -> AppResult { let user_id = match state.jwt.verify_tenant_selector_token(bearer.token()) { Ok(claims) => claims.sub, Err(_) => state .jwt .verify_token(bearer.token()) .map(|claims| claims.sub) .map_err(|_| AppError::unauthorized())?, }; let mut conn = state.db_unscoped()?; let membership_exists = memberships_dsl::user_memberships .filter(memberships_dsl::user_id.eq(user_id)) .filter(memberships_dsl::tenant_id.eq(payload.tenant_id)) .inner_join(tenant_dsl::tenants) .select(memberships_dsl::id) .first::(&mut conn) .optional()?; if membership_exists.is_none() { return Err(AppError::unauthorized()); } let user: User = dsl::users .find(user_id) .first(&mut conn) .map_err(AppError::from)?; issue_session(&state, &mut conn, &user, payload.tenant_id) } pub async fn logout( State(state): State, user: AuthenticatedUser, jar: Option>, ) -> AppResult<(HeaderMap, StatusCode)> { let mut conn = state.db_unscoped()?; let now = Utc::now().naive_utc(); let mut rows_affected = 0; if let Some(cookies) = jar { if let Some(value) = cookies.get(REFRESH_COOKIE_NAME) { let hashed = hash_refresh_token(value); rows_affected = diesel::update( refresh_dsl::refresh_tokens .filter(refresh_dsl::token_hash.eq(hashed)) .filter(refresh_dsl::user_id.eq(user.user_id)) .filter(refresh_dsl::revoked_at.is_null()), ) .set(( refresh_dsl::revoked_at.eq(now), refresh_dsl::updated_at.eq(now), )) .execute(&mut conn) .unwrap_or(0); } } if rows_affected == 0 { let _ = diesel::update( refresh_dsl::refresh_tokens .filter(refresh_dsl::user_id.eq(user.user_id)) .filter(refresh_dsl::revoked_at.is_null()), ) .set(( refresh_dsl::revoked_at.eq(now), refresh_dsl::updated_at.eq(now), )) .execute(&mut conn); } let mut headers = HeaderMap::new(); headers.insert(SET_COOKIE, build_clear_refresh_cookie(&state)); Ok((headers, StatusCode::NO_CONTENT)) } pub async fn me(user: AuthenticatedUser) -> Json { Json(user) } pub async fn list_tenants( State(state): State, auth: Option>>, ) -> AppResult> { let bearer = auth.ok_or_else(AppError::unauthorized)?; let token = bearer.token(); let user_id = match state.jwt.verify_token(token) { Ok(claims) => claims.sub, Err(_) => { let claims = state .jwt .verify_tenant_selector_token(token) .map_err(|_| AppError::unauthorized())?; claims.sub } }; let mut conn = state.db_unscoped()?; let tenants = memberships_dsl::user_memberships .inner_join(tenant_dsl::tenants) .filter(memberships_dsl::user_id.eq(user_id)) .select((tenant_dsl::id, tenant_dsl::slug)) .load::<(Uuid, String)>(&mut conn)? .into_iter() .map(|(id, slug)| TenantSnippet { id, slug }) .collect(); Ok(Json(TenantListResponse { tenants })) } fn issue_session( state: &AppState, conn: &mut PgConnection, user: &User, tenant_id: Uuid, ) -> AppResult { let now = Utc::now(); let access_token = state .jwt .generate_token(user.id, tenant_id, &user.username) .map_err(AppError::from)?; let tenant_slug: String = tenant_dsl::tenants .find(tenant_id) .select(tenant_dsl::slug) .first(conn) .map_err(AppError::from)?; let refresh_value = generate_refresh_token(); let refresh_hash = hash_refresh_token(&refresh_value); let refresh_expires_at = now + ChronoDuration::days(state.config.refresh_token_expiry_days); let new_refresh = NewRefreshToken { id: Uuid::new_v4(), user_id: user.id, token_hash: refresh_hash, issued_at: now.naive_utc(), expires_at: refresh_expires_at.naive_utc(), tenant_id, }; diesel::insert_into(refresh_tokens::table) .values(&new_refresh) .execute(conn)?; let mut response = Json(LoginResponse { access_token, token_type: "Bearer".to_string(), expires_in: state.config.jwt_expiry_minutes * 60, tenant: TenantSnippet { id: tenant_id, slug: tenant_slug, }, }) .into_response(); response.headers_mut().insert( SET_COOKIE, build_refresh_cookie(state, &refresh_value, refresh_expires_at), ); Ok(response) } fn hash_refresh_token(token: &str) -> String { let mut hasher = Sha256::new(); hasher.update(token.as_bytes()); hex::encode(hasher.finalize()) } fn generate_refresh_token() -> String { let mut bytes = [0u8; 32]; OsRng.fill_bytes(&mut bytes); hex::encode(bytes) } fn build_refresh_cookie( state: &AppState, token: &str, expires_at: chrono::DateTime, ) -> HeaderValue { let max_age = ChronoDuration::days(state.config.refresh_token_expiry_days).num_seconds(); let mut parts = vec![format!("{}={}", REFRESH_COOKIE_NAME, token)]; parts.push("Path=/".into()); parts.push("HttpOnly".into()); parts.push("SameSite=Strict".into()); parts.push(format!("Max-Age={}", max_age)); parts.push(format!("Expires={}", expires_at.to_rfc2822())); if state.config.refresh_cookie_secure { parts.push("Secure".into()); } if let Some(domain) = &state.config.refresh_cookie_domain { parts.push(format!("Domain={}", domain)); } HeaderValue::from_str(&parts.join("; ")).expect("valid refresh cookie") } fn build_clear_refresh_cookie(state: &AppState) -> HeaderValue { let mut parts = vec![format!("{}=", REFRESH_COOKIE_NAME)]; parts.push("Path=/".into()); parts.push("HttpOnly".into()); parts.push("SameSite=Strict".into()); parts.push("Max-Age=0".into()); parts.push("Expires=Thu, 01 Jan 1970 00:00:00 GMT".into()); if state.config.refresh_cookie_secure { parts.push("Secure".into()); } if let Some(domain) = &state.config.refresh_cookie_domain { parts.push(format!("Domain={}", domain)); } HeaderValue::from_str(&parts.join("; ")).expect("valid refresh cookie") }