use std::sync::Arc;
use std::time::{Duration, Instant};
use async_trait::async_trait;
use futures::stream::{self, StreamExt};
use serde_json::{Value as JsonValue, json};
use time::OffsetDateTime;
use tokio_util::sync::CancellationToken;
use toolkit_macros::domain_model;
use tracing::{debug, info, instrument, warn};
use uuid::Uuid;
use chat_engine_sdk::models::{
Capability, CapabilityValue, LifecycleState, MessagePartInput, TenantId, UserId, VariantInfo,
};
use chat_engine_sdk::plugin::{PluginCallContext, SessionPluginCtx};
use crate::domain::error::{ChatEngineError, Result};
use crate::domain::message::{Message, MessageRole, StreamingEvent};
use crate::domain::ports::SessionRepo;
use crate::domain::ports::{MessageRepo, SessionTypeRepo};
use crate::domain::service::message_service::{
MessageEventKind, MessageService, SendMessageStream,
};
use crate::domain::service::plugin_service::PluginService;
use crate::domain::service::session_service::{Identity, merge_plugin_metadata};
use crate::domain::session::{Session, SessionType};
pub const DEFAULT_SWITCH_TYPE_DEADLINE: Duration = Duration::from_secs(10);
#[domain_model]
#[derive(Debug, Clone)]
pub struct VariantListing {
pub variants: Vec<VariantEntry>,
pub current_index: Option<u32>,
}
#[domain_model]
#[derive(Debug, Clone)]
pub struct VariantEntry {
pub message: Message,
pub info: VariantInfo,
}
#[async_trait]
pub trait VariantRepo: Send + Sync {
async fn list_siblings(
&self,
session_id: Uuid,
parent_message_id: Option<Uuid>,
) -> Result<Vec<Message>>;
async fn insert_user_and_assistant_stub_for_branch(
&self,
session_id: Uuid,
parent_message_id: Uuid,
parts: Vec<MessagePartInput>,
file_ids: Option<Vec<Uuid>>,
tenant_id: Option<String>,
user_id: Option<String>,
) -> Result<(Uuid, i32, Uuid)>;
async fn ancestor_chain(&self, session_id: Uuid, message_id: Uuid) -> Result<Vec<Uuid>>;
async fn collect_descendants(&self, session_id: Uuid, message_id: Uuid) -> Result<Vec<Uuid>>;
async fn apply_active_flips(
&self,
session_id: Uuid,
activate_ids: Vec<Uuid>,
deactivate_ids: Vec<Uuid>,
) -> Result<()>;
async fn update_session_type(
&self,
tenant_id: &str,
user_id: &str,
session_id: Uuid,
new_session_type_id: Uuid,
new_capabilities: JsonValue,
) -> Result<Session>;
}
#[domain_model]
#[derive(Clone)]
pub struct VariantService {
sessions: Arc<dyn SessionRepo>,
session_types: Arc<dyn SessionTypeRepo>,
messages: Arc<dyn MessageRepo>,
variants: Arc<dyn VariantRepo>,
plugins: PluginService,
message_service: Arc<MessageService>,
plugin_timeout: Duration,
}
impl VariantService {
#[must_use]
pub fn new(
sessions: Arc<dyn SessionRepo>,
session_types: Arc<dyn SessionTypeRepo>,
messages: Arc<dyn MessageRepo>,
variants: Arc<dyn VariantRepo>,
plugins: PluginService,
message_service: Arc<MessageService>,
) -> Self {
Self {
sessions,
session_types,
messages,
variants,
plugins,
message_service,
plugin_timeout: DEFAULT_SWITCH_TYPE_DEADLINE,
}
}
#[must_use]
pub fn with_plugin_timeout(mut self, timeout: Duration) -> Self {
self.plugin_timeout = timeout;
self
}
#[instrument(
skip(self, identity),
fields(
session_id = %session_id,
message_id = %message_id,
operation = "navigate",
),
)]
pub async fn list_variants(
&self,
identity: &Identity,
session_id: Uuid,
message_id: Uuid,
) -> Result<VariantListing> {
let started = OffsetDateTime::now_utc();
let session = self.load_session(identity, session_id).await?;
self.gate_lifecycle_navigation(&session)?;
let target = self
.messages
.find_message_in_session(session_id, message_id)
.await?
.ok_or_else(|| ChatEngineError::not_found("message", message_id))?;
let siblings = self
.variants
.list_siblings(session_id, target.parent_message_id)
.await?;
let total = u32::try_from(siblings.len()).unwrap_or(u32::MAX);
let mut variants = Vec::with_capacity(siblings.len());
let mut current_index: Option<u32> = None;
for (idx, m) in siblings.into_iter().enumerate() {
let is_active = m.is_active;
let info = VariantInfo {
message_id: m.message_id,
variant_index: m.variant_index,
total_variants: total,
is_active,
};
if is_active {
current_index = Some(u32::try_from(idx).unwrap_or(0));
}
variants.push(VariantEntry { message: m, info });
}
log_op_finished(started, "navigate", session_id, message_id, None);
Ok(VariantListing {
variants,
current_index,
})
}
#[instrument(
skip(self, identity),
fields(
session_id = %session_id,
message_id = %message_id,
operation = "set_active",
),
)]
pub async fn set_active_variant(
&self,
identity: &Identity,
session_id: Uuid,
message_id: Uuid,
) -> Result<VariantEntry> {
let started = OffsetDateTime::now_utc();
let session = self.load_session(identity, session_id).await?;
self.gate_lifecycle_mutation(&session)?;
let target = self
.messages
.find_message_in_session(session_id, message_id)
.await?
.ok_or_else(|| ChatEngineError::not_found("message", message_id))?;
self.update_active_path(session_id, target.message_id)
.await?;
let refreshed = self
.messages
.find_message_in_session(session_id, message_id)
.await?
.ok_or_else(|| ChatEngineError::not_found("message", message_id))?;
let total = u32::try_from(
self.variants
.list_siblings(session_id, refreshed.parent_message_id)
.await?
.len(),
)
.unwrap_or(u32::MAX);
let info = VariantInfo {
message_id: refreshed.message_id,
variant_index: refreshed.variant_index,
total_variants: total,
is_active: refreshed.is_active,
};
log_op_finished(
started,
"set_active",
session_id,
message_id,
Some(refreshed.variant_index),
);
increment_variant_creation_total("set_active");
Ok(VariantEntry {
message: refreshed,
info,
})
}
pub async fn update_active_path(&self, session_id: Uuid, message_id: Uuid) -> Result<()> {
let chain = self.variants.ancestor_chain(session_id, message_id).await?;
if chain.is_empty() {
return Err(ChatEngineError::not_found("message", message_id));
}
let chain_set: std::collections::HashSet<Uuid> = chain.iter().copied().collect();
let mut deactivate: Vec<Uuid> = Vec::new();
for ancestor_id in &chain {
let ancestor = self
.messages
.find_message_in_session(session_id, *ancestor_id)
.await?
.ok_or_else(|| ChatEngineError::not_found("message", *ancestor_id))?;
let siblings = self
.variants
.list_siblings(session_id, ancestor.parent_message_id)
.await?;
for sibling in siblings {
if !chain_set.contains(&sibling.message_id) {
deactivate.push(sibling.message_id);
let descendants = self
.variants
.collect_descendants(session_id, sibling.message_id)
.await?;
deactivate.extend(descendants);
}
}
}
deactivate.sort();
deactivate.dedup();
deactivate.retain(|id| !chain_set.contains(id));
self.variants
.apply_active_flips(session_id, chain, deactivate)
.await
}
#[instrument(
skip(self, identity, cancel),
fields(
session_id = %session_id,
message_id = %message_id,
operation = "recreate",
),
)]
pub async fn recreate_variant(
&self,
identity: &Identity,
session_id: Uuid,
message_id: Uuid,
capabilities: Option<Vec<CapabilityValue>>,
cancel: CancellationToken,
) -> Result<SendMessageStream> {
let started = OffsetDateTime::now_utc();
let session = self.load_session(identity, session_id).await?;
self.gate_lifecycle_mutation(&session)?;
let target = self
.messages
.find_message_in_session(session_id, message_id)
.await?
.ok_or_else(|| ChatEngineError::not_found("message", message_id))?;
if !matches!(target.role, MessageRole::Assistant) {
return Err(ChatEngineError::bad_request(
"recreate only applies to assistant messages",
));
}
let parent_message_id = target.parent_message_id.ok_or_else(|| {
ChatEngineError::bad_request(
"target assistant message has no parent \u{2014} cannot recreate",
)
})?;
let session_type_id = session.session_type_id.ok_or_else(|| {
ChatEngineError::bad_request(
"session has no session_type bound; recreate cannot be routed",
)
})?;
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.clone().ok_or_else(|| {
ChatEngineError::bad_request(
"session_type has no plugin_instance_id; recreate cannot be routed",
)
})?;
let inserted = self
.message_service
.prepare_recreate_stub(
session_id,
parent_message_id,
Some(identity.tenant_id.clone()),
)
.await
.map_err(map_unique_violation_to_conflict)?;
let new_message_id = inserted.assistant_message_id;
let new_variant_index = inserted.user_variant_index;
let history = self
.build_branched_history(session_id, parent_message_id)
.await?;
let siblings_now = self
.variants
.list_siblings(session_id, Some(parent_message_id))
.await?;
let total_variants = u32::try_from(siblings_now.len()).unwrap_or(u32::MAX);
let new_variant_info = VariantInfo {
message_id: new_message_id,
variant_index: u32::try_from(new_variant_index).unwrap_or(0),
total_variants,
is_active: true,
};
let stream = self
.message_service
.dispatch_to_plugin(
identity,
session_id,
session_type_id,
plugin_instance_id,
new_message_id,
history,
capabilities,
MessageEventKind::Recreate,
cancel,
)
.await?;
let variants_repo = Arc::clone(&self.variants);
let messages_repo = Arc::clone(&self.messages);
let wrapped = wrap_stream_with_finalizer(
stream,
new_variant_info,
session_id,
new_message_id,
move || {
let variants = Arc::clone(&variants_repo);
let messages = Arc::clone(&messages_repo);
async move {
update_active_path_with_repos(variants, messages, session_id, new_message_id)
.await
}
},
);
log_op_finished(
started,
"recreate",
session_id,
new_message_id,
Some(u32::try_from(new_variant_index).unwrap_or(0)),
);
increment_variant_creation_total("recreate");
Ok(wrapped)
}
#[instrument(
skip(self, identity, parts, cancel),
fields(
session_id = %session_id,
branch_point_message_id = %branch_point_message_id,
operation = "branch",
),
)]
pub async fn branch_message(
&self,
identity: &Identity,
session_id: Uuid,
branch_point_message_id: Uuid,
parts: Vec<MessagePartInput>,
file_ids: Option<Vec<Uuid>>,
capabilities: Option<Vec<CapabilityValue>>,
cancel: CancellationToken,
) -> Result<SendMessageStream> {
let started = OffsetDateTime::now_utc();
let session = self.load_session(identity, session_id).await?;
self.gate_lifecycle_mutation(&session)?;
let _branch_point = self
.messages
.find_message_in_session(session_id, branch_point_message_id)
.await?
.ok_or_else(|| ChatEngineError::not_found("message", branch_point_message_id))?;
let session_type_id = session.session_type_id.ok_or_else(|| {
ChatEngineError::bad_request(
"session has no session_type bound; branch cannot be routed",
)
})?;
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.clone().ok_or_else(|| {
ChatEngineError::bad_request(
"session_type has no plugin_instance_id; branch cannot be routed",
)
})?;
let (user_message_id, _user_variant_index, assistant_message_id) = self
.variants
.insert_user_and_assistant_stub_for_branch(
session_id,
branch_point_message_id,
parts,
file_ids,
Some(identity.tenant_id.clone()),
Some(identity.user_id.clone()),
)
.await
.map_err(map_unique_violation_to_conflict)?;
let history = self
.build_branched_history(session_id, branch_point_message_id)
.await?;
let stream = self
.message_service
.dispatch_to_plugin(
identity,
session_id,
session_type_id,
plugin_instance_id,
assistant_message_id,
history,
capabilities,
MessageEventKind::New,
cancel,
)
.await?;
let variants_repo = Arc::clone(&self.variants);
let messages_repo = Arc::clone(&self.messages);
let wrapped = wrap_stream_simple(stream, move || {
let variants = Arc::clone(&variants_repo);
let messages = Arc::clone(&messages_repo);
async move {
update_active_path_with_repos(variants, messages, session_id, assistant_message_id)
.await
}
});
log_op_finished(started, "branch", session_id, user_message_id, None);
increment_variant_creation_total("branch");
Ok(wrapped)
}
async fn build_branched_history(
&self,
session_id: Uuid,
branch_point_message_id: Uuid,
) -> Result<Vec<Message>> {
let chain = self
.variants
.ancestor_chain(session_id, branch_point_message_id)
.await?;
let mut out: Vec<Message> = Vec::with_capacity(chain.len());
for id in chain.iter().rev() {
if let Some(m) = self
.messages
.find_message_in_session(session_id, *id)
.await?
{
if m.is_hidden_from_backend || !m.is_complete {
continue;
}
out.push(m);
}
}
Ok(out)
}
pub async fn validate_session_type_switch(
&self,
identity: &Identity,
session_id: Uuid,
target_session_type_id: Uuid,
) -> Result<(SessionType, String, Vec<Capability>, Option<JsonValue>)> {
let session = self.load_session(identity, session_id).await?;
self.gate_lifecycle_mutation(&session)?;
let target = self
.session_types
.find_by_id(target_session_type_id)
.await?
.ok_or_else(|| ChatEngineError::not_found("session_type", target_session_type_id))?;
let plugin_instance_id = target.plugin_instance_id.clone().ok_or_else(|| {
ChatEngineError::bad_request("target session_type has no plugin_instance_id")
})?;
let plugin = self
.plugins
.resolve(&plugin_instance_id)
.map_err(|err| match err {
ChatEngineError::NotFound { .. } => ChatEngineError::BackendUnavailable {
reason: format!("target plugin '{plugin_instance_id}' is not registered"),
retry_after: None,
source: None,
},
other => other,
})?;
let plugin_config = self
.plugins
.load_config(&plugin_instance_id, target_session_type_id)
.await?;
let cancel = CancellationToken::new();
let call_ctx = PluginCallContext {
request_id: Uuid::new_v4(),
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: target_session_type_id,
plugin_config,
enabled_capabilities: None,
deadline: Some(Instant::now() + self.plugin_timeout),
cancel: cancel.clone(),
};
let session_ctx = SessionPluginCtx {
session_type_id: target_session_type_id,
session_id: Some(session_id),
call_ctx,
};
let response =
tokio::time::timeout(self.plugin_timeout, plugin.on_session_updated(session_ctx))
.await
.map_err(|_| {
cancel.cancel();
ChatEngineError::BackendUnavailable {
reason: "plugin on_session_updated deadline elapsed".into(),
retry_after: None,
source: None,
}
})?
.map_err(ChatEngineError::from)?;
let returned_metadata = response.metadata;
let available = response.capabilities;
let current_names: Vec<String> = enabled_capability_names(&session);
let available_names: std::collections::HashSet<&str> =
available.iter().map(|c| c.name.as_str()).collect();
for name in ¤t_names {
if !available_names.contains(name.as_str()) {
return Err(ChatEngineError::conflict(format!(
"new session type's available_capabilities is not a superset of \
enabled_capabilities (missing '{name}')",
)));
}
}
Ok((target, plugin_instance_id, available, returned_metadata))
}
#[instrument(
skip(self, identity),
fields(
session_id = %session_id,
target_session_type_id = %target_session_type_id,
operation = "switch_type",
),
)]
pub async fn switch_session_type(
&self,
identity: &Identity,
session_id: Uuid,
target_session_type_id: Uuid,
) -> Result<Session> {
let started = OffsetDateTime::now_utc();
let (_target_type, _plugin_instance_id, capabilities, plugin_metadata) = self
.validate_session_type_switch(identity, session_id, target_session_type_id)
.await?;
let caps_json = serde_json::to_value(&capabilities).unwrap_or(JsonValue::Array(Vec::new()));
let mut updated = self
.variants
.update_session_type(
&identity.tenant_id,
&identity.user_id,
session_id,
target_session_type_id,
caps_json,
)
.await?;
if let Some(plugin_meta) = plugin_metadata {
let merged = merge_plugin_metadata(updated.metadata.clone(), plugin_meta);
updated = self
.sessions
.update_metadata(
&identity.tenant_id,
&identity.user_id,
session_id,
Some(merged),
)
.await?;
}
log_op_finished(started, "switch_type", session_id, session_id, None);
increment_variant_creation_total("switch_type");
Ok(updated)
}
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 gate_lifecycle_mutation(&self, session: &Session) -> Result<()> {
if matches!(session.lifecycle_state, LifecycleState::Active) {
Ok(())
} else {
Err(ChatEngineError::conflict(format!(
"session is {} and does not accept variant mutations",
session.lifecycle_state,
)))
}
}
fn gate_lifecycle_navigation(&self, session: &Session) -> Result<()> {
if matches!(
session.lifecycle_state,
LifecycleState::Active | LifecycleState::Archived
) {
Ok(())
} else {
Err(ChatEngineError::conflict(format!(
"session is {} and does not accept variant navigation",
session.lifecycle_state,
)))
}
}
}
fn map_unique_violation_to_conflict(err: ChatEngineError) -> ChatEngineError {
if let ChatEngineError::Internal { reason, source } = &err {
let lower = reason.to_lowercase();
if lower.contains("exhausted")
|| lower.contains("uq_messages_session_parent_variant")
|| lower.contains("unique constraint")
{
return ChatEngineError::Conflict {
reason: format!("concurrent variant creation: {reason}"),
};
}
let _ = source;
}
err
}
async fn update_active_path_with_repos(
variants: Arc<dyn VariantRepo>,
messages: Arc<dyn MessageRepo>,
session_id: Uuid,
message_id: Uuid,
) -> Result<()> {
let chain = variants.ancestor_chain(session_id, message_id).await?;
if chain.is_empty() {
return Err(ChatEngineError::not_found("message", message_id));
}
let chain_set: std::collections::HashSet<Uuid> = chain.iter().copied().collect();
let mut deactivate: Vec<Uuid> = Vec::new();
for ancestor_id in &chain {
let ancestor = messages
.find_message_in_session(session_id, *ancestor_id)
.await?
.ok_or_else(|| ChatEngineError::not_found("message", *ancestor_id))?;
let siblings = variants
.list_siblings(session_id, ancestor.parent_message_id)
.await?;
for sibling in siblings {
if !chain_set.contains(&sibling.message_id) {
deactivate.push(sibling.message_id);
let descendants = variants
.collect_descendants(session_id, sibling.message_id)
.await?;
deactivate.extend(descendants);
}
}
}
deactivate.sort();
deactivate.dedup();
deactivate.retain(|id| !chain_set.contains(id));
variants
.apply_active_flips(session_id, chain, deactivate)
.await
}
fn enabled_capability_names(session: &Session) -> Vec<String> {
let Some(JsonValue::Array(arr)) = session.enabled_capabilities.as_ref() else {
return Vec::new();
};
arr.iter()
.filter_map(|entry| match entry {
JsonValue::Object(map) => map.get("name").and_then(|n| n.as_str()).map(str::to_owned),
_ => None,
})
.collect()
}
fn wrap_stream_with_finalizer<F, Fut>(
upstream: SendMessageStream,
variant_info: VariantInfo,
session_id: Uuid,
new_message_id: Uuid,
finalizer: F,
) -> SendMessageStream
where
F: FnOnce() -> Fut + Send + 'static,
Fut: std::future::Future<Output = Result<()>> + Send + 'static,
{
let info_json = serde_json::to_value(&variant_info).unwrap_or(JsonValue::Null);
let mapped = upstream.map(move |evt| augment_complete_event(evt, &info_json));
let sentinel = stream::once(async move {
if let Err(err) = finalizer().await {
warn!(
session_id = %session_id,
message_id = %new_message_id,
error = %err,
"active-path update after stream end failed (variant retained, but is_active state may be stale)"
);
}
None::<StreamingEvent>
})
.filter_map(|v: Option<StreamingEvent>| async move { v });
mapped.chain(sentinel).boxed()
}
fn wrap_stream_simple<F, Fut>(upstream: SendMessageStream, finalizer: F) -> SendMessageStream
where
F: FnOnce() -> Fut + Send + 'static,
Fut: std::future::Future<Output = Result<()>> + Send + 'static,
{
let sentinel = stream::once(async move {
if let Err(err) = finalizer().await {
warn!(error = %err, "post-stream active-path update failed");
}
None::<StreamingEvent>
})
.filter_map(|v: Option<StreamingEvent>| async move { v });
upstream.chain(sentinel).boxed()
}
fn augment_complete_event(evt: StreamingEvent, variant_info: &JsonValue) -> StreamingEvent {
match evt {
StreamingEvent::Complete(mut c) => {
let merged = match c.metadata.take() {
Some(JsonValue::Object(mut map)) => {
map.insert("variant_info".to_string(), variant_info.clone());
Some(JsonValue::Object(map))
}
Some(other) => {
Some(json!({ "inner": other, "variant_info": variant_info }))
}
None => Some(json!({ "variant_info": variant_info })),
};
c.metadata = merged;
StreamingEvent::Complete(c)
}
other => other,
}
}
fn log_op_finished(
started: OffsetDateTime,
operation: &'static str,
session_id: Uuid,
message_id: Uuid,
variant_index: Option<u32>,
) {
let now = OffsetDateTime::now_utc();
let duration_ms = (now - started).whole_milliseconds().max(0);
if let Some(idx) = variant_index {
info!(
session_id = %session_id,
message_id = %message_id,
variant_index = idx,
operation,
duration_ms,
"variant operation completed"
);
} else {
info!(
session_id = %session_id,
message_id = %message_id,
operation,
duration_ms,
"variant operation completed"
);
}
}
fn increment_variant_creation_total(operation: &'static str) {
debug!(
operation,
"variant_creation_total += 1 (no metrics facade yet)"
);
}
#[cfg(test)]
#[path = "variant_service_tests.rs"]
mod variant_service_tests;