fix(providers): drop stale signed thinking blocks after a mid-conversation model switch (#10007)
Signed-off-by: Kyle De Freitas <kdefreitas@squareup.com> Co-authored-by: Douwe M Osinga <douwe@sidewalklabs.com>
This commit is contained in:
@@ -46,11 +46,12 @@ macro_rules! string_enum {
|
||||
|
||||
string_enum!(ThinkingType { Adaptive => "adaptive", Enabled => "enabled", Disabled => "disabled" });
|
||||
|
||||
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
|
||||
#[derive(Debug, Clone, Default, PartialEq, Eq)]
|
||||
pub struct AnthropicFormatOptions {
|
||||
pub preserve_unsigned_thinking: bool,
|
||||
pub preserve_thinking_context: bool,
|
||||
pub thinking_disabled: bool,
|
||||
pub current_model: Option<String>,
|
||||
}
|
||||
|
||||
impl AnthropicFormatOptions {
|
||||
@@ -69,10 +70,28 @@ impl AnthropicFormatOptions {
|
||||
preserve_unsigned_thinking,
|
||||
preserve_thinking_context,
|
||||
thinking_disabled,
|
||||
current_model: self
|
||||
.current_model
|
||||
.or_else(|| Some(model_config.model_name.clone())),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub fn thinking_block_is_stale(message: &Message, current_model: Option<&str>) -> bool {
|
||||
let Some(current_model) = current_model else {
|
||||
return false;
|
||||
};
|
||||
let Some(inference) = message.metadata.inference.as_ref() else {
|
||||
return false;
|
||||
};
|
||||
let requested = inference.requested_model.as_str();
|
||||
let resolved = inference.resolved_model.as_deref().unwrap_or("");
|
||||
if requested.is_empty() && resolved.is_empty() {
|
||||
return false;
|
||||
}
|
||||
current_model != requested && current_model != resolved
|
||||
}
|
||||
|
||||
fn canonical_thinking_mode(provider_name: &str, model_name: &str) -> Option<ThinkingMode> {
|
||||
maybe_get_canonical_model(provider_name, model_name).and_then(|model| model.thinking_mode)
|
||||
}
|
||||
@@ -177,12 +196,12 @@ fn args_to_input_value(arguments: Option<JsonObject>) -> Value {
|
||||
|
||||
/// Convert internal Message format to Anthropic's API message specification
|
||||
pub fn format_messages(messages: &[Message]) -> Vec<Value> {
|
||||
format_messages_with_options(messages, AnthropicFormatOptions::default())
|
||||
format_messages_with_options(messages, &AnthropicFormatOptions::default())
|
||||
}
|
||||
|
||||
fn format_messages_with_options(
|
||||
messages: &[Message],
|
||||
options: AnthropicFormatOptions,
|
||||
options: &AnthropicFormatOptions,
|
||||
) -> Vec<Value> {
|
||||
let mut anthropic_messages = Vec::new();
|
||||
|
||||
@@ -192,6 +211,8 @@ fn format_messages_with_options(
|
||||
Role::Assistant => ASSISTANT_ROLE,
|
||||
};
|
||||
|
||||
let thinking_is_stale = thinking_block_is_stale(message, options.current_model.as_deref());
|
||||
|
||||
let mut content = Vec::new();
|
||||
for msg_content in &message.content {
|
||||
match msg_content {
|
||||
@@ -346,11 +367,13 @@ fn format_messages_with_options(
|
||||
// Anthropic rejects thinking blocks sent without a matching thinking config.
|
||||
if !options.thinking_disabled {
|
||||
if !thinking.signature.is_empty() {
|
||||
content.push(json!({
|
||||
TYPE_FIELD: THINKING_TYPE,
|
||||
THINKING_TYPE: thinking.thinking,
|
||||
SIGNATURE_FIELD: thinking.signature
|
||||
}));
|
||||
if !thinking_is_stale {
|
||||
content.push(json!({
|
||||
TYPE_FIELD: THINKING_TYPE,
|
||||
THINKING_TYPE: thinking.thinking,
|
||||
SIGNATURE_FIELD: thinking.signature
|
||||
}));
|
||||
}
|
||||
} else if options.preserve_unsigned_thinking
|
||||
&& !thinking.thinking.is_empty()
|
||||
{
|
||||
@@ -362,7 +385,7 @@ fn format_messages_with_options(
|
||||
}
|
||||
}
|
||||
MessageContentBlock::RedactedThinking(redacted) => {
|
||||
if !options.thinking_disabled {
|
||||
if !options.thinking_disabled && !thinking_is_stale {
|
||||
content.push(json!({
|
||||
TYPE_FIELD: REDACTED_THINKING_TYPE,
|
||||
DATA_FIELD: redacted.data
|
||||
@@ -741,7 +764,7 @@ pub fn create_request_for_model(
|
||||
options: AnthropicFormatOptions,
|
||||
) -> Result<Value> {
|
||||
let options = options.for_model(model_config);
|
||||
let anthropic_messages = format_messages_with_options(messages, options);
|
||||
let anthropic_messages = format_messages_with_options(messages, &options);
|
||||
let tool_specs = format_tools(tools);
|
||||
let system_spec = format_system(system);
|
||||
|
||||
@@ -1120,7 +1143,7 @@ where
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::conversation::message::Message;
|
||||
use crate::conversation::message::{Message, MessageContent};
|
||||
use crate::model::ModelConfig;
|
||||
use rmcp::object;
|
||||
use serde_json::json;
|
||||
@@ -1316,10 +1339,11 @@ mod tests {
|
||||
|
||||
let spec = format_messages_with_options(
|
||||
&messages,
|
||||
AnthropicFormatOptions {
|
||||
&AnthropicFormatOptions {
|
||||
preserve_unsigned_thinking: true,
|
||||
preserve_thinking_context: false,
|
||||
thinking_disabled: false,
|
||||
current_model: None,
|
||||
},
|
||||
);
|
||||
|
||||
@@ -1331,6 +1355,64 @@ mod tests {
|
||||
assert_eq!(spec[1]["content"][0]["text"], "Hi there");
|
||||
}
|
||||
|
||||
fn signed_thinking_from_model(model: &str) -> Message {
|
||||
use crate::conversation::message::InferenceMetadata;
|
||||
Message::assistant()
|
||||
.with_content(MessageContent::thinking("internal", "sig-abc"))
|
||||
.with_text("answer")
|
||||
.with_inference(InferenceMetadata {
|
||||
provider: "anthropic".to_string(),
|
||||
requested_model: model.to_string(),
|
||||
resolved_model: None,
|
||||
provider_session_id: None,
|
||||
})
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn drops_signed_thinking_from_a_different_model() {
|
||||
let messages = vec![signed_thinking_from_model("claude-opus-4-1")];
|
||||
let opts = AnthropicFormatOptions {
|
||||
current_model: Some("claude-sonnet-4-5".to_string()),
|
||||
..Default::default()
|
||||
};
|
||||
let spec = format_messages_with_options(&messages, &opts);
|
||||
let types: Vec<&str> = spec[0]["content"]
|
||||
.as_array()
|
||||
.unwrap()
|
||||
.iter()
|
||||
.map(|c| c["type"].as_str().unwrap())
|
||||
.collect();
|
||||
assert!(
|
||||
!types.contains(&"thinking"),
|
||||
"stale thinking must be dropped"
|
||||
);
|
||||
assert!(types.contains(&"text"), "text content must be preserved");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn keeps_signed_thinking_from_the_same_model() {
|
||||
let messages = vec![signed_thinking_from_model("claude-sonnet-4-5")];
|
||||
let opts = AnthropicFormatOptions {
|
||||
current_model: Some("claude-sonnet-4-5".to_string()),
|
||||
..Default::default()
|
||||
};
|
||||
let spec = format_messages_with_options(&messages, &opts);
|
||||
assert_eq!(spec[0]["content"][0]["type"], "thinking");
|
||||
assert_eq!(spec[0]["content"][0]["signature"], "sig-abc");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn keeps_signed_thinking_when_provenance_unknown() {
|
||||
let messages =
|
||||
vec![Message::assistant().with_content(MessageContent::thinking("internal", "sig"))];
|
||||
let opts = AnthropicFormatOptions {
|
||||
current_model: Some("claude-sonnet-4-5".to_string()),
|
||||
..Default::default()
|
||||
};
|
||||
let spec = format_messages_with_options(&messages, &opts);
|
||||
assert_eq!(spec[0]["content"][0]["type"], "thinking");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_tools_to_anthropic_spec() {
|
||||
let tools = vec![
|
||||
@@ -1554,6 +1636,7 @@ mod tests {
|
||||
preserve_unsigned_thinking: true,
|
||||
preserve_thinking_context: true,
|
||||
thinking_disabled: false,
|
||||
current_model: None,
|
||||
},
|
||||
)?;
|
||||
|
||||
|
||||
@@ -1,8 +1,8 @@
|
||||
use crate::cache_semantics::{apply_chat_payload_breakpoints, CacheSemantics};
|
||||
use crate::conversation::message::{Message, MessageContentBlock};
|
||||
use crate::formats::anthropic::{
|
||||
adaptive_output_effort, model_supports_temperature, thinking_budget_tokens,
|
||||
thinking_type_for_provider, ThinkingType,
|
||||
adaptive_output_effort, model_supports_temperature, thinking_block_is_stale,
|
||||
thinking_budget_tokens, thinking_type_for_provider, ThinkingType,
|
||||
};
|
||||
use crate::model::ModelConfig;
|
||||
|
||||
@@ -104,13 +104,14 @@ fn format_tool_response(
|
||||
result
|
||||
}
|
||||
|
||||
/// Convert internal Message format to Databricks' API message specification
|
||||
/// Databricks is mostly OpenAI compatible, but has some differences (reasoning type, etc)
|
||||
/// some openai compatible endpoints use the anthropic image spec at the content level
|
||||
/// even though the message structure is otherwise following openai, the enum switches this
|
||||
fn format_messages(messages: &[Message], image_format: &ImageFormat) -> Vec<DatabricksMessage> {
|
||||
fn format_messages(
|
||||
messages: &[Message],
|
||||
image_format: &ImageFormat,
|
||||
current_model: Option<&str>,
|
||||
) -> Vec<DatabricksMessage> {
|
||||
let mut result = Vec::new();
|
||||
for message in messages {
|
||||
let thinking_is_stale = thinking_block_is_stale(message, current_model);
|
||||
let mut converted = DatabricksMessage {
|
||||
content: Value::Null,
|
||||
role: match message.role {
|
||||
@@ -137,22 +138,26 @@ fn format_messages(messages: &[Message], image_format: &ImageFormat) -> Vec<Data
|
||||
}
|
||||
}
|
||||
MessageContentBlock::Thinking(content) => {
|
||||
has_multiple_content = true;
|
||||
content_array.push(json!({
|
||||
"type": "reasoning",
|
||||
"summary": [{
|
||||
"type": "summary_text",
|
||||
"text": content.thinking,
|
||||
"signature": content.signature
|
||||
}]
|
||||
}));
|
||||
if !thinking_is_stale {
|
||||
has_multiple_content = true;
|
||||
content_array.push(json!({
|
||||
"type": "reasoning",
|
||||
"summary": [{
|
||||
"type": "summary_text",
|
||||
"text": content.thinking,
|
||||
"signature": content.signature
|
||||
}]
|
||||
}));
|
||||
}
|
||||
}
|
||||
MessageContentBlock::RedactedThinking(content) => {
|
||||
has_multiple_content = true;
|
||||
content_array.push(json!({
|
||||
"type": "reasoning",
|
||||
"summary": [{"type": "summary_encrypted_text", "data": content.data}]
|
||||
}));
|
||||
if !thinking_is_stale {
|
||||
has_multiple_content = true;
|
||||
content_array.push(json!({
|
||||
"type": "reasoning",
|
||||
"summary": [{"type": "summary_encrypted_text", "data": content.data}]
|
||||
}));
|
||||
}
|
||||
}
|
||||
MessageContentBlock::ToolRequest(request) => {
|
||||
has_tool_calls = true;
|
||||
@@ -527,7 +532,7 @@ pub fn create_request_for_provider(
|
||||
tool_call_id: None,
|
||||
};
|
||||
|
||||
let messages_spec = format_messages(messages, image_format);
|
||||
let messages_spec = format_messages(messages, image_format, Some(&model_config.model_name));
|
||||
let mut tools_spec = if !tools.is_empty() {
|
||||
format_tools(tools, &model_config.model_name)?
|
||||
} else {
|
||||
@@ -601,7 +606,7 @@ pub fn create_request_for_provider(
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::conversation::message::Message;
|
||||
use crate::conversation::message::{Message, MessageContent};
|
||||
use rmcp::model::CallToolResult;
|
||||
use rmcp::object;
|
||||
use serde_json::json;
|
||||
@@ -629,7 +634,7 @@ mod tests {
|
||||
#[test]
|
||||
fn test_format_messages() -> anyhow::Result<()> {
|
||||
let message = Message::user().with_text("Hello");
|
||||
let spec = format_messages(&[message], &ImageFormat::OpenAi);
|
||||
let spec = format_messages(&[message], &ImageFormat::OpenAi, None);
|
||||
|
||||
assert_eq!(spec.len(), 1);
|
||||
assert_eq!(spec[0].role, "user");
|
||||
@@ -637,6 +642,113 @@ mod tests {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn keeps_reasoning_block_from_the_same_model() {
|
||||
use crate::conversation::message::InferenceMetadata;
|
||||
let message = Message::assistant()
|
||||
.with_content(MessageContent::thinking("internal", "sig-xyz"))
|
||||
.with_text("answer")
|
||||
.with_inference(InferenceMetadata {
|
||||
provider: "databricks".to_string(),
|
||||
requested_model: "databricks-claude-opus-4-1".to_string(),
|
||||
resolved_model: None,
|
||||
provider_session_id: None,
|
||||
});
|
||||
|
||||
let spec = format_messages(
|
||||
&[message],
|
||||
&ImageFormat::OpenAi,
|
||||
Some("databricks-claude-opus-4-1"),
|
||||
);
|
||||
let has_reasoning = spec[0]
|
||||
.content
|
||||
.as_array()
|
||||
.map(|a| a.iter().any(|c| c["type"] == "reasoning"))
|
||||
.unwrap_or(false);
|
||||
assert!(has_reasoning, "same-model reasoning must be kept");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn drops_reasoning_block_from_a_different_model() {
|
||||
use crate::conversation::message::InferenceMetadata;
|
||||
let message = Message::assistant()
|
||||
.with_content(MessageContent::thinking("internal", "sig-xyz"))
|
||||
.with_text("answer")
|
||||
.with_inference(InferenceMetadata {
|
||||
provider: "databricks".to_string(),
|
||||
requested_model: "databricks-claude-opus-4-1".to_string(),
|
||||
resolved_model: None,
|
||||
provider_session_id: None,
|
||||
});
|
||||
|
||||
let spec = format_messages(
|
||||
&[message],
|
||||
&ImageFormat::OpenAi,
|
||||
Some("databricks-claude-sonnet-4-5"),
|
||||
);
|
||||
let has_reasoning = spec[0]
|
||||
.content
|
||||
.as_array()
|
||||
.map(|a| a.iter().any(|c| c["type"] == "reasoning"))
|
||||
.unwrap_or(false);
|
||||
assert!(!has_reasoning, "stale reasoning block must be dropped");
|
||||
assert_eq!(spec[0].content, Value::String("answer".to_string()));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn keeps_reasoning_when_endpoint_matches_despite_upstream_resolved_name() {
|
||||
use crate::conversation::message::InferenceMetadata;
|
||||
let message = Message::assistant()
|
||||
.with_content(MessageContent::thinking("internal", "sig-xyz"))
|
||||
.with_text("answer")
|
||||
.with_inference(InferenceMetadata {
|
||||
provider: "databricks".to_string(),
|
||||
requested_model: "databricks-claude-opus-4-1".to_string(),
|
||||
resolved_model: Some("claude-opus-4.1".to_string()),
|
||||
provider_session_id: None,
|
||||
});
|
||||
|
||||
let spec = format_messages(
|
||||
&[message],
|
||||
&ImageFormat::OpenAi,
|
||||
Some("databricks-claude-opus-4-1"),
|
||||
);
|
||||
let has_reasoning = spec[0]
|
||||
.content
|
||||
.as_array()
|
||||
.map(|a| a.iter().any(|c| c["type"] == "reasoning"))
|
||||
.unwrap_or(false);
|
||||
assert!(
|
||||
has_reasoning,
|
||||
"same-endpoint reasoning must be kept even when resolved_model differs"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn keeps_reasoning_when_current_model_matches_upstream_resolved_name() {
|
||||
use crate::conversation::message::InferenceMetadata;
|
||||
let message = Message::assistant()
|
||||
.with_content(MessageContent::thinking("internal", "sig-xyz"))
|
||||
.with_text("answer")
|
||||
.with_inference(InferenceMetadata {
|
||||
provider: "databricks".to_string(),
|
||||
requested_model: "my-claude-endpoint".to_string(),
|
||||
resolved_model: Some("claude-opus-4.1".to_string()),
|
||||
provider_session_id: None,
|
||||
});
|
||||
|
||||
let spec = format_messages(&[message], &ImageFormat::OpenAi, Some("claude-opus-4.1"));
|
||||
let has_reasoning = spec[0]
|
||||
.content
|
||||
.as_array()
|
||||
.map(|a| a.iter().any(|c| c["type"] == "reasoning"))
|
||||
.unwrap_or(false);
|
||||
assert!(
|
||||
has_reasoning,
|
||||
"reasoning must be kept when current_model matches the upstream resolved_model"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_format_messages_sanitizes_resource_tool_response() {
|
||||
let message = Message::user().with_tool_response(
|
||||
@@ -647,7 +759,7 @@ mod tests {
|
||||
)])),
|
||||
);
|
||||
|
||||
let spec = format_messages(&[message], &ImageFormat::OpenAi);
|
||||
let spec = format_messages(&[message], &ImageFormat::OpenAi, None);
|
||||
|
||||
assert_eq!(spec[0].content, "visibletext");
|
||||
}
|
||||
@@ -715,7 +827,7 @@ mod tests {
|
||||
));
|
||||
|
||||
let as_value =
|
||||
serde_json::to_value(format_messages(&messages, &ImageFormat::OpenAi)).unwrap();
|
||||
serde_json::to_value(format_messages(&messages, &ImageFormat::OpenAi, None)).unwrap();
|
||||
let spec = as_value.as_array().unwrap();
|
||||
|
||||
assert_eq!(spec.len(), 4);
|
||||
@@ -751,7 +863,7 @@ mod tests {
|
||||
));
|
||||
|
||||
let as_value =
|
||||
serde_json::to_value(format_messages(&messages, &ImageFormat::OpenAi)).unwrap();
|
||||
serde_json::to_value(format_messages(&messages, &ImageFormat::OpenAi, None)).unwrap();
|
||||
let spec = as_value.as_array().unwrap();
|
||||
|
||||
assert_eq!(spec.len(), 2);
|
||||
@@ -821,7 +933,7 @@ mod tests {
|
||||
// Create message with image path
|
||||
let message = Message::user().with_text(format!("Here is an image: {}", png_path_str));
|
||||
let as_value =
|
||||
serde_json::to_value(format_messages(&[message], &ImageFormat::OpenAi)).unwrap();
|
||||
serde_json::to_value(format_messages(&[message], &ImageFormat::OpenAi, None)).unwrap();
|
||||
let spec = as_value.as_array().unwrap();
|
||||
|
||||
assert_eq!(spec.len(), 1);
|
||||
@@ -1350,7 +1462,7 @@ mod tests {
|
||||
let message = Message::assistant()
|
||||
.with_tool_request("tool1", Ok(CallToolRequestParams::new("test_tool")));
|
||||
|
||||
let spec = format_messages(&[message], &ImageFormat::OpenAi);
|
||||
let spec = format_messages(&[message], &ImageFormat::OpenAi, None);
|
||||
let as_value = serde_json::to_value(spec)?;
|
||||
let spec_array = as_value.as_array().unwrap();
|
||||
|
||||
@@ -1387,7 +1499,7 @@ mod tests {
|
||||
final_resp,
|
||||
];
|
||||
|
||||
let spec = serde_json::to_value(format_messages(&messages, &ImageFormat::OpenAi))?;
|
||||
let spec = serde_json::to_value(format_messages(&messages, &ImageFormat::OpenAi, None))?;
|
||||
let mut open = std::collections::HashSet::new();
|
||||
for m in spec.as_array().unwrap() {
|
||||
match m.get("role").and_then(|v| v.as_str()) {
|
||||
@@ -1422,7 +1534,7 @@ mod tests {
|
||||
.with_arguments(object!({"param": "value", "number": 42}))),
|
||||
);
|
||||
|
||||
let spec = format_messages(&[message], &ImageFormat::OpenAi);
|
||||
let spec = format_messages(&[message], &ImageFormat::OpenAi, None);
|
||||
let as_value = serde_json::to_value(spec)?;
|
||||
let spec_array = as_value.as_array().unwrap();
|
||||
|
||||
@@ -1469,7 +1581,7 @@ mod tests {
|
||||
None,
|
||||
);
|
||||
|
||||
let spec = format_messages(&[message], &ImageFormat::OpenAi);
|
||||
let spec = format_messages(&[message], &ImageFormat::OpenAi, None);
|
||||
let as_value = serde_json::to_value(spec)?;
|
||||
let spec_array = as_value.as_array().unwrap();
|
||||
|
||||
@@ -1601,7 +1713,7 @@ mod tests {
|
||||
None,
|
||||
);
|
||||
|
||||
let spec = format_messages(&[message], &ImageFormat::OpenAi);
|
||||
let spec = format_messages(&[message], &ImageFormat::OpenAi, None);
|
||||
let as_value = serde_json::to_value(spec)?;
|
||||
let spec_array = as_value.as_array().unwrap();
|
||||
|
||||
@@ -1641,7 +1753,7 @@ mod tests {
|
||||
];
|
||||
|
||||
let as_value =
|
||||
serde_json::to_value(format_messages(&messages, &ImageFormat::OpenAi)).unwrap();
|
||||
serde_json::to_value(format_messages(&messages, &ImageFormat::OpenAi, None)).unwrap();
|
||||
let spec = as_value.as_array().unwrap();
|
||||
let roles: Vec<&str> = spec.iter().map(|m| m["role"].as_str().unwrap()).collect();
|
||||
|
||||
@@ -1675,7 +1787,7 @@ mod tests {
|
||||
];
|
||||
|
||||
let as_value =
|
||||
serde_json::to_value(format_messages(&messages, &ImageFormat::OpenAi)).unwrap();
|
||||
serde_json::to_value(format_messages(&messages, &ImageFormat::OpenAi, None)).unwrap();
|
||||
let spec = as_value.as_array().unwrap();
|
||||
let roles: Vec<&str> = spec.iter().map(|m| m["role"].as_str().unwrap()).collect();
|
||||
|
||||
|
||||
@@ -178,7 +178,7 @@ impl AnthropicProvider {
|
||||
system,
|
||||
messages,
|
||||
tools,
|
||||
self.format_options,
|
||||
self.format_options.clone(),
|
||||
)?;
|
||||
payload["stream"] = Value::Bool(true);
|
||||
let mut log = start_log(model_config, &payload)?;
|
||||
@@ -376,6 +376,7 @@ fn format_options_for_provider(preserves_thinking: bool) -> AnthropicFormatOptio
|
||||
preserve_unsigned_thinking: preserves_thinking,
|
||||
preserve_thinking_context: preserves_thinking,
|
||||
thinking_disabled: false,
|
||||
current_model: None,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -2182,13 +2182,11 @@ impl Agent {
|
||||
.ok()
|
||||
.and_then(|model_info| model_info.resolved_model);
|
||||
let provider_session_id = provider.provider_session_id();
|
||||
let inference = (resolved_model.is_some() || provider_session_id.is_some()).then(|| {
|
||||
InferenceMetadata {
|
||||
provider: provider_name.clone(),
|
||||
requested_model,
|
||||
resolved_model,
|
||||
provider_session_id,
|
||||
}
|
||||
let inference = Some(InferenceMetadata {
|
||||
provider: provider_name.clone(),
|
||||
requested_model,
|
||||
resolved_model,
|
||||
provider_session_id,
|
||||
});
|
||||
let session_manager = self.config.session_manager.clone();
|
||||
let session_id = session_config.id.clone();
|
||||
|
||||
@@ -501,13 +501,11 @@ impl Inference for InferenceRunner<'_> {
|
||||
.ok()
|
||||
.and_then(|model_info| model_info.resolved_model);
|
||||
let provider_session_id = self.provider.provider_session_id();
|
||||
let inference = (resolved_model.is_some() || provider_session_id.is_some()).then(|| {
|
||||
InferenceMetadata {
|
||||
provider: self.provider.get_name().to_string(),
|
||||
requested_model,
|
||||
resolved_model,
|
||||
provider_session_id,
|
||||
}
|
||||
let inference = Some(InferenceMetadata {
|
||||
provider: self.provider.get_name().to_string(),
|
||||
requested_model,
|
||||
resolved_model,
|
||||
provider_session_id,
|
||||
});
|
||||
|
||||
let mut accumulator = Conversation::empty();
|
||||
|
||||
@@ -319,3 +319,23 @@ async fn usage_and_provider_errors_survive_persistence() -> Result<()> {
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn requested_model_is_recorded_without_resolved_model() -> Result<()> {
|
||||
let (pipeline, api) = test_pipeline().await?;
|
||||
api.on("hello").reply("hi there");
|
||||
|
||||
let result = pipeline.run(["hello"]).await?;
|
||||
let requested_model = &result.session.model_config.as_ref().unwrap().model_name;
|
||||
let inference = result
|
||||
.conversation()
|
||||
.messages()
|
||||
.iter()
|
||||
.find(|message| message.role == rmcp::model::Role::Assistant)
|
||||
.and_then(|message| message.metadata.inference.as_ref())
|
||||
.expect("assistant inference metadata");
|
||||
|
||||
assert_eq!(&inference.requested_model, requested_model);
|
||||
assert_eq!(inference.resolved_model, None);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
@@ -323,13 +323,32 @@ impl BedrockProvider {
|
||||
let visible_messages: Vec<&Message> =
|
||||
messages.iter().filter(|m| m.is_agent_visible()).collect();
|
||||
|
||||
let last_idx = visible_messages.len().saturating_sub(1);
|
||||
let mut bedrock_messages: Vec<bedrock::Message> = Vec::new();
|
||||
for message in visible_messages {
|
||||
let formatted =
|
||||
to_bedrock_message_with_caching(message, false, Some(&model.model_name))?;
|
||||
if formatted.content().is_empty() {
|
||||
continue;
|
||||
}
|
||||
if let Some(previous) = bedrock_messages.last_mut() {
|
||||
if previous.role() == formatted.role() {
|
||||
previous.content.extend(formatted.content);
|
||||
continue;
|
||||
}
|
||||
}
|
||||
bedrock_messages.push(formatted);
|
||||
}
|
||||
|
||||
let bedrock_messages = visible_messages
|
||||
.iter()
|
||||
.enumerate()
|
||||
.map(|(idx, m)| to_bedrock_message_with_caching(m, enable_caching && idx == last_idx))
|
||||
.collect::<Result<Vec<_>>>()?;
|
||||
if enable_caching {
|
||||
if let Some(last) = bedrock_messages.last_mut() {
|
||||
last.content.push(bedrock::ContentBlock::CachePoint(
|
||||
bedrock::CachePointBlock::builder()
|
||||
.r#type(bedrock::CachePointType::Default)
|
||||
.build()
|
||||
.map_err(|error| ProviderError::ExecutionError(error.to_string()))?,
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
let tool_config = if tools.is_empty() {
|
||||
None
|
||||
@@ -1026,6 +1045,41 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn stale_reasoning_only_turn_is_removed_and_neighboring_roles_are_merged() {
|
||||
use crate::conversation::message::{InferenceMetadata, MessageContent};
|
||||
|
||||
let (provider, model) = create_mock_provider_and_model("anthropic.claude-sonnet-4");
|
||||
let messages = vec![
|
||||
Message::user().with_text("first"),
|
||||
Message::assistant()
|
||||
.with_content(MessageContent::thinking("internal", "sig-abc"))
|
||||
.with_inference(InferenceMetadata {
|
||||
provider: "aws_bedrock".to_string(),
|
||||
requested_model: "anthropic.claude-opus-4".to_string(),
|
||||
resolved_model: None,
|
||||
provider_session_id: None,
|
||||
}),
|
||||
Message::user().with_text("second"),
|
||||
];
|
||||
|
||||
let parts = provider
|
||||
.build_request_parts(&model, "system", &messages, &[])
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(parts.messages.len(), 1);
|
||||
assert_eq!(parts.messages[0].role(), &bedrock::ConversationRole::User);
|
||||
let text: Vec<&str> = parts.messages[0]
|
||||
.content()
|
||||
.iter()
|
||||
.filter_map(|content| match content {
|
||||
bedrock::ContentBlock::Text(text) => Some(text.as_str()),
|
||||
_ => None,
|
||||
})
|
||||
.collect();
|
||||
assert_eq!(text, vec!["first", "second"]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[serial]
|
||||
fn test_caching_enabled_for_claude_model() {
|
||||
|
||||
@@ -17,8 +17,9 @@ use crate::conversation::message::{Message, MessageContent};
|
||||
use crate::providers::bedrock::BEDROCK_PROVIDER_NAME;
|
||||
use crate::providers::canonical::maybe_get_canonical_model;
|
||||
use crate::providers::formats::anthropic::{
|
||||
adaptive_output_effort, model_supports_temperature, thinking_budget_tokens,
|
||||
thinking_type_for_provider, ThinkingType, ANTHROPIC_PROVIDER_NAME, MIN_ANSWER_TOKENS,
|
||||
adaptive_output_effort, model_supports_temperature, thinking_block_is_stale,
|
||||
thinking_budget_tokens, thinking_type_for_provider, ThinkingType, ANTHROPIC_PROVIDER_NAME,
|
||||
MIN_ANSWER_TOKENS,
|
||||
};
|
||||
use crate::utils::sanitize_unicode_tags;
|
||||
use goose_providers::conversation::token_usage::Usage;
|
||||
@@ -153,10 +154,22 @@ fn bedrock_model_supports_temperature(model_config: &ModelConfig) -> bool {
|
||||
pub fn to_bedrock_message_with_caching(
|
||||
message: &Message,
|
||||
enable_caching: bool,
|
||||
current_model: Option<&str>,
|
||||
) -> Result<bedrock::Message> {
|
||||
let thinking_is_stale = thinking_block_is_stale(message, current_model);
|
||||
let mut content_blocks: Vec<bedrock::ContentBlock> = message
|
||||
.content
|
||||
.iter()
|
||||
.filter(|content| {
|
||||
if !thinking_is_stale {
|
||||
return true;
|
||||
}
|
||||
match content {
|
||||
MessageContent::Thinking(thinking) => thinking.signature.is_empty(),
|
||||
MessageContent::RedactedThinking(_) => false,
|
||||
_ => true,
|
||||
}
|
||||
})
|
||||
.map(to_bedrock_message_content)
|
||||
.collect::<Result<_>>()?;
|
||||
|
||||
@@ -871,7 +884,7 @@ mod tests {
|
||||
MessageContent::text("Second text"),
|
||||
],
|
||||
);
|
||||
let bedrock_message = to_bedrock_message_with_caching(&message, true)?;
|
||||
let bedrock_message = to_bedrock_message_with_caching(&message, true, None)?;
|
||||
assert_eq!(bedrock_message.content.len(), 3);
|
||||
if let bedrock::ContentBlock::Text(text) = &bedrock_message.content[0] {
|
||||
assert_eq!(text, "First text");
|
||||
@@ -889,7 +902,7 @@ mod tests {
|
||||
));
|
||||
|
||||
// Caching disabled: no cache point added
|
||||
let no_cache = to_bedrock_message_with_caching(&message, false)?;
|
||||
let no_cache = to_bedrock_message_with_caching(&message, false, None)?;
|
||||
assert_eq!(no_cache.content.len(), 2);
|
||||
for block in &no_cache.content {
|
||||
assert!(!matches!(block, bedrock::ContentBlock::CachePoint(_)));
|
||||
@@ -897,12 +910,57 @@ mod tests {
|
||||
|
||||
// Empty content: no cache point added even with caching enabled
|
||||
let empty = Message::new(Role::User, Utc::now().timestamp(), vec![]);
|
||||
let empty_msg = to_bedrock_message_with_caching(&empty, true)?;
|
||||
let empty_msg = to_bedrock_message_with_caching(&empty, true, None)?;
|
||||
assert_eq!(empty_msg.content.len(), 0);
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn signed_thinking_from_model(model: &str) -> Message {
|
||||
use crate::conversation::message::InferenceMetadata;
|
||||
|
||||
Message::assistant()
|
||||
.with_content(MessageContent::thinking("internal", "sig-abc"))
|
||||
.with_text("answer")
|
||||
.with_inference(InferenceMetadata {
|
||||
provider: "aws_bedrock".to_string(),
|
||||
requested_model: model.to_string(),
|
||||
resolved_model: None,
|
||||
provider_session_id: None,
|
||||
})
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn keeps_signed_thinking_from_the_same_model() -> Result<()> {
|
||||
let message = signed_thinking_from_model("anthropic.claude-sonnet-4");
|
||||
let formatted =
|
||||
to_bedrock_message_with_caching(&message, false, Some("anthropic.claude-sonnet-4"))?;
|
||||
|
||||
assert!(matches!(
|
||||
formatted.content[0],
|
||||
bedrock::ContentBlock::ReasoningContent(_)
|
||||
));
|
||||
assert!(matches!(
|
||||
formatted.content[1],
|
||||
bedrock::ContentBlock::Text(_)
|
||||
));
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn drops_signed_thinking_from_a_different_model() -> Result<()> {
|
||||
let message = signed_thinking_from_model("anthropic.claude-opus-4");
|
||||
let formatted =
|
||||
to_bedrock_message_with_caching(&message, false, Some("anthropic.claude-sonnet-4"))?;
|
||||
|
||||
assert_eq!(formatted.content.len(), 1);
|
||||
assert!(matches!(
|
||||
formatted.content[0],
|
||||
bedrock::ContentBlock::Text(_)
|
||||
));
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_from_bedrock_usage_folds_cache_tokens_into_input() {
|
||||
let usage = bedrock::TokenUsage::builder()
|
||||
@@ -1244,7 +1302,7 @@ mod tests {
|
||||
],
|
||||
);
|
||||
|
||||
let bedrock_message = to_bedrock_message_with_caching(&message, true)?;
|
||||
let bedrock_message = to_bedrock_message_with_caching(&message, true, None)?;
|
||||
|
||||
// Verify cache point is added after all content blocks (text + tool request + cache point)
|
||||
assert_eq!(bedrock_message.content.len(), 3);
|
||||
@@ -1345,7 +1403,7 @@ mod tests {
|
||||
)],
|
||||
);
|
||||
|
||||
let bedrock_message = to_bedrock_message_with_caching(&message, true)?;
|
||||
let bedrock_message = to_bedrock_message_with_caching(&message, true, None)?;
|
||||
|
||||
// Verify cache point is added after tool response content
|
||||
assert_eq!(bedrock_message.content.len(), 2);
|
||||
@@ -1385,7 +1443,7 @@ mod tests {
|
||||
],
|
||||
);
|
||||
|
||||
let bedrock_message = to_bedrock_message_with_caching(&message, true)?;
|
||||
let bedrock_message = to_bedrock_message_with_caching(&message, true, None)?;
|
||||
|
||||
// Verify cache point is added at the end after all tool requests
|
||||
assert_eq!(bedrock_message.content.len(), 4);
|
||||
|
||||
Reference in New Issue
Block a user