use std::collections::HashMap;
use std::sync::{Arc, RwLock};
use std::time::Instant;
use tracing;
use crate::error::{PepError, Result};
use crate::oidc_client::{OidcClient, TokenResponse};
use crate::token_provider::{compute_expires_at, seconds_until_expiry};
use crate::token_store::StoredToken;
#[derive(Clone, Debug)]
struct SessionEntry {
token: StoredToken,
last_accessed: Instant,
}
#[derive(Debug, Default)]
pub struct InMemoryTokenStore {
sessions: RwLock<HashMap<String, StoredToken>>,
}
impl InMemoryTokenStore {
pub fn new() -> Self {
Self::default()
}
}
impl crate::token_store::TokenStore for InMemoryTokenStore {
fn load(&self, name: &str) -> Result<Option<StoredToken>> {
let sessions = self.sessions.read().unwrap();
Ok(sessions.get(name).cloned())
}
fn save(&self, name: &str, token: &StoredToken) -> Result<()> {
let mut sessions = self.sessions.write().unwrap();
sessions.insert(name.to_string(), token.clone());
Ok(())
}
fn delete(&self, name: &str) -> Result<()> {
let mut sessions = self.sessions.write().unwrap();
sessions.remove(name);
Ok(())
}
}
pub struct WebSessionManager {
sessions: Arc<RwLock<HashMap<String, SessionEntry>>>,
oidc_client: OidcClient,
issuer_url: String,
client_id: String,
client_secret: Option<String>,
scope: String,
refresh_buffer_secs: u64,
idle_timeout_secs: u64,
sweep_threshold: usize,
}
impl std::fmt::Debug for WebSessionManager {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("WebSessionManager")
.field("issuer_url", &self.issuer_url)
.field("client_id", &self.client_id)
.field("scope", &self.scope)
.field("idle_timeout_secs", &self.idle_timeout_secs)
.finish()
}
}
const DEFAULT_REFRESH_BUFFER_SECS: u64 = 60;
const DEFAULT_IDLE_TIMEOUT_SECS: u64 = 3600; const DEFAULT_SWEEP_THRESHOLD: usize = 64;
impl WebSessionManager {
pub fn new(
oidc_client: OidcClient,
issuer_url: String,
client_id: String,
client_secret: Option<String>,
scope: String,
) -> Self {
Self {
sessions: Arc::new(RwLock::new(HashMap::new())),
oidc_client,
issuer_url,
client_id,
client_secret,
scope,
refresh_buffer_secs: DEFAULT_REFRESH_BUFFER_SECS,
idle_timeout_secs: DEFAULT_IDLE_TIMEOUT_SECS,
sweep_threshold: DEFAULT_SWEEP_THRESHOLD,
}
}
pub fn with_refresh_buffer(mut self, secs: u64) -> Self {
self.refresh_buffer_secs = secs;
self
}
pub fn with_idle_timeout(mut self, secs: u64) -> Self {
self.idle_timeout_secs = secs;
self
}
pub fn with_sweep_threshold(mut self, threshold: usize) -> Self {
self.sweep_threshold = threshold;
self
}
pub async fn create_session(&self, token_response: &TokenResponse) -> Result<String> {
let session_id = uuid::Uuid::new_v4().to_string();
let expires_at = compute_expires_at(token_response.expires_in);
let stored = StoredToken::new(
&token_response.access_token,
token_response.refresh_token.clone(),
&token_response.token_type,
&expires_at,
token_response.scope.clone(),
);
let entry = SessionEntry {
token: stored,
last_accessed: Instant::now(),
};
{
let mut sessions = self.sessions.write().unwrap();
sessions.insert(session_id.clone(), entry);
}
tracing::debug!(
session_id = %session_id,
expires_at = %expires_at,
"Created web session"
);
Ok(session_id)
}
pub async fn get_token(&self, session_id: &str) -> Result<String> {
let stored = {
let mut sessions = self.sessions.write().unwrap();
if sessions.len() > self.sweep_threshold {
self.sweep_idle_sessions(&mut sessions);
}
let entry = match sessions.get_mut(session_id) {
Some(e) => e,
None => {
tracing::debug!(session_id = %session_id, "Session not found");
return Err(PepError::AuthenticationRequired);
}
};
let idle_secs = entry.last_accessed.elapsed().as_secs();
if idle_secs > self.idle_timeout_secs {
tracing::debug!(
session_id = %session_id,
idle_secs = idle_secs,
idle_timeout = self.idle_timeout_secs,
"Session idle-expired"
);
sessions.remove(session_id);
return Err(PepError::AuthenticationRequired);
}
entry.last_accessed = Instant::now();
entry.token.clone()
};
let remaining = seconds_until_expiry(&stored.expires_at);
if remaining > self.refresh_buffer_secs {
return Ok(stored.access_token);
}
tracing::debug!(
session_id = %session_id,
remaining_secs = remaining,
"Token near expiry, attempting refresh"
);
self.refresh_session(session_id, &stored).await
}
pub fn destroy_session(&self, session_id: &str) -> Result<()> {
let mut sessions = self.sessions.write().unwrap();
sessions.remove(session_id);
tracing::debug!(session_id = %session_id, "Session destroyed");
Ok(())
}
pub fn session_count(&self) -> usize {
self.sessions.read().unwrap().len()
}
async fn refresh_session(
&self,
session_id: &str,
stored: &StoredToken,
) -> Result<String> {
let refresh_token = match &stored.refresh_token {
Some(rt) => rt.clone(),
None => {
tracing::debug!(session_id = %session_id, "No refresh token, cannot refresh");
let mut sessions = self.sessions.write().unwrap();
sessions.remove(session_id);
return Err(PepError::AuthenticationRequired);
}
};
let response = self
.oidc_client
.refresh_access_token(
&self.issuer_url,
&self.client_id,
self.client_secret.as_deref(),
&refresh_token,
Some(&self.scope),
)
.await;
match response {
Ok(token_response) => {
let new_expires_at = compute_expires_at(token_response.expires_in);
let updated = StoredToken::new(
&token_response.access_token,
token_response
.refresh_token
.clone()
.or(Some(refresh_token)),
&token_response.token_type,
&new_expires_at,
token_response.scope.clone().or(stored.scope.clone()),
);
let access_token = updated.access_token.clone();
let mut sessions = self.sessions.write().unwrap();
if let Some(entry) = sessions.get_mut(session_id) {
entry.token = updated;
}
tracing::info!(
session_id = %session_id,
expires_at = %new_expires_at,
"Token refreshed successfully"
);
Ok(access_token)
}
Err(PepError::TokenRefreshFailed { status, detail }) => {
tracing::warn!(
session_id = %session_id,
status = status,
"Token refresh failed: {}", detail
);
let mut sessions = self.sessions.write().unwrap();
sessions.remove(session_id);
Err(PepError::AuthenticationRequired)
}
Err(e) => {
tracing::warn!(
session_id = %session_id,
"Token refresh error: {}", e
);
Err(e)
}
}
}
fn sweep_idle_sessions(&self, sessions: &mut HashMap<String, SessionEntry>) {
let before = sessions.len();
let timeout = std::time::Duration::from_secs(self.idle_timeout_secs);
sessions.retain(|_, entry| entry.last_accessed.elapsed() < timeout);
let swept = before - sessions.len();
if swept > 0 {
tracing::info!(
swept = swept,
remaining = sessions.len(),
"Idle session sweep complete"
);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::token_store::TokenStore;
#[test]
fn test_in_memory_store_round_trip() {
let store = InMemoryTokenStore::new();
let token = StoredToken::new(
"access123",
Some("refresh456".to_string()),
"Bearer",
"2099-01-01T00:00:00Z",
Some("openid profile".to_string()),
);
store.save("session1", &token).unwrap();
let loaded = store.load("session1").unwrap().expect("should exist");
assert_eq!(loaded.access_token, "access123");
assert_eq!(loaded.refresh_token.as_deref(), Some("refresh456"));
}
#[test]
fn test_in_memory_store_delete() {
let store = InMemoryTokenStore::new();
let token = StoredToken::new("a", None, "Bearer", "2099-01-01T00:00:00Z", None);
store.save("temp", &token).unwrap();
assert!(store.load("temp").unwrap().is_some());
store.delete("temp").unwrap();
assert!(store.load("temp").unwrap().is_none());
}
#[test]
fn test_in_memory_store_load_nonexistent() {
let store = InMemoryTokenStore::new();
assert!(store.load("ghost").unwrap().is_none());
}
#[test]
fn test_in_memory_store_overwrite() {
let store = InMemoryTokenStore::new();
let token1 = StoredToken::new("first", None, "Bearer", "2099-01-01T00:00:00Z", None);
store.save("key", &token1).unwrap();
let token2 = StoredToken::new("second", None, "Bearer", "2099-01-01T00:00:00Z", None);
store.save("key", &token2).unwrap();
let loaded = store.load("key").unwrap().unwrap();
assert_eq!(loaded.access_token, "second");
}
fn make_test_mgr() -> WebSessionManager {
WebSessionManager::new(
OidcClient::new(),
"https://idm.example.com/oauth2/openid/pdt-api".to_string(),
"pdt-api".to_string(),
None,
"openid profile".to_string(),
)
}
fn make_token_response(access: &str, refresh: Option<&str>) -> TokenResponse {
TokenResponse {
access_token: access.to_string(),
token_type: "Bearer".to_string(),
expires_in: Some(900),
refresh_token: refresh.map(|s| s.to_string()),
id_token: None,
scope: Some("openid profile".to_string()),
}
}
#[test]
fn test_create_session_stores_token() {
let mgr = make_test_mgr();
let rt = tokio::runtime::Runtime::new().unwrap();
let session_id = rt
.block_on(mgr.create_session(&make_token_response("access-abc", Some("refresh-xyz"))))
.unwrap();
assert!(!session_id.is_empty());
assert!(uuid::Uuid::parse_str(&session_id).is_ok());
}
#[test]
fn test_create_two_sessions_have_different_ids() {
let mgr = make_test_mgr();
let rt = tokio::runtime::Runtime::new().unwrap();
let id1 = rt.block_on(mgr.create_session(&make_token_response("a", None))).unwrap();
let id2 = rt.block_on(mgr.create_session(&make_token_response("b", None))).unwrap();
assert_ne!(id1, id2);
}
#[test]
fn test_get_token_valid_returns_access_token() {
let mgr = make_test_mgr();
let rt = tokio::runtime::Runtime::new().unwrap();
let session_id = rt
.block_on(mgr.create_session(&make_token_response("my-access-token", Some("my-refresh"))))
.unwrap();
let token = rt.block_on(mgr.get_token(&session_id)).unwrap();
assert_eq!(token, "my-access-token");
}
#[test]
fn test_get_token_unknown_session_errors() {
let mgr = make_test_mgr();
let rt = tokio::runtime::Runtime::new().unwrap();
let result = rt.block_on(mgr.get_token("nonexistent-session"));
assert!(matches!(result, Err(PepError::AuthenticationRequired)));
}
#[test]
fn test_destroy_session() {
let mgr = make_test_mgr();
let rt = tokio::runtime::Runtime::new().unwrap();
let session_id = rt
.block_on(mgr.create_session(&make_token_response("test-token", None)))
.unwrap();
let token = rt.block_on(mgr.get_token(&session_id)).unwrap();
assert_eq!(token, "test-token");
mgr.destroy_session(&session_id).unwrap();
let result = rt.block_on(mgr.get_token(&session_id));
assert!(matches!(result, Err(PepError::AuthenticationRequired)));
}
#[test]
fn test_get_token_expired_no_refresh_errors() {
let mgr = make_test_mgr();
let rt = tokio::runtime::Runtime::new().unwrap();
let stored = StoredToken::new(
"expired-access",
None,
"Bearer",
"2020-01-01T00:00:00Z",
None,
);
{
let mut sessions = mgr.sessions.write().unwrap();
sessions.insert(
"expired-session".to_string(),
SessionEntry {
token: stored,
last_accessed: Instant::now(),
},
);
}
let result = rt.block_on(mgr.get_token("expired-session"));
assert!(matches!(result, Err(PepError::AuthenticationRequired)));
}
#[test]
fn test_idle_timeout_evicts_session() {
let mgr = make_test_mgr().with_idle_timeout(0);
let rt = tokio::runtime::Runtime::new().unwrap();
let session_id = rt
.block_on(mgr.create_session(&make_token_response("will-expire", None)))
.unwrap();
std::thread::sleep(std::time::Duration::from_secs(1));
let result = rt.block_on(mgr.get_token(&session_id));
assert!(matches!(result, Err(PepError::AuthenticationRequired)));
}
#[test]
fn test_idle_timeout_does_not_evict_active_session() {
let mgr = make_test_mgr().with_idle_timeout(3600);
let rt = tokio::runtime::Runtime::new().unwrap();
let session_id = rt
.block_on(mgr.create_session(&make_token_response("active", None)))
.unwrap();
let token = rt.block_on(mgr.get_token(&session_id)).unwrap();
assert_eq!(token, "active");
let token = rt.block_on(mgr.get_token(&session_id)).unwrap();
assert_eq!(token, "active");
}
#[test]
fn test_sweep_removes_idle_sessions() {
let mgr = make_test_mgr()
.with_idle_timeout(0) .with_sweep_threshold(2);
let rt = tokio::runtime::Runtime::new().unwrap();
let id1 = rt.block_on(mgr.create_session(&make_token_response("a", None))).unwrap();
let id2 = rt.block_on(mgr.create_session(&make_token_response("b", None))).unwrap();
let id3 = rt.block_on(mgr.create_session(&make_token_response("c", None))).unwrap();
assert_eq!(mgr.session_count(), 3);
std::thread::sleep(std::time::Duration::from_secs(1));
let result = rt.block_on(mgr.get_token(&id1));
assert!(matches!(result, Err(PepError::AuthenticationRequired)));
assert_eq!(mgr.session_count(), 0);
let _ = (id2, id3);
}
#[test]
fn test_sweep_preserves_active_sessions() {
let mgr = make_test_mgr()
.with_idle_timeout(3600) .with_sweep_threshold(2);
let rt = tokio::runtime::Runtime::new().unwrap();
let id1 = rt.block_on(mgr.create_session(&make_token_response("a", None))).unwrap();
let _id2 = rt.block_on(mgr.create_session(&make_token_response("b", None))).unwrap();
let _id3 = rt.block_on(mgr.create_session(&make_token_response("c", None))).unwrap();
assert_eq!(mgr.session_count(), 3);
let token = rt.block_on(mgr.get_token(&id1)).unwrap();
assert_eq!(token, "a");
assert_eq!(mgr.session_count(), 3); }
#[test]
fn test_session_count() {
let mgr = make_test_mgr();
assert_eq!(mgr.session_count(), 0);
let rt = tokio::runtime::Runtime::new().unwrap();
let _ = rt.block_on(mgr.create_session(&make_token_response("a", None))).unwrap();
let _ = rt.block_on(mgr.create_session(&make_token_response("b", None))).unwrap();
assert_eq!(mgr.session_count(), 2);
mgr.destroy_session(&rt.block_on(mgr.create_session(&make_token_response("c", None))).unwrap()).unwrap();
assert_eq!(mgr.session_count(), 2);
}
#[test]
fn test_session_manager_debug() {
let mgr = make_test_mgr();
let debug_str = format!("{:?}", mgr);
assert!(debug_str.contains("WebSessionManager"));
assert!(debug_str.contains("pdt-api"));
assert!(debug_str.contains("idle_timeout_secs"));
}
#[test]
fn test_get_token_near_expiry_no_refresh() {
let mgr = make_test_mgr();
let stored = StoredToken::new(
"still-valid",
None,
"Bearer",
"2099-01-01T00:00:00Z",
None,
);
{
let mut sessions = mgr.sessions.write().unwrap();
sessions.insert(
"valid-session".to_string(),
SessionEntry {
token: stored,
last_accessed: Instant::now(),
},
);
}
let rt = tokio::runtime::Runtime::new().unwrap();
let token = rt.block_on(mgr.get_token("valid-session")).unwrap();
assert_eq!(token, "still-valid");
}
#[test]
fn test_with_idle_timeout_builder() {
let mgr = make_test_mgr().with_idle_timeout(7200);
assert_eq!(mgr.idle_timeout_secs, 7200);
}
#[test]
fn test_with_sweep_threshold_builder() {
let mgr = make_test_mgr().with_sweep_threshold(128);
assert_eq!(mgr.sweep_threshold, 128);
}
#[test]
fn test_with_refresh_buffer_builder() {
let mgr = make_test_mgr().with_refresh_buffer(120);
assert_eq!(mgr.refresh_buffer_secs, 120);
}
}