Enable bedrock prompt cache (#6710)
Signed-off-by: fbalicchia <fbalicchia@cuebiq.com> Co-authored-by: jh-block <jhugo@block.xyz>
This commit is contained in:
committed by
GitHub
parent
2265cd7858
commit
b58144632b
@@ -18,9 +18,9 @@ use reqwest::header::HeaderValue;
|
|||||||
use rmcp::model::Tool;
|
use rmcp::model::Tool;
|
||||||
use serde_json::Value;
|
use serde_json::Value;
|
||||||
|
|
||||||
// Import the migrated helper functions from providers/formats/bedrock.rs
|
|
||||||
use super::formats::bedrock::{
|
use super::formats::bedrock::{
|
||||||
from_bedrock_message, from_bedrock_usage, to_bedrock_message, to_bedrock_tool_config,
|
from_bedrock_message, from_bedrock_usage, to_bedrock_message_with_caching,
|
||||||
|
to_bedrock_tool_config,
|
||||||
};
|
};
|
||||||
use crate::session_context::SESSION_ID_HEADER;
|
use crate::session_context::SESSION_ID_HEADER;
|
||||||
|
|
||||||
@@ -189,6 +189,15 @@ impl BedrockProvider {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn should_enable_caching(&self) -> bool {
|
||||||
|
let config = crate::config::Config::global();
|
||||||
|
|
||||||
|
let enabled = config
|
||||||
|
.get_param::<bool>("BEDROCK_ENABLE_CACHING")
|
||||||
|
.unwrap_or(false);
|
||||||
|
enabled && self.model.model_name.contains("anthropic.claude")
|
||||||
|
}
|
||||||
|
|
||||||
async fn converse(
|
async fn converse(
|
||||||
&self,
|
&self,
|
||||||
session_id: Option<&str>,
|
session_id: Option<&str>,
|
||||||
@@ -198,16 +207,51 @@ impl BedrockProvider {
|
|||||||
) -> Result<(bedrock::Message, Option<bedrock::TokenUsage>), ProviderError> {
|
) -> Result<(bedrock::Message, Option<bedrock::TokenUsage>), ProviderError> {
|
||||||
let model_name = &self.model.model_name;
|
let model_name = &self.model.model_name;
|
||||||
|
|
||||||
|
let enable_caching = self.should_enable_caching();
|
||||||
|
|
||||||
|
let system_blocks = if enable_caching {
|
||||||
|
vec![
|
||||||
|
bedrock::SystemContentBlock::Text(system.to_string()),
|
||||||
|
// Add cache point AFTER the system prompt content
|
||||||
|
bedrock::SystemContentBlock::CachePoint(
|
||||||
|
bedrock::CachePointBlock::builder()
|
||||||
|
.r#type(bedrock::CachePointType::Default)
|
||||||
|
.build()
|
||||||
|
.map_err(|e| {
|
||||||
|
ProviderError::ExecutionError(format!(
|
||||||
|
"Failed to build cache point: {}",
|
||||||
|
e
|
||||||
|
))
|
||||||
|
})?,
|
||||||
|
),
|
||||||
|
]
|
||||||
|
} else {
|
||||||
|
vec![bedrock::SystemContentBlock::Text(system.to_string())]
|
||||||
|
};
|
||||||
|
|
||||||
|
let visible_messages: Vec<&Message> =
|
||||||
|
messages.iter().filter(|m| m.is_agent_visible()).collect();
|
||||||
|
|
||||||
|
// Cache the earliest messages (not most recent) because prompt caching
|
||||||
|
// requires exact prefix matching — caching recent messages would shift
|
||||||
|
// positions each turn, causing misses.
|
||||||
|
const MESSAGE_CACHE_BUDGET: usize = 3;
|
||||||
|
let cache_count = if enable_caching {
|
||||||
|
visible_messages.len().min(MESSAGE_CACHE_BUDGET)
|
||||||
|
} else {
|
||||||
|
0
|
||||||
|
};
|
||||||
|
|
||||||
let mut request = self
|
let mut request = self
|
||||||
.client
|
.client
|
||||||
.converse()
|
.converse()
|
||||||
.system(bedrock::SystemContentBlock::Text(system.to_string()))
|
.set_system(Some(system_blocks))
|
||||||
.model_id(model_name.to_string())
|
.model_id(model_name.to_string())
|
||||||
.set_messages(Some(
|
.set_messages(Some(
|
||||||
messages
|
visible_messages
|
||||||
.iter()
|
.iter()
|
||||||
.filter(|m| m.is_agent_visible())
|
.enumerate()
|
||||||
.map(to_bedrock_message)
|
.map(|(idx, m)| to_bedrock_message_with_caching(m, idx < cache_count))
|
||||||
.collect::<Result<_>>()?,
|
.collect::<Result<_>>()?,
|
||||||
));
|
));
|
||||||
|
|
||||||
@@ -272,7 +316,7 @@ impl ProviderDef for BedrockProvider {
|
|||||||
ProviderMetadata::new(
|
ProviderMetadata::new(
|
||||||
BEDROCK_PROVIDER_NAME,
|
BEDROCK_PROVIDER_NAME,
|
||||||
"Amazon Bedrock",
|
"Amazon Bedrock",
|
||||||
"Run models through Amazon Bedrock. Supports AWS SSO profiles - run 'aws sso login --profile <profile-name>' before using. Configure with AWS_PROFILE and AWS_REGION, use environment variables/credentials, or use AWS_BEARER_TOKEN_BEDROCK for bearer token authentication. Region is required for bearer token auth (can be set via AWS_REGION, AWS_DEFAULT_REGION, or AWS profile).",
|
"Run models through Amazon Bedrock. Supports AWS SSO profiles - run 'aws sso login --profile <profile-name>' before using. Configure with AWS_PROFILE and AWS_REGION, use environment variables/credentials, or use AWS_BEARER_TOKEN_BEDROCK for bearer token authentication. Region is required for bearer token auth (can be set via AWS_REGION, AWS_DEFAULT_REGION, or AWS profile). Prompt caching can be enabled for Anthropic Claude models by setting BEDROCK_ENABLE_CACHING=true.",
|
||||||
BEDROCK_DEFAULT_MODEL,
|
BEDROCK_DEFAULT_MODEL,
|
||||||
BEDROCK_KNOWN_MODELS.to_vec(),
|
BEDROCK_KNOWN_MODELS.to_vec(),
|
||||||
BEDROCK_DOC_LINK,
|
BEDROCK_DOC_LINK,
|
||||||
@@ -280,6 +324,7 @@ impl ProviderDef for BedrockProvider {
|
|||||||
ConfigKey::new("AWS_PROFILE", false, false, Some("default"), true),
|
ConfigKey::new("AWS_PROFILE", false, false, Some("default"), true),
|
||||||
ConfigKey::new("AWS_REGION", false, false, None, true),
|
ConfigKey::new("AWS_REGION", false, false, None, true),
|
||||||
ConfigKey::new("AWS_BEARER_TOKEN_BEDROCK", false, true, None, true),
|
ConfigKey::new("AWS_BEARER_TOKEN_BEDROCK", false, true, None, true),
|
||||||
|
ConfigKey::new("BEDROCK_ENABLE_CACHING", false, false, Some("false"), false),
|
||||||
],
|
],
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
@@ -363,6 +408,31 @@ impl Provider for BedrockProvider {
|
|||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
mod tests {
|
mod tests {
|
||||||
use super::*;
|
use super::*;
|
||||||
|
use serial_test::serial;
|
||||||
|
|
||||||
|
fn create_mock_provider(model_name: &str) -> BedrockProvider {
|
||||||
|
let sdk_config = aws_config::SdkConfig::builder()
|
||||||
|
.behavior_version(aws_config::BehaviorVersion::latest())
|
||||||
|
.region(aws_config::Region::new("us-east-1"))
|
||||||
|
.build();
|
||||||
|
let client = Client::new(&sdk_config);
|
||||||
|
|
||||||
|
BedrockProvider {
|
||||||
|
client,
|
||||||
|
model: ModelConfig {
|
||||||
|
model_name: model_name.to_string(),
|
||||||
|
context_limit: None,
|
||||||
|
temperature: None,
|
||||||
|
max_tokens: None,
|
||||||
|
toolshim: false,
|
||||||
|
toolshim_model: None,
|
||||||
|
fast_model_config: None,
|
||||||
|
request_params: None,
|
||||||
|
},
|
||||||
|
retry_config: RetryConfig::default(),
|
||||||
|
name: "aws_bedrock".to_string(),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn test_metadata_config_keys_have_expected_flags() {
|
fn test_metadata_config_keys_have_expected_flags() {
|
||||||
@@ -403,5 +473,55 @@ mod tests {
|
|||||||
bearer_token.secret,
|
bearer_token.secret,
|
||||||
"AWS_BEARER_TOKEN_BEDROCK should be marked as secret"
|
"AWS_BEARER_TOKEN_BEDROCK should be marked as secret"
|
||||||
);
|
);
|
||||||
|
|
||||||
|
let caching = meta
|
||||||
|
.config_keys
|
||||||
|
.iter()
|
||||||
|
.find(|k| k.name == "BEDROCK_ENABLE_CACHING")
|
||||||
|
.expect("BEDROCK_ENABLE_CACHING config key should exist");
|
||||||
|
assert!(
|
||||||
|
!caching.required,
|
||||||
|
"BEDROCK_ENABLE_CACHING should not be required"
|
||||||
|
);
|
||||||
|
assert!(
|
||||||
|
!caching.secret,
|
||||||
|
"BEDROCK_ENABLE_CACHING should not be marked as secret"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
#[serial]
|
||||||
|
fn test_caching_disabled_by_default() {
|
||||||
|
// Ensure clean environment
|
||||||
|
std::env::remove_var("BEDROCK_ENABLE_CACHING");
|
||||||
|
|
||||||
|
let provider = create_mock_provider("us.anthropic.claude-sonnet-4-5-20250929-v1:0");
|
||||||
|
assert!(
|
||||||
|
!provider.should_enable_caching(),
|
||||||
|
"Caching should be disabled by default"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_caching_disabled_for_non_claude_models() {
|
||||||
|
let provider = create_mock_provider("amazon.titan-text-express-v1");
|
||||||
|
assert!(
|
||||||
|
!provider.should_enable_caching(),
|
||||||
|
"Caching should be disabled for non-Claude models"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
#[serial]
|
||||||
|
fn test_caching_enabled_for_claude_model() {
|
||||||
|
std::env::set_var("BEDROCK_ENABLE_CACHING", "true");
|
||||||
|
|
||||||
|
let provider = create_mock_provider("us.anthropic.claude-sonnet-4-5-20250929-v1:0");
|
||||||
|
assert!(
|
||||||
|
provider.should_enable_caching(),
|
||||||
|
"Caching should be enabled for Claude models when BEDROCK_ENABLE_CACHING=true"
|
||||||
|
);
|
||||||
|
|
||||||
|
std::env::remove_var("BEDROCK_ENABLE_CACHING");
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -17,16 +17,28 @@ use serde_json::Value;
|
|||||||
use super::super::base::Usage;
|
use super::super::base::Usage;
|
||||||
use crate::conversation::message::{Message, MessageContent};
|
use crate::conversation::message::{Message, MessageContent};
|
||||||
|
|
||||||
pub fn to_bedrock_message(message: &Message) -> Result<bedrock::Message> {
|
pub fn to_bedrock_message_with_caching(
|
||||||
|
message: &Message,
|
||||||
|
enable_caching: bool,
|
||||||
|
) -> Result<bedrock::Message> {
|
||||||
|
let mut content_blocks: Vec<bedrock::ContentBlock> = message
|
||||||
|
.content
|
||||||
|
.iter()
|
||||||
|
.map(to_bedrock_message_content)
|
||||||
|
.collect::<Result<_>>()?;
|
||||||
|
|
||||||
|
if enable_caching && !content_blocks.is_empty() {
|
||||||
|
content_blocks.push(bedrock::ContentBlock::CachePoint(
|
||||||
|
bedrock::CachePointBlock::builder()
|
||||||
|
.r#type(bedrock::CachePointType::Default)
|
||||||
|
.build()
|
||||||
|
.map_err(|e| anyhow!("Failed to build cache point for message: {}", e))?,
|
||||||
|
));
|
||||||
|
}
|
||||||
|
|
||||||
bedrock::Message::builder()
|
bedrock::Message::builder()
|
||||||
.role(to_bedrock_role(&message.role))
|
.role(to_bedrock_role(&message.role))
|
||||||
.set_content(Some(
|
.set_content(Some(content_blocks))
|
||||||
message
|
|
||||||
.content
|
|
||||||
.iter()
|
|
||||||
.map(to_bedrock_message_content)
|
|
||||||
.collect::<Result<_>>()?,
|
|
||||||
))
|
|
||||||
.build()
|
.build()
|
||||||
.map_err(|err| anyhow!("Failed to construct Bedrock message: {}", err))
|
.map_err(|err| anyhow!("Failed to construct Bedrock message: {}", err))
|
||||||
}
|
}
|
||||||
@@ -159,9 +171,9 @@ pub fn to_bedrock_role(role: &Role) -> bedrock::ConversationRole {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn to_bedrock_image(data: &String, mime_type: &String) -> Result<bedrock::ImageBlock> {
|
pub fn to_bedrock_image(data: &str, mime_type: &str) -> Result<bedrock::ImageBlock> {
|
||||||
// Extract format from MIME type
|
// Extract format from MIME type
|
||||||
let format = match mime_type.as_str() {
|
let format = match mime_type {
|
||||||
"image/png" => bedrock::ImageFormat::Png,
|
"image/png" => bedrock::ImageFormat::Png,
|
||||||
"image/jpeg" | "image/jpg" => bedrock::ImageFormat::Jpeg,
|
"image/jpeg" | "image/jpg" => bedrock::ImageFormat::Jpeg,
|
||||||
"image/gif" => bedrock::ImageFormat::Gif,
|
"image/gif" => bedrock::ImageFormat::Gif,
|
||||||
@@ -287,6 +299,7 @@ pub fn from_bedrock_message(message: &bedrock::Message) -> Result<Message> {
|
|||||||
let content = message
|
let content = message
|
||||||
.content()
|
.content()
|
||||||
.iter()
|
.iter()
|
||||||
|
.filter(|block| !matches!(block, bedrock::ContentBlock::CachePoint(_)))
|
||||||
.map(from_bedrock_content_block)
|
.map(from_bedrock_content_block)
|
||||||
.collect::<Result<Vec<_>>>()?;
|
.collect::<Result<Vec<_>>>()?;
|
||||||
let created = Utc::now().timestamp();
|
let created = Utc::now().timestamp();
|
||||||
@@ -328,6 +341,10 @@ pub fn from_bedrock_content_block(block: &bedrock::ContentBlock) -> Result<Messa
|
|||||||
})
|
})
|
||||||
},
|
},
|
||||||
),
|
),
|
||||||
|
bedrock::ContentBlock::CachePoint(_) => {
|
||||||
|
// Filtered upstream in from_bedrock_message
|
||||||
|
bail!("CachePoint blocks should have been filtered out during message processing")
|
||||||
|
}
|
||||||
_ => bail!("Unsupported content block type from Bedrock"),
|
_ => bail!("Unsupported content block type from Bedrock"),
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
@@ -477,4 +494,242 @@ mod tests {
|
|||||||
|
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_to_bedrock_message_with_caching() -> Result<()> {
|
||||||
|
use chrono::Utc;
|
||||||
|
use rmcp::model::Role;
|
||||||
|
|
||||||
|
// Multiple content blocks: cache point appended at end, order preserved
|
||||||
|
let message = Message::new(
|
||||||
|
Role::User,
|
||||||
|
Utc::now().timestamp(),
|
||||||
|
vec![
|
||||||
|
MessageContent::text("First text"),
|
||||||
|
MessageContent::text("Second text"),
|
||||||
|
],
|
||||||
|
);
|
||||||
|
let bedrock_message = to_bedrock_message_with_caching(&message, true)?;
|
||||||
|
assert_eq!(bedrock_message.content.len(), 3);
|
||||||
|
if let bedrock::ContentBlock::Text(text) = &bedrock_message.content[0] {
|
||||||
|
assert_eq!(text, "First text");
|
||||||
|
} else {
|
||||||
|
panic!("Expected text content block");
|
||||||
|
}
|
||||||
|
if let bedrock::ContentBlock::Text(text) = &bedrock_message.content[1] {
|
||||||
|
assert_eq!(text, "Second text");
|
||||||
|
} else {
|
||||||
|
panic!("Expected text content block");
|
||||||
|
}
|
||||||
|
assert!(matches!(
|
||||||
|
bedrock_message.content[2],
|
||||||
|
bedrock::ContentBlock::CachePoint(_)
|
||||||
|
));
|
||||||
|
|
||||||
|
// Caching disabled: no cache point added
|
||||||
|
let no_cache = to_bedrock_message_with_caching(&message, false)?;
|
||||||
|
assert_eq!(no_cache.content.len(), 2);
|
||||||
|
for block in &no_cache.content {
|
||||||
|
assert!(!matches!(block, bedrock::ContentBlock::CachePoint(_)));
|
||||||
|
}
|
||||||
|
|
||||||
|
// 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)?;
|
||||||
|
assert_eq!(empty_msg.content.len(), 0);
|
||||||
|
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_from_bedrock_content_block_cache_point() {
|
||||||
|
// Create a cache point block with the required type field
|
||||||
|
let cache_point = bedrock::CachePointBlock::builder()
|
||||||
|
.r#type(bedrock::CachePointType::Default)
|
||||||
|
.build()
|
||||||
|
.unwrap();
|
||||||
|
let content_block = bedrock::ContentBlock::CachePoint(cache_point);
|
||||||
|
|
||||||
|
// Verify that converting a cache point results in an error
|
||||||
|
let result = from_bedrock_content_block(&content_block);
|
||||||
|
assert!(result.is_err());
|
||||||
|
assert!(result
|
||||||
|
.unwrap_err()
|
||||||
|
.to_string()
|
||||||
|
.contains("CachePoint blocks should have been filtered out"));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_from_bedrock_message_filters_cache_points() -> Result<()> {
|
||||||
|
use rmcp::model::Role;
|
||||||
|
|
||||||
|
// Create a Bedrock message with mixed content including CachePoint
|
||||||
|
let cache_point = bedrock::CachePointBlock::builder()
|
||||||
|
.r#type(bedrock::CachePointType::Default)
|
||||||
|
.build()
|
||||||
|
.unwrap();
|
||||||
|
|
||||||
|
let bedrock_message = bedrock::Message::builder()
|
||||||
|
.role(bedrock::ConversationRole::Assistant)
|
||||||
|
.content(bedrock::ContentBlock::Text("First text".to_string()))
|
||||||
|
.content(bedrock::ContentBlock::CachePoint(cache_point))
|
||||||
|
.content(bedrock::ContentBlock::Text("Second text".to_string()))
|
||||||
|
.build()
|
||||||
|
.unwrap();
|
||||||
|
|
||||||
|
// Convert from Bedrock format
|
||||||
|
let message = from_bedrock_message(&bedrock_message)?;
|
||||||
|
|
||||||
|
// Verify that CachePoint was filtered out and only text content remains
|
||||||
|
assert_eq!(message.content.len(), 2);
|
||||||
|
assert_eq!(message.role, Role::Assistant);
|
||||||
|
|
||||||
|
if let MessageContent::Text(text) = &message.content[0] {
|
||||||
|
assert_eq!(text.text, "First text");
|
||||||
|
} else {
|
||||||
|
panic!("Expected first text content");
|
||||||
|
}
|
||||||
|
|
||||||
|
if let MessageContent::Text(text) = &message.content[1] {
|
||||||
|
assert_eq!(text.text, "Second text");
|
||||||
|
} else {
|
||||||
|
panic!("Expected second text content");
|
||||||
|
}
|
||||||
|
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_cache_points_with_tool_request_messages() -> Result<()> {
|
||||||
|
use chrono::Utc;
|
||||||
|
use rmcp::model::{CallToolRequestParams, Role};
|
||||||
|
use serde_json::json;
|
||||||
|
|
||||||
|
let message = Message::new(
|
||||||
|
Role::Assistant,
|
||||||
|
Utc::now().timestamp(),
|
||||||
|
vec![
|
||||||
|
MessageContent::text("I'll use a tool"),
|
||||||
|
MessageContent::tool_request(
|
||||||
|
"tool_1".to_string(),
|
||||||
|
Ok(CallToolRequestParams {
|
||||||
|
meta: None,
|
||||||
|
task: None,
|
||||||
|
name: "test_tool".into(),
|
||||||
|
arguments: Some(object(json!({"param": "value"}))),
|
||||||
|
}),
|
||||||
|
),
|
||||||
|
],
|
||||||
|
);
|
||||||
|
|
||||||
|
let bedrock_message = to_bedrock_message_with_caching(&message, true)?;
|
||||||
|
|
||||||
|
// Verify cache point is added after all content blocks (text + tool request + cache point)
|
||||||
|
assert_eq!(bedrock_message.content.len(), 3);
|
||||||
|
assert!(matches!(
|
||||||
|
bedrock_message.content[0],
|
||||||
|
bedrock::ContentBlock::Text(_)
|
||||||
|
));
|
||||||
|
assert!(matches!(
|
||||||
|
bedrock_message.content[1],
|
||||||
|
bedrock::ContentBlock::ToolUse(_)
|
||||||
|
));
|
||||||
|
assert!(matches!(
|
||||||
|
bedrock_message.content[2],
|
||||||
|
bedrock::ContentBlock::CachePoint(_)
|
||||||
|
));
|
||||||
|
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_cache_points_with_tool_response_messages() -> Result<()> {
|
||||||
|
use chrono::Utc;
|
||||||
|
use rmcp::model::{CallToolResult, Role};
|
||||||
|
|
||||||
|
let message = Message::new(
|
||||||
|
Role::User,
|
||||||
|
Utc::now().timestamp(),
|
||||||
|
vec![MessageContent::tool_response(
|
||||||
|
"tool_1".to_string(),
|
||||||
|
Ok(CallToolResult {
|
||||||
|
content: vec![Content::text("Tool result text".to_string())],
|
||||||
|
structured_content: None,
|
||||||
|
is_error: Some(false),
|
||||||
|
meta: None,
|
||||||
|
}),
|
||||||
|
)],
|
||||||
|
);
|
||||||
|
|
||||||
|
let bedrock_message = to_bedrock_message_with_caching(&message, true)?;
|
||||||
|
|
||||||
|
// Verify cache point is added after tool response content
|
||||||
|
assert_eq!(bedrock_message.content.len(), 2);
|
||||||
|
assert!(matches!(
|
||||||
|
bedrock_message.content[0],
|
||||||
|
bedrock::ContentBlock::ToolResult(_)
|
||||||
|
));
|
||||||
|
assert!(matches!(
|
||||||
|
bedrock_message.content[1],
|
||||||
|
bedrock::ContentBlock::CachePoint(_)
|
||||||
|
));
|
||||||
|
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_cache_points_with_mixed_tool_content() -> Result<()> {
|
||||||
|
use chrono::Utc;
|
||||||
|
use rmcp::model::{CallToolRequestParams, Role};
|
||||||
|
use serde_json::json;
|
||||||
|
|
||||||
|
let message = Message::new(
|
||||||
|
Role::Assistant,
|
||||||
|
Utc::now().timestamp(),
|
||||||
|
vec![
|
||||||
|
MessageContent::text("Using tools"),
|
||||||
|
MessageContent::tool_request(
|
||||||
|
"tool_1".to_string(),
|
||||||
|
Ok(CallToolRequestParams {
|
||||||
|
meta: None,
|
||||||
|
task: None,
|
||||||
|
name: "tool_a".into(),
|
||||||
|
arguments: Some(object(json!({"key": "val"}))),
|
||||||
|
}),
|
||||||
|
),
|
||||||
|
MessageContent::tool_request(
|
||||||
|
"tool_2".to_string(),
|
||||||
|
Ok(CallToolRequestParams {
|
||||||
|
meta: None,
|
||||||
|
task: None,
|
||||||
|
name: "tool_b".into(),
|
||||||
|
arguments: Some(object(json!({"key": "val"}))),
|
||||||
|
}),
|
||||||
|
),
|
||||||
|
],
|
||||||
|
);
|
||||||
|
|
||||||
|
let bedrock_message = to_bedrock_message_with_caching(&message, true)?;
|
||||||
|
|
||||||
|
// Verify cache point is added at the end after all tool requests
|
||||||
|
assert_eq!(bedrock_message.content.len(), 4);
|
||||||
|
assert!(matches!(
|
||||||
|
bedrock_message.content[0],
|
||||||
|
bedrock::ContentBlock::Text(_)
|
||||||
|
));
|
||||||
|
assert!(matches!(
|
||||||
|
bedrock_message.content[1],
|
||||||
|
bedrock::ContentBlock::ToolUse(_)
|
||||||
|
));
|
||||||
|
assert!(matches!(
|
||||||
|
bedrock_message.content[2],
|
||||||
|
bedrock::ContentBlock::ToolUse(_)
|
||||||
|
));
|
||||||
|
assert!(matches!(
|
||||||
|
bedrock_message.content[3],
|
||||||
|
bedrock::ContentBlock::CachePoint(_)
|
||||||
|
));
|
||||||
|
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user