mirror of
https://github.com/farion1231/cc-switch.git
synced 2026-07-28 00:35:32 +08:00
feat(proxy): add Gemini Native tool argument rectification
This commit is contained in:
@@ -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<axum::response::Response, ProxyError> {
|
||||
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)
|
||||
|
||||
@@ -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::<String>()
|
||||
}
|
||||
|
||||
fn extract_tool_calls(parts: &[Value]) -> Vec<GeminiToolCallMeta> {
|
||||
parts
|
||||
fn extract_tool_calls(
|
||||
parts: &[Value],
|
||||
tool_schema_hints: Option<&AnthropicToolSchemaHints>,
|
||||
) -> Vec<GeminiToolCallMeta> {
|
||||
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<E: std::error::Error + Send + 'st
|
||||
shadow_store: Option<Arc<GeminiShadowStore>>,
|
||||
provider_id: Option<String>,
|
||||
session_id: Option<String>,
|
||||
tool_schema_hints: Option<AnthropicToolSchemaHints>,
|
||||
) -> impl Stream<Item = Result<Bytes, std::io::Error>> + Send {
|
||||
async_stream::stream! {
|
||||
let mut buffer = String::new();
|
||||
@@ -278,14 +286,16 @@ pub fn create_anthropic_sse_stream_from_gemini<E: std::error::Error + Send + 'st
|
||||
.and_then(|value| value.get("parts"))
|
||||
.and_then(|value| value.as_array())
|
||||
{
|
||||
let mut rectified_parts = parts.clone();
|
||||
rectify_tool_call_parts(&mut rectified_parts, tool_schema_hints.as_ref());
|
||||
if let Some(signature) = extract_text_thought_signature(parts) {
|
||||
text_thought_signature = Some(signature);
|
||||
}
|
||||
merge_tool_call_snapshots(
|
||||
&mut tool_call_snapshots,
|
||||
extract_tool_calls(parts),
|
||||
extract_tool_calls(&rectified_parts, tool_schema_hints.as_ref()),
|
||||
);
|
||||
let visible_text = extract_visible_text(parts);
|
||||
let visible_text = extract_visible_text(&rectified_parts);
|
||||
if !visible_text.is_empty() {
|
||||
let is_cumulative = visible_text.starts_with(&accumulated_text);
|
||||
let delta = if is_cumulative {
|
||||
@@ -491,7 +501,7 @@ mod tests {
|
||||
.into_iter()
|
||||
.map(|chunk| Ok::<Bytes, std::io::Error>(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::<Vec<_>>()
|
||||
@@ -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, std::io::Error>(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::<Vec<_>>()
|
||||
.await
|
||||
.into_iter()
|
||||
.map(|item| String::from_utf8(item.unwrap().to_vec()).unwrap())
|
||||
.collect::<Vec<_>>()
|
||||
.join("")
|
||||
});
|
||||
|
||||
assert!(output.contains("\"partial_json\":\"{\\\"command\\\":\\\"git status\\\"}\""));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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<String>,
|
||||
required_keys: Vec<String>,
|
||||
}
|
||||
|
||||
pub type AnthropicToolSchemaHints = HashMap<String, AnthropicToolSchemaHint>;
|
||||
|
||||
pub fn anthropic_to_gemini(body: Value) -> Result<Value, ProxyError> {
|
||||
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<Value, ProxyError> {
|
||||
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<Value, ProxyError> {
|
||||
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<Vec<Value>, ProxyError> {
|
||||
let mut contents = Vec::new();
|
||||
let mut tool_name_by_id = std::collections::HashMap::<String, String>::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<Vec<Value>> {
|
||||
.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::<Vec<_>>();
|
||||
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::<Vec<_>>()
|
||||
})
|
||||
.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::<HashSet<_>>();
|
||||
let unexpected_keys = args_object
|
||||
.keys()
|
||||
.filter(|key| !expected_key_set.contains(key.as_str()))
|
||||
.cloned()
|
||||
.collect::<Vec<_>>();
|
||||
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<String, String>,
|
||||
tool_name_by_id: &mut HashMap<String, String>,
|
||||
) {
|
||||
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<String, String> {
|
||||
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<String, String>) {
|
||||
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<GeminiToolCallMeta> {
|
||||
@@ -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!({
|
||||
|
||||
Reference in New Issue
Block a user