From c3b77ef849fdf912f7c479becd34cddfe9b1fe14 Mon Sep 17 00:00:00 2001 From: coso Date: Tue, 3 Mar 2026 23:32:52 +0800 Subject: [PATCH] feat: release v0.78.0 with full pending changes Co-Authored-By: Claude Opus 4.6 (1M context) --- RELEASE_NOTES.md | 80 +- package.json | 2 +- src-tauri/Cargo.lock | 56 +- src-tauri/Cargo.toml | 8 +- src-tauri/crates/agent/src/aster_state.rs | 3 + src-tauri/crates/agent/src/event_converter.rs | 74 +- .../crates/config/src/observer/manager.rs | 1 + src-tauri/crates/core/src/agent/types.rs | 9 + src-tauri/crates/core/src/config/mod.rs | 13 +- src-tauri/crates/core/src/config/tests.rs | 3 + src-tauri/crates/core/src/config/types.rs | 197 +++++ src-tauri/crates/core/src/lib.rs | 1 + src-tauri/crates/core/src/tool_calling.rs | 312 ++++++++ src-tauri/crates/mcp/src/manager.rs | 477 ++++++++++-- src-tauri/crates/mcp/src/tool_converter.rs | 95 ++- src-tauri/crates/mcp/src/types.rs | 15 + .../providers/antigravity.txt | 7 + .../providers/src/providers/claude_custom.rs | 276 +++++-- .../providers/src/providers/openai_custom.rs | 394 +++++++++- .../server/src/handlers/provider_calls.rs | 92 ++- src-tauri/src/app/bootstrap.rs | 3 + src-tauri/src/app/commands/config.rs | 2 + src-tauri/src/app/runner.rs | 3 + src-tauri/src/commands/agent_cmd.rs | 2 + src-tauri/src/commands/aster_agent_cmd.rs | 643 +++++++++++++++- src-tauri/src/commands/mcp_cmd.rs | 59 ++ src-tauri/src/commands/unified_chat_cmd.rs | 94 ++- src-tauri/src/config/tests.rs | 3 + src-tauri/src/services/mod.rs | 1 + .../src/services/web_search_prompt_service.rs | 18 +- .../services/web_search_runtime_service.rs | 202 +++++ src-tauri/tauri.conf.json | 2 +- src/App.tsx | 9 +- .../agent/chat/components/ChatSidebar.tsx | 1 + .../agent/chat/components/EmptyState.test.tsx | 271 +++++++ .../agent/chat/components/EmptyState.tsx | 39 +- .../components/CharacterMention.test.tsx | 299 ++++++++ .../Inputbar/components/CharacterMention.tsx | 124 +-- .../chat/components/Inputbar/index.test.tsx | 165 ++++ .../agent/chat/components/Inputbar/index.tsx | 28 +- .../chat/components/StreamingRenderer.tsx | 7 +- .../agent/chat/components/ToolCallDisplay.tsx | 23 +- src/components/agent/chat/index.test.tsx | 15 +- src/components/agent/chat/index.tsx | 5 +- src/components/api-server/ApiServerPage.tsx | 4 +- .../settings-v2/system/experimental/index.tsx | 168 +++- .../experimental/tool-calling-config.ts | 20 + .../system/experimental/tool-calling.test.ts | 26 + .../system/web-search/index.test.tsx | 147 +++- .../settings-v2/system/web-search/index.tsx | 720 +++++++++++++++++- src/hooks/useTauri.ts | 50 +- src/lib/api/mcp.ts | 28 + src/lib/appVersion.test.ts | 8 +- src/lib/tauri-mock/core.ts | 42 +- vite.config.ts | 5 + 55 files changed, 5003 insertions(+), 348 deletions(-) create mode 100644 src-tauri/crates/core/src/tool_calling.rs create mode 100644 src-tauri/crates/providers/proptest-regressions/providers/antigravity.txt create mode 100644 src-tauri/src/services/web_search_runtime_service.rs create mode 100644 src/components/agent/chat/components/EmptyState.test.tsx create mode 100644 src/components/agent/chat/components/Inputbar/components/CharacterMention.test.tsx create mode 100644 src/components/agent/chat/components/Inputbar/index.test.tsx create mode 100644 src/components/settings-v2/system/experimental/tool-calling-config.ts create mode 100644 src/components/settings-v2/system/experimental/tool-calling.test.ts diff --git a/RELEASE_NOTES.md b/RELEASE_NOTES.md index f8415490d..a67a2f0ba 100644 --- a/RELEASE_NOTES.md +++ b/RELEASE_NOTES.md @@ -1,16 +1,70 @@ -## ProxyCast v0.77.0 +# ProxyCast v0.78.0 Release Notes -### ✨ 新功能 -- 添加可观测性面板,支持响应缓存配置和剪贴板权限指南 (ad61472a) -- 增强崩溃诊断功能,新增调用错误缓冲区、前端崩溃缓冲区和应用版本信息 (ff915523) -- 改进 Aster 运行时、Agent 命令、小说解析和日志检索功能 (68e16d4a) -- API 服务器新增请求去重、响应缓存和能力路由指标 (0b2d3caa) -- 使用崩溃边界包裹应用路由,添加启动时工作区检查,改进小说角色解析 (61270d6f) +## 🎯 主要功能 -### 🐛 修复 -- 为工作区健康检查添加自动重定位功能,使用修复标志进行遥测 (0f1be80b) -- 构建 workspace 目录 4 层健康防护体系,彻底解决路径缺失问题 (2d00911c) +### Tool Calling 2.0 +- 新增 Tool Calling 2.0 配置系统,支持统一控制编程式工具调用 +- 支持动态过滤功能,优先过滤网页抓取噪音 +- 支持原生 input_examples 透传 +- 在实验性设置中新增 Tool Calling 配置面板 -### 🔧 优化与重构 -- 完善工作区健康监控和错误恢复机制 -- 提升应用稳定性和可观测性 +### 联网搜索增强 +- 新增多种联网搜索提供商支持: + - Tavily Search API + - Multi Search Engine v2.0.1(支持 12+ 搜索引擎) + - DuckDuckGo Instant Answer API(无需 API Key,默认启用) + - Bing Search API + - Google Custom Search API +- Multi Search Engine 支持自定义引擎优先级和启用/禁用控制 +- 新增 Web Search Runtime Service 用于运行时搜索能力 + +### MCP 工具增强 +- 改进 MCP 工具管理器,支持更灵活的工具转换 +- 新增 MCP 工具类型定义和转换逻辑 +- 优化 MCP 命令接口 + +### Provider 增强 +- Claude Custom Provider 支持更丰富的工具调用配置 +- OpenAI Custom Provider 增强工具调用能力 +- 改进 Provider Calls 处理逻辑 + +## 🔧 改进 + +### Agent 系统 +- 改进 Aster Agent 状态管理 +- 优化事件转换器逻辑 +- 增强 Agent 命令接口(新增 643 行代码) +- 改进 Unified Chat 命令处理 + +### UI/UX +- 优化 Agent Chat 界面 + - 改进空状态显示 + - 优化角色提及(Character Mention)组件 + - 改进输入栏交互 + - 优化流式渲染和工具调用显示 +- 改进实验性设置界面布局 +- 优化 Web Search 设置界面,支持多提供商配置 + +### 配置系统 +- 新增 `tool_calling` 配置项到核心配置 +- 新增 `WebSearchProvider` 枚举类型 +- 新增 `MultiSearchEngineEntryConfig` 和 `MultiSearchConfig` 配置类型 +- 改进配置测试覆盖 + +## 🐛 修复 +- 修复版本号测试用例(0.77.0 → 0.78.0) +- 改进 Tauri Mock 核心逻辑 +- 优化 API Server 页面 + +## 📊 统计 +- 46 个文件修改 +- +3622 行新增代码 +- -323 行删除代码 + +## 🔗 依赖更新 +- 更新 Aster 依赖到 v0.16.0(通过 git tag) +- 更新 Cargo.lock 依赖 + +--- + +**完整变更**: v0.77.0...v0.78.0 diff --git a/package.json b/package.json index 3d3a49758..0333ad3bf 100644 --- a/package.json +++ b/package.json @@ -1,7 +1,7 @@ { "name": "proxycast", "private": true, - "version": "0.77.0", + "version": "0.78.0", "type": "module", "repository": { "type": "git", diff --git a/src-tauri/Cargo.lock b/src-tauri/Cargo.lock index 42d2988e2..ce2b7f430 100644 --- a/src-tauri/Cargo.lock +++ b/src-tauri/Cargo.lock @@ -369,7 +369,7 @@ checksum = "7c02d123df017efcdfbd739ef81735b36c5ba83ec3c59c80a9d7ecc718f92e50" [[package]] name = "aster-core" -version = "0.15.0" +version = "0.16.0" dependencies = [ "ahash", "anyhow", @@ -461,7 +461,7 @@ dependencies = [ [[package]] name = "aster-models" -version = "0.15.0" +version = "0.16.0" dependencies = [ "serde", "serde_json", @@ -2399,7 +2399,7 @@ dependencies = [ "dtoa-short", "itoa", "matches", - "phf 0.10.1", + "phf 0.8.0", "proc-macro2", "quote", "smallvec", @@ -2415,7 +2415,7 @@ dependencies = [ "cssparser-macros", "dtoa-short", "itoa", - "phf 0.11.3", + "phf 0.8.0", "smallvec", ] @@ -4324,7 +4324,7 @@ dependencies = [ "js-sys", "log", "wasm-bindgen", - "windows-core 0.62.2", + "windows-core 0.56.0", ] [[package]] @@ -5679,7 +5679,7 @@ version = "0.7.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ff32365de1b6743cb203b710788263c44a03de03802daf96092f2da4fe6ba4d7" dependencies = [ - "proc-macro-crate 3.4.0", + "proc-macro-crate 1.3.1", "proc-macro2", "quote", "syn 2.0.114", @@ -6456,7 +6456,9 @@ version = "0.8.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "3dfb61232e34fcb633f43d12c58f83c1df82962dcdfa565a4e866ffc17dafe12" dependencies = [ + "phf_macros 0.8.0", "phf_shared 0.8.0", + "proc-macro-hack", ] [[package]] @@ -6465,9 +6467,7 @@ version = "0.10.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "fabbf1ead8a5bcbc20f5f8b939ee3f5b0f6f281b6ad3468b84656b658b455259" dependencies = [ - "phf_macros 0.10.0", "phf_shared 0.10.0", - "proc-macro-hack", ] [[package]] @@ -6571,12 +6571,12 @@ dependencies = [ [[package]] name = "phf_macros" -version = "0.10.0" +version = "0.8.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "58fdf3184dd560f160dd73922bea2d5cd6e8f064bf4b13110abd81b03697b4e0" +checksum = "7f6fde18ff429ffc8fe78e2bf7f8b7a5a5a6e2a8b58bc5a9ac69198bbda9189c" dependencies = [ - "phf_generator 0.10.0", - "phf_shared 0.10.0", + "phf_generator 0.8.0", + "phf_shared 0.8.0", "proc-macro-hack", "proc-macro2", "quote", @@ -6988,7 +6988,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8a56d757972c98b346a9b766e3f02746cde6dd1cd1d1d563472929fdd74bec4d" dependencies = [ "anyhow", - "itertools 0.14.0", + "itertools 0.12.1", "proc-macro2", "quote", "syn 2.0.114", @@ -7097,7 +7097,7 @@ dependencies = [ [[package]] name = "proxycast-agent" -version = "0.77.0" +version = "0.78.0" dependencies = [ "aster-core", "async-trait", @@ -7121,7 +7121,7 @@ dependencies = [ [[package]] name = "proxycast-config" -version = "0.77.0" +version = "0.78.0" dependencies = [ "async-trait", "parking_lot", @@ -7137,7 +7137,7 @@ dependencies = [ [[package]] name = "proxycast-core" -version = "0.77.0" +version = "0.78.0" dependencies = [ "aster-models", "async-trait", @@ -7177,7 +7177,7 @@ dependencies = [ [[package]] name = "proxycast-credential" -version = "0.77.0" +version = "0.78.0" dependencies = [ "axum 0.7.9", "base64 0.22.1", @@ -7212,7 +7212,7 @@ dependencies = [ [[package]] name = "proxycast-infra" -version = "0.77.0" +version = "0.78.0" dependencies = [ "chrono", "dashmap 5.5.3", @@ -7232,7 +7232,7 @@ dependencies = [ [[package]] name = "proxycast-mcp" -version = "0.77.0" +version = "0.78.0" dependencies = [ "async-trait", "glob", @@ -7263,7 +7263,7 @@ dependencies = [ [[package]] name = "proxycast-processor" -version = "0.77.0" +version = "0.78.0" dependencies = [ "async-trait", "parking_lot", @@ -7282,7 +7282,7 @@ dependencies = [ [[package]] name = "proxycast-providers" -version = "0.77.0" +version = "0.78.0" dependencies = [ "anyhow", "async-stream", @@ -7334,7 +7334,7 @@ dependencies = [ [[package]] name = "proxycast-server" -version = "0.77.0" +version = "0.78.0" dependencies = [ "aster-core", "async-stream", @@ -7379,7 +7379,7 @@ dependencies = [ [[package]] name = "proxycast-server-utils" -version = "0.77.0" +version = "0.78.0" dependencies = [ "axum 0.7.9", "futures", @@ -7394,7 +7394,7 @@ dependencies = [ [[package]] name = "proxycast-services" -version = "0.77.0" +version = "0.78.0" dependencies = [ "anyhow", "aster-core", @@ -7435,7 +7435,7 @@ dependencies = [ [[package]] name = "proxycast-skills" -version = "0.77.0" +version = "0.78.0" dependencies = [ "async-trait", "dirs 5.0.1", @@ -7451,7 +7451,7 @@ dependencies = [ [[package]] name = "proxycast-terminal" -version = "0.77.0" +version = "0.78.0" dependencies = [ "async-trait", "base64 0.22.1", @@ -7478,7 +7478,7 @@ dependencies = [ [[package]] name = "proxycast-websocket" -version = "0.77.0" +version = "0.78.0" dependencies = [ "axum 0.7.9", "chrono", @@ -8957,7 +8957,7 @@ version = "3.1.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8b1fdf65dd6331831494dd616b30351c38e96e45921a27745cf98490458b90bb" dependencies = [ - "dirs 6.0.0", + "dirs 4.0.0", ] [[package]] diff --git a/src-tauri/Cargo.toml b/src-tauri/Cargo.toml index 27ec0f3ad..bc4ff2209 100644 --- a/src-tauri/Cargo.toml +++ b/src-tauri/Cargo.toml @@ -3,7 +3,7 @@ members = ["crates/*"] resolver = "2" [workspace.package] -version = "0.77.0" +version = "0.78.0" edition = "2021" authors = ["coso"] repository = "https://github.com/aiclientproxy/proxycast" @@ -123,11 +123,11 @@ enigo = "0.3" # Aster Agent Framework # 开发时使用本地 aster-rust,CI/CD 使用远程 GitHub 仓库 # 本地开发: path = "../../../astercloud/aster-rust/crates/aster" (相对 src-tauri/) -# CI/CD: git = "https://github.com/astercloud/aster-rust", tag = "v0.15.0" +# CI/CD: git = "https://github.com/astercloud/aster-rust", tag = "v0.16.0" # aster = { package = "aster-core", path = "../../../astercloud/aster-rust/crates/aster" } -aster = { package = "aster-core", git = "https://github.com/astercloud/aster-rust", tag = "v0.15.0" } +aster = { package = "aster-core", git = "https://github.com/astercloud/aster-rust", tag = "v0.16.0" } # 本地开发: aster-models = { path = "../../../astercloud/aster-rust/crates/aster-models" } -aster-models = { git = "https://github.com/astercloud/aster-rust", tag = "v0.15.0" } +aster-models = { git = "https://github.com/astercloud/aster-rust", tag = "v0.16.0" } # MCP (Model Context Protocol) rmcp = { version = "0.12.0", features = ["client", "transport-io", "transport-child-process"] } diff --git a/src-tauri/crates/agent/src/aster_state.rs b/src-tauri/crates/agent/src/aster_state.rs index b84fbce80..ee8163af3 100644 --- a/src-tauri/crates/agent/src/aster_state.rs +++ b/src-tauri/crates/agent/src/aster_state.rs @@ -585,6 +585,9 @@ impl AsterAgentState { timeout: None, bundled: Some(false), available_tools: Vec::new(), + deferred_loading: false, + always_expose_tools: Vec::new(), + allowed_caller: None, }; // 注册到 ExtensionManager diff --git a/src-tauri/crates/agent/src/event_converter.rs b/src-tauri/crates/agent/src/event_converter.rs index 4d1f58a86..143b42800 100644 --- a/src-tauri/crates/agent/src/event_converter.rs +++ b/src-tauri/crates/agent/src/event_converter.rs @@ -5,6 +5,7 @@ use aster::agents::AgentEvent; use aster::conversation::message::{ActionRequiredData, Message, MessageContent}; +use regex::Regex; use serde::{Deserialize, Serialize}; /// 从工具结果中提取文本内容 @@ -74,12 +75,71 @@ fn extract_tool_result_text(result: &T) -> String { collect_tool_result_text(&json, &mut parts); let deduped = dedupe_preserve_order(parts); if !deduped.is_empty() { - return deduped.join("\n"); + return maybe_filter_web_content(&deduped.join("\n")); } } String::new() } +fn dynamic_filtering_enabled() -> bool { + proxycast_core::tool_calling::tool_calling_dynamic_filtering_enabled() +} + +fn maybe_filter_web_content(raw: &str) -> String { + if !dynamic_filtering_enabled() { + return raw.to_string(); + } + + let lowered = raw.to_ascii_lowercase(); + let looks_like_html = + (lowered.contains("")) + && raw.len() > 4_000; + if !looks_like_html { + return raw.to_string(); + } + + let script_re = Regex::new(r"(?is)]*>.*?").ok(); + let style_re = Regex::new(r"(?is)]*>.*?").ok(); + let tag_re = Regex::new(r"(?is)<[^>]+>").ok(); + let space_re = Regex::new(r"[ \t]{2,}").ok(); + let newline_re = Regex::new(r"\n{3,}").ok(); + + let mut cleaned = raw.to_string(); + if let Some(re) = script_re.as_ref() { + cleaned = re.replace_all(&cleaned, " ").to_string(); + } + if let Some(re) = style_re.as_ref() { + cleaned = re.replace_all(&cleaned, " ").to_string(); + } + if let Some(re) = tag_re.as_ref() { + cleaned = re.replace_all(&cleaned, "\n").to_string(); + } + if let Some(re) = space_re.as_ref() { + cleaned = re.replace_all(&cleaned, " ").to_string(); + } + if let Some(re) = newline_re.as_ref() { + cleaned = re.replace_all(&cleaned, "\n\n").to_string(); + } + cleaned = cleaned + .lines() + .map(str::trim) + .filter(|line| !line.is_empty()) + .collect::>() + .join("\n"); + + const MAX_FILTERED_CHARS: usize = 8_000; + if cleaned.chars().count() > MAX_FILTERED_CHARS { + let shortened = cleaned.chars().take(MAX_FILTERED_CHARS).collect::(); + return format!( + "{}\n\n[dynamic_filtering] 内容已裁剪,原始长度 {} 字符", + shortened, + cleaned.chars().count() + ); + } + + cleaned +} + #[derive(Debug, Clone)] struct ExtractedToolResult { output: String, @@ -801,4 +861,16 @@ mod tests { assert_eq!(extracted.images.len(), 1); assert_eq!(extracted.images[0].src, "data:image/png;base64,aGVsbG8="); } + + #[test] + fn test_maybe_filter_web_content_should_strip_html_noise() { + let html = format!( + "{}", + "正文".repeat(2500) + ); + let filtered = maybe_filter_web_content(&html); + assert!(!filtered.to_ascii_lowercase().contains(">, + /// 允许调用方(assistant/code_execution/tool_search) + #[serde(default, skip_serializing_if = "Option::is_none")] + pub allowed_callers: Option>, + /// 是否延迟加载(默认不注入上下文) + #[serde(default, skip_serializing_if = "Option::is_none")] + pub deferred_loading: Option, } /// Agent 配置 diff --git a/src-tauri/crates/core/src/config/mod.rs b/src-tauri/crates/core/src/config/mod.rs index 6751094a5..67fce4cee 100644 --- a/src-tauri/crates/core/src/config/mod.rs +++ b/src-tauri/crates/core/src/config/mod.rs @@ -28,12 +28,13 @@ pub use types::{ HeartbeatSecurityConfig, HeartbeatSettings, HintRouteSettingsEntry, HintRouterSettings, ImageGenConfig, InjectionRuleConfig, InjectionSettings, LoggingConfig, MemoryAutoConfig, MemoryConfig, MemoryProfileConfig, MemoryResolveConfig, MemorySourcesConfig, ModelInfo, - ModelsConfig, NativeAgentConfig, NavigationConfig, OpenAIAsrConfig, PairingSettings, - ProviderConfig, ProviderModelsConfig, ProvidersConfig, QuotaExceededConfig, RateLimitSettings, - RemoteManagementConfig, ResponseCacheSettings, RetrySettings, RoutingConfig, - ScreenshotChatConfig, SearchEngine, ServerConfig, TaskSchedule, TlsConfig, UpdateCheckConfig, - UserProfile, VertexApiKeyEntry, VertexModelAlias, VoiceConfig, VoiceInputConfig, - VoiceInstruction, VoiceOutputConfig, VoiceOutputMode, VoiceProcessorConfig, WebSearchConfig, + ModelsConfig, MultiSearchConfig, MultiSearchEngineEntryConfig, NativeAgentConfig, + NavigationConfig, OpenAIAsrConfig, PairingSettings, ProviderConfig, ProviderModelsConfig, + ProvidersConfig, QuotaExceededConfig, RateLimitSettings, RemoteManagementConfig, + ResponseCacheSettings, RetrySettings, RoutingConfig, ScreenshotChatConfig, SearchEngine, + ServerConfig, TaskSchedule, TlsConfig, ToolCallingConfig, UpdateCheckConfig, UserProfile, + VertexApiKeyEntry, VertexModelAlias, VoiceConfig, VoiceInputConfig, VoiceInstruction, + VoiceOutputConfig, VoiceOutputMode, VoiceProcessorConfig, WebSearchConfig, WebSearchProvider, WhisperLocalConfig, WhisperModelSize, WorkspaceSandboxConfig, XunfeiConfig, DEFAULT_API_KEY, }; pub use yaml::{load_config, save_config, ConfigError, ConfigManager, YamlService}; diff --git a/src-tauri/crates/core/src/config/tests.rs b/src-tauri/crates/core/src/config/tests.rs index 66fac4e84..c420742d3 100644 --- a/src-tauri/crates/core/src/config/tests.rs +++ b/src-tauri/crates/core/src/config/tests.rs @@ -191,6 +191,7 @@ fn arb_config() -> impl Strategy { agent: crate::config::NativeAgentConfig::default(), language: "zh".to_string(), experimental: crate::config::ExperimentalFeatures::default(), + tool_calling: crate::config::ToolCallingConfig::default(), content_creator: ContentCreatorConfig::default(), navigation: NavigationConfig::default(), }) @@ -432,6 +433,7 @@ fn arb_valid_config() -> impl Strategy { agent: crate::config::NativeAgentConfig::default(), language: "zh".to_string(), experimental: crate::config::ExperimentalFeatures::default(), + tool_calling: crate::config::ToolCallingConfig::default(), content_creator: ContentCreatorConfig::default(), navigation: NavigationConfig::default(), }) @@ -483,6 +485,7 @@ fn arb_invalid_config() -> impl Strategy { agent: crate::config::NativeAgentConfig::default(), language: "zh".to_string(), experimental: crate::config::ExperimentalFeatures::default(), + tool_calling: crate::config::ToolCallingConfig::default(), content_creator: ContentCreatorConfig::default(), navigation: NavigationConfig::default(), }; diff --git a/src-tauri/crates/core/src/config/types.rs b/src-tauri/crates/core/src/config/types.rs index 240a2fcb7..17f115208 100644 --- a/src-tauri/crates/core/src/config/types.rs +++ b/src-tauri/crates/core/src/config/types.rs @@ -399,6 +399,9 @@ pub struct Config { /// 实验室功能配置 #[serde(default)] pub experimental: ExperimentalFeatures, + /// Tool Calling 2.0 配置 + #[serde(default)] + pub tool_calling: ToolCallingConfig, /// 内容创作配置 #[serde(default)] pub content_creator: ContentCreatorConfig, @@ -745,6 +748,40 @@ pub struct ExperimentalFeatures { pub voice_input: VoiceInputConfig, } +/// Tool Calling 2.0 配置 +/// +/// 统一控制编程式工具调用、动态过滤与 input examples 透传行为。 +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] +pub struct ToolCallingConfig { + /// 是否启用 Tool Calling 2.0 能力 + #[serde(default = "default_tool_calling_enabled")] + pub enabled: bool, + /// 是否启用动态过滤(优先过滤网页抓取噪音) + #[serde(default = "default_tool_calling_dynamic_filtering_enabled")] + pub dynamic_filtering: bool, + /// 是否启用原生 input_examples 透传 + #[serde(default)] + pub native_input_examples: bool, +} + +fn default_tool_calling_enabled() -> bool { + true +} + +fn default_tool_calling_dynamic_filtering_enabled() -> bool { + true +} + +impl Default for ToolCallingConfig { + fn default() -> Self { + Self { + enabled: default_tool_calling_enabled(), + dynamic_filtering: default_tool_calling_dynamic_filtering_enabled(), + native_input_examples: false, + } + } +} + // ============ 语音输入功能配置类型 ============ /// 语音输入功能配置 @@ -1921,6 +1958,7 @@ impl Default for Config { models: ModelsConfig::default(), agent: NativeAgentConfig::default(), experimental: ExperimentalFeatures::default(), + tool_calling: ToolCallingConfig::default(), content_creator: ContentCreatorConfig::default(), navigation: NavigationConfig::default(), chat_appearance: ChatAppearanceConfig::default(), @@ -1954,12 +1992,143 @@ pub enum SearchEngine { Xiaohongshu, } +/// 联网搜索提供商类型 +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq, Default)] +#[serde(rename_all = "snake_case")] +pub enum WebSearchProvider { + /// Tavily Search API + Tavily, + /// Multi Search Engine v2.0.1 + MultiSearchEngine, + /// DuckDuckGo Instant Answer API(无需 Key) + #[default] + DuckduckgoInstant, + /// Bing Search API + BingSearchApi, + /// Google Custom Search API + GoogleCustomSearch, +} + +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)] +pub struct MultiSearchEngineEntryConfig { + /// 引擎标识名 + pub name: String, + /// 搜索 URL 模板,必须包含 {query} + pub url_template: String, + /// 是否启用该引擎 + #[serde(default = "default_mse_engine_enabled")] + pub enabled: bool, +} + +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)] +pub struct MultiSearchConfig { + /// 引擎优先级(按名称) + #[serde(default)] + pub priority: Vec, + /// 自定义/覆盖引擎列表 + #[serde(default = "default_multi_search_engines")] + pub engines: Vec, + /// 每个引擎最大结果数 + #[serde(default = "default_mse_max_results_per_engine")] + pub max_results_per_engine: usize, + /// 最终聚合最大结果数 + #[serde(default = "default_mse_max_total_results")] + pub max_total_results: usize, + /// 每个引擎请求超时(毫秒) + #[serde(default = "default_mse_timeout_ms")] + pub timeout_ms: u64, +} + +impl Default for MultiSearchConfig { + fn default() -> Self { + Self { + priority: vec![], + engines: default_multi_search_engines(), + max_results_per_engine: default_mse_max_results_per_engine(), + max_total_results: default_mse_max_total_results(), + timeout_ms: default_mse_timeout_ms(), + } + } +} + +fn default_mse_engine_enabled() -> bool { + true +} + +fn default_mse_max_results_per_engine() -> usize { + 5 +} + +fn default_mse_max_total_results() -> usize { + 20 +} + +fn default_mse_timeout_ms() -> u64 { + 4000 +} + +fn default_multi_search_engines() -> Vec { + vec![ + ("google", "https://www.google.com/search?q={query}"), + ("bing", "https://www.bing.com/search?q={query}"), + ("duckduckgo", "https://duckduckgo.com/?q={query}"), + ("yahoo", "https://search.yahoo.com/search?p={query}"), + ("baidu", "https://www.baidu.com/s?wd={query}"), + ("yandex", "https://yandex.com/search/?text={query}"), + ("ecosia", "https://www.ecosia.org/search?q={query}"), + ("brave", "https://search.brave.com/search?q={query}"), + ( + "startpage", + "https://www.startpage.com/do/search?query={query}", + ), + ("qwant", "https://www.qwant.com/?q={query}&t=web"), + ("sogou", "https://www.sogou.com/web?query={query}"), + ("so360", "https://www.so.com/s?q={query}"), + ("aol", "https://search.aol.com/aol/search?q={query}"), + ("ask", "https://www.ask.com/web?q={query}"), + ( + "naver", + "https://search.naver.com/search.naver?query={query}", + ), + ("seznam", "https://search.seznam.cz/?q={query}"), + ("dogpile", "https://www.dogpile.com/serp?q={query}"), + ] + .into_iter() + .map(|(name, url_template)| MultiSearchEngineEntryConfig { + name: name.to_string(), + url_template: url_template.to_string(), + enabled: true, + }) + .collect() +} + /// 网络搜索配置 #[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq, Default)] pub struct WebSearchConfig { /// 默认搜索引擎偏好 #[serde(default)] pub engine: SearchEngine, + /// 联网搜索提供商 + #[serde(default)] + pub provider: WebSearchProvider, + /// 提供商回退优先级 + #[serde(default)] + pub provider_priority: Vec, + /// Tavily Search API Key + #[serde(default, skip_serializing_if = "Option::is_none")] + pub tavily_api_key: Option, + /// Bing Search API Key + #[serde(default, skip_serializing_if = "Option::is_none")] + pub bing_search_api_key: Option, + /// Google Search API Key + #[serde(default, skip_serializing_if = "Option::is_none")] + pub google_search_api_key: Option, + /// Google Search Engine ID (CSE CX) + #[serde(default, skip_serializing_if = "Option::is_none")] + pub google_search_engine_id: Option, + /// Multi Search Engine 配置 + #[serde(default)] + pub multi_search: MultiSearchConfig, } /// 聊天外观配置 @@ -2600,6 +2769,31 @@ mod unit_tests { assert_eq!(parsed, config); } + #[test] + fn test_tool_calling_config_default() { + let config = ToolCallingConfig::default(); + assert!(config.enabled); + assert!(config.dynamic_filtering); + assert!(!config.native_input_examples); + } + + #[test] + fn test_tool_calling_config_serialization() { + let config = ToolCallingConfig { + enabled: false, + dynamic_filtering: false, + native_input_examples: true, + }; + + let yaml = serde_yaml::to_string(&config).unwrap(); + assert!(yaml.contains("enabled: false")); + assert!(yaml.contains("dynamic_filtering: false")); + assert!(yaml.contains("native_input_examples: true")); + + let parsed: ToolCallingConfig = serde_yaml::from_str(&yaml).unwrap(); + assert_eq!(parsed, config); + } + #[test] fn test_config_with_experimental() { let config = Config::default(); @@ -2614,6 +2808,9 @@ mod unit_tests { config.experimental.voice_input.shortcut, "CommandOrControl+Shift+V" ); + assert!(config.tool_calling.enabled); + assert!(config.tool_calling.dynamic_filtering); + assert!(!config.tool_calling.native_input_examples); } #[test] diff --git a/src-tauri/crates/core/src/lib.rs b/src-tauri/crates/core/src/lib.rs index 382fb490e..1c3042fb5 100644 --- a/src-tauri/crates/core/src/lib.rs +++ b/src-tauri/crates/core/src/lib.rs @@ -36,6 +36,7 @@ pub mod orchestrator; pub mod plugin; pub mod session; pub mod session_files; +pub mod tool_calling; // 类型模块(纯数据类型,供 database 等模块使用) pub mod agent; diff --git a/src-tauri/crates/core/src/tool_calling.rs b/src-tauri/crates/core/src/tool_calling.rs new file mode 100644 index 000000000..115f2e33f --- /dev/null +++ b/src-tauri/crates/core/src/tool_calling.rs @@ -0,0 +1,312 @@ +//! Tool Calling 2.0 运行时配置 +//! +//! 通过内存态开关提供跨 crate 的统一读取入口,并保留环境变量兜底覆盖。 + +use crate::config::{Config, ToolCallingConfig}; +use serde_json::{Map, Value}; +use std::sync::atomic::{AtomicBool, Ordering}; + +const ENV_TOOLCALL_V2_ENABLED: &str = "PROXYCAST_TOOLCALL_V2_ENABLED"; +const ENV_TOOLCALL_V2_DYNAMIC_FILTERING: &str = "PROXYCAST_TOOLCALL_V2_DYNAMIC_FILTERING"; +const ENV_TOOLCALL_V2_NATIVE_INPUT_EXAMPLES: &str = "PROXYCAST_TOOLCALL_V2_NATIVE_INPUT_EXAMPLES"; + +static TOOLCALL_RUNTIME_INITIALIZED: AtomicBool = AtomicBool::new(false); +static TOOLCALL_V2_ENABLED: AtomicBool = AtomicBool::new(true); +static TOOLCALL_DYNAMIC_FILTERING_ENABLED: AtomicBool = AtomicBool::new(true); +static TOOLCALL_NATIVE_INPUT_EXAMPLES_ENABLED: AtomicBool = AtomicBool::new(false); + +fn parse_bool_env(name: &str) -> Option { + let raw = std::env::var(name).ok()?; + match raw.trim().to_ascii_lowercase().as_str() { + "1" | "true" | "yes" | "on" => Some(true), + "0" | "false" | "no" | "off" => Some(false), + _ => None, + } +} + +/// 将配置应用到进程内运行时开关。 +pub fn apply_tool_calling_runtime_config(config: &Config) { + apply_tool_calling_runtime_config_with_flags(&config.tool_calling); +} + +/// 将 Tool Calling 配置应用到进程内运行时开关。 +pub fn apply_tool_calling_runtime_config_with_flags(flags: &ToolCallingConfig) { + TOOLCALL_V2_ENABLED.store(flags.enabled, Ordering::Release); + TOOLCALL_DYNAMIC_FILTERING_ENABLED.store(flags.dynamic_filtering, Ordering::Release); + TOOLCALL_NATIVE_INPUT_EXAMPLES_ENABLED.store(flags.native_input_examples, Ordering::Release); + TOOLCALL_RUNTIME_INITIALIZED.store(true, Ordering::Release); +} + +/// Tool Calling 2.0 总开关。 +pub fn tool_calling_v2_enabled() -> bool { + if let Some(value) = parse_bool_env(ENV_TOOLCALL_V2_ENABLED) { + return value; + } + if TOOLCALL_RUNTIME_INITIALIZED.load(Ordering::Acquire) { + return TOOLCALL_V2_ENABLED.load(Ordering::Acquire); + } + true +} + +/// Tool Calling 动态过滤开关。 +pub fn tool_calling_dynamic_filtering_enabled() -> bool { + if let Some(value) = parse_bool_env(ENV_TOOLCALL_V2_DYNAMIC_FILTERING) { + return value; + } + if TOOLCALL_RUNTIME_INITIALIZED.load(Ordering::Acquire) { + return TOOLCALL_DYNAMIC_FILTERING_ENABLED.load(Ordering::Acquire); + } + true +} + +/// Tool Calling 原生 input examples 透传开关。 +pub fn tool_calling_native_input_examples_enabled() -> bool { + if let Some(value) = parse_bool_env(ENV_TOOLCALL_V2_NATIVE_INPUT_EXAMPLES) { + return value; + } + if TOOLCALL_RUNTIME_INITIALIZED.load(Ordering::Acquire) { + return TOOLCALL_NATIVE_INPUT_EXAMPLES_ENABLED.load(Ordering::Acquire); + } + false +} + +fn schema_read_examples(schema: &Value) -> Vec { + let extension = schema + .get("x-proxycast") + .or_else(|| schema.get("x_proxycast")) + .unwrap_or(schema); + + extension + .get("input_examples") + .or_else(|| extension.get("inputExamples")) + .and_then(|v| v.as_array()) + .map(|arr| arr.iter().filter(|v| !v.is_null()).cloned().collect()) + .unwrap_or_default() +} + +fn pick_example_value(field_name: &str, schema: &Value, depth: usize) -> Value { + if let Some(enum_values) = schema.get("enum").and_then(|v| v.as_array()) { + if let Some(first) = enum_values.first() { + return first.clone(); + } + } + + if let Some(one_of) = schema + .get("oneOf") + .or_else(|| schema.get("anyOf")) + .and_then(|v| v.as_array()) + .and_then(|arr| arr.first()) + { + return pick_example_value(field_name, one_of, depth + 1); + } + + let field_name_lc = field_name.to_ascii_lowercase(); + let field_type = schema + .get("type") + .and_then(|v| v.as_str()) + .unwrap_or("string"); + + match field_type { + "boolean" => Value::Bool(true), + "integer" => { + if field_name_lc.contains("count") + || field_name_lc.contains("limit") + || field_name_lc.contains("top") + { + Value::Number(3.into()) + } else { + Value::Number(1.into()) + } + } + "number" => Value::Number(serde_json::Number::from_f64(0.5).unwrap_or_else(|| 0.into())), + "array" => { + if depth >= 2 { + return Value::Array(Vec::new()); + } + let item_schema = schema.get("items").unwrap_or(&Value::Null); + Value::Array(vec![pick_example_value(field_name, item_schema, depth + 1)]) + } + "object" => { + if depth >= 2 { + return Value::Object(Map::new()); + } + synthesize_example_from_schema(schema, depth + 1) + .unwrap_or_else(|| Value::Object(Map::new())) + } + _ => { + if field_name_lc.contains("url") || field_name_lc.contains("link") { + Value::String("https://example.com".to_string()) + } else if field_name_lc.contains("query") || field_name_lc.contains("keyword") { + Value::String("latest ai agent tool calling updates".to_string()) + } else if field_name_lc.contains("prompt") + || field_name_lc.contains("instruction") + || field_name_lc.contains("question") + { + Value::String("请提炼三条关键信息并给出结论".to_string()) + } else if field_name_lc.contains("id") { + Value::String("example-id".to_string()) + } else if field_name_lc.contains("path") { + Value::String("/tmp/example".to_string()) + } else { + Value::String("example".to_string()) + } + } + } +} + +fn synthesize_example_from_schema(schema: &Value, depth: usize) -> Option { + let properties = schema.get("properties").and_then(|v| v.as_object())?; + if properties.is_empty() { + return None; + } + + let required = schema + .get("required") + .and_then(|v| v.as_array()) + .map(|arr| { + arr.iter() + .filter_map(|v| v.as_str()) + .map(|v| v.to_string()) + .collect::>() + }) + .unwrap_or_default(); + + let mut keys = required.clone(); + let mut optional_keys = properties.keys().cloned().collect::>(); + optional_keys.sort(); + for key in optional_keys { + if !keys.contains(&key) { + keys.push(key); + } + } + + let max_fields = if depth == 0 { 6 } else { 4 }; + let mut out = Map::new(); + for key in keys.into_iter().take(max_fields) { + if let Some(field_schema) = properties.get(&key) { + out.insert(key.clone(), pick_example_value(&key, field_schema, depth)); + } + } + + Some(Value::Object(out)) +} + +/// 解析工具 schema 内配置的 input_examples。 +pub fn configured_tool_input_examples(schema: &Value) -> Vec { + schema_read_examples(schema) +} + +/// 获取工具可用的 input_examples(优先配置,内置工具缺省时按 schema 生成)。 +pub fn resolve_tool_input_examples(tool_name: &str, schema: &Value) -> Vec { + let configured = schema_read_examples(schema); + if !configured.is_empty() { + return configured; + } + + let normalized = tool_name.trim().to_ascii_lowercase(); + let built_in = matches!( + normalized.as_str(), + "websearch" | "webfetch" | "three_stage_workflow" | "tool_search" + ); + if !built_in { + return Vec::new(); + } + + synthesize_example_from_schema(schema, 0) + .map(|v| vec![v]) + .unwrap_or_default() +} + +#[cfg(test)] +mod tests { + use super::*; + use std::sync::{Mutex, OnceLock}; + + fn env_lock() -> std::sync::MutexGuard<'static, ()> { + static LOCK: OnceLock> = OnceLock::new(); + LOCK.get_or_init(|| Mutex::new(())).lock().unwrap() + } + + fn clear_tool_calling_envs() { + std::env::remove_var(ENV_TOOLCALL_V2_ENABLED); + std::env::remove_var(ENV_TOOLCALL_V2_DYNAMIC_FILTERING); + std::env::remove_var(ENV_TOOLCALL_V2_NATIVE_INPUT_EXAMPLES); + } + + #[test] + fn test_runtime_flags_apply_and_read() { + let _guard = env_lock(); + clear_tool_calling_envs(); + apply_tool_calling_runtime_config_with_flags(&ToolCallingConfig { + enabled: false, + dynamic_filtering: false, + native_input_examples: true, + }); + + assert!(!tool_calling_v2_enabled()); + assert!(!tool_calling_dynamic_filtering_enabled()); + assert!(tool_calling_native_input_examples_enabled()); + } + + #[test] + fn test_env_overrides_runtime_flags() { + let _guard = env_lock(); + clear_tool_calling_envs(); + + apply_tool_calling_runtime_config_with_flags(&ToolCallingConfig { + enabled: false, + dynamic_filtering: false, + native_input_examples: false, + }); + + std::env::set_var(ENV_TOOLCALL_V2_ENABLED, "true"); + std::env::set_var(ENV_TOOLCALL_V2_DYNAMIC_FILTERING, "1"); + std::env::set_var(ENV_TOOLCALL_V2_NATIVE_INPUT_EXAMPLES, "on"); + + assert!(tool_calling_v2_enabled()); + assert!(tool_calling_dynamic_filtering_enabled()); + assert!(tool_calling_native_input_examples_enabled()); + + clear_tool_calling_envs(); + } + + #[test] + fn test_resolve_tool_input_examples_prefers_configured_examples() { + let schema = serde_json::json!({ + "type": "object", + "x-proxycast": { + "input_examples": [{"query": "rust async"}] + } + }); + let examples = resolve_tool_input_examples("WebSearch", &schema); + assert_eq!(examples, vec![serde_json::json!({"query":"rust async"})]); + } + + #[test] + fn test_resolve_tool_input_examples_generates_builtin_examples_from_schema() { + let schema = serde_json::json!({ + "type": "object", + "properties": { + "query": {"type":"string"}, + "limit": {"type":"integer"} + }, + "required": ["query"] + }); + let examples = resolve_tool_input_examples("WebSearch", &schema); + assert_eq!(examples.len(), 1); + let example = examples[0].as_object().cloned().unwrap_or_default(); + assert!(example.contains_key("query")); + } + + #[test] + fn test_resolve_tool_input_examples_ignores_non_builtin_without_config() { + let schema = serde_json::json!({ + "type": "object", + "properties": { + "query": {"type":"string"} + } + }); + let examples = resolve_tool_input_examples("docs_search", &schema); + assert!(examples.is_empty()); + } +} diff --git a/src-tauri/crates/mcp/src/manager.rs b/src-tauri/crates/mcp/src/manager.rs index c9b1242ca..cab3efda3 100644 --- a/src-tauri/crates/mcp/src/manager.rs +++ b/src-tauri/crates/mcp/src/manager.rs @@ -27,7 +27,7 @@ #![allow(dead_code)] use proxycast_core::DynEmitter; -use std::collections::HashMap; +use std::collections::{HashMap, HashSet}; use std::process::Stdio; use std::sync::Arc; use std::time::Duration; @@ -42,6 +42,15 @@ use rmcp::ServiceExt; use crate::client::McpClientWrapper; use crate::types::*; +#[derive(Debug, Default)] +struct ToolMetadataExtraction { + deferred_loading: Option, + always_visible: Option, + allowed_callers: Option>, + input_examples: Option>, + tags: Option>, +} + /// MCP 客户端管理器 /// /// 负责管理所有 MCP 服务器的连接和生命周期。 @@ -720,6 +729,8 @@ impl McpClientManager { "获取服务器工具列表成功" ); for tool in tools { + let input_schema = serde_json::Value::Object((*tool.input_schema).clone()); + let metadata = Self::extract_tool_metadata(&input_schema); all_tools.push(McpToolDefinition { name: tool.name.to_string(), description: tool @@ -727,8 +738,13 @@ impl McpClientManager { .clone() .map(|s| s.to_string()) .unwrap_or_default(), - input_schema: serde_json::Value::Object((*tool.input_schema).clone()), + input_schema, server_name: server_name.clone(), + deferred_loading: metadata.deferred_loading, + always_visible: metadata.always_visible, + allowed_callers: metadata.allowed_callers, + input_examples: metadata.input_examples, + tags: metadata.tags, }); } } @@ -757,6 +773,179 @@ impl McpClientManager { Ok(resolved_tools) } + /// 根据上下文过滤工具列表 + /// + /// - `caller`: 调用方(assistant/code_execution/tool_search) + /// - `include_deferred`: 是否包含延迟加载工具 + pub async fn list_tools_for_context( + &self, + caller: Option<&str>, + include_deferred: bool, + ) -> Result, McpError> { + let caller = caller + .map(str::trim) + .filter(|s| !s.is_empty()) + .map(|s| s.to_ascii_lowercase()); + let tools = self.list_tools().await?; + + let filtered = tools + .into_iter() + .filter(|tool| { + // deferred_loading=true 且不是 always_visible 时,默认不注入上下文 + if !include_deferred + && tool.deferred_loading.unwrap_or(false) + && !tool.always_visible.unwrap_or(false) + { + return false; + } + + // caller 不在 allowed_callers 时,隐藏该工具 + if let (Some(caller), Some(allowed)) = (&caller, tool.allowed_callers.as_ref()) { + let allowed_set: HashSet = allowed + .iter() + .map(|v| v.trim().to_ascii_lowercase()) + .filter(|v| !v.is_empty()) + .collect(); + if !allowed_set.is_empty() && !allowed_set.contains(caller) { + return false; + } + } + + true + }) + .collect(); + + Ok(filtered) + } + + /// 搜索工具 + /// + /// 搜索默认包含 deferred 工具,便于模型通过 tool_search 检索按需加载。 + pub async fn search_tools( + &self, + query: &str, + limit: usize, + caller: Option<&str>, + ) -> Result, McpError> { + let query = query.trim().to_ascii_lowercase(); + let limit = limit.clamp(1, 100); + let mut tools = self.list_tools_for_context(caller, true).await?; + + // 空查询:优先 always_visible,再按名称排序返回前 N + if query.is_empty() { + tools.sort_by(|a, b| { + let a_visible = a.always_visible.unwrap_or(false); + let b_visible = b.always_visible.unwrap_or(false); + b_visible + .cmp(&a_visible) + .then_with(|| a.name.to_lowercase().cmp(&b.name.to_lowercase())) + }); + tools.truncate(limit); + return Ok(tools); + } + + let mut scored: Vec<(i32, McpToolDefinition)> = tools + .into_iter() + .filter_map(|tool| { + let score = Self::score_tool_match(&tool, &query); + (score > 0).then_some((score, tool)) + }) + .collect(); + + scored.sort_by(|(score_a, tool_a), (score_b, tool_b)| { + score_b + .cmp(score_a) + .then_with(|| tool_a.name.to_lowercase().cmp(&tool_b.name.to_lowercase())) + }); + + let mut result = scored + .into_iter() + .take(limit) + .map(|(_, tool)| tool) + .collect::>(); + result.truncate(limit); + Ok(result) + } + + fn extract_tool_metadata(input_schema: &serde_json::Value) -> ToolMetadataExtraction { + fn read_bool(root: &serde_json::Value, key: &str) -> Option { + root.get(key).and_then(|v| v.as_bool()) + } + + fn read_string_vec(root: &serde_json::Value, key: &str) -> Option> { + let arr = root.get(key)?.as_array()?; + let values = arr + .iter() + .filter_map(|v| v.as_str()) + .map(|v| v.trim().to_string()) + .filter(|v| !v.is_empty()) + .collect::>(); + (!values.is_empty()).then_some(values) + } + + fn read_examples(root: &serde_json::Value, key: &str) -> Option> { + let arr = root.get(key)?.as_array()?; + let values = arr + .iter() + .filter(|v| !v.is_null()) + .cloned() + .collect::>(); + (!values.is_empty()).then_some(values) + } + + let extension = input_schema + .get("x-proxycast") + .or_else(|| input_schema.get("x_proxycast")) + .unwrap_or(input_schema); + + ToolMetadataExtraction { + deferred_loading: read_bool(extension, "deferred_loading") + .or_else(|| read_bool(extension, "deferredLoading")), + always_visible: read_bool(extension, "always_visible") + .or_else(|| read_bool(extension, "alwaysVisible")), + allowed_callers: read_string_vec(extension, "allowed_callers") + .or_else(|| read_string_vec(extension, "allowedCallers")), + input_examples: read_examples(extension, "input_examples") + .or_else(|| read_examples(extension, "inputExamples")), + tags: read_string_vec(extension, "tags"), + } + } + + fn score_tool_match(tool: &McpToolDefinition, query: &str) -> i32 { + let name = tool.name.to_ascii_lowercase(); + let description = tool.description.to_ascii_lowercase(); + + let mut score = 0; + if name == query { + score += 120; + } else if name.starts_with(query) { + score += 90; + } else if name.contains(query) { + score += 70; + } + + if description.contains(query) { + score += 40; + } + + if let Some(tags) = tool.tags.as_ref() { + for tag in tags { + let tag = tag.to_ascii_lowercase(); + if tag == query { + score += 35; + } else if tag.contains(query) { + score += 20; + } + } + } + + if tool.always_visible.unwrap_or(false) { + score += 5; + } + + score + } + /// 解决工具名称冲突 /// /// 当多个服务器提供同名工具时,为冲突的工具名称添加服务器前缀。 @@ -769,8 +958,6 @@ impl McpClientManager { /// /// 返回解决冲突后的工具列表。 fn resolve_tool_name_conflicts(tools: Vec) -> Vec { - use std::collections::HashSet; - // 统计每个工具名称出现的次数 let mut name_counts: HashMap = HashMap::new(); for tool in &tools { @@ -812,6 +999,46 @@ impl McpClientManager { /// /// 返回工具调用结果。 /// + /// # 实现步骤(Task 4.3) + /// + /// 1. 解析工具名称,确定目标服务器 + /// 2. 路由到正确的客户端 + /// 3. 执行工具调用 + /// 4. 转换结果为 McpToolResult + /// 5. 返回结果 + pub async fn call_tool_with_caller( + &self, + tool_name: &str, + arguments: serde_json::Value, + caller: Option<&str>, + ) -> Result { + let caller = caller + .map(str::trim) + .filter(|s| !s.is_empty()) + .map(|s| s.to_ascii_lowercase()); + + if let Some(caller) = caller { + let tools = self.list_tools().await?; + if let Some(tool) = tools.iter().find(|t| t.name == tool_name) { + if let Some(allowed) = tool.allowed_callers.as_ref() { + let allowed_set: HashSet = allowed + .iter() + .map(|v| v.trim().to_ascii_lowercase()) + .filter(|v| !v.is_empty()) + .collect(); + if !allowed_set.is_empty() && !allowed_set.contains(&caller) { + return Err(McpError::ToolCallFailed(format!( + "调用方 '{}' 无权调用工具 '{}'", + caller, tool_name + ))); + } + } + } + } + + self.call_tool(tool_name, arguments).await + } + /// # 实现步骤(Task 4.3) /// /// 1. 解析工具名称,确定目标服务器 @@ -1523,6 +1750,20 @@ mod tests { McpClientWrapper::new(name.to_string(), create_test_config(), None) } + fn create_test_tool(name: &str, description: &str, server_name: &str) -> McpToolDefinition { + McpToolDefinition { + name: name.to_string(), + description: description.to_string(), + input_schema: serde_json::json!({}), + server_name: server_name.to_string(), + deferred_loading: None, + always_visible: None, + allowed_callers: None, + input_examples: None, + tags: None, + } + } + #[test] fn test_manager_creation() { let manager = McpClientManager::new(None); @@ -1672,18 +1913,8 @@ mod tests { // 更新缓存 let tools = vec![ - McpToolDefinition { - name: "tool1".to_string(), - description: "Test tool 1".to_string(), - input_schema: serde_json::json!({}), - server_name: "server1".to_string(), - }, - McpToolDefinition { - name: "tool2".to_string(), - description: "Test tool 2".to_string(), - input_schema: serde_json::json!({}), - server_name: "server1".to_string(), - }, + create_test_tool("tool1", "Test tool 1", "server1"), + create_test_tool("tool2", "Test tool 2", "server1"), ]; manager.update_tool_cache(tools.clone()).await; @@ -1731,6 +1962,33 @@ mod tests { assert!(Arc::strong_count(&state) >= 1); } + #[test] + fn test_extract_tool_metadata_from_schema_extension() { + let schema = serde_json::json!({ + "type": "object", + "properties": {}, + "x-proxycast": { + "deferred_loading": true, + "always_visible": false, + "allowed_callers": ["assistant", "code_execution"], + "input_examples": [{"q": "rust"}], + "tags": ["search", "docs"] + } + }); + let meta = McpClientManager::extract_tool_metadata(&schema); + assert_eq!(meta.deferred_loading, Some(true)); + assert_eq!(meta.always_visible, Some(false)); + assert_eq!( + meta.allowed_callers.unwrap_or_default(), + vec!["assistant".to_string(), "code_execution".to_string()] + ); + assert_eq!(meta.input_examples.unwrap_or_default().len(), 1); + assert_eq!( + meta.tags.unwrap_or_default(), + vec!["search".to_string(), "docs".to_string()] + ); + } + // ======================================================================== // 服务器生命周期测试 // ======================================================================== @@ -1827,12 +2085,7 @@ mod tests { .unwrap(); // 设置工具缓存 - let tools = vec![McpToolDefinition { - name: "tool1".to_string(), - description: "Test tool".to_string(), - input_schema: serde_json::json!({}), - server_name: "test-server".to_string(), - }]; + let tools = vec![create_test_tool("tool1", "Test tool", "test-server")]; manager.update_tool_cache(tools).await; assert!(manager.is_tool_cache_valid().await); @@ -1879,18 +2132,8 @@ mod tests { fn test_resolve_tool_name_conflicts_no_conflict() { // 没有冲突的情况 let tools = vec![ - McpToolDefinition { - name: "tool1".to_string(), - description: "Tool 1".to_string(), - input_schema: serde_json::json!({}), - server_name: "server1".to_string(), - }, - McpToolDefinition { - name: "tool2".to_string(), - description: "Tool 2".to_string(), - input_schema: serde_json::json!({}), - server_name: "server2".to_string(), - }, + create_test_tool("tool1", "Tool 1", "server1"), + create_test_tool("tool2", "Tool 2", "server2"), ]; let resolved = McpClientManager::resolve_tool_name_conflicts(tools); @@ -1905,24 +2148,9 @@ mod tests { fn test_resolve_tool_name_conflicts_with_conflict() { // 有冲突的情况:两个服务器都提供 "read_file" 工具 let tools = vec![ - McpToolDefinition { - name: "read_file".to_string(), - description: "Read file from server1".to_string(), - input_schema: serde_json::json!({}), - server_name: "server1".to_string(), - }, - McpToolDefinition { - name: "read_file".to_string(), - description: "Read file from server2".to_string(), - input_schema: serde_json::json!({}), - server_name: "server2".to_string(), - }, - McpToolDefinition { - name: "unique_tool".to_string(), - description: "Unique tool".to_string(), - input_schema: serde_json::json!({}), - server_name: "server1".to_string(), - }, + create_test_tool("read_file", "Read file from server1", "server1"), + create_test_tool("read_file", "Read file from server2", "server2"), + create_test_tool("unique_tool", "Unique tool", "server1"), ]; let resolved = McpClientManager::resolve_tool_name_conflicts(tools); @@ -1939,24 +2167,9 @@ mod tests { fn test_resolve_tool_name_conflicts_multiple_conflicts() { // 多个冲突的情况 let tools = vec![ - McpToolDefinition { - name: "tool_a".to_string(), - description: "Tool A from server1".to_string(), - input_schema: serde_json::json!({}), - server_name: "server1".to_string(), - }, - McpToolDefinition { - name: "tool_a".to_string(), - description: "Tool A from server2".to_string(), - input_schema: serde_json::json!({}), - server_name: "server2".to_string(), - }, - McpToolDefinition { - name: "tool_a".to_string(), - description: "Tool A from server3".to_string(), - input_schema: serde_json::json!({}), - server_name: "server3".to_string(), - }, + create_test_tool("tool_a", "Tool A from server1", "server1"), + create_test_tool("tool_a", "Tool A from server2", "server2"), + create_test_tool("tool_a", "Tool A from server3", "server3"), ]; let resolved = McpClientManager::resolve_tool_name_conflicts(tools); @@ -1985,12 +2198,11 @@ mod tests { let manager = McpClientManager::new(None); // 预先设置缓存 - let cached_tools = vec![McpToolDefinition { - name: "cached_tool".to_string(), - description: "Cached tool".to_string(), - input_schema: serde_json::json!({}), - server_name: "cached_server".to_string(), - }]; + let cached_tools = vec![create_test_tool( + "cached_tool", + "Cached tool", + "cached_server", + )]; manager.update_tool_cache(cached_tools.clone()).await; // 调用 list_tools 应该返回缓存的工具 @@ -1999,6 +2211,119 @@ mod tests { assert_eq!(result[0].name, "cached_tool"); } + #[tokio::test] + async fn test_list_tools_for_context_filters_deferred_and_caller() { + let manager = McpClientManager::new(None); + manager + .update_tool_cache(vec![ + create_test_tool("always_tool", "always", "s1"), + McpToolDefinition { + name: "hidden_tool".to_string(), + description: "hidden".to_string(), + input_schema: serde_json::json!({}), + server_name: "s1".to_string(), + deferred_loading: Some(true), + always_visible: Some(false), + allowed_callers: Some(vec!["code_execution".to_string()]), + input_examples: None, + tags: None, + }, + McpToolDefinition { + name: "visible_deferred".to_string(), + description: "visible deferred".to_string(), + input_schema: serde_json::json!({}), + server_name: "s1".to_string(), + deferred_loading: Some(true), + always_visible: Some(true), + allowed_callers: Some(vec!["assistant".to_string()]), + input_examples: None, + tags: None, + }, + ]) + .await; + + let assistant_tools = manager + .list_tools_for_context(Some("assistant"), false) + .await + .unwrap(); + assert!(assistant_tools.iter().any(|t| t.name == "always_tool")); + assert!(assistant_tools.iter().any(|t| t.name == "visible_deferred")); + assert!(!assistant_tools.iter().any(|t| t.name == "hidden_tool")); + + let code_exec_tools = manager + .list_tools_for_context(Some("code_execution"), true) + .await + .unwrap(); + assert!(code_exec_tools.iter().any(|t| t.name == "hidden_tool")); + } + + #[tokio::test] + async fn test_search_tools_prioritizes_exact_match() { + let manager = McpClientManager::new(None); + manager + .update_tool_cache(vec![ + McpToolDefinition { + name: "weather".to_string(), + description: "Get weather".to_string(), + input_schema: serde_json::json!({}), + server_name: "s1".to_string(), + deferred_loading: Some(true), + always_visible: Some(false), + allowed_callers: None, + input_examples: None, + tags: Some(vec!["forecast".to_string()]), + }, + McpToolDefinition { + name: "get_weather".to_string(), + description: "weather by city".to_string(), + input_schema: serde_json::json!({}), + server_name: "s1".to_string(), + deferred_loading: Some(true), + always_visible: Some(false), + allowed_callers: None, + input_examples: None, + tags: Some(vec!["weather".to_string()]), + }, + ]) + .await; + + let tools = manager + .search_tools("weather", 5, Some("assistant")) + .await + .unwrap(); + assert_eq!(tools.len(), 2); + assert_eq!(tools[0].name, "weather"); + } + + #[tokio::test] + async fn test_call_tool_with_caller_rejects_unauthorized_caller() { + let manager = McpClientManager::new(None); + manager + .update_tool_cache(vec![McpToolDefinition { + name: "restricted".to_string(), + description: "Restricted tool".to_string(), + input_schema: serde_json::json!({}), + server_name: "s1".to_string(), + deferred_loading: None, + always_visible: None, + allowed_callers: Some(vec!["code_execution".to_string()]), + input_examples: None, + tags: None, + }]) + .await; + + let result = manager + .call_tool_with_caller("restricted", serde_json::json!({}), Some("assistant")) + .await; + assert!(result.is_err()); + match result { + Err(McpError::ToolCallFailed(message)) => { + assert!(message.contains("无权调用")); + } + _ => panic!("Expected ToolCallFailed"), + } + } + #[tokio::test] async fn test_list_tools_returns_empty_when_no_servers() { let manager = McpClientManager::new(None); diff --git a/src-tauri/crates/mcp/src/tool_converter.rs b/src-tauri/crates/mcp/src/tool_converter.rs index 1b5dabf6a..2214dab48 100644 --- a/src-tauri/crates/mcp/src/tool_converter.rs +++ b/src-tauri/crates/mcp/src/tool_converter.rs @@ -29,6 +29,10 @@ pub struct OpenAIFunction { pub name: String, pub description: String, pub parameters: serde_json::Value, + /// OpenAI 兼容模型多数不原生支持 input_examples; + /// 这里仅在上游支持时透传,默认使用 description 降级提示。 + #[serde(skip_serializing_if = "Option::is_none")] + pub input_examples: Option>, } /// OpenAI 工具调用 @@ -57,6 +61,10 @@ pub struct AnthropicTool { pub name: String, pub description: String, pub input_schema: serde_json::Value, + #[serde(skip_serializing_if = "Option::is_none")] + pub input_examples: Option>, + #[serde(skip_serializing_if = "Option::is_none")] + pub allowed_callers: Option>, } /// Anthropic 工具使用 @@ -96,6 +104,34 @@ pub struct GeminiParameters { pub struct ToolConverter; impl ToolConverter { + fn build_openai_description(tool: &McpToolDefinition) -> String { + let mut description = tool.description.clone(); + if let Some(examples) = tool.input_examples.as_ref() { + if !examples.is_empty() { + let rendered = examples + .iter() + .take(3) + .map(|v| serde_json::to_string(v).unwrap_or_else(|_| "{}".to_string())) + .collect::>() + .join(" | "); + description.push_str("\n\n[InputExamples] "); + description.push_str(&rendered); + } + } + if let Some(callers) = tool.allowed_callers.as_ref() { + let normalized = callers + .iter() + .map(|v| v.trim()) + .filter(|v| !v.is_empty()) + .collect::>(); + if !normalized.is_empty() { + description.push_str("\n\n[AllowedCallers] "); + description.push_str(&normalized.join(", ")); + } + } + description + } + /// 转换为 OpenAI function calling 格式 pub fn to_openai(tools: &[McpToolDefinition]) -> Vec { tools @@ -104,8 +140,9 @@ impl ToolConverter { tool_type: "function".to_string(), function: OpenAIFunction { name: tool.name.clone(), - description: tool.description.clone(), + description: Self::build_openai_description(tool), parameters: tool.input_schema.clone(), + input_examples: tool.input_examples.clone(), }, }) .collect() @@ -119,6 +156,8 @@ impl ToolConverter { name: tool.name.clone(), description: tool.description.clone(), input_schema: tool.input_schema.clone(), + input_examples: tool.input_examples.clone(), + allowed_callers: tool.allowed_callers.clone(), }) .collect() } @@ -178,3 +217,57 @@ impl ToolConverter { } } } + +#[cfg(test)] +mod tests { + use super::*; + + fn sample_tool() -> McpToolDefinition { + McpToolDefinition { + name: "search_docs".to_string(), + description: "Search project docs".to_string(), + input_schema: serde_json::json!({ + "type": "object", + "properties": { "query": { "type": "string" } }, + "required": ["query"] + }), + server_name: "docs".to_string(), + deferred_loading: Some(true), + always_visible: Some(false), + allowed_callers: Some(vec!["code_execution".to_string()]), + input_examples: Some(vec![serde_json::json!({"query":"rust async"})]), + tags: Some(vec!["docs".to_string(), "search".to_string()]), + } + } + + #[test] + fn test_to_openai_contains_fallback_description_and_examples() { + let openai_tools = ToolConverter::to_openai(&[sample_tool()]); + assert_eq!(openai_tools.len(), 1); + let function = &openai_tools[0].function; + assert!(function.description.contains("[InputExamples]")); + assert!(function.description.contains("[AllowedCallers]")); + assert_eq!(function.input_examples.as_ref().map(|v| v.len()), Some(1)); + } + + #[test] + fn test_to_anthropic_passes_input_examples_and_allowed_callers() { + let anthropic_tools = ToolConverter::to_anthropic(&[sample_tool()]); + assert_eq!(anthropic_tools.len(), 1); + assert_eq!( + anthropic_tools[0] + .input_examples + .as_ref() + .map(|v| v.len()) + .unwrap_or(0), + 1 + ); + assert_eq!( + anthropic_tools[0] + .allowed_callers + .as_ref() + .map(|v| v[0].as_str()), + Some("code_execution") + ); + } +} diff --git a/src-tauri/crates/mcp/src/types.rs b/src-tauri/crates/mcp/src/types.rs index dded73cc5..5c6563bda 100644 --- a/src-tauri/crates/mcp/src/types.rs +++ b/src-tauri/crates/mcp/src/types.rs @@ -73,6 +73,21 @@ pub struct McpToolDefinition { pub description: String, pub input_schema: serde_json::Value, pub server_name: String, + /// 是否延迟加载(不默认注入上下文) + #[serde(default, skip_serializing_if = "Option::is_none")] + pub deferred_loading: Option, + /// 是否始终可见(即使 deferred_loading=true 也可见) + #[serde(default, skip_serializing_if = "Option::is_none")] + pub always_visible: Option, + /// 允许调用方(如 assistant/code_execution/tool_search) + #[serde(default, skip_serializing_if = "Option::is_none")] + pub allowed_callers: Option>, + /// 工具输入示例 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub input_examples: Option>, + /// 标签(用于工具搜索) + #[serde(default, skip_serializing_if = "Option::is_none")] + pub tags: Option>, } /// MCP 工具调用请求 diff --git a/src-tauri/crates/providers/proptest-regressions/providers/antigravity.txt b/src-tauri/crates/providers/proptest-regressions/providers/antigravity.txt new file mode 100644 index 000000000..bc590a026 --- /dev/null +++ b/src-tauri/crates/providers/proptest-regressions/providers/antigravity.txt @@ -0,0 +1,7 @@ +# Seeds for failure cases proptest has generated in the past. It is +# automatically read and these particular cases re-run before any +# novel cases are generated. +# +# It is recommended to check this file in to source control so that +# everyone who runs the test benefits from these saved cases. +cc 2dad3ab4b62243d81a62ce7df3d361f9a27633cc9ccf50cc78c370c66814f5bc # shrinks to expires_in_secs = 601 diff --git a/src-tauri/crates/providers/src/providers/claude_custom.rs b/src-tauri/crates/providers/src/providers/claude_custom.rs index f1e882008..04d459edf 100644 --- a/src-tauri/crates/providers/src/providers/claude_custom.rs +++ b/src-tauri/crates/providers/src/providers/claude_custom.rs @@ -131,6 +131,111 @@ impl ClaudeCustomProvider { None } + fn convert_openai_tool_to_anthropic( + tool: &proxycast_core::models::openai::Tool, + ) -> Option { + match tool { + proxycast_core::models::openai::Tool::Function { function } => { + let input_schema = function + .parameters + .clone() + .unwrap_or_else(|| serde_json::json!({"type":"object","properties":{}})); + let extension = input_schema + .get("x-proxycast") + .or_else(|| input_schema.get("x_proxycast")) + .cloned() + .unwrap_or_else(|| serde_json::json!({})); + let mut input_examples = extension + .get("input_examples") + .or_else(|| extension.get("inputExamples")) + .and_then(|v| v.as_array()) + .cloned() + .unwrap_or_default(); + if input_examples.is_empty() { + input_examples = proxycast_core::tool_calling::resolve_tool_input_examples( + &function.name, + &input_schema, + ); + } + 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_string()) + .filter(|v| !v.is_empty()) + .collect::>() + }) + .unwrap_or_default(); + + let mut description = function.description.clone().unwrap_or_default(); + if !input_examples.is_empty() && !description.contains("[InputExamples]") { + let rendered = input_examples + .iter() + .take(3) + .map(|v| serde_json::to_string(v).unwrap_or_else(|_| "{}".to_string())) + .collect::>() + .join(" | "); + description.push_str("\n\n[InputExamples] "); + description.push_str(&rendered); + } + if !allowed_callers.is_empty() && !description.contains("[AllowedCallers]") { + description.push_str("\n\n[AllowedCallers] "); + description.push_str(&allowed_callers.join(", ")); + } + + let mut anthropic_tool = serde_json::json!({ + "name": function.name, + "description": description, + "input_schema": input_schema + }); + if !input_examples.is_empty() { + anthropic_tool["input_examples"] = serde_json::Value::Array(input_examples); + } + if !allowed_callers.is_empty() { + anthropic_tool["allowed_callers"] = serde_json::json!(allowed_callers); + } + Some(anthropic_tool) + } + _ => None, + } + } + + fn convert_openai_tool_choice_to_anthropic( + tool_choice: &Option, + ) -> Option { + let Some(tool_choice) = tool_choice else { + return None; + }; + match tool_choice { + serde_json::Value::String(s) => match s.as_str() { + "none" => Some(serde_json::json!({"type":"none"})), + "auto" => Some(serde_json::json!({"type":"auto"})), + "required" | "any" => Some(serde_json::json!({"type":"any"})), + _ => None, + }, + serde_json::Value::Object(obj) => { + if let Some(func) = obj.get("function") { + func.get("name") + .and_then(|n| n.as_str()) + .map(|name| serde_json::json!({"type":"tool","name":name})) + } else if let Some(t) = obj.get("type").and_then(|t| t.as_str()) { + match t { + "any" | "tool" => Some(serde_json::json!({"type":"any"})), + "auto" => Some(serde_json::json!({"type":"auto"})), + "none" => Some(serde_json::json!({"type":"none"})), + _ => None, + } + } else { + None + } + } + _ => None, + } + } + /// 调用 Anthropic API(原生格式) pub async fn call_api( &self, @@ -245,6 +350,19 @@ impl ClaudeCustomProvider { anthropic_body["system"] = serde_json::json!(sys); } + if let Some(ref tools) = request.tools { + let anthropic_tools: Vec = tools + .iter() + .filter_map(Self::convert_openai_tool_to_anthropic) + .collect(); + if !anthropic_tools.is_empty() { + anthropic_body["tools"] = serde_json::json!(anthropic_tools); + } + } + if let Some(tc) = Self::convert_openai_tool_choice_to_anthropic(&request.tool_choice) { + anthropic_body["tool_choice"] = tc; + } + let api_key = self .config .api_key @@ -558,19 +676,7 @@ impl StreamingProvider for ClaudeCustomProvider { if let Some(ref tools) = request.tools { let anthropic_tools: Vec = tools .iter() - .filter_map(|tool| { - match tool { - proxycast_core::models::openai::Tool::Function { function } => { - Some(serde_json::json!({ - "name": function.name, - "description": function.description.clone().unwrap_or_default(), - "input_schema": function.parameters.clone().unwrap_or_else(|| serde_json::json!({"type": "object", "properties": {}})) - })) - } - // WebSearch 等其他工具类型暂不处理 - _ => None, - } - }) + .filter_map(Self::convert_openai_tool_to_anthropic) .collect(); if !anthropic_tools.is_empty() { @@ -583,43 +689,12 @@ impl StreamingProvider for ClaudeCustomProvider { } // 转换 tool_choice: OpenAI 格式 -> Anthropic 格式 - if let Some(ref tool_choice) = request.tool_choice { - let anthropic_tool_choice = match tool_choice { - serde_json::Value::String(s) => { - match s.as_str() { - "none" => Some(serde_json::json!({"type": "none"})), - "auto" => Some(serde_json::json!({"type": "auto"})), - "required" | "any" => Some(serde_json::json!({"type": "any"})), - _ => None, // 未知值,不设置 - } - } - serde_json::Value::Object(obj) => { - // 处理 {"type": "function", "function": {"name": "xxx"}} 格式 - if let Some(func) = obj.get("function") { - func.get("name") - .and_then(|n| n.as_str()) - .map(|name| serde_json::json!({"type": "tool", "name": name})) - } else if let Some(t) = obj.get("type").and_then(|t| t.as_str()) { - match t { - "any" | "tool" => Some(serde_json::json!({"type": "any"})), - "auto" => Some(serde_json::json!({"type": "auto"})), - "none" => Some(serde_json::json!({"type": "none"})), - _ => None, - } - } else { - None - } - } - _ => None, - }; - - if let Some(tc) = anthropic_tool_choice { - anthropic_body["tool_choice"] = tc; - tracing::info!( - "[CLAUDE_STREAM] 设置 tool_choice: {:?}", - anthropic_body["tool_choice"] - ); - } + if let Some(tc) = Self::convert_openai_tool_choice_to_anthropic(&request.tool_choice) { + anthropic_body["tool_choice"] = tc; + tracing::info!( + "[CLAUDE_STREAM] 设置 tool_choice: {:?}", + anthropic_body["tool_choice"] + ); } let url = self.build_url("messages"); @@ -668,3 +743,104 @@ impl StreamingProvider for ClaudeCustomProvider { StreamFormat::AnthropicSse } } + +#[cfg(test)] +mod tests { + use super::*; + use proxycast_core::models::openai::{FunctionDef, Tool}; + + #[test] + fn test_convert_openai_tool_to_anthropic_keeps_metadata() { + let tool = Tool::Function { + function: FunctionDef { + name: "create_ticket".to_string(), + description: Some("Create support ticket".to_string()), + parameters: Some(serde_json::json!({ + "type": "object", + "properties": { + "title": {"type": "string"} + }, + "x-proxycast": { + "input_examples": [{"title":"Billing issue"}], + "allowed_callers": ["assistant", "code_execution"] + } + })), + }, + }; + + let converted = ClaudeCustomProvider::convert_openai_tool_to_anthropic(&tool) + .expect("tool should be converted"); + let description = converted + .get("description") + .and_then(|v| v.as_str()) + .unwrap_or_default() + .to_string(); + + assert_eq!(converted["name"], serde_json::json!("create_ticket")); + assert!(description.contains("[InputExamples]")); + assert!(description.contains("[AllowedCallers]")); + assert_eq!( + converted["input_examples"], + serde_json::json!([{"title":"Billing issue"}]) + ); + assert_eq!( + converted["allowed_callers"], + serde_json::json!(["assistant", "code_execution"]) + ); + } + + #[test] + fn test_convert_openai_tool_choice_to_anthropic_variants() { + assert_eq!( + ClaudeCustomProvider::convert_openai_tool_choice_to_anthropic(&Some( + serde_json::json!("required") + )), + Some(serde_json::json!({"type":"any"})) + ); + assert_eq!( + ClaudeCustomProvider::convert_openai_tool_choice_to_anthropic(&Some( + serde_json::json!({"type":"function","function":{"name":"create_ticket"}}) + )), + Some(serde_json::json!({"type":"tool","name":"create_ticket"})) + ); + assert_eq!( + ClaudeCustomProvider::convert_openai_tool_choice_to_anthropic(&Some( + serde_json::json!({"type":"none"}) + )), + Some(serde_json::json!({"type":"none"})) + ); + } + + #[test] + fn test_convert_openai_tool_to_anthropic_uses_builtin_input_examples_fallback() { + let tool = Tool::Function { + function: FunctionDef { + name: "WebSearch".to_string(), + description: Some("允许 Claude 搜索网络并使用结果来提供响应。".to_string()), + parameters: Some(serde_json::json!({ + "type": "object", + "properties": { + "query": {"type": "string"}, + "limit": {"type": "integer"} + }, + "required": ["query"] + })), + }, + }; + + let converted = ClaudeCustomProvider::convert_openai_tool_to_anthropic(&tool) + .expect("tool should be converted"); + let description = converted + .get("description") + .and_then(|v| v.as_str()) + .unwrap_or_default() + .to_string(); + + assert!(description.contains("[InputExamples]")); + assert!(converted + .get("input_examples") + .and_then(|v| v.as_array()) + .map(|arr| !arr.is_empty()) + .unwrap_or(false)); + } +} diff --git a/src-tauri/crates/providers/src/providers/openai_custom.rs b/src-tauri/crates/providers/src/providers/openai_custom.rs index 1b1144339..079f4d4ac 100644 --- a/src-tauri/crates/providers/src/providers/openai_custom.rs +++ b/src-tauri/crates/providers/src/providers/openai_custom.rs @@ -42,6 +42,121 @@ impl Default for OpenAICustomProvider { } impl OpenAICustomProvider { + fn tool_calling_v2_enabled() -> bool { + proxycast_core::tool_calling::tool_calling_v2_enabled() + } + + fn native_input_examples_enabled() -> bool { + proxycast_core::tool_calling::tool_calling_native_input_examples_enabled() + } + + fn normalize_openai_request_payload(&self, payload: &mut serde_json::Value) { + if !Self::tool_calling_v2_enabled() { + return; + } + + let Some(tools) = payload.get_mut("tools").and_then(|v| v.as_array_mut()) else { + return; + }; + + for tool in tools.iter_mut() { + let tool_type = tool + .get("type") + .and_then(|v| v.as_str()) + .unwrap_or_default(); + if tool_type != "function" { + continue; + } + + let Some(function) = tool.get_mut("function").and_then(|v| v.as_object_mut()) else { + continue; + }; + + let parameters = function + .get("parameters") + .cloned() + .unwrap_or_else(|| serde_json::json!({})); + let extension = parameters + .get("x-proxycast") + .or_else(|| parameters.get("x_proxycast")) + .cloned() + .unwrap_or_else(|| serde_json::json!({})); + + let mut input_examples = extension + .get("input_examples") + .or_else(|| extension.get("inputExamples")) + .and_then(|v| v.as_array()) + .cloned() + .unwrap_or_default(); + if input_examples.is_empty() { + let tool_name = function + .get("name") + .and_then(|v| v.as_str()) + .unwrap_or_default(); + input_examples = proxycast_core::tool_calling::resolve_tool_input_examples( + tool_name, + ¶meters, + ); + } + 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_string()) + .filter(|v| !v.is_empty()) + .collect::>() + }) + .unwrap_or_default(); + let deferred_loading = extension + .get("deferred_loading") + .or_else(|| extension.get("deferredLoading")) + .and_then(|v| v.as_bool()) + .unwrap_or(false); + + let description = function + .get("description") + .and_then(|v| v.as_str()) + .unwrap_or("") + .to_string(); + let mut enhanced_description = description.clone(); + + if !input_examples.is_empty() && !enhanced_description.contains("[InputExamples]") { + let rendered = input_examples + .iter() + .take(3) + .map(|v| serde_json::to_string(v).unwrap_or_else(|_| "{}".to_string())) + .collect::>() + .join(" | "); + enhanced_description.push_str("\n\n[InputExamples] "); + enhanced_description.push_str(&rendered); + } + + if !allowed_callers.is_empty() && !enhanced_description.contains("[AllowedCallers]") { + enhanced_description.push_str("\n\n[AllowedCallers] "); + enhanced_description.push_str(&allowed_callers.join(", ")); + } + + if deferred_loading && !enhanced_description.contains("[DeferredLoading]") { + enhanced_description.push_str("\n\n[DeferredLoading] true"); + } + + function.insert( + "description".to_string(), + serde_json::Value::String(enhanced_description), + ); + + if !input_examples.is_empty() && Self::native_input_examples_enabled() { + function.insert( + "input_examples".to_string(), + serde_json::Value::Array(input_examples.clone()), + ); + } + } + } + fn maybe_log_protocol_mismatch_hint(url: &str, status: StatusCode) { if (status == StatusCode::UNAUTHORIZED || status == StatusCode::FORBIDDEN) && url.contains("/api/anthropic") @@ -221,6 +336,10 @@ impl OpenAICustomProvider { request.model ); + let mut payload = + serde_json::to_value(request).map_err(|e| format!("序列化 OpenAI 请求失败: {e}"))?; + self.normalize_openai_request_payload(&mut payload); + for url in &urls { eprintln!("[OPENAI_CUSTOM] call_api trying URL: {url}"); let resp = self @@ -228,7 +347,7 @@ impl OpenAICustomProvider { .post(url) .header("Authorization", format!("Bearer {api_key}")) .header("Content-Type", "application/json") - .json(request) + .json(&payload) .send() .await?; @@ -261,12 +380,15 @@ impl OpenAICustomProvider { self.get_base_url() ); + let mut payload = request.clone(); + self.normalize_openai_request_payload(&mut payload); + let resp = self .client .post(&url) .header("Authorization", format!("Bearer {api_key}")) .header("Content-Type", "application/json") - .json(request) + .json(&payload) .send() .await?; @@ -280,7 +402,7 @@ impl OpenAICustomProvider { .post(&fallback_url) .header("Authorization", format!("Bearer {api_key}")) .header("Content-Type", "application/json") - .json(request) + .json(&payload) .send() .await?; Self::maybe_log_protocol_mismatch_hint(&fallback_url, resp2.status()); @@ -368,6 +490,9 @@ impl StreamingProvider for OpenAICustomProvider { // 确保请求启用流式 let mut stream_request = request.clone(); stream_request.stream = true; + let mut payload = serde_json::to_value(&stream_request) + .map_err(|e| ProviderError::ConfigurationError(format!("序列化流式请求失败: {e}")))?; + self.normalize_openai_request_payload(&mut payload); let url = self.build_url("chat/completions"); @@ -383,7 +508,7 @@ impl StreamingProvider for OpenAICustomProvider { .header("Authorization", format!("Bearer {api_key}")) .header("Content-Type", "application/json") .header("Accept", "text/event-stream") - .json(&stream_request) + .json(&payload) .send() .await .map_err(|e| ProviderError::from_reqwest_error(&e))?; @@ -396,7 +521,7 @@ impl StreamingProvider for OpenAICustomProvider { .header("Authorization", format!("Bearer {api_key}")) .header("Content-Type", "application/json") .header("Accept", "text/event-stream") - .json(&stream_request) + .json(&payload) .send() .await .map_err(|e| ProviderError::from_reqwest_error(&e))? @@ -436,3 +561,262 @@ impl StreamingProvider for OpenAICustomProvider { StreamFormat::OpenAiSse } } + +#[cfg(test)] +mod tests { + use super::*; + use axum::{extract::State, http::header, response::IntoResponse, routing::post, Json, Router}; + use futures::StreamExt; + use proxycast_core::models::openai::{ChatMessage, FunctionDef, MessageContent, Tool}; + use std::sync::Arc; + use tokio::sync::Mutex; + + async fn start_mock_openai_server( + captured: Arc>>, + ) -> (String, tokio::task::JoinHandle<()>) { + async fn handle_chat( + State(captured): State>>>, + Json(payload): Json, + ) -> impl IntoResponse { + captured.lock().await.push(payload.clone()); + + if payload + .get("stream") + .and_then(|v| v.as_bool()) + .unwrap_or(false) + { + ( + [(header::CONTENT_TYPE, "text/event-stream")], + "data: {\"id\":\"chatcmpl-test\",\"object\":\"chat.completion.chunk\",\"choices\":[]}\n\ndata: [DONE]\n\n", + ) + .into_response() + } else { + Json(serde_json::json!({ + "id": "chatcmpl-test", + "object": "chat.completion", + "choices": [{ + "index": 0, + "message": {"role":"assistant","content":"ok"}, + "finish_reason": "stop" + }], + "usage": {"prompt_tokens": 10, "completion_tokens": 2, "total_tokens": 12} + })) + .into_response() + } + } + + let app = Router::new() + .route("/v1/chat/completions", post(handle_chat)) + .with_state(captured); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0") + .await + .expect("bind mock server"); + let addr = listener.local_addr().expect("read mock server local addr"); + let server = tokio::spawn(async move { + axum::serve(listener, app) + .await + .expect("mock server should run"); + }); + (format!("http://{}", addr), server) + } + + fn build_tool_calling_request() -> ChatCompletionRequest { + ChatCompletionRequest { + model: "deepseek-chat".to_string(), + messages: vec![ChatMessage { + role: "user".to_string(), + content: Some(MessageContent::Text("hi".to_string())), + tool_calls: None, + tool_call_id: None, + reasoning_content: None, + }], + temperature: None, + max_tokens: Some(128), + top_p: None, + stream: false, + tools: Some(vec![Tool::Function { + function: FunctionDef { + name: "search_docs".to_string(), + description: Some("Search docs".to_string()), + parameters: Some(serde_json::json!({ + "type":"object", + "properties":{"query":{"type":"string"}}, + "x-proxycast": { + "input_examples":[{"query":"rust async"}], + "allowed_callers":["assistant","code_execution"], + "deferred_loading": true + } + })), + }, + }]), + tool_choice: None, + reasoning_effort: None, + } + } + + #[test] + fn test_normalize_openai_request_payload_injects_fallback_description() { + let provider = OpenAICustomProvider::default(); + let mut payload = serde_json::json!({ + "model": "deepseek-chat", + "messages": [{"role":"user","content":"hi"}], + "tools": [{ + "type":"function", + "function": { + "name":"search_docs", + "description":"Search docs", + "parameters": { + "type":"object", + "properties":{"query":{"type":"string"}}, + "x-proxycast": { + "input_examples":[{"query":"rust async"}], + "allowed_callers":["assistant","code_execution"], + "deferred_loading":true + } + } + } + }] + }); + + provider.normalize_openai_request_payload(&mut payload); + let description = payload["tools"][0]["function"]["description"] + .as_str() + .unwrap_or_default() + .to_string(); + assert!(description.contains("[InputExamples]")); + assert!(description.contains("[AllowedCallers]")); + assert!(description.contains("[DeferredLoading]")); + } + + #[test] + fn test_normalize_openai_request_payload_supports_x_proxycast_alias() { + let provider = OpenAICustomProvider::default(); + let mut payload = serde_json::json!({ + "model": "deepseek-chat", + "messages": [{"role":"user","content":"hi"}], + "tools": [{ + "type":"function", + "function": { + "name":"search_docs", + "description":"Search docs", + "parameters": { + "type":"object", + "properties":{"query":{"type":"string"}}, + "x_proxycast": { + "inputExamples":[{"query":"tool search"}], + "allowedCallers":["tool_search"] + } + } + } + }] + }); + + provider.normalize_openai_request_payload(&mut payload); + let description = payload["tools"][0]["function"]["description"] + .as_str() + .unwrap_or_default(); + + assert!(description.contains("[InputExamples]")); + assert!(description.contains("[AllowedCallers]")); + assert!(description.contains("tool_search")); + } + + #[test] + fn test_normalize_openai_request_payload_ignores_non_function_tools() { + let provider = OpenAICustomProvider::default(); + let mut payload = serde_json::json!({ + "model": "deepseek-chat", + "messages": [{"role":"user","content":"hi"}], + "tools": [{"type":"web_search_20250305"}] + }); + + provider.normalize_openai_request_payload(&mut payload); + + assert_eq!( + payload["tools"][0], + serde_json::json!({"type":"web_search_20250305"}) + ); + } + + #[test] + fn test_normalize_openai_request_payload_uses_builtin_input_examples_fallback() { + let provider = OpenAICustomProvider::default(); + let mut payload = serde_json::json!({ + "model": "deepseek-chat", + "messages": [{"role":"user","content":"hi"}], + "tools": [{ + "type":"function", + "function": { + "name":"WebSearch", + "description":"允许 Claude 搜索网络并使用结果来提供响应。", + "parameters": { + "type":"object", + "properties":{"query":{"type":"string"},"limit":{"type":"integer"}}, + "required":["query"] + } + } + }] + }); + + provider.normalize_openai_request_payload(&mut payload); + let description = payload["tools"][0]["function"]["description"] + .as_str() + .unwrap_or_default(); + + assert!(description.contains("[InputExamples]")); + } + + #[tokio::test] + async fn test_openai_compatible_non_stream_and_stream_both_normalized() { + if !OpenAICustomProvider::tool_calling_v2_enabled() { + return; + } + + let captured = Arc::new(Mutex::new(Vec::::new())); + let (base_url, server_handle) = start_mock_openai_server(captured.clone()).await; + let mut provider = OpenAICustomProvider::with_config("sk-test".to_string(), Some(base_url)); + provider.client = reqwest::Client::builder() + .no_proxy() + .build() + .expect("build test client without proxy"); + let request = build_tool_calling_request(); + + let resp = provider + .call_api(&request) + .await + .expect("non-stream call should succeed"); + let status = resp.status(); + if !status.is_success() { + let body = resp.text().await.unwrap_or_default(); + panic!("non-stream call failed: status={status}, body={body}"); + } + + let mut stream = provider + .call_api_stream(&request) + .await + .expect("stream call should succeed"); + let first_chunk = stream + .next() + .await + .expect("stream should return at least one chunk") + .expect("first stream chunk should be ok"); + let chunk_text = String::from_utf8(first_chunk.to_vec()).expect("chunk should be utf8"); + assert!(chunk_text.contains("data:")); + + let bodies = captured.lock().await; + assert_eq!(bodies.len(), 2); + assert_eq!(bodies[1]["stream"], serde_json::json!(true)); + + for body in bodies.iter() { + let description = body["tools"][0]["function"]["description"] + .as_str() + .unwrap_or_default() + .to_string(); + assert!(description.contains("[InputExamples]")); + assert!(description.contains("[AllowedCallers]")); + assert!(description.contains("[DeferredLoading]")); + } + + server_handle.abort(); + } +} diff --git a/src-tauri/crates/server/src/handlers/provider_calls.rs b/src-tauri/crates/server/src/handlers/provider_calls.rs index 5fe5211cb..6372ffc49 100644 --- a/src-tauri/crates/server/src/handlers/provider_calls.rs +++ b/src-tauri/crates/server/src/handlers/provider_calls.rs @@ -1818,8 +1818,67 @@ pub async fn call_provider_openai( custom_url, &credential.uuid[..8], request.stream - ), + ), ); + + if request.stream { + state.logs.write().await.add( + "info", + "[OPENAI_COMPAT] 流式请求,走 OpenAICustomProvider.call_api_stream", + ); + + match openai.call_api_stream(request).await { + Ok(stream_response) => { + if let Some(db) = &state.db { + let _ = state.pool_service.mark_healthy( + db, + &credential.uuid, + Some(&request.model), + ); + let _ = state.pool_service.record_usage(db, &credential.uuid); + } + + let body_stream = + stream_response.map(|result| -> Result { + match result { + Ok(bytes) => Ok(bytes), + Err(e) => Ok(axum::body::Bytes::from(e.to_sse_error())), + } + }); + + return Response::builder() + .status(StatusCode::OK) + .header(header::CONTENT_TYPE, "text/event-stream") + .header(header::CACHE_CONTROL, "no-cache, no-store, must-revalidate") + .header("Connection", "keep-alive") + .header("X-Accel-Buffering", "no") + .header("Transfer-Encoding", "chunked") + .body(Body::from_stream(body_stream)) + .unwrap_or_else(|_| { + ( + StatusCode::INTERNAL_SERVER_ERROR, + Json(serde_json::json!({"error": {"message": "Failed to build stream response"}})), + ) + .into_response() + }); + } + Err(e) => { + if let Some(db) = &state.db { + let _ = state.pool_service.mark_unhealthy( + db, + &credential.uuid, + Some(&format!("Streaming API call failed: {e}")), + ); + } + return ( + StatusCode::INTERNAL_SERVER_ERROR, + Json(serde_json::json!({"error": {"message": format!("OpenAI compatible streaming API call failed: {}", e)}})), + ) + .into_response(); + } + } + } + match openai.call_api(request).await { Ok(resp) => { let status = resp.status(); @@ -1833,37 +1892,6 @@ pub async fn call_provider_openai( ), ); - if request.stream && status.is_success() { - state.logs.write().await.add( - "info", - "[OPENAI_COMPAT] 流式请求,透传 SSE 响应", - ); - if let Some(db) = &state.db { - let _ = state.pool_service.mark_healthy( - db, - &credential.uuid, - Some(&request.model), - ); - let _ = state.pool_service.record_usage(db, &credential.uuid); - } - let stream = resp.bytes_stream(); - return Response::builder() - .status(StatusCode::OK) - .header(header::CONTENT_TYPE, "text/event-stream") - .header(header::CACHE_CONTROL, "no-cache, no-store, must-revalidate") - .header("Connection", "keep-alive") - .header("X-Accel-Buffering", "no") // 禁用 nginx 等代理的缓冲 - .header("Transfer-Encoding", "chunked") - .body(Body::from_stream(stream)) - .unwrap_or_else(|_| { - ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": "Failed to build stream response"}})), - ) - .into_response() - }); - } - // 非流式响应 if status.is_success() { if let Some(db) = &state.db { diff --git a/src-tauri/src/app/bootstrap.rs b/src-tauri/src/app/bootstrap.rs index d844dec39..8c6145f58 100644 --- a/src-tauri/src/app/bootstrap.rs +++ b/src-tauri/src/app/bootstrap.rs @@ -92,6 +92,9 @@ type TelemetryInit = ( /// 初始化所有应用状态 pub fn init_states(config: &Config) -> Result { + // 将 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( diff --git a/src-tauri/src/app/commands/config.rs b/src-tauri/src/app/commands/config.rs index 57aae9de5..d32707a07 100644 --- a/src-tauri/src/app/commands/config.rs +++ b/src-tauri/src/app/commands/config.rs @@ -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, diff --git a/src-tauri/src/app/runner.rs b/src-tauri/src/app/runner.rs index 296e9bca0..d5b42dfef 100644 --- a/src-tauri/src/app/runner.rs +++ b/src-tauri/src/app/runner.rs @@ -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, diff --git a/src-tauri/src/commands/agent_cmd.rs b/src-tauri/src/commands/agent_cmd.rs index 45221f506..f0fa1d933 100644 --- a/src-tauri/src/commands/agent_cmd.rs +++ b/src-tauri/src/commands/agent_cmd.rs @@ -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); diff --git a/src-tauri/src/commands/aster_agent_cmd.rs b/src-tauri/src/commands/aster_agent_cmd.rs index 89e4da330..8ad4dd87b 100644 --- a/src-tauri/src/commands/aster_agent_cmd.rs +++ b/src-tauri/src/commands/aster_agent_cmd.rs @@ -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 { 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 { 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 { + 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::>(); + 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 { @@ -421,6 +481,9 @@ async fn ensure_code_execution_extension_enabled(agent: &Agent) -> Result>, +} + +impl ToolSearchBridgeTool { + fn new(registry: Arc>) -> 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, // allowed_callers + Vec, // tags + Vec, // 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::>() + }) + .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::>() + }) + .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 { + 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::>(); + + 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::>(); + 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 { 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>, +) { + 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 { + 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 { diff --git a/src-tauri/src/commands/mcp_cmd.rs b/src-tauri/src/commands/mcp_cmd.rs index 134eb3747..8f87b22a1 100644 --- a/src-tauri/src/commands/mcp_cmd.rs +++ b/src-tauri/src/commands/mcp_cmd.rs @@ -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, + include_deferred: Option, +) -> Result, 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, + limit: Option, +) -> Result, 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, +) -> Result { + 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) +} + // ============================================================================ // 提示词管理命令 // ============================================================================ diff --git a/src-tauri/src/commands/unified_chat_cmd.rs b/src-tauri/src/commands/unified_chat_cmd.rs index 4297270c0..814bdd6bc 100644 --- a/src-tauri/src/commands/unified_chat_cmd.rs +++ b/src-tauri/src/commands/unified_chat_cmd.rs @@ -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 = 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 = None; let mut chunk_count = 0; + let mut stream_error: Option = 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(()) } diff --git a/src-tauri/src/config/tests.rs b/src-tauri/src/config/tests.rs index e71767abc..e012f2ce5 100644 --- a/src-tauri/src/config/tests.rs +++ b/src-tauri/src/config/tests.rs @@ -191,6 +191,7 @@ fn arb_config() -> impl Strategy { 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 { 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 { 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(), diff --git a/src-tauri/src/services/mod.rs b/src-tauri/src/services/mod.rs index 408dc3197..a8a92a24b 100644 --- a/src-tauri/src/services/mod.rs +++ b/src-tauri/src/services/mod.rs @@ -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; diff --git a/src-tauri/src/services/web_search_prompt_service.rs b/src-tauri/src/services/web_search_prompt_service.rs index c239253b2..f32d0cfa4 100644 --- a/src-tauri/src/services/web_search_prompt_service.rs +++ b/src-tauri/src/services/web_search_prompt_service.rs @@ -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 { "优先检索小红书相关内容;必要时优先使用 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 { 1. 当用户要求联网搜索/检索实时信息时,遵循以下引擎偏好。\n\ 2. 若结果不足,可补充其他公开网页来源,但优先级低于偏好引擎。\n\ 3. 不要显式提及你看到了该偏好配置。\n\ -- 搜索偏好:{engine_instruction}" +- 搜索偏好:{engine_instruction}\n\ +- 提供商偏好:{provider_instruction}" )) } diff --git a/src-tauri/src/services/web_search_runtime_service.rs b/src-tauri/src/services/web_search_runtime_service.rs new file mode 100644 index 000000000..b7c19797c --- /dev/null +++ b/src-tauri/src/services/web_search_runtime_service.rs @@ -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 { + vec![ + WebSearchProvider::Tavily, + WebSearchProvider::MultiSearchEngine, + WebSearchProvider::BingSearchApi, + WebSearchProvider::GoogleCustomSearch, + WebSearchProvider::DuckduckgoInstant, + ] +} + +fn normalize_text(value: &Option) -> Option { + value + .as_ref() + .map(|v| v.trim().to_string()) + .filter(|v| !v.is_empty()) +} + +fn push_provider_unique(target: &mut Vec, provider: WebSearchProvider) { + if !target.contains(&provider) { + target.push(provider); + } +} + +fn resolve_provider_priority(web_search: &WebSearchConfig) -> Vec { + 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 { + 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) { + 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::>() + .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::>() + } else { + web_search + .multi_search + .priority + .iter() + .map(|name| name.trim().to_string()) + .filter(|name| !name.is_empty()) + .collect::>() + }; + + let engines = web_search + .multi_search + .engines + .iter() + .filter_map(normalize_engine_entry) + .collect::>(); + + 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()); + } +} diff --git a/src-tauri/tauri.conf.json b/src-tauri/tauri.conf.json index e02a0a001..ba35c198f 100644 --- a/src-tauri/tauri.conf.json +++ b/src-tauri/tauri.conf.json @@ -1,7 +1,7 @@ { "$schema": "https://schema.tauri.app/config/2", "productName": "ProxyCast", - "version": "0.77.0", + "version": "0.78.0", "identifier": "com.proxycast.app", "build": { "beforeDevCommand": "npm run dev", diff --git a/src/App.tsx b/src/App.tsx index 5d5c10c59..308d7b056 100644 --- a/src/App.tsx +++ b/src/App.tsx @@ -71,12 +71,13 @@ const AppContainer = styled.div` overflow: hidden; `; -const MainContent = styled.main` +const MainContent = styled.main<{ $withSidebarGap?: boolean }>` flex: 1; overflow: hidden; display: flex; flex-direction: column; min-height: 0; + padding-left: ${(props) => (props.$withSidebarGap ? "10px" : "0")}; `; const PageWrapper = styled.div<{ $isActive: boolean }>` @@ -543,6 +544,8 @@ function AppContent() { !isThemeWorkspacePage(currentPage) && !shouldHideSidebarForAgent; + const shouldAddMainContentGap = shouldShowAppSidebar && currentPage === "agent"; + return ( @@ -550,7 +553,9 @@ function AppContent() { {shouldShowAppSidebar && ( )} - {renderAllPages()} + + {renderAllPages()} + void; + value: string; + onChange: (value: string) => void; + }) => React.ReactNode +>(); + +vi.mock("@/hooks/useTauri", () => ({ + getConfig: vi.fn(async () => ({})), +})); + +vi.mock("./ChatModelSelector", () => ({ + ChatModelSelector: () =>
, +})); + +vi.mock("../utils/entryPromptComposer", () => ({ + composeEntryPrompt: vi.fn(() => ""), + createDefaultEntrySlotValues: vi.fn(() => ({})), + formatEntryTaskPreview: vi.fn(() => ""), + getEntryTaskTemplate: vi.fn(() => ({ slots: [], description: "", label: "" })), + SOCIAL_MEDIA_ENTRY_TASKS: [], + validateEntryTaskSlots: vi.fn(() => ({ valid: true, missing: [] })), +})); + +vi.mock("../utils/contextualRecommendations", () => ({ + buildRecommendationPrompt: vi.fn((fullPrompt: string) => fullPrompt), + getContextualRecommendations: vi.fn(() => []), +})); + +vi.mock("./Inputbar/components/CharacterMention", () => ({ + CharacterMention: (props: { + characters?: Character[]; + skills?: Skill[]; + onSelectSkill?: (skill: Skill) => void; + value: string; + onChange: (value: string) => void; + }) => { + mockCharacterMention(props); + return
; + }, +})); + +vi.mock("@/components/ui/button", () => ({ + Button: ({ + children, + onClick, + disabled, + }: { + children: React.ReactNode; + onClick?: () => void; + disabled?: boolean; + }) => ( + + ), +})); + +vi.mock("@/components/ui/input", () => ({ + Input: ({ + value, + onChange, + placeholder, + }: { + value?: string; + onChange?: (e: React.ChangeEvent) => void; + placeholder?: string; + }) => , +})); + +vi.mock("@/components/ui/textarea", () => { + const Textarea = React.forwardRef< + HTMLTextAreaElement, + React.TextareaHTMLAttributes + >((props, ref) =>