use chrono::Utc;
use std::sync::Arc;
use crate::config::AuthConfig;
use crate::entity::{AuthSession, AuthUser};
use crate::error::AuthResult;
use crate::schema::AuthSchema;
use crate::store::AuthStore;
use crate::types::CreateSession;
pub struct SessionManager<S: AuthSchema> {
config: Arc<AuthConfig>,
database: Arc<dyn AuthStore<S>>,
}
impl<S: AuthSchema> Clone for SessionManager<S> {
fn clone(&self) -> Self {
Self {
config: self.config.clone(),
database: self.database.clone(),
}
}
}
impl<S: AuthSchema> SessionManager<S> {
pub fn new(config: Arc<AuthConfig>, database: Arc<dyn AuthStore<S>>) -> Self {
Self { config, database }
}
pub async fn create_session(
&self,
user: &impl AuthUser,
ip_address: Option<String>,
user_agent: Option<String>,
) -> AuthResult<S::Session> {
let expires_at = Utc::now() + self.config.session.expires_in;
let create_session = CreateSession {
user_id: user.id().to_string(),
expires_at,
ip_address,
user_agent,
impersonated_by: None,
active_organization_id: None,
};
let session = self.database.create_session(create_session).await?;
Ok(session)
}
pub async fn get_session(&self, token: &str) -> AuthResult<Option<S::Session>> {
let mut session = self.database.get_session(token).await?;
let should_refresh = if let Some(ref s) = session {
let now = Utc::now();
if s.expires_at() < now || !s.active() {
if let Err(err) = self.database.delete_session(token).await {
tracing::warn!(
error = %err,
"Failed to delete expired session; will be retried later"
);
}
return Ok(None);
}
if !self.config.session.disable_session_refresh {
match self.config.session.update_age {
Some(age) => {
let updated = s.updated_at();
Utc::now().signed_duration_since(updated) >= age
}
None => true,
}
} else {
false
}
} else {
false
};
if should_refresh {
let new_expires_at = Utc::now() + self.config.session.expires_in;
match self
.database
.update_session_expiry(token, new_expires_at)
.await
{
Ok(()) => {
match self.database.get_session(token).await {
Ok(Some(refreshed)) => session = Some(refreshed),
Ok(None) => {
tracing::warn!(
"Session re-read after refresh returned None (concurrent revoke?); returning pre-refresh value"
);
}
Err(err) => {
tracing::warn!(
error = %err,
"Session re-read after refresh failed; returning pre-refresh value"
);
}
}
}
Err(err) => {
tracing::warn!(
error = %err,
"Failed to refresh session expiry; returning pre-refresh session"
);
}
}
}
Ok(session)
}
pub async fn delete_session(&self, token: &str) -> AuthResult<()> {
self.database.delete_session(token).await?;
Ok(())
}
pub async fn delete_user_sessions(&self, user_id: impl AsRef<str>) -> AuthResult<()> {
self.database.delete_user_sessions(user_id.as_ref()).await?;
Ok(())
}
pub async fn list_user_sessions(
&self,
user_id: impl AsRef<str>,
) -> AuthResult<Vec<S::Session>> {
let sessions = self.database.get_user_sessions(user_id.as_ref()).await?;
let now = Utc::now();
let active_sessions = sessions
.into_iter()
.filter(|session| session.expires_at() > now && session.active())
.collect();
Ok(active_sessions)
}
pub async fn revoke_session(&self, token: &str) -> AuthResult<bool> {
let session_exists = self.get_session(token).await?.is_some();
if session_exists {
self.delete_session(token).await?;
Ok(true)
} else {
Ok(false)
}
}
pub async fn revoke_all_user_sessions(&self, user_id: impl AsRef<str>) -> AuthResult<usize> {
let user_id = user_id.as_ref();
let sessions = self.list_user_sessions(user_id).await?;
let count = sessions.len();
self.delete_user_sessions(user_id).await?;
Ok(count)
}
pub async fn revoke_other_user_sessions(
&self,
user_id: impl AsRef<str>,
current_token: &str,
) -> AuthResult<usize> {
let sessions = self.list_user_sessions(user_id).await?;
let mut count = 0;
for session in sessions {
if session.token() != current_token {
self.delete_session(session.token()).await?;
count += 1;
}
}
Ok(count)
}
pub async fn cleanup_expired_sessions(&self) -> AuthResult<usize> {
let count = self.database.delete_expired_sessions().await?;
Ok(count)
}
pub fn is_session_fresh(&self, session: &impl AuthSession) -> bool {
match self.config.session.fresh_age {
Some(fresh_age) => session.created_at() + fresh_age > Utc::now(),
None => false,
}
}
pub fn validate_token_format(&self, token: &str) -> bool {
token.starts_with("session_") && token.len() > 40
}
pub fn extract_session_token(&self, req: &crate::types::AuthRequest) -> Option<String> {
if let Some(auth_header) = req.headers.get("authorization")
&& let Some(token) = auth_header.strip_prefix("Bearer ")
{
return Some(token.to_string());
}
if let Some(cookie_header) = req.headers.get("cookie") {
let cookie_name = &self.config.session.cookie_name;
for c in cookie::Cookie::split_parse(cookie_header).flatten() {
if c.name() == cookie_name && !c.value().is_empty() {
return Some(c.value().to_string());
}
}
}
None
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::entity::AuthSession;
use crate::test_store::{BundledSchema, test_config, test_database};
use crate::types::AuthRequest;
use crate::types::HttpMethod;
use crate::wire::SessionView;
use chrono::Duration;
fn test_manager() -> SessionManager<BundledSchema> {
let runtime = tokio::runtime::Runtime::new().expect("runtime should build");
SessionManager::new(test_config(), runtime.block_on(test_database()))
}
#[test]
fn valid_token_format() {
let mgr = test_manager();
let token = "session_abcdefghijklmnopqrstuvwxyz1234567890";
assert!(mgr.validate_token_format(token));
}
#[test]
fn invalid_token_no_prefix() {
let mgr = test_manager();
assert!(!mgr.validate_token_format("abcdefghijklmnopqrstuvwxyz1234567890"));
}
#[test]
fn invalid_token_too_short() {
let mgr = test_manager();
assert!(!mgr.validate_token_format("session_short"));
}
#[test]
fn extract_from_bearer() {
let mgr = test_manager();
let mut req = AuthRequest::new(HttpMethod::Get, "/test");
let _ = req
.headers
.insert("authorization".into(), "Bearer my-token".into());
assert_eq!(mgr.extract_session_token(&req), Some("my-token".into()));
}
#[test]
fn extract_from_cookie() {
let mgr = test_manager();
let mut req = AuthRequest::new(HttpMethod::Get, "/test");
let _ = req.headers.insert(
"cookie".into(),
"better-auth.session_token=tok123; other=val".into(),
);
assert_eq!(mgr.extract_session_token(&req), Some("tok123".into()));
}
#[test]
fn extract_bearer_takes_precedence_over_cookie() {
let mgr = test_manager();
let mut req = AuthRequest::new(HttpMethod::Get, "/test");
let _ = req
.headers
.insert("authorization".into(), "Bearer bearer-tok".into());
let _ = req.headers.insert(
"cookie".into(),
"better-auth.session_token=cookie-tok".into(),
);
assert_eq!(mgr.extract_session_token(&req), Some("bearer-tok".into()));
}
#[test]
fn extract_returns_none_without_auth() {
let mgr = test_manager();
let req = AuthRequest::new(HttpMethod::Get, "/test");
assert_eq!(mgr.extract_session_token(&req), None);
}
#[test]
fn extract_skips_empty_cookie_value() {
let mgr = test_manager();
let mut req = AuthRequest::new(HttpMethod::Get, "/test");
let _ = req
.headers
.insert("cookie".into(), "better-auth.session_token=".into());
assert_eq!(mgr.extract_session_token(&req), None);
}
#[test]
fn session_fresh_when_within_window() {
let mut config = AuthConfig::new("test-secret-min-32-chars-1234567");
config.session.fresh_age = Some(Duration::minutes(10));
let runtime = tokio::runtime::Runtime::new().expect("runtime should build");
let mgr = SessionManager::new(Arc::new(config), runtime.block_on(test_database()));
let session = SessionView {
id: "s1".into(),
expires_at: Utc::now() + Duration::hours(1),
token: "tok".into(),
created_at: Utc::now(),
updated_at: Utc::now(),
ip_address: None,
user_agent: None,
user_id: "u1".into(),
impersonated_by: None,
active_organization_id: None,
active: true,
};
assert!(mgr.is_session_fresh(&session));
}
#[test]
fn session_not_fresh_when_old() {
let mut config = AuthConfig::new("test-secret-min-32-chars-1234567");
config.session.fresh_age = Some(Duration::minutes(10));
let runtime = tokio::runtime::Runtime::new().expect("runtime should build");
let mgr = SessionManager::new(Arc::new(config), runtime.block_on(test_database()));
let session = SessionView {
id: "s1".into(),
expires_at: Utc::now() + Duration::hours(1),
token: "tok".into(),
created_at: Utc::now() - Duration::minutes(20),
updated_at: Utc::now(),
ip_address: None,
user_agent: None,
user_id: "u1".into(),
impersonated_by: None,
active_organization_id: None,
active: true,
};
assert!(!mgr.is_session_fresh(&session));
}
#[test]
fn session_never_fresh_when_no_fresh_age() {
let mgr = test_manager(); let session = SessionView {
id: "s1".into(),
expires_at: Utc::now() + Duration::hours(1),
token: "tok".into(),
created_at: Utc::now(),
updated_at: Utc::now(),
ip_address: None,
user_agent: None,
user_id: "u1".into(),
impersonated_by: None,
active_organization_id: None,
active: true,
};
assert!(!mgr.is_session_fresh(&session));
}
#[tokio::test]
async fn create_and_get_session() {
let db = test_database().await;
let mgr = SessionManager::new(test_config(), db.clone());
let user = db
.create_user(crate::types::CreateUser::new().with_email("test@test.com"))
.await
.unwrap();
let session = mgr.create_session(&user, None, None).await.unwrap();
let token = session.token().to_string();
let retrieved = mgr.get_session(&token).await.unwrap();
assert!(retrieved.is_some());
}
#[tokio::test]
async fn refresh_returns_the_persisted_expiry() {
let db = test_database().await;
let mut config = AuthConfig::new("test-secret-min-32-chars-1234567");
config.session.update_age = None;
let mgr = SessionManager::new(Arc::new(config), db.clone());
let user = db
.create_user(crate::types::CreateUser::new().with_email("refresh@test.com"))
.await
.unwrap();
let session = mgr.create_session(&user, None, None).await.unwrap();
let token = session.token().to_string();
let stale = session.expires_at() - Duration::minutes(30);
db.update_session_expiry(&token, stale).await.unwrap();
let returned = mgr
.get_session(&token)
.await
.unwrap()
.expect("session should still be live");
let stored = db
.get_session(&token)
.await
.unwrap()
.expect("session should still be stored");
assert!(
returned.expires_at() > stale,
"refresh should have extended the expiry"
);
assert_eq!(
returned.expires_at(),
stored.expires_at(),
"returned session must reflect the persisted expiry, not the pre-refresh value"
);
}
#[tokio::test]
async fn create_session_without_metadata_uses_empty_strings() {
let db = test_database().await;
let mgr = SessionManager::new(test_config(), db.clone());
let user = db
.create_user(crate::types::CreateUser::new().with_email("test@test.com"))
.await
.unwrap();
let session = mgr.create_session(&user, None, None).await.unwrap();
assert_eq!(session.ip_address.as_deref(), Some(""));
assert_eq!(session.user_agent.as_deref(), Some(""));
}
#[tokio::test]
async fn delete_session_removes_it() {
let db = test_database().await;
let mgr = SessionManager::new(test_config(), db.clone());
let user = db
.create_user(crate::types::CreateUser::new().with_email("test@test.com"))
.await
.unwrap();
let session = mgr.create_session(&user, None, None).await.unwrap();
let token = session.token().to_string();
mgr.delete_session(&token).await.unwrap();
let retrieved = mgr.get_session(&token).await.unwrap();
assert!(retrieved.is_none());
}
#[tokio::test]
async fn revoke_session_returns_true_when_found() {
let db = test_database().await;
let mgr = SessionManager::new(test_config(), db.clone());
let user = db
.create_user(crate::types::CreateUser::new().with_email("test@test.com"))
.await
.unwrap();
let session = mgr.create_session(&user, None, None).await.unwrap();
let result = mgr.revoke_session(session.token()).await.unwrap();
assert!(result);
}
#[tokio::test]
async fn revoke_session_returns_false_when_not_found() {
let mgr = SessionManager::new(test_config(), test_database().await);
let result = mgr.revoke_session("nonexistent-token").await.unwrap();
assert!(!result);
}
#[tokio::test]
async fn list_user_sessions_excludes_expired() {
let db = test_database().await;
let mgr = SessionManager::new(test_config(), db.clone());
let user = db
.create_user(crate::types::CreateUser::new().with_email("test@test.com"))
.await
.unwrap();
let _ = mgr.create_session(&user, None, None).await.unwrap();
let _ = mgr.create_session(&user, None, None).await.unwrap();
let sessions = mgr.list_user_sessions(user.id()).await.unwrap();
assert_eq!(sessions.len(), 2);
}
#[tokio::test]
async fn revoke_all_user_sessions() {
let db = test_database().await;
let mgr = SessionManager::new(test_config(), db.clone());
let user = db
.create_user(crate::types::CreateUser::new().with_email("test@test.com"))
.await
.unwrap();
let _ = mgr.create_session(&user, None, None).await.unwrap();
let _ = mgr.create_session(&user, None, None).await.unwrap();
let count = mgr.revoke_all_user_sessions(user.id()).await.unwrap();
assert_eq!(count, 2);
let sessions = mgr.list_user_sessions(user.id()).await.unwrap();
assert!(sessions.is_empty());
}
#[tokio::test]
async fn revoke_other_sessions_keeps_current() {
let db = test_database().await;
let mgr = SessionManager::new(test_config(), db.clone());
let user = db
.create_user(crate::types::CreateUser::new().with_email("test@test.com"))
.await
.unwrap();
let current = mgr.create_session(&user, None, None).await.unwrap();
let _ = mgr.create_session(&user, None, None).await.unwrap();
let _ = mgr.create_session(&user, None, None).await.unwrap();
let count = mgr
.revoke_other_user_sessions(user.id(), current.token())
.await
.unwrap();
assert_eq!(count, 2);
let remaining = mgr.list_user_sessions(user.id()).await.unwrap();
assert_eq!(remaining.len(), 1);
assert_eq!(remaining[0].token(), current.token());
}
}