mirror of
https://github.com/aiclientproxy/proxycast.git
synced 2026-09-24 23:10:56 +08:00
chore(release): v0.79.0
This commit is contained in:
Generated
+17
-15
@@ -6996,7 +6996,7 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "proxycast"
|
||||
version = "0.77.0"
|
||||
version = "0.79.0"
|
||||
dependencies = [
|
||||
"anyhow",
|
||||
"arboard",
|
||||
@@ -7097,12 +7097,13 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "proxycast-agent"
|
||||
version = "0.78.0"
|
||||
version = "0.79.0"
|
||||
dependencies = [
|
||||
"aster-core",
|
||||
"async-trait",
|
||||
"chrono",
|
||||
"dirs 5.0.1",
|
||||
"futures",
|
||||
"proxycast-core",
|
||||
"proxycast-mcp",
|
||||
"proxycast-providers",
|
||||
@@ -7121,7 +7122,7 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "proxycast-config"
|
||||
version = "0.78.0"
|
||||
version = "0.79.0"
|
||||
dependencies = [
|
||||
"async-trait",
|
||||
"parking_lot",
|
||||
@@ -7137,7 +7138,7 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "proxycast-core"
|
||||
version = "0.78.0"
|
||||
version = "0.79.0"
|
||||
dependencies = [
|
||||
"aster-models",
|
||||
"async-trait",
|
||||
@@ -7177,7 +7178,7 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "proxycast-credential"
|
||||
version = "0.78.0"
|
||||
version = "0.79.0"
|
||||
dependencies = [
|
||||
"axum 0.7.9",
|
||||
"base64 0.22.1",
|
||||
@@ -7212,7 +7213,7 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "proxycast-infra"
|
||||
version = "0.78.0"
|
||||
version = "0.79.0"
|
||||
dependencies = [
|
||||
"chrono",
|
||||
"dashmap 5.5.3",
|
||||
@@ -7232,9 +7233,10 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "proxycast-mcp"
|
||||
version = "0.78.0"
|
||||
version = "0.79.0"
|
||||
dependencies = [
|
||||
"async-trait",
|
||||
"dirs 5.0.1",
|
||||
"glob",
|
||||
"proxycast-core",
|
||||
"rmcp",
|
||||
@@ -7263,7 +7265,7 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "proxycast-processor"
|
||||
version = "0.78.0"
|
||||
version = "0.79.0"
|
||||
dependencies = [
|
||||
"async-trait",
|
||||
"parking_lot",
|
||||
@@ -7282,7 +7284,7 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "proxycast-providers"
|
||||
version = "0.78.0"
|
||||
version = "0.79.0"
|
||||
dependencies = [
|
||||
"anyhow",
|
||||
"async-stream",
|
||||
@@ -7334,7 +7336,7 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "proxycast-server"
|
||||
version = "0.78.0"
|
||||
version = "0.79.0"
|
||||
dependencies = [
|
||||
"aster-core",
|
||||
"async-stream",
|
||||
@@ -7379,7 +7381,7 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "proxycast-server-utils"
|
||||
version = "0.78.0"
|
||||
version = "0.79.0"
|
||||
dependencies = [
|
||||
"axum 0.7.9",
|
||||
"futures",
|
||||
@@ -7394,7 +7396,7 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "proxycast-services"
|
||||
version = "0.78.0"
|
||||
version = "0.79.0"
|
||||
dependencies = [
|
||||
"anyhow",
|
||||
"aster-core",
|
||||
@@ -7435,7 +7437,7 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "proxycast-skills"
|
||||
version = "0.78.0"
|
||||
version = "0.79.0"
|
||||
dependencies = [
|
||||
"async-trait",
|
||||
"dirs 5.0.1",
|
||||
@@ -7451,7 +7453,7 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "proxycast-terminal"
|
||||
version = "0.78.0"
|
||||
version = "0.79.0"
|
||||
dependencies = [
|
||||
"async-trait",
|
||||
"base64 0.22.1",
|
||||
@@ -7478,7 +7480,7 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "proxycast-websocket"
|
||||
version = "0.78.0"
|
||||
version = "0.79.0"
|
||||
dependencies = [
|
||||
"axum 0.7.9",
|
||||
"chrono",
|
||||
|
||||
@@ -3,7 +3,7 @@ members = ["crates/*"]
|
||||
resolver = "2"
|
||||
|
||||
[workspace.package]
|
||||
version = "0.78.0"
|
||||
version = "0.79.0"
|
||||
edition = "2021"
|
||||
authors = ["coso"]
|
||||
repository = "https://github.com/aiclientproxy/proxycast"
|
||||
@@ -190,7 +190,7 @@ version = "2.4"
|
||||
|
||||
[package]
|
||||
name = "proxycast"
|
||||
version = "0.77.0"
|
||||
version.workspace = true
|
||||
description = "AI API Proxy Desktop App"
|
||||
authors = ["you"]
|
||||
edition = "2021"
|
||||
|
||||
@@ -26,3 +26,4 @@ regex.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
tempfile.workspace = true
|
||||
futures.workspace = true
|
||||
|
||||
@@ -17,6 +17,7 @@ use proxycast_core::database::DbConnection;
|
||||
use proxycast_core::models::provider_pool_model::{
|
||||
CredentialData, PoolProviderType, ProviderCredential,
|
||||
};
|
||||
use proxycast_core::models::provider_type::is_custom_provider_id;
|
||||
use proxycast_services::api_key_provider_service::ApiKeyProviderService;
|
||||
use proxycast_services::provider_pool_service::ProviderPoolService;
|
||||
use std::sync::Arc;
|
||||
@@ -63,6 +64,8 @@ pub struct AsterProviderConfig {
|
||||
pub base_url: Option<String>,
|
||||
/// 凭证 UUID(用于记录使用和健康状态)
|
||||
pub credential_uuid: String,
|
||||
/// 是否强制 OpenAI provider 使用 Responses API(用于 Codex 等兼容链路)
|
||||
pub force_responses_api: bool,
|
||||
}
|
||||
|
||||
/// 凭证池桥接器
|
||||
@@ -128,6 +131,39 @@ impl CredentialBridge {
|
||||
.await
|
||||
}
|
||||
|
||||
fn resolve_api_provider_type_hint(
|
||||
&self,
|
||||
db: &DbConnection,
|
||||
provider_type_hint: &str,
|
||||
) -> Option<ApiProviderType> {
|
||||
if let Ok(api_type) = provider_type_hint.parse::<ApiProviderType>() {
|
||||
return Some(api_type);
|
||||
}
|
||||
|
||||
if !is_custom_provider_id(provider_type_hint) {
|
||||
return None;
|
||||
}
|
||||
|
||||
match self.api_key_service.get_provider(db, provider_type_hint) {
|
||||
Ok(Some(provider_with_keys)) => Some(provider_with_keys.provider.provider_type),
|
||||
Ok(None) => {
|
||||
tracing::warn!(
|
||||
"[CredentialBridge] custom provider 不存在: {}, 使用默认映射",
|
||||
provider_type_hint
|
||||
);
|
||||
None
|
||||
}
|
||||
Err(error) => {
|
||||
tracing::warn!(
|
||||
"[CredentialBridge] 读取 custom provider 失败: {} ({}),使用默认映射",
|
||||
provider_type_hint,
|
||||
error
|
||||
);
|
||||
None
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 将 ProxyCast 凭证转换为 Aster Provider 配置
|
||||
async fn credential_to_config(
|
||||
&self,
|
||||
@@ -142,20 +178,23 @@ impl CredentialBridge {
|
||||
credential.provider_type
|
||||
);
|
||||
|
||||
let (provider_name, api_key, base_url) = match &credential.credential {
|
||||
let (provider_name, api_key, base_url, force_responses_api) = match &credential.credential {
|
||||
// OpenAI API Key - 根据 provider_type_hint 确定实际的 Provider
|
||||
CredentialData::OpenAIKey { api_key, base_url } => {
|
||||
// 使用 provider_type_hint 来确定 aster provider 名称
|
||||
let provider = map_provider_type_to_aster(provider_type_hint);
|
||||
let resolved_api_type = self.resolve_api_provider_type_hint(db, provider_type_hint);
|
||||
let provider =
|
||||
map_provider_type_to_aster_with_api_type(provider_type_hint, resolved_api_type);
|
||||
tracing::info!(
|
||||
"[CredentialBridge] OpenAIKey: provider_type_hint={} -> aster_provider={}",
|
||||
"[CredentialBridge] OpenAIKey: provider_type_hint={}, resolved_api_type={:?} -> aster_provider={}",
|
||||
provider_type_hint,
|
||||
resolved_api_type,
|
||||
provider
|
||||
);
|
||||
(
|
||||
provider.to_string(),
|
||||
Some(api_key.clone()),
|
||||
base_url.clone(),
|
||||
resolved_api_type == Some(ApiProviderType::Codex),
|
||||
)
|
||||
}
|
||||
|
||||
@@ -165,6 +204,7 @@ impl CredentialBridge {
|
||||
"anthropic".to_string(),
|
||||
Some(api_key.clone()),
|
||||
base_url.clone(),
|
||||
false,
|
||||
),
|
||||
|
||||
// Kiro OAuth - 需要获取 access_token
|
||||
@@ -173,7 +213,7 @@ impl CredentialBridge {
|
||||
.get_kiro_token(creds_file_path, db, &credential.uuid)
|
||||
.await?;
|
||||
// Kiro 使用 CodeWhisperer API,映射到 bedrock provider
|
||||
("bedrock".to_string(), Some(token), None)
|
||||
("bedrock".to_string(), Some(token), None, false)
|
||||
}
|
||||
|
||||
// Gemini OAuth
|
||||
@@ -181,7 +221,7 @@ impl CredentialBridge {
|
||||
creds_file_path, ..
|
||||
} => {
|
||||
let token = self.get_oauth_token(creds_file_path).await?;
|
||||
("google".to_string(), Some(token), None)
|
||||
("google".to_string(), Some(token), None, false)
|
||||
}
|
||||
|
||||
// Gemini API Key
|
||||
@@ -191,6 +231,7 @@ impl CredentialBridge {
|
||||
"google".to_string(),
|
||||
Some(api_key.clone()),
|
||||
base_url.clone(),
|
||||
false,
|
||||
),
|
||||
|
||||
// Vertex AI
|
||||
@@ -200,6 +241,7 @@ impl CredentialBridge {
|
||||
"gcpvertexai".to_string(),
|
||||
Some(api_key.clone()),
|
||||
base_url.clone(),
|
||||
false,
|
||||
),
|
||||
|
||||
// Codex OAuth
|
||||
@@ -208,13 +250,19 @@ impl CredentialBridge {
|
||||
api_base_url,
|
||||
} => {
|
||||
let token = self.get_codex_token(creds_file_path).await?;
|
||||
("codex".to_string(), Some(token), api_base_url.clone())
|
||||
(
|
||||
// 统一走 OpenAI provider,保证 tools/stream 事件链路一致
|
||||
"openai".to_string(),
|
||||
Some(token),
|
||||
api_base_url.clone(),
|
||||
true,
|
||||
)
|
||||
}
|
||||
|
||||
// Claude OAuth
|
||||
CredentialData::ClaudeOAuth { creds_file_path } => {
|
||||
let token = self.get_oauth_token(creds_file_path).await?;
|
||||
("anthropic".to_string(), Some(token), None)
|
||||
("anthropic".to_string(), Some(token), None, false)
|
||||
}
|
||||
|
||||
// Antigravity OAuth
|
||||
@@ -222,7 +270,7 @@ impl CredentialBridge {
|
||||
creds_file_path, ..
|
||||
} => {
|
||||
let token = self.get_oauth_token(creds_file_path).await?;
|
||||
("google".to_string(), Some(token), None)
|
||||
("google".to_string(), Some(token), None, false)
|
||||
}
|
||||
};
|
||||
|
||||
@@ -232,6 +280,7 @@ impl CredentialBridge {
|
||||
api_key,
|
||||
base_url,
|
||||
credential_uuid: credential.uuid.clone(),
|
||||
force_responses_api,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -418,12 +467,32 @@ fn set_provider_env_vars(config: &AsterProviderConfig) {
|
||||
std::env::set_var(env_key, api_key);
|
||||
}
|
||||
|
||||
if config.provider_name == "openai" {
|
||||
if config.force_responses_api {
|
||||
std::env::set_var("OPENAI_FORCE_RESPONSES_API", "1");
|
||||
} else {
|
||||
std::env::remove_var("OPENAI_FORCE_RESPONSES_API");
|
||||
}
|
||||
}
|
||||
|
||||
// 设置 base_url
|
||||
// Aster 的 OpenAI Provider 使用 OPENAI_HOST(仅 scheme+host+port)和
|
||||
// OPENAI_BASE_PATH(路径部分 + /chat/completions)环境变量
|
||||
if let Some(base_url) = &config.base_url {
|
||||
match config.provider_name.as_str() {
|
||||
"openai" => {
|
||||
// 当显式强制 responses 模式时,需要将路径前缀保留在 OPENAI_HOST 中,
|
||||
// 因为 Aster OpenAI provider 在 responses 模式下固定请求 v1/responses,
|
||||
// 不会读取 OPENAI_BASE_PATH。
|
||||
if config.force_responses_api {
|
||||
std::env::set_var("OPENAI_HOST", base_url);
|
||||
std::env::remove_var("OPENAI_BASE_PATH");
|
||||
tracing::info!(
|
||||
"[CredentialBridge] 强制 Responses 模式: 设置 OPENAI_HOST={}, 清理 OPENAI_BASE_PATH",
|
||||
base_url
|
||||
);
|
||||
return;
|
||||
}
|
||||
// 解析 base_url,将路径部分拆分到 OPENAI_BASE_PATH
|
||||
// 例如 https://open.bigmodel.cn/api/paas/v4
|
||||
// -> OPENAI_HOST = https://open.bigmodel.cn
|
||||
@@ -525,6 +594,21 @@ fn map_provider_type_to_aster(provider_type: &str) -> &'static str {
|
||||
}
|
||||
}
|
||||
|
||||
fn map_provider_type_to_aster_with_api_type(
|
||||
provider_type: &str,
|
||||
resolved_api_type: Option<ApiProviderType>,
|
||||
) -> &'static str {
|
||||
if let Some(api_type) = resolved_api_type {
|
||||
// Codex API Key 在 Aster 中应走 OpenAI provider(支持标准 tools + responses 转换逻辑),
|
||||
// 避免误走 codex CLI provider 导致工具事件丢失。
|
||||
if api_type == ApiProviderType::Codex {
|
||||
return "openai";
|
||||
}
|
||||
return api_type.runtime_spec().aster_provider_name;
|
||||
}
|
||||
map_provider_type_to_aster(provider_type)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
@@ -540,6 +624,56 @@ mod tests {
|
||||
assert_eq!(map_pool_type_to_aster(&PoolProviderType::Kiro), "bedrock");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_map_provider_type_to_aster_with_api_type() {
|
||||
assert_eq!(
|
||||
map_provider_type_to_aster_with_api_type(
|
||||
"custom-a32774c6-6fd0-433b-8b81-e95340e08793",
|
||||
Some(ApiProviderType::Codex),
|
||||
),
|
||||
"openai"
|
||||
);
|
||||
assert_eq!(
|
||||
map_provider_type_to_aster_with_api_type(
|
||||
"custom-a32774c6-6fd0-433b-8b81-e95340e08793",
|
||||
Some(ApiProviderType::AnthropicCompatible),
|
||||
),
|
||||
"anthropic"
|
||||
);
|
||||
assert_eq!(
|
||||
map_provider_type_to_aster_with_api_type("deepseek", None),
|
||||
"openai"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_set_provider_env_vars_openai_codex_responses_keeps_full_base_url() {
|
||||
std::env::remove_var("OPENAI_HOST");
|
||||
std::env::remove_var("OPENAI_BASE_PATH");
|
||||
std::env::remove_var("OPENAI_FORCE_RESPONSES_API");
|
||||
|
||||
let config = AsterProviderConfig {
|
||||
provider_name: "openai".to_string(),
|
||||
model_name: "gpt-5.3-codex".to_string(),
|
||||
api_key: Some("test-key".to_string()),
|
||||
base_url: Some("https://example.com/openai".to_string()),
|
||||
credential_uuid: "test-uuid".to_string(),
|
||||
force_responses_api: true,
|
||||
};
|
||||
|
||||
set_provider_env_vars(&config);
|
||||
|
||||
assert_eq!(
|
||||
std::env::var("OPENAI_HOST").ok(),
|
||||
Some("https://example.com/openai".to_string())
|
||||
);
|
||||
assert!(std::env::var("OPENAI_BASE_PATH").is_err());
|
||||
assert_eq!(
|
||||
std::env::var("OPENAI_FORCE_RESPONSES_API").ok().as_deref(),
|
||||
Some("1")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_credential_bridge_error_display() {
|
||||
let err = CredentialBridgeError::NoCredentials("test".to_string());
|
||||
@@ -582,6 +716,7 @@ mod tests {
|
||||
api_key: Some("test-key".to_string()),
|
||||
base_url: Some("https://open.bigmodel.cn/api/anthropic".to_string()),
|
||||
credential_uuid: "test-uuid".to_string(),
|
||||
force_responses_api: false,
|
||||
};
|
||||
|
||||
set_provider_env_vars(&config);
|
||||
|
||||
@@ -56,6 +56,74 @@ fn default_timeout() -> u64 {
|
||||
10
|
||||
}
|
||||
|
||||
fn shell_command_flag(shell: &str) -> &'static str {
|
||||
let executable = std::path::Path::new(shell)
|
||||
.file_name()
|
||||
.and_then(|name| name.to_str())
|
||||
.unwrap_or(shell)
|
||||
.to_ascii_lowercase();
|
||||
|
||||
if executable == "cmd" || executable == "cmd.exe" {
|
||||
"/C"
|
||||
} else if executable.contains("powershell") || executable == "pwsh" || executable == "pwsh.exe"
|
||||
{
|
||||
"-Command"
|
||||
} else {
|
||||
"-c"
|
||||
}
|
||||
}
|
||||
|
||||
fn resolve_command_shell() -> String {
|
||||
let shell_from_env = std::env::var("SHELL").ok().and_then(|value| {
|
||||
let cleaned = value
|
||||
.split('\0')
|
||||
.next()
|
||||
.unwrap_or_default()
|
||||
.trim()
|
||||
.to_string();
|
||||
(!cleaned.is_empty()).then_some(cleaned)
|
||||
});
|
||||
|
||||
#[cfg(target_os = "windows")]
|
||||
{
|
||||
if let Some(shell) = shell_from_env {
|
||||
let path = std::path::Path::new(&shell);
|
||||
if path.is_absolute() && path.exists() {
|
||||
return shell;
|
||||
}
|
||||
}
|
||||
|
||||
if let Ok(comspec) = std::env::var("COMSPEC") {
|
||||
let cleaned = comspec
|
||||
.split('\0')
|
||||
.next()
|
||||
.unwrap_or_default()
|
||||
.trim()
|
||||
.to_string();
|
||||
if !cleaned.is_empty() && std::path::Path::new(&cleaned).exists() {
|
||||
return cleaned;
|
||||
}
|
||||
}
|
||||
|
||||
"cmd.exe".to_string()
|
||||
}
|
||||
|
||||
#[cfg(not(target_os = "windows"))]
|
||||
{
|
||||
if let Some(shell) = shell_from_env {
|
||||
if std::path::Path::new(&shell).exists() {
|
||||
return shell;
|
||||
}
|
||||
}
|
||||
|
||||
if std::path::Path::new("/bin/sh").exists() {
|
||||
"/bin/sh".to_string()
|
||||
} else {
|
||||
"sh".to_string()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Hook 执行结果
|
||||
#[derive(Debug)]
|
||||
pub struct HookResult {
|
||||
@@ -209,9 +277,11 @@ impl HookManager {
|
||||
|
||||
async fn execute_hook(hook: &HookDefinition, context: &HookContext) -> HookResult {
|
||||
let context_json = serde_json::to_string(context).unwrap_or_default();
|
||||
let shell = resolve_command_shell();
|
||||
let shell_flag = shell_command_flag(&shell);
|
||||
|
||||
let child = Command::new("sh")
|
||||
.arg("-c")
|
||||
let child = Command::new(&shell)
|
||||
.arg(shell_flag)
|
||||
.arg(&hook.command)
|
||||
.env(
|
||||
"HOOK_EVENT",
|
||||
@@ -411,6 +481,19 @@ mod tests {
|
||||
assert!(!HookManager::is_blocked(&results_ok));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_shell_command_flag_for_common_shells() {
|
||||
assert_eq!(shell_command_flag("cmd.exe"), "/C");
|
||||
assert_eq!(shell_command_flag("pwsh"), "-Command");
|
||||
assert_eq!(shell_command_flag("/bin/sh"), "-c");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_resolve_command_shell_not_empty() {
|
||||
let shell = resolve_command_shell();
|
||||
assert!(!shell.trim().is_empty());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_trigger_executes_matching_hooks() {
|
||||
let mut mgr = HookManager::new();
|
||||
|
||||
@@ -0,0 +1,167 @@
|
||||
use futures::StreamExt;
|
||||
use proxycast_agent::{
|
||||
convert_agent_event, AsterAgentState, SessionConfigBuilder, TauriAgentEvent,
|
||||
};
|
||||
use proxycast_core::database::dao::api_key_provider::ApiProviderType;
|
||||
use proxycast_core::database::init_database;
|
||||
use proxycast_services::api_key_provider_service::ApiKeyProviderService;
|
||||
use uuid::Uuid;
|
||||
|
||||
fn should_run_real_test() -> bool {
|
||||
std::env::var("PROXYCAST_REAL_API_TEST").ok().as_deref() == Some("1")
|
||||
}
|
||||
|
||||
fn resolve_model_name(
|
||||
explicit: Option<String>,
|
||||
provider_models: &[String],
|
||||
) -> Result<String, String> {
|
||||
if let Some(model) = explicit
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
{
|
||||
return Ok(model.to_string());
|
||||
}
|
||||
|
||||
if let Some(model) = provider_models
|
||||
.iter()
|
||||
.map(|value| value.trim())
|
||||
.find(|value| !value.is_empty())
|
||||
{
|
||||
return Ok(model.to_string());
|
||||
}
|
||||
|
||||
Err(
|
||||
"未找到可用模型:请设置 PROXYCAST_REAL_MODEL,或在 Provider custom_models 中配置模型。"
|
||||
.to_string(),
|
||||
)
|
||||
}
|
||||
|
||||
fn resolve_codex_provider_and_model(
|
||||
db: &proxycast_core::database::DbConnection,
|
||||
) -> Result<(String, String), String> {
|
||||
let explicit_model = std::env::var("PROXYCAST_REAL_MODEL").ok();
|
||||
|
||||
if let Ok(explicit) = std::env::var("PROXYCAST_REAL_PROVIDER_ID") {
|
||||
let trimmed = explicit.trim();
|
||||
if !trimmed.is_empty() {
|
||||
let service = ApiKeyProviderService::new();
|
||||
let provider = service
|
||||
.get_provider(db, trimmed)?
|
||||
.ok_or_else(|| format!("未找到指定 Provider: {trimmed}"))?;
|
||||
let model = resolve_model_name(explicit_model, &provider.provider.custom_models)?;
|
||||
return Ok((trimmed.to_string(), model));
|
||||
}
|
||||
}
|
||||
|
||||
let service = ApiKeyProviderService::new();
|
||||
let providers = service.get_all_providers(db)?;
|
||||
providers
|
||||
.into_iter()
|
||||
.find(|item| {
|
||||
item.provider.enabled
|
||||
&& item.provider.provider_type == ApiProviderType::Codex
|
||||
&& item.api_keys.iter().any(|key| key.enabled)
|
||||
})
|
||||
.map(|item| {
|
||||
let model = resolve_model_name(explicit_model, &item.provider.custom_models)?;
|
||||
Ok((item.provider.id, model))
|
||||
})
|
||||
.transpose()?
|
||||
.ok_or_else(|| "未找到启用且含可用 Key 的 Codex Provider".to_string())
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
#[ignore = "真实联网测试:设置 PROXYCAST_REAL_API_TEST=1 后执行"]
|
||||
async fn test_real_codex_stream_emits_tool_events() {
|
||||
if !should_run_real_test() {
|
||||
return;
|
||||
}
|
||||
|
||||
let db = init_database().expect("初始化数据库失败");
|
||||
let (provider_id, model_name) =
|
||||
resolve_codex_provider_and_model(&db).expect("解析 Codex Provider/模型失败");
|
||||
let session_id = format!("real-codex-tool-{}", Uuid::new_v4());
|
||||
|
||||
let state = AsterAgentState::new();
|
||||
state
|
||||
.configure_provider_from_pool(&db, &provider_id, &model_name, &session_id)
|
||||
.await
|
||||
.expect("配置 Provider 失败");
|
||||
|
||||
let agent_arc = state.get_agent_arc();
|
||||
let agent_guard = agent_arc.read().await;
|
||||
let agent = agent_guard.as_ref().expect("Agent 未初始化");
|
||||
|
||||
let tools = agent.list_tools(None).await;
|
||||
assert!(
|
||||
!tools.is_empty(),
|
||||
"工具列表为空,无法验证 tool_start/tool_end"
|
||||
);
|
||||
|
||||
let preferred_tool = tools
|
||||
.iter()
|
||||
.map(|tool| tool.name.to_string())
|
||||
.find(|name| name.contains("list_tools"))
|
||||
.unwrap_or_else(|| "bash".to_string());
|
||||
|
||||
let prompt = if preferred_tool == "bash" {
|
||||
"请严格执行以下步骤:\
|
||||
1) 必须调用工具 bash,执行命令 `echo PROXYCAST_REAL_TOOL_EVENT`; \
|
||||
2) 然后只回复 `REAL_TOOL_OK`。"
|
||||
.to_string()
|
||||
} else {
|
||||
format!(
|
||||
"请严格执行以下步骤:\
|
||||
1) 必须调用工具 `{}` 一次;\
|
||||
2) 如果需要参数请传空对象;\
|
||||
3) 然后只回复 `REAL_TOOL_OK`。",
|
||||
preferred_tool
|
||||
)
|
||||
};
|
||||
|
||||
let user_message = aster::conversation::message::Message::user().with_text(prompt);
|
||||
let session_config = SessionConfigBuilder::new(&session_id).build();
|
||||
let mut stream = agent
|
||||
.reply(user_message, session_config, None)
|
||||
.await
|
||||
.expect("创建流式回复失败");
|
||||
|
||||
let mut tool_start_count = 0usize;
|
||||
let mut tool_end_count = 0usize;
|
||||
let mut error_messages: Vec<String> = Vec::new();
|
||||
let mut text_buffer = String::new();
|
||||
|
||||
while let Some(event_result) = stream.next().await {
|
||||
match event_result {
|
||||
Ok(agent_event) => {
|
||||
for event in convert_agent_event(agent_event) {
|
||||
match event {
|
||||
TauriAgentEvent::ToolStart { .. } => tool_start_count += 1,
|
||||
TauriAgentEvent::ToolEnd { .. } => tool_end_count += 1,
|
||||
TauriAgentEvent::TextDelta { text } => text_buffer.push_str(&text),
|
||||
TauriAgentEvent::Error { message } => error_messages.push(message),
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
}
|
||||
Err(err) => error_messages.push(format!("stream_error: {err}")),
|
||||
}
|
||||
}
|
||||
|
||||
assert!(
|
||||
error_messages.is_empty(),
|
||||
"流式过程中出现错误: {:?}",
|
||||
error_messages
|
||||
);
|
||||
assert!(
|
||||
tool_start_count > 0,
|
||||
"未收到 tool_start 事件,文本输出: {}",
|
||||
text_buffer
|
||||
);
|
||||
assert!(
|
||||
tool_end_count > 0,
|
||||
"未收到 tool_end 事件,文本输出: {}",
|
||||
text_buffer
|
||||
);
|
||||
}
|
||||
@@ -123,7 +123,7 @@ impl ProviderType {
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::ProviderType;
|
||||
use super::{ProviderType, DEFAULT_SYSTEM_PROMPT};
|
||||
|
||||
#[test]
|
||||
fn test_custom_provider_does_not_force_anthropic_protocol() {
|
||||
@@ -147,6 +147,18 @@ mod tests {
|
||||
ProviderType::Gemini
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_default_system_prompt_mentions_web_tools() {
|
||||
assert!(DEFAULT_SYSTEM_PROMPT.contains("WebSearch"));
|
||||
assert!(DEFAULT_SYSTEM_PROMPT.contains("WebFetch"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_default_system_prompt_requires_authorization_for_local_ops_only() {
|
||||
assert!(DEFAULT_SYSTEM_PROMPT.contains("本地操作需授权"));
|
||||
assert!(DEFAULT_SYSTEM_PROMPT.contains("实时信息"));
|
||||
}
|
||||
}
|
||||
|
||||
/// Agent 会话状态
|
||||
@@ -331,14 +343,16 @@ pub const DEFAULT_SYSTEM_PROMPT: &str = r#"你是 ProxyCast 内置的 AI 助手
|
||||
|
||||
<core_principles>
|
||||
1. **自然交流优先**:问候、闲聊、问答类对话,直接用文字回复
|
||||
2. **显式授权操作**:只有当用户明确提供路径或命令时,才能执行工具
|
||||
3. **不主动探索**:不要未经请求就读取文件或执行命令
|
||||
2. **按任务选择工具**:当工具能显著提升准确性或完成度时,应主动调用工具
|
||||
3. **本地操作需授权**:读取/修改本地文件、执行本地命令前,需有用户明确意图
|
||||
4. **不无依据臆测**:涉及实时信息或外部事实时,优先通过工具检索再回答
|
||||
</core_principles>
|
||||
|
||||
<tool_use_rules>
|
||||
## 何时使用工具
|
||||
|
||||
✅ **使用工具的情况**:
|
||||
- 用户需要实时信息、新闻、外部来源或网页事实核验
|
||||
- 用户明确提供了文件路径(如 "读取 /path/to/file")
|
||||
- 用户明确要求执行命令(如 "运行 npm install")
|
||||
- 用户要求创建或修改文件
|
||||
@@ -347,14 +361,16 @@ pub const DEFAULT_SYSTEM_PROMPT: &str = r#"你是 ProxyCast 内置的 AI 助手
|
||||
- 用户说 "你好"、"嗨"、"hello" 等问候语
|
||||
- 用户进行闲聊或一般性提问
|
||||
- 用户没有提供具体路径时猜测路径
|
||||
- 为了 "了解环境" 或 "打招呼" 而读取文件
|
||||
- 为了 "了解环境" 或 "打招呼" 而读取本地文件/执行本地命令
|
||||
|
||||
## 可用工具
|
||||
|
||||
- **read_file**:读取用户指定的文件或目录内容
|
||||
- **write_file**:创建或覆盖用户指定的文件
|
||||
- **edit_file**:修改用户指定文件的特定内容
|
||||
- **read**:读取用户指定的文件或目录内容
|
||||
- **write**:创建或覆盖用户指定的文件
|
||||
- **edit**:修改用户指定文件的特定内容
|
||||
- **bash**:执行用户要求的 shell 命令
|
||||
- **WebSearch**:联网搜索公开网页信息
|
||||
- **WebFetch**:抓取并提取指定网页内容
|
||||
</tool_use_rules>
|
||||
|
||||
<response_examples>
|
||||
|
||||
@@ -15,3 +15,4 @@ tracing.workspace = true
|
||||
thiserror.workspace = true
|
||||
glob.workspace = true
|
||||
rmcp.workspace = true
|
||||
dirs.workspace = true
|
||||
|
||||
@@ -445,8 +445,8 @@ impl McpClientManager {
|
||||
}
|
||||
}
|
||||
|
||||
// 设置工作目录
|
||||
if let Some(ref cwd) = config.cwd {
|
||||
// 设置工作目录(清洗 `\0` 和无效空白)
|
||||
if let Some(cwd) = config.sanitized_cwd() {
|
||||
command.current_dir(cwd);
|
||||
}
|
||||
|
||||
|
||||
@@ -10,6 +10,7 @@
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::collections::HashMap;
|
||||
use std::path::PathBuf;
|
||||
|
||||
// ============================================================================
|
||||
// 服务器配置和状态
|
||||
@@ -37,6 +38,29 @@ fn default_timeout() -> u64 {
|
||||
30
|
||||
}
|
||||
|
||||
impl McpServerConfig {
|
||||
/// 获取清洗后的工作目录(去除 `\0`、首尾空白,并展开 `~`)
|
||||
pub fn sanitized_cwd(&self) -> Option<PathBuf> {
|
||||
let cwd = self.cwd.as_deref()?;
|
||||
let cleaned = cwd.split('\0').next().unwrap_or_default().trim();
|
||||
if cleaned.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
if cleaned == "~" {
|
||||
return Some(dirs::home_dir().unwrap_or_else(|| PathBuf::from(cleaned)));
|
||||
}
|
||||
|
||||
if cleaned.starts_with("~/") || cleaned.starts_with("~\\") {
|
||||
if let Some(home) = dirs::home_dir() {
|
||||
return Some(home.join(&cleaned[2..]));
|
||||
}
|
||||
}
|
||||
|
||||
Some(PathBuf::from(cleaned))
|
||||
}
|
||||
}
|
||||
|
||||
/// MCP 服务器信息(包含运行状态)
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct McpServerInfo {
|
||||
@@ -257,3 +281,32 @@ use tokio::sync::Mutex;
|
||||
///
|
||||
/// 使用 Arc<Mutex<McpClientManager>> 包装,支持跨线程共享和异步访问。
|
||||
pub type McpManagerState = Arc<Mutex<super::manager::McpClientManager>>;
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::McpServerConfig;
|
||||
use std::collections::HashMap;
|
||||
use std::path::PathBuf;
|
||||
|
||||
fn sample_config(cwd: Option<String>) -> McpServerConfig {
|
||||
McpServerConfig {
|
||||
command: "npx".to_string(),
|
||||
args: vec!["-y".to_string(), "some-server".to_string()],
|
||||
env: HashMap::new(),
|
||||
cwd,
|
||||
timeout: 30,
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn sanitized_cwd_should_strip_nul_suffix() {
|
||||
let config = sample_config(Some(" /tmp/demo\0ignored ".to_string()));
|
||||
assert_eq!(config.sanitized_cwd(), Some(PathBuf::from("/tmp/demo")));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn sanitized_cwd_should_reject_empty_value() {
|
||||
let config = sample_config(Some(" \0 ".to_string()));
|
||||
assert!(config.sanitized_cwd().is_none());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1229,155 +1229,58 @@ fn parse_jwt_claims(token: &str) -> (Option<String>, Option<String>) {
|
||||
(account_id, email)
|
||||
}
|
||||
|
||||
/// 根据模型名称获取对应的 Codex instructions
|
||||
/// 参考 CLIProxyAPI: internal/misc/codex_instructions.go
|
||||
fn get_codex_instructions_for_model(model_name: &str) -> &'static str {
|
||||
let model_lower = model_name.to_lowercase();
|
||||
|
||||
if model_lower.contains("codex-max") {
|
||||
// GPT-5.1 Codex Max 专用 prompt
|
||||
CODEX_MAX_INSTRUCTIONS
|
||||
} else if model_lower.contains("5.2-codex") {
|
||||
// GPT-5.2 Codex 专用 prompt
|
||||
CODEX_52_INSTRUCTIONS
|
||||
} else if model_lower.contains("codex") {
|
||||
// GPT-5 Codex 通用 prompt
|
||||
CODEX_INSTRUCTIONS
|
||||
} else if model_lower.contains("5.1") {
|
||||
// GPT-5.1 通用 prompt
|
||||
GPT_51_INSTRUCTIONS
|
||||
} else if model_lower.contains("5.2") {
|
||||
// GPT-5.2 通用 prompt
|
||||
GPT_52_INSTRUCTIONS
|
||||
} else {
|
||||
// 默认使用 Codex prompt
|
||||
CODEX_INSTRUCTIONS
|
||||
fn extract_text_fragments(content: &serde_json::Value) -> Vec<String> {
|
||||
if let Some(text) = content
|
||||
.as_str()
|
||||
.map(str::trim)
|
||||
.filter(|text| !text.is_empty())
|
||||
{
|
||||
return vec![text.to_string()];
|
||||
}
|
||||
|
||||
let mut fragments = Vec::new();
|
||||
if let Some(parts) = content.as_array() {
|
||||
for part in parts {
|
||||
let text = if let Some(kind) = part.get("type").and_then(|v| v.as_str()) {
|
||||
match kind {
|
||||
"text" | "input_text" | "output_text" => {
|
||||
part.get("text").and_then(|v| v.as_str())
|
||||
}
|
||||
_ => part.get("text").and_then(|v| v.as_str()),
|
||||
}
|
||||
} else {
|
||||
part.get("text").and_then(|v| v.as_str())
|
||||
};
|
||||
if let Some(text) = text.map(str::trim).filter(|text| !text.is_empty()) {
|
||||
fragments.push(text.to_string());
|
||||
}
|
||||
}
|
||||
}
|
||||
fragments
|
||||
}
|
||||
|
||||
// GPT-5 Codex 通用 prompt(最新版本)
|
||||
// 来源: CLIProxyAPI/internal/misc/codex_instructions/gpt_5_codex_prompt.md-009
|
||||
const CODEX_INSTRUCTIONS: &str = r#"You are Codex, based on GPT-5. You are running as a coding agent in the Codex CLI on a user's computer.
|
||||
fn resolve_codex_instructions(
|
||||
request: &serde_json::Value,
|
||||
system_instructions: &[String],
|
||||
) -> Option<String> {
|
||||
if let Some(request_instructions) = request
|
||||
.get("instructions")
|
||||
.and_then(|value| value.as_str())
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
{
|
||||
return Some(request_instructions.to_string());
|
||||
}
|
||||
|
||||
## General
|
||||
if !system_instructions.is_empty() {
|
||||
return Some(system_instructions.join("\n\n"));
|
||||
}
|
||||
|
||||
- When searching for text or files, prefer using `rg` or `rg --files` respectively because `rg` is much faster than alternatives like `grep`. (If the `rg` command is not found, then use alternatives.)
|
||||
|
||||
## Editing constraints
|
||||
|
||||
- Default to ASCII when editing or creating files. Only introduce non-ASCII or other Unicode characters when there is a clear justification and the file already uses them.
|
||||
- Add succinct code comments that explain what is going on if code is not self-explanatory. You should not add comments like "Assigns the value to the variable", but a brief comment might be useful ahead of a complex code block that the user would otherwise have to spend time parsing out. Usage of these comments should be rare.
|
||||
- Try to use apply_patch for single file edits, but it is fine to explore other options to make the edit if it does not work well. Do not use apply_patch for changes that are auto-generated (i.e. generating package.json or running a lint or format command like gofmt) or when scripting is more efficient (such as search and replacing a string across a codebase).
|
||||
- You may be in a dirty git worktree.
|
||||
* NEVER revert existing changes you did not make unless explicitly requested, since these changes were made by the user.
|
||||
* If asked to make a commit or code edits and there are unrelated changes to your work or changes that you didn't make in those files, don't revert those changes.
|
||||
* If the changes are in files you've touched recently, you should read carefully and understand how you can work with the changes rather than reverting them.
|
||||
* If the changes are in unrelated files, just ignore them and don't revert them.
|
||||
- Do not amend a commit unless explicitly requested to do so.
|
||||
- While you are working, you might notice unexpected changes that you didn't make. If this happens, STOP IMMEDIATELY and ask the user how they would like to proceed.
|
||||
- **NEVER** use destructive commands like `git reset --hard` or `git checkout --` unless specifically requested or approved by the user.
|
||||
|
||||
## Plan tool
|
||||
|
||||
When using the planning tool:
|
||||
- Skip using the planning tool for straightforward tasks (roughly the easiest 25%).
|
||||
- Do not make single-step plans.
|
||||
- When you made a plan, update it after having performed one of the sub-tasks that you shared on the plan.
|
||||
|
||||
## Codex CLI harness, sandboxing, and approvals
|
||||
|
||||
The Codex CLI harness supports several different configurations for sandboxing and escalation approvals that the user can choose from.
|
||||
|
||||
Filesystem sandboxing defines which files can be read or written. The options for `sandbox_mode` are:
|
||||
- **read-only**: The sandbox only permits reading files.
|
||||
- **workspace-write**: The sandbox permits reading files, and editing files in `cwd` and `writable_roots`. Editing files in other directories requires approval.
|
||||
- **danger-full-access**: No filesystem sandboxing - all commands are permitted.
|
||||
|
||||
Network sandboxing defines whether network can be accessed without approval. Options for `network_access` are:
|
||||
- **restricted**: Requires approval
|
||||
- **enabled**: No approval needed
|
||||
|
||||
Approvals are your mechanism to get user consent to run shell commands without the sandbox. Possible configuration options for `approval_policy` are
|
||||
- **untrusted**: The harness will escalate most commands for user approval, apart from a limited allowlist of safe "read" commands.
|
||||
- **on-failure**: The harness will allow all commands to run in the sandbox (if enabled), and failures will be escalated to the user for approval to run again without the sandbox.
|
||||
- **on-request**: Commands will be run in the sandbox by default, and you can specify in your tool call if you want to escalate a command to run without sandboxing. (Note that this mode is not always available. If it is, you'll see parameters for it in the `shell` command description.)
|
||||
- **never**: This is a non-interactive mode where you may NEVER ask the user for approval to run commands. Instead, you must always persist and work around constraints to solve the task for the user. You MUST do your utmost best to finish the task and validate your work before yielding. If this mode is paired with `danger-full-access`, take advantage of it to deliver the best outcome for the user. Further, in this mode, your default testing philosophy is overridden: Even if you don't see local patterns for testing, you may add tests and scripts to validate your work. Just remove them before yielding.
|
||||
|
||||
When you are running with `approval_policy == on-request`, and sandboxing enabled, here are scenarios where you'll need to request approval:
|
||||
- You need to run a command that writes to a directory that requires it (e.g. running tests that write to /var)
|
||||
- You need to run a GUI app (e.g., open/xdg-open/osascript) to open browsers or files.
|
||||
- You are running sandboxed and need to run a command that requires network access (e.g. installing packages)
|
||||
- If you run a command that is important to solving the user's query, but it fails because of sandboxing, rerun the command with approval. ALWAYS proceed to use the `sandbox_permissions` and `justification` parameters - do not message the user before requesting approval for the command.
|
||||
- You are about to take a potentially destructive action such as an `rm` or `git reset` that the user did not explicitly ask for
|
||||
- (for all of these, you should weigh alternative paths that do not require approval)
|
||||
|
||||
When `sandbox_mode` is set to read-only, you'll need to request approval for any command that isn't a read.
|
||||
|
||||
You will be told what filesystem sandboxing, network sandboxing, and approval mode are active in a developer or user message. If you are not told about this, assume that you are running with workspace-write, network sandboxing enabled, and approval on-failure.
|
||||
|
||||
Although they introduce friction to the user because your work is paused until the user responds, you should leverage them when necessary to accomplish important work. If the completing the task requires escalated permissions, Do not let these settings or the sandbox deter you from attempting to accomplish the user's task unless it is set to "never", in which case never ask for approvals.
|
||||
|
||||
When requesting approval to execute a command that will require escalated privileges:
|
||||
- Provide the `sandbox_permissions` parameter with the value `"require_escalated"`
|
||||
- Include a short, 1 sentence explanation for why you need escalated permissions in the justification parameter
|
||||
|
||||
## Special user requests
|
||||
|
||||
- If the user makes a simple request (such as asking for the time) which you can fulfill by running a terminal command (such as `date`), you should do so.
|
||||
- If the user asks for a "review", default to a code review mindset: prioritise identifying bugs, risks, behavioural regressions, and missing tests. Findings must be the primary focus of the response - keep summaries or overviews brief and only after enumerating the issues. Present findings first (ordered by severity with file/line references), follow with open questions or assumptions, and offer a change-summary only as a secondary detail. If no findings are discovered, state that explicitly and mention any residual risks or testing gaps.
|
||||
|
||||
## Presenting your work and final message
|
||||
|
||||
You are producing plain text that will later be styled by the CLI. Follow these rules exactly. Formatting should make results easy to scan, but not feel mechanical. Use judgment to decide how much structure adds value.
|
||||
|
||||
- Default: be very concise; friendly coding teammate tone.
|
||||
- Ask only when needed; suggest ideas; mirror the user's style.
|
||||
- For substantial work, summarize clearly; follow final-answer formatting.
|
||||
- Skip heavy formatting for simple confirmations.
|
||||
- Don't dump large files you've written; reference paths only.
|
||||
- No "save/copy this file" - User is on the same machine.
|
||||
- Offer logical next steps (tests, commits, build) briefly; add verify steps if you couldn't do something.
|
||||
- For code changes:
|
||||
* Lead with a quick explanation of the change, and then give more details on the context covering where and why a change was made. Do not start this explanation with "summary", just jump right in.
|
||||
* If there are natural next steps the user may want to take, suggest them at the end of your response. Do not make suggestions if there are no natural next steps.
|
||||
* When suggesting multiple options, use numeric lists for the suggestions so the user can quickly respond with a single number.
|
||||
- The user does not command execution outputs. When asked to show the output of a command (e.g. `git show`), relay the important details in your answer or summarize the key lines so the user understands the result.
|
||||
|
||||
### Final answer structure and style guidelines
|
||||
|
||||
- Plain text; CLI handles styling. Use structure only when it helps scanability.
|
||||
- Headers: optional; short Title Case (1-3 words) wrapped in **...**; no blank line before the first bullet; add only if they truly help.
|
||||
- Bullets: use - ; merge related points; keep to one line when possible; 4-6 per list ordered by importance; keep phrasing consistent.
|
||||
- Monospace: backticks for commands/paths/env vars/code ids and inline examples; use for literal keyword bullets; never combine with **.
|
||||
- Code samples or multi-line snippets should be wrapped in fenced code blocks; include an info string as often as possible.
|
||||
- Structure: group related bullets; order sections general -> specific -> supporting; for subsections, start with a bolded keyword bullet, then items; match complexity to the task.
|
||||
- Tone: collaborative, concise, factual; present tense, active voice; self-contained; no "above/below"; parallel wording.
|
||||
- Don'ts: no nested bullets/hierarchies; no ANSI codes; don't cram unrelated keywords; keep keyword lists short-wrap/reformat if long; avoid naming formatting styles in answers.
|
||||
- Adaptation: code explanations -> precise, structured with code refs; simple tasks -> lead with outcome; big changes -> logical walkthrough + rationale + next actions; casual one-offs -> plain sentences, no headers/bullets.
|
||||
- File References: When referencing files in your response, make sure to include the relevant start line and always follow the below rules:
|
||||
* Use inline code to make file paths clickable.
|
||||
* Each reference should have a stand alone path. Even if it's the same file.
|
||||
* Accepted: absolute, workspace-relative, a/ or b/ diff prefixes, or bare filename/suffix.
|
||||
* Line/column (1-based, optional): :line[:column] or #Lline[Ccolumn] (column defaults to 1).
|
||||
* Do not use URIs like file://, vscode://, or https://.
|
||||
* Do not provide range of lines
|
||||
* Examples: src/app.ts, src/app.ts:42, b/server/index.js#L10, C:\repo\project\main.rs:12:5"#;
|
||||
|
||||
// GPT-5.1 Codex Max 专用 prompt
|
||||
// 来源: CLIProxyAPI/internal/misc/codex_instructions/gpt-5.1-codex-max_prompt.md-002
|
||||
const CODEX_MAX_INSTRUCTIONS: &str = CODEX_INSTRUCTIONS;
|
||||
|
||||
// GPT-5.2 Codex 专用 prompt
|
||||
// 来源: CLIProxyAPI/internal/misc/codex_instructions/gpt-5.2-codex_prompt.md-001
|
||||
const CODEX_52_INSTRUCTIONS: &str = CODEX_INSTRUCTIONS;
|
||||
|
||||
// GPT-5.1 通用 prompt
|
||||
// 来源: CLIProxyAPI/internal/misc/codex_instructions/gpt_5_1_prompt.md-004
|
||||
const GPT_51_INSTRUCTIONS: &str = CODEX_INSTRUCTIONS;
|
||||
|
||||
// GPT-5.2 通用 prompt
|
||||
// 来源: CLIProxyAPI/internal/misc/codex_instructions/gpt_5_2_prompt.md-001
|
||||
const GPT_52_INSTRUCTIONS: &str = CODEX_INSTRUCTIONS;
|
||||
std::env::var("PROXYCAST_CODEX_DEFAULT_INSTRUCTIONS")
|
||||
.ok()
|
||||
.map(|value| value.trim().to_string())
|
||||
.filter(|value| !value.is_empty())
|
||||
}
|
||||
|
||||
/// Transform OpenAI chat completion request to Codex format
|
||||
/// 参考 CLIProxyAPI: internal/translator/codex/openai/chat-completions/codex_openai_request.go
|
||||
@@ -1390,6 +1293,7 @@ fn transform_to_codex_format(
|
||||
|
||||
// Build input array from messages
|
||||
let mut input = Vec::new();
|
||||
let mut system_instructions: Vec<String> = Vec::new();
|
||||
|
||||
if let Some(msgs) = messages {
|
||||
for msg in msgs {
|
||||
@@ -1398,16 +1302,8 @@ fn transform_to_codex_format(
|
||||
|
||||
match role {
|
||||
"system" => {
|
||||
// System messages 转换为 user message(Codex 使用 instructions 而不是 system role)
|
||||
if let Some(text) = content.as_str() {
|
||||
if !text.is_empty() {
|
||||
input.push(serde_json::json!({
|
||||
"type": "message",
|
||||
"role": "user",
|
||||
"content": [{"type": "input_text", "text": text}]
|
||||
}));
|
||||
}
|
||||
}
|
||||
// system 消息统一转换到 instructions,避免污染用户输入。
|
||||
system_instructions.extend(extract_text_fragments(content));
|
||||
}
|
||||
"user" => {
|
||||
let content_parts = if let Some(text) = content.as_str() {
|
||||
@@ -1510,10 +1406,9 @@ fn transform_to_codex_format(
|
||||
"include": ["reasoning.encrypted_content"]
|
||||
});
|
||||
|
||||
// 根据模型名称选择正确的 instructions
|
||||
// 参考 CLIProxyAPI: internal/misc/codex_instructions.go
|
||||
let instructions = get_codex_instructions_for_model(model);
|
||||
codex_request["instructions"] = serde_json::json!(instructions);
|
||||
if let Some(instructions) = resolve_codex_instructions(request, &system_instructions) {
|
||||
codex_request["instructions"] = serde_json::json!(instructions);
|
||||
}
|
||||
|
||||
// 处理可选参数:temperature, max_tokens (-> max_output_tokens), top_p
|
||||
if let Some(temp) = request.get("temperature") {
|
||||
@@ -1908,17 +1803,45 @@ mod tests {
|
||||
|
||||
assert_eq!(result["model"], "gpt-4o");
|
||||
assert_eq!(result["stream"], true);
|
||||
// instructions 字段存在,使用 Codex 默认 prompt
|
||||
// system 消息应映射到 instructions
|
||||
assert!(result.get("instructions").is_some());
|
||||
// 验证 instructions 以正确的前缀开始
|
||||
let instructions = result["instructions"].as_str().unwrap();
|
||||
assert!(instructions.starts_with("You are Codex, based on GPT-5."));
|
||||
assert_eq!(instructions, "You are a helpful assistant.");
|
||||
|
||||
let input = result["input"].as_array().unwrap();
|
||||
// system message 被转换为 user message,所以有 2 条消息
|
||||
assert_eq!(input.len(), 2);
|
||||
assert_eq!(input[0]["role"], "user"); // system -> user
|
||||
assert_eq!(input[1]["role"], "user"); // original user
|
||||
// system message 不应污染输入,保留 user 消息即可
|
||||
assert_eq!(input.len(), 1);
|
||||
assert_eq!(input[0]["role"], "user");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_transform_to_codex_format_without_system_does_not_inject_instructions() {
|
||||
let request = serde_json::json!({
|
||||
"model": "gpt-4o",
|
||||
"messages": [
|
||||
{"role": "user", "content": "Hello!"}
|
||||
]
|
||||
});
|
||||
|
||||
let result = transform_to_codex_format(&request).unwrap();
|
||||
assert!(result.get("instructions").is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_transform_to_codex_format_uses_explicit_instructions() {
|
||||
let request = serde_json::json!({
|
||||
"model": "gpt-4o",
|
||||
"instructions": "You are a general assistant.",
|
||||
"messages": [
|
||||
{"role": "user", "content": "Hello!"}
|
||||
]
|
||||
});
|
||||
|
||||
let result = transform_to_codex_format(&request).unwrap();
|
||||
assert_eq!(
|
||||
result["instructions"].as_str(),
|
||||
Some("You are a general assistant.")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
||||
@@ -44,6 +44,38 @@ pub struct ConnectionTestResult {
|
||||
mod tests {
|
||||
use super::ApiKeyProviderService;
|
||||
use proxycast_core::database::dao::api_key_provider::ApiProviderType;
|
||||
use proxycast_core::database::init_database;
|
||||
use rusqlite::OptionalExtension;
|
||||
|
||||
fn resolve_real_codex_provider_id(
|
||||
db: &proxycast_core::database::DbConnection,
|
||||
) -> Result<String, String> {
|
||||
if let Ok(explicit) = std::env::var("PROXYCAST_REAL_PROVIDER_ID") {
|
||||
let trimmed = explicit.trim();
|
||||
if !trimmed.is_empty() {
|
||||
return Ok(trimmed.to_string());
|
||||
}
|
||||
}
|
||||
|
||||
let conn = db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?;
|
||||
conn.query_row(
|
||||
r#"
|
||||
SELECT p.id
|
||||
FROM api_key_providers p
|
||||
JOIN api_keys k ON k.provider_id = p.id
|
||||
WHERE p.enabled = 1
|
||||
AND k.enabled = 1
|
||||
AND p.type = 'codex'
|
||||
ORDER BY p.updated_at DESC
|
||||
LIMIT 1
|
||||
"#,
|
||||
[],
|
||||
|row| row.get::<_, String>(0),
|
||||
)
|
||||
.optional()
|
||||
.map_err(|e| format!("查询 Codex Provider 失败: {e}"))?
|
||||
.ok_or_else(|| "未找到可用的 Codex Provider,请先在设置中配置并启用".to_string())
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_build_codex_responses_request_input_list() {
|
||||
@@ -101,6 +133,38 @@ data: [DONE]\n";
|
||||
let none = ApiKeyProviderService::pick_test_model(None, &[], &[]);
|
||||
assert!(none.is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
#[ignore = "真实联网测试:设置 PROXYCAST_REAL_API_TEST=1 后执行"]
|
||||
async fn test_real_codex_provider_chat_gpt_5_3_codex() {
|
||||
if std::env::var("PROXYCAST_REAL_API_TEST").ok().as_deref() != Some("1") {
|
||||
return;
|
||||
}
|
||||
|
||||
let db = init_database().expect("初始化数据库失败");
|
||||
let service = ApiKeyProviderService::new();
|
||||
let provider_id = resolve_real_codex_provider_id(&db).expect("解析 Codex Provider 失败");
|
||||
let model =
|
||||
std::env::var("PROXYCAST_REAL_MODEL").unwrap_or_else(|_| "gpt-5.3-codex".to_string());
|
||||
let prompt = std::env::var("PROXYCAST_REAL_PROMPT")
|
||||
.unwrap_or_else(|_| "请仅回复 REAL_OK".to_string());
|
||||
|
||||
let result = service
|
||||
.test_chat(&db, &provider_id, Some(model.clone()), prompt)
|
||||
.await
|
||||
.expect("真实调用失败");
|
||||
|
||||
assert!(
|
||||
result.success,
|
||||
"真实调用未成功: provider_id={provider_id}, model={model}, error={:?}, raw={:?}",
|
||||
result.error, result.raw
|
||||
);
|
||||
assert!(
|
||||
result.error.is_none(),
|
||||
"真实调用返回错误: {:?}",
|
||||
result.error
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
|
||||
@@ -86,35 +86,136 @@ fn should_create_backup() -> bool {
|
||||
true
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
#[cfg_attr(not(target_os = "windows"), allow(dead_code))]
|
||||
enum ShellConfigSyntax {
|
||||
Posix,
|
||||
PowerShell,
|
||||
}
|
||||
|
||||
/// 获取当前 shell 配置文件路径
|
||||
/// 优先级:zsh > bash
|
||||
fn get_shell_config_path() -> Result<PathBuf, Box<dyn std::error::Error + Send + Sync>> {
|
||||
fn get_shell_config_target(
|
||||
) -> Result<(PathBuf, ShellConfigSyntax), Box<dyn std::error::Error + Send + Sync>> {
|
||||
let home = dirs::home_dir().ok_or("Cannot find home directory")?;
|
||||
|
||||
// 检查 SHELL 环境变量
|
||||
if let Ok(shell) = std::env::var("SHELL") {
|
||||
if shell.contains("zsh") {
|
||||
let zshrc = home.join(".zshrc");
|
||||
return Ok(zshrc);
|
||||
} else if shell.contains("bash") {
|
||||
let bashrc = home.join(".bashrc");
|
||||
return Ok(bashrc);
|
||||
#[cfg(target_os = "windows")]
|
||||
{
|
||||
let documents = dirs::document_dir().unwrap_or_else(|| home.join("Documents"));
|
||||
let ps7_profile = documents
|
||||
.join("PowerShell")
|
||||
.join("Microsoft.PowerShell_profile.ps1");
|
||||
let winps_profile = documents
|
||||
.join("WindowsPowerShell")
|
||||
.join("Microsoft.PowerShell_profile.ps1");
|
||||
|
||||
if ps7_profile.exists() {
|
||||
return Ok((ps7_profile, ShellConfigSyntax::PowerShell));
|
||||
}
|
||||
if winps_profile.exists() {
|
||||
return Ok((winps_profile, ShellConfigSyntax::PowerShell));
|
||||
}
|
||||
|
||||
return Ok((ps7_profile, ShellConfigSyntax::PowerShell));
|
||||
}
|
||||
|
||||
// 默认检查文件是否存在
|
||||
let zshrc = home.join(".zshrc");
|
||||
if zshrc.exists() {
|
||||
return Ok(zshrc);
|
||||
#[cfg(not(target_os = "windows"))]
|
||||
{
|
||||
// 检查 SHELL 环境变量
|
||||
if let Ok(shell) = std::env::var("SHELL") {
|
||||
if shell.contains("zsh") {
|
||||
let zshrc = home.join(".zshrc");
|
||||
return Ok((zshrc, ShellConfigSyntax::Posix));
|
||||
} else if shell.contains("bash") {
|
||||
let bashrc = home.join(".bashrc");
|
||||
return Ok((bashrc, ShellConfigSyntax::Posix));
|
||||
}
|
||||
}
|
||||
|
||||
// 默认检查文件是否存在
|
||||
let zshrc = home.join(".zshrc");
|
||||
if zshrc.exists() {
|
||||
return Ok((zshrc, ShellConfigSyntax::Posix));
|
||||
}
|
||||
|
||||
let bashrc = home.join(".bashrc");
|
||||
if bashrc.exists() {
|
||||
return Ok((bashrc, ShellConfigSyntax::Posix));
|
||||
}
|
||||
|
||||
// 如果都不存在,默认使用 .zshrc(macOS 默认)
|
||||
Ok((zshrc, ShellConfigSyntax::Posix))
|
||||
}
|
||||
}
|
||||
|
||||
fn get_shell_config_path() -> Result<PathBuf, Box<dyn std::error::Error + Send + Sync>> {
|
||||
Ok(get_shell_config_target()?.0)
|
||||
}
|
||||
|
||||
fn escape_shell_env_value(value: &str, syntax: ShellConfigSyntax) -> String {
|
||||
match syntax {
|
||||
ShellConfigSyntax::Posix => value.replace('\\', "\\\\").replace('"', "\\\""),
|
||||
ShellConfigSyntax::PowerShell => value.replace('`', "``").replace('"', "`\""),
|
||||
}
|
||||
}
|
||||
|
||||
fn format_shell_env_line(key: &str, value: &str, syntax: ShellConfigSyntax) -> String {
|
||||
let escaped_value = escape_shell_env_value(value, syntax);
|
||||
match syntax {
|
||||
ShellConfigSyntax::Posix => format!("export {key}=\"{escaped_value}\""),
|
||||
ShellConfigSyntax::PowerShell => format!("$env:{key} = \"{escaped_value}\""),
|
||||
}
|
||||
}
|
||||
|
||||
fn parse_key_value(expr: &str) -> Option<(String, String)> {
|
||||
let eq_pos = expr.find('=')?;
|
||||
let key = expr[..eq_pos].trim();
|
||||
if key.is_empty() {
|
||||
return None;
|
||||
}
|
||||
let value = expr[eq_pos + 1..].trim().to_string();
|
||||
Some((key.to_string(), value))
|
||||
}
|
||||
|
||||
fn unquote_shell_value(raw: &str) -> String {
|
||||
let value = if (raw.starts_with('"') && raw.ends_with('"'))
|
||||
|| (raw.starts_with('\'') && raw.ends_with('\''))
|
||||
{
|
||||
raw[1..raw.len() - 1].to_string()
|
||||
} else {
|
||||
raw.to_string()
|
||||
};
|
||||
value.replace("\\\"", "\"").replace("\\\\", "\\")
|
||||
}
|
||||
|
||||
fn unquote_powershell_value(raw: &str) -> String {
|
||||
let value = if (raw.starts_with('"') && raw.ends_with('"'))
|
||||
|| (raw.starts_with('\'') && raw.ends_with('\''))
|
||||
{
|
||||
raw[1..raw.len() - 1].to_string()
|
||||
} else {
|
||||
raw.to_string()
|
||||
};
|
||||
value.replace("`\"", "\"").replace("``", "`")
|
||||
}
|
||||
|
||||
fn parse_shell_env_line(line: &str) -> Option<(String, String)> {
|
||||
let trimmed = line.trim();
|
||||
|
||||
if let Some(export_line) = trimmed.strip_prefix("export ") {
|
||||
let (key, value) = parse_key_value(export_line)?;
|
||||
return Some((key, unquote_shell_value(&value)));
|
||||
}
|
||||
|
||||
let bashrc = home.join(".bashrc");
|
||||
if bashrc.exists() {
|
||||
return Ok(bashrc);
|
||||
if let Some(env_line) = trimmed
|
||||
.strip_prefix("$env:")
|
||||
.or_else(|| trimmed.strip_prefix("$Env:"))
|
||||
{
|
||||
let (key, value) = parse_key_value(env_line)?;
|
||||
return Some((key, unquote_powershell_value(&value)));
|
||||
}
|
||||
|
||||
// 如果都不存在,默认使用 .zshrc(macOS 默认)
|
||||
Ok(zshrc)
|
||||
None
|
||||
}
|
||||
|
||||
/// 将环境变量写入 shell 配置文件
|
||||
@@ -124,13 +225,17 @@ fn get_shell_config_path() -> Result<PathBuf, Box<dyn std::error::Error + Send +
|
||||
pub fn write_env_to_shell_config(
|
||||
env_vars: &[(String, String)],
|
||||
) -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
|
||||
let config_path = get_shell_config_path()?;
|
||||
let (config_path, syntax) = get_shell_config_target()?;
|
||||
|
||||
tracing::info!(
|
||||
"Writing environment variables to: {}",
|
||||
config_path.display()
|
||||
);
|
||||
|
||||
if let Some(parent) = config_path.parent() {
|
||||
fs::create_dir_all(parent)?;
|
||||
}
|
||||
|
||||
// 读取现有配置
|
||||
let existing_content = if config_path.exists() {
|
||||
fs::read_to_string(&config_path)?
|
||||
@@ -170,9 +275,9 @@ pub fn write_env_to_shell_config(
|
||||
new_content.push_str("# Do not edit this block manually\n");
|
||||
|
||||
for (key, value) in env_vars {
|
||||
// 转义值中的特殊字符
|
||||
let escaped_value = value.replace('\\', "\\\\").replace('"', "\\\"");
|
||||
new_content.push_str(&format!("export {key}=\"{escaped_value}\"\n"));
|
||||
let line = format_shell_env_line(key, value, syntax);
|
||||
new_content.push_str(&line);
|
||||
new_content.push('\n');
|
||||
}
|
||||
|
||||
new_content.push_str(ENV_BLOCK_END);
|
||||
@@ -222,22 +327,8 @@ fn read_env_from_shell_config(
|
||||
continue;
|
||||
}
|
||||
|
||||
if in_proxycast_block && trimmed.starts_with("export ") {
|
||||
// 解析 export KEY="VALUE" 格式
|
||||
let export_line = trimmed.strip_prefix("export ").unwrap_or(trimmed);
|
||||
if let Some(eq_pos) = export_line.find('=') {
|
||||
let key = export_line[..eq_pos].trim().to_string();
|
||||
let value_part = export_line[eq_pos + 1..].trim();
|
||||
|
||||
// 移除引号
|
||||
let value = if (value_part.starts_with('"') && value_part.ends_with('"'))
|
||||
|| (value_part.starts_with('\'') && value_part.ends_with('\''))
|
||||
{
|
||||
value_part[1..value_part.len() - 1].to_string()
|
||||
} else {
|
||||
value_part.to_string()
|
||||
};
|
||||
|
||||
if in_proxycast_block {
|
||||
if let Some((key, value)) = parse_shell_env_line(trimmed) {
|
||||
env_vars.push((key, value));
|
||||
}
|
||||
}
|
||||
@@ -758,7 +849,16 @@ pub fn read_live_settings_for_display(
|
||||
// 获取 shell 配置文件路径
|
||||
let shell_config_path = get_shell_config_path()
|
||||
.map(|p| p.display().to_string())
|
||||
.unwrap_or_else(|_| "~/.zshrc or ~/.bashrc".to_string());
|
||||
.unwrap_or_else(|_| {
|
||||
#[cfg(target_os = "windows")]
|
||||
{
|
||||
"Documents/PowerShell/Microsoft.PowerShell_profile.ps1".to_string()
|
||||
}
|
||||
#[cfg(not(target_os = "windows"))]
|
||||
{
|
||||
"~/.zshrc or ~/.bashrc".to_string()
|
||||
}
|
||||
});
|
||||
|
||||
// 返回包含两部分的结构
|
||||
Ok(json!({
|
||||
|
||||
@@ -288,6 +288,7 @@ mod tests {
|
||||
|
||||
#[cfg(test)]
|
||||
mod shell_config_write_tests {
|
||||
use super::*;
|
||||
|
||||
/// **Feature: shell-write, Property 1: 特殊字符转义**
|
||||
#[test]
|
||||
@@ -314,13 +315,39 @@ mod tests {
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_parse_shell_env_line_supports_posix_export() {
|
||||
let parsed = parse_shell_env_line(r#"export OPENAI_API_KEY="abc123""#)
|
||||
.expect("Should parse posix export line");
|
||||
assert_eq!(parsed.0, "OPENAI_API_KEY");
|
||||
assert_eq!(parsed.1, "abc123");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_parse_shell_env_line_supports_powershell_env() {
|
||||
let parsed = parse_shell_env_line(r#"$env:OPENAI_BASE_URL = "https://example.com""#)
|
||||
.expect("Should parse PowerShell env line");
|
||||
assert_eq!(parsed.0, "OPENAI_BASE_URL");
|
||||
assert_eq!(parsed.1, "https://example.com");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_format_shell_env_line_powershell_style() {
|
||||
let line = format_shell_env_line(
|
||||
"TEST_KEY",
|
||||
r#"value with "quotes""#,
|
||||
ShellConfigSyntax::PowerShell,
|
||||
);
|
||||
assert_eq!(line, "$env:TEST_KEY = \"value with `\"quotes`\"\"");
|
||||
}
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// 总结
|
||||
// ============================================================================
|
||||
//
|
||||
// 本测试模块包含 3 个子模块,共 10 个单元测试:
|
||||
// 本测试模块包含 3 个子模块,共 13 个单元测试:
|
||||
//
|
||||
// 1. **原子写入测试** (3 个测试)
|
||||
// - 正常写入、备份创建、JSON 往返
|
||||
@@ -328,8 +355,11 @@ mod tests {
|
||||
// 2. **认证冲突清理测试** (5 个测试)
|
||||
// - 单独 TOKEN、单独 KEY、冲突处理、都为空、空值处理
|
||||
//
|
||||
// 3. **Shell 配置写入测试** (1 个测试)
|
||||
// 3. **Shell 配置写入测试** (4 个测试)
|
||||
// - 特殊字符转义验证
|
||||
// - POSIX export 解析
|
||||
// - PowerShell 环境变量解析
|
||||
// - PowerShell 写入格式验证
|
||||
//
|
||||
// **注意**:由于 `sync_claude_settings`、`write_env_to_shell_config` 等函数
|
||||
// 依赖于真实的文件系统路径(如 ~/.claude、~/.zshrc),完整的集成测试
|
||||
|
||||
@@ -167,6 +167,13 @@ impl BlockMeta {
|
||||
_ => String::new(),
|
||||
}
|
||||
}
|
||||
|
||||
/// 获取清洗后的工作目录(去除 `\0`、首尾空白)
|
||||
pub fn sanitized_cmd_cwd(&self) -> Option<String> {
|
||||
let cwd = self.cmd_cwd.as_deref()?;
|
||||
let cleaned = cwd.split('\0').next().unwrap_or_default().trim();
|
||||
(!cleaned.is_empty()).then_some(cleaned.to_string())
|
||||
}
|
||||
}
|
||||
|
||||
/// 运行时选项
|
||||
@@ -294,4 +301,19 @@ mod tests {
|
||||
assert_eq!(meta.get_string("cmd"), "");
|
||||
assert_eq!(meta.get_string("term_mode"), "term");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_block_meta_sanitized_cmd_cwd() {
|
||||
let meta = BlockMeta {
|
||||
cmd_cwd: Some(" /tmp/demo\0ignored ".to_string()),
|
||||
..Default::default()
|
||||
};
|
||||
assert_eq!(meta.sanitized_cmd_cwd(), Some("/tmp/demo".to_string()));
|
||||
|
||||
let empty = BlockMeta {
|
||||
cmd_cwd: Some(" \0 ".to_string()),
|
||||
..Default::default()
|
||||
};
|
||||
assert!(empty.sanitized_cmd_cwd().is_none());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -21,6 +21,7 @@
|
||||
//! - 17.10: fish 使用 -C 参数 source 集成脚本
|
||||
|
||||
use std::io::{Read, Write};
|
||||
use std::path::{Path, PathBuf};
|
||||
use std::sync::atomic::{AtomicBool, AtomicI32, Ordering};
|
||||
use std::sync::Arc;
|
||||
|
||||
@@ -38,6 +39,102 @@ use crate::events::{event_names, SessionStatus, TerminalOutputEvent, TerminalSta
|
||||
use crate::integration::{ShellLaunchBuilder, ShellType};
|
||||
use crate::persistence::BlockFile;
|
||||
|
||||
fn resolve_default_shell() -> String {
|
||||
let shell_from_env = std::env::var("SHELL").ok().and_then(|value| {
|
||||
let cleaned = value
|
||||
.split('\0')
|
||||
.next()
|
||||
.unwrap_or_default()
|
||||
.trim()
|
||||
.to_string();
|
||||
(!cleaned.is_empty()).then_some(cleaned)
|
||||
});
|
||||
|
||||
#[cfg(target_os = "windows")]
|
||||
{
|
||||
if let Some(shell) = shell_from_env {
|
||||
let path = Path::new(&shell);
|
||||
if path.is_absolute() && path.exists() {
|
||||
return shell;
|
||||
}
|
||||
}
|
||||
|
||||
if let Ok(comspec) = std::env::var("COMSPEC") {
|
||||
let cleaned = comspec
|
||||
.split('\0')
|
||||
.next()
|
||||
.unwrap_or_default()
|
||||
.trim()
|
||||
.to_string();
|
||||
if !cleaned.is_empty() && Path::new(&cleaned).exists() {
|
||||
return cleaned;
|
||||
}
|
||||
}
|
||||
|
||||
"cmd.exe".to_string()
|
||||
}
|
||||
|
||||
#[cfg(not(target_os = "windows"))]
|
||||
{
|
||||
if let Some(shell) = shell_from_env {
|
||||
let path = Path::new(&shell);
|
||||
if path.exists() {
|
||||
return shell;
|
||||
}
|
||||
}
|
||||
|
||||
if Path::new("/bin/bash").exists() {
|
||||
"/bin/bash".to_string()
|
||||
} else {
|
||||
"/bin/sh".to_string()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn resolve_working_dir(cwd: Option<&str>) -> Option<PathBuf> {
|
||||
let dir = cwd?;
|
||||
let cleaned = dir
|
||||
.split('\0')
|
||||
.next()
|
||||
.unwrap_or_default()
|
||||
.trim()
|
||||
.to_string();
|
||||
if cleaned.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let expanded = if cleaned.starts_with("~/") {
|
||||
if let Some(home) = dirs::home_dir() {
|
||||
home.join(&cleaned[2..])
|
||||
} else {
|
||||
PathBuf::from(&cleaned)
|
||||
}
|
||||
} else if cleaned == "~" {
|
||||
dirs::home_dir().unwrap_or_else(|| PathBuf::from(&cleaned))
|
||||
} else {
|
||||
PathBuf::from(&cleaned)
|
||||
};
|
||||
|
||||
(expanded.exists() && expanded.is_dir()).then_some(expanded)
|
||||
}
|
||||
|
||||
fn shell_command_flag(shell: &str) -> &'static str {
|
||||
let executable = Path::new(shell)
|
||||
.file_name()
|
||||
.and_then(|name| name.to_str())
|
||||
.unwrap_or(shell)
|
||||
.to_ascii_lowercase();
|
||||
|
||||
if executable == "cmd" || executable == "cmd.exe" {
|
||||
"/C"
|
||||
} else if executable.contains("powershell") || executable == "pwsh" || executable == "pwsh.exe"
|
||||
{
|
||||
"-Command"
|
||||
} else {
|
||||
"-c"
|
||||
}
|
||||
}
|
||||
|
||||
/// Shell 进程封装
|
||||
///
|
||||
/// 封装 PTY 进程,提供输入输出和生命周期管理。
|
||||
@@ -199,8 +296,17 @@ impl<E: TerminalEventEmitter> ShellProc<E> {
|
||||
};
|
||||
|
||||
// 设置工作目录
|
||||
if let Some(cwd) = &block_meta.cmd_cwd {
|
||||
cmd.cwd(cwd);
|
||||
let sanitized_cwd = block_meta.sanitized_cmd_cwd();
|
||||
if let Some(resolved_cwd) = resolve_working_dir(sanitized_cwd.as_deref()) {
|
||||
cmd.cwd(resolved_cwd);
|
||||
} else if let Some(raw_cwd) = block_meta.cmd_cwd.as_deref() {
|
||||
tracing::warn!(
|
||||
"[ShellProc] 工作目录无效或不存在: {:?}, 使用主目录",
|
||||
raw_cwd
|
||||
);
|
||||
if let Some(home) = dirs::home_dir() {
|
||||
cmd.cwd(home);
|
||||
}
|
||||
} else if let Some(home) = dirs::home_dir() {
|
||||
cmd.cwd(home);
|
||||
}
|
||||
@@ -219,7 +325,7 @@ impl<E: TerminalEventEmitter> ShellProc<E> {
|
||||
block_id: &str,
|
||||
) -> Result<CommandBuilder, TerminalError> {
|
||||
// 获取用户默认 shell
|
||||
let shell = std::env::var("SHELL").unwrap_or_else(|_| "/bin/bash".to_string());
|
||||
let shell = resolve_default_shell();
|
||||
tracing::info!("[ShellProc] 使用 shell: {}", shell);
|
||||
|
||||
// 获取应用数据目录
|
||||
@@ -267,9 +373,9 @@ impl<E: TerminalEventEmitter> ShellProc<E> {
|
||||
tracing::info!("[ShellProc] 执行命令: {}", cmd_str);
|
||||
|
||||
// 使用 shell 执行命令
|
||||
let shell = std::env::var("SHELL").unwrap_or_else(|_| "/bin/bash".to_string());
|
||||
let shell = resolve_default_shell();
|
||||
let mut cmd = CommandBuilder::new(&shell);
|
||||
cmd.arg("-c");
|
||||
cmd.arg(shell_command_flag(&shell));
|
||||
|
||||
// 构建完整命令字符串
|
||||
let full_cmd = if let Some(args) = &block_meta.cmd_args {
|
||||
@@ -556,3 +662,37 @@ impl<E: TerminalEventEmitter> Drop for ShellProc<E> {
|
||||
tracing::debug!("[ShellProc] 进程已销毁: block_id={}", self.block_id);
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{resolve_default_shell, resolve_working_dir, shell_command_flag};
|
||||
|
||||
#[test]
|
||||
fn resolve_working_dir_should_strip_nul_suffix() {
|
||||
let temp_dir = tempfile::tempdir().expect("create temp dir");
|
||||
let raw = format!("{}\0", temp_dir.path().to_string_lossy());
|
||||
|
||||
let resolved = resolve_working_dir(Some(&raw));
|
||||
|
||||
assert_eq!(resolved.as_deref(), Some(temp_dir.path()));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolve_working_dir_should_reject_invalid_path() {
|
||||
let resolved = resolve_working_dir(Some("/path/not-exists\0"));
|
||||
assert!(resolved.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolve_default_shell_should_not_be_empty() {
|
||||
let shell = resolve_default_shell();
|
||||
assert!(!shell.trim().is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn shell_command_flag_should_match_common_shells() {
|
||||
assert_eq!(shell_command_flag("cmd.exe"), "/C");
|
||||
assert_eq!(shell_command_flag("pwsh"), "-Command");
|
||||
assert_eq!(shell_command_flag("/bin/bash"), "-c");
|
||||
}
|
||||
}
|
||||
|
||||
@@ -227,8 +227,8 @@ impl SSHShellProc {
|
||||
let mut full_cmd = String::new();
|
||||
|
||||
// 如果指定了工作目录,先 cd 到该目录
|
||||
if let Some(cwd) = &block_meta.cmd_cwd {
|
||||
full_cmd.push_str(&format!("cd {} && ", shell_escape(cwd)));
|
||||
if let Some(cwd) = block_meta.sanitized_cmd_cwd() {
|
||||
full_cmd.push_str(&format!("cd {} && ", shell_escape(&cwd)));
|
||||
}
|
||||
|
||||
// 设置环境变量
|
||||
|
||||
@@ -757,7 +757,7 @@ impl WSLShellProc {
|
||||
if let Some(ref path) = opts.initial_path {
|
||||
cmd.arg("--cd");
|
||||
cmd.arg(path);
|
||||
} else if let Some(ref cwd) = block_meta.cmd_cwd {
|
||||
} else if let Some(cwd) = block_meta.sanitized_cmd_cwd() {
|
||||
cmd.arg("--cd");
|
||||
cmd.arg(cwd);
|
||||
}
|
||||
|
||||
@@ -14,6 +14,7 @@
|
||||
//! 输出历史保存在循环缓冲区中,前端连接时可以获取历史数据。
|
||||
|
||||
use std::io::{Read, Write};
|
||||
use std::path::{Path, PathBuf};
|
||||
use std::sync::atomic::{AtomicBool, Ordering};
|
||||
use std::sync::Arc;
|
||||
|
||||
@@ -89,6 +90,85 @@ pub struct PtySession {
|
||||
}
|
||||
|
||||
impl PtySession {
|
||||
fn resolve_default_shell() -> String {
|
||||
let shell_from_env = std::env::var("SHELL").ok().and_then(|value| {
|
||||
let cleaned = value
|
||||
.split('\0')
|
||||
.next()
|
||||
.unwrap_or_default()
|
||||
.trim()
|
||||
.to_string();
|
||||
(!cleaned.is_empty()).then_some(cleaned)
|
||||
});
|
||||
|
||||
#[cfg(target_os = "windows")]
|
||||
{
|
||||
if let Some(shell) = shell_from_env {
|
||||
let path = Path::new(&shell);
|
||||
if path.is_absolute() && path.exists() {
|
||||
return shell;
|
||||
}
|
||||
}
|
||||
|
||||
if let Ok(comspec) = std::env::var("COMSPEC") {
|
||||
let cleaned = comspec
|
||||
.split('\0')
|
||||
.next()
|
||||
.unwrap_or_default()
|
||||
.trim()
|
||||
.to_string();
|
||||
if !cleaned.is_empty() && Path::new(&cleaned).exists() {
|
||||
return cleaned;
|
||||
}
|
||||
}
|
||||
|
||||
"cmd.exe".to_string()
|
||||
}
|
||||
|
||||
#[cfg(not(target_os = "windows"))]
|
||||
{
|
||||
if let Some(shell) = shell_from_env {
|
||||
let path = Path::new(&shell);
|
||||
if path.exists() {
|
||||
return shell;
|
||||
}
|
||||
}
|
||||
|
||||
if Path::new("/bin/bash").exists() {
|
||||
"/bin/bash".to_string()
|
||||
} else {
|
||||
"/bin/sh".to_string()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn resolve_working_dir(cwd: Option<String>) -> Option<PathBuf> {
|
||||
let dir = cwd?;
|
||||
let cleaned = dir
|
||||
.split('\0')
|
||||
.next()
|
||||
.unwrap_or_default()
|
||||
.trim()
|
||||
.to_string();
|
||||
if cleaned.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let expanded = if cleaned.starts_with("~/") {
|
||||
if let Some(home) = dirs::home_dir() {
|
||||
home.join(&cleaned[2..])
|
||||
} else {
|
||||
PathBuf::from(&cleaned)
|
||||
}
|
||||
} else if cleaned == "~" {
|
||||
dirs::home_dir().unwrap_or_else(|| PathBuf::from(&cleaned))
|
||||
} else {
|
||||
PathBuf::from(&cleaned)
|
||||
};
|
||||
|
||||
(expanded.exists() && expanded.is_dir()).then_some(expanded)
|
||||
}
|
||||
|
||||
/// 创建新的 PTY 会话(使用默认大小)
|
||||
///
|
||||
/// PTY 使用默认大小 (24x80) 预创建,前端连接后通过 resize 同步实际大小。
|
||||
@@ -163,8 +243,8 @@ impl PtySession {
|
||||
})
|
||||
.map_err(|e| TerminalError::PtyCreationFailed(e.to_string()))?;
|
||||
|
||||
// 获取用户默认 shell
|
||||
let shell = std::env::var("SHELL").unwrap_or_else(|_| "/bin/bash".to_string());
|
||||
// 获取用户默认 shell(Windows 优先使用 COMSPEC/cmd.exe)
|
||||
let shell = Self::resolve_default_shell();
|
||||
tracing::info!("[终端] 使用 shell: {}", shell);
|
||||
|
||||
// 构建命令
|
||||
@@ -172,31 +252,13 @@ impl PtySession {
|
||||
cmd.env("TERM", "xterm-256color");
|
||||
|
||||
// 设置工作目录
|
||||
if let Some(dir) = cwd {
|
||||
// 展开 ~ 为用户主目录
|
||||
let expanded_dir = if dir.starts_with("~/") {
|
||||
if let Some(home) = dirs::home_dir() {
|
||||
home.join(&dir[2..])
|
||||
} else {
|
||||
std::path::PathBuf::from(&dir)
|
||||
}
|
||||
} else if dir == "~" {
|
||||
dirs::home_dir().unwrap_or_else(|| std::path::PathBuf::from(&dir))
|
||||
} else {
|
||||
std::path::PathBuf::from(&dir)
|
||||
};
|
||||
|
||||
if expanded_dir.exists() && expanded_dir.is_dir() {
|
||||
tracing::info!("[终端] 设置工作目录: {:?}", expanded_dir);
|
||||
cmd.cwd(expanded_dir);
|
||||
} else {
|
||||
tracing::warn!(
|
||||
"[终端] 工作目录不存在或不是目录: {:?}, 使用主目录",
|
||||
expanded_dir
|
||||
);
|
||||
if let Some(home) = dirs::home_dir() {
|
||||
cmd.cwd(home);
|
||||
}
|
||||
if let Some(expanded_dir) = Self::resolve_working_dir(cwd.clone()) {
|
||||
tracing::info!("[终端] 设置工作目录: {:?}", expanded_dir);
|
||||
cmd.cwd(expanded_dir);
|
||||
} else if let Some(raw) = cwd {
|
||||
tracing::warn!("[终端] 工作目录无效或不存在: {:?}, 使用主目录", raw);
|
||||
if let Some(home) = dirs::home_dir() {
|
||||
cmd.cwd(home);
|
||||
}
|
||||
} else if let Some(home) = dirs::home_dir() {
|
||||
cmd.cwd(home);
|
||||
@@ -380,3 +442,30 @@ impl PtySession {
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::PtySession;
|
||||
|
||||
#[test]
|
||||
fn resolve_working_dir_should_strip_nul_suffix() {
|
||||
let temp_dir = tempfile::tempdir().expect("create temp dir");
|
||||
let raw = format!("{}\0", temp_dir.path().to_string_lossy());
|
||||
|
||||
let resolved = PtySession::resolve_working_dir(Some(raw));
|
||||
|
||||
assert_eq!(resolved.as_deref(), Some(temp_dir.path()));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolve_working_dir_should_reject_invalid_path() {
|
||||
let resolved = PtySession::resolve_working_dir(Some("/path/not-exists\0".to_string()));
|
||||
assert!(resolved.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolve_default_shell_should_not_be_empty() {
|
||||
let shell = PtySession::resolve_default_shell();
|
||||
assert!(!shell.trim().is_empty());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -19,6 +19,12 @@ use crate::mcp::{McpManagerState, McpServerConfig};
|
||||
use crate::services::execution_tracker_service::{ExecutionTracker, RunFinalizeOptions, RunSource};
|
||||
use crate::services::heartbeat_service::HeartbeatServiceState;
|
||||
use crate::services::memory_profile_prompt_service::merge_system_prompt_with_memory_profile;
|
||||
#[cfg(test)]
|
||||
use crate::services::request_tool_policy_prompt_service::REQUEST_TOOL_POLICY_MARKER;
|
||||
use crate::services::request_tool_policy_prompt_service::{
|
||||
execute_web_search_preflight_if_needed, merge_system_prompt_with_request_tool_policy,
|
||||
resolve_request_tool_policy, RequestToolPolicy, WebSearchExecutionTracker,
|
||||
};
|
||||
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;
|
||||
@@ -46,7 +52,7 @@ use proxycast_agent::event_converter::convert_agent_event;
|
||||
use proxycast_services::mcp_service::McpService;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::collections::HashMap;
|
||||
use std::path::PathBuf;
|
||||
use std::path::{Path, PathBuf};
|
||||
use std::sync::{Arc, OnceLock};
|
||||
use std::time::Duration;
|
||||
use tauri::{AppHandle, Emitter, State};
|
||||
@@ -158,6 +164,8 @@ pub struct AsterAgentStatus {
|
||||
/// Provider 配置请求
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub struct ConfigureProviderRequest {
|
||||
#[serde(default)]
|
||||
pub provider_id: Option<String>,
|
||||
pub provider_name: String,
|
||||
pub model_name: String,
|
||||
#[serde(default)]
|
||||
@@ -310,21 +318,27 @@ pub async fn aster_agent_reset(
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub struct AsterChatRequest {
|
||||
pub message: String,
|
||||
#[serde(alias = "sessionId")]
|
||||
pub session_id: String,
|
||||
#[serde(alias = "eventName")]
|
||||
pub event_name: String,
|
||||
#[serde(default)]
|
||||
#[allow(dead_code)]
|
||||
pub images: Option<Vec<ImageInput>>,
|
||||
/// Provider 配置(可选,如果未配置则使用当前配置)
|
||||
#[serde(default)]
|
||||
#[serde(default, alias = "providerConfig")]
|
||||
pub provider_config: Option<ConfigureProviderRequest>,
|
||||
/// 项目 ID(可选,用于注入项目上下文到 System Prompt)
|
||||
#[serde(default)]
|
||||
#[serde(default, alias = "projectId")]
|
||||
pub project_id: Option<String>,
|
||||
/// Workspace ID(必填,用于校验会话与工作区一致性)
|
||||
#[serde(alias = "workspaceId")]
|
||||
pub workspace_id: String,
|
||||
/// 是否强制开启联网搜索工具策略
|
||||
#[serde(default, alias = "webSearch")]
|
||||
pub web_search: Option<bool>,
|
||||
/// 执行策略(react / code_orchestrated / auto)
|
||||
#[serde(default)]
|
||||
#[serde(default, alias = "executionStrategy")]
|
||||
pub execution_strategy: Option<AsterExecutionStrategy>,
|
||||
}
|
||||
|
||||
@@ -395,14 +409,6 @@ fn should_force_react_for_message(message: &str) -> bool {
|
||||
"webfetch",
|
||||
"web fetch",
|
||||
"web_fetch",
|
||||
"联网搜索",
|
||||
"网络搜索",
|
||||
"实时新闻",
|
||||
"最新新闻",
|
||||
"今日要闻",
|
||||
"时事新闻",
|
||||
"breaking news",
|
||||
"news today",
|
||||
];
|
||||
resolve_intent_hints("PROXYCAST_FORCE_REACT_HINTS", &default_hints)
|
||||
.iter()
|
||||
@@ -499,9 +505,41 @@ async fn stream_reply_once(
|
||||
app: &AppHandle,
|
||||
event_name: &str,
|
||||
message_text: &str,
|
||||
working_directory: Option<&Path>,
|
||||
session_config: aster::agents::SessionConfig,
|
||||
cancel_token: CancellationToken,
|
||||
request_tool_policy: &RequestToolPolicy,
|
||||
) -> Result<(), ReplyAttemptError> {
|
||||
let mut web_search_tracker = WebSearchExecutionTracker::default();
|
||||
let preflight = execute_web_search_preflight_if_needed(
|
||||
agent,
|
||||
&session_config.id,
|
||||
message_text,
|
||||
working_directory,
|
||||
Some(cancel_token.clone()),
|
||||
request_tool_policy,
|
||||
&mut web_search_tracker,
|
||||
)
|
||||
.await;
|
||||
match preflight {
|
||||
Ok(preflight_execution) => {
|
||||
for event in preflight_execution.events {
|
||||
if let Err(error) = app.emit(event_name, &event) {
|
||||
tracing::error!("[AsterAgent] 发送预调用事件失败: {}", error);
|
||||
}
|
||||
}
|
||||
}
|
||||
Err(error) => {
|
||||
return Err(ReplyAttemptError {
|
||||
message: format!(
|
||||
"{error}\n尝试记录: {}",
|
||||
web_search_tracker.format_attempts()
|
||||
),
|
||||
emitted_any: false,
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
let user_message = Message::user().with_text(message_text);
|
||||
let mut stream = agent
|
||||
.reply(user_message, session_config, Some(cancel_token))
|
||||
@@ -522,6 +560,23 @@ async fn stream_reply_once(
|
||||
};
|
||||
let tauri_events = convert_agent_event(agent_event);
|
||||
for tauri_event in tauri_events {
|
||||
match &tauri_event {
|
||||
TauriAgentEvent::ToolStart {
|
||||
tool_name, tool_id, ..
|
||||
} => web_search_tracker.record_tool_start(
|
||||
request_tool_policy,
|
||||
tool_id,
|
||||
tool_name,
|
||||
),
|
||||
TauriAgentEvent::ToolEnd { tool_id, result } => web_search_tracker
|
||||
.record_tool_end(
|
||||
request_tool_policy,
|
||||
tool_id,
|
||||
result.success,
|
||||
result.error.as_deref(),
|
||||
),
|
||||
_ => {}
|
||||
}
|
||||
if let Err(e) = app.emit(event_name, &tauri_event) {
|
||||
tracing::error!("[AsterAgent] 发送事件失败: {}", e);
|
||||
}
|
||||
@@ -543,6 +598,15 @@ async fn stream_reply_once(
|
||||
}
|
||||
}
|
||||
|
||||
if let Err(validation_error) =
|
||||
web_search_tracker.validate_web_search_requirement(request_tool_policy)
|
||||
{
|
||||
return Err(ReplyAttemptError {
|
||||
message: validation_error,
|
||||
emitted_any,
|
||||
});
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
@@ -2082,6 +2146,15 @@ pub async fn aster_agent_chat_stream(
|
||||
);
|
||||
}
|
||||
|
||||
// 构建请求级工具策略:effective_web_search = request.web_search ?? mode_default(false)
|
||||
let request_tool_policy = resolve_request_tool_policy(request.web_search, false);
|
||||
tracing::info!(
|
||||
"[AsterAgent][WebSearchGuard] session={}, request_web_search={:?}, mode_default_web_search=false, effective_web_search={}",
|
||||
session_id,
|
||||
request.web_search,
|
||||
request_tool_policy.effective_web_search
|
||||
);
|
||||
|
||||
// 构建 system_prompt:优先使用项目上下文,其次使用 session 的 system_prompt
|
||||
// 同时读取会话已持久化的 execution_strategy
|
||||
let (system_prompt, persisted_strategy) = {
|
||||
@@ -2138,9 +2211,12 @@ pub async fn aster_agent_chat_stream(
|
||||
}
|
||||
};
|
||||
|
||||
let merged_prompt = merge_system_prompt_with_web_search(
|
||||
merge_system_prompt_with_memory_profile(resolved_prompt, &runtime_config),
|
||||
&runtime_config,
|
||||
let merged_prompt = merge_system_prompt_with_request_tool_policy(
|
||||
merge_system_prompt_with_web_search(
|
||||
merge_system_prompt_with_memory_profile(resolved_prompt, &runtime_config),
|
||||
&runtime_config,
|
||||
),
|
||||
&request_tool_policy,
|
||||
);
|
||||
|
||||
(merged_prompt, persisted)
|
||||
@@ -2176,7 +2252,8 @@ pub async fn aster_agent_chat_stream(
|
||||
// 如果提供了 Provider 配置,则配置 Provider
|
||||
if let Some(provider_config) = &request.provider_config {
|
||||
tracing::info!(
|
||||
"[AsterAgent] 收到 provider_config: provider_name={}, model_name={}, has_api_key={}, base_url={:?}",
|
||||
"[AsterAgent] 收到 provider_config: provider_id={:?}, provider_name={}, model_name={}, has_api_key={}, base_url={:?}",
|
||||
provider_config.provider_id,
|
||||
provider_config.provider_name,
|
||||
provider_config.model_name,
|
||||
provider_config.api_key.is_some(),
|
||||
@@ -2193,11 +2270,15 @@ pub async fn aster_agent_chat_stream(
|
||||
if provider_config.api_key.is_some() {
|
||||
state.configure_provider(config, session_id, &db).await?;
|
||||
} else {
|
||||
// 没有 api_key,使用凭证池(provider_name 作为 provider_type)
|
||||
// 没有 api_key,使用凭证池(优先 provider_id,其次 provider_name)
|
||||
let provider_selector = provider_config
|
||||
.provider_id
|
||||
.as_deref()
|
||||
.unwrap_or(&provider_config.provider_name);
|
||||
state
|
||||
.configure_provider_from_pool(
|
||||
&db,
|
||||
&provider_config.provider_name,
|
||||
provider_selector,
|
||||
&provider_config.model_name,
|
||||
session_id,
|
||||
)
|
||||
@@ -2287,16 +2368,19 @@ pub async fn aster_agent_chat_stream(
|
||||
"event_name": request.event_name.clone(),
|
||||
"execution_strategy": format!("{:?}", effective_strategy).to_lowercase(),
|
||||
"message_length": request.message.chars().count(),
|
||||
"web_search_enabled": request_tool_policy.effective_web_search,
|
||||
})),
|
||||
RunFinalizeOptions {
|
||||
success_metadata: Some(serde_json::json!({
|
||||
"execution_strategy": format!("{:?}", effective_strategy).to_lowercase(),
|
||||
"workspace_id": workspace_id.clone(),
|
||||
"web_search_enabled": request_tool_policy.effective_web_search,
|
||||
})),
|
||||
error_code: Some("chat_stream_failed".to_string()),
|
||||
error_metadata: Some(serde_json::json!({
|
||||
"execution_strategy": format!("{:?}", effective_strategy).to_lowercase(),
|
||||
"workspace_id": workspace_id.clone(),
|
||||
"web_search_enabled": request_tool_policy.effective_web_search,
|
||||
})),
|
||||
},
|
||||
async {
|
||||
@@ -2310,8 +2394,10 @@ pub async fn aster_agent_chat_stream(
|
||||
&app,
|
||||
&request.event_name,
|
||||
&request.message,
|
||||
Some(Path::new(&workspace_root)),
|
||||
build_session_config(),
|
||||
cancel_token.clone(),
|
||||
&request_tool_policy,
|
||||
)
|
||||
.await;
|
||||
|
||||
@@ -2341,8 +2427,10 @@ pub async fn aster_agent_chat_stream(
|
||||
&app,
|
||||
&request.event_name,
|
||||
&request.message,
|
||||
Some(Path::new(&workspace_root)),
|
||||
build_session_config(),
|
||||
cancel_token.clone(),
|
||||
&request_tool_policy,
|
||||
)
|
||||
.await
|
||||
.map_err(|fallback_err| fallback_err.message)
|
||||
@@ -2699,6 +2787,20 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_aster_chat_request_deserialize_with_web_search_flag() {
|
||||
let json = r#"{
|
||||
"message": "Hello",
|
||||
"session_id": "test-session",
|
||||
"event_name": "agent_stream",
|
||||
"workspace_id": "workspace-test",
|
||||
"web_search": true
|
||||
}"#;
|
||||
|
||||
let request: AsterChatRequest = serde_json::from_str(json).unwrap();
|
||||
assert_eq!(request.web_search, Some(true));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_aster_execution_strategy_default_is_auto() {
|
||||
assert_eq!(
|
||||
@@ -2747,7 +2849,7 @@ mod tests {
|
||||
#[test]
|
||||
fn test_aster_execution_strategy_code_orchestrated_still_prefers_react_for_web_search() {
|
||||
let strategy = AsterExecutionStrategy::CodeOrchestrated
|
||||
.effective_for_message("请联网搜索今天的 AI 新闻并给出来源");
|
||||
.effective_for_message("请使用 WebSearch 工具检索并给出来源");
|
||||
assert_eq!(strategy, AsterExecutionStrategy::React);
|
||||
}
|
||||
|
||||
@@ -2758,6 +2860,32 @@ mod tests {
|
||||
assert_eq!(strategy, AsterExecutionStrategy::React);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_merge_system_prompt_with_request_tool_policy_adds_policy_when_enabled() {
|
||||
let policy = resolve_request_tool_policy(Some(true), false);
|
||||
let merged =
|
||||
merge_system_prompt_with_request_tool_policy(Some("你是助手".to_string()), &policy)
|
||||
.expect("should have merged prompt");
|
||||
assert!(merged.contains(REQUEST_TOOL_POLICY_MARKER));
|
||||
assert!(merged.contains("WebSearch"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_merge_system_prompt_with_request_tool_policy_keeps_original_when_disabled() {
|
||||
let base = Some("你好".to_string());
|
||||
let policy = resolve_request_tool_policy(Some(false), false);
|
||||
let merged = merge_system_prompt_with_request_tool_policy(base.clone(), &policy);
|
||||
assert_eq!(merged, base);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_merge_system_prompt_with_request_tool_policy_no_duplicate_marker() {
|
||||
let base = Some(format!("{REQUEST_TOOL_POLICY_MARKER}\n已有策略"));
|
||||
let policy = resolve_request_tool_policy(Some(true), false);
|
||||
let merged = merge_system_prompt_with_request_tool_policy(base.clone(), &policy);
|
||||
assert_eq!(merged, base);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_should_fallback_to_react_from_code_orchestrated_when_no_event_emitted() {
|
||||
let error = ReplyAttemptError {
|
||||
|
||||
@@ -56,34 +56,64 @@ pub struct PythonEnvInfo {
|
||||
pub missing_packages: Vec<String>,
|
||||
}
|
||||
|
||||
fn python_candidates() -> &'static [&'static str] {
|
||||
#[cfg(target_os = "windows")]
|
||||
{
|
||||
&["python", "py", "python3"]
|
||||
}
|
||||
|
||||
#[cfg(not(target_os = "windows"))]
|
||||
{
|
||||
&["python3", "python"]
|
||||
}
|
||||
}
|
||||
|
||||
fn detect_python_command() -> Option<String> {
|
||||
for candidate in python_candidates() {
|
||||
if let Ok(output) = Command::new(candidate).arg("--version").output() {
|
||||
if output.status.success() {
|
||||
return Some((*candidate).to_string());
|
||||
}
|
||||
}
|
||||
}
|
||||
None
|
||||
}
|
||||
|
||||
fn extract_python_version(output: &std::process::Output) -> String {
|
||||
let stdout = String::from_utf8_lossy(&output.stdout).trim().to_string();
|
||||
if !stdout.is_empty() {
|
||||
return stdout;
|
||||
}
|
||||
String::from_utf8_lossy(&output.stderr).trim().to_string()
|
||||
}
|
||||
|
||||
/// 检查 Python 环境
|
||||
#[tauri::command]
|
||||
pub async fn check_python_env() -> Result<PythonEnvInfo, String> {
|
||||
// 检查 Python 是否安装
|
||||
let python_check = Command::new("python3").arg("--version").output();
|
||||
|
||||
let (python_installed, python_version) = match python_check {
|
||||
Ok(output) => {
|
||||
let version = String::from_utf8_lossy(&output.stdout).trim().to_string();
|
||||
(true, Some(version))
|
||||
let python_command = detect_python_command();
|
||||
let (python_installed, python_version, python_cmd) = match python_command {
|
||||
Some(cmd) => {
|
||||
let output = Command::new(&cmd).arg("--version").output();
|
||||
match output {
|
||||
Ok(output) => (true, Some(extract_python_version(&output)), cmd),
|
||||
Err(_) => (true, None, cmd),
|
||||
}
|
||||
}
|
||||
None => {
|
||||
return Ok(PythonEnvInfo {
|
||||
python_installed: false,
|
||||
python_version: None,
|
||||
missing_packages: vec![],
|
||||
});
|
||||
}
|
||||
Err(_) => (false, None),
|
||||
};
|
||||
|
||||
if !python_installed {
|
||||
return Ok(PythonEnvInfo {
|
||||
python_installed: false,
|
||||
python_version: None,
|
||||
missing_packages: vec![],
|
||||
});
|
||||
}
|
||||
|
||||
// 检查必需的 Python 包
|
||||
let required_packages = vec!["mido", "music21", "numpy", "demucs", "basic-pitch"];
|
||||
let mut missing_packages = Vec::new();
|
||||
|
||||
for package in required_packages {
|
||||
let check = Command::new("python3")
|
||||
let check = Command::new(&python_cmd)
|
||||
.arg("-c")
|
||||
.arg(format!("import {}", package.replace("-", "_")))
|
||||
.output();
|
||||
@@ -105,9 +135,11 @@ pub async fn check_python_env() -> Result<PythonEnvInfo, String> {
|
||||
pub async fn analyze_midi(midi_path: String) -> Result<MidiAnalysisResult, String> {
|
||||
// 获取 Python 脚本路径
|
||||
let script_path = get_resource_path("scripts/midi_analyzer.py")?;
|
||||
let python_cmd = detect_python_command()
|
||||
.ok_or_else(|| "Python is not installed or not found in PATH".to_string())?;
|
||||
|
||||
// 调用 Python 脚本
|
||||
let output = Command::new("python3")
|
||||
let output = Command::new(&python_cmd)
|
||||
.arg(&script_path)
|
||||
.arg(&midi_path)
|
||||
.output()
|
||||
@@ -128,9 +160,11 @@ pub async fn analyze_midi(midi_path: String) -> Result<MidiAnalysisResult, Strin
|
||||
pub async fn convert_mp3_to_midi(mp3_path: String, output_path: String) -> Result<String, String> {
|
||||
// 获取 Python 脚本路径
|
||||
let script_path = get_resource_path("scripts/audio_to_midi.py")?;
|
||||
let python_cmd = detect_python_command()
|
||||
.ok_or_else(|| "Python is not installed or not found in PATH".to_string())?;
|
||||
|
||||
// 调用 Python 脚本
|
||||
let output = Command::new("python3")
|
||||
let output = Command::new(&python_cmd)
|
||||
.arg(&script_path)
|
||||
.arg(&mp3_path)
|
||||
.arg(&output_path)
|
||||
@@ -194,8 +228,12 @@ fn get_resource_path(relative_path: &str) -> Result<PathBuf, String> {
|
||||
#[tauri::command]
|
||||
pub async fn install_python_dependencies() -> Result<String, String> {
|
||||
let packages = vec!["mido", "music21", "numpy", "demucs", "basic-pitch"];
|
||||
let python_cmd = detect_python_command()
|
||||
.ok_or_else(|| "Python is not installed or not found in PATH".to_string())?;
|
||||
|
||||
let output = Command::new("pip3")
|
||||
let output = Command::new(&python_cmd)
|
||||
.arg("-m")
|
||||
.arg("pip")
|
||||
.arg("install")
|
||||
.args(&packages)
|
||||
.output()
|
||||
@@ -208,3 +246,25 @@ pub async fn install_python_dependencies() -> Result<String, String> {
|
||||
|
||||
Ok("Dependencies installed successfully".to_string())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::python_candidates;
|
||||
|
||||
#[test]
|
||||
fn python_candidates_should_not_be_empty() {
|
||||
assert!(!python_candidates().is_empty());
|
||||
}
|
||||
|
||||
#[cfg(target_os = "windows")]
|
||||
#[test]
|
||||
fn windows_python_candidates_should_prioritize_python() {
|
||||
assert_eq!(python_candidates().first().copied(), Some("python"));
|
||||
}
|
||||
|
||||
#[cfg(not(target_os = "windows"))]
|
||||
#[test]
|
||||
fn unix_python_candidates_should_prioritize_python3() {
|
||||
assert_eq!(python_candidates().first().copied(), Some("python3"));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -20,6 +20,10 @@ use crate::config::GlobalConfigManagerState;
|
||||
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::request_tool_policy_prompt_service::{
|
||||
execute_web_search_preflight_if_needed, merge_system_prompt_with_request_tool_policy,
|
||||
resolve_request_tool_policy, RequestToolPolicy, WebSearchExecutionTracker,
|
||||
};
|
||||
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;
|
||||
@@ -56,14 +60,19 @@ pub struct CreateSessionRequest {
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub struct SendMessageRequest {
|
||||
/// 会话 ID
|
||||
#[serde(alias = "sessionId")]
|
||||
pub session_id: String,
|
||||
/// 消息内容
|
||||
pub message: String,
|
||||
/// 事件名称(用于前端监听)
|
||||
#[serde(alias = "eventName")]
|
||||
pub event_name: String,
|
||||
/// 图片输入(可选,用于多模态对话)
|
||||
/// TODO: 实现图片处理逻辑,将图片转换为 Aster Message 的 ImageContent
|
||||
pub images: Option<Vec<ImageInput>>,
|
||||
/// 请求级联网搜索开关
|
||||
#[serde(default, alias = "webSearch")]
|
||||
pub web_search: Option<bool>,
|
||||
}
|
||||
|
||||
/// 图片输入
|
||||
@@ -361,46 +370,30 @@ pub async fn chat_send_message(
|
||||
&config,
|
||||
);
|
||||
|
||||
let prefer_web_search_tools = matches!(session.mode, ChatMode::General);
|
||||
let mode_default_web_search = matches!(session.mode, ChatMode::General);
|
||||
let request_tool_policy =
|
||||
resolve_request_tool_policy(request.web_search, mode_default_web_search);
|
||||
tracing::info!(
|
||||
"[UnifiedChat][WebSearchGuard] session={}, mode={:?}, prefer_web_search_tools={}",
|
||||
"[UnifiedChat][WebSearchGuard] session={}, mode={:?}, request_web_search={:?}, mode_default_web_search={}, effective_web_search={}",
|
||||
request.session_id,
|
||||
session.mode,
|
||||
prefer_web_search_tools
|
||||
request.web_search,
|
||||
mode_default_web_search,
|
||||
request_tool_policy.effective_web_search
|
||||
);
|
||||
|
||||
let result = match session.mode {
|
||||
ChatMode::Agent | ChatMode::Creator => {
|
||||
// 使用 Aster Agent 处理
|
||||
send_message_with_aster(
|
||||
&app,
|
||||
&db,
|
||||
&agent_state,
|
||||
&request.session_id,
|
||||
&request.message,
|
||||
&request.event_name,
|
||||
merged_system_prompt.as_deref(),
|
||||
config.memory.enabled,
|
||||
false,
|
||||
)
|
||||
.await
|
||||
}
|
||||
ChatMode::General => {
|
||||
// 通用模式:也使用 Aster Agent,但不启用工具
|
||||
send_message_with_aster(
|
||||
&app,
|
||||
&db,
|
||||
&agent_state,
|
||||
&request.session_id,
|
||||
&request.message,
|
||||
&request.event_name,
|
||||
merged_system_prompt.as_deref(),
|
||||
config.memory.enabled,
|
||||
true,
|
||||
)
|
||||
.await
|
||||
}
|
||||
};
|
||||
let result = send_message_with_aster(
|
||||
&app,
|
||||
&db,
|
||||
&agent_state,
|
||||
&request.session_id,
|
||||
&request.message,
|
||||
&request.event_name,
|
||||
merged_system_prompt.as_deref(),
|
||||
config.memory.enabled,
|
||||
&request_tool_policy,
|
||||
)
|
||||
.await;
|
||||
|
||||
let total_elapsed = start_time.elapsed();
|
||||
tracing::info!(
|
||||
@@ -422,13 +415,13 @@ async fn send_message_with_aster(
|
||||
event_name: &str,
|
||||
system_prompt: Option<&str>,
|
||||
include_context_trace: bool,
|
||||
prefer_web_search_tools: bool,
|
||||
request_tool_policy: &RequestToolPolicy,
|
||||
) -> Result<(), String> {
|
||||
let start_time = std::time::Instant::now();
|
||||
tracing::info!(
|
||||
"[UnifiedChat][WebSearchGuard] session={}, prefer_web_search_tools={}",
|
||||
"[UnifiedChat][WebSearchGuard] session={}, effective_web_search={}",
|
||||
session_id,
|
||||
prefer_web_search_tools
|
||||
request_tool_policy.effective_web_search
|
||||
);
|
||||
|
||||
// 确保 Agent 已初始化
|
||||
@@ -454,26 +447,17 @@ async fn send_message_with_aster(
|
||||
// 创建取消令牌
|
||||
let cancel_token = agent_state.create_cancel_token(session_id).await;
|
||||
|
||||
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()
|
||||
};
|
||||
let effective_system_prompt = merge_system_prompt_with_request_tool_policy(
|
||||
system_prompt.map(|prompt| prompt.to_string()),
|
||||
request_tool_policy,
|
||||
);
|
||||
|
||||
// 构建消息(如果有 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)
|
||||
let user_message = Message::user().with_text(message);
|
||||
let mut session_config_builder = SessionConfigBuilder::new(session_id);
|
||||
if let Some(prompt) = effective_system_prompt {
|
||||
session_config_builder = session_config_builder.system_prompt(prompt);
|
||||
}
|
||||
let session_config = session_config_builder
|
||||
.include_context_trace(include_context_trace)
|
||||
.build();
|
||||
|
||||
@@ -483,7 +467,7 @@ async fn send_message_with_aster(
|
||||
let agent = guard.as_ref().ok_or("Agent 未初始化")?;
|
||||
|
||||
let mut removed_extension: Option<ExtensionConfig> = None;
|
||||
if prefer_web_search_tools {
|
||||
if request_tool_policy.effective_web_search {
|
||||
let extension_configs = agent.get_extension_configs().await;
|
||||
if let Some(extension) = extension_configs
|
||||
.into_iter()
|
||||
@@ -516,6 +500,47 @@ async fn send_message_with_aster(
|
||||
|
||||
// 调用 Agent
|
||||
let reply_start = std::time::Instant::now();
|
||||
let mut web_search_tracker = WebSearchExecutionTracker::default();
|
||||
let preflight = execute_web_search_preflight_if_needed(
|
||||
agent,
|
||||
session_id,
|
||||
message,
|
||||
None,
|
||||
Some(cancel_token.clone()),
|
||||
request_tool_policy,
|
||||
&mut web_search_tracker,
|
||||
)
|
||||
.await;
|
||||
match preflight {
|
||||
Ok(preflight_execution) => {
|
||||
for event in preflight_execution.events {
|
||||
if let Err(error) = app.emit(event_name, &event) {
|
||||
tracing::error!("[UnifiedChat] 发送预调用事件失败: {}", error);
|
||||
}
|
||||
}
|
||||
}
|
||||
Err(error) => {
|
||||
let error_event = TauriAgentEvent::Error {
|
||||
message: format!(
|
||||
"{error}\n尝试记录: {}",
|
||||
web_search_tracker.format_attempts()
|
||||
),
|
||||
};
|
||||
let _ = app.emit(event_name, &error_event);
|
||||
agent_state.remove_cancel_token(session_id).await;
|
||||
if let Some(extension) = removed_extension {
|
||||
if let Err(restore_error) = agent.add_extension(extension).await {
|
||||
tracing::warn!(
|
||||
"[UnifiedChat] 预调用失败后恢复 {} 扩展失败: {}",
|
||||
CODE_EXECUTION_EXTENSION_NAME,
|
||||
restore_error
|
||||
);
|
||||
}
|
||||
}
|
||||
return Err(error);
|
||||
}
|
||||
}
|
||||
|
||||
let stream_result = agent
|
||||
.reply(user_message, session_config, Some(cancel_token.clone()))
|
||||
.await;
|
||||
@@ -539,6 +564,24 @@ async fn send_message_with_aster(
|
||||
|
||||
let tauri_events = convert_agent_event(agent_event);
|
||||
for tauri_event in tauri_events {
|
||||
match &tauri_event {
|
||||
TauriAgentEvent::ToolStart {
|
||||
tool_name, tool_id, ..
|
||||
} => web_search_tracker.record_tool_start(
|
||||
request_tool_policy,
|
||||
tool_id,
|
||||
tool_name,
|
||||
),
|
||||
TauriAgentEvent::ToolEnd { tool_id, result } => {
|
||||
web_search_tracker.record_tool_end(
|
||||
request_tool_policy,
|
||||
tool_id,
|
||||
result.success,
|
||||
result.error.as_deref(),
|
||||
);
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
if let Err(e) = app.emit(event_name, &tauri_event) {
|
||||
tracing::error!("[UnifiedChat] 发送事件失败: {}", e);
|
||||
}
|
||||
@@ -555,9 +598,23 @@ async fn send_message_with_aster(
|
||||
}
|
||||
}
|
||||
|
||||
// 发送完成事件
|
||||
let done_event = TauriAgentEvent::FinalDone { usage: None };
|
||||
let _ = app.emit(event_name, &done_event);
|
||||
if stream_error.is_none() {
|
||||
if let Err(validation_error) =
|
||||
web_search_tracker.validate_web_search_requirement(request_tool_policy)
|
||||
{
|
||||
let error_event = TauriAgentEvent::Error {
|
||||
message: validation_error.clone(),
|
||||
};
|
||||
let _ = app.emit(event_name, &error_event);
|
||||
stream_error = Some(validation_error);
|
||||
}
|
||||
}
|
||||
|
||||
if stream_error.is_none() {
|
||||
// 发送完成事件
|
||||
let done_event = TauriAgentEvent::FinalDone { usage: None };
|
||||
let _ = app.emit(event_name, &done_event);
|
||||
}
|
||||
|
||||
let stream_elapsed = start_time.elapsed();
|
||||
tracing::info!(
|
||||
@@ -633,3 +690,44 @@ pub async fn chat_configure_provider(
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::services::request_tool_policy_prompt_service::resolve_request_tool_policy;
|
||||
|
||||
#[test]
|
||||
fn test_send_message_request_deserialize_web_search_camel_case() {
|
||||
let payload = serde_json::json!({
|
||||
"sessionId": "session-1",
|
||||
"message": "hello",
|
||||
"eventName": "event-1",
|
||||
"webSearch": true
|
||||
});
|
||||
let request: SendMessageRequest =
|
||||
serde_json::from_value(payload).expect("deserialize request");
|
||||
assert_eq!(request.web_search, Some(true));
|
||||
assert_eq!(request.session_id, "session-1");
|
||||
assert_eq!(request.event_name, "event-1");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_send_message_request_deserialize_web_search_snake_case() {
|
||||
let payload = serde_json::json!({
|
||||
"session_id": "session-1",
|
||||
"message": "hello",
|
||||
"event_name": "event-1",
|
||||
"web_search": false
|
||||
});
|
||||
let request: SendMessageRequest =
|
||||
serde_json::from_value(payload).expect("deserialize request");
|
||||
assert_eq!(request.web_search, Some(false));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_unified_effective_web_search_uses_request_override() {
|
||||
let mode_default = true;
|
||||
let policy = resolve_request_tool_policy(Some(false), mode_default);
|
||||
assert!(!policy.effective_web_search);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -14,6 +14,7 @@ pub mod memory_profile_prompt_service;
|
||||
pub mod memory_rules_loader_service;
|
||||
pub mod memory_source_resolver_service;
|
||||
pub mod novel_service;
|
||||
pub mod request_tool_policy_prompt_service;
|
||||
pub mod sysinfo_service;
|
||||
pub mod update_check_service;
|
||||
pub mod update_window;
|
||||
|
||||
@@ -0,0 +1,566 @@
|
||||
//! 请求级工具策略提示词服务
|
||||
//!
|
||||
//! 将本次请求的工具偏好(例如“开启联网搜索”)统一转换为系统提示词附加项,
|
||||
//! 避免通过改写用户原始消息来注入策略。
|
||||
//!
|
||||
//! 设计目标:
|
||||
//! - 单一策略入口:Aster 聊天入口与统一执行入口复用同一策略解析/校验逻辑
|
||||
//! - 请求级优先:`effective_web_search = request.web_search ?? mode_default`
|
||||
//! - 配置驱动:工具白/黑名单支持环境变量覆盖,保留扩展性
|
||||
|
||||
use aster::agents::Agent;
|
||||
use aster::tools::ToolContext;
|
||||
use proxycast_agent::event_converter::{TauriAgentEvent, TauriToolResult};
|
||||
use std::collections::HashMap;
|
||||
use std::path::Path;
|
||||
use tokio_util::sync::CancellationToken;
|
||||
use uuid::Uuid;
|
||||
|
||||
pub const REQUEST_TOOL_POLICY_MARKER: &str = "【请求级工具策略】";
|
||||
|
||||
const DEFAULT_REQUIRED_TOOLS: &[&str] = &["WebSearch"];
|
||||
const DEFAULT_ALLOWED_TOOLS: &[&str] = &["WebSearch", "WebFetch"];
|
||||
const WEB_SEARCH_REQUIRED_TOOLS_ENV: &str = "PROXYCAST_WEB_SEARCH_REQUIRED_TOOLS";
|
||||
const WEB_SEARCH_ALLOWED_TOOLS_ENV: &str = "PROXYCAST_WEB_SEARCH_ALLOWED_TOOLS";
|
||||
const WEB_SEARCH_DISALLOWED_TOOLS_ENV: &str = "PROXYCAST_WEB_SEARCH_DISALLOWED_TOOLS";
|
||||
const WEB_SEARCH_PREFLIGHT_ENABLED_ENV: &str = "PROXYCAST_WEB_SEARCH_PREFLIGHT_ENABLED";
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct RequestToolPolicy {
|
||||
/// 本次请求是否开启联网搜索策略
|
||||
pub effective_web_search: bool,
|
||||
/// 必须至少成功一次的工具(默认 WebSearch)
|
||||
pub required_tools: Vec<String>,
|
||||
/// 允许的联网工具集合(默认 WebSearch/WebFetch)
|
||||
pub allowed_tools: Vec<String>,
|
||||
/// 禁止工具集合(可配置)
|
||||
pub disallowed_tools: Vec<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct ToolAttemptRecord {
|
||||
pub tool_id: String,
|
||||
pub tool_name: String,
|
||||
pub success: Option<bool>,
|
||||
pub error: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Default)]
|
||||
pub struct WebSearchExecutionTracker {
|
||||
ordered_tool_ids: Vec<String>,
|
||||
attempts_by_id: HashMap<String, ToolAttemptRecord>,
|
||||
}
|
||||
|
||||
impl WebSearchExecutionTracker {
|
||||
pub fn record_tool_start(
|
||||
&mut self,
|
||||
policy: &RequestToolPolicy,
|
||||
tool_id: &str,
|
||||
tool_name: &str,
|
||||
) {
|
||||
if !policy.effective_web_search || tool_id.trim().is_empty() || tool_name.trim().is_empty()
|
||||
{
|
||||
return;
|
||||
}
|
||||
|
||||
if !self.attempts_by_id.contains_key(tool_id) {
|
||||
self.ordered_tool_ids.push(tool_id.to_string());
|
||||
self.attempts_by_id.insert(
|
||||
tool_id.to_string(),
|
||||
ToolAttemptRecord {
|
||||
tool_id: tool_id.to_string(),
|
||||
tool_name: tool_name.to_string(),
|
||||
success: None,
|
||||
error: None,
|
||||
},
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
pub fn record_tool_end(
|
||||
&mut self,
|
||||
policy: &RequestToolPolicy,
|
||||
tool_id: &str,
|
||||
success: bool,
|
||||
error: Option<&str>,
|
||||
) {
|
||||
if !policy.effective_web_search || tool_id.trim().is_empty() {
|
||||
return;
|
||||
}
|
||||
if let Some(record) = self.attempts_by_id.get_mut(tool_id) {
|
||||
record.success = Some(success);
|
||||
record.error = error
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(|value| value.to_string());
|
||||
}
|
||||
}
|
||||
|
||||
pub fn validate_web_search_requirement(
|
||||
&self,
|
||||
policy: &RequestToolPolicy,
|
||||
) -> Result<(), String> {
|
||||
if !policy.effective_web_search {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let disallowed_attempts: Vec<&ToolAttemptRecord> = self
|
||||
.ordered_tool_ids
|
||||
.iter()
|
||||
.filter_map(|tool_id| self.attempts_by_id.get(tool_id))
|
||||
.filter(|record| matches_tool_list(&record.tool_name, &policy.disallowed_tools))
|
||||
.collect();
|
||||
if !disallowed_attempts.is_empty() {
|
||||
let disallowed_names = disallowed_attempts
|
||||
.iter()
|
||||
.map(|record| record.tool_name.clone())
|
||||
.collect::<Vec<_>>()
|
||||
.join(", ");
|
||||
return Err(format!(
|
||||
"联网搜索策略阻止了禁止工具调用: {}。\n尝试记录: {}",
|
||||
disallowed_names,
|
||||
self.format_attempts()
|
||||
));
|
||||
}
|
||||
|
||||
let required_attempts: Vec<&ToolAttemptRecord> = self
|
||||
.ordered_tool_ids
|
||||
.iter()
|
||||
.filter_map(|tool_id| self.attempts_by_id.get(tool_id))
|
||||
.filter(|record| policy.matches_any_required_tool(&record.tool_name))
|
||||
.collect();
|
||||
|
||||
if required_attempts.is_empty() {
|
||||
return Err(format!(
|
||||
"联网搜索已开启,但未检测到必需工具调用。必须先调用 {} 至少一次后再给出最终答复。\n尝试记录: {}",
|
||||
policy.required_tools.join(", "),
|
||||
self.format_attempts()
|
||||
));
|
||||
}
|
||||
|
||||
if required_attempts
|
||||
.iter()
|
||||
.any(|record| record.success.unwrap_or(false))
|
||||
{
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
Err(format!(
|
||||
"联网搜索已开启,但必需工具调用全部失败,无法给出符合约束的最终答复。\n失败原因与尝试记录: {}",
|
||||
self.format_attempts()
|
||||
))
|
||||
}
|
||||
|
||||
pub fn format_attempts(&self) -> String {
|
||||
if self.ordered_tool_ids.is_empty() {
|
||||
return "无工具调用".to_string();
|
||||
}
|
||||
|
||||
self.ordered_tool_ids
|
||||
.iter()
|
||||
.filter_map(|tool_id| self.attempts_by_id.get(tool_id))
|
||||
.map(|record| {
|
||||
let status = match record.success {
|
||||
Some(true) => "success".to_string(),
|
||||
Some(false) => {
|
||||
format!("failed({})", record.error.as_deref().unwrap_or("unknown"))
|
||||
}
|
||||
None => "pending".to_string(),
|
||||
};
|
||||
format!("{}#{}:{}", record.tool_name, record.tool_id, status)
|
||||
})
|
||||
.collect::<Vec<_>>()
|
||||
.join("; ")
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct PreflightToolExecution {
|
||||
pub events: Vec<TauriAgentEvent>,
|
||||
}
|
||||
|
||||
impl PreflightToolExecution {
|
||||
fn none() -> Self {
|
||||
Self { events: Vec::new() }
|
||||
}
|
||||
}
|
||||
|
||||
impl RequestToolPolicy {
|
||||
pub fn matches_any_required_tool(&self, tool_name: &str) -> bool {
|
||||
matches_tool_list(tool_name, &self.required_tools)
|
||||
}
|
||||
|
||||
pub fn matches_any_allowed_tool(&self, tool_name: &str) -> bool {
|
||||
matches_tool_list(tool_name, &self.allowed_tools)
|
||||
}
|
||||
}
|
||||
|
||||
/// 解析请求级工具策略
|
||||
///
|
||||
/// 规则:
|
||||
/// - `effective_web_search = request_web_search.unwrap_or(mode_default)`
|
||||
/// - 白/黑名单支持环境变量覆盖:
|
||||
/// - `PROXYCAST_WEB_SEARCH_REQUIRED_TOOLS`
|
||||
/// - `PROXYCAST_WEB_SEARCH_ALLOWED_TOOLS`
|
||||
/// - `PROXYCAST_WEB_SEARCH_DISALLOWED_TOOLS`
|
||||
pub fn resolve_request_tool_policy(
|
||||
request_web_search: Option<bool>,
|
||||
mode_default: bool,
|
||||
) -> RequestToolPolicy {
|
||||
let effective_web_search = request_web_search.unwrap_or(mode_default);
|
||||
let required_tools = parse_tool_list_env(WEB_SEARCH_REQUIRED_TOOLS_ENV, DEFAULT_REQUIRED_TOOLS);
|
||||
let mut allowed_tools =
|
||||
parse_tool_list_env(WEB_SEARCH_ALLOWED_TOOLS_ENV, DEFAULT_ALLOWED_TOOLS);
|
||||
let disallowed_tools = parse_tool_list_env(WEB_SEARCH_DISALLOWED_TOOLS_ENV, &[]);
|
||||
|
||||
for required in &required_tools {
|
||||
if !allowed_tools
|
||||
.iter()
|
||||
.any(|candidate| is_same_tool(candidate, required))
|
||||
{
|
||||
allowed_tools.push(required.clone());
|
||||
}
|
||||
}
|
||||
|
||||
RequestToolPolicy {
|
||||
effective_web_search,
|
||||
required_tools,
|
||||
allowed_tools,
|
||||
disallowed_tools,
|
||||
}
|
||||
}
|
||||
|
||||
/// 合并请求级工具策略到系统提示词
|
||||
///
|
||||
/// - `effective_web_search=false`:保持原始 system prompt 不变
|
||||
/// - 已包含 marker 时:不重复追加
|
||||
pub fn merge_system_prompt_with_request_tool_policy(
|
||||
base_prompt: Option<String>,
|
||||
policy: &RequestToolPolicy,
|
||||
) -> Option<String> {
|
||||
if !policy.effective_web_search {
|
||||
return base_prompt;
|
||||
}
|
||||
|
||||
let disallowed_line = if policy.disallowed_tools.is_empty() {
|
||||
"无".to_string()
|
||||
} else {
|
||||
policy.disallowed_tools.join(", ")
|
||||
};
|
||||
|
||||
let policy_prompt = format!(
|
||||
"{REQUEST_TOOL_POLICY_MARKER}\n\
|
||||
- 用户在本次请求中已开启“联网搜索”开关。\n\
|
||||
- 必须先调用 {} 至少一次(必要时再调用 WebFetch),再输出最终答复。\n\
|
||||
- 若工具调用失败,必须返回失败原因与尝试记录;不要在未完成必需工具调用前直接给最终结论。\n\
|
||||
- 允许工具: {}\n\
|
||||
- 禁止工具: {}",
|
||||
policy.required_tools.join(", "),
|
||||
policy.allowed_tools.join(", "),
|
||||
disallowed_line
|
||||
);
|
||||
|
||||
match base_prompt {
|
||||
Some(base) => {
|
||||
if base.contains(REQUEST_TOOL_POLICY_MARKER) {
|
||||
Some(base)
|
||||
} else if base.trim().is_empty() {
|
||||
Some(policy_prompt)
|
||||
} else {
|
||||
Some(format!("{base}\n\n{policy_prompt}"))
|
||||
}
|
||||
}
|
||||
None => Some(policy_prompt),
|
||||
}
|
||||
}
|
||||
|
||||
fn parse_tool_list_env(key: &str, default_values: &[&str]) -> Vec<String> {
|
||||
let from_env = std::env::var(key)
|
||||
.ok()
|
||||
.map(|raw| parse_tool_list(&raw))
|
||||
.filter(|tools| !tools.is_empty());
|
||||
|
||||
let values =
|
||||
from_env.unwrap_or_else(|| default_values.iter().map(|item| item.to_string()).collect());
|
||||
dedup_tools(values)
|
||||
}
|
||||
|
||||
fn parse_tool_list(raw: &str) -> Vec<String> {
|
||||
raw.split(',')
|
||||
.map(str::trim)
|
||||
.filter(|item| !item.is_empty())
|
||||
.map(|item| item.to_string())
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn dedup_tools(values: Vec<String>) -> Vec<String> {
|
||||
let mut result: Vec<String> = Vec::new();
|
||||
for value in values {
|
||||
if !result.iter().any(|existing| is_same_tool(existing, &value)) {
|
||||
result.push(value);
|
||||
}
|
||||
}
|
||||
result
|
||||
}
|
||||
|
||||
fn matches_tool_list(tool_name: &str, list: &[String]) -> bool {
|
||||
list.iter()
|
||||
.any(|candidate| is_same_tool(tool_name, candidate))
|
||||
}
|
||||
|
||||
fn is_same_tool(a: &str, b: &str) -> bool {
|
||||
let normalized_a = normalize_tool_name(a);
|
||||
let normalized_b = normalize_tool_name(b);
|
||||
if normalized_a.is_empty() || normalized_b.is_empty() {
|
||||
return false;
|
||||
}
|
||||
normalized_a == normalized_b
|
||||
|| normalized_a.contains(&normalized_b)
|
||||
|| normalized_b.contains(&normalized_a)
|
||||
}
|
||||
|
||||
fn normalize_tool_name(value: &str) -> String {
|
||||
value
|
||||
.chars()
|
||||
.filter(|ch| ch.is_ascii_alphanumeric())
|
||||
.flat_map(|ch| ch.to_lowercase())
|
||||
.collect::<String>()
|
||||
}
|
||||
|
||||
/// 当开启联网搜索时,在正式回复前执行一次 WebSearch 预调用。
|
||||
///
|
||||
/// 目标:
|
||||
/// - 通过执行层保证至少一次 WebSearch 调用(而非仅依赖提示词)
|
||||
/// - 统一生成 tool_start/tool_end 事件,供前端落地
|
||||
/// - 若预调用失败,返回失败原因并由上层中断本次回答
|
||||
pub async fn execute_web_search_preflight_if_needed(
|
||||
agent: &Agent,
|
||||
session_id: &str,
|
||||
message_text: &str,
|
||||
working_directory: Option<&Path>,
|
||||
cancel_token: Option<CancellationToken>,
|
||||
policy: &RequestToolPolicy,
|
||||
tracker: &mut WebSearchExecutionTracker,
|
||||
) -> Result<PreflightToolExecution, String> {
|
||||
if !policy.effective_web_search || !is_web_search_preflight_enabled() {
|
||||
return Ok(PreflightToolExecution::none());
|
||||
}
|
||||
|
||||
let registry_arc = agent.tool_registry().clone();
|
||||
let registry = registry_arc.read().await;
|
||||
let available_tools = registry.get_definitions();
|
||||
let preflight_tool = available_tools
|
||||
.iter()
|
||||
.find(|definition| {
|
||||
policy.matches_any_required_tool(&definition.name)
|
||||
&& normalize_tool_name(&definition.name).contains("websearch")
|
||||
})
|
||||
.ok_or_else(|| {
|
||||
format!(
|
||||
"联网搜索已开启,但未找到可执行的必需工具定义。required_tools={}, available_tools={}",
|
||||
policy.required_tools.join(", "),
|
||||
available_tools
|
||||
.iter()
|
||||
.map(|definition| definition.name.clone())
|
||||
.collect::<Vec<_>>()
|
||||
.join(", ")
|
||||
)
|
||||
})?;
|
||||
|
||||
let query = derive_preflight_query(message_text);
|
||||
let params = serde_json::json!({ "query": query });
|
||||
let arguments = serde_json::to_string(¶ms).ok();
|
||||
let tool_id = format!("preflight-websearch-{}", Uuid::new_v4());
|
||||
tracker.record_tool_start(policy, &tool_id, &preflight_tool.name);
|
||||
|
||||
let mut context = ToolContext::new(
|
||||
working_directory
|
||||
.map(Path::to_path_buf)
|
||||
.or_else(|| std::env::current_dir().ok())
|
||||
.unwrap_or_default(),
|
||||
)
|
||||
.with_session_id(session_id.to_string());
|
||||
if let Some(token) = cancel_token {
|
||||
context = context.with_cancellation_token(token);
|
||||
}
|
||||
|
||||
let mut events = vec![TauriAgentEvent::ToolStart {
|
||||
tool_name: preflight_tool.name.clone(),
|
||||
tool_id: tool_id.clone(),
|
||||
arguments,
|
||||
}];
|
||||
|
||||
let result = registry
|
||||
.execute(&preflight_tool.name, params, &context, None)
|
||||
.await
|
||||
.map_err(|error| format!("执行 WebSearch 预调用失败: {}", error.to_string()));
|
||||
|
||||
match result {
|
||||
Ok(tool_result) => {
|
||||
tracker.record_tool_end(
|
||||
policy,
|
||||
&tool_id,
|
||||
tool_result.success,
|
||||
tool_result.error.as_deref(),
|
||||
);
|
||||
let event = TauriAgentEvent::ToolEnd {
|
||||
tool_id,
|
||||
result: TauriToolResult {
|
||||
success: tool_result.success,
|
||||
output: tool_result.output.unwrap_or_default(),
|
||||
error: tool_result.error,
|
||||
images: None,
|
||||
},
|
||||
};
|
||||
events.push(event);
|
||||
|
||||
if events
|
||||
.last()
|
||||
.and_then(|event| match event {
|
||||
TauriAgentEvent::ToolEnd { result, .. } => Some(result.success),
|
||||
_ => None,
|
||||
})
|
||||
.unwrap_or(false)
|
||||
{
|
||||
Ok(PreflightToolExecution { events })
|
||||
} else {
|
||||
let failure = events.last().and_then(|event| match event {
|
||||
TauriAgentEvent::ToolEnd { result, .. } => result.error.clone(),
|
||||
_ => None,
|
||||
});
|
||||
Err(format!(
|
||||
"联网搜索预调用失败: {}",
|
||||
failure.unwrap_or_else(|| "unknown".to_string())
|
||||
))
|
||||
}
|
||||
}
|
||||
Err(error) => {
|
||||
tracker.record_tool_end(policy, &tool_id, false, Some(error.as_str()));
|
||||
events.push(TauriAgentEvent::ToolEnd {
|
||||
tool_id,
|
||||
result: TauriToolResult {
|
||||
success: false,
|
||||
output: String::new(),
|
||||
error: Some(error.clone()),
|
||||
images: None,
|
||||
},
|
||||
});
|
||||
Err(error)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn is_web_search_preflight_enabled() -> bool {
|
||||
match std::env::var(WEB_SEARCH_PREFLIGHT_ENABLED_ENV) {
|
||||
Ok(raw) => match raw.trim().to_ascii_lowercase().as_str() {
|
||||
"0" | "false" | "no" | "off" => false,
|
||||
_ => true,
|
||||
},
|
||||
Err(_) => true,
|
||||
}
|
||||
}
|
||||
|
||||
fn derive_preflight_query(message_text: &str) -> String {
|
||||
let trimmed = message_text.trim();
|
||||
if trimmed.chars().count() >= 2 {
|
||||
return trimmed.to_string();
|
||||
}
|
||||
if trimmed.is_empty() {
|
||||
return "最新信息".to_string();
|
||||
}
|
||||
|
||||
// 兜底补齐最短长度,避免触发 WebSearch.query minLength 校验失败
|
||||
let mut fallback = trimmed.to_string();
|
||||
while fallback.chars().count() < 2 {
|
||||
fallback.push_str(" 信息");
|
||||
}
|
||||
fallback
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn resolves_effective_web_search_with_request_override() {
|
||||
let policy = resolve_request_tool_policy(Some(false), true);
|
||||
assert!(!policy.effective_web_search);
|
||||
|
||||
let policy = resolve_request_tool_policy(Some(true), false);
|
||||
assert!(policy.effective_web_search);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolves_effective_web_search_with_mode_default() {
|
||||
let policy = resolve_request_tool_policy(None, true);
|
||||
assert!(policy.effective_web_search);
|
||||
|
||||
let policy = resolve_request_tool_policy(None, false);
|
||||
assert!(!policy.effective_web_search);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn keeps_original_prompt_when_disabled() {
|
||||
let base = Some("base".to_string());
|
||||
let policy = resolve_request_tool_policy(Some(false), false);
|
||||
assert_eq!(
|
||||
merge_system_prompt_with_request_tool_policy(base.clone(), &policy),
|
||||
base
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn appends_policy_prompt_when_enabled() {
|
||||
let policy = resolve_request_tool_policy(Some(true), false);
|
||||
let merged =
|
||||
merge_system_prompt_with_request_tool_policy(Some("base".to_string()), &policy)
|
||||
.expect("merged prompt should exist");
|
||||
assert!(merged.contains(REQUEST_TOOL_POLICY_MARKER));
|
||||
assert!(merged.contains("必须先调用"));
|
||||
assert!(merged.contains("WebSearch"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn no_duplicate_when_marker_exists() {
|
||||
let base = Some(format!("{REQUEST_TOOL_POLICY_MARKER}\nexists"));
|
||||
let policy = resolve_request_tool_policy(Some(true), false);
|
||||
assert_eq!(
|
||||
merge_system_prompt_with_request_tool_policy(base.clone(), &policy),
|
||||
base
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn tracker_requires_websearch_when_enabled() {
|
||||
let policy = resolve_request_tool_policy(Some(true), false);
|
||||
let mut tracker = WebSearchExecutionTracker::default();
|
||||
tracker.record_tool_start(&policy, "tool-1", "WebFetch");
|
||||
tracker.record_tool_end(&policy, "tool-1", true, None);
|
||||
let err = tracker
|
||||
.validate_web_search_requirement(&policy)
|
||||
.expect_err("missing web search should fail");
|
||||
assert!(err.contains("未检测到必需工具调用"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn tracker_accepts_successful_websearch() {
|
||||
let policy = resolve_request_tool_policy(Some(true), false);
|
||||
let mut tracker = WebSearchExecutionTracker::default();
|
||||
tracker.record_tool_start(&policy, "tool-1", "WebSearch");
|
||||
tracker.record_tool_end(&policy, "tool-1", true, None);
|
||||
assert!(tracker.validate_web_search_requirement(&policy).is_ok());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn tracker_reports_failure_record() {
|
||||
let policy = resolve_request_tool_policy(Some(true), false);
|
||||
let mut tracker = WebSearchExecutionTracker::default();
|
||||
tracker.record_tool_start(&policy, "tool-1", "WebSearch");
|
||||
tracker.record_tool_end(&policy, "tool-1", false, Some("network timeout"));
|
||||
let err = tracker
|
||||
.validate_web_search_requirement(&policy)
|
||||
.expect_err("failed required tool should fail");
|
||||
assert!(err.contains("network timeout"));
|
||||
assert!(err.contains("尝试记录"));
|
||||
}
|
||||
}
|
||||
@@ -1,7 +1,7 @@
|
||||
{
|
||||
"$schema": "https://schema.tauri.app/config/2",
|
||||
"productName": "ProxyCast",
|
||||
"version": "0.78.0",
|
||||
"version": "0.79.0",
|
||||
"identifier": "com.proxycast.app",
|
||||
"build": {
|
||||
"beforeDevCommand": "npm run dev",
|
||||
|
||||
@@ -0,0 +1,268 @@
|
||||
use futures::StreamExt;
|
||||
use proxycast_agent::{
|
||||
convert_agent_event, AsterAgentState, SessionConfigBuilder, TauriAgentEvent,
|
||||
};
|
||||
use proxycast_core::database::dao::api_key_provider::ApiProviderType;
|
||||
use proxycast_core::database::init_database;
|
||||
use proxycast_lib::services::request_tool_policy_prompt_service::{
|
||||
merge_system_prompt_with_request_tool_policy, resolve_request_tool_policy,
|
||||
WebSearchExecutionTracker,
|
||||
};
|
||||
use proxycast_services::api_key_provider_service::ApiKeyProviderService;
|
||||
use uuid::Uuid;
|
||||
|
||||
fn should_run_real_test() -> bool {
|
||||
std::env::var("PROXYCAST_REAL_API_TEST").ok().as_deref() == Some("1")
|
||||
}
|
||||
|
||||
fn resolve_model_name(
|
||||
explicit: Option<String>,
|
||||
provider_models: &[String],
|
||||
) -> Result<String, String> {
|
||||
if let Some(model) = explicit
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
{
|
||||
return Ok(model.to_string());
|
||||
}
|
||||
|
||||
if let Some(model) = provider_models
|
||||
.iter()
|
||||
.map(|value| value.trim())
|
||||
.find(|value| !value.is_empty())
|
||||
{
|
||||
return Ok(model.to_string());
|
||||
}
|
||||
|
||||
Err(
|
||||
"未找到可用模型:请设置 PROXYCAST_REAL_MODEL,或在 Provider custom_models 中配置模型。"
|
||||
.to_string(),
|
||||
)
|
||||
}
|
||||
|
||||
fn resolve_codex_provider_and_model(
|
||||
db: &proxycast_core::database::DbConnection,
|
||||
) -> Result<(String, String), String> {
|
||||
let explicit_model = std::env::var("PROXYCAST_REAL_MODEL").ok();
|
||||
|
||||
if let Ok(explicit) = std::env::var("PROXYCAST_REAL_PROVIDER_ID") {
|
||||
let trimmed = explicit.trim();
|
||||
if !trimmed.is_empty() {
|
||||
let service = ApiKeyProviderService::new();
|
||||
let provider = service
|
||||
.get_provider(db, trimmed)?
|
||||
.ok_or_else(|| format!("未找到指定 Provider: {trimmed}"))?;
|
||||
let model = resolve_model_name(explicit_model, &provider.provider.custom_models)?;
|
||||
return Ok((trimmed.to_string(), model));
|
||||
}
|
||||
}
|
||||
|
||||
let service = ApiKeyProviderService::new();
|
||||
let providers = service.get_all_providers(db)?;
|
||||
providers
|
||||
.into_iter()
|
||||
.find(|item| {
|
||||
item.provider.enabled
|
||||
&& item.provider.provider_type == ApiProviderType::Codex
|
||||
&& item.api_keys.iter().any(|key| key.enabled)
|
||||
})
|
||||
.map(|item| -> Result<(String, String), String> {
|
||||
let model = resolve_model_name(explicit_model, &item.provider.custom_models)?;
|
||||
Ok((item.provider.id, model))
|
||||
})
|
||||
.transpose()?
|
||||
.ok_or_else(|| "未找到启用且含可用 Key 的 Codex Provider".to_string())
|
||||
}
|
||||
|
||||
#[derive(Debug, Default)]
|
||||
struct RealRunSummary {
|
||||
session_id: String,
|
||||
web_search: bool,
|
||||
model: String,
|
||||
tool_start_count: usize,
|
||||
tool_end_count: usize,
|
||||
web_search_tool_names: Vec<String>,
|
||||
errors: Vec<String>,
|
||||
final_text_preview: String,
|
||||
}
|
||||
|
||||
async fn run_real_case(
|
||||
state: &AsterAgentState,
|
||||
db: &proxycast_core::database::DbConnection,
|
||||
provider_id: &str,
|
||||
model_name: &str,
|
||||
web_search: bool,
|
||||
prompt: &str,
|
||||
) -> Result<RealRunSummary, String> {
|
||||
let session_id = format!("real-web-policy-{}", Uuid::new_v4());
|
||||
state
|
||||
.configure_provider_from_pool(db, provider_id, model_name, &session_id)
|
||||
.await
|
||||
.map_err(|e| format!("配置 Provider 失败: {e}"))?;
|
||||
|
||||
let agent_arc = state.get_agent_arc();
|
||||
let guard = agent_arc.read().await;
|
||||
let agent = guard.as_ref().ok_or_else(|| "Agent 未初始化".to_string())?;
|
||||
|
||||
let policy = resolve_request_tool_policy(Some(web_search), false);
|
||||
let merged_prompt = merge_system_prompt_with_request_tool_policy(None, &policy);
|
||||
let mut session_config_builder = SessionConfigBuilder::new(&session_id);
|
||||
if let Some(system_prompt) = merged_prompt {
|
||||
session_config_builder = session_config_builder.system_prompt(system_prompt);
|
||||
}
|
||||
let session_config = session_config_builder.build();
|
||||
|
||||
let user_message = aster::conversation::message::Message::user().with_text(prompt);
|
||||
let mut stream = agent
|
||||
.reply(user_message, session_config, None)
|
||||
.await
|
||||
.map_err(|e| format!("创建流式回复失败: {e}"))?;
|
||||
|
||||
let mut summary = RealRunSummary {
|
||||
session_id,
|
||||
web_search,
|
||||
model: model_name.to_string(),
|
||||
..RealRunSummary::default()
|
||||
};
|
||||
let mut tracker = WebSearchExecutionTracker::default();
|
||||
let mut text_buffer = String::new();
|
||||
|
||||
while let Some(event_result) = stream.next().await {
|
||||
match event_result {
|
||||
Ok(agent_event) => {
|
||||
for event in convert_agent_event(agent_event) {
|
||||
match &event {
|
||||
TauriAgentEvent::ToolStart {
|
||||
tool_name, tool_id, ..
|
||||
} => {
|
||||
summary.tool_start_count += 1;
|
||||
tracker.record_tool_start(&policy, tool_id, tool_name);
|
||||
if tool_name.to_ascii_lowercase().contains("websearch") {
|
||||
summary.web_search_tool_names.push(tool_name.clone());
|
||||
}
|
||||
}
|
||||
TauriAgentEvent::ToolEnd { tool_id, result } => {
|
||||
summary.tool_end_count += 1;
|
||||
tracker.record_tool_end(
|
||||
&policy,
|
||||
tool_id,
|
||||
result.success,
|
||||
result.error.as_deref(),
|
||||
);
|
||||
}
|
||||
TauriAgentEvent::TextDelta { text } => text_buffer.push_str(text),
|
||||
TauriAgentEvent::Error { message } => summary.errors.push(message.clone()),
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
}
|
||||
Err(error) => summary.errors.push(format!("stream_error: {error}")),
|
||||
}
|
||||
}
|
||||
|
||||
if let Err(error) = tracker.validate_web_search_requirement(&policy) {
|
||||
summary.errors.push(error);
|
||||
}
|
||||
|
||||
summary.final_text_preview = text_buffer.chars().take(280).collect();
|
||||
Ok(summary)
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
#[ignore = "真实联网测试:设置 PROXYCAST_REAL_API_TEST=1 后执行"]
|
||||
async fn test_real_gpt53_codex_web_search_scenarios() {
|
||||
if !should_run_real_test() {
|
||||
return;
|
||||
}
|
||||
|
||||
let db = init_database().expect("初始化数据库失败");
|
||||
let (provider_id, resolved_model) =
|
||||
resolve_codex_provider_and_model(&db).expect("解析 Codex Provider/模型失败");
|
||||
let model_name = std::env::var("PROXYCAST_REAL_MODEL").unwrap_or_else(|_| {
|
||||
if resolved_model.trim().is_empty() {
|
||||
"gpt-5.3-codex".to_string()
|
||||
} else {
|
||||
resolved_model
|
||||
}
|
||||
});
|
||||
|
||||
assert_eq!(
|
||||
model_name.trim(),
|
||||
"gpt-5.3-codex",
|
||||
"本测试仅允许使用 gpt-5.3-codex"
|
||||
);
|
||||
|
||||
let state = AsterAgentState::new();
|
||||
|
||||
let scenario_a = run_real_case(
|
||||
&state,
|
||||
&db,
|
||||
&provider_id,
|
||||
&model_name,
|
||||
false,
|
||||
"场景A:webSearch=false。请简要解释什么是 Rust 的所有权模型。",
|
||||
)
|
||||
.await
|
||||
.expect("场景A调用失败");
|
||||
|
||||
println!(
|
||||
"[ScenarioA] request={{model:{}, web_search:{}, session:{}}} events={{tool_start:{}, tool_end:{}, web_search_tools:{:?}}} errors={:?} final_preview={}",
|
||||
scenario_a.model,
|
||||
scenario_a.web_search,
|
||||
scenario_a.session_id,
|
||||
scenario_a.tool_start_count,
|
||||
scenario_a.tool_end_count,
|
||||
scenario_a.web_search_tool_names,
|
||||
scenario_a.errors,
|
||||
scenario_a.final_text_preview
|
||||
);
|
||||
assert!(
|
||||
scenario_a.errors.is_empty(),
|
||||
"场景A出现错误: {:?}",
|
||||
scenario_a.errors
|
||||
);
|
||||
|
||||
let scenario_b = run_real_case(
|
||||
&state,
|
||||
&db,
|
||||
&provider_id,
|
||||
&model_name,
|
||||
true,
|
||||
"场景B:webSearch=true。请搜索并总结2026年3月4日全球重要新闻,给出来源链接。",
|
||||
)
|
||||
.await
|
||||
.expect("场景B调用失败");
|
||||
|
||||
println!(
|
||||
"[ScenarioB] request={{model:{}, web_search:{}, session:{}}} events={{tool_start:{}, tool_end:{}, web_search_tools:{:?}}} errors={:?} final_preview={}",
|
||||
scenario_b.model,
|
||||
scenario_b.web_search,
|
||||
scenario_b.session_id,
|
||||
scenario_b.tool_start_count,
|
||||
scenario_b.tool_end_count,
|
||||
scenario_b.web_search_tool_names,
|
||||
scenario_b.errors,
|
||||
scenario_b.final_text_preview
|
||||
);
|
||||
|
||||
assert!(
|
||||
scenario_b.errors.is_empty(),
|
||||
"场景B出现错误: {:?}",
|
||||
scenario_b.errors
|
||||
);
|
||||
assert!(
|
||||
scenario_b
|
||||
.web_search_tool_names
|
||||
.iter()
|
||||
.any(|name| name.to_ascii_lowercase().contains("websearch")),
|
||||
"场景B必须包含 WebSearch 工具调用,实际: {:?}",
|
||||
scenario_b.web_search_tool_names
|
||||
);
|
||||
assert!(
|
||||
scenario_b.tool_start_count > 0 && scenario_b.tool_end_count > 0,
|
||||
"场景B必须出现 tool_start/tool_end 事件,实际: start={}, end={}",
|
||||
scenario_b.tool_start_count,
|
||||
scenario_b.tool_end_count
|
||||
);
|
||||
}
|
||||
@@ -0,0 +1,134 @@
|
||||
use proxycast_agent::AsterAgentState;
|
||||
use proxycast_core::database::dao::api_key_provider::ApiProviderType;
|
||||
use proxycast_core::database::init_database;
|
||||
use proxycast_lib::services::request_tool_policy_prompt_service::{
|
||||
execute_web_search_preflight_if_needed, resolve_request_tool_policy, WebSearchExecutionTracker,
|
||||
};
|
||||
use proxycast_services::api_key_provider_service::ApiKeyProviderService;
|
||||
use uuid::Uuid;
|
||||
|
||||
fn should_run_real_test() -> bool {
|
||||
std::env::var("PROXYCAST_REAL_API_TEST").ok().as_deref() == Some("1")
|
||||
}
|
||||
|
||||
fn resolve_model_name(
|
||||
explicit: Option<String>,
|
||||
provider_models: &[String],
|
||||
) -> Result<String, String> {
|
||||
if let Some(model) = explicit
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
{
|
||||
return Ok(model.to_string());
|
||||
}
|
||||
|
||||
if let Some(model) = provider_models
|
||||
.iter()
|
||||
.map(|value| value.trim())
|
||||
.find(|value| !value.is_empty())
|
||||
{
|
||||
return Ok(model.to_string());
|
||||
}
|
||||
|
||||
Err("未找到可用模型".to_string())
|
||||
}
|
||||
|
||||
fn resolve_codex_provider_and_model(
|
||||
db: &proxycast_core::database::DbConnection,
|
||||
) -> Result<(String, String), String> {
|
||||
let explicit_model = std::env::var("PROXYCAST_REAL_MODEL").ok();
|
||||
let service = ApiKeyProviderService::new();
|
||||
let providers = service.get_all_providers(db)?;
|
||||
|
||||
providers
|
||||
.into_iter()
|
||||
.find(|item| {
|
||||
item.provider.enabled
|
||||
&& item.provider.provider_type == ApiProviderType::Codex
|
||||
&& item.api_keys.iter().any(|key| key.enabled)
|
||||
})
|
||||
.map(|item| -> Result<(String, String), String> {
|
||||
let model = resolve_model_name(explicit_model, &item.provider.custom_models)?;
|
||||
Ok((item.provider.id, model))
|
||||
})
|
||||
.transpose()?
|
||||
.ok_or_else(|| "未找到启用且含可用 Key 的 Codex Provider".to_string())
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
#[ignore = "真实联网测试:设置 PROXYCAST_REAL_API_TEST=1 后执行"]
|
||||
async fn test_real_web_search_preflight_short_input_continue() {
|
||||
if !should_run_real_test() {
|
||||
return;
|
||||
}
|
||||
|
||||
let db = init_database().expect("初始化数据库失败");
|
||||
let (provider_id, resolved_model) =
|
||||
resolve_codex_provider_and_model(&db).expect("解析 Codex Provider/模型失败");
|
||||
let model_name = std::env::var("PROXYCAST_REAL_MODEL").unwrap_or(resolved_model);
|
||||
assert_eq!(
|
||||
model_name.trim(),
|
||||
"gpt-5.3-codex",
|
||||
"本测试仅允许使用 gpt-5.3-codex"
|
||||
);
|
||||
|
||||
let state = AsterAgentState::new();
|
||||
let session_id = format!("real-web-preflight-{}", Uuid::new_v4());
|
||||
state
|
||||
.configure_provider_from_pool(&db, &provider_id, &model_name, &session_id)
|
||||
.await
|
||||
.expect("配置 Provider 失败");
|
||||
|
||||
let agent_arc = state.get_agent_arc();
|
||||
let guard = agent_arc.read().await;
|
||||
let agent = guard.as_ref().expect("Agent 未初始化");
|
||||
|
||||
let policy = resolve_request_tool_policy(Some(true), false);
|
||||
let mut tracker = WebSearchExecutionTracker::default();
|
||||
let execution = execute_web_search_preflight_if_needed(
|
||||
agent,
|
||||
&session_id,
|
||||
"继续",
|
||||
None,
|
||||
None,
|
||||
&policy,
|
||||
&mut tracker,
|
||||
)
|
||||
.await
|
||||
.expect("预调用失败");
|
||||
|
||||
let mut tool_start_count = 0usize;
|
||||
let mut tool_end_count = 0usize;
|
||||
let mut tool_names = Vec::new();
|
||||
for event in execution.events {
|
||||
match event {
|
||||
proxycast_agent::TauriAgentEvent::ToolStart { tool_name, .. } => {
|
||||
tool_start_count += 1;
|
||||
tool_names.push(tool_name);
|
||||
}
|
||||
proxycast_agent::TauriAgentEvent::ToolEnd { .. } => {
|
||||
tool_end_count += 1;
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
println!(
|
||||
"[PreflightContinue] request={{model:{}, web_search:true, prompt:\"继续\", session:{}}} events={{tool_start:{}, tool_end:{}, tools:{:?}}}",
|
||||
model_name, session_id, tool_start_count, tool_end_count, tool_names
|
||||
);
|
||||
|
||||
assert!(
|
||||
tool_names
|
||||
.iter()
|
||||
.any(|name| name.to_ascii_lowercase().contains("websearch")),
|
||||
"预调用必须包含 WebSearch,实际: {:?}",
|
||||
tool_names
|
||||
);
|
||||
assert!(tool_start_count > 0, "必须出现 tool_start");
|
||||
assert!(tool_end_count > 0, "必须出现 tool_end");
|
||||
tracker
|
||||
.validate_web_search_requirement(&policy)
|
||||
.expect("预调用后应满足必需工具约束");
|
||||
}
|
||||
Reference in New Issue
Block a user