fix: stabilize v1.7.0

This commit is contained in:
coso
2026-04-11 02:53:07 +08:00
parent 26751bc117
commit 9412768f39
9 changed files with 279 additions and 27 deletions
+95 -1
View File
@@ -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 {