Significant progress. Performing cleanup now
This commit is contained in:
@@ -1,4 +1,4 @@
|
||||
use std::sync::Arc;
|
||||
use std::sync::{Arc, Mutex};
|
||||
|
||||
use anyhow::Result;
|
||||
use aws_config::BehaviorVersion;
|
||||
@@ -113,6 +113,7 @@ impl BedrockClient {
|
||||
temperature: Option<f32>,
|
||||
cross_region_inference: bool,
|
||||
diagnostic_logger: Option<Arc<BedrockDiagnosticLogger>>,
|
||||
messages_sent: Arc<Mutex<Vec<ConversationMessage>>>,
|
||||
) -> Result<ResponseStream, BedrockError> {
|
||||
let effective_model_id = if cross_region_inference {
|
||||
apply_cross_region_prefix(model_id, &self.region)
|
||||
@@ -192,6 +193,7 @@ impl BedrockClient {
|
||||
task_id.to_string(),
|
||||
needs_create_task,
|
||||
diagnostic_logger,
|
||||
messages_sent,
|
||||
)))
|
||||
}
|
||||
|
||||
|
||||
@@ -1,9 +1,9 @@
|
||||
use std::collections::HashMap;
|
||||
|
||||
use aws_sdk_bedrockruntime::types::{
|
||||
ContentBlock, ConversationRole, InferenceConfiguration, Message as BedrockMessage,
|
||||
SystemContentBlock, Tool, ToolConfiguration, ToolInputSchema, ToolResultBlock,
|
||||
ToolResultContentBlock, ToolResultStatus, ToolSpecification, ToolUseBlock,
|
||||
CachePointBlock, CachePointType, 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;
|
||||
@@ -15,7 +15,7 @@ pub struct ConvertedRequest {
|
||||
pub tool_config: Option<ToolConfiguration>,
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
#[derive(Clone, Debug)]
|
||||
pub struct ConversationMessage {
|
||||
pub role: MessageRole,
|
||||
pub content: MessageContent,
|
||||
@@ -208,7 +208,30 @@ fn convert_messages(messages: Vec<ConversationMessage>) -> Vec<BedrockMessage> {
|
||||
result.push(message);
|
||||
}
|
||||
|
||||
coalesce_consecutive_roles(result)
|
||||
let mut messages = coalesce_consecutive_roles(result);
|
||||
|
||||
// Add a cache point to the second-to-last message (the conversation prefix
|
||||
// that is stable between requests). This allows Bedrock to cache all prior
|
||||
// context and only process the latest message as new input tokens.
|
||||
if messages.len() >= 2 {
|
||||
let cache_idx = messages.len() - 2;
|
||||
let msg = messages.remove(cache_idx);
|
||||
let mut content = msg.content().to_vec();
|
||||
content.push(ContentBlock::CachePoint(
|
||||
CachePointBlock::builder()
|
||||
.r#type(CachePointType::Default)
|
||||
.build()
|
||||
.expect("valid cache point"),
|
||||
));
|
||||
let cached_msg = BedrockMessage::builder()
|
||||
.role(msg.role().clone())
|
||||
.set_content(Some(content))
|
||||
.build()
|
||||
.expect("valid message with cache point");
|
||||
messages.insert(cache_idx, cached_msg);
|
||||
}
|
||||
|
||||
messages
|
||||
}
|
||||
|
||||
fn coalesce_consecutive_roles(messages: Vec<BedrockMessage>) -> Vec<BedrockMessage> {
|
||||
@@ -245,7 +268,15 @@ fn coalesce_consecutive_roles(messages: Vec<BedrockMessage>) -> Vec<BedrockMessa
|
||||
fn convert_system_prompt(system_prompt: Option<String>) -> Vec<SystemContentBlock> {
|
||||
match system_prompt {
|
||||
Some(prompt) if !prompt.is_empty() => {
|
||||
vec![SystemContentBlock::Text(prompt)]
|
||||
vec![
|
||||
SystemContentBlock::Text(prompt),
|
||||
SystemContentBlock::CachePoint(
|
||||
CachePointBlock::builder()
|
||||
.r#type(CachePointType::Default)
|
||||
.build()
|
||||
.expect("valid cache point"),
|
||||
),
|
||||
]
|
||||
}
|
||||
_ => vec![],
|
||||
}
|
||||
@@ -277,7 +308,7 @@ fn build_tool_config(tools: Vec<ToolDefinition>) -> Option<ToolConfiguration> {
|
||||
return None;
|
||||
}
|
||||
|
||||
let tool_specs: Vec<Tool> = tools
|
||||
let mut tool_specs: Vec<Tool> = tools
|
||||
.into_iter()
|
||||
.map(|tool| {
|
||||
let input_schema_doc = json_to_document(tool.input_schema);
|
||||
@@ -292,6 +323,13 @@ fn build_tool_config(tools: Vec<ToolDefinition>) -> Option<ToolConfiguration> {
|
||||
})
|
||||
.collect();
|
||||
|
||||
tool_specs.push(Tool::CachePoint(
|
||||
CachePointBlock::builder()
|
||||
.r#type(CachePointType::Default)
|
||||
.build()
|
||||
.expect("valid cache point"),
|
||||
));
|
||||
|
||||
Some(
|
||||
ToolConfiguration::builder()
|
||||
.set_tools(Some(tool_specs))
|
||||
|
||||
@@ -1,79 +1,366 @@
|
||||
use warp_multi_agent_api as api;
|
||||
|
||||
use super::convert::{ConversationMessage, MessageContent, MessageRole, ToolDefinition};
|
||||
use super::convert::{ConversationMessage, ContentPart, MessageContent, MessageRole, ToolDefinition};
|
||||
|
||||
/// Extract new input messages from the current request and convert them directly
|
||||
/// to ConversationMessage format for the Bedrock message history.
|
||||
/// This extracts UserQuery and ToolCallResult from request.input only.
|
||||
pub fn extract_new_input_messages(request: &api::Request) -> Vec<ConversationMessage> {
|
||||
let mut results = Vec::new();
|
||||
let Some(input) = &request.input else {
|
||||
return results;
|
||||
};
|
||||
let Some(input_type) = &input.r#type else {
|
||||
return results;
|
||||
};
|
||||
|
||||
match input_type {
|
||||
api::request::input::Type::UserInputs(user_inputs) => {
|
||||
let mut tool_results: Vec<ConversationMessage> = Vec::new();
|
||||
for user_input in &user_inputs.inputs {
|
||||
match &user_input.input {
|
||||
Some(api::request::input::user_inputs::user_input::Input::ToolCallResult(
|
||||
result,
|
||||
)) => {
|
||||
if !result.tool_call_id.is_empty() {
|
||||
let content = extract_tool_result_content(result);
|
||||
tool_results.push(ConversationMessage {
|
||||
role: MessageRole::User,
|
||||
content: MessageContent::ToolResult {
|
||||
tool_use_id: result.tool_call_id.clone(),
|
||||
content,
|
||||
is_error: false,
|
||||
},
|
||||
});
|
||||
}
|
||||
}
|
||||
Some(api::request::input::user_inputs::user_input::Input::UserQuery(
|
||||
query,
|
||||
)) => {
|
||||
if !query.query.is_empty() {
|
||||
results.push(ConversationMessage {
|
||||
role: MessageRole::User,
|
||||
content: MessageContent::Text(query.query.clone()),
|
||||
});
|
||||
}
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
if !tool_results.is_empty() {
|
||||
if tool_results.len() == 1 {
|
||||
results.extend(tool_results);
|
||||
} else {
|
||||
let parts: Vec<ContentPart> = tool_results
|
||||
.into_iter()
|
||||
.map(|tr| match tr.content {
|
||||
MessageContent::ToolResult { tool_use_id, content, is_error } => {
|
||||
ContentPart::ToolResult { tool_use_id, content, is_error }
|
||||
}
|
||||
_ => unreachable!(),
|
||||
})
|
||||
.collect();
|
||||
results.push(ConversationMessage {
|
||||
role: MessageRole::User,
|
||||
content: MessageContent::MultiPart(parts),
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
#[allow(deprecated)]
|
||||
api::request::input::Type::UserQuery(query) => {
|
||||
if !query.query.is_empty() {
|
||||
results.push(ConversationMessage {
|
||||
role: MessageRole::User,
|
||||
content: MessageContent::Text(query.query.clone()),
|
||||
});
|
||||
}
|
||||
}
|
||||
api::request::input::Type::InitProjectRules(_) => {
|
||||
results.push(ConversationMessage {
|
||||
role: MessageRole::User,
|
||||
content: MessageContent::Text(
|
||||
"Initialize this project. Analyze the codebase structure and files, \
|
||||
generate an AGENTS.md file documenting project conventions and setup \
|
||||
instructions, and offer to create a development environment configuration. \
|
||||
Use the available tools to inspect the project before responding."
|
||||
.to_string(),
|
||||
),
|
||||
});
|
||||
}
|
||||
api::request::input::Type::CreateEnvironment(env) => {
|
||||
let repo_info = if env.repo_paths.is_empty() {
|
||||
String::new()
|
||||
} else {
|
||||
format!(" Repositories: {}", env.repo_paths.join(", "))
|
||||
};
|
||||
results.push(ConversationMessage {
|
||||
role: MessageRole::User,
|
||||
content: MessageContent::Text(format!(
|
||||
"Create a development environment for this project. \
|
||||
Set up necessary dependencies, configuration files, and tooling.{repo_info}"
|
||||
)),
|
||||
});
|
||||
}
|
||||
api::request::input::Type::CreateNewProject(project) => {
|
||||
results.push(ConversationMessage {
|
||||
role: MessageRole::User,
|
||||
content: MessageContent::Text(format!(
|
||||
"Create a new project: {}",
|
||||
project.query,
|
||||
)),
|
||||
});
|
||||
}
|
||||
api::request::input::Type::CloneRepository(repo) => {
|
||||
results.push(ConversationMessage {
|
||||
role: MessageRole::User,
|
||||
content: MessageContent::Text(format!(
|
||||
"Clone the repository at {} and set it up for development.",
|
||||
repo.url,
|
||||
)),
|
||||
});
|
||||
}
|
||||
api::request::input::Type::ResumeConversation(_) => {
|
||||
results.push(ConversationMessage {
|
||||
role: MessageRole::User,
|
||||
content: MessageContent::Text(
|
||||
"Continue where we left off. Review the conversation history and proceed with the next steps."
|
||||
.to_string(),
|
||||
),
|
||||
});
|
||||
}
|
||||
api::request::input::Type::CodeReview(_) => {
|
||||
results.push(ConversationMessage {
|
||||
role: MessageRole::User,
|
||||
content: MessageContent::Text(
|
||||
"Review the following code changes and provide detailed feedback on correctness, style, and potential issues."
|
||||
.to_string(),
|
||||
),
|
||||
});
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
|
||||
for msg in &results {
|
||||
let desc = match &msg.content {
|
||||
MessageContent::Text(t) => format!("Text({}chars)", t.len()),
|
||||
MessageContent::ToolResult { tool_use_id, .. } => format!("ToolResult({})", tool_use_id),
|
||||
MessageContent::MultiPart(parts) => format!("MultiPart({} parts)", parts.len()),
|
||||
_ => "Other".to_string(),
|
||||
};
|
||||
log::info!("[bedrock] New input message: role={:?}, content={}", msg.role, desc);
|
||||
}
|
||||
|
||||
results
|
||||
}
|
||||
|
||||
/// For the Bedrock direct path: inject all input messages (user queries and tool call
|
||||
/// results) into the task's messages so they persist in conversation history for future
|
||||
/// requests. Without this, inputs are lost after the current request cycle because they
|
||||
/// only exist in `request.input` and are never stored in `task_context.tasks[].messages`.
|
||||
pub fn inject_input_messages_into_task(request: &mut api::Request) {
|
||||
let input_messages: Vec<api::Message> = extract_input_messages(request);
|
||||
if input_messages.is_empty() {
|
||||
return;
|
||||
}
|
||||
|
||||
log::info!(
|
||||
"[bedrock] Injecting {} input messages into task history",
|
||||
input_messages.len()
|
||||
);
|
||||
|
||||
if let Some(task_context) = &mut request.task_context {
|
||||
if let Some(task) = task_context.tasks.first_mut() {
|
||||
task.messages.extend(input_messages);
|
||||
} else {
|
||||
// No task exists yet — create one to hold the messages
|
||||
let task_id = uuid::Uuid::new_v4().to_string();
|
||||
task_context.tasks.push(api::Task {
|
||||
id: task_id,
|
||||
messages: input_messages,
|
||||
..Default::default()
|
||||
});
|
||||
}
|
||||
} else {
|
||||
let task_id = uuid::Uuid::new_v4().to_string();
|
||||
request.task_context = Some(api::request::TaskContext {
|
||||
tasks: vec![api::Task {
|
||||
id: task_id,
|
||||
messages: input_messages,
|
||||
..Default::default()
|
||||
}],
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
fn extract_input_messages(request: &api::Request) -> Vec<api::Message> {
|
||||
let mut results = Vec::new();
|
||||
let Some(input) = &request.input else {
|
||||
return results;
|
||||
};
|
||||
let Some(input_type) = &input.r#type else {
|
||||
return results;
|
||||
};
|
||||
|
||||
let task_id = request
|
||||
.task_context
|
||||
.as_ref()
|
||||
.and_then(|tc| tc.tasks.first())
|
||||
.map(|t| t.id.clone())
|
||||
.unwrap_or_default();
|
||||
|
||||
match input_type {
|
||||
api::request::input::Type::UserInputs(user_inputs) => {
|
||||
for user_input in &user_inputs.inputs {
|
||||
match &user_input.input {
|
||||
Some(api::request::input::user_inputs::user_input::Input::ToolCallResult(
|
||||
result,
|
||||
)) => {
|
||||
if !result.tool_call_id.is_empty() {
|
||||
let content = extract_tool_result_content(result);
|
||||
results.push(api::Message {
|
||||
id: uuid::Uuid::new_v4().to_string(),
|
||||
task_id: task_id.clone(),
|
||||
request_id: String::new(),
|
||||
timestamp: None,
|
||||
server_message_data: String::new(),
|
||||
citations: vec![],
|
||||
message: Some(api::message::Message::ToolCallResult(
|
||||
api::message::ToolCallResult {
|
||||
tool_call_id: result.tool_call_id.clone(),
|
||||
context: None,
|
||||
result: Some(
|
||||
api::message::tool_call_result::Result::Server(
|
||||
api::message::tool_call_result::ServerResult {
|
||||
serialized_result: content,
|
||||
},
|
||||
),
|
||||
),
|
||||
},
|
||||
)),
|
||||
});
|
||||
}
|
||||
}
|
||||
Some(api::request::input::user_inputs::user_input::Input::UserQuery(
|
||||
query,
|
||||
)) => {
|
||||
if !query.query.is_empty() {
|
||||
results.push(api::Message {
|
||||
id: uuid::Uuid::new_v4().to_string(),
|
||||
task_id: task_id.clone(),
|
||||
request_id: String::new(),
|
||||
timestamp: None,
|
||||
server_message_data: String::new(),
|
||||
citations: vec![],
|
||||
message: Some(api::message::Message::UserQuery(
|
||||
api::message::UserQuery {
|
||||
query: query.query.clone(),
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
});
|
||||
}
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
}
|
||||
api::request::input::Type::UserQuery(query) => {
|
||||
if !query.query.is_empty() {
|
||||
results.push(api::Message {
|
||||
id: uuid::Uuid::new_v4().to_string(),
|
||||
task_id: task_id.clone(),
|
||||
request_id: String::new(),
|
||||
timestamp: None,
|
||||
server_message_data: String::new(),
|
||||
citations: vec![],
|
||||
message: Some(api::message::Message::UserQuery(
|
||||
api::message::UserQuery {
|
||||
query: query.query.clone(),
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
});
|
||||
}
|
||||
}
|
||||
api::request::input::Type::ToolCallResult(result) => {
|
||||
if !result.tool_call_id.is_empty() {
|
||||
let content = extract_tool_result_content(result);
|
||||
results.push(api::Message {
|
||||
id: uuid::Uuid::new_v4().to_string(),
|
||||
task_id: task_id.clone(),
|
||||
request_id: String::new(),
|
||||
timestamp: None,
|
||||
server_message_data: String::new(),
|
||||
citations: vec![],
|
||||
message: Some(api::message::Message::ToolCallResult(
|
||||
api::message::ToolCallResult {
|
||||
tool_call_id: result.tool_call_id.clone(),
|
||||
context: None,
|
||||
result: Some(
|
||||
api::message::tool_call_result::Result::Server(
|
||||
api::message::tool_call_result::ServerResult {
|
||||
serialized_result: content,
|
||||
},
|
||||
),
|
||||
),
|
||||
},
|
||||
)),
|
||||
});
|
||||
}
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
|
||||
results
|
||||
}
|
||||
|
||||
pub fn extract_messages_from_request(request: &api::Request) -> Vec<ConversationMessage> {
|
||||
let mut messages = Vec::new();
|
||||
|
||||
if let Some(task_context) = &request.task_context {
|
||||
log::info!(
|
||||
"[bedrock-debug] extract_messages: {} tasks in task_context",
|
||||
task_context.tasks.len()
|
||||
);
|
||||
for task in &task_context.tasks {
|
||||
log::info!(
|
||||
"[bedrock-debug] extract_messages: task '{}' has {} messages",
|
||||
task.id,
|
||||
task.messages.len()
|
||||
);
|
||||
for msg in &task.messages {
|
||||
let msg_type = msg.message.as_ref().map(|m| match m {
|
||||
api::message::Message::UserQuery(_) => "UserQuery",
|
||||
api::message::Message::AgentOutput(_) => "AgentOutput",
|
||||
api::message::Message::ToolCall(_) => "ToolCall",
|
||||
api::message::Message::ToolCallResult(_) => "ToolCallResult",
|
||||
api::message::Message::AgentReasoning(_) => "AgentReasoning",
|
||||
_ => "Other",
|
||||
}).unwrap_or("None");
|
||||
log::info!(
|
||||
"[bedrock-debug] extract_messages: msg id='{}' type={}",
|
||||
msg.id,
|
||||
msg_type
|
||||
);
|
||||
if let Some(converted) = convert_proto_message(msg) {
|
||||
messages.push(converted);
|
||||
}
|
||||
}
|
||||
}
|
||||
} else {
|
||||
log::warn!("[bedrock-debug] extract_messages: NO task_context in request!");
|
||||
}
|
||||
|
||||
if let Some(input) = &request.input {
|
||||
if let Some(input_type) = &input.r#type {
|
||||
#[allow(deprecated)]
|
||||
match input_type {
|
||||
api::request::input::Type::UserInputs(user_inputs) => {
|
||||
for user_input in &user_inputs.inputs {
|
||||
if let Some(input_variant) = &user_input.input {
|
||||
match input_variant {
|
||||
api::request::input::user_inputs::user_input::Input::UserQuery(
|
||||
query,
|
||||
) => {
|
||||
if !query.query.is_empty() {
|
||||
messages.push(ConversationMessage {
|
||||
role: MessageRole::User,
|
||||
content: MessageContent::Text(query.query.clone()),
|
||||
});
|
||||
}
|
||||
}
|
||||
api::request::input::user_inputs::user_input::Input::ToolCallResult(
|
||||
result,
|
||||
) => {
|
||||
let content = extract_tool_result_content(result);
|
||||
if !result.tool_call_id.is_empty() {
|
||||
messages.push(ConversationMessage {
|
||||
role: MessageRole::User,
|
||||
content: MessageContent::ToolResult {
|
||||
tool_use_id: result.tool_call_id.clone(),
|
||||
content,
|
||||
is_error: false,
|
||||
},
|
||||
});
|
||||
}
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
api::request::input::Type::UserQuery(query) => {
|
||||
if !query.query.is_empty() {
|
||||
messages.push(ConversationMessage {
|
||||
role: MessageRole::User,
|
||||
content: MessageContent::Text(query.query.clone()),
|
||||
});
|
||||
}
|
||||
}
|
||||
api::request::input::Type::ToolCallResult(result) => {
|
||||
let content = extract_tool_result_content(result);
|
||||
if !result.tool_call_id.is_empty() {
|
||||
messages.push(ConversationMessage {
|
||||
role: MessageRole::User,
|
||||
content: MessageContent::ToolResult {
|
||||
tool_use_id: result.tool_call_id.clone(),
|
||||
content,
|
||||
is_error: false,
|
||||
},
|
||||
});
|
||||
}
|
||||
}
|
||||
// UserInputs (UserQuery + ToolCallResult) are already injected
|
||||
// into task messages by inject_input_messages_into_task().
|
||||
api::request::input::Type::UserInputs(_) => {}
|
||||
api::request::input::Type::UserQuery(_) => {}
|
||||
api::request::input::Type::ToolCallResult(_) => {}
|
||||
api::request::input::Type::InitProjectRules(_) => {
|
||||
messages.push(ConversationMessage {
|
||||
role: MessageRole::User,
|
||||
@@ -159,6 +446,14 @@ pub fn extract_messages_from_request(request: &api::Request) -> Vec<Conversation
|
||||
|
||||
ensure_starts_with_user_message(&mut messages);
|
||||
ensure_tool_results_paired(&mut messages);
|
||||
log::info!(
|
||||
"[bedrock-debug] extract_messages: TOTAL {} messages to send to Bedrock (User={}, Assistant={}, ToolResult={}, ToolUse={})",
|
||||
messages.len(),
|
||||
messages.iter().filter(|m| m.role == MessageRole::User && matches!(&m.content, MessageContent::Text(_))).count(),
|
||||
messages.iter().filter(|m| m.role == MessageRole::Assistant && matches!(&m.content, MessageContent::Text(_))).count(),
|
||||
messages.iter().filter(|m| matches!(&m.content, MessageContent::ToolResult { .. })).count(),
|
||||
messages.iter().filter(|m| matches!(&m.content, MessageContent::ToolUse { .. })).count(),
|
||||
);
|
||||
messages
|
||||
}
|
||||
|
||||
@@ -331,11 +626,11 @@ fn tool_definition_for_name(name: &str) -> ToolDefinition {
|
||||
},
|
||||
"read_files" => ToolDefinition {
|
||||
name: "read_files".to_string(),
|
||||
description: "Read the contents of one or more files.".to_string(),
|
||||
description: "Read the contents of one or more files. ALWAYS pass all files you need in a single call rather than making multiple separate calls.".to_string(),
|
||||
input_schema: serde_json::json!({
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"files": { "type": "array", "items": { "type": "string" }, "description": "File paths to read" }
|
||||
"files": { "type": "array", "items": { "type": "string" }, "description": "File paths to read. Include ALL files you need in one call for efficiency." }
|
||||
},
|
||||
"required": ["files"]
|
||||
}),
|
||||
@@ -353,11 +648,11 @@ fn tool_definition_for_name(name: &str) -> ToolDefinition {
|
||||
},
|
||||
"grep" => ToolDefinition {
|
||||
name: "grep".to_string(),
|
||||
description: "Search for patterns in files using grep.".to_string(),
|
||||
description: "Search for patterns in files. Pass all search patterns in one call.".to_string(),
|
||||
input_schema: serde_json::json!({
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"queries": { "type": "array", "items": { "type": "string" }, "description": "Search patterns" },
|
||||
"queries": { "type": "array", "items": { "type": "string" }, "description": "Search patterns. Include ALL patterns you need in one call." },
|
||||
"path": { "type": "string", "description": "Directory to search in" }
|
||||
},
|
||||
"required": ["queries"]
|
||||
@@ -365,7 +660,7 @@ fn tool_definition_for_name(name: &str) -> ToolDefinition {
|
||||
},
|
||||
"file_glob" => ToolDefinition {
|
||||
name: "file_glob".to_string(),
|
||||
description: "Find files matching glob patterns.".to_string(),
|
||||
description: "Find files matching glob patterns. Pass all patterns in one call.".to_string(),
|
||||
input_schema: serde_json::json!({
|
||||
"type": "object",
|
||||
"properties": {
|
||||
@@ -544,6 +839,9 @@ fn format_tool_call_result(result: &api::message::ToolCallResult) -> String {
|
||||
_ => "Read files completed.".to_string(),
|
||||
}
|
||||
}
|
||||
api::message::tool_call_result::Result::Server(server_result) => {
|
||||
server_result.serialized_result.clone()
|
||||
}
|
||||
_ => "Tool completed successfully.".to_string(),
|
||||
}
|
||||
} else {
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
use std::sync::{Arc, Mutex};
|
||||
|
||||
use futures::StreamExt;
|
||||
use serde_json::json;
|
||||
use std::path::PathBuf;
|
||||
@@ -553,6 +555,7 @@ impl AgentSimulation {
|
||||
8192,
|
||||
None,
|
||||
false,
|
||||
Arc::new(Mutex::new(Vec::new())),
|
||||
)
|
||||
.await
|
||||
.expect("converse_stream should succeed");
|
||||
@@ -1137,6 +1140,7 @@ async fn test_reasoning_model_produces_substantial_output() {
|
||||
8192,
|
||||
None,
|
||||
false,
|
||||
Arc::new(Mutex::new(Vec::new())),
|
||||
)
|
||||
.await
|
||||
.expect("converse_stream should succeed");
|
||||
@@ -1257,6 +1261,7 @@ async fn test_event_sequence_matches_controller_expectations() {
|
||||
100,
|
||||
None,
|
||||
false,
|
||||
Arc::new(Mutex::new(Vec::new())),
|
||||
)
|
||||
.await
|
||||
.expect("should connect");
|
||||
@@ -1369,6 +1374,7 @@ async fn test_followup_turn_does_not_send_create_task() {
|
||||
100,
|
||||
None,
|
||||
false,
|
||||
Arc::new(Mutex::new(Vec::new())),
|
||||
)
|
||||
.await
|
||||
.expect("should connect");
|
||||
@@ -1446,6 +1452,7 @@ async fn run_slash_command_test(
|
||||
4096,
|
||||
None,
|
||||
true,
|
||||
Arc::new(Mutex::new(Vec::new())),
|
||||
)
|
||||
.await
|
||||
.expect("stream should connect");
|
||||
@@ -1671,6 +1678,7 @@ async fn test_slash_resume_conversation() {
|
||||
256,
|
||||
None,
|
||||
true,
|
||||
Arc::new(Mutex::new(Vec::new())),
|
||||
)
|
||||
.await
|
||||
.expect("stream should connect");
|
||||
@@ -1780,6 +1788,7 @@ async fn test_empty_messages_safety_check() {
|
||||
100,
|
||||
None,
|
||||
true,
|
||||
Arc::new(Mutex::new(Vec::new())),
|
||||
)
|
||||
.await
|
||||
.expect("safety fallback message should work");
|
||||
@@ -1929,6 +1938,7 @@ async fn test_full_proto_round_trip_with_tool_history() {
|
||||
1024,
|
||||
None,
|
||||
true,
|
||||
Arc::new(Mutex::new(Vec::new())),
|
||||
)
|
||||
.await;
|
||||
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
use std::sync::{Arc, Mutex};
|
||||
|
||||
use futures::StreamExt;
|
||||
use serde_json::json;
|
||||
|
||||
@@ -62,6 +64,7 @@ async fn collect_stream_output(
|
||||
8192,
|
||||
None,
|
||||
false,
|
||||
Arc::new(Mutex::new(Vec::new())),
|
||||
)
|
||||
.await
|
||||
.expect("converse_stream should succeed");
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
use std::sync::Arc;
|
||||
use std::sync::{Arc, Mutex};
|
||||
|
||||
use aws_sdk_bedrockruntime::operation::converse_stream::ConverseStreamOutput;
|
||||
use aws_sdk_bedrockruntime::types::{
|
||||
@@ -13,6 +13,7 @@ use warp_multi_agent_api::{self as api, ClientAction, ResponseEvent};
|
||||
use crate::ai::agent::api::Event;
|
||||
use crate::server::server_api::AIApiError;
|
||||
|
||||
use super::convert::{ContentPart, MessageContent, MessageRole, ConversationMessage};
|
||||
use super::diagnostic::BedrockDiagnosticLogger;
|
||||
|
||||
pub fn bedrock_stream_to_response_events(
|
||||
@@ -20,6 +21,7 @@ pub fn bedrock_stream_to_response_events(
|
||||
task_id: String,
|
||||
needs_create_task: bool,
|
||||
diagnostic_logger: Option<Arc<BedrockDiagnosticLogger>>,
|
||||
messages_sent: Arc<Mutex<Vec<ConversationMessage>>>,
|
||||
) -> BoxStream<'static, Event> {
|
||||
let request_id = Uuid::new_v4().to_string();
|
||||
let conversation_id = Uuid::new_v4().to_string();
|
||||
@@ -55,6 +57,10 @@ pub fn bedrock_stream_to_response_events(
|
||||
let mut output_tokens: i32 = 0;
|
||||
let mut stop_reason = stream_finished::Reason::Done(stream_finished::Done {});
|
||||
|
||||
// Track full assistant text and tool calls for bedrock_message_history
|
||||
let mut history_text = String::new();
|
||||
let mut history_tool_calls: Vec<ContentPart> = Vec::new();
|
||||
|
||||
let mut event_count: u32 = 0;
|
||||
loop {
|
||||
match output.stream.recv().await {
|
||||
@@ -109,6 +115,7 @@ pub fn bedrock_stream_to_response_events(
|
||||
match d {
|
||||
ContentBlockDelta::Text(text) => {
|
||||
log::info!("[bedrock-debug] Event #{event_count}: TextDelta ({} chars): {:?}", text.len(), &text[..text.len().min(80)]);
|
||||
history_text.push_str(text);
|
||||
if text_flushed {
|
||||
let msg_id = current_text_message_id.as_ref().unwrap();
|
||||
let append = build_append_text(
|
||||
@@ -149,6 +156,14 @@ pub fn bedrock_stream_to_response_events(
|
||||
log::info!("[bedrock-debug] Event #{event_count}: ContentBlockStop (tool_use_id={:?})", if current_tool_use_id.is_empty() { "none" } else { ¤t_tool_use_id });
|
||||
if !current_tool_use_id.is_empty() {
|
||||
log::debug!("[bedrock] Tool call complete: {} ({})", current_tool_name, current_tool_use_id);
|
||||
// Track for bedrock_message_history
|
||||
let input_json: serde_json::Value = serde_json::from_str(¤t_tool_input_json)
|
||||
.unwrap_or(serde_json::Value::Object(serde_json::Map::new()));
|
||||
history_tool_calls.push(ContentPart::ToolUse {
|
||||
tool_use_id: current_tool_use_id.clone(),
|
||||
name: current_tool_name.clone(),
|
||||
input: input_json,
|
||||
});
|
||||
if let Some(ref logger) = diagnostic_logger {
|
||||
logger.log_stream_event(&format!(
|
||||
"ToolCall: name={}, id={}, input={}",
|
||||
@@ -242,6 +257,49 @@ pub fn bedrock_stream_to_response_events(
|
||||
}
|
||||
|
||||
log::info!("[bedrock] Stream finished: input_tokens={input_tokens}, output_tokens={output_tokens}");
|
||||
|
||||
// Build and store the assistant message into bedrock_messages_sent
|
||||
// so the controller can persist it as part of conversation history.
|
||||
{
|
||||
let mut parts: Vec<ContentPart> = Vec::new();
|
||||
if !history_text.is_empty() {
|
||||
parts.push(ContentPart::Text(history_text));
|
||||
}
|
||||
parts.extend(history_tool_calls);
|
||||
|
||||
if !parts.is_empty() {
|
||||
let assistant_msg = if parts.len() == 1 {
|
||||
match parts.remove(0) {
|
||||
ContentPart::Text(t) => ConversationMessage {
|
||||
role: MessageRole::Assistant,
|
||||
content: MessageContent::Text(t),
|
||||
},
|
||||
ContentPart::ToolUse { tool_use_id, name, input } => ConversationMessage {
|
||||
role: MessageRole::Assistant,
|
||||
content: MessageContent::ToolUse { tool_use_id, name, input },
|
||||
},
|
||||
other => ConversationMessage {
|
||||
role: MessageRole::Assistant,
|
||||
content: MessageContent::MultiPart(vec![other]),
|
||||
},
|
||||
}
|
||||
} else {
|
||||
ConversationMessage {
|
||||
role: MessageRole::Assistant,
|
||||
content: MessageContent::MultiPart(parts),
|
||||
}
|
||||
};
|
||||
|
||||
if let Ok(mut sent) = messages_sent.lock() {
|
||||
sent.push(assistant_msg);
|
||||
log::info!(
|
||||
"[bedrock] Stored assistant message in history. Total messages: {}",
|
||||
sent.len()
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if let Some(ref logger) = diagnostic_logger {
|
||||
let stop_reason_str = match &stop_reason {
|
||||
stream_finished::Reason::Done(_) => "EndTurn",
|
||||
|
||||
Reference in New Issue
Block a user