chore: openai reasoning model cleanup (#7529)

This commit is contained in:
Lifei Zhou
2026-03-03 02:36:52 +11:00
committed by GitHub
parent 1350542aae
commit 82facd4ebf
10 changed files with 142 additions and 52 deletions
+1
View File
@@ -429,6 +429,7 @@ mod tests {
toolshim_model: None,
fast_model_config: None,
request_params: None,
reasoning: None,
},
retry_config: RetryConfig::default(),
name: "aws_bedrock".to_string(),
-2
View File
@@ -46,10 +46,8 @@ pub const DATABRICKS_DEFAULT_MODEL: &str = "databricks-claude-sonnet-4";
const DATABRICKS_DEFAULT_FAST_MODEL: &str = "databricks-claude-haiku-4-5";
pub const DATABRICKS_KNOWN_MODELS: &[&str] = &[
"databricks-claude-sonnet-4-5",
"databricks-claude-3-7-sonnet",
"databricks-meta-llama-3-3-70b-instruct",
"databricks-meta-llama-3-1-405b-instruct",
"databricks-dbrx-instruct",
];
pub const DATABRICKS_DOC_URL: &str =
@@ -545,14 +545,10 @@ pub fn create_request(
}
let model_name = model_config.model_name.to_string();
let is_o1 = model_name.starts_with("o1") || model_name.starts_with("goose-o1");
let is_o3 = model_name.starts_with("o3") || model_name.starts_with("goose-o3");
let is_gpt_5 = model_name.starts_with("gpt-5") || model_name.starts_with("goose-gpt-5");
let is_openai_reasoning_model = is_o1 || is_o3 || is_gpt_5;
let is_openai_reasoning_model = model_config.is_openai_reasoning_model();
let is_claude_sonnet =
model_name.contains("claude-3-7-sonnet") || model_name.contains("claude-4-sonnet"); // can be goose- or databricks-
// Only extract reasoning effort for O1/O3 models
let (model_name, reasoning_effort) = if is_openai_reasoning_model {
let parts: Vec<&str> = model_config.model_name.split('-').collect();
let last_part = parts.last().unwrap();
@@ -568,7 +564,6 @@ pub fn create_request(
),
}
} else {
// For non-O family models, use the model name as is and no reasoning effort
(model_config.model_name.to_string(), None)
};
@@ -650,16 +645,10 @@ pub fn create_request(
}
}
// OpenAI reasoning models use max_completion_tokens instead of max_tokens
let key = if is_openai_reasoning_model {
"max_completion_tokens"
} else {
"max_tokens"
};
payload
.as_object_mut()
.unwrap()
.insert(key.to_string(), json!(model_config.max_output_tokens()));
payload.as_object_mut().unwrap().insert(
"max_completion_tokens".to_string(),
json!(model_config.max_output_tokens()),
);
}
// Apply cache control for Claude models to enable prompt caching
@@ -1056,6 +1045,7 @@ mod tests {
toolshim_model: None,
fast_model_config: None,
request_params: None,
reasoning: None,
};
let request = create_request(&model_config, "system", &[], &[], &ImageFormat::OpenAi)?;
let obj = request.as_object().unwrap();
@@ -1067,7 +1057,7 @@ mod tests {
"content": "system"
}
],
"max_tokens": 1024
"max_completion_tokens": 1024
});
for (key, value) in expected.as_object().unwrap() {
@@ -1088,6 +1078,7 @@ mod tests {
toolshim_model: None,
fast_model_config: None,
request_params: None,
reasoning: None,
};
let request = create_request(&model_config, "system", &[], &[], &ImageFormat::OpenAi)?;
assert_eq!(request["reasoning_effort"], "high");
@@ -1440,6 +1431,7 @@ mod tests {
toolshim_model: None,
fast_model_config: None,
request_params: None,
reasoning: None,
};
let messages = vec![
@@ -1492,6 +1484,7 @@ mod tests {
toolshim_model: None,
fast_model_config: None,
request_params: None,
reasoning: None,
};
let messages = vec![Message::user().with_text("Hello")];
+12 -21
View File
@@ -769,14 +769,9 @@ pub fn create_request(
));
}
let is_ox_model = model_config.model_name.starts_with("o1")
|| model_config.model_name.starts_with("o2")
|| model_config.model_name.starts_with("o3")
|| model_config.model_name.starts_with("o4")
|| model_config.model_name.starts_with("gpt-5");
let is_reasoning_model = model_config.is_openai_reasoning_model();
// Only extract reasoning effort for O-series models
let (model_name, reasoning_effort) = if is_ox_model {
let (model_name, reasoning_effort) = if is_reasoning_model {
let parts: Vec<&str> = model_config.model_name.split('-').collect();
let last_part = parts.last().unwrap();
@@ -791,12 +786,11 @@ pub fn create_request(
),
}
} else {
// For non-O family models, use the model name as is and no reasoning effort
(model_config.model_name.to_string(), None)
};
let system_message = json!({
"role": if is_ox_model { "developer" } else { "system" },
"role": if is_reasoning_model { "developer" } else { "system" },
"content": system
});
@@ -822,22 +816,16 @@ pub fn create_request(
}
// o1, o3 models currently don't support temperature
if !is_ox_model {
if !is_reasoning_model {
if let Some(temp) = model_config.temperature {
payload["temperature"] = json!(temp);
}
}
// o1/o3 models use max_completion_tokens instead of max_tokens
let key = if is_ox_model {
"max_completion_tokens"
} else {
"max_tokens"
};
payload
.as_object_mut()
.unwrap()
.insert(key.to_string(), json!(model_config.max_output_tokens()));
payload.as_object_mut().unwrap().insert(
"max_completion_tokens".to_string(),
json!(model_config.max_output_tokens()),
);
if for_streaming {
payload["stream"] = json!(true);
@@ -1500,6 +1488,7 @@ mod tests {
toolshim_model: None,
fast_model_config: None,
request_params: None,
reasoning: None,
};
let request = create_request(
&model_config,
@@ -1518,7 +1507,7 @@ mod tests {
"content": "system"
}
],
"max_tokens": 1024
"max_completion_tokens": 1024
});
for (key, value) in expected.as_object().unwrap() {
@@ -1540,6 +1529,7 @@ mod tests {
toolshim_model: None,
fast_model_config: None,
request_params: None,
reasoning: None,
};
let request = create_request(
&model_config,
@@ -1581,6 +1571,7 @@ mod tests {
toolshim_model: None,
fast_model_config: None,
request_params: None,
reasoning: None,
};
let request = create_request(
&model_config,
@@ -827,6 +827,7 @@ mod tests {
toolshim_model: None,
fast_model_config: None,
request_params: None,
reasoning: None,
};
let messages = vec![