use std::collections::HashMap;
use std::collections::VecDeque;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::{Arc, Mutex, RwLock, Weak};
use std::time::Instant;
use tokio::sync::{broadcast, watch};
use crate::terminal::Terminal;
pub const DEFAULT_ORPHAN_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(60);
pub type AttachResult = (
Vec<ScrollbackEvent>,
broadcast::Receiver<Vec<u8>>,
watch::Receiver<(u16, u16)>,
);
#[derive(Clone, Debug, PartialEq)]
pub enum ScrollbackEvent {
Output(Vec<u8>),
WindowSize(u16, u16),
}
impl ScrollbackEvent {
fn byte_cost(&self) -> usize {
match self {
Self::Output(data) => data.len(),
Self::WindowSize(_, _) => 4,
}
}
}
pub struct Session {
id: String,
pub terminal: Terminal,
scrollback: Mutex<VecDeque<ScrollbackEvent>>,
scrollback_bytes: Mutex<usize>,
scrollback_limit: usize,
clients: AtomicUsize,
detached_at: Mutex<Option<Instant>>,
window_size: watch::Sender<(u16, u16)>,
orphan_timeout: std::time::Duration,
}
impl Session {
pub fn new(
terminal: Terminal,
output_rx: broadcast::Receiver<Vec<u8>>,
scrollback_limit: usize,
orphan_timeout: std::time::Duration,
) -> Arc<Self> {
let id = uuid::Uuid::new_v4().to_string();
let (ws_tx, _) = watch::channel((24, 80));
let session = Arc::new(Self {
id,
terminal,
scrollback: Mutex::new(VecDeque::new()),
scrollback_bytes: Mutex::new(0),
scrollback_limit,
clients: AtomicUsize::new(0),
detached_at: Mutex::new(None),
window_size: ws_tx,
orphan_timeout,
});
let weak: Weak<Session> = Arc::downgrade(&session);
let mut rx = output_rx;
tokio::spawn(async move {
loop {
match rx.recv().await {
Ok(data) => {
let Some(s) = weak.upgrade() else {
break;
};
s.push_scrollback(ScrollbackEvent::Output(data));
}
Err(broadcast::error::RecvError::Lagged(_)) => {
continue;
}
Err(broadcast::error::RecvError::Closed) => break,
}
}
});
session
}
pub fn id(&self) -> &str {
&self.id
}
fn push_scrollback(&self, event: ScrollbackEvent) {
let cost = event.byte_cost();
let mut sb = self.scrollback.lock().unwrap();
let mut bytes = self.scrollback_bytes.lock().unwrap();
*bytes += cost;
sb.push_back(event);
while *bytes > self.scrollback_limit {
if let Some(old) = sb.pop_front() {
*bytes -= old.byte_cost();
} else {
break;
}
}
}
pub fn attach(&self) -> AttachResult {
self.clients.fetch_add(1, Ordering::Relaxed);
*self.detached_at.lock().unwrap() = None;
let sb = self.scrollback.lock().unwrap();
let rx = self.terminal.subscribe();
let ws_rx = self.window_size.subscribe();
let events: Vec<ScrollbackEvent> = sb.iter().cloned().collect();
(events, rx, ws_rx)
}
pub fn set_window_size(&self, rows: u16, cols: u16) {
let _ = self.window_size.send((rows, cols));
self.push_scrollback(ScrollbackEvent::WindowSize(rows, cols));
}
pub fn detach(&self) {
if self.clients.fetch_sub(1, Ordering::Relaxed) == 1 {
*self.detached_at.lock().unwrap() = Some(Instant::now());
}
}
pub fn client_count(&self) -> usize {
self.clients.load(Ordering::Relaxed)
}
fn is_orphaned(&self) -> bool {
self.clients.load(Ordering::Relaxed) == 0
&& self
.detached_at
.lock()
.unwrap()
.is_some_and(|t| t.elapsed() >= self.orphan_timeout)
}
}
pub struct SessionStore {
sessions: RwLock<HashMap<String, Arc<Session>>>,
}
impl SessionStore {
pub fn new() -> Arc<Self> {
Arc::new(Self {
sessions: RwLock::new(HashMap::new()),
})
}
pub fn insert(self: &Arc<Self>, session: Arc<Session>) {
let sid = session.id().to_owned();
self.sessions
.write()
.unwrap()
.insert(sid.clone(), session.clone());
let store = Arc::downgrade(self);
let closed_rx = session.terminal.closed();
tokio::spawn(async move {
loop {
tokio::time::sleep(std::time::Duration::from_secs(1)).await;
let Some(store) = store.upgrade() else { return };
let should_remove = {
let sessions = store.sessions.read().unwrap();
match sessions.get(&sid) {
Some(s) => {
s.is_orphaned()
|| (*closed_rx.borrow() && s.clients.load(Ordering::Relaxed) == 0)
}
None => return,
}
};
if should_remove {
store.sessions.write().unwrap().remove(&sid);
tracing::info!("removed session {sid}");
return;
}
}
});
}
pub fn get(&self, id: &str) -> Option<Arc<Session>> {
self.sessions.read().unwrap().get(id).cloned()
}
pub fn is_empty(&self) -> bool {
self.sessions.read().unwrap().is_empty()
}
}
#[cfg(test)]
mod tests {
use super::*;
const TEST_SCROLLBACK_LIMIT: usize = 256 * 1024;
fn spawn_session() -> Arc<Session> {
let (terminal, output_rx) = Terminal::spawn("/bin/sh", None).expect("spawn /bin/sh");
Session::new(
terminal,
output_rx,
TEST_SCROLLBACK_LIMIT,
DEFAULT_ORPHAN_TIMEOUT,
)
}
#[tokio::test]
async fn test_attach_detach_clients() {
let session = spawn_session();
let (_sb1, _rx1, _ws1) = session.attach();
assert_eq!(session.clients.load(Ordering::Relaxed), 1);
let (_sb2, _rx2, _ws2) = session.attach();
assert_eq!(session.clients.load(Ordering::Relaxed), 2);
session.detach();
assert_eq!(session.clients.load(Ordering::Relaxed), 1);
}
#[tokio::test]
async fn test_not_orphaned_with_clients() {
let session = spawn_session();
let (_sb, _rx, _ws) = session.attach();
assert!(!session.is_orphaned());
}
#[tokio::test]
async fn test_not_orphaned_immediately_after_detach() {
let session = spawn_session();
let (_sb, _rx, _ws) = session.attach();
session.detach();
assert!(!session.is_orphaned());
}
#[tokio::test]
async fn test_orphaned_after_timeout() {
let session = spawn_session();
let (_sb, _rx, _ws) = session.attach();
session.detach();
*session.detached_at.lock().unwrap() =
Some(Instant::now() - session.orphan_timeout - std::time::Duration::from_secs(1));
assert!(session.is_orphaned());
}
#[tokio::test]
async fn test_scrollback_captures_output() {
let session = spawn_session();
session
.terminal
.write(b"echo scrollback_test_marker\n".to_vec())
.await
.unwrap();
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
let (events, _rx, _ws) = session.attach();
let has_marker = events.iter().any(|e| match e {
ScrollbackEvent::Output(data) => {
String::from_utf8_lossy(data).contains("scrollback_test_marker")
}
_ => false,
});
assert!(has_marker, "scrollback should contain Output with marker");
}
#[tokio::test]
async fn test_session_store_insert_and_get() {
let store = SessionStore::new();
let session = spawn_session();
let id = session.id().to_owned();
store.insert(session);
assert!(store.get(&id).is_some());
assert!(store.get("nonexistent").is_none());
}
#[tokio::test]
async fn test_scrollback_eviction_removes_whole_events() {
let (terminal, output_rx) = Terminal::spawn("/bin/sh", None).expect("spawn");
let session = Session::new(terminal, output_rx, 10, DEFAULT_ORPHAN_TIMEOUT);
session.push_scrollback(ScrollbackEvent::Output(b"aaaaa".to_vec())); session.push_scrollback(ScrollbackEvent::Output(b"bbbbb".to_vec())); session.push_scrollback(ScrollbackEvent::Output(b"ccc".to_vec()));
let sb = session.scrollback.lock().unwrap();
let bytes = *session.scrollback_bytes.lock().unwrap();
assert!(bytes <= 10, "bytes {bytes} should be within limit");
assert!(
sb.iter().all(|e| matches!(e, ScrollbackEvent::Output(_))),
"all events should be Output"
);
assert_ne!(
sb.front(),
Some(&ScrollbackEvent::Output(b"aaaaa".to_vec())),
"oldest event should have been evicted"
);
}
#[tokio::test]
async fn test_set_window_size_records_event() {
let (terminal, output_rx) = Terminal::spawn("/bin/sh", None).expect("spawn");
let session = Session::new(
terminal,
output_rx,
TEST_SCROLLBACK_LIMIT,
DEFAULT_ORPHAN_TIMEOUT,
);
session.set_window_size(40, 120);
let sb = session.scrollback.lock().unwrap();
let has_ws = sb
.iter()
.any(|e| matches!(e, ScrollbackEvent::WindowSize(40, 120)));
assert!(has_ws, "scrollback should contain WindowSize(40, 120)");
}
}