diff --git a/apps/desktop/src/components/editor/EditorSettingsDialog.vue b/apps/desktop/src/components/editor/EditorSettingsDialog.vue index 754e21fe5..c753230cd 100644 --- a/apps/desktop/src/components/editor/EditorSettingsDialog.vue +++ b/apps/desktop/src/components/editor/EditorSettingsDialog.vue @@ -942,6 +942,7 @@ const aiCompletionsMode = computed(() => aiEditApiStyle.value === "completions") const aiTesting = ref(false); const aiTestResult = ref<"" | "success" | "error">(""); const aiTestError = ref(""); +const aiTestLatency = ref(null); const aiRequiresApiKey = computed(() => AI_PROVIDER_PRESETS[aiEditProvider.value].requiresApiKey); const aiSupportsAuthMethod = computed(() => aiEditProvider.value === "claude"); const aiCredentialLabel = computed(() => (aiSupportsAuthMethod.value && aiEditAuthMethod.value === "bearer" ? "Auth Token" : "API Key")); @@ -1077,6 +1078,7 @@ function syncAiEditState() { aiEditEnableThinking.value = settingsStore.aiConfig.enableThinking ?? true; aiTestResult.value = ""; aiTestError.value = ""; + aiTestLatency.value = null; clearAiModelOptions(); } @@ -1087,6 +1089,9 @@ function aiSelectProvider(provider: AiProvider) { aiEditApiStyle.value = AI_PROVIDER_PRESETS[provider].apiStyle; aiEditAuthMethod.value = AI_PROVIDER_PRESETS[provider].authMethod; if (!AI_PROVIDER_PRESETS[provider].requiresApiKey) aiEditApiKey.value = ""; + aiTestResult.value = ""; + aiTestError.value = ""; + aiTestLatency.value = null; clearAiModelOptions(); } @@ -1113,9 +1118,11 @@ async function aiTestConn() { aiTesting.value = true; aiTestResult.value = ""; aiTestError.value = ""; + aiTestLatency.value = null; try { - await aiTestConnection(currentAiEditConfig()); + const result = await aiTestConnection(currentAiEditConfig()); aiTestResult.value = "success"; + aiTestLatency.value = result.latencyMs ?? null; } catch (e: any) { aiTestResult.value = "error"; aiTestError.value = e?.message || String(e); @@ -2455,8 +2462,9 @@ watch( {{ t("connection.test") }} - - {{ t("connection.testSuccess") }} + + {{ t("connection.testSuccess") }} + {{ aiTestLatency }}ms {{ aiTestError }} diff --git a/apps/desktop/src/lib/http.ts b/apps/desktop/src/lib/http.ts index b91d73274..db1e37c18 100644 --- a/apps/desktop/src/lib/http.ts +++ b/apps/desktop/src/lib/http.ts @@ -27,7 +27,7 @@ import type { } from "@/types/database"; import type { SchemaDiffPreparation, SchemaDiffPreparationOptions, TableDiff, FunctionDiff, SequenceDiff, RuleDiff, OwnerDiff } from "@/lib/schemaDiff"; import type { SidebarObjectKind } from "@/lib/databaseObjectCapabilities"; -import type { AiConfig } from "@/stores/settingsStore"; +import type { AiConfig, AiTestConnectionResult } from "@/stores/settingsStore"; import type { AgentDriverInfo, AiCompletionRequest, @@ -778,7 +778,7 @@ export async function aiCancelStream(sessionId: string): Promise { return post("/api/ai/cancel-stream", { sessionId }); } -export async function aiTestConnection(config: AiConfig): Promise { +export async function aiTestConnection(config: AiConfig): Promise { return post("/api/ai/test-connection", { config }); } diff --git a/apps/desktop/src/lib/tauri.ts b/apps/desktop/src/lib/tauri.ts index 5e83b8264..18e4098fd 100644 --- a/apps/desktop/src/lib/tauri.ts +++ b/apps/desktop/src/lib/tauri.ts @@ -27,7 +27,7 @@ import type { SavedSqlLibrary, } from "@/types/database"; import type { SidebarObjectKind } from "@/lib/databaseObjectCapabilities"; -import type { AiConfig } from "@/stores/settingsStore"; +import type { AiConfig, AiTestConnectionResult } from "@/stores/settingsStore"; import type { QueryEditability } from "@/lib/sqlAnalysis"; import type { DataGridColumnValueFilterConditionOptions, DataGridContextFilterConditionOptions, DataGridCountSqlOptions, DataGridCopyInsertStatementOptions, DataGridCopyUpdateStatementOptions, DataGridSaveStatementOptions, HiveTablePropertiesSqlOptions } from "@/lib/dataGridSql"; import type { DataCompareFromTablesOptions, DataCompareFromTablesPreparation, DataCompareSyncPlan, DataCompareSyncPlanOptions, DataComparePreparation, DataComparePreparationOptions } from "@/lib/dataCompare"; @@ -307,7 +307,7 @@ export async function saveAiConfig(config: AiConfig): Promise { return invoke("save_ai_config", { config }); } -export async function aiTestConnection(config: AiConfig): Promise { +export async function aiTestConnection(config: AiConfig): Promise { return invoke("ai_test_connection", { config }); } diff --git a/apps/desktop/src/stores/settingsStore.ts b/apps/desktop/src/stores/settingsStore.ts index 71700cc06..63f4071f4 100644 --- a/apps/desktop/src/stores/settingsStore.ts +++ b/apps/desktop/src/stores/settingsStore.ts @@ -27,6 +27,14 @@ export interface AiConfig { enableThinking?: boolean; } +export interface AiTestConnectionResult { + success: boolean; + message: string; + latencyMs?: number; + modelUsed: string; + errorCategory?: string; +} + export interface DesktopSettings { show_tray_icon: boolean; icon_theme: DesktopIconTheme; diff --git a/crates/dbx-core/src/ai.rs b/crates/dbx-core/src/ai.rs index 9ba3918e7..11f645e2d 100644 --- a/crates/dbx-core/src/ai.rs +++ b/crates/dbx-core/src/ai.rs @@ -169,6 +169,20 @@ pub struct AiModelInfo { pub display_name: Option, } +/// Result of an AI connection test (mirrors CC-Switch's StreamCheckResult). +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct AiTestConnectionResult { + pub success: bool, + pub message: String, + /// First-chunk latency in milliseconds, if successful. + pub latency_ms: Option, + pub model_used: String, + /// Error category for the frontend to render specific guidance. + #[serde(skip_serializing_if = "Option::is_none")] + pub error_category: Option, +} + // --------------------------------------------------------------------------- // Pure helpers // --------------------------------------------------------------------------- @@ -721,41 +735,184 @@ pub async fn call_gemini(client: &reqwest::Client, request: AiCompletionRequest) // High-level: test_connection_core / complete // --------------------------------------------------------------------------- -pub async fn test_connection_core(config: &AiConfig) -> Result { - validate_config(config)?; +/// Read the SSE byte stream until the first content-bearing chunk arrives, +/// then return its latency and the delta text. Used by `test_connection_core` +/// to mirror CC-Switch's streaming probe approach. +async fn measure_first_stream_chunk( + mut byte_stream: impl futures::Stream> + Unpin, + start: std::time::Instant, + is_claude: bool, + is_gemini: bool, +) -> Result<(u64, String), String> { + let mut buf = String::new(); + while let Some(chunk) = byte_stream.next().await { + let chunk = chunk.map_err(|e| format!("stream read error: {e}"))?; + buf.push_str(&String::from_utf8_lossy(&chunk)); - let client = build_ai_http_client(config, 15)?; + while let Some(pos) = buf.find('\n') { + let line = buf[..pos].to_string(); + buf = buf[pos + 1..].to_string(); - let request = AiCompletionRequest { - config: config.clone(), - system_prompt: String::new(), - messages: vec![AiMessage { - role: "user".into(), - content: "hi".into(), - tool_call_id: None, - tool_calls: Vec::new(), - }], - max_tokens: Some(16), - temperature: Some(0.0), - }; + let Some(data) = stream_data_payload(&line) else { continue }; + if data == "[DONE]" { + // stream finished without content — not a real failure but rare + return Err("no content in response".to_string()); + } - match request.config.provider { - AiProvider::Claude => call_claude(&client, request).await, - AiProvider::Gemini => call_gemini(&client, request).await, - AiProvider::Openai - | AiProvider::Deepseek - | AiProvider::Qwen - | AiProvider::Ollama - | AiProvider::OpenaiCompatible - | AiProvider::Custom => { - if request.config.api_style == AiApiStyle::Responses { - call_responses_api(&client, request).await + // Parse the JSON to extract the text delta + let parsed: serde_json::Value = serde_json::from_str(data).map_err(|e| format!("JSON parse error: {e}"))?; + + let delta = if is_claude { + claude_stream_text(&parsed).map(|s| s.to_string()) + } else if is_gemini { + let text = gemini_text(&parsed); + if text.is_empty() { + None + } else { + Some(text) + } } else { - call_openai_compatible(&client, request).await + openai_stream_text(&parsed).or_else(|| Some(responses_text(&parsed))) + }; + + if let Some(text) = delta { + if !text.trim().is_empty() { + let latency = start.elapsed().as_millis() as u64; + return Ok((latency, text)); + } } } } - .map(|_| "OK".to_string()) + Err("stream ended without content".to_string()) +} + +const TEST_PROMPT: &str = "Who are you?"; + +pub async fn test_connection_core(config: &AiConfig) -> Result { + validate_config(config)?; + + let client = build_ai_http_client(config, 15)?; + let start = std::time::Instant::now(); + + let is_claude = matches!(config.provider, AiProvider::Claude); + let is_gemini = matches!(config.provider, AiProvider::Gemini); + let model = config.model.clone(); + + // Build the streaming request and get the byte stream + let byte_stream = match config.provider { + AiProvider::Claude => { + let body = json!({ + "model": &model, + "max_tokens": 16, + "temperature": 0.0, + "system": "", + "messages": [{ "role": "user", "content": TEST_PROMPT }], + "stream": true, + }); + let res = client + .post(resolve_endpoint(config)) + .headers(claude_headers(config)?) + .json(&body) + .send() + .await + .map_err(|e| format!("Claude request failed: {e}"))?; + if !res.status().is_success() { + let data: serde_json::Value = res.json().await.map_err(|e| e.to_string())?; + return Err(categorize_error(&data, config)); + } + res.bytes_stream() + } + AiProvider::Gemini => { + let ep = resolve_endpoint(config); + let res = client + .post(&ep) + .header(CONTENT_TYPE, "application/json") + .query(&[("key", config.api_key.as_str()), ("alt", "sse")]) + .json(&json!({ + "contents": [{ "parts": [{ "text": TEST_PROMPT }], "role": "user" }], + "generationConfig": { "maxOutputTokens": 16, "temperature": 0.0 }, + })) + .send() + .await + .map_err(|e| format!("Gemini request failed: {e}"))?; + if !res.status().is_success() { + let data: serde_json::Value = res.json().await.map_err(|e| e.to_string())?; + return Err(categorize_error(&data, config)); + } + res.bytes_stream() + } + _ => { + // OpenAI-compatible providers + let messages = vec![json!({ "role": "user", "content": TEST_PROMPT })]; + let mut body_obj = json!({ + "model": &model, + "messages": messages, + "max_tokens": 16, + "temperature": 0.0, + "stream": true, + }); + if !config.enable_thinking { + body_obj["extra_body"] = json!({ + "chat_template_kwargs": { "enable_thinking": false } + }); + } + let ep = resolve_endpoint(config); + let res = client + .post(&ep) + .headers(maybe_bearer_headers(config)?) + .json(&body_obj) + .send() + .await + .map_err(|e| format!("AI request failed: {e}"))?; + if !res.status().is_success() { + let data: serde_json::Value = res.json().await.map_err(|e| e.to_string())?; + return Err(categorize_error(&data, config)); + } + res.bytes_stream() + } + }; + + match measure_first_stream_chunk(byte_stream, start, is_claude, is_gemini).await { + Ok((latency, _delta)) => Ok(AiTestConnectionResult { + success: true, + message: format!("OK — {}ms", latency), + latency_ms: Some(latency), + model_used: model, + error_category: None, + }), + Err(e) => { + let category = classify_error(&e); + Err(format!("[{category}] {e}")) + } + } +} + +/// Map known API error bodies to a short category string. +fn categorize_error(data: &serde_json::Value, _config: &AiConfig) -> String { + let raw = extract_error(data).unwrap_or_else(|| "API error".to_string()); + let category = classify_error(&raw); + format!("[{category}] {raw}") +} + +fn classify_error(msg: &str) -> &'static str { + let lower = msg.to_ascii_lowercase(); + if lower.contains("401") + || lower.contains("unauthorized") + || lower.contains("invalid api key") + || lower.contains("incorrect api key") + { + "auth" + } else if lower.contains("404") || lower.contains("not found") || lower.contains("model not found") { + "modelNotFound" + } else if lower.contains("429") || lower.contains("rate limit") || lower.contains("too many requests") { + "rateLimit" + } else if lower.contains("timeout") || lower.contains("timed out") { + "timeout" + } else if lower.contains("connect") || lower.contains("dns") || lower.contains("resolve") { + "network" + } else { + "unknown" + } } pub async fn complete(request: &AiCompletionRequest) -> Result { diff --git a/crates/dbx-web/src/routes/ai.rs b/crates/dbx-web/src/routes/ai.rs index c983bea94..8e9ca0021 100644 --- a/crates/dbx-web/src/routes/ai.rs +++ b/crates/dbx-web/src/routes/ai.rs @@ -8,7 +8,7 @@ use serde::Deserialize; use dbx_core::agent_events::AgentEvent; use dbx_core::agent_loop::{run_agent_loop, AgentLoopContext}; -use dbx_core::ai::{AiCompletionRequest, AiConfig, AiConversation, AiModelInfo, AiStreamChunk}; +use dbx_core::ai::{AiCompletionRequest, AiConfig, AiConversation, AiModelInfo, AiStreamChunk, AiTestConnectionResult}; use dbx_core::models::connection::DatabaseType; use crate::error::AppError; @@ -126,7 +126,9 @@ pub async fn ai_complete(Json(body): Json) -> Result) -> Result, AppError> { +pub async fn ai_test_connection( + Json(body): Json, +) -> Result, AppError> { let result = dbx_core::ai::test_connection_core(&body.config).await.map_err(AppError)?; Ok(Json(result)) } diff --git a/src-tauri/src/commands/ai.rs b/src-tauri/src/commands/ai.rs index b048b0469..11d8b6ab4 100644 --- a/src-tauri/src/commands/ai.rs +++ b/src-tauri/src/commands/ai.rs @@ -5,7 +5,7 @@ use super::connection::AppState; pub use dbx_core::ai::*; #[tauri::command] -pub async fn ai_test_connection(config: AiConfig) -> Result { +pub async fn ai_test_connection(config: AiConfig) -> Result { dbx_core::ai::test_connection_core(&config).await }