use std::sync::Arc;
use axum::Extension;
use axum::extract::Path;
use axum::response::{Json, Response};
use chat_engine_sdk::models::{CapabilityValue, MessagePart, MessagePartInput, VariantInfo};
use serde::{Deserialize, Serialize};
use serde_json::Value as JsonValue;
use tokio_util::sync::CancellationToken;
use tracing::field::Empty;
use uuid::Uuid;
use toolkit_security::SecurityContext;
use crate::api::rest::handlers::sessions::{identity_from_ctx, reject_body_identity};
use crate::domain::error::{ChatEngineError, Result};
use crate::domain::service::variant_service::{VariantEntry, VariantListing, VariantService};
#[derive(Debug, Deserialize, Default)]
pub struct RecreateBody {
#[serde(default)]
pub enabled_capabilities: Option<Vec<CapabilityValue>>,
pub tenant_id: Option<JsonValue>,
pub user_id: Option<JsonValue>,
}
#[derive(Debug, Deserialize)]
pub struct BranchBody {
#[serde(default)]
pub parts: Vec<MessagePartInput>,
#[serde(default)]
pub file_ids: Option<Vec<Uuid>>,
#[serde(default)]
pub enabled_capabilities: Option<Vec<CapabilityValue>>,
pub tenant_id: Option<JsonValue>,
pub user_id: Option<JsonValue>,
}
#[derive(Debug, Deserialize)]
pub struct ActiveVariantBody {
pub message_id: Uuid,
#[serde(default)]
pub variant_index: Option<u32>,
pub tenant_id: Option<JsonValue>,
pub user_id: Option<JsonValue>,
}
#[derive(Debug, Deserialize)]
pub struct ActiveVariantCompatBody {
pub variant_index: u32,
pub tenant_id: Option<JsonValue>,
pub user_id: Option<JsonValue>,
}
#[derive(Debug, Deserialize)]
pub struct SwitchSessionTypeBody {
pub session_type_id: Uuid,
pub tenant_id: Option<JsonValue>,
pub user_id: Option<JsonValue>,
}
#[derive(Debug, Serialize)]
pub struct ListVariantsResponse {
pub variants: Vec<ListVariantsEntry>,
pub current_index: Option<u32>,
}
#[derive(Debug, Serialize)]
pub struct ListVariantsEntry {
pub message_id: Uuid,
pub variant_index: u32,
pub total_variants: u32,
pub is_active: bool,
pub is_complete: bool,
pub parts: Vec<MessagePart>,
pub metadata: Option<JsonValue>,
#[serde(with = "time::serde::rfc3339")]
pub created_at: time::OffsetDateTime,
}
impl From<VariantListing> for ListVariantsResponse {
fn from(value: VariantListing) -> Self {
Self {
current_index: value.current_index,
variants: value
.variants
.into_iter()
.map(|VariantEntry { message, info }| ListVariantsEntry {
message_id: info.message_id,
variant_index: info.variant_index,
total_variants: info.total_variants,
is_active: info.is_active,
is_complete: message.is_complete,
parts: message.parts,
metadata: message.metadata,
created_at: message.created_at,
})
.collect(),
}
}
}
#[tracing::instrument(
skip(svc, ctx, body),
fields(
request_id = Empty,
session_id = %session_id,
message_id = %message_id,
),
)]
pub async fn recreate_variant(
Extension(ctx): Extension<SecurityContext>,
Extension(svc): Extension<Arc<VariantService>>,
Path((session_id, message_id)): Path<(Uuid, Uuid)>,
Json(body): Json<RecreateBody>,
) -> Result<Response> {
reject_body_identity(&body.tenant_id, &body.user_id)?;
let identity = identity_from_ctx(&ctx)?;
let cancel = CancellationToken::new();
let stream = svc
.recreate_variant(
&identity,
session_id,
message_id,
body.enabled_capabilities,
cancel,
)
.await?;
stream_to_sse_response(stream)
}
#[tracing::instrument(
skip(svc, ctx, body),
fields(
session_id = %session_id,
branch_point_message_id = %message_id,
),
)]
pub async fn branch_message(
Extension(ctx): Extension<SecurityContext>,
Extension(svc): Extension<Arc<VariantService>>,
Path((session_id, message_id)): Path<(Uuid, Uuid)>,
Json(body): Json<BranchBody>,
) -> Result<Response> {
reject_body_identity(&body.tenant_id, &body.user_id)?;
let identity = identity_from_ctx(&ctx)?;
let cancel = CancellationToken::new();
let stream = svc
.branch_message(
&identity,
session_id,
message_id,
body.parts,
body.file_ids,
body.enabled_capabilities,
cancel,
)
.await?;
stream_to_sse_response(stream)
}
#[tracing::instrument(
skip(svc, ctx),
fields(session_id = %session_id, message_id = %message_id),
)]
pub async fn list_variants(
Extension(ctx): Extension<SecurityContext>,
Extension(svc): Extension<Arc<VariantService>>,
Path((session_id, message_id)): Path<(Uuid, Uuid)>,
) -> Result<Json<ListVariantsResponse>> {
let identity = identity_from_ctx(&ctx)?;
let listing = svc.list_variants(&identity, session_id, message_id).await?;
Ok(Json(listing.into()))
}
#[tracing::instrument(
skip(svc, ctx, body),
fields(session_id = %session_id),
)]
pub async fn set_active_variant(
Extension(ctx): Extension<SecurityContext>,
Extension(svc): Extension<Arc<VariantService>>,
Path(session_id): Path<Uuid>,
Json(body): Json<ActiveVariantBody>,
) -> Result<Json<VariantInfo>> {
reject_body_identity(&body.tenant_id, &body.user_id)?;
let identity = identity_from_ctx(&ctx)?;
let entry = svc
.set_active_variant(&identity, session_id, body.message_id)
.await?;
Ok(Json(entry.info))
}
#[tracing::instrument(
skip(svc, ctx, body),
fields(session_id = %session_id, message_id = %message_id),
)]
pub async fn set_active_variant_compat(
Extension(ctx): Extension<SecurityContext>,
Extension(svc): Extension<Arc<VariantService>>,
Path((session_id, message_id)): Path<(Uuid, Uuid)>,
Json(body): Json<ActiveVariantCompatBody>,
) -> Result<Json<VariantInfo>> {
reject_body_identity(&body.tenant_id, &body.user_id)?;
let identity = identity_from_ctx(&ctx)?;
let listing = svc.list_variants(&identity, session_id, message_id).await?;
let target = listing
.variants
.into_iter()
.find(|e| e.info.variant_index == body.variant_index)
.ok_or_else(|| {
ChatEngineError::not_found(
"variant",
format!("{}:variant_index={}", message_id, body.variant_index),
)
})?;
let entry = svc
.set_active_variant(&identity, session_id, target.message.message_id)
.await?;
Ok(Json(entry.info))
}
#[tracing::instrument(
skip(svc, ctx, body),
fields(session_id = %session_id),
)]
pub async fn switch_session_type(
Extension(ctx): Extension<SecurityContext>,
Extension(svc): Extension<Arc<VariantService>>,
Path(session_id): Path<Uuid>,
Json(body): Json<SwitchSessionTypeBody>,
) -> Result<Json<chat_engine_sdk::models::Session>> {
reject_body_identity(&body.tenant_id, &body.user_id)?;
let identity = identity_from_ctx(&ctx)?;
let session = svc
.switch_session_type(&identity, session_id, body.session_type_id)
.await?;
Ok(Json(session))
}
#[tracing::instrument(
skip(svc, ctx, body),
fields(session_id = %session_id),
)]
pub async fn switch_session_type_compat(
Extension(ctx): Extension<SecurityContext>,
Extension(svc): Extension<Arc<VariantService>>,
Path(session_id): Path<Uuid>,
Json(body): Json<SwitchSessionTypeBody>,
) -> Result<Json<chat_engine_sdk::models::Session>> {
switch_session_type(Extension(ctx), Extension(svc), Path(session_id), Json(body)).await
}
fn stream_to_sse_response(
stream: crate::domain::service::message_service::SendMessageStream,
) -> Result<Response> {
Ok(crate::api::rest::sse_delta_stream_response(stream))
}
#[cfg(test)]
#[path = "variants_tests.rs"]
mod variants_tests;