From b0ad07f6f2e8e394ed1988ae7922bf1fd713a9a6 Mon Sep 17 00:00:00 2001 From: Ryan Ward Date: Tue, 4 Aug 2026 17:25:19 -0500 Subject: [PATCH] Add unified models UI and Rig Bedrock runtime --- Cargo.lock | 100 ++++ Cargo.toml | 3 + app/Cargo.toml | 2 +- app/src/ai/agent/api/impl.rs | 28 +- app/src/ai/bedrock/client.rs | 29 +- app/src/ai/bedrock/convert.rs | 14 +- app/src/ai/bedrock/diagnostic.rs | 8 + app/src/ai/bedrock/e2e_tests.rs | 1 + app/src/ai/bedrock/external_config.rs | 1 + app/src/ai/bedrock/integration_tests.rs | 1 + app/src/ai/bedrock/models.rs | 29 + app/src/ai/bedrock/models_tests.rs | 34 +- app/src/ai/bedrock/request_translator.rs | 17 +- app/src/ai/bedrock/translator.rs | 7 + app/src/ai/blocklist/controller.rs | 7 + .../blocklist/controller/response_stream.rs | 22 +- .../ai/blocklist/usage/context_window_view.rs | 9 + app/src/ai/crosscheck/reviewer.rs | 1 + app/src/ai/llms.rs | 133 ++++- app/src/ai/llms_tests.rs | 66 ++- app/src/ai/openai/convert.rs | 9 + app/src/ai/runtime/mod.rs | 2 +- app/src/ai/runtime/rig.rs | 130 ++++- app/src/ai/runtime/rig_request.rs | 58 +- app/src/ai/runtime/rig_request_tests.rs | 36 +- app/src/ai/runtime/rig_tests.rs | 45 +- app/src/settings/ai.rs | 11 +- app/src/settings_view/ai_page.rs | 548 +++++++++++++----- app/src/settings_view/mod.rs | 17 +- app/src/settings_view/mod_tests.rs | 13 +- crates/galaxy_agent_core/src/tool_policy.rs | 4 +- crates/galaxy_agent_core/src/types.rs | 42 +- crates/galaxy_agent_rig/Cargo.toml | 4 + crates/galaxy_agent_rig/src/bedrock.rs | 162 ++++++ crates/galaxy_agent_rig/src/bedrock_tests.rs | 279 +++++++++ crates/galaxy_agent_rig/src/lib.rs | 4 + .../galaxy_agent_rig/src/openai_compatible.rs | 355 +----------- .../src/openai_compatible_tests.rs | 5 +- crates/galaxy_agent_rig/src/request.rs | 211 +++++++ crates/galaxy_agent_rig/src/stream.rs | 208 +++++++ plans/galaxy-local-first-rig.md | 31 +- 41 files changed, 2122 insertions(+), 564 deletions(-) create mode 100644 crates/galaxy_agent_rig/src/bedrock.rs create mode 100644 crates/galaxy_agent_rig/src/bedrock_tests.rs create mode 100644 crates/galaxy_agent_rig/src/request.rs create mode 100644 crates/galaxy_agent_rig/src/stream.rs diff --git a/Cargo.lock b/Cargo.lock index 931bb13e..b457131c 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1745,17 +1745,23 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "635d23afda0a6ab48d666c4d447c4873e8d1e83518a2be2093122397e50b838e" dependencies = [ "aws-smithy-async", + "aws-smithy-protocol-test", "aws-smithy-runtime-api", "aws-smithy-types", + "bytes", "h2", "http 1.5.0", + "http-body 1.1.0", "hyper", "hyper-rustls", "hyper-util", + "indexmap 2.14.0", "pin-project-lite", "rustls", "rustls-native-certs", "rustls-pki-types", + "serde", + "serde_json", "tokio", "tokio-rustls", "tower", @@ -1782,6 +1788,25 @@ dependencies = [ "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]] name = "aws-smithy-query" version = "0.62.0" @@ -2698,6 +2723,25 @@ dependencies = [ "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]] name = "cc" version = "1.4.0" @@ -4386,6 +4430,12 @@ dependencies = [ "syn 2.0.119", ] +[[package]] +name = "diff" +version = "0.1.13" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "56254986775e3233ffa9c4d7d3faaf6d36a2c09d30b20687e9f88bc8bafc16c8" + [[package]] name = "difflib" version = "0.4.0" @@ -5897,9 +5947,13 @@ version = "0.1.0" dependencies = [ "async-stream", "async-trait", + "aws-sdk-bedrockruntime", + "aws-smithy-http-client", + "base64 0.22.1", "bytes", "futures", "galaxy_agent_core", + "rig-bedrock", "rig-core", "serde_json", "tokio", @@ -11556,6 +11610,16 @@ version = "0.3.1" source = "registry+https://github.com/rust-lang/crates.io-index" 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]] name = "prettyplease" version = "0.2.37" @@ -12778,6 +12842,27 @@ dependencies = [ "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]] name = "rig-core" version = "0.40.0" @@ -12942,6 +13027,15 @@ dependencies = [ "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]] name = "roxmltree" version = "0.20.0" @@ -13525,6 +13619,12 @@ version = "0.6.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "cd0b0ec5f1c1ca621c432a25813d8d60c88abe6d3e08a3eb9cf37d97a0fe3d73" +[[package]] +name = "separator" +version = "0.4.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f97841a747eef040fcd2e7b3b9a220a7205926e60488e673d9e4926d27772ce5" + [[package]] name = "seq-macro" version = "0.3.6" diff --git a/Cargo.toml b/Cargo.toml index e719cc72..7794aa9c 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -138,6 +138,8 @@ async-stream = "0.3.5" async-task = "4.2.0" async-trait = "0.1.89" 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" base64 = "0.22" bincode = "1.3.3" @@ -260,6 +262,7 @@ reqwest = { version = "0.13", features = [ ] } reqwest-eventsource = { package = "aha-reqwest-eventsource", version = "0.1" } rig-core = "=0.40.0" +rig-bedrock = "=0.40.0" resvg = "0.47.0" rust-embed = { version = "8.7.0", features = ["include-exclude"] } rustc-hash = "2.1.1" diff --git a/app/Cargo.toml b/app/Cargo.toml index 91091d99..847d639f 100644 --- a/app/Cargo.toml +++ b/app/Cargo.toml @@ -328,7 +328,7 @@ tracing-subscriber.workspace = true # AWS SDK (loading credentials for BYO LLM) aws-config = { version = "1.8.16", features = ["credentials-login"] } 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-smithy-types = "1" aws-types = "1" diff --git a/app/src/ai/agent/api/impl.rs b/app/src/ai/agent/api/impl.rs index 1b2529c4..cd46d917 100644 --- a/app/src/ai/agent/api/impl.rs +++ b/app/src/ai/agent/api/impl.rs @@ -28,8 +28,8 @@ pub async fn generate_multi_agent_output( redaction::redact_inputs(&mut params.input); } - if let ProviderConfig::OpenAI(config) = &provider_config { - if config.use_rig { + match &provider_config { + ProviderConfig::OpenAI(config) if config.use_rig => { return Ok(crate::ai::runtime::rig_openai_response_stream( config.clone(), params, @@ -38,6 +38,30 @@ pub async fn generate_multi_agent_output( 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(); diff --git a/app/src/ai/bedrock/client.rs b/app/src/ai/bedrock/client.rs index fd133056..44a4bb9f 100644 --- a/app/src/ai/bedrock/client.rs +++ b/app/src/ai/bedrock/client.rs @@ -5,6 +5,8 @@ use aws_config::BehaviorVersion; use aws_credential_types::provider::ProvideCredentials; use aws_sdk_bedrockruntime::config::Region; 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::diagnostic::BedrockDiagnosticLogger; @@ -38,6 +40,7 @@ pub struct BedrockClientConfig { pub secret_access_key: String, pub session_token: Option, pub cross_region_inference: bool, + pub use_rig: bool, } impl BedrockClientConfig { @@ -141,8 +144,8 @@ impl BedrockClient { match provider.provide_credentials().await { Ok(creds) => { log::info!( - "[bedrock] Resolved AWS credentials successfully: access_key_id={:?}, has_session_token={}, expiry={:?}", - creds.access_key_id(), + "[bedrock] Resolved AWS credentials successfully: has_access_key_id={}, has_session_token={}, expiry={:?}", + !creds.access_key_id().is_empty(), creds.session_token().is_some(), 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, + ) -> Result { + 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)] pub async fn converse_stream( &self, diff --git a/app/src/ai/bedrock/convert.rs b/app/src/ai/bedrock/convert.rs index a4306367..dd3b86b8 100644 --- a/app/src/ai/bedrock/convert.rs +++ b/app/src/ai/bedrock/convert.rs @@ -3,8 +3,9 @@ use std::collections::HashMap; use aws_sdk_bedrockruntime::types::{ CachePointBlock, CachePointType, CacheTtl, ContentBlock, ConversationRole, ImageBlock, ImageFormat, ImageSource, InferenceConfiguration, Message as BedrockMessage, - SystemContentBlock, Tool, ToolConfiguration, ToolInputSchema, ToolResultBlock, - ToolResultContentBlock, ToolResultStatus, ToolSpecification, ToolUseBlock, + ReasoningContentBlock, ReasoningTextBlock, SystemContentBlock, Tool, ToolConfiguration, + ToolInputSchema, ToolResultBlock, ToolResultContentBlock, ToolResultStatus, ToolSpecification, + ToolUseBlock, }; use aws_smithy_types::{Blob, Document}; use serde_json::Value as JsonValue; @@ -148,6 +149,15 @@ fn convert_messages( .into_iter() .map(|part| match part { 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::ToolUse { tool_use_id, diff --git a/app/src/ai/bedrock/diagnostic.rs b/app/src/ai/bedrock/diagnostic.rs index a346cb73..d696a631 100644 --- a/app/src/ai/bedrock/diagnostic.rs +++ b/app/src/ai/bedrock/diagnostic.rs @@ -492,6 +492,14 @@ fn serialize_messages(messages: &[ConversationMessage]) -> JsonValue { super::convert::ContentPart::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 } => { serde_json::json!({ "type": "image", diff --git a/app/src/ai/bedrock/e2e_tests.rs b/app/src/ai/bedrock/e2e_tests.rs index 1fac4a10..3cfa7111 100644 --- a/app/src/ai/bedrock/e2e_tests.rs +++ b/app/src/ai/bedrock/e2e_tests.rs @@ -269,6 +269,7 @@ fn get_test_config() -> Option { secret_access_key: String::new(), session_token: None, cross_region_inference: false, + use_rig: false, }) } diff --git a/app/src/ai/bedrock/external_config.rs b/app/src/ai/bedrock/external_config.rs index 3cff89b6..2f9cd94c 100644 --- a/app/src/ai/bedrock/external_config.rs +++ b/app/src/ai/bedrock/external_config.rs @@ -144,6 +144,7 @@ fn parse_claude_code_model_map( model_id: arn, display_name, vision_supported: true, + use_rig: false, } }) .collect() diff --git a/app/src/ai/bedrock/integration_tests.rs b/app/src/ai/bedrock/integration_tests.rs index 155b3424..27ceb715 100644 --- a/app/src/ai/bedrock/integration_tests.rs +++ b/app/src/ai/bedrock/integration_tests.rs @@ -24,6 +24,7 @@ fn get_test_config() -> Option { secret_access_key: String::new(), session_token: None, cross_region_inference: false, + use_rig: false, }) } diff --git a/app/src/ai/bedrock/models.rs b/app/src/ai/bedrock/models.rs index e4bf7dfe..acbc5abd 100644 --- a/app/src/ai/bedrock/models.rs +++ b/app/src/ai/bedrock/models.rs @@ -129,6 +129,7 @@ pub fn get_effective_models(user_models: &[BedrockModelConfig]) -> Vec Vec 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 { if model_id.starts_with("arn:") { return model_id.to_string(); diff --git a/app/src/ai/bedrock/models_tests.rs b/app/src/ai/bedrock/models_tests.rs index f65b7423..c283a0a3 100644 --- a/app/src/ai/bedrock/models_tests.rs +++ b/app/src/ai/bedrock/models_tests.rs @@ -82,8 +82,8 @@ fn test_cross_region_prefix_unknown_region() { fn test_get_effective_models_empty_returns_defaults() { let models = get_effective_models(&[]); assert_eq!(models.len(), DEFAULT_BEDROCK_MODELS.len()); - assert_eq!(models[0].model_id, "anthropic.claude-opus-4-6[1m]"); - assert_eq!(models[0].display_name, "Claude Opus 4.6"); + assert_eq!(models[0].model_id, "us.anthropic.claude-opus-4-6-v1[1m]"); + assert_eq!(models[0].display_name, "Claude Opus 4.6 (1M)"); } #[test] @@ -92,6 +92,7 @@ fn test_get_effective_models_custom_overrides() { model_id: "custom.model-v1:0".to_string(), display_name: "Custom Model".to_string(), vision_supported: false, + use_rig: true, }]; let models = get_effective_models(&custom); 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"; 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, + )); +} diff --git a/app/src/ai/bedrock/request_translator.rs b/app/src/ai/bedrock/request_translator.rs index e5d56bcf..bd5fb26c 100644 --- a/app/src/ai/bedrock/request_translator.rs +++ b/app/src/ai/bedrock/request_translator.rs @@ -750,6 +750,7 @@ fn persist_input_images_on_latest_user_message( Some(api::input_context::Image { data, mime_type }) } Some(ContentPart::Text(_)) + | Some(ContentPart::Reasoning { .. }) | Some(ContentPart::ToolUse { .. }) | Some(ContentPart::ToolResult { .. }) | None => None, @@ -873,11 +874,17 @@ fn is_pure_tool_result(content: &MessageContent) -> bool { fn strip_tool_result_parts(content: &mut MessageContent) { if let MessageContent::MultiPart(parts) = content { 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); *content = match part { ContentPart::Text(t) => MessageContent::Text(t), ContentPart::Image { .. } => unreachable!(), + ContentPart::Reasoning { .. } => unreachable!(), ContentPart::ToolUse { tool_use_id, name, @@ -916,11 +923,17 @@ fn strip_orphaned_tool_result_parts( ContentPart::ToolResult { tool_use_id, .. } => valid_ids.contains(tool_use_id), _ => 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); *content = match part { ContentPart::Text(t) => MessageContent::Text(t), ContentPart::Image { .. } => unreachable!(), + ContentPart::Reasoning { .. } => unreachable!(), ContentPart::ToolUse { tool_use_id, name, diff --git a/app/src/ai/bedrock/translator.rs b/app/src/ai/bedrock/translator.rs index 53817b74..52c87801 100644 --- a/app/src/ai/bedrock/translator.rs +++ b/app/src/ai/bedrock/translator.rs @@ -180,6 +180,13 @@ fn describe_message_content(content: &crate::ai::bedrock::convert::MessageConten .iter() .map(|p| match p { 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 } => { format!("Image({mime_type},{}bytes)", data.len()) } diff --git a/app/src/ai/blocklist/controller.rs b/app/src/ai/blocklist/controller.rs index 425b66d7..c5083d38 100644 --- a/app/src/ai/blocklist/controller.rs +++ b/app/src/ai/blocklist/controller.rs @@ -4346,6 +4346,7 @@ impl BlocklistAIController { secret_access_key: settings.bedrock_secret_access_key.value().clone(), session_token: None, cross_region_inference: *settings.bedrock_cross_region_inference.value(), + use_rig: false, }; if let ::ai::api_keys::AwsCredentialsState::Loaded { credentials, .. } = @@ -4431,6 +4432,9 @@ impl BlocklistAIController { .iter() .map(|p| match p { ContentPart::Text(t) => t.clone(), + ContentPart::Reasoning { text, .. } => { + format!("[Reasoning] {text}") + } ContentPart::Image { .. } => "[Image attachment]".to_string(), ContentPart::ToolUse { name, input, .. } => { format!("[Tool: {}] {}", name, input) @@ -4540,6 +4544,9 @@ impl BlocklistAIController { .iter() .map(|p| match p { ContentPart::Text(t) => (t.len() / 4) as u32, + ContentPart::Reasoning { text, .. } => { + (text.len() / 4) as u32 + } ContentPart::Image { .. } => 1_600, ContentPart::ToolUse { input, .. } => { (input.to_string().len() / 4) as u32 diff --git a/app/src/ai/blocklist/controller/response_stream.rs b/app/src/ai/blocklist/controller/response_stream.rs index fade9c19..f0fab496 100644 --- a/app/src/ai/blocklist/controller/response_stream.rs +++ b/app/src/ai/blocklist/controller/response_stream.rs @@ -243,15 +243,33 @@ impl ResponseStream { // Fall back to Bedrock if *settings.bedrock_enabled.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(), + ®ion, + 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 mut config = BedrockClientConfig { auth_method, profile: settings.bedrock_profile.value().clone(), - region: settings.bedrock_region.value().clone(), + region, access_key_id: settings.bedrock_access_key_id.value().clone(), secret_access_key: settings.bedrock_secret_access_key.value().clone(), 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, .. } = diff --git a/app/src/ai/blocklist/usage/context_window_view.rs b/app/src/ai/blocklist/usage/context_window_view.rs index 0b876455..c28e042d 100644 --- a/app/src/ai/blocklist/usage/context_window_view.rs +++ b/app/src/ai/blocklist/usage/context_window_view.rs @@ -45,6 +45,7 @@ impl View for ContextWindowView { .iter() .map(|p| match p { ContentPart::Text(t) => t.len(), + ContentPart::Reasoning { text, .. } => text.len(), ContentPart::Image { .. } => 6_400, ContentPart::ToolUse { input, .. } => input.to_string().len(), ContentPart::ToolResult { content, .. } => content.len(), @@ -112,6 +113,14 @@ impl View for ContextWindowView { ContentPart::Text(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 } => { out.push_str(&format!( "[Part {} Image] mime_type={}, bytes={}\n", diff --git a/app/src/ai/crosscheck/reviewer.rs b/app/src/ai/crosscheck/reviewer.rs index 27845d69..07cfef74 100644 --- a/app/src/ai/crosscheck/reviewer.rs +++ b/app/src/ai/crosscheck/reviewer.rs @@ -186,6 +186,7 @@ impl CrosscheckReviewer { secret_access_key: settings.bedrock_secret_access_key.value().clone(), session_token: None, cross_region_inference: *settings.bedrock_cross_region_inference.value(), + use_rig: false, }; if let ::ai::api_keys::AwsCredentialsState::Loaded { credentials, .. } = diff --git a/app/src/ai/llms.rs b/app/src/ai/llms.rs index 950eff9e..f845b941 100644 --- a/app/src/ai/llms.rs +++ b/app/src/ai/llms.rs @@ -721,6 +721,7 @@ impl LLMPreferences { model_id: default.model_id.to_string(), display_name: default.display_name.to_string(), vision_supported: default.vision_supported, + use_rig: false, }); 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, + ) { + 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. pub fn get_active_base_model<'a>( &'a self, @@ -1994,6 +2069,52 @@ fn openai_model_context_size(model: &OpenAIModelConfig) -> u32 { 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, +) -> Vec { + 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"))] fn openai_model_context_window(model: &OpenAIModelConfig) -> LLMContextWindow { let context_size = openai_model_context_size(model); @@ -2118,7 +2239,11 @@ async fn fetch_from_litellm_model_info( max_output_tokens, provider, 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(); @@ -2241,7 +2366,11 @@ async fn fetch_from_openai_models( max_output_tokens, provider, 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(); diff --git a/app/src/ai/llms_tests.rs b/app/src/ai/llms_tests.rs index a3a11304..687d2e09 100644 --- a/app/src/ai/llms_tests.rs +++ b/app/src/ai/llms_tests.rs @@ -10,7 +10,7 @@ use crate::network::NetworkStatus; use crate::server::cloud_objects::update_manager::UpdateManager; use crate::server::server_api::ServerApiProvider; 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::workspaces::team_tester::TeamTesterStatus; use crate::workspaces::user_workspaces::UserWorkspaces; @@ -138,3 +138,67 @@ fn llm_info_round_trip_serializes_and_deserializes() { 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"); +} diff --git a/app/src/ai/openai/convert.rs b/app/src/ai/openai/convert.rs index 8521bf48..f82e464f 100644 --- a/app/src/ai/openai/convert.rs +++ b/app/src/ai/openai/convert.rs @@ -121,6 +121,9 @@ fn convert_user_message(content: MessageContent) -> ConvertedMessages { ContentPart::Text(text) => { user_content_parts.push(UserContentPart::Text(text)); } + ContentPart::Reasoning { text, .. } => { + user_content_parts.push(UserContentPart::Text(text)); + } ContentPart::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); } + ContentPart::Reasoning { text, .. } => { + if !text_content.is_empty() { + text_content.push('\n'); + } + text_content.push_str(&text); + } ContentPart::ToolUse { tool_use_id, name, diff --git a/app/src/ai/runtime/mod.rs b/app/src/ai/runtime/mod.rs index 8a6797b7..333c9c25 100644 --- a/app/src/ai/runtime/mod.rs +++ b/app/src/ai/runtime/mod.rs @@ -4,4 +4,4 @@ mod rig_request; mod rig_tool; pub(crate) use provider::ProviderRuntime; -pub(crate) use rig::rig_openai_response_stream; +pub(crate) use rig::{rig_bedrock_response_stream, rig_openai_response_stream}; diff --git a/app/src/ai/runtime/rig.rs b/app/src/ai/runtime/rig.rs index d48eae5e..dadba13d 100644 --- a/app/src/ai/runtime/rig.rs +++ b/app/src/ai/runtime/rig.rs @@ -11,10 +11,12 @@ use uuid::Uuid; use warp_multi_agent_api::response_event::stream_finished; use warp_multi_agent_api::{self as api, ClientAction, ResponseEvent, ToolType}; -use super::rig_request::{prepare_rig_turn, PreparedRigTurn}; +use super::rig_request::{prepare_bedrock_rig_turn, prepare_rig_turn, PreparedRigTurn}; use super::rig_tool::action_from_tool_call; use crate::ai::agent::api::{Event, RequestParams, ResponseStream, StreamEvent}; use crate::ai::agent::AIAgentAction; +use crate::ai::bedrock::client::{BedrockClient, BedrockClientConfig}; +use crate::ai::bedrock::external_config::ExternalBedrockConfig; use crate::ai::bedrock::response_translator::{ build_add_agent_output_message, build_append_text, build_create_task, build_stream_init, build_user_query_message, @@ -32,6 +34,75 @@ pub(crate) fn rig_openai_response_stream( cancellation_rx: oneshot::Receiver<()>, ) -> ResponseStream { let skill_path_origin = params.session_context.skill_path_origin(); + let prepared = prepare_rig_turn(&config, params, supported_tools, supported_cli_agent_tools); + let model_id = prepared.request.model.as_str().to_string(); + let runtime = OpenAICompatibleRuntime::new(OpenAICompatibleRuntimeConfig { + base_url: config.base_url, + api_key: config.api_key, + model: model_id.clone(), + max_output_tokens: config.max_output_tokens.map(u64::from), + supports_system_messages: config.supports_system_messages, + }); + rig_response_stream( + runtime, + prepared, + skill_path_origin, + config.max_input_tokens, + "rig_openai_compatible", + cancellation_rx, + ) +} + +pub(crate) async fn rig_bedrock_response_stream( + config: BedrockClientConfig, + params: RequestParams, + supported_tools: Vec, + supported_cli_agent_tools: Vec, + cancellation_rx: oneshot::Receiver<()>, +) -> anyhow::Result { + 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( + runtime: R, + prepared: PreparedRigTurn, + skill_path_origin: ai::skills::SkillPathOrigin, + max_context_tokens: Option, + stream_type: &'static str, + cancellation_rx: oneshot::Receiver<()>, +) -> ResponseStream +where + R: AgentRuntime + Send + Sync + 'static, +{ let PreparedRigTurn { task_id, needs_create_task, @@ -40,21 +111,12 @@ pub(crate) fn rig_openai_response_stream( persistent_messages, tool_result_archive, messages_sent, - } = prepare_rig_turn(&config, params, supported_tools, supported_cli_agent_tools); + } = prepared; store_messages_sent(&messages_sent, &persistent_messages); let conversation_id = turn_request.conversation_id.clone(); let model_id = turn_request.model.as_str().to_string(); let tool_policy = ToolPolicy::new(&turn_request.tools); - - let runtime = OpenAICompatibleRuntime::new(OpenAICompatibleRuntimeConfig { - base_url: config.base_url, - api_key: config.api_key, - model: model_id.clone(), - max_output_tokens: config.max_output_tokens.map(u64::from), - supports_system_messages: config.supports_system_messages, - }); - let max_context_tokens = config.max_input_tokens; let stream = async_stream::stream! { let (control_sender, control) = turn_control(); let start_future = runtime.start_turn(turn_request, control).fuse(); @@ -67,7 +129,7 @@ pub(crate) fn rig_openai_response_stream( match start_future.await { Ok(stream) => stream, Err(error) => { - yield Err(agent_error(error)); + yield Err(agent_error(error, stream_type)); return; } } @@ -75,7 +137,7 @@ pub(crate) fn rig_openai_response_stream( result = start_future => match result { Ok(stream) => stream, Err(error) => { - yield Err(agent_error(error)); + yield Err(agent_error(error, stream_type)); return; } }, @@ -87,6 +149,8 @@ pub(crate) fn rig_openai_response_stream( let mut current_text_message_id: Option = None; let mut current_reasoning_message_id: Option = None; let mut full_text = String::new(); + let mut full_reasoning = String::new(); + let mut reasoning_signature = None; let mut proposed_tools = Vec::new(); let mut assistant_history_index = None; let mut usage = Usage::default(); @@ -106,7 +170,7 @@ pub(crate) fn rig_openai_response_stream( let event = match event { Ok(event) => event, Err(error) => { - yield Err(agent_error(error)); + yield Err(agent_error(error, stream_type)); return; } }; @@ -133,6 +197,7 @@ pub(crate) fn rig_openai_response_stream( } } AgentEvent::ReasoningDelta { text } => { + full_reasoning.push_str(&text); if let Some(message_id) = ¤t_reasoning_message_id { yield Ok(StreamEvent::Response(build_append_reasoning(&task_id, message_id, &text))); } else { @@ -141,6 +206,17 @@ pub(crate) fn rig_openai_response_stream( current_reasoning_message_id = Some(message_id); } } + AgentEvent::ReasoningCompleted { text, signature } => { + if current_reasoning_message_id.is_none() && !text.is_empty() { + let message_id = Uuid::new_v4().to_string(); + yield Ok(StreamEvent::Response(build_add_reasoning(&task_id, &message_id, &text))); + current_reasoning_message_id = Some(message_id); + } + if !text.is_empty() { + full_reasoning = text; + } + reasoning_signature = signature; + } AgentEvent::UsageUpdated { usage: updated } => usage = updated, AgentEvent::Tool { event: ToolEvent::Proposed { call }, @@ -148,6 +224,8 @@ pub(crate) fn rig_openai_response_stream( proposed_tools.push(call.clone()); sync_assistant_turn( &messages_sent, + &full_reasoning, + reasoning_signature.as_deref(), &full_text, &proposed_tools, &mut assistant_history_index, @@ -164,7 +242,7 @@ pub(crate) fn rig_openai_response_stream( yield Err(agent_error(AgentError::new( galaxy_agent_core::AgentErrorKind::Protocol, message, - ))); + ), stream_type)); return; } } @@ -198,6 +276,8 @@ pub(crate) fn rig_openai_response_stream( } sync_assistant_turn( &messages_sent, + &full_reasoning, + reasoning_signature.as_deref(), &full_text, &proposed_tools, &mut assistant_history_index, @@ -222,7 +302,7 @@ pub(crate) fn rig_openai_response_stream( yield Err(agent_error(AgentError::new( galaxy_agent_core::AgentErrorKind::Protocol, "the provider runtime attempted to execute a tool outside Galaxy's permission boundary", - ))); + ), stream_type)); return; } } @@ -264,11 +344,22 @@ fn append_tool_result( fn sync_assistant_turn( messages_sent: &std::sync::Arc>>, + reasoning_text: &str, + reasoning_signature: Option<&str>, text: &str, tool_calls: &[ToolCall], history_index: &mut Option, ) { - let mut parts = Vec::with_capacity(usize::from(!text.is_empty()) + tool_calls.len()); + let has_reasoning = !reasoning_text.is_empty() || reasoning_signature.is_some(); + let mut parts = Vec::with_capacity( + usize::from(has_reasoning) + usize::from(!text.is_empty()) + tool_calls.len(), + ); + if has_reasoning { + parts.push(ContentPart::Reasoning { + text: reasoning_text.to_string(), + signature: reasoning_signature.map(str::to_string), + }); + } if !text.is_empty() { parts.push(ContentPart::Text(text.to_string())); } @@ -293,6 +384,7 @@ fn sync_assistant_turn( name, input, }, + reasoning @ ContentPart::Reasoning { .. } => MessageContent::MultiPart(vec![reasoning]), ContentPart::Image { .. } | ContentPart::ToolResult { .. } => unreachable!(), } } else { @@ -395,10 +487,10 @@ fn saturating_i32(value: u64) -> i32 { i32::try_from(value).unwrap_or(i32::MAX) } -fn agent_error(error: AgentError) -> Arc { +fn agent_error(error: AgentError, stream_type: &'static str) -> Arc { Arc::new( AIApiError::Stream { - stream_type: "rig_openai_compatible", + stream_type, source: anyhow::anyhow!(error), } .into_quota_limit_if_provider_budget_exhausted(), diff --git a/app/src/ai/runtime/rig_request.rs b/app/src/ai/runtime/rig_request.rs index b0b82fe5..b5fa0922 100644 --- a/app/src/ai/runtime/rig_request.rs +++ b/app/src/ai/runtime/rig_request.rs @@ -13,7 +13,9 @@ use warp_multi_agent_api::ToolType; use crate::ai::agent::api::RequestParams; use crate::ai::agent::{AIAgentContext, AIAgentInput, MCPContext, UserQueryMode}; -use crate::ai::bedrock::request_translator::{default_tool_definitions, tool_name_is_supported}; +use crate::ai::bedrock::request_translator::{ + default_tool_definitions, sanitize_messages_for_bedrock, tool_name_is_supported, +}; use crate::ai::openai::client::OpenAIClientConfig; use crate::ai::openai::request_translator::sanitize_messages_for_openai; @@ -32,6 +34,47 @@ pub(crate) fn prepare_rig_turn( params: RequestParams, supported_tools: Vec, supported_cli_agent_tools: Vec, +) -> 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, + params: RequestParams, + supported_tools: Vec, + supported_cli_agent_tools: Vec, +) -> 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, + max_output_tokens: Option, + sanitizer: RigRequestSanitizer, + params: RequestParams, + supported_tools: Vec, + supported_cli_agent_tools: Vec, ) -> PreparedRigTurn { let RequestParams { input, @@ -70,7 +113,12 @@ pub(crate) fn prepare_rig_turn( for message in &mut persistent_messages { message.truncate_tool_results_for_provider_request(); } - sanitize_messages_for_openai(&mut persistent_messages); + match sanitizer { + RigRequestSanitizer::OpenAICompatible => { + sanitize_messages_for_openai(&mut persistent_messages) + } + RigRequestSanitizer::Bedrock => sanitize_messages_for_bedrock(&mut persistent_messages), + } let mut turn_messages = Vec::new(); if let Some(summary) = progressive_summary { @@ -91,16 +139,14 @@ pub(crate) fn prepare_rig_turn( } turn_messages.extend(persistent_messages.clone()); - let model_id = config - .model - .clone() + let model_id = model_override .filter(|model| !model.is_empty() && model != "auto") .unwrap_or_else(|| model.as_str().to_string()); let mut request = TurnRequest::new(model_id, turn_messages); request.conversation_id = conversation_token.map(|token| token.as_str().to_string()); request.system_prompt = Some(system_prompt); request.tools = tools; - request.max_output_tokens = config.max_output_tokens.map(u64::from); + request.max_output_tokens = max_output_tokens; PreparedRigTurn { task_id, diff --git a/app/src/ai/runtime/rig_request_tests.rs b/app/src/ai/runtime/rig_request_tests.rs index af228c83..a96d9857 100644 --- a/app/src/ai/runtime/rig_request_tests.rs +++ b/app/src/ai/runtime/rig_request_tests.rs @@ -4,7 +4,7 @@ use std::sync::Arc; use galaxy_agent_core::{ContentPart, MessageContent, MessageRole, ToolResult, ToolResultStatus}; use warp_multi_agent_api::ToolType; -use super::{input_messages, prepare_rig_turn, tool_definitions}; +use super::{input_messages, prepare_bedrock_rig_turn, prepare_rig_turn, tool_definitions}; use crate::ai::agent::api::RequestParams; use crate::ai::agent::{ AIAgentContext, AIAgentInput, AnyFileContent, FileContext, MCPContext, MCPServer, UserQueryMode, @@ -114,6 +114,40 @@ fn builds_a_rig_turn_directly_from_galaxy_request_state() { )); } +#[test] +fn bedrock_rig_turn_uses_bedrock_history_invariants_without_a_proto_round_trip() { + let mut params = RequestParams::new_for_test(); + params.message_history = vec![galaxy_agent_core::ConversationMessage { + role: MessageRole::Assistant, + content: MessageContent::Text("Prior assistant message".to_string()), + }]; + params.input = vec![user_query("Continue safely")]; + + let prepared = prepare_bedrock_rig_turn( + "anthropic.claude-test".to_string(), + Some(64_000), + params, + Vec::new(), + Vec::new(), + ); + + assert_eq!(prepared.request.model.as_str(), "anthropic.claude-test"); + assert_eq!(prepared.request.max_output_tokens, Some(64_000)); + assert_eq!( + prepared + .request + .messages + .first() + .map(|message| message.role), + Some(MessageRole::User) + ); + assert_eq!( + prepared.request.messages.last().map(|message| message.role), + Some(MessageRole::User) + ); + assert_eq!(prepared.request.messages, prepared.persistent_messages); +} + #[test] #[allow(deprecated)] fn grouped_mcp_tool_names_use_the_installation_id_not_the_display_name() { diff --git a/app/src/ai/runtime/rig_tests.rs b/app/src/ai/runtime/rig_tests.rs index 22f0278f..b18f7089 100644 --- a/app/src/ai/runtime/rig_tests.rs +++ b/app/src/ai/runtime/rig_tests.rs @@ -2,7 +2,7 @@ use std::sync::{Arc, Mutex}; use ai::skills::SkillPathOrigin; use galaxy_agent_core::{ - MessageContent, MessageRole, StopReason, ToolCall, ToolResult, ToolResultStatus, + ContentPart, MessageContent, MessageRole, StopReason, ToolCall, ToolResult, ToolResultStatus, }; use warp_multi_agent_api::response_event::stream_finished; @@ -134,12 +134,16 @@ fn assistant_history_is_updated_before_fast_tool_execution_can_continue() { sync_assistant_turn( &messages, + "", + None, "I'll inspect both.", std::slice::from_ref(&first_call), &mut history_index, ); sync_assistant_turn( &messages, + "", + None, "I'll inspect both.", &[first_call, second_call], &mut history_index, @@ -162,6 +166,43 @@ fn assistant_history_is_updated_before_fast_tool_execution_can_continue() { ); } +#[test] +fn signed_reasoning_is_persisted_before_the_tool_call() { + let messages = Arc::new(Mutex::new(Vec::new())); + let mut history_index = None; + let call = ToolCall { + id: "call-1".to_string(), + name: "read_files".to_string(), + arguments: serde_json::json!({"files": ["Cargo.toml"]}), + }; + + sync_assistant_turn( + &messages, + "I should inspect the manifest.", + Some("signed-reasoning"), + "", + std::slice::from_ref(&call), + &mut history_index, + ); + + let messages = messages.lock().unwrap(); + let MessageContent::MultiPart(parts) = &messages[0].content else { + panic!("expected reasoning and tool call parts"); + }; + assert!(matches!( + parts.as_slice(), + [ + ContentPart::Reasoning { + text, + signature: Some(signature), + }, + ContentPart::ToolUse { tool_use_id, .. }, + ] if text == "I should inspect the manifest." + && signature == "signed-reasoning" + && tool_use_id == "call-1" + )); +} + #[test] fn rejected_tool_result_is_paired_with_the_assistant_call_in_history() { let messages = Arc::new(Mutex::new(Vec::new())); @@ -174,6 +215,8 @@ fn rejected_tool_result_is_paired_with_the_assistant_call_in_history() { sync_assistant_turn( &messages, "", + None, + "", std::slice::from_ref(&call), &mut history_index, ); diff --git a/app/src/settings/ai.rs b/app/src/settings/ai.rs index 127664d4..4215e4ee 100644 --- a/app/src/settings/ai.rs +++ b/app/src/settings/ai.rs @@ -828,6 +828,11 @@ pub struct BedrockModelConfig { #[serde(default)] #[schemars(description = "Whether the model supports image/vision input.")] 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 {} @@ -890,8 +895,10 @@ impl settings_value::SettingsValue for OpenAIModelConfig {} impl OpenAIModelConfig { pub fn supports_system_messages(&self) -> bool { - self.supports_system_messages - .unwrap_or_else(|| !self.model_id.starts_with("codex-gpt-")) + if self.model_id.starts_with("codex-gpt-") { + return false; + } + self.supports_system_messages.unwrap_or(true) } } diff --git a/app/src/settings_view/ai_page.rs b/app/src/settings_view/ai_page.rs index df9529c4..e82ab63d 100644 --- a/app/src/settings_view/ai_page.rs +++ b/app/src/settings_view/ai_page.rs @@ -88,8 +88,8 @@ use crate::settings::{ GitOperationsAutogenEnabled, IncludeAgentCommandsInHistory, InputSettings, IntelligentAutosuggestionsEnabled, LongRunningCommandSubmissionMode, MemoryEnabled, NLDInTerminalEnabled, NaturalLanguageAutosuggestionsEnabled, OpenAIEnabled, - OrchestrationMessageDisplayMode, PromptSubmissionMode, RuleSuggestionsEnabled, - SharedBlockTitleGenerationEnabled, ShouldRenderCLIAgentToolbar, + OpenAIProviderConfig, OrchestrationMessageDisplayMode, PromptSubmissionMode, + RuleSuggestionsEnabled, SharedBlockTitleGenerationEnabled, ShouldRenderCLIAgentToolbar, ShouldRenderUseAgentToolbarForUserCommands, ShowAgentTips, ShowConversationHistory, ShowHintText, ThinkingDisplayMode, VoiceInputEnabled, WarpDriveContextEnabled, }; @@ -117,10 +117,8 @@ pub enum AISubpage { Knowledge, /// Third-party CLI agent settings. ThirdPartyCLIAgents, - /// AWS Bedrock direct provider configuration. - Bedrock, - /// OpenAI-compatible (LiteLLM) provider configuration. - OpenAI, + /// Unified model and provider configuration. + Models, /// Experimental features. Experiments, } @@ -132,8 +130,7 @@ impl AISubpage { SettingsSection::AgentProfiles => Some(Self::Profiles), SettingsSection::Knowledge => Some(Self::Knowledge), SettingsSection::ThirdPartyCLIAgents => Some(Self::ThirdPartyCLIAgents), - SettingsSection::Bedrock => Some(Self::Bedrock), - SettingsSection::OpenAI => Some(Self::OpenAI), + SettingsSection::Models => Some(Self::Models), SettingsSection::Experiments => Some(Self::Experiments), // AgentMCPServers renders the standalone MCPServers page, not an AI subpage. _ => None, @@ -1908,15 +1905,18 @@ impl AISettingsPageView { } } - /// Fetches models from the LiteLLM endpoint and stores them in memory via LLMPreferences. - fn fetch_litellm_models(&mut self, ctx: &mut ViewContext) { - use crate::ai::llms::LLMPreferences; - + fn fetch_openai_provider_models(&mut self, provider_index: usize, ctx: &mut ViewContext) { 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) { + let (page, _) = Self::build_page(self.active_subpage, ctx); + self.page = page; + ctx.notify(); + } + fn build_page( subpage: Option, ctx: &mut ViewContext, @@ -2034,15 +2034,10 @@ impl AISettingsPageView { Some(AISubpage::ThirdPartyCLIAgents) => { widgets.push(Box::new(CLIAgentWidget::default())); } - Some(AISubpage::Bedrock) => { - let widget = BedrockSettingsWidget::new(ctx); - widgets.push(Box::new(widget)); - let title: Option<&str> = None; - return (PageType::new_uncategorized(widgets, title), None); - } - Some(AISubpage::OpenAI) => { - let widget = OpenAISettingsWidget::new(ctx); - widgets.push(Box::new(widget)); + Some(AISubpage::Models) => { + widgets.push(Box::new(ModelsOverviewWidget)); + widgets.push(Box::new(OpenAISettingsWidget::new(ctx))); + widgets.push(Box::new(BedrockSettingsWidget::new(ctx))); let title: Option<&str> = None; return (PageType::new_uncategorized(widgets, title), None); } @@ -2807,10 +2802,13 @@ pub enum AISettingsPageAction { SetBedrockAuthMethod(BedrockAuthMethod), SetBedrockProfile(String), ToggleBedrockCrossRegionInference, + ToggleBedrockModelRig(usize), ToggleOpenAIEnabled, ToggleAcpEnabled, RefreshAcpDiscovery, - FetchOpenAIModels, + FetchOpenAIProviderModels(usize), + AddOpenAIProvider, + RemoveOpenAIProvider(usize), ToggleFileBasedMcp, ToggleIncludeAgentCommandsInHistory, ToggleAgentAttribution, @@ -3562,6 +3560,17 @@ impl TypedActionView for AISettingsPageView { }); 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 => { AISettings::handle(ctx).update(ctx, |settings, 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"))] self.refresh_acp_discovery(ctx); } - AISettingsPageAction::FetchOpenAIModels => { - // Trigger a fetch of models from the LiteLLM endpoint - self.fetch_litellm_models(ctx); + AISettingsPageAction::FetchOpenAIProviderModels(provider_index) => { + self.fetch_openai_provider_models(*provider_index, 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 => { 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 { + 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::(); + 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 { enabled_toggle: SwitchStateHandle, auto_login_toggle: SwitchStateHandle, @@ -7239,6 +7315,7 @@ struct BedrockSettingsWidget { auth_refresh_command_editor: ViewHandle, access_key_editor: ViewHandle, secret_key_editor: ViewHandle, + model_rig_toggles: RefCell>, } impl BedrockSettingsWidget { @@ -7250,6 +7327,7 @@ impl BedrockSettingsWidget { 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 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 mut dropdown = Dropdown::new(ctx); @@ -7475,6 +7553,11 @@ impl BedrockSettingsWidget { auth_refresh_command_editor, access_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.); + 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_")); 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(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::( + format!("{} — Rig", model.display_name), + Some(styles::header_font_color(is_enabled, app)), + None, + LocalOnlyIconState::Hidden, + ToggleState::Enabled, + appearance, + ), + toggle, + appearance, + None, + )); + } } else { column.add_child(render_ai_setting_description( "No models configured. Add models to ~/.galaxy/settings.toml under [ai.bedrock].", @@ -8003,104 +8129,149 @@ impl SettingsWidget for ACPSettingsWidget { } } -struct OpenAISettingsWidget { - enabled_toggle: SwitchStateHandle, +struct OpenAIProviderEditor { + name_editor: ViewHandle, base_url_editor: ViewHandle, api_key_editor: ViewHandle, fetch_button: MouseStateHandle, + remove_button: MouseStateHandle, +} + +struct OpenAISettingsWidget { + enabled_toggle: SwitchStateHandle, + provider_editors: Vec, + add_provider_button: MouseStateHandle, } impl OpenAISettingsWidget { + fn create_editor( + value: String, + placeholder: &'static str, + is_password: bool, + ctx: &mut ViewContext<::View>, + ) -> ViewHandle { + 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<::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(); - let api_key_val = ai_settings.openai_api_key.value().clone(); + for (provider_index, provider) in providers.into_iter().enumerate() { + 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 appearance = Appearance::as_ref(ctx); - let options = SingleLineEditorOptions { - is_password: false, - 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("http://localhost:4000/v1", ctx); - editor.set_buffer_text(&base_url_val, ctx); - editor - }); - ctx.subscribe_to_view(&base_url_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_base_url.set_value(value, ctx); - }); + let base_url_editor = + Self::create_editor(provider.base_url, "http://localhost:4000/v1", false, ctx); + ctx.subscribe_to_view(&base_url_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.base_url = value; + report_if_error!(settings.openai_providers.set_value(providers, ctx)); + } + }); + } + }); + + let api_key_editor = Self::create_editor( + provider.api_key.unwrap_or_default(), + "sk-... (optional)", + true, + ctx, + ); + ctx.subscribe_to_view(&api_key_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.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| { - let appearance = Appearance::as_ref(ctx); - let options = SingleLineEditorOptions { - is_password: true, - 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("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); - }); - } - }); + provider_editors.push(OpenAIProviderEditor { + name_editor, + base_url_editor, + api_key_editor, + fetch_button: MouseStateHandle::default(), + remove_button: MouseStateHandle::default(), + }); + } - let base_url_editor_clone = base_url_editor.clone(); - let api_key_editor_clone = api_key_editor.clone(); + let editor_handles = provider_editors + .iter() + .flat_map(|provider| { + [ + provider.name_editor.clone(), + provider.base_url_editor.clone(), + provider.api_key_editor.clone(), + ] + }) + .collect::>(); ctx.subscribe_to_model(&AISettings::handle(ctx), move |_, _, event, ctx| { if matches!(event, AISettingsChangedEvent::OpenAIEnabled { .. }) { let is_enabled = *AISettings::as_ref(ctx).openai_enabled.value(); - AISettingsPageView::update_editor_interaction_state( - base_url_editor_clone.clone(), - is_enabled, - ctx, - ); - AISettingsPageView::update_editor_interaction_state( - api_key_editor_clone.clone(), - is_enabled, - ctx, - ); + for editor in &editor_handles { + AISettingsPageView::update_editor_interaction_state( + editor.clone(), + is_enabled, + ctx, + ); + } ctx.notify(); } }); Self { enabled_toggle: SwitchStateHandle::default(), - base_url_editor, - api_key_editor, - fetch_button: MouseStateHandle::default(), + provider_editors, + add_provider_button: MouseStateHandle::default(), } } @@ -8164,8 +8335,11 @@ impl SettingsWidget for OpenAISettingsWidget { 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::( - "Enable OpenAI-Compatible Provider", + "Enable model providers", AISettingsPageAction::ToggleOpenAIEnabled, is_enabled, true, @@ -8174,65 +8348,153 @@ impl SettingsWidget for OpenAISettingsWidget { app, )); 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, 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( - appearance, - "Base URL", - 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, - )); + for (provider_index, provider) in ai_settings.openai_providers.value().iter().enumerate() { + let Some(editors) = self.provider_editors.get(provider_index) else { + continue; + }; - column.add_child(Self::render_input( - appearance, - "API Key", - self.api_key_editor.clone(), - is_enabled, - app, - )); - column.add_child(render_ai_setting_description( - "Optional. Leave empty if the proxy handles authentication.", - is_enabled, - app, - )); + column.add_child(render_separator(appearance)); + column.add_child( + build_sub_header( + appearance, + format!("Provider {}: {}", provider_index + 1, provider.name), + None, + ) + .finish(), + ); + column.add_child(Self::render_input( + 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)); - - // Fetch models button - let fetch_button = appearance + let add_provider_button = appearance .ui_builder() - .button(ButtonVariant::Secondary, self.fetch_button.clone()) - .with_text_label("Fetch Models from Endpoint".to_owned()) + .button(ButtonVariant::Secondary, self.add_provider_button.clone()) + .with_text_label("Add Provider".to_owned()) .build() .on_click(move |ctx, _, _| { - ctx.dispatch_typed_action(AISettingsPageAction::FetchOpenAIModels); + ctx.dispatch_typed_action(AISettingsPageAction::AddOpenAIProvider); }) .finish(); - column.add_child(fetch_button); + column.add_child(add_provider_button); 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, app, )); column.add_child(render_separator(appearance)); - // Show configured models count - let configured_models: Vec<_> = ai_settings.openai_models.value().clone(); + let mut configured_models = ai_settings + .openai_providers + .value() + .iter() + .flat_map(|provider| provider.models.iter()) + .collect::>(); + configured_models.extend(ai_settings.openai_models.value().iter()); if !configured_models.is_empty() { let description = format!( - "{} model{} configured via settings.toml.", + "{} model{} configured across all OpenAI-compatible providers.", configured_models.len(), if configured_models.len() == 1 { "" @@ -8246,7 +8508,7 @@ impl SettingsWidget for OpenAISettingsWidget { let preview: String = configured_models .iter() .take(5) - .map(|m| m.display_name.as_str()) + .map(|model| model.display_name.as_str()) .collect::>() .join(", "); let suffix = if configured_models.len() > 5 { @@ -8261,7 +8523,7 @@ impl SettingsWidget for OpenAISettingsWidget { )); } else { 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, app, )); diff --git a/app/src/settings_view/mod.rs b/app/src/settings_view/mod.rs index 3da8f059..9254a2af 100644 --- a/app/src/settings_view/mod.rs +++ b/app/src/settings_view/mod.rs @@ -245,8 +245,7 @@ pub enum SettingsSection { AgentMCPServers, Knowledge, ThirdPartyCLIAgents, - Bedrock, - OpenAI, + Models, Experiments, /// Internal backing-page identifier for CodeSettingsPageView. Multiple subpages /// (CodeIndexing, EditorAndCodeReview) share this single backing page, @@ -274,8 +273,7 @@ impl Display for SettingsSection { SettingsSection::AgentMCPServers => write!(f, "MCP servers"), SettingsSection::Knowledge => write!(f, "Knowledge"), SettingsSection::ThirdPartyCLIAgents => write!(f, "Third party CLI agents"), - SettingsSection::Bedrock => write!(f, "AWS Bedrock"), - SettingsSection::OpenAI => write!(f, "OpenAI / LiteLLM"), + SettingsSection::Models => write!(f, "Models"), SettingsSection::Experiments => write!(f, "Experiments"), SettingsSection::Warpify => write!(f, "Wormhole"), SettingsSection::CodeIndexing => write!(f, "Indexing and projects"), @@ -300,8 +298,7 @@ impl SettingsSection { | Self::AgentMCPServers | Self::Knowledge | Self::ThirdPartyCLIAgents - | Self::Bedrock - | Self::OpenAI + | Self::Models | Self::Experiments ) } @@ -329,12 +326,11 @@ impl SettingsSection { pub fn ai_subpages() -> &'static [Self] { &[ Self::WarpAgent, + Self::Models, Self::AgentProfiles, Self::AgentMCPServers, Self::Knowledge, Self::ThirdPartyCLIAgents, - Self::Bedrock, - Self::OpenAI, Self::Experiments, ] } @@ -367,8 +363,9 @@ impl FromStr for SettingsSection { "MCP servers" | "AgentMCPServers" => Ok(Self::AgentMCPServers), "Knowledge" => Ok(Self::Knowledge), "Third party CLI agents" | "ThirdPartyCLIAgents" => Ok(Self::ThirdPartyCLIAgents), - "AWS Bedrock" | "Bedrock" => Ok(Self::Bedrock), - "OpenAI / LiteLLM" | "OpenAI" => Ok(Self::OpenAI), + "Models" | "AWS Bedrock" | "Bedrock" | "OpenAI / LiteLLM" | "OpenAI" => { + Ok(Self::Models) + } "Indexing and projects" | "CodeIndexing" => Ok(Self::CodeIndexing), "Editor and Code Review" | "EditorAndCodeReview" => Ok(Self::EditorAndCodeReview), "Experiments" => Ok(Self::Experiments), diff --git a/app/src/settings_view/mod_tests.rs b/app/src/settings_view/mod_tests.rs index 53dab041..4e711d0e 100644 --- a/app/src/settings_view/mod_tests.rs +++ b/app/src/settings_view/mod_tests.rs @@ -86,8 +86,7 @@ fn current_settings_display_names_round_trip() { SettingsSection::ThirdPartyCLIAgents, "Third party CLI agents", ), - (SettingsSection::Bedrock, "AWS Bedrock"), - (SettingsSection::OpenAI, "OpenAI / LiteLLM"), + (SettingsSection::Models, "Models"), (SettingsSection::Experiments, "Experiments"), (SettingsSection::CodeIndexing, "Indexing and projects"), ( @@ -111,8 +110,10 @@ fn legacy_settings_names_remain_parseable() { ("AgentProfiles", SettingsSection::AgentProfiles), ("AgentMCPServers", SettingsSection::AgentMCPServers), ("ThirdPartyCLIAgents", SettingsSection::ThirdPartyCLIAgents), - ("Bedrock", SettingsSection::Bedrock), - ("OpenAI", SettingsSection::OpenAI), + ("AWS Bedrock", SettingsSection::Models), + ("Bedrock", SettingsSection::Models), + ("OpenAI / LiteLLM", SettingsSection::Models), + ("OpenAI", SettingsSection::Models), ("CodeIndexing", SettingsSection::CodeIndexing), ("EditorAndCodeReview", SettingsSection::EditorAndCodeReview), ] { @@ -215,7 +216,7 @@ fn collapsed_umbrella_uses_first_and_last_visible_subpages() { let stops = build_nav_stops(&nav_items, |section| { !matches!( 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 { nav_index: 0, first_subpage: SettingsSection::AgentProfiles, - last_subpage: SettingsSection::Bedrock, + last_subpage: SettingsSection::ThirdPartyCLIAgents, } ); } diff --git a/crates/galaxy_agent_core/src/tool_policy.rs b/crates/galaxy_agent_core/src/tool_policy.rs index f510370c..678a01b0 100644 --- a/crates/galaxy_agent_core/src/tool_policy.rs +++ b/crates/galaxy_agent_core/src/tool_policy.rs @@ -250,7 +250,9 @@ fn collect_tool_entries(messages: &[ConversationMessage], entries: &mut Vec pair_result(tool_use_id, content, &mut pending, entries), - ContentPart::Text(_) | ContentPart::Image { .. } => {} + ContentPart::Text(_) + | ContentPart::Reasoning { .. } + | ContentPart::Image { .. } => {} } } } diff --git a/crates/galaxy_agent_core/src/types.rs b/crates/galaxy_agent_core/src/types.rs index f802e8ef..6531466e 100644 --- a/crates/galaxy_agent_core/src/types.rs +++ b/crates/galaxy_agent_core/src/types.rs @@ -42,6 +42,10 @@ pub enum MessageContent { #[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] pub enum ContentPart { Text(String), + Reasoning { + text: String, + signature: Option, + }, Image { data: Vec, mime_type: String, @@ -231,12 +235,28 @@ pub enum StopReason { #[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] pub enum AgentEvent { - TurnStarted { runtime_request_id: String }, - TextDelta { text: String }, - ReasoningDelta { text: String }, - Tool { event: ToolEvent }, - UsageUpdated { usage: Usage }, - TurnStopped { reason: StopReason }, + TurnStarted { + runtime_request_id: String, + }, + TextDelta { + text: String, + }, + ReasoningDelta { + text: String, + }, + ReasoningCompleted { + text: String, + signature: Option, + }, + Tool { + event: ToolEvent, + }, + UsageUpdated { + usage: Usage, + }, + TurnStopped { + reason: StopReason, + }, } 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::MultiPart(parts) => { for part in parts { - if let ContentPart::ToolResult { content, .. } = part { - truncate_tool_result_text(content); + match part { + ContentPart::ToolResult { content, .. } => { + truncate_tool_result_text(content); + } + ContentPart::Text(_) + | ContentPart::Reasoning { .. } + | ContentPart::Image { .. } + | ContentPart::ToolUse { .. } => {} } } } diff --git a/crates/galaxy_agent_rig/Cargo.toml b/crates/galaxy_agent_rig/Cargo.toml index 90bd9358..f7687ed8 100644 --- a/crates/galaxy_agent_rig/Cargo.toml +++ b/crates/galaxy_agent_rig/Cargo.toml @@ -8,13 +8,17 @@ license.workspace = true [dependencies] async-stream.workspace = true async-trait.workspace = true +aws-sdk-bedrockruntime.workspace = true +base64.workspace = true futures.workspace = true galaxy_agent_core.workspace = true rig-core.workspace = true +rig-bedrock.workspace = true serde_json.workspace = true uuid.workspace = true [dev-dependencies] +aws-smithy-http-client.workspace = true bytes.workspace = true rig-core = { workspace = true, features = ["test-utils"] } tokio = { workspace = true, features = ["macros", "rt"] } diff --git a/crates/galaxy_agent_rig/src/bedrock.rs b/crates/galaxy_agent_rig/src/bedrock.rs new file mode 100644 index 00000000..0ecb59df --- /dev/null +++ b/crates/galaxy_agent_rig/src/bedrock.rs @@ -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, +} + +/// 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 { + 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 { + 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, +) -> Result { + 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 { + 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; diff --git a/crates/galaxy_agent_rig/src/bedrock_tests.rs b/crates/galaxy_agent_rig/src/bedrock_tests.rs new file mode 100644 index 00000000..ac5428a7 --- /dev/null +++ b/crates/galaxy_agent_rig/src/bedrock_tests.rs @@ -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::>() + .await + .into_iter() + .collect::, _>>() + .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::>(); + 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::>(); + 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::>(); + 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" + )); +} diff --git a/crates/galaxy_agent_rig/src/lib.rs b/crates/galaxy_agent_rig/src/lib.rs index b00371d5..9fb4a485 100644 --- a/crates/galaxy_agent_rig/src/lib.rs +++ b/crates/galaxy_agent_rig/src/lib.rs @@ -1,5 +1,9 @@ //! Rig-backed implementations of Galaxy's provider-neutral agent runtime. +mod bedrock; mod openai_compatible; +mod request; +mod stream; +pub use bedrock::*; pub use openai_compatible::*; diff --git a/crates/galaxy_agent_rig/src/openai_compatible.rs b/crates/galaxy_agent_rig/src/openai_compatible.rs index 05f83d66..963f6d63 100644 --- a/crates/galaxy_agent_rig/src/openai_compatible.rs +++ b/crates/galaxy_agent_rig/src/openai_compatible.rs @@ -1,22 +1,14 @@ use async_trait::async_trait; -use futures::{FutureExt, StreamExt}; use galaxy_agent_core::{ - AgentError, AgentErrorKind, AgentEvent, AgentEventStream, AgentRuntime, ContentPart, - ConversationMessage, MessageContent, MessageRole, RuntimeCapabilities, RuntimeDescriptor, - RuntimeKind, StopReason, ToolCall, TurnCommand, TurnControl, TurnRequest, Usage, + AgentError, AgentErrorKind, AgentEventStream, AgentRuntime, RuntimeCapabilities, + RuntimeDescriptor, RuntimeKind, TurnControl, TurnRequest, }; -use rig_core::OneOrMany; use rig_core::client::CompletionClient; -use rig_core::completion::{ - AssistantContent, CompletionError, CompletionModel, CompletionRequest, GetTokenUsage, Message, - ToolDefinition, -}; -use rig_core::message::{ - DocumentSourceKind, Image, ImageMediaType, MimeType, ToolResultContent, UserContent, -}; +use rig_core::completion::{CompletionModel, CompletionRequest}; 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)] pub struct OpenAICompatibleRuntimeConfig { @@ -91,143 +83,13 @@ where M: CompletionModel + 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 completion_request = build_completion_request( request, configured_max_output_tokens, supports_system_messages, )?; - 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 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, - }), - ])) + start_provider_model_turn(model, completion_request, control, max_output_tokens).await } fn build_completion_request( @@ -235,208 +97,17 @@ fn build_completion_request( configured_max_output_tokens: Option, supports_system_messages: bool, ) -> Result { - 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)?); - } - - 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!({ + build_provider_completion_request( + request, + configured_max_output_tokens, + supports_system_messages, + false, + Some(serde_json::json!({ "stream_options": { "include_usage": true } })), - output_schema: None, - }) -} - -fn convert_message(message: ConversationMessage) -> Result { - 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, 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::, _>>()?, - MessageContent::ToolUse { .. } => { - return Err(invalid_role("tool use", "user")); - } - }; - one_or_many(parts, "user") -} - -fn assistant_content(content: MessageContent) -> Result, 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::, _>>()?, - MessageContent::ToolResult { .. } => { - return Err(invalid_role("tool result", "assistant")); - } - }; - one_or_many(parts, "assistant") -} - -fn convert_user_part(part: ContentPart) -> Result { - 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 { - 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(parts: Vec, role: &str) -> Result, 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)] #[path = "openai_compatible_tests.rs"] mod tests; diff --git a/crates/galaxy_agent_rig/src/openai_compatible_tests.rs b/crates/galaxy_agent_rig/src/openai_compatible_tests.rs index c45300a9..f94fde67 100644 --- a/crates/galaxy_agent_rig/src/openai_compatible_tests.rs +++ b/crates/galaxy_agent_rig/src/openai_compatible_tests.rs @@ -1,8 +1,11 @@ use futures::StreamExt; 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::completion::{AssistantContent, Message}; +use rig_core::message::{ToolResultContent, UserContent}; use rig_core::providers::openai; use rig_core::test_utils::MockStreamingClient; diff --git a/crates/galaxy_agent_rig/src/request.rs b/crates/galaxy_agent_rig/src/request.rs new file mode 100644 index 00000000..5b85c005 --- /dev/null +++ b/crates/galaxy_agent_rig/src/request.rs @@ -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, + supports_system_messages: bool, + encode_images_as_base64: bool, + additional_params: Option, +) -> Result { + 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 { + 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, 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::, _>>()?, + 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, 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::, _>>()?, + 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 { + 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 { + 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(parts: Vec, role: &str) -> Result, 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"), + ) +} diff --git a/crates/galaxy_agent_rig/src/stream.rs b/crates/galaxy_agent_rig/src/stream.rs new file mode 100644 index 00000000..3c6eb6b7 --- /dev/null +++ b/crates/galaxy_agent_rig/src/stream.rs @@ -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( + model: M, + completion_request: CompletionRequest, + control: TurnControl, + max_output_tokens: Option, +) -> Result +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 { + 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 +} diff --git a/plans/galaxy-local-first-rig.md b/plans/galaxy-local-first-rig.md index b66ec4d0..3bc77531 100644 --- a/plans/galaxy-local-first-rig.md +++ b/plans/galaxy-local-first-rig.md @@ -336,9 +336,26 @@ rendering and non-Rig compatibility runtimes, not in Rig's executable tool path. ### Phase 4 — Bedrock through Rig -- Implement Bedrock client construction and model resolution through `rig-bedrock`. -- Compare request behavior for system prompts, images, tool schemas, cache controls, reasoning, - inference profiles, token usage, and context limits. +- [x] Pin `rig-bedrock` 0.40.0 and construct it from Galaxy's already-resolved AWS SDK client so + profile, SSO, static-key, region, and egress ownership stay at Galaxy's explicit boundary. +- [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. - 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 -Begin Phase 4 with a focused `rig-bedrock` compatibility spike. Establish client construction and -model/inference-profile resolution first, then add semantic parity fixtures for system prompts, -images, tool schemas, cache controls, reasoning, usage, and context limits before routing any -configured Bedrock model away from the existing compatibility implementation. +Finish Phase 4 with opt-in live Bedrock semantic comparisons for system prompts, images, tools, +signed reasoning, usage, cancellation, and context limits. Keep per-model Rig routing opt-in until +those live fixtures pass, then make Rig the default for supported models and retain the compatibility +runtime only for explicitly unsupported cache behavior.