Add unified models UI and Rig Bedrock runtime

This commit is contained in:
2026-08-04 17:25:19 -05:00
parent a3c68e9c30
commit b0ad07f6f2
41 changed files with 2122 additions and 564 deletions
+1 -1
View File
@@ -4,4 +4,4 @@ mod rig_request;
mod rig_tool;
pub(crate) use provider::ProviderRuntime;
pub(crate) use rig::rig_openai_response_stream;
pub(crate) use rig::{rig_bedrock_response_stream, rig_openai_response_stream};
+111 -19
View File
@@ -11,10 +11,12 @@ use uuid::Uuid;
use warp_multi_agent_api::response_event::stream_finished;
use warp_multi_agent_api::{self as api, ClientAction, ResponseEvent, ToolType};
use super::rig_request::{prepare_rig_turn, PreparedRigTurn};
use super::rig_request::{prepare_bedrock_rig_turn, prepare_rig_turn, PreparedRigTurn};
use super::rig_tool::action_from_tool_call;
use crate::ai::agent::api::{Event, RequestParams, ResponseStream, StreamEvent};
use crate::ai::agent::AIAgentAction;
use crate::ai::bedrock::client::{BedrockClient, BedrockClientConfig};
use crate::ai::bedrock::external_config::ExternalBedrockConfig;
use crate::ai::bedrock::response_translator::{
build_add_agent_output_message, build_append_text, build_create_task, build_stream_init,
build_user_query_message,
@@ -32,6 +34,75 @@ pub(crate) fn rig_openai_response_stream(
cancellation_rx: oneshot::Receiver<()>,
) -> ResponseStream {
let skill_path_origin = params.session_context.skill_path_origin();
let prepared = prepare_rig_turn(&config, params, supported_tools, supported_cli_agent_tools);
let model_id = prepared.request.model.as_str().to_string();
let runtime = OpenAICompatibleRuntime::new(OpenAICompatibleRuntimeConfig {
base_url: config.base_url,
api_key: config.api_key,
model: model_id.clone(),
max_output_tokens: config.max_output_tokens.map(u64::from),
supports_system_messages: config.supports_system_messages,
});
rig_response_stream(
runtime,
prepared,
skill_path_origin,
config.max_input_tokens,
"rig_openai_compatible",
cancellation_rx,
)
}
pub(crate) async fn rig_bedrock_response_stream(
config: BedrockClientConfig,
params: RequestParams,
supported_tools: Vec<ToolType>,
supported_cli_agent_tools: Vec<ToolType>,
cancellation_rx: oneshot::Receiver<()>,
) -> anyhow::Result<ResponseStream> {
let skill_path_origin = params.session_context.skill_path_origin();
let max_context_tokens = params.context_window_limit;
let model = params.model.as_str().to_string();
let max_output_tokens = Some(64_000);
let cross_region_inference = config.cross_region_inference;
let external_config = ExternalBedrockConfig::load();
let prompt_caching = !external_config.disable_prompt_caching;
let client = BedrockClient::from_config(config).await?;
let runtime = client.rig_runtime(
model.clone(),
cross_region_inference,
prompt_caching,
max_output_tokens,
)?;
let prepared = prepare_bedrock_rig_turn(
model,
max_output_tokens,
params,
supported_tools,
supported_cli_agent_tools,
);
Ok(rig_response_stream(
runtime,
prepared,
skill_path_origin,
max_context_tokens,
"rig_bedrock",
cancellation_rx,
))
}
fn rig_response_stream<R>(
runtime: R,
prepared: PreparedRigTurn,
skill_path_origin: ai::skills::SkillPathOrigin,
max_context_tokens: Option<u32>,
stream_type: &'static str,
cancellation_rx: oneshot::Receiver<()>,
) -> ResponseStream
where
R: AgentRuntime + Send + Sync + 'static,
{
let PreparedRigTurn {
task_id,
needs_create_task,
@@ -40,21 +111,12 @@ pub(crate) fn rig_openai_response_stream(
persistent_messages,
tool_result_archive,
messages_sent,
} = prepare_rig_turn(&config, params, supported_tools, supported_cli_agent_tools);
} = prepared;
store_messages_sent(&messages_sent, &persistent_messages);
let conversation_id = turn_request.conversation_id.clone();
let model_id = turn_request.model.as_str().to_string();
let tool_policy = ToolPolicy::new(&turn_request.tools);
let runtime = OpenAICompatibleRuntime::new(OpenAICompatibleRuntimeConfig {
base_url: config.base_url,
api_key: config.api_key,
model: model_id.clone(),
max_output_tokens: config.max_output_tokens.map(u64::from),
supports_system_messages: config.supports_system_messages,
});
let max_context_tokens = config.max_input_tokens;
let stream = async_stream::stream! {
let (control_sender, control) = turn_control();
let start_future = runtime.start_turn(turn_request, control).fuse();
@@ -67,7 +129,7 @@ pub(crate) fn rig_openai_response_stream(
match start_future.await {
Ok(stream) => stream,
Err(error) => {
yield Err(agent_error(error));
yield Err(agent_error(error, stream_type));
return;
}
}
@@ -75,7 +137,7 @@ pub(crate) fn rig_openai_response_stream(
result = start_future => match result {
Ok(stream) => stream,
Err(error) => {
yield Err(agent_error(error));
yield Err(agent_error(error, stream_type));
return;
}
},
@@ -87,6 +149,8 @@ pub(crate) fn rig_openai_response_stream(
let mut current_text_message_id: Option<String> = None;
let mut current_reasoning_message_id: Option<String> = None;
let mut full_text = String::new();
let mut full_reasoning = String::new();
let mut reasoning_signature = None;
let mut proposed_tools = Vec::new();
let mut assistant_history_index = None;
let mut usage = Usage::default();
@@ -106,7 +170,7 @@ pub(crate) fn rig_openai_response_stream(
let event = match event {
Ok(event) => event,
Err(error) => {
yield Err(agent_error(error));
yield Err(agent_error(error, stream_type));
return;
}
};
@@ -133,6 +197,7 @@ pub(crate) fn rig_openai_response_stream(
}
}
AgentEvent::ReasoningDelta { text } => {
full_reasoning.push_str(&text);
if let Some(message_id) = &current_reasoning_message_id {
yield Ok(StreamEvent::Response(build_append_reasoning(&task_id, message_id, &text)));
} else {
@@ -141,6 +206,17 @@ pub(crate) fn rig_openai_response_stream(
current_reasoning_message_id = Some(message_id);
}
}
AgentEvent::ReasoningCompleted { text, signature } => {
if current_reasoning_message_id.is_none() && !text.is_empty() {
let message_id = Uuid::new_v4().to_string();
yield Ok(StreamEvent::Response(build_add_reasoning(&task_id, &message_id, &text)));
current_reasoning_message_id = Some(message_id);
}
if !text.is_empty() {
full_reasoning = text;
}
reasoning_signature = signature;
}
AgentEvent::UsageUpdated { usage: updated } => usage = updated,
AgentEvent::Tool {
event: ToolEvent::Proposed { call },
@@ -148,6 +224,8 @@ pub(crate) fn rig_openai_response_stream(
proposed_tools.push(call.clone());
sync_assistant_turn(
&messages_sent,
&full_reasoning,
reasoning_signature.as_deref(),
&full_text,
&proposed_tools,
&mut assistant_history_index,
@@ -164,7 +242,7 @@ pub(crate) fn rig_openai_response_stream(
yield Err(agent_error(AgentError::new(
galaxy_agent_core::AgentErrorKind::Protocol,
message,
)));
), stream_type));
return;
}
}
@@ -198,6 +276,8 @@ pub(crate) fn rig_openai_response_stream(
}
sync_assistant_turn(
&messages_sent,
&full_reasoning,
reasoning_signature.as_deref(),
&full_text,
&proposed_tools,
&mut assistant_history_index,
@@ -222,7 +302,7 @@ pub(crate) fn rig_openai_response_stream(
yield Err(agent_error(AgentError::new(
galaxy_agent_core::AgentErrorKind::Protocol,
"the provider runtime attempted to execute a tool outside Galaxy's permission boundary",
)));
), stream_type));
return;
}
}
@@ -264,11 +344,22 @@ fn append_tool_result(
fn sync_assistant_turn(
messages_sent: &std::sync::Arc<std::sync::Mutex<Vec<ConversationMessage>>>,
reasoning_text: &str,
reasoning_signature: Option<&str>,
text: &str,
tool_calls: &[ToolCall],
history_index: &mut Option<usize>,
) {
let mut parts = Vec::with_capacity(usize::from(!text.is_empty()) + tool_calls.len());
let has_reasoning = !reasoning_text.is_empty() || reasoning_signature.is_some();
let mut parts = Vec::with_capacity(
usize::from(has_reasoning) + usize::from(!text.is_empty()) + tool_calls.len(),
);
if has_reasoning {
parts.push(ContentPart::Reasoning {
text: reasoning_text.to_string(),
signature: reasoning_signature.map(str::to_string),
});
}
if !text.is_empty() {
parts.push(ContentPart::Text(text.to_string()));
}
@@ -293,6 +384,7 @@ fn sync_assistant_turn(
name,
input,
},
reasoning @ ContentPart::Reasoning { .. } => MessageContent::MultiPart(vec![reasoning]),
ContentPart::Image { .. } | ContentPart::ToolResult { .. } => unreachable!(),
}
} else {
@@ -395,10 +487,10 @@ fn saturating_i32(value: u64) -> i32 {
i32::try_from(value).unwrap_or(i32::MAX)
}
fn agent_error(error: AgentError) -> Arc<AIApiError> {
fn agent_error(error: AgentError, stream_type: &'static str) -> Arc<AIApiError> {
Arc::new(
AIApiError::Stream {
stream_type: "rig_openai_compatible",
stream_type,
source: anyhow::anyhow!(error),
}
.into_quota_limit_if_provider_budget_exhausted(),
+52 -6
View File
@@ -13,7 +13,9 @@ use warp_multi_agent_api::ToolType;
use crate::ai::agent::api::RequestParams;
use crate::ai::agent::{AIAgentContext, AIAgentInput, MCPContext, UserQueryMode};
use crate::ai::bedrock::request_translator::{default_tool_definitions, tool_name_is_supported};
use crate::ai::bedrock::request_translator::{
default_tool_definitions, sanitize_messages_for_bedrock, tool_name_is_supported,
};
use crate::ai::openai::client::OpenAIClientConfig;
use crate::ai::openai::request_translator::sanitize_messages_for_openai;
@@ -32,6 +34,47 @@ pub(crate) fn prepare_rig_turn(
params: RequestParams,
supported_tools: Vec<ToolType>,
supported_cli_agent_tools: Vec<ToolType>,
) -> PreparedRigTurn {
prepare_rig_turn_for_provider(
config.model.clone(),
config.max_output_tokens.map(u64::from),
RigRequestSanitizer::OpenAICompatible,
params,
supported_tools,
supported_cli_agent_tools,
)
}
pub(crate) fn prepare_bedrock_rig_turn(
model: String,
max_output_tokens: Option<u64>,
params: RequestParams,
supported_tools: Vec<ToolType>,
supported_cli_agent_tools: Vec<ToolType>,
) -> PreparedRigTurn {
prepare_rig_turn_for_provider(
Some(model),
max_output_tokens,
RigRequestSanitizer::Bedrock,
params,
supported_tools,
supported_cli_agent_tools,
)
}
#[derive(Clone, Copy)]
enum RigRequestSanitizer {
OpenAICompatible,
Bedrock,
}
fn prepare_rig_turn_for_provider(
model_override: Option<String>,
max_output_tokens: Option<u64>,
sanitizer: RigRequestSanitizer,
params: RequestParams,
supported_tools: Vec<ToolType>,
supported_cli_agent_tools: Vec<ToolType>,
) -> PreparedRigTurn {
let RequestParams {
input,
@@ -70,7 +113,12 @@ pub(crate) fn prepare_rig_turn(
for message in &mut persistent_messages {
message.truncate_tool_results_for_provider_request();
}
sanitize_messages_for_openai(&mut persistent_messages);
match sanitizer {
RigRequestSanitizer::OpenAICompatible => {
sanitize_messages_for_openai(&mut persistent_messages)
}
RigRequestSanitizer::Bedrock => sanitize_messages_for_bedrock(&mut persistent_messages),
}
let mut turn_messages = Vec::new();
if let Some(summary) = progressive_summary {
@@ -91,16 +139,14 @@ pub(crate) fn prepare_rig_turn(
}
turn_messages.extend(persistent_messages.clone());
let model_id = config
.model
.clone()
let model_id = model_override
.filter(|model| !model.is_empty() && model != "auto")
.unwrap_or_else(|| model.as_str().to_string());
let mut request = TurnRequest::new(model_id, turn_messages);
request.conversation_id = conversation_token.map(|token| token.as_str().to_string());
request.system_prompt = Some(system_prompt);
request.tools = tools;
request.max_output_tokens = config.max_output_tokens.map(u64::from);
request.max_output_tokens = max_output_tokens;
PreparedRigTurn {
task_id,
+35 -1
View File
@@ -4,7 +4,7 @@ use std::sync::Arc;
use galaxy_agent_core::{ContentPart, MessageContent, MessageRole, ToolResult, ToolResultStatus};
use warp_multi_agent_api::ToolType;
use super::{input_messages, prepare_rig_turn, tool_definitions};
use super::{input_messages, prepare_bedrock_rig_turn, prepare_rig_turn, tool_definitions};
use crate::ai::agent::api::RequestParams;
use crate::ai::agent::{
AIAgentContext, AIAgentInput, AnyFileContent, FileContext, MCPContext, MCPServer, UserQueryMode,
@@ -114,6 +114,40 @@ fn builds_a_rig_turn_directly_from_galaxy_request_state() {
));
}
#[test]
fn bedrock_rig_turn_uses_bedrock_history_invariants_without_a_proto_round_trip() {
let mut params = RequestParams::new_for_test();
params.message_history = vec![galaxy_agent_core::ConversationMessage {
role: MessageRole::Assistant,
content: MessageContent::Text("Prior assistant message".to_string()),
}];
params.input = vec![user_query("Continue safely")];
let prepared = prepare_bedrock_rig_turn(
"anthropic.claude-test".to_string(),
Some(64_000),
params,
Vec::new(),
Vec::new(),
);
assert_eq!(prepared.request.model.as_str(), "anthropic.claude-test");
assert_eq!(prepared.request.max_output_tokens, Some(64_000));
assert_eq!(
prepared
.request
.messages
.first()
.map(|message| message.role),
Some(MessageRole::User)
);
assert_eq!(
prepared.request.messages.last().map(|message| message.role),
Some(MessageRole::User)
);
assert_eq!(prepared.request.messages, prepared.persistent_messages);
}
#[test]
#[allow(deprecated)]
fn grouped_mcp_tool_names_use_the_installation_id_not_the_display_name() {
+44 -1
View File
@@ -2,7 +2,7 @@ use std::sync::{Arc, Mutex};
use ai::skills::SkillPathOrigin;
use galaxy_agent_core::{
MessageContent, MessageRole, StopReason, ToolCall, ToolResult, ToolResultStatus,
ContentPart, MessageContent, MessageRole, StopReason, ToolCall, ToolResult, ToolResultStatus,
};
use warp_multi_agent_api::response_event::stream_finished;
@@ -134,12 +134,16 @@ fn assistant_history_is_updated_before_fast_tool_execution_can_continue() {
sync_assistant_turn(
&messages,
"",
None,
"I'll inspect both.",
std::slice::from_ref(&first_call),
&mut history_index,
);
sync_assistant_turn(
&messages,
"",
None,
"I'll inspect both.",
&[first_call, second_call],
&mut history_index,
@@ -162,6 +166,43 @@ fn assistant_history_is_updated_before_fast_tool_execution_can_continue() {
);
}
#[test]
fn signed_reasoning_is_persisted_before_the_tool_call() {
let messages = Arc::new(Mutex::new(Vec::new()));
let mut history_index = None;
let call = ToolCall {
id: "call-1".to_string(),
name: "read_files".to_string(),
arguments: serde_json::json!({"files": ["Cargo.toml"]}),
};
sync_assistant_turn(
&messages,
"I should inspect the manifest.",
Some("signed-reasoning"),
"",
std::slice::from_ref(&call),
&mut history_index,
);
let messages = messages.lock().unwrap();
let MessageContent::MultiPart(parts) = &messages[0].content else {
panic!("expected reasoning and tool call parts");
};
assert!(matches!(
parts.as_slice(),
[
ContentPart::Reasoning {
text,
signature: Some(signature),
},
ContentPart::ToolUse { tool_use_id, .. },
] if text == "I should inspect the manifest."
&& signature == "signed-reasoning"
&& tool_use_id == "call-1"
));
}
#[test]
fn rejected_tool_result_is_paired_with_the_assistant_call_in_history() {
let messages = Arc::new(Mutex::new(Vec::new()));
@@ -174,6 +215,8 @@ fn rejected_tool_result_is_paired_with_the_assistant_call_in_history() {
sync_assistant_turn(
&messages,
"",
None,
"",
std::slice::from_ref(&call),
&mut history_index,
);