use modelplease::{
BedrockProvider, BedrockProviderConfig, BedrockProviderDeps, CacheTtl, ContentPart,
GenerateRequest, LanguageModelConfig, LanguageModelProvider, Message, ModelId, RetryConfig,
Role,
};
const DEFAULT_CACHING_MODEL: &str = "us.anthropic.claude-haiku-4-5-20251001-v1:0";
const DEFAULT_NONCACHING_MODEL: &str = "us.meta.llama3-3-70b-instruct-v1:0";
fn big_system_prompt(nonce: &str) -> String {
use std::fmt::Write as _;
let mut s = format!("Cache test run {nonce}. You are a precise field extractor.\n");
for i in 0..600 {
let _ = writeln!(
s,
"Rule {i}: extract each requested value exactly and respond using the documented \
field markers, never inventing data that is not present in the input."
);
}
s
}
async fn make_provider() -> BedrockProvider {
let sdk = aws_config::defaults(aws_config::BehaviorVersion::latest())
.load()
.await;
let deps = BedrockProviderDeps {
runtime_client: aws_sdk_bedrockruntime::Client::new(&sdk),
control_client: aws_sdk_bedrock::Client::new(&sdk),
};
BedrockProvider::new(
deps,
BedrockProviderConfig {
region: None,
retry_config: RetryConfig::default(),
},
)
}
fn cached_messages(nonce: &str) -> Vec<Message> {
vec![
Message::with_parts(
Role::System,
vec![
ContentPart::text(big_system_prompt(nonce)),
ContentPart::cache_breakpoint(),
],
),
Message::user("Reply with the single word: ok."),
]
}
#[tokio::test(flavor = "multi_thread")]
#[ignore = "hits real AWS Bedrock; run via ... --run-ignored all"]
async fn bedrock_live_prompt_cache_creation_then_read() {
let model = ModelId::new(
std::env::var("BEDROCK_TEST_MODEL").unwrap_or_else(|_| DEFAULT_CACHING_MODEL.to_owned()),
);
let nonce = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos()
.to_string();
let messages = cached_messages(&nonce);
let config = LanguageModelConfig {
max_tokens: Some(16),
..Default::default()
};
let provider = make_provider().await;
let first = provider
.generate(GenerateRequest {
model: &model,
messages: &messages,
config: &config,
})
.await
.unwrap();
let u1 = first.usage.unwrap();
assert!(
u1.cache_creation_input_tokens > 0,
"call 1 should write the cache; got {u1:?} (is the prompt above the model min, and does \
this model support cachePoint?)"
);
assert_eq!(u1.cache_read_input_tokens, 0, "call 1 is a cold miss");
let second = provider
.generate(GenerateRequest {
model: &model,
messages: &messages,
config: &config,
})
.await
.unwrap();
let u2 = second.usage.unwrap();
assert!(
u2.cache_read_input_tokens > 0,
"call 2 should read the cache; got {u2:?} (ran within the 5-minute TTL?)"
);
assert!(
u2.cache_read_input_tokens > u2.input_tokens,
"warm call should serve most of the prompt from cache; got {u2:?}"
);
}
#[tokio::test(flavor = "multi_thread")]
#[ignore = "hits real AWS Bedrock; run via ... --run-ignored all"]
async fn bedrock_live_non_caching_model_succeeds() {
let model = ModelId::new(
std::env::var("BEDROCK_TEST_NONCACHING_MODEL")
.unwrap_or_else(|_| DEFAULT_NONCACHING_MODEL.to_owned()),
);
let messages = cached_messages("noncaching");
let config = LanguageModelConfig {
max_tokens: Some(16),
..Default::default()
};
let provider = make_provider().await;
let result = provider
.generate(GenerateRequest {
model: &model,
messages: &messages,
config: &config,
})
.await;
assert!(
result.is_ok(),
"non-caching model should still succeed: {result:?}"
);
let usage = result.unwrap().usage.unwrap();
assert_eq!(usage.cache_creation_input_tokens, 0);
assert_eq!(usage.cache_read_input_tokens, 0);
}
fn big_user_head(nonce: &str) -> String {
use std::fmt::Write as _;
let mut s =
format!("Stable context for run {nonce}. The following rules describe the extraction:\n");
for i in 0..600 {
let _ = writeln!(
s,
"Rule {i}: extract each requested value exactly using the documented field markers."
);
}
s
}
#[tokio::test(flavor = "multi_thread")]
#[ignore = "hits real AWS Bedrock; run via ... --run-ignored all"]
async fn bedrock_live_below_min_system_folds_into_user_head_cache() {
let model = ModelId::new(
std::env::var("BEDROCK_TEST_MODEL").unwrap_or_else(|_| DEFAULT_CACHING_MODEL.to_owned()),
);
let nonce = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos()
.to_string();
let head = big_user_head(&nonce);
let messages = vec![
Message::with_parts(
Role::System,
vec![
ContentPart::text("You are a precise field extractor."),
ContentPart::cache_breakpoint(),
],
),
Message::with_parts(
Role::User,
vec![
ContentPart::text(head),
ContentPart::cache_breakpoint(),
ContentPart::text("Reply with the single word: ok."),
],
),
];
let config = LanguageModelConfig {
max_tokens: Some(16),
..Default::default()
};
let provider = make_provider().await;
let first = provider
.generate(GenerateRequest {
model: &model,
messages: &messages,
config: &config,
})
.await
.unwrap();
let u1 = first.usage.unwrap();
assert!(
u1.cache_creation_input_tokens > 0,
"call 1 must cache the cumulative system+head span; got {u1:?}. \
If 0, either Bedrock rejected the below-min system breakpoint (fail-safe stripped \
all cache) or the user head wasn't large enough to clear the model min."
);
let second = provider
.generate(GenerateRequest {
model: &model,
messages: &messages,
config: &config,
})
.await
.unwrap();
let u2 = second.usage.unwrap();
assert!(
u2.cache_read_input_tokens > 0,
"call 2 must read the cumulative cache; got {u2:?}. \
A zero here means caching was stripped — Bedrock likely errored on the below-min \
breakpoint-1, and the fix is to suppress breakpoint-1 when a later cacheable \
breakpoint exists."
);
}
#[tokio::test(flavor = "multi_thread")]
#[ignore = "hits real AWS Bedrock; run via ... --run-ignored all"]
async fn bedrock_live_one_hour_ttl_request_accepted() {
let model = ModelId::new(
std::env::var("BEDROCK_TEST_MODEL").unwrap_or_else(|_| DEFAULT_CACHING_MODEL.to_owned()),
);
let nonce = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos()
.to_string();
let messages = cached_messages(&nonce);
let config = LanguageModelConfig {
max_tokens: Some(16),
cache_ttl: CacheTtl::OneHour,
..Default::default()
};
let provider = make_provider().await;
let resp = provider
.generate(GenerateRequest {
model: &model,
messages: &messages,
config: &config,
})
.await
.expect("bedrock must accept a 1h-TTL cachePoint request");
let usage = resp.usage.unwrap();
assert!(
usage.cache_creation_input_tokens > 0,
"1h-TTL call must still write the cache; got {usage:?}"
);
}