From 9682d405b6fe34397bdc090eeb4994988ef73162 Mon Sep 17 00:00:00 2001 From: coso Date: Sun, 14 Dec 2025 11:37:35 +0800 Subject: [PATCH] feat: add Claude Code streaming support and enhanced logging - Add SSE streaming response for Anthropic /v1/messages endpoint - Add tool_calls support in response parsing - Add detailed request/response logging for debugging - Enhance API compatibility check with tool_call test - Add Claude Code compatibility check UI in Settings --- src-tauri/src/lib.rs | 120 +++-- src-tauri/src/server.rs | 929 +++++++++++++++++++++++++++++++----- src/components/Settings.tsx | 32 +- 3 files changed, 922 insertions(+), 159 deletions(-) diff --git a/src-tauri/src/lib.rs b/src-tauri/src/lib.rs index d1fe4a706..0d2d8ceca 100644 --- a/src-tauri/src/lib.rs +++ b/src-tauri/src/lib.rs @@ -792,39 +792,87 @@ async fn check_api_compatibility( ) -> Result { logs.write().await.add( "info", - &format!("[API检测] 开始检测 {provider} API 兼容性..."), + &format!("[API检测] 开始检测 {provider} API 兼容性 (Claude Code 功能测试)..."), ); let s = state.read().await; let mut results: Vec = Vec::new(); let mut warnings: Vec = Vec::new(); - let models_to_check = match provider.as_str() { - "kiro" => vec!["claude-sonnet-4-5", "claude-3-7-sonnet-20250219"], - "gemini" => vec!["gemini-2.5-flash", "gemini-2.5-pro"], - "qwen" => vec!["qwen3-coder-plus", "qwen3-coder-flash"], + // Claude Code 需要的测试项目 + let test_cases: Vec<(&str, &str)> = match provider.as_str() { + "kiro" => vec![ + ("claude-sonnet-4-5", "basic"), // 基础对话 + ("claude-sonnet-4-5", "tool_call"), // Tool Calls 支持 + ], + "gemini" => vec![("gemini-2.5-flash", "basic"), ("gemini-2.5-pro", "basic")], + "qwen" => vec![ + ("qwen3-coder-plus", "basic"), + ("qwen3-coder-flash", "basic"), + ], _ => vec![], }; - for model in models_to_check { + for (model, test_type) in test_cases { let start = std::time::Instant::now(); + let test_name = format!("{model} ({test_type})"); - // 构建简单的测试请求 - let test_request = crate::models::openai::ChatCompletionRequest { - model: model.to_string(), - messages: vec![crate::models::openai::ChatMessage { - role: "user".to_string(), - content: Some(crate::models::openai::MessageContent::Text( - "Hi".to_string(), - )), - tool_calls: None, - tool_call_id: None, - }], - temperature: None, - max_tokens: Some(10), - stream: false, - tools: None, - tool_choice: None, + // 根据测试类型构建不同的请求 + let test_request = match test_type { + "tool_call" => { + // 测试 Tool Calls - Claude Code 核心功能 + crate::models::openai::ChatCompletionRequest { + model: model.to_string(), + messages: vec![crate::models::openai::ChatMessage { + role: "user".to_string(), + content: Some(crate::models::openai::MessageContent::Text( + "What is 2+2? Use the calculator tool to compute this.".to_string(), + )), + tool_calls: None, + tool_call_id: None, + }], + temperature: None, + max_tokens: Some(100), + stream: false, + tools: Some(vec![crate::models::openai::Tool { + tool_type: "function".to_string(), + function: crate::models::openai::FunctionDef { + name: "calculator".to_string(), + description: Some("Perform basic arithmetic calculations".to_string()), + parameters: Some(serde_json::json!({ + "type": "object", + "properties": { + "expression": { + "type": "string", + "description": "The math expression to evaluate" + } + }, + "required": ["expression"] + })), + }, + }]), + tool_choice: None, + } + } + _ => { + // 基础对话测试 + crate::models::openai::ChatCompletionRequest { + model: model.to_string(), + messages: vec![crate::models::openai::ChatMessage { + role: "user".to_string(), + content: Some(crate::models::openai::MessageContent::Text( + "Say 'OK' only.".to_string(), + )), + tool_calls: None, + tool_call_id: None, + }], + temperature: None, + max_tokens: Some(10), + stream: false, + tools: None, + tool_choice: None, + } + } }; let result = match provider.as_str() { @@ -840,33 +888,43 @@ async fn check_api_compatibility( let body = resp.text().await.unwrap_or_default(); let (available, error_type, error_message) = if (200..300).contains(&status) { + // 对于 tool_call 测试,额外检查响应是否包含 tool use + if test_type == "tool_call" { + let has_tool_use = + body.contains("\"name\"") && body.contains("\"toolUseId\""); + if !has_tool_use { + warnings.push(format!( + "{test_name}: 响应未包含 tool_use,Claude Code 可能无法正常工作" + )); + } + } (true, None, None) } else { let err_type = match status { 401 => { - warnings.push(format!("模型 {model} 返回 401: Token 可能已过期或无效")); + warnings.push(format!("{test_name} 返回 401: Token 可能已过期或无效")); Some("AUTH_ERROR".to_string()) } 403 => { warnings.push(format!( - "模型 {model} 返回 403: 无权访问,可能需要刷新 Token" + "{test_name} 返回 403: 无权访问,可能需要刷新 Token" )); Some("FORBIDDEN".to_string()) } 400 => { - warnings.push(format!("模型 {model} 返回 400: 请求格式可能已变更")); + warnings.push(format!("{test_name} 返回 400: 请求格式可能已变更")); Some("BAD_REQUEST".to_string()) } 404 => { - warnings.push(format!("模型 {model} 返回 404: 模型或接口可能已下线")); + warnings.push(format!("{test_name} 返回 404: 模型或接口可能已下线")); Some("NOT_FOUND".to_string()) } 429 => { - warnings.push(format!("模型 {model} 返回 429: 请求过于频繁")); + warnings.push(format!("{test_name} 返回 429: 请求过于频繁")); Some("RATE_LIMITED".to_string()) } 500..=599 => { - warnings.push(format!("模型 {model} 返回 {status}: 服务端错误")); + warnings.push(format!("{test_name} 返回 {status}: 服务端错误")); Some("SERVER_ERROR".to_string()) } _ => Some("UNKNOWN_ERROR".to_string()), @@ -879,7 +937,7 @@ async fn check_api_compatibility( }; results.push(ApiCheckResult { - model: model.to_string(), + model: test_name, available, status, error_type, @@ -888,9 +946,9 @@ async fn check_api_compatibility( }); } Err(e) => { - warnings.push(format!("模型 {model} 请求失败: {e}")); + warnings.push(format!("{test_name} 请求失败: {e}")); results.push(ApiCheckResult { - model: model.to_string(), + model: test_name, available: false, status: 0, error_type: Some("REQUEST_FAILED".to_string()), diff --git a/src-tauri/src/server.rs b/src-tauri/src/server.rs index 1e5d81cb9..c01d1bd4f 100644 --- a/src-tauri/src/server.rs +++ b/src-tauri/src/server.rs @@ -10,12 +10,14 @@ use crate::providers::kiro::KiroProvider; use crate::providers::openai_custom::OpenAICustomProvider; use crate::providers::qwen::QwenProvider; use axum::{ + body::Body, extract::State, - http::{HeaderMap, StatusCode}, + http::{header, HeaderMap, StatusCode}, response::{IntoResponse, Response}, routing::{get, post}, Json, Router, }; +use futures::stream; use serde::{Deserialize, Serialize}; use std::sync::Arc; use tokio::sync::{oneshot, RwLock}; @@ -246,22 +248,73 @@ async fn chat_completions( state.logs.write().await.add( "info", - &format!("POST /v1/chat/completions model={}", request.model), + &format!( + "POST /v1/chat/completions model={} stream={}", + request.model, request.stream + ), ); + // 检查是否需要刷新 token + { + let mut kiro = state.kiro.write().await; + if kiro.credentials.access_token.is_none() { + if let Err(e) = kiro.refresh_token().await { + state + .logs + .write() + .await + .add("error", &format!("Token refresh failed: {e}")); + return ( + StatusCode::UNAUTHORIZED, + Json(serde_json::json!({"error": {"message": format!("Token refresh failed: {e}")}})), + ).into_response(); + } + } + } + let kiro = state.kiro.read().await; match kiro.call_api(&request).await { Ok(resp) => { - if resp.status().is_success() { - // 解析 CodeWhisperer 响应并转换 + let status = resp.status(); + if status.is_success() { match resp.text().await { Ok(body) => { - state - .logs - .write() - .await - .add("info", "Request completed successfully"); + let parsed = parse_cw_response(&body); + let has_tool_calls = !parsed.tool_calls.is_empty(); + + state.logs.write().await.add( + "info", + &format!( + "Request completed: content_len={}, tool_calls={}", + parsed.content.len(), + parsed.tool_calls.len() + ), + ); + + // 构建消息 + let message = if has_tool_calls { + serde_json::json!({ + "role": "assistant", + "content": if parsed.content.is_empty() { serde_json::Value::Null } else { serde_json::json!(parsed.content) }, + "tool_calls": parsed.tool_calls.iter().map(|tc| { + serde_json::json!({ + "id": tc.id, + "type": "function", + "function": { + "name": tc.function.name, + "arguments": tc.function.arguments + } + }) + }).collect::>() + }) + } else { + serde_json::json!({ + "role": "assistant", + "content": parsed.content + }) + }; + let response = serde_json::json!({ "id": format!("chatcmpl-{}", uuid::Uuid::new_v4()), "object": "chat.completion", @@ -272,11 +325,8 @@ async fn chat_completions( "model": request.model, "choices": [{ "index": 0, - "message": { - "role": "assistant", - "content": extract_content_from_cw_response(&body) - }, - "finish_reason": "stop" + "message": message, + "finish_reason": if has_tool_calls { "tool_calls" } else { "stop" } }], "usage": { "prompt_tokens": 0, @@ -292,87 +342,96 @@ async fn chat_completions( ) .into_response(), } - } else { - let status = resp.status(); - let body = resp.text().await.unwrap_or_default(); - ( - StatusCode::from_u16(status.as_u16()).unwrap_or(StatusCode::INTERNAL_SERVER_ERROR), - Json(serde_json::json!({"error": {"message": format!("Upstream error: {}", body)}})) - ).into_response() - } - } - Err(e) => ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": e.to_string()}})), - ) - .into_response(), - } -} + } else if status.as_u16() == 403 { + // Token 过期,尝试刷新 + drop(kiro); + let mut kiro = state.kiro.write().await; + state + .logs + .write() + .await + .add("warn", "Got 403, attempting token refresh"); -async fn anthropic_messages( - State(state): State, - headers: HeaderMap, - Json(request): Json, -) -> Response { - if let Err(e) = verify_api_key(&headers, &state.api_key).await { - state - .logs - .write() - .await - .add("warn", "Unauthorized request to /v1/messages"); - return e.into_response(); - } + match kiro.refresh_token().await { + Ok(_) => { + // 重试请求 + drop(kiro); + let kiro = state.kiro.read().await; + match kiro.call_api(&request).await { + Ok(retry_resp) => { + if retry_resp.status().is_success() { + match retry_resp.text().await { + Ok(body) => { + let parsed = parse_cw_response(&body); + let has_tool_calls = !parsed.tool_calls.is_empty(); - state.logs.write().await.add( - "info", - &format!("POST /v1/messages (Anthropic) model={}", request.model), - ); + let message = if has_tool_calls { + serde_json::json!({ + "role": "assistant", + "content": if parsed.content.is_empty() { serde_json::Value::Null } else { serde_json::json!(parsed.content) }, + "tool_calls": parsed.tool_calls.iter().map(|tc| { + serde_json::json!({ + "id": tc.id, + "type": "function", + "function": { + "name": tc.function.name, + "arguments": tc.function.arguments + } + }) + }).collect::>() + }) + } else { + serde_json::json!({ + "role": "assistant", + "content": parsed.content + }) + }; - // 转换为 OpenAI 格式 - let openai_request = convert_anthropic_to_openai(&request); - let kiro = state.kiro.read().await; - - match kiro.call_api(&openai_request).await { - Ok(resp) => { - if resp.status().is_success() { - match resp.text().await { - Ok(body) => { - let content = extract_content_from_cw_response(&body); - state - .logs - .write() - .await - .add("info", "Anthropic request completed successfully"); - // 返回 Anthropic 格式响应 - let response = serde_json::json!({ - "id": format!("msg_{}", uuid::Uuid::new_v4()), - "type": "message", - "role": "assistant", - "content": [{"type": "text", "text": content}], - "model": request.model, - "stop_reason": "end_turn", - "usage": { - "input_tokens": 0, - "output_tokens": 0 + let response = serde_json::json!({ + "id": format!("chatcmpl-{}", uuid::Uuid::new_v4()), + "object": "chat.completion", + "created": std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap() + .as_secs(), + "model": request.model, + "choices": [{ + "index": 0, + "message": message, + "finish_reason": if has_tool_calls { "tool_calls" } else { "stop" } + }], + "usage": { + "prompt_tokens": 0, + "completion_tokens": 0, + "total_tokens": 0 + } + }); + return Json(response).into_response(); + } + Err(e) => return ( + StatusCode::INTERNAL_SERVER_ERROR, + Json(serde_json::json!({"error": {"message": e.to_string()}})), + ).into_response(), + } + } + let body = retry_resp.text().await.unwrap_or_default(); + ( + StatusCode::INTERNAL_SERVER_ERROR, + Json(serde_json::json!({"error": {"message": format!("Retry failed: {}", body)}})), + ).into_response() } - }); - Json(response).into_response() - } - Err(e) => { - state - .logs - .write() - .await - .add("error", &format!("Response parse error: {e}")); - ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": e.to_string()}})), - ) - .into_response() + Err(e) => ( + StatusCode::INTERNAL_SERVER_ERROR, + Json(serde_json::json!({"error": {"message": e.to_string()}})), + ).into_response(), + } } + Err(e) => ( + StatusCode::UNAUTHORIZED, + Json(serde_json::json!({"error": {"message": format!("Token refresh failed: {e}")}})), + ).into_response(), } } else { - let status = resp.status(); let body = resp.text().await.unwrap_or_default(); state.logs.write().await.add( "error", @@ -403,6 +462,494 @@ async fn anthropic_messages( } } +async fn anthropic_messages( + State(state): State, + headers: HeaderMap, + Json(request): Json, +) -> Response { + if let Err(e) = verify_api_key(&headers, &state.api_key).await { + state + .logs + .write() + .await + .add("warn", "Unauthorized request to /v1/messages"); + return e.into_response(); + } + + // 详细记录请求信息 + let msg_count = request.messages.len(); + let has_tools = request.tools.as_ref().map(|t| t.len()).unwrap_or(0); + let has_system = request.system.is_some(); + state.logs.write().await.add( + "info", + &format!( + "[REQ] POST /v1/messages model={} stream={} messages={} tools={} has_system={}", + request.model, request.stream, msg_count, has_tools, has_system + ), + ); + + // 记录最后一条消息的角色和内容预览 + if let Some(last_msg) = request.messages.last() { + let content_preview = match &last_msg.content { + serde_json::Value::String(s) => s.chars().take(100).collect::(), + serde_json::Value::Array(arr) => { + if let Some(first) = arr.first() { + if let Some(text) = first.get("text").and_then(|t| t.as_str()) { + text.chars().take(100).collect::() + } else { + format!("[{} blocks]", arr.len()) + } + } else { + "[empty]".to_string() + } + } + _ => "[unknown]".to_string(), + }; + state.logs.write().await.add( + "debug", + &format!( + "[REQ] Last message: role={} content={}", + last_msg.role, content_preview + ), + ); + } + + // 检查是否需要刷新 token + { + let mut kiro = state.kiro.write().await; + if kiro.credentials.access_token.is_none() { + state + .logs + .write() + .await + .add("info", "[AUTH] No access token, attempting refresh..."); + if let Err(e) = kiro.refresh_token().await { + state + .logs + .write() + .await + .add("error", &format!("[AUTH] Token refresh failed: {e}")); + return ( + StatusCode::UNAUTHORIZED, + Json(serde_json::json!({"error": {"message": format!("Token refresh failed: {e}")}})), + ) + .into_response(); + } + state + .logs + .write() + .await + .add("info", "[AUTH] Token refreshed successfully"); + } + } + + // 转换为 OpenAI 格式 + let openai_request = convert_anthropic_to_openai(&request); + + // 记录转换后的请求信息 + state.logs.write().await.add( + "debug", + &format!( + "[CONVERT] OpenAI format: messages={} tools={} stream={}", + openai_request.messages.len(), + openai_request.tools.as_ref().map(|t| t.len()).unwrap_or(0), + openai_request.stream + ), + ); + + let kiro = state.kiro.read().await; + + match kiro.call_api(&openai_request).await { + Ok(resp) => { + let status = resp.status(); + state + .logs + .write() + .await + .add("info", &format!("[RESP] Upstream status: {status}")); + + 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 parsed = parse_cw_response(&body); + + // 详细记录解析结果 + state.logs.write().await.add( + "info", + &format!( + "[RESP] Parsed: content_len={}, tool_calls={}, content_preview={}", + parsed.content.len(), + parsed.tool_calls.len(), + parsed.content.chars().take(100).collect::() + ), + ); + + // 记录 tool calls 详情 + for (i, tc) in parsed.tool_calls.iter().enumerate() { + state.logs.write().await.add( + "debug", + &format!( + "[RESP] Tool call {}: name={} id={}", + i, tc.function.name, tc.id + ), + ); + } + + // 如果请求流式响应,返回 SSE 格式 + if request.stream { + return build_anthropic_stream_response( + &request.model, + &parsed.content, + &parsed.tool_calls, + ); + } + + // 非流式响应 + build_anthropic_response( + &request.model, + &parsed.content, + &parsed.tool_calls, + ) + } + Err(e) => { + state + .logs + .write() + .await + .add("error", &format!("[ERROR] Response body read failed: {e}")); + ( + StatusCode::INTERNAL_SERVER_ERROR, + Json(serde_json::json!({"error": {"message": e.to_string()}})), + ) + .into_response() + } + } + } else if status.as_u16() == 403 { + // Token 过期,尝试刷新 + drop(kiro); + let mut kiro = state.kiro.write().await; + state.logs.write().await.add( + "warn", + "[AUTH] Got 403 Forbidden, attempting token refresh...", + ); + + match kiro.refresh_token().await { + Ok(_) => { + state + .logs + .write() + .await + .add("info", "[AUTH] Token refreshed, retrying request..."); + drop(kiro); + let kiro = state.kiro.read().await; + match kiro.call_api(&openai_request).await { + Ok(retry_resp) => { + let retry_status = retry_resp.status(); + state.logs.write().await.add( + "info", + &format!("[RETRY] Response status: {retry_status}"), + ); + if retry_resp.status().is_success() { + match retry_resp.text().await { + Ok(body) => { + let parsed = parse_cw_response(&body); + state.logs.write().await.add( + "info", + &format!( + "[RETRY] Success: content_len={}, tool_calls={}", + parsed.content.len(), parsed.tool_calls.len() + ), + ); + if request.stream { + return build_anthropic_stream_response( + &request.model, + &parsed.content, + &parsed.tool_calls, + ); + } + return build_anthropic_response( + &request.model, + &parsed.content, + &parsed.tool_calls, + ); + } + Err(e) => { + state.logs.write().await.add( + "error", + &format!("[RETRY] Body read failed: {e}"), + ); + return ( + StatusCode::INTERNAL_SERVER_ERROR, + Json(serde_json::json!({"error": {"message": e.to_string()}})), + ) + .into_response(); + } + } + } + let body = retry_resp.text().await.unwrap_or_default(); + state.logs.write().await.add( + "error", + &format!( + "[RETRY] Failed with status {retry_status}: {}", + &body[..body.len().min(500)] + ), + ); + ( + StatusCode::INTERNAL_SERVER_ERROR, + Json(serde_json::json!({"error": {"message": format!("Retry failed: {}", body)}})), + ) + .into_response() + } + Err(e) => { + state + .logs + .write() + .await + .add("error", &format!("[RETRY] Request failed: {e}")); + ( + StatusCode::INTERNAL_SERVER_ERROR, + Json(serde_json::json!({"error": {"message": e.to_string()}})), + ) + .into_response() + } + } + } + Err(e) => { + state + .logs + .write() + .await + .add("error", &format!("[AUTH] Token refresh failed: {e}")); + ( + StatusCode::UNAUTHORIZED, + Json(serde_json::json!({"error": {"message": format!("Token refresh failed: {e}")}})), + ) + .into_response() + } + } + } else { + let body = resp.text().await.unwrap_or_default(); + state.logs.write().await.add( + "error", + &format!( + "[ERROR] Upstream error HTTP {}: {}", + status, + &body[..body.len().min(500)] + ), + ); + ( + StatusCode::from_u16(status.as_u16()) + .unwrap_or(StatusCode::INTERNAL_SERVER_ERROR), + Json( + serde_json::json!({"error": {"message": format!("Upstream error: {}", body)}}), + ), + ) + .into_response() + } + } + Err(e) => { + // 详细记录网络/连接错误 + let error_details = format!("{e:?}"); + state + .logs + .write() + .await + .add("error", &format!("[ERROR] Kiro API call failed: {e}")); + state.logs.write().await.add( + "debug", + &format!("[ERROR] Full error details: {error_details}"), + ); + ( + StatusCode::INTERNAL_SERVER_ERROR, + Json(serde_json::json!({"error": {"message": e.to_string()}})), + ) + .into_response() + } + } +} + +/// 构建 Anthropic 非流式响应 +fn build_anthropic_response(model: &str, content: &str, tool_calls: &[ToolCall]) -> Response { + let has_tool_calls = !tool_calls.is_empty(); + let mut content_array: Vec = Vec::new(); + + if !content.is_empty() { + content_array.push(serde_json::json!({ + "type": "text", + "text": content + })); + } + + for tc in tool_calls { + let input: serde_json::Value = + serde_json::from_str(&tc.function.arguments).unwrap_or(serde_json::json!({})); + content_array.push(serde_json::json!({ + "type": "tool_use", + "id": tc.id, + "name": tc.function.name, + "input": input + })); + } + + if content_array.is_empty() { + content_array.push(serde_json::json!({"type": "text", "text": ""})); + } + + let response = serde_json::json!({ + "id": format!("msg_{}", uuid::Uuid::new_v4()), + "type": "message", + "role": "assistant", + "content": content_array, + "model": model, + "stop_reason": if has_tool_calls { "tool_use" } else { "end_turn" }, + "stop_sequence": null, + "usage": {"input_tokens": 0, "output_tokens": 0} + }); + Json(response).into_response() +} + +/// 构建 Anthropic 流式响应 (SSE) +fn build_anthropic_stream_response( + model: &str, + content: &str, + tool_calls: &[ToolCall], +) -> Response { + let has_tool_calls = !tool_calls.is_empty(); + let message_id = format!("msg_{}", uuid::Uuid::new_v4()); + let model = model.to_string(); + let content = content.to_string(); + let tool_calls = tool_calls.to_vec(); + + // 构建 SSE 事件流 + let mut events: Vec = Vec::new(); + + // 1. message_start + let message_start = serde_json::json!({ + "type": "message_start", + "message": { + "id": message_id, + "type": "message", + "role": "assistant", + "model": model, + "content": [], + "stop_reason": null, + "stop_sequence": null, + "usage": {"input_tokens": 0, "output_tokens": 0} + } + }); + events.push(format!("event: message_start\ndata: {message_start}\n\n")); + + 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" + )); + + // content_block_delta - 发送完整内容 + let block_delta = serde_json::json!({ + "type": "content_block_delta", + "index": block_index, + "delta": {"type": "text_delta", "text": content} + }); + 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; + } + + // 3. Tool use 块 + for tc in &tool_calls { + let input: serde_json::Value = + serde_json::from_str(&tc.function.arguments).unwrap_or(serde_json::json!({})); + + // content_block_start + let block_start = serde_json::json!({ + "type": "content_block_start", + "index": block_index, + "content_block": { + "type": "tool_use", + "id": tc.id, + "name": tc.function.name, + "input": {} + } + }); + events.push(format!( + "event: content_block_start\ndata: {block_start}\n\n" + )); + + // content_block_delta - input_json_delta + let block_delta = serde_json::json!({ + "type": "content_block_delta", + "index": block_index, + "delta": { + "type": "input_json_delta", + "partial_json": serde_json::to_string(&input).unwrap_or_default() + } + }); + 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; + } + + // 4. message_delta + let message_delta = serde_json::json!({ + "type": "message_delta", + "delta": { + "stop_reason": if has_tool_calls { "tool_use" } else { "end_turn" }, + "stop_sequence": null + }, + "usage": {"output_tokens": 0} + }); + events.push(format!("event: message_delta\ndata: {message_delta}\n\n")); + + // 5. message_stop + let message_stop = serde_json::json!({"type": "message_stop"}); + events.push(format!("event: message_stop\ndata: {message_stop}\n\n")); + + // 创建 SSE 响应 + let body_stream = stream::iter(events.into_iter().map(Ok::<_, std::convert::Infallible>)); + let body = Body::from_stream(body_stream); + + Response::builder() + .status(StatusCode::OK) + .header(header::CONTENT_TYPE, "text/event-stream") + .header(header::CACHE_CONTROL, "no-cache") + .header(header::CONNECTION, "keep-alive") + .body(body) + .unwrap() +} + async fn count_tokens( State(state): State, headers: HeaderMap, @@ -419,42 +966,186 @@ async fn count_tokens( .into_response() } -fn extract_content_from_cw_response(body: &str) -> String { - // CodeWhisperer 返回 AWS Event Stream 格式 - // 使用正则提取 JSON 内容 - let mut content = String::new(); +/// CodeWhisperer 响应解析结果 +#[derive(Debug, Default)] +struct CWParsedResponse { + content: String, + tool_calls: Vec, +} - // 查找所有 {"content":"..."} 模式 - let re = regex::Regex::new(r#"\{"content":"([^"\\]*(\\.[^"\\]*)*)"\}"#).ok(); +/// 解析 CodeWhisperer AWS Event Stream 响应 +fn parse_cw_response(body: &str) -> CWParsedResponse { + let mut result = CWParsedResponse::default(); + let mut current_tool: Option<(String, String, String)> = None; // (id, name, input) + + // 解析所有 JSON 事件 + let patterns = [ + r#"\{"content":"#, + r#"\{"name":"#, + r#"\{"input":"#, + r#"\{"stop":"#, + ]; + + let mut pos = 0; + while pos < body.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); + } + } + + if next_start >= body.len() { + break; + } + + // 找到匹配的 } + if let Some(json_str) = extract_json_object(&body[next_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); + } + } + // 处理 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 + .get("input") + .and_then(|v| v.as_str()) + .unwrap_or("") + .to_string(); + current_tool = Some((tool_use_id.to_string(), name.to_string(), input)); + + // 如果同时有 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, + }, + }); + } + } + } + // 处理 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 事件 + 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(); + } else { + pos = next_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, + }, + }); + } + + // 解析 bracket 格式的 tool calls: [Called xxx with args: {...}] + parse_bracket_tool_calls(&mut result); + + result +} + +/// 从字符串中提取完整的 JSON 对象 +fn extract_json_object(s: &str) -> Option<&str> { + if !s.starts_with('{') { + return None; + } + + let mut brace_count = 0; + let mut in_string = false; + let mut escape_next = false; + + for (i, c) in s.char_indices() { + if escape_next { + escape_next = false; + continue; + } + + match c { + '\\' if in_string => escape_next = true, + '"' => in_string = !in_string, + '{' if !in_string => brace_count += 1, + '}' if !in_string => { + brace_count -= 1; + if brace_count == 0 { + return Some(&s[..=i]); + } + } + _ => {} + } + } + None +} + +/// 解析 bracket 格式的 tool calls +fn parse_bracket_tool_calls(result: &mut CWParsedResponse) { + let re = + regex::Regex::new(r"\[Called\s+(\w+)\s+with\s+args:\s*(\{[^}]*(?:\{[^}]*\}[^}]*)*\})\]") + .ok(); if let Some(re) = re { - for cap in re.captures_iter(body) { - if let Some(text) = cap.get(1) { - // 处理转义字符 - let unescaped = text - .as_str() - .replace("\\n", "\n") - .replace("\\t", "\t") - .replace("\\\"", "\"") - .replace("\\\\", "\\"); - content.push_str(&unescaped); + let mut to_remove = Vec::new(); + for cap in re.captures_iter(&result.content) { + if let (Some(name), Some(args)) = (cap.get(1), cap.get(2)) { + let tool_id = format!( + "call_{}", + &uuid::Uuid::new_v4().to_string().replace('-', "")[..8] + ); + result.tool_calls.push(ToolCall { + id: tool_id, + call_type: "function".to_string(), + function: FunctionCall { + name: name.as_str().to_string(), + arguments: args.as_str().to_string(), + }, + }); + to_remove.push(cap.get(0).unwrap().as_str().to_string()); } } - } - - if content.is_empty() { - // 备用方案:查找 assistantResponseEvent - if let Some(start) = body.find(r#""content":""#) { - let rest = &body[start + 11..]; - if let Some(end) = rest.find('"') { - content = rest[..end].to_string(); - } + // 从 content 中移除 tool call 文本 + for s in to_remove { + result.content = result.content.replace(&s, ""); } - } - - if content.is_empty() { - "Response received but could not parse content".to_string() - } else { - content + result.content = result.content.trim().to_string(); } } diff --git a/src/components/Settings.tsx b/src/components/Settings.tsx index 04e0a7fae..929d8f7e8 100644 --- a/src/components/Settings.tsx +++ b/src/components/Settings.tsx @@ -213,28 +213,36 @@ export function Settings() { - {/* API 兼容性检测 */} + {/* Claude Code 兼容性检测 */}
- -

API 兼容性检测

+ +

Claude Code 兼容性检测

- 检测当前配置的模型是否可用,识别 API 变更或认证问题 + 检测 API 是否支持 Claude Code 所需的功能:基础对话、Tool Calls 等

+
+

检测项目:

+
    +
  • 基础对话能力 (basic)
  • +
  • Tool Calls 支持 (tool_call) - Claude Code 核心功能
  • +
+
+
@@ -259,9 +267,9 @@ export function Settings() {
- {/* 模型检测结果 */} + {/* 检测结果 */}
-

模型状态:

+

检测结果:

{checkResult.results.map((r) => (
)} - {r.model} + + {r.model.includes("tool_call") ? ( + {r.model} + ) : ( + r.model + )} +
{r.status > 0 && HTTP {r.status}}