fix(google): use includeThoughts/part.thought for thinking handling (#7593)

Signed-off-by: rabi <ramishra@redhat.com>
This commit is contained in:
Rabi Mishra
2026-03-11 23:55:23 +05:30
committed by GitHub
parent 7247d61648
commit d04b761a9c
+245 -133
View File
@@ -15,6 +15,8 @@ use serde_json::{json, Map, Value};
use std::ops::Deref; use std::ops::Deref;
pub const THOUGHT_SIGNATURE_KEY: &str = "thoughtSignature"; pub const THOUGHT_SIGNATURE_KEY: &str = "thoughtSignature";
const SYNTHETIC_THOUGHT_SIGNATURE: &str = "skip_thought_signature_validator";
const GEMINI25_DEFAULT_THINKING_BUDGET: i32 = 8192;
pub fn metadata_with_signature(signature: &str) -> ProviderMetadata { pub fn metadata_with_signature(signature: &str) -> ProviderMetadata {
let mut map = ProviderMetadata::new(); let mut map = ProviderMetadata::new();
@@ -29,6 +31,36 @@ pub fn get_thought_signature(metadata: &Option<ProviderMetadata>) -> Option<&str
.and_then(|v| v.as_str()) .and_then(|v| v.as_str())
} }
fn is_user_loop_boundary(message: &Message) -> bool {
message.role == Role::User
&& message
.content
.iter()
.any(|content| !matches!(content, MessageContent::ToolResponse(_)))
}
fn insert_thought_signature(part: &mut Map<String, Value>, signature: &str) {
part.insert(THOUGHT_SIGNATURE_KEY.to_string(), json!(signature));
}
fn maybe_insert_signature_from_metadata(
part: &mut Map<String, Value>,
metadata: &Option<ProviderMetadata>,
) {
if let Some(signature) = get_thought_signature(metadata) {
insert_thought_signature(part, signature);
}
}
fn build_function_response_part(name: &str, text: String) -> Map<String, Value> {
let mut part = Map::new();
let mut function_response = Map::new();
function_response.insert("name".to_string(), json!(name));
function_response.insert("response".to_string(), json!({"content": {"text": text}}));
part.insert("functionResponse".to_string(), json!(function_response));
part
}
/// Convert internal Message format to Google's API message specification /// Convert internal Message format to Google's API message specification
pub fn format_messages(messages: &[Message]) -> Vec<Value> { pub fn format_messages(messages: &[Message]) -> Vec<Value> {
let filtered: Vec<_> = messages let filtered: Vec<_> = messages
@@ -44,26 +76,27 @@ pub fn format_messages(messages: &[Message]) -> Vec<Value> {
}) })
.collect(); .collect();
let last_assistant_idx = filtered let active_loop_start_idx = filtered
.iter() .iter()
.enumerate() .enumerate()
.filter(|(_, m)| m.role != Role::User) .rev()
.map(|(i, _)| i) .find(|(_, m)| is_user_loop_boundary(m))
.next_back(); .map(|(i, _)| i);
filtered filtered
.iter() .iter()
.enumerate() .enumerate()
.map(|(idx, message)| { .filter_map(|(idx, message)| {
let role = if message.role == Role::User { let role = if message.role == Role::User {
"user" "user"
} else { } else {
"model" "model"
}; };
let include_signature = match last_assistant_idx { let include_signature = active_loop_start_idx.is_none_or(|start_idx| idx >= start_idx);
Some(last_idx) => idx >= last_idx, // Only the first model tool call in a turn is guaranteed to carry
None => false, // a signature for loop continuity.
}; let mut needs_synthetic_for_first_model_tool_call =
include_signature && message.role != Role::User;
let mut parts = Vec::new(); let mut parts = Vec::new();
for message_content in message.content.iter() { for message_content in message.content.iter() {
match message_content { match message_content {
@@ -92,12 +125,15 @@ pub fn format_messages(messages: &[Message]) -> Vec<Value> {
if include_signature { if include_signature {
if let Some(signature) = get_thought_signature(&request.metadata) { if let Some(signature) = get_thought_signature(&request.metadata) {
part.insert( insert_thought_signature(&mut part, signature);
THOUGHT_SIGNATURE_KEY.to_string(), } else if needs_synthetic_for_first_model_tool_call {
json!(signature), insert_thought_signature(
&mut part,
SYNTHETIC_THOUGHT_SIGNATURE,
); );
} }
} }
needs_synthetic_for_first_model_tool_call = false;
parts.push(json!(part)); parts.push(json!(part));
} }
@@ -138,50 +174,22 @@ pub fn format_messages(messages: &[Message]) -> Vec<Value> {
if text.is_empty() { if text.is_empty() {
text = "Tool call is done.".to_string(); text = "Tool call is done.".to_string();
} }
let mut part = Map::new(); let mut part = build_function_response_part(&response.id, text);
let mut function_response = Map::new();
function_response.insert("name".to_string(), json!(response.id));
function_response
.insert("response".to_string(), json!({"content": {"text": text}}));
part.insert("functionResponse".to_string(), json!(function_response));
if include_signature { if include_signature {
if let Some(signature) = get_thought_signature(&response.metadata) { maybe_insert_signature_from_metadata(&mut part, &response.metadata);
part.insert(
THOUGHT_SIGNATURE_KEY.to_string(),
json!(signature),
);
}
} }
parts.push(json!(part)); parts.push(json!(part));
} }
Err(e) => { Err(e) => {
let mut part = Map::new(); let mut part =
let mut function_response = Map::new(); build_function_response_part(&response.id, format!("Error: {}", e));
function_response.insert("name".to_string(), json!(response.id));
function_response.insert(
"response".to_string(),
json!({"content": {"text": format!("Error: {}", e)}}),
);
part.insert("functionResponse".to_string(), json!(function_response));
if include_signature { if include_signature {
if let Some(signature) = get_thought_signature(&response.metadata) { maybe_insert_signature_from_metadata(&mut part, &response.metadata);
part.insert(
THOUGHT_SIGNATURE_KEY.to_string(),
json!(signature),
);
}
} }
parts.push(json!(part)); parts.push(json!(part));
} }
}, },
MessageContent::Thinking(thinking) => { MessageContent::Thinking(_) => {}
let mut part = Map::new();
part.insert("text".to_string(), json!(thinking.thinking));
if include_signature {
part.insert("thoughtSignature".to_string(), json!(thinking.signature));
}
parts.push(json!(part));
}
MessageContent::Image(image) => { MessageContent::Image(image) => {
parts.push(json!({ parts.push(json!({
"inline_data": { "inline_data": {
@@ -194,7 +202,11 @@ pub fn format_messages(messages: &[Message]) -> Vec<Value> {
_ => {} _ => {}
} }
} }
json!({"role": role, "parts": parts}) if parts.is_empty() {
None
} else {
Some(json!({"role": role, "parts": parts}))
}
}) })
.collect() .collect()
} }
@@ -221,32 +233,15 @@ pub fn format_tools(tools: &[Tool]) -> Vec<Value> {
.collect() .collect()
} }
#[derive(Clone, Copy)]
enum SignedTextHandling {
SignedTextAsThinking,
SignedTextAsRegularText,
}
fn process_response_part_non_streaming(
part: &Value,
last_signature: &mut Option<String>,
has_function_calls: bool,
) -> Option<MessageContent> {
// For non-streaming: signed text is thinking only if there are function calls
let handling = if has_function_calls {
SignedTextHandling::SignedTextAsThinking
} else {
SignedTextHandling::SignedTextAsRegularText
};
process_response_part_impl(part, last_signature, handling)
}
fn process_response_part_impl( fn process_response_part_impl(
part: &Value, part: &Value,
last_signature: &mut Option<String>, last_signature: &mut Option<String>,
signed_text_handling: SignedTextHandling,
) -> Option<MessageContent> { ) -> Option<MessageContent> {
let signature = part.get(THOUGHT_SIGNATURE_KEY).and_then(|v| v.as_str()); let signature = part.get(THOUGHT_SIGNATURE_KEY).and_then(|v| v.as_str());
let is_thought = part
.get("thought")
.and_then(|value| value.as_bool())
.unwrap_or(false);
if let Some(sig) = signature { if let Some(sig) = signature {
*last_signature = Some(sig.to_string()); *last_signature = Some(sig.to_string());
@@ -257,11 +252,13 @@ fn process_response_part_impl(
if text.is_empty() { if text.is_empty() {
return None; return None;
} }
match (signature, signed_text_handling) { if is_thought {
(Some(sig), SignedTextHandling::SignedTextAsThinking) => { match signature {
Some(MessageContent::thinking(text.to_string(), sig.to_string())) Some(sig) => Some(MessageContent::thinking(text.to_string(), sig.to_string())),
None => Some(MessageContent::reasoning(text.to_string())),
} }
_ => Some(MessageContent::text(text.to_string())), } else {
Some(MessageContent::text(text.to_string()))
} }
} else if text_value.is_some() { } else if text_value.is_some() {
tracing::warn!( tracing::warn!(
@@ -323,15 +320,11 @@ pub fn response_to_message(response: Value) -> Result<Message> {
return Ok(Message::new(role, created, Vec::new())); return Ok(Message::new(role, created, Vec::new()));
}; };
let has_function_calls = parts.iter().any(|p| p.get("functionCall").is_some());
let mut content = Vec::new(); let mut content = Vec::new();
let mut last_signature: Option<String> = None; let mut last_signature: Option<String> = None;
for part in parts { for part in parts {
if let Some(msg_content) = if let Some(msg_content) = process_response_part_impl(part, &mut last_signature) {
process_response_part_non_streaming(part, &mut last_signature, has_function_calls)
{
content.push(msg_content); content.push(msg_content);
} }
} }
@@ -467,9 +460,7 @@ where
if let Some(parts) = parts { if let Some(parts) = parts {
for part in parts { for part in parts {
// Always emit text as regular text during streaming — we can't if let Some(content) = process_response_part_impl(part, &mut last_signature) {
// know yet whether function calls will follow.
if let Some(content) = process_response_part_impl(part, &mut last_signature, SignedTextHandling::SignedTextAsRegularText) {
let message = Message::new( let message = Message::new(
Role::Assistant, Role::Assistant,
chrono::Utc::now().timestamp(), chrono::Utc::now().timestamp(),
@@ -524,7 +515,11 @@ enum ThinkingLevel {
#[derive(Serialize)] #[derive(Serialize)]
#[serde(rename_all = "camelCase")] #[serde(rename_all = "camelCase")]
struct ThinkingConfig { struct ThinkingConfig {
thinking_level: ThinkingLevel, #[serde(skip_serializing_if = "Option::is_none")]
thinking_level: Option<ThinkingLevel>,
#[serde(skip_serializing_if = "Option::is_none")]
thinking_budget: Option<i32>,
include_thoughts: bool,
} }
#[derive(Serialize)] #[derive(Serialize)]
@@ -539,33 +534,59 @@ struct GoogleRequest<'a> {
} }
fn get_thinking_config(model_config: &ModelConfig) -> Option<ThinkingConfig> { fn get_thinking_config(model_config: &ModelConfig) -> Option<ThinkingConfig> {
if !model_config let model_name = model_config.model_name.to_lowercase();
.model_name let is_gemini_3 = model_name.starts_with("gemini-3");
.to_lowercase() let is_gemini_25 = model_name.starts_with("gemini-2.5");
.starts_with("gemini-3") if !is_gemini_3 && !is_gemini_25 {
{
return None; return None;
} }
let thinking_level_str = model_config if is_gemini_3 {
.get_config_param::<String>("thinking_level", "GEMINI3_THINKING_LEVEL") let thinking_level_str = model_config
.map(|s| s.to_lowercase()) .get_config_param::<String>("thinking_level", "GEMINI3_THINKING_LEVEL")
.unwrap_or_else(|| "low".to_string()); .map(|s| s.to_lowercase())
.unwrap_or_else(|| "low".to_string());
let thinking_level = match thinking_level_str.as_str() { let thinking_level = match thinking_level_str.as_str() {
"high" => ThinkingLevel::High, "high" => ThinkingLevel::High,
"low" => ThinkingLevel::Low, "low" => ThinkingLevel::Low,
invalid => { invalid => {
tracing::warn!( tracing::warn!(
"Invalid thinking level '{}' for model '{}'. Valid levels: low, high. Using 'low'.", "Invalid thinking level '{}' for model '{}'. Valid levels: low, high. Using 'low'.",
invalid, invalid,
model_config.model_name, model_config.model_name,
); );
ThinkingLevel::Low ThinkingLevel::Low
} }
}; };
Some(ThinkingConfig { thinking_level }) Some(ThinkingConfig {
thinking_level: Some(thinking_level),
thinking_budget: None,
include_thoughts: true,
})
} else {
let thinking_budget = match model_config
.get_config_param::<i32>("thinking_budget", "GEMINI25_THINKING_BUDGET")
{
Some(budget) if budget >= 0 => budget,
Some(budget) => {
tracing::warn!(
"Invalid thinking budget '{}' for model '{}'. Must be >= 0. Using '{}'.",
budget,
model_config.model_name,
GEMINI25_DEFAULT_THINKING_BUDGET,
);
GEMINI25_DEFAULT_THINKING_BUDGET
}
None => GEMINI25_DEFAULT_THINKING_BUDGET,
};
Some(ThinkingConfig {
thinking_level: None,
thinking_budget: Some(thinking_budget),
include_thoughts: true,
})
}
} }
pub fn create_request( pub fn create_request(
@@ -609,6 +630,7 @@ mod tests {
use rmcp::model::{CallToolRequestParams, CallToolResult}; use rmcp::model::{CallToolRequestParams, CallToolResult};
use rmcp::{model::Content, object}; use rmcp::{model::Content, object};
use serde_json::json; use serde_json::json;
use std::collections::HashMap;
fn set_up_text_message(text: &str, role: Role) -> Message { fn set_up_text_message(text: &str, role: Role) -> Message {
Message::new(role, 0, vec![MessageContent::text(text.to_string())]) Message::new(role, 0, vec![MessageContent::text(text.to_string())])
@@ -964,7 +986,7 @@ mod tests {
const SIG: &str = "thought_sig_abc"; const SIG: &str = "thought_sig_abc";
let response_with_tools = google_response(vec![ let response_with_tools = google_response(vec![
json!({"text": "Let me think...", "thoughtSignature": SIG}), json!({"text": "Let me think...", "thought": true, "thoughtSignature": SIG}),
json!({"functionCall": {"name": "shell", "args": {"cmd": "ls"}}, "thoughtSignature": SIG}), json!({"functionCall": {"name": "shell", "args": {"cmd": "ls"}}, "thoughtSignature": SIG}),
json!({"functionCall": {"name": "read", "args": {}}}), json!({"functionCall": {"name": "read", "args": {}}}),
]); ]);
@@ -995,23 +1017,22 @@ mod tests {
Ok(tool_result("output")), Ok(tool_result("output")),
req1.metadata.as_ref(), req1.metadata.as_ref(),
); );
let google_out = format_messages(&[native.clone(), tool_response.clone()]); let user_prompt = set_up_text_message("List files", Role::User);
assert_eq!(google_out[0]["parts"][0]["thoughtSignature"], SIG); let google_out =
format_messages(&[user_prompt.clone(), native.clone(), tool_response.clone()]);
assert_eq!(google_out[1]["parts"][0]["thoughtSignature"], SIG); assert_eq!(google_out[1]["parts"][0]["thoughtSignature"], SIG);
assert_eq!(google_out[2]["parts"][0]["thoughtSignature"], SIG);
let second_assistant = let second_assistant = response_to_message(google_response(vec![json!({
Message::assistant().with_thinking("More thinking".to_string(), "sig_456".to_string()); "functionCall": {"name": "echo", "args": {}},
let google_multi = format_messages(&[native, tool_response, second_assistant]); "thoughtSignature": "sig_456"
assert!(google_multi[0]["parts"][0] })]))
.get("thoughtSignature") .unwrap();
.is_none()); let google_multi = format_messages(&[user_prompt, native, tool_response, second_assistant]);
assert!(google_multi[1]["parts"][0] assert_eq!(google_multi[1]["parts"][0]["thoughtSignature"], SIG);
.get("thoughtSignature") assert_eq!(google_multi[2]["parts"][0]["thoughtSignature"], SIG);
.is_none()); assert_eq!(google_multi[3]["parts"][0]["thoughtSignature"], "sig_456");
assert_eq!(google_multi[2]["parts"][0]["thoughtSignature"], "sig_456");
// Text-only response WITH signature but WITHOUT function calls should be regular text
// (per original behavior: thinking is only when reasoning before tool calls)
let final_response_with_sig = let final_response_with_sig =
google_response(vec![json!({"text": "Done!", "thoughtSignature": SIG})]); google_response(vec![json!({"text": "Done!", "thoughtSignature": SIG})]);
let final_native_with_sig = response_to_message(final_response_with_sig).unwrap(); let final_native_with_sig = response_to_message(final_response_with_sig).unwrap();
@@ -1028,6 +1049,49 @@ mod tests {
); );
} }
#[test]
fn test_thought_without_signature_maps_to_reasoning() {
let response = google_response(vec![json!({
"text": "Working through options...",
"thought": true
})]);
let native = response_to_message(response).unwrap();
assert_eq!(native.content.len(), 1);
assert!(native.content[0].as_reasoning().is_some());
}
#[test]
fn test_format_messages_omits_messages_with_empty_parts() {
let user_prompt = set_up_text_message("hello", Role::User);
let thinking_only =
Message::assistant().with_thinking("internal".to_string(), "sig_123".to_string());
let reasoning_only = response_to_message(google_response(vec![json!({
"text": "deliberating",
"thought": true
})]))
.unwrap();
let formatted = format_messages(&[user_prompt, thinking_only, reasoning_only]);
assert_eq!(formatted.len(), 1);
assert_eq!(formatted[0]["role"], "user");
assert_eq!(formatted[0]["parts"][0]["text"], "hello");
}
#[test]
fn test_active_loop_injects_synthetic_signature_for_first_model_tool_call() {
let user_prompt = set_up_text_message("Find a restaurant", Role::User);
let assistant_tool = response_to_message(google_response(vec![json!({
"functionCall": {"name": "find_restaurant", "args": {"cuisine": "italian"}}
})]))
.unwrap();
let formatted = format_messages(&[user_prompt, assistant_tool]);
assert_eq!(
formatted[1]["parts"][0][THOUGHT_SIGNATURE_KEY],
SYNTHETIC_THOUGHT_SIGNATURE
);
}
const GOOGLE_TEXT_STREAM: &str = concat!( const GOOGLE_TEXT_STREAM: &str = concat!(
r#"data: {"candidates": [{"content": {"role": "model", "#, r#"data: {"candidates": [{"content": {"role": "model", "#,
r#""parts": [{"text": "Hello"}]}}]}"#, r#""parts": [{"text": "Hello"}]}}]}"#,
@@ -1082,7 +1146,6 @@ mod tests {
assert_eq!(usage.usage.input_tokens, Some(10)); assert_eq!(usage.usage.input_tokens, Some(10));
assert_eq!(usage.usage.output_tokens, Some(3)); assert_eq!(usage.usage.output_tokens, Some(3));
// Verify all streaming messages have consistent IDs for UI aggregation
assert!( assert!(
message_ids.iter().all(|id| id.is_some()), message_ids.iter().all(|id| id.is_some()),
"All streaming messages should have an ID" "All streaming messages should have an ID"
@@ -1125,29 +1188,30 @@ mod tests {
async fn test_streaming_with_thought_signature() { async fn test_streaming_with_thought_signature() {
use futures::StreamExt; use futures::StreamExt;
async fn collect_streaming_text(raw: &str) -> (String, usize) { async fn collect_streaming_text(raw: &str) -> (String, usize, usize) {
let lines: Vec<Result<String, anyhow::Error>> = let lines: Vec<Result<String, anyhow::Error>> =
raw.lines().map(|l| Ok(l.to_string())).collect(); raw.lines().map(|l| Ok(l.to_string())).collect();
let stream = Box::pin(futures::stream::iter(lines)); let stream = Box::pin(futures::stream::iter(lines));
let mut msg_stream = std::pin::pin!(response_to_streaming_message(stream)); let mut msg_stream = std::pin::pin!(response_to_streaming_message(stream));
let mut text = String::new(); let mut text = String::new();
let mut thinking = 0usize; let mut thinking = 0usize;
let mut reasoning = 0usize;
while let Some(Ok((message, _))) = msg_stream.next().await { while let Some(Ok((message, _))) = msg_stream.next().await {
if let Some(msg) = message { if let Some(msg) = message {
for c in &msg.content { for c in &msg.content {
match c { match c {
MessageContent::Text(t) => text.push_str(&t.text), MessageContent::Text(t) => text.push_str(&t.text),
MessageContent::Thinking(_) => thinking += 1, MessageContent::Thinking(_) => thinking += 1,
MessageContent::Reasoning(_) => reasoning += 1,
_ => {} _ => {}
} }
} }
} }
} }
(text, thinking) (text, thinking, reasoning)
} }
// First chunk signed let (text, thinking, reasoning) = collect_streaming_text(concat!(
let (text, thinking) = collect_streaming_text(concat!(
r#"data: {"candidates": [{"content": {"role": "model", "#, r#"data: {"candidates": [{"content": {"role": "model", "#,
r#""parts": [{"text": "Hello", "thoughtSignature": "sig1"}]}}], "#, r#""parts": [{"text": "Hello", "thoughtSignature": "sig1"}]}}], "#,
r#""modelVersion": "gemini-3-flash-preview"}"#, r#""modelVersion": "gemini-3-flash-preview"}"#,
@@ -1157,10 +1221,10 @@ mod tests {
)) ))
.await; .await;
assert_eq!(thinking, 0); assert_eq!(thinking, 0);
assert_eq!(reasoning, 0);
assert_eq!(text, "Hello world"); assert_eq!(text, "Hello world");
// Last chunk signed (the reported truncation bug) let (text, thinking, reasoning) = collect_streaming_text(concat!(
let (text, thinking) = collect_streaming_text(concat!(
r#"data: {"candidates": [{"content": {"role": "model", "#, r#"data: {"candidates": [{"content": {"role": "model", "#,
r#""parts": [{"text": "SECURITY.md: Project"}]}}], "#, r#""parts": [{"text": "SECURITY.md: Project"}]}}], "#,
r#""modelVersion": "gemini-3-flash-preview"}"#, r#""modelVersion": "gemini-3-flash-preview"}"#,
@@ -1171,10 +1235,10 @@ mod tests {
)) ))
.await; .await;
assert_eq!(thinking, 0); assert_eq!(thinking, 0);
assert_eq!(reasoning, 0);
assert_eq!(text, "SECURITY.md: Project policies.\n\nRead it?"); assert_eq!(text, "SECURITY.md: Project policies.\n\nRead it?");
// Intermediate chunk signed let (text, thinking, reasoning) = collect_streaming_text(concat!(
let (text, thinking) = collect_streaming_text(concat!(
r#"data: {"candidates": [{"content": {"role": "model", "#, r#"data: {"candidates": [{"content": {"role": "model", "#,
r#""parts": [{"text": "one "}]}}], "modelVersion": "gemini-3-flash-preview"}"#, r#""parts": [{"text": "one "}]}}], "modelVersion": "gemini-3-flash-preview"}"#,
"\n", "\n",
@@ -1186,7 +1250,20 @@ mod tests {
)) ))
.await; .await;
assert_eq!(thinking, 0); assert_eq!(thinking, 0);
assert_eq!(reasoning, 0);
assert_eq!(text, "one two three"); assert_eq!(text, "one two three");
let (text, thinking, reasoning) = collect_streaming_text(concat!(
r#"data: {"candidates": [{"content": {"role": "model", "#,
r#""parts": [{"text": "internal chain", "thought": true, "thoughtSignature": "sig4"}]}}]}"#,
"\n",
r#"data: {"candidates": [{"content": {"role": "model", "#,
r#""parts": [{"text": "visible"}]}}]}"#
))
.await;
assert_eq!(thinking, 1);
assert_eq!(reasoning, 0);
assert_eq!(text, "visible");
} }
#[tokio::test] #[tokio::test]
@@ -1215,7 +1292,6 @@ mod tests {
async fn test_streaming_with_sse_event_lines() { async fn test_streaming_with_sse_event_lines() {
use futures::StreamExt; use futures::StreamExt;
// SSE format can include event: lines which should be skipped
let sse_stream = r#"event: message let sse_stream = r#"event: message
data: {"candidates": [{"content": {"role": "model", "parts": [{"text": "Hello"}]}}]} data: {"candidates": [{"content": {"role": "model", "parts": [{"text": "Hello"}]}}]}
@@ -1272,7 +1348,6 @@ data: [DONE]"#;
} }
} }
// Only "Complete" should be captured, stream should stop at [DONE]
assert_eq!(text_parts, vec!["Complete"]); assert_eq!(text_parts, vec!["Complete"]);
} }
@@ -1308,19 +1383,56 @@ data: [DONE]"#;
fn test_get_thinking_config() { fn test_get_thinking_config() {
use crate::model::ModelConfig; use crate::model::ModelConfig;
// Test 1: Gemini 3 model defaults to low thinking level
let config = ModelConfig::new("gemini-3-pro").unwrap(); let config = ModelConfig::new("gemini-3-pro").unwrap();
let result = get_thinking_config(&config); let result = get_thinking_config(&config);
assert!(result.is_some()); assert!(result.is_some());
let thinking_config = result.unwrap(); let thinking_config = result.unwrap();
assert!(matches!(thinking_config.thinking_level, ThinkingLevel::Low)); assert!(thinking_config.thinking_level.is_some());
assert!(thinking_config.thinking_budget.is_none());
assert!(thinking_config.include_thoughts);
// Test 2: Case-insensitive model detection
let config = ModelConfig::new("Gemini-3-Flash").unwrap(); let config = ModelConfig::new("Gemini-3-Flash").unwrap();
let result = get_thinking_config(&config); let result = get_thinking_config(&config);
assert!(result.is_some()); assert!(result.is_some());
// Test 3: Non-Gemini 3 model returns None let config = ModelConfig::new("gemini-2.5-flash").unwrap();
let result = get_thinking_config(&config);
assert!(result.is_some());
let thinking_config = result.unwrap();
assert!(thinking_config.include_thoughts);
assert!(thinking_config.thinking_level.is_none());
assert_eq!(
thinking_config.thinking_budget,
Some(GEMINI25_DEFAULT_THINKING_BUDGET)
);
let mut params = HashMap::new();
params.insert("thinking_budget".to_string(), json!(4096));
let config = ModelConfig::new("gemini-2.5-flash")
.unwrap()
.with_request_params(Some(params));
let result = get_thinking_config(&config);
assert!(result.is_some());
let thinking_config = result.unwrap();
assert_eq!(thinking_config.thinking_budget, Some(4096));
let mut params = HashMap::new();
params.insert("thinking_budget".to_string(), json!(-1));
let config = ModelConfig::new("gemini-2.5-flash")
.unwrap()
.with_request_params(Some(params));
let result = get_thinking_config(&config);
assert!(result.is_some());
let thinking_config = result.unwrap();
assert_eq!(
thinking_config.thinking_budget,
Some(GEMINI25_DEFAULT_THINKING_BUDGET)
);
let config = ModelConfig::new("gemini-2.0-flash").unwrap();
let result = get_thinking_config(&config);
assert!(result.is_none());
let config = ModelConfig::new("gpt-4o").unwrap(); let config = ModelConfig::new("gpt-4o").unwrap();
let result = get_thinking_config(&config); let result = get_thinking_config(&config);
assert!(result.is_none()); assert!(result.is_none());