use std::collections::HashMap;
use std::sync::Arc;
use anyhow::{Context as _, Result};
use async_trait::async_trait;
use parking_lot::Mutex;
use theway_core::{Session, SessionTreeEntry};
use theway_transport::session_observability::{
ListSessionMessagesRequest, SessionMessagePage, SessionObservabilityOps,
};
use theway_transport::transport::SessionOps;
use theway_transport::wire::{WireSessionSnapshot, WireStatus};
use crate::feed_replay::session_tree_entry_wire_blocks;
use crate::runtime_storage::SessionRepository;
pub(crate) struct DaemonSessionObservability {
session_ops: Arc<dyn SessionOps>,
session_states: Arc<Mutex<HashMap<String, WireStatus>>>,
latest: Arc<Mutex<WireStatus>>,
repo: Arc<dyn SessionRepository>,
}
impl DaemonSessionObservability {
pub(crate) fn new(
session_ops: Arc<dyn SessionOps>,
session_states: Arc<Mutex<HashMap<String, WireStatus>>>,
latest: Arc<Mutex<WireStatus>>,
repo: Arc<dyn SessionRepository>,
) -> Self {
Self {
session_ops,
session_states,
latest,
repo,
}
}
fn live_status(&self, session_id: &str) -> Option<WireStatus> {
self.session_states
.lock()
.get(session_id)
.cloned()
.or_else(|| {
let latest = self.latest.lock();
(latest.session_id == session_id).then(|| latest.clone())
})
}
}
#[async_trait]
impl SessionObservabilityOps for DaemonSessionObservability {
async fn authoritative_snapshot(&self, session_id: &str) -> Result<WireSessionSnapshot> {
let resource = self
.session_ops
.session_snapshot(session_id)
.await
.with_context(|| format!("load resource snapshot for session {session_id}"));
let Some(live) = self.live_status(session_id) else {
let mut resource = resource?;
if let Ok(page) = self
.list_session_messages(&ListSessionMessagesRequest {
session_id: session_id.to_string(),
before_entry_id: None,
limit: u32::MAX,
})
.await
{
resource.feed.blocks = page.blocks;
}
return Ok(resource);
};
let Ok(mut resource) = resource else {
return Ok(WireSessionSnapshot::from(&live));
};
let mut merged = WireSessionSnapshot::from(&live);
merged.session_id = if resource.session_id.is_empty() {
merged.session_id
} else {
resource.session_id.clone()
};
merged.info = resource.info;
merged.info.sidebar = live.sidebar.clone();
merged.graph_state.nodes = std::mem::take(&mut resource.graph_state.nodes);
merged.graph_state.active_node_id = resource.graph_state.active_node_id.take();
merged.lineage = resource.lineage;
Ok(merged)
}
async fn list_session_messages(
&self,
request: &ListSessionMessagesRequest,
) -> Result<SessionMessagePage> {
let limit = request.effective_limit() as usize;
let store = self
.repo
.open(&request.session_id)
.await?
.with_context(|| format!("no session matches id {}", request.session_id))?;
let session = Session::from_store(store);
let branch = session.branch(None).await?;
let messages: Vec<&SessionTreeEntry> = branch
.iter()
.filter(|entry| matches!(entry, SessionTreeEntry::Message { .. }))
.collect();
let total = messages.len() as u64;
let before = match request.before_entry_id.as_deref() {
Some(id) => match messages.iter().position(|entry| entry.id() == id) {
Some(position) => position,
None => {
return Ok(SessionMessagePage {
session_id: request.session_id.clone(),
blocks: Vec::new(),
next_before_entry_id: None,
has_more: false,
total,
});
}
},
None => messages.len(),
};
let start = before.saturating_sub(limit);
let page = &messages[start..before];
let mut blocks = Vec::new();
for entry in page {
if let Some(entry_blocks) = session_tree_entry_wire_blocks(entry) {
blocks.extend(entry_blocks);
}
}
Ok(SessionMessagePage {
session_id: request.session_id.clone(),
blocks,
next_before_entry_id: page.first().map(|entry| entry.id().to_string()),
has_more: start > 0,
total,
})
}
}
#[cfg(test)]
tests_bridge_macro::tests_bridge!("session_observability");