mirror of
https://github.com/aiclientproxy/proxycast.git
synced 2026-09-24 23:10:56 +08:00
release: v1.7.0
This commit is contained in:
@@ -406,7 +406,10 @@ fn build_provider_env_vars(
|
||||
]
|
||||
}
|
||||
// Ollama 本地部署
|
||||
"ollama" => vec![("OLLAMA_BASE_URL".to_string(), api_host.to_string())],
|
||||
"ollama" => vec![
|
||||
("OLLAMA_BASE_URL".to_string(), api_host.to_string()),
|
||||
("OLLAMA_HOST".to_string(), api_host.to_string()),
|
||||
],
|
||||
_ => {
|
||||
// 未知类型尽量按已注册 ApiProviderType 的协议族生成环境变量
|
||||
if let Ok(api_type) = provider_type.parse::<ApiProviderType>() {
|
||||
|
||||
@@ -39,13 +39,17 @@ pub async fn aster_agent_init(
|
||||
pub async fn aster_agent_configure_provider(
|
||||
state: State<'_, AsterAgentState>,
|
||||
db: State<'_, DbConnection>,
|
||||
request: ConfigureProviderRequest,
|
||||
mut request: ConfigureProviderRequest,
|
||||
session_id: String,
|
||||
) -> Result<AsterAgentStatus, String> {
|
||||
let runtime_tool_call_decision =
|
||||
enrich_provider_config_with_runtime_tool_strategy(&mut request).await;
|
||||
tracing::info!(
|
||||
"[AsterAgent] 配置 Provider: {} / {}",
|
||||
"[AsterAgent] 配置 Provider: {} / {},tool_call_strategy={:?},toolshim_model={:?}",
|
||||
request.provider_name,
|
||||
request.model_name
|
||||
request.model_name,
|
||||
runtime_tool_call_decision.strategy,
|
||||
runtime_tool_call_decision.toolshim_model
|
||||
);
|
||||
|
||||
let provider_selector = request
|
||||
@@ -61,6 +65,11 @@ pub async fn aster_agent_configure_provider(
|
||||
credential_uuid: None,
|
||||
force_responses_api: false,
|
||||
credential_path: None,
|
||||
toolshim: matches!(
|
||||
request.tool_call_strategy,
|
||||
Some(RuntimeToolCallStrategy::ToolShim)
|
||||
),
|
||||
toolshim_model: request.toolshim_model.clone(),
|
||||
};
|
||||
|
||||
state
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
use super::*;
|
||||
use chrono::{DateTime, Utc};
|
||||
use lime_core::models::model_registry::ModelCapabilities;
|
||||
|
||||
/// Aster Agent 状态信息
|
||||
#[derive(Debug, Serialize)]
|
||||
@@ -26,6 +27,12 @@ pub struct ConfigureProviderRequest {
|
||||
pub api_key: Option<String>,
|
||||
#[serde(default)]
|
||||
pub base_url: Option<String>,
|
||||
#[serde(default, alias = "modelCapabilities")]
|
||||
pub model_capabilities: Option<ModelCapabilities>,
|
||||
#[serde(default, alias = "toolCallStrategy")]
|
||||
pub tool_call_strategy: Option<RuntimeToolCallStrategy>,
|
||||
#[serde(default, alias = "toolshimModel")]
|
||||
pub toolshim_model: Option<String>,
|
||||
}
|
||||
|
||||
/// 从凭证池配置 Provider 的请求
|
||||
|
||||
@@ -270,6 +270,7 @@ mod mcp_bridge;
|
||||
mod pdf_read_skill_launch;
|
||||
mod presentation_skill_launch;
|
||||
mod prompt_context;
|
||||
mod provider_runtime_strategy;
|
||||
mod reply_runtime;
|
||||
mod report_skill_launch;
|
||||
mod request_model_resolution;
|
||||
@@ -401,6 +402,9 @@ pub(crate) use prompt_context::{
|
||||
merge_system_prompt_with_service_skill_launch_preload,
|
||||
merge_system_prompt_with_team_preference,
|
||||
};
|
||||
pub(crate) use provider_runtime_strategy::{
|
||||
enrich_provider_config_with_runtime_tool_strategy, RuntimeToolCallStrategy,
|
||||
};
|
||||
#[cfg(test)]
|
||||
use reply_runtime::message_suggests_live_search;
|
||||
use reply_runtime::{
|
||||
|
||||
@@ -579,6 +579,48 @@ fn build_markdown_bundle_translation_followup(
|
||||
lines
|
||||
}
|
||||
|
||||
fn build_markdown_bundle_source_material_followup(
|
||||
execution: &ServiceSkillLaunchPreloadExecution,
|
||||
) -> Vec<String> {
|
||||
if !execution.result.ok {
|
||||
return Vec::new();
|
||||
}
|
||||
|
||||
let Some(saved_content) = execution.result.saved_content.as_ref() else {
|
||||
return Vec::new();
|
||||
};
|
||||
let Some(markdown_relative_path) = saved_content
|
||||
.markdown_relative_path
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
else {
|
||||
return Vec::new();
|
||||
};
|
||||
|
||||
let export_kind = execution
|
||||
.result
|
||||
.data
|
||||
.as_ref()
|
||||
.and_then(serde_json::Value::as_object)
|
||||
.and_then(|data| extract_object_string(data, &["export_kind", "exportKind"]));
|
||||
if export_kind.as_deref() != Some("markdown_bundle") {
|
||||
return Vec::new();
|
||||
}
|
||||
|
||||
let markdown_path = build_prompt_file_path(
|
||||
saved_content.project_root_path.as_deref(),
|
||||
markdown_relative_path,
|
||||
);
|
||||
vec![
|
||||
format!("- 当前已导出的 Markdown bundle 位于 {markdown_path}。"),
|
||||
"- 这个 bundle 是系统侧采集得到的源材料,不默认等于用户要的最终交付结果。".to_string(),
|
||||
"- 如果用户当前目标是继续提炼、分析、改写、生成技能包、报告、方案、脚本或其他正式成果,必须先基于这份已保存的 Markdown 继续完成原任务,而不是只重复导出成功摘要。".to_string(),
|
||||
"- 后续处理优先使用本地文件工具(Read / Write / Edit / Glob)围绕已保存 bundle 展开;不要再次抓站点,也不要把保存路径、图片数量或采集摘要原样复述后就停止。".to_string(),
|
||||
"- 除非当前任务本身就是翻译、校对或回写源 Markdown,否则不要把 exports 下的源 bundle 直接当成最终结果目录;需要新增正式结果时,应另外写出真正可交付的工作区文件。".to_string(),
|
||||
]
|
||||
}
|
||||
|
||||
fn build_service_skill_launch_preload_prompt(
|
||||
execution: &ServiceSkillLaunchPreloadExecution,
|
||||
) -> String {
|
||||
@@ -652,6 +694,7 @@ fn build_service_skill_launch_preload_prompt(
|
||||
format!("- 已预执行请求(JSON):{request_json}。"),
|
||||
format!("- 已预执行结果(JSON):{result_json}。"),
|
||||
];
|
||||
lines.extend(build_markdown_bundle_source_material_followup(execution));
|
||||
lines.extend(build_markdown_bundle_translation_followup(execution));
|
||||
lines.push(
|
||||
"- 除非用户明确要求“重跑一次 / 换关键词 / 换筛选条件 / 重新抓取”,否则本回合不要再次调用任何站点执行工具。"
|
||||
|
||||
@@ -0,0 +1,409 @@
|
||||
use super::dto::ConfigureProviderRequest;
|
||||
use lime_core::models::model_registry::ModelCapabilities;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::collections::HashSet;
|
||||
use std::time::Duration;
|
||||
use url::Url;
|
||||
|
||||
const OLLAMA_RUNTIME_PROBE_TIMEOUT_SECS: u64 = 5;
|
||||
|
||||
#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub(crate) enum RuntimeToolCallStrategy {
|
||||
Native,
|
||||
ToolShim,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub(crate) struct RuntimeToolCallDecision {
|
||||
pub capabilities: ModelCapabilities,
|
||||
pub strategy: RuntimeToolCallStrategy,
|
||||
pub toolshim_model: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct OllamaShowResponse {
|
||||
#[serde(default)]
|
||||
capabilities: Vec<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct OllamaTagsResponse {
|
||||
#[serde(default)]
|
||||
models: Vec<OllamaTagModel>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct OllamaTagModel {
|
||||
name: String,
|
||||
}
|
||||
|
||||
fn normalize_provider_identity(value: &str) -> String {
|
||||
value.trim().to_ascii_lowercase()
|
||||
}
|
||||
|
||||
fn is_ollama_provider(provider_selector: Option<&str>, provider_name: &str) -> bool {
|
||||
provider_selector
|
||||
.map(normalize_provider_identity)
|
||||
.or_else(|| Some(normalize_provider_identity(provider_name)))
|
||||
.is_some_and(|identity| identity == "ollama")
|
||||
}
|
||||
|
||||
fn normalize_optional_text(value: Option<&str>) -> Option<String> {
|
||||
let trimmed = value?.trim();
|
||||
if trimmed.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(trimmed.to_string())
|
||||
}
|
||||
}
|
||||
|
||||
fn normalize_ollama_base_url(base_url: Option<&str>) -> String {
|
||||
let candidate =
|
||||
normalize_optional_text(base_url).unwrap_or_else(|| "http://127.0.0.1:11434".to_string());
|
||||
let raw = if candidate.starts_with("http://") || candidate.starts_with("https://") {
|
||||
candidate
|
||||
} else {
|
||||
format!("http://{candidate}")
|
||||
};
|
||||
|
||||
let mut parsed = Url::parse(&raw).unwrap_or_else(|_| {
|
||||
Url::parse("http://127.0.0.1:11434").expect("hardcoded ollama fallback url is valid")
|
||||
});
|
||||
|
||||
if matches!(parsed.host_str(), Some("localhost")) {
|
||||
let _ = parsed.set_host(Some("127.0.0.1"));
|
||||
}
|
||||
if parsed.port().is_none() && parsed.scheme() == "http" {
|
||||
let _ = parsed.set_port(Some(11434));
|
||||
}
|
||||
|
||||
let trimmed_path = parsed.path().trim_end_matches('/').to_string();
|
||||
if trimmed_path.is_empty() {
|
||||
parsed.set_path("");
|
||||
} else {
|
||||
parsed.set_path(&trimmed_path);
|
||||
}
|
||||
parsed.to_string().trim_end_matches('/').to_string()
|
||||
}
|
||||
|
||||
fn default_runtime_model_capabilities() -> ModelCapabilities {
|
||||
ModelCapabilities {
|
||||
vision: false,
|
||||
tools: true,
|
||||
streaming: true,
|
||||
json_mode: true,
|
||||
function_calling: true,
|
||||
reasoning: false,
|
||||
}
|
||||
}
|
||||
|
||||
fn build_runtime_tool_call_decision(
|
||||
provider_selector: Option<&str>,
|
||||
provider_name: &str,
|
||||
model_name: &str,
|
||||
capabilities: ModelCapabilities,
|
||||
toolshim_model: Option<String>,
|
||||
) -> RuntimeToolCallDecision {
|
||||
let supports_native_tools = capabilities.tools || capabilities.function_calling;
|
||||
if is_ollama_provider(provider_selector, provider_name) && !supports_native_tools {
|
||||
return RuntimeToolCallDecision {
|
||||
capabilities,
|
||||
strategy: RuntimeToolCallStrategy::ToolShim,
|
||||
toolshim_model: Some(toolshim_model.unwrap_or_else(|| model_name.to_string())),
|
||||
};
|
||||
}
|
||||
|
||||
RuntimeToolCallDecision {
|
||||
capabilities,
|
||||
strategy: RuntimeToolCallStrategy::Native,
|
||||
toolshim_model: None,
|
||||
}
|
||||
}
|
||||
|
||||
fn parse_ollama_show_capabilities(
|
||||
response: OllamaShowResponse,
|
||||
fallback: Option<&ModelCapabilities>,
|
||||
) -> ModelCapabilities {
|
||||
let capability_set = response
|
||||
.capabilities
|
||||
.into_iter()
|
||||
.map(|capability| capability.trim().to_ascii_lowercase())
|
||||
.collect::<HashSet<_>>();
|
||||
let fallback = fallback.cloned().unwrap_or_default();
|
||||
let supports_tools = capability_set.contains("tools");
|
||||
|
||||
ModelCapabilities {
|
||||
vision: capability_set.contains("vision") || fallback.vision,
|
||||
tools: supports_tools,
|
||||
streaming: true,
|
||||
json_mode: supports_tools || fallback.json_mode,
|
||||
function_calling: supports_tools,
|
||||
reasoning: capability_set.contains("thinking") || fallback.reasoning,
|
||||
}
|
||||
}
|
||||
|
||||
async fn fetch_ollama_show_capabilities(
|
||||
client: &reqwest::Client,
|
||||
base_url: &str,
|
||||
model_name: &str,
|
||||
fallback: Option<&ModelCapabilities>,
|
||||
) -> Option<ModelCapabilities> {
|
||||
let url = format!("{base_url}/api/show");
|
||||
let response = client
|
||||
.post(url)
|
||||
.json(&serde_json::json!({ "name": model_name }))
|
||||
.send()
|
||||
.await
|
||||
.ok()?;
|
||||
let response = response.error_for_status().ok()?;
|
||||
let parsed = response.json::<OllamaShowResponse>().await.ok()?;
|
||||
Some(parse_ollama_show_capabilities(parsed, fallback))
|
||||
}
|
||||
|
||||
async fn fetch_ollama_toolshim_interpreter_model(
|
||||
client: &reqwest::Client,
|
||||
base_url: &str,
|
||||
selected_model: &str,
|
||||
) -> Option<String> {
|
||||
let url = format!("{base_url}/api/tags");
|
||||
let response = client.get(url).send().await.ok()?;
|
||||
let response = response.error_for_status().ok()?;
|
||||
let parsed = response.json::<OllamaTagsResponse>().await.ok()?;
|
||||
|
||||
for model in parsed
|
||||
.models
|
||||
.into_iter()
|
||||
.map(|model| model.name)
|
||||
.filter(|name| {
|
||||
normalize_provider_identity(name) != normalize_provider_identity(selected_model)
|
||||
})
|
||||
{
|
||||
let Some(capabilities) =
|
||||
fetch_ollama_show_capabilities(client, base_url, &model, None).await
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
if capabilities.tools || capabilities.function_calling {
|
||||
return Some(model);
|
||||
}
|
||||
}
|
||||
|
||||
None
|
||||
}
|
||||
|
||||
pub(crate) fn resolve_runtime_tool_call_decision_from_capabilities(
|
||||
provider_selector: Option<&str>,
|
||||
provider_name: &str,
|
||||
model_name: &str,
|
||||
capabilities: Option<&ModelCapabilities>,
|
||||
) -> RuntimeToolCallDecision {
|
||||
build_runtime_tool_call_decision(
|
||||
provider_selector,
|
||||
provider_name,
|
||||
model_name,
|
||||
capabilities
|
||||
.cloned()
|
||||
.unwrap_or_else(default_runtime_model_capabilities),
|
||||
None,
|
||||
)
|
||||
}
|
||||
|
||||
pub(crate) async fn resolve_runtime_tool_call_decision(
|
||||
provider_selector: Option<&str>,
|
||||
provider_name: &str,
|
||||
model_name: &str,
|
||||
base_url: Option<&str>,
|
||||
capabilities: Option<&ModelCapabilities>,
|
||||
) -> RuntimeToolCallDecision {
|
||||
if !is_ollama_provider(provider_selector, provider_name) {
|
||||
return resolve_runtime_tool_call_decision_from_capabilities(
|
||||
provider_selector,
|
||||
provider_name,
|
||||
model_name,
|
||||
capabilities,
|
||||
);
|
||||
}
|
||||
|
||||
let fallback_capabilities = capabilities
|
||||
.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()
|
||||
{
|
||||
Ok(client) => client,
|
||||
Err(error) => {
|
||||
tracing::warn!(
|
||||
"[AsterAgent] 创建 Ollama 运行时能力探测客户端失败,回退 catalog 能力: {}",
|
||||
error
|
||||
);
|
||||
return build_runtime_tool_call_decision(
|
||||
provider_selector,
|
||||
provider_name,
|
||||
model_name,
|
||||
fallback_capabilities,
|
||||
None,
|
||||
);
|
||||
}
|
||||
};
|
||||
|
||||
let Some(live_capabilities) = fetch_ollama_show_capabilities(
|
||||
&client,
|
||||
&base_url,
|
||||
model_name,
|
||||
Some(&fallback_capabilities),
|
||||
)
|
||||
.await
|
||||
else {
|
||||
tracing::warn!(
|
||||
"[AsterAgent] 读取 Ollama 模型能力失败,回退 catalog 能力: model={}, base_url={}",
|
||||
model_name,
|
||||
base_url
|
||||
);
|
||||
return build_runtime_tool_call_decision(
|
||||
provider_selector,
|
||||
provider_name,
|
||||
model_name,
|
||||
fallback_capabilities,
|
||||
None,
|
||||
);
|
||||
};
|
||||
|
||||
let toolshim_model = if live_capabilities.tools || live_capabilities.function_calling {
|
||||
None
|
||||
} else {
|
||||
fetch_ollama_toolshim_interpreter_model(&client, &base_url, model_name)
|
||||
.await
|
||||
.or_else(|| Some(model_name.to_string()))
|
||||
};
|
||||
|
||||
build_runtime_tool_call_decision(
|
||||
provider_selector,
|
||||
provider_name,
|
||||
model_name,
|
||||
live_capabilities,
|
||||
toolshim_model,
|
||||
)
|
||||
}
|
||||
|
||||
pub(crate) async fn enrich_provider_config_with_runtime_tool_strategy(
|
||||
provider_config: &mut ConfigureProviderRequest,
|
||||
) -> RuntimeToolCallDecision {
|
||||
let decision = resolve_runtime_tool_call_decision(
|
||||
provider_config.provider_id.as_deref(),
|
||||
&provider_config.provider_name,
|
||||
&provider_config.model_name,
|
||||
provider_config.base_url.as_deref(),
|
||||
provider_config.model_capabilities.as_ref(),
|
||||
)
|
||||
.await;
|
||||
|
||||
provider_config.model_capabilities = Some(decision.capabilities.clone());
|
||||
provider_config.tool_call_strategy = Some(decision.strategy);
|
||||
provider_config.toolshim_model = decision.toolshim_model.clone();
|
||||
|
||||
decision
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn native_strategy_keeps_tool_capable_model_on_native_path() {
|
||||
let capabilities = ModelCapabilities {
|
||||
vision: false,
|
||||
tools: true,
|
||||
streaming: true,
|
||||
json_mode: true,
|
||||
function_calling: true,
|
||||
reasoning: false,
|
||||
};
|
||||
|
||||
let decision = resolve_runtime_tool_call_decision_from_capabilities(
|
||||
Some("ollama"),
|
||||
"ollama",
|
||||
"glm-5.1:cloud",
|
||||
Some(&capabilities),
|
||||
);
|
||||
|
||||
assert_eq!(decision.strategy, RuntimeToolCallStrategy::Native);
|
||||
assert!(decision.toolshim_model.is_none());
|
||||
assert!(decision.capabilities.tools);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn toolshim_strategy_wraps_ollama_model_without_native_tools() {
|
||||
let capabilities = ModelCapabilities {
|
||||
vision: false,
|
||||
tools: false,
|
||||
streaming: true,
|
||||
json_mode: false,
|
||||
function_calling: false,
|
||||
reasoning: true,
|
||||
};
|
||||
|
||||
let decision = resolve_runtime_tool_call_decision_from_capabilities(
|
||||
Some("ollama"),
|
||||
"ollama",
|
||||
"deepseek-r1:latest",
|
||||
Some(&capabilities),
|
||||
);
|
||||
|
||||
assert_eq!(decision.strategy, RuntimeToolCallStrategy::ToolShim);
|
||||
assert_eq!(
|
||||
decision.toolshim_model.as_deref(),
|
||||
Some("deepseek-r1:latest")
|
||||
);
|
||||
assert!(!decision.capabilities.tools);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_ollama_show_capabilities_uses_live_capabilities() {
|
||||
let parsed = parse_ollama_show_capabilities(
|
||||
OllamaShowResponse {
|
||||
capabilities: vec![
|
||||
"completion".to_string(),
|
||||
"thinking".to_string(),
|
||||
"tools".to_string(),
|
||||
],
|
||||
},
|
||||
None,
|
||||
);
|
||||
|
||||
assert!(parsed.tools);
|
||||
assert!(parsed.function_calling);
|
||||
assert!(parsed.reasoning);
|
||||
assert!(parsed.json_mode);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn enrich_provider_config_sets_runtime_strategy_fields() {
|
||||
let mut provider_config = ConfigureProviderRequest {
|
||||
provider_id: Some("openai".to_string()),
|
||||
provider_name: "openai".to_string(),
|
||||
model_name: "gpt-4o".to_string(),
|
||||
api_key: None,
|
||||
base_url: None,
|
||||
model_capabilities: None,
|
||||
tool_call_strategy: None,
|
||||
toolshim_model: None,
|
||||
};
|
||||
|
||||
let decision =
|
||||
enrich_provider_config_with_runtime_tool_strategy(&mut provider_config).await;
|
||||
|
||||
assert_eq!(decision.strategy, RuntimeToolCallStrategy::Native);
|
||||
assert_eq!(
|
||||
provider_config.tool_call_strategy,
|
||||
Some(RuntimeToolCallStrategy::Native)
|
||||
);
|
||||
assert!(provider_config.toolshim_model.is_none());
|
||||
assert!(provider_config
|
||||
.model_capabilities
|
||||
.as_ref()
|
||||
.is_some_and(|capabilities| capabilities.tools));
|
||||
}
|
||||
}
|
||||
@@ -1,5 +1,6 @@
|
||||
use super::*;
|
||||
use crate::commands::model_registry_cmd::ModelRegistryState;
|
||||
use lime_core::database::dao::api_key_provider::{ApiProviderType, ProviderGroup};
|
||||
use lime_core::models::model_registry::{
|
||||
EnhancedModelMetadata, ModelCapabilities, ModelSource, ModelTier, ProviderAliasConfig,
|
||||
};
|
||||
@@ -9,11 +10,22 @@ use tauri::Manager;
|
||||
#[derive(Debug, Clone)]
|
||||
struct ProviderResolutionContext {
|
||||
provider_selector: String,
|
||||
aster_provider_name: String,
|
||||
compatibility_provider_key: String,
|
||||
registry_provider_ids: Vec<String>,
|
||||
alias_key: String,
|
||||
custom_models: Vec<String>,
|
||||
is_custom_provider: bool,
|
||||
provider_type: Option<ApiProviderType>,
|
||||
provider_group: Option<ProviderGroup>,
|
||||
configured_api_host: Option<String>,
|
||||
has_credentials: bool,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
enum RuntimeProviderConfigurationStrategy {
|
||||
Manual { base_url: Option<String> },
|
||||
CredentialPool,
|
||||
}
|
||||
|
||||
fn normalize_identifier(value: &str) -> String {
|
||||
@@ -45,6 +57,77 @@ fn provider_registry_id_from_key(provider_key: &str) -> String {
|
||||
}
|
||||
}
|
||||
|
||||
fn provider_type_from_key(provider_key: &str) -> Option<ApiProviderType> {
|
||||
match normalize_identifier(provider_key).as_str() {
|
||||
"openai" | "iflow" => Some(ApiProviderType::Openai),
|
||||
"anthropic" | "claude" | "claude_oauth" => Some(ApiProviderType::Anthropic),
|
||||
"anthropic-compatible" => Some(ApiProviderType::AnthropicCompatible),
|
||||
"gemini" | "gemini_api_key" => Some(ApiProviderType::Gemini),
|
||||
"azure-openai" => Some(ApiProviderType::AzureOpenai),
|
||||
"vertexai" => Some(ApiProviderType::Vertexai),
|
||||
"aws-bedrock" | "bedrock" => Some(ApiProviderType::AwsBedrock),
|
||||
"ollama" => Some(ApiProviderType::Ollama),
|
||||
"fal" => Some(ApiProviderType::Fal),
|
||||
"new-api" => Some(ApiProviderType::NewApi),
|
||||
"gateway" => Some(ApiProviderType::Gateway),
|
||||
"codex" => Some(ApiProviderType::Codex),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn normalize_runtime_provider_base_url(
|
||||
provider_type: Option<ApiProviderType>,
|
||||
base_url: Option<String>,
|
||||
) -> Option<String> {
|
||||
let normalized = normalize_optional_text(base_url)?;
|
||||
if provider_type != Some(ApiProviderType::Ollama) {
|
||||
return Some(normalized);
|
||||
}
|
||||
|
||||
let trimmed = normalized.trim_end_matches('/').to_string();
|
||||
if let Some(without_version) = trimmed.strip_suffix("/v1") {
|
||||
return normalize_optional_text(Some(without_version.to_string())).or_else(|| {
|
||||
Some(
|
||||
ApiProviderType::Ollama
|
||||
.runtime_spec()
|
||||
.default_api_host
|
||||
.to_string(),
|
||||
)
|
||||
});
|
||||
}
|
||||
|
||||
Some(trimmed)
|
||||
}
|
||||
|
||||
fn resolve_runtime_provider_configuration_strategy(
|
||||
context: &ProviderResolutionContext,
|
||||
) -> RuntimeProviderConfigurationStrategy {
|
||||
let configured_base_url = normalize_runtime_provider_base_url(
|
||||
context.provider_type,
|
||||
context.configured_api_host.clone(),
|
||||
);
|
||||
let is_credentialless_local_provider =
|
||||
matches!(context.provider_group, Some(ProviderGroup::Local)) && !context.has_credentials;
|
||||
let should_use_manual_provider =
|
||||
is_credentialless_local_provider || context.provider_type == Some(ApiProviderType::Ollama);
|
||||
|
||||
if !should_use_manual_provider {
|
||||
return RuntimeProviderConfigurationStrategy::CredentialPool;
|
||||
}
|
||||
|
||||
let fallback_base_url = context.provider_type.map(|provider_type| {
|
||||
normalize_runtime_provider_base_url(
|
||||
Some(provider_type),
|
||||
Some(provider_type.runtime_spec().default_api_host.to_string()),
|
||||
)
|
||||
.unwrap_or_else(|| provider_type.runtime_spec().default_api_host.to_string())
|
||||
});
|
||||
|
||||
RuntimeProviderConfigurationStrategy::Manual {
|
||||
base_url: configured_base_url.or(fallback_base_url),
|
||||
}
|
||||
}
|
||||
|
||||
fn infer_reasoning_capability(model_id: &str) -> bool {
|
||||
let normalized = normalize_identifier(model_id);
|
||||
normalized.contains("thinking") || normalized.contains("reasoning")
|
||||
@@ -227,18 +310,37 @@ fn build_provider_resolution_context(
|
||||
let provider_selector = normalize_identifier(provider_selector);
|
||||
let is_custom_provider =
|
||||
lime_core::models::provider_type::is_custom_provider_id(&provider_selector);
|
||||
let mut provider_type = provider_type_from_key(&provider_selector);
|
||||
let mut aster_provider_name = provider_type
|
||||
.map(|provider_type| provider_type.runtime_spec().aster_provider_name.to_string())
|
||||
.unwrap_or_else(|| provider_selector.clone());
|
||||
let mut compatibility_provider_key = provider_selector.clone();
|
||||
let mut registry_provider_ids = vec![
|
||||
provider_selector.clone(),
|
||||
provider_registry_id_from_key(&provider_selector),
|
||||
];
|
||||
let mut custom_models = Vec::new();
|
||||
let mut provider_group = None;
|
||||
let mut configured_api_host = None;
|
||||
let mut has_credentials = false;
|
||||
|
||||
if is_custom_provider {
|
||||
if let Some(provider_with_keys) = api_key_provider_service
|
||||
.0
|
||||
.get_provider(db, &provider_selector)?
|
||||
{
|
||||
if let Some(provider_with_keys) = api_key_provider_service
|
||||
.0
|
||||
.get_provider(db, &provider_selector)?
|
||||
{
|
||||
provider_type = Some(provider_with_keys.provider.provider_type);
|
||||
aster_provider_name = provider_with_keys
|
||||
.provider
|
||||
.provider_type
|
||||
.runtime_spec()
|
||||
.aster_provider_name
|
||||
.to_string();
|
||||
provider_group = Some(provider_with_keys.provider.group);
|
||||
configured_api_host =
|
||||
normalize_optional_text(Some(provider_with_keys.provider.api_host.clone()));
|
||||
has_credentials = !provider_with_keys.api_keys.is_empty();
|
||||
|
||||
if is_custom_provider {
|
||||
compatibility_provider_key = provider_with_keys.provider.provider_type.to_string();
|
||||
registry_provider_ids.push(provider_registry_id_from_key(&compatibility_provider_key));
|
||||
custom_models = provider_with_keys.provider.custom_models;
|
||||
@@ -251,10 +353,15 @@ fn build_provider_resolution_context(
|
||||
});
|
||||
|
||||
Ok(ProviderResolutionContext {
|
||||
aster_provider_name,
|
||||
alias_key: provider_alias_config_key(&provider_selector),
|
||||
compatibility_provider_key,
|
||||
custom_models,
|
||||
configured_api_host,
|
||||
has_credentials,
|
||||
is_custom_provider,
|
||||
provider_group,
|
||||
provider_type,
|
||||
provider_selector,
|
||||
registry_provider_ids,
|
||||
})
|
||||
@@ -904,18 +1011,38 @@ pub(super) async fn resolve_runtime_request_provider_config(
|
||||
);
|
||||
}
|
||||
|
||||
let provider_strategy = resolve_runtime_provider_configuration_strategy(&context);
|
||||
let base_url = match provider_strategy {
|
||||
RuntimeProviderConfigurationStrategy::Manual { base_url } => base_url,
|
||||
RuntimeProviderConfigurationStrategy::CredentialPool => None,
|
||||
};
|
||||
let model_capabilities = find_model_meta(&resolved_model, &catalog)
|
||||
.map(|model| model.capabilities.clone())
|
||||
.unwrap_or_else(|| {
|
||||
infer_model_capabilities(
|
||||
&resolved_model,
|
||||
Some(&context.provider_selector),
|
||||
None,
|
||||
None,
|
||||
)
|
||||
});
|
||||
|
||||
Ok(Some(ConfigureProviderRequest {
|
||||
provider_id: Some(context.provider_selector.clone()),
|
||||
provider_name: context.provider_selector,
|
||||
provider_name: context.aster_provider_name,
|
||||
model_name: resolved_model,
|
||||
api_key: None,
|
||||
base_url: None,
|
||||
base_url,
|
||||
model_capabilities: Some(model_capabilities),
|
||||
tool_call_strategy: None,
|
||||
toolshim_model: None,
|
||||
}))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use lime_core::database::dao::api_key_provider::ProviderGroup;
|
||||
|
||||
fn build_model(
|
||||
id: &str,
|
||||
@@ -1214,4 +1341,39 @@ mod tests {
|
||||
("gemini".to_string(), RequestPreferenceSource::Request)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn runtime_provider_strategy_prefers_manual_mode_for_credentialless_local_provider() {
|
||||
let context = ProviderResolutionContext {
|
||||
provider_selector: "ollama".to_string(),
|
||||
aster_provider_name: "ollama".to_string(),
|
||||
compatibility_provider_key: "ollama".to_string(),
|
||||
registry_provider_ids: vec!["ollama".to_string()],
|
||||
alias_key: "ollama".to_string(),
|
||||
custom_models: vec![],
|
||||
is_custom_provider: false,
|
||||
provider_type: Some(ApiProviderType::Ollama),
|
||||
provider_group: Some(ProviderGroup::Local),
|
||||
configured_api_host: Some("http://127.0.0.1:11434".to_string()),
|
||||
has_credentials: false,
|
||||
};
|
||||
|
||||
assert_eq!(
|
||||
resolve_runtime_provider_configuration_strategy(&context),
|
||||
RuntimeProviderConfigurationStrategy::Manual {
|
||||
base_url: Some("http://127.0.0.1:11434".to_string()),
|
||||
}
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn normalize_runtime_provider_base_url_strips_ollama_v1_suffix() {
|
||||
assert_eq!(
|
||||
normalize_runtime_provider_base_url(
|
||||
Some(ApiProviderType::Ollama),
|
||||
Some("http://127.0.0.1:11434/v1/".to_string()),
|
||||
),
|
||||
Some("http://127.0.0.1:11434".to_string())
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -13,6 +13,37 @@ const CONTEXT_COMPACTION_NOT_NEEDED_WARNING_CODE: &str = "context_compaction_not
|
||||
const LIME_RUNTIME_METADATA_KEY: &str = "lime_runtime";
|
||||
const LIME_RUNTIME_AUTO_COMPACT_KEY: &str = "auto_compact";
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
enum ProviderConfigApplyMode {
|
||||
Direct,
|
||||
CredentialPool,
|
||||
}
|
||||
|
||||
fn normalize_provider_identity(value: &str) -> String {
|
||||
value.trim().to_ascii_lowercase()
|
||||
}
|
||||
|
||||
fn resolve_provider_config_apply_mode(
|
||||
provider_config: &ConfigureProviderRequest,
|
||||
) -> ProviderConfigApplyMode {
|
||||
if provider_config.api_key.is_some() || provider_config.base_url.is_some() {
|
||||
return ProviderConfigApplyMode::Direct;
|
||||
}
|
||||
|
||||
let provider_selector = provider_config
|
||||
.provider_id
|
||||
.as_deref()
|
||||
.unwrap_or(&provider_config.provider_name);
|
||||
let normalized_selector = normalize_provider_identity(provider_selector);
|
||||
let normalized_provider_name = normalize_provider_identity(&provider_config.provider_name);
|
||||
|
||||
if normalized_selector == "ollama" || normalized_provider_name == "ollama" {
|
||||
return ProviderConfigApplyMode::Direct;
|
||||
}
|
||||
|
||||
ProviderConfigApplyMode::CredentialPool
|
||||
}
|
||||
|
||||
fn emit_runtime_side_event(
|
||||
app: &AppHandle,
|
||||
event_name: &str,
|
||||
@@ -441,6 +472,21 @@ async fn execute_aster_chat_request(
|
||||
if let Some(resolved_provider_config) = resolved_provider_config {
|
||||
request.provider_config = Some(resolved_provider_config);
|
||||
}
|
||||
if let Some(provider_config) = request.provider_config.as_mut() {
|
||||
let runtime_tool_call_decision =
|
||||
enrich_provider_config_with_runtime_tool_strategy(provider_config).await;
|
||||
tracing::info!(
|
||||
"[AsterAgent] provider_config 运行时工具策略: provider_id={:?}, provider_name={}, model_name={}, strategy={:?}, toolshim_model={:?}, tools={}, function_calling={}, reasoning={}",
|
||||
provider_config.provider_id,
|
||||
provider_config.provider_name,
|
||||
provider_config.model_name,
|
||||
runtime_tool_call_decision.strategy,
|
||||
runtime_tool_call_decision.toolshim_model,
|
||||
runtime_tool_call_decision.capabilities.tools,
|
||||
runtime_tool_call_decision.capabilities.function_calling,
|
||||
runtime_tool_call_decision.capabilities.reasoning
|
||||
);
|
||||
}
|
||||
normalize_runtime_turn_request_metadata(
|
||||
&mut request,
|
||||
session_recent_harness_context.theme.as_deref(),
|
||||
@@ -1015,6 +1061,7 @@ async fn execute_aster_chat_request(
|
||||
provider_config.api_key.is_some(),
|
||||
provider_config.base_url
|
||||
);
|
||||
let apply_mode = resolve_provider_config_apply_mode(provider_config);
|
||||
let config = ProviderConfig {
|
||||
provider_name: provider_config.provider_name.clone(),
|
||||
provider_selector: provider_config
|
||||
@@ -1027,31 +1074,39 @@ async fn execute_aster_chat_request(
|
||||
credential_uuid: None,
|
||||
force_responses_api: false,
|
||||
credential_path: None,
|
||||
toolshim: matches!(
|
||||
provider_config.tool_call_strategy,
|
||||
Some(RuntimeToolCallStrategy::ToolShim)
|
||||
),
|
||||
toolshim_model: provider_config.toolshim_model.clone(),
|
||||
};
|
||||
// 如果前端提供了 api_key,直接使用;否则从凭证池选择凭证
|
||||
if provider_config.api_key.is_some() {
|
||||
state.configure_provider(config, session_id, db).await?;
|
||||
let provider_selector = provider_config
|
||||
.provider_id
|
||||
.as_deref()
|
||||
.unwrap_or(&provider_config.provider_name);
|
||||
persist_session_provider_routing(session_id, provider_selector).await?;
|
||||
} else {
|
||||
// 没有 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_selector,
|
||||
&provider_config.model_name,
|
||||
session_id,
|
||||
)
|
||||
.await?;
|
||||
persist_session_provider_routing(session_id, provider_selector).await?;
|
||||
let provider_selector = provider_config
|
||||
.provider_id
|
||||
.as_deref()
|
||||
.unwrap_or(&provider_config.provider_name);
|
||||
tracing::info!(
|
||||
"[AsterAgent] provider_config 应用策略: provider_selector={}, mode={:?}, tool_call_strategy={:?}, toolshim_model={:?}",
|
||||
provider_selector,
|
||||
apply_mode,
|
||||
provider_config.tool_call_strategy,
|
||||
provider_config.toolshim_model
|
||||
);
|
||||
match apply_mode {
|
||||
ProviderConfigApplyMode::Direct => {
|
||||
state.configure_provider(config, session_id, db).await?;
|
||||
}
|
||||
ProviderConfigApplyMode::CredentialPool => {
|
||||
state
|
||||
.configure_provider_from_pool(
|
||||
db,
|
||||
provider_selector,
|
||||
&provider_config.model_name,
|
||||
session_id,
|
||||
)
|
||||
.await?;
|
||||
}
|
||||
}
|
||||
persist_session_provider_routing(session_id, provider_selector).await?;
|
||||
}
|
||||
|
||||
// 检查 Provider 是否已配置
|
||||
@@ -2738,6 +2793,25 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolve_provider_config_apply_mode_prefers_direct_for_ollama_without_api_key() {
|
||||
let provider_config = ConfigureProviderRequest {
|
||||
provider_id: Some("ollama".to_string()),
|
||||
provider_name: "ollama".to_string(),
|
||||
model_name: "deepseek-r1:latest".to_string(),
|
||||
api_key: None,
|
||||
base_url: None,
|
||||
model_capabilities: None,
|
||||
tool_call_strategy: None,
|
||||
toolshim_model: None,
|
||||
};
|
||||
|
||||
assert_eq!(
|
||||
resolve_provider_config_apply_mode(&provider_config),
|
||||
ProviderConfigApplyMode::Direct
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn service_skill_launch_stage_should_preserve_simple_user_message_and_force_site_run_first() {
|
||||
let user_message = "请帮我使用 GitHub 查一下 AI Agent 项目";
|
||||
|
||||
@@ -1,4 +1,7 @@
|
||||
use super::*;
|
||||
use crate::agent_tools::catalog::{
|
||||
LIME_SITE_INFO_TOOL_NAME, LIME_SITE_RUN_TOOL_NAME, LIME_SITE_SEARCH_TOOL_NAME,
|
||||
};
|
||||
use crate::services::site_capability_service::{
|
||||
get_site_adapter, run_site_adapter_with_optional_save, RunSiteAdapterRequest,
|
||||
SiteAdapterDefinition, SiteAdapterRunResult,
|
||||
@@ -13,6 +16,12 @@ const SERVICE_SKILL_LAUNCH_BROWSER_DENY_PATTERNS: &[&str] = &[
|
||||
"playwright*",
|
||||
];
|
||||
|
||||
const SERVICE_SKILL_LAUNCH_REPEAT_SITE_TOOL_DENY_PATTERNS: &[&str] = &[
|
||||
LIME_SITE_RUN_TOOL_NAME,
|
||||
LIME_SITE_SEARCH_TOOL_NAME,
|
||||
LIME_SITE_INFO_TOOL_NAME,
|
||||
];
|
||||
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
pub(crate) struct ServiceSkillLaunchSiteAdapterContext {
|
||||
pub(crate) adapter_name: String,
|
||||
@@ -374,6 +383,10 @@ pub(crate) fn service_skill_launch_browser_deny_patterns() -> &'static [&'static
|
||||
SERVICE_SKILL_LAUNCH_BROWSER_DENY_PATTERNS
|
||||
}
|
||||
|
||||
pub(crate) fn service_skill_launch_repeat_site_tool_deny_patterns() -> &'static [&'static str] {
|
||||
SERVICE_SKILL_LAUNCH_REPEAT_SITE_TOOL_DENY_PATTERNS
|
||||
}
|
||||
|
||||
pub(crate) fn build_service_skill_launch_run_request(
|
||||
request_metadata: Option<&serde_json::Value>,
|
||||
) -> Option<RunSiteAdapterRequest> {
|
||||
@@ -449,4 +462,20 @@ pub(crate) fn append_service_skill_launch_session_permissions(
|
||||
metadata: HashMap::new(),
|
||||
});
|
||||
}
|
||||
|
||||
for pattern in service_skill_launch_repeat_site_tool_deny_patterns() {
|
||||
permissions.push(ToolPermission {
|
||||
tool: (*pattern).to_string(),
|
||||
allowed: false,
|
||||
priority: 1250,
|
||||
conditions: conditions.clone(),
|
||||
parameter_restrictions: Vec::new(),
|
||||
scope: PermissionScope::Session,
|
||||
reason: Some(
|
||||
"站点技能启动回合已完成系统预执行,禁止在同回合重复调用站点执行工具".to_string(),
|
||||
),
|
||||
expires_at: None,
|
||||
metadata: HashMap::new(),
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
@@ -760,6 +760,18 @@ mod tests {
|
||||
deny_rule.conditions[0].value,
|
||||
serde_json::json!("session-service-skill-1")
|
||||
);
|
||||
|
||||
let site_run_rule = permissions
|
||||
.iter()
|
||||
.find(|permission| permission.tool == "lime_site_run")
|
||||
.expect("should add site run deny rule");
|
||||
assert!(!site_run_rule.allowed);
|
||||
assert_eq!(site_run_rule.priority, 1250);
|
||||
assert_eq!(site_run_rule.conditions.len(), 1);
|
||||
assert_eq!(
|
||||
site_run_rule.conditions[0].value,
|
||||
serde_json::json!("session-service-skill-1")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
@@ -5484,6 +5496,11 @@ mod tests {
|
||||
)
|
||||
.expect("should contain preload prompt");
|
||||
|
||||
assert!(merged.contains("这个 bundle 是系统侧采集得到的源材料"));
|
||||
assert!(merged.contains("不默认等于用户要的最终交付结果"));
|
||||
assert!(merged.contains("必须先基于这份已保存的 Markdown 继续完成原任务"));
|
||||
assert!(merged.contains("不要把保存路径、图片数量或采集摘要原样复述后就停止"));
|
||||
assert!(merged.contains("不要把 exports 下的源 bundle 直接当成最终结果目录"));
|
||||
assert!(merged.contains("Markdown 正文翻译成中文"));
|
||||
assert!(merged.contains("/tmp/project/saved/x-article-export/index.md"));
|
||||
assert!(merged.contains("只允许新增 Read / Write / Edit"));
|
||||
@@ -6754,7 +6771,7 @@ mod tests {
|
||||
.description()
|
||||
.contains("Fetches full schema definitions"));
|
||||
|
||||
super::tool_runtime::search_bridge::register_tool_search_tool_to_registry(
|
||||
super::tool_runtime::register_tool_search_tool_to_registry(
|
||||
&mut guard,
|
||||
registry.clone(),
|
||||
None,
|
||||
|
||||
@@ -13,7 +13,7 @@ pub(crate) mod media_cli_bridge;
|
||||
#[path = "tool_runtime/resource_search_tools.rs"]
|
||||
mod resource_search_tools;
|
||||
#[path = "tool_runtime/search_bridge.rs"]
|
||||
mod search_bridge;
|
||||
pub(crate) mod search_bridge;
|
||||
#[path = "tool_runtime/service_skill_tools.rs"]
|
||||
mod service_skill_tools;
|
||||
#[path = "tool_runtime/site_tools.rs"]
|
||||
@@ -33,6 +33,8 @@ pub(crate) use mcp_resource_tools::ensure_mcp_resource_tools_registered;
|
||||
#[allow(unused_imports)]
|
||||
pub(crate) use mcp_resource_tools::{ListMcpResourcesBridgeTool, ReadMcpResourceBridgeTool};
|
||||
pub(crate) use search_bridge::ensure_tool_search_tool_registered;
|
||||
#[cfg(test)]
|
||||
pub(crate) use search_bridge::register_tool_search_tool_to_registry;
|
||||
#[allow(unused_imports)]
|
||||
pub(crate) use search_bridge::ToolSearchBridgeTool;
|
||||
#[allow(unused_imports)]
|
||||
|
||||
@@ -373,7 +373,7 @@ impl Tool for ToolSearchBridgeTool {
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn register_tool_search_tool_to_registry(
|
||||
pub(crate) fn register_tool_search_tool_to_registry(
|
||||
registry: &mut aster::tools::ToolRegistry,
|
||||
registry_arc: Arc<tokio::sync::RwLock<aster::tools::ToolRegistry>>,
|
||||
extension_manager: Option<Arc<aster::agents::extension_manager::ExtensionManager>>,
|
||||
|
||||
@@ -121,7 +121,7 @@ pub async fn handle_command(
|
||||
return Ok(result);
|
||||
}
|
||||
|
||||
if let Some(result) = models::try_handle(state, cmd).await? {
|
||||
if let Some(result) = models::try_handle(state, cmd, args.as_ref()).await? {
|
||||
return Ok(result);
|
||||
}
|
||||
|
||||
@@ -594,4 +594,21 @@ mod tests {
|
||||
|
||||
assert!(error.to_string().contains("Dev Bridge 未持有 AppHandle"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn fetch_provider_models_auto_is_bridged() {
|
||||
let state = make_test_state();
|
||||
|
||||
let error = handle_command(
|
||||
&state,
|
||||
"fetch_provider_models_auto",
|
||||
Some(serde_json::json!({
|
||||
"providerId": "ollama"
|
||||
})),
|
||||
)
|
||||
.await
|
||||
.expect_err("missing provider should fail after bridge routing");
|
||||
|
||||
assert!(error.to_string().contains("Provider 不存在: ollama"));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
use super::{args_or_default, get_string_arg};
|
||||
use crate::dev_bridge::DevBridgeState;
|
||||
use lime_server_utils::load_model_registry_provider_ids_from_resources;
|
||||
use lime_services::model_registry_service::ModelRegistryService;
|
||||
use serde_json::Value as JsonValue;
|
||||
|
||||
type DynError = Box<dyn std::error::Error>;
|
||||
@@ -7,6 +9,7 @@ type DynError = Box<dyn std::error::Error>;
|
||||
pub(super) async fn try_handle(
|
||||
state: &DevBridgeState,
|
||||
cmd: &str,
|
||||
args: Option<&JsonValue>,
|
||||
) -> Result<Option<JsonValue>, DynError> {
|
||||
let result: JsonValue = match cmd {
|
||||
"get_models" => serde_json::json!({
|
||||
@@ -56,6 +59,58 @@ pub(super) async fn try_handle(
|
||||
"get_model_registry_provider_ids" => {
|
||||
serde_json::to_value(load_model_registry_provider_ids_from_resources()?)?
|
||||
}
|
||||
"fetch_provider_models_auto" => {
|
||||
let args = args_or_default(args);
|
||||
let provider_id = get_string_arg(&args, "providerId", "provider_id")?;
|
||||
let db = state
|
||||
.db
|
||||
.as_ref()
|
||||
.ok_or_else(|| "Database not initialized".to_string())?;
|
||||
let provider = state
|
||||
.api_key_provider_service
|
||||
.get_provider(db, &provider_id)?
|
||||
.ok_or_else(|| format!("Provider 不存在: {provider_id}"))?;
|
||||
|
||||
let api_host = provider.provider.api_host.clone();
|
||||
if api_host.is_empty() {
|
||||
return Err("Provider 没有配置 API Host".into());
|
||||
}
|
||||
|
||||
let provider_type = provider.provider.provider_type;
|
||||
let requires_api_key = ModelRegistryService::requires_api_key_for_model_fetch(
|
||||
&provider_id,
|
||||
&api_host,
|
||||
provider_type,
|
||||
);
|
||||
let api_key = if requires_api_key {
|
||||
state
|
||||
.api_key_provider_service
|
||||
.get_next_api_key(db, &provider_id)?
|
||||
.ok_or_else(|| format!("Provider {provider_id} 没有可用的 API Key"))?
|
||||
} else {
|
||||
state
|
||||
.api_key_provider_service
|
||||
.get_next_api_key(db, &provider_id)?
|
||||
.unwrap_or_default()
|
||||
};
|
||||
|
||||
let guard = state.model_registry.read().await;
|
||||
let service = guard
|
||||
.as_ref()
|
||||
.ok_or_else(|| "模型注册服务未初始化".to_string())?;
|
||||
|
||||
serde_json::to_value(
|
||||
service
|
||||
.fetch_models_from_api_with_hints(
|
||||
&provider_id,
|
||||
&api_host,
|
||||
&api_key,
|
||||
Some(provider_type),
|
||||
&provider.provider.custom_models,
|
||||
)
|
||||
.await?,
|
||||
)?
|
||||
}
|
||||
_ => return Ok(None),
|
||||
};
|
||||
|
||||
|
||||
Reference in New Issue
Block a user