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:
fmq
2026-06-15 17:30:45 +08:00
parent 22e7e1dcee
commit b1fb884f21
23 changed files with 4174 additions and 347 deletions
+6
View File
@@ -0,0 +1,6 @@
// src/agent/mod.rs
// 科研智能体模块
// 基于 ReAct 框架实现 Thought -> Action -> Observation 循环
pub mod tools;
pub mod runtime;
+786
View File
@@ -0,0 +1,786 @@
// src/agent/runtime.rs
//
// 科研智能体运行时核心模块。
// 实现 ReAct 循环:Thought -> Action (工具调用) -> Observation -> Thought...
// 支持会话持久化、上下文压缩、死循环检测和 SSE 流式输出。
use std::sync::Arc;
use tracing::{info, warn, error};
use serde::Serialize;
use sqlx::SqlitePool;
use tokio::sync::mpsc;
use crate::api::AppState;
use crate::clients::llm::{
ChatMessage, LlmClient, MessageRole, StreamEvent,
};
use super::tools::{ToolContext, ToolOutput, ToolRegistry};
/// Agent 配置参数
#[derive(Debug, Clone)]
pub struct AgentConfig {
/// 最大 ReAct 迭代次数
pub max_steps: usize,
/// 同质调用检测阈值(连续相同调用次数)
pub duplicate_call_threshold: usize,
/// 工具执行超时时间(秒)
pub tool_timeout_secs: u64,
/// 工具输出最大字符数
pub max_tool_output_chars: usize,
/// 上下文 Token 估算上限(触发自动摘要压缩)
pub context_char_limit: usize,
}
impl Default for AgentConfig {
fn default() -> Self {
AgentConfig {
max_steps: std::env::var("AGENT_MAX_STEPS")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(8),
duplicate_call_threshold: 3,
tool_timeout_secs: std::env::var("AGENT_TOOL_TIMEOUT_SECS")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(120),
max_tool_output_chars: std::env::var("AGENT_MAX_TOOL_OUTPUT_CHARS")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(4000),
context_char_limit: std::env::var("AGENT_CONTEXT_CHAR_LIMIT")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(16000),
}
}
}
/// SSE 流式事件(发送给前端)
#[derive(Debug, Clone, Serialize)]
#[serde(tag = "type")]
pub enum AgentStreamEvent {
/// 会话创建/恢复
#[serde(rename = "session")]
Session {
session_id: String,
title: String,
},
/// 智能体思考过程
#[serde(rename = "thought")]
Thought {
content: String,
step: usize,
},
/// 工具调用开始
#[serde(rename = "tool_call")]
ToolCall {
name: String,
arguments: serde_json::Value,
step: usize,
},
/// 工具执行结果(Observation
#[serde(rename = "tool_result")]
ToolResult {
name: String,
output: String,
is_error: bool,
metadata: serde_json::Value,
step: usize,
},
/// 文本增量流式输出(最终回答)
#[serde(rename = "text_delta")]
TextDelta {
content: String,
},
/// Token 使用统计
#[serde(rename = "usage")]
Usage {
prompt_tokens: u32,
completion_tokens: u32,
total_tokens: u32,
},
/// 错误通知
#[serde(rename = "error")]
Error {
message: String,
},
/// 完成标记
#[serde(rename = "done")]
Done,
}
/// 同质调用检测器
#[derive(Debug, Default)]
struct DuplicateDetector {
last_call: Option<(String, String)>, // (tool_name, arguments)
consecutive_count: usize,
}
impl DuplicateDetector {
/// 记录一次调用,返回是否检测到死循环
fn record(&mut self, tool_name: &str, arguments: &str, threshold: usize) -> bool {
let key = (tool_name.to_string(), arguments.to_string());
if self.last_call.as_ref() == Some(&key) {
self.consecutive_count += 1;
if self.consecutive_count >= threshold {
return true;
}
} else {
self.last_call = Some(key);
self.consecutive_count = 1;
}
false
}
}
/// 智能体运行时
pub struct AgentRuntime {
app_state: Arc<AppState>,
config: AgentConfig,
tool_registry: ToolRegistry,
}
impl AgentRuntime {
/// 创建新的运行时实例
pub fn new(app_state: Arc<AppState>) -> Self {
AgentRuntime {
app_state,
config: AgentConfig::default(),
tool_registry: ToolRegistry::new(),
}
}
/// 创建带自定义配置的运行时实例
pub fn with_config(app_state: Arc<AppState>, config: AgentConfig) -> Self {
AgentRuntime {
app_state,
config,
tool_registry: ToolRegistry::new(),
}
}
/// 执行完整的智能体对话回合(流式 SSE 输出)
///
/// 流程:
/// 1. 加载或创建会话
/// 2. 构建消息上下文
/// 3. ReAct 循环:LLM 调用 -> 工具执行 -> 结果注入 -> 再次调用 ...
/// 4. 最终回答流式输出
/// 5. 持久化所有消息
pub async fn run_turn(
&self,
session_id: Option<String>,
question: &str,
tx: mpsc::UnboundedSender<AgentStreamEvent>,
) -> anyhow::Result<String> {
let db = &self.app_state.db;
let llm = &self.app_state.llm;
// 1. 创建或恢复会话
let sid = match session_id {
Some(id) => {
// 验证会话存在
let exists: bool = sqlx::query_scalar(
"SELECT EXISTS(SELECT 1 FROM agent_sessions WHERE session_id = ? AND deleted_at IS NULL)"
)
.bind(&id)
.fetch_one(db)
.await
.unwrap_or(false);
if !exists {
return Err(anyhow::anyhow!("会话 {} 不存在或已删除", id));
}
id
}
None => {
let new_id = uuid::Uuid::new_v4().to_string();
sqlx::query(
"INSERT INTO agent_sessions (session_id, title, model) VALUES (?, ?, ?)"
)
.bind(&new_id)
.bind("")
.bind(llm.model())
.execute(db)
.await?;
new_id
}
};
let _ = tx.send(AgentStreamEvent::Session {
session_id: sid.clone(),
title: String::new(),
});
// 2. 加载历史消息(过滤掉 thought 字段,仅保留纯对话上下文)
let mut messages = self.load_history_for_llm(db, &sid).await?;
// 获取当前轮次号
let turn_index: i32 = sqlx::query_scalar(
"SELECT COALESCE(MAX(turn_index), -1) + 1 FROM agent_messages WHERE session_id = ?"
)
.bind(&sid)
.fetch_one(db)
.await
.unwrap_or(0);
// 3. 构建系统提示词
if messages.is_empty() || messages[0].role != MessageRole::System {
messages.insert(0, ChatMessage::system(self.system_prompt()));
}
// 4. 添加用户消息
messages.push(ChatMessage::user(question));
self.save_message(db, &sid, turn_index, 0, &ChatMessage::user(question), None).await?;
// 5. ReAct 循环
let tool_defs = self.tool_registry.definitions();
let tool_ctx = ToolContext {
app_state: Arc::clone(&self.app_state),
};
let mut duplicate_detector = DuplicateDetector::default();
let mut step = 0;
loop {
step += 1;
// 检查是否被用户手动中止
if let Ok(mut cancelled) = self.app_state.cancelled_runs.lock() {
if cancelled.remove(&sid) {
warn!("[AgentRuntime] 用户手动中止了会话 {} 的智能体执行", sid);
let _ = tx.send(AgentStreamEvent::Error {
message: "用户已手动中止执行。".to_string(),
});
break;
}
}
// 上下文安全检查
let context_chars: usize = messages.iter()
.filter_map(|m| m.content.as_ref())
.map(|c| c.len())
.sum();
if context_chars > self.config.context_char_limit {
info!("[AgentRuntime] 上下文超限 ({} > {}),触发压缩", context_chars, self.config.context_char_limit);
self.compress_context(&mut messages, llm).await;
}
// 调用 LLM(使用 `chat_stream` 实时流式读取,支持思维过程/最终回答的流式发送和中止检测)
let mut stream_rx = match llm.chat_stream(&messages, &tool_defs).await {
Ok(rx) => rx,
Err(e) => {
error!("[AgentRuntime] LLM stream 调用失败: {}", e);
let _ = tx.send(AgentStreamEvent::Error {
message: format!("大模型流式调用失败: {}", e),
});
break;
}
};
let mut accumulated_content = String::new();
let mut accumulated_reasoning = String::new();
let mut accumulated_tool_calls: Option<Vec<crate::clients::llm::ToolCall>> = None;
let mut usage: Option<crate::clients::llm::TokenUsage> = None;
let mut is_tool_call_step = false;
let cancel_fut = async {
loop {
tokio::time::sleep(std::time::Duration::from_millis(250)).await;
if let Ok(cancelled) = self.app_state.cancelled_runs.lock() {
if cancelled.contains(&sid) {
return;
}
}
}
};
enum StreamLoopResult {
Success,
Error(String),
Cancelled,
}
let stream_loop_res = {
let mut cancel_pinned = Box::pin(cancel_fut);
let mut error_msg = None;
let mut cancelled = false;
loop {
tokio::select! {
event_opt = stream_rx.recv() => {
match event_opt {
Some(event) => {
match event {
StreamEvent::ReasoningDelta(delta) => {
accumulated_reasoning.push_str(&delta);
// 实时流式发送思考过程给前端
let _ = tx.send(AgentStreamEvent::Thought {
content: accumulated_reasoning.clone(),
step,
});
}
StreamEvent::TextDelta(delta) => {
accumulated_content.push_str(&delta);
// 如果目前还没发现是工具调用步骤,就实时流式发送文本给前端作为最终回答
if !is_tool_call_step {
let _ = tx.send(AgentStreamEvent::TextDelta {
content: delta,
});
}
}
StreamEvent::ToolCallsComplete(tool_calls) => {
is_tool_call_step = true;
accumulated_tool_calls = Some(tool_calls);
}
StreamEvent::ToolCallDelta { .. } => {}
StreamEvent::Usage(u) => {
usage = Some(u);
}
StreamEvent::Done => {
break;
}
StreamEvent::Error(e) => {
error_msg = Some(e);
break;
}
}
}
None => break,
}
}
_ = &mut cancel_pinned => {
cancelled = true;
break;
}
}
}
if cancelled {
StreamLoopResult::Cancelled
} else if let Some(e) = error_msg {
StreamLoopResult::Error(e)
} else {
StreamLoopResult::Success
}
};
match stream_loop_res {
StreamLoopResult::Success => {}
StreamLoopResult::Error(e_str) => {
error!("[AgentRuntime] 流式读取错误: {}", e_str);
let _ = tx.send(AgentStreamEvent::Error {
message: format!("大模型流式读取失败: {}", e_str),
});
break;
}
StreamLoopResult::Cancelled => {
if let Ok(mut cancelled) = self.app_state.cancelled_runs.lock() {
cancelled.remove(&sid);
}
warn!("[AgentRuntime] 在流式调用期间被用户手动中止,会话 ID: {}", sid);
let _ = tx.send(AgentStreamEvent::Error {
message: "用户已手动中止执行。".to_string(),
});
break;
}
}
// 处理 Thought(优先使用原生推理内容,否则如果属于工具调用步骤,使用 accumulated_content 存储)
let mut thought_content = None;
if !accumulated_reasoning.is_empty() {
thought_content = Some(accumulated_reasoning.clone());
}
if thought_content.is_none() && is_tool_call_step {
if !accumulated_content.is_empty() {
// 有工具调用时,content 被视为 Thought
thought_content = Some(accumulated_content.clone());
}
}
// 如果是在工具调用步骤中产生的前言描述,而我们之前没实时以 Thought 发送过,此时统一作为 Thought 发送给前端展示
if is_tool_call_step {
if let Some(ref thought_text) = thought_content {
let _ = tx.send(AgentStreamEvent::Thought {
content: thought_text.clone(),
step,
});
}
}
let reasoning_option = if accumulated_reasoning.is_empty() { None } else { Some(accumulated_reasoning.clone()) };
// 无工具调用 = 最终回答
if accumulated_tool_calls.is_none() || accumulated_tool_calls.as_ref().unwrap().is_empty() {
// 保存助手最终回答消息
let assistant_msg = ChatMessage::assistant_with_reasoning(
Some(accumulated_content.clone()),
reasoning_option.clone(),
None,
);
self.save_message(db, &sid, turn_index, step as i32, &assistant_msg, reasoning_option.as_deref()).await?;
messages.push(assistant_msg);
// 如果有原生推理内容且之前没发送过,发送给前端展示最终思维链
if let Some(ref thought_text) = reasoning_option {
if thought_content.is_none() {
let _ = tx.send(AgentStreamEvent::Thought {
content: thought_text.clone(),
step,
});
}
}
// 流式发送 Done 或 Token 消耗
if let Some(u) = usage {
let _ = tx.send(AgentStreamEvent::Usage {
prompt_tokens: u.prompt_tokens,
completion_tokens: u.completion_tokens,
total_tokens: u.total_tokens,
});
}
break;
}
let tool_calls = accumulated_tool_calls.unwrap();
// 有工具调用 —— 构建 assistant 消息(含 tool_calls 和 reasoning_content
let assistant_msg = ChatMessage::assistant_with_reasoning(
if accumulated_content.is_empty() { None } else { Some(accumulated_content.clone()) },
reasoning_option.clone(),
Some(tool_calls.clone()),
);
self.save_message(db, &sid, turn_index, step as i32, &assistant_msg, reasoning_option.as_deref()).await?;
messages.push(assistant_msg);
// 逐个执行工具
for tool_call in &tool_calls {
let tool_name = &tool_call.function.name;
let tool_args_str = &tool_call.function.arguments;
// 死循环检测
if duplicate_detector.record(tool_name, tool_args_str, self.config.duplicate_call_threshold) {
warn!("[AgentRuntime] 检测到死循环:{} 连续调用 {} 次", tool_name, self.config.duplicate_call_threshold);
let _ = tx.send(AgentStreamEvent::Error {
message: format!("检测到工具 {} 的重复调用,已自动终止循环。", tool_name),
});
// 注入错误 tool result 让 LLM 知道要停止
let error_msg = ChatMessage::tool_result(
&tool_call.id,
format!("错误:工具 {} 被连续重复调用 {} 次,参数完全相同。请停止重复调用并直接给出目前收集到的答案。", tool_name, self.config.duplicate_call_threshold),
);
messages.push(error_msg);
continue;
}
// 解析参数
let args: serde_json::Value = match serde_json::from_str(tool_args_str) {
Ok(v) => v,
Err(e) => {
let error_output = format!("工具参数 JSON 解析失败: {}", e);
let _ = tx.send(AgentStreamEvent::ToolResult {
name: tool_name.clone(),
output: error_output.clone(),
is_error: true,
metadata: serde_json::json!({}),
step,
});
let tool_msg = ChatMessage::tool_result(&tool_call.id, &error_output);
self.save_message(db, &sid, turn_index, step as i32, &tool_msg, None).await?;
messages.push(tool_msg);
continue;
}
};
// 发送工具调用事件
let _ = tx.send(AgentStreamEvent::ToolCall {
name: tool_name.clone(),
arguments: args.clone(),
step,
});
enum ToolResultEnum {
Success(ToolOutput),
Cancelled,
}
// 执行工具(带超时保护和手动中止检测)
let tool_res = match self.tool_registry.get(tool_name) {
Some(tool) => {
let timeout = std::time::Duration::from_secs(self.config.tool_timeout_secs);
let tool_fut = tool.execute(args, &tool_ctx);
let cancel_fut = async {
loop {
tokio::time::sleep(std::time::Duration::from_millis(250)).await;
if let Ok(cancelled) = self.app_state.cancelled_runs.lock() {
if cancelled.contains(&sid) {
return;
}
}
}
};
tokio::select! {
res = tokio::time::timeout(timeout, tool_fut) => {
match res {
Ok(output) => ToolResultEnum::Success(output),
Err(_) => ToolResultEnum::Success(ToolOutput::error(format!("工具 {} 执行超时({}秒)", tool_name, self.config.tool_timeout_secs))),
}
}
_ = cancel_fut => {
ToolResultEnum::Cancelled
}
}
}
None => ToolResultEnum::Success(ToolOutput::error(format!("未知工具: {}", tool_name))),
};
let output = match tool_res {
ToolResultEnum::Success(out) => out,
ToolResultEnum::Cancelled => {
if let Ok(mut cancelled) = self.app_state.cancelled_runs.lock() {
cancelled.remove(&sid);
}
warn!("[AgentRuntime] 在工具 {} 执行期间被用户手动中止,会话 ID: {}", tool_name, sid);
let _ = tx.send(AgentStreamEvent::Error {
message: "用户已手动中止执行。".to_string(),
});
break;
}
};
// 发送工具结果事件
let _ = tx.send(AgentStreamEvent::ToolResult {
name: tool_name.clone(),
output: output.content.clone(),
is_error: output.is_error,
metadata: output.metadata.clone(),
step,
});
// 截断工具输出
let truncated_content = if output.content.len() > self.config.max_tool_output_chars {
let truncated: String = output.content.chars().take(self.config.max_tool_output_chars).collect();
format!("{}...\n[输出已截断,原始长度: {} 字符]", truncated, output.content.len())
} else {
output.content.clone()
};
// 构建 tool result 消息
let tool_msg = ChatMessage::tool_result(&tool_call.id, &truncated_content);
self.save_message(db, &sid, turn_index, step as i32, &tool_msg, None).await?;
messages.push(tool_msg);
}
}
// 6. 更新会话元信息
let new_turn_count: i32 = sqlx::query_scalar(
"SELECT COUNT(DISTINCT turn_index) FROM agent_messages WHERE session_id = ?"
)
.bind(&sid)
.fetch_one(db)
.await
.unwrap_or(0);
// 首轮自动生成标题
if new_turn_count <= 1 {
let title = self.generate_title(question);
sqlx::query("UPDATE agent_sessions SET title = ?, turn_count = ?, updated_at = CURRENT_TIMESTAMP WHERE session_id = ?")
.bind(&title)
.bind(new_turn_count)
.bind(&sid)
.execute(db)
.await?;
} else {
sqlx::query("UPDATE agent_sessions SET turn_count = ?, updated_at = CURRENT_TIMESTAMP WHERE session_id = ?")
.bind(new_turn_count)
.bind(&sid)
.execute(db)
.await?;
}
let _ = tx.send(AgentStreamEvent::Done);
Ok(sid)
}
/// 系统提示词
fn system_prompt(&self) -> String {
"你是一位专业的天体物理学研究助手,具备丰富的天文学知识。你可以使用以下工具帮助用户进行科研工作:\n\
\n\
- search_papers: 检索天文学文献(ADS/arXiv\n\
- get_paper_content: 获取文献全文内容(自动下载、解析)\n\
- read_local_file: 快速读取已解析的本地文献\n\
- rag_search: 在已向量化的文献库中进行语义检索\n\
- query_target: 查询天体物理属性(坐标、光谱型等)\n\
\n\
请遵循以下原则:\n\
1. 先思考用户的问题需要什么信息,再决定调用哪些工具。\n\
2. 优先使用已有的本地文献资源(read_local_file / rag_search),必要时再检索新文献。\n\
3. 回答时引用具体文献来源,使用 ADS bibcode 标注。\n\
4. 对于数学公式,使用标准 LaTeX 格式。\n\
5. 用中文回答用户的问题,但保持科学术语的准确性(可附带英文原文)。\n\
6. 如果一个工具调用失败,不要重复使用完全相同的参数重试,尝试换一种方式。".to_string()
}
/// 从数据库加载历史消息(包含 thought 作为 reasoning_content,以备原生思考模型使用)
async fn load_history_for_llm(
&self,
db: &SqlitePool,
session_id: &str,
) -> anyhow::Result<Vec<ChatMessage>> {
let rows: Vec<(String, String, Option<String>, Option<String>, Option<String>)> = sqlx::query_as(
"SELECT role, content, tool_calls, tool_call_id, thought FROM agent_messages \
WHERE session_id = ? ORDER BY id ASC"
)
.bind(session_id)
.fetch_all(db)
.await?;
let mut messages = Vec::new();
for (role_str, content, tool_calls_json, tool_call_id, _thought) in rows {
let role = match role_str.as_str() {
"system" => MessageRole::System,
"user" => MessageRole::User,
"assistant" => MessageRole::Assistant,
"tool" => MessageRole::Tool,
_ => continue,
};
let tool_calls: Option<Vec<crate::clients::llm::ToolCall>> = tool_calls_json
.and_then(|json_str| serde_json::from_str(&json_str).ok());
messages.push(ChatMessage {
role,
content: if content.is_empty() { None } else { Some(content) },
tool_call_id,
tool_calls,
name: None,
reasoning_content: None,
});
}
Ok(messages)
}
/// 保存消息到数据库
async fn save_message(
&self,
db: &SqlitePool,
session_id: &str,
turn_index: i32,
step_index: i32,
msg: &ChatMessage,
thought: Option<&str>,
) -> anyhow::Result<()> {
let role = match msg.role {
MessageRole::System => "system",
MessageRole::User => "user",
MessageRole::Assistant => "assistant",
MessageRole::Tool => "tool",
};
let content = msg.content.as_deref().unwrap_or("");
let tool_calls_json = msg.tool_calls.as_ref()
.map(|tc| serde_json::to_string(tc).unwrap_or_default());
let tool_call_id = msg.tool_call_id.as_deref();
let token_count = content.len() as i32 / 4; // 粗略估算
sqlx::query(
"INSERT INTO agent_messages (session_id, turn_index, step_index, role, content, thought, tool_calls, tool_call_id, token_count) \
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)"
)
.bind(session_id)
.bind(turn_index)
.bind(step_index)
.bind(role)
.bind(content)
.bind(thought)
.bind(&tool_calls_json)
.bind(tool_call_id)
.bind(token_count)
.execute(db)
.await?;
Ok(())
}
/// 上下文压缩:保留系统提示、最近的 user 消息、以及最近的 tool_calls/tool 对
async fn compress_context(
&self,
messages: &mut Vec<ChatMessage>,
llm: &LlmClient,
) {
if messages.len() <= 4 {
return;
}
// 保留系统消息
let system_msg = messages.first().cloned();
// 找到安全切割点:必须保证 assistant(tool_calls) 和后续 tool(result) 不被切断
// 策略:保留最近 6 条消息 + 系统消息
let keep_count = 6.min(messages.len() - 1);
let to_summarize = &messages[1..messages.len() - keep_count];
if to_summarize.is_empty() {
return;
}
// 生成摘要
let summary_content: String = to_summarize.iter()
.filter_map(|m| {
let role = match m.role {
MessageRole::User => "用户",
MessageRole::Assistant => "助手",
MessageRole::Tool => "工具",
_ => return None,
};
m.content.as_ref().map(|c| {
let preview: String = c.chars().take(200).collect();
format!("[{}] {}", role, preview)
})
})
.collect::<Vec<_>>()
.join("\n");
let summary_prompt = format!(
"请用简洁的中文总结以下对话历史的要点(不超过500字):\n\n{}",
summary_content
);
let summary = match llm.chat_completion(
"你是一个对话摘要助手。请提取对话的关键信息和结论。",
&summary_prompt,
).await {
Ok(s) => s,
Err(e) => {
warn!("[AgentRuntime] 上下文摘要生成失败: {},回退为简单截断", e);
format!("[历史摘要] 此前进行了 {} 轮对话交互", to_summarize.len())
}
};
// 重建消息列表
let recent = messages[messages.len() - keep_count..].to_vec();
messages.clear();
if let Some(sys) = system_msg {
messages.push(sys);
}
messages.push(ChatMessage::user(format!("[历史对话摘要]\n{}", summary)));
messages.extend(recent);
info!("[AgentRuntime] 上下文压缩完成,消息数: {}", messages.len());
}
/// 根据用户首条问题生成会话标题
fn generate_title(&self, question: &str) -> String {
let chars: String = question.chars().take(50).collect();
if question.len() > 50 {
format!("{}...", chars)
} else {
chars
}
}
}
+659
View File
@@ -0,0 +1,659 @@
// src/agent/tools.rs
//
// 科研智能体工具集定义与实现。
// 每个工具遵循 AgentTool trait,向大模型声明 JSON Schema 参数定义,
// 并在 execute 中调用已有的服务层完成实际业务操作。
use async_trait::async_trait;
use serde_json::json;
use std::sync::Arc;
use tracing::{info, error};
use crate::api::AppState;
use crate::clients::llm::ToolDefinition;
/// 工具执行上下文,封装全局共享状态
pub struct ToolContext {
pub app_state: Arc<AppState>,
}
/// 工具执行结果
#[derive(Debug, Clone)]
pub struct ToolOutput {
/// 给大模型阅读的截断文本
pub content: String,
/// 是否为错误
pub is_error: bool,
/// 结构化元数据(给前端 Timeline 直接渲染)
pub metadata: serde_json::Value,
}
impl ToolOutput {
/// 创建成功结果
pub fn success(content: impl Into<String>, metadata: serde_json::Value) -> Self {
ToolOutput {
content: content.into(),
is_error: false,
metadata,
}
}
/// 创建错误结果
pub fn error(msg: impl Into<String>) -> Self {
ToolOutput {
content: msg.into(),
is_error: true,
metadata: json!({}),
}
}
}
/// 智能体工具 trait
#[async_trait]
pub trait AgentTool: Send + Sync {
/// 工具名称(与 LLM function calling 的 name 保持一致)
fn name(&self) -> &str;
/// 工具描述(告知 LLM 何时应该调用该工具)
fn description(&self) -> &str;
/// JSON Schema 格式的参数定义
fn parameters(&self) -> serde_json::Value;
/// 执行工具逻辑
async fn execute(&self, args: serde_json::Value, ctx: &ToolContext) -> ToolOutput;
}
/// 工具注册表,管理所有可用工具
pub struct ToolRegistry {
tools: Vec<Box<dyn AgentTool>>,
}
impl ToolRegistry {
/// 创建默认工具注册表(包含全部科研工具)
pub fn new() -> Self {
let tools: Vec<Box<dyn AgentTool>> = vec![
Box::new(SearchPapersTool),
Box::new(GetPaperMetadataTool),
Box::new(DownloadPaperTool),
Box::new(ParsePaperTool),
Box::new(GetPaperContentTool),
Box::new(RagSearchTool),
Box::new(QueryTargetTool),
];
ToolRegistry { tools }
}
/// 根据名称查找工具
pub fn get(&self, name: &str) -> Option<&dyn AgentTool> {
self.tools.iter().find(|t| t.name() == name).map(|t| t.as_ref())
}
/// 生成所有工具的 ToolDefinition 列表(用于发送给 LLM
pub fn definitions(&self) -> Vec<ToolDefinition> {
self.tools.iter().map(|t| {
ToolDefinition::new(t.name(), t.description(), t.parameters())
}).collect()
}
}
/// 截断文本到指定最大字符数
fn truncate_content(s: &str, max_chars: usize) -> String {
if s.len() <= max_chars {
s.to_string()
} else {
let truncated: String = s.chars().take(max_chars).collect();
format!("{}\n\n[... 内容已截断,共 {} 字符 ...]", truncated, s.len())
}
}
// ────────────────────────── 1. SearchPapersTool ──────────────────────────
/// 文献搜索工具:调用 ADS/arXiv 进行跨库检索
pub struct SearchPapersTool;
#[async_trait]
impl AgentTool for SearchPapersTool {
fn name(&self) -> &str { "search_papers" }
fn description(&self) -> &str {
"搜索天文学文献。支持 NASA ADS 和 arXiv 跨平台联合检索,结果自动合并去重,并关联本地馆藏状态与引用关系网。输入关键词或高级检索式,返回匹配的文献列表。\
适用于:查找相关文献、了解研究领域现状、获取特定主题的论文。"
}
fn parameters(&self) -> serde_json::Value {
json!({
"type": "object",
"properties": {
"query": {
"type": "string",
"description": "搜索关键词或高级检索式。支持语法:\
1. 字段限定:au:\"作者\" 或 author:\"作者\"、ti:\"标题\" 或 title:\"标题\"、abs:\"摘要关键字\"\
2. 年份限定:year:2020(单年)或 year:2020-2025(年份区间);\
3. 逻辑运算:支持 AND、OR、NOT 逻辑组合及括号分组,如 '(ti:subdwarf OR ti:\"white dwarf\") AND year:2020-2025'\
4. 短语匹配:用双引号 \"\" 包含精确匹配短语,如 '\"Gaia BH1\"'。\
所有的中文标点符号(如“”(),;)在后台均会自动清洗转换。"
},
"rows": {
"type": "integer",
"description": "返回结果数量,默认5,最大20",
"default": 5
}
},
"required": ["query"]
})
}
async fn execute(&self, args: serde_json::Value, ctx: &ToolContext) -> ToolOutput {
let query = match args.get("query").and_then(|q| q.as_str()) {
Some(q) => q.to_string(),
None => return ToolOutput::error("缺少必需参数 'query'"),
};
let rows = args.get("rows").and_then(|r| r.as_i64()).unwrap_or(5).min(20) as i32;
info!("[SearchPapersTool] 执行文献搜索: query='{}', rows={}", query, rows);
let state = &ctx.app_state;
match crate::services::search::search_papers(state, &query, "all", 0, rows, "relevance").await {
Ok(results) => {
if results.is_empty() {
return ToolOutput::success("未找到匹配的文献。请尝试调整搜索关键词。", json!({ "count": 0 }));
}
// 格式化结果(保留详细信息给 LLM,但不含摘要且不作截断)
let display_results: Vec<serde_json::Value> = results.iter().map(|p| {
let first_author = p.authors.first().cloned().unwrap_or_else(|| "未知".to_string());
json!({
"bibcode": p.bibcode,
"title": p.title,
"first_author": first_author,
"year": p.year,
"citation_count": p.citation_count,
})
}).collect();
let content = display_results.iter().enumerate().map(|(i, r)| {
format!(
"{}. [{}] {} ({})\n 第一作者: {}\n 被引: {}",
i + 1,
r["bibcode"].as_str().unwrap_or(""),
r["title"].as_str().unwrap_or(""),
r["year"].as_str().unwrap_or(""),
r["first_author"].as_str().unwrap_or("未知"),
r["citation_count"].as_i64().unwrap_or(0)
)
}).collect::<Vec<_>>().join("\n\n");
ToolOutput::success(
content,
json!({
"count": results.len(),
"papers": display_results
})
)
}
Err(e) => {
error!("[SearchPapersTool] 检索失败: {}", e);
ToolOutput::error(format!("文献检索失败: {}", e))
}
}
}
}
// ────────────────────────── 1b. GetPaperMetadataTool ──────────────────────────
/// 获取文献元数据工具:获取指定文献的完整元数据(包含完整标题、所有作者、出版期刊、关键字、引用数、完整摘要等)
pub struct GetPaperMetadataTool;
#[async_trait]
impl AgentTool for GetPaperMetadataTool {
fn name(&self) -> &str { "get_paper_metadata" }
fn description(&self) -> &str {
"获取指定文献的完整元数据信息(包括完整标题、所有作者、出版期刊、关键字、引用数、完整摘要等)。\
适用于:需要查看某篇文献的详细信息、阅读完整摘要以评估文献相关性。"
}
fn parameters(&self) -> serde_json::Value {
json!({
"type": "object",
"properties": {
"bibcode": {
"type": "string",
"description": "文献的唯一标识符,支持 ADS Bibcode(如 '2024ApJ...960..123A')、DOI(如 '10.3847/1538-4357/ad0c5a')或 arXiv ID(如 '2401.12345'"
}
},
"required": ["bibcode"]
})
}
async fn execute(&self, args: serde_json::Value, ctx: &ToolContext) -> ToolOutput {
let bibcode = match args.get("bibcode").and_then(|b| b.as_str()) {
Some(b) => b.to_string(),
None => return ToolOutput::error("缺少必需参数 'bibcode'"),
};
info!("[GetPaperMetadataTool] 获取文献元数据: {}", bibcode);
let state = &ctx.app_state;
match crate::api::helpers::get_paper_from_db(&state.db, &state.config.library_dir, &bibcode).await {
Ok(paper) => {
let content = format!(
"文献元数据 [{}]\n\
标题: {}\n\
作者: {}\n\
年份: {}\n\
期刊: {}\n\
关键字: {}\n\
DOI: {}\n\
arXiv ID: {}\n\
引用数: {}\n\
参考文献数: {}\n\
文献类型: {}\n\
已下载: {}\n\
已解析为 Markdown: {}\n\
摘要:\n{}",
paper.bibcode,
paper.title,
paper.authors.join(", "),
paper.year,
paper.pub_journal,
paper.keywords.join(", "),
paper.doi,
paper.arxiv_id,
paper.citation_count,
paper.reference_count,
paper.doctype,
paper.is_downloaded,
paper.has_markdown,
paper.abstract_text
);
ToolOutput::success(content, json!(paper))
}
Err(e) => ToolOutput::error(format!("获取文献 {} 元数据失败: {}", bibcode, e)),
}
}
}
// ────────────────────────── 2. GetPaperContentTool ──────────────────────────
/// 获取文献内容工具:仅从本地读取并获取已解析的文献 Markdown 全文内容
pub struct GetPaperContentTool;
#[async_trait]
impl AgentTool for GetPaperContentTool {
fn name(&self) -> &str { "get_paper_content" }
fn description(&self) -> &str {
"读取并在本地库中获取已解析的文献 Markdown 完整文本内容。\
注意:本工具仅能读取已在数据库注册且已解析的文献,不会自动触发下载或解析。若文献未下载或未解析,本工具会返回详细指引,提示先依次调用 download_paper 和 parse_paper。"
}
fn parameters(&self) -> serde_json::Value {
json!({
"type": "object",
"properties": {
"bibcode": {
"type": "string",
"description": "文献的唯一标识符,支持 ADS Bibcode(如 '2024ApJ...960..123A')、DOI(如 '10.3847/1538-4357/ad0c5a')或 arXiv ID(如 '2401.12345'"
}
},
"required": ["bibcode"]
})
}
async fn execute(&self, args: serde_json::Value, ctx: &ToolContext) -> ToolOutput {
let bibcode = match args.get("bibcode").and_then(|b| b.as_str()) {
Some(b) => b.to_string(),
None => return ToolOutput::error("缺少必需参数 'bibcode'"),
};
info!("[GetPaperContentTool] 获取文献内容: {}", bibcode);
let state = &ctx.app_state;
let paths = crate::api::helpers::check_paper_paths_in_db(&state.db, &state.config.library_dir, &bibcode).await;
let md_opt = match paths {
Ok(Some((_, _, md_opt, _))) => md_opt,
Ok(None) => return ToolOutput::error(format!("获取文献内容失败:该文献未在本地数据库中注册,请先使用 search_papers 搜索该文献。")),
Err(e) => return ToolOutput::error(format!("获取文献内容失败: {}", e)),
};
let md_rel = match md_opt {
Some(rel) => rel,
None => return ToolOutput::error(format!("获取文献内容失败:该文献尚未完成结构化解析。如果未下载,请先调用 download_paper;如果已下载,请先调用 parse_paper 进行解析。")),
};
let md_abs = state.config.library_dir.join(&md_rel);
if !md_abs.exists() {
return ToolOutput::error(format!("获取文献内容失败:文献本地 Markdown 文件已丢失,请重新调用 parse_paper 进行解析。"));
}
match std::fs::read_to_string(&md_abs) {
Ok(content) => ToolOutput::success(
content.clone(),
json!({ "bibcode": bibcode, "chars": content.len() })
),
Err(e) => ToolOutput::error(format!("获取文献内容失败,读取本地文件错误: {}", e)),
}
}
}
// ────────────────────────── 2a. DownloadPaperTool ──────────────────────────
/// 下载文献全文资源工具:仅下载文献全文资源(PDF/HTML)至本地图书馆
pub struct DownloadPaperTool;
#[async_trait]
impl AgentTool for DownloadPaperTool {
fn name(&self) -> &str { "download_paper" }
fn description(&self) -> &str {
"下载指定文献的全文资源(PDF 或 HTML)至本地图书馆,为后续的结构化解析做好准备。\
适用于:需要阅读或分析新搜寻到的、尚未下载的文献。"
}
fn parameters(&self) -> serde_json::Value {
json!({
"type": "object",
"properties": {
"bibcode": {
"type": "string",
"description": "文献的唯一标识符,支持 ADS Bibcode(如 '2024ApJ...960..123A')、DOI(如 '10.3847/1538-4357/ad0c5a')或 arXiv ID(如 '2401.12345'"
},
"force": {
"type": "boolean",
"description": "是否强制重新下载(即使本地已下载该文献)"
}
},
"required": ["bibcode"]
})
}
async fn execute(&self, args: serde_json::Value, ctx: &ToolContext) -> ToolOutput {
let bibcode = match args.get("bibcode").and_then(|b| b.as_str()) {
Some(b) => b.to_string(),
None => return ToolOutput::error("缺少必需参数 'bibcode'"),
};
let force = args.get("force").and_then(|f| f.as_bool()).unwrap_or(false);
info!("[DownloadPaperTool] 下载文献全文资源: {}, 强制重下: {}", bibcode, force);
let state = &ctx.app_state;
match state.downloader.download_paper_service(
&state.db,
&state.config.library_dir,
&bibcode,
force,
)
.await {
Ok(paper) => ToolOutput::success(
format!("文献 {} 全文资源下载成功。格式 - PDF: {}, HTML: {}", bibcode, paper.has_pdf, paper.has_html),
json!({ "bibcode": bibcode, "has_pdf": paper.has_pdf, "has_html": paper.has_html })
),
Err(e) => ToolOutput::error(format!("文献 {} 下载失败: {}", bibcode, e)),
}
}
}
// ────────────────────────── 2b. ParsePaperTool ──────────────────────────
/// 结构化解析文献内容工具:仅对已下载的物理资源进行结构化解析生成 Markdown
pub struct ParsePaperTool;
#[async_trait]
impl AgentTool for ParsePaperTool {
fn name(&self) -> &str { "parse_paper" }
fn description(&self) -> &str {
"将指定文献本地已下载的 HTML 或 PDF 资源解析为结构化的 Markdown 文本,并保存至本地 Markdown 文件夹。\
注意:调用此工具前必须确保文献已被成功下载(已执行 download_paper)。"
}
fn parameters(&self) -> serde_json::Value {
json!({
"type": "object",
"properties": {
"bibcode": {
"type": "string",
"description": "文献的唯一标识符,支持 ADS Bibcode(如 '2024ApJ...960..123A')、DOI(如 '10.3847/1538-4357/ad0c5a')或 arXiv ID(如 '2401.12345'"
},
"force": {
"type": "boolean",
"description": "是否强制重新解析(即使本地已解析过该文献)"
}
},
"required": ["bibcode"]
})
}
async fn execute(&self, args: serde_json::Value, ctx: &ToolContext) -> ToolOutput {
let bibcode = match args.get("bibcode").and_then(|b| b.as_str()) {
Some(b) => b.to_string(),
None => return ToolOutput::error("缺少必需参数 'bibcode'"),
};
let force = args.get("force").and_then(|f| f.as_bool()).unwrap_or(false);
info!("[ParsePaperTool] 结构化解析文献内容: {}, 强制重析: {}", bibcode, force);
let state = &ctx.app_state;
match crate::services::parser::parse_paper_service(
&state.db,
&state.config.library_dir,
&state.qiniu,
&state.config,
&bibcode,
force,
)
.await {
Ok(markdown) => ToolOutput::success(
format!("文献 {} 结构化解析成功。解析后 Markdown 字符总数: {}", bibcode, markdown.len()),
json!({ "bibcode": bibcode, "chars": markdown.len() })
),
Err(e) => {
let msg = e.to_string();
if msg.contains("请先下载") {
ToolOutput::error(format!("文献 {} 解析失败:未检测到已下载的本地资源文件,请先调用 download_paper 工具进行下载。", bibcode))
} else {
ToolOutput::error(format!("文献 {} 解析失败: {}", bibcode, msg))
}
}
}
}
}
// ────────────────────────── 4. RagSearchTool ──────────────────────────
/// RAG 向量检索工具:基于语义相似度检索文献切片
pub struct RagSearchTool;
#[async_trait]
impl AgentTool for RagSearchTool {
fn name(&self) -> &str { "rag_search" }
fn description(&self) -> &str {
"在已向量化的文献库中进行语义检索。输入自然语言问题,返回最相关的文献片段。\
适用于:跨多篇文献查找特定信息、回答需要综合多个来源的问题。要求文献已完成向量化(embed)。"
}
fn parameters(&self) -> serde_json::Value {
json!({
"type": "object",
"properties": {
"question": {
"type": "string",
"description": "用于语义检索的自然语言问题"
},
"top_k": {
"type": "integer",
"description": "返回最相关的片段数量,默认5",
"default": 5
}
},
"required": ["question"]
})
}
async fn execute(&self, args: serde_json::Value, ctx: &ToolContext) -> ToolOutput {
let question = match args.get("question").and_then(|q| q.as_str()) {
Some(q) => q.to_string(),
None => return ToolOutput::error("缺少必需参数 'question'"),
};
let top_k = args.get("top_k").and_then(|k| k.as_u64()).unwrap_or(5) as usize;
info!("[RagSearchTool] 执行语义检索: question='{}', top_k={}", question, top_k);
let state = &ctx.app_state;
match crate::services::rag::retrieve(&state.db, &state.embedding, &question, top_k).await {
Ok(results) => {
if results.is_empty() {
return ToolOutput::success(
"未找到相关的文献片段。文献库中可能尚无向量化数据,请先对目标文献执行向量化操作。",
json!({ "count": 0 })
);
}
let content = results.iter().enumerate().map(|(i, r)| {
format!(
"[片段 {} | 来源: {} | 段落: {} | 相似度距离: {:.4}]\n{}",
i + 1, r.bibcode, r.paragraph_index, r.distance, r.content
)
}).collect::<Vec<_>>().join("\n\n---\n\n");
let sources: Vec<serde_json::Value> = results.iter().map(|r| {
json!({
"bibcode": r.bibcode,
"paragraph_index": r.paragraph_index,
"distance": r.distance,
"preview": truncate_content(&r.content, 100)
})
}).collect();
ToolOutput::success(
truncate_content(&content, 4000),
json!({ "count": results.len(), "sources": sources })
)
}
Err(e) => ToolOutput::error(format!("RAG 语义检索失败: {}", e)),
}
}
}
// ────────────────────────── 5. QueryTargetTool ──────────────────────────
/// 天体信息查询工具:通过 CDS Sesame 查询天体物理属性
pub struct QueryTargetTool;
#[async_trait]
impl AgentTool for QueryTargetTool {
fn name(&self) -> &str { "query_target" }
fn description(&self) -> &str {
"查询天体的基本物理属性信息。输入天体名称,返回坐标 (RA/Dec)、视星等、光谱型、视差等属性。\
数据来源为 CDS SIMBAD/Sesame 名称解析服务。适用于:获取天体基本参数、验证天体身份。"
}
fn parameters(&self) -> serde_json::Value {
json!({
"type": "object",
"properties": {
"object_name": {
"type": "string",
"description": "天体名称,如 'NGC 6752', 'GD 358', 'HD 209458', 'M 31' 等"
}
},
"required": ["object_name"]
})
}
async fn execute(&self, args: serde_json::Value, ctx: &ToolContext) -> ToolOutput {
let object_name = match args.get("object_name").and_then(|n| n.as_str()) {
Some(n) => n.to_string(),
None => return ToolOutput::error("缺少必需参数 'object_name'"),
};
info!("[QueryTargetTool] 查询天体信息: {}", object_name);
let state = &ctx.app_state;
let client = reqwest::Client::new();
match crate::services::target::query_target_cached(&state.db, &object_name, None, &client).await {
Ok(info) => {
let content = format!(
"天体: {}\nRA: {}\nDec: {}\n视差: {}\n光谱型: {}\nV星等: {}\n别名: {}",
info.target_name,
info.ra.as_deref().unwrap_or("未知"),
info.dec.as_deref().unwrap_or("未知"),
info.parallax.map(|p| format!("{:.4} mas", p)).unwrap_or_else(|| "未知".to_string()),
info.spectral_type.as_deref().unwrap_or("未知"),
info.v_magnitude.map(|v| format!("{:.2}", v)).unwrap_or_else(|| "未知".to_string()),
if info.aliases.is_empty() { "".to_string() } else { info.aliases.join(", ") }
);
ToolOutput::success(
content,
json!({
"target_name": info.target_name,
"ra": info.ra,
"dec": info.dec,
"parallax": info.parallax,
"spectral_type": info.spectral_type,
"v_magnitude": info.v_magnitude,
"aliases": info.aliases
})
)
}
Err(e) => ToolOutput::error(format!("天体 '{}' 查询失败: {}", object_name, e)),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_truncate_content_short() {
let text = "Hello, world!";
assert_eq!(truncate_content(text, 100), text);
}
#[test]
fn test_truncate_content_long() {
let text = "a".repeat(5000);
let result = truncate_content(&text, 100);
assert!(result.contains("内容已截断"));
assert!(result.contains("5000"));
}
#[test]
fn test_tool_output_success() {
let output = ToolOutput::success("ok", json!({"key": "value"}));
assert!(!output.is_error);
assert_eq!(output.content, "ok");
}
#[test]
fn test_tool_output_error() {
let output = ToolOutput::error("something went wrong");
assert!(output.is_error);
}
#[test]
fn test_tool_registry_definitions() {
let registry = ToolRegistry::new();
let defs = registry.definitions();
assert_eq!(defs.len(), 7);
assert!(defs.iter().any(|d| d.function.name == "search_papers"));
assert!(defs.iter().any(|d| d.function.name == "get_paper_metadata"));
assert!(defs.iter().any(|d| d.function.name == "download_paper"));
assert!(defs.iter().any(|d| d.function.name == "parse_paper"));
assert!(defs.iter().any(|d| d.function.name == "get_paper_content"));
assert!(defs.iter().any(|d| d.function.name == "rag_search"));
assert!(defs.iter().any(|d| d.function.name == "query_target"));
}
#[test]
fn test_tool_registry_get() {
let registry = ToolRegistry::new();
assert!(registry.get("search_papers").is_some());
assert!(registry.get("nonexistent").is_none());
}
}