diff --git a/src-tauri/src/provider.rs b/src-tauri/src/provider.rs index 39ece0ed8..19d8d10e6 100644 --- a/src-tauri/src/provider.rs +++ b/src-tauri/src/provider.rs @@ -283,6 +283,14 @@ pub struct ProviderMeta { /// If not set, provider ID is used automatically during format conversion. #[serde(rename = "promptCacheKey", skip_serializing_if = "Option::is_none")] pub prompt_cache_key: Option, + /// 累加模式应用中,该 provider 是否已写入 live config。 + /// 用于区分仅保存在数据库中的 provider,避免其更新被 live config 解析错误误伤。 + #[serde( + rename = "liveConfigManaged", + default, + skip_serializing_if = "std::ops::Not::not" + )] + pub live_config_managed: bool, /// 供应商类型标识(用于特殊供应商检测) /// - "github_copilot": GitHub Copilot 供应商 #[serde(rename = "providerType", skip_serializing_if = "Option::is_none")] diff --git a/src-tauri/src/proxy/providers/streaming_responses.rs b/src-tauri/src/proxy/providers/streaming_responses.rs index 4ad941f09..ea9274ff8 100644 --- a/src-tauri/src/proxy/providers/streaming_responses.rs +++ b/src-tauri/src/proxy/providers/streaming_responses.rs @@ -974,7 +974,9 @@ mod tests { "data: {\"type\":\"response.completed\",\"response\":{\"status\":\"completed\",\"usage\":{\"input_tokens\":5,\"output_tokens\":2}}}\n\n" ); - let upstream = stream::iter(vec![Ok(Bytes::from(input.as_bytes().to_vec()))]); + let upstream = stream::iter(vec![Ok::<_, std::io::Error>(Bytes::from( + input.as_bytes().to_vec(), + ))]); let converted = create_anthropic_sse_stream_from_responses(upstream); let chunks: Vec<_> = converted.collect().await; let events: Vec = chunks diff --git a/src-tauri/src/services/provider/mod.rs b/src-tauri/src/services/provider/mod.rs index 7e8ab80ad..4c8425e1a 100644 --- a/src-tauri/src/services/provider/mod.rs +++ b/src-tauri/src/services/provider/mod.rs @@ -52,7 +52,66 @@ pub struct SwitchResult { #[cfg(test)] mod tests { use super::*; + use crate::database::Database; + use crate::provider::ProviderMeta; + use crate::store::AppState; use serde_json::json; + use std::fs; + use std::path::Path; + use std::sync::{Arc, Mutex, OnceLock}; + + fn test_guard() -> std::sync::MutexGuard<'static, ()> { + static LOCK: OnceLock> = OnceLock::new(); + LOCK.get_or_init(|| Mutex::new(())) + .lock() + .unwrap_or_else(|err| err.into_inner()) + } + + fn with_test_home(test: impl FnOnce(&AppState, &Path) -> T) -> T { + let _guard = test_guard(); + let temp = tempfile::tempdir().expect("tempdir"); + let old_test_home = std::env::var_os("CC_SWITCH_TEST_HOME"); + let old_home = std::env::var_os("HOME"); + std::env::set_var("CC_SWITCH_TEST_HOME", temp.path()); + std::env::set_var("HOME", temp.path()); + + let db = Arc::new(Database::memory().expect("in-memory database")); + let state = AppState::new(db); + let result = test(&state, temp.path()); + + match old_test_home { + Some(value) => std::env::set_var("CC_SWITCH_TEST_HOME", value), + None => std::env::remove_var("CC_SWITCH_TEST_HOME"), + } + match old_home { + Some(value) => std::env::set_var("HOME", value), + None => std::env::remove_var("HOME"), + } + + result + } + + fn openclaw_provider(id: &str) -> Provider { + Provider { + id: id.to_string(), + name: format!("Provider {id}"), + settings_config: json!({ + "baseUrl": "https://api.deepseek.com", + "apiKey": "test-key", + "api": "openai-completions", + "models": [], + }), + website_url: None, + category: Some("custom".to_string()), + created_at: Some(1), + sort_index: Some(0), + notes: None, + meta: None, + icon: None, + icon_color: None, + in_failover_queue: false, + } + } #[test] fn validate_provider_settings_rejects_missing_auth() { @@ -129,6 +188,78 @@ base_url = "http://localhost:8080" "should keep mcp_servers.* base_url" ); } + + #[test] + fn rename_rejects_missing_original_provider() { + with_test_home(|state, _| { + let original = openclaw_provider("deepseek"); + ProviderService::add(state, AppType::OpenClaw, original.clone(), false) + .expect("seed db-only provider"); + + let mut renamed = original.clone(); + renamed.id = "deepseek-copy".to_string(); + + let err = ProviderService::update( + state, + AppType::OpenClaw, + Some("missing-provider"), + renamed, + ) + .expect_err("stale originalId should be rejected"); + + assert!( + err.to_string().contains("Original provider"), + "expected missing original provider error, got {err:?}" + ); + assert!( + state + .db + .get_provider_by_id("deepseek-copy", AppType::OpenClaw.as_str()) + .expect("query renamed provider") + .is_none(), + "rename must not create a new row when originalId is stale" + ); + }); + } + + #[test] + fn db_only_additive_update_survives_live_config_parse_errors() { + with_test_home(|state, home| { + let provider = openclaw_provider("deepseek"); + ProviderService::add(state, AppType::OpenClaw, provider.clone(), false) + .expect("seed db-only provider"); + + let stored = state + .db + .get_provider_by_id("deepseek", AppType::OpenClaw.as_str()) + .expect("query stored provider") + .expect("provider should exist"); + assert_eq!( + stored.meta.as_ref().map(|meta| meta.live_config_managed), + Some(false), + "db-only provider should be marked as not live-managed" + ); + + let openclaw_dir = home.join(".openclaw"); + fs::create_dir_all(&openclaw_dir).expect("create openclaw dir"); + fs::write(openclaw_dir.join("openclaw.json"), "{ invalid json5") + .expect("write malformed config"); + + let mut updated = stored.clone(); + updated.name = "DeepSeek Edited".to_string(); + updated.meta.get_or_insert_with(ProviderMeta::default); + + ProviderService::update(state, AppType::OpenClaw, None, updated) + .expect("db-only update should ignore live parse errors"); + + let saved = state + .db + .get_provider_by_id("deepseek", AppType::OpenClaw.as_str()) + .expect("query updated provider") + .expect("updated provider should exist"); + assert_eq!(saved.name, "DeepSeek Edited"); + }); + } } impl ProviderService { @@ -141,6 +272,20 @@ impl ProviderService { } } + /// Check whether a provider exists in live config, tolerating parse errors + /// for providers that have never been written to live (`live_config_managed == false`). + fn check_live_config_exists( + app_type: &AppType, + provider_id: &str, + is_db_only: bool, + ) -> Result { + if is_db_only { + Ok(provider_exists_in_live_config(app_type, provider_id).unwrap_or(false)) + } else { + provider_exists_in_live_config(app_type, provider_id) + } + } + /// List all providers for an app type pub fn list( state: &AppState, @@ -177,6 +322,12 @@ impl ProviderService { Self::normalize_provider_if_claude(&app_type, &mut provider); Self::validate_provider_settings(&app_type, &provider)?; normalize_provider_common_config_for_storage(state.db.as_ref(), &app_type, &mut provider)?; + if app_type.is_additive_mode() { + provider + .meta + .get_or_insert_with(Default::default) + .live_config_managed = add_to_live; + } // Save to database state.db.save_provider(app_type.as_str(), &provider)?; @@ -221,6 +372,9 @@ impl ProviderService { let mut provider = provider; let original_id = original_id.unwrap_or(provider.id.as_str()).to_string(); let provider_id_changed = original_id != provider.id; + let existing_provider = state + .db + .get_provider_by_id(&original_id, app_type.as_str())?; // Normalize Claude model keys Self::normalize_provider_if_claude(&app_type, &mut provider); Self::validate_provider_settings(&app_type, &provider)?; @@ -233,18 +387,33 @@ impl ProviderService { )); } - if provider_exists_in_live_config(&app_type, &original_id)? { + let Some(existing_provider) = existing_provider else { + return Err(AppError::Message(format!( + "Original provider '{}' does not exist in app '{}'", + original_id, + app_type.as_str() + ))); + }; + let is_db_only = !existing_provider + .meta + .as_ref() + .is_some_and(|m| m.live_config_managed); + let original_in_live = + Self::check_live_config_exists(&app_type, &original_id, is_db_only)?; + if original_in_live { return Err(AppError::Message( "Provider key cannot be changed after the provider has been added to the app config" .to_string(), )); } + let next_id_in_live = + Self::check_live_config_exists(&app_type, &provider.id, is_db_only)?; if state .db .get_provider_by_id(&provider.id, app_type.as_str())? .is_some() - || provider_exists_in_live_config(&app_type, &provider.id)? + || next_id_in_live { return Err(AppError::Message(format!( "Provider '{}' already exists in app '{}'", @@ -253,6 +422,10 @@ impl ProviderService { ))); } + provider + .meta + .get_or_insert_with(Default::default) + .live_config_managed = false; state.db.save_provider(app_type.as_str(), &provider)?; state.db.delete_provider(app_type.as_str(), &original_id)?; @@ -263,9 +436,6 @@ impl ProviderService { return Ok(true); } - // Save to database - state.db.save_provider(app_type.as_str(), &provider)?; - // Additive mode apps (OpenCode, OpenClaw): only sync to live when the provider // already exists in live config. Editing a DB-only provider must not auto-add it. if app_type.is_additive_mode() { @@ -299,13 +469,31 @@ impl ProviderService { } return Ok(true); } - if !provider_exists_in_live_config(&app_type, &provider.id)? { + let is_db_only = !provider + .meta + .as_ref() + .is_some_and(|m| m.live_config_managed); + let live_config_managed = + Self::check_live_config_exists(&app_type, &provider.id, is_db_only)?; + provider + .meta + .get_or_insert_with(Default::default) + .live_config_managed = live_config_managed; + + // Save to database after live-config presence is resolved so parse errors + // do not report failure after already mutating DB state. + state.db.save_provider(app_type.as_str(), &provider)?; + + if !live_config_managed { return Ok(true); } write_live_with_common_config(state.db.as_ref(), &app_type, &provider)?; return Ok(true); } + // Save to database + state.db.save_provider(app_type.as_str(), &provider)?; + // For other apps: Check if this is current provider (use effective current, not just DB) let effective_current = crate::settings::get_effective_current_provider(&state.db, &app_type)?; @@ -474,6 +662,15 @@ impl ProviderService { ))); } } + + if let Some(mut provider) = state.db.get_provider_by_id(id, app_type.as_str())? { + provider + .meta + .get_or_insert_with(Default::default) + .live_config_managed = false; + state.db.save_provider(app_type.as_str(), &provider)?; + } + Ok(()) } diff --git a/src/App.tsx b/src/App.tsx index 68fc878cb..20e88b336 100644 --- a/src/App.tsx +++ b/src/App.tsx @@ -615,7 +615,19 @@ function App() { }; if (activeApp === "opencode" || activeApp === "openclaw") { - const existingKeys = Object.keys(providers); + const liveProviderIds = + activeApp === "opencode" + ? await queryClient.ensureQueryData({ + queryKey: ["opencodeLiveProviderIds"], + queryFn: () => providersApi.getOpenCodeLiveProviderIds(), + }) + : await queryClient.ensureQueryData({ + queryKey: openclawKeys.liveProviderIds, + queryFn: () => providersApi.getOpenClawLiveProviderIds(), + }); + const existingKeys = Array.from( + new Set([...Object.keys(providers), ...liveProviderIds]), + ); duplicatedProvider.providerKey = generateUniqueProviderCopyKey( provider.id, existingKeys, diff --git a/tests/hooks/useProviderActions.test.tsx b/tests/hooks/useProviderActions.test.tsx index d6d858d72..8a7a2e8b4 100644 --- a/tests/hooks/useProviderActions.test.tsx +++ b/tests/hooks/useProviderActions.test.tsx @@ -169,7 +169,10 @@ describe("useProviderActions", () => { await result.current.updateProvider(provider); }); - expect(updateProviderMutateAsync).toHaveBeenCalledWith(provider); + expect(updateProviderMutateAsync).toHaveBeenCalledWith({ + provider, + originalId: undefined, + }); expect(providersApiUpdateTrayMenuMock).toHaveBeenCalledTimes(1); }); diff --git a/tests/integration/App.test.tsx b/tests/integration/App.test.tsx index 76fd3c253..2524be0c0 100644 --- a/tests/integration/App.test.tsx +++ b/tests/integration/App.test.tsx @@ -5,6 +5,7 @@ import { describe, it, expect, beforeEach, vi } from "vitest"; import { resetProviderState, setCurrentProviderId, + setLiveProviderIds, setProviders, } from "../msw/state"; import { emitTauriEvent } from "../msw/tauriMocks"; @@ -239,7 +240,7 @@ describe("App integration with MSW", () => { }); }); - it("duplicates openclaw providers with a generated provider key", async () => { + it("duplicates openclaw providers with a generated key that avoids live-only ids", async () => { setProviders("openclaw", { deepseek: { id: "deepseek", @@ -256,6 +257,7 @@ describe("App integration with MSW", () => { }, }); setCurrentProviderId("openclaw", "deepseek"); + setLiveProviderIds("openclaw", ["deepseek-copy"]); const { default: App } = await import("@/App"); renderApp(App); @@ -272,7 +274,7 @@ describe("App integration with MSW", () => { await waitFor(() => { const providerList = screen.getByTestId("provider-list").textContent; - expect(providerList).toContain("deepseek-copy"); + expect(providerList).toContain("deepseek-copy-2"); expect(providerList).toContain("DeepSeek copy"); }); diff --git a/tests/msw/handlers.ts b/tests/msw/handlers.ts index 3bd410bc3..41212cb2c 100644 --- a/tests/msw/handlers.ts +++ b/tests/msw/handlers.ts @@ -6,6 +6,7 @@ import { deleteProvider, deleteSession, getCurrentProviderId, + getLiveProviderIds, getSessionMessages, getProviders, listProviders, @@ -67,8 +68,12 @@ export const handlers = [ http.post(`${TAURI_ENDPOINT}/update_tray_menu`, () => success(true)), + http.post(`${TAURI_ENDPOINT}/get_opencode_live_provider_ids`, () => + success(getLiveProviderIds("opencode")), + ), + http.post(`${TAURI_ENDPOINT}/get_openclaw_live_provider_ids`, () => - success([]), + success(getLiveProviderIds("openclaw")), ), http.post(`${TAURI_ENDPOINT}/get_openclaw_default_model`, () => diff --git a/tests/msw/state.ts b/tests/msw/state.ts index 291d5201b..9768914d9 100644 --- a/tests/msw/state.ts +++ b/tests/msw/state.ts @@ -10,6 +10,7 @@ import type { type ProvidersByApp = Record>; type CurrentProviderState = Record; type McpConfigState = Record>; +type LiveProviderIdsByApp = Record<"opencode" | "openclaw", string[]>; const createDefaultProviders = (): ProvidersByApp => ({ claude: { @@ -77,6 +78,10 @@ const createDefaultCurrent = (): CurrentProviderState => ({ let providers = createDefaultProviders(); let current = createDefaultCurrent(); +let liveProviderIds: LiveProviderIdsByApp = { + opencode: [], + openclaw: [], +}; let settingsState: Settings = { showInTray: true, minimizeToTrayOnClose: true, @@ -184,6 +189,10 @@ const cloneProviders = (value: ProvidersByApp) => export const resetProviderState = () => { providers = createDefaultProviders(); current = createDefaultCurrent(); + liveProviderIds = { + opencode: [], + openclaw: [], + }; sessionsState = createDefaultSessions(); sessionMessagesState = createDefaultSessionMessages(); settingsState = { @@ -243,6 +252,17 @@ export const getProviders = (appType: AppId) => export const getCurrentProviderId = (appType: AppId) => current[appType] ?? ""; +export const getLiveProviderIds = (appType: "opencode" | "openclaw") => [ + ...liveProviderIds[appType], +]; + +export const setLiveProviderIds = ( + appType: "opencode" | "openclaw", + ids: string[], +) => { + liveProviderIds[appType] = [...ids]; +}; + export const setCurrentProviderId = (appType: AppId, providerId: string) => { current[appType] = providerId; };