299 lines
8.9 KiB
Rust
299 lines
8.9 KiB
Rust
use std::collections::HashMap;
|
|
|
|
use aws_sdk_bedrockruntime::types::{
|
|
ContentBlock, ConversationRole, InferenceConfiguration, Message as BedrockMessage,
|
|
SystemContentBlock, Tool, ToolConfiguration, ToolInputSchema, ToolResultBlock,
|
|
ToolResultContentBlock, ToolResultStatus, ToolSpecification, ToolUseBlock,
|
|
};
|
|
use aws_smithy_types::Document;
|
|
use serde_json::Value as JsonValue;
|
|
|
|
pub struct ConvertedRequest {
|
|
pub messages: Vec<BedrockMessage>,
|
|
pub system: Vec<SystemContentBlock>,
|
|
pub inference_config: InferenceConfiguration,
|
|
pub tool_config: Option<ToolConfiguration>,
|
|
}
|
|
|
|
pub struct ConversationMessage {
|
|
pub role: MessageRole,
|
|
pub content: MessageContent,
|
|
}
|
|
|
|
pub enum MessageRole {
|
|
User,
|
|
Assistant,
|
|
}
|
|
|
|
pub enum MessageContent {
|
|
Text(String),
|
|
ToolUse {
|
|
tool_use_id: String,
|
|
name: String,
|
|
input: JsonValue,
|
|
},
|
|
ToolResult {
|
|
tool_use_id: String,
|
|
content: String,
|
|
is_error: bool,
|
|
},
|
|
MultiPart(Vec<ContentPart>),
|
|
}
|
|
|
|
pub enum ContentPart {
|
|
Text(String),
|
|
ToolUse {
|
|
tool_use_id: String,
|
|
name: String,
|
|
input: JsonValue,
|
|
},
|
|
ToolResult {
|
|
tool_use_id: String,
|
|
content: String,
|
|
is_error: bool,
|
|
},
|
|
}
|
|
|
|
pub struct ToolDefinition {
|
|
pub name: String,
|
|
pub description: String,
|
|
pub input_schema: JsonValue,
|
|
}
|
|
|
|
pub fn build_converse_request(
|
|
messages: Vec<ConversationMessage>,
|
|
system_prompt: Option<String>,
|
|
tools: Vec<ToolDefinition>,
|
|
max_tokens: i32,
|
|
temperature: Option<f32>,
|
|
top_p: Option<f32>,
|
|
stop_sequences: Option<Vec<String>>,
|
|
) -> ConvertedRequest {
|
|
let bedrock_messages = convert_messages(messages);
|
|
let system = convert_system_prompt(system_prompt);
|
|
let inference_config = build_inference_config(max_tokens, temperature, top_p, stop_sequences);
|
|
let tool_config = build_tool_config(tools);
|
|
|
|
ConvertedRequest {
|
|
messages: bedrock_messages,
|
|
system,
|
|
inference_config,
|
|
tool_config,
|
|
}
|
|
}
|
|
|
|
fn json_to_document(value: JsonValue) -> Document {
|
|
match value {
|
|
JsonValue::Null => Document::Null,
|
|
JsonValue::Bool(b) => Document::Bool(b),
|
|
JsonValue::Number(n) => {
|
|
if let Some(i) = n.as_i64() {
|
|
Document::Number(aws_smithy_types::Number::PosInt(i as u64))
|
|
} else if let Some(f) = n.as_f64() {
|
|
Document::Number(aws_smithy_types::Number::Float(f))
|
|
} else {
|
|
Document::Null
|
|
}
|
|
}
|
|
JsonValue::String(s) => Document::String(s),
|
|
JsonValue::Array(arr) => {
|
|
Document::Array(arr.into_iter().map(json_to_document).collect())
|
|
}
|
|
JsonValue::Object(obj) => {
|
|
let map: HashMap<String, Document> = obj
|
|
.into_iter()
|
|
.map(|(k, v)| (k, json_to_document(v)))
|
|
.collect();
|
|
Document::Object(map)
|
|
}
|
|
}
|
|
}
|
|
|
|
fn convert_messages(messages: Vec<ConversationMessage>) -> Vec<BedrockMessage> {
|
|
let mut result = Vec::new();
|
|
|
|
for msg in messages {
|
|
let role = match msg.role {
|
|
MessageRole::User => ConversationRole::User,
|
|
MessageRole::Assistant => ConversationRole::Assistant,
|
|
};
|
|
|
|
let content_blocks = match msg.content {
|
|
MessageContent::Text(text) => vec![ContentBlock::Text(text)],
|
|
MessageContent::ToolUse {
|
|
tool_use_id,
|
|
name,
|
|
input,
|
|
} => {
|
|
let input_doc = json_to_document(input);
|
|
vec![ContentBlock::ToolUse(
|
|
ToolUseBlock::builder()
|
|
.tool_use_id(tool_use_id)
|
|
.name(name)
|
|
.input(input_doc)
|
|
.build()
|
|
.expect("valid tool use block"),
|
|
)]
|
|
}
|
|
MessageContent::ToolResult {
|
|
tool_use_id,
|
|
content,
|
|
is_error,
|
|
} => {
|
|
let status = if is_error {
|
|
ToolResultStatus::Error
|
|
} else {
|
|
ToolResultStatus::Success
|
|
};
|
|
vec![ContentBlock::ToolResult(
|
|
ToolResultBlock::builder()
|
|
.tool_use_id(tool_use_id)
|
|
.status(status)
|
|
.content(ToolResultContentBlock::Text(content))
|
|
.build()
|
|
.expect("valid tool result block"),
|
|
)]
|
|
}
|
|
MessageContent::MultiPart(parts) => parts
|
|
.into_iter()
|
|
.map(|part| match part {
|
|
ContentPart::Text(text) => ContentBlock::Text(text),
|
|
ContentPart::ToolUse {
|
|
tool_use_id,
|
|
name,
|
|
input,
|
|
} => {
|
|
let input_doc = json_to_document(input);
|
|
ContentBlock::ToolUse(
|
|
ToolUseBlock::builder()
|
|
.tool_use_id(tool_use_id)
|
|
.name(name)
|
|
.input(input_doc)
|
|
.build()
|
|
.expect("valid tool use block"),
|
|
)
|
|
}
|
|
ContentPart::ToolResult {
|
|
tool_use_id,
|
|
content,
|
|
is_error,
|
|
} => {
|
|
let status = if is_error {
|
|
ToolResultStatus::Error
|
|
} else {
|
|
ToolResultStatus::Success
|
|
};
|
|
ContentBlock::ToolResult(
|
|
ToolResultBlock::builder()
|
|
.tool_use_id(tool_use_id)
|
|
.status(status)
|
|
.content(ToolResultContentBlock::Text(content))
|
|
.build()
|
|
.expect("valid tool result block"),
|
|
)
|
|
}
|
|
})
|
|
.collect(),
|
|
};
|
|
|
|
let message = BedrockMessage::builder()
|
|
.role(role)
|
|
.set_content(Some(content_blocks))
|
|
.build()
|
|
.expect("valid message");
|
|
|
|
result.push(message);
|
|
}
|
|
|
|
coalesce_consecutive_roles(result)
|
|
}
|
|
|
|
fn coalesce_consecutive_roles(messages: Vec<BedrockMessage>) -> Vec<BedrockMessage> {
|
|
if messages.is_empty() {
|
|
return messages;
|
|
}
|
|
|
|
let mut result: Vec<BedrockMessage> = Vec::new();
|
|
|
|
for msg in messages {
|
|
let should_merge = result
|
|
.last()
|
|
.map(|last| last.role() == msg.role())
|
|
.unwrap_or(false);
|
|
|
|
if should_merge {
|
|
let last = result.pop().unwrap();
|
|
let mut combined_content: Vec<ContentBlock> = last.content().to_vec();
|
|
combined_content.extend(msg.content().to_vec());
|
|
let merged = BedrockMessage::builder()
|
|
.role(last.role().clone())
|
|
.set_content(Some(combined_content))
|
|
.build()
|
|
.expect("valid merged message");
|
|
result.push(merged);
|
|
} else {
|
|
result.push(msg);
|
|
}
|
|
}
|
|
|
|
result
|
|
}
|
|
|
|
fn convert_system_prompt(system_prompt: Option<String>) -> Vec<SystemContentBlock> {
|
|
match system_prompt {
|
|
Some(prompt) if !prompt.is_empty() => {
|
|
vec![SystemContentBlock::Text(prompt)]
|
|
}
|
|
_ => vec![],
|
|
}
|
|
}
|
|
|
|
fn build_inference_config(
|
|
max_tokens: i32,
|
|
temperature: Option<f32>,
|
|
top_p: Option<f32>,
|
|
stop_sequences: Option<Vec<String>>,
|
|
) -> InferenceConfiguration {
|
|
let mut builder = InferenceConfiguration::builder().max_tokens(max_tokens);
|
|
|
|
if let Some(temp) = temperature {
|
|
builder = builder.temperature(temp);
|
|
}
|
|
if let Some(p) = top_p {
|
|
builder = builder.top_p(p);
|
|
}
|
|
if let Some(stops) = stop_sequences {
|
|
builder = builder.set_stop_sequences(Some(stops));
|
|
}
|
|
|
|
builder.build()
|
|
}
|
|
|
|
fn build_tool_config(tools: Vec<ToolDefinition>) -> Option<ToolConfiguration> {
|
|
if tools.is_empty() {
|
|
return None;
|
|
}
|
|
|
|
let tool_specs: Vec<Tool> = tools
|
|
.into_iter()
|
|
.map(|tool| {
|
|
let input_schema_doc = json_to_document(tool.input_schema);
|
|
Tool::ToolSpec(
|
|
ToolSpecification::builder()
|
|
.name(tool.name)
|
|
.description(tool.description)
|
|
.input_schema(ToolInputSchema::Json(input_schema_doc))
|
|
.build()
|
|
.expect("valid tool spec"),
|
|
)
|
|
})
|
|
.collect();
|
|
|
|
Some(
|
|
ToolConfiguration::builder()
|
|
.set_tools(Some(tool_specs))
|
|
.build()
|
|
.expect("valid tool config"),
|
|
)
|
|
}
|