From f62937d4820d6010099b974df797f4d4fc0a557a Mon Sep 17 00:00:00 2001 From: Zaki Manian Date: Sun, 1 Mar 2026 00:32:59 -0800 Subject: [PATCH] fix: persist model name to .env so dotted names survive restart (#426) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * fix: persist model name to .env so dotted names survive restart (#400) The setup wizard saved selected_model to the DB but not to .env. Since Config::from_env_with_toml() runs before the DB connects, the model name was lost on restart -- backends fell back to hardcoded defaults, truncating names like "llama3.2" to "llama3". - Add LlmBackend::model_env_var() as single source of truth for the backend-to-env-var mapping - Write the model env var in write_bootstrap_env() using the new method - Add selected_model fallback to all 6 backends (was missing from OpenAI, Anthropic, Ollama, and Tinfoil) Co-Authored-By: Claude Opus 4.6 * refactor: extract resolve_model() helper to reduce duplication Address review feedback: the env → settings → default model resolution pattern was repeated across all 6 backends. Centralise it in a single LlmConfig::resolve_model() helper. Co-Authored-By: Claude Opus 4.6 --------- Co-authored-by: Claude Opus 4.6 --- src/config/llm.rs | 122 ++++++++++++++++++++++++++++++++++++++++---- src/setup/wizard.rs | 14 +++++ 2 files changed, 125 insertions(+), 11 deletions(-) diff --git a/src/config/llm.rs b/src/config/llm.rs index 2db37f98..ba42ed9d 100644 --- a/src/config/llm.rs +++ b/src/config/llm.rs @@ -60,6 +60,23 @@ impl std::fmt::Display for LlmBackend { } } +impl LlmBackend { + /// The environment variable that configures the model name for this backend. + /// + /// Used by both `LlmConfig::resolve()` (reads the var) and the setup wizard + /// (writes the var to `.env`). Centralised here so the two stay in sync. + pub fn model_env_var(&self) -> &'static str { + match self { + Self::NearAi => "NEARAI_MODEL", + Self::OpenAi => "OPENAI_MODEL", + Self::Anthropic => "ANTHROPIC_MODEL", + Self::Ollama => "OLLAMA_MODEL", + Self::OpenAiCompatible => "LLM_MODEL", + Self::Tinfoil => "TINFOIL_MODEL", + } + } +} + /// Configuration for direct OpenAI API access. #[derive(Debug, Clone)] pub struct OpenAiDirectConfig { @@ -178,6 +195,17 @@ pub struct NearAiConfig { } impl LlmConfig { + /// Resolve a model name from env var → settings.selected_model → hardcoded default. + fn resolve_model( + env_var: &str, + settings: &Settings, + default: &str, + ) -> Result { + Ok(optional_env(env_var)? + .or_else(|| settings.selected_model.clone()) + .unwrap_or_else(|| default.to_string())) + } + pub(crate) fn resolve(settings: &Settings) -> Result { // Determine backend: env var > settings > default (NearAi) let backend: LlmBackend = if let Some(b) = optional_env("LLM_BACKEND")? { @@ -205,9 +233,7 @@ impl LlmConfig { let nearai_api_key = optional_env("NEARAI_API_KEY")?.map(SecretString::from); let nearai = NearAiConfig { - model: optional_env("NEARAI_MODEL")? - .or_else(|| settings.selected_model.clone()) - .unwrap_or_else(|| "zai-org/GLM-latest".to_string()), + model: Self::resolve_model("NEARAI_MODEL", settings, "zai-org/GLM-latest")?, cheap_model: optional_env("NEARAI_CHEAP_MODEL")?, base_url: optional_env("NEARAI_BASE_URL")?.unwrap_or_else(|| { if nearai_api_key.is_some() { @@ -248,7 +274,7 @@ impl LlmConfig { key: "OPENAI_API_KEY".to_string(), hint: "Set OPENAI_API_KEY when LLM_BACKEND=openai".to_string(), })?; - let model = optional_env("OPENAI_MODEL")?.unwrap_or_else(|| "gpt-4o".to_string()); + let model = Self::resolve_model("OPENAI_MODEL", settings, "gpt-4o")?; let base_url = optional_env("OPENAI_BASE_URL")?; Some(OpenAiDirectConfig { api_key, @@ -266,8 +292,8 @@ impl LlmConfig { key: "ANTHROPIC_API_KEY".to_string(), hint: "Set ANTHROPIC_API_KEY when LLM_BACKEND=anthropic".to_string(), })?; - let model = optional_env("ANTHROPIC_MODEL")? - .unwrap_or_else(|| "claude-sonnet-4-20250514".to_string()); + let model = + Self::resolve_model("ANTHROPIC_MODEL", settings, "claude-sonnet-4-20250514")?; let base_url = optional_env("ANTHROPIC_BASE_URL")?; Some(AnthropicDirectConfig { api_key, @@ -282,7 +308,7 @@ impl LlmConfig { let base_url = optional_env("OLLAMA_BASE_URL")? .or_else(|| settings.ollama_base_url.clone()) .unwrap_or_else(|| "http://localhost:11434".to_string()); - let model = optional_env("OLLAMA_MODEL")?.unwrap_or_else(|| "llama3".to_string()); + let model = Self::resolve_model("OLLAMA_MODEL", settings, "llama3")?; Some(OllamaConfig { base_url, model }) } else { None @@ -296,9 +322,7 @@ impl LlmConfig { hint: "Set LLM_BASE_URL when LLM_BACKEND=openai_compatible".to_string(), })?; let api_key = optional_env("LLM_API_KEY")?.map(SecretString::from); - let model = optional_env("LLM_MODEL")? - .or_else(|| settings.selected_model.clone()) - .unwrap_or_else(|| "default".to_string()); + let model = Self::resolve_model("LLM_MODEL", settings, "default")?; let extra_headers = optional_env("LLM_EXTRA_HEADERS")? .map(|val| parse_extra_headers(&val)) .transpose()? @@ -320,7 +344,7 @@ impl LlmConfig { key: "TINFOIL_API_KEY".to_string(), hint: "Set TINFOIL_API_KEY when LLM_BACKEND=tinfoil".to_string(), })?; - let model = optional_env("TINFOIL_MODEL")?.unwrap_or_else(|| "kimi-k2-5".to_string()); + let model = Self::resolve_model("TINFOIL_MODEL", settings, "kimi-k2-5")?; Some(TinfoilConfig { api_key, model }) } else { None @@ -506,4 +530,80 @@ mod tests { ] ); } + + /// Clear all ollama-related env vars. + fn clear_ollama_env() { + // SAFETY: Only called under ENV_MUTEX in tests. + unsafe { + std::env::remove_var("LLM_BACKEND"); + std::env::remove_var("OLLAMA_BASE_URL"); + std::env::remove_var("OLLAMA_MODEL"); + } + } + + #[test] + fn ollama_uses_selected_model_when_ollama_model_unset() { + let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + clear_ollama_env(); + + let settings = Settings { + llm_backend: Some("ollama".to_string()), + selected_model: Some("llama3.2".to_string()), + ..Default::default() + }; + + let cfg = LlmConfig::resolve(&settings).expect("resolve should succeed"); + let ollama = cfg.ollama.expect("ollama config should be present"); + + assert_eq!(ollama.model, "llama3.2"); + } + + #[test] + fn ollama_model_env_overrides_selected_model() { + let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + clear_ollama_env(); + // SAFETY: Under ENV_MUTEX. + unsafe { + std::env::set_var("OLLAMA_MODEL", "mistral:latest"); + } + + let settings = Settings { + llm_backend: Some("ollama".to_string()), + selected_model: Some("llama3.2".to_string()), + ..Default::default() + }; + + let cfg = LlmConfig::resolve(&settings).expect("resolve should succeed"); + let ollama = cfg.ollama.expect("ollama config should be present"); + + assert_eq!(ollama.model, "mistral:latest"); + + // SAFETY: Under ENV_MUTEX. + unsafe { + std::env::remove_var("OLLAMA_MODEL"); + } + } + + #[test] + fn openai_compatible_preserves_dotted_model_name() { + let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + clear_openai_compatible_env(); + + let settings = Settings { + llm_backend: Some("openai_compatible".to_string()), + openai_compatible_base_url: Some("http://localhost:11434/v1".to_string()), + selected_model: Some("llama3.2".to_string()), + ..Default::default() + }; + + let cfg = LlmConfig::resolve(&settings).expect("resolve should succeed"); + let compat = cfg + .openai_compatible + .expect("openai-compatible config should be present"); + + assert_eq!( + compat.model, "llama3.2", + "model name with dot must not be truncated" + ); + } } diff --git a/src/setup/wizard.rs b/src/setup/wizard.rs index 55efb68c..58647a86 100644 --- a/src/setup/wizard.rs +++ b/src/setup/wizard.rs @@ -2040,6 +2040,20 @@ impl SetupWizard { env_vars.push(("OLLAMA_BASE_URL", url.clone())); } + // Model name: same chicken-and-egg — Config::from_env() resolves the + // model before the DB is connected, so we must persist it to .env. + // Write the backend-specific env var so the correct resolution path + // picks it up. + if let Some(ref model) = self.settings.selected_model { + let backend: crate::config::LlmBackend = self + .settings + .llm_backend + .as_deref() + .and_then(|s| s.parse().ok()) + .unwrap_or_default(); + env_vars.push((backend.model_env_var(), model.clone())); + } + // Preserve NEARAI_API_KEY if present (set by API key auth flow) if let Ok(api_key) = std::env::var("NEARAI_API_KEY") && !api_key.is_empty()