Files
galaxy/app/src/ai/bedrock/convert.rs
T

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"),
)
}