Rebasing, going about this another way
This commit is contained in:
@@ -0,0 +1,298 @@
|
||||
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"),
|
||||
)
|
||||
}
|
||||
Reference in New Issue
Block a user