use crate::AuthState;
use crate::BaseClaims;
use crate::client::IssuerData;
use cookie::Expiration;
use openidconnect::{AuthenticationFlow, CsrfToken, Nonce, Scope};
use openidconnect::{AuthorizationCode, OAuth2TokenResponse, core::CoreResponseType};
use rocket::http::SameSite;
use rocket::http::{Cookie, CookieJar};
use rocket::{Route, State, response::Redirect, routes};
use time::OffsetDateTime;
#[get("/keycloak")]
pub async fn keycloak(auth_state: &State<AuthState>) -> Redirect {
let (authorize_url, _csrf_state, _nonce) = auth_state
.client
.client
.authorize_url(
AuthenticationFlow::<CoreResponseType>::AuthorizationCode,
CsrfToken::new_random,
Nonce::new_random,
)
.add_scope(Scope::new("email".to_string()))
.add_scope(Scope::new("profile".to_string()))
.url();
Redirect::to(authorize_url.to_string())
}
fn check_expiration(cookie: &Cookie<'_>) -> (Option<OffsetDateTime>, bool) {
match cookie.expires() {
Some(Expiration::Session) => (None, false),
Some(Expiration::DateTime(offset)) => {
let ts = OffsetDateTime::now_utc();
if offset > ts {
return (Some(offset), false);
} else {
return (Some(offset), true);
}
}
None => (None, false),
}
}
#[get("/callback?<code>&<state>&<iss>&<session_state>")]
pub async fn callback(
jar: &CookieJar<'_>,
auth_state: &State<AuthState>,
code: String,
state: String,
session_state: String,
iss: String,
) -> Result<Redirect, crate::Error> {
if let Some(cookie) = jar.get("access_token") {
let (expiration, expired) = check_expiration(&cookie);
if !expired {
return Ok(Redirect::to(auth_state.config.post_login().to_string()));
}
}
let token = auth_state
.client
.exchange_code(AuthorizationCode::new(code))
.await?;
let expires_at: OffsetDateTime = match token.expires_in() {
Some(expires_in) => OffsetDateTime::now_utc() + expires_in,
None => {
let token_data = auth_state
.validator
.decode::<BaseClaims>(token.access_token().secret())
.unwrap();
OffsetDateTime::from_unix_timestamp(token_data.claims.exp as i64)
.unwrap_or_else(|_| OffsetDateTime::now_utc())
}
};
let validator = &auth_state.validator;
let supported_algs = validator
.get_supported_algorithms_for_issuer(&iss)
.ok_or_else(|| {
eprintln!("unknown issuer: {}", iss);
crate::Error::MissingIssuerUrl
})?;
let chosen_alg = if supported_algs.contains(&"RS256".to_string()) {
"RS256".to_string()
} else if let Some(first) = supported_algs.first() {
first.clone()
} else {
return Err(crate::Error::MissingAlgoForIssuer(iss.into()));
};
jar.add(
Cookie::build(("access_token", token.access_token().secret().to_string()))
.secure(false)
.expires(expires_at)
.http_only(true)
.same_site(SameSite::Lax),
);
let issuer_data = IssuerData {
issuer: iss,
algorithm: chosen_alg,
};
let json = serde_json::to_string(&issuer_data).unwrap();
jar.add(
Cookie::build(("issuer_data", json))
.secure(false)
.expires(expires_at)
.http_only(true)
.same_site(SameSite::Lax),
);
Ok(Redirect::to(auth_state.config.post_login().to_string()))
}
pub fn get_routes() -> Vec<Route> {
routes![keycloak, callback]
}