diff --git a/index.html b/index.html
index 8b9bf5e..4d96169 100644
--- a/index.html
+++ b/index.html
@@ -351,7 +351,6 @@
if (code) {
// Exchange code for token
- // NOTE: This assumes the Authentik client is "Public" and doesn't require a secret
try {
const response = await fetch("https://idm.flegr.me/application/o/token/", {
method: "POST",
@@ -366,6 +365,9 @@
const data = await response.json();
if (data.access_token) {
localStorage.setItem("access_token", data.access_token);
+ if (data.refresh_token) {
+ localStorage.setItem("refresh_token", data.refresh_token);
+ }
window.history.replaceState({}, document.title, window.location.pathname);
} else {
console.error("Token exchange failed:", data);
@@ -375,13 +377,64 @@
}
}
- const storedToken = localStorage.getItem("access_token");
- if (!storedToken) {
- const url = `${AUTH_URL}?client_id=${CLIENT_ID}&response_type=code&redirect_uri=${encodeURIComponent(REDIRECT_URI)}&scope=openid+profile+email`;
+ let accessToken = localStorage.getItem("access_token");
+ const refreshToken = localStorage.getItem("refresh_token");
+
+ // Check if token is expired or about to expire (within 60 seconds)
+ if (accessToken && isTokenExpired(accessToken)) {
+ console.log("Access token expired, attempting refresh...");
+ if (refreshToken) {
+ accessToken = await performTokenRefresh(refreshToken);
+ } else {
+ accessToken = null;
+ }
+ }
+
+ if (!accessToken) {
+ // Redirect to Authentik with offline_access scope to get a refresh_token
+ const url = `${AUTH_URL}?client_id=${CLIENT_ID}&response_type=code&redirect_uri=${encodeURIComponent(REDIRECT_URI)}&scope=openid+profile+email+offline_access`;
window.location.href = url;
return null;
}
- return storedToken;
+ return accessToken;
+ }
+
+ function isTokenExpired(token) {
+ try {
+ const payload = JSON.parse(atob(token.split('.')[1]));
+ const now = Math.floor(Date.now() / 1000);
+ return payload.exp < (now + 60);
+ } catch (e) {
+ return true;
+ }
+ }
+
+ async function performTokenRefresh(refreshToken) {
+ try {
+ const response = await fetch("https://idm.flegr.me/application/o/token/", {
+ method: "POST",
+ headers: { "Content-Type": "application/x-www-form-urlencoded" },
+ body: new URLSearchParams({
+ grant_type: "refresh_token",
+ client_id: CLIENT_ID,
+ refresh_token: refreshToken,
+ })
+ });
+ const data = await response.json();
+ if (data.access_token) {
+ localStorage.setItem("access_token", data.access_token);
+ if (data.refresh_token) {
+ localStorage.setItem("refresh_token", data.refresh_token);
+ }
+ console.log("Token refreshed successfully");
+ return data.access_token;
+ }
+ console.error("Token refresh failed:", data);
+ return null;
+ } catch (e) {
+ console.error("Error during token refresh:", e);
+ return null;
+ }
}
(async () => {
diff --git a/src/main.rs b/src/main.rs
index fb9106f..2dfcb7a 100644
--- a/src/main.rs
+++ b/src/main.rs
@@ -10,6 +10,7 @@ use axum::{
};
use dotenvy::dotenv;
use futures::{sink::SinkExt, stream::StreamExt};
+use jsonwebtoken::{Algorithm, DecodingKey, Validation, decode, decode_header};
use serde::{Deserialize, Serialize};
use std::{collections::HashMap, sync::Arc};
use tokio::sync::{RwLock, mpsc};
@@ -88,15 +89,22 @@ impl ChatState {
// Application state
struct AppState {
chat: ChatState,
+ jwks: RwLock>,
}
#[tokio::main]
async fn main() {
dotenv().ok();
- let app_state = Arc::new(AppState {
- chat: ChatState::new(),
- });
+ let chat = ChatState::new();
+ let jwks = RwLock::new(HashMap::new());
+
+ let app_state = Arc::new(AppState { chat, jwks });
+
+ // Initial JWKS fetch
+ if let Err(e) = fetch_jwks(&app_state).await {
+ eprintln!("Warning: Failed to fetch initial JWKS: {}", e);
+ }
// Build application with routes
let app = Router::new()
@@ -113,6 +121,34 @@ async fn main() {
axum::serve(listener, app).await.unwrap();
}
+async fn fetch_jwks(state: &AppState) -> Result<(), Box> {
+ let resp: JwkSet = reqwest::get("https://idm.flegr.me/application/o/chat/jwks/")
+ .await?
+ .json()
+ .await?;
+
+ let mut keys = state.jwks.write().await;
+ for jwk in resp.keys {
+ if let (Some(kid), Some(n), Some(e)) = (jwk.kid, jwk.n, jwk.e) {
+ keys.insert(kid, (n, e));
+ }
+ }
+ println!("Fetched {} public keys from Authentik", keys.len());
+ Ok(())
+}
+
+#[derive(Deserialize)]
+struct JwkSet {
+ keys: Vec,
+}
+
+#[derive(Deserialize)]
+struct Jwk {
+ kid: Option,
+ n: Option,
+ e: Option,
+}
+
// Handlers
async fn index() -> impl IntoResponse {
@@ -135,26 +171,59 @@ async fn websocket_handler(
Query(params): Query,
State(state): State>,
) -> impl IntoResponse {
- // For now, we'll verify the token by calling the userinfo endpoint.
- // In a production app, you should verify the JWT signature locally using JWKS.
- let client = reqwest::Client::new();
- let user_info_resp = client
- .get("https://idm.flegr.me/application/o/userinfo/")
- .header("User-Agent", "axum-chat-app")
- .bearer_auth(¶ms.token)
- .send()
- .await;
+ let header = match decode_header(¶ms.token) {
+ Ok(h) => h,
+ Err(_) => return StatusCode::BAD_REQUEST.into_response(),
+ };
- match user_info_resp {
- Ok(resp) if resp.status().is_success() => {
- let user: User = match resp.json().await {
- Ok(u) => u,
- Err(_) => return StatusCode::INTERNAL_SERVER_ERROR.into_response(),
- };
- ws.on_upgrade(move |socket| handle_websocket(socket, state, user))
+ let kid = match header.kid {
+ Some(k) => k,
+ None => return StatusCode::UNAUTHORIZED.into_response(),
+ };
+
+ let components = {
+ let jwks = state.jwks.read().await;
+ if let Some(c) = jwks.get(&kid).cloned() {
+ c
+ } else {
+ // Key not found, might need to refresh JWKS
+ drop(jwks);
+ if let Err(e) = fetch_jwks(&state).await {
+ eprintln!("Failed to refresh JWKS: {}", e);
+ return StatusCode::UNAUTHORIZED.into_response();
+ }
+ let jwks = state.jwks.read().await;
+ match jwks.get(&kid).cloned() {
+ Some(c) => c,
+ None => return StatusCode::UNAUTHORIZED.into_response(),
+ }
}
- _ => StatusCode::UNAUTHORIZED.into_response(),
- }
+ };
+
+ let key = match DecodingKey::from_rsa_components(&components.0, &components.1) {
+ Ok(k) => k,
+ Err(_) => return StatusCode::INTERNAL_SERVER_ERROR.into_response(),
+ };
+
+ let mut validation = Validation::new(Algorithm::RS256);
+ validation.set_audience(&["xLL8a23JBnBvEarjQUY6Jgd4zLnnqIi2y3oB6laM"]);
+ validation.set_issuer(&["https://idm.flegr.me/application/o/chat/"]);
+
+ let token_data = match decode::(¶ms.token, &key, &validation) {
+ Ok(c) => c,
+ Err(e) => {
+ eprintln!("Token validation failed: {}", e);
+ return StatusCode::UNAUTHORIZED.into_response();
+ }
+ };
+
+ let user = User {
+ login: token_data.claims.preferred_username,
+ avatar_url: String::new(),
+ };
+
+ let state_clone = state.clone();
+ ws.on_upgrade(move |socket| handle_websocket(socket, state_clone, user))
}
// WebSocket connection handler