mod memory;
mod postgres;
mod store;
mod summarizer;
use std::sync::Arc;
use chrono::{Duration as ChronoDuration, Utc};
pub use memory::InMemorySessionStateStore;
pub use postgres::PostgresSessionStateStore;
pub use store::{MAX_SESSION_VALUE_BYTES, SUMMARY_KEY, SessionStateEntry};
pub use summarizer::{NoOpSummarizer, SummarizeFuture, Summarizer};
use uuid::Uuid;
use crate::error::{AuthError, Result};
#[cfg(test)]
mod tests;
pub enum SessionStateBackend {
InMemory(InMemorySessionStateStore),
Postgres(PostgresSessionStateStore),
}
impl std::fmt::Debug for SessionStateBackend {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::InMemory(_) => f.write_str("SessionStateBackend::InMemory"),
Self::Postgres(_) => f.write_str("SessionStateBackend::Postgres"),
}
}
}
impl SessionStateBackend {
async fn get(
&self,
session_id: Uuid,
thread_id: &str,
key: &str,
) -> Result<Option<SessionStateEntry>> {
match self {
Self::InMemory(s) => s.get(session_id, thread_id, key),
Self::Postgres(s) => s.get(session_id, thread_id, key).await,
}
}
async fn set(&self, entry: &SessionStateEntry) -> Result<()> {
match self {
Self::InMemory(s) => s.set(entry.clone()),
Self::Postgres(s) => s.set(entry).await,
}
}
async fn delete(&self, session_id: Uuid, thread_id: &str, key: &str) -> Result<()> {
match self {
Self::InMemory(s) => s.delete(session_id, thread_id, key),
Self::Postgres(s) => s.delete(session_id, thread_id, key).await,
}
}
async fn list_thread(
&self,
session_id: Uuid,
thread_id: &str,
) -> Result<Vec<SessionStateEntry>> {
match self {
Self::InMemory(s) => s.list_thread(session_id, thread_id),
Self::Postgres(s) => s.list_thread(session_id, thread_id).await,
}
}
async fn expire_thread(&self, session_id: Uuid, thread_id: &str) -> Result<()> {
match self {
Self::InMemory(s) => s.expire_thread(session_id, thread_id),
Self::Postgres(s) => s.expire_thread(session_id, thread_id).await,
}
}
async fn thread_count(
&self,
session_id: Uuid,
thread_id: &str,
exclude_key: &str,
) -> Result<usize> {
match self {
Self::InMemory(s) => s.thread_count(session_id, thread_id, exclude_key),
Self::Postgres(s) => s.thread_count(session_id, thread_id, exclude_key).await,
}
}
async fn replace_thread(&self, summary: &SessionStateEntry) -> Result<()> {
match self {
Self::InMemory(s) => s.replace_thread(summary.clone()),
Self::Postgres(s) => s.replace_thread(summary).await,
}
}
pub async fn evict_expired(&self) -> Result<u64> {
match self {
Self::InMemory(s) => s.evict_expired(),
Self::Postgres(s) => s.evict_expired().await,
}
}
}
pub struct SessionState {
backend: SessionStateBackend,
default_ttl: ChronoDuration,
summarize_after: Option<usize>,
summarizer: Option<Arc<dyn Summarizer>>,
}
impl std::fmt::Debug for SessionState {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("SessionState")
.field("backend", &self.backend)
.field("default_ttl", &self.default_ttl)
.field("summarize_after", &self.summarize_after)
.field("summarizer_installed", &self.summarizer.is_some())
.finish()
}
}
impl SessionState {
#[must_use]
pub fn new(backend: SessionStateBackend, default_ttl_secs: u64) -> Self {
let secs = i64::try_from(default_ttl_secs.max(1)).unwrap_or(i64::MAX);
Self {
backend,
default_ttl: ChronoDuration::seconds(secs),
summarize_after: None,
summarizer: None,
}
}
#[must_use]
pub fn with_summarizer(
mut self,
summarizer: Arc<dyn Summarizer>,
after_entries: usize,
) -> Self {
self.summarizer = Some(summarizer);
self.summarize_after = Some(after_entries.max(1));
self
}
pub async fn get(
&self,
session_id: Uuid,
thread_id: &str,
key: &str,
) -> Result<Option<SessionStateEntry>> {
self.backend.get(session_id, thread_id, key).await
}
pub async fn set(
&self,
session_id: Uuid,
thread_id: &str,
key: &str,
value: serde_json::Value,
) -> Result<()> {
if key == SUMMARY_KEY {
return Err(AuthError::SessionError {
message: format!(
"session-state key '{SUMMARY_KEY}' is reserved for the summarisation collapse"
),
});
}
let serialized_len = value.to_string().len();
if serialized_len > MAX_SESSION_VALUE_BYTES {
return Err(AuthError::SessionError {
message: format!(
"session-state value for '{key}' is {serialized_len} bytes serialized — the \
cap is {MAX_SESSION_VALUE_BYTES}. Store working state, not blobs."
),
});
}
let now = Utc::now();
self.backend
.set(&SessionStateEntry {
session_id,
thread_id: thread_id.to_string(),
key: key.to_string(),
value,
updated_at: now,
expires_at: now + self.default_ttl,
})
.await?;
self.maybe_summarize(session_id, thread_id).await
}
pub async fn delete(&self, session_id: Uuid, thread_id: &str, key: &str) -> Result<()> {
self.backend.delete(session_id, thread_id, key).await
}
pub async fn list_thread(
&self,
session_id: Uuid,
thread_id: &str,
) -> Result<Vec<SessionStateEntry>> {
self.backend.list_thread(session_id, thread_id).await
}
pub async fn expire_thread(&self, session_id: Uuid, thread_id: &str) -> Result<()> {
self.backend.expire_thread(session_id, thread_id).await
}
pub async fn evict_expired(&self) -> Result<u64> {
self.backend.evict_expired().await
}
async fn maybe_summarize(&self, session_id: Uuid, thread_id: &str) -> Result<()> {
let (Some(summarizer), Some(after)) = (self.summarizer.as_ref(), self.summarize_after)
else {
return Ok(());
};
if self.backend.thread_count(session_id, thread_id, SUMMARY_KEY).await? <= after {
return Ok(());
}
let entries = self.backend.list_thread(session_id, thread_id).await?;
let summary_value = summarizer.summarize(entries).await?;
let now = Utc::now();
self.backend
.replace_thread(&SessionStateEntry {
session_id,
thread_id: thread_id.to_string(),
key: SUMMARY_KEY.to_string(),
value: summary_value,
updated_at: now,
expires_at: now + self.default_ttl,
})
.await
}
}