From ccb3667a0c2dece6a90fc0b564c460dd1eb45fff Mon Sep 17 00:00:00 2001 From: t8y2 <1156263951@qq.com> Date: Wed, 29 Apr 2026 21:32:35 +0800 Subject: [PATCH] feat: Enhance AI Assistant with new actions and context handling --- src-tauri/src/commands/ai.rs | 174 ++++++++++++++++++ src-tauri/src/commands/mod.rs | 1 + src-tauri/src/lib.rs | 3 + src/App.vue | 215 ++++++++++++++++++++-- src/components/editor/AiAssistant.vue | 254 ++++++++++++++++++++++---- src/components/editor/QueryEditor.vue | 88 ++++++++- src/components/grid/DataGrid.vue | 25 ++- src/composables/useToast.ts | 16 ++ src/i18n/locales/en.ts | 33 ++++ src/i18n/locales/zh-CN.ts | 33 ++++ src/lib/ai.ts | 232 ++++++++++++++++++----- src/lib/tauri.ts | 26 +++ src/stores/queryStore.ts | 19 ++ src/stores/settingsStore.ts | 32 +++- 14 files changed, 1045 insertions(+), 106 deletions(-) create mode 100644 src-tauri/src/commands/ai.rs create mode 100644 src/composables/useToast.ts diff --git a/src-tauri/src/commands/ai.rs b/src-tauri/src/commands/ai.rs new file mode 100644 index 000000000..e12e20e79 --- /dev/null +++ b/src-tauri/src/commands/ai.rs @@ -0,0 +1,174 @@ +use reqwest::header::{HeaderMap, HeaderValue, AUTHORIZATION, CONTENT_TYPE}; +use serde::{Deserialize, Serialize}; +use serde_json::json; +use tauri::{AppHandle, Manager}; + +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(rename_all = "lowercase")] +pub enum AiProvider { + Claude, + Openai, + Custom, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct AiConfig { + pub provider: AiProvider, + pub api_key: String, + pub endpoint: String, + pub model: String, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct AiMessage { + pub role: String, + pub content: String, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct AiCompletionRequest { + pub config: AiConfig, + pub system_prompt: String, + pub messages: Vec, + pub max_tokens: Option, + pub temperature: Option, +} + +fn ai_config_file(app: &AppHandle) -> Result { + let dir = app.path().app_data_dir().map_err(|e| e.to_string())?; + std::fs::create_dir_all(&dir).map_err(|e| e.to_string())?; + Ok(dir.join("ai_config.json")) +} + +#[tauri::command] +pub async fn save_ai_config(app: AppHandle, config: AiConfig) -> Result<(), String> { + let json = serde_json::to_string_pretty(&config).map_err(|e| e.to_string())?; + std::fs::write(ai_config_file(&app)?, json).map_err(|e| e.to_string()) +} + +#[tauri::command] +pub async fn load_ai_config(app: AppHandle) -> Result, String> { + let path = ai_config_file(&app)?; + if !path.exists() { + return Ok(None); + } + let json = std::fs::read_to_string(path).map_err(|e| e.to_string())?; + serde_json::from_str(&json).map(Some).map_err(|e| e.to_string()) +} + +#[tauri::command] +pub async fn ai_complete(request: AiCompletionRequest) -> Result { + if request.config.api_key.trim().is_empty() { + return Err("API key is required".to_string()); + } + if request.config.endpoint.trim().is_empty() { + return Err("Endpoint is required".to_string()); + } + if request.config.model.trim().is_empty() { + return Err("Model is required".to_string()); + } + + let client = reqwest::Client::builder() + .timeout(std::time::Duration::from_secs(60)) + .build() + .map_err(|e| e.to_string())?; + + match request.config.provider { + AiProvider::Claude => call_claude(&client, request).await, + AiProvider::Openai | AiProvider::Custom => call_openai_compatible(&client, request).await, + } +} + +async fn call_claude(client: &reqwest::Client, request: AiCompletionRequest) -> Result { + let mut headers = HeaderMap::new(); + headers.insert(CONTENT_TYPE, HeaderValue::from_static("application/json")); + headers.insert( + "x-api-key", + HeaderValue::from_str(&request.config.api_key).map_err(|e| e.to_string())?, + ); + headers.insert("anthropic-version", HeaderValue::from_static("2023-06-01")); + + let body = json!({ + "model": request.config.model, + "max_tokens": request.max_tokens.unwrap_or(2048), + "temperature": request.temperature.unwrap_or(0.2), + "system": request.system_prompt, + "messages": request.messages, + }); + + let res = client + .post(&request.config.endpoint) + .headers(headers) + .json(&body) + .send() + .await + .map_err(|e| format!("Claude request failed: {e}"))?; + + let status = res.status(); + let data: serde_json::Value = res.json().await.map_err(|e| e.to_string())?; + if !status.is_success() { + return Err(extract_error(&data).unwrap_or_else(|| format!("Claude API error: {status}"))); + } + + Ok(data["content"] + .as_array() + .and_then(|items| items.iter().find_map(|item| item["text"].as_str())) + .unwrap_or_default() + .to_string()) +} + +async fn call_openai_compatible( + client: &reqwest::Client, + request: AiCompletionRequest, +) -> Result { + let mut headers = HeaderMap::new(); + headers.insert(CONTENT_TYPE, HeaderValue::from_static("application/json")); + headers.insert( + AUTHORIZATION, + HeaderValue::from_str(&format!("Bearer {}", request.config.api_key)).map_err(|e| e.to_string())?, + ); + + let mut messages = vec![json!({ "role": "system", "content": request.system_prompt })]; + messages.extend( + request + .messages + .iter() + .map(|message| json!({ "role": message.role, "content": message.content })), + ); + + let body = json!({ + "model": request.config.model, + "messages": messages, + "max_tokens": request.max_tokens.unwrap_or(2048), + "temperature": request.temperature.unwrap_or(0.2), + }); + + let res = client + .post(&request.config.endpoint) + .headers(headers) + .json(&body) + .send() + .await + .map_err(|e| format!("AI request failed: {e}"))?; + + let status = res.status(); + let data: serde_json::Value = res.json().await.map_err(|e| e.to_string())?; + if !status.is_success() { + return Err(extract_error(&data).unwrap_or_else(|| format!("API error: {status}"))); + } + + Ok(data["choices"][0]["message"]["content"] + .as_str() + .unwrap_or_default() + .to_string()) +} + +fn extract_error(data: &serde_json::Value) -> Option { + data["error"]["message"] + .as_str() + .or_else(|| data["error"].as_str()) + .map(ToString::to_string) +} diff --git a/src-tauri/src/commands/mod.rs b/src-tauri/src/commands/mod.rs index a38a9b21f..82aa9206d 100644 --- a/src-tauri/src/commands/mod.rs +++ b/src-tauri/src/commands/mod.rs @@ -1,3 +1,4 @@ +pub mod ai; pub mod connection; pub mod history; pub mod mongo_cmd; diff --git a/src-tauri/src/lib.rs b/src-tauri/src/lib.rs index 3435fa6f0..092454c7e 100644 --- a/src-tauri/src/lib.rs +++ b/src-tauri/src/lib.rs @@ -31,6 +31,9 @@ pub fn run() { } }) .invoke_handler(tauri::generate_handler![ + commands::ai::ai_complete, + commands::ai::save_ai_config, + commands::ai::load_ai_config, commands::connection::test_connection, commands::connection::connect_db, commands::connection::disconnect_db, diff --git a/src/App.vue b/src/App.vue index e5d96ebc3..49b9c43ac 100644 --- a/src/App.vue +++ b/src/App.vue @@ -1,11 +1,17 @@