后端: - 权限系统重写: 全局→按 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 状态批处理
816 lines
26 KiB
Rust
816 lines
26 KiB
Rust
// 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,
|
||
}))
|
||
}
|