Complete local-first Rig provider migration
This commit is contained in:
+237
-79
@@ -17,14 +17,16 @@ use warp_multi_agent_api as api;
|
||||
use super::custom_model_routers::{self, CustomModelRouter, ModelConfigError};
|
||||
use super::execution_profiles::profiles::AIExecutionProfilesModel;
|
||||
use crate::ai::acp::{acp_launch_fingerprint, acp_selection_identity};
|
||||
use crate::ai::bedrock::models::get_effective_models;
|
||||
use crate::auth::auth_manager::{AuthManager, AuthManagerEvent};
|
||||
use crate::auth::AuthStateProvider;
|
||||
use crate::network::{NetworkStatus, NetworkStatusEvent, NetworkStatusKind};
|
||||
#[cfg(not(target_family = "wasm"))]
|
||||
use crate::persistence::model::{AcpConversationData, AgentBackend};
|
||||
use crate::server::server_api::ServerApiProvider;
|
||||
use crate::settings::{AcpConfigValueSettings, BedrockModelConfig, OpenAIModelConfig};
|
||||
use crate::settings::{
|
||||
AcpConfigValueSettings, BedrockModelConfig, OpenAIModelConfig, OpenAIProviderConfig,
|
||||
OpenAIProviderKind,
|
||||
};
|
||||
use crate::user_config::{WarpConfig, WarpConfigUpdateEvent};
|
||||
use crate::workspaces::user_workspaces::{UserWorkspaces, UserWorkspacesEvent};
|
||||
use crate::{report_error, AISettings};
|
||||
@@ -659,6 +661,7 @@ impl LLMPreferences {
|
||||
| AISettingsChangedEvent::OpenAIProviders { .. }
|
||||
| AISettingsChangedEvent::AcpAgents { .. }
|
||||
| AISettingsChangedEvent::AcpAgentId { .. }
|
||||
| AISettingsChangedEvent::BedrockModels { .. }
|
||||
) {
|
||||
me.inject_bedrock_models(ctx);
|
||||
me.inject_openai_models(ctx);
|
||||
@@ -670,6 +673,11 @@ impl LLMPreferences {
|
||||
) {
|
||||
me.fetch_openai_models_from_endpoint(ctx);
|
||||
}
|
||||
if matches!(event, AISettingsChangedEvent::BedrockEnabled { .. })
|
||||
&& *AISettings::as_ref(ctx).bedrock_enabled.value()
|
||||
{
|
||||
me.refresh_bedrock_models(ctx);
|
||||
}
|
||||
// Safety: ensure the default model is still present in choices.
|
||||
// If all provider models were removed, the default_id would dangle.
|
||||
me.ensure_default_model_present();
|
||||
@@ -709,8 +717,8 @@ impl LLMPreferences {
|
||||
|
||||
#[cfg(not(target_family = "wasm"))]
|
||||
{
|
||||
Self::ensure_default_models_in_settings(ctx);
|
||||
me.inject_bedrock_models(ctx);
|
||||
Self::ensure_default_chatgpt_models_in_settings(ctx);
|
||||
me.refresh_bedrock_models(ctx);
|
||||
me.inject_openai_models(ctx);
|
||||
me.ensure_default_model_present();
|
||||
me.fetch_openai_models_from_endpoint(ctx);
|
||||
@@ -720,39 +728,85 @@ impl LLMPreferences {
|
||||
}
|
||||
|
||||
#[cfg(not(target_family = "wasm"))]
|
||||
fn ensure_default_models_in_settings(ctx: &mut ModelContext<Self>) {
|
||||
use crate::ai::bedrock::models::DEFAULT_BEDROCK_MODELS;
|
||||
fn ensure_default_chatgpt_models_in_settings(ctx: &mut ModelContext<Self>) {
|
||||
let mut providers = AISettings::as_ref(ctx).openai_providers.value().clone();
|
||||
let default_chatgpt_models = crate::settings::ai::default_chatgpt_provider().models;
|
||||
let mut providers_changed = false;
|
||||
for provider in &mut providers {
|
||||
if provider.kind != OpenAIProviderKind::ChatGPTSubscription {
|
||||
continue;
|
||||
}
|
||||
|
||||
let settings = AISettings::as_ref(ctx);
|
||||
let mut current_models: Vec<BedrockModelConfig> = settings.bedrock_models.value().clone();
|
||||
for default_model in &default_chatgpt_models {
|
||||
if !provider
|
||||
.models
|
||||
.iter()
|
||||
.any(|model| model.model_id == default_model.model_id)
|
||||
{
|
||||
provider.models.push(default_model.clone());
|
||||
providers_changed = true;
|
||||
}
|
||||
}
|
||||
|
||||
let existing_ids: std::collections::HashSet<String> =
|
||||
current_models.iter().map(|m| m.model_id.clone()).collect();
|
||||
|
||||
let mut added = false;
|
||||
for default in DEFAULT_BEDROCK_MODELS {
|
||||
if !existing_ids.contains(default.model_id as &str) {
|
||||
current_models.push(BedrockModelConfig {
|
||||
model_id: default.model_id.to_string(),
|
||||
display_name: default.display_name.to_string(),
|
||||
vision_supported: default.vision_supported,
|
||||
use_rig: false,
|
||||
});
|
||||
added = true;
|
||||
for model in &mut provider.models {
|
||||
if !model.reasoning_efforts.is_empty() {
|
||||
continue;
|
||||
}
|
||||
if let Some(default_model) = default_chatgpt_models
|
||||
.iter()
|
||||
.find(|default_model| default_model.model_id == model.model_id)
|
||||
{
|
||||
if !default_model.reasoning_efforts.is_empty() {
|
||||
model.reasoning_efforts = default_model.reasoning_efforts.clone();
|
||||
providers_changed = true;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if added {
|
||||
log::info!(
|
||||
"[bedrock] Added missing default models to settings — now {} total",
|
||||
current_models.len()
|
||||
);
|
||||
if providers_changed {
|
||||
AISettings::handle(ctx).update(ctx, |settings, ctx| {
|
||||
let _ = settings.bedrock_models.set_value(current_models, ctx);
|
||||
let _ = settings.openai_providers.set_value(providers, ctx);
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(not(target_family = "wasm"))]
|
||||
fn refresh_bedrock_models(&mut self, ctx: &mut ModelContext<Self>) {
|
||||
let settings = AISettings::as_ref(ctx);
|
||||
if !*settings.bedrock_enabled.value() {
|
||||
return;
|
||||
}
|
||||
let config = crate::ai::bedrock::client::BedrockClientConfig {
|
||||
auth_method: *settings.bedrock_auth_method.value(),
|
||||
profile: settings.bedrock_profile.value().clone(),
|
||||
region: settings.bedrock_region.value().clone(),
|
||||
access_key_id: settings.bedrock_access_key_id.value().clone(),
|
||||
secret_access_key: settings.bedrock_secret_access_key.value().clone(),
|
||||
session_token: None,
|
||||
cross_region_inference: *settings.bedrock_cross_region_inference.value(),
|
||||
use_rig: false,
|
||||
};
|
||||
|
||||
let _ = ctx.spawn(
|
||||
async move { crate::ai::bedrock::discovery::discover_available_models(config).await },
|
||||
|me, result, ctx| match result {
|
||||
Ok(models) => {
|
||||
AISettings::handle(ctx).update(ctx, |settings, ctx| {
|
||||
if let Err(error) = settings.bedrock_models.set_value(models, ctx) {
|
||||
log::warn!("[bedrock] Failed to persist discovered models: {error}");
|
||||
}
|
||||
});
|
||||
me.inject_bedrock_models(ctx);
|
||||
me.ensure_default_model_present();
|
||||
ctx.emit(LLMPreferencesEvent::UpdatedAvailableLLMs);
|
||||
}
|
||||
Err(error) => {
|
||||
log::debug!("[bedrock] Startup model discovery unavailable: {error}");
|
||||
}
|
||||
},
|
||||
);
|
||||
}
|
||||
|
||||
#[cfg(not(target_family = "wasm"))]
|
||||
fn inject_bedrock_models(&mut self, ctx: &AppContext) {
|
||||
// Galaxy's runtime inventory is rebuilt exclusively from enabled local
|
||||
@@ -768,7 +822,9 @@ impl LLMPreferences {
|
||||
return;
|
||||
}
|
||||
|
||||
let user_models: Vec<BedrockModelConfig> = settings.bedrock_models.value().clone();
|
||||
// Bedrock models are populated only by the control-plane discovery
|
||||
// flow. Never fall back to a static catalog or external config here.
|
||||
let discovered_models: Vec<BedrockModelConfig> = settings.bedrock_models.value().clone();
|
||||
let region = settings.bedrock_region.value().clone();
|
||||
let cross_region = *settings.bedrock_cross_region_inference.value();
|
||||
|
||||
@@ -777,7 +833,7 @@ impl LLMPreferences {
|
||||
let external_config = ExternalBedrockConfig::load();
|
||||
let require_1h_cache = external_config.enable_prompt_caching_1h;
|
||||
|
||||
let mut effective = get_effective_models(&user_models);
|
||||
let mut effective = discovered_models;
|
||||
|
||||
// Filter out models that don't support 1-hour caching if required
|
||||
if require_1h_cache {
|
||||
@@ -943,8 +999,15 @@ impl LLMPreferences {
|
||||
return;
|
||||
}
|
||||
|
||||
let mut provider_entries: Vec<(String, String, Option<String>, Vec<OpenAIModelConfig>)> =
|
||||
Vec::new();
|
||||
type OpenAIProviderEntry = (
|
||||
String,
|
||||
OpenAIProviderKind,
|
||||
bool,
|
||||
String,
|
||||
Option<String>,
|
||||
Vec<OpenAIModelConfig>,
|
||||
);
|
||||
let mut provider_entries: Vec<OpenAIProviderEntry> = Vec::new();
|
||||
|
||||
let configured_models = settings.openai_models.value().clone();
|
||||
let single_provider_models = if configured_models.is_empty() {
|
||||
@@ -968,7 +1031,14 @@ impl LLMPreferences {
|
||||
} else {
|
||||
"LiteLLM".to_string()
|
||||
};
|
||||
provider_entries.push((name, base_url, api_key, single_provider_models));
|
||||
provider_entries.push((
|
||||
name,
|
||||
OpenAIProviderKind::OpenAICompatible,
|
||||
true,
|
||||
base_url,
|
||||
api_key,
|
||||
single_provider_models,
|
||||
));
|
||||
}
|
||||
|
||||
provider_entries.extend(
|
||||
@@ -977,11 +1047,17 @@ impl LLMPreferences {
|
||||
.value()
|
||||
.iter()
|
||||
.filter_map(|provider| {
|
||||
if provider.base_url.trim().is_empty() || provider.models.is_empty() {
|
||||
if !provider.enabled
|
||||
|| (provider.kind == OpenAIProviderKind::OpenAICompatible
|
||||
&& provider.base_url.trim().is_empty())
|
||||
|| provider.models.is_empty()
|
||||
{
|
||||
return None;
|
||||
}
|
||||
Some((
|
||||
provider.name.clone(),
|
||||
provider.kind,
|
||||
provider.enabled,
|
||||
provider.base_url.clone(),
|
||||
provider.api_key.clone(),
|
||||
provider.models.clone(),
|
||||
@@ -995,58 +1071,93 @@ impl LLMPreferences {
|
||||
|
||||
let mut total_injected = 0;
|
||||
let mut seen_model_ids: HashSet<String> = HashSet::new();
|
||||
for (provider_name, base_url, api_key, models) in provider_entries {
|
||||
for (provider_name, provider_kind, provider_enabled, base_url, api_key, models) in
|
||||
provider_entries
|
||||
{
|
||||
if !provider_enabled {
|
||||
continue;
|
||||
}
|
||||
for model in &models {
|
||||
if !model.enabled {
|
||||
continue;
|
||||
}
|
||||
if !seen_model_ids.insert(model.model_id.clone()) {
|
||||
continue;
|
||||
}
|
||||
|
||||
// Register the routing entry
|
||||
let client_config = OpenAIClientConfig {
|
||||
base_url: base_url.clone(),
|
||||
api_key: api_key.clone(),
|
||||
model: None, // filled per-request from model_id
|
||||
max_input_tokens: Some(openai_model_context_size(model)),
|
||||
max_output_tokens: model.max_output_tokens,
|
||||
use_rig: model.use_rig,
|
||||
supports_system_messages: model.supports_system_messages(),
|
||||
};
|
||||
self.openai_provider_routing
|
||||
.insert(model.model_id.clone(), client_config);
|
||||
let reasoning_efforts: Vec<Option<&String>> =
|
||||
if provider_kind == OpenAIProviderKind::ChatGPTSubscription {
|
||||
// Keep the base model as the provider-default mode, then expose each
|
||||
// explicitly supported effort as a separate selectable variant.
|
||||
std::iter::once(None)
|
||||
.chain(model.reasoning_efforts.iter().map(Some))
|
||||
.collect()
|
||||
} else {
|
||||
vec![None]
|
||||
};
|
||||
|
||||
let llm_info = LLMInfo {
|
||||
id: LLMId::from(model.model_id.as_str()),
|
||||
display_name: model.display_name.clone(),
|
||||
base_model_name: model.display_name.clone(),
|
||||
reasoning_level: None,
|
||||
usage_metadata: LLMUsageMetadata {
|
||||
request_multiplier: 1,
|
||||
credit_multiplier: None,
|
||||
},
|
||||
description: Some(provider_name.clone()),
|
||||
disable_reason: None,
|
||||
vision_supported: model.vision_supported,
|
||||
spec: None,
|
||||
provider: LLMProvider::LiteLLM,
|
||||
host_configs: HashMap::from([(
|
||||
LLMModelHost::DirectApi,
|
||||
RoutingHostConfig {
|
||||
enabled: true,
|
||||
model_routing_host: LLMModelHost::DirectApi,
|
||||
for reasoning_effort in reasoning_efforts {
|
||||
let reasoning_effort = reasoning_effort.cloned();
|
||||
let model_key = reasoning_effort.as_deref().map_or_else(
|
||||
|| model.model_id.clone(),
|
||||
|effort| openai_model_variant_id(&model.model_id, effort),
|
||||
);
|
||||
|
||||
// Register the routing entry. Reasoning variants keep the provider's
|
||||
// actual model ID while using their synthetic key only for selection.
|
||||
let client_config = OpenAIClientConfig {
|
||||
kind: provider_kind,
|
||||
base_url: base_url.clone(),
|
||||
api_key: api_key.clone(),
|
||||
model: Some(model.model_id.clone()),
|
||||
reasoning_effort: reasoning_effort.clone(),
|
||||
max_input_tokens: Some(openai_model_context_size(model)),
|
||||
max_output_tokens: model.max_output_tokens,
|
||||
use_rig: model.use_rig
|
||||
|| provider_kind == OpenAIProviderKind::ChatGPTSubscription,
|
||||
supports_system_messages: model.supports_system_messages(),
|
||||
};
|
||||
self.openai_provider_routing
|
||||
.insert(model_key.clone(), client_config);
|
||||
|
||||
let display_name = reasoning_effort.as_deref().map_or_else(
|
||||
|| model.display_name.clone(),
|
||||
|effort| format!("{} ({effort})", model.display_name),
|
||||
);
|
||||
let llm_info = LLMInfo {
|
||||
id: LLMId::from(model_key.as_str()),
|
||||
display_name,
|
||||
base_model_name: model.display_name.clone(),
|
||||
reasoning_level: reasoning_effort,
|
||||
usage_metadata: LLMUsageMetadata {
|
||||
request_multiplier: 1,
|
||||
credit_multiplier: None,
|
||||
},
|
||||
)]),
|
||||
discount_percentage: None,
|
||||
context_window: openai_model_context_window(model),
|
||||
};
|
||||
self.models_by_feature
|
||||
.agent_mode
|
||||
.choices
|
||||
.push(llm_info.clone());
|
||||
self.models_by_feature.coding.choices.push(llm_info.clone());
|
||||
if let Some(ref mut cli) = self.models_by_feature.cli_agent {
|
||||
cli.choices.push(llm_info);
|
||||
description: Some(provider_name.clone()),
|
||||
disable_reason: None,
|
||||
vision_supported: model.vision_supported,
|
||||
spec: None,
|
||||
provider: LLMProvider::LiteLLM,
|
||||
host_configs: HashMap::from([(
|
||||
LLMModelHost::DirectApi,
|
||||
RoutingHostConfig {
|
||||
enabled: true,
|
||||
model_routing_host: LLMModelHost::DirectApi,
|
||||
},
|
||||
)]),
|
||||
discount_percentage: None,
|
||||
context_window: openai_model_context_window(model),
|
||||
};
|
||||
self.models_by_feature
|
||||
.agent_mode
|
||||
.choices
|
||||
.push(llm_info.clone());
|
||||
self.models_by_feature.coding.choices.push(llm_info.clone());
|
||||
if let Some(ref mut cli) = self.models_by_feature.cli_agent {
|
||||
cli.choices.push(llm_info);
|
||||
}
|
||||
total_injected += 1;
|
||||
}
|
||||
total_injected += 1;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1438,6 +1549,43 @@ impl LLMPreferences {
|
||||
);
|
||||
}
|
||||
|
||||
/// Discovers models for a provider draft without persisting or injecting it.
|
||||
///
|
||||
/// The provider setup modal uses this to keep configuration changes atomic
|
||||
/// until the user clicks Save.
|
||||
#[cfg(not(target_family = "wasm"))]
|
||||
pub(crate) async fn discover_openai_provider_models(
|
||||
provider: OpenAIProviderConfig,
|
||||
) -> Result<Vec<OpenAIModelConfig>, String> {
|
||||
if provider.base_url.trim().is_empty() {
|
||||
return Err("Enter a provider URL before testing the connection.".to_string());
|
||||
}
|
||||
|
||||
let base_url = provider.base_url.trim_end_matches('/').to_string();
|
||||
let client = reqwest::Client::builder()
|
||||
.timeout(std::time::Duration::from_secs(10))
|
||||
.build()
|
||||
.map_err(|error| format!("Could not create the provider client: {error}"))?;
|
||||
let api_key = provider.api_key.as_deref().filter(|key| !key.is_empty());
|
||||
|
||||
let models = if let Some(models) =
|
||||
fetch_from_litellm_model_info(&base_url, api_key, &client).await
|
||||
{
|
||||
models
|
||||
} else {
|
||||
fetch_from_openai_models(&base_url, api_key, &client).await
|
||||
};
|
||||
|
||||
if models.is_empty() {
|
||||
return Err(
|
||||
"The provider responded, but no models were found at /model/info or /models."
|
||||
.to_string(),
|
||||
);
|
||||
}
|
||||
|
||||
Ok(models)
|
||||
}
|
||||
|
||||
/// Returns the `LLMInfo` for the base LLM to be used for an Agent Mode request.
|
||||
pub fn get_active_base_model<'a>(
|
||||
&'a self,
|
||||
@@ -2228,7 +2376,7 @@ fn openai_model_context_size(model: &OpenAIModelConfig) -> u32 {
|
||||
/// Merges endpoint metadata into a provider's configured models without
|
||||
/// discarding local routing choices or manually configured models.
|
||||
#[cfg(not(target_family = "wasm"))]
|
||||
fn merge_discovered_provider_models(
|
||||
pub(crate) fn merge_discovered_provider_models(
|
||||
existing_models: &[OpenAIModelConfig],
|
||||
discovered_models: Vec<OpenAIModelConfig>,
|
||||
) -> Vec<OpenAIModelConfig> {
|
||||
@@ -2245,6 +2393,7 @@ fn merge_discovered_provider_models(
|
||||
.find(|model| model.model_id == discovered.model_id)
|
||||
{
|
||||
discovered.display_name = existing.display_name.clone();
|
||||
discovered.enabled = existing.enabled;
|
||||
discovered.use_rig = existing.use_rig;
|
||||
if existing.supports_system_messages.is_some() {
|
||||
discovered.supports_system_messages = existing.supports_system_messages;
|
||||
@@ -2271,6 +2420,11 @@ fn merge_discovered_provider_models(
|
||||
merged
|
||||
}
|
||||
|
||||
#[cfg(not(target_family = "wasm"))]
|
||||
fn openai_model_variant_id(model_id: &str, reasoning_effort: &str) -> String {
|
||||
format!("{model_id}::reasoning::{reasoning_effort}")
|
||||
}
|
||||
|
||||
#[cfg(not(target_family = "wasm"))]
|
||||
fn openai_model_context_window(model: &OpenAIModelConfig) -> LLMContextWindow {
|
||||
let context_size = openai_model_context_size(model);
|
||||
@@ -2400,6 +2554,8 @@ async fn fetch_from_litellm_model_info(
|
||||
} else {
|
||||
model_info["supports_system_messages"].as_bool()
|
||||
},
|
||||
reasoning_efforts: Vec::new(),
|
||||
enabled: true,
|
||||
})
|
||||
})
|
||||
.collect();
|
||||
@@ -2527,6 +2683,8 @@ async fn fetch_from_openai_models(
|
||||
} else {
|
||||
m["supports_system_messages"].as_bool()
|
||||
},
|
||||
reasoning_efforts: Vec::new(),
|
||||
enabled: true,
|
||||
})
|
||||
})
|
||||
.collect();
|
||||
|
||||
Reference in New Issue
Block a user