feat: ReAct Agent 自主科研引擎、LLM 流式工具调用、论文搜索服务化
核心新增 —— ReAct 智能体引擎 (src/agent/): - ReAct 主循环:自主推理→工具调用→观察→迭代,支持流式 SSE 推送、 上下文自动压缩、重复调用死循环检测、手动取消与超时控制 - 7 个内置工具:ADS/arXiv 论文搜索、元数据查询、全文读取、PDF/HTML 下载解析、RAG 向量语义检索、CDS Sesame 天体目标查询 LLM 客户端增强 (src/clients/llm.rs): - 新增流式/非流式 Function Calling (tool calling) 支持 - 原生 reasoning_content 推理链字段 (DeepSeek/QwQ) - 多模态图片+文本对话、TokenUsage 用量统计 搜索服务重构 (src/services/search.rs): - 抽取 ADS/arXiv 搜索为统一 service,跨平台去重、本地状态合并 - API handler 瘦身至薄封装层 下载器增强 (src/services/download.rs): - Obscura 无头浏览器反爬下载通道 (进程内/CLI 双模式) - arXiv/Springer 专用策略、ADS Scan PDF 兜底、CAPTCHA 绕过 解析器增强 (src/services/parser/mod.rs): - Markdown 输出自动附加 YAML 前端元数据头 (标题/作者/来源) 新增 API (src/api/agent.rs): - POST /api/chat/agent — SSE 流式智能体对话 - GET/DELETE /api/chat/sessions — 会话管理 CRUD - POST /api/chat/sessions/:id/stop — 紧急停止 数据库: agent_sessions + agent_messages 表 前端 (dashboard/): - 全新"智能科研"标签页 (ResearchAgentPanel) — 会话列表、ReAct 步骤时间线可视化、SSE 实时流式渲染、Markdown+KaTeX 答案展示 - Reader AI 助手集成 Agent 工作流,含工具调用步骤展示与引用溯源 - 侧边栏新增"智能科研"导航入口
This commit is contained in:
@@ -0,0 +1,244 @@
|
||||
// src/api/agent.rs
|
||||
//
|
||||
// 科研智能体 API 控制器。
|
||||
// 提供 SSE 流式对话接口和会话管理 CRUD 接口。
|
||||
|
||||
use axum::{
|
||||
extract::{Path, Query, State},
|
||||
http::StatusCode,
|
||||
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::{info, error};
|
||||
|
||||
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>,
|
||||
}
|
||||
|
||||
pub async fn chat_agent(
|
||||
State(state): State<Arc<AppState>>,
|
||||
Json(req): Json<AgentChatRequest>,
|
||||
) -> Result<Sse<impl Stream<Item = Result<Event, Infallible>>>, (StatusCode, String)> {
|
||||
info!("接收到智能体对话请求: question='{}', session_id={:?}", req.question, req.session_id);
|
||||
|
||||
let runtime = AgentRuntime::new(Arc::clone(&state));
|
||||
let (tx, mut rx) = tokio::sync::mpsc::unbounded_channel::<AgentStreamEvent>();
|
||||
|
||||
let question = req.question.clone();
|
||||
let session_id = req.session_id.clone();
|
||||
|
||||
// 在后台 tokio 任务中执行 Agent 循环
|
||||
tokio::spawn(async move {
|
||||
match runtime.run_turn(session_id, &question, 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 事件流
|
||||
let stream = async_stream::stream! {
|
||||
while let Some(event) = rx.recv().await {
|
||||
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(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 turn_count: i32,
|
||||
pub created_at: String,
|
||||
pub updated_at: String,
|
||||
}
|
||||
|
||||
pub async fn list_sessions(
|
||||
State(state): State<Arc<AppState>>,
|
||||
Query(params): Query<SessionListParams>,
|
||||
) -> Result<Json<Vec<SessionSummary>>, (StatusCode, String)> {
|
||||
let limit = params.limit.unwrap_or(50);
|
||||
let offset = params.offset.unwrap_or(0);
|
||||
|
||||
let rows = sqlx::query(
|
||||
"SELECT session_id, title, model, turn_count, 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| (StatusCode::INTERNAL_SERVER_ERROR, format!("查询会话列表失败: {}", e)))?;
|
||||
|
||||
let sessions: Vec<SessionSummary> = rows.iter().map(|r| {
|
||||
SessionSummary {
|
||||
session_id: r.get(0),
|
||||
title: r.get(1),
|
||||
model: r.get(2),
|
||||
turn_count: r.get(3),
|
||||
created_at: r.get(4),
|
||||
updated_at: r.get(5),
|
||||
}
|
||||
}).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 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>,
|
||||
) -> Result<Json<SessionDetail>, (StatusCode, String)> {
|
||||
// 查询会话元信息
|
||||
let session_row = sqlx::query(
|
||||
"SELECT session_id, title, model, turn_count, 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| (StatusCode::INTERNAL_SERVER_ERROR, format!("查询会话失败: {}", e)))?
|
||||
.ok_or((StatusCode::NOT_FOUND, format!("会话 {} 不存在", session_id)))?;
|
||||
|
||||
let session = SessionSummary {
|
||||
session_id: session_row.get(0),
|
||||
title: session_row.get(1),
|
||||
model: session_row.get(2),
|
||||
turn_count: session_row.get(3),
|
||||
created_at: session_row.get(4),
|
||||
updated_at: session_row.get(5),
|
||||
};
|
||||
|
||||
// 查询消息列表
|
||||
let msg_rows = sqlx::query(
|
||||
"SELECT id, turn_index, step_index, role, content, thought, tool_calls, tool_call_id, token_count, metadata, created_at \
|
||||
FROM agent_messages \
|
||||
WHERE session_id = ? \
|
||||
ORDER BY id ASC"
|
||||
)
|
||||
.bind(&session_id)
|
||||
.fetch_all(&state.db)
|
||||
.await
|
||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, format!("查询消息列表失败: {}", e)))?;
|
||||
|
||||
let messages: Vec<MessageRecord> = msg_rows.iter().map(|r| {
|
||||
let tool_calls_json: Option<String> = r.get(6);
|
||||
let metadata_json: Option<String> = r.get(9);
|
||||
|
||||
MessageRecord {
|
||||
id: r.get(0),
|
||||
turn_index: r.get(1),
|
||||
step_index: r.get(2),
|
||||
role: r.get(3),
|
||||
content: r.get(4),
|
||||
thought: r.get(5),
|
||||
tool_calls: tool_calls_json.and_then(|s| serde_json::from_str(&s).ok()),
|
||||
tool_call_id: r.get(7),
|
||||
token_count: r.get(8),
|
||||
metadata: metadata_json.and_then(|s| serde_json::from_str(&s).ok()),
|
||||
created_at: r.get(10),
|
||||
}
|
||||
}).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>,
|
||||
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
|
||||
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| (StatusCode::INTERNAL_SERVER_ERROR, format!("删除会话失败: {}", e)))?;
|
||||
|
||||
if result.rows_affected() == 0 {
|
||||
return Err((StatusCode::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>,
|
||||
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
|
||||
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 })))
|
||||
}
|
||||
Reference in New Issue
Block a user