Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
113 changes: 79 additions & 34 deletions src/app.rs
Original file line number Diff line number Diff line change
Expand Up @@ -157,6 +157,23 @@ fn has_single_user_message(chat: &Chat) -> bool {
== 1
}

fn models_dialog_provider_ids() -> Option<Vec<String>> {
let mut signature = crate::persistence::AuthDAO::new()
.and_then(|dao| dao.load())
.ok()?
.into_keys()
.map(|provider_id| format!("auth:{provider_id}"))
.collect::<Vec<_>>();

if let Ok(discovery) = crate::model::discovery::Discovery::new() {
signature.extend(discovery.custom_provider_dialog_signature());
}

signature.sort();
signature.dedup();
Some(signature)
}

#[derive(Debug, Clone, PartialEq, Eq)]
pub struct StreamingRetryStatus {
pub attempt: usize,
Expand Down Expand Up @@ -315,6 +332,7 @@ enum ModelsTaskKind {
struct ModelsTaskMessage {
kind: ModelsTaskKind,
result: crate::command::registry::CommandResult,
provider_signature: Option<Vec<String>>,
}

#[derive(Debug)]
Expand Down Expand Up @@ -1063,7 +1081,6 @@ impl App {
);
let agent_steps = agent_registry.max_steps_map();
let provider_timeouts = loaded_config.merged_config.provider_timeouts.clone();

let theme_for_colors = themes
.get(current_theme_index)
.or_else(|| themes.first())
Expand All @@ -1082,7 +1099,10 @@ impl App {
.with_permission_rules(loaded_config.merged_config.permission_rules.clone())
.with_agent_permission_rules(agent_registry.permission_rules_map());

let discovery = crate::model::discovery::Discovery::new().ok();
let discovery = crate::model::discovery::Discovery::new_with_custom(Some(
loaded_config.merged_config.custom_providers.clone(),
))
.ok();
let now = std::time::Instant::now();

Ok(Self {
Expand Down Expand Up @@ -6741,6 +6761,16 @@ impl App {
Ok(providers) => providers,
Err(_) => return,
};
let connected_provider_ids = connected_providers
.keys()
.cloned()
.collect::<std::collections::HashSet<String>>();

let discovery = Discovery::new();
let configured_provider_ids = discovery
.as_ref()
.map(Discovery::custom_provider_ids)
.unwrap_or_default();

let include_runtime = crate::model::extensions::ModelExtensions::runtime()
.iter()
Expand All @@ -6760,17 +6790,22 @@ impl App {
)
});

if connected_providers.is_empty() && !include_runtime && !include_unauthenticated_free {
if connected_providers.is_empty()
&& configured_provider_ids.is_empty()
&& !include_runtime
&& !include_unauthenticated_free
{
return;
}

let has_persistent = connected_providers.keys().any(|provider_id| {
!crate::model::extensions::ModelExtensions::is_runtime_provider(provider_id)
}) || include_unauthenticated_free
}) || !configured_provider_ids.is_empty()
|| include_unauthenticated_free
|| connected_providers.is_empty();

let models = if has_persistent {
match Discovery::new() {
match discovery.as_ref() {
Ok(discovery) => match tokio::task::block_in_place(|| {
let rt = tokio::runtime::Handle::current();
rt.block_on(discovery.fetch_models())
Expand Down Expand Up @@ -6806,6 +6841,19 @@ impl App {
} else {
return;
};
let mut models = models;
if include_runtime {
let runtime_models = tokio::task::block_in_place(|| {
let rt = tokio::runtime::Handle::current();
rt.block_on(
crate::model::extensions::ModelExtensions::runtime_models_for_dialog_cached_or_empty(),
)
});
crate::model::discovery::merge_dialog_models(&mut models, runtime_models);
}
if let Ok(discovery) = discovery.as_ref() {
discovery.apply_custom_models_to_dialog(&mut models);
}

self.model_reasoning_options = models
.iter()
Expand All @@ -6826,8 +6874,11 @@ impl App {
std::collections::HashMap::new();

let is_model_selectable = |model: &ModelType| {
connected_providers.contains_key(&model.provider_id)
|| crate::model::extensions::ModelExtensions::is_available_without_connection(model)
crate::model::discovery::is_model_selectable(
model,
&connected_provider_ids,
&configured_provider_ids,
)
};

for model in &models {
Expand Down Expand Up @@ -7081,14 +7132,7 @@ impl App {
return true;
}

let connected_provider_ids = crate::persistence::AuthDAO::new()
.and_then(|dao| dao.load())
.map(|providers| {
let mut ids = providers.into_keys().collect::<Vec<_>>();
ids.sort();
ids
})
.ok();
let connected_provider_ids = models_dialog_provider_ids();

if kind == ModelsTaskKind::Load
&& parsed.args.is_empty()
Expand Down Expand Up @@ -7129,12 +7173,17 @@ impl App {
}

let parsed = parsed.clone();
let provider_signature = models_dialog_provider_ids();
tokio::spawn(async move {
let result = match kind {
ModelsTaskKind::Load => crate::command::handlers::load_models(parsed).await,
ModelsTaskKind::Refresh => crate::command::handlers::refresh_models().await,
};
let _ = sender.send(ModelsTaskMessage { kind, result });
let _ = sender.send(ModelsTaskMessage {
kind,
result,
provider_signature,
});
});
true
}
Expand Down Expand Up @@ -8076,7 +8125,12 @@ impl App {
}

for event in events {
match (event.kind, event.result) {
let ModelsTaskMessage {
kind,
result,
provider_signature,
} = event;
match (kind, result) {
(
ModelsTaskKind::Load,
crate::command::registry::CommandResult::ShowDialog { title, items },
Expand All @@ -8094,14 +8148,10 @@ impl App {
})
.collect();
self.models_dialog_state.finish_loading();
self.models_dialog_provider_ids = crate::persistence::AuthDAO::new()
.and_then(|dao| dao.load())
.map(|providers| {
let mut ids = providers.into_keys().collect::<Vec<_>>();
ids.sort();
ids
})
.ok();
let current_signature = models_dialog_provider_ids();
self.models_dialog_provider_ids = (provider_signature == current_signature)
.then_some(provider_signature)
.flatten();
if self.overlay_focus == OverlayFocus::ModelsDialog {
self.show_models_dialog(title, dialog_items);
}
Expand Down Expand Up @@ -11872,6 +11922,7 @@ mod tests {
active: false,
}],
},
provider_signature: models_dialog_provider_ids(),
})
.unwrap();
app.models_receiver = Some(receiver);
Expand Down Expand Up @@ -11900,15 +11951,8 @@ mod tests {
active: false,
}],
);
// Match the connected-provider snapshot used by the reopen cache check.
app.models_dialog_provider_ids = crate::persistence::AuthDAO::new()
.and_then(|dao| dao.load())
.map(|providers| {
let mut ids = providers.into_keys().collect::<Vec<_>>();
ids.sort();
ids
})
.ok();
// Match the auth/config snapshot used by the reopen cache check.
app.models_dialog_provider_ids = models_dialog_provider_ids();
app.models_dialog_state.dialog.hide();
app.overlay_focus = OverlayFocus::None;

Expand Down Expand Up @@ -11945,6 +11989,7 @@ mod tests {
.send(ModelsTaskMessage {
kind: ModelsTaskKind::Refresh,
result: crate::command::registry::CommandResult::Success(String::new()),
provider_signature: models_dialog_provider_ids(),
})
.unwrap();
app.models_receiver = Some(receiver);
Expand Down
65 changes: 43 additions & 22 deletions src/command/handlers.rs
Original file line number Diff line number Diff line change
Expand Up @@ -106,7 +106,6 @@ pub fn handle_connect<'a>(
Ok(providers) => providers,
Err(e) => return CommandResult::Error(format!("Failed to load providers: {}", e)),
};

fn fallback_providers(
) -> std::collections::HashMap<String, crate::model::discovery::Provider> {
use crate::model::discovery::Provider;
Expand Down Expand Up @@ -274,6 +273,10 @@ pub async fn load_models(parsed: ParsedCommand) -> CommandResult {
Ok(providers) => providers,
Err(e) => return CommandResult::Error(format!("Failed to load providers: {}", e)),
};
let connected_provider_ids = connected_providers
.keys()
.cloned()
.collect::<std::collections::HashSet<String>>();

let provider_filter_matches_runtime = provider_filter.as_deref().is_some_and(|filter| {
let filter = filter.to_ascii_lowercase();
Expand All @@ -288,6 +291,16 @@ pub async fn load_models(parsed: ParsedCommand) -> CommandResult {
})
});

let discovery = Discovery::new();
let configured_provider_ids = discovery
.as_ref()
.map(Discovery::custom_provider_ids)
.unwrap_or_default();
let provider_filter_matches_configured = provider_filter.as_deref().is_some_and(|filter| {
discovery
.as_ref()
.is_ok_and(|discovery| discovery.custom_provider_matches_filter(filter))
});
let provider_filter_matches_unauthenticated_free = provider_filter.as_deref().is_some_and(
crate::model::extensions::ModelExtensions::unauthenticated_free_provider_matches_filter,
);
Expand All @@ -300,16 +313,16 @@ pub async fn load_models(parsed: ParsedCommand) -> CommandResult {
let has_persistent = connected_providers.keys().any(|provider_id| {
!crate::model::extensions::ModelExtensions::is_runtime_provider(provider_id)
}) || provider_filter.is_none()
|| provider_filter_matches_configured
|| provider_filter_matches_unauthenticated_free;

let snapshot_models = crate::model::effective_catalog::models_for_dialog()
.ok()
.flatten();
let discovery = Discovery::new();
let mut models: Vec<ModelType> = if let Some(models) = snapshot_models.as_ref() {
models.clone()
} else if has_persistent {
match discovery {
match discovery.as_ref() {
Ok(d) => match d.fetch_models().await {
Ok(models) => models
.into_iter()
Expand Down Expand Up @@ -352,23 +365,31 @@ pub async fn load_models(parsed: ParsedCommand) -> CommandResult {
Vec::new()
};

if let Ok(discovery) = discovery.as_ref() {
discovery.apply_custom_models_to_dialog(&mut models);
}

let mut runtime_errors = Vec::new();
if snapshot_models.is_none() && has_runtime {
if has_runtime {
let runtime_result =
crate::model::extensions::ModelExtensions::runtime_models_for_dialog_cached().await;
models.extend(runtime_result.models);
crate::model::discovery::merge_dialog_models(&mut models, runtime_result.models);
runtime_errors = runtime_result.errors;
}

if snapshot_models.is_none() && !models.is_empty() {
if let Err(err) =
crate::model::effective_catalog::publish_refreshed_models(models.clone())
{
push_toast(Toast::new(
format!("Failed to seed model catalog cache: {}", err),
ToastLevel::Warning,
Some(std::time::Duration::from_secs(3)),
));
if let Ok(discovery) = Discovery::new_with_custom(None) {
if let Ok(snapshot_models) = discovery.fetch_models().await {
if let Err(err) =
crate::model::effective_catalog::publish_refreshed_models(snapshot_models)
{
push_toast(Toast::new(
format!("Failed to seed model catalog cache: {}", err),
ToastLevel::Warning,
Some(std::time::Duration::from_secs(3)),
));
}
}
}
}

Expand All @@ -378,14 +399,14 @@ pub async fn load_models(parsed: ParsedCommand) -> CommandResult {
std::collections::HashMap::new();

let is_model_selectable = |model: &ModelType| {
(connected_providers.contains_key(&model.provider_id)
|| crate::model::extensions::ModelExtensions::is_available_without_connection(
model,
))
&& crate::model::extensions::ModelExtensions::model_matches_provider_filter(
model,
provider_filter.as_deref(),
)
crate::model::discovery::is_model_selectable(
model,
&connected_provider_ids,
&configured_provider_ids,
) && crate::model::extensions::ModelExtensions::model_matches_provider_filter(
model,
provider_filter.as_deref(),
)
};

for model in &models {
Expand Down Expand Up @@ -798,7 +819,7 @@ pub async fn refresh_models() -> CommandResult {
.sum::<usize>()
+ runtime_model_count;

let models = match crate::model::discovery::Discovery::new() {
let models = match crate::model::discovery::Discovery::new_with_custom(None) {
Ok(discovery) => match discovery.fetch_models().await {
Ok(models) => models,
Err(err) => {
Expand Down
Loading