diff --git a/src-tauri/src/proxy/providers/streaming_gemini.rs b/src-tauri/src/proxy/providers/streaming_gemini.rs index 885d9668c..19684d266 100644 --- a/src-tauri/src/proxy/providers/streaming_gemini.rs +++ b/src-tauri/src/proxy/providers/streaming_gemini.rs @@ -408,6 +408,32 @@ pub fn create_anthropic_sse_stream_from_gemini(Bytes::from(chunk))), + ); + let mut converted = Box::pin(create_anthropic_sse_stream_from_gemini( + stream, + Some(store.clone()), + Some("provider-a".to_string()), + Some("session-1".to_string()), + None, + )); + + futures::executor::block_on(async { + while let Some(item) = converted.next().await { + let event = String::from_utf8(item.unwrap().to_vec()).unwrap(); + if event.contains("\"type\":\"tool_use\"") { + break; + } + } + }); + + let shadow = store + .latest_assistant_content("provider-a", "session-1") + .unwrap(); + assert_eq!(shadow["parts"][0]["functionCall"]["name"], "Bash"); + assert_eq!(shadow["parts"][0]["thoughtSignature"], "sig-tool-1"); + } + #[test] fn rectifies_streamed_tool_call_args_from_tool_schema_hints() { let owned_chunks = vec![ diff --git a/src-tauri/src/proxy/providers/transform_gemini.rs b/src-tauri/src/proxy/providers/transform_gemini.rs index c4a1f9c1c..c7c65fb4d 100644 --- a/src-tauri/src/proxy/providers/transform_gemini.rs +++ b/src-tauri/src/proxy/providers/transform_gemini.rs @@ -268,6 +268,7 @@ fn convert_messages_to_contents( shadow_turns: &[GeminiAssistantTurn], ) -> Result, ProxyError> { let mut contents = Vec::new(); + let mut used_shadow_indices = HashSet::new(); let total_assistant_messages = messages .iter() .filter(|message| message.get("role").and_then(|value| value.as_str()) == Some("assistant")) @@ -290,12 +291,20 @@ fn convert_messages_to_contents( let gemini_role = if role == "assistant" { "model" } else { "user" }; let parts = if role == "assistant" { - let shadow_index = assistant_seen_index + let positional_shadow_index = assistant_seen_index .checked_sub(shadow_start_index) - .filter(|index| *index < effective_shadow_turns.len()); + .filter(|index| *index < effective_shadow_turns.len()) + .filter(|index| !used_shadow_indices.contains(index)); + let tool_use_match_index = find_matching_shadow_turn_for_assistant_message( + message.get("content"), + effective_shadow_turns, + ) + .filter(|index| !used_shadow_indices.contains(index)); assistant_seen_index += 1; + let shadow_index = tool_use_match_index.or(positional_shadow_index); if let Some(index) = shadow_index { + used_shadow_indices.insert(index); let shadow_turn = &effective_shadow_turns[index]; merge_tool_names_from_shadow(shadow_turn, &mut tool_name_by_id); if let Some(parts) = shadow_parts(&shadow_turn.assistant_content) { @@ -331,6 +340,67 @@ fn convert_messages_to_contents( Ok(contents) } +fn find_matching_shadow_turn_for_assistant_message( + content: Option<&Value>, + shadow_turns: &[GeminiAssistantTurn], +) -> Option { + let (tool_use_ids, tool_use_names) = extract_assistant_tool_use_keys(content); + if tool_use_ids.is_empty() && tool_use_names.is_empty() { + return None; + } + + shadow_turns.iter().enumerate().find_map(|(index, turn)| { + turn.tool_calls + .iter() + .any(|tool_call| { + tool_call + .id + .as_deref() + .is_some_and(|id| tool_use_ids.contains(id)) + || tool_use_names.contains(tool_call.name.as_str()) + || tool_use_names.contains(normalize_tool_name(&tool_call.name)) + }) + .then_some(index) + }) +} + +fn extract_assistant_tool_use_keys(content: Option<&Value>) -> (HashSet, HashSet) { + let mut tool_use_ids = HashSet::new(); + let mut tool_use_names = HashSet::new(); + let Some(blocks) = content.and_then(|value| value.as_array()) else { + return (tool_use_ids, tool_use_names); + }; + + for block in blocks { + if block.get("type").and_then(|value| value.as_str()) != Some("tool_use") { + continue; + } + + if let Some(id) = block + .get("id") + .and_then(|value| value.as_str()) + .filter(|id| !id.is_empty()) + { + tool_use_ids.insert(id.to_string()); + } + + if let Some(name) = block + .get("name") + .and_then(|value| value.as_str()) + .filter(|name| !name.is_empty()) + { + tool_use_names.insert(name.to_string()); + tool_use_names.insert(normalize_tool_name(name).to_string()); + } + } + + (tool_use_ids, tool_use_names) +} + +fn normalize_tool_name(name: &str) -> &str { + name.rsplit(':').next().unwrap_or(name) +} + fn convert_message_content_to_parts( content: Option<&Value>, role: &str, @@ -1259,4 +1329,70 @@ mod tests { "tool_2" ); } + + #[test] + fn shadow_replay_matches_tool_use_turn_by_id_when_position_drifts() { + let store = GeminiShadowStore::with_limits(8, 4); + store.record_assistant_turn( + "prov", + "sess", + json!({ + "parts": [{ + "functionCall": { + "id": "call_1", + "name": "Bash", + "args": { "command": "ls -R" } + }, + "thoughtSignature": "sig-tool-1" + }] + }), + vec![GeminiToolCallMeta::new( + Some("call_1"), + "Bash", + json!({ "command": "ls -R" }), + Some("sig-tool-1"), + )], + ); + + let input = json!({ + "messages": [ + { + "role": "assistant", + "content": [ + { + "type": "tool_use", + "id": "call_1", + "name": "default_api:Bash", + "input": { "command": "ls -R" } + } + ] + }, + { + "role": "user", + "content": [ + { "type": "tool_result", "tool_use_id": "call_1", "content": "ok" } + ] + }, + { + "role": "assistant", + "content": [ + { "type": "text", "text": "local-only assistant turn without Gemini shadow" } + ] + } + ] + }); + + let result = + anthropic_to_gemini_with_shadow(input, Some(&store), Some("prov"), Some("sess")) + .unwrap(); + + assert_eq!( + result["contents"][0]["parts"][0]["functionCall"]["name"], + "Bash" + ); + assert_eq!( + result["contents"][0]["parts"][0]["thoughtSignature"], + "sig-tool-1" + ); + } } diff --git a/src-tauri/src/proxy/session.rs b/src-tauri/src/proxy/session.rs index cd94a034d..c5a786aa0 100644 --- a/src-tauri/src/proxy/session.rs +++ b/src-tauri/src/proxy/session.rs @@ -242,6 +242,12 @@ pub fn extract_session_id( body: &serde_json::Value, client_format: &str, ) -> SessionIdResult { + if client_format == "claude" { + if let Some(result) = extract_claude_session(headers, body) { + return result; + } + } + // Codex 请求特殊处理 if client_format == "codex" || client_format == "openai" { if let Some(result) = extract_codex_session(headers, body) { @@ -258,6 +264,28 @@ pub fn extract_session_id( generate_new_session_id() } +/// 提取 Claude Session ID +fn extract_claude_session( + headers: &HeaderMap, + body: &serde_json::Value, +) -> Option { + for header_name in &["x-claude-code-session-id", "claude-code-session-id"] { + if let Some(value) = headers.get(*header_name) { + if let Ok(session_id) = value.to_str() { + if !session_id.is_empty() { + return Some(SessionIdResult { + session_id: session_id.to_string(), + source: SessionIdSource::Header, + client_provided: true, + }); + } + } + } + } + + extract_from_metadata(body) +} + /// 提取 Codex Session ID fn extract_codex_session(headers: &HeaderMap, body: &serde_json::Value) -> Option { // 1. 从 headers 提取 @@ -515,6 +543,47 @@ mod tests { assert!(result.client_provided); } + #[test] + fn test_extract_session_from_claude_header() { + let mut headers = HeaderMap::new(); + headers.insert( + "x-claude-code-session-id", + "d937243f-2702-4f20-97b6-c9682235ab81".parse().unwrap(), + ); + let body = json!({ + "model": "claude-3-5-sonnet", + "messages": [{"role": "user", "content": "Hello"}] + }); + + let result = extract_session_id(&headers, &body, "claude"); + + assert_eq!(result.session_id, "d937243f-2702-4f20-97b6-c9682235ab81"); + assert_eq!(result.source, SessionIdSource::Header); + assert!(result.client_provided); + } + + #[test] + fn test_extract_session_from_claude_header_precedes_metadata() { + let mut headers = HeaderMap::new(); + headers.insert( + "x-claude-code-session-id", + "header-session-123".parse().unwrap(), + ); + let body = json!({ + "model": "claude-3-5-sonnet", + "messages": [{"role": "user", "content": "Hello"}], + "metadata": { + "session_id": "my-session-123" + } + }); + + let result = extract_session_id(&headers, &body, "claude"); + + assert_eq!(result.session_id, "header-session-123"); + assert_eq!(result.source, SessionIdSource::Header); + assert!(result.client_provided); + } + #[test] fn test_extract_session_from_codex_previous_response_id() { let headers = HeaderMap::new();