use std::collections::HashMap;
use std::sync::Arc;
use std::time::{Duration, Instant};
use parking_lot::Mutex;
use chat_engine_sdk::models::{LifecycleState, TenantId, UserId};
use chat_engine_sdk::plugin::{PluginCallContext, SessionPluginCtx};
use futures::stream::{self, BoxStream, StreamExt};
use serde_json::Value as JsonValue;
use time::OffsetDateTime;
use tokio::sync::mpsc;
use tokio_util::sync::CancellationToken;
use toolkit_macros::domain_model;
use tracing::{debug, info, instrument, warn};
use uuid::Uuid;
use crate::domain::error::{ChatEngineError, Result};
use crate::domain::message::{
StreamingChunkEvent, StreamingCompleteEvent, StreamingErrorEvent, StreamingEvent,
StreamingStartEvent,
};
use crate::domain::ports::MessageRepo;
use crate::domain::ports::SessionRepo;
use crate::domain::ports::SessionTypeRepo;
use crate::domain::retention::RetentionPolicy;
use crate::domain::service::plugin_service::PluginService;
use crate::domain::service::session_service::Identity;
use crate::domain::session::{Session, get_retention_policy, set_retention_policy};
pub const DEFAULT_SUMMARY_DEADLINE: Duration = Duration::from_mins(2);
pub const DEFAULT_SUMMARY_BUFFER_SIZE: usize = 64;
pub type SummaryStream = BoxStream<'static, StreamingEvent>;
#[domain_model]
#[derive(Debug, Clone)]
pub struct SessionCleanupOutcome {
pub session_id: Uuid,
pub policy_type: &'static str,
pub messages_deleted: u64,
pub duration_ms: u64,
pub skipped_locked: bool,
}
#[domain_model]
#[derive(Debug, Clone, Default)]
pub struct RetentionCleanupReport {
pub sessions: Vec<SessionCleanupOutcome>,
}
impl RetentionCleanupReport {
#[must_use]
pub fn total_messages_deleted(&self) -> u64 {
self.sessions.iter().map(|s| s.messages_deleted).sum()
}
#[must_use]
pub fn skipped_count(&self) -> usize {
self.sessions.iter().filter(|s| s.skipped_locked).count()
}
}
#[domain_model]
#[derive(Debug, Clone)]
pub struct ValidatedPolicy(RetentionPolicy);
impl From<ValidatedPolicy> for RetentionPolicy {
fn from(v: ValidatedPolicy) -> Self {
v.0
}
}
#[domain_model]
#[derive(Clone)]
pub struct IntelligenceService {
sessions: Arc<dyn SessionRepo>,
session_types: Arc<dyn SessionTypeRepo>,
messages: Arc<dyn MessageRepo>,
plugins: PluginService,
summary_buffer_size: usize,
summary_deadline: Duration,
retention_max_sessions_per_tick: u32,
retention_max_deletes_per_session: u32,
retention_cursor: Arc<Mutex<HashMap<String, Uuid>>>,
}
pub const DEFAULT_RETENTION_MAX_SESSIONS_PER_TICK: u32 = 1000;
pub const DEFAULT_RETENTION_MAX_DELETES_PER_SESSION: u32 = 1000;
impl IntelligenceService {
#[must_use]
pub fn new(
sessions: Arc<dyn SessionRepo>,
session_types: Arc<dyn SessionTypeRepo>,
messages: Arc<dyn MessageRepo>,
plugins: PluginService,
) -> Self {
Self {
sessions,
session_types,
messages,
plugins,
summary_buffer_size: DEFAULT_SUMMARY_BUFFER_SIZE,
summary_deadline: DEFAULT_SUMMARY_DEADLINE,
retention_max_sessions_per_tick: DEFAULT_RETENTION_MAX_SESSIONS_PER_TICK,
retention_max_deletes_per_session: DEFAULT_RETENTION_MAX_DELETES_PER_SESSION,
retention_cursor: Arc::new(Mutex::new(HashMap::new())),
}
}
#[must_use]
pub fn with_buffer_size(mut self, size: usize) -> Self {
self.summary_buffer_size = size.max(1);
self
}
#[must_use]
pub fn with_summary_deadline(mut self, deadline: Duration) -> Self {
self.summary_deadline = deadline;
self
}
#[must_use]
pub fn with_retention_caps(
mut self,
max_sessions_per_tick: u32,
max_deletes_per_session: u32,
) -> Self {
self.retention_max_sessions_per_tick = max_sessions_per_tick.max(1);
self.retention_max_deletes_per_session = max_deletes_per_session.max(1);
self
}
#[instrument(skip(self), fields(session_id = %session_id))]
pub async fn get_effective_retention_policy(
&self,
identity: &Identity,
session_id: Uuid,
) -> Result<RetentionPolicy> {
let session = self.load_session(identity, session_id).await?;
Ok(resolve_effective_policy(&session))
}
#[instrument(skip(self), fields(session_id = %session_id))]
pub async fn update_session_retention_policy(
&self,
identity: &Identity,
session_id: Uuid,
policy: RetentionPolicy,
) -> Result<RetentionPolicy> {
let validated = validate_retention_policy(policy)?;
let mut session = self.load_session(identity, session_id).await?;
if matches!(
session.lifecycle_state,
LifecycleState::SoftDeleted | LifecycleState::HardDeleted
) {
return Err(ChatEngineError::conflict(format!(
"session is {} and cannot accept retention_policy updates",
session.lifecycle_state
)));
}
let persisted_policy = validated.0.clone();
set_retention_policy(&mut session, validated.0);
let new_metadata = session.metadata.clone();
let _persisted = self
.sessions
.update_metadata(
&identity.tenant_id,
&identity.user_id,
session_id,
new_metadata,
)
.await?;
info!(
session_id = %session_id,
policy_type = %retention_policy_label(&persisted_policy),
"persisted per-session retention policy"
);
Ok(persisted_policy)
}
#[instrument(
skip(self, identity, cancel),
fields(
session_id = %session_id,
user_id = %identity.user_id,
request_id,
summary_message_id,
),
)]
pub async fn summarize_session(
&self,
identity: &Identity,
session_id: Uuid,
cancel: CancellationToken,
) -> Result<SummaryStream> {
let session = self.load_session(identity, session_id).await?;
if !matches!(session.lifecycle_state, LifecycleState::Active) {
return Err(ChatEngineError::conflict(format!(
"session is {} and cannot be summarized",
session.lifecycle_state
)));
}
let session_type_id =
session
.session_type_id
.ok_or_else(|| ChatEngineError::BadRequest {
reason: "session has no session_type bound; summary cannot be generated"
.to_string(),
})?;
let session_type = self
.session_types
.find_by_id(session_type_id)
.await?
.ok_or_else(|| ChatEngineError::not_found("session_type", session_type_id))?;
let plugin_instance_id = session_type.plugin_instance_id.ok_or_else(|| {
ChatEngineError::BackendUnavailable {
reason: "session_type has no plugin_instance_id; summarization unsupported"
.to_string(),
retry_after: None,
source: None,
}
})?;
let plugin = self
.plugins
.resolve(&plugin_instance_id)
.map_err(|err| match err {
ChatEngineError::NotFound { .. } => ChatEngineError::BackendUnavailable {
reason: format!(
"plugin '{plugin_instance_id}' is not registered; summarization unsupported"
),
retry_after: None,
source: None,
},
other => other,
})?;
let plugin_config = self
.plugins
.load_config(&plugin_instance_id, session_type_id)
.await?;
let history = self.messages.fetch_active_history(session_id, None).await?;
let request_id = Uuid::new_v4();
tracing::Span::current().record("request_id", tracing::field::display(request_id));
let plugin_cancel = cancel.child_token();
let deadline = Instant::now() + self.summary_deadline;
let call_ctx = PluginCallContext {
request_id,
tenant_id: TenantId::new(identity.tenant_id.as_str()),
user_id: UserId::new(identity.user_id.as_str()),
plugin_instance_id: plugin_instance_id.clone(),
session_type_id,
plugin_config,
enabled_capabilities: session
.enabled_capabilities
.as_ref()
.and_then(|v| serde_json::from_value(v.clone()).ok()),
deadline: Some(deadline),
cancel: plugin_cancel.clone(),
};
let plugin_ctx = SessionPluginCtx {
session_type_id,
session_id: Some(session_id),
call_ctx,
};
info!(
session_id = %session_id,
history_len = history.len(),
plugin_instance_id = %plugin_instance_id,
"invoking on_session_summary"
);
let plugin_stream = match plugin.on_session_summary(plugin_ctx).await {
Ok(s) => s,
Err(err) => {
warn!(
session_id = %session_id,
error = %err,
"on_session_summary returned pre-stream failure"
);
return Err(err.into());
}
};
let summary_message_id = Uuid::new_v4();
tracing::Span::current().record(
"summary_message_id",
tracing::field::display(summary_message_id),
);
let stream = self.spawn_summary_driver(
session_id,
summary_message_id,
plugin_stream,
cancel,
plugin_cancel,
deadline,
identity.tenant_id.clone(),
);
Ok(stream)
}
#[instrument(skip(self), fields(tenant_id = %tenant_id))]
pub async fn run_retention_cleanup_for_tenant(
&self,
tenant_id: &str,
) -> Result<RetentionCleanupReport> {
let cap = self.retention_max_sessions_per_tick;
let after = self.retention_cursor.lock().get(tenant_id).copied();
let mut active = self
.sessions
.list_active_sessions_for_tenant(tenant_id, after, cap)
.await?;
if active.is_empty() && after.is_some() {
active = self
.sessions
.list_active_sessions_for_tenant(tenant_id, None, cap)
.await?;
}
let batch_len = active.len();
let last_id = active.last().map(|r| r.session_id);
let mut outcomes: Vec<SessionCleanupOutcome> = Vec::with_capacity(active.len());
for row in active {
let session: Session = row;
let policy = resolve_effective_policy(&session);
let label = retention_policy_label(&policy);
if matches!(policy, RetentionPolicy::None) {
outcomes.push(SessionCleanupOutcome {
session_id: session.session_id,
policy_type: label,
messages_deleted: 0,
duration_ms: 0,
skipped_locked: false,
});
continue;
}
let start = Instant::now();
let lock_acquired = true;
if !lock_acquired {
outcomes.push(SessionCleanupOutcome {
session_id: session.session_id,
policy_type: label,
messages_deleted: 0,
duration_ms: start.elapsed().as_millis() as u64,
skipped_locked: true,
});
continue;
}
let eligible = self
.evaluate_retention_policy(session.session_id, &policy)
.await?;
let mut removed: u64 = 0;
for id in eligible {
let n = self
.messages
.delete_message_subtree(session.session_id, id)
.await?;
removed += n;
}
let duration_ms = start.elapsed().as_millis() as u64;
info!(
session_id = %session.session_id,
messages_deleted = removed,
policy_type = label,
duration_ms = duration_ms,
"retention cleanup completed for session"
);
outcomes.push(SessionCleanupOutcome {
session_id: session.session_id,
policy_type: label,
messages_deleted: removed,
duration_ms,
skipped_locked: false,
});
}
match last_id {
Some(id) if batch_len == cap as usize => {
self.retention_cursor
.lock()
.insert(tenant_id.to_owned(), id);
debug!(
tenant_id,
cap,
next_after = %id,
"retention sweep filled a full batch; more sessions deferred to next tick",
);
}
_ => {
self.retention_cursor.lock().remove(tenant_id);
}
}
outcomes.sort_by_key(|o| o.session_id);
Ok(RetentionCleanupReport { sessions: outcomes })
}
#[instrument(skip(self))]
pub async fn run_retention_cleanup_all_tenants(&self) -> Result<RetentionCleanupReport> {
let tenants = self.sessions.list_tenants_with_active_sessions().await?;
let mut aggregated: Vec<SessionCleanupOutcome> = Vec::new();
for tenant_id in tenants {
match self.run_retention_cleanup_for_tenant(&tenant_id).await {
Ok(report) => aggregated.extend(report.sessions),
Err(err) => warn!(
%tenant_id,
error = %err,
"retention cleanup failed for tenant; continuing with next tenant",
),
}
}
Ok(RetentionCleanupReport {
sessions: aggregated,
})
}
pub(crate) async fn evaluate_retention_policy(
&self,
session_id: Uuid,
policy: &RetentionPolicy,
) -> Result<Vec<Uuid>> {
match policy {
RetentionPolicy::None => Ok(Vec::new()),
RetentionPolicy::AgeBased { max_age_days } => {
let cutoff = OffsetDateTime::now_utc()
- Duration::from_secs(u64::from(*max_age_days) * 86_400);
self.messages
.list_non_root_message_ids_older_than(
session_id,
cutoff,
self.retention_max_deletes_per_session,
)
.await
}
RetentionPolicy::CountBased { max_message_count } => {
let max = u64::from(*max_message_count);
let total = self.messages.count_non_root_messages(session_id).await?;
if total <= max {
return Ok(Vec::new());
}
let surplus = total - max;
let limit = surplus
.min(u64::from(self.retention_max_deletes_per_session))
.try_into()
.unwrap_or(self.retention_max_deletes_per_session);
self.messages
.list_oldest_non_root_message_ids(session_id, limit)
.await
}
}
}
async fn load_session(&self, identity: &Identity, session_id: Uuid) -> Result<Session> {
let row = self
.sessions
.find_by_id(&identity.tenant_id, &identity.user_id, session_id)
.await?
.ok_or_else(|| ChatEngineError::not_found("session", session_id))?;
Ok(row)
}
fn spawn_summary_driver(
&self,
session_id: Uuid,
summary_message_id: Uuid,
mut plugin_stream: chat_engine_sdk::plugin::PluginStream,
cancel: CancellationToken,
plugin_cancel: CancellationToken,
deadline: Instant,
tenant_id: String,
) -> SummaryStream {
let (tx, rx) = mpsc::channel::<StreamingEvent>(self.summary_buffer_size);
let messages = Arc::clone(&self.messages);
let plugin_cancel_for_deadline = plugin_cancel.clone();
tokio::spawn(async move {
tokio::time::sleep_until(tokio::time::Instant::from_std(deadline)).await;
plugin_cancel_for_deadline.cancel();
});
let tx_for_driver = tx.clone();
tokio::spawn(async move {
let start = StreamingEvent::Start(StreamingStartEvent {
message_id: summary_message_id,
});
if tx_for_driver.send(start).await.is_err() {
cancel.cancel();
return;
}
let mut accumulator = String::new();
let mut last_metadata: Option<JsonValue> = None;
let mut summarized_ids: Vec<Uuid> = Vec::new();
let mut completed = false;
let mut errored: Option<String> = None;
loop {
tokio::select! {
biased;
_ = cancel.cancelled() => {
plugin_cancel.cancel();
break;
}
next = plugin_stream.next() => {
let Some(item) = next else {
break;
};
match item {
Ok(StreamingEvent::Start(_)) => {
}
Ok(StreamingEvent::Chunk(c)) => {
accumulator.push_str(&c.chunk);
let evt = StreamingEvent::Chunk(StreamingChunkEvent {
message_id: summary_message_id,
chunk: c.chunk,
});
if tx_for_driver.send(evt).await.is_err() {
plugin_cancel.cancel();
break;
}
}
Ok(StreamingEvent::Complete(c)) => {
if let Some(ref meta) = c.metadata {
summarized_ids = extract_summarized_ids(meta);
}
last_metadata = c.metadata;
completed = true;
break;
}
Ok(StreamingEvent::Error(e)) => {
let evt = StreamingEvent::Error(StreamingErrorEvent {
message_id: summary_message_id,
error: e.error.clone(),
});
tx_for_driver.send(evt).await.ok();
errored = Some(e.error);
break;
}
Ok(_) => {}
Err(err) => {
let s = err.to_string();
let evt = StreamingEvent::Error(StreamingErrorEvent {
message_id: summary_message_id,
error: s.clone(),
});
tx_for_driver.send(evt).await.ok();
errored = Some(s);
break;
}
}
}
}
}
if completed {
match messages
.insert_summary_message(
session_id,
accumulator,
last_metadata.clone(),
summarized_ids,
Some(tenant_id.clone()),
)
.await
{
Ok(_) => {
let evt = StreamingEvent::Complete(StreamingCompleteEvent {
message_id: summary_message_id,
metadata: last_metadata,
file_citations: vec![],
link_citations: vec![],
references: vec![],
});
tx_for_driver.send(evt).await.ok();
}
Err(err) => {
warn!(
session_id = %session_id,
summary_message_id = %summary_message_id,
error = %err,
"failed to persist session summary after stream complete",
);
let evt = StreamingEvent::Error(StreamingErrorEvent {
message_id: summary_message_id,
error: format!("failed to persist session summary: {err}"),
});
tx_for_driver.send(evt).await.ok();
}
}
} else if let Some(err) = errored {
warn!(
session_id = %session_id,
summary_message_id = %summary_message_id,
error = %err,
"summary stream errored mid-flight; no summary persisted"
);
}
});
stream::unfold(
rx,
|mut rx| async move { rx.recv().await.map(|evt| (evt, rx)) },
)
.boxed()
}
}
pub fn validate_retention_policy(policy: RetentionPolicy) -> Result<ValidatedPolicy> {
match &policy {
RetentionPolicy::None => {}
RetentionPolicy::AgeBased { max_age_days } => {
if *max_age_days < 1 {
return Err(ChatEngineError::bad_request(
"max_age_days required and must be >= 1",
));
}
}
RetentionPolicy::CountBased { max_message_count } => {
if *max_message_count < 1 {
return Err(ChatEngineError::bad_request(
"max_message_count required and must be >= 1",
));
}
}
}
Ok(ValidatedPolicy(policy))
}
#[must_use]
pub fn resolve_effective_policy(session: &Session) -> RetentionPolicy {
get_retention_policy(session).unwrap_or(RetentionPolicy::None)
}
#[must_use]
pub fn retention_policy_label(p: &RetentionPolicy) -> &'static str {
match p {
RetentionPolicy::None => "none",
RetentionPolicy::AgeBased { .. } => "age_based",
RetentionPolicy::CountBased { .. } => "count_based",
}
}
fn extract_summarized_ids(meta: &JsonValue) -> Vec<Uuid> {
let Some(arr) = meta
.get("summarized_message_ids")
.and_then(|v| v.as_array())
else {
return Vec::new();
};
arr.iter()
.filter_map(|v| v.as_str().and_then(|s| Uuid::parse_str(s).ok()))
.collect()
}
#[cfg(test)]
#[path = "intelligence_service_tests.rs"]
mod intelligence_service_tests;