mirror of
https://github.com/aiclientproxy/proxycast.git
synced 2026-09-24 23:10:56 +08:00
fix: stabilize v1.7.0
This commit is contained in:
@@ -2,6 +2,7 @@ use aster::conversation::message::{Message, MessageContent};
|
||||
use aster::model::ModelConfig;
|
||||
use aster::providers::base::{
|
||||
LeadWorkerProviderTrait, MessageStream, Provider, ProviderMetadata, ProviderUsage,
|
||||
SessionNameGenerationExecutionStrategy,
|
||||
};
|
||||
use aster::providers::errors::ProviderError;
|
||||
use aster::providers::RetryConfig;
|
||||
@@ -201,6 +202,17 @@ impl Provider for ProviderSafety {
|
||||
self.inner.get_active_model_name()
|
||||
}
|
||||
|
||||
async fn generate_session_name(
|
||||
&self,
|
||||
messages: &aster::conversation::Conversation,
|
||||
) -> Result<String, ProviderError> {
|
||||
self.inner.generate_session_name(messages).await
|
||||
}
|
||||
|
||||
fn session_name_generation_execution_strategy(&self) -> SessionNameGenerationExecutionStrategy {
|
||||
self.inner.session_name_generation_execution_strategy()
|
||||
}
|
||||
|
||||
async fn configure_oauth(&self) -> Result<(), ProviderError> {
|
||||
self.inner.configure_oauth().await
|
||||
}
|
||||
@@ -212,8 +224,11 @@ mod tests {
|
||||
normalize_provider_messages, normalize_provider_model_config, wrap_provider_with_safety,
|
||||
};
|
||||
use aster::conversation::message::{Message, MessageContent};
|
||||
use aster::conversation::Conversation;
|
||||
use aster::model::ModelConfig;
|
||||
use aster::providers::base::{Provider, ProviderMetadata, ProviderUsage, Usage};
|
||||
use aster::providers::base::{
|
||||
Provider, ProviderMetadata, ProviderUsage, SessionNameGenerationExecutionStrategy, Usage,
|
||||
};
|
||||
use aster::providers::errors::ProviderError;
|
||||
use async_trait::async_trait;
|
||||
use rmcp::model::{CallToolRequestParam, CallToolResult, ErrorCode, ErrorData, Tool};
|
||||
@@ -323,6 +338,55 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
struct SessionNamingProvider {
|
||||
model_config: ModelConfig,
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl Provider for SessionNamingProvider {
|
||||
fn metadata() -> ProviderMetadata
|
||||
where
|
||||
Self: Sized,
|
||||
{
|
||||
ProviderMetadata::empty()
|
||||
}
|
||||
|
||||
fn get_name(&self) -> &str {
|
||||
"session-naming"
|
||||
}
|
||||
|
||||
async fn complete_with_model(
|
||||
&self,
|
||||
model_config: &ModelConfig,
|
||||
_system: &str,
|
||||
_messages: &[Message],
|
||||
_tools: &[Tool],
|
||||
) -> Result<(Message, ProviderUsage), ProviderError> {
|
||||
Ok((
|
||||
Message::assistant().with_text("ok"),
|
||||
ProviderUsage::new(model_config.model_name.clone(), Usage::default()),
|
||||
))
|
||||
}
|
||||
|
||||
fn get_model_config(&self) -> ModelConfig {
|
||||
self.model_config.clone()
|
||||
}
|
||||
|
||||
async fn generate_session_name(
|
||||
&self,
|
||||
_messages: &aster::conversation::Conversation,
|
||||
) -> Result<String, ProviderError> {
|
||||
Ok("wrapped-title".to_string())
|
||||
}
|
||||
|
||||
fn session_name_generation_execution_strategy(
|
||||
&self,
|
||||
) -> SessionNameGenerationExecutionStrategy {
|
||||
SessionNameGenerationExecutionStrategy::AfterReply
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn wrap_provider_with_safety_should_disable_fast_model_for_complete_fast() {
|
||||
let seen_models = Arc::new(Mutex::new(Vec::new()));
|
||||
@@ -370,6 +434,36 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn wrap_provider_with_safety_should_forward_session_name_generation() {
|
||||
let provider = Arc::new(SessionNamingProvider {
|
||||
model_config: ModelConfig::new("deepseek-r1:latest").expect("create model config"),
|
||||
});
|
||||
let wrapped = wrap_provider_with_safety(provider, false);
|
||||
let messages =
|
||||
Conversation::new(vec![Message::user().with_text("你好")]).expect("conversation");
|
||||
|
||||
let generated = wrapped
|
||||
.generate_session_name(&messages)
|
||||
.await
|
||||
.expect("generate session name");
|
||||
|
||||
assert_eq!(generated, "wrapped-title");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn wrap_provider_with_safety_should_forward_session_name_strategy() {
|
||||
let provider = Arc::new(SessionNamingProvider {
|
||||
model_config: ModelConfig::new("deepseek-r1:latest").expect("create model config"),
|
||||
});
|
||||
let wrapped = wrap_provider_with_safety(provider, false);
|
||||
|
||||
assert_eq!(
|
||||
wrapped.session_name_generation_execution_strategy(),
|
||||
SessionNameGenerationExecutionStrategy::AfterReply
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn normalize_provider_messages_should_remove_orphan_tool_response() {
|
||||
let messages = vec![
|
||||
|
||||
@@ -40,7 +40,7 @@ use crate::model::ModelConfig;
|
||||
use crate::permission::permission_inspector::PermissionInspector;
|
||||
use crate::permission::permission_judge::PermissionCheckResult;
|
||||
use crate::permission::PermissionConfirmation;
|
||||
use crate::providers::base::Provider;
|
||||
use crate::providers::base::{Provider, SessionNameGenerationExecutionStrategy};
|
||||
use crate::providers::errors::ProviderError;
|
||||
use crate::recipe::{Author, Recipe, Response, Settings, SubRecipe};
|
||||
use crate::scheduler_trait::SchedulerTrait;
|
||||
@@ -3330,17 +3330,26 @@ impl Agent {
|
||||
let provider = self.provider().await?;
|
||||
let session_for_name = session.clone().without_messages();
|
||||
let conversation_for_name = conversation.clone();
|
||||
tokio::spawn(async move {
|
||||
if let Err(e) = SessionManager::maybe_update_name_for_session(
|
||||
&session_for_name,
|
||||
&conversation_for_name,
|
||||
provider,
|
||||
)
|
||||
.await
|
||||
{
|
||||
warn!("Failed to generate session description: {}", e);
|
||||
}
|
||||
});
|
||||
let deferred_session_name_generation =
|
||||
match provider.session_name_generation_execution_strategy() {
|
||||
SessionNameGenerationExecutionStrategy::Background => {
|
||||
tokio::spawn(async move {
|
||||
if let Err(e) = SessionManager::maybe_update_name_for_session(
|
||||
&session_for_name,
|
||||
&conversation_for_name,
|
||||
provider,
|
||||
)
|
||||
.await
|
||||
{
|
||||
warn!("Failed to generate session description: {}", e);
|
||||
}
|
||||
});
|
||||
None
|
||||
}
|
||||
SessionNameGenerationExecutionStrategy::AfterReply => {
|
||||
Some((session_for_name, conversation_for_name, provider))
|
||||
}
|
||||
};
|
||||
let working_dir = session.working_dir.clone();
|
||||
|
||||
Ok(Box::pin(async_stream::try_stream! {
|
||||
@@ -3801,6 +3810,22 @@ impl Agent {
|
||||
|
||||
tokio::task::yield_now().await;
|
||||
}
|
||||
|
||||
if let Some((session_for_name, conversation_for_name, provider)) =
|
||||
deferred_session_name_generation
|
||||
{
|
||||
tokio::spawn(async move {
|
||||
if let Err(e) = SessionManager::maybe_update_name_for_session(
|
||||
&session_for_name,
|
||||
&conversation_for_name,
|
||||
provider,
|
||||
)
|
||||
.await
|
||||
{
|
||||
warn!("Failed to generate session description: {}", e);
|
||||
}
|
||||
});
|
||||
}
|
||||
}))
|
||||
}
|
||||
|
||||
|
||||
@@ -186,6 +186,31 @@ pub fn should_bypass_proxy(target_url: &str, no_proxy: &[String]) -> bool {
|
||||
false
|
||||
}
|
||||
|
||||
/// 检查目标 URL 是否应直接绕过系统代理。
|
||||
///
|
||||
/// 主要用于本地 loopback / unspecified 地址,避免本机服务请求被系统代理截流。
|
||||
pub fn should_bypass_system_proxy_for_url(target_url: &str) -> bool {
|
||||
let Ok(url) = Url::parse(target_url) else {
|
||||
return false;
|
||||
};
|
||||
|
||||
let Some(hostname) = url.host_str() else {
|
||||
return false;
|
||||
};
|
||||
|
||||
if matches!(
|
||||
hostname,
|
||||
"localhost" | "127.0.0.1" | "::1" | "0.0.0.0" | "host.docker.internal"
|
||||
) {
|
||||
return true;
|
||||
}
|
||||
|
||||
hostname
|
||||
.parse::<std::net::IpAddr>()
|
||||
.map(|ip| ip.is_loopback() || ip.is_unspecified())
|
||||
.unwrap_or(false)
|
||||
}
|
||||
|
||||
/// 获取目标 URL 的代理 URL
|
||||
pub fn get_proxy_for_url(target_url: &str, config: &ProxyConfig) -> Option<String> {
|
||||
// 检查是否绕过代理
|
||||
|
||||
@@ -48,6 +48,25 @@ fn test_should_bypass_proxy_all() {
|
||||
assert!(should_bypass_proxy("http://any.domain.com", &no_proxy));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_should_bypass_system_proxy_for_loopback_url() {
|
||||
assert!(should_bypass_system_proxy_for_url(
|
||||
"http://127.0.0.1:11434/api/chat"
|
||||
));
|
||||
assert!(should_bypass_system_proxy_for_url(
|
||||
"http://localhost:11434/api/tags"
|
||||
));
|
||||
assert!(should_bypass_system_proxy_for_url("http://0.0.0.0:3000"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_should_not_bypass_system_proxy_for_remote_url() {
|
||||
assert!(!should_bypass_system_proxy_for_url(
|
||||
"https://api.openai.com/v1/chat/completions"
|
||||
));
|
||||
assert!(!should_bypass_system_proxy_for_url("https://example.com"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_timeout_config_default() {
|
||||
let config = TimeoutConfig::default();
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
use crate::network::should_bypass_system_proxy_for_url;
|
||||
use crate::session_context::SESSION_ID_HEADER;
|
||||
use anyhow::Result;
|
||||
use async_trait::async_trait;
|
||||
@@ -205,6 +206,10 @@ impl ApiClient {
|
||||
|
||||
pub fn with_timeout(host: String, auth: AuthMethod, timeout: Duration) -> Result<Self> {
|
||||
let mut client_builder = Client::builder().timeout(timeout);
|
||||
if should_bypass_system_proxy_for_url(&host) {
|
||||
tracing::info!("[ApiClient] 本地地址绕过系统代理: {}", host);
|
||||
client_builder = client_builder.no_proxy();
|
||||
}
|
||||
|
||||
// Configure TLS if needed
|
||||
let tls_config = TlsConfig::from_config()?;
|
||||
@@ -228,6 +233,10 @@ impl ApiClient {
|
||||
let mut client_builder = Client::builder()
|
||||
.timeout(self.timeout)
|
||||
.default_headers(self.default_headers.clone());
|
||||
if should_bypass_system_proxy_for_url(&self.host) {
|
||||
tracing::info!("[ApiClient] 重建客户端时绕过系统代理: {}", self.host);
|
||||
client_builder = client_builder.no_proxy();
|
||||
}
|
||||
|
||||
// Configure TLS if needed
|
||||
if let Some(ref tls_config) = self.tls_config {
|
||||
@@ -408,6 +417,23 @@ impl fmt::Debug for ApiClient {
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn should_bypass_proxy_for_loopback_host() {
|
||||
assert!(should_bypass_system_proxy_for_url("http://127.0.0.1:11434"));
|
||||
assert!(should_bypass_system_proxy_for_url(
|
||||
"http://localhost:11434/api"
|
||||
));
|
||||
assert!(should_bypass_system_proxy_for_url("http://0.0.0.0:1234"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn should_not_bypass_proxy_for_remote_host() {
|
||||
assert!(!should_bypass_system_proxy_for_url(
|
||||
"https://api.openai.com/v1"
|
||||
));
|
||||
assert!(!should_bypass_system_proxy_for_url("https://example.com"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_session_id_header_injection() {
|
||||
let client = ApiClient::new(
|
||||
|
||||
@@ -35,6 +35,12 @@ pub fn get_current_model() -> Option<String> {
|
||||
|
||||
pub static MSG_COUNT_FOR_SESSION_NAME_GENERATION: usize = 3;
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum SessionNameGenerationExecutionStrategy {
|
||||
Background,
|
||||
AfterReply,
|
||||
}
|
||||
|
||||
/// Information about a model's capabilities
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, ToSchema, PartialEq)]
|
||||
pub struct ModelInfo {
|
||||
@@ -577,6 +583,10 @@ pub trait Provider: Send + Sync {
|
||||
Ok(safe_truncate(&description, 100))
|
||||
}
|
||||
|
||||
fn session_name_generation_execution_strategy(&self) -> SessionNameGenerationExecutionStrategy {
|
||||
SessionNameGenerationExecutionStrategy::Background
|
||||
}
|
||||
|
||||
// Generate a prompt for a session name based on the conversation history
|
||||
fn create_session_name_prompt(&self, context: &[String]) -> String {
|
||||
// Create a prompt for a concise description
|
||||
|
||||
@@ -1,5 +1,8 @@
|
||||
use super::api_client::{ApiClient, AuthMethod};
|
||||
use super::base::{ConfigKey, MessageStream, Provider, ProviderMetadata, ProviderUsage, Usage};
|
||||
use super::base::{
|
||||
ConfigKey, MessageStream, Provider, ProviderMetadata, ProviderUsage,
|
||||
SessionNameGenerationExecutionStrategy, Usage,
|
||||
};
|
||||
use super::errors::ProviderError;
|
||||
use super::retry::ProviderRetry;
|
||||
use super::utils::{
|
||||
@@ -278,6 +281,10 @@ impl Provider for OllamaProvider {
|
||||
Ok(safe_truncate(&description, 100))
|
||||
}
|
||||
|
||||
fn session_name_generation_execution_strategy(&self) -> SessionNameGenerationExecutionStrategy {
|
||||
SessionNameGenerationExecutionStrategy::AfterReply
|
||||
}
|
||||
|
||||
fn supports_streaming(&self) -> bool {
|
||||
self.supports_streaming
|
||||
}
|
||||
|
||||
@@ -37,6 +37,7 @@ use super::ollama::OLLAMA_HOST;
|
||||
use crate::conversation::message::{Message, MessageContent};
|
||||
use crate::conversation::Conversation;
|
||||
use crate::model::ModelConfig;
|
||||
use crate::network::should_bypass_system_proxy_for_url;
|
||||
use crate::providers::formats::openai::create_request;
|
||||
use anyhow::Result;
|
||||
use reqwest::Client;
|
||||
@@ -76,13 +77,19 @@ impl OllamaInterpreter {
|
||||
}
|
||||
|
||||
pub fn new_with_model(interpreter_model: Option<String>) -> Result<Self, ProviderError> {
|
||||
let client = Client::builder()
|
||||
.timeout(Duration::from_secs(600))
|
||||
let base_url = Self::get_ollama_base_url()?;
|
||||
let mut client_builder = Client::builder().timeout(Duration::from_secs(600));
|
||||
if should_bypass_system_proxy_for_url(&base_url) {
|
||||
tracing::info!(
|
||||
"[ToolShim] 本地 Ollama 结构化请求绕过系统代理: {}",
|
||||
base_url
|
||||
);
|
||||
client_builder = client_builder.no_proxy();
|
||||
}
|
||||
let client = client_builder
|
||||
.build()
|
||||
.expect("Failed to create HTTP client");
|
||||
|
||||
let base_url = Self::get_ollama_base_url()?;
|
||||
|
||||
Ok(Self {
|
||||
client,
|
||||
base_url,
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
use super::dto::ConfigureProviderRequest;
|
||||
use aster::network::should_bypass_system_proxy_for_url;
|
||||
use lime_core::models::model_registry::ModelCapabilities;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::collections::HashSet;
|
||||
@@ -98,6 +99,17 @@ fn default_runtime_model_capabilities() -> ModelCapabilities {
|
||||
}
|
||||
}
|
||||
|
||||
fn conservative_ollama_fallback_capabilities(fallback: &ModelCapabilities) -> ModelCapabilities {
|
||||
ModelCapabilities {
|
||||
vision: fallback.vision,
|
||||
tools: false,
|
||||
streaming: true,
|
||||
json_mode: false,
|
||||
function_calling: false,
|
||||
reasoning: fallback.reasoning,
|
||||
}
|
||||
}
|
||||
|
||||
fn build_runtime_tool_call_decision(
|
||||
provider_selector: Option<&str>,
|
||||
provider_name: &str,
|
||||
@@ -229,10 +241,16 @@ pub(crate) async fn resolve_runtime_tool_call_decision(
|
||||
.cloned()
|
||||
.unwrap_or_else(default_runtime_model_capabilities);
|
||||
let base_url = normalize_ollama_base_url(base_url);
|
||||
let client = match reqwest::Client::builder()
|
||||
.timeout(Duration::from_secs(OLLAMA_RUNTIME_PROBE_TIMEOUT_SECS))
|
||||
.build()
|
||||
{
|
||||
let mut client_builder =
|
||||
reqwest::Client::builder().timeout(Duration::from_secs(OLLAMA_RUNTIME_PROBE_TIMEOUT_SECS));
|
||||
if should_bypass_system_proxy_for_url(&base_url) {
|
||||
tracing::info!(
|
||||
"[AsterAgent] Ollama 运行时能力探测绕过系统代理: {}",
|
||||
base_url
|
||||
);
|
||||
client_builder = client_builder.no_proxy();
|
||||
}
|
||||
let client = match client_builder.build() {
|
||||
Ok(client) => client,
|
||||
Err(error) => {
|
||||
tracing::warn!(
|
||||
@@ -243,8 +261,8 @@ pub(crate) async fn resolve_runtime_tool_call_decision(
|
||||
provider_selector,
|
||||
provider_name,
|
||||
model_name,
|
||||
fallback_capabilities,
|
||||
None,
|
||||
conservative_ollama_fallback_capabilities(&fallback_capabilities),
|
||||
Some(model_name.to_string()),
|
||||
);
|
||||
}
|
||||
};
|
||||
@@ -258,7 +276,7 @@ pub(crate) async fn resolve_runtime_tool_call_decision(
|
||||
.await
|
||||
else {
|
||||
tracing::warn!(
|
||||
"[AsterAgent] 读取 Ollama 模型能力失败,回退 catalog 能力: model={}, base_url={}",
|
||||
"[AsterAgent] 读取 Ollama 模型能力失败,保守降级到 toolshim: model={}, base_url={}",
|
||||
model_name,
|
||||
base_url
|
||||
);
|
||||
@@ -266,8 +284,8 @@ pub(crate) async fn resolve_runtime_tool_call_decision(
|
||||
provider_selector,
|
||||
provider_name,
|
||||
model_name,
|
||||
fallback_capabilities,
|
||||
None,
|
||||
conservative_ollama_fallback_capabilities(&fallback_capabilities),
|
||||
Some(model_name.to_string()),
|
||||
);
|
||||
};
|
||||
|
||||
@@ -379,6 +397,27 @@ mod tests {
|
||||
assert!(parsed.json_mode);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn conservative_ollama_fallback_disables_native_tool_flags() {
|
||||
let fallback = ModelCapabilities {
|
||||
vision: true,
|
||||
tools: true,
|
||||
streaming: true,
|
||||
json_mode: true,
|
||||
function_calling: true,
|
||||
reasoning: true,
|
||||
};
|
||||
|
||||
let parsed = conservative_ollama_fallback_capabilities(&fallback);
|
||||
|
||||
assert!(parsed.vision);
|
||||
assert!(parsed.streaming);
|
||||
assert!(parsed.reasoning);
|
||||
assert!(!parsed.tools);
|
||||
assert!(!parsed.function_calling);
|
||||
assert!(!parsed.json_mode);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn enrich_provider_config_sets_runtime_strategy_fields() {
|
||||
let mut provider_config = ConfigureProviderRequest {
|
||||
|
||||
Reference in New Issue
Block a user