use crate::agent::provider::LlmProvider;
use crate::agent::types::{Block, Msg, Role, Stop, ToolSpec, Turn};
use async_trait::async_trait;
use aws_sdk_bedrockruntime::types::{
ContentBlock, ConversationRole, InferenceConfiguration, Message, StopReason, SystemContentBlock,
Tool, ToolConfiguration, ToolInputSchema, ToolResultBlock, ToolResultContentBlock,
ToolResultStatus, ToolSpecification, ToolUseBlock,
};
use aws_sdk_bedrockruntime::Client;
use aws_smithy_types::{Document, Number};
use serde_json::Value;
use std::collections::HashMap;
pub struct BedrockProvider {
client: Client,
model_id: String,
}
impl BedrockProvider {
pub async fn new(model_id: String, region: Option<String>) -> Result<BedrockProvider, String> {
let mut loader = aws_config::defaults(aws_config::BehaviorVersion::latest());
if let Some(r) = region {
loader = loader.region(aws_config::Region::new(r));
}
let conf = loader.load().await;
Ok(BedrockProvider { client: Client::new(&conf), model_id })
}
}
fn value_to_document(v: &Value) -> Document {
match v {
Value::Null => Document::Null,
Value::Bool(b) => Document::Bool(*b),
Value::Number(n) => {
if let Some(u) = n.as_u64() {
Document::Number(Number::PosInt(u))
} else if let Some(i) = n.as_i64() {
Document::Number(Number::NegInt(i))
} else {
Document::Number(Number::Float(n.as_f64().unwrap_or(0.0)))
}
}
Value::String(s) => Document::String(s.clone()),
Value::Array(a) => Document::Array(a.iter().map(value_to_document).collect()),
Value::Object(o) => {
let m: HashMap<String, Document> = o.iter().map(|(k, v)| (k.clone(), value_to_document(v))).collect();
Document::Object(m)
}
}
}
fn document_to_value(d: &Document) -> Value {
match d {
Document::Null => Value::Null,
Document::Bool(b) => Value::Bool(*b),
Document::Number(Number::PosInt(u)) => Value::from(*u),
Document::Number(Number::NegInt(i)) => Value::from(*i),
Document::Number(Number::Float(f)) => Value::from(*f),
Document::String(s) => Value::String(s.clone()),
Document::Array(a) => Value::Array(a.iter().map(document_to_value).collect()),
Document::Object(o) => {
Value::Object(o.iter().map(|(k, v)| (k.clone(), document_to_value(v))).collect())
}
}
}
fn to_message(m: &Msg) -> Result<Message, String> {
let role = match m.role {
Role::User => ConversationRole::User,
Role::Assistant => ConversationRole::Assistant,
};
let mut content = Vec::new();
for b in &m.blocks {
match b {
Block::Text(t) => content.push(ContentBlock::Text(t.clone())),
Block::ToolUse { id, name, input } => {
let tu = ToolUseBlock::builder()
.tool_use_id(id.clone())
.name(name.clone())
.input(value_to_document(input))
.build()
.map_err(|e| e.to_string())?;
content.push(ContentBlock::ToolUse(tu));
}
Block::ToolResult { id, content: c, is_error } => {
let tr = ToolResultBlock::builder()
.tool_use_id(id.clone())
.content(ToolResultContentBlock::Text(c.clone()))
.status(if *is_error { ToolResultStatus::Error } else { ToolResultStatus::Success })
.build()
.map_err(|e| e.to_string())?;
content.push(ContentBlock::ToolResult(tr));
}
}
}
Message::builder().role(role).set_content(Some(content)).build().map_err(|e| e.to_string())
}
#[async_trait]
impl LlmProvider for BedrockProvider {
fn name(&self) -> &str {
"bedrock"
}
async fn chat(&self, system: &str, msgs: &[Msg], tools: &[ToolSpec]) -> Result<Turn, String> {
let mut req = self
.client
.converse()
.model_id(&self.model_id)
.system(SystemContentBlock::Text(system.to_string()))
.inference_config(InferenceConfiguration::builder().max_tokens(2048).build());
if !tools.is_empty() {
let mut tc = ToolConfiguration::builder();
for t in tools {
let spec = ToolSpecification::builder()
.name(t.name.clone())
.description(t.description.clone())
.input_schema(ToolInputSchema::Json(value_to_document(&t.schema)))
.build()
.map_err(|e| e.to_string())?;
tc = tc.tools(Tool::ToolSpec(spec));
}
req = req.tool_config(tc.build().map_err(|e| e.to_string())?);
}
for m in msgs {
req = req.messages(to_message(m)?);
}
let resp = req.send().await.map_err(|e| format!("bedrock converse: {e:?}"))?;
let stop = match resp.stop_reason() {
StopReason::ToolUse => Stop::ToolUse,
StopReason::EndTurn => Stop::EndTurn,
_ => Stop::Other,
};
let out = resp.output().ok_or("bedrock: no output")?;
let message = out.as_message().map_err(|_| "bedrock: output was not a message".to_string())?;
let mut text = String::new();
let mut tool_uses = Vec::new();
for block in message.content() {
match block {
ContentBlock::Text(t) => text.push_str(t),
ContentBlock::ToolUse(tu) => {
tool_uses.push((tu.tool_use_id().to_string(), tu.name().to_string(), document_to_value(tu.input())));
}
_ => {}
}
}
Ok(Turn { text, tool_uses, stop })
}
async fn chat_json(
&self,
system: &str,
msgs: &[Msg],
schema: &serde_json::Value,
name: &str,
) -> Result<Option<serde_json::Value>, String> {
use aws_sdk_bedrockruntime::types::{ToolChoice, SpecificToolChoice};
let spec = ToolSpecification::builder()
.name(name)
.description("Emit the result as JSON matching the schema.")
.input_schema(ToolInputSchema::Json(value_to_document(schema)))
.build()
.map_err(|e| e.to_string())?;
let tool_choice = ToolChoice::Tool(
SpecificToolChoice::builder().name(name).build().map_err(|e| e.to_string())?,
);
let tc = ToolConfiguration::builder()
.tools(Tool::ToolSpec(spec))
.tool_choice(tool_choice)
.build()
.map_err(|e| e.to_string())?;
let mut req = self
.client
.converse()
.model_id(&self.model_id)
.system(SystemContentBlock::Text(system.to_string()))
.inference_config(
InferenceConfiguration::builder().max_tokens(3000).temperature(0.0).build(),
)
.tool_config(tc);
for m in msgs {
req = req.messages(to_message(m)?);
}
let resp = req.send().await.map_err(|e| format!("bedrock converse: {e:?}"))?;
let out = resp.output().ok_or("bedrock: no output")?;
let message = out.as_message().map_err(|_| "bedrock: output was not a message".to_string())?;
for block in message.content() {
if let ContentBlock::ToolUse(tu) = block {
return Ok(Some(document_to_value(tu.input())));
}
}
Ok(None)
}
}