use std::collections::HashMap;
use std::sync::Mutex;
use std::time::{Duration, Instant};
use chrono::Utc;
use crate::core::context::{get_current_task, get_dexcost_context_sync};
use crate::core::models::{Task, TaskStatus};
struct SessionEntry {
task: Task,
last_activity: Instant,
}
pub struct SessionManager {
sessions: Mutex<HashMap<u64, SessionEntry>>,
idle_timeout: Duration,
}
impl SessionManager {
pub fn new(idle_timeout: Duration) -> Self {
Self {
sessions: Mutex::new(HashMap::new()),
idle_timeout,
}
}
pub fn get_or_create_session(&self, call_type: &str) -> Task {
if let Some(existing) = get_current_task() {
let ctx_id = Self::context_id();
let mut sessions = self.sessions.lock().unwrap_or_else(|e| {
eprintln!("[dexcost] mutex poisoned, recovering: {}", e);
e.into_inner()
});
if let Some(entry) = sessions.get_mut(&ctx_id) {
entry.last_activity = Instant::now();
}
return existing;
}
let ctx_id = Self::context_id();
let mut sessions = self.sessions.lock().unwrap_or_else(|e| {
eprintln!("[dexcost] mutex poisoned, recovering: {}", e);
e.into_inner()
});
if let Some(entry) = sessions.get_mut(&ctx_id) {
entry.last_activity = Instant::now();
return entry.task.clone();
}
let task_type = if call_type.is_empty() {
"agent_session".to_string()
} else {
call_type.to_string()
};
let mut task = Task::new(&task_type);
task.metadata
.insert("session".to_string(), serde_json::Value::Bool(true));
task.metadata.insert(
"initiated_by".to_string(),
serde_json::Value::String(call_type.to_string()),
);
if let Some(ctx) = get_dexcost_context_sync() {
task.customer_id = ctx.customer_id;
task.project_id = ctx.project_id;
}
sessions.insert(
ctx_id,
SessionEntry {
task: task.clone(),
last_activity: Instant::now(),
},
);
task
}
pub fn finalize_idle_sessions(&self) -> Vec<Task> {
let now = Instant::now();
let mut finalized = Vec::new();
let mut sessions = self.sessions.lock().unwrap_or_else(|e| {
eprintln!("[dexcost] mutex poisoned, recovering: {}", e);
e.into_inner()
});
let idle_ids: Vec<u64> = sessions
.iter()
.filter(|(_, entry)| now.duration_since(entry.last_activity) >= self.idle_timeout)
.map(|(id, _)| *id)
.collect();
for ctx_id in idle_ids {
if let Some(mut entry) = sessions.remove(&ctx_id) {
entry.task.status = TaskStatus::Success;
entry.task.ended_at = Some(Utc::now());
finalized.push(entry.task);
}
}
finalized
}
pub fn active_session_count(&self) -> usize {
let sessions = self.sessions.lock().unwrap_or_else(|e| {
eprintln!("[dexcost] mutex poisoned, recovering: {}", e);
e.into_inner()
});
sessions.len()
}
pub fn clear(&self) {
let mut sessions = self.sessions.lock().unwrap_or_else(|e| {
eprintln!("[dexcost] mutex poisoned, recovering: {}", e);
e.into_inner()
});
sessions.clear();
}
fn context_id() -> u64 {
let id = std::thread::current().id();
let debug = format!("{:?}", id);
debug
.chars()
.filter(|c| c.is_ascii_digit())
.collect::<String>()
.parse::<u64>()
.unwrap_or(0)
}
}
static SESSION_MANAGER: std::sync::OnceLock<SessionManager> = std::sync::OnceLock::new();
pub fn get_session_manager() -> &'static SessionManager {
SESSION_MANAGER.get_or_init(|| SessionManager::new(Duration::from_secs(30)))
}
pub fn reset_session_manager() {
if let Some(mgr) = SESSION_MANAGER.get() {
mgr.clear();
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::time::Duration;
#[test]
fn test_new_session_created() {
let mgr = SessionManager::new(Duration::from_secs(30));
let task = mgr.get_or_create_session("llm_call");
assert_eq!(task.task_type, "llm_call");
assert_eq!(task.status, TaskStatus::Pending);
assert!(!task.task_id.is_empty());
assert_eq!(mgr.active_session_count(), 1);
}
#[test]
fn test_reuses_existing_session() {
let mgr = SessionManager::new(Duration::from_secs(30));
let first = mgr.get_or_create_session("llm_call");
let second = mgr.get_or_create_session("http_call");
assert_eq!(first.task_id, second.task_id);
assert_eq!(mgr.active_session_count(), 1);
}
#[test]
fn test_session_metadata() {
let mgr = SessionManager::new(Duration::from_secs(30));
let task = mgr.get_or_create_session("llm_call");
assert_eq!(
task.metadata.get("session"),
Some(&serde_json::Value::Bool(true))
);
assert_eq!(
task.metadata.get("initiated_by"),
Some(&serde_json::Value::String("llm_call".to_string()))
);
}
#[test]
fn test_finalize_idle_sessions() {
let mgr = SessionManager::new(Duration::from_millis(1));
let _task = mgr.get_or_create_session("llm_call");
assert_eq!(mgr.active_session_count(), 1);
std::thread::sleep(Duration::from_millis(10));
let finalized = mgr.finalize_idle_sessions();
assert_eq!(finalized.len(), 1);
assert_eq!(finalized[0].status, TaskStatus::Success);
assert!(finalized[0].ended_at.is_some());
assert_eq!(mgr.active_session_count(), 0);
}
#[test]
fn test_finalize_does_not_remove_active() {
let mgr = SessionManager::new(Duration::from_secs(300));
let _task = mgr.get_or_create_session("llm_call");
let finalized = mgr.finalize_idle_sessions();
assert!(finalized.is_empty());
assert_eq!(mgr.active_session_count(), 1);
}
#[test]
fn test_clear() {
let mgr = SessionManager::new(Duration::from_secs(30));
let _task = mgr.get_or_create_session("llm_call");
assert_eq!(mgr.active_session_count(), 1);
mgr.clear();
assert_eq!(mgr.active_session_count(), 0);
}
#[test]
fn test_empty_call_type_defaults_to_agent_session() {
let mgr = SessionManager::new(Duration::from_secs(30));
let task = mgr.get_or_create_session("");
assert_eq!(task.task_type, "agent_session");
}
#[test]
fn test_global_session_manager() {
reset_session_manager();
let mgr = get_session_manager();
assert_eq!(mgr.active_session_count(), 0);
}
}