diff --git a/src-tauri/src/proxy/handlers.rs b/src-tauri/src/proxy/handlers.rs index a672589e4..b8a233e52 100644 --- a/src-tauri/src/proxy/handlers.rs +++ b/src-tauri/src/proxy/handlers.rs @@ -145,11 +145,13 @@ async fn handle_claude_transform( response: super::hyper_client::ProxyResponse, ctx: &RequestContext, state: &ProxyState, - _original_body: &Value, + original_body: &Value, is_stream: bool, api_format: &str, ) -> Result { let status = response.status(); + let tool_schema_hints = transform_gemini::extract_anthropic_tool_schema_hints(original_body); + let tool_schema_hints = (!tool_schema_hints.is_empty()).then_some(tool_schema_hints); if is_stream { // 根据 api_format 选择流式转换器 @@ -164,6 +166,7 @@ async fn handle_claude_transform( Some(state.gemini_shadow.clone()), Some(ctx.provider.id.clone()), Some(ctx.session_id.clone()), + tool_schema_hints.clone(), ))) } else { Box::new(Box::pin(create_anthropic_sse_stream(stream))) @@ -254,11 +257,12 @@ async fn handle_claude_transform( let anthropic_response = if api_format == "openai_responses" { transform_responses::responses_to_anthropic(upstream_response) } else if api_format == "gemini_native" { - transform_gemini::gemini_to_anthropic_with_shadow( + transform_gemini::gemini_to_anthropic_with_shadow_and_hints( upstream_response, Some(state.gemini_shadow.as_ref()), Some(&ctx.provider.id), Some(&ctx.session_id), + tool_schema_hints.as_ref(), ) } else { transform::openai_to_anthropic(upstream_response) diff --git a/src-tauri/src/proxy/providers/streaming_gemini.rs b/src-tauri/src/proxy/providers/streaming_gemini.rs index d68662167..76435a8ff 100644 --- a/src-tauri/src/proxy/providers/streaming_gemini.rs +++ b/src-tauri/src/proxy/providers/streaming_gemini.rs @@ -4,6 +4,7 @@ //! SSE events for Claude-compatible clients. use super::gemini_shadow::{GeminiShadowStore, GeminiToolCallMeta}; +use super::transform_gemini::{rectify_tool_call_parts, AnthropicToolSchemaHints}; use crate::proxy::sse::{strip_sse_field, take_sse_block}; use bytes::Bytes; use futures::stream::{Stream, StreamExt}; @@ -69,8 +70,14 @@ fn extract_visible_text(parts: &[Value]) -> String { .collect::() } -fn extract_tool_calls(parts: &[Value]) -> Vec { - parts +fn extract_tool_calls( + parts: &[Value], + tool_schema_hints: Option<&AnthropicToolSchemaHints>, +) -> Vec { + let mut rectified_parts = parts.to_vec(); + rectify_tool_call_parts(&mut rectified_parts, tool_schema_hints); + + rectified_parts .iter() .filter_map(|part| { let function_call = part.get("functionCall")?; @@ -174,6 +181,7 @@ pub fn create_anthropic_sse_stream_from_gemini>, provider_id: Option, session_id: Option, + tool_schema_hints: Option, ) -> impl Stream> + Send { async_stream::stream! { let mut buffer = String::new(); @@ -278,14 +286,16 @@ pub fn create_anthropic_sse_stream_from_gemini(Bytes::from(chunk))), ); - let converted = create_anthropic_sse_stream_from_gemini(stream, None, None, None); + let converted = create_anthropic_sse_stream_from_gemini(stream, None, None, None, None); futures::executor::block_on(async move { converted .collect::>() @@ -520,6 +530,7 @@ mod tests { Some(store), Some(provider_id.to_string()), Some(session_id.to_string()), + None, ); futures::executor::block_on(async move { converted @@ -616,4 +627,42 @@ mod tests { "sig-1" ); } + + #[test] + fn rectifies_streamed_tool_call_args_from_tool_schema_hints() { + let owned_chunks = vec![ + "data: {\"responseId\":\"resp_5\",\"modelVersion\":\"gemini-2.5-pro\",\"candidates\":[{\"finishReason\":\"STOP\",\"content\":{\"parts\":[{\"functionCall\":{\"id\":\"call_1\",\"name\":\"Bash\",\"args\":{\"args\":\"git status\"}}}]}}],\"usageMetadata\":{\"promptTokenCount\":5,\"totalTokenCount\":8}}\n\n".to_string(), + ]; + let stream = futures::stream::iter( + owned_chunks + .into_iter() + .map(|chunk| Ok::(Bytes::from(chunk))), + ); + let hints = super::super::transform_gemini::extract_anthropic_tool_schema_hints(&json!({ + "tools": [{ + "name": "Bash", + "input_schema": { + "type": "object", + "properties": { + "command": { "type": "string" }, + "timeout": { "type": "number" } + }, + "required": ["command"] + } + }] + })); + let converted = + create_anthropic_sse_stream_from_gemini(stream, None, None, None, Some(hints)); + let output = futures::executor::block_on(async move { + converted + .collect::>() + .await + .into_iter() + .map(|item| String::from_utf8(item.unwrap().to_vec()).unwrap()) + .collect::>() + .join("") + }); + + assert!(output.contains("\"partial_json\":\"{\\\"command\\\":\\\"git status\\\"}\"")); + } } diff --git a/src-tauri/src/proxy/providers/transform_gemini.rs b/src-tauri/src/proxy/providers/transform_gemini.rs index 2b0a601f4..96becd2ae 100644 --- a/src-tauri/src/proxy/providers/transform_gemini.rs +++ b/src-tauri/src/proxy/providers/transform_gemini.rs @@ -8,6 +8,15 @@ use super::gemini_schema::build_gemini_function_declaration; use super::gemini_shadow::{GeminiAssistantTurn, GeminiShadowStore, GeminiToolCallMeta}; use crate::proxy::error::ProxyError; use serde_json::{json, Map, Value}; +use std::collections::{HashMap, HashSet}; + +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub struct AnthropicToolSchemaHint { + expected_keys: Vec, + required_keys: Vec, +} + +pub type AnthropicToolSchemaHints = HashMap; pub fn anthropic_to_gemini(body: Value) -> Result { anthropic_to_gemini_with_shadow(body, None, None, None) @@ -77,6 +86,16 @@ pub fn gemini_to_anthropic_with_shadow( shadow_store: Option<&GeminiShadowStore>, provider_id: Option<&str>, session_id: Option<&str>, +) -> Result { + gemini_to_anthropic_with_shadow_and_hints(body, shadow_store, provider_id, session_id, None) +} + +pub fn gemini_to_anthropic_with_shadow_and_hints( + body: Value, + shadow_store: Option<&GeminiShadowStore>, + provider_id: Option<&str>, + session_id: Option<&str>, + tool_schema_hints: Option<&AnthropicToolSchemaHints>, ) -> Result { if let Some(block_reason) = body .get("promptFeedback") @@ -111,10 +130,13 @@ pub fn gemini_to_anthropic_with_shadow( .cloned() .unwrap_or_default(); + let mut rectified_parts = parts.clone(); + rectify_tool_call_parts(&mut rectified_parts, tool_schema_hints); + let mut content = Vec::new(); let mut has_tool_use = false; - for part in &parts { + for part in &rectified_parts { if part.get("thought").and_then(|value| value.as_bool()) == Some(true) { continue; } @@ -164,11 +186,15 @@ pub fn gemini_to_anthropic_with_shadow( session_id, candidate.get("content"), ) { + let mut shadow_content = content.clone(); + if let Some(parts_value) = shadow_content.get_mut("parts") { + *parts_value = json!(rectified_parts.clone()); + } store.record_assistant_turn( provider_id, session_id, - content.clone(), - extract_tool_call_meta(&parts), + shadow_content, + extract_tool_call_meta(&rectified_parts), ); } @@ -242,7 +268,7 @@ fn convert_messages_to_contents( shadow_turns: &[GeminiAssistantTurn], ) -> Result, ProxyError> { let mut contents = Vec::new(); - let mut tool_name_by_id = std::collections::HashMap::::new(); + let mut tool_name_by_id = build_tool_name_map_from_shadow_turns(shadow_turns); let total_assistant_messages = messages .iter() .filter(|message| message.get("role").and_then(|value| value.as_str()) == Some("assistant")) @@ -287,6 +313,10 @@ fn convert_messages_to_contents( convert_message_content_to_parts(message.get("content"), role, &mut tool_name_by_id)? }; + if role == "assistant" { + merge_tool_names_from_parts(&parts, &mut tool_name_by_id); + } + contents.push(json!({ "role": gemini_role, "parts": parts @@ -410,7 +440,11 @@ fn convert_message_content_to_parts( let name = tool_name_by_id .get(tool_use_id) .cloned() - .unwrap_or_default(); + .ok_or_else(|| { + ProxyError::TransformError(format!( + "Unable to resolve Gemini functionResponse.name for tool_use_id `{tool_use_id}`" + )) + })?; parts.push(json!({ "functionResponse": { @@ -457,15 +491,173 @@ fn shadow_parts(content: &Value) -> Option> { .or_else(|| content.as_array().cloned()) } +pub fn extract_anthropic_tool_schema_hints(body: &Value) -> AnthropicToolSchemaHints { + body.get("tools") + .and_then(|value| value.as_array()) + .into_iter() + .flatten() + .filter_map(|tool| { + let name = tool.get("name").and_then(|value| value.as_str())?; + let input_schema = tool + .get("input_schema") + .and_then(|value| value.as_object())?; + let properties = input_schema + .get("properties") + .and_then(|value| value.as_object())?; + if properties.is_empty() { + return None; + } + + let expected_keys = properties.keys().cloned().collect::>(); + let required_keys = input_schema + .get("required") + .and_then(|value| value.as_array()) + .map(|values| { + values + .iter() + .filter_map(|value| value.as_str().map(ToString::to_string)) + .collect::>() + }) + .unwrap_or_default(); + + Some(( + name.to_string(), + AnthropicToolSchemaHint { + expected_keys, + required_keys, + }, + )) + }) + .collect() +} + +pub fn rectify_tool_call_parts( + parts: &mut [Value], + tool_schema_hints: Option<&AnthropicToolSchemaHints>, +) { + for part in parts { + let Some(function_call) = part + .get_mut("functionCall") + .and_then(|value| value.as_object_mut()) + else { + continue; + }; + let Some(name) = function_call + .get("name") + .and_then(|value| value.as_str()) + .map(ToString::to_string) + else { + continue; + }; + let Some(args) = function_call.get_mut("args") else { + continue; + }; + + if rectify_tool_call_args(&name, args, tool_schema_hints) { + log::info!("[Claude/Gemini] Rectified tool args for `{name}`"); + } + } +} + +pub fn rectify_tool_call_args( + tool_name: &str, + args: &mut Value, + tool_schema_hints: Option<&AnthropicToolSchemaHints>, +) -> bool { + let Some(tool_schema_hints) = tool_schema_hints else { + return false; + }; + let Some(hint) = tool_schema_hints.get(tool_name) else { + return false; + }; + let Some(args_object) = args.as_object_mut() else { + return false; + }; + if args_object.is_empty() || hint.expected_keys.is_empty() { + return false; + } + + let expected_key_set = hint + .expected_keys + .iter() + .map(String::as_str) + .collect::>(); + let unexpected_keys = args_object + .keys() + .filter(|key| !expected_key_set.contains(key.as_str())) + .cloned() + .collect::>(); + if unexpected_keys.len() != 1 { + return false; + } + + let target_key = hint + .required_keys + .iter() + .find(|key| !args_object.contains_key(key.as_str())) + .cloned() + .or_else(|| { + if hint.expected_keys.len() == 1 && args_object.len() == 1 { + hint.expected_keys.first().cloned() + } else { + None + } + }); + let Some(target_key) = target_key else { + return false; + }; + if args_object.contains_key(&target_key) { + return false; + } + + let source_key = &unexpected_keys[0]; + let Some(value) = args_object.remove(source_key) else { + return false; + }; + args_object.insert(target_key, value); + true +} + fn merge_tool_names_from_shadow( turn: &GeminiAssistantTurn, - tool_name_by_id: &mut std::collections::HashMap, + tool_name_by_id: &mut HashMap, ) { for tool_call in &turn.tool_calls { if let Some(id) = &tool_call.id { tool_name_by_id.insert(id.clone(), tool_call.name.clone()); } } + + if let Some(parts) = shadow_parts(&turn.assistant_content) { + merge_tool_names_from_parts(&parts, tool_name_by_id); + } +} + +fn build_tool_name_map_from_shadow_turns( + shadow_turns: &[GeminiAssistantTurn], +) -> HashMap { + let mut tool_name_by_id = HashMap::new(); + for turn in shadow_turns { + merge_tool_names_from_shadow(turn, &mut tool_name_by_id); + } + tool_name_by_id +} + +fn merge_tool_names_from_parts(parts: &[Value], tool_name_by_id: &mut HashMap) { + for part in parts { + let Some(function_call) = part.get("functionCall") else { + continue; + }; + let Some(id) = function_call.get("id").and_then(|value| value.as_str()) else { + continue; + }; + let Some(name) = function_call.get("name").and_then(|value| value.as_str()) else { + continue; + }; + if !id.is_empty() && !name.is_empty() { + tool_name_by_id.insert(id.to_string(), name.to_string()); + } + } } fn extract_tool_call_meta(parts: &[Value]) -> Vec { @@ -676,6 +868,68 @@ mod tests { ); } + #[test] + fn anthropic_to_gemini_resolves_tool_result_name_from_shadow_content() { + let store = GeminiShadowStore::with_limits(8, 4); + store.record_assistant_turn( + "provider-a", + "session-1", + json!({ + "parts": [{ + "functionCall": { + "id": "call_1", + "name": "get_weather", + "args": { "city": "Tokyo" } + } + }] + }), + vec![], + ); + + let input = json!({ + "messages": [ + { + "role": "user", + "content": [ + { "type": "tool_result", "tool_use_id": "call_1", "content": "Sunny" } + ] + } + ] + }); + + let result = anthropic_to_gemini_with_shadow( + input, + Some(&store), + Some("provider-a"), + Some("session-1"), + ) + .unwrap(); + + assert_eq!( + result["contents"][0]["parts"][0]["functionResponse"]["name"], + "get_weather" + ); + } + + #[test] + fn anthropic_to_gemini_rejects_tool_result_without_resolvable_name() { + let input = json!({ + "messages": [ + { + "role": "user", + "content": [ + { "type": "tool_result", "tool_use_id": "call_1", "content": "Sunny" } + ] + } + ] + }); + + let error = anthropic_to_gemini(input).unwrap_err(); + assert!(error + .to_string() + .contains("Unable to resolve Gemini functionResponse.name")); + } + #[test] fn anthropic_to_gemini_uses_parameters_json_schema_for_rich_tool_schema() { let input = json!({ @@ -765,6 +1019,46 @@ mod tests { assert_eq!(result["stop_reason"], "tool_use"); } + #[test] + fn gemini_to_anthropic_rectifies_tool_args_from_schema_hints() { + let input = json!({ + "responseId": "resp_2", + "modelVersion": "gemini-2.5-pro", + "candidates": [{ + "finishReason": "STOP", + "content": { + "parts": [{ + "functionCall": { + "id": "call_1", + "name": "Skill", + "args": { "name": "git-commit" } + } + }] + } + }] + }); + let hints = extract_anthropic_tool_schema_hints(&json!({ + "tools": [{ + "name": "Skill", + "input_schema": { + "type": "object", + "properties": { + "skill": { "type": "string" }, + "args": { "type": "string" } + }, + "required": ["skill"] + } + }] + })); + + let result = + gemini_to_anthropic_with_shadow_and_hints(input, None, None, None, Some(&hints)) + .unwrap(); + + assert_eq!(result["content"][0]["input"]["skill"], "git-commit"); + assert!(result["content"][0]["input"].get("name").is_none()); + } + #[test] fn gemini_to_anthropic_maps_blocked_prompt_to_refusal() { let input = json!({