AstroResearch/src/api/agent.rs
Asfmq c5fd5b0d66 refactor: 全栈质量硬化
后端:
  - 权限系统重写: 全局→按 Session 隔离, 新增规则查询 API
  - 安全加固: 登录 IP 限流, Token 仅存 Cookie, bibcode 白名单校验
  - SSE 超时保护, 异步 I/O 迁移, 10+ 处静默 DB 错误改为显式日志
  - ar5iv 下标解析修复, parse_paper_row 去重, 优雅关闭

  前端:
  - useSyncScroll 重写: 段落 ID 映射修复中英错位
  - 全局竞态修复 (active 标志), libraryRef 闭包过期修复
  - ErrorBoundary + vitest 测试基础设施
  - Logo 组件提取, CustomSelect 泛型化, TabId 类型统一
  - ReaderPanel 自动视图模式, AIAssistantPanel 状态批处理
2026-06-28 14:40:26 +08:00

816 lines
26 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

// src/api/agent.rs
//
// 科研智能体 API 控制器。
// 提供 SSE 流式对话接口和会话管理 CRUD 接口。
use axum::{
extract::{Path, Query, State},
response::sse::{Event, Sse},
Json,
};
use futures_util::stream::Stream;
use serde::{Deserialize, Serialize};
use sqlx::Row;
use std::convert::Infallible;
use std::sync::Arc;
use tracing::{error, info};
use super::error::{ApiResult, AppError};
use super::AppState;
use crate::agent::runtime::{AgentRuntime, AgentStreamEvent};
// ── POST /api/chat/agent ──
// SSE 流式智能体对话接口
#[derive(Debug, Deserialize)]
pub struct AgentChatRequest {
pub question: String,
pub session_id: Option<String>,
/// Agent 运行模式: "default" / "deep-research" / "literature-reader"
#[serde(default = "default_mode")]
pub mode: String,
/// 是否启用 LLM 思考模式。None = 由 mode 决定Some(true/false) = 用户显式覆盖。
#[serde(default)]
pub thinking: Option<bool>,
/// 可选的图片附件base64 编码 + MIME 类型)
#[serde(default)]
pub image: Option<AttachedImage>,
}
/// 用户附带的图片,用于多模态 Agent 提问。
#[derive(Debug, Deserialize)]
pub struct AttachedImage {
/// base64 编码的图片数据(不含 data:xxx;base64, 前缀)。与 path 互斥。
#[serde(default)]
pub data: String,
/// MIME 类型,如 "image/png"、"image/jpeg"
#[serde(default)]
pub mime_type: String,
/// 已有图片的相对路径(重试时复用已有文件,不再重新 base64 解码存盘)
#[serde(default)]
pub path: Option<String>,
}
fn default_mode() -> String {
"default".to_string()
}
#[derive(Debug, Serialize)]
pub struct AgentModeDto {
pub id: &'static str,
pub name: &'static str,
pub description: &'static str,
pub icon: &'static str,
}
// ── GET /api/chat/modes ──
// 获取可用的智能体运行模式列表
pub async fn get_agent_modes() -> Json<Vec<AgentModeDto>> {
use crate::agent::modes::ModeRegistry;
let registry = ModeRegistry::builtins();
let modes = registry
.list()
.iter()
.map(|m| AgentModeDto {
id: m.id,
name: m.name,
description: m.description,
icon: m.icon,
})
.collect();
Json(modes)
}
pub async fn chat_agent(
State(state): State<Arc<AppState>>,
Json(req): Json<AgentChatRequest>,
) -> ApiResult<Sse<impl Stream<Item = Result<Event, Infallible>>>> {
info!(
"接收到智能体对话请求: question='{}', session_id={:?}, has_image={}",
req.question,
req.session_id,
req.image.is_some()
);
// 处理图片附件:保存到磁盘,路径注入 Agent 上下文,前端和 DB 保留原始问题
let question = req.question.clone();
let (image_context, image_path_for_db): (Option<String>, Option<String>) = match req.image {
Some(ref img) => {
// 如果带有 path 字段(重试场景),复用已有文件,不重新保存
let relative_path: String = if let Some(ref existing_path) = img.path {
let full = state.config.library_dir.join(existing_path);
if full.exists() {
info!("重试复用已有图片: {}", existing_path);
existing_path.clone()
} else {
return Err(AppError::bad_request(format!(
"图片文件不存在: {}",
existing_path
)));
}
} else {
if img.data.is_empty() {
return Err(AppError::bad_request("图片数据为空"));
}
if !img.mime_type.starts_with("image/") {
return Err(AppError::bad_request(format!(
"不支持的图片类型: {}",
img.mime_type
)));
}
if state.vision_llm.is_none() {
return Err(AppError::bad_request(
"图片分析功能未启用。请配置 LLM_VISION_MODEL 环境变量后重试。",
));
}
let ext = img.mime_type.strip_prefix("image/").unwrap_or("png");
let upload_dir = state
.config
.library_dir
.join(".agent_images")
.join("uploads");
std::fs::create_dir_all(&upload_dir)
.map_err(|e| AppError::internal(format!("创建上传目录失败: {}", e)))?;
let filename = format!("{}.{}", uuid::Uuid::new_v4(), ext);
let filepath = upload_dir.join(&filename);
use base64::{engine::general_purpose, Engine as _};
let bytes = general_purpose::STANDARD
.decode(&img.data)
.map_err(|e| AppError::bad_request(format!("图片 base64 解码失败: {}", e)))?;
std::fs::write(&filepath, &bytes)
.map_err(|e| AppError::internal(format!("保存图片失败: {}", e)))?;
let rel = filepath
.strip_prefix(&state.config.library_dir)
.unwrap_or(&filepath)
.display()
.to_string();
info!("用户图片已保存: {}", rel);
rel
};
let ctx = format!(
"用户上传了一张图片,已保存到: {}\n如需分析此图片,请使用 analyze_image 工具,传入 image_path=\"{}\"",
relative_path, relative_path
);
(Some(ctx), Some(relative_path))
}
None => (None, None),
};
let mut runtime = AgentRuntime::new(Arc::clone(&state)).with_mode(&req.mode);
// 只有 mode 未强制固定 thinking 时,用户才可以覆盖
if runtime.mode_fixed_thinking().is_none() {
if let Some(thinking) = req.thinking {
runtime = runtime.with_thinking(thinking);
}
}
let (tx, mut rx) = tokio::sync::mpsc::unbounded_channel::<AgentStreamEvent>();
let session_id = req.session_id.clone();
// 在后台 tokio 任务中执行 Agent 循环
tokio::spawn(async move {
match runtime
.run_turn_with_image_context(
session_id,
&question,
image_context,
image_path_for_db,
tx.clone(),
)
.await
{
Ok(sid) => {
info!("智能体对话完成: session_id={}", sid);
}
Err(e) => {
error!("智能体对话执行出错: {}", e);
let _ = tx.send(AgentStreamEvent::Error {
message: format!("智能体执行错误: {}", e),
});
let _ = tx.send(AgentStreamEvent::Done);
}
}
});
// 将 mpsc 通道转换为 SSE 事件流(带 10 分钟超时)
const SSE_TIMEOUT_SECS: u64 = 600;
let stream = async_stream::stream! {
let deadline = tokio::time::Instant::now() + std::time::Duration::from_secs(SSE_TIMEOUT_SECS);
loop {
let remaining = deadline.saturating_duration_since(tokio::time::Instant::now());
if remaining.is_zero() {
let timeout_event = AgentStreamEvent::Error {
message: "Agent 执行超时10 分钟),请重试。".to_string(),
};
let data = serde_json::to_string(&timeout_event).unwrap_or_default();
yield Ok(Event::default().data(data));
break;
}
match tokio::time::timeout(remaining, rx.recv()).await {
Ok(Some(event)) => {
let data = serde_json::to_string(&event).unwrap_or_default();
let is_done = matches!(event, AgentStreamEvent::Done);
yield Ok(Event::default().data(data));
if is_done {
break;
}
}
Ok(None) => break, // channel closed
Err(_) => {
let timeout_event = AgentStreamEvent::Error {
message: "Agent 执行超时10 分钟),请重试。".to_string(),
};
let data = serde_json::to_string(&timeout_event).unwrap_or_default();
yield Ok(Event::default().data(data));
break;
}
}
}
};
Ok(Sse::new(stream))
}
// ── GET /api/chat/sessions ──
// 获取会话列表
#[derive(Debug, Deserialize)]
pub struct SessionListParams {
pub limit: Option<i64>,
pub offset: Option<i64>,
}
#[derive(Debug, Serialize)]
pub struct SessionSummary {
pub session_id: String,
pub title: String,
pub model: String,
pub mode: String,
pub turn_count: i32,
pub summary: Option<String>,
pub created_at: String,
pub updated_at: String,
}
pub async fn list_sessions(
State(state): State<Arc<AppState>>,
Query(params): Query<SessionListParams>,
) -> ApiResult<Json<Vec<SessionSummary>>> {
let limit = params.limit.unwrap_or(50);
let offset = params.offset.unwrap_or(0);
let rows = sqlx::query(
"SELECT session_id, title, model, mode, turn_count, summary, created_at, updated_at \
FROM agent_sessions \
WHERE deleted_at IS NULL \
ORDER BY updated_at DESC \
LIMIT ? OFFSET ?",
)
.bind(limit)
.bind(offset)
.fetch_all(&state.db)
.await
.map_err(|e| AppError::internal(format!("查询会话列表失败: {}", e)))?;
let sessions: Vec<SessionSummary> = rows
.iter()
.map(|r| SessionSummary {
session_id: r.get(0),
title: r.get(1),
model: r.get(2),
mode: r.get(3),
turn_count: r.get(4),
summary: r.get(5),
created_at: r.get(6),
updated_at: r.get(7),
})
.collect();
Ok(Json(sessions))
}
// ── GET /api/chat/sessions/:id ──
// 获取单个会话的全部消息历史
#[derive(Debug, Serialize)]
pub struct SessionDetail {
pub session: SessionSummary,
pub messages: Vec<MessageRecord>,
}
#[derive(Debug, Serialize)]
pub struct MessageRecord {
pub id: i64,
pub agent_name: String,
pub turn_index: i32,
pub step_index: i32,
pub role: String,
pub content: String,
pub thought: Option<String>,
pub tool_calls: Option<serde_json::Value>,
pub tool_call_id: Option<String>,
pub token_count: i32,
pub metadata: Option<serde_json::Value>,
pub created_at: String,
}
pub async fn get_session(
State(state): State<Arc<AppState>>,
Path(session_id): Path<String>,
) -> ApiResult<Json<SessionDetail>> {
// 查询会话元信息
let session_row = sqlx::query(
"SELECT session_id, title, model, mode, turn_count, summary, created_at, updated_at \
FROM agent_sessions \
WHERE session_id = ? AND deleted_at IS NULL",
)
.bind(&session_id)
.fetch_optional(&state.db)
.await
.map_err(|e| AppError::internal(format!("查询会话失败: {}", e)))?
.ok_or_else(|| AppError::not_found(format!("会话 {} 不存在", session_id)))?;
let session = SessionSummary {
session_id: session_row.get(0),
title: session_row.get(1),
model: session_row.get(2),
mode: session_row.get(3),
turn_count: session_row.get(4),
summary: session_row.get(5),
created_at: session_row.get(6),
updated_at: session_row.get(7),
};
// 查询消息列表(包含 lead 和 subagent 消息,前端按 agent_name/metadata 区分渲染)
let msg_rows = sqlx::query(
"SELECT id, agent_name, turn_index, step_index, role, content, thought, tool_calls, tool_call_id, token_count, metadata, created_at \
FROM agent_messages \
WHERE session_id = ? AND active = 1 \
ORDER BY id ASC"
)
.bind(&session_id)
.fetch_all(&state.db)
.await
.map_err(|e| AppError::internal(format!("查询消息列表失败: {}", e)))?;
let messages: Vec<MessageRecord> = msg_rows
.iter()
.map(|r| {
let tool_calls_json: Option<String> = r.get(7);
let metadata_json: Option<String> = r.get(10);
MessageRecord {
id: r.get(0),
agent_name: r.get(1),
turn_index: r.get(2),
step_index: r.get(3),
role: r.get(4),
content: r.get(5),
thought: r.get(6),
tool_calls: tool_calls_json.and_then(|s| serde_json::from_str(&s).ok()),
tool_call_id: r.get(8),
token_count: r.get(9),
metadata: metadata_json.and_then(|s| serde_json::from_str(&s).ok()),
created_at: r.get(11),
}
})
.collect();
Ok(Json(SessionDetail { session, messages }))
}
// ── DELETE /api/chat/sessions/:id ──
// 软删除会话
pub async fn delete_session(
State(state): State<Arc<AppState>>,
Path(session_id): Path<String>,
) -> ApiResult<Json<serde_json::Value>> {
let result = sqlx::query(
"UPDATE agent_sessions SET deleted_at = CURRENT_TIMESTAMP WHERE session_id = ? AND deleted_at IS NULL"
)
.bind(&session_id)
.execute(&state.db)
.await
.map_err(|e| AppError::internal(format!("删除会话失败: {}", e)))?;
if result.rows_affected() == 0 {
return Err(AppError::not_found(format!(
"会话 {} 不存在或已删除",
session_id
)));
}
info!("会话已软删除: {}", session_id);
Ok(Json(
serde_json::json!({ "status": "deleted", "session_id": session_id }),
))
}
// ── POST /api/chat/sessions/:id/stop ──
// 手动停止智能体执行接口
pub async fn stop_agent(
State(state): State<Arc<AppState>>,
Path(session_id): Path<String>,
) -> ApiResult<Json<serde_json::Value>> {
if let Ok(mut cancelled) = state.cancelled_runs.lock() {
cancelled.insert(session_id.clone());
}
info!("已接收并记录手动中止请求,会话 ID: {}", session_id);
Ok(Json(
serde_json::json!({ "status": "stopping", "session_id": session_id }),
))
}
// ── GET /api/chat/metrics ──
// 返回聚合的智能体运行指标
#[derive(Debug, Serialize)]
pub struct AgentMetricsResponse {
pub total_sessions: i64,
pub total_tool_calls: i64,
pub tool_call_breakdown: serde_json::Value,
pub avg_steps_per_session: f64,
pub error_rate: f64,
}
pub async fn get_agent_metrics(
State(state): State<Arc<AppState>>,
) -> ApiResult<Json<AgentMetricsResponse>> {
// 总会话数
let total_sessions: i64 =
sqlx::query_scalar("SELECT COUNT(*) FROM agent_sessions WHERE deleted_at IS NULL")
.fetch_one(&state.db)
.await
.unwrap_or(0);
// 工具调用统计(从审计日志聚合)
let total_tool_calls: i64 =
sqlx::query_scalar("SELECT COUNT(*) FROM agent_audit_log WHERE status IN ('OK', 'FAIL')")
.fetch_one(&state.db)
.await
.unwrap_or(0);
// 各工具调用次数
let breakdown_rows: Vec<(String, i64)> = sqlx::query_as(
"SELECT COALESCE(tool_name, 'unknown'), COUNT(*) as cnt \
FROM agent_audit_log \
WHERE status IN ('OK', 'FAIL') \
GROUP BY tool_name \
ORDER BY cnt DESC",
)
.fetch_all(&state.db)
.await
.unwrap_or_default();
let tool_call_breakdown: serde_json::Value = breakdown_rows
.iter()
.map(|(name, cnt)| serde_json::json!({ name: cnt }))
.fold(serde_json::json!({}), |mut acc, v| {
if let serde_json::Value::Object(map) = &mut acc {
if let serde_json::Value::Object(v_map) = v {
for (k, val) in v_map {
map.insert(k.clone(), val.clone());
}
}
}
acc
});
// 平均步数
let avg_steps: f64 = sqlx::query_scalar(
"SELECT COALESCE(AVG(CAST(turn_count AS REAL)), 0.0) \
FROM agent_sessions WHERE deleted_at IS NULL",
)
.fetch_one(&state.db)
.await
.unwrap_or(0.0);
// 错误率
let total_errors: i64 =
sqlx::query_scalar("SELECT COUNT(*) FROM agent_audit_log WHERE status = 'FAIL'")
.fetch_one(&state.db)
.await
.unwrap_or(0);
let error_rate = if total_tool_calls > 0 {
total_errors as f64 / total_tool_calls as f64
} else {
0.0
};
Ok(Json(AgentMetricsResponse {
total_sessions,
total_tool_calls,
tool_call_breakdown,
avg_steps_per_session: avg_steps,
error_rate,
}))
}
// ── GET /api/chat/sessions/:id/audit ──
// 返回指定会话的审计日志
#[derive(Debug, Serialize)]
pub struct AuditLogEntry {
pub id: i64,
pub step: i32,
pub tool_name: Option<String>,
pub status: String,
pub elapsed_ms: i32,
pub output_preview: Option<String>,
pub created_at: String,
}
pub async fn get_session_audit(
State(state): State<Arc<AppState>>,
Path(session_id): Path<String>,
) -> ApiResult<Json<Vec<AuditLogEntry>>> {
let rows = sqlx::query(
"SELECT id, step, tool_name, status, elapsed_ms, output_preview, created_at \
FROM agent_audit_log \
WHERE session_id = ? \
ORDER BY id ASC",
)
.bind(&session_id)
.fetch_all(&state.db)
.await
.map_err(|e| AppError::internal(format!("查询审计日志失败: {}", e)))?;
let entries: Vec<AuditLogEntry> = rows
.iter()
.map(|r| AuditLogEntry {
id: r.get(0),
step: r.get(1),
tool_name: r.get(2),
status: r.get(3),
elapsed_ms: r.get(4),
output_preview: r.get(5),
created_at: r.get(6),
})
.collect();
Ok(Json(entries))
}
// ── POST /api/chat/answer_question ──
// 用户回答 Agent 的提问ask_user 工具配合使用)
#[derive(Debug, Deserialize)]
pub struct AnswerQuestionRequest {
pub question_id: String,
pub answers: Vec<String>,
pub free_text: Option<String>,
}
pub async fn answer_question(
State(state): State<Arc<AppState>>,
Json(req): Json<AnswerQuestionRequest>,
) -> ApiResult<Json<serde_json::Value>> {
use crate::agent::tools::ask_user::UserAnswer;
let mut pending = match state.pending_questions.lock() {
Ok(p) => p,
Err(_) => {
return Err(AppError::internal("服务器内部状态异常,请稍后重试"));
}
};
let question_id = req.question_id.clone();
match pending.remove(&question_id) {
Some(pq) => {
let answer = UserAnswer {
question_id: question_id.clone(),
answers: req.answers.clone(),
free_text: req.free_text.clone(),
};
match pq.answer_tx.send(answer) {
Ok(()) => {
info!("[API] 用户回答了问题: id={}", question_id);
Ok(Json(
serde_json::json!({"status": "ok", "question_id": question_id}),
))
}
Err(_) => Err(AppError::gone("问题已超时或已被回答")),
}
}
None => Err(AppError::not_found(format!(
"未找到待回答问题: {}",
question_id
))),
}
}
// ── GET /api/chat/pending_questions ──
// 获取当前待回答的问题(前端轮询或初始化)
pub async fn get_pending_questions(
State(state): State<Arc<AppState>>,
) -> Json<Vec<serde_json::Value>> {
let pending = match state.pending_questions.lock() {
Ok(p) => p,
Err(_) => return Json(Vec::new()),
};
let questions: Vec<serde_json::Value> = pending
.iter()
.map(|(id, pq)| {
serde_json::from_str::<serde_json::Value>(&pq.question_json)
.unwrap_or(serde_json::json!({"question_id": id}))
})
.collect();
Json(questions)
}
// ── POST /api/chat/sessions/:id/permissions/respond ──
// 用户响应权限请求
pub async fn respond_permission(
State(state): State<Arc<AppState>>,
Path(session_id): Path<String>,
Json(req): Json<super::PermissionResponse>,
) -> ApiResult<Json<serde_json::Value>> {
let mut perms = match state.pending_permissions.lock() {
Ok(p) => p,
Err(_) => return Err(AppError::internal("内部状态异常")),
};
// 按 tool_call_id 查找匹配的权限请求
let perm_id = perms
.iter()
.find(|(_, p)| p.tool_call_id == req.tool_call_id)
.map(|(id, _)| id.clone());
match perm_id {
Some(id) => {
let perm = perms.remove(&id).unwrap();
match perm.response_tx.send(req) {
Ok(()) => {
info!(
"[API] 用户响应了权限请求: session={} tool_call_id={}",
session_id, perm.tool_call_id
);
Ok(Json(serde_json::json!({"status": "ok"})))
}
Err(_) => Err(AppError::gone("权限请求已超时或已处理")),
}
}
None => Err(AppError::not_found(
"未找到该权限请求(可能已超时或已处理)",
)),
}
}
// ── GET /api/chat/sessions/:id/permissions ──
// 获取当前待处理的权限请求(前端轮询)
pub async fn get_pending_permissions(
State(state): State<Arc<AppState>>,
) -> Json<Vec<serde_json::Value>> {
let perms = match state.pending_permissions.lock() {
Ok(p) => p,
Err(_) => return Json(Vec::new()),
};
let result: Vec<serde_json::Value> = perms
.iter()
.map(|(id, p)| {
serde_json::json!({
"permission_id": id,
"tool_call_id": p.tool_call_id,
"tool_name": p.tool_name,
"message": p.message,
"arguments": p.arguments,
})
})
.collect();
Json(result)
}
// ── POST /api/chat/sessions/:id/branch ──
// 创建会话分叉(复制所有 active=1 消息到新会话)
#[derive(Debug, Serialize)]
pub struct BranchResponse {
pub branch_session_id: String,
pub forked_at_message_id: i64,
pub copied_count: usize,
}
pub async fn branch_session(
State(state): State<Arc<AppState>>,
Path(session_id): Path<String>,
) -> ApiResult<Json<BranchResponse>> {
let result = crate::agent::runtime::session::branch_session(&state.db, &session_id)
.await
.map_err(|e| AppError::bad_request(e.to_string()))?;
Ok(Json(BranchResponse {
branch_session_id: result.branch_session_id,
forked_at_message_id: result.forked_at_message_id,
copied_count: result.copied_count,
}))
}
// ── POST /api/chat/sessions/:id/retry ──
// 重试最后一次对话(硬删除 + 返回消息文本供前端重提交)
#[derive(Debug, Serialize)]
pub struct RetryResponse {
/// 被删除的用户消息文本(前端可自动重提交)
pub retried_message: String,
pub new_turn_index: i32,
pub deleted_count: i64,
pub session_id: String,
/// 原消息附带的图片路径(如果有)
pub image_path: Option<String>,
}
pub async fn retry_session(
State(state): State<Arc<AppState>>,
Path(session_id): Path<String>,
) -> ApiResult<Json<RetryResponse>> {
let (retried_message, new_turn_index, image_path) =
crate::agent::runtime::session::retry_last_turn(&state.db, &session_id)
.await
.map_err(|e| AppError::bad_request(e.to_string()))?;
Ok(Json(RetryResponse {
retried_message,
new_turn_index,
deleted_count: 0,
session_id,
image_path,
}))
}
// ── POST /api/chat/sessions/:id/rewind ──
// 回退会话到指定的消息之前(软删除)
//
// 回退后的消息标记为 active=0审计保留
// 在未产生新对话前可通过 /rewind/restore 恢复。
// 如果已产生新对话,使用 /branch 分叉探索替代路径。
#[derive(Debug, Deserialize)]
pub struct RewindRequest {
/// 回退 N 个用户轮次(默认 1
pub n: Option<usize>,
/// 或者指定回退到的消息 ID
pub message_id: Option<i64>,
}
#[derive(Debug, Serialize)]
pub struct RewindResponse {
pub rewound_count: usize,
pub target_preview: String,
pub new_turn_index: i32,
pub session_id: String,
}
pub async fn rewind_session(
State(state): State<Arc<AppState>>,
Path(session_id): Path<String>,
Json(req): Json<RewindRequest>,
) -> ApiResult<Json<RewindResponse>> {
let result = if let Some(msg_id) = req.message_id {
crate::agent::runtime::session::rewind_to_message(&state.db, &session_id, msg_id)
.await
.map_err(|e| AppError::bad_request(e.to_string()))?
} else {
let n = req.n.unwrap_or(1);
crate::agent::runtime::session::rewind_n_turns(&state.db, &session_id, n)
.await
.map_err(|e| AppError::bad_request(e.to_string()))?
};
Ok(Json(RewindResponse {
rewound_count: result.rewound_count,
target_preview: result.target_preview,
new_turn_index: result.new_turn_index,
session_id: session_id.clone(),
}))
}
// ── POST /api/chat/sessions/:id/rewind/restore ──
// 恢复最近一次回退undo-of-undo
// 仅当回退后未产生新对话时才可恢复。
#[derive(Debug, Serialize)]
pub struct RestoreResponse {
pub restored_count: usize,
pub session_id: String,
}
pub async fn restore_rewound_session(
State(state): State<Arc<AppState>>,
Path(session_id): Path<String>,
) -> ApiResult<Json<RestoreResponse>> {
let count = crate::agent::runtime::session::restore_rewound(&state.db, &session_id)
.await
.map_err(|e| AppError::conflict(e.to_string()))?;
Ok(Json(RestoreResponse {
restored_count: count,
session_id,
}))
}