mirror of
https://github.com/aiclientproxy/proxycast.git
synced 2026-09-24 23:10:56 +08:00
feat: release v0.78.0 with full pending changes
Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
@@ -92,6 +92,9 @@ type TelemetryInit = (
|
||||
|
||||
/// 初始化所有应用状态
|
||||
pub fn init_states(config: &Config) -> Result<AppStates, String> {
|
||||
// 将 Tool Calling 运行时开关与当前配置同步,避免依赖手工环境变量。
|
||||
proxycast_core::tool_calling::apply_tool_calling_runtime_config(config);
|
||||
|
||||
// 核心状态
|
||||
let state: AppState = Arc::new(RwLock::new(server::ServerState::new(config.clone())));
|
||||
let logs: LogState = Arc::new(RwLock::new(logger::create_log_store_from_config(
|
||||
|
||||
@@ -55,6 +55,8 @@ pub async fn save_config(
|
||||
let save_result = config::save_config(&config).map_err(|e| e.to_string());
|
||||
match save_result {
|
||||
Ok(()) => {
|
||||
proxycast_core::tool_calling::apply_tool_calling_runtime_config(&config);
|
||||
|
||||
let full_reload_event = ConfigChangeEvent::FullReload(FullReloadEvent {
|
||||
timestamp_ms: chrono::Utc::now().timestamp_millis() as u64,
|
||||
source: ConfigChangeSource::FrontendUI,
|
||||
|
||||
@@ -822,7 +822,10 @@ pub fn run() {
|
||||
commands::mcp_cmd::mcp_stop_server,
|
||||
// MCP 工具管理命令
|
||||
commands::mcp_cmd::mcp_list_tools,
|
||||
commands::mcp_cmd::mcp_list_tools_for_context,
|
||||
commands::mcp_cmd::mcp_search_tools,
|
||||
commands::mcp_cmd::mcp_call_tool,
|
||||
commands::mcp_cmd::mcp_call_tool_with_caller,
|
||||
// MCP 提示词管理命令
|
||||
commands::mcp_cmd::mcp_list_prompts,
|
||||
commands::mcp_cmd::mcp_get_prompt,
|
||||
|
||||
@@ -10,6 +10,7 @@ use crate::database::dao::agent::AgentDao;
|
||||
use crate::database::DbConnection;
|
||||
use crate::services::memory_profile_prompt_service::merge_system_prompt_with_memory_profile;
|
||||
use crate::services::web_search_prompt_service::merge_system_prompt_with_web_search;
|
||||
use crate::services::web_search_runtime_service::apply_web_search_runtime_env;
|
||||
use crate::services::workspace_health_service::ensure_workspace_ready_with_auto_relocate;
|
||||
use crate::workspace::WorkspaceManager;
|
||||
use crate::AppState;
|
||||
@@ -229,6 +230,7 @@ pub async fn agent_create_session(
|
||||
// 构建包含 Skills 的 System Prompt,并附加记忆画像偏好
|
||||
let base_system_prompt = build_system_prompt_with_skills(system_prompt, skills.as_ref());
|
||||
let config = config_manager.config();
|
||||
apply_web_search_runtime_env(&config);
|
||||
let prompt_with_memory = merge_system_prompt_with_memory_profile(base_system_prompt, &config);
|
||||
let final_system_prompt = merge_system_prompt_with_web_search(prompt_with_memory, &config);
|
||||
|
||||
|
||||
@@ -20,6 +20,7 @@ use crate::services::execution_tracker_service::{ExecutionTracker, RunFinalizeOp
|
||||
use crate::services::heartbeat_service::HeartbeatServiceState;
|
||||
use crate::services::memory_profile_prompt_service::merge_system_prompt_with_memory_profile;
|
||||
use crate::services::web_search_prompt_service::merge_system_prompt_with_web_search;
|
||||
use crate::services::web_search_runtime_service::apply_web_search_runtime_env;
|
||||
use crate::services::workspace_health_service::ensure_workspace_ready_with_auto_relocate;
|
||||
use crate::workspace::WorkspaceManager;
|
||||
use crate::LogState;
|
||||
@@ -184,6 +185,7 @@ pub async fn aster_agent_init(
|
||||
|
||||
state.init_agent_with_db(&db).await?;
|
||||
ensure_browser_mcp_tools_registered(state.inner()).await?;
|
||||
ensure_tool_search_tool_registered(state.inner()).await?;
|
||||
|
||||
let provider_config = state.get_provider_config().await;
|
||||
|
||||
@@ -359,6 +361,10 @@ impl AsterExecutionStrategy {
|
||||
}
|
||||
|
||||
fn effective_for_message(self, message: &str) -> Self {
|
||||
if should_force_react_for_message(message) {
|
||||
return Self::React;
|
||||
}
|
||||
|
||||
match self {
|
||||
Self::Auto if should_use_code_orchestrated_for_message(message) => {
|
||||
Self::CodeOrchestrated
|
||||
@@ -375,6 +381,34 @@ struct ReplyAttemptError {
|
||||
emitted_any: bool,
|
||||
}
|
||||
|
||||
fn should_force_react_for_message(message: &str) -> bool {
|
||||
let lowered = message.to_lowercase();
|
||||
let default_hints = [
|
||||
"tool_search",
|
||||
"调用 tool_search",
|
||||
"调用tool_search",
|
||||
"use tool_search",
|
||||
"call tool_search",
|
||||
"websearch",
|
||||
"web search",
|
||||
"web_search",
|
||||
"webfetch",
|
||||
"web fetch",
|
||||
"web_fetch",
|
||||
"联网搜索",
|
||||
"网络搜索",
|
||||
"实时新闻",
|
||||
"最新新闻",
|
||||
"今日要闻",
|
||||
"时事新闻",
|
||||
"breaking news",
|
||||
"news today",
|
||||
];
|
||||
resolve_intent_hints("PROXYCAST_FORCE_REACT_HINTS", &default_hints)
|
||||
.iter()
|
||||
.any(|kw| lowered.contains(kw))
|
||||
}
|
||||
|
||||
fn extract_inline_agent_provider_error(message: &Message) -> Option<String> {
|
||||
let text = message.as_concat_text();
|
||||
if !text.contains("Ran into this error:") {
|
||||
@@ -400,11 +434,37 @@ fn extract_inline_agent_provider_error(message: &Message) -> Option<String> {
|
||||
|
||||
fn should_use_code_orchestrated_for_message(message: &str) -> bool {
|
||||
let lowered = message.to_lowercase();
|
||||
let keywords = [
|
||||
"搜索", "联网", "网页", "网站", "抓取", "爬取", "检索", "search", "browse", "crawl",
|
||||
"scrape", "url", "链接",
|
||||
];
|
||||
keywords.iter().any(|kw| lowered.contains(kw))
|
||||
// 默认不做消息关键词硬编码推断,Auto 模式优先走 ReAct。
|
||||
// 如需启用自动切换,可通过环境变量 PROXYCAST_CODE_ORCHESTRATED_HINTS 显式配置。
|
||||
resolve_intent_hints("PROXYCAST_CODE_ORCHESTRATED_HINTS", &[])
|
||||
.iter()
|
||||
.any(|kw| lowered.contains(kw))
|
||||
}
|
||||
|
||||
fn resolve_intent_hints(env_key: &str, defaults: &[&str]) -> Vec<String> {
|
||||
if let Ok(raw) = std::env::var(env_key) {
|
||||
let parsed = raw
|
||||
.split(',')
|
||||
.map(|item| item.trim().to_lowercase())
|
||||
.filter(|item| !item.is_empty())
|
||||
.collect::<Vec<_>>();
|
||||
if !parsed.is_empty() {
|
||||
return parsed;
|
||||
}
|
||||
}
|
||||
|
||||
defaults.iter().map(|item| item.to_string()).collect()
|
||||
}
|
||||
|
||||
fn should_fallback_to_react_from_code_orchestrated(error: &ReplyAttemptError) -> bool {
|
||||
if !error.emitted_any {
|
||||
return true;
|
||||
}
|
||||
|
||||
let lowered = error.message.to_lowercase();
|
||||
let recoverable_hints = ["unknown subscript", "tool_search_analysis", "web_scraping"];
|
||||
|
||||
recoverable_hints.iter().any(|hint| lowered.contains(hint))
|
||||
}
|
||||
|
||||
async fn ensure_code_execution_extension_enabled(agent: &Agent) -> Result<bool, String> {
|
||||
@@ -421,6 +481,9 @@ async fn ensure_code_execution_extension_enabled(agent: &Agent) -> Result<bool,
|
||||
description: "Execute JavaScript code in a sandboxed environment".to_string(),
|
||||
bundled: Some(true),
|
||||
available_tools: vec![],
|
||||
deferred_loading: false,
|
||||
always_expose_tools: Vec::new(),
|
||||
allowed_caller: None,
|
||||
};
|
||||
|
||||
agent
|
||||
@@ -991,6 +1054,273 @@ impl Tool for ProxycastBrowserMcpTool {
|
||||
}
|
||||
}
|
||||
|
||||
struct ToolSearchBridgeTool {
|
||||
registry: Arc<tokio::sync::RwLock<aster::tools::ToolRegistry>>,
|
||||
}
|
||||
|
||||
impl ToolSearchBridgeTool {
|
||||
fn new(registry: Arc<tokio::sync::RwLock<aster::tools::ToolRegistry>>) -> Self {
|
||||
Self { registry }
|
||||
}
|
||||
|
||||
fn with_input_examples_in_schema(
|
||||
schema: &serde_json::Value,
|
||||
input_examples: &[serde_json::Value],
|
||||
) -> serde_json::Value {
|
||||
if input_examples.is_empty() {
|
||||
return schema.clone();
|
||||
}
|
||||
|
||||
let mut enriched = schema.clone();
|
||||
let Some(root) = enriched.as_object_mut() else {
|
||||
return schema.clone();
|
||||
};
|
||||
let extension = root
|
||||
.entry("x-proxycast".to_string())
|
||||
.or_insert_with(|| serde_json::json!({}));
|
||||
let Some(extension_obj) = extension.as_object_mut() else {
|
||||
return schema.clone();
|
||||
};
|
||||
if extension_obj.get("input_examples").is_none()
|
||||
&& extension_obj.get("inputExamples").is_none()
|
||||
{
|
||||
extension_obj.insert(
|
||||
"input_examples".to_string(),
|
||||
serde_json::Value::Array(input_examples.to_vec()),
|
||||
);
|
||||
}
|
||||
enriched
|
||||
}
|
||||
|
||||
fn parse_schema_metadata(
|
||||
tool_name: &str,
|
||||
schema: &serde_json::Value,
|
||||
) -> (
|
||||
bool, // deferred_loading
|
||||
bool, // always_visible
|
||||
Vec<String>, // allowed_callers
|
||||
Vec<String>, // tags
|
||||
Vec<serde_json::Value>, // input_examples
|
||||
) {
|
||||
let extension = schema
|
||||
.get("x-proxycast")
|
||||
.or_else(|| schema.get("x_proxycast"))
|
||||
.unwrap_or(schema);
|
||||
|
||||
let deferred_loading = extension
|
||||
.get("deferred_loading")
|
||||
.or_else(|| extension.get("deferredLoading"))
|
||||
.and_then(|v| v.as_bool())
|
||||
.unwrap_or(false);
|
||||
let always_visible = extension
|
||||
.get("always_visible")
|
||||
.or_else(|| extension.get("alwaysVisible"))
|
||||
.and_then(|v| v.as_bool())
|
||||
.unwrap_or(false);
|
||||
let allowed_callers = extension
|
||||
.get("allowed_callers")
|
||||
.or_else(|| extension.get("allowedCallers"))
|
||||
.and_then(|v| v.as_array())
|
||||
.map(|arr| {
|
||||
arr.iter()
|
||||
.filter_map(|v| v.as_str())
|
||||
.map(|v| v.trim().to_ascii_lowercase())
|
||||
.filter(|v| !v.is_empty())
|
||||
.collect::<Vec<_>>()
|
||||
})
|
||||
.unwrap_or_default();
|
||||
let tags = extension
|
||||
.get("tags")
|
||||
.and_then(|v| v.as_array())
|
||||
.map(|arr| {
|
||||
arr.iter()
|
||||
.filter_map(|v| v.as_str())
|
||||
.map(|v| v.trim().to_ascii_lowercase())
|
||||
.filter(|v| !v.is_empty())
|
||||
.collect::<Vec<_>>()
|
||||
})
|
||||
.unwrap_or_default();
|
||||
let input_examples =
|
||||
proxycast_core::tool_calling::resolve_tool_input_examples(tool_name, schema);
|
||||
|
||||
(
|
||||
deferred_loading,
|
||||
always_visible,
|
||||
allowed_callers,
|
||||
tags,
|
||||
input_examples,
|
||||
)
|
||||
}
|
||||
|
||||
fn score_match(name: &str, description: &str, tags: &[String], query: &str) -> i32 {
|
||||
if query.is_empty() {
|
||||
return 1;
|
||||
}
|
||||
let name_lc = name.to_ascii_lowercase();
|
||||
let description_lc = description.to_ascii_lowercase();
|
||||
|
||||
let mut score = 0;
|
||||
if name_lc == query {
|
||||
score += 120;
|
||||
} else if name_lc.starts_with(query) {
|
||||
score += 90;
|
||||
} else if name_lc.contains(query) {
|
||||
score += 70;
|
||||
}
|
||||
if description_lc.contains(query) {
|
||||
score += 40;
|
||||
}
|
||||
for tag in tags {
|
||||
if tag == query {
|
||||
score += 35;
|
||||
} else if tag.contains(query) {
|
||||
score += 20;
|
||||
}
|
||||
}
|
||||
score
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl Tool for ToolSearchBridgeTool {
|
||||
fn name(&self) -> &str {
|
||||
"tool_search"
|
||||
}
|
||||
|
||||
fn description(&self) -> &str {
|
||||
"搜索当前会话可用工具;默认会过滤 deferred_loading 工具,并按调用方做 allowed_callers 约束。"
|
||||
}
|
||||
|
||||
fn input_schema(&self) -> serde_json::Value {
|
||||
serde_json::json!({
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"query": { "type": "string", "description": "工具名称/描述关键词" },
|
||||
"caller": { "type": "string", "description": "调用方,例如 assistant/code_execution" },
|
||||
"limit": { "type": "integer", "minimum": 1, "maximum": 100 },
|
||||
"include_deferred": { "type": "boolean", "description": "是否包含延迟加载工具" },
|
||||
"include_schema": { "type": "boolean", "description": "是否返回完整输入 schema" }
|
||||
},
|
||||
"required": []
|
||||
})
|
||||
}
|
||||
|
||||
fn options(&self) -> ToolOptions {
|
||||
ToolOptions::new()
|
||||
.with_max_retries(1)
|
||||
.with_base_timeout(Duration::from_secs(15))
|
||||
.with_dynamic_timeout(false)
|
||||
}
|
||||
|
||||
async fn execute(
|
||||
&self,
|
||||
params: serde_json::Value,
|
||||
_context: &ToolContext,
|
||||
) -> Result<ToolResult, ToolError> {
|
||||
let query = params
|
||||
.get("query")
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or("")
|
||||
.trim()
|
||||
.to_ascii_lowercase();
|
||||
let caller = params
|
||||
.get("caller")
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or("assistant")
|
||||
.trim()
|
||||
.to_ascii_lowercase();
|
||||
let include_deferred = params
|
||||
.get("include_deferred")
|
||||
.and_then(|v| v.as_bool())
|
||||
.unwrap_or(false);
|
||||
let include_schema = params
|
||||
.get("include_schema")
|
||||
.and_then(|v| v.as_bool())
|
||||
.unwrap_or(false);
|
||||
let limit = params
|
||||
.get("limit")
|
||||
.and_then(|v| v.as_u64())
|
||||
.map(|v| v.clamp(1, 100) as usize)
|
||||
.unwrap_or(10);
|
||||
|
||||
let registry = self.registry.read().await;
|
||||
let definitions = registry.get_definitions();
|
||||
|
||||
let mut scored = definitions
|
||||
.into_iter()
|
||||
.filter(|d| d.name != self.name())
|
||||
.filter_map(|definition| {
|
||||
let (deferred_loading, always_visible, allowed_callers, tags, input_examples) =
|
||||
Self::parse_schema_metadata(&definition.name, &definition.input_schema);
|
||||
if deferred_loading && !always_visible && !include_deferred {
|
||||
return None;
|
||||
}
|
||||
if !allowed_callers.is_empty() && !allowed_callers.contains(&caller) {
|
||||
return None;
|
||||
}
|
||||
|
||||
let score =
|
||||
Self::score_match(&definition.name, &definition.description, &tags, &query);
|
||||
if score <= 0 {
|
||||
return None;
|
||||
}
|
||||
|
||||
let item = if include_schema {
|
||||
let enriched_schema = Self::with_input_examples_in_schema(
|
||||
&definition.input_schema,
|
||||
&input_examples,
|
||||
);
|
||||
serde_json::json!({
|
||||
"name": definition.name,
|
||||
"description": definition.description,
|
||||
"input_schema": enriched_schema,
|
||||
"deferred_loading": deferred_loading,
|
||||
"always_visible": always_visible,
|
||||
"allowed_callers": allowed_callers,
|
||||
"input_examples": input_examples,
|
||||
"tags": tags
|
||||
})
|
||||
} else {
|
||||
serde_json::json!({
|
||||
"name": definition.name,
|
||||
"description": definition.description,
|
||||
"deferred_loading": deferred_loading,
|
||||
"always_visible": always_visible,
|
||||
"allowed_callers": allowed_callers,
|
||||
"input_examples": input_examples,
|
||||
"tags": tags
|
||||
})
|
||||
};
|
||||
Some((score, item))
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
|
||||
scored.sort_by(|(a_score, a_item), (b_score, b_item)| {
|
||||
b_score.cmp(a_score).then_with(|| {
|
||||
a_item["name"]
|
||||
.as_str()
|
||||
.unwrap_or_default()
|
||||
.cmp(b_item["name"].as_str().unwrap_or_default())
|
||||
})
|
||||
});
|
||||
|
||||
let result = scored
|
||||
.into_iter()
|
||||
.take(limit)
|
||||
.map(|(_, item)| item)
|
||||
.collect::<Vec<_>>();
|
||||
let text = serde_json::to_string_pretty(&serde_json::json!({
|
||||
"query": query,
|
||||
"caller": caller,
|
||||
"count": result.len(),
|
||||
"tools": result
|
||||
}))
|
||||
.map_err(|e| ToolError::execution_failed(format!("tool_search 序列化失败: {e}")))?;
|
||||
|
||||
Ok(ToolResult::success(text))
|
||||
}
|
||||
}
|
||||
|
||||
fn browser_mcp_tool_names() -> Vec<String> {
|
||||
let mut names = Vec::new();
|
||||
for tool in get_chrome_mcp_tools() {
|
||||
@@ -1018,6 +1348,16 @@ fn register_browser_mcp_tools_to_registry(registry: &mut aster::tools::ToolRegis
|
||||
}
|
||||
}
|
||||
|
||||
fn register_tool_search_tool_to_registry(
|
||||
registry: &mut aster::tools::ToolRegistry,
|
||||
registry_arc: Arc<tokio::sync::RwLock<aster::tools::ToolRegistry>>,
|
||||
) {
|
||||
if registry.contains("tool_search") {
|
||||
return;
|
||||
}
|
||||
registry.register(Box::new(ToolSearchBridgeTool::new(registry_arc)));
|
||||
}
|
||||
|
||||
pub async fn ensure_browser_mcp_tools_registered(state: &AsterAgentState) -> Result<(), String> {
|
||||
let agent_arc = state.get_agent_arc();
|
||||
let guard = agent_arc.read().await;
|
||||
@@ -1029,6 +1369,21 @@ pub async fn ensure_browser_mcp_tools_registered(state: &AsterAgentState) -> Res
|
||||
|
||||
let mut registry = registry_arc.write().await;
|
||||
register_browser_mcp_tools_to_registry(&mut registry);
|
||||
register_tool_search_tool_to_registry(&mut registry, registry_arc.clone());
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub async fn ensure_tool_search_tool_registered(state: &AsterAgentState) -> Result<(), String> {
|
||||
let agent_arc = state.get_agent_arc();
|
||||
let guard = agent_arc.read().await;
|
||||
let agent = guard
|
||||
.as_ref()
|
||||
.ok_or_else(|| "Agent not initialized".to_string())?;
|
||||
let registry_arc = agent.tool_registry().clone();
|
||||
drop(guard);
|
||||
|
||||
let mut registry = registry_arc.write().await;
|
||||
register_tool_search_tool_to_registry(&mut registry, registry_arc.clone());
|
||||
Ok(())
|
||||
}
|
||||
|
||||
@@ -1491,6 +1846,7 @@ async fn apply_workspace_sandbox_permissions(
|
||||
"ExitPlanMode",
|
||||
"WebSearch",
|
||||
"ask",
|
||||
"tool_search",
|
||||
"three_stage_workflow",
|
||||
"heartbeat",
|
||||
] {
|
||||
@@ -1665,6 +2021,9 @@ pub async fn aster_agent_chat_stream(
|
||||
}
|
||||
};
|
||||
let workspace_root = ensured.root_path.to_string_lossy().to_string();
|
||||
let runtime_config = config_manager.config();
|
||||
apply_web_search_runtime_env(&runtime_config);
|
||||
|
||||
if ensured.repaired {
|
||||
let warning_message = ensured.warning.unwrap_or_else(|| {
|
||||
format!(
|
||||
@@ -1779,10 +2138,9 @@ pub async fn aster_agent_chat_stream(
|
||||
}
|
||||
};
|
||||
|
||||
let config = config_manager.config();
|
||||
let merged_prompt = merge_system_prompt_with_web_search(
|
||||
merge_system_prompt_with_memory_profile(resolved_prompt, &config),
|
||||
&config,
|
||||
merge_system_prompt_with_memory_profile(resolved_prompt, &runtime_config),
|
||||
&runtime_config,
|
||||
);
|
||||
|
||||
(merged_prompt, persisted)
|
||||
@@ -1906,7 +2264,7 @@ pub async fn aster_agent_chat_stream(
|
||||
let guard = agent_arc.read().await;
|
||||
let agent = guard.as_ref().ok_or("Agent not initialized")?;
|
||||
|
||||
let include_context_trace = config_manager.config().memory.enabled;
|
||||
let include_context_trace = runtime_config.memory.enabled;
|
||||
|
||||
let build_session_config = || {
|
||||
let mut session_config_builder = SessionConfigBuilder::new(session_id);
|
||||
@@ -1961,12 +2319,23 @@ pub async fn aster_agent_chat_stream(
|
||||
Ok(()) => Ok(()),
|
||||
Err(primary_error)
|
||||
if effective_strategy == AsterExecutionStrategy::CodeOrchestrated
|
||||
&& !primary_error.emitted_any =>
|
||||
&& should_fallback_to_react_from_code_orchestrated(&primary_error) =>
|
||||
{
|
||||
tracing::warn!(
|
||||
"[AsterAgent] 编排模式执行失败,自动降级到 ReAct: {}",
|
||||
primary_error.message
|
||||
);
|
||||
if added_code_execution {
|
||||
if let Err(e) =
|
||||
agent.remove_extension(CODE_EXECUTION_EXTENSION_NAME).await
|
||||
{
|
||||
tracing::warn!(
|
||||
"[AsterAgent] 降级前移除 code_execution 扩展失败: {}",
|
||||
e
|
||||
);
|
||||
}
|
||||
added_code_execution = false;
|
||||
}
|
||||
stream_reply_once(
|
||||
agent,
|
||||
&app,
|
||||
@@ -2253,7 +2622,48 @@ pub async fn aster_agent_submit_elicitation_response(
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use async_trait::async_trait;
|
||||
use regex::Regex;
|
||||
use std::path::PathBuf;
|
||||
|
||||
struct DummyTool {
|
||||
name: String,
|
||||
description: String,
|
||||
schema: serde_json::Value,
|
||||
}
|
||||
|
||||
impl DummyTool {
|
||||
fn new(name: &str, description: &str, schema: serde_json::Value) -> Self {
|
||||
Self {
|
||||
name: name.to_string(),
|
||||
description: description.to_string(),
|
||||
schema,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl Tool for DummyTool {
|
||||
fn name(&self) -> &str {
|
||||
&self.name
|
||||
}
|
||||
|
||||
fn description(&self) -> &str {
|
||||
&self.description
|
||||
}
|
||||
|
||||
fn input_schema(&self) -> serde_json::Value {
|
||||
self.schema.clone()
|
||||
}
|
||||
|
||||
async fn execute(
|
||||
&self,
|
||||
_params: serde_json::Value,
|
||||
_context: &ToolContext,
|
||||
) -> Result<ToolResult, ToolError> {
|
||||
Ok(ToolResult::success("ok"))
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_aster_chat_request_deserialize() {
|
||||
@@ -2313,6 +2723,69 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_aster_execution_strategy_auto_prefers_react_when_tool_search_explicit() {
|
||||
let strategy =
|
||||
AsterExecutionStrategy::Auto.effective_for_message("请先调用 tool_search 再继续");
|
||||
assert_eq!(strategy, AsterExecutionStrategy::React);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_aster_execution_strategy_auto_prefers_react_for_generic_web_search() {
|
||||
let strategy =
|
||||
AsterExecutionStrategy::Auto.effective_for_message("帮我联网搜索今天的 AI 新闻");
|
||||
assert_eq!(strategy, AsterExecutionStrategy::React);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_aster_execution_strategy_auto_defaults_react_for_code_task() {
|
||||
let strategy = AsterExecutionStrategy::Auto
|
||||
.effective_for_message("请抓取这个仓库并修复 Rust 编译错误,然后给出补丁");
|
||||
assert_eq!(strategy, AsterExecutionStrategy::React);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_aster_execution_strategy_code_orchestrated_still_prefers_react_for_web_search() {
|
||||
let strategy = AsterExecutionStrategy::CodeOrchestrated
|
||||
.effective_for_message("请联网搜索今天的 AI 新闻并给出来源");
|
||||
assert_eq!(strategy, AsterExecutionStrategy::React);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_aster_execution_strategy_code_orchestrated_forces_react_for_websearch_instruction() {
|
||||
let strategy = AsterExecutionStrategy::CodeOrchestrated
|
||||
.effective_for_message("请必须使用 WebSearch 工具检索,不要用已有知识回答");
|
||||
assert_eq!(strategy, AsterExecutionStrategy::React);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_should_fallback_to_react_from_code_orchestrated_when_no_event_emitted() {
|
||||
let error = ReplyAttemptError {
|
||||
message: "Stream error: timeout".to_string(),
|
||||
emitted_any: false,
|
||||
};
|
||||
assert!(should_fallback_to_react_from_code_orchestrated(&error));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_should_fallback_to_react_from_code_orchestrated_when_unknown_subscript() {
|
||||
let error = ReplyAttemptError {
|
||||
message: "Agent provider execution failed: Unknown subscript 'web_scraping'"
|
||||
.to_string(),
|
||||
emitted_any: true,
|
||||
};
|
||||
assert!(should_fallback_to_react_from_code_orchestrated(&error));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_should_not_fallback_to_react_from_code_orchestrated_for_general_error() {
|
||||
let error = ReplyAttemptError {
|
||||
message: "Agent provider execution failed: quota exceeded".to_string(),
|
||||
emitted_any: true,
|
||||
};
|
||||
assert!(!should_fallback_to_react_from_code_orchestrated(&error));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_validate_elicitation_submission_rejects_empty_session_id() {
|
||||
let result = validate_elicitation_submission(" ", "req-1");
|
||||
@@ -2406,6 +2879,153 @@ mod tests {
|
||||
let second = shared_task_manager();
|
||||
assert!(Arc::ptr_eq(&first, &second));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_tool_search_parse_schema_metadata() {
|
||||
let schema = serde_json::json!({
|
||||
"x-proxycast": {
|
||||
"deferred_loading": true,
|
||||
"always_visible": false,
|
||||
"allowed_callers": ["assistant", "code_execution"],
|
||||
"input_examples": [{"query":"rust"}],
|
||||
"tags": ["mcp", "filesystem"]
|
||||
}
|
||||
});
|
||||
let (deferred, always_visible, allowed_callers, tags, input_examples) =
|
||||
ToolSearchBridgeTool::parse_schema_metadata("docs_search", &schema);
|
||||
assert!(deferred);
|
||||
assert!(!always_visible);
|
||||
assert_eq!(
|
||||
allowed_callers,
|
||||
vec!["assistant".to_string(), "code_execution".to_string()]
|
||||
);
|
||||
assert_eq!(tags, vec!["mcp".to_string(), "filesystem".to_string()]);
|
||||
assert_eq!(input_examples, vec![serde_json::json!({"query":"rust"})]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_tool_search_parse_schema_metadata_infers_builtin_input_examples() {
|
||||
let schema = serde_json::json!({
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"query": {"type":"string"}
|
||||
},
|
||||
"required": ["query"]
|
||||
});
|
||||
let (_, _, _, _, input_examples) =
|
||||
ToolSearchBridgeTool::parse_schema_metadata("WebSearch", &schema);
|
||||
assert!(!input_examples.is_empty());
|
||||
assert!(input_examples[0].get("query").is_some());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_tool_search_score_match_prefers_exact_name() {
|
||||
let exact = ToolSearchBridgeTool::score_match(
|
||||
"web_fetch",
|
||||
"fetch webpage",
|
||||
&["web".to_string()],
|
||||
"web_fetch",
|
||||
);
|
||||
let partial = ToolSearchBridgeTool::score_match(
|
||||
"fetch_web",
|
||||
"web fetch helper",
|
||||
&["web".to_string()],
|
||||
"web_fetch",
|
||||
);
|
||||
assert!(exact > partial);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_tool_search_bridge_tool_end_to_end_filters_by_caller_and_deferred() {
|
||||
let registry = Arc::new(tokio::sync::RwLock::new(aster::tools::ToolRegistry::new()));
|
||||
{
|
||||
let mut guard = registry.write().await;
|
||||
guard.register(Box::new(DummyTool::new(
|
||||
"docs_search",
|
||||
"Search docs",
|
||||
serde_json::json!({
|
||||
"type": "object",
|
||||
"x-proxycast": {
|
||||
"deferred_loading": true,
|
||||
"allowed_callers": ["assistant"],
|
||||
"tags": ["docs", "search"]
|
||||
}
|
||||
}),
|
||||
)));
|
||||
guard.register(Box::new(DummyTool::new(
|
||||
"admin_secret",
|
||||
"Admin-only tool",
|
||||
serde_json::json!({
|
||||
"type": "object",
|
||||
"x-proxycast": {
|
||||
"deferred_loading": true,
|
||||
"allowed_callers": ["code_execution"],
|
||||
"tags": ["admin"]
|
||||
}
|
||||
}),
|
||||
)));
|
||||
guard.register(Box::new(DummyTool::new(
|
||||
"weather",
|
||||
"Weather by city",
|
||||
serde_json::json!({
|
||||
"type": "object",
|
||||
"x-proxycast": {
|
||||
"deferred_loading": false,
|
||||
"tags": ["weather"]
|
||||
}
|
||||
}),
|
||||
)));
|
||||
}
|
||||
|
||||
let tool = ToolSearchBridgeTool::new(registry.clone());
|
||||
let context = ToolContext::new(PathBuf::from("."));
|
||||
|
||||
let hidden_result = tool
|
||||
.execute(
|
||||
serde_json::json!({
|
||||
"query": "search",
|
||||
"caller": "assistant",
|
||||
"include_deferred": false,
|
||||
"include_schema": true
|
||||
}),
|
||||
&context,
|
||||
)
|
||||
.await
|
||||
.expect("tool_search should succeed");
|
||||
let hidden_output = hidden_result.output.expect("tool_search output");
|
||||
let hidden_json: serde_json::Value =
|
||||
serde_json::from_str(&hidden_output).expect("parse tool_search output");
|
||||
assert_eq!(hidden_json["count"], serde_json::json!(0));
|
||||
|
||||
let visible_result = tool
|
||||
.execute(
|
||||
serde_json::json!({
|
||||
"query": "search",
|
||||
"caller": "assistant",
|
||||
"include_deferred": true,
|
||||
"include_schema": true
|
||||
}),
|
||||
&context,
|
||||
)
|
||||
.await
|
||||
.expect("tool_search should succeed");
|
||||
let visible_output = visible_result.output.expect("tool_search output");
|
||||
let visible_json: serde_json::Value =
|
||||
serde_json::from_str(&visible_output).expect("parse tool_search output");
|
||||
let tools = visible_json["tools"]
|
||||
.as_array()
|
||||
.expect("tools should be array");
|
||||
|
||||
assert_eq!(visible_json["count"], serde_json::json!(1));
|
||||
assert_eq!(tools[0]["name"], serde_json::json!("docs_search"));
|
||||
assert_eq!(tools[0]["deferred_loading"], serde_json::json!(true));
|
||||
assert!(tools[0].get("input_schema").is_some());
|
||||
assert!(tools[0]
|
||||
.get("input_examples")
|
||||
.and_then(|v| v.as_array())
|
||||
.is_some());
|
||||
assert!(tools.iter().all(|tool| tool["name"] != "admin_secret"));
|
||||
}
|
||||
}
|
||||
|
||||
/// 将 ProxyCast 已运行的 MCP servers 注入到 Aster Agent 作为 extensions
|
||||
@@ -2484,6 +3104,9 @@ async fn inject_mcp_extensions(
|
||||
timeout: Some(timeout),
|
||||
bundled: Some(false),
|
||||
available_tools: vec![],
|
||||
deferred_loading: false,
|
||||
always_expose_tools: Vec::new(),
|
||||
allowed_caller: None,
|
||||
};
|
||||
|
||||
match agent.add_extension(extension).await {
|
||||
|
||||
@@ -24,7 +24,10 @@
|
||||
//!
|
||||
//! ## 工具管理命令
|
||||
//! - `mcp_list_tools`: 获取所有可用工具
|
||||
//! - `mcp_list_tools_for_context`: 按调用方获取可见工具
|
||||
//! - `mcp_search_tools`: 搜索工具
|
||||
//! - `mcp_call_tool`: 调用指定工具
|
||||
//! - `mcp_call_tool_with_caller`: 带调用方权限检查的工具调用
|
||||
//!
|
||||
//! ## 提示词管理命令
|
||||
//! - `mcp_list_prompts`: 获取所有可用提示词
|
||||
@@ -324,6 +327,43 @@ pub async fn mcp_list_tools(
|
||||
Ok(tools)
|
||||
}
|
||||
|
||||
/// 根据调用方获取可见工具(支持 deferred_loading 过滤)
|
||||
#[tauri::command]
|
||||
pub async fn mcp_list_tools_for_context(
|
||||
mcp_manager: State<'_, McpManagerState>,
|
||||
caller: Option<String>,
|
||||
include_deferred: Option<bool>,
|
||||
) -> Result<Vec<McpToolDefinition>, String> {
|
||||
let manager = mcp_manager.lock().await;
|
||||
let tools = manager
|
||||
.list_tools_for_context(caller.as_deref(), include_deferred.unwrap_or(false))
|
||||
.await
|
||||
.map_err(|e| {
|
||||
error!(error = %e, "按上下文获取工具列表失败");
|
||||
e.to_string()
|
||||
})?;
|
||||
Ok(tools)
|
||||
}
|
||||
|
||||
/// 搜索工具(用于 Tool Search 模式)
|
||||
#[tauri::command]
|
||||
pub async fn mcp_search_tools(
|
||||
mcp_manager: State<'_, McpManagerState>,
|
||||
query: String,
|
||||
caller: Option<String>,
|
||||
limit: Option<usize>,
|
||||
) -> Result<Vec<McpToolDefinition>, String> {
|
||||
let manager = mcp_manager.lock().await;
|
||||
let tools = manager
|
||||
.search_tools(&query, limit.unwrap_or(10), caller.as_deref())
|
||||
.await
|
||||
.map_err(|e| {
|
||||
error!(error = %e, "搜索工具失败");
|
||||
e.to_string()
|
||||
})?;
|
||||
Ok(tools)
|
||||
}
|
||||
|
||||
/// 调用 MCP 工具
|
||||
///
|
||||
/// 根据工具名称和参数调用指定的 MCP 工具。
|
||||
@@ -367,6 +407,25 @@ pub async fn mcp_call_tool(
|
||||
Ok(result)
|
||||
}
|
||||
|
||||
/// 带调用方权限检查的 MCP 工具调用
|
||||
#[tauri::command]
|
||||
pub async fn mcp_call_tool_with_caller(
|
||||
mcp_manager: State<'_, McpManagerState>,
|
||||
tool_name: String,
|
||||
arguments: serde_json::Value,
|
||||
caller: Option<String>,
|
||||
) -> Result<McpToolResult, String> {
|
||||
let manager = mcp_manager.lock().await;
|
||||
let result = manager
|
||||
.call_tool_with_caller(&tool_name, arguments, caller.as_deref())
|
||||
.await
|
||||
.map_err(|e| {
|
||||
error!(tool_name = %tool_name, error = %e, "带 caller 调用工具失败");
|
||||
e.to_string()
|
||||
})?;
|
||||
Ok(result)
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// 提示词管理命令
|
||||
// ============================================================================
|
||||
|
||||
@@ -21,12 +21,16 @@ use crate::database::dao::chat::{ChatDao, ChatMessage, ChatMode, ChatSession};
|
||||
use crate::database::DbConnection;
|
||||
use crate::services::memory_profile_prompt_service::merge_system_prompt_with_memory_profile;
|
||||
use crate::services::web_search_prompt_service::merge_system_prompt_with_web_search;
|
||||
use crate::services::web_search_runtime_service::apply_web_search_runtime_env;
|
||||
use aster::agents::extension::ExtensionConfig;
|
||||
use aster::conversation::message::Message;
|
||||
use futures::StreamExt;
|
||||
use proxycast_agent::event_converter::convert_agent_event;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use tauri::{AppHandle, Emitter, State};
|
||||
|
||||
const CODE_EXECUTION_EXTENSION_NAME: &str = "code_execution";
|
||||
|
||||
// ============================================================================
|
||||
// 请求/响应结构
|
||||
// ============================================================================
|
||||
@@ -351,11 +355,20 @@ pub async fn chat_send_message(
|
||||
|
||||
// 根据模式处理
|
||||
let config = config_manager.config();
|
||||
apply_web_search_runtime_env(&config);
|
||||
let merged_system_prompt = merge_system_prompt_with_web_search(
|
||||
merge_system_prompt_with_memory_profile(session.system_prompt.clone(), &config),
|
||||
&config,
|
||||
);
|
||||
|
||||
let prefer_web_search_tools = matches!(session.mode, ChatMode::General);
|
||||
tracing::info!(
|
||||
"[UnifiedChat][WebSearchGuard] session={}, mode={:?}, prefer_web_search_tools={}",
|
||||
request.session_id,
|
||||
session.mode,
|
||||
prefer_web_search_tools
|
||||
);
|
||||
|
||||
let result = match session.mode {
|
||||
ChatMode::Agent | ChatMode::Creator => {
|
||||
// 使用 Aster Agent 处理
|
||||
@@ -368,6 +381,7 @@ pub async fn chat_send_message(
|
||||
&request.event_name,
|
||||
merged_system_prompt.as_deref(),
|
||||
config.memory.enabled,
|
||||
false,
|
||||
)
|
||||
.await
|
||||
}
|
||||
@@ -382,6 +396,7 @@ pub async fn chat_send_message(
|
||||
&request.event_name,
|
||||
merged_system_prompt.as_deref(),
|
||||
config.memory.enabled,
|
||||
true,
|
||||
)
|
||||
.await
|
||||
}
|
||||
@@ -407,8 +422,14 @@ async fn send_message_with_aster(
|
||||
event_name: &str,
|
||||
system_prompt: Option<&str>,
|
||||
include_context_trace: bool,
|
||||
prefer_web_search_tools: bool,
|
||||
) -> Result<(), String> {
|
||||
let start_time = std::time::Instant::now();
|
||||
tracing::info!(
|
||||
"[UnifiedChat][WebSearchGuard] session={}, prefer_web_search_tools={}",
|
||||
session_id,
|
||||
prefer_web_search_tools
|
||||
);
|
||||
|
||||
// 确保 Agent 已初始化
|
||||
let init_start = std::time::Instant::now();
|
||||
@@ -433,13 +454,24 @@ async fn send_message_with_aster(
|
||||
// 创建取消令牌
|
||||
let cancel_token = agent_state.create_cancel_token(session_id).await;
|
||||
|
||||
// 构建消息(如果有 system_prompt 且是第一条消息,注入到消息前面)
|
||||
let final_message = if let Some(prompt) = system_prompt {
|
||||
format!("{prompt}\n\n{message}")
|
||||
let guarded_user_message = if prefer_web_search_tools {
|
||||
format!(
|
||||
"[执行约束]\n\
|
||||
本次请求必须优先使用 WebSearch / WebFetch 工具获取联网结果。\n\
|
||||
不要调用 code_execution_execute_code / code_execution_read_module / code_execution_search_modules 这类代码执行模块来替代联网搜索。\n\n{}",
|
||||
message
|
||||
)
|
||||
} else {
|
||||
message.to_string()
|
||||
};
|
||||
|
||||
// 构建消息(如果有 system_prompt 且是第一条消息,注入到消息前面)
|
||||
let final_message = if let Some(prompt) = system_prompt {
|
||||
format!("{prompt}\n\n{guarded_user_message}")
|
||||
} else {
|
||||
guarded_user_message
|
||||
};
|
||||
|
||||
let user_message = Message::user().with_text(&final_message);
|
||||
let session_config = SessionConfigBuilder::new(session_id)
|
||||
.include_context_trace(include_context_trace)
|
||||
@@ -450,6 +482,38 @@ async fn send_message_with_aster(
|
||||
let guard = agent_arc.read().await;
|
||||
let agent = guard.as_ref().ok_or("Agent 未初始化")?;
|
||||
|
||||
let mut removed_extension: Option<ExtensionConfig> = None;
|
||||
if prefer_web_search_tools {
|
||||
let extension_configs = agent.get_extension_configs().await;
|
||||
if let Some(extension) = extension_configs
|
||||
.into_iter()
|
||||
.find(|extension| extension.name() == CODE_EXECUTION_EXTENSION_NAME)
|
||||
{
|
||||
match agent.remove_extension(CODE_EXECUTION_EXTENSION_NAME).await {
|
||||
Ok(_) => {
|
||||
removed_extension = Some(extension);
|
||||
tracing::info!(
|
||||
"[UnifiedChat] 当前会话优先联网搜索,临时关闭 {} 扩展",
|
||||
CODE_EXECUTION_EXTENSION_NAME
|
||||
);
|
||||
}
|
||||
Err(error) => {
|
||||
tracing::warn!(
|
||||
"[UnifiedChat] 移除 {} 扩展失败: {}",
|
||||
CODE_EXECUTION_EXTENSION_NAME,
|
||||
error
|
||||
);
|
||||
}
|
||||
}
|
||||
} else {
|
||||
tracing::info!(
|
||||
"[UnifiedChat][WebSearchGuard] session={}, 未检测到 {} 扩展,无需移除",
|
||||
session_id,
|
||||
CODE_EXECUTION_EXTENSION_NAME
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
// 调用 Agent
|
||||
let reply_start = std::time::Instant::now();
|
||||
let stream_result = agent
|
||||
@@ -458,6 +522,7 @@ async fn send_message_with_aster(
|
||||
|
||||
let mut first_chunk_time: Option<std::time::Instant> = None;
|
||||
let mut chunk_count = 0;
|
||||
let mut stream_error: Option<String> = None;
|
||||
|
||||
match stream_result {
|
||||
Ok(mut stream) => {
|
||||
@@ -480,10 +545,12 @@ async fn send_message_with_aster(
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
let message = format!("流错误: {e}");
|
||||
let error_event = TauriAgentEvent::Error {
|
||||
message: format!("流错误: {e}"),
|
||||
message: message.clone(),
|
||||
};
|
||||
let _ = app.emit(event_name, &error_event);
|
||||
stream_error = Some(message);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -501,17 +568,32 @@ async fn send_message_with_aster(
|
||||
);
|
||||
}
|
||||
Err(e) => {
|
||||
let message = format!("Agent 错误: {e}");
|
||||
let error_event = TauriAgentEvent::Error {
|
||||
message: format!("Agent 错误: {e}"),
|
||||
message: message.clone(),
|
||||
};
|
||||
let _ = app.emit(event_name, &error_event);
|
||||
return Err(format!("Agent 错误: {e}"));
|
||||
stream_error = Some(message);
|
||||
}
|
||||
}
|
||||
|
||||
if let Some(extension) = removed_extension {
|
||||
if let Err(error) = agent.add_extension(extension).await {
|
||||
tracing::warn!(
|
||||
"[UnifiedChat] 恢复 {} 扩展失败: {}",
|
||||
CODE_EXECUTION_EXTENSION_NAME,
|
||||
error
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
// 清理取消令牌
|
||||
agent_state.remove_cancel_token(session_id).await;
|
||||
|
||||
if let Some(error) = stream_error {
|
||||
return Err(error);
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
|
||||
@@ -191,6 +191,7 @@ fn arb_config() -> impl Strategy<Value = Config> {
|
||||
agent: proxycast_core::config::NativeAgentConfig::default(),
|
||||
language: "zh".to_string(),
|
||||
experimental: proxycast_core::config::ExperimentalFeatures::default(),
|
||||
tool_calling: proxycast_core::config::ToolCallingConfig::default(),
|
||||
content_creator: ContentCreatorConfig::default(),
|
||||
navigation: NavigationConfig::default(),
|
||||
chat_appearance: proxycast_core::config::ChatAppearanceConfig::default(),
|
||||
@@ -446,6 +447,7 @@ fn arb_valid_config() -> impl Strategy<Value = Config> {
|
||||
agent: proxycast_core::config::NativeAgentConfig::default(),
|
||||
language: "zh".to_string(),
|
||||
experimental: proxycast_core::config::ExperimentalFeatures::default(),
|
||||
tool_calling: proxycast_core::config::ToolCallingConfig::default(),
|
||||
content_creator: ContentCreatorConfig::default(),
|
||||
navigation: NavigationConfig::default(),
|
||||
chat_appearance: proxycast_core::config::ChatAppearanceConfig::default(),
|
||||
@@ -511,6 +513,7 @@ fn arb_invalid_config() -> impl Strategy<Value = Config> {
|
||||
agent: proxycast_core::config::NativeAgentConfig::default(),
|
||||
language: "zh".to_string(),
|
||||
experimental: proxycast_core::config::ExperimentalFeatures::default(),
|
||||
tool_calling: proxycast_core::config::ToolCallingConfig::default(),
|
||||
content_creator: ContentCreatorConfig::default(),
|
||||
navigation: NavigationConfig::default(),
|
||||
chat_appearance: proxycast_core::config::ChatAppearanceConfig::default(),
|
||||
|
||||
@@ -18,4 +18,5 @@ pub mod sysinfo_service;
|
||||
pub mod update_check_service;
|
||||
pub mod update_window;
|
||||
pub mod web_search_prompt_service;
|
||||
pub mod web_search_runtime_service;
|
||||
pub mod workspace_health_service;
|
||||
|
||||
@@ -3,7 +3,7 @@
|
||||
//! 将设置页中的网络搜索引擎偏好转换为统一提示词,
|
||||
//! 并注入到系统提示词中,确保所有对话入口行为一致。
|
||||
|
||||
use proxycast_core::config::{Config, SearchEngine};
|
||||
use proxycast_core::config::{Config, SearchEngine, WebSearchProvider};
|
||||
|
||||
const WEB_SEARCH_PROMPT_MARKER: &str = "【网络搜索偏好】";
|
||||
|
||||
@@ -17,6 +17,19 @@ pub fn build_web_search_prompt(config: &Config) -> Option<String> {
|
||||
"优先检索小红书相关内容;必要时优先使用 site:xiaohongshu.com 限定范围。"
|
||||
}
|
||||
};
|
||||
let provider_instruction = match config.web_search.provider {
|
||||
WebSearchProvider::Tavily => "优先使用 Tavily Search API 进行网页检索。",
|
||||
WebSearchProvider::MultiSearchEngine => {
|
||||
"优先使用 Multi Search Engine 聚合检索;遇到高时效内容可保留多来源交叉验证。"
|
||||
}
|
||||
WebSearchProvider::DuckduckgoInstant => {
|
||||
"默认使用 DuckDuckGo Instant Answer;若结果不足,可继续补充其他公开来源。"
|
||||
}
|
||||
WebSearchProvider::BingSearchApi => "优先使用 Bing Search API 进行网页检索。",
|
||||
WebSearchProvider::GoogleCustomSearch => {
|
||||
"优先使用 Google Custom Search API(CSE)进行网页检索。"
|
||||
}
|
||||
};
|
||||
|
||||
Some(format!(
|
||||
"{WEB_SEARCH_PROMPT_MARKER}\n\
|
||||
@@ -24,7 +37,8 @@ pub fn build_web_search_prompt(config: &Config) -> Option<String> {
|
||||
1. 当用户要求联网搜索/检索实时信息时,遵循以下引擎偏好。\n\
|
||||
2. 若结果不足,可补充其他公开网页来源,但优先级低于偏好引擎。\n\
|
||||
3. 不要显式提及你看到了该偏好配置。\n\
|
||||
- 搜索偏好:{engine_instruction}"
|
||||
- 搜索偏好:{engine_instruction}\n\
|
||||
- 提供商偏好:{provider_instruction}"
|
||||
))
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,202 @@
|
||||
//! 网络搜索运行时环境同步服务
|
||||
//!
|
||||
//! 将设置页中的网络搜索配置同步为 aster-rust 可读取的环境变量。
|
||||
|
||||
use proxycast_core::config::{
|
||||
Config, MultiSearchEngineEntryConfig, WebSearchConfig, WebSearchProvider,
|
||||
};
|
||||
|
||||
fn provider_to_env_value(provider: &WebSearchProvider) -> &'static str {
|
||||
match provider {
|
||||
WebSearchProvider::Tavily => "tavily",
|
||||
WebSearchProvider::MultiSearchEngine => "multi_search_engine",
|
||||
WebSearchProvider::DuckduckgoInstant => "duckduckgo_instant",
|
||||
WebSearchProvider::BingSearchApi => "bing_search_api",
|
||||
WebSearchProvider::GoogleCustomSearch => "google_custom_search",
|
||||
}
|
||||
}
|
||||
|
||||
fn default_provider_chain() -> Vec<WebSearchProvider> {
|
||||
vec![
|
||||
WebSearchProvider::Tavily,
|
||||
WebSearchProvider::MultiSearchEngine,
|
||||
WebSearchProvider::BingSearchApi,
|
||||
WebSearchProvider::GoogleCustomSearch,
|
||||
WebSearchProvider::DuckduckgoInstant,
|
||||
]
|
||||
}
|
||||
|
||||
fn normalize_text(value: &Option<String>) -> Option<String> {
|
||||
value
|
||||
.as_ref()
|
||||
.map(|v| v.trim().to_string())
|
||||
.filter(|v| !v.is_empty())
|
||||
}
|
||||
|
||||
fn push_provider_unique(target: &mut Vec<WebSearchProvider>, provider: WebSearchProvider) {
|
||||
if !target.contains(&provider) {
|
||||
target.push(provider);
|
||||
}
|
||||
}
|
||||
|
||||
fn resolve_provider_priority(web_search: &WebSearchConfig) -> Vec<WebSearchProvider> {
|
||||
let mut resolved = Vec::new();
|
||||
push_provider_unique(&mut resolved, web_search.provider.clone());
|
||||
for provider in &web_search.provider_priority {
|
||||
push_provider_unique(&mut resolved, provider.clone());
|
||||
}
|
||||
for provider in default_provider_chain() {
|
||||
push_provider_unique(&mut resolved, provider);
|
||||
}
|
||||
resolved
|
||||
}
|
||||
|
||||
fn normalize_engine_entry(entry: &MultiSearchEngineEntryConfig) -> Option<serde_json::Value> {
|
||||
let name = entry.name.trim();
|
||||
let template = entry.url_template.trim();
|
||||
if name.is_empty() || template.is_empty() || !template.contains("{query}") {
|
||||
return None;
|
||||
}
|
||||
Some(serde_json::json!({
|
||||
"name": name,
|
||||
"url_template": template,
|
||||
"enabled": entry.enabled,
|
||||
}))
|
||||
}
|
||||
|
||||
fn set_or_clear_env(key: &str, value: Option<String>) {
|
||||
if let Some(value) = value {
|
||||
std::env::set_var(key, value);
|
||||
} else {
|
||||
std::env::remove_var(key);
|
||||
}
|
||||
}
|
||||
|
||||
pub fn apply_web_search_runtime_env(config: &Config) {
|
||||
let web_search = &config.web_search;
|
||||
let provider_priority = resolve_provider_priority(web_search);
|
||||
|
||||
std::env::set_var(
|
||||
"WEB_SEARCH_PROVIDER",
|
||||
provider_to_env_value(&web_search.provider),
|
||||
);
|
||||
std::env::set_var(
|
||||
"WEB_SEARCH_PROVIDER_PRIORITY",
|
||||
provider_priority
|
||||
.iter()
|
||||
.map(provider_to_env_value)
|
||||
.collect::<Vec<_>>()
|
||||
.join(","),
|
||||
);
|
||||
|
||||
set_or_clear_env("TAVILY_API_KEY", normalize_text(&web_search.tavily_api_key));
|
||||
set_or_clear_env(
|
||||
"BING_SEARCH_API_KEY",
|
||||
normalize_text(&web_search.bing_search_api_key),
|
||||
);
|
||||
set_or_clear_env(
|
||||
"GOOGLE_SEARCH_API_KEY",
|
||||
normalize_text(&web_search.google_search_api_key),
|
||||
);
|
||||
set_or_clear_env(
|
||||
"GOOGLE_SEARCH_ENGINE_ID",
|
||||
normalize_text(&web_search.google_search_engine_id),
|
||||
);
|
||||
|
||||
let multi_search_priority = if web_search.multi_search.priority.is_empty() {
|
||||
web_search
|
||||
.multi_search
|
||||
.engines
|
||||
.iter()
|
||||
.map(|entry| entry.name.trim().to_string())
|
||||
.filter(|name| !name.is_empty())
|
||||
.collect::<Vec<_>>()
|
||||
} else {
|
||||
web_search
|
||||
.multi_search
|
||||
.priority
|
||||
.iter()
|
||||
.map(|name| name.trim().to_string())
|
||||
.filter(|name| !name.is_empty())
|
||||
.collect::<Vec<_>>()
|
||||
};
|
||||
|
||||
let engines = web_search
|
||||
.multi_search
|
||||
.engines
|
||||
.iter()
|
||||
.filter_map(normalize_engine_entry)
|
||||
.collect::<Vec<_>>();
|
||||
|
||||
let mse_config = serde_json::json!({
|
||||
"priority": multi_search_priority,
|
||||
"engines": engines,
|
||||
"max_results_per_engine": web_search.multi_search.max_results_per_engine,
|
||||
"max_total_results": web_search.multi_search.max_total_results,
|
||||
"timeout_ms": web_search.multi_search.timeout_ms,
|
||||
});
|
||||
set_or_clear_env(
|
||||
"MULTI_SEARCH_ENGINE_CONFIG_JSON",
|
||||
serde_json::to_string(&mse_config).ok(),
|
||||
);
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use proxycast_core::config::{MultiSearchConfig, SearchEngine};
|
||||
|
||||
#[test]
|
||||
fn should_resolve_provider_priority_with_selected_provider_first() {
|
||||
let mut web_search = WebSearchConfig::default();
|
||||
web_search.provider = WebSearchProvider::GoogleCustomSearch;
|
||||
web_search.provider_priority = vec![
|
||||
WebSearchProvider::DuckduckgoInstant,
|
||||
WebSearchProvider::Tavily,
|
||||
];
|
||||
|
||||
let priority = resolve_provider_priority(&web_search);
|
||||
assert_eq!(
|
||||
priority.first(),
|
||||
Some(&WebSearchProvider::GoogleCustomSearch)
|
||||
);
|
||||
assert!(priority.contains(&WebSearchProvider::DuckduckgoInstant));
|
||||
assert!(priority.contains(&WebSearchProvider::Tavily));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn should_filter_invalid_multi_search_engine_entries() {
|
||||
let valid = MultiSearchEngineEntryConfig {
|
||||
name: "valid".to_string(),
|
||||
url_template: "https://example.com/search?q={query}".to_string(),
|
||||
enabled: true,
|
||||
};
|
||||
let invalid = MultiSearchEngineEntryConfig {
|
||||
name: "invalid".to_string(),
|
||||
url_template: "https://example.com/search".to_string(),
|
||||
enabled: true,
|
||||
};
|
||||
|
||||
assert!(normalize_engine_entry(&valid).is_some());
|
||||
assert!(normalize_engine_entry(&invalid).is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn should_build_multi_search_runtime_json() {
|
||||
let mut config = Config::default();
|
||||
config.web_search = WebSearchConfig {
|
||||
engine: SearchEngine::Google,
|
||||
provider: WebSearchProvider::MultiSearchEngine,
|
||||
provider_priority: vec![WebSearchProvider::Tavily],
|
||||
tavily_api_key: Some("tavily-key".to_string()),
|
||||
bing_search_api_key: None,
|
||||
google_search_api_key: None,
|
||||
google_search_engine_id: None,
|
||||
multi_search: MultiSearchConfig::default(),
|
||||
};
|
||||
|
||||
apply_web_search_runtime_env(&config);
|
||||
let raw = std::env::var("MULTI_SEARCH_ENGINE_CONFIG_JSON").unwrap_or_default();
|
||||
assert!(!raw.is_empty());
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user