feat: refresh token

This commit is contained in:
2026-08-29 12:58:34 +02:00
parent ec1bbf729f
commit a9194fa66c
22 changed files with 537 additions and 15 deletions

View File

@@ -0,0 +1,10 @@
CREATE TABLE IF NOT EXISTS refresh_sessions (
id TEXT PRIMARY KEY NOT NULL,
user_id TEXT NOT NULL REFERENCES users(id) ON DELETE CASCADE,
token TEXT NOT NULL UNIQUE,
expires_at TEXT NOT NULL,
created_at TEXT NOT NULL
);
CREATE INDEX IF NOT EXISTS idx_refresh_sessions_token ON refresh_sessions(token);
CREATE INDEX IF NOT EXISTS idx_refresh_sessions_user_id ON refresh_sessions(user_id);

View File

@@ -1,5 +1,6 @@
pub mod db; pub mod db;
pub mod link; pub mod link;
pub mod note; pub mod note;
pub mod refresh_session;
pub mod tag; pub mod tag;
pub mod user; pub mod user;

View File

@@ -0,0 +1,107 @@
use async_trait::async_trait;
use sqlx::{FromRow, SqlitePool};
use domain::{
auth::{RefreshSession, RefreshSessionId, ports::RefreshSessionRepository},
errors::DomainResult,
user::UserId,
};
use crate::db::{RepoExt, parse_dt};
pub struct SqliteRefreshSessionRepository {
pool: SqlitePool,
}
impl SqliteRefreshSessionRepository {
pub fn new(pool: SqlitePool) -> Self {
Self { pool }
}
}
#[derive(FromRow)]
struct RefreshSessionRow {
id: String,
user_id: String,
token: String,
expires_at: String,
created_at: String,
}
impl TryFrom<RefreshSessionRow> for RefreshSession {
type Error = domain::errors::DomainError;
fn try_from(row: RefreshSessionRow) -> Result<Self, Self::Error> {
use domain::errors::DomainError;
let id = RefreshSessionId::from_uuid(
uuid::Uuid::parse_str(&row.id)
.map_err(|e| DomainError::Repository(format!("invalid session uuid: {e}")))?,
);
let user_id = UserId::from_uuid(
uuid::Uuid::parse_str(&row.user_id)
.map_err(|e| DomainError::Repository(format!("invalid user uuid: {e}")))?,
);
let expires_at = parse_dt(&row.expires_at)?;
let created_at = parse_dt(&row.created_at)?;
Ok(RefreshSession::from_persistence(
id, user_id, row.token, expires_at, created_at,
))
}
}
#[async_trait]
impl RefreshSessionRepository for SqliteRefreshSessionRepository {
async fn create(&self, session: &RefreshSession) -> DomainResult<()> {
sqlx::query(
"INSERT INTO refresh_sessions (id, user_id, token, expires_at, created_at) VALUES (?, ?, ?, ?, ?)",
)
.bind(session.id().as_uuid().to_string())
.bind(session.user_id().as_uuid().to_string())
.bind(session.token())
.bind(session.expires_at().to_rfc3339())
.bind(session.created_at().to_rfc3339())
.execute(&self.pool)
.await
.repo()
.map(|_| ())
}
async fn find_by_token(&self, token: &str) -> DomainResult<Option<RefreshSession>> {
sqlx::query_as::<_, RefreshSessionRow>(
"SELECT id, user_id, token, expires_at, created_at FROM refresh_sessions WHERE token = ?",
)
.bind(token)
.fetch_optional(&self.pool)
.await
.repo()?
.map(RefreshSession::try_from)
.transpose()
}
async fn revoke(&self, token: &str) -> DomainResult<()> {
sqlx::query("DELETE FROM refresh_sessions WHERE token = ?")
.bind(token)
.execute(&self.pool)
.await
.repo()
.map(|_| ())
}
async fn revoke_all_for_user(&self, user_id: &UserId) -> DomainResult<()> {
sqlx::query("DELETE FROM refresh_sessions WHERE user_id = ?")
.bind(user_id.as_uuid().to_string())
.execute(&self.pool)
.await
.repo()
.map(|_| ())
}
async fn delete_expired(&self) -> DomainResult<u64> {
let result = sqlx::query("DELETE FROM refresh_sessions WHERE expires_at < ?")
.bind(chrono::Utc::now().to_rfc3339())
.execute(&self.pool)
.await
.repo()?;
Ok(result.rows_affected())
}
}

View File

@@ -14,6 +14,16 @@ pub struct RegisterRequest {
pub password: String, pub password: String,
} }
#[derive(Debug, Deserialize, utoipa::ToSchema)]
pub struct RefreshRequest {
pub refresh_token: String,
}
#[derive(Debug, Deserialize, utoipa::ToSchema)]
pub struct LogoutRequest {
pub refresh_token: String,
}
#[derive(Debug, Serialize, utoipa::ToSchema)] #[derive(Debug, Serialize, utoipa::ToSchema)]
pub struct UserResponse { pub struct UserResponse {
pub id: Uuid, pub id: Uuid,
@@ -26,4 +36,5 @@ pub struct UserResponse {
pub struct AuthResponse { pub struct AuthResponse {
pub user: UserResponse, pub user: UserResponse,
pub access_token: String, pub access_token: String,
pub refresh_token: String,
} }

View File

@@ -0,0 +1,7 @@
use domain::errors::DomainResult;
use crate::context::AppContext;
pub async fn execute(ctx: &AppContext, refresh_token: &str) -> DomainResult<()> {
ctx.repos.refresh_session.revoke(refresh_token).await
}

View File

@@ -1,3 +1,5 @@
pub mod commands; pub mod commands;
pub mod login; pub mod login;
pub mod logout;
pub mod refresh;
pub mod register; pub mod register;

View File

@@ -0,0 +1,43 @@
use domain::{
auth::RefreshSession,
errors::{DomainError, DomainResult},
user::entity::User,
};
use crate::context::AppContext;
pub struct RefreshResult {
pub user: User,
pub refresh_token: String,
}
pub async fn execute(ctx: &AppContext, old_refresh_token: &str) -> DomainResult<RefreshResult> {
let session = ctx
.repos
.refresh_session
.find_by_token(old_refresh_token)
.await?
.ok_or_else(|| DomainError::Unauthorized("invalid refresh token".into()))?;
ctx.repos.refresh_session.revoke(old_refresh_token).await?;
if session.is_expired() {
return Err(DomainError::Unauthorized("refresh token expired".into()));
}
let user = ctx
.repos
.user
.find_by_id(&session.user_id())
.await?
.ok_or_else(|| DomainError::NotFound("user no longer exists".into()))?;
let new_session = RefreshSession::new(session.user_id(), ctx.config.refresh_token_ttl_seconds);
let refresh_token = new_session.token().to_string();
ctx.repos.refresh_session.create(&new_session).await?;
Ok(RefreshResult {
user,
refresh_token,
})
}

View File

@@ -5,6 +5,8 @@ pub struct AppConfig {
pub smart: SmartConfig, pub smart: SmartConfig,
/// When false the `/auth/register` endpoint returns 403. /// When false the `/auth/register` endpoint returns 403.
pub allow_registration: bool, pub allow_registration: bool,
/// Refresh token time-to-live in seconds. Default: 30 days.
pub refresh_token_ttl_seconds: i64,
} }
#[derive(Debug, Clone)] #[derive(Debug, Clone)]
@@ -28,6 +30,7 @@ impl AppConfig {
base_url: std::env::var("BASE_URL").unwrap_or_else(|_| "http://localhost:3000".into()), base_url: std::env::var("BASE_URL").unwrap_or_else(|_| "http://localhost:3000".into()),
smart: SmartConfig::default(), smart: SmartConfig::default(),
allow_registration: true, allow_registration: true,
refresh_token_ttl_seconds: 30 * 24 * 3600,
} }
} }
} }

View File

@@ -1,6 +1,7 @@
use std::sync::Arc; use std::sync::Arc;
use domain::{ use domain::{
auth::ports::RefreshSessionRepository,
events::{EventConsumer, EventPublisher}, events::{EventConsumer, EventPublisher},
note::ports::{LinkRepository, NoteRepository}, note::ports::{LinkRepository, NoteRepository},
smart::ports::{EmbeddingGenerator, VectorStore}, smart::ports::{EmbeddingGenerator, VectorStore},
@@ -16,6 +17,7 @@ pub struct Repositories {
pub tag: Arc<dyn TagRepository>, pub tag: Arc<dyn TagRepository>,
pub user: Arc<dyn UserRepository>, pub user: Arc<dyn UserRepository>,
pub link: Arc<dyn LinkRepository>, pub link: Arc<dyn LinkRepository>,
pub refresh_session: Arc<dyn RefreshSessionRepository>,
} }
#[derive(Clone)] #[derive(Clone)]

View File

@@ -253,6 +253,54 @@ impl LinkRepository for MemoryLinkRepo {
} }
} }
// ── RefreshSession ──────────────────────────────────────────────────────────
#[derive(Default)]
pub struct MemoryRefreshSessionRepo {
sessions: Mutex<Vec<domain::auth::RefreshSession>>,
}
#[async_trait]
impl domain::auth::ports::RefreshSessionRepository for MemoryRefreshSessionRepo {
async fn create(&self, session: &domain::auth::RefreshSession) -> DomainResult<()> {
self.sessions.lock().unwrap().push(session.clone());
Ok(())
}
async fn find_by_token(
&self,
token: &str,
) -> DomainResult<Option<domain::auth::RefreshSession>> {
Ok(self
.sessions
.lock()
.unwrap()
.iter()
.find(|s| s.token() == token)
.cloned())
}
async fn revoke(&self, token: &str) -> DomainResult<()> {
self.sessions.lock().unwrap().retain(|s| s.token() != token);
Ok(())
}
async fn revoke_all_for_user(&self, user_id: &UserId) -> DomainResult<()> {
self.sessions
.lock()
.unwrap()
.retain(|s| s.user_id() != *user_id);
Ok(())
}
async fn delete_expired(&self) -> DomainResult<u64> {
let before = self.sessions.lock().unwrap().len();
self.sessions.lock().unwrap().retain(|s| !s.is_expired());
let after = self.sessions.lock().unwrap().len();
Ok((before - after) as u64)
}
}
// ── PasswordHasher ─────────────────────────────────────────────────────────── // ── PasswordHasher ───────────────────────────────────────────────────────────
pub struct PlaintextHasher; pub struct PlaintextHasher;
@@ -310,6 +358,7 @@ impl TestContext {
tag: Arc::new(MemoryTagRepo::default()), tag: Arc::new(MemoryTagRepo::default()),
user: Arc::new(MemoryUserRepo::default()), user: Arc::new(MemoryUserRepo::default()),
link: Arc::new(MemoryLinkRepo::default()), link: Arc::new(MemoryLinkRepo::default()),
refresh_session: Arc::new(MemoryRefreshSessionRepo::default()),
}, },
services: Services { services: Services {
password_hasher: Arc::new(PlaintextHasher), password_hasher: Arc::new(PlaintextHasher),
@@ -322,6 +371,7 @@ impl TestContext {
base_url: "http://localhost:3000".into(), base_url: "http://localhost:3000".into(),
smart: SmartConfig::default(), smart: SmartConfig::default(),
allow_registration: true, allow_registration: true,
refresh_token_ttl_seconds: 2_592_000,
}, },
}; };

View File

@@ -0,0 +1,4 @@
pub mod ports;
pub mod refresh_session;
pub use refresh_session::{RefreshSession, RefreshSessionId};

View File

@@ -0,0 +1,14 @@
use async_trait::async_trait;
use super::RefreshSession;
use crate::errors::DomainResult;
use crate::user::UserId;
#[async_trait]
pub trait RefreshSessionRepository: Send + Sync {
async fn create(&self, session: &RefreshSession) -> DomainResult<()>;
async fn find_by_token(&self, token: &str) -> DomainResult<Option<RefreshSession>>;
async fn revoke(&self, token: &str) -> DomainResult<()>;
async fn revoke_all_for_user(&self, user_id: &UserId) -> DomainResult<()>;
async fn delete_expired(&self) -> DomainResult<u64>;
}

View File

@@ -0,0 +1,83 @@
use chrono::{DateTime, Duration, Utc};
use uuid::Uuid;
use crate::user::UserId;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct RefreshSessionId(Uuid);
impl RefreshSessionId {
pub fn generate() -> Self {
Self(Uuid::new_v4())
}
pub fn from_uuid(id: Uuid) -> Self {
Self(id)
}
pub fn as_uuid(self) -> Uuid {
self.0
}
}
#[derive(Debug, Clone)]
pub struct RefreshSession {
id: RefreshSessionId,
user_id: UserId,
token: String,
expires_at: DateTime<Utc>,
created_at: DateTime<Utc>,
}
impl RefreshSession {
pub fn new(user_id: UserId, ttl_seconds: i64) -> Self {
let now = Utc::now();
Self {
id: RefreshSessionId::generate(),
user_id,
token: Uuid::new_v4().to_string(),
expires_at: now + Duration::seconds(ttl_seconds),
created_at: now,
}
}
pub fn from_persistence(
id: RefreshSessionId,
user_id: UserId,
token: String,
expires_at: DateTime<Utc>,
created_at: DateTime<Utc>,
) -> Self {
Self {
id,
user_id,
token,
expires_at,
created_at,
}
}
pub fn id(&self) -> RefreshSessionId {
self.id
}
pub fn user_id(&self) -> UserId {
self.user_id
}
pub fn token(&self) -> &str {
&self.token
}
pub fn expires_at(&self) -> DateTime<Utc> {
self.expires_at
}
pub fn created_at(&self) -> DateTime<Utc> {
self.created_at
}
pub fn is_expired(&self) -> bool {
Utc::now() > self.expires_at
}
}

View File

@@ -8,6 +8,8 @@ pub enum DomainError {
Conflict(String), Conflict(String),
#[error("forbidden: {0}")] #[error("forbidden: {0}")]
Forbidden(String), Forbidden(String),
#[error("unauthorized: {0}")]
Unauthorized(String),
#[error("validation: {0}")] #[error("validation: {0}")]
Validation(String), Validation(String),
#[error("repository: {0}")] #[error("repository: {0}")]

View File

@@ -1,3 +1,4 @@
pub mod auth;
pub mod errors; pub mod errors;
pub mod events; pub mod events;
pub mod note; pub mod note;

View File

@@ -30,6 +30,7 @@ impl From<DomainError> for ApiError {
match e { match e {
DomainError::NotFound(msg) => Self::NotFound(msg), DomainError::NotFound(msg) => Self::NotFound(msg),
DomainError::Forbidden(msg) => Self::Forbidden(msg), DomainError::Forbidden(msg) => Self::Forbidden(msg),
DomainError::Unauthorized(_) => Self::Unauthorized,
DomainError::Conflict(msg) => Self::Conflict(msg), DomainError::Conflict(msg) => Self::Conflict(msg),
DomainError::Validation(msg) => Self::Validation(msg), DomainError::Validation(msg) => Self::Validation(msg),
DomainError::Repository(msg) => { DomainError::Repository(msg) => {

View File

@@ -1,4 +1,6 @@
use api_types::auth::{AuthResponse, LoginRequest, RegisterRequest, UserResponse}; use api_types::auth::{
AuthResponse, LoginRequest, LogoutRequest, RefreshRequest, RegisterRequest, UserResponse,
};
use utoipa::OpenApi; use utoipa::OpenApi;
#[derive(OpenApi)] #[derive(OpenApi)]
@@ -6,8 +8,17 @@ use utoipa::OpenApi;
paths( paths(
crate::routes::auth::login_handler, crate::routes::auth::login_handler,
crate::routes::auth::register_handler, crate::routes::auth::register_handler,
crate::routes::auth::refresh_handler,
crate::routes::auth::logout_handler,
crate::routes::auth::me_handler, crate::routes::auth::me_handler,
), ),
components(schemas(LoginRequest, RegisterRequest, AuthResponse, UserResponse)) components(schemas(
LoginRequest,
RegisterRequest,
RefreshRequest,
LogoutRequest,
AuthResponse,
UserResponse
))
)] )]
pub struct AuthDoc; pub struct AuthDoc;

View File

@@ -5,11 +5,14 @@ use axum::{
routing::{get, post}, routing::{get, post},
}; };
use api_types::auth::{AuthResponse, LoginRequest, RegisterRequest, UserResponse}; use api_types::auth::{
AuthResponse, LoginRequest, LogoutRequest, RefreshRequest, RegisterRequest, UserResponse,
};
use application::auth::{ use application::auth::{
commands::{LoginCommand, RegisterCommand}, commands::{LoginCommand, RegisterCommand},
login, register, login, logout, refresh, register,
}; };
use domain::auth::RefreshSession;
use crate::{ use crate::{
error::{ApiError, ApiResult}, error::{ApiError, ApiResult},
@@ -22,6 +25,8 @@ pub fn router() -> Router<PresentationState> {
Router::new() Router::new()
.route("/login", post(login_handler)) .route("/login", post(login_handler))
.route("/register", post(register_handler)) .route("/register", post(register_handler))
.route("/refresh", post(refresh_handler))
.route("/logout", post(logout_handler))
.route("/me", get(me_handler)) .route("/me", get(me_handler))
} }
@@ -52,9 +57,20 @@ pub async fn login_handler(
.create_token(&user) .create_token(&user)
.map_err(|e| ApiError::internal(format!("jwt error: {e}")))?; .map_err(|e| ApiError::internal(format!("jwt error: {e}")))?;
let session = RefreshSession::new(user.id, state.ctx.config.refresh_token_ttl_seconds);
let refresh_token = session.token().to_string();
state
.ctx
.repos
.refresh_session
.create(&session)
.await
.map_err(ApiError::from)?;
Ok(Json(AuthResponse { Ok(Json(AuthResponse {
user: user_response(user), user: user_response(user),
access_token: token, access_token: token,
refresh_token,
})) }))
} }
@@ -90,15 +106,72 @@ pub async fn register_handler(
.create_token(&user) .create_token(&user)
.map_err(|e| ApiError::internal(format!("jwt error: {e}")))?; .map_err(|e| ApiError::internal(format!("jwt error: {e}")))?;
let session = RefreshSession::new(user.id, state.ctx.config.refresh_token_ttl_seconds);
let refresh_token = session.token().to_string();
state
.ctx
.repos
.refresh_session
.create(&session)
.await
.map_err(ApiError::from)?;
Ok(( Ok((
StatusCode::CREATED, StatusCode::CREATED,
Json(AuthResponse { Json(AuthResponse {
user: user_response(user), user: user_response(user),
access_token: token, access_token: token,
refresh_token,
}), }),
)) ))
} }
#[utoipa::path(
post, path = "/api/v1/auth/refresh",
request_body = RefreshRequest,
responses(
(status = 200, body = AuthResponse),
(status = 401, body = api_types::errors::ErrorResponse, description = "Invalid or expired refresh token"),
)
)]
pub async fn refresh_handler(
State(state): State<PresentationState>,
Json(payload): Json<RefreshRequest>,
) -> ApiResult<Json<AuthResponse>> {
let result = refresh::execute(&state.ctx, &payload.refresh_token)
.await
.map_err(ApiError::from)?;
let token = state
.jwt_validator
.create_token(&result.user)
.map_err(|e| ApiError::internal(format!("jwt error: {e}")))?;
Ok(Json(AuthResponse {
user: user_response(result.user),
access_token: token,
refresh_token: result.refresh_token,
}))
}
#[utoipa::path(
post, path = "/api/v1/auth/logout",
request_body = LogoutRequest,
responses(
(status = 204, description = "Logged out"),
)
)]
pub async fn logout_handler(
State(state): State<PresentationState>,
Json(payload): Json<LogoutRequest>,
) -> ApiResult<StatusCode> {
logout::execute(&state.ctx, &payload.refresh_token)
.await
.map_err(ApiError::from)?;
Ok(StatusCode::NO_CONTENT)
}
#[utoipa::path( #[utoipa::path(
get, path = "/api/v1/auth/me", get, path = "/api/v1/auth/me",
responses( responses(

View File

@@ -56,6 +56,10 @@ pub struct WiringConfig {
/// `VectorStore` (for querying related notes), not `EmbeddingGenerator`. /// `VectorStore` (for querying related notes), not `EmbeddingGenerator`.
/// Default: `false`. /// Default: `false`.
pub enable_embeddings: bool, pub enable_embeddings: bool,
/// `REFRESH_TOKEN_TTL_SECONDS` — refresh token time-to-live in seconds.
/// Default: `2592000` (30 days).
pub refresh_token_ttl_seconds: i64,
} }
impl WiringConfig { impl WiringConfig {
@@ -79,6 +83,7 @@ impl WiringConfig {
enable_embeddings: optional_env("ENABLE_EMBEDDINGS") enable_embeddings: optional_env("ENABLE_EMBEDDINGS")
.map(|s| s == "true" || s == "1") .map(|s| s == "true" || s == "1")
.unwrap_or(false), .unwrap_or(false),
refresh_token_ttl_seconds: parse_env("REFRESH_TOKEN_TTL_SECONDS", 2_592_000i64)?,
}) })
} }
@@ -90,6 +95,7 @@ impl WiringConfig {
min_similarity: self.smart_min_similarity, min_similarity: self.smart_min_similarity,
}, },
allow_registration: self.allow_registration, allow_registration: self.allow_registration,
refresh_token_ttl_seconds: self.refresh_token_ttl_seconds,
} }
} }

View File

@@ -20,6 +20,7 @@ use sqlite::{
db::{connect, run_migrations}, db::{connect, run_migrations},
link::SqliteLinkRepository, link::SqliteLinkRepository,
note::SqliteNoteRepository, note::SqliteNoteRepository,
refresh_session::SqliteRefreshSessionRepository,
tag::SqliteTagRepository, tag::SqliteTagRepository,
user::SqliteUserRepository, user::SqliteUserRepository,
}; };
@@ -43,6 +44,7 @@ pub async fn build_context(cfg: &WiringConfig) -> anyhow::Result<AppContext> {
tag: Arc::new(SqliteTagRepository::new(pool.clone())), tag: Arc::new(SqliteTagRepository::new(pool.clone())),
user: Arc::new(SqliteUserRepository::new(pool.clone())), user: Arc::new(SqliteUserRepository::new(pool.clone())),
link: Arc::new(SqliteLinkRepository::new(pool.clone())), link: Arc::new(SqliteLinkRepository::new(pool.clone())),
refresh_session: Arc::new(SqliteRefreshSessionRepository::new(pool.clone())),
}; };
// ── Auth ────────────────────────────────────────────────────────────────── // ── Auth ──────────────────────────────────────────────────────────────────

View File

@@ -1,5 +1,5 @@
import { useQuery, useMutation, useQueryClient } from "@tanstack/react-query"; import { useQuery, useMutation, useQueryClient } from "@tanstack/react-query";
import { api, setAuthToken, clearAuthToken } from "@/lib/api"; import { api, setAuthToken, setRefreshToken, clearAllTokens, getRefreshToken } from "@/lib/api";
import { useNavigate } from "react-router-dom"; import { useNavigate } from "react-router-dom";
export interface User { export interface User {
@@ -11,6 +11,7 @@ export interface User {
export interface AuthResponse { export interface AuthResponse {
user: User; user: User;
access_token: string; access_token: string;
refresh_token: string;
} }
async function fetchUser(): Promise<User | null> { async function fetchUser(): Promise<User | null> {
@@ -40,6 +41,7 @@ export function useLogin() {
api.post("/auth/login", credentials), api.post("/auth/login", credentials),
onSuccess: (result: AuthResponse) => { onSuccess: (result: AuthResponse) => {
setAuthToken(result.access_token); setAuthToken(result.access_token);
setRefreshToken(result.refresh_token);
queryClient.invalidateQueries({ queryKey: ["user"] }); queryClient.invalidateQueries({ queryKey: ["user"] });
navigate("/"); navigate("/");
}, },
@@ -55,6 +57,7 @@ export function useRegister() {
api.post("/auth/register", credentials), api.post("/auth/register", credentials),
onSuccess: (result: AuthResponse) => { onSuccess: (result: AuthResponse) => {
setAuthToken(result.access_token); setAuthToken(result.access_token);
setRefreshToken(result.refresh_token);
queryClient.invalidateQueries({ queryKey: ["user"] }); queryClient.invalidateQueries({ queryKey: ["user"] });
navigate("/"); navigate("/");
}, },
@@ -66,10 +69,14 @@ export function useLogout() {
const navigate = useNavigate(); const navigate = useNavigate();
return useMutation({ return useMutation({
// JWT logout is client-side: discard the token. mutationFn: async () => {
mutationFn: () => Promise.resolve(), const refreshToken = getRefreshToken();
onSuccess: () => { if (refreshToken) {
clearAuthToken(); await api.post("/auth/logout", { refresh_token: refreshToken }).catch(() => {});
}
},
onSettled: () => {
clearAllTokens();
queryClient.setQueryData(["user"], null); queryClient.setQueryData(["user"], null);
navigate("/login"); navigate("/login");
}, },

View File

@@ -6,19 +6,36 @@ declare global {
} }
} }
const TOKEN_STORAGE_KEY = 'k_notes_auth_token'; const ACCESS_TOKEN_KEY = 'k_notes_auth_token';
const REFRESH_TOKEN_KEY = 'k_notes_refresh_token';
// JWT Token management
export function setAuthToken(token: string): void { export function setAuthToken(token: string): void {
localStorage.setItem(TOKEN_STORAGE_KEY, token); localStorage.setItem(ACCESS_TOKEN_KEY, token);
} }
export function getAuthToken(): string | null { export function getAuthToken(): string | null {
return localStorage.getItem(TOKEN_STORAGE_KEY); return localStorage.getItem(ACCESS_TOKEN_KEY);
} }
export function clearAuthToken(): void { export function clearAuthToken(): void {
localStorage.removeItem(TOKEN_STORAGE_KEY); localStorage.removeItem(ACCESS_TOKEN_KEY);
}
export function setRefreshToken(token: string): void {
localStorage.setItem(REFRESH_TOKEN_KEY, token);
}
export function getRefreshToken(): string | null {
return localStorage.getItem(REFRESH_TOKEN_KEY);
}
export function clearRefreshToken(): void {
localStorage.removeItem(REFRESH_TOKEN_KEY);
}
export function clearAllTokens(): void {
clearAuthToken();
clearRefreshToken();
} }
const getApiUrl = () => { const getApiUrl = () => {
@@ -53,6 +70,31 @@ export class ApiError extends Error {
} }
} }
let refreshPromise: Promise<boolean> | null = null;
async function tryRefreshToken(): Promise<boolean> {
const refreshToken = getRefreshToken();
if (!refreshToken) return false;
try {
const url = `${getApiUrl()}/auth/refresh`;
const response = await fetch(url, {
method: "POST",
headers: { "Content-Type": "application/json" },
body: JSON.stringify({ refresh_token: refreshToken }),
});
if (!response.ok) return false;
const data = await response.json();
setAuthToken(data.access_token);
setRefreshToken(data.refresh_token);
return true;
} catch {
return false;
}
}
async function fetchWithAuth(endpoint: string, options: RequestInit = {}) { async function fetchWithAuth(endpoint: string, options: RequestInit = {}) {
const url = `${getApiUrl()}${endpoint}`; const url = `${getApiUrl()}${endpoint}`;
const token = getAuthToken(); const token = getAuthToken();
@@ -81,6 +123,47 @@ async function fetchWithAuth(endpoint: string, options: RequestInit = {}) {
const response = (await Promise.race([fetchPromise, timeoutPromise])) as Response; const response = (await Promise.race([fetchPromise, timeoutPromise])) as Response;
if (response.status === 401) {
if (!refreshPromise) {
refreshPromise = tryRefreshToken().finally(() => {
refreshPromise = null;
});
}
const refreshed = await refreshPromise;
if (refreshed) {
// Retry with the new token
const newToken = getAuthToken();
if (newToken) {
headers["Authorization"] = `Bearer ${newToken}`;
}
const retryResponse = await fetch(url, { ...options, headers, credentials: "include" });
if (!retryResponse.ok) {
let errorMessage = "An error occurred";
try {
const errorData = await retryResponse.json();
errorMessage = errorData.error?.message || errorData.message || errorMessage;
} catch {
// failed to parse json
}
throw new ApiError(retryResponse.status, errorMessage);
}
if (retryResponse.status === 204) return null;
try {
return await retryResponse.json();
} catch {
return null;
}
}
// Refresh failed — clear tokens and redirect
clearAllTokens();
window.location.href = "/login";
throw new ApiError(401, "Session expired");
}
if (!response.ok) { if (!response.ok) {
// Try to parse error message // Try to parse error message
let errorMessage = "An error occurred"; let errorMessage = "An error occurred";
@@ -141,4 +224,3 @@ export const api = {
}, },
importData: (data: any) => api.post("/import", data), importData: (data: any) => api.post("/import", data),
}; };