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
Generated
+100
View File
@@ -1745,17 +1745,23 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "635d23afda0a6ab48d666c4d447c4873e8d1e83518a2be2093122397e50b838e" checksum = "635d23afda0a6ab48d666c4d447c4873e8d1e83518a2be2093122397e50b838e"
dependencies = [ dependencies = [
"aws-smithy-async", "aws-smithy-async",
"aws-smithy-protocol-test",
"aws-smithy-runtime-api", "aws-smithy-runtime-api",
"aws-smithy-types", "aws-smithy-types",
"bytes",
"h2", "h2",
"http 1.5.0", "http 1.5.0",
"http-body 1.1.0",
"hyper", "hyper",
"hyper-rustls", "hyper-rustls",
"hyper-util", "hyper-util",
"indexmap 2.14.0",
"pin-project-lite", "pin-project-lite",
"rustls", "rustls",
"rustls-native-certs", "rustls-native-certs",
"rustls-pki-types", "rustls-pki-types",
"serde",
"serde_json",
"tokio", "tokio",
"tokio-rustls", "tokio-rustls",
"tower", "tower",
@@ -1782,6 +1788,25 @@ dependencies = [
"aws-smithy-runtime-api", "aws-smithy-runtime-api",
] ]
[[package]]
name = "aws-smithy-protocol-test"
version = "0.64.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "f76511a0e223ce78deb6a78b8afebda99cb737cfbc8a58d96dcb190f012dd40a"
dependencies = [
"assert-json-diff",
"aws-smithy-runtime-api",
"base64-simd",
"cbor-diag",
"ciborium",
"http 0.2.12",
"pretty_assertions",
"regex-lite",
"roxmltree 0.14.1",
"serde_json",
"thiserror 2.0.19",
]
[[package]] [[package]]
name = "aws-smithy-query" name = "aws-smithy-query"
version = "0.62.0" version = "0.62.0"
@@ -2698,6 +2723,25 @@ dependencies = [
"cipher", "cipher",
] ]
[[package]]
name = "cbor-diag"
version = "0.1.12"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "dc245b6ecd09b23901a4fbad1ad975701fd5061ceaef6afa93a2d70605a64429"
dependencies = [
"bs58",
"chrono",
"data-encoding",
"half",
"nom 7.1.3",
"num-bigint",
"num-rational",
"num-traits",
"separator",
"url",
"uuid",
]
[[package]] [[package]]
name = "cc" name = "cc"
version = "1.4.0" version = "1.4.0"
@@ -4386,6 +4430,12 @@ dependencies = [
"syn 2.0.119", "syn 2.0.119",
] ]
[[package]]
name = "diff"
version = "0.1.13"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "56254986775e3233ffa9c4d7d3faaf6d36a2c09d30b20687e9f88bc8bafc16c8"
[[package]] [[package]]
name = "difflib" name = "difflib"
version = "0.4.0" version = "0.4.0"
@@ -5897,9 +5947,13 @@ version = "0.1.0"
dependencies = [ dependencies = [
"async-stream", "async-stream",
"async-trait", "async-trait",
"aws-sdk-bedrockruntime",
"aws-smithy-http-client",
"base64 0.22.1",
"bytes", "bytes",
"futures", "futures",
"galaxy_agent_core", "galaxy_agent_core",
"rig-bedrock",
"rig-core", "rig-core",
"serde_json", "serde_json",
"tokio", "tokio",
@@ -11556,6 +11610,16 @@ version = "0.3.1"
source = "registry+https://github.com/rust-lang/crates.io-index" source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "e8cf8e6a8aa66ce33f63993ffc4ea4271eb5b0530a9002db8455ea6050c77bfa" checksum = "e8cf8e6a8aa66ce33f63993ffc4ea4271eb5b0530a9002db8455ea6050c77bfa"
[[package]]
name = "pretty_assertions"
version = "1.4.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "3ae130e2f271fbc2ac3a40fb1d07180839cdbbe443c7a27e1e3c13c5cac0116d"
dependencies = [
"diff",
"yansi",
]
[[package]] [[package]]
name = "prettyplease" name = "prettyplease"
version = "0.2.37" version = "0.2.37"
@@ -12778,6 +12842,27 @@ dependencies = [
"bytemuck", "bytemuck",
] ]
[[package]]
name = "rig-bedrock"
version = "0.40.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "10e8ee8d206e78398eca2db97cb0cf27f43c903a99c60a730bc7a0cfeaf3ee83"
dependencies = [
"async-stream",
"aws-config",
"aws-sdk-bedrockruntime",
"aws-smithy-types",
"base64 0.22.1",
"rig-core",
"rig-derive",
"schemars 1.2.2",
"serde",
"serde_json",
"tokio",
"tracing",
"uuid",
]
[[package]] [[package]]
name = "rig-core" name = "rig-core"
version = "0.40.0" version = "0.40.0"
@@ -12942,6 +13027,15 @@ dependencies = [
"syn 2.0.119", "syn 2.0.119",
] ]
[[package]]
name = "roxmltree"
version = "0.14.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "921904a62e410e37e215c40381b7117f830d9d89ba60ab5236170541dd25646b"
dependencies = [
"xmlparser",
]
[[package]] [[package]]
name = "roxmltree" name = "roxmltree"
version = "0.20.0" version = "0.20.0"
@@ -13525,6 +13619,12 @@ version = "0.6.0"
source = "registry+https://github.com/rust-lang/crates.io-index" source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "cd0b0ec5f1c1ca621c432a25813d8d60c88abe6d3e08a3eb9cf37d97a0fe3d73" checksum = "cd0b0ec5f1c1ca621c432a25813d8d60c88abe6d3e08a3eb9cf37d97a0fe3d73"
[[package]]
name = "separator"
version = "0.4.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "f97841a747eef040fcd2e7b3b9a220a7205926e60488e673d9e4926d27772ce5"
[[package]] [[package]]
name = "seq-macro" name = "seq-macro"
version = "0.3.6" version = "0.3.6"
+3
View File
@@ -138,6 +138,8 @@ async-stream = "0.3.5"
async-task = "4.2.0" async-task = "4.2.0"
async-trait = "0.1.89" async-trait = "0.1.89"
async-fs = "2.1.2" async-fs = "2.1.2"
aws-sdk-bedrockruntime = { version = "1", default-features = false, features = ["default-https-client", "rt-tokio"] }
aws-smithy-http-client = { version = "1", features = ["test-util"] }
backtrace = "0.3.76" backtrace = "0.3.76"
base64 = "0.22" base64 = "0.22"
bincode = "1.3.3" bincode = "1.3.3"
@@ -260,6 +262,7 @@ reqwest = { version = "0.13", features = [
] } ] }
reqwest-eventsource = { package = "aha-reqwest-eventsource", version = "0.1" } reqwest-eventsource = { package = "aha-reqwest-eventsource", version = "0.1" }
rig-core = "=0.40.0" rig-core = "=0.40.0"
rig-bedrock = "=0.40.0"
resvg = "0.47.0" resvg = "0.47.0"
rust-embed = { version = "8.7.0", features = ["include-exclude"] } rust-embed = { version = "8.7.0", features = ["include-exclude"] }
rustc-hash = "2.1.1" rustc-hash = "2.1.1"
+1 -1
View File
@@ -328,7 +328,7 @@ tracing-subscriber.workspace = true
# AWS SDK (loading credentials for BYO LLM) # AWS SDK (loading credentials for BYO LLM)
aws-config = { version = "1.8.16", features = ["credentials-login"] } aws-config = { version = "1.8.16", features = ["credentials-login"] }
aws-credential-types = "1" aws-credential-types = "1"
aws-sdk-bedrockruntime = { version = "1", default-features = false, features = ["default-https-client", "rt-tokio"] } aws-sdk-bedrockruntime.workspace = true
aws-sdk-sts = { version = "1", default-features = false, features = ["default-https-client", "rt-tokio"] } aws-sdk-sts = { version = "1", default-features = false, features = ["default-https-client", "rt-tokio"] }
aws-smithy-types = "1" aws-smithy-types = "1"
aws-types = "1" aws-types = "1"
+26 -2
View File
@@ -28,8 +28,8 @@ pub async fn generate_multi_agent_output(
redaction::redact_inputs(&mut params.input); redaction::redact_inputs(&mut params.input);
} }
if let ProviderConfig::OpenAI(config) = &provider_config { match &provider_config {
if config.use_rig { ProviderConfig::OpenAI(config) if config.use_rig => {
return Ok(crate::ai::runtime::rig_openai_response_stream( return Ok(crate::ai::runtime::rig_openai_response_stream(
config.clone(), config.clone(),
params, params,
@@ -38,6 +38,30 @@ pub async fn generate_multi_agent_output(
cancellation_rx, cancellation_rx,
)); ));
} }
ProviderConfig::Bedrock(config) if config.use_rig => {
return match crate::ai::runtime::rig_bedrock_response_stream(
config.clone(),
params,
supported_tools,
supported_cli_agent_tools,
cancellation_rx,
)
.await
{
Ok(stream) => Ok(stream),
Err(error) => {
log::error!("[rig/bedrock] Runtime error: {error}");
let error = Arc::new(crate::server::server_api::AIApiError::Stream {
stream_type: "rig_bedrock",
source: error,
});
let (sender, receiver) = async_channel::unbounded();
let _ = sender.send(Err(error)).await;
Ok(Box::pin(receiver))
}
};
}
ProviderConfig::OpenAI(_) | ProviderConfig::Bedrock(_) | ProviderConfig::None => {}
} }
let mut logging_metadata = HashMap::new(); let mut logging_metadata = HashMap::new();
+27 -2
View File
@@ -5,6 +5,8 @@ use aws_config::BehaviorVersion;
use aws_credential_types::provider::ProvideCredentials; use aws_credential_types::provider::ProvideCredentials;
use aws_sdk_bedrockruntime::config::Region; use aws_sdk_bedrockruntime::config::Region;
use aws_sdk_bedrockruntime::Client as BedrockRuntimeClient; use aws_sdk_bedrockruntime::Client as BedrockRuntimeClient;
use galaxy_agent_core::AgentError;
use galaxy_agent_rig::{BedrockRigConfig, BedrockRuntime};
use super::convert::{build_converse_request, CachingConfig, ConversationMessage, ToolDefinition}; use super::convert::{build_converse_request, CachingConfig, ConversationMessage, ToolDefinition};
use super::diagnostic::BedrockDiagnosticLogger; use super::diagnostic::BedrockDiagnosticLogger;
@@ -38,6 +40,7 @@ pub struct BedrockClientConfig {
pub secret_access_key: String, pub secret_access_key: String,
pub session_token: Option<String>, pub session_token: Option<String>,
pub cross_region_inference: bool, pub cross_region_inference: bool,
pub use_rig: bool,
} }
impl BedrockClientConfig { impl BedrockClientConfig {
@@ -141,8 +144,8 @@ impl BedrockClient {
match provider.provide_credentials().await { match provider.provide_credentials().await {
Ok(creds) => { Ok(creds) => {
log::info!( log::info!(
"[bedrock] Resolved AWS credentials successfully: access_key_id={:?}, has_session_token={}, expiry={:?}", "[bedrock] Resolved AWS credentials successfully: has_access_key_id={}, has_session_token={}, expiry={:?}",
creds.access_key_id(), !creds.access_key_id().is_empty(),
creds.session_token().is_some(), creds.session_token().is_some(),
creds.expiry(), creds.expiry(),
); );
@@ -168,6 +171,28 @@ impl BedrockClient {
}) })
} }
/// Builds the Phase 4 Rig runtime from the AWS SDK client whose region and
/// credentials Galaxy already resolved. This does not change production
/// routing; callers opt in only after the Bedrock parity suite passes.
pub fn rig_runtime(
&self,
model: String,
cross_region_inference: bool,
prompt_caching: bool,
max_output_tokens: Option<u64>,
) -> Result<BedrockRuntime, AgentError> {
BedrockRuntime::from_aws_client(
self.runtime_client.clone(),
BedrockRigConfig {
model,
region: self.region.clone(),
cross_region_inference,
prompt_caching,
max_output_tokens,
},
)
}
#[allow(clippy::too_many_arguments)] #[allow(clippy::too_many_arguments)]
pub async fn converse_stream( pub async fn converse_stream(
&self, &self,
+12 -2
View File
@@ -3,8 +3,9 @@ use std::collections::HashMap;
use aws_sdk_bedrockruntime::types::{ use aws_sdk_bedrockruntime::types::{
CachePointBlock, CachePointType, CacheTtl, ContentBlock, ConversationRole, ImageBlock, CachePointBlock, CachePointType, CacheTtl, ContentBlock, ConversationRole, ImageBlock,
ImageFormat, ImageSource, InferenceConfiguration, Message as BedrockMessage, ImageFormat, ImageSource, InferenceConfiguration, Message as BedrockMessage,
SystemContentBlock, Tool, ToolConfiguration, ToolInputSchema, ToolResultBlock, ReasoningContentBlock, ReasoningTextBlock, SystemContentBlock, Tool, ToolConfiguration,
ToolResultContentBlock, ToolResultStatus, ToolSpecification, ToolUseBlock, ToolInputSchema, ToolResultBlock, ToolResultContentBlock, ToolResultStatus, ToolSpecification,
ToolUseBlock,
}; };
use aws_smithy_types::{Blob, Document}; use aws_smithy_types::{Blob, Document};
use serde_json::Value as JsonValue; use serde_json::Value as JsonValue;
@@ -148,6 +149,15 @@ fn convert_messages(
.into_iter() .into_iter()
.map(|part| match part { .map(|part| match part {
ContentPart::Text(text) => ContentBlock::Text(text), ContentPart::Text(text) => ContentBlock::Text(text),
ContentPart::Reasoning { text, signature } => {
ContentBlock::ReasoningContent(ReasoningContentBlock::ReasoningText(
ReasoningTextBlock::builder()
.text(text)
.set_signature(signature)
.build()
.expect("valid reasoning text block"),
))
}
ContentPart::Image { data, mime_type } => image_content_block(data, &mime_type), ContentPart::Image { data, mime_type } => image_content_block(data, &mime_type),
ContentPart::ToolUse { ContentPart::ToolUse {
tool_use_id, tool_use_id,
+8
View File
@@ -492,6 +492,14 @@ fn serialize_messages(messages: &[ConversationMessage]) -> JsonValue {
super::convert::ContentPart::Text(t) => { super::convert::ContentPart::Text(t) => {
serde_json::json!({"type": "text", "text": t}) serde_json::json!({"type": "text", "text": t})
} }
super::convert::ContentPart::Reasoning { text, signature } => {
serde_json::json!({
"type": "reasoning",
"char_length": text.len(),
"has_signature": signature.is_some(),
"text": "REDACTED",
})
}
super::convert::ContentPart::Image { data, mime_type } => { super::convert::ContentPart::Image { data, mime_type } => {
serde_json::json!({ serde_json::json!({
"type": "image", "type": "image",
+1
View File
@@ -269,6 +269,7 @@ fn get_test_config() -> Option<BedrockClientConfig> {
secret_access_key: String::new(), secret_access_key: String::new(),
session_token: None, session_token: None,
cross_region_inference: false, cross_region_inference: false,
use_rig: false,
}) })
} }
+1
View File
@@ -144,6 +144,7 @@ fn parse_claude_code_model_map(
model_id: arn, model_id: arn,
display_name, display_name,
vision_supported: true, vision_supported: true,
use_rig: false,
} }
}) })
.collect() .collect()
+1
View File
@@ -24,6 +24,7 @@ fn get_test_config() -> Option<BedrockClientConfig> {
secret_access_key: String::new(), secret_access_key: String::new(),
session_token: None, session_token: None,
cross_region_inference: false, cross_region_inference: false,
use_rig: false,
}) })
} }
+29
View File
@@ -129,6 +129,7 @@ pub fn get_effective_models(user_models: &[BedrockModelConfig]) -> Vec<BedrockMo
model_id: m.model_id.to_string(), model_id: m.model_id.to_string(),
display_name: m.display_name.to_string(), display_name: m.display_name.to_string(),
vision_supported: m.vision_supported, vision_supported: m.vision_supported,
use_rig: false,
}) })
.collect(); .collect();
for default in defaults { for default in defaults {
@@ -145,10 +146,38 @@ pub fn get_effective_models(user_models: &[BedrockModelConfig]) -> Vec<BedrockMo
model_id: m.model_id.to_string(), model_id: m.model_id.to_string(),
display_name: m.display_name.to_string(), display_name: m.display_name.to_string(),
vision_supported: m.vision_supported, vision_supported: m.vision_supported,
use_rig: false,
}) })
.collect() .collect()
} }
pub fn configured_model_uses_rig(
selected_model_id: &str,
configured_models: &[BedrockModelConfig],
region: &str,
cross_region_inference: bool,
) -> bool {
let selected_model_id = strip_context_marker(selected_model_id);
configured_models.iter().any(|model| {
if !model.use_rig {
return false;
}
let configured_model_id = strip_context_marker(&model.model_id);
if configured_model_id == selected_model_id {
return true;
}
galaxy_agent_rig::resolve_bedrock_model_id(&model.model_id, region, cross_region_inference)
.is_ok_and(|resolved| strip_context_marker(&resolved) == selected_model_id)
})
}
fn strip_context_marker(model_id: &str) -> &str {
model_id
.strip_suffix("[1m]")
.or_else(|| model_id.strip_suffix("[1M]"))
.unwrap_or(model_id)
}
pub fn apply_cross_region_prefix(model_id: &str, region: &str) -> String { pub fn apply_cross_region_prefix(model_id: &str, region: &str) -> String {
if model_id.starts_with("arn:") { if model_id.starts_with("arn:") {
return model_id.to_string(); return model_id.to_string();
+32 -2
View File
@@ -82,8 +82,8 @@ fn test_cross_region_prefix_unknown_region() {
fn test_get_effective_models_empty_returns_defaults() { fn test_get_effective_models_empty_returns_defaults() {
let models = get_effective_models(&[]); let models = get_effective_models(&[]);
assert_eq!(models.len(), DEFAULT_BEDROCK_MODELS.len()); assert_eq!(models.len(), DEFAULT_BEDROCK_MODELS.len());
assert_eq!(models[0].model_id, "anthropic.claude-opus-4-6[1m]"); assert_eq!(models[0].model_id, "us.anthropic.claude-opus-4-6-v1[1m]");
assert_eq!(models[0].display_name, "Claude Opus 4.6"); assert_eq!(models[0].display_name, "Claude Opus 4.6 (1M)");
} }
#[test] #[test]
@@ -92,6 +92,7 @@ fn test_get_effective_models_custom_overrides() {
model_id: "custom.model-v1:0".to_string(), model_id: "custom.model-v1:0".to_string(),
display_name: "Custom Model".to_string(), display_name: "Custom Model".to_string(),
vision_supported: false, vision_supported: false,
use_rig: true,
}]; }];
let models = get_effective_models(&custom); let models = get_effective_models(&custom);
assert_eq!(models.len(), 1); assert_eq!(models.len(), 1);
@@ -103,3 +104,32 @@ fn test_cross_region_prefix_skips_arn() {
let arn = "arn:aws:bedrock:us-east-1:156729649053:application-inference-profile/fma5bsw4oxhy"; let arn = "arn:aws:bedrock:us-east-1:156729649053:application-inference-profile/fma5bsw4oxhy";
assert_eq!(apply_cross_region_prefix(arn, "us-east-1"), arn); assert_eq!(apply_cross_region_prefix(arn, "us-east-1"), arn);
} }
#[test]
fn rig_opt_in_matches_context_markers_and_resolved_inference_profiles() {
let configured = vec![BedrockModelConfig {
model_id: "anthropic.claude-test[1m]".to_string(),
display_name: "Claude Test".to_string(),
vision_supported: false,
use_rig: true,
}];
assert!(configured_model_uses_rig(
"us.anthropic.claude-test",
&configured,
"us-east-1",
true,
));
assert!(configured_model_uses_rig(
"anthropic.claude-test[1M]",
&configured,
"us-east-1",
false,
));
assert!(!configured_model_uses_rig(
"anthropic.other-model",
&configured,
"us-east-1",
false,
));
}
+15 -2
View File
@@ -750,6 +750,7 @@ fn persist_input_images_on_latest_user_message(
Some(api::input_context::Image { data, mime_type }) Some(api::input_context::Image { data, mime_type })
} }
Some(ContentPart::Text(_)) Some(ContentPart::Text(_))
| Some(ContentPart::Reasoning { .. })
| Some(ContentPart::ToolUse { .. }) | Some(ContentPart::ToolUse { .. })
| Some(ContentPart::ToolResult { .. }) | Some(ContentPart::ToolResult { .. })
| None => None, | None => None,
@@ -873,11 +874,17 @@ fn is_pure_tool_result(content: &MessageContent) -> bool {
fn strip_tool_result_parts(content: &mut MessageContent) { fn strip_tool_result_parts(content: &mut MessageContent) {
if let MessageContent::MultiPart(parts) = content { if let MessageContent::MultiPart(parts) = content {
parts.retain(|p| !matches!(p, ContentPart::ToolResult { .. })); parts.retain(|p| !matches!(p, ContentPart::ToolResult { .. }));
if parts.len() == 1 && !matches!(parts.first(), Some(ContentPart::Image { .. })) { if parts.len() == 1
&& !matches!(
parts.first(),
Some(ContentPart::Image { .. } | ContentPart::Reasoning { .. })
)
{
let part = parts.remove(0); let part = parts.remove(0);
*content = match part { *content = match part {
ContentPart::Text(t) => MessageContent::Text(t), ContentPart::Text(t) => MessageContent::Text(t),
ContentPart::Image { .. } => unreachable!(), ContentPart::Image { .. } => unreachable!(),
ContentPart::Reasoning { .. } => unreachable!(),
ContentPart::ToolUse { ContentPart::ToolUse {
tool_use_id, tool_use_id,
name, name,
@@ -916,11 +923,17 @@ fn strip_orphaned_tool_result_parts(
ContentPart::ToolResult { tool_use_id, .. } => valid_ids.contains(tool_use_id), ContentPart::ToolResult { tool_use_id, .. } => valid_ids.contains(tool_use_id),
_ => true, _ => true,
}); });
if parts.len() == 1 && !matches!(parts.first(), Some(ContentPart::Image { .. })) { if parts.len() == 1
&& !matches!(
parts.first(),
Some(ContentPart::Image { .. } | ContentPart::Reasoning { .. })
)
{
let part = parts.remove(0); let part = parts.remove(0);
*content = match part { *content = match part {
ContentPart::Text(t) => MessageContent::Text(t), ContentPart::Text(t) => MessageContent::Text(t),
ContentPart::Image { .. } => unreachable!(), ContentPart::Image { .. } => unreachable!(),
ContentPart::Reasoning { .. } => unreachable!(),
ContentPart::ToolUse { ContentPart::ToolUse {
tool_use_id, tool_use_id,
name, name,
+7
View File
@@ -180,6 +180,13 @@ fn describe_message_content(content: &crate::ai::bedrock::convert::MessageConten
.iter() .iter()
.map(|p| match p { .map(|p| match p {
ContentPart::Text(t) => format!("Text({})", t.len()), ContentPart::Text(t) => format!("Text({})", t.len()),
ContentPart::Reasoning { text, signature } => {
format!(
"Reasoning({}chars,signed={})",
text.len(),
signature.is_some()
)
}
ContentPart::Image { data, mime_type } => { ContentPart::Image { data, mime_type } => {
format!("Image({mime_type},{}bytes)", data.len()) format!("Image({mime_type},{}bytes)", data.len())
} }
+7
View File
@@ -4346,6 +4346,7 @@ impl BlocklistAIController {
secret_access_key: settings.bedrock_secret_access_key.value().clone(), secret_access_key: settings.bedrock_secret_access_key.value().clone(),
session_token: None, session_token: None,
cross_region_inference: *settings.bedrock_cross_region_inference.value(), cross_region_inference: *settings.bedrock_cross_region_inference.value(),
use_rig: false,
}; };
if let ::ai::api_keys::AwsCredentialsState::Loaded { credentials, .. } = if let ::ai::api_keys::AwsCredentialsState::Loaded { credentials, .. } =
@@ -4431,6 +4432,9 @@ impl BlocklistAIController {
.iter() .iter()
.map(|p| match p { .map(|p| match p {
ContentPart::Text(t) => t.clone(), ContentPart::Text(t) => t.clone(),
ContentPart::Reasoning { text, .. } => {
format!("[Reasoning] {text}")
}
ContentPart::Image { .. } => "[Image attachment]".to_string(), ContentPart::Image { .. } => "[Image attachment]".to_string(),
ContentPart::ToolUse { name, input, .. } => { ContentPart::ToolUse { name, input, .. } => {
format!("[Tool: {}] {}", name, input) format!("[Tool: {}] {}", name, input)
@@ -4540,6 +4544,9 @@ impl BlocklistAIController {
.iter() .iter()
.map(|p| match p { .map(|p| match p {
ContentPart::Text(t) => (t.len() / 4) as u32, ContentPart::Text(t) => (t.len() / 4) as u32,
ContentPart::Reasoning { text, .. } => {
(text.len() / 4) as u32
}
ContentPart::Image { .. } => 1_600, ContentPart::Image { .. } => 1_600,
ContentPart::ToolUse { input, .. } => { ContentPart::ToolUse { input, .. } => {
(input.to_string().len() / 4) as u32 (input.to_string().len() / 4) as u32
@@ -243,15 +243,33 @@ impl ResponseStream {
// Fall back to Bedrock // Fall back to Bedrock
if *settings.bedrock_enabled.value() { if *settings.bedrock_enabled.value() {
let auth_method = *settings.bedrock_auth_method.value(); let auth_method = *settings.bedrock_auth_method.value();
let region = settings.bedrock_region.value().clone();
let cross_region_inference = *settings.bedrock_cross_region_inference.value();
let mut use_rig = crate::ai::bedrock::models::configured_model_uses_rig(
model_id,
settings.bedrock_models.value(),
&region,
cross_region_inference,
);
if use_rig
&& crate::ai::bedrock::external_config::ExternalBedrockConfig::load()
.enable_prompt_caching_1h
{
log::warn!(
"[rig/bedrock] Using the compatibility runtime because Rig does not yet expose Bedrock's one-hour cache TTL"
);
use_rig = false;
}
let api_key_manager = ::ai::api_keys::ApiKeyManager::as_ref(ctx); let api_key_manager = ::ai::api_keys::ApiKeyManager::as_ref(ctx);
let mut config = BedrockClientConfig { let mut config = BedrockClientConfig {
auth_method, auth_method,
profile: settings.bedrock_profile.value().clone(), profile: settings.bedrock_profile.value().clone(),
region: settings.bedrock_region.value().clone(), region,
access_key_id: settings.bedrock_access_key_id.value().clone(), access_key_id: settings.bedrock_access_key_id.value().clone(),
secret_access_key: settings.bedrock_secret_access_key.value().clone(), secret_access_key: settings.bedrock_secret_access_key.value().clone(),
session_token: None, session_token: None,
cross_region_inference: *settings.bedrock_cross_region_inference.value(), cross_region_inference,
use_rig,
}; };
if let ::ai::api_keys::AwsCredentialsState::Loaded { credentials, .. } = if let ::ai::api_keys::AwsCredentialsState::Loaded { credentials, .. } =
@@ -45,6 +45,7 @@ impl View for ContextWindowView {
.iter() .iter()
.map(|p| match p { .map(|p| match p {
ContentPart::Text(t) => t.len(), ContentPart::Text(t) => t.len(),
ContentPart::Reasoning { text, .. } => text.len(),
ContentPart::Image { .. } => 6_400, ContentPart::Image { .. } => 6_400,
ContentPart::ToolUse { input, .. } => input.to_string().len(), ContentPart::ToolUse { input, .. } => input.to_string().len(),
ContentPart::ToolResult { content, .. } => content.len(), ContentPart::ToolResult { content, .. } => content.len(),
@@ -112,6 +113,14 @@ impl View for ContextWindowView {
ContentPart::Text(t) => { ContentPart::Text(t) => {
out.push_str(&format!("[Part {} Text] {}\n", pi, t)); out.push_str(&format!("[Part {} Text] {}\n", pi, t));
} }
ContentPart::Reasoning { text, signature } => {
out.push_str(&format!(
"[Part {} Reasoning] signed={}\n{}\n",
pi,
signature.is_some(),
text
));
}
ContentPart::Image { data, mime_type } => { ContentPart::Image { data, mime_type } => {
out.push_str(&format!( out.push_str(&format!(
"[Part {} Image] mime_type={}, bytes={}\n", "[Part {} Image] mime_type={}, bytes={}\n",
+1
View File
@@ -186,6 +186,7 @@ impl CrosscheckReviewer {
secret_access_key: settings.bedrock_secret_access_key.value().clone(), secret_access_key: settings.bedrock_secret_access_key.value().clone(),
session_token: None, session_token: None,
cross_region_inference: *settings.bedrock_cross_region_inference.value(), cross_region_inference: *settings.bedrock_cross_region_inference.value(),
use_rig: false,
}; };
if let ::ai::api_keys::AwsCredentialsState::Loaded { credentials, .. } = if let ::ai::api_keys::AwsCredentialsState::Loaded { credentials, .. } =
+131 -2
View File
@@ -721,6 +721,7 @@ impl LLMPreferences {
model_id: default.model_id.to_string(), model_id: default.model_id.to_string(),
display_name: default.display_name.to_string(), display_name: default.display_name.to_string(),
vision_supported: default.vision_supported, vision_supported: default.vision_supported,
use_rig: false,
}); });
added = true; added = true;
} }
@@ -1234,6 +1235,80 @@ impl LLMPreferences {
); );
} }
/// Explicitly refreshes the models for one entry in the OpenAI-compatible
/// provider registry. Unlike the legacy endpoint refresh, this is only
/// called from a user action so configured remote endpoints are never
/// contacted merely because Galaxy started.
#[cfg(not(target_family = "wasm"))]
pub fn fetch_openai_provider_models(
&mut self,
provider_index: usize,
ctx: &mut ModelContext<Self>,
) {
let settings = AISettings::as_ref(ctx);
if !*settings.openai_enabled.value() {
return;
}
let Some(provider) = settings
.openai_providers
.value()
.get(provider_index)
.cloned()
else {
return;
};
if provider.base_url.trim().is_empty() {
return;
}
let requested_base_url = provider.base_url;
let api_key = provider.api_key.filter(|key| !key.is_empty());
let request_base_url = requested_base_url.clone();
let _ = ctx.spawn(
async move {
let base = request_base_url.trim_end_matches('/');
let client = reqwest::Client::builder()
.timeout(std::time::Duration::from_secs(10))
.build()
.unwrap_or_default();
if let Some(models) =
fetch_from_litellm_model_info(base, api_key.as_deref(), &client).await
{
return models;
}
fetch_from_openai_models(base, api_key.as_deref(), &client).await
},
move |_, discovered_models, ctx| {
if discovered_models.is_empty() {
return;
}
AISettings::handle(ctx).update(ctx, |settings, ctx| {
let mut providers = settings.openai_providers.value().clone();
let Some(provider) = providers.get_mut(provider_index) else {
return;
};
// Do not apply a response to an entry that was edited or
// reordered while its discovery request was in flight.
if provider.base_url != requested_base_url {
return;
}
provider.models =
merge_discovered_provider_models(&provider.models, discovered_models);
if let Err(err) = settings.openai_providers.set_value(providers, ctx) {
report_error!(err.context("Failed to persist discovered provider models"));
}
});
},
);
}
/// Returns the `LLMInfo` for the base LLM to be used for an Agent Mode request. /// Returns the `LLMInfo` for the base LLM to be used for an Agent Mode request.
pub fn get_active_base_model<'a>( pub fn get_active_base_model<'a>(
&'a self, &'a self,
@@ -1994,6 +2069,52 @@ fn openai_model_context_size(model: &OpenAIModelConfig) -> u32 {
model.max_input_tokens.unwrap_or(model.context_size) model.max_input_tokens.unwrap_or(model.context_size)
} }
/// Merges endpoint metadata into a provider's configured models without
/// discarding local routing choices or manually configured models.
#[cfg(not(target_family = "wasm"))]
fn merge_discovered_provider_models(
existing_models: &[OpenAIModelConfig],
discovered_models: Vec<OpenAIModelConfig>,
) -> Vec<OpenAIModelConfig> {
let mut merged = Vec::with_capacity(discovered_models.len() + existing_models.len());
let mut discovered_ids = HashSet::new();
for mut discovered in discovered_models {
if !discovered_ids.insert(discovered.model_id.clone()) {
continue;
}
if let Some(existing) = existing_models
.iter()
.find(|model| model.model_id == discovered.model_id)
{
discovered.display_name = existing.display_name.clone();
discovered.use_rig = existing.use_rig;
if existing.supports_system_messages.is_some() {
discovered.supports_system_messages = existing.supports_system_messages;
}
if discovered.provider.is_none() {
discovered.provider = existing.provider.clone();
}
} else {
discovered.use_rig = true;
}
if discovered.model_id.starts_with("codex-gpt-") {
discovered.supports_system_messages = Some(false);
}
merged.push(discovered);
}
merged.extend(
existing_models
.iter()
.filter(|model| !discovered_ids.contains(&model.model_id))
.cloned(),
);
merged
}
#[cfg(not(target_family = "wasm"))] #[cfg(not(target_family = "wasm"))]
fn openai_model_context_window(model: &OpenAIModelConfig) -> LLMContextWindow { fn openai_model_context_window(model: &OpenAIModelConfig) -> LLMContextWindow {
let context_size = openai_model_context_size(model); let context_size = openai_model_context_size(model);
@@ -2118,7 +2239,11 @@ async fn fetch_from_litellm_model_info(
max_output_tokens, max_output_tokens,
provider, provider,
use_rig: false, use_rig: false,
supports_system_messages: model_info["supports_system_messages"].as_bool(), supports_system_messages: if model_name.starts_with("codex-gpt-") {
Some(false)
} else {
model_info["supports_system_messages"].as_bool()
},
}) })
}) })
.collect(); .collect();
@@ -2241,7 +2366,11 @@ async fn fetch_from_openai_models(
max_output_tokens, max_output_tokens,
provider, provider,
use_rig: false, use_rig: false,
supports_system_messages: m["supports_system_messages"].as_bool(), supports_system_messages: if id.starts_with("codex-gpt-") {
Some(false)
} else {
m["supports_system_messages"].as_bool()
},
}) })
}) })
.collect(); .collect();
+65 -1
View File
@@ -10,7 +10,7 @@ use crate::network::NetworkStatus;
use crate::server::cloud_objects::update_manager::UpdateManager; use crate::server::cloud_objects::update_manager::UpdateManager;
use crate::server::server_api::ServerApiProvider; use crate::server::server_api::ServerApiProvider;
use crate::server::sync_queue::SyncQueue; use crate::server::sync_queue::SyncQueue;
use crate::settings::{OpenAIModelConfig, OpenAIProviderConfig}; use crate::settings::OpenAIModelConfig;
use crate::test_util::settings::initialize_settings_for_tests; use crate::test_util::settings::initialize_settings_for_tests;
use crate::workspaces::team_tester::TeamTesterStatus; use crate::workspaces::team_tester::TeamTesterStatus;
use crate::workspaces::user_workspaces::UserWorkspaces; use crate::workspaces::user_workspaces::UserWorkspaces;
@@ -138,3 +138,67 @@ fn llm_info_round_trip_serializes_and_deserializes() {
assert_eq!(info, round_tripped); assert_eq!(info, round_tripped);
} }
fn openai_model(model_id: &str) -> OpenAIModelConfig {
OpenAIModelConfig {
model_id: model_id.to_string(),
display_name: model_id.to_string(),
vision_supported: false,
context_size: 200_000,
max_input_tokens: None,
max_output_tokens: None,
provider: None,
use_rig: false,
supports_system_messages: None,
}
}
#[test]
fn provider_discovery_preserves_local_model_overrides() {
let mut existing = openai_model("codex-gpt-5.6-sol-xhigh");
existing.display_name = "My Codex".to_string();
existing.context_size = 100_000;
existing.provider = Some("openai".to_string());
existing.use_rig = true;
// Even stale or incorrect endpoint metadata must not opt ChatGPT-backed
// Codex models back into the system role.
existing.supports_system_messages = Some(true);
let mut discovered = openai_model("codex-gpt-5.6-sol-xhigh");
discovered.display_name = "Codex from endpoint".to_string();
discovered.context_size = 400_000;
discovered.max_output_tokens = Some(32_000);
discovered.supports_system_messages = Some(true);
let merged = merge_discovered_provider_models(&[existing], vec![discovered]);
assert_eq!(merged.len(), 1);
assert_eq!(merged[0].display_name, "My Codex");
assert_eq!(merged[0].context_size, 400_000);
assert_eq!(merged[0].max_output_tokens, Some(32_000));
assert!(merged[0].use_rig);
assert_eq!(merged[0].supports_system_messages, Some(false));
assert_eq!(merged[0].provider.as_deref(), Some("openai"));
}
#[test]
fn codex_models_reject_system_messages_even_with_stale_true_metadata() {
let mut model = openai_model("codex-gpt-5.6-sol-xhigh");
model.supports_system_messages = Some(true);
assert!(!model.supports_system_messages());
}
#[test]
fn provider_discovery_enables_rig_for_new_models_and_keeps_manual_models() {
let manual = openai_model("manual-model");
let discovered = openai_model("codex-gpt-new");
let merged = merge_discovered_provider_models(&[manual], vec![discovered]);
assert_eq!(merged.len(), 2);
assert_eq!(merged[0].model_id, "codex-gpt-new");
assert!(merged[0].use_rig);
assert_eq!(merged[0].supports_system_messages, Some(false));
assert_eq!(merged[1].model_id, "manual-model");
}
+9
View File
@@ -121,6 +121,9 @@ fn convert_user_message(content: MessageContent) -> ConvertedMessages {
ContentPart::Text(text) => { ContentPart::Text(text) => {
user_content_parts.push(UserContentPart::Text(text)); user_content_parts.push(UserContentPart::Text(text));
} }
ContentPart::Reasoning { text, .. } => {
user_content_parts.push(UserContentPart::Text(text));
}
ContentPart::Image { data, mime_type } => { ContentPart::Image { data, mime_type } => {
user_content_parts.push(UserContentPart::Image { data, mime_type }); user_content_parts.push(UserContentPart::Image { data, mime_type });
} }
@@ -201,6 +204,12 @@ fn convert_assistant_message(content: MessageContent) -> ConvertedMessages {
} }
text_content.push_str(&text); text_content.push_str(&text);
} }
ContentPart::Reasoning { text, .. } => {
if !text_content.is_empty() {
text_content.push('\n');
}
text_content.push_str(&text);
}
ContentPart::ToolUse { ContentPart::ToolUse {
tool_use_id, tool_use_id,
name, name,
+1 -1
View File
@@ -4,4 +4,4 @@ mod rig_request;
mod rig_tool; mod rig_tool;
pub(crate) use provider::ProviderRuntime; 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::response_event::stream_finished;
use warp_multi_agent_api::{self as api, ClientAction, ResponseEvent, ToolType}; 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 super::rig_tool::action_from_tool_call;
use crate::ai::agent::api::{Event, RequestParams, ResponseStream, StreamEvent}; use crate::ai::agent::api::{Event, RequestParams, ResponseStream, StreamEvent};
use crate::ai::agent::AIAgentAction; 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::{ use crate::ai::bedrock::response_translator::{
build_add_agent_output_message, build_append_text, build_create_task, build_stream_init, build_add_agent_output_message, build_append_text, build_create_task, build_stream_init,
build_user_query_message, build_user_query_message,
@@ -32,6 +34,75 @@ pub(crate) fn rig_openai_response_stream(
cancellation_rx: oneshot::Receiver<()>, cancellation_rx: oneshot::Receiver<()>,
) -> ResponseStream { ) -> ResponseStream {
let skill_path_origin = params.session_context.skill_path_origin(); 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 { let PreparedRigTurn {
task_id, task_id,
needs_create_task, needs_create_task,
@@ -40,21 +111,12 @@ pub(crate) fn rig_openai_response_stream(
persistent_messages, persistent_messages,
tool_result_archive, tool_result_archive,
messages_sent, messages_sent,
} = prepare_rig_turn(&config, params, supported_tools, supported_cli_agent_tools); } = prepared;
store_messages_sent(&messages_sent, &persistent_messages); store_messages_sent(&messages_sent, &persistent_messages);
let conversation_id = turn_request.conversation_id.clone(); let conversation_id = turn_request.conversation_id.clone();
let model_id = turn_request.model.as_str().to_string(); let model_id = turn_request.model.as_str().to_string();
let tool_policy = ToolPolicy::new(&turn_request.tools); 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 stream = async_stream::stream! {
let (control_sender, control) = turn_control(); let (control_sender, control) = turn_control();
let start_future = runtime.start_turn(turn_request, control).fuse(); 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 { match start_future.await {
Ok(stream) => stream, Ok(stream) => stream,
Err(error) => { Err(error) => {
yield Err(agent_error(error)); yield Err(agent_error(error, stream_type));
return; return;
} }
} }
@@ -75,7 +137,7 @@ pub(crate) fn rig_openai_response_stream(
result = start_future => match result { result = start_future => match result {
Ok(stream) => stream, Ok(stream) => stream,
Err(error) => { Err(error) => {
yield Err(agent_error(error)); yield Err(agent_error(error, stream_type));
return; return;
} }
}, },
@@ -87,6 +149,8 @@ pub(crate) fn rig_openai_response_stream(
let mut current_text_message_id: Option<String> = None; let mut current_text_message_id: Option<String> = None;
let mut current_reasoning_message_id: Option<String> = None; let mut current_reasoning_message_id: Option<String> = None;
let mut full_text = String::new(); 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 proposed_tools = Vec::new();
let mut assistant_history_index = None; let mut assistant_history_index = None;
let mut usage = Usage::default(); let mut usage = Usage::default();
@@ -106,7 +170,7 @@ pub(crate) fn rig_openai_response_stream(
let event = match event { let event = match event {
Ok(event) => event, Ok(event) => event,
Err(error) => { Err(error) => {
yield Err(agent_error(error)); yield Err(agent_error(error, stream_type));
return; return;
} }
}; };
@@ -133,6 +197,7 @@ pub(crate) fn rig_openai_response_stream(
} }
} }
AgentEvent::ReasoningDelta { text } => { AgentEvent::ReasoningDelta { text } => {
full_reasoning.push_str(&text);
if let Some(message_id) = &current_reasoning_message_id { if let Some(message_id) = &current_reasoning_message_id {
yield Ok(StreamEvent::Response(build_append_reasoning(&task_id, message_id, &text))); yield Ok(StreamEvent::Response(build_append_reasoning(&task_id, message_id, &text)));
} else { } else {
@@ -141,6 +206,17 @@ pub(crate) fn rig_openai_response_stream(
current_reasoning_message_id = Some(message_id); 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::UsageUpdated { usage: updated } => usage = updated,
AgentEvent::Tool { AgentEvent::Tool {
event: ToolEvent::Proposed { call }, event: ToolEvent::Proposed { call },
@@ -148,6 +224,8 @@ pub(crate) fn rig_openai_response_stream(
proposed_tools.push(call.clone()); proposed_tools.push(call.clone());
sync_assistant_turn( sync_assistant_turn(
&messages_sent, &messages_sent,
&full_reasoning,
reasoning_signature.as_deref(),
&full_text, &full_text,
&proposed_tools, &proposed_tools,
&mut assistant_history_index, &mut assistant_history_index,
@@ -164,7 +242,7 @@ pub(crate) fn rig_openai_response_stream(
yield Err(agent_error(AgentError::new( yield Err(agent_error(AgentError::new(
galaxy_agent_core::AgentErrorKind::Protocol, galaxy_agent_core::AgentErrorKind::Protocol,
message, message,
))); ), stream_type));
return; return;
} }
} }
@@ -198,6 +276,8 @@ pub(crate) fn rig_openai_response_stream(
} }
sync_assistant_turn( sync_assistant_turn(
&messages_sent, &messages_sent,
&full_reasoning,
reasoning_signature.as_deref(),
&full_text, &full_text,
&proposed_tools, &proposed_tools,
&mut assistant_history_index, &mut assistant_history_index,
@@ -222,7 +302,7 @@ pub(crate) fn rig_openai_response_stream(
yield Err(agent_error(AgentError::new( yield Err(agent_error(AgentError::new(
galaxy_agent_core::AgentErrorKind::Protocol, galaxy_agent_core::AgentErrorKind::Protocol,
"the provider runtime attempted to execute a tool outside Galaxy's permission boundary", "the provider runtime attempted to execute a tool outside Galaxy's permission boundary",
))); ), stream_type));
return; return;
} }
} }
@@ -264,11 +344,22 @@ fn append_tool_result(
fn sync_assistant_turn( fn sync_assistant_turn(
messages_sent: &std::sync::Arc<std::sync::Mutex<Vec<ConversationMessage>>>, messages_sent: &std::sync::Arc<std::sync::Mutex<Vec<ConversationMessage>>>,
reasoning_text: &str,
reasoning_signature: Option<&str>,
text: &str, text: &str,
tool_calls: &[ToolCall], tool_calls: &[ToolCall],
history_index: &mut Option<usize>, 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() { if !text.is_empty() {
parts.push(ContentPart::Text(text.to_string())); parts.push(ContentPart::Text(text.to_string()));
} }
@@ -293,6 +384,7 @@ fn sync_assistant_turn(
name, name,
input, input,
}, },
reasoning @ ContentPart::Reasoning { .. } => MessageContent::MultiPart(vec![reasoning]),
ContentPart::Image { .. } | ContentPart::ToolResult { .. } => unreachable!(), ContentPart::Image { .. } | ContentPart::ToolResult { .. } => unreachable!(),
} }
} else { } else {
@@ -395,10 +487,10 @@ fn saturating_i32(value: u64) -> i32 {
i32::try_from(value).unwrap_or(i32::MAX) 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( Arc::new(
AIApiError::Stream { AIApiError::Stream {
stream_type: "rig_openai_compatible", stream_type,
source: anyhow::anyhow!(error), source: anyhow::anyhow!(error),
} }
.into_quota_limit_if_provider_budget_exhausted(), .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::api::RequestParams;
use crate::ai::agent::{AIAgentContext, AIAgentInput, MCPContext, UserQueryMode}; 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::client::OpenAIClientConfig;
use crate::ai::openai::request_translator::sanitize_messages_for_openai; use crate::ai::openai::request_translator::sanitize_messages_for_openai;
@@ -32,6 +34,47 @@ pub(crate) fn prepare_rig_turn(
params: RequestParams, params: RequestParams,
supported_tools: Vec<ToolType>, supported_tools: Vec<ToolType>,
supported_cli_agent_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 { ) -> PreparedRigTurn {
let RequestParams { let RequestParams {
input, input,
@@ -70,7 +113,12 @@ pub(crate) fn prepare_rig_turn(
for message in &mut persistent_messages { for message in &mut persistent_messages {
message.truncate_tool_results_for_provider_request(); 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(); let mut turn_messages = Vec::new();
if let Some(summary) = progressive_summary { if let Some(summary) = progressive_summary {
@@ -91,16 +139,14 @@ pub(crate) fn prepare_rig_turn(
} }
turn_messages.extend(persistent_messages.clone()); turn_messages.extend(persistent_messages.clone());
let model_id = config let model_id = model_override
.model
.clone()
.filter(|model| !model.is_empty() && model != "auto") .filter(|model| !model.is_empty() && model != "auto")
.unwrap_or_else(|| model.as_str().to_string()); .unwrap_or_else(|| model.as_str().to_string());
let mut request = TurnRequest::new(model_id, turn_messages); let mut request = TurnRequest::new(model_id, turn_messages);
request.conversation_id = conversation_token.map(|token| token.as_str().to_string()); request.conversation_id = conversation_token.map(|token| token.as_str().to_string());
request.system_prompt = Some(system_prompt); request.system_prompt = Some(system_prompt);
request.tools = tools; request.tools = tools;
request.max_output_tokens = config.max_output_tokens.map(u64::from); request.max_output_tokens = max_output_tokens;
PreparedRigTurn { PreparedRigTurn {
task_id, task_id,
+35 -1
View File
@@ -4,7 +4,7 @@ use std::sync::Arc;
use galaxy_agent_core::{ContentPart, MessageContent, MessageRole, ToolResult, ToolResultStatus}; use galaxy_agent_core::{ContentPart, MessageContent, MessageRole, ToolResult, ToolResultStatus};
use warp_multi_agent_api::ToolType; 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::api::RequestParams;
use crate::ai::agent::{ use crate::ai::agent::{
AIAgentContext, AIAgentInput, AnyFileContent, FileContext, MCPContext, MCPServer, UserQueryMode, 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] #[test]
#[allow(deprecated)] #[allow(deprecated)]
fn grouped_mcp_tool_names_use_the_installation_id_not_the_display_name() { 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 ai::skills::SkillPathOrigin;
use galaxy_agent_core::{ 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; 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( sync_assistant_turn(
&messages, &messages,
"",
None,
"I'll inspect both.", "I'll inspect both.",
std::slice::from_ref(&first_call), std::slice::from_ref(&first_call),
&mut history_index, &mut history_index,
); );
sync_assistant_turn( sync_assistant_turn(
&messages, &messages,
"",
None,
"I'll inspect both.", "I'll inspect both.",
&[first_call, second_call], &[first_call, second_call],
&mut history_index, &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] #[test]
fn rejected_tool_result_is_paired_with_the_assistant_call_in_history() { fn rejected_tool_result_is_paired_with_the_assistant_call_in_history() {
let messages = Arc::new(Mutex::new(Vec::new())); 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( sync_assistant_turn(
&messages, &messages,
"", "",
None,
"",
std::slice::from_ref(&call), std::slice::from_ref(&call),
&mut history_index, &mut history_index,
); );
+9 -2
View File
@@ -828,6 +828,11 @@ pub struct BedrockModelConfig {
#[serde(default)] #[serde(default)]
#[schemars(description = "Whether the model supports image/vision input.")] #[schemars(description = "Whether the model supports image/vision input.")]
pub vision_supported: bool, pub vision_supported: bool,
#[serde(default)]
#[schemars(
description = "Route this model through Galaxy's Rig Bedrock runtime. Disabled by default while compatibility validation is in progress."
)]
pub use_rig: bool,
} }
impl settings_value::SettingsValue for BedrockModelConfig {} impl settings_value::SettingsValue for BedrockModelConfig {}
@@ -890,8 +895,10 @@ impl settings_value::SettingsValue for OpenAIModelConfig {}
impl OpenAIModelConfig { impl OpenAIModelConfig {
pub fn supports_system_messages(&self) -> bool { pub fn supports_system_messages(&self) -> bool {
self.supports_system_messages if self.model_id.starts_with("codex-gpt-") {
.unwrap_or_else(|| !self.model_id.starts_with("codex-gpt-")) return false;
}
self.supports_system_messages.unwrap_or(true)
} }
} }
+405 -143
View File
@@ -88,8 +88,8 @@ use crate::settings::{
GitOperationsAutogenEnabled, IncludeAgentCommandsInHistory, InputSettings, GitOperationsAutogenEnabled, IncludeAgentCommandsInHistory, InputSettings,
IntelligentAutosuggestionsEnabled, LongRunningCommandSubmissionMode, MemoryEnabled, IntelligentAutosuggestionsEnabled, LongRunningCommandSubmissionMode, MemoryEnabled,
NLDInTerminalEnabled, NaturalLanguageAutosuggestionsEnabled, OpenAIEnabled, NLDInTerminalEnabled, NaturalLanguageAutosuggestionsEnabled, OpenAIEnabled,
OrchestrationMessageDisplayMode, PromptSubmissionMode, RuleSuggestionsEnabled, OpenAIProviderConfig, OrchestrationMessageDisplayMode, PromptSubmissionMode,
SharedBlockTitleGenerationEnabled, ShouldRenderCLIAgentToolbar, RuleSuggestionsEnabled, SharedBlockTitleGenerationEnabled, ShouldRenderCLIAgentToolbar,
ShouldRenderUseAgentToolbarForUserCommands, ShowAgentTips, ShowConversationHistory, ShouldRenderUseAgentToolbarForUserCommands, ShowAgentTips, ShowConversationHistory,
ShowHintText, ThinkingDisplayMode, VoiceInputEnabled, WarpDriveContextEnabled, ShowHintText, ThinkingDisplayMode, VoiceInputEnabled, WarpDriveContextEnabled,
}; };
@@ -117,10 +117,8 @@ pub enum AISubpage {
Knowledge, Knowledge,
/// Third-party CLI agent settings. /// Third-party CLI agent settings.
ThirdPartyCLIAgents, ThirdPartyCLIAgents,
/// AWS Bedrock direct provider configuration. /// Unified model and provider configuration.
Bedrock, Models,
/// OpenAI-compatible (LiteLLM) provider configuration.
OpenAI,
/// Experimental features. /// Experimental features.
Experiments, Experiments,
} }
@@ -132,8 +130,7 @@ impl AISubpage {
SettingsSection::AgentProfiles => Some(Self::Profiles), SettingsSection::AgentProfiles => Some(Self::Profiles),
SettingsSection::Knowledge => Some(Self::Knowledge), SettingsSection::Knowledge => Some(Self::Knowledge),
SettingsSection::ThirdPartyCLIAgents => Some(Self::ThirdPartyCLIAgents), SettingsSection::ThirdPartyCLIAgents => Some(Self::ThirdPartyCLIAgents),
SettingsSection::Bedrock => Some(Self::Bedrock), SettingsSection::Models => Some(Self::Models),
SettingsSection::OpenAI => Some(Self::OpenAI),
SettingsSection::Experiments => Some(Self::Experiments), SettingsSection::Experiments => Some(Self::Experiments),
// AgentMCPServers renders the standalone MCPServers page, not an AI subpage. // AgentMCPServers renders the standalone MCPServers page, not an AI subpage.
_ => None, _ => None,
@@ -1908,15 +1905,18 @@ impl AISettingsPageView {
} }
} }
/// Fetches models from the LiteLLM endpoint and stores them in memory via LLMPreferences. fn fetch_openai_provider_models(&mut self, provider_index: usize, ctx: &mut ViewContext<Self>) {
fn fetch_litellm_models(&mut self, ctx: &mut ViewContext<Self>) {
use crate::ai::llms::LLMPreferences;
LLMPreferences::handle(ctx).update(ctx, |llm_prefs, ctx| { LLMPreferences::handle(ctx).update(ctx, |llm_prefs, ctx| {
llm_prefs.fetch_openai_models_from_endpoint(ctx); llm_prefs.fetch_openai_provider_models(provider_index, ctx);
}); });
} }
fn rebuild_active_subpage(&mut self, ctx: &mut ViewContext<Self>) {
let (page, _) = Self::build_page(self.active_subpage, ctx);
self.page = page;
ctx.notify();
}
fn build_page( fn build_page(
subpage: Option<AISubpage>, subpage: Option<AISubpage>,
ctx: &mut ViewContext<Self>, ctx: &mut ViewContext<Self>,
@@ -2034,15 +2034,10 @@ impl AISettingsPageView {
Some(AISubpage::ThirdPartyCLIAgents) => { Some(AISubpage::ThirdPartyCLIAgents) => {
widgets.push(Box::new(CLIAgentWidget::default())); widgets.push(Box::new(CLIAgentWidget::default()));
} }
Some(AISubpage::Bedrock) => { Some(AISubpage::Models) => {
let widget = BedrockSettingsWidget::new(ctx); widgets.push(Box::new(ModelsOverviewWidget));
widgets.push(Box::new(widget)); widgets.push(Box::new(OpenAISettingsWidget::new(ctx)));
let title: Option<&str> = None; widgets.push(Box::new(BedrockSettingsWidget::new(ctx)));
return (PageType::new_uncategorized(widgets, title), None);
}
Some(AISubpage::OpenAI) => {
let widget = OpenAISettingsWidget::new(ctx);
widgets.push(Box::new(widget));
let title: Option<&str> = None; let title: Option<&str> = None;
return (PageType::new_uncategorized(widgets, title), None); return (PageType::new_uncategorized(widgets, title), None);
} }
@@ -2807,10 +2802,13 @@ pub enum AISettingsPageAction {
SetBedrockAuthMethod(BedrockAuthMethod), SetBedrockAuthMethod(BedrockAuthMethod),
SetBedrockProfile(String), SetBedrockProfile(String),
ToggleBedrockCrossRegionInference, ToggleBedrockCrossRegionInference,
ToggleBedrockModelRig(usize),
ToggleOpenAIEnabled, ToggleOpenAIEnabled,
ToggleAcpEnabled, ToggleAcpEnabled,
RefreshAcpDiscovery, RefreshAcpDiscovery,
FetchOpenAIModels, FetchOpenAIProviderModels(usize),
AddOpenAIProvider,
RemoveOpenAIProvider(usize),
ToggleFileBasedMcp, ToggleFileBasedMcp,
ToggleIncludeAgentCommandsInHistory, ToggleIncludeAgentCommandsInHistory,
ToggleAgentAttribution, ToggleAgentAttribution,
@@ -3562,6 +3560,17 @@ impl TypedActionView for AISettingsPageView {
}); });
ctx.notify(); ctx.notify();
} }
AISettingsPageAction::ToggleBedrockModelRig(index) => {
AISettings::handle(ctx).update(ctx, |settings, ctx| {
let mut models = settings.bedrock_models.value().clone();
let Some(model) = models.get_mut(*index) else {
return;
};
model.use_rig = !model.use_rig;
report_if_error!(settings.bedrock_models.set_value(models, ctx));
});
ctx.notify();
}
AISettingsPageAction::ToggleOpenAIEnabled => { AISettingsPageAction::ToggleOpenAIEnabled => {
AISettings::handle(ctx).update(ctx, |settings, ctx| { AISettings::handle(ctx).update(ctx, |settings, ctx| {
report_if_error!(settings.openai_enabled.toggle_and_save_value(ctx)); report_if_error!(settings.openai_enabled.toggle_and_save_value(ctx));
@@ -3580,9 +3589,32 @@ impl TypedActionView for AISettingsPageView {
#[cfg(not(target_family = "wasm"))] #[cfg(not(target_family = "wasm"))]
self.refresh_acp_discovery(ctx); self.refresh_acp_discovery(ctx);
} }
AISettingsPageAction::FetchOpenAIModels => { AISettingsPageAction::FetchOpenAIProviderModels(provider_index) => {
// Trigger a fetch of models from the LiteLLM endpoint self.fetch_openai_provider_models(*provider_index, ctx);
self.fetch_litellm_models(ctx); }
AISettingsPageAction::AddOpenAIProvider => {
AISettings::handle(ctx).update(ctx, |settings, ctx| {
let mut providers = settings.openai_providers.value().clone();
let provider_number = providers.len() + 1;
providers.push(OpenAIProviderConfig {
name: format!("Provider {provider_number}"),
base_url: "http://localhost:4000/v1".to_string(),
api_key: None,
models: Vec::new(),
});
report_if_error!(settings.openai_providers.set_value(providers, ctx));
});
self.rebuild_active_subpage(ctx);
}
AISettingsPageAction::RemoveOpenAIProvider(provider_index) => {
AISettings::handle(ctx).update(ctx, |settings, ctx| {
let mut providers = settings.openai_providers.value().clone();
if *provider_index < providers.len() {
providers.remove(*provider_index);
report_if_error!(settings.openai_providers.set_value(providers, ctx));
}
});
self.rebuild_active_subpage(ctx);
} }
AISettingsPageAction::ToggleFileBasedMcp => { AISettingsPageAction::ToggleFileBasedMcp => {
AISettings::handle(ctx).update(ctx, |settings, ctx| { AISettings::handle(ctx).update(ctx, |settings, ctx| {
@@ -7230,6 +7262,50 @@ impl SettingsWidget for CloudHandoffWidget {
} }
} }
struct ModelsOverviewWidget;
impl SettingsWidget for ModelsOverviewWidget {
type View = AISettingsPageView;
fn search_terms(&self) -> &str {
"models providers rig litellm openai compatible ollama lm studio bedrock"
}
fn render(
&self,
_view: &Self::View,
appearance: &Appearance,
app: &AppContext,
) -> Box<dyn Element> {
let settings = AISettings::as_ref(app);
let endpoint_count = settings.openai_providers.value().len();
let endpoint_model_count = settings
.openai_providers
.value()
.iter()
.map(|provider| provider.models.len())
.sum::<usize>();
let bedrock_model_count = settings.bedrock_models.value().len();
Flex::column()
.with_spacing(8.)
.with_child(build_sub_header(appearance, "Models", None).finish())
.with_child(render_ai_setting_description(
"Configure the model providers available to Galaxy. OpenAI-compatible endpoints and opted-in Bedrock models share the same Rig conversation, tool, and UI runtime; Bedrock's compatibility path remains available during validation.",
true,
app,
))
.with_child(render_ai_setting_description(
format!(
"{endpoint_count} OpenAI-compatible provider(s) with {endpoint_model_count} model(s); {bedrock_model_count} Bedrock model(s)."
),
true,
app,
))
.finish()
}
}
struct BedrockSettingsWidget { struct BedrockSettingsWidget {
enabled_toggle: SwitchStateHandle, enabled_toggle: SwitchStateHandle,
auto_login_toggle: SwitchStateHandle, auto_login_toggle: SwitchStateHandle,
@@ -7239,6 +7315,7 @@ struct BedrockSettingsWidget {
auth_refresh_command_editor: ViewHandle<EditorView>, auth_refresh_command_editor: ViewHandle<EditorView>,
access_key_editor: ViewHandle<EditorView>, access_key_editor: ViewHandle<EditorView>,
secret_key_editor: ViewHandle<EditorView>, secret_key_editor: ViewHandle<EditorView>,
model_rig_toggles: RefCell<Vec<SwitchStateHandle>>,
} }
impl BedrockSettingsWidget { impl BedrockSettingsWidget {
@@ -7250,6 +7327,7 @@ impl BedrockSettingsWidget {
let auth_cmd_val = ai_settings.bedrock_auth_refresh_command.value().clone(); let auth_cmd_val = ai_settings.bedrock_auth_refresh_command.value().clone();
let access_key_val = ai_settings.bedrock_access_key_id.value().clone(); let access_key_val = ai_settings.bedrock_access_key_id.value().clone();
let secret_key_val = ai_settings.bedrock_secret_access_key.value().clone(); let secret_key_val = ai_settings.bedrock_secret_access_key.value().clone();
let bedrock_model_count = ai_settings.bedrock_models.value().len();
let auth_method_dropdown = ctx.add_typed_action_view(|ctx| { let auth_method_dropdown = ctx.add_typed_action_view(|ctx| {
let mut dropdown = Dropdown::new(ctx); let mut dropdown = Dropdown::new(ctx);
@@ -7475,6 +7553,11 @@ impl BedrockSettingsWidget {
auth_refresh_command_editor, auth_refresh_command_editor,
access_key_editor, access_key_editor,
secret_key_editor, secret_key_editor,
model_rig_toggles: RefCell::new(
(0..bedrock_model_count)
.map(|_| SwitchStateHandle::default())
.collect(),
),
} }
} }
@@ -7540,6 +7623,8 @@ impl SettingsWidget for BedrockSettingsWidget {
let mut column = Flex::column().with_spacing(16.); let mut column = Flex::column().with_spacing(16.);
column.add_child(build_sub_header(appearance, "AWS Bedrock", None).finish());
let has_aws_env = std::env::vars_os().any(|(k, _)| k.to_string_lossy().starts_with("AWS_")); let has_aws_env = std::env::vars_os().any(|(k, _)| k.to_string_lossy().starts_with("AWS_"));
if has_aws_env { if has_aws_env {
@@ -7673,6 +7758,47 @@ impl SettingsWidget for BedrockSettingsWidget {
} }
); );
column.add_child(render_ai_setting_description(description, is_enabled, app)); column.add_child(render_ai_setting_description(description, is_enabled, app));
column.add_child(build_sub_header(appearance, "Bedrock runtime", None).finish());
column.add_child(render_ai_setting_description(
"Opt individual Bedrock models into the shared Rig runtime. Models left off continue through the compatibility runtime; one-hour prompt-cache TTL requests always fall back automatically.",
is_enabled,
app,
));
let toggle_handles = {
let mut toggles = self.model_rig_toggles.borrow_mut();
while toggles.len() < configured_models.len() {
toggles.push(SwitchStateHandle::default());
}
toggles.clone()
};
for (index, model) in configured_models.iter().enumerate() {
let toggle = appearance
.ui_builder()
.switch(toggle_handles[index].clone())
.check(model.use_rig)
.with_disabled(!is_enabled)
.build()
.on_click(move |ctx, _, _| {
ctx.dispatch_typed_action(AISettingsPageAction::ToggleBedrockModelRig(
index,
));
})
.finish();
column.add_child(build_toggle_element(
render_body_item_label::<AISettingsPageAction>(
format!("{} — Rig", model.display_name),
Some(styles::header_font_color(is_enabled, app)),
None,
LocalOnlyIconState::Hidden,
ToggleState::Enabled,
appearance,
),
toggle,
appearance,
None,
));
}
} else { } else {
column.add_child(render_ai_setting_description( column.add_child(render_ai_setting_description(
"No models configured. Add models to ~/.galaxy/settings.toml under [ai.bedrock].", "No models configured. Add models to ~/.galaxy/settings.toml under [ai.bedrock].",
@@ -8003,104 +8129,149 @@ impl SettingsWidget for ACPSettingsWidget {
} }
} }
struct OpenAISettingsWidget { struct OpenAIProviderEditor {
enabled_toggle: SwitchStateHandle, name_editor: ViewHandle<EditorView>,
base_url_editor: ViewHandle<EditorView>, base_url_editor: ViewHandle<EditorView>,
api_key_editor: ViewHandle<EditorView>, api_key_editor: ViewHandle<EditorView>,
fetch_button: MouseStateHandle, fetch_button: MouseStateHandle,
remove_button: MouseStateHandle,
}
struct OpenAISettingsWidget {
enabled_toggle: SwitchStateHandle,
provider_editors: Vec<OpenAIProviderEditor>,
add_provider_button: MouseStateHandle,
} }
impl OpenAISettingsWidget { impl OpenAISettingsWidget {
fn create_editor(
value: String,
placeholder: &'static str,
is_password: bool,
ctx: &mut ViewContext<<Self as SettingsWidget>::View>,
) -> ViewHandle<EditorView> {
ctx.add_typed_action_view(move |ctx| {
let appearance = Appearance::as_ref(ctx);
let options = SingleLineEditorOptions {
is_password,
text: TextOptions {
font_size_override: Some(appearance.ui_font_size()),
font_family_override: Some(appearance.monospace_font_family()),
text_colors_override: Some(TextColors {
default_color: appearance.theme().active_ui_text_color(),
disabled_color: appearance.theme().disabled_ui_text_color(),
hint_color: appearance.theme().disabled_ui_text_color(),
}),
..Default::default()
},
..Default::default()
};
let mut editor = EditorView::single_line(options, ctx);
editor.set_placeholder_text(placeholder, ctx);
editor.set_buffer_text(&value, ctx);
editor
})
}
fn new(ctx: &mut ViewContext<<Self as SettingsWidget>::View>) -> Self { fn new(ctx: &mut ViewContext<<Self as SettingsWidget>::View>) -> Self {
let ai_settings = AISettings::as_ref(ctx); let providers = AISettings::as_ref(ctx).openai_providers.value().clone();
let is_enabled = *AISettings::as_ref(ctx).openai_enabled.value();
let mut provider_editors = Vec::with_capacity(providers.len());
let base_url_val = ai_settings.openai_base_url.value().clone(); for (provider_index, provider) in providers.into_iter().enumerate() {
let api_key_val = ai_settings.openai_api_key.value().clone(); let name_editor = Self::create_editor(provider.name, "Provider name", false, ctx);
ctx.subscribe_to_view(&name_editor, move |_, editor, event, ctx| {
if matches!(event, EditorEvent::Blurred | EditorEvent::Enter) {
let value = editor.as_ref(ctx).buffer_text(ctx);
AISettings::handle(ctx).update(ctx, |settings, ctx| {
let mut providers = settings.openai_providers.value().clone();
if let Some(provider) = providers.get_mut(provider_index) {
provider.name = value;
report_if_error!(settings.openai_providers.set_value(providers, ctx));
}
});
}
});
let base_url_editor = ctx.add_typed_action_view(move |ctx| { let base_url_editor =
let appearance = Appearance::as_ref(ctx); Self::create_editor(provider.base_url, "http://localhost:4000/v1", false, ctx);
let options = SingleLineEditorOptions { ctx.subscribe_to_view(&base_url_editor, move |_, editor, event, ctx| {
is_password: false, if matches!(event, EditorEvent::Blurred | EditorEvent::Enter) {
text: TextOptions { let value = editor.as_ref(ctx).buffer_text(ctx);
font_size_override: Some(appearance.ui_font_size()), AISettings::handle(ctx).update(ctx, |settings, ctx| {
font_family_override: Some(appearance.monospace_font_family()), let mut providers = settings.openai_providers.value().clone();
text_colors_override: Some(TextColors { if let Some(provider) = providers.get_mut(provider_index) {
default_color: appearance.theme().active_ui_text_color(), provider.base_url = value;
disabled_color: appearance.theme().disabled_ui_text_color(), report_if_error!(settings.openai_providers.set_value(providers, ctx));
hint_color: appearance.theme().disabled_ui_text_color(), }
}), });
..Default::default() }
}, });
..Default::default()
}; let api_key_editor = Self::create_editor(
let mut editor = EditorView::single_line(options, ctx); provider.api_key.unwrap_or_default(),
editor.set_placeholder_text("http://localhost:4000/v1", ctx); "sk-... (optional)",
editor.set_buffer_text(&base_url_val, ctx); true,
editor ctx,
}); );
ctx.subscribe_to_view(&base_url_editor, |_, editor, event, ctx| { ctx.subscribe_to_view(&api_key_editor, move |_, editor, event, ctx| {
if matches!(event, EditorEvent::Blurred | EditorEvent::Enter) { if matches!(event, EditorEvent::Blurred | EditorEvent::Enter) {
let value = editor.as_ref(ctx).buffer_text(ctx); let value = editor.as_ref(ctx).buffer_text(ctx);
AISettings::handle(ctx).update(ctx, |settings, ctx| { AISettings::handle(ctx).update(ctx, |settings, ctx| {
let _ = settings.openai_base_url.set_value(value, ctx); let mut providers = settings.openai_providers.value().clone();
}); if let Some(provider) = providers.get_mut(provider_index) {
provider.api_key = (!value.is_empty()).then_some(value);
report_if_error!(settings.openai_providers.set_value(providers, ctx));
}
});
}
});
for editor in [&name_editor, &base_url_editor, &api_key_editor] {
AISettingsPageView::update_editor_interaction_state(
editor.clone(),
is_enabled,
ctx,
);
} }
});
let api_key_editor = ctx.add_typed_action_view(move |ctx| { provider_editors.push(OpenAIProviderEditor {
let appearance = Appearance::as_ref(ctx); name_editor,
let options = SingleLineEditorOptions { base_url_editor,
is_password: true, api_key_editor,
text: TextOptions { fetch_button: MouseStateHandle::default(),
font_size_override: Some(appearance.ui_font_size()), remove_button: MouseStateHandle::default(),
font_family_override: Some(appearance.monospace_font_family()), });
text_colors_override: Some(TextColors { }
default_color: appearance.theme().active_ui_text_color(),
disabled_color: appearance.theme().disabled_ui_text_color(),
hint_color: appearance.theme().disabled_ui_text_color(),
}),
..Default::default()
},
..Default::default()
};
let mut editor = EditorView::single_line(options, ctx);
editor.set_placeholder_text("sk-... (optional)", ctx);
editor.set_buffer_text(&api_key_val, ctx);
editor
});
ctx.subscribe_to_view(&api_key_editor, |_, editor, event, ctx| {
if matches!(event, EditorEvent::Blurred | EditorEvent::Enter) {
let value = editor.as_ref(ctx).buffer_text(ctx);
AISettings::handle(ctx).update(ctx, |settings, ctx| {
let _ = settings.openai_api_key.set_value(value, ctx);
});
}
});
let base_url_editor_clone = base_url_editor.clone(); let editor_handles = provider_editors
let api_key_editor_clone = api_key_editor.clone(); .iter()
.flat_map(|provider| {
[
provider.name_editor.clone(),
provider.base_url_editor.clone(),
provider.api_key_editor.clone(),
]
})
.collect::<Vec<_>>();
ctx.subscribe_to_model(&AISettings::handle(ctx), move |_, _, event, ctx| { ctx.subscribe_to_model(&AISettings::handle(ctx), move |_, _, event, ctx| {
if matches!(event, AISettingsChangedEvent::OpenAIEnabled { .. }) { if matches!(event, AISettingsChangedEvent::OpenAIEnabled { .. }) {
let is_enabled = *AISettings::as_ref(ctx).openai_enabled.value(); let is_enabled = *AISettings::as_ref(ctx).openai_enabled.value();
AISettingsPageView::update_editor_interaction_state( for editor in &editor_handles {
base_url_editor_clone.clone(), AISettingsPageView::update_editor_interaction_state(
is_enabled, editor.clone(),
ctx, is_enabled,
); ctx,
AISettingsPageView::update_editor_interaction_state( );
api_key_editor_clone.clone(), }
is_enabled,
ctx,
);
ctx.notify(); ctx.notify();
} }
}); });
Self { Self {
enabled_toggle: SwitchStateHandle::default(), enabled_toggle: SwitchStateHandle::default(),
base_url_editor, provider_editors,
api_key_editor, add_provider_button: MouseStateHandle::default(),
fetch_button: MouseStateHandle::default(),
} }
} }
@@ -8164,8 +8335,11 @@ impl SettingsWidget for OpenAISettingsWidget {
let mut column = Flex::column().with_spacing(16.); let mut column = Flex::column().with_spacing(16.);
column
.add_child(build_sub_header(appearance, "OpenAI-compatible providers", None).finish());
column.add_child(render_ai_setting_toggle::<OpenAIEnabled>( column.add_child(render_ai_setting_toggle::<OpenAIEnabled>(
"Enable OpenAI-Compatible Provider", "Enable model providers",
AISettingsPageAction::ToggleOpenAIEnabled, AISettingsPageAction::ToggleOpenAIEnabled,
is_enabled, is_enabled,
true, true,
@@ -8174,65 +8348,153 @@ impl SettingsWidget for OpenAISettingsWidget {
app, app,
)); ));
column.add_child(render_ai_setting_description( column.add_child(render_ai_setting_description(
"Route AI requests through an OpenAI-compatible endpoint (e.g. LiteLLM proxy).", "Route configured LiteLLM, Ollama, LM Studio, vLLM, and other OpenAI-compatible models through Galaxy's provider registry.",
true, true,
app, app,
)); ));
column.add_child(render_separator(appearance)); if ai_settings.openai_providers.value().is_empty() {
column.add_child(render_ai_setting_description(
"No providers configured. Add a provider to connect a local or private OpenAI-compatible endpoint.",
is_enabled,
app,
));
}
column.add_child(Self::render_input( for (provider_index, provider) in ai_settings.openai_providers.value().iter().enumerate() {
appearance, let Some(editors) = self.provider_editors.get(provider_index) else {
"Base URL", continue;
self.base_url_editor.clone(), };
is_enabled,
app,
));
column.add_child(render_ai_setting_description(
"The OpenAI-compatible API base URL (e.g. http://localhost:4000/v1).",
is_enabled,
app,
));
column.add_child(Self::render_input( column.add_child(render_separator(appearance));
appearance, column.add_child(
"API Key", build_sub_header(
self.api_key_editor.clone(), appearance,
is_enabled, format!("Provider {}: {}", provider_index + 1, provider.name),
app, None,
)); )
column.add_child(render_ai_setting_description( .finish(),
"Optional. Leave empty if the proxy handles authentication.", );
is_enabled, column.add_child(Self::render_input(
app, appearance,
)); "Name",
editors.name_editor.clone(),
is_enabled,
app,
));
column.add_child(Self::render_input(
appearance,
"Base URL",
editors.base_url_editor.clone(),
is_enabled,
app,
));
column.add_child(Self::render_input(
appearance,
"API Key",
editors.api_key_editor.clone(),
is_enabled,
app,
));
column.add_child(render_ai_setting_description(
"The API key is optional, stored only in ~/.galaxy/settings.toml, and never synced to the cloud.",
is_enabled,
app,
));
let fetch_button = appearance
.ui_builder()
.button(ButtonVariant::Secondary, editors.fetch_button.clone())
.with_text_label("Discover Models".to_owned());
let fetch_button = if !is_enabled || provider.base_url.trim().is_empty() {
fetch_button.disabled().build().finish()
} else {
fetch_button
.build()
.on_click(move |ctx, _, _| {
ctx.dispatch_typed_action(AISettingsPageAction::FetchOpenAIProviderModels(
provider_index,
));
})
.finish()
};
let remove_button = appearance
.ui_builder()
.button(ButtonVariant::Error, editors.remove_button.clone())
.with_text_label("Remove Provider".to_owned())
.build()
.on_click(move |ctx, _, _| {
ctx.dispatch_typed_action(AISettingsPageAction::RemoveOpenAIProvider(
provider_index,
));
})
.finish();
column.add_child(
Flex::row()
.with_spacing(8.)
.with_child(fetch_button)
.with_child(remove_button)
.finish(),
);
let model_names = provider
.models
.iter()
.take(5)
.map(|model| model.display_name.as_str())
.join(", ");
let overflow = provider.models.len().saturating_sub(5);
let overflow = if overflow > 0 {
format!(" (+{overflow} more)")
} else {
String::new()
};
let models_description = if provider.models.is_empty() {
"No models configured. Discover models from this endpoint.".to_string()
} else {
format!(
"{} model{}: {model_names}{overflow}",
provider.models.len(),
if provider.models.len() == 1 { "" } else { "s" },
)
};
column.add_child(render_ai_setting_description(
models_description,
is_enabled,
app,
));
}
column.add_child(render_separator(appearance)); column.add_child(render_separator(appearance));
let add_provider_button = appearance
// Fetch models button
let fetch_button = appearance
.ui_builder() .ui_builder()
.button(ButtonVariant::Secondary, self.fetch_button.clone()) .button(ButtonVariant::Secondary, self.add_provider_button.clone())
.with_text_label("Fetch Models from Endpoint".to_owned()) .with_text_label("Add Provider".to_owned())
.build() .build()
.on_click(move |ctx, _, _| { .on_click(move |ctx, _, _| {
ctx.dispatch_typed_action(AISettingsPageAction::FetchOpenAIModels); ctx.dispatch_typed_action(AISettingsPageAction::AddOpenAIProvider);
}) })
.finish(); .finish();
column.add_child(fetch_button); column.add_child(add_provider_button);
column.add_child(render_ai_setting_description( column.add_child(render_ai_setting_description(
"Queries the /models endpoint and populates the model list with available models and their context window sizes.", "Model discovery only contacts an endpoint when you click Discover Models.",
is_enabled, is_enabled,
app, app,
)); ));
column.add_child(render_separator(appearance)); column.add_child(render_separator(appearance));
// Show configured models count let mut configured_models = ai_settings
let configured_models: Vec<_> = ai_settings.openai_models.value().clone(); .openai_providers
.value()
.iter()
.flat_map(|provider| provider.models.iter())
.collect::<Vec<_>>();
configured_models.extend(ai_settings.openai_models.value().iter());
if !configured_models.is_empty() { if !configured_models.is_empty() {
let description = format!( let description = format!(
"{} model{} configured via settings.toml.", "{} model{} configured across all OpenAI-compatible providers.",
configured_models.len(), configured_models.len(),
if configured_models.len() == 1 { if configured_models.len() == 1 {
"" ""
@@ -8246,7 +8508,7 @@ impl SettingsWidget for OpenAISettingsWidget {
let preview: String = configured_models let preview: String = configured_models
.iter() .iter()
.take(5) .take(5)
.map(|m| m.display_name.as_str()) .map(|model| model.display_name.as_str())
.collect::<Vec<_>>() .collect::<Vec<_>>()
.join(", "); .join(", ");
let suffix = if configured_models.len() > 5 { let suffix = if configured_models.len() > 5 {
@@ -8261,7 +8523,7 @@ impl SettingsWidget for OpenAISettingsWidget {
)); ));
} else { } else {
column.add_child(render_ai_setting_description( column.add_child(render_ai_setting_description(
"No models configured. Use 'Fetch Models' or add them to ~/.galaxy/settings.toml under [ai.openai].", "No models configured. Add a provider and discover its models, or configure [[ai.providers.models]] in ~/.galaxy/settings.toml.",
is_enabled, is_enabled,
app, app,
)); ));
+7 -10
View File
@@ -245,8 +245,7 @@ pub enum SettingsSection {
AgentMCPServers, AgentMCPServers,
Knowledge, Knowledge,
ThirdPartyCLIAgents, ThirdPartyCLIAgents,
Bedrock, Models,
OpenAI,
Experiments, Experiments,
/// Internal backing-page identifier for CodeSettingsPageView. Multiple subpages /// Internal backing-page identifier for CodeSettingsPageView. Multiple subpages
/// (CodeIndexing, EditorAndCodeReview) share this single backing page, /// (CodeIndexing, EditorAndCodeReview) share this single backing page,
@@ -274,8 +273,7 @@ impl Display for SettingsSection {
SettingsSection::AgentMCPServers => write!(f, "MCP servers"), SettingsSection::AgentMCPServers => write!(f, "MCP servers"),
SettingsSection::Knowledge => write!(f, "Knowledge"), SettingsSection::Knowledge => write!(f, "Knowledge"),
SettingsSection::ThirdPartyCLIAgents => write!(f, "Third party CLI agents"), SettingsSection::ThirdPartyCLIAgents => write!(f, "Third party CLI agents"),
SettingsSection::Bedrock => write!(f, "AWS Bedrock"), SettingsSection::Models => write!(f, "Models"),
SettingsSection::OpenAI => write!(f, "OpenAI / LiteLLM"),
SettingsSection::Experiments => write!(f, "Experiments"), SettingsSection::Experiments => write!(f, "Experiments"),
SettingsSection::Warpify => write!(f, "Wormhole"), SettingsSection::Warpify => write!(f, "Wormhole"),
SettingsSection::CodeIndexing => write!(f, "Indexing and projects"), SettingsSection::CodeIndexing => write!(f, "Indexing and projects"),
@@ -300,8 +298,7 @@ impl SettingsSection {
| Self::AgentMCPServers | Self::AgentMCPServers
| Self::Knowledge | Self::Knowledge
| Self::ThirdPartyCLIAgents | Self::ThirdPartyCLIAgents
| Self::Bedrock | Self::Models
| Self::OpenAI
| Self::Experiments | Self::Experiments
) )
} }
@@ -329,12 +326,11 @@ impl SettingsSection {
pub fn ai_subpages() -> &'static [Self] { pub fn ai_subpages() -> &'static [Self] {
&[ &[
Self::WarpAgent, Self::WarpAgent,
Self::Models,
Self::AgentProfiles, Self::AgentProfiles,
Self::AgentMCPServers, Self::AgentMCPServers,
Self::Knowledge, Self::Knowledge,
Self::ThirdPartyCLIAgents, Self::ThirdPartyCLIAgents,
Self::Bedrock,
Self::OpenAI,
Self::Experiments, Self::Experiments,
] ]
} }
@@ -367,8 +363,9 @@ impl FromStr for SettingsSection {
"MCP servers" | "AgentMCPServers" => Ok(Self::AgentMCPServers), "MCP servers" | "AgentMCPServers" => Ok(Self::AgentMCPServers),
"Knowledge" => Ok(Self::Knowledge), "Knowledge" => Ok(Self::Knowledge),
"Third party CLI agents" | "ThirdPartyCLIAgents" => Ok(Self::ThirdPartyCLIAgents), "Third party CLI agents" | "ThirdPartyCLIAgents" => Ok(Self::ThirdPartyCLIAgents),
"AWS Bedrock" | "Bedrock" => Ok(Self::Bedrock), "Models" | "AWS Bedrock" | "Bedrock" | "OpenAI / LiteLLM" | "OpenAI" => {
"OpenAI / LiteLLM" | "OpenAI" => Ok(Self::OpenAI), Ok(Self::Models)
}
"Indexing and projects" | "CodeIndexing" => Ok(Self::CodeIndexing), "Indexing and projects" | "CodeIndexing" => Ok(Self::CodeIndexing),
"Editor and Code Review" | "EditorAndCodeReview" => Ok(Self::EditorAndCodeReview), "Editor and Code Review" | "EditorAndCodeReview" => Ok(Self::EditorAndCodeReview),
"Experiments" => Ok(Self::Experiments), "Experiments" => Ok(Self::Experiments),
+7 -6
View File
@@ -86,8 +86,7 @@ fn current_settings_display_names_round_trip() {
SettingsSection::ThirdPartyCLIAgents, SettingsSection::ThirdPartyCLIAgents,
"Third party CLI agents", "Third party CLI agents",
), ),
(SettingsSection::Bedrock, "AWS Bedrock"), (SettingsSection::Models, "Models"),
(SettingsSection::OpenAI, "OpenAI / LiteLLM"),
(SettingsSection::Experiments, "Experiments"), (SettingsSection::Experiments, "Experiments"),
(SettingsSection::CodeIndexing, "Indexing and projects"), (SettingsSection::CodeIndexing, "Indexing and projects"),
( (
@@ -111,8 +110,10 @@ fn legacy_settings_names_remain_parseable() {
("AgentProfiles", SettingsSection::AgentProfiles), ("AgentProfiles", SettingsSection::AgentProfiles),
("AgentMCPServers", SettingsSection::AgentMCPServers), ("AgentMCPServers", SettingsSection::AgentMCPServers),
("ThirdPartyCLIAgents", SettingsSection::ThirdPartyCLIAgents), ("ThirdPartyCLIAgents", SettingsSection::ThirdPartyCLIAgents),
("Bedrock", SettingsSection::Bedrock), ("AWS Bedrock", SettingsSection::Models),
("OpenAI", SettingsSection::OpenAI), ("Bedrock", SettingsSection::Models),
("OpenAI / LiteLLM", SettingsSection::Models),
("OpenAI", SettingsSection::Models),
("CodeIndexing", SettingsSection::CodeIndexing), ("CodeIndexing", SettingsSection::CodeIndexing),
("EditorAndCodeReview", SettingsSection::EditorAndCodeReview), ("EditorAndCodeReview", SettingsSection::EditorAndCodeReview),
] { ] {
@@ -215,7 +216,7 @@ fn collapsed_umbrella_uses_first_and_last_visible_subpages() {
let stops = build_nav_stops(&nav_items, |section| { let stops = build_nav_stops(&nav_items, |section| {
!matches!( !matches!(
section, section,
SettingsSection::WarpAgent | SettingsSection::OpenAI | SettingsSection::Experiments SettingsSection::WarpAgent | SettingsSection::Models | SettingsSection::Experiments
) )
}); });
@@ -224,7 +225,7 @@ fn collapsed_umbrella_uses_first_and_last_visible_subpages() {
NavStop::CollapsedUmbrella { NavStop::CollapsedUmbrella {
nav_index: 0, nav_index: 0,
first_subpage: SettingsSection::AgentProfiles, first_subpage: SettingsSection::AgentProfiles,
last_subpage: SettingsSection::Bedrock, last_subpage: SettingsSection::ThirdPartyCLIAgents,
} }
); );
} }
+3 -1
View File
@@ -250,7 +250,9 @@ fn collect_tool_entries(messages: &[ConversationMessage], entries: &mut Vec<Tool
content, content,
.. ..
} => pair_result(tool_use_id, content, &mut pending, entries), } => pair_result(tool_use_id, content, &mut pending, entries),
ContentPart::Text(_) | ContentPart::Image { .. } => {} ContentPart::Text(_)
| ContentPart::Reasoning { .. }
| ContentPart::Image { .. } => {}
} }
} }
} }
+34 -8
View File
@@ -42,6 +42,10 @@ pub enum MessageContent {
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] #[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub enum ContentPart { pub enum ContentPart {
Text(String), Text(String),
Reasoning {
text: String,
signature: Option<String>,
},
Image { Image {
data: Vec<u8>, data: Vec<u8>,
mime_type: String, mime_type: String,
@@ -231,12 +235,28 @@ pub enum StopReason {
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] #[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub enum AgentEvent { pub enum AgentEvent {
TurnStarted { runtime_request_id: String }, TurnStarted {
TextDelta { text: String }, runtime_request_id: String,
ReasoningDelta { text: String }, },
Tool { event: ToolEvent }, TextDelta {
UsageUpdated { usage: Usage }, text: String,
TurnStopped { reason: StopReason }, },
ReasoningDelta {
text: String,
},
ReasoningCompleted {
text: String,
signature: Option<String>,
},
Tool {
event: ToolEvent,
},
UsageUpdated {
usage: Usage,
},
TurnStopped {
reason: StopReason,
},
} }
fn truncate_tool_results_in_content(content: &mut MessageContent) { fn truncate_tool_results_in_content(content: &mut MessageContent) {
@@ -245,8 +265,14 @@ fn truncate_tool_results_in_content(content: &mut MessageContent) {
MessageContent::ToolResult { content, .. } => truncate_tool_result_text(content), MessageContent::ToolResult { content, .. } => truncate_tool_result_text(content),
MessageContent::MultiPart(parts) => { MessageContent::MultiPart(parts) => {
for part in parts { for part in parts {
if let ContentPart::ToolResult { content, .. } = part { match part {
truncate_tool_result_text(content); ContentPart::ToolResult { content, .. } => {
truncate_tool_result_text(content);
}
ContentPart::Text(_)
| ContentPart::Reasoning { .. }
| ContentPart::Image { .. }
| ContentPart::ToolUse { .. } => {}
} }
} }
} }
+4
View File
@@ -8,13 +8,17 @@ license.workspace = true
[dependencies] [dependencies]
async-stream.workspace = true async-stream.workspace = true
async-trait.workspace = true async-trait.workspace = true
aws-sdk-bedrockruntime.workspace = true
base64.workspace = true
futures.workspace = true futures.workspace = true
galaxy_agent_core.workspace = true galaxy_agent_core.workspace = true
rig-core.workspace = true rig-core.workspace = true
rig-bedrock.workspace = true
serde_json.workspace = true serde_json.workspace = true
uuid.workspace = true uuid.workspace = true
[dev-dependencies] [dev-dependencies]
aws-smithy-http-client.workspace = true
bytes.workspace = true bytes.workspace = true
rig-core = { workspace = true, features = ["test-utils"] } rig-core = { workspace = true, features = ["test-utils"] }
tokio = { workspace = true, features = ["macros", "rt"] } tokio = { workspace = true, features = ["macros", "rt"] }
+162
View File
@@ -0,0 +1,162 @@
use async_trait::async_trait;
use aws_sdk_bedrockruntime::Client as AwsBedrockClient;
use galaxy_agent_core::{
AgentError, AgentErrorKind, AgentEventStream, AgentRuntime, RuntimeCapabilities,
RuntimeDescriptor, RuntimeKind, TurnControl, TurnRequest,
};
use rig_bedrock::client::Client as RigBedrockClient;
use rig_bedrock::completion::CompletionModel;
use rig_core::client::CompletionClient;
use rig_core::completion::CompletionRequest;
use crate::request::build_completion_request;
use crate::stream::start_model_turn;
const INFERENCE_PROFILE_PREFIXES: &[&str] = &["us.", "eu.", "apac.", "jp.", "au.", "global."];
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct BedrockRigConfig {
pub model: String,
pub region: String,
pub cross_region_inference: bool,
pub prompt_caching: bool,
pub max_output_tokens: Option<u64>,
}
/// A Rig Bedrock client built from Galaxy's already-resolved AWS SDK client.
///
/// Credential/profile/SSO resolution remains in Galaxy's explicit Bedrock
/// configuration boundary. Rig receives the resulting SDK client and owns the
/// Converse request/stream conversion from that point onward.
#[derive(Clone)]
pub struct BedrockRuntime {
client: RigBedrockClient,
config: BedrockRigConfig,
resolved_model: String,
descriptor: RuntimeDescriptor,
}
impl BedrockRuntime {
pub fn from_aws_client(
client: AwsBedrockClient,
config: BedrockRigConfig,
) -> Result<Self, AgentError> {
let resolved_model =
resolve_bedrock_model_id(&config.model, &config.region, config.cross_region_inference)?;
let descriptor = RuntimeDescriptor {
id: format!("rig-bedrock:{resolved_model}"),
display_name: format!("Rig / Bedrock / {resolved_model}"),
kind: RuntimeKind::Provider,
capabilities: RuntimeCapabilities {
model_selection: true,
session_resume: false,
steering: false,
tool_permissions: false,
},
};
Ok(Self {
client: RigBedrockClient::from(client),
config,
resolved_model,
descriptor,
})
}
pub fn resolved_model(&self) -> &str {
&self.resolved_model
}
pub fn completion_model(&self) -> CompletionModel {
let model = self.client.completion_model(&self.resolved_model);
if self.config.prompt_caching {
model.with_prompt_caching()
} else {
model
}
}
}
#[async_trait]
impl AgentRuntime for BedrockRuntime {
fn descriptor(&self) -> &RuntimeDescriptor {
&self.descriptor
}
async fn start_turn(
&self,
request: TurnRequest,
control: TurnControl,
) -> Result<AgentEventStream, AgentError> {
let max_output_tokens = request.max_output_tokens.or(self.config.max_output_tokens);
let mut completion_request =
build_bedrock_completion_request(request, self.config.max_output_tokens)?;
// Context markers and inference-profile expansion are Galaxy model
// configuration, not identifiers that Rig should send unchanged.
completion_request.model = Some(self.resolved_model.clone());
start_model_turn(
self.completion_model(),
completion_request,
control,
max_output_tokens,
)
.await
}
}
pub fn build_bedrock_completion_request(
request: TurnRequest,
configured_max_output_tokens: Option<u64>,
) -> Result<CompletionRequest, AgentError> {
build_completion_request(request, configured_max_output_tokens, true, true, None)
}
pub fn resolve_bedrock_model_id(
configured_model: &str,
region: &str,
cross_region_inference: bool,
) -> Result<String, AgentError> {
let model = strip_context_marker(configured_model.trim());
if model.is_empty() {
return Err(AgentError::new(
AgentErrorKind::Configuration,
"Bedrock model ID is empty",
));
}
if !cross_region_inference
|| model.starts_with("arn:")
|| INFERENCE_PROFILE_PREFIXES
.iter()
.any(|prefix| model.starts_with(prefix))
{
return Ok(model.to_string());
}
let prefix = inference_profile_prefix(region);
Ok(prefix
.map(|prefix| format!("{prefix}.{model}"))
.unwrap_or_else(|| model.to_string()))
}
fn strip_context_marker(model: &str) -> &str {
model
.get(..model.len().saturating_sub(4))
.filter(|_| model.ends_with("[1m]") || model.ends_with("[1M]"))
.unwrap_or(model)
}
fn inference_profile_prefix(region: &str) -> Option<&'static str> {
match region {
region if region.starts_with("us-") || region.starts_with("ca-") => Some("us"),
region if region.starts_with("eu-") || region == "il-central-1" => Some("eu"),
"ap-northeast-1" | "ap-northeast-3" => Some("jp"),
"ap-southeast-2" | "ap-southeast-4" | "ap-southeast-6" => Some("au"),
region if region.starts_with("ap-") => Some("apac"),
_ => None,
}
}
#[cfg(test)]
#[path = "bedrock_tests.rs"]
mod tests;
@@ -0,0 +1,279 @@
use aws_sdk_bedrockruntime::config::Region;
use aws_smithy_http_client::test_util::NeverClient;
use futures::StreamExt;
use galaxy_agent_core::{
AgentEvent, AgentRuntime, ContentPart, ConversationMessage, MessageContent, MessageRole,
StopReason, ToolDefinition, TurnCommand, TurnRequest, Usage,
};
use rig_bedrock::streaming::{BedrockStreamingResponse, BedrockUsage};
use rig_core::completion::{AssistantContent, CompletionError, GetTokenUsage, Message};
use rig_core::message::{DocumentSourceKind, ToolResultContent, UserContent};
use super::*;
use crate::stream::{completion_error_stop_reason, map_usage};
#[test]
fn resolves_context_marker_and_us_inference_profile() {
assert_eq!(
resolve_bedrock_model_id("anthropic.claude-sonnet-4-6[1m]", "us-east-1", true,).unwrap(),
"us.anthropic.claude-sonnet-4-6"
);
}
#[test]
fn resolves_each_supported_inference_geography() {
for (region, expected_prefix) in [
("eu-west-1", "eu"),
("il-central-1", "eu"),
("ap-northeast-1", "jp"),
("ap-southeast-2", "au"),
("ap-southeast-1", "apac"),
("ca-central-1", "us"),
] {
assert_eq!(
resolve_bedrock_model_id("anthropic.claude-test", region, true).unwrap(),
format!("{expected_prefix}.anthropic.claude-test")
);
}
}
#[test]
fn preserves_arns_existing_profiles_and_unknown_regions() {
let arn = "arn:aws:bedrock:us-east-1:123:application-inference-profile/example";
assert_eq!(
resolve_bedrock_model_id(arn, "us-east-1", true).unwrap(),
arn
);
assert_eq!(
resolve_bedrock_model_id("global.anthropic.claude-test", "us-east-1", true).unwrap(),
"global.anthropic.claude-test"
);
assert_eq!(
resolve_bedrock_model_id("anthropic.claude-test", "me-south-1", true).unwrap(),
"anthropic.claude-test"
);
}
#[test]
fn prefixes_amazon_models_instead_of_mistaking_provider_for_geography() {
assert_eq!(
resolve_bedrock_model_id("amazon.nova-pro-v1:0", "us-east-1", true).unwrap(),
"us.amazon.nova-pro-v1:0"
);
}
#[test]
fn rejects_an_empty_model_id() {
let error = resolve_bedrock_model_id(" ", "us-east-1", false).unwrap_err();
assert_eq!(error.kind, AgentErrorKind::Configuration);
}
#[test]
fn normalizes_bedrock_usage_and_max_token_stop() {
let response = BedrockStreamingResponse {
usage: Some(BedrockUsage {
input_tokens: 100,
output_tokens: 25,
total_tokens: 125,
cache_read_input_tokens: Some(40),
cache_write_input_tokens: Some(10),
}),
};
assert_eq!(
map_usage(response.token_usage()),
Usage {
input_tokens: 100,
output_tokens: 25,
cached_input_tokens: 40,
cache_creation_input_tokens: 10,
}
);
assert_eq!(
completion_error_stop_reason(&CompletionError::ProviderError(
"Exceeded max tokens".to_string(),
)),
Some(StopReason::MaxTokens)
);
}
#[test]
fn constructs_rig_client_from_galaxys_resolved_aws_client_without_network() {
let sdk_config = aws_sdk_bedrockruntime::Config::builder()
.behavior_version_latest()
.region(Region::new("us-east-1"))
.http_client(NeverClient::new())
.build();
let aws_client = AwsBedrockClient::from_conf(sdk_config);
let client = BedrockRuntime::from_aws_client(
aws_client,
BedrockRigConfig {
model: "anthropic.claude-test[1M]".to_string(),
region: "us-east-1".to_string(),
cross_region_inference: true,
prompt_caching: true,
max_output_tokens: Some(8_192),
},
)
.unwrap();
assert_eq!(client.resolved_model(), "us.anthropic.claude-test");
let completion_model = client.completion_model();
assert_eq!(completion_model.model, client.resolved_model());
assert!(completion_model.prompt_caching);
}
#[tokio::test]
async fn cancellation_before_bedrock_stream_start_never_contacts_aws() {
let never_client = NeverClient::new();
let sdk_config = aws_sdk_bedrockruntime::Config::builder()
.behavior_version_latest()
.region(Region::new("us-east-1"))
.http_client(never_client.clone())
.build();
let runtime = BedrockRuntime::from_aws_client(
AwsBedrockClient::from_conf(sdk_config),
BedrockRigConfig {
model: "anthropic.claude-test".to_string(),
region: "us-east-1".to_string(),
cross_region_inference: false,
prompt_caching: false,
max_output_tokens: None,
},
)
.unwrap();
let (sender, control) = galaxy_agent_core::turn_control();
sender.send(TurnCommand::Cancel).await.unwrap();
let events = runtime
.start_turn(
TurnRequest::new(
"anthropic.claude-test",
vec![ConversationMessage {
role: MessageRole::User,
content: MessageContent::Text("Hello".to_string()),
}],
),
control,
)
.await
.unwrap()
.collect::<Vec<_>>()
.await
.into_iter()
.collect::<Result<Vec<_>, _>>()
.unwrap();
assert!(matches!(events[0], AgentEvent::TurnStarted { .. }));
assert_eq!(
events[1],
AgentEvent::TurnStopped {
reason: StopReason::Cancelled,
}
);
assert_eq!(never_client.num_calls(), 0);
}
#[test]
fn bedrock_request_preserves_system_image_reasoning_tool_and_token_semantics() {
let mut request = TurnRequest::new(
"anthropic.claude-test",
vec![
ConversationMessage {
role: MessageRole::User,
content: MessageContent::MultiPart(vec![
ContentPart::Text("Describe the image".to_string()),
ContentPart::Image {
data: vec![1, 2, 3, 4],
mime_type: "image/png".to_string(),
},
]),
},
ConversationMessage {
role: MessageRole::Assistant,
content: MessageContent::MultiPart(vec![
ContentPart::Reasoning {
text: "I should inspect the manifest.".to_string(),
signature: Some("signed-reasoning".to_string()),
},
ContentPart::ToolUse {
tool_use_id: "call-1".to_string(),
name: "read_files".to_string(),
input: serde_json::json!({"files": ["Cargo.toml"]}),
},
]),
},
ConversationMessage {
role: MessageRole::User,
content: MessageContent::ToolResult {
tool_use_id: "call-1".to_string(),
content: "permission denied".to_string(),
is_error: true,
},
},
],
);
request.system_prompt = Some("Use Galaxy tools safely".to_string());
request.max_output_tokens = Some(4_096);
request.tools.push(ToolDefinition {
name: "read_files".to_string(),
description: "Read project files".to_string(),
input_schema: serde_json::json!({
"type": "object",
"properties": {"files": {"type": "array"}}
}),
});
let converted = build_bedrock_completion_request(request, Some(8_192)).unwrap();
assert!(converted.additional_params.is_none());
assert_eq!(converted.max_tokens, Some(4_096));
assert_eq!(converted.tools.len(), 1);
assert_eq!(converted.tools[0].name, "read_files");
let messages = converted.chat_history.iter().collect::<Vec<_>>();
let [system, user, assistant, result] = messages.as_slice() else {
panic!("expected system, user, assistant, and tool-result messages");
};
assert!(matches!(
system,
Message::System { content } if content == "Use Galaxy tools safely"
));
let Message::User { content } = user else {
panic!("expected a user image message");
};
let user_content = content.iter().collect::<Vec<_>>();
assert!(matches!(
user_content.as_slice(),
[UserContent::Text(text), UserContent::Image(image)]
if text.text == "Describe the image"
&& matches!(&image.data, DocumentSourceKind::Base64(data) if data == "AQIDBA==")
));
let Message::Assistant { content, .. } = assistant else {
panic!("expected an assistant tool call");
};
let assistant_content = content.iter().collect::<Vec<_>>();
let [
AssistantContent::Reasoning(reasoning),
AssistantContent::ToolCall(call),
] = assistant_content.as_slice()
else {
panic!("expected signed reasoning followed by a tool call");
};
assert_eq!(reasoning.display_text(), "I should inspect the manifest.");
assert_eq!(reasoning.first_signature(), Some("signed-reasoning"));
assert_eq!(call.id, "call-1");
assert_eq!(call.function.name, "read_files");
let Message::User { content } = result else {
panic!("expected a user tool result");
};
let Some(UserContent::ToolResult(result)) = content.iter().next() else {
panic!("expected tool result content");
};
assert_eq!(result.id, "call-1");
assert!(matches!(
result.content.iter().next(),
Some(ToolResultContent::Text(text)) if text.text == "[ERROR] permission denied"
));
}
+4
View File
@@ -1,5 +1,9 @@
//! Rig-backed implementations of Galaxy's provider-neutral agent runtime. //! Rig-backed implementations of Galaxy's provider-neutral agent runtime.
mod bedrock;
mod openai_compatible; mod openai_compatible;
mod request;
mod stream;
pub use bedrock::*;
pub use openai_compatible::*; pub use openai_compatible::*;
+13 -342
View File
@@ -1,22 +1,14 @@
use async_trait::async_trait; use async_trait::async_trait;
use futures::{FutureExt, StreamExt};
use galaxy_agent_core::{ use galaxy_agent_core::{
AgentError, AgentErrorKind, AgentEvent, AgentEventStream, AgentRuntime, ContentPart, AgentError, AgentErrorKind, AgentEventStream, AgentRuntime, RuntimeCapabilities,
ConversationMessage, MessageContent, MessageRole, RuntimeCapabilities, RuntimeDescriptor, RuntimeDescriptor, RuntimeKind, TurnControl, TurnRequest,
RuntimeKind, StopReason, ToolCall, TurnCommand, TurnControl, TurnRequest, Usage,
}; };
use rig_core::OneOrMany;
use rig_core::client::CompletionClient; use rig_core::client::CompletionClient;
use rig_core::completion::{ use rig_core::completion::{CompletionModel, CompletionRequest};
AssistantContent, CompletionError, CompletionModel, CompletionRequest, GetTokenUsage, Message,
ToolDefinition,
};
use rig_core::message::{
DocumentSourceKind, Image, ImageMediaType, MimeType, ToolResultContent, UserContent,
};
use rig_core::providers::openai; use rig_core::providers::openai;
use rig_core::streaming::StreamedAssistantContent;
use uuid::Uuid; use crate::request::build_completion_request as build_provider_completion_request;
use crate::stream::start_model_turn as start_provider_model_turn;
#[derive(Clone, Debug, PartialEq, Eq)] #[derive(Clone, Debug, PartialEq, Eq)]
pub struct OpenAICompatibleRuntimeConfig { pub struct OpenAICompatibleRuntimeConfig {
@@ -91,143 +83,13 @@ where
M: CompletionModel + Send + Sync + 'static, M: CompletionModel + Send + Sync + 'static,
M::StreamingResponse: Send + Sync + 'static, M::StreamingResponse: Send + Sync + 'static,
{ {
let runtime_request_id = Uuid::new_v4().to_string();
let max_output_tokens = request.max_output_tokens.or(configured_max_output_tokens); let max_output_tokens = request.max_output_tokens.or(configured_max_output_tokens);
let completion_request = build_completion_request( let completion_request = build_completion_request(
request, request,
configured_max_output_tokens, configured_max_output_tokens,
supports_system_messages, supports_system_messages,
)?; )?;
let stream_future = model.stream(completion_request).fuse(); start_provider_model_turn(model, completion_request, control, max_output_tokens).await
let initial_control = control.clone();
let control_future = initial_control.receive().fuse();
futures::pin_mut!(stream_future, control_future);
let mut rig_stream = futures::select_biased! {
command = control_future => match command {
Ok(TurnCommand::Cancel) => {
return Ok(stopped_before_stream(runtime_request_id));
}
Ok(TurnCommand::Steer { .. }) | Err(_) => {
stream_future.await.map_err(map_completion_error)?
}
},
result = stream_future => result.map_err(map_completion_error)?,
};
let events = async_stream::stream! {
yield Ok(AgentEvent::TurnStarted {
runtime_request_id,
});
let mut control_open = true;
let mut last_output_tokens = 0;
loop {
let next_item = rig_stream.next().fuse();
let next_command = if control_open {
futures::future::Either::Left(control.receive())
} else {
futures::future::Either::Right(futures::future::pending())
}
.fuse();
futures::pin_mut!(next_item, next_command);
futures::select_biased! {
command = next_command => {
match command {
Ok(TurnCommand::Cancel) => {
rig_stream.cancel();
yield Ok(AgentEvent::TurnStopped {
reason: StopReason::Cancelled,
});
return;
}
Ok(TurnCommand::Steer { .. }) => {
// Steering is not advertised by this runtime yet.
}
Err(_) => control_open = false,
}
}
item = next_item => {
let Some(item) = item else {
yield Ok(AgentEvent::TurnStopped {
reason: if max_output_tokens.is_some_and(|max| {
last_output_tokens >= max
}) {
StopReason::MaxTokens
} else {
StopReason::Completed
},
});
return;
};
match item {
Ok(StreamedAssistantContent::Text(text)) => {
if !text.text.is_empty() {
yield Ok(AgentEvent::TextDelta { text: text.text });
}
}
Ok(StreamedAssistantContent::Reasoning(reasoning)) => {
let text = reasoning.display_text();
if !text.is_empty() {
yield Ok(AgentEvent::ReasoningDelta { text });
}
}
Ok(StreamedAssistantContent::ReasoningDelta { reasoning, .. }) => {
if !reasoning.is_empty() {
yield Ok(AgentEvent::ReasoningDelta { text: reasoning });
}
}
Ok(StreamedAssistantContent::ToolCall { tool_call, .. }) => {
yield Ok(AgentEvent::Tool {
event: galaxy_agent_core::ToolEvent::Proposed {
call: ToolCall {
id: tool_call.id,
name: tool_call.function.name,
arguments: tool_call.function.arguments,
},
},
});
}
Ok(StreamedAssistantContent::ToolCallDelta { .. }) => {
// Rig emits a complete ToolCall after its deltas, which
// is the canonical event Galaxy consumes.
}
Ok(StreamedAssistantContent::Final(response)) => {
let mapped_usage = map_usage(response.token_usage());
last_output_tokens = mapped_usage.output_tokens;
yield Ok(AgentEvent::UsageUpdated {
usage: mapped_usage,
});
}
Ok(StreamedAssistantContent::Unknown(value)) => {
yield Err(AgentError::new(
AgentErrorKind::Protocol,
format!("Rig returned an unsupported provider event: {value}"),
));
return;
}
Err(error) => {
yield Err(map_completion_error(error));
return;
}
}
}
}
}
};
Ok(Box::pin(events))
}
fn stopped_before_stream(runtime_request_id: String) -> AgentEventStream {
Box::pin(futures::stream::iter([
Ok(AgentEvent::TurnStarted { runtime_request_id }),
Ok(AgentEvent::TurnStopped {
reason: StopReason::Cancelled,
}),
]))
} }
fn build_completion_request( fn build_completion_request(
@@ -235,208 +97,17 @@ fn build_completion_request(
configured_max_output_tokens: Option<u64>, configured_max_output_tokens: Option<u64>,
supports_system_messages: bool, supports_system_messages: bool,
) -> Result<CompletionRequest, AgentError> { ) -> Result<CompletionRequest, AgentError> {
let mut messages = Vec::new(); build_provider_completion_request(
if let Some(system_prompt) = request.system_prompt { request,
if supports_system_messages { configured_max_output_tokens,
messages.push(Message::System { supports_system_messages,
content: system_prompt, false,
}); Some(serde_json::json!({
} else {
messages.push(Message::User {
content: OneOrMany::one(UserContent::text(system_prompt)),
});
}
}
for message in request.messages {
messages.push(convert_message(message)?);
}
let chat_history = OneOrMany::many(messages).map_err(|_| {
AgentError::new(
AgentErrorKind::InvalidRequest,
"a Rig turn requires at least one conversation message",
)
})?;
Ok(CompletionRequest {
model: Some(request.model.as_str().to_string()),
preamble: None,
chat_history,
documents: Vec::new(),
tools: request
.tools
.into_iter()
.map(|tool| ToolDefinition {
name: tool.name,
description: tool.description,
parameters: tool.input_schema,
})
.collect(),
temperature: None,
max_tokens: request.max_output_tokens.or(configured_max_output_tokens),
tool_choice: None,
additional_params: Some(serde_json::json!({
"stream_options": { "include_usage": true } "stream_options": { "include_usage": true }
})), })),
output_schema: None,
})
}
fn convert_message(message: ConversationMessage) -> Result<Message, AgentError> {
match message.role {
MessageRole::User => Ok(Message::User {
content: user_content(message.content)?,
}),
MessageRole::Assistant => Ok(Message::Assistant {
id: None,
content: assistant_content(message.content)?,
}),
}
}
fn user_content(content: MessageContent) -> Result<OneOrMany<UserContent>, AgentError> {
let parts = match content {
MessageContent::Text(text) => vec![UserContent::text(text)],
MessageContent::ToolResult {
tool_use_id,
content,
is_error,
} => vec![UserContent::tool_result(
tool_use_id,
OneOrMany::one(ToolResultContent::text(tool_result_text(content, is_error))),
)],
MessageContent::MultiPart(parts) => parts
.into_iter()
.map(convert_user_part)
.collect::<Result<Vec<_>, _>>()?,
MessageContent::ToolUse { .. } => {
return Err(invalid_role("tool use", "user"));
}
};
one_or_many(parts, "user")
}
fn assistant_content(content: MessageContent) -> Result<OneOrMany<AssistantContent>, AgentError> {
let parts = match content {
MessageContent::Text(text) => vec![AssistantContent::text(text)],
MessageContent::ToolUse {
tool_use_id,
name,
input,
} => vec![AssistantContent::tool_call(tool_use_id, name, input)],
MessageContent::MultiPart(parts) => parts
.into_iter()
.map(convert_assistant_part)
.collect::<Result<Vec<_>, _>>()?,
MessageContent::ToolResult { .. } => {
return Err(invalid_role("tool result", "assistant"));
}
};
one_or_many(parts, "assistant")
}
fn convert_user_part(part: ContentPart) -> Result<UserContent, AgentError> {
match part {
ContentPart::Text(text) => Ok(UserContent::text(text)),
ContentPart::Image { data, mime_type } => Ok(UserContent::image_raw(
data,
ImageMediaType::from_mime_type(&mime_type),
None,
)),
ContentPart::ToolResult {
tool_use_id,
content,
is_error,
} => Ok(UserContent::tool_result(
tool_use_id,
OneOrMany::one(ToolResultContent::text(tool_result_text(content, is_error))),
)),
ContentPart::ToolUse { .. } => Err(invalid_role("tool use", "user")),
}
}
fn convert_assistant_part(part: ContentPart) -> Result<AssistantContent, AgentError> {
match part {
ContentPart::Text(text) => Ok(AssistantContent::text(text)),
ContentPart::Image { data, mime_type } => Ok(AssistantContent::Image(Image {
data: DocumentSourceKind::Raw(data),
media_type: ImageMediaType::from_mime_type(&mime_type),
detail: None,
additional_params: None,
})),
ContentPart::ToolUse {
tool_use_id,
name,
input,
} => Ok(AssistantContent::tool_call(tool_use_id, name, input)),
ContentPart::ToolResult { .. } => Err(invalid_role("tool result", "assistant")),
}
}
fn tool_result_text(content: String, is_error: bool) -> String {
if is_error {
format!("[ERROR] {content}")
} else {
content
}
}
fn one_or_many<T: Clone>(parts: Vec<T>, role: &str) -> Result<OneOrMany<T>, AgentError> {
OneOrMany::many(parts).map_err(|_| {
AgentError::new(
AgentErrorKind::InvalidRequest,
format!("{role} message has no content"),
)
})
}
fn invalid_role(content: &str, role: &str) -> AgentError {
AgentError::new(
AgentErrorKind::InvalidRequest,
format!("{content} content cannot appear in a {role} message"),
) )
} }
fn map_usage(usage: rig_core::completion::Usage) -> Usage {
Usage {
input_tokens: usage.input_tokens,
output_tokens: usage.output_tokens,
cached_input_tokens: usage.cached_input_tokens,
cache_creation_input_tokens: usage.cache_creation_input_tokens,
}
}
fn map_completion_error(error: CompletionError) -> AgentError {
let status = error
.provider_response_status()
.map(|status| status.as_u16());
let kind = match status {
Some(401 | 403) => AgentErrorKind::Authentication,
Some(429) => AgentErrorKind::RateLimited,
Some(400 | 404 | 413 | 422) => AgentErrorKind::InvalidRequest,
Some(500..=599) => AgentErrorKind::Provider,
Some(_) => AgentErrorKind::Provider,
None => match &error {
CompletionError::HttpError(_)
| CompletionError::UrlError(_)
| CompletionError::RequestError(_) => AgentErrorKind::Transport,
CompletionError::JsonError(_) | CompletionError::ResponseError(_) => {
AgentErrorKind::Protocol
}
CompletionError::ProviderError(_) | CompletionError::ProviderResponse(_) => {
AgentErrorKind::Provider
}
_ => AgentErrorKind::Provider,
},
};
let mut mapped = AgentError::new(kind, error.to_string());
mapped.recoverable = matches!(
kind,
AgentErrorKind::RateLimited | AgentErrorKind::Transport
);
mapped
}
#[cfg(test)] #[cfg(test)]
#[path = "openai_compatible_tests.rs"] #[path = "openai_compatible_tests.rs"]
mod tests; mod tests;
@@ -1,8 +1,11 @@
use futures::StreamExt; use futures::StreamExt;
use galaxy_agent_core::{ use galaxy_agent_core::{
AgentEvent, AgentRuntime, ConversationMessage, MessageContent, MessageRole, ToolEvent, AgentEvent, AgentRuntime, ConversationMessage, MessageContent, MessageRole, StopReason,
ToolEvent, TurnCommand, Usage,
}; };
use rig_core::client::CompletionClient; use rig_core::client::CompletionClient;
use rig_core::completion::{AssistantContent, Message};
use rig_core::message::{ToolResultContent, UserContent};
use rig_core::providers::openai; use rig_core::providers::openai;
use rig_core::test_utils::MockStreamingClient; use rig_core::test_utils::MockStreamingClient;
+211
View File
@@ -0,0 +1,211 @@
use base64::Engine;
use base64::prelude::BASE64_STANDARD;
use galaxy_agent_core::{
AgentError, AgentErrorKind, ContentPart, ConversationMessage, MessageContent, MessageRole,
TurnRequest,
};
use rig_core::OneOrMany;
use rig_core::completion::{AssistantContent, CompletionRequest, Message, ToolDefinition};
use rig_core::message::{
DocumentSourceKind, Image, ImageMediaType, MimeType, Reasoning, ToolResultContent, UserContent,
};
pub(crate) fn build_completion_request(
request: TurnRequest,
configured_max_output_tokens: Option<u64>,
supports_system_messages: bool,
encode_images_as_base64: bool,
additional_params: Option<serde_json::Value>,
) -> Result<CompletionRequest, AgentError> {
let mut messages = Vec::new();
if let Some(system_prompt) = request.system_prompt {
if supports_system_messages {
messages.push(Message::System {
content: system_prompt,
});
} else {
messages.push(Message::User {
content: OneOrMany::one(UserContent::text(system_prompt)),
});
}
}
for message in request.messages {
messages.push(convert_message(message, encode_images_as_base64)?);
}
let chat_history = OneOrMany::many(messages).map_err(|_| {
AgentError::new(
AgentErrorKind::InvalidRequest,
"a Rig turn requires at least one conversation message",
)
})?;
Ok(CompletionRequest {
model: Some(request.model.as_str().to_string()),
preamble: None,
chat_history,
documents: Vec::new(),
tools: request
.tools
.into_iter()
.map(|tool| ToolDefinition {
name: tool.name,
description: tool.description,
parameters: tool.input_schema,
})
.collect(),
temperature: None,
max_tokens: request.max_output_tokens.or(configured_max_output_tokens),
tool_choice: None,
additional_params,
output_schema: None,
})
}
fn convert_message(
message: ConversationMessage,
encode_images_as_base64: bool,
) -> Result<Message, AgentError> {
match message.role {
MessageRole::User => Ok(Message::User {
content: user_content(message.content, encode_images_as_base64)?,
}),
MessageRole::Assistant => Ok(Message::Assistant {
id: None,
content: assistant_content(message.content, encode_images_as_base64)?,
}),
}
}
fn user_content(
content: MessageContent,
encode_images_as_base64: bool,
) -> Result<OneOrMany<UserContent>, AgentError> {
let parts = match content {
MessageContent::Text(text) => vec![UserContent::text(text)],
MessageContent::ToolResult {
tool_use_id,
content,
is_error,
} => vec![UserContent::tool_result(
tool_use_id,
OneOrMany::one(ToolResultContent::text(tool_result_text(content, is_error))),
)],
MessageContent::MultiPart(parts) => parts
.into_iter()
.map(|part| convert_user_part(part, encode_images_as_base64))
.collect::<Result<Vec<_>, _>>()?,
MessageContent::ToolUse { .. } => {
return Err(invalid_role("tool use", "user"));
}
};
one_or_many(parts, "user")
}
fn assistant_content(
content: MessageContent,
encode_images_as_base64: bool,
) -> Result<OneOrMany<AssistantContent>, AgentError> {
let parts = match content {
MessageContent::Text(text) => vec![AssistantContent::text(text)],
MessageContent::ToolUse {
tool_use_id,
name,
input,
} => vec![AssistantContent::tool_call(tool_use_id, name, input)],
MessageContent::MultiPart(parts) => parts
.into_iter()
.map(|part| convert_assistant_part(part, encode_images_as_base64))
.collect::<Result<Vec<_>, _>>()?,
MessageContent::ToolResult { .. } => {
return Err(invalid_role("tool result", "assistant"));
}
};
one_or_many(parts, "assistant")
}
fn convert_user_part(
part: ContentPart,
encode_images_as_base64: bool,
) -> Result<UserContent, AgentError> {
match part {
ContentPart::Text(text) => Ok(UserContent::text(text)),
ContentPart::Reasoning { .. } => Err(invalid_role("reasoning", "user")),
ContentPart::Image { data, mime_type } => {
let media_type = ImageMediaType::from_mime_type(&mime_type);
if encode_images_as_base64 {
Ok(UserContent::image_base64(
BASE64_STANDARD.encode(data),
media_type,
None,
))
} else {
Ok(UserContent::image_raw(data, media_type, None))
}
}
ContentPart::ToolResult {
tool_use_id,
content,
is_error,
} => Ok(UserContent::tool_result(
tool_use_id,
OneOrMany::one(ToolResultContent::text(tool_result_text(content, is_error))),
)),
ContentPart::ToolUse { .. } => Err(invalid_role("tool use", "user")),
}
}
fn convert_assistant_part(
part: ContentPart,
encode_images_as_base64: bool,
) -> Result<AssistantContent, AgentError> {
match part {
ContentPart::Text(text) => Ok(AssistantContent::text(text)),
ContentPart::Reasoning { text, signature } => Ok(AssistantContent::Reasoning(
Reasoning::new_with_signature(&text, signature),
)),
ContentPart::Image { data, mime_type } => Ok(AssistantContent::Image(Image {
data: if encode_images_as_base64 {
DocumentSourceKind::Base64(BASE64_STANDARD.encode(data))
} else {
DocumentSourceKind::Raw(data)
},
media_type: ImageMediaType::from_mime_type(&mime_type),
detail: None,
additional_params: None,
})),
ContentPart::ToolUse {
tool_use_id,
name,
input,
} => Ok(AssistantContent::tool_call(tool_use_id, name, input)),
ContentPart::ToolResult { .. } => Err(invalid_role("tool result", "assistant")),
}
}
fn tool_result_text(content: String, is_error: bool) -> String {
// Rig core does not yet carry Bedrock's optional ToolResultStatus. Keep
// Galaxy's structured error state in the domain model and make the error
// semantic explicit in the provider-visible result text.
if is_error {
format!("[ERROR] {content}")
} else {
content
}
}
fn one_or_many<T: Clone>(parts: Vec<T>, role: &str) -> Result<OneOrMany<T>, AgentError> {
OneOrMany::many(parts).map_err(|_| {
AgentError::new(
AgentErrorKind::InvalidRequest,
format!("{role} message has no content"),
)
})
}
fn invalid_role(content: &str, role: &str) -> AgentError {
AgentError::new(
AgentErrorKind::InvalidRequest,
format!("{content} content cannot appear in a {role} message"),
)
}
+208
View File
@@ -0,0 +1,208 @@
use futures::{FutureExt, StreamExt};
use galaxy_agent_core::{
AgentError, AgentErrorKind, AgentEvent, AgentEventStream, StopReason, ToolCall, TurnCommand,
TurnControl, Usage,
};
use rig_core::completion::{CompletionError, CompletionModel, CompletionRequest, GetTokenUsage};
use rig_core::streaming::StreamedAssistantContent;
use uuid::Uuid;
pub(crate) async fn start_model_turn<M>(
model: M,
completion_request: CompletionRequest,
control: TurnControl,
max_output_tokens: Option<u64>,
) -> Result<AgentEventStream, AgentError>
where
M: CompletionModel + Send + Sync + 'static,
M::StreamingResponse: Send + Sync + 'static,
{
let runtime_request_id = Uuid::new_v4().to_string();
let stream_future = model.stream(completion_request).fuse();
let initial_control = control.clone();
let control_future = initial_control.receive().fuse();
futures::pin_mut!(stream_future, control_future);
let mut rig_stream = futures::select_biased! {
command = control_future => match command {
Ok(TurnCommand::Cancel) => {
return Ok(stopped_before_stream(runtime_request_id));
}
Ok(TurnCommand::Steer { .. }) | Err(_) => {
stream_future.await.map_err(map_completion_error)?
}
},
result = stream_future => result.map_err(map_completion_error)?,
};
let events = async_stream::stream! {
yield Ok(AgentEvent::TurnStarted {
runtime_request_id,
});
let mut control_open = true;
let mut last_output_tokens = 0;
loop {
let next_item = rig_stream.next().fuse();
let next_command = if control_open {
futures::future::Either::Left(control.receive())
} else {
futures::future::Either::Right(futures::future::pending())
}
.fuse();
futures::pin_mut!(next_item, next_command);
futures::select_biased! {
command = next_command => {
match command {
Ok(TurnCommand::Cancel) => {
rig_stream.cancel();
yield Ok(AgentEvent::TurnStopped {
reason: StopReason::Cancelled,
});
return;
}
Ok(TurnCommand::Steer { .. }) => {
// Steering is not advertised by provider runtimes yet.
}
Err(_) => control_open = false,
}
}
item = next_item => {
let Some(item) = item else {
yield Ok(AgentEvent::TurnStopped {
reason: if max_output_tokens.is_some_and(|max| {
last_output_tokens >= max
}) {
StopReason::MaxTokens
} else {
StopReason::Completed
},
});
return;
};
match item {
Ok(StreamedAssistantContent::Text(text)) => {
if !text.text.is_empty() {
yield Ok(AgentEvent::TextDelta { text: text.text });
}
}
Ok(StreamedAssistantContent::Reasoning(reasoning)) => {
let text = reasoning.display_text();
yield Ok(AgentEvent::ReasoningCompleted {
text,
signature: reasoning.first_signature().map(str::to_string),
});
}
Ok(StreamedAssistantContent::ReasoningDelta { reasoning, .. }) => {
if !reasoning.is_empty() {
yield Ok(AgentEvent::ReasoningDelta { text: reasoning });
}
}
Ok(StreamedAssistantContent::ToolCall { tool_call, .. }) => {
yield Ok(AgentEvent::Tool {
event: galaxy_agent_core::ToolEvent::Proposed {
call: ToolCall {
id: tool_call.id,
name: tool_call.function.name,
arguments: tool_call.function.arguments,
},
},
});
}
Ok(StreamedAssistantContent::ToolCallDelta { .. }) => {
// Rig emits a complete ToolCall after its deltas, which
// is the canonical event Galaxy consumes.
}
Ok(StreamedAssistantContent::Final(response)) => {
let mapped_usage = map_usage(response.token_usage());
last_output_tokens = mapped_usage.output_tokens;
yield Ok(AgentEvent::UsageUpdated {
usage: mapped_usage,
});
}
Ok(StreamedAssistantContent::Unknown(value)) => {
yield Err(AgentError::new(
AgentErrorKind::Protocol,
format!("Rig returned an unsupported provider event: {value}"),
));
return;
}
Err(error) => {
if let Some(reason) = completion_error_stop_reason(&error) {
yield Ok(AgentEvent::TurnStopped { reason });
return;
}
yield Err(map_completion_error(error));
return;
}
}
}
}
}
};
Ok(Box::pin(events))
}
fn stopped_before_stream(runtime_request_id: String) -> AgentEventStream {
Box::pin(futures::stream::iter([
Ok(AgentEvent::TurnStarted { runtime_request_id }),
Ok(AgentEvent::TurnStopped {
reason: StopReason::Cancelled,
}),
]))
}
pub(crate) fn map_usage(usage: rig_core::completion::Usage) -> Usage {
Usage {
input_tokens: usage.input_tokens,
output_tokens: usage.output_tokens,
cached_input_tokens: usage.cached_input_tokens,
cache_creation_input_tokens: usage.cache_creation_input_tokens,
}
}
pub(crate) fn completion_error_stop_reason(error: &CompletionError) -> Option<StopReason> {
match error {
// rig-bedrock 0.40 currently surfaces Bedrock's MaxTokens stop as a
// provider error. Normalize it here so the UI sees the same semantic
// stop reason as every other Rig-backed provider.
CompletionError::ProviderError(message) if message == "Exceeded max tokens" => {
Some(StopReason::MaxTokens)
}
_ => None,
}
}
fn map_completion_error(error: CompletionError) -> AgentError {
let status = error
.provider_response_status()
.map(|status| status.as_u16());
let kind = match status {
Some(401 | 403) => AgentErrorKind::Authentication,
Some(429) => AgentErrorKind::RateLimited,
Some(400 | 404 | 413 | 422) => AgentErrorKind::InvalidRequest,
Some(500..=599) => AgentErrorKind::Provider,
Some(_) => AgentErrorKind::Provider,
None => match &error {
CompletionError::HttpError(_)
| CompletionError::UrlError(_)
| CompletionError::RequestError(_) => AgentErrorKind::Transport,
CompletionError::JsonError(_) | CompletionError::ResponseError(_) => {
AgentErrorKind::Protocol
}
CompletionError::ProviderError(_) | CompletionError::ProviderResponse(_) => {
AgentErrorKind::Provider
}
_ => AgentErrorKind::Provider,
},
};
let mut mapped = AgentError::new(kind, error.to_string());
mapped.recoverable = matches!(
kind,
AgentErrorKind::RateLimited | AgentErrorKind::Transport
);
mapped
}
+24 -7
View File
@@ -336,9 +336,26 @@ rendering and non-Rig compatibility runtimes, not in Rig's executable tool path.
### Phase 4 — Bedrock through Rig ### Phase 4 — Bedrock through Rig
- Implement Bedrock client construction and model resolution through `rig-bedrock`. - [x] Pin `rig-bedrock` 0.40.0 and construct it from Galaxy's already-resolved AWS SDK client so
- Compare request behavior for system prompts, images, tool schemas, cache controls, reasoning, profile, SSO, static-key, region, and egress ownership stay at Galaxy's explicit boundary.
inference profiles, token usage, and context limits. - [x] Resolve context markers, ARNs, existing inference profiles, and regional inference-profile
prefixes before passing a model ID to Rig.
- [x] Reuse one Galaxy-to-Rig request adapter and one Rig-to-`AgentEvent` streaming lifecycle for
OpenAI-compatible and Bedrock providers; handle Bedrock's required base64 image representation at
that single request boundary.
- [x] Add hermetic compatibility fixtures for system prompts, images, tool calls/results, cache
enablement, cancellation, inference profiles, token limits, usage/cache usage, and max-token stop
normalization without contacting AWS.
- [x] Preserve signed Bedrock reasoning blocks in Galaxy conversation history so adaptive-thinking
tool-call turns can be replayed without losing their signatures.
- [x] Define the Rig 0.40 parity policy: Galaxy retains structured tool-result error state locally
and sends an explicit `[ERROR]` result prefix because Rig core has no Bedrock status field;
Rig owns system/message cache checkpoints, tool-schema caching is treated as an optimization,
and one-hour cache-TTL requests stay on the compatibility runtime.
- [x] Add a model-by-model Rig switch to the unified Models page and route opted-in Bedrock models
through the same request, event, permission, history, and UI adapter as OpenAI-compatible models.
- [ ] Run opt-in live semantic comparisons for system prompts, images, tools, reasoning, usage, and
context limits before selecting the Rig runtime for any configured Bedrock model.
- Keep a short-lived compatibility fallback for unsupported Bedrock behavior, measured by tests. - Keep a short-lived compatibility fallback for unsupported Bedrock behavior, measured by tests.
- Delete custom Bedrock translation code only after parity is proven. - Delete custom Bedrock translation code only after parity is proven.
@@ -418,7 +435,7 @@ contract is what the UI and persistence observe.
## Immediate next vertical slice ## Immediate next vertical slice
Begin Phase 4 with a focused `rig-bedrock` compatibility spike. Establish client construction and Finish Phase 4 with opt-in live Bedrock semantic comparisons for system prompts, images, tools,
model/inference-profile resolution first, then add semantic parity fixtures for system prompts, signed reasoning, usage, cancellation, and context limits. Keep per-model Rig routing opt-in until
images, tool schemas, cache controls, reasoning, usage, and context limits before routing any those live fixtures pass, then make Rig the default for supported models and retain the compatibility
configured Bedrock model away from the existing compatibility implementation. runtime only for explicitly unsupported cache behavior.