use axum::{
Router,
extract::{Path, Query, State},
http::{HeaderMap, Uri},
response::Redirect,
routing::{get, post},
};
use ocre::{
Ctx, Error, Result, Session,
oauth::{self, Pkce, Provider},
};
use serde::{Deserialize, Serialize};
use crate::{
auth::{OAUTH_PROVIDERS, origin, sign_in},
models::{identity, user},
};
const PENDING: &str = "oauth_pending";
pub fn routes() -> Router<Ctx> {
Router::new().route("/auth/{provider}", post(start)).route("/auth/{provider}/callback", get(callback))
}
#[derive(Serialize, Deserialize)]
struct Pending {
provider: String,
state: String,
verifier: String,
}
#[derive(Deserialize)]
struct Callback {
code: Option<String>,
state: Option<String>,
error: Option<String>,
}
fn provider(name: &str) -> Result<&'static Provider> {
oauth::provider(name).filter(|provider| OAUTH_PROVIDERS.contains(&provider.name)).ok_or(Error::NotFound)
}
fn callback_url(uri: &Uri, provider: &Provider) -> String {
format!("{}/auth/{}/callback", origin(uri), provider.name)
}
async fn start(State(ctx): State<Ctx>, session: Session, uri: Uri, Path(name): Path<String>) -> Result<Redirect> {
let provider = provider(&name)?;
let client_id = ctx.secret(provider.client_id_secret).await?;
let pkce = Pkce::new();
let state = ocre::token::generate();
let url = oauth::authorize_url(provider, &client_id, &callback_url(&uri, provider), &state, &pkce.challenge);
session.insert(PENDING, Pending { provider: provider.name.to_owned(), state, verifier: pkce.verifier })?;
Ok(Redirect::to(&url))
}
async fn callback(
State(ctx): State<Ctx>,
session: Session,
headers: HeaderMap,
uri: Uri,
Path(name): Path<String>,
Query(params): Query<Callback>,
) -> Result<Redirect> {
let provider = provider(&name)?;
let pending: Option<Pending> = session.get(PENDING)?;
session.remove(PENDING)?;
let failed = |message: &str| -> Result<Redirect> {
session.flash("alert", message)?;
Ok(Redirect::to("/login"))
};
let (Some(pending), Some(code), Some(state)) = (pending, params.code, params.state) else {
let reason = if params.error.is_some() { "Sign-in was cancelled." } else { "Sign-in failed. Please try again." };
return failed(reason);
};
if pending.provider != provider.name || !ocre::token::constant_time_eq(pending.state.as_bytes(), state.as_bytes()) {
return failed("Sign-in failed. Please try again.");
}
let token = match oauth::exchange_code(&ctx, provider, &callback_url(&uri, provider), &code, &pending.verifier).await {
Ok(token) => token,
Err(Error::Unauthorized) => return failed("Sign-in failed. Please try again."),
Err(err) => return Err(err),
};
let profile = oauth::profile(provider, &token).await?;
let user = match identity::find_user(&ctx, provider.name, &profile.uid).await? {
Some(user) => user,
None => {
let Some(email) = profile.email.as_deref() else {
return failed("Your account there has no verified email address.");
};
let user = match user::find_by_email(&ctx, email).await? {
Some(user) => {
user::confirm(&ctx, user.id).await?;
user
}
None => user::create_confirmed(&ctx, email).await?,
};
identity::link(&ctx, user.id, provider.name, &profile.uid).await?;
user
}
};
let next = sign_in(&ctx, &session, &headers, &user, false).await?;
session.flash("notice", "Signed in.")?;
Ok(Redirect::to(&next))
}