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
1 change: 1 addition & 0 deletions src/apps/cli/src/main.rs
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@ mod daemon;
mod diagnostics;
mod logging;
mod management;
mod model_selection;
mod modes;
mod peer_host;
mod plugin_diagnostics;
Expand Down
22 changes: 7 additions & 15 deletions src/apps/cli/src/management.rs
Original file line number Diff line number Diff line change
Expand Up @@ -83,6 +83,7 @@ pub(crate) async fn print_models() -> Result<()> {
config_service.get_config(None).await?;

let primary_model_id = global_config.ai.default_models.primary.clone();
let mode_model_id = crate::model_selection::resolve_mode_model_id(&global_config.ai);

println!("AI models");
println!();
Expand All @@ -93,12 +94,7 @@ pub(crate) async fn print_models() -> Result<()> {

for model in models {
let is_primary = primary_model_id.as_deref() == Some(model.id.as_str());
let current_modes: Vec<String> = global_config
.ai
.agent_models
.iter()
.filter_map(|(mode, model_id)| (model_id == &model.id).then_some(mode.clone()))
.collect();
let is_mode_default = mode_model_id.as_deref() == Some(model.id.as_str());

println!(
"- {}{} ({})",
Expand All @@ -109,8 +105,8 @@ pub(crate) async fn print_models() -> Result<()> {
println!(" Name: {}", model.name);
println!(" Provider: {}", model.provider);
println!(" Model: {}", model.model_name);
if !current_modes.is_empty() {
println!(" Used by modes: {}", current_modes.join(", "));
if is_mode_default {
println!(" Used by modes: all");
}
}

Expand Down Expand Up @@ -170,16 +166,12 @@ pub(crate) async fn print_mcp_servers() -> Result<()> {

pub(crate) async fn set_default_model(model_id: &str) -> Result<()> {
let config_service = ensure_global_config_service().await?;
let agent_registry = get_agent_registry();
let modes = agent_registry.get_modes_info().await;

config_service
.set_config("ai.default_models.primary", model_id)
.await?;
for mode in modes {
let path = format!("ai.agent_models.{}", mode.id);
config_service.set_config(&path, model_id).await?;
}
config_service
.set_config("ai.agent_model_defaults.mode", model_id)
.await?;

println!("Default model set to: {}", model_id);
Ok(())
Expand Down
68 changes: 68 additions & 0 deletions src/apps/cli/src/model_selection.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,68 @@
use bitfun_core::service::config::AIConfig;

/// Resolve the shared future-mode selector to the concrete enabled model shown
/// by CLI model pickers and status surfaces.
pub(crate) fn resolve_mode_model_id(ai_config: &AIConfig) -> Option<String> {
let selector = ai_config.agent_model_defaults.mode.trim();
match selector {
"" | "auto" | "default" => ai_config.resolve_model_selection("primary"),
selector => ai_config.resolve_model_selection(selector),
}
}

#[cfg(test)]
mod tests {
use super::*;

fn config_with_selector(selector: &str) -> AIConfig {
serde_json::from_value(serde_json::json!({
"models": [
{
"id": "primary-model",
"name": "Primary",
"provider": "openai",
"model_name": "primary-model",
"enabled": true
},
{
"id": "fast-model",
"name": "Fast",
"provider": "openai",
"model_name": "fast-model",
"enabled": true
},
{
"id": "explicit-model",
"name": "Explicit",
"provider": "openai",
"model_name": "explicit-model",
"enabled": true
}
],
"default_models": {
"primary": "primary-model",
"fast": "fast-model"
},
"agent_model_defaults": {
"mode": selector
}
}))
.expect("test AI config should deserialize")
}

#[test]
fn resolves_symbolic_and_explicit_mode_defaults_for_cli_display() {
assert_eq!(
resolve_mode_model_id(&config_with_selector("auto")).as_deref(),
Some("primary-model")
);
assert_eq!(
resolve_mode_model_id(&config_with_selector("fast")).as_deref(),
Some("fast-model")
);
assert_eq!(
resolve_mode_model_id(&config_with_selector("explicit-model")).as_deref(),
Some("explicit-model")
);
}
}
62 changes: 19 additions & 43 deletions src/apps/cli/src/modes/chat.rs
Original file line number Diff line number Diff line change
Expand Up @@ -3881,7 +3881,6 @@ impl ChatMode {
chat_state: &mut ChatState,
rt_handle: &tokio::runtime::Handle,
) {
let agent_type = self.agent_type.clone();
let result: Option<String> = tokio::task::block_in_place(|| {
rt_handle.block_on(async {
let config_service = GlobalConfigManager::get_service().await.ok()?;
Expand All @@ -3890,14 +3889,7 @@ impl ChatMode {
let global_config: bitfun_core::service::config::GlobalConfig =
config_service.get_config(None).await.ok()?;

// Resolve model ID for the current agent
let model_id = global_config
.ai
.agent_models
.get(&agent_type)
.cloned()
.or_else(|| global_config.ai.default_models.primary.clone())
.unwrap_or_else(|| "primary".to_string());
let model_id = crate::model_selection::resolve_mode_model_id(&global_config.ai)?;

fn provider_display_name(
model: &bitfun_core::service::config::AIModelConfig,
Expand Down Expand Up @@ -3927,20 +3919,10 @@ impl ChatMode {
format!("{} / {}", model.model_name, provider_display_name(model))
}

// Find model name
let model_name = if model_id == "primary" {
// Resolve primary model
let primary_id = global_config.ai.default_models.primary.as_deref()?;
models
.iter()
.find(|m| m.id == primary_id)
.map(model_display_name)
} else {
models
.iter()
.find(|m| m.id == model_id)
.map(model_display_name)
};
let model_name = models
.iter()
.find(|model| model.id == model_id)
.map(model_display_name);

model_name
})
Expand All @@ -3958,7 +3940,6 @@ impl ChatMode {
chat_state: &mut ChatState,
rt_handle: &tokio::runtime::Handle,
) {
let agent_type = self.agent_type.clone();
let result = tokio::task::block_in_place(|| {
rt_handle.block_on(async {
let config_service = match GlobalConfigManager::get_service().await {
Expand All @@ -3974,13 +3955,8 @@ impl ChatMode {
let global_config: bitfun_core::service::config::GlobalConfig =
config_service.get_config(None).await.ok()?;

// Get current model ID
let current_model_id = global_config
.ai
.agent_models
.get(&agent_type)
.cloned()
.or_else(|| global_config.ai.default_models.primary.clone());
let current_model_id =
crate::model_selection::resolve_mode_model_id(&global_config.ai);

// Convert to ModelItem list (only enabled models)
let model_items: Vec<ModelItem> = models
Expand Down Expand Up @@ -4020,10 +3996,19 @@ impl ChatMode {
) {
let selected_id = selected.id.clone();
let selected_display_name = format!("{} / {}", selected.model_name, selected.name);
let modes = self.get_mode_agents(rt_handle);
let session_id = chat_state.core_session_id.clone();

let success = tokio::task::block_in_place(|| {
rt_handle.block_on(async {
if let Err(e) = self
.agent
.update_session_model(&session_id, &selected_id)
.await
{
tracing::error!("Failed to update current session model: {}", e);
return false;
}

let config_service = match GlobalConfigManager::get_service().await {
Ok(s) => s,
Err(e) => {
Expand All @@ -4032,23 +4017,14 @@ impl ChatMode {
}
};

// Update default primary model
if let Err(e) = config_service
.set_config("ai.default_models.primary", &selected_id)
.set_config("ai.agent_model_defaults.mode", &selected_id)
.await
{
tracing::error!("Failed to set default primary model: {}", e);
tracing::error!("Failed to set future mode model: {}", e);
return false;
}

// Update agent_models for all modes
for mode in &modes {
let path = format!("ai.agent_models.{}", mode.id);
if let Err(e) = config_service.set_config(&path, &selected_id).await {
tracing::error!("Failed to set model for mode '{}': {}", mode.id, e);
}
}

true
})
});
Expand Down
36 changes: 21 additions & 15 deletions src/apps/cli/src/modes/exec.rs
Original file line number Diff line number Diff line change
Expand Up @@ -83,7 +83,7 @@ impl ExecTokenUsage {
) -> Option<&'a str> {
let AgenticEvent::TokenUsageUpdated {
turn_id,
model_id,
model_config_id,
input_tokens,
output_tokens,
total_tokens,
Expand All @@ -108,7 +108,7 @@ impl ExecTokenUsage {
} else {
*aggregate = Some(round);
}
Some(model_id)
Some(model_config_id)
}
}

Expand Down Expand Up @@ -760,24 +760,26 @@ impl ExecMode {

self.emit_stream_envelope(&envelope)?;

if let Some(model_id) =
if let Some(model_config_id) =
ExecTokenUsage::accumulate_event(&mut usage, event, &turn_id)
{
self.record_resolved_model_id(&session_id, model_id).await;
self.record_resolved_model_config_id(&session_id, model_config_id)
.await;
}

match event {
AgenticEvent::ModelRoundStarted {
turn_id: event_turn_id,
model_id: Some(model_id),
model_config_id,
..
}
| AgenticEvent::ModelRoundCompleted {
turn_id: event_turn_id,
model_id: Some(model_id),
model_config_id,
..
} if event_turn_id == &turn_id => {
self.record_resolved_model_id(&session_id, model_id).await;
self.record_resolved_model_config_id(&session_id, model_config_id)
.await;
}

AgenticEvent::TextChunk {
Expand Down Expand Up @@ -1071,15 +1073,15 @@ impl ExecMode {
.unwrap_or_else(|| Err(anyhow::anyhow!("Execution ended without a terminal event")))
}

async fn record_resolved_model_id(&self, session_id: &str, model_id: &str) {
let trimmed = model_id.trim();
async fn record_resolved_model_config_id(&self, session_id: &str, model_config_id: &str) {
let trimmed = model_config_id.trim();
if trimmed.is_empty() || matches!(trimmed, "auto" | "default" | "primary" | "fast") {
return;
}

if let Err(error) = self.agent.update_session_model(session_id, trimmed).await {
tracing::debug!(
"Failed to persist resolved CLI model id: session_id={}, model_id={}, error={}",
"Failed to persist resolved CLI model config id: session_id={}, model_config_id={}, error={}",
session_id,
trimmed,
error
Expand Down Expand Up @@ -1508,7 +1510,8 @@ mod patch_tests {
AgenticEvent::TokenUsageUpdated {
session_id: "session-1".to_string(),
turn_id: "turn-1".to_string(),
model_id: "model".to_string(),
model_config_id: "model-config".to_string(),
effective_model_name: "provider-model".to_string(),
input_tokens: 100,
output_tokens: Some(25),
total_tokens: 125,
Expand All @@ -1520,7 +1523,8 @@ mod patch_tests {
AgenticEvent::TokenUsageUpdated {
session_id: "session-1".to_string(),
turn_id: "turn-1".to_string(),
model_id: "model".to_string(),
model_config_id: "model-config".to_string(),
effective_model_name: "provider-model".to_string(),
input_tokens: 200,
output_tokens: Some(50),
total_tokens: 250,
Expand All @@ -1535,7 +1539,7 @@ mod patch_tests {
for event in &events {
assert_eq!(
ExecTokenUsage::accumulate_event(&mut usage, event, "turn-1"),
Some("model")
Some("model-config")
);
}

Expand All @@ -1558,7 +1562,8 @@ mod patch_tests {
AgenticEvent::TokenUsageUpdated {
session_id: "session-1".to_string(),
turn_id: "turn-1".to_string(),
model_id: "model".to_string(),
model_config_id: "model-config".to_string(),
effective_model_name: "provider-model".to_string(),
input_tokens: 100,
output_tokens: None,
total_tokens: 100,
Expand All @@ -1570,7 +1575,8 @@ mod patch_tests {
AgenticEvent::TokenUsageUpdated {
session_id: "session-1".to_string(),
turn_id: "turn-1".to_string(),
model_id: "model".to_string(),
model_config_id: "model-config".to_string(),
effective_model_name: "provider-model".to_string(),
input_tokens: 50,
output_tokens: Some(10),
total_tokens: 60,
Expand Down
Loading
Loading