// 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, /// Agent 运行模式: "default" / "deep-research" / "literature-reader" #[serde(default = "default_mode")] pub mode: String, /// 是否启用 LLM 思考模式。None = 由 mode 决定,Some(true/false) = 用户显式覆盖。 #[serde(default)] pub thinking: Option, /// 可选的图片附件(base64 编码 + MIME 类型) #[serde(default)] pub image: Option, } /// 用户附带的图片,用于多模态 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, } 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> { 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>, Json(req): Json, ) -> ApiResult>>> { 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, Option) = 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::(); 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, pub offset: Option, } #[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, pub created_at: String, pub updated_at: String, } pub async fn list_sessions( State(state): State>, Query(params): Query, ) -> ApiResult>> { 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 = 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, } #[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, pub tool_calls: Option, pub tool_call_id: Option, pub token_count: i32, pub metadata: Option, pub created_at: String, } pub async fn get_session( State(state): State>, Path(session_id): Path, ) -> ApiResult> { // 查询会话元信息 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 = msg_rows .iter() .map(|r| { let tool_calls_json: Option = r.get(7); let metadata_json: Option = 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>, Path(session_id): Path, ) -> ApiResult> { 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>, Path(session_id): Path, ) -> ApiResult> { 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>, ) -> ApiResult> { // 总会话数 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, pub status: String, pub elapsed_ms: i32, pub output_preview: Option, pub created_at: String, } pub async fn get_session_audit( State(state): State>, Path(session_id): Path, ) -> ApiResult>> { 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 = 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, pub free_text: Option, } pub async fn answer_question( State(state): State>, Json(req): Json, ) -> ApiResult> { 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>, ) -> Json> { let pending = match state.pending_questions.lock() { Ok(p) => p, Err(_) => return Json(Vec::new()), }; let questions: Vec = pending .iter() .map(|(id, pq)| { serde_json::from_str::(&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>, Path(session_id): Path, Json(req): Json, ) -> ApiResult> { 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>, ) -> Json> { let perms = match state.pending_permissions.lock() { Ok(p) => p, Err(_) => return Json(Vec::new()), }; let result: Vec = 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>, Path(session_id): Path, ) -> ApiResult> { 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, } pub async fn retry_session( State(state): State>, Path(session_id): Path, ) -> ApiResult> { 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, /// 或者指定回退到的消息 ID pub message_id: Option, } #[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>, Path(session_id): Path, Json(req): Json, ) -> ApiResult> { 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>, Path(session_id): Path, ) -> ApiResult> { 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, })) }