Additional cleanup of caching

This commit is contained in:
Ryan Ward
2026-06-01 15:21:03 -05:00
parent 7f039d2c08
commit 54e4e8abf5
12 changed files with 100 additions and 59 deletions
+2
View File
@@ -133,6 +133,7 @@ impl BedrockClient {
needs_create_task: bool,
messages: Vec<ConversationMessage>,
system_prompt: Option<String>,
compact_summary: Option<String>,
tools: Vec<ToolDefinition>,
max_tokens: i32,
temperature: Option<f32>,
@@ -158,6 +159,7 @@ impl BedrockClient {
let converted = build_converse_request(
messages.clone(),
system_prompt.clone(),
compact_summary,
tools.clone(),
max_tokens,
temperature,
+30 -16
View File
@@ -59,7 +59,7 @@ pub enum ContentPart {
},
}
#[derive(Clone)]
#[derive(Debug, Clone)]
pub struct ToolDefinition {
pub name: String,
pub description: String,
@@ -69,6 +69,7 @@ pub struct ToolDefinition {
pub fn build_converse_request(
messages: Vec<ConversationMessage>,
system_prompt: Option<String>,
compact_summary: Option<String>,
tools: Vec<ToolDefinition>,
max_tokens: i32,
temperature: Option<f32>,
@@ -76,7 +77,7 @@ pub fn build_converse_request(
stop_sequences: Option<Vec<String>>,
) -> ConvertedRequest {
let bedrock_messages = convert_messages(messages);
let system = convert_system_prompt(system_prompt);
let system = convert_system_prompt(system_prompt, compact_summary);
let inference_config = build_inference_config(max_tokens, temperature, top_p, stop_sequences);
let tool_config = build_tool_config(tools);
@@ -267,22 +268,35 @@ fn coalesce_consecutive_roles(messages: Vec<BedrockMessage>) -> Vec<BedrockMessa
result
}
fn convert_system_prompt(system_prompt: Option<String>) -> Vec<SystemContentBlock> {
match system_prompt {
Some(prompt) if !prompt.is_empty() => {
vec![
SystemContentBlock::Text(prompt),
SystemContentBlock::CachePoint(
CachePointBlock::builder()
.r#type(CachePointType::Default)
.ttl(CacheTtl::OneHour)
.build()
.expect("valid cache point"),
),
]
fn convert_system_prompt(
system_prompt: Option<String>,
compact_summary: Option<String>,
) -> Vec<SystemContentBlock> {
let mut blocks = Vec::new();
if let Some(prompt) = system_prompt {
if !prompt.is_empty() {
blocks.push(SystemContentBlock::Text(prompt));
}
_ => vec![],
}
if let Some(summary) = compact_summary {
blocks.push(SystemContentBlock::Text(format!(
"<conversation-summary>\n{summary}\n</conversation-summary>"
)));
}
if !blocks.is_empty() {
blocks.push(SystemContentBlock::CachePoint(
CachePointBlock::builder()
.r#type(CachePointType::Default)
.ttl(CacheTtl::OneHour)
.build()
.expect("valid cache point"),
));
}
blocks
}
fn build_inference_config(
+15 -13
View File
@@ -10,7 +10,7 @@ fn test_text_message_converts_to_single_block() {
content: MessageContent::Text("Hello".to_string()),
}];
let result = build_converse_request(messages, None, vec![], 4096, None, None, None);
let result = build_converse_request(messages, None, None, vec![], 4096, None, None, None);
assert_eq!(result.messages.len(), 1);
assert_eq!(result.messages[0].role(), &ConversationRole::User);
@@ -29,7 +29,7 @@ fn test_tool_use_produces_valid_json_input() {
},
}];
let result = build_converse_request(messages, None, vec![], 4096, None, None, None);
let result = build_converse_request(messages, None, None, vec![], 4096, None, None, None);
assert_eq!(result.messages.len(), 1);
assert_eq!(result.messages[0].role(), &ConversationRole::Assistant);
@@ -53,7 +53,7 @@ fn test_tool_result_with_matching_id() {
},
}];
let result = build_converse_request(messages, None, vec![], 4096, None, None, None);
let result = build_converse_request(messages, None, None, vec![], 4096, None, None, None);
assert_eq!(result.messages.len(), 1);
match &result.messages[0].content()[0] {
@@ -75,7 +75,7 @@ fn test_tool_result_error_status() {
},
}];
let result = build_converse_request(messages, None, vec![], 4096, None, None, None);
let result = build_converse_request(messages, None, None, vec![], 4096, None, None, None);
match &result.messages[0].content()[0] {
ContentBlock::ToolResult(block) => {
@@ -101,7 +101,7 @@ fn test_consecutive_same_role_messages_coalesced() {
},
];
let result = build_converse_request(messages, None, vec![], 4096, None, None, None);
let result = build_converse_request(messages, None, None, vec![], 4096, None, None, None);
assert_eq!(result.messages.len(), 1);
assert_eq!(result.messages[0].content().len(), 2);
@@ -126,7 +126,7 @@ fn test_alternating_roles_not_coalesced() {
},
];
let result = build_converse_request(messages, None, vec![], 4096, None, None, None);
let result = build_converse_request(messages, None, None, vec![], 4096, None, None, None);
assert_eq!(result.messages.len(), 3);
assert_eq!(result.messages[0].role(), &ConversationRole::User);
@@ -144,6 +144,7 @@ fn test_system_prompt_separated_from_messages() {
let result = build_converse_request(
messages,
Some("You are a helpful assistant.".to_string()),
None,
vec![],
4096,
None,
@@ -159,16 +160,16 @@ fn test_system_prompt_separated_from_messages() {
#[test]
fn test_empty_system_prompt_produces_empty_vec() {
let result =
build_converse_request(vec![], Some("".to_string()), vec![], 4096, None, None, None);
build_converse_request(vec![], Some("".to_string()), None, vec![], 4096, None, None, None);
assert!(result.system.is_empty());
let result2 = build_converse_request(vec![], None, vec![], 4096, None, None, None);
let result2 = build_converse_request(vec![], None, None, vec![], 4096, None, None, None);
assert!(result2.system.is_empty());
}
#[test]
fn test_empty_tools_produce_none_config() {
let result = build_converse_request(vec![], None, vec![], 4096, None, None, None);
let result = build_converse_request(vec![], None, None, vec![], 4096, None, None, None);
assert!(result.tool_config.is_none());
}
@@ -186,7 +187,7 @@ fn test_tool_definitions_produce_tool_config() {
}),
}];
let result = build_converse_request(vec![], None, tools, 4096, None, None, None);
let result = build_converse_request(vec![], None, None, tools, 4096, None, None, None);
assert!(result.tool_config.is_some());
let config = result.tool_config.unwrap();
@@ -195,7 +196,7 @@ fn test_tool_definitions_produce_tool_config() {
#[test]
fn test_inference_config_max_tokens_only() {
let result = build_converse_request(vec![], None, vec![], 8192, None, None, None);
let result = build_converse_request(vec![], None, None, vec![], 8192, None, None, None);
assert_eq!(result.inference_config.max_tokens(), Some(8192));
assert_eq!(result.inference_config.temperature(), None);
assert_eq!(result.inference_config.top_p(), None);
@@ -207,6 +208,7 @@ fn test_inference_config_all_params() {
let result = build_converse_request(
vec![],
None,
None,
vec![],
4096,
Some(0.7),
@@ -233,7 +235,7 @@ fn test_multipart_content_produces_multiple_blocks() {
]),
}];
let result = build_converse_request(messages, None, vec![], 4096, None, None, None);
let result = build_converse_request(messages, None, None, vec![], 4096, None, None, None);
assert_eq!(result.messages[0].content().len(), 2);
assert!(matches!(
@@ -271,7 +273,7 @@ fn test_tool_result_after_tool_use_coalesced_into_user_message() {
},
];
let result = build_converse_request(messages, None, vec![], 4096, None, None, None);
let result = build_converse_request(messages, None, None, vec![], 4096, None, None, None);
assert_eq!(result.messages.len(), 3);
assert_eq!(result.messages[0].role(), &ConversationRole::User);
+16
View File
@@ -550,6 +550,7 @@ impl AgentSimulation {
needs_create_task,
self.conversation.clone(),
Some(self.system_prompt.clone()),
None,
self.tools.clone(),
8192,
None,
@@ -557,6 +558,7 @@ impl AgentSimulation {
None,
None,
Arc::new(Mutex::new(Vec::new())),
false,
)
.await
.expect("converse_stream should succeed");
@@ -1137,6 +1139,7 @@ async fn test_reasoning_model_produces_substantial_output() {
true,
messages,
system,
None,
agent_tools(),
8192,
None,
@@ -1144,6 +1147,7 @@ async fn test_reasoning_model_produces_substantial_output() {
None,
None,
Arc::new(Mutex::new(Vec::new())),
false,
)
.await
.expect("converse_stream should succeed");
@@ -1260,6 +1264,7 @@ async fn test_event_sequence_matches_controller_expectations() {
true,
messages,
None,
None,
vec![],
100,
None,
@@ -1267,6 +1272,7 @@ async fn test_event_sequence_matches_controller_expectations() {
None,
None,
Arc::new(Mutex::new(Vec::new())),
false,
)
.await
.expect("should connect");
@@ -1375,6 +1381,7 @@ async fn test_followup_turn_does_not_send_create_task() {
false,
messages,
None,
None,
vec![],
100,
None,
@@ -1382,6 +1389,7 @@ async fn test_followup_turn_does_not_send_create_task() {
None,
None,
Arc::new(Mutex::new(Vec::new())),
false,
)
.await
.expect("should connect");
@@ -1459,6 +1467,7 @@ async fn run_slash_command_test(
true,
messages,
system_prompt,
None,
tools,
4096,
None,
@@ -1466,6 +1475,7 @@ async fn run_slash_command_test(
None,
None,
Arc::new(Mutex::new(Vec::new())),
false,
)
.await
.expect("stream should connect");
@@ -1687,6 +1697,7 @@ async fn test_slash_resume_conversation() {
true,
messages,
None,
None,
vec![],
256,
None,
@@ -1694,6 +1705,7 @@ async fn test_slash_resume_conversation() {
None,
None,
Arc::new(Mutex::new(Vec::new())),
false,
)
.await
.expect("stream should connect");
@@ -1799,6 +1811,7 @@ async fn test_empty_messages_safety_check() {
true,
messages,
None,
None,
vec![],
100,
None,
@@ -1806,6 +1819,7 @@ async fn test_empty_messages_safety_check() {
None,
None,
Arc::new(Mutex::new(Vec::new())),
false,
)
.await
.expect("safety fallback message should work");
@@ -1960,6 +1974,7 @@ async fn test_full_proto_round_trip_with_tool_history() {
false,
messages,
system_prompt,
None,
tools,
1024,
None,
@@ -1967,6 +1982,7 @@ async fn test_full_proto_round_trip_with_tool_history() {
None,
None,
Arc::new(Mutex::new(Vec::new())),
false,
)
.await;
+2
View File
@@ -59,6 +59,7 @@ async fn collect_stream_output(
true,
messages,
system_prompt,
None,
tools,
8192,
None,
@@ -66,6 +67,7 @@ async fn collect_stream_output(
None,
None,
Arc::new(Mutex::new(Vec::new())),
false,
)
.await
.expect("converse_stream should succeed");
+5 -1
View File
@@ -13,6 +13,7 @@ pub struct TranslatorRequest {
pub model_id: String,
pub root_task_id: Option<String>,
pub bedrock_message_history: Vec<ConversationMessage>,
pub bedrock_compact_summary: Option<String>,
pub bedrock_messages_sent: Arc<Mutex<Vec<ConversationMessage>>>,
pub is_summarization: bool,
}
@@ -71,12 +72,14 @@ pub async fn execute(
request_translator::sanitize_messages_for_bedrock(&mut messages);
let system_prompt = request_translator::extract_system_prompt(request);
let compact_summary = params.bedrock_compact_summary;
let tools = request_translator::extract_tools(request);
log::info!(
"[bedrock] Sending {} messages, system_prompt={}, tools={}",
"[bedrock] Sending {} messages, system_prompt={}, compact_summary={}, tools={}",
messages.len(),
system_prompt.is_some(),
compact_summary.is_some(),
tools.len()
);
@@ -94,6 +97,7 @@ pub async fn execute(
needs_create_task,
messages.clone(),
system_prompt,
compact_summary,
tools,
64000,
None,