diff --git a/package.json b/package.json index 3fb23b77d..3ee919df9 100644 --- a/package.json +++ b/package.json @@ -1,7 +1,7 @@ { "name": "proxycast", "private": true, - "version": "0.3.0", + "version": "0.4.0", "type": "module", "scripts": { "dev": "vite", diff --git a/scripts/update-version.sh b/scripts/update-version.sh new file mode 100755 index 000000000..19ff3ce1d --- /dev/null +++ b/scripts/update-version.sh @@ -0,0 +1,39 @@ +#!/bin/bash +# 统一版本管理脚本 +# 用法: ./scripts/update-version.sh +# 示例: ./scripts/update-version.sh 0.5.0 + +set -e + +if [ -z "$1" ]; then + echo "用法: $0 " + echo "示例: $0 0.5.0" + exit 1 +fi + +NEW_VERSION="$1" +SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd)" +PROJECT_ROOT="$(dirname "$SCRIPT_DIR")" + +echo "更新版本到 $NEW_VERSION..." + +# 更新 package.json +sed -i '' "s/\"version\": \"[^\"]*\"/\"version\": \"$NEW_VERSION\"/" "$PROJECT_ROOT/package.json" +echo "✓ package.json" + +# 更新 Cargo.toml +sed -i '' "s/^version = \"[^\"]*\"/version = \"$NEW_VERSION\"/" "$PROJECT_ROOT/src-tauri/Cargo.toml" +echo "✓ src-tauri/Cargo.toml" + +# 更新 tauri.conf.json +sed -i '' "s/\"version\": \"[^\"]*\"/\"version\": \"$NEW_VERSION\"/" "$PROJECT_ROOT/src-tauri/tauri.conf.json" +echo "✓ src-tauri/tauri.conf.json" + +echo "" +echo "版本已更新到 $NEW_VERSION" +echo "" +echo "下一步:" +echo " git add -A" +echo " git commit -m 'chore: bump version to $NEW_VERSION'" +echo " git tag v$NEW_VERSION" +echo " git push origin main --tags" diff --git a/src-tauri/Cargo.lock b/src-tauri/Cargo.lock index 7acba54d2..fdf8011ab 100644 --- a/src-tauri/Cargo.lock +++ b/src-tauri/Cargo.lock @@ -2789,7 +2789,7 @@ dependencies = [ [[package]] name = "proxycast" -version = "0.3.0" +version = "0.4.0" dependencies = [ "async-stream", "axum", diff --git a/src-tauri/Cargo.toml b/src-tauri/Cargo.toml index 4eb716327..ad7915163 100644 --- a/src-tauri/Cargo.toml +++ b/src-tauri/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "proxycast" -version = "0.3.0" +version = "0.4.0" description = "AI API Proxy Desktop App" authors = ["you"] edition = "2021" diff --git a/src-tauri/src/converter/cw_to_openai.rs b/src-tauri/src/converter/cw_to_openai.rs index d26527c63..77c9ec0f8 100644 --- a/src-tauri/src/converter/cw_to_openai.rs +++ b/src-tauri/src/converter/cw_to_openai.rs @@ -14,7 +14,7 @@ pub fn convert_cw_event_to_openai_chunk( ) -> Option { let created = SystemTime::now() .duration_since(UNIX_EPOCH) - .unwrap() + .unwrap_or_default() .as_secs(); if let Some(resp_event) = &event.assistant_response_event { @@ -78,7 +78,7 @@ pub fn create_openai_response( ) -> ChatCompletionResponse { let created = SystemTime::now() .duration_since(UNIX_EPOCH) - .unwrap() + .unwrap_or_default() .as_secs(); let finish_reason = if tool_calls.is_some() { @@ -117,7 +117,7 @@ pub fn create_openai_response( pub fn create_stream_end_chunk(model: &str, response_id: &str) -> ChatCompletionChunk { let created = SystemTime::now() .duration_since(UNIX_EPOCH) - .unwrap() + .unwrap_or_default() .as_secs(); ChatCompletionChunk { diff --git a/src-tauri/src/converter/openai_to_cw.rs b/src-tauri/src/converter/openai_to_cw.rs index 0634a7938..dac31c77d 100644 --- a/src-tauri/src/converter/openai_to_cw.rs +++ b/src-tauri/src/converter/openai_to_cw.rs @@ -8,7 +8,9 @@ use uuid::Uuid; pub fn get_model_map() -> HashMap<&'static str, &'static str> { let mut map = HashMap::new(); map.insert("claude-opus-4-5", "claude-opus-4.5"); + map.insert("claude-opus-4-5-20251101", "claude-opus-4.5"); map.insert("claude-haiku-4-5", "claude-haiku-4.5"); + map.insert("claude-haiku-4-5-20251001", "claude-haiku-4.5"); map.insert("claude-sonnet-4-5", "CLAUDE_SONNET_4_5_20250929_V1_0"); map.insert( "claude-sonnet-4-5-20250929", @@ -32,6 +34,111 @@ pub fn get_model_map() -> HashMap<&'static str, &'static str> { pub const DEFAULT_MODEL: &str = "CLAUDE_SONNET_4_5_20250929_V1_0"; +/// 预处理消息:合并连续的 tool 消息到前一个 assistant 消息后的 user 消息 +fn preprocess_messages(messages: &[&ChatMessage]) -> Vec { + let mut result: Vec = Vec::new(); + let mut pending_tool_results: Vec = Vec::new(); + + for msg in messages { + match msg.role.as_str() { + "tool" => { + // 收集 tool 结果 + let content = msg.get_content_text(); + let tool_id = msg.tool_call_id.clone().unwrap_or_default(); + pending_tool_results.push(CWToolResult { + content: vec![CWTextContent { text: content }], + status: "success".to_string(), + tool_use_id: tool_id, + }); + } + "user" => { + // 如果有待处理的 tool results,合并到这个 user 消息 + let content = msg.get_content_text(); + let mut tool_results = pending_tool_results.clone(); + pending_tool_results.clear(); + + // 去重 tool_results + let mut seen_ids = std::collections::HashSet::new(); + tool_results.retain(|tr| seen_ids.insert(tr.tool_use_id.clone())); + + result.push(ProcessedMessage { + role: "user".to_string(), + content, + tool_calls: None, + tool_results: if tool_results.is_empty() { + None + } else { + Some(tool_results) + }, + }); + } + "assistant" => { + // 如果有待处理的 tool results,先创建一个 user 消息 + if !pending_tool_results.is_empty() { + let mut tool_results = pending_tool_results.clone(); + pending_tool_results.clear(); + + // 去重 tool_results + let mut seen_ids = std::collections::HashSet::new(); + tool_results.retain(|tr| seen_ids.insert(tr.tool_use_id.clone())); + + result.push(ProcessedMessage { + role: "user".to_string(), + content: "Tool results provided.".to_string(), + tool_calls: None, + tool_results: Some(tool_results), + }); + } + + let content = msg.get_content_text(); + let tool_calls = msg.tool_calls.as_ref().map(|calls| { + calls + .iter() + .map(|tc| CWToolUse { + input: serde_json::from_str(&tc.function.arguments) + .unwrap_or(serde_json::json!({})), + name: tc.function.name.clone(), + tool_use_id: tc.id.clone(), + }) + .collect() + }); + + result.push(ProcessedMessage { + role: "assistant".to_string(), + content, + tool_calls, + tool_results: None, + }); + } + _ => {} + } + } + + // 处理末尾的 tool results + if !pending_tool_results.is_empty() { + let mut tool_results = pending_tool_results; + let mut seen_ids = std::collections::HashSet::new(); + tool_results.retain(|tr| seen_ids.insert(tr.tool_use_id.clone())); + + result.push(ProcessedMessage { + role: "user".to_string(), + content: "Tool results provided.".to_string(), + tool_calls: None, + tool_results: Some(tool_results), + }); + } + + result +} + +#[derive(Debug, Clone)] +struct ProcessedMessage { + role: String, + content: String, + tool_calls: Option>, + tool_results: Option>, +} + /// 将 OpenAI ChatCompletionRequest 转换为 CodeWhisperer 请求 pub fn convert_openai_to_codewhisperer( request: &ChatCompletionRequest, @@ -47,32 +154,46 @@ pub fn convert_openai_to_codewhisperer( // 提取 system prompt 和消息 let mut system_prompt = String::new(); - let mut messages: Vec<&ChatMessage> = Vec::new(); + let mut raw_messages: Vec<&ChatMessage> = Vec::new(); for msg in &request.messages { if msg.role == "system" { system_prompt = msg.get_content_text(); } else { - messages.push(msg); + raw_messages.push(msg); } } + // 预处理消息:合并 tool 消息 + let messages = preprocess_messages(&raw_messages); + // 构建历史记录 let mut history: Vec = Vec::new(); let mut start_idx = 0; // 处理 system prompt - 合并到第一条用户消息 if !system_prompt.is_empty() && !messages.is_empty() && messages[0].role == "user" { - let first_content = messages[0].get_content_text(); + let first_content = &messages[0].content; let combined = format!("{system_prompt}\n\n{first_content}"); + + let mut user_msg = UserInputMessage { + content: combined, + model_id: cw_model.clone(), + origin: "AI_EDITOR".to_string(), + images: None, + user_input_message_context: None, + }; + + // 如果第一条消息有 tool_results,也要包含 + if let Some(ref tool_results) = messages[0].tool_results { + user_msg.user_input_message_context = Some(UserInputMessageContext { + tools: None, + tool_results: Some(tool_results.clone()), + }); + } + history.push(HistoryItem::User(UserHistoryItem { - user_input_message: UserInputMessage { - content: combined, - model_id: cw_model.clone(), - origin: "AI_EDITOR".to_string(), - images: None, - user_input_message_context: None, - }, + user_input_message: user_msg, })); start_idx = 1; } @@ -85,29 +206,28 @@ pub fn convert_openai_to_codewhisperer( { match msg.role.as_str() { "user" => { - let content = msg.get_content_text(); - let tool_results = extract_tool_results(msg); + let content = if msg.content.is_empty() { + if msg.tool_results.is_some() { + "Tool results provided.".to_string() + } else { + "Continue".to_string() + } + } else { + msg.content.clone() + }; let mut user_msg = UserInputMessage { - content: if content.is_empty() { - if tool_results.is_some() { - "Tool results provided.".to_string() - } else { - "Continue".to_string() - } - } else { - content - }, + content, model_id: cw_model.clone(), origin: "AI_EDITOR".to_string(), images: None, user_input_message_context: None, }; - if tool_results.is_some() { + if let Some(ref tool_results) = msg.tool_results { user_msg.user_input_message_context = Some(UserInputMessageContext { tools: None, - tool_results, + tool_results: Some(tool_results.clone()), }); } @@ -116,41 +236,16 @@ pub fn convert_openai_to_codewhisperer( })); } "assistant" => { - let content = msg.get_content_text(); - let tool_uses = extract_tool_uses(msg); + let content = if msg.content.is_empty() { + "I understand.".to_string() + } else { + msg.content.clone() + }; history.push(HistoryItem::Assistant(AssistantHistoryItem { assistant_response_message: AssistantResponseMessage { - content: if content.is_empty() { - "I understand.".to_string() - } else { - content - }, - tool_uses, - }, - })); - } - "tool" => { - let tool_content = msg.get_content_text(); - let tool_id = msg.tool_call_id.clone().unwrap_or_default(); - - history.push(HistoryItem::User(UserHistoryItem { - user_input_message: UserInputMessage { - content: format!( - "Tool result: {}", - &tool_content[..tool_content.len().min(200)] - ), - model_id: cw_model.clone(), - origin: "AI_EDITOR".to_string(), - images: None, - user_input_message_context: Some(UserInputMessageContext { - tools: None, - tool_results: Some(vec![CWToolResult { - content: vec![CWTextContent { text: tool_content }], - status: "success".to_string(), - tool_use_id: tool_id, - }]), - }), + content, + tool_uses: msg.tool_calls.clone(), }, })); } @@ -162,20 +257,23 @@ pub fn convert_openai_to_codewhisperer( let history = fix_history_alternation(history, &cw_model); // 构建当前消息 - let current_content = if messages.is_empty() { - "Continue".to_string() - } else { - let last_msg = messages.last().unwrap(); + let (current_content, current_tool_results) = if let Some(last_msg) = messages.last() { if last_msg.role == "assistant" { - "Continue".to_string() + ("Continue".to_string(), None) } else { - let content = last_msg.get_content_text(); - if content.is_empty() { - "Continue".to_string() + let content = if last_msg.content.is_empty() { + if last_msg.tool_results.is_some() { + "Tool results provided.".to_string() + } else { + "Continue".to_string() + } } else { - content - } + last_msg.content.clone() + }; + (content, last_msg.tool_results.clone()) } + } else { + ("Continue".to_string(), None) }; // 构建 tools @@ -211,13 +309,6 @@ pub fn convert_openai_to_codewhisperer( .collect() }); - let current_tool_results = if !messages.is_empty() { - let last_msg = messages.last().unwrap(); - extract_tool_results(last_msg) - } else { - None - }; - let user_input_message_context = if tools.is_some() || current_tool_results.is_some() { Some(UserInputMessageContext { tools, @@ -250,36 +341,6 @@ pub fn convert_openai_to_codewhisperer( } } -fn extract_tool_results(msg: &ChatMessage) -> Option> { - if msg.role == "tool" { - let content = msg.get_content_text(); - let tool_id = msg.tool_call_id.clone().unwrap_or_default(); - return Some(vec![CWToolResult { - content: vec![CWTextContent { text: content }], - status: "success".to_string(), - tool_use_id: tool_id, - }]); - } - None -} - -fn extract_tool_uses(msg: &ChatMessage) -> Option> { - msg.tool_calls.as_ref().map(|calls| { - calls - .iter() - .map(|tc| { - let input: serde_json::Value = - serde_json::from_str(&tc.function.arguments).unwrap_or(serde_json::json!({})); - CWToolUse { - input, - name: tc.function.name.clone(), - tool_use_id: tc.id.clone(), - } - }) - .collect() - }) -} - /// 修复历史记录,确保 user/assistant 严格交替 fn fix_history_alternation(history: Vec, model_id: &str) -> Vec { if history.is_empty() { @@ -290,15 +351,51 @@ fn fix_history_alternation(history: Vec, model_id: &str) -> Vec { - // 如果上一条也是 user,插入占位 assistant - if let Some(HistoryItem::User(_)) = fixed.last() { - fixed.push(HistoryItem::Assistant(AssistantHistoryItem { - assistant_response_message: AssistantResponseMessage { - content: "I understand.".to_string(), - tool_uses: None, - }, - })); + HistoryItem::User(user_item) => { + // 如果上一条也是 user,合并 tool_results 或插入占位 assistant + if let Some(HistoryItem::User(last_user)) = fixed.last_mut() { + // 尝试合并 tool_results + let has_tool_results = user_item + .user_input_message + .user_input_message_context + .as_ref() + .map(|ctx| ctx.tool_results.is_some()) + .unwrap_or(false); + + if has_tool_results { + // 合并 tool_results 到上一个 user 消息 + let new_results = user_item + .user_input_message + .user_input_message_context + .as_ref() + .and_then(|ctx| ctx.tool_results.clone()) + .unwrap_or_default(); + + if let Some(ref mut ctx) = + last_user.user_input_message.user_input_message_context + { + if let Some(ref mut existing) = ctx.tool_results { + existing.extend(new_results); + } else { + ctx.tool_results = Some(new_results); + } + } else { + last_user.user_input_message.user_input_message_context = + Some(UserInputMessageContext { + tools: None, + tool_results: Some(new_results), + }); + } + continue; + } else { + // 插入占位 assistant + fixed.push(HistoryItem::Assistant(AssistantHistoryItem { + assistant_response_message: AssistantResponseMessage { + content: "I understand.".to_string(), + tool_uses: None, + }, + })); + } } fixed.push(item); } diff --git a/src-tauri/src/lib.rs b/src-tauri/src/lib.rs index 0d2d8ceca..bfecea0fb 100644 --- a/src-tauri/src/lib.rs +++ b/src-tauri/src/lib.rs @@ -228,21 +228,24 @@ async fn get_env_variables(state: tauri::State<'_, AppState>) -> Result String { - if token.len() <= 12 { + let chars: Vec = token.chars().collect(); + if chars.len() <= 12 { "****".to_string() } else { - format!("{}****{}", &token[..6], &token[token.len() - 4..]) + let prefix: String = chars[..6].iter().collect(); + let suffix: String = chars[chars.len() - 4..].iter().collect(); + format!("{}****{}", prefix, suffix) } } #[tauri::command] async fn get_token_file_hash() -> Result { let path = providers::kiro::KiroProvider::default_creds_path(); - if !path.exists() { + if !tokio::fs::try_exists(&path).await.unwrap_or(false) { return Ok("".to_string()); } - let content = std::fs::read(&path).map_err(|e| e.to_string())?; + let content = tokio::fs::read(&path).await.map_err(|e| e.to_string())?; let hash = format!("{:x}", md5::compute(&content)); Ok(hash) } @@ -256,7 +259,7 @@ async fn check_and_reload_credentials( ) -> Result { let path = providers::kiro::KiroProvider::default_creds_path(); - if !path.exists() { + if !tokio::fs::try_exists(&path).await.unwrap_or(false) { return Ok(CheckResult { changed: false, new_hash: "".to_string(), @@ -264,7 +267,7 @@ async fn check_and_reload_credentials( }); } - let content = std::fs::read(&path).map_err(|e| e.to_string())?; + let content = tokio::fs::read(&path).await.map_err(|e| e.to_string())?; let new_hash = format!("{:x}", md5::compute(&content)); if !last_hash.is_empty() && new_hash != last_hash { @@ -414,11 +417,11 @@ async fn get_gemini_env_variables( #[tauri::command] async fn get_gemini_token_file_hash() -> Result { let path = providers::gemini::GeminiProvider::default_creds_path(); - if !path.exists() { + if !tokio::fs::try_exists(&path).await.unwrap_or(false) { return Ok("".to_string()); } - let content = std::fs::read(&path).map_err(|e| e.to_string())?; + let content = tokio::fs::read(&path).await.map_err(|e| e.to_string())?; let hash = format!("{:x}", md5::compute(&content)); Ok(hash) } @@ -431,7 +434,7 @@ async fn check_and_reload_gemini_credentials( ) -> Result { let path = providers::gemini::GeminiProvider::default_creds_path(); - if !path.exists() { + if !tokio::fs::try_exists(&path).await.unwrap_or(false) { return Ok(CheckResult { changed: false, new_hash: "".to_string(), @@ -439,7 +442,7 @@ async fn check_and_reload_gemini_credentials( }); } - let content = std::fs::read(&path).map_err(|e| e.to_string())?; + let content = tokio::fs::read(&path).await.map_err(|e| e.to_string())?; let new_hash = format!("{:x}", md5::compute(&content)); if !last_hash.is_empty() && new_hash != last_hash { @@ -590,11 +593,11 @@ async fn get_qwen_env_variables( #[tauri::command] async fn get_qwen_token_file_hash() -> Result { let path = providers::qwen::QwenProvider::default_creds_path(); - if !path.exists() { + if !tokio::fs::try_exists(&path).await.unwrap_or(false) { return Ok("".to_string()); } - let content = std::fs::read(&path).map_err(|e| e.to_string())?; + let content = tokio::fs::read(&path).await.map_err(|e| e.to_string())?; let hash = format!("{:x}", md5::compute(&content)); Ok(hash) } diff --git a/src-tauri/src/logger.rs b/src-tauri/src/logger.rs index 747175d21..7203120b2 100644 --- a/src-tauri/src/logger.rs +++ b/src-tauri/src/logger.rs @@ -1,6 +1,9 @@ //! 日志管理模块 -use chrono::Utc; +use chrono::{Local, Utc}; use serde::{Deserialize, Serialize}; +use std::fs::{self, OpenOptions}; +use std::io::Write; +use std::path::PathBuf; use std::sync::Arc; use tokio::sync::RwLock; @@ -14,13 +17,26 @@ pub struct LogEntry { pub struct LogStore { logs: Vec, max_logs: usize, + log_file_path: Option, } impl Default for LogStore { fn default() -> Self { + // 默认日志文件路径: ~/.proxycast/logs/proxycast.log + let log_dir = dirs::home_dir() + .unwrap_or_else(|| PathBuf::from(".")) + .join(".proxycast") + .join("logs"); + + // 创建日志目录 + let _ = fs::create_dir_all(&log_dir); + + let log_file = log_dir.join("proxycast.log"); + Self { logs: Vec::new(), max_logs: 1000, + log_file_path: Some(log_file), } } } @@ -31,13 +47,24 @@ impl LogStore { } pub fn add(&mut self, level: &str, message: &str) { + let now = Utc::now(); let entry = LogEntry { - timestamp: Utc::now().to_rfc3339(), + timestamp: now.to_rfc3339(), level: level.to_string(), message: message.to_string(), }; - self.logs.push(entry); + self.logs.push(entry.clone()); + + // 写入日志文件 + if let Some(ref path) = self.log_file_path { + let local_time = Local::now().format("%Y-%m-%d %H:%M:%S%.3f"); + let log_line = format!("{} [{}] {}\n", local_time, level.to_uppercase(), message); + + if let Ok(mut file) = OpenOptions::new().create(true).append(true).open(path) { + let _ = file.write_all(log_line.as_bytes()); + } + } // 保持日志数量在限制内 if self.logs.len() > self.max_logs { @@ -45,6 +72,23 @@ impl LogStore { } } + /// 记录原始响应到单独的文件(用于调试) + pub fn log_raw_response(&self, request_id: &str, body: &str) { + if let Some(ref log_path) = self.log_file_path { + let log_dir = log_path.parent().unwrap_or(std::path::Path::new(".")); + let raw_file = log_dir.join(format!("raw_response_{request_id}.txt")); + + if let Ok(mut file) = OpenOptions::new() + .create(true) + .truncate(true) + .write(true) + .open(&raw_file) + { + let _ = file.write_all(body.as_bytes()); + } + } + } + pub fn get_logs(&self) -> Vec { self.logs.clone() } @@ -52,6 +96,12 @@ impl LogStore { pub fn clear(&mut self) { self.logs.clear(); } + + pub fn get_log_file_path(&self) -> Option { + self.log_file_path + .as_ref() + .map(|p| p.to_string_lossy().to_string()) + } } #[allow(dead_code)] diff --git a/src-tauri/src/providers/gemini.rs b/src-tauri/src/providers/gemini.rs index 3918ff287..0a0d319d5 100644 --- a/src-tauri/src/providers/gemini.rs +++ b/src-tauri/src/providers/gemini.rs @@ -145,11 +145,11 @@ impl GeminiProvider { .join(CREDENTIALS_FILE) } - pub fn load_credentials(&mut self) -> Result<(), Box> { + pub async fn load_credentials(&mut self) -> Result<(), Box> { let path = Self::default_creds_path(); - if path.exists() { - let content = std::fs::read_to_string(&path)?; + if tokio::fs::try_exists(&path).await.unwrap_or(false) { + let content = tokio::fs::read_to_string(&path).await?; let creds: GeminiCredentials = serde_json::from_str(&content)?; self.credentials = creds; } @@ -157,13 +157,13 @@ impl GeminiProvider { Ok(()) } - pub fn save_credentials(&self) -> Result<(), Box> { + pub async fn save_credentials(&self) -> Result<(), Box> { let path = Self::default_creds_path(); if let Some(parent) = path.parent() { - std::fs::create_dir_all(parent)?; + tokio::fs::create_dir_all(parent).await?; } let content = serde_json::to_string_pretty(&self.credentials)?; - std::fs::write(&path, content)?; + tokio::fs::write(&path, content).await?; Ok(()) } diff --git a/src-tauri/src/providers/kiro.rs b/src-tauri/src/providers/kiro.rs index c8386f399..ccc9f42f2 100644 --- a/src-tauri/src/providers/kiro.rs +++ b/src-tauri/src/providers/kiro.rs @@ -62,27 +62,27 @@ impl KiroProvider { .join("kiro-auth-token.json") } - pub fn load_credentials(&mut self) -> Result<(), Box> { + pub async fn load_credentials(&mut self) -> Result<(), Box> { let path = Self::default_creds_path(); - let dir = path.parent().unwrap(); + let dir = path.parent().ok_or("Invalid path: no parent directory")?; let mut merged = KiroCredentials::default(); // 读取主凭证文件 - if path.exists() { - let content = std::fs::read_to_string(&path)?; + if tokio::fs::try_exists(&path).await.unwrap_or(false) { + let content = tokio::fs::read_to_string(&path).await?; let creds: KiroCredentials = serde_json::from_str(&content)?; merge_credentials(&mut merged, &creds); } // 读取目录中其他 JSON 文件 - if dir.is_dir() { - for entry in std::fs::read_dir(dir)? { - let entry = entry?; + if tokio::fs::try_exists(dir).await.unwrap_or(false) { + let mut entries = tokio::fs::read_dir(dir).await?; + while let Some(entry) = entries.next_entry().await? { let file_path = entry.path(); if file_path.extension().map(|e| e == "json").unwrap_or(false) && file_path != path { - if let Ok(content) = std::fs::read_to_string(&file_path) { + if let Ok(content) = tokio::fs::read_to_string(&file_path).await { if let Ok(creds) = serde_json::from_str::(&content) { merge_credentials(&mut merged, &creds); } @@ -177,12 +177,12 @@ impl KiroProvider { Ok(new_token.to_string()) } - pub fn save_credentials(&self) -> Result<(), Box> { + pub async fn save_credentials(&self) -> Result<(), Box> { let path = Self::default_creds_path(); // 读取现有文件内容 - let mut existing: serde_json::Value = if path.exists() { - let content = std::fs::read_to_string(&path)?; + let mut existing: serde_json::Value = if tokio::fs::try_exists(&path).await.unwrap_or(false) { + let content = tokio::fs::read_to_string(&path).await?; serde_json::from_str(&content).unwrap_or(serde_json::json!({})) } else { serde_json::json!({}) @@ -201,7 +201,7 @@ impl KiroProvider { // 写回文件 let content = serde_json::to_string_pretty(&existing)?; - std::fs::write(&path, content)?; + tokio::fs::write(&path, content).await?; Ok(()) } @@ -225,6 +225,45 @@ impl KiroProvider { let cw_request = convert_openai_to_codewhisperer(request, profile_arn); let url = self.get_base_url(); + // Debug: 记录转换后的请求 + if let Ok(json_str) = serde_json::to_string_pretty(&cw_request) { + // 保存到文件用于调试 + let uuid_prefix = uuid::Uuid::new_v4() + .to_string() + .split('-') + .next() + .unwrap_or("unknown") + .to_string(); + let debug_path = dirs::home_dir() + .unwrap_or_default() + .join(".proxycast") + .join("logs") + .join(format!("cw_request_{}.json", uuid_prefix)); + let _ = tokio::fs::write(&debug_path, &json_str).await; + tracing::debug!("[CW_REQ] Request saved to {:?}", debug_path); + + // 记录历史消息数量和 tool_results 情况 + let history_len = cw_request + .conversation_state + .history + .as_ref() + .map(|h| h.len()) + .unwrap_or(0); + let current_has_tools = cw_request + .conversation_state + .current_message + .user_input_message + .user_input_message_context + .as_ref() + .map(|ctx| ctx.tool_results.as_ref().map(|tr| tr.len()).unwrap_or(0)) + .unwrap_or(0); + tracing::info!( + "[CW_REQ] history={} current_tool_results={}", + history_len, + current_has_tools + ); + } + let resp = self .client .post(&url) diff --git a/src-tauri/src/providers/qwen.rs b/src-tauri/src/providers/qwen.rs index 008f5f36e..40d530c75 100644 --- a/src-tauri/src/providers/qwen.rs +++ b/src-tauri/src/providers/qwen.rs @@ -58,11 +58,11 @@ impl QwenProvider { .join(CREDENTIALS_FILE) } - pub fn load_credentials(&mut self) -> Result<(), Box> { + pub async fn load_credentials(&mut self) -> Result<(), Box> { let path = Self::default_creds_path(); - if path.exists() { - let content = std::fs::read_to_string(&path)?; + if tokio::fs::try_exists(&path).await.unwrap_or(false) { + let content = tokio::fs::read_to_string(&path).await?; let creds: QwenCredentials = serde_json::from_str(&content)?; self.credentials = creds; } @@ -70,13 +70,13 @@ impl QwenProvider { Ok(()) } - pub fn save_credentials(&self) -> Result<(), Box> { + pub async fn save_credentials(&self) -> Result<(), Box> { let path = Self::default_creds_path(); if let Some(parent) = path.parent() { - std::fs::create_dir_all(parent)?; + tokio::fs::create_dir_all(parent).await?; } let content = serde_json::to_string_pretty(&self.credentials)?; - std::fs::write(&path, content)?; + tokio::fs::write(&path, content).await?; Ok(()) } @@ -119,9 +119,7 @@ impl QwenProvider { .ok_or("No refresh token available")?; let client_id = std::env::var("QWEN_OAUTH_CLIENT_ID") - .ok() - .or_else(|| Some("f0304373b74a44d2b584a3fb70ca9e56".to_string())) - .unwrap(); + .unwrap_or_else(|_| "f0304373b74a44d2b584a3fb70ca9e56".to_string()); let body = serde_json::json!({ "grant_type": "refresh_token", diff --git a/src-tauri/src/server.rs b/src-tauri/src/server.rs index c01d1bd4f..02a83cb95 100644 --- a/src-tauri/src/server.rs +++ b/src-tauri/src/server.rs @@ -22,6 +22,16 @@ use serde::{Deserialize, Serialize}; use std::sync::Arc; use tokio::sync::{oneshot, RwLock}; +/// 安全截断字符串到指定字符数,避免 UTF-8 边界问题 +fn safe_truncate(s: &str, max_chars: usize) -> String { + let chars: Vec = s.chars().collect(); + if chars.len() <= max_chars { + s.to_string() + } else { + chars[..max_chars].iter().collect() + } +} + #[derive(Debug, Clone, Serialize, Deserialize)] pub struct ServerStatus { pub running: bool, @@ -135,6 +145,9 @@ struct AppState { api_key: String, kiro: Arc>, logs: Arc>, + kiro_refresh_lock: Arc>, + gemini_refresh_lock: Arc>, + qwen_refresh_lock: Arc>, } async fn run_server( @@ -149,6 +162,9 @@ async fn run_server( api_key: api_key.to_string(), kiro: Arc::new(RwLock::new(kiro)), logs, + kiro_refresh_lock: Arc::new(tokio::sync::Mutex::new(())), + gemini_refresh_lock: Arc::new(tokio::sync::Mutex::new(())), + qwen_refresh_lock: Arc::new(tokio::sync::Mutex::new(())), }; let app = Router::new() @@ -256,6 +272,7 @@ async fn chat_completions( // 检查是否需要刷新 token { + let _guard = state.kiro_refresh_lock.lock().await; let mut kiro = state.kiro.write().await; if kiro.credentials.access_token.is_none() { if let Err(e) = kiro.refresh_token().await { @@ -320,7 +337,7 @@ async fn chat_completions( "object": "chat.completion", "created": std::time::SystemTime::now() .duration_since(std::time::UNIX_EPOCH) - .unwrap() + .unwrap_or_default() .as_secs(), "model": request.model, "choices": [{ @@ -345,6 +362,7 @@ async fn chat_completions( } else if status.as_u16() == 403 { // Token 过期,尝试刷新 drop(kiro); + let _guard = state.kiro_refresh_lock.lock().await; let mut kiro = state.kiro.write().await; state .logs @@ -392,7 +410,7 @@ async fn chat_completions( "object": "chat.completion", "created": std::time::SystemTime::now() .duration_since(std::time::UNIX_EPOCH) - .unwrap() + .unwrap_or_default() .as_secs(), "model": request.model, "choices": [{ @@ -435,11 +453,7 @@ async fn chat_completions( let body = resp.text().await.unwrap_or_default(); state.logs.write().await.add( "error", - &format!( - "Upstream error {}: {}", - status, - &body[..body.len().min(200)] - ), + &format!("Upstream error {}: {}", status, safe_truncate(&body, 200)), ); ( StatusCode::from_u16(status.as_u16()).unwrap_or(StatusCode::INTERNAL_SERVER_ERROR), @@ -516,6 +530,7 @@ async fn anthropic_messages( // 检查是否需要刷新 token { + let _guard = state.kiro_refresh_lock.lock().await; let mut kiro = state.kiro.write().await; if kiro.credentials.access_token.is_none() { state @@ -571,12 +586,29 @@ async fn anthropic_messages( if status.is_success() { match resp.text().await { Ok(body) => { - // 记录原始响应长度和预览 + // 记录原始响应长度 state.logs.write().await.add( "debug", &format!("[RESP] Raw body length: {} bytes", body.len()), ); + // 保存原始响应到文件用于调试 + let request_id = uuid::Uuid::new_v4().to_string()[..8].to_string(); + state.logs.read().await.log_raw_response(&request_id, &body); + state.logs.write().await.add( + "debug", + &format!("[RESP] Raw response saved to raw_response_{request_id}.txt"), + ); + + // 记录响应的前500字符用于调试 + state.logs.write().await.add( + "debug", + &format!( + "[RESP] Body preview: {}", + &body.chars().take(500).collect::() + ), + ); + let parsed = parse_cw_response(&body); // 详细记录解析结果 @@ -633,6 +665,7 @@ async fn anthropic_messages( } else if status.as_u16() == 403 { // Token 过期,尝试刷新 drop(kiro); + let _guard = state.kiro_refresh_lock.lock().await; let mut kiro = state.kiro.write().await; state.logs.write().await.add( "warn", @@ -697,7 +730,7 @@ async fn anthropic_messages( "error", &format!( "[RETRY] Failed with status {retry_status}: {}", - &body[..body.len().min(500)] + safe_truncate(&body, 500) ), ); ( @@ -740,7 +773,7 @@ async fn anthropic_messages( &format!( "[ERROR] Upstream error HTTP {}: {}", status, - &body[..body.len().min(500)] + safe_truncate(&body, 500) ), ); ( @@ -847,18 +880,18 @@ fn build_anthropic_stream_response( let mut block_index = 0; - // 2. 文本内容块 - if !content.is_empty() { - // content_block_start - let block_start = serde_json::json!({ - "type": "content_block_start", - "index": block_index, - "content_block": {"type": "text", "text": ""} - }); - events.push(format!( - "event: content_block_start\ndata: {block_start}\n\n" - )); + // 2. 文本内容块 - 即使为空也要发送,Claude Code 需要至少一个 content block + // content_block_start + let block_start = serde_json::json!({ + "type": "content_block_start", + "index": block_index, + "content_block": {"type": "text", "text": ""} + }); + events.push(format!( + "event: content_block_start\ndata: {block_start}\n\n" + )); + if !content.is_empty() { // content_block_delta - 发送完整内容 let block_delta = serde_json::json!({ "type": "content_block_delta", @@ -868,17 +901,17 @@ fn build_anthropic_stream_response( events.push(format!( "event: content_block_delta\ndata: {block_delta}\n\n" )); - - // content_block_stop - let block_stop = serde_json::json!({ - "type": "content_block_stop", - "index": block_index - }); - events.push(format!("event: content_block_stop\ndata: {block_stop}\n\n")); - - block_index += 1; } + // content_block_stop + let block_stop = serde_json::json!({ + "type": "content_block_stop", + "index": block_index + }); + events.push(format!("event: content_block_stop\ndata: {block_stop}\n\n")); + + block_index += 1; + // 3. Tool use 块 for tc in &tool_calls { let input: serde_json::Value = @@ -947,7 +980,13 @@ fn build_anthropic_stream_response( .header(header::CACHE_CONTROL, "no-cache") .header(header::CONNECTION, "keep-alive") .body(body) - .unwrap() + .unwrap_or_else(|e| { + tracing::error!("Failed to build SSE response: {}", e); + Response::builder() + .status(StatusCode::INTERNAL_SERVER_ERROR) + .body(Body::empty()) + .unwrap_or_default() + }) } async fn count_tokens( @@ -974,109 +1013,123 @@ struct CWParsedResponse { } /// 解析 CodeWhisperer AWS Event Stream 响应 +/// AWS Event Stream 是二进制格式,JSON payload 嵌入在二进制头部之间 fn parse_cw_response(body: &str) -> CWParsedResponse { let mut result = CWParsedResponse::default(); - let mut current_tool: Option<(String, String, String)> = None; // (id, name, input) + // 使用 HashMap 来跟踪多个并发的 tool calls + // key: toolUseId, value: (name, input_accumulated) + let mut tool_map: std::collections::HashMap = + std::collections::HashMap::new(); - // 解析所有 JSON 事件 - let patterns = [ - r#"\{"content":"#, - r#"\{"name":"#, - r#"\{"input":"#, - r#"\{"stop":"#, + // 将字符串转换为字节,因为 AWS Event Stream 包含二进制数据 + let bytes = body.as_bytes(); + + // 搜索所有 JSON 对象的模式 + // AWS Event Stream 格式: [binary headers]{"content":"..."}[binary trailer] + let json_patterns: &[&[u8]] = &[ + b"{\"content\":", + b"{\"name\":", + b"{\"input\":", + b"{\"stop\":", + b"{\"followupPrompt\":", + b"{\"toolUseId\":", ]; let mut pos = 0; - while pos < body.len() { + while pos < bytes.len() { // 找到下一个 JSON 对象的开始 - let mut next_start = body.len(); - for pattern in &patterns { - if let Some(idx) = body[pos..].find(pattern) { - next_start = next_start.min(pos + idx); + let mut next_start: Option = None; + + for pattern in json_patterns { + if let Some(idx) = find_subsequence(&bytes[pos..], pattern) { + let abs_pos = pos + idx; + if next_start.map_or(true, |start| abs_pos < start) { + next_start = Some(abs_pos); + } } } - if next_start >= body.len() { - break; - } + let start = match next_start { + Some(s) => s, + None => break, + }; - // 找到匹配的 } - if let Some(json_str) = extract_json_object(&body[next_start..]) { - if let Ok(value) = serde_json::from_str::(json_str) { + // 从 start 位置提取完整的 JSON 对象 + if let Some(json_str) = extract_json_from_bytes(&bytes[start..]) { + if let Ok(value) = serde_json::from_str::(&json_str) { // 处理 content 事件 if let Some(content) = value.get("content").and_then(|v| v.as_str()) { // 跳过 followupPrompt if value.get("followupPrompt").is_none() { - let unescaped = content - .replace("\\n", "\n") - .replace("\\t", "\t") - .replace("\\\"", "\"") - .replace("\\\\", "\\"); - result.content.push_str(&unescaped); + result.content.push_str(content); } } - // 处理 tool use 开始事件 - else if let (Some(name), Some(tool_use_id)) = ( - value.get("name").and_then(|v| v.as_str()), - value.get("toolUseId").and_then(|v| v.as_str()), - ) { - let input = value + // 处理 tool use 事件 (包含 toolUseId) + else if let Some(tool_use_id) = value.get("toolUseId").and_then(|v| v.as_str()) { + let name = value + .get("name") + .and_then(|v| v.as_str()) + .unwrap_or("") + .to_string(); + let input_chunk = value .get("input") .and_then(|v| v.as_str()) .unwrap_or("") .to_string(); - current_tool = Some((tool_use_id.to_string(), name.to_string(), input)); + let is_stop = value.get("stop").and_then(|v| v.as_bool()).unwrap_or(false); - // 如果同时有 stop,直接完成 - if value.get("stop").and_then(|v| v.as_bool()).unwrap_or(false) { - if let Some((id, name, input)) = current_tool.take() { - result.tool_calls.push(ToolCall { - id, - call_type: "function".to_string(), - function: FunctionCall { - name, - arguments: input, - }, - }); + // 获取或创建 tool entry + let entry = tool_map + .entry(tool_use_id.to_string()) + .or_insert_with(|| (String::new(), String::new())); + + // 更新 name(如果有) + if !name.is_empty() { + entry.0 = name; + } + + // 累积 input + entry.1.push_str(&input_chunk); + + // 如果是 stop 事件,完成这个 tool call + if is_stop { + if let Some((name, input)) = tool_map.remove(tool_use_id) { + if !name.is_empty() { + result.tool_calls.push(ToolCall { + id: tool_use_id.to_string(), + call_type: "function".to_string(), + function: FunctionCall { + name, + arguments: input, + }, + }); + } } } } - // 处理 input 续传事件 - else if let Some(input) = value.get("input").and_then(|v| v.as_str()) { - if let Some((_, _, ref mut current_input)) = current_tool { - current_input.push_str(input); - } - } - // 处理 stop 事件 + // 处理独立的 stop 事件(没有 toolUseId) else if value.get("stop").and_then(|v| v.as_bool()).unwrap_or(false) { - if let Some((id, name, input)) = current_tool.take() { - result.tool_calls.push(ToolCall { - id, - call_type: "function".to_string(), - function: FunctionCall { - name, - arguments: input, - }, - }); - } + // 这种情况不应该发生,但以防万一 } } - pos = next_start + json_str.len(); + pos = start + json_str.len(); } else { - pos = next_start + 1; + pos = start + 1; } } - // 处理未完成的 tool call - if let Some((id, name, input)) = current_tool { - result.tool_calls.push(ToolCall { - id, - call_type: "function".to_string(), - function: FunctionCall { - name, - arguments: input, - }, - }); + // 处理未完成的 tool calls(没有收到 stop 事件的) + for (id, (name, input)) in tool_map { + if !name.is_empty() { + result.tool_calls.push(ToolCall { + id, + call_type: "function".to_string(), + function: FunctionCall { + name, + arguments: input, + }, + }); + } } // 解析 bracket 格式的 tool calls: [Called xxx with args: {...}] @@ -1085,7 +1138,50 @@ fn parse_cw_response(body: &str) -> CWParsedResponse { result } -/// 从字符串中提取完整的 JSON 对象 +/// 在字节数组中查找子序列 +fn find_subsequence(haystack: &[u8], needle: &[u8]) -> Option { + haystack + .windows(needle.len()) + .position(|window| window == needle) +} + +/// 从字节数组中提取 JSON 对象字符串 +fn extract_json_from_bytes(bytes: &[u8]) -> Option { + if bytes.is_empty() || bytes[0] != b'{' { + return None; + } + + let mut brace_count = 0; + let mut in_string = false; + let mut escape_next = false; + let mut end_pos = None; + + for (i, &b) in bytes.iter().enumerate() { + if escape_next { + escape_next = false; + continue; + } + + match b { + b'\\' if in_string => escape_next = true, + b'"' => in_string = !in_string, + b'{' if !in_string => brace_count += 1, + b'}' if !in_string => { + brace_count -= 1; + if brace_count == 0 { + end_pos = Some(i + 1); + break; + } + } + _ => {} + } + } + + end_pos.and_then(|end| String::from_utf8(bytes[..end].to_vec()).ok()) +} + +/// 从字符串中提取完整的 JSON 对象 (保留用于兼容) +#[allow(dead_code)] fn extract_json_object(s: &str) -> Option<&str> { if !s.starts_with('{') { return None; @@ -1139,7 +1235,9 @@ fn parse_bracket_tool_calls(result: &mut CWParsedResponse) { arguments: args.as_str().to_string(), }, }); - to_remove.push(cap.get(0).unwrap().as_str().to_string()); + if let Some(full_match) = cap.get(0) { + to_remove.push(full_match.as_str().to_string()); + } } } // 从 content 中移除 tool call 文本 diff --git a/src-tauri/tauri.conf.json b/src-tauri/tauri.conf.json index 01273f79b..13ce5f449 100644 --- a/src-tauri/tauri.conf.json +++ b/src-tauri/tauri.conf.json @@ -1,7 +1,7 @@ { "$schema": "https://schema.tauri.app/config/2", "productName": "ProxyCast", - "version": "0.3.0", + "version": "0.4.0", "identifier": "com.proxycast.app", "build": { "beforeDevCommand": "npm run dev", diff --git a/src/components/Settings.tsx b/src/components/Settings.tsx index 929d8f7e8..41c1d8d12 100644 --- a/src/components/Settings.tsx +++ b/src/components/Settings.tsx @@ -285,7 +285,9 @@ export function Settings() { )} {r.model.includes("tool_call") ? ( - {r.model} + + {r.model} + ) : ( r.model )}