pub(crate) mod cache;
pub(crate) mod convert;
use async_trait::async_trait;
use aws_config::BehaviorVersion;
use aws_config::meta::region::RegionProviderChain;
use aws_sdk_bedrockruntime::Client as BedrockRuntimeClient;
use aws_sdk_bedrockruntime::types::InferenceConfiguration;
use aws_smithy_types::error::display::DisplayErrorContext;
use aws_types::region::Region;
use serde_json::{Value, json};
use tokio::sync::OnceCell;
use crate::inference::adapter::InferenceAdapter;
use crate::inference::configurator::{Configurator, ResolvedProvider};
use crate::inference::error::InferenceError;
use crate::inference::registry::{ProviderCapabilities, ProviderId, capabilities};
use crate::inference::types::{ChatRequest, ChatResponse, ToolChoice};
const ENV_REGION_TRUSTY: &str = "TRUSTY_AWS_REGION";
const ENV_REGION_AWS: &str = "AWS_REGION";
const DEFAULT_REGION: &str = "us-east-1";
pub(crate) fn resolve_bedrock_region(explicit: Option<&str>) -> String {
if let Some(r) = explicit.filter(|s| !s.is_empty()) {
return r.to_string();
}
for var in [ENV_REGION_TRUSTY, ENV_REGION_AWS] {
if let Ok(val) = std::env::var(var) {
let val = val.trim().to_string();
if !val.is_empty() {
return val;
}
}
}
DEFAULT_REGION.to_string()
}
#[derive(Debug)]
pub struct BedrockAdapter {
region: String,
client: OnceCell<BedrockRuntimeClient>,
capabilities: ProviderCapabilities,
}
impl BedrockAdapter {
pub fn new(region: Option<&str>) -> Self {
Self {
region: resolve_bedrock_region(region),
client: OnceCell::new(),
capabilities: *capabilities(ProviderId::Bedrock),
}
}
pub fn region(&self) -> &str {
&self.region
}
async fn client(&self) -> Result<&BedrockRuntimeClient, InferenceError> {
self.client
.get_or_try_init(|| async {
let config = aws_config::defaults(BehaviorVersion::latest())
.region(RegionProviderChain::first_try(Region::new(
self.region.clone(),
)))
.load()
.await;
Ok::<_, InferenceError>(BedrockRuntimeClient::new(&config))
})
.await
}
}
#[async_trait]
impl InferenceAdapter for BedrockAdapter {
fn name(&self) -> &str {
ProviderId::Bedrock.as_str()
}
fn capabilities(&self) -> &ProviderCapabilities {
&self.capabilities
}
async fn chat(&self, request: &ChatRequest) -> Result<ChatResponse, InferenceError> {
let client = self.client().await?;
let (system_blocks, messages) = convert::build_converse_messages(request)?;
let inference = InferenceConfiguration::builder()
.set_max_tokens(request.max_tokens.map(|v| v as i32))
.set_temperature(request.temperature)
.build();
let mut sdk_req = client
.converse()
.model_id(bedrock_model_id(&request.model))
.inference_config(inference)
.set_messages(Some(messages));
if !system_blocks.is_empty() {
sdk_req = sdk_req.set_system(Some(system_blocks));
}
if let Some(tools) = &request.tools {
let tool_config = convert::build_tool_config(tools, request.tool_choice.as_ref())?;
if let Some(tool_config) = tool_config {
sdk_req = sdk_req.tool_config(tool_config);
}
}
let resp = sdk_req.send().await.map_err(|e| {
InferenceError::Provider(format!(
"Converse call failed (model={}, region={}): {}",
request.model,
self.region,
DisplayErrorContext(&e)
))
})?;
Ok(convert::converse_output_to_chat_response(
&resp,
&request.model,
))
}
fn map_tool_choice(&self, choice: ToolChoice) -> Value {
match choice {
ToolChoice::None => json!("none"),
ToolChoice::Auto => json!({"auto": {}}),
ToolChoice::Required => json!({"any": {}}),
ToolChoice::Function(name) => json!({"tool": {"name": name}}),
}
}
}
pub(crate) fn bedrock_model_id(slug: &str) -> &str {
slug.strip_prefix("bedrock/").unwrap_or(slug)
}
pub fn build(resolved: &ResolvedProvider) -> Result<Box<dyn InferenceAdapter>, InferenceError> {
debug_assert_eq!(resolved.provider(), ProviderId::Bedrock);
Ok(Box::new(BedrockAdapter::new(None)))
}
pub fn factory(resolved: &ResolvedProvider) -> Result<Box<dyn InferenceAdapter>, InferenceError> {
build(resolved)
}
pub fn register_bedrock_factory(cfg: &mut Configurator) {
cfg.register(ProviderId::Bedrock, Box::new(factory));
}
#[cfg(test)]
#[path = "tests.rs"]
mod tests;