From f4e2475c60c595d04603f4e5dbb0491360a5eab2 Mon Sep 17 00:00:00 2001 From: Ryan Ward Date: Wed, 6 May 2026 07:02:12 -0500 Subject: [PATCH] Rebasing, going about this another way --- Cargo.lock | 64 ++ app/Cargo.toml | 3 + app/assets/bundled/svg/bedrock.svg | 3 + app/src/ai/agent/api.rs | 10 +- app/src/ai/agent/api/impl.rs | 97 +++ app/src/ai/agent/api/impl_tests.rs | 1 + app/src/ai/aws_credentials.rs | 12 +- app/src/ai/bedrock/client.rs | 174 ++++ app/src/ai/bedrock/convert.rs | 298 +++++++ app/src/ai/bedrock/convert_request.rs | 378 +++++++++ app/src/ai/bedrock/convert_tests.rs | 280 ++++++ app/src/ai/bedrock/discovery.rs | 176 ++++ app/src/ai/bedrock/mod.rs | 13 + app/src/ai/bedrock/models.rs | 107 +++ app/src/ai/bedrock/models_tests.rs | 158 ++++ app/src/ai/bedrock/stream.rs | 343 ++++++++ app/src/ai/bedrock/stream_tests.rs | 69 ++ app/src/ai/blocklist/block.rs | 12 +- .../blocklist/controller/response_stream.rs | 43 +- .../aws_bedrock_credentials_error.rs | 8 +- .../ai/blocklist/passive_suggestions/maa.rs | 2 +- .../ai/execution_profiles/model_menu_items.rs | 9 +- app/src/ai/llms.rs | 136 ++- app/src/ai/mod.rs | 3 + app/src/settings/ai.rs | 180 +++- app/src/settings_view/ai_page.rs | 800 ++++++++++-------- app/src/settings_view/mod.rs | 5 + app/src/terminal/input/slash_commands/mod.rs | 78 +- app/src/terminal/view.rs | 10 +- .../view/inline_banner/session_state.rs | 6 +- app/src/workspace/action.rs | 4 +- app/src/workspace/mod.rs | 4 +- app/src/workspace/view.rs | 4 +- app/src/workspaces/user_workspaces.rs | 4 +- app/src/workspaces/user_workspaces_tests.rs | 10 +- crates/warp_core/src/ui/icons.rs | 2 + 36 files changed, 3040 insertions(+), 466 deletions(-) create mode 100644 app/assets/bundled/svg/bedrock.svg create mode 100644 app/src/ai/bedrock/client.rs create mode 100644 app/src/ai/bedrock/convert.rs create mode 100644 app/src/ai/bedrock/convert_request.rs create mode 100644 app/src/ai/bedrock/convert_tests.rs create mode 100644 app/src/ai/bedrock/discovery.rs create mode 100644 app/src/ai/bedrock/mod.rs create mode 100644 app/src/ai/bedrock/models.rs create mode 100644 app/src/ai/bedrock/models_tests.rs create mode 100644 app/src/ai/bedrock/stream.rs create mode 100644 app/src/ai/bedrock/stream_tests.rs diff --git a/Cargo.lock b/Cargo.lock index 5487866f..4432e1ca 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1550,6 +1550,7 @@ dependencies = [ "aws-credential-types", "aws-sigv4", "aws-smithy-async", + "aws-smithy-eventstream", "aws-smithy-http", "aws-smithy-runtime", "aws-smithy-runtime-api", @@ -1565,6 +1566,53 @@ dependencies = [ "uuid", ] +[[package]] +name = "aws-sdk-bedrock" +version = "1.128.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "793d3212661f68a7a26f8713cd033d4d74f473cacca297126c21abccfae0dc70" +dependencies = [ + "aws-credential-types", + "aws-runtime", + "aws-smithy-async", + "aws-smithy-http", + "aws-smithy-json", + "aws-smithy-runtime", + "aws-smithy-runtime-api", + "aws-smithy-types", + "aws-types", + "bytes", + "fastrand 2.3.0", + "http 0.2.12", + "regex-lite", + "tracing", +] + +[[package]] +name = "aws-sdk-bedrockruntime" +version = "1.120.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "15b8dcf42378ab2d5accac1652cdd059114fb071baf53250ceafb76fcdde347f" +dependencies = [ + "aws-credential-types", + "aws-runtime", + "aws-sigv4", + "aws-smithy-async", + "aws-smithy-eventstream", + "aws-smithy-http", + "aws-smithy-json", + "aws-smithy-runtime", + "aws-smithy-runtime-api", + "aws-smithy-types", + "aws-types", + "bytes", + "fastrand 2.3.0", + "http 0.2.12", + "hyper 0.14.32", + "regex-lite", + "tracing", +] + [[package]] name = "aws-sdk-signin" version = "1.2.0" @@ -1661,6 +1709,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "69e523e1c4e8e7e8ff219d732988e22bfeae8a1cafdbe6d9eca1546fa080be7c" dependencies = [ "aws-credential-types", + "aws-smithy-eventstream", "aws-smithy-http", "aws-smithy-runtime-api", "aws-smithy-types", @@ -1687,12 +1736,24 @@ dependencies = [ "tokio", ] +[[package]] +name = "aws-smithy-eventstream" +version = "0.60.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "dc12f8b310e38cad85cf3bef45ad236f470717393c613266ce0a89512286b650" +dependencies = [ + "aws-smithy-types", + "bytes", + "crc32fast", +] + [[package]] name = "aws-smithy-http" version = "0.62.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "826141069295752372f8203c17f28e30c464d22899a43a0c9fd9c458d469c88b" dependencies = [ + "aws-smithy-eventstream", "aws-smithy-runtime-api", "aws-smithy-types", "bytes", @@ -14162,7 +14223,10 @@ dependencies = [ "async-trait", "aws-config", "aws-credential-types", + "aws-sdk-bedrock", + "aws-sdk-bedrockruntime", "aws-sdk-sts", + "aws-smithy-types", "aws-types", "axum", "base64 0.22.1", diff --git a/app/Cargo.toml b/app/Cargo.toml index 6cd6009f..9270045d 100644 --- a/app/Cargo.toml +++ b/app/Cargo.toml @@ -287,7 +287,10 @@ tokio-util.workspace = true # AWS SDK (loading credentials for BYO LLM) aws-config = { version = "1.8.12", features = ["credentials-login"] } aws-credential-types = "1" +aws-sdk-bedrock = "1" +aws-sdk-bedrockruntime = "1" aws-sdk-sts = "1" +aws-smithy-types = "1" aws-types = "1" rmcp = { workspace = true, features = [ diff --git a/app/assets/bundled/svg/bedrock.svg b/app/assets/bundled/svg/bedrock.svg new file mode 100644 index 00000000..22454afc --- /dev/null +++ b/app/assets/bundled/svg/bedrock.svg @@ -0,0 +1,3 @@ + + + diff --git a/app/src/ai/agent/api.rs b/app/src/ai/agent/api.rs index 2cb27725..584a2fa7 100644 --- a/app/src/ai/agent/api.rs +++ b/app/src/ai/agent/api.rs @@ -124,6 +124,9 @@ pub struct RequestParams { pub research_agent_enabled: bool, pub orchestration_enabled: bool, pub supported_tools_override: Option>, + /// The root task ID for the conversation — needed for direct Bedrock streaming + /// since optimistic tasks don't appear in the proto task_context. + pub root_task_id: Option, /// The conversation ID of the parent agent that spawned this child agent, if any. pub parent_agent_id: Option, /// The display name for this agent (e.g. "Agent 1"), assigned by the orchestrator. @@ -235,7 +238,7 @@ impl RequestParams { let user_workspaces = UserWorkspaces::as_ref(app); let api_keys = ApiKeyManager::as_ref(app).api_keys_for_request( user_workspaces.is_byo_api_key_enabled(), - user_workspaces.is_aws_bedrock_credentials_enabled(app), + user_workspaces.is_bedrock_enabled(app), ); let allow_use_of_warp_credits_with_byok = *AISettings::as_ref(app).can_use_warp_credits_with_byok; @@ -307,6 +310,11 @@ impl RequestParams { research_agent_enabled, orchestration_enabled, supported_tools_override: request_input.supported_tools_override.clone(), + root_task_id: request_input + .input_messages + .keys() + .next() + .map(|id| id.to_string()), parent_agent_id: None, agent_name: None, } diff --git a/app/src/ai/agent/api/impl.rs b/app/src/ai/agent/api/impl.rs index d5390b8b..e2e392a4 100644 --- a/app/src/ai/agent/api/impl.rs +++ b/app/src/ai/agent/api/impl.rs @@ -5,12 +5,14 @@ use futures_util::StreamExt; use warp_core::features::FeatureFlag; use warp_multi_agent_api as api; +use crate::ai::bedrock::client::{BedrockClient, BedrockClientConfig}; use crate::server::server_api::ServerApi; use super::{convert_to::convert_input, ConvertToAPITypeError, RequestParams, ResponseStream}; pub async fn generate_multi_agent_output( server_api: Arc, + bedrock_config: Option, mut params: RequestParams, cancellation_rx: futures::channel::oneshot::Receiver<()>, ) -> Result { @@ -129,6 +131,101 @@ pub async fn generate_multi_agent_output( mcp_context: params.mcp_context.map(Into::into), }; + if let Some(config) = bedrock_config { + let model_id_for_fallback_check = request + .settings + .as_ref() + .and_then(|s| s.model_config.as_ref()) + .map(|mc| mc.base.clone()) + .unwrap_or_default(); + let is_arn = model_id_for_fallback_check.starts_with("arn:"); + let fallback_to_warp = config.fallback_to_warp && !is_arn; + if is_arn && config.fallback_to_warp { + log::info!("[bedrock] Fallback disabled for ARN-based model (not available on Warp server)"); + } + match BedrockClient::from_config(config).await { + Ok(bedrock) => { + let task_id = params.root_task_id.clone().unwrap_or_else(|| { + request + .task_context + .as_ref() + .and_then(|tc| tc.tasks.first()) + .map(|t| t.id.clone()) + .unwrap_or_else(|| uuid::Uuid::new_v4().to_string()) + }); + + log::info!("[bedrock] Starting stream with task_id={task_id}"); + + let model_id = request + .settings + .as_ref() + .and_then(|s| s.model_config.as_ref()) + .map(|mc| mc.base.clone()) + .unwrap_or_default(); + + log::info!("[bedrock] Model: {model_id}"); + + let messages = + crate::ai::bedrock::convert_request::extract_messages_from_request(&request); + let system_prompt = + crate::ai::bedrock::convert_request::extract_system_prompt(&request); + let tools = crate::ai::bedrock::convert_request::extract_tools(&request); + + log::info!( + "[bedrock] Sending {} messages, system_prompt={}, tools={}", + messages.len(), + system_prompt.is_some(), + tools.len() + ); + + match bedrock + .converse_stream( + &model_id, + &task_id, + messages, + system_prompt, + tools, + 8192, + None, + true, + ) + .await + { + Ok(stream) => { + let output_stream = stream.take_until(cancellation_rx); + return Ok(Box::pin(output_stream)); + } + Err(e) => { + if fallback_to_warp { + log::warn!("Bedrock stream failed, falling back to server: {e}"); + } else { + let err = Arc::new(crate::server::server_api::AIApiError::Stream { + stream_type: "bedrock_converse", + source: anyhow::anyhow!("{e}"), + }); + let (tx, rx) = async_channel::unbounded(); + let _ = tx.send(Err(err)).await; + return Ok(Box::pin(rx)); + } + } + } + } + Err(e) => { + if fallback_to_warp { + log::warn!("Bedrock client creation failed, falling back to server: {e}"); + } else { + let err = Arc::new(crate::server::server_api::AIApiError::Stream { + stream_type: "bedrock_converse", + source: anyhow::anyhow!("{e}"), + }); + let (tx, rx) = async_channel::unbounded(); + let _ = tx.send(Err(err)).await; + return Ok(Box::pin(rx)); + } + } + } + } + let response_stream = server_api.generate_multi_agent_output(&request).await; match response_stream { Ok(stream) => { diff --git a/app/src/ai/agent/api/impl_tests.rs b/app/src/ai/agent/api/impl_tests.rs index c219ff5c..73c2679a 100644 --- a/app/src/ai/agent/api/impl_tests.rs +++ b/app/src/ai/agent/api/impl_tests.rs @@ -37,6 +37,7 @@ fn request_params_with_ask_user_question_enabled(ask_user_question_enabled: bool research_agent_enabled: false, orchestration_enabled: false, supported_tools_override: None, + root_task_id: None, parent_agent_id: None, agent_name: None, } diff --git a/app/src/ai/aws_credentials.rs b/app/src/ai/aws_credentials.rs index 3ab3fe1b..a7b56fb4 100644 --- a/app/src/ai/aws_credentials.rs +++ b/app/src/ai/aws_credentials.rs @@ -189,7 +189,7 @@ impl AwsCredentialRefresher for ApiKeyManager { .. }) = event { - let auth_command = &AISettings::as_ref(ctx).aws_bedrock_auth_refresh_command; + let auth_command = &AISettings::as_ref(ctx).bedrock_auth_refresh_command; if command.trim().starts_with(auth_command.trim()) { log::debug!("Detected AWS auth command completion, refreshing credentials"); drop(refresh_aws_credentials(manager, ctx)); @@ -215,9 +215,9 @@ impl AwsCredentialRefresher for ApiKeyManager { ctx.subscribe_to_model(&AISettings::handle(ctx), |manager, event, ctx| { if matches!( event, - AISettingsChangedEvent::AwsBedrockProfile { .. } - | AISettingsChangedEvent::AwsBedrockAuthRefreshCommand { .. } - | AISettingsChangedEvent::AwsBedrockCredentialsEnabled { .. } + AISettingsChangedEvent::BedrockProfile { .. } + | AISettingsChangedEvent::BedrockAuthRefreshCommand { .. } + | AISettingsChangedEvent::BedrockEnabled { .. } ) { drop(refresh_aws_credentials(manager, ctx)); } @@ -248,14 +248,14 @@ fn refresh_aws_credentials_local_chain( manager: &mut ApiKeyManager, ctx: &mut ModelContext, ) -> BoxFuture<'static, Result<(), String>> { - let is_available = UserWorkspaces::as_ref(ctx).is_aws_bedrock_credentials_enabled(ctx); + let is_available = UserWorkspaces::as_ref(ctx).is_bedrock_enabled(ctx); if !is_available { manager.set_aws_credentials_state(AwsCredentialsState::Disabled, ctx); return Box::pin(async { Ok(()) }); } - let profile = (*AISettings::as_ref(ctx).aws_bedrock_profile).clone(); + let profile = (*AISettings::as_ref(ctx).bedrock_profile).clone(); manager.set_aws_credentials_state(AwsCredentialsState::Refreshing, ctx); diff --git a/app/src/ai/bedrock/client.rs b/app/src/ai/bedrock/client.rs new file mode 100644 index 00000000..fa1c2947 --- /dev/null +++ b/app/src/ai/bedrock/client.rs @@ -0,0 +1,174 @@ +use anyhow::Result; +use aws_config::BehaviorVersion; +use aws_sdk_bedrockruntime::config::Region; +use aws_sdk_bedrockruntime::Client as BedrockRuntimeClient; + +use crate::settings::ai::BedrockAuthMethod; + +use super::convert::{build_converse_request, ConversationMessage, ToolDefinition}; +use super::models::apply_cross_region_prefix; +use super::stream::bedrock_stream_to_response_events; +use crate::ai::agent::api::ResponseStream; + +pub struct BedrockClient { + runtime_client: BedrockRuntimeClient, + region: String, +} + +#[derive(Debug, Clone)] +pub struct BedrockClientConfig { + pub auth_method: BedrockAuthMethod, + pub profile: String, + pub region: String, + pub access_key_id: String, + pub secret_access_key: String, + pub cross_region_inference: bool, + pub fallback_to_warp: bool, +} + +#[derive(Debug, thiserror::Error)] +pub enum BedrockError { + #[error("Bedrock credentials not configured")] + CredentialsNotConfigured, + #[error("Bedrock region not configured and could not be auto-detected")] + RegionNotConfigured, + #[error("AWS credential error: {0}")] + CredentialError(String), + #[error("Bedrock API error: {0}")] + ApiError(String), + #[error("Model not found: {0}")] + ModelNotFound(String), + #[error("Access denied: {0}")] + AccessDenied(String), + #[error("Throttling: {0}")] + Throttling(String), + #[error("Validation error: {0}")] + ValidationError(String), +} + +impl BedrockClient { + pub async fn from_config(config: BedrockClientConfig) -> Result { + let aws_config = match config.auth_method { + BedrockAuthMethod::Profile | BedrockAuthMethod::Sso => { + let mut loader = + aws_config::defaults(BehaviorVersion::latest()).profile_name(&config.profile); + + if !config.region.is_empty() { + loader = loader.region(Region::new(config.region.clone())); + } + + loader.load().await + } + BedrockAuthMethod::StaticKeys => { + if config.access_key_id.is_empty() || config.secret_access_key.is_empty() { + return Err(BedrockError::CredentialsNotConfigured); + } + + let creds = aws_credential_types::Credentials::new( + &config.access_key_id, + &config.secret_access_key, + None, + None, + "warp-bedrock-static", + ); + + let mut loader = + aws_config::defaults(BehaviorVersion::latest()).credentials_provider(creds); + + if !config.region.is_empty() { + loader = loader.region(Region::new(config.region.clone())); + } else { + loader = loader.region(Region::new("us-east-1".to_string())); + } + + loader.load().await + } + }; + + let region = aws_config + .region() + .map(|r| r.to_string()) + .ok_or(BedrockError::RegionNotConfigured)?; + + let runtime_client = BedrockRuntimeClient::new(&aws_config); + + Ok(Self { + runtime_client, + region, + }) + } + + pub async fn converse_stream( + &self, + model_id: &str, + task_id: &str, + messages: Vec, + system_prompt: Option, + tools: Vec, + max_tokens: i32, + temperature: Option, + cross_region_inference: bool, + ) -> Result { + let effective_model_id = if cross_region_inference { + apply_cross_region_prefix(model_id, &self.region) + } else { + model_id.to_string() + }; + + log::info!( + "[bedrock] converse_stream: model={effective_model_id}, region={}, messages={}, tools={}", + self.region, + messages.len(), + tools.len() + ); + + let converted = + build_converse_request(messages, system_prompt, tools, max_tokens, temperature, None, None); + + let mut request = self + .runtime_client + .converse_stream() + .model_id(&effective_model_id) + .set_system(Some(converted.system)) + .set_messages(Some(converted.messages)) + .inference_config(converted.inference_config); + + if let Some(tool_config) = converted.tool_config { + request = request.tool_config(tool_config); + } + + let output = request.send().await.map_err(|e| { + let debug_msg = format!("{:?}", e); + let display_msg = format!("{}", e); + log::error!("[bedrock] API error (display): {display_msg}"); + log::error!("[bedrock] API error (debug): {debug_msg}"); + let msg = if debug_msg.len() > display_msg.len() { + debug_msg + } else { + display_msg + }; + if msg.contains("AccessDenied") || msg.contains("access denied") { + BedrockError::AccessDenied(msg) + } else if msg.contains("ThrottlingException") || msg.contains("throttl") { + BedrockError::Throttling(msg) + } else if msg.contains("ValidationException") || msg.contains("validation") { + BedrockError::ValidationError(msg) + } else if msg.contains("ResourceNotFoundException") { + BedrockError::ModelNotFound(effective_model_id.clone()) + } else { + BedrockError::ApiError(msg) + } + })?; + + log::info!("[bedrock] Stream connected successfully"); + Ok(Box::pin(bedrock_stream_to_response_events(output, task_id.to_string()))) + } + + pub fn runtime_client(&self) -> &BedrockRuntimeClient { + &self.runtime_client + } + + pub fn region(&self) -> &str { + &self.region + } +} diff --git a/app/src/ai/bedrock/convert.rs b/app/src/ai/bedrock/convert.rs new file mode 100644 index 00000000..afe190e2 --- /dev/null +++ b/app/src/ai/bedrock/convert.rs @@ -0,0 +1,298 @@ +use std::collections::HashMap; + +use aws_sdk_bedrockruntime::types::{ + ContentBlock, ConversationRole, InferenceConfiguration, Message as BedrockMessage, + SystemContentBlock, Tool, ToolConfiguration, ToolInputSchema, ToolResultBlock, + ToolResultContentBlock, ToolResultStatus, ToolSpecification, ToolUseBlock, +}; +use aws_smithy_types::Document; +use serde_json::Value as JsonValue; + +pub struct ConvertedRequest { + pub messages: Vec, + pub system: Vec, + pub inference_config: InferenceConfiguration, + pub tool_config: Option, +} + +pub struct ConversationMessage { + pub role: MessageRole, + pub content: MessageContent, +} + +pub enum MessageRole { + User, + Assistant, +} + +pub enum MessageContent { + Text(String), + ToolUse { + tool_use_id: String, + name: String, + input: JsonValue, + }, + ToolResult { + tool_use_id: String, + content: String, + is_error: bool, + }, + MultiPart(Vec), +} + +pub enum ContentPart { + Text(String), + ToolUse { + tool_use_id: String, + name: String, + input: JsonValue, + }, + ToolResult { + tool_use_id: String, + content: String, + is_error: bool, + }, +} + +pub struct ToolDefinition { + pub name: String, + pub description: String, + pub input_schema: JsonValue, +} + +pub fn build_converse_request( + messages: Vec, + system_prompt: Option, + tools: Vec, + max_tokens: i32, + temperature: Option, + top_p: Option, + stop_sequences: Option>, +) -> ConvertedRequest { + let bedrock_messages = convert_messages(messages); + let system = convert_system_prompt(system_prompt); + let inference_config = build_inference_config(max_tokens, temperature, top_p, stop_sequences); + let tool_config = build_tool_config(tools); + + ConvertedRequest { + messages: bedrock_messages, + system, + inference_config, + tool_config, + } +} + +fn json_to_document(value: JsonValue) -> Document { + match value { + JsonValue::Null => Document::Null, + JsonValue::Bool(b) => Document::Bool(b), + JsonValue::Number(n) => { + if let Some(i) = n.as_i64() { + Document::Number(aws_smithy_types::Number::PosInt(i as u64)) + } else if let Some(f) = n.as_f64() { + Document::Number(aws_smithy_types::Number::Float(f)) + } else { + Document::Null + } + } + JsonValue::String(s) => Document::String(s), + JsonValue::Array(arr) => { + Document::Array(arr.into_iter().map(json_to_document).collect()) + } + JsonValue::Object(obj) => { + let map: HashMap = obj + .into_iter() + .map(|(k, v)| (k, json_to_document(v))) + .collect(); + Document::Object(map) + } + } +} + +fn convert_messages(messages: Vec) -> Vec { + let mut result = Vec::new(); + + for msg in messages { + let role = match msg.role { + MessageRole::User => ConversationRole::User, + MessageRole::Assistant => ConversationRole::Assistant, + }; + + let content_blocks = match msg.content { + MessageContent::Text(text) => vec![ContentBlock::Text(text)], + MessageContent::ToolUse { + tool_use_id, + name, + input, + } => { + let input_doc = json_to_document(input); + vec![ContentBlock::ToolUse( + ToolUseBlock::builder() + .tool_use_id(tool_use_id) + .name(name) + .input(input_doc) + .build() + .expect("valid tool use block"), + )] + } + MessageContent::ToolResult { + tool_use_id, + content, + is_error, + } => { + let status = if is_error { + ToolResultStatus::Error + } else { + ToolResultStatus::Success + }; + vec![ContentBlock::ToolResult( + ToolResultBlock::builder() + .tool_use_id(tool_use_id) + .status(status) + .content(ToolResultContentBlock::Text(content)) + .build() + .expect("valid tool result block"), + )] + } + MessageContent::MultiPart(parts) => parts + .into_iter() + .map(|part| match part { + ContentPart::Text(text) => ContentBlock::Text(text), + ContentPart::ToolUse { + tool_use_id, + name, + input, + } => { + let input_doc = json_to_document(input); + ContentBlock::ToolUse( + ToolUseBlock::builder() + .tool_use_id(tool_use_id) + .name(name) + .input(input_doc) + .build() + .expect("valid tool use block"), + ) + } + ContentPart::ToolResult { + tool_use_id, + content, + is_error, + } => { + let status = if is_error { + ToolResultStatus::Error + } else { + ToolResultStatus::Success + }; + ContentBlock::ToolResult( + ToolResultBlock::builder() + .tool_use_id(tool_use_id) + .status(status) + .content(ToolResultContentBlock::Text(content)) + .build() + .expect("valid tool result block"), + ) + } + }) + .collect(), + }; + + let message = BedrockMessage::builder() + .role(role) + .set_content(Some(content_blocks)) + .build() + .expect("valid message"); + + result.push(message); + } + + coalesce_consecutive_roles(result) +} + +fn coalesce_consecutive_roles(messages: Vec) -> Vec { + if messages.is_empty() { + return messages; + } + + let mut result: Vec = Vec::new(); + + for msg in messages { + let should_merge = result + .last() + .map(|last| last.role() == msg.role()) + .unwrap_or(false); + + if should_merge { + let last = result.pop().unwrap(); + let mut combined_content: Vec = last.content().to_vec(); + combined_content.extend(msg.content().to_vec()); + let merged = BedrockMessage::builder() + .role(last.role().clone()) + .set_content(Some(combined_content)) + .build() + .expect("valid merged message"); + result.push(merged); + } else { + result.push(msg); + } + } + + result +} + +fn convert_system_prompt(system_prompt: Option) -> Vec { + match system_prompt { + Some(prompt) if !prompt.is_empty() => { + vec![SystemContentBlock::Text(prompt)] + } + _ => vec![], + } +} + +fn build_inference_config( + max_tokens: i32, + temperature: Option, + top_p: Option, + stop_sequences: Option>, +) -> InferenceConfiguration { + let mut builder = InferenceConfiguration::builder().max_tokens(max_tokens); + + if let Some(temp) = temperature { + builder = builder.temperature(temp); + } + if let Some(p) = top_p { + builder = builder.top_p(p); + } + if let Some(stops) = stop_sequences { + builder = builder.set_stop_sequences(Some(stops)); + } + + builder.build() +} + +fn build_tool_config(tools: Vec) -> Option { + if tools.is_empty() { + return None; + } + + let tool_specs: Vec = tools + .into_iter() + .map(|tool| { + let input_schema_doc = json_to_document(tool.input_schema); + Tool::ToolSpec( + ToolSpecification::builder() + .name(tool.name) + .description(tool.description) + .input_schema(ToolInputSchema::Json(input_schema_doc)) + .build() + .expect("valid tool spec"), + ) + }) + .collect(); + + Some( + ToolConfiguration::builder() + .set_tools(Some(tool_specs)) + .build() + .expect("valid tool config"), + ) +} diff --git a/app/src/ai/bedrock/convert_request.rs b/app/src/ai/bedrock/convert_request.rs new file mode 100644 index 00000000..5149ea98 --- /dev/null +++ b/app/src/ai/bedrock/convert_request.rs @@ -0,0 +1,378 @@ +use warp_multi_agent_api as api; + +use super::convert::{ConversationMessage, MessageContent, MessageRole, ToolDefinition}; + +pub fn extract_messages_from_request(request: &api::Request) -> Vec { + let mut messages = Vec::new(); + + if let Some(task_context) = &request.task_context { + for task in &task_context.tasks { + for msg in &task.messages { + if let Some(converted) = convert_proto_message(msg) { + messages.push(converted); + } + } + } + } + + if let Some(input) = &request.input { + if let Some(input_type) = &input.r#type { + #[allow(deprecated)] + match input_type { + api::request::input::Type::UserInputs(user_inputs) => { + for user_input in &user_inputs.inputs { + if let Some(input_variant) = &user_input.input { + match input_variant { + api::request::input::user_inputs::user_input::Input::UserQuery( + query, + ) => { + if !query.query.is_empty() { + messages.push(ConversationMessage { + role: MessageRole::User, + content: MessageContent::Text(query.query.clone()), + }); + } + } + api::request::input::user_inputs::user_input::Input::ToolCallResult( + result, + ) => { + let content = extract_tool_result_content(result); + if !result.tool_call_id.is_empty() { + messages.push(ConversationMessage { + role: MessageRole::User, + content: MessageContent::ToolResult { + tool_use_id: result.tool_call_id.clone(), + content, + is_error: false, + }, + }); + } + } + _ => {} + } + } + } + } + api::request::input::Type::UserQuery(query) => { + if !query.query.is_empty() { + messages.push(ConversationMessage { + role: MessageRole::User, + content: MessageContent::Text(query.query.clone()), + }); + } + } + api::request::input::Type::ToolCallResult(result) => { + let content = extract_tool_result_content(result); + if !result.tool_call_id.is_empty() { + messages.push(ConversationMessage { + role: MessageRole::User, + content: MessageContent::ToolResult { + tool_use_id: result.tool_call_id.clone(), + content, + is_error: false, + }, + }); + } + } + _ => {} + } + } + } + + messages +} + +pub fn extract_system_prompt(_request: &api::Request) -> Option { + Some("You are a helpful AI coding assistant. You help users with software engineering tasks including writing code, debugging, and explaining concepts.".to_string()) +} + +pub fn extract_tools(request: &api::Request) -> Vec { + let mut tools = Vec::new(); + let mut seen_names = std::collections::HashSet::new(); + + if let Some(task_context) = &request.task_context { + for task in &task_context.tasks { + for msg in &task.messages { + if let Some(api::message::Message::ToolCall(tool_call)) = &msg.message { + let (name, _) = extract_tool_call_info(tool_call); + if name != "unknown_tool" && seen_names.insert(name.clone()) { + tools.push(tool_definition_for_name(&name)); + } + } + } + } + } + + if let Some(input) = &request.input { + if let Some(input_type) = &input.r#type { + #[allow(deprecated)] + match input_type { + api::request::input::Type::UserInputs(user_inputs) => { + for user_input in &user_inputs.inputs { + if let Some(api::request::input::user_inputs::user_input::Input::ToolCallResult(_)) = &user_input.input { + break; + } + } + } + api::request::input::Type::ToolCallResult(_) => {} + _ => {} + } + } + } + + if tools.is_empty() { + let messages = extract_messages_from_request(request); + let has_tool_content = messages.iter().any(|m| { + matches!( + m.content, + MessageContent::ToolUse { .. } | MessageContent::ToolResult { .. } + ) + }); + if has_tool_content { + tools = default_tool_definitions(); + } + } + + tools +} + +fn tool_definition_for_name(name: &str) -> ToolDefinition { + match name { + "run_shell_command" => ToolDefinition { + name: "run_shell_command".to_string(), + description: "Execute a shell command and return its output.".to_string(), + input_schema: serde_json::json!({ + "type": "object", + "properties": { + "command": { "type": "string", "description": "The shell command to execute" } + }, + "required": ["command"] + }), + }, + "read_files" => ToolDefinition { + name: "read_files".to_string(), + description: "Read the contents of one or more files.".to_string(), + input_schema: serde_json::json!({ + "type": "object", + "properties": { + "files": { "type": "array", "items": { "type": "string" }, "description": "File paths to read" } + }, + "required": ["files"] + }), + }, + "apply_file_diffs" => ToolDefinition { + name: "apply_file_diffs".to_string(), + description: "Apply search/replace diffs to files.".to_string(), + input_schema: serde_json::json!({ + "type": "object", + "properties": { + "diffs": { "type": "array", "items": { "type": "object", "properties": { "file_path": { "type": "string" }, "search": { "type": "string" }, "replace": { "type": "string" } }, "required": ["file_path", "search", "replace"] } } + }, + "required": ["diffs"] + }), + }, + "grep" => ToolDefinition { + name: "grep".to_string(), + description: "Search for patterns in files using grep.".to_string(), + input_schema: serde_json::json!({ + "type": "object", + "properties": { + "queries": { "type": "array", "items": { "type": "string" }, "description": "Search patterns" }, + "path": { "type": "string", "description": "Directory to search in" } + }, + "required": ["queries"] + }), + }, + "file_glob" => ToolDefinition { + name: "file_glob".to_string(), + description: "Find files matching glob patterns.".to_string(), + input_schema: serde_json::json!({ + "type": "object", + "properties": { + "patterns": { "type": "array", "items": { "type": "string" }, "description": "Glob patterns to match" } + }, + "required": ["patterns"] + }), + }, + _ => ToolDefinition { + name: name.to_string(), + description: format!("Tool: {}", name), + input_schema: serde_json::json!({ + "type": "object", + "properties": {} + }), + }, + } +} + +fn default_tool_definitions() -> Vec { + vec![ + tool_definition_for_name("run_shell_command"), + tool_definition_for_name("read_files"), + tool_definition_for_name("apply_file_diffs"), + tool_definition_for_name("grep"), + tool_definition_for_name("file_glob"), + ] +} + +fn convert_proto_message(msg: &api::Message) -> Option { + let message_content = msg.message.as_ref()?; + + match message_content { + api::message::Message::UserQuery(query) => Some(ConversationMessage { + role: MessageRole::User, + content: MessageContent::Text(query.query.clone()), + }), + api::message::Message::AgentOutput(output) => Some(ConversationMessage { + role: MessageRole::Assistant, + content: MessageContent::Text(output.text.clone()), + }), + api::message::Message::ToolCall(tool_call) => { + let (name, input) = extract_tool_call_info(tool_call); + Some(ConversationMessage { + role: MessageRole::Assistant, + content: MessageContent::ToolUse { + tool_use_id: msg.id.clone(), + name, + input, + }, + }) + } + api::message::Message::ToolCallResult(result) => { + let content = format_tool_call_result(result); + Some(ConversationMessage { + role: MessageRole::User, + content: MessageContent::ToolResult { + tool_use_id: result.tool_call_id.clone(), + content, + is_error: false, + }, + }) + } + api::message::Message::AgentReasoning(_) => None, + _ => None, + } +} + +fn extract_tool_call_info(tool_call: &api::message::ToolCall) -> (String, serde_json::Value) { + if let Some(tool) = &tool_call.tool { + match tool { + api::message::tool_call::Tool::RunShellCommand(cmd) => ( + "run_shell_command".to_string(), + serde_json::json!({ "command": cmd.command }), + ), + api::message::tool_call::Tool::ReadFiles(read) => ( + "read_files".to_string(), + serde_json::json!({ "files": read.files.iter().map(|f| &f.name).collect::>() }), + ), + api::message::tool_call::Tool::ApplyFileDiffs(diffs) => ( + "apply_file_diffs".to_string(), + serde_json::json!({ "diffs": diffs.diffs.iter().map(|d| { + serde_json::json!({ + "file_path": d.file_path, + "search": d.search, + "replace": d.replace + }) + }).collect::>() }), + ), + api::message::tool_call::Tool::Grep(grep) => ( + "grep".to_string(), + serde_json::json!({ "queries": grep.queries, "path": grep.path }), + ), + #[allow(deprecated)] + api::message::tool_call::Tool::FileGlob(glob) => ( + "file_glob".to_string(), + serde_json::json!({ "patterns": glob.patterns }), + ), + _ => ("unknown_tool".to_string(), serde_json::json!({})), + } + } else { + ("unknown_tool".to_string(), serde_json::json!({})) + } +} + +fn extract_tool_result_content(result: &api::request::input::ToolCallResult) -> String { + if let Some(result_type) = &result.result { + match result_type { + api::request::input::tool_call_result::Result::RunShellCommand(cmd_result) => { + match &cmd_result.result { + Some( + api::run_shell_command_result::Result::CommandFinished(finished), + ) => finished.output.clone(), + Some( + api::run_shell_command_result::Result::LongRunningCommandSnapshot( + snapshot, + ), + ) => snapshot.output.clone(), + _ => "Command completed.".to_string(), + } + } + api::request::input::tool_call_result::Result::ReadFiles(read_result) => { + match &read_result.result { + Some(api::read_files_result::Result::TextFilesSuccess(success)) => success + .files + .iter() + .map(|f| format!("{}:\n{}", f.file_path, f.content)) + .collect::>() + .join("\n\n"), + Some(api::read_files_result::Result::AnyFilesSuccess(success)) => success + .files + .iter() + .filter_map(|f| match &f.content { + Some(api::any_file_content::Content::TextContent(t)) => { + Some(format!("{}:\n{}", t.file_path, t.content)) + } + _ => None, + }) + .collect::>() + .join("\n\n"), + _ => "Failed to read files.".to_string(), + } + } + _ => "Tool completed successfully.".to_string(), + } + } else { + "Tool completed.".to_string() + } +} + +fn format_tool_call_result(result: &api::message::ToolCallResult) -> String { + if let Some(result_type) = &result.result { + match result_type { + api::message::tool_call_result::Result::RunShellCommand(cmd_result) => { + match &cmd_result.result { + Some( + api::run_shell_command_result::Result::CommandFinished(finished), + ) => { + format!( + "Exit code: {}\nOutput: {}", + finished.exit_code, finished.output + ) + } + Some( + api::run_shell_command_result::Result::LongRunningCommandSnapshot( + snapshot, + ), + ) => { + format!("Output (running): {}", snapshot.output) + } + _ => "Command completed.".to_string(), + } + } + api::message::tool_call_result::Result::ReadFiles(read_result) => { + match &read_result.result { + Some(api::read_files_result::Result::TextFilesSuccess(success)) => success + .files + .iter() + .map(|f| format!("{}:\n{}", f.file_path, f.content)) + .collect::>() + .join("\n\n"), + _ => "Read files completed.".to_string(), + } + } + _ => "Tool completed successfully.".to_string(), + } + } else { + "Tool completed.".to_string() + } +} diff --git a/app/src/ai/bedrock/convert_tests.rs b/app/src/ai/bedrock/convert_tests.rs new file mode 100644 index 00000000..c989f644 --- /dev/null +++ b/app/src/ai/bedrock/convert_tests.rs @@ -0,0 +1,280 @@ +use aws_sdk_bedrockruntime::types::{ContentBlock, ConversationRole}; +use serde_json::json; + +use super::convert::*; + +#[test] +fn test_text_message_converts_to_single_block() { + let messages = vec![ConversationMessage { + role: MessageRole::User, + content: MessageContent::Text("Hello".to_string()), + }]; + + let result = build_converse_request(messages, None, vec![], 4096, None, None, None); + + assert_eq!(result.messages.len(), 1); + assert_eq!(result.messages[0].role(), &ConversationRole::User); + assert_eq!(result.messages[0].content().len(), 1); + assert!(matches!(&result.messages[0].content()[0], ContentBlock::Text(t) if t == "Hello")); +} + +#[test] +fn test_tool_use_produces_valid_json_input() { + let messages = vec![ConversationMessage { + role: MessageRole::Assistant, + content: MessageContent::ToolUse { + tool_use_id: "tool_123".to_string(), + name: "read_file".to_string(), + input: json!({"path": "/tmp/test.txt"}), + }, + }]; + + let result = build_converse_request(messages, None, vec![], 4096, None, None, None); + + assert_eq!(result.messages.len(), 1); + assert_eq!(result.messages[0].role(), &ConversationRole::Assistant); + match &result.messages[0].content()[0] { + ContentBlock::ToolUse(block) => { + assert_eq!(block.tool_use_id(), "tool_123"); + assert_eq!(block.name(), "read_file"); + } + other => panic!("Expected ToolUse, got {:?}", other), + } +} + +#[test] +fn test_tool_result_with_matching_id() { + let messages = vec![ConversationMessage { + role: MessageRole::User, + content: MessageContent::ToolResult { + tool_use_id: "tool_123".to_string(), + content: "file contents here".to_string(), + is_error: false, + }, + }]; + + let result = build_converse_request(messages, None, vec![], 4096, None, None, None); + + assert_eq!(result.messages.len(), 1); + match &result.messages[0].content()[0] { + ContentBlock::ToolResult(block) => { + assert_eq!(block.tool_use_id(), "tool_123"); + } + other => panic!("Expected ToolResult, got {:?}", other), + } +} + +#[test] +fn test_tool_result_error_status() { + let messages = vec![ConversationMessage { + role: MessageRole::User, + content: MessageContent::ToolResult { + tool_use_id: "tool_456".to_string(), + content: "permission denied".to_string(), + is_error: true, + }, + }]; + + let result = build_converse_request(messages, None, vec![], 4096, None, None, None); + + match &result.messages[0].content()[0] { + ContentBlock::ToolResult(block) => { + assert_eq!( + block.status(), + Some(&aws_sdk_bedrockruntime::types::ToolResultStatus::Error) + ); + } + other => panic!("Expected ToolResult, got {:?}", other), + } +} + +#[test] +fn test_consecutive_same_role_messages_coalesced() { + let messages = vec![ + ConversationMessage { + role: MessageRole::User, + content: MessageContent::Text("first".to_string()), + }, + ConversationMessage { + role: MessageRole::User, + content: MessageContent::Text("second".to_string()), + }, + ]; + + let result = build_converse_request(messages, None, vec![], 4096, None, None, None); + + assert_eq!(result.messages.len(), 1); + assert_eq!(result.messages[0].content().len(), 2); + assert!(matches!(&result.messages[0].content()[0], ContentBlock::Text(t) if t == "first")); + assert!(matches!(&result.messages[0].content()[1], ContentBlock::Text(t) if t == "second")); +} + +#[test] +fn test_alternating_roles_not_coalesced() { + let messages = vec![ + ConversationMessage { + role: MessageRole::User, + content: MessageContent::Text("question".to_string()), + }, + ConversationMessage { + role: MessageRole::Assistant, + content: MessageContent::Text("answer".to_string()), + }, + ConversationMessage { + role: MessageRole::User, + content: MessageContent::Text("followup".to_string()), + }, + ]; + + let result = build_converse_request(messages, None, vec![], 4096, None, None, None); + + assert_eq!(result.messages.len(), 3); + assert_eq!(result.messages[0].role(), &ConversationRole::User); + assert_eq!(result.messages[1].role(), &ConversationRole::Assistant); + assert_eq!(result.messages[2].role(), &ConversationRole::User); +} + +#[test] +fn test_system_prompt_separated_from_messages() { + let messages = vec![ConversationMessage { + role: MessageRole::User, + content: MessageContent::Text("hi".to_string()), + }]; + + let result = build_converse_request( + messages, + Some("You are a helpful assistant.".to_string()), + vec![], + 4096, + None, + None, + None, + ); + + assert_eq!(result.system.len(), 1); + assert_eq!(result.messages.len(), 1); + assert_eq!(result.messages[0].role(), &ConversationRole::User); +} + +#[test] +fn test_empty_system_prompt_produces_empty_vec() { + let result = build_converse_request(vec![], Some("".to_string()), vec![], 4096, None, None, None); + assert!(result.system.is_empty()); + + let result2 = build_converse_request(vec![], None, vec![], 4096, None, None, None); + assert!(result2.system.is_empty()); +} + +#[test] +fn test_empty_tools_produce_none_config() { + let result = build_converse_request(vec![], None, vec![], 4096, None, None, None); + assert!(result.tool_config.is_none()); +} + +#[test] +fn test_tool_definitions_produce_tool_config() { + let tools = vec![ToolDefinition { + name: "read_file".to_string(), + description: "Read a file from disk".to_string(), + input_schema: json!({ + "type": "object", + "properties": { + "path": {"type": "string", "description": "File path"} + }, + "required": ["path"] + }), + }]; + + let result = build_converse_request(vec![], None, tools, 4096, None, None, None); + + assert!(result.tool_config.is_some()); + let config = result.tool_config.unwrap(); + assert_eq!(config.tools().len(), 1); +} + +#[test] +fn test_inference_config_max_tokens_only() { + let result = build_converse_request(vec![], None, vec![], 8192, None, None, None); + assert_eq!(result.inference_config.max_tokens(), Some(8192)); + assert_eq!(result.inference_config.temperature(), None); + assert_eq!(result.inference_config.top_p(), None); + assert!(result.inference_config.stop_sequences().is_empty()); +} + +#[test] +fn test_inference_config_all_params() { + let result = build_converse_request( + vec![], + None, + vec![], + 4096, + Some(0.7), + Some(0.9), + Some(vec!["STOP".to_string()]), + ); + assert_eq!(result.inference_config.max_tokens(), Some(4096)); + assert_eq!(result.inference_config.temperature(), Some(0.7)); + assert_eq!(result.inference_config.top_p(), Some(0.9)); + assert_eq!(result.inference_config.stop_sequences(), &["STOP"]); +} + +#[test] +fn test_multipart_content_produces_multiple_blocks() { + let messages = vec![ConversationMessage { + role: MessageRole::Assistant, + content: MessageContent::MultiPart(vec![ + ContentPart::Text("Let me help.".to_string()), + ContentPart::ToolUse { + tool_use_id: "tu_1".to_string(), + name: "run_command".to_string(), + input: json!({"command": "ls"}), + }, + ]), + }]; + + let result = build_converse_request(messages, None, vec![], 4096, None, None, None); + + assert_eq!(result.messages[0].content().len(), 2); + assert!(matches!(&result.messages[0].content()[0], ContentBlock::Text(_))); + assert!(matches!( + &result.messages[0].content()[1], + ContentBlock::ToolUse(_) + )); +} + +#[test] +fn test_tool_result_after_tool_use_coalesced_into_user_message() { + let messages = vec![ + ConversationMessage { + role: MessageRole::User, + content: MessageContent::Text("Do something".to_string()), + }, + ConversationMessage { + role: MessageRole::Assistant, + content: MessageContent::ToolUse { + tool_use_id: "tu_1".to_string(), + name: "cmd".to_string(), + input: json!({}), + }, + }, + ConversationMessage { + role: MessageRole::User, + content: MessageContent::ToolResult { + tool_use_id: "tu_1".to_string(), + content: "done".to_string(), + is_error: false, + }, + }, + ]; + + let result = build_converse_request(messages, None, vec![], 4096, None, None, None); + + assert_eq!(result.messages.len(), 3); + assert_eq!(result.messages[0].role(), &ConversationRole::User); + assert_eq!(result.messages[1].role(), &ConversationRole::Assistant); + assert_eq!(result.messages[2].role(), &ConversationRole::User); + assert!(matches!( + &result.messages[2].content()[0], + ContentBlock::ToolResult(_) + )); +} diff --git a/app/src/ai/bedrock/discovery.rs b/app/src/ai/bedrock/discovery.rs new file mode 100644 index 00000000..7fbfe407 --- /dev/null +++ b/app/src/ai/bedrock/discovery.rs @@ -0,0 +1,176 @@ +use std::collections::BTreeSet; +use std::path::PathBuf; + +use anyhow::Result; +use aws_config::BehaviorVersion; +use aws_sdk_bedrock::types::InferenceProfileType; +use aws_sdk_bedrock::Client as BedrockControlClient; +use aws_sdk_bedrockruntime::config::Region; + +use crate::settings::ai::{BedrockAuthMethod, BedrockModelConfig}; + +use super::client::BedrockClientConfig; + +pub fn list_aws_profiles() -> Vec { + let mut profiles = BTreeSet::new(); + + if let Some(home) = dirs::home_dir() { + parse_config_file(home.join(".aws").join("config"), &mut profiles); + parse_credentials_file(home.join(".aws").join("credentials"), &mut profiles); + } + + profiles.into_iter().collect() +} + +fn parse_config_file(path: PathBuf, profiles: &mut BTreeSet) { + let contents = match std::fs::read_to_string(&path) { + Ok(c) => c, + Err(_) => return, + }; + + for line in contents.lines() { + let trimmed = line.trim(); + if trimmed.starts_with('[') && trimmed.ends_with(']') { + let section = &trimmed[1..trimmed.len() - 1]; + if section == "default" { + profiles.insert("default".to_string()); + } else if let Some(name) = section.strip_prefix("profile ") { + profiles.insert(name.trim().to_string()); + } + } + } +} + +fn parse_credentials_file(path: PathBuf, profiles: &mut BTreeSet) { + let contents = match std::fs::read_to_string(&path) { + Ok(c) => c, + Err(_) => return, + }; + + for line in contents.lines() { + let trimmed = line.trim(); + if trimmed.starts_with('[') && trimmed.ends_with(']') { + let name = &trimmed[1..trimmed.len() - 1]; + profiles.insert(name.trim().to_string()); + } + } +} + +pub async fn discover_inference_profiles( + config: &BedrockClientConfig, +) -> Result> { + let aws_config = build_aws_config(config).await; + let client = BedrockControlClient::new(&aws_config); + + let mut models = Vec::new(); + + fetch_profiles_by_type(&client, InferenceProfileType::SystemDefined, &mut models).await?; + fetch_profiles_by_type(&client, InferenceProfileType::Application, &mut models).await?; + + models.sort_by(|a, b| a.display_name.cmp(&b.display_name)); + models.dedup_by(|a, b| a.model_id == b.model_id); + + Ok(models) +} + +async fn fetch_profiles_by_type( + client: &BedrockControlClient, + profile_type: InferenceProfileType, + models: &mut Vec, +) -> Result<()> { + let mut next_token: Option = None; + + loop { + let mut req = client + .list_inference_profiles() + .type_equals(profile_type.clone()) + .max_results(100); + if let Some(token) = next_token.take() { + req = req.next_token(token); + } + + let resp = req.send().await?; + + for summary in resp.inference_profile_summaries() { + let profile_id = summary.inference_profile_id(); + let profile_name = summary.inference_profile_name(); + let profile_arn = summary.inference_profile_arn(); + + let model_id = if profile_arn.contains(":application-inference-profile/") { + profile_arn.to_string() + } else { + profile_id.to_string() + }; + + if should_skip_model(profile_name) { + continue; + } + + models.push(BedrockModelConfig { + model_id, + display_name: profile_name.to_string(), + vision_supported: true, + }); + } + + next_token = resp.next_token().map(|s| s.to_string()); + if next_token.is_none() { + break; + } + } + + Ok(()) +} + +fn should_skip_model(name: &str) -> bool { + let lower = name.to_lowercase(); + lower.contains("embed") + || lower.contains("stable image") + || lower.contains("stable-image") + || lower.contains("upscale") + || lower.contains("outpaint") + || lower.contains("inpaint") + || lower.contains("recolor") + || lower.contains("erase") + || lower.contains("style transfer") + || lower.contains("style guide") + || lower.contains("remove background") + || lower.contains("search and replace") + || lower.contains("control sketch") + || lower.contains("control structure") +} + +async fn build_aws_config(config: &BedrockClientConfig) -> aws_config::SdkConfig { + match config.auth_method { + BedrockAuthMethod::Profile | BedrockAuthMethod::Sso => { + let mut loader = + aws_config::defaults(BehaviorVersion::latest()).profile_name(&config.profile); + + if !config.region.is_empty() { + loader = loader.region(Region::new(config.region.clone())); + } + + loader.load().await + } + BedrockAuthMethod::StaticKeys => { + let creds = aws_credential_types::Credentials::new( + &config.access_key_id, + &config.secret_access_key, + None, + None, + "warp-bedrock-discovery", + ); + + let mut loader = + aws_config::defaults(BehaviorVersion::latest()).credentials_provider(creds); + + if !config.region.is_empty() { + loader = loader.region(Region::new(config.region.clone())); + } else { + loader = loader.region(Region::new("us-east-1".to_string())); + } + + loader.load().await + } + } +} diff --git a/app/src/ai/bedrock/mod.rs b/app/src/ai/bedrock/mod.rs new file mode 100644 index 00000000..3a8352d7 --- /dev/null +++ b/app/src/ai/bedrock/mod.rs @@ -0,0 +1,13 @@ +pub mod client; +pub mod convert; +pub mod convert_request; +pub mod discovery; +pub mod models; +pub mod stream; + +#[cfg(test)] +mod convert_tests; +#[cfg(test)] +mod models_tests; +#[cfg(test)] +mod stream_tests; diff --git a/app/src/ai/bedrock/models.rs b/app/src/ai/bedrock/models.rs new file mode 100644 index 00000000..364a1a2f --- /dev/null +++ b/app/src/ai/bedrock/models.rs @@ -0,0 +1,107 @@ +use crate::settings::ai::BedrockModelConfig; + +pub struct DefaultModel { + pub model_id: &'static str, + pub display_name: &'static str, + pub vision_supported: bool, +} + +pub const DEFAULT_BEDROCK_MODELS: &[DefaultModel] = &[ + DefaultModel { + model_id: "anthropic.claude-opus-4-7", + display_name: "Claude Opus 4.7", + vision_supported: true, + }, + DefaultModel { + model_id: "anthropic.claude-sonnet-4-6", + display_name: "Claude Sonnet 4.6", + vision_supported: true, + }, + DefaultModel { + model_id: "anthropic.claude-sonnet-4-20250514-v1:0", + display_name: "Claude Sonnet 4", + vision_supported: true, + }, + DefaultModel { + model_id: "anthropic.claude-haiku-4-5-20251001-v1:0", + display_name: "Claude Haiku 4.5", + vision_supported: true, + }, + DefaultModel { + model_id: "amazon.nova-pro-v1:0", + display_name: "Amazon Nova Pro", + vision_supported: true, + }, + DefaultModel { + model_id: "amazon.nova-lite-v1:0", + display_name: "Amazon Nova Lite", + vision_supported: true, + }, + DefaultModel { + model_id: "amazon.nova-micro-v1:0", + display_name: "Amazon Nova Micro", + vision_supported: false, + }, + DefaultModel { + model_id: "deepseek.r1-v1:0", + display_name: "DeepSeek R1", + vision_supported: false, + }, +]; + +pub fn get_effective_models(user_models: &[BedrockModelConfig]) -> Vec { + if user_models.is_empty() { + DEFAULT_BEDROCK_MODELS + .iter() + .map(|m| BedrockModelConfig { + model_id: m.model_id.to_string(), + display_name: m.display_name.to_string(), + vision_supported: m.vision_supported, + }) + .collect() + } else { + user_models.to_vec() + } +} + +pub fn apply_cross_region_prefix(model_id: &str, region: &str) -> String { + if model_id.starts_with("arn:") { + return model_id.to_string(); + } + + if model_id.contains('.') && model_id.split('.').next().unwrap_or("").len() <= 6 { + return model_id.to_string(); + } + + let prefix = match region { + r if r.starts_with("us-") || r.starts_with("ca-") => "us", + r if r.starts_with("eu-") || r == "il-central-1" => "eu", + r if r == "ap-northeast-1" || r == "ap-northeast-3" => "jp", + r if r == "ap-southeast-2" || r == "ap-southeast-4" || r == "ap-southeast-6" => "au", + r if r.starts_with("ap-") => "apac", + _ => return model_id.to_string(), + }; + format!("{}.{}", prefix, model_id) +} + +pub fn is_bedrock_model(model_id: &str, configured_models: &[BedrockModelConfig]) -> bool { + if model_id.starts_with("arn:aws:bedrock:") { + return true; + } + + let effective = get_effective_models(configured_models); + effective.iter().any(|m| m.model_id == model_id) + || model_id.starts_with("anthropic.") + || model_id.starts_with("amazon.") + || model_id.starts_with("meta.") + || model_id.starts_with("mistral.") + || model_id.starts_with("cohere.") + || model_id.starts_with("ai21.") + || model_id.starts_with("deepseek.") + || has_cross_region_prefix(model_id) +} + +fn has_cross_region_prefix(model_id: &str) -> bool { + let prefixes = ["us.", "eu.", "jp.", "apac.", "au.", "global."]; + prefixes.iter().any(|p| model_id.starts_with(p)) +} diff --git a/app/src/ai/bedrock/models_tests.rs b/app/src/ai/bedrock/models_tests.rs new file mode 100644 index 00000000..df9e590f --- /dev/null +++ b/app/src/ai/bedrock/models_tests.rs @@ -0,0 +1,158 @@ +use super::models::*; +use crate::settings::ai::BedrockModelConfig; + +#[test] +fn test_cross_region_prefix_us_east() { + assert_eq!( + apply_cross_region_prefix("anthropic.claude-3-5-sonnet-20241022-v1:0", "us-east-1"), + "us.anthropic.claude-3-5-sonnet-20241022-v1:0" + ); +} + +#[test] +fn test_cross_region_prefix_eu_west() { + assert_eq!( + apply_cross_region_prefix("anthropic.claude-3-5-sonnet-20241022-v1:0", "eu-west-1"), + "eu.anthropic.claude-3-5-sonnet-20241022-v1:0" + ); +} + +#[test] +fn test_cross_region_prefix_ap_northeast_1() { + assert_eq!( + apply_cross_region_prefix("anthropic.claude-3-5-sonnet-20241022-v1:0", "ap-northeast-1"), + "jp.anthropic.claude-3-5-sonnet-20241022-v1:0" + ); +} + +#[test] +fn test_cross_region_prefix_ap_southeast_2() { + assert_eq!( + apply_cross_region_prefix("anthropic.claude-3-5-sonnet-20241022-v1:0", "ap-southeast-2"), + "au.anthropic.claude-3-5-sonnet-20241022-v1:0" + ); +} + +#[test] +fn test_cross_region_prefix_ap_southeast_1() { + assert_eq!( + apply_cross_region_prefix("anthropic.claude-3-5-sonnet-20241022-v1:0", "ap-southeast-1"), + "apac.anthropic.claude-3-5-sonnet-20241022-v1:0" + ); +} + +#[test] +fn test_cross_region_prefix_already_prefixed() { + assert_eq!( + apply_cross_region_prefix("us.anthropic.claude-sonnet-4-6", "us-east-1"), + "us.anthropic.claude-sonnet-4-6" + ); + assert_eq!( + apply_cross_region_prefix("global.anthropic.claude-sonnet-4-6", "us-east-1"), + "global.anthropic.claude-sonnet-4-6" + ); +} + +#[test] +fn test_cross_region_prefix_canada_maps_to_us() { + assert_eq!( + apply_cross_region_prefix("anthropic.claude-sonnet-4-6", "ca-central-1"), + "us.anthropic.claude-sonnet-4-6" + ); +} + +#[test] +fn test_cross_region_prefix_unknown_region() { + assert_eq!( + apply_cross_region_prefix("anthropic.claude-3-5-sonnet-20241022-v1:0", "me-south-1"), + "anthropic.claude-3-5-sonnet-20241022-v1:0" + ); +} + +#[test] +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-7"); + assert_eq!(models[0].display_name, "Claude Opus 4.7"); +} + +#[test] +fn test_get_effective_models_custom_overrides() { + let custom = vec![BedrockModelConfig { + model_id: "custom.model-v1:0".to_string(), + display_name: "Custom Model".to_string(), + vision_supported: false, + }]; + let models = get_effective_models(&custom); + assert_eq!(models.len(), 1); + assert_eq!(models[0].model_id, "custom.model-v1:0"); +} + +#[test] +fn test_is_bedrock_model_known_prefix() { + assert!(is_bedrock_model("anthropic.claude-sonnet-4-6", &[])); + assert!(is_bedrock_model("amazon.nova-pro-v1:0", &[])); + assert!(is_bedrock_model("meta.llama3-70b-instruct-v1:0", &[])); + assert!(is_bedrock_model("mistral.mistral-large-v1:0", &[])); + assert!(is_bedrock_model("deepseek.r1-v1:0", &[])); +} + +#[test] +fn test_is_bedrock_model_cross_region_prefix() { + assert!(is_bedrock_model("us.anthropic.claude-sonnet-4-6", &[])); + assert!(is_bedrock_model("eu.anthropic.claude-sonnet-4-6", &[])); + assert!(is_bedrock_model("global.anthropic.claude-opus-4-7", &[])); +} + +#[test] +fn test_is_bedrock_model_unknown() { + assert!(!is_bedrock_model("gpt-4o", &[])); + assert!(!is_bedrock_model("gemini-pro", &[])); +} + +#[test] +fn test_is_bedrock_model_custom_config() { + let custom = vec![BedrockModelConfig { + model_id: "custom.my-model-v1:0".to_string(), + display_name: "Custom".to_string(), + vision_supported: false, + }]; + assert!(is_bedrock_model("custom.my-model-v1:0", &custom)); +} + +#[test] +fn test_is_bedrock_model_arn() { + assert!(is_bedrock_model( + "arn:aws:bedrock:us-east-1:156729649053:application-inference-profile/fma5bsw4oxhy", + &[], + )); +} + +#[test] +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 test_is_bedrock_model_coding_agent_arn() { + assert!(is_bedrock_model( + "arn:aws:bedrock:us-east-1:156729649053:application-inference-profile/coding-agent-anthropic-claude-opus-4-6-lt2v72", + &[], + )); +} + +#[test] +fn test_is_bedrock_model_custom_config_with_arn() { + let custom = vec![BedrockModelConfig { + model_id: "arn:aws:bedrock:us-east-1:156729649053:application-inference-profile/coding-assistant-inference-profile".to_string(), + display_name: "Coding Assistant".to_string(), + vision_supported: true, + }]; + assert!(is_bedrock_model( + "arn:aws:bedrock:us-east-1:156729649053:application-inference-profile/coding-assistant-inference-profile", + &custom, + )); +} diff --git a/app/src/ai/bedrock/stream.rs b/app/src/ai/bedrock/stream.rs new file mode 100644 index 00000000..e6b6739c --- /dev/null +++ b/app/src/ai/bedrock/stream.rs @@ -0,0 +1,343 @@ +use std::sync::Arc; + +use aws_sdk_bedrockruntime::operation::converse_stream::ConverseStreamOutput; +use aws_sdk_bedrockruntime::types::{ + ContentBlockDelta, ContentBlockStart, ConverseStreamOutput as StreamEvent, + ReasoningContentBlockDelta, StopReason, +}; +use futures::stream::BoxStream; +use uuid::Uuid; +use warp_multi_agent_api::response_event::stream_finished; +use warp_multi_agent_api::{self as api, ClientAction, ResponseEvent}; + +use crate::ai::agent::api::Event; +use crate::server::server_api::AIApiError; + +pub fn bedrock_stream_to_response_events( + mut output: ConverseStreamOutput, + task_id: String, +) -> BoxStream<'static, Event> { + let request_id = Uuid::new_v4().to_string(); + let conversation_id = Uuid::new_v4().to_string(); + + let stream = async_stream::stream! { + log::info!("[bedrock] Stream started: task_id={task_id}, request_id={request_id}"); + let init_event = build_stream_init(&request_id, &conversation_id); + yield Ok(init_event); + + let mut current_text_message_id: Option = None; + let mut current_tool_use_id = String::new(); + let mut current_tool_name = String::new(); + let mut current_tool_input_json = String::new(); + let mut input_tokens: i32 = 0; + let mut output_tokens: i32 = 0; + let mut stop_reason = stream_finished::Reason::Done(stream_finished::Done {}); + + loop { + match output.stream.recv().await { + Ok(Some(event)) => match event { + StreamEvent::MessageStart(_) => {} + StreamEvent::ContentBlockStart(block_start) => { + if let Some(start) = block_start.start() { + match start { + ContentBlockStart::ToolUse(tool_start) => { + current_tool_use_id = tool_start.tool_use_id().to_string(); + current_tool_name = tool_start.name().to_string(); + current_tool_input_json.clear(); + } + _ => {} + } + } + } + StreamEvent::ContentBlockDelta(delta) => { + if let Some(d) = delta.delta() { + match d { + ContentBlockDelta::Text(text) => { + log::trace!("[bedrock] Text delta ({} chars): {:?}", text.len(), &text[..text.len().min(100)]); + if current_text_message_id.is_none() { + let msg_id = Uuid::new_v4().to_string(); + current_text_message_id = Some(msg_id.clone()); + log::debug!("[bedrock] First text chunk, creating message msg_id={msg_id}"); + let add_msg = build_add_agent_output_message( + &task_id, + &msg_id, + text, + ); + yield Ok(add_msg); + } else { + let msg_id = current_text_message_id.as_ref().unwrap(); + let append = build_append_text( + &task_id, + msg_id, + text, + ); + yield Ok(append); + } + } + ContentBlockDelta::ReasoningContent(reasoning) => { + if let ReasoningContentBlockDelta::Text(text) = reasoning { + log::trace!("[bedrock] Reasoning delta ({} chars)", text.len()); + if current_text_message_id.is_none() { + let msg_id = Uuid::new_v4().to_string(); + current_text_message_id = Some(msg_id.clone()); + log::debug!("[bedrock] First reasoning chunk, creating message msg_id={msg_id}"); + let add_msg = build_add_agent_output_message( + &task_id, + &msg_id, + text, + ); + yield Ok(add_msg); + } else { + let msg_id = current_text_message_id.as_ref().unwrap(); + let append = build_append_text( + &task_id, + msg_id, + text, + ); + yield Ok(append); + } + } + } + ContentBlockDelta::ToolUse(tool_delta) => { + current_tool_input_json.push_str(tool_delta.input()); + } + _ => {} + } + } + } + StreamEvent::ContentBlockStop(_) => { + if !current_tool_use_id.is_empty() { + log::debug!("[bedrock] Tool call complete: {} ({})", current_tool_name, current_tool_use_id); + let tool_msg = build_tool_call_message( + &task_id, + ¤t_tool_use_id, + ¤t_tool_name, + ¤t_tool_input_json, + ); + yield Ok(tool_msg); + current_tool_use_id.clear(); + current_tool_name.clear(); + current_tool_input_json.clear(); + } + } + StreamEvent::MessageStop(stop) => { + stop_reason = match stop.stop_reason() { + StopReason::EndTurn => { + stream_finished::Reason::Done(stream_finished::Done {}) + } + StopReason::MaxTokens => { + stream_finished::Reason::MaxTokenLimit( + stream_finished::ReachedMaxTokenLimit {}, + ) + } + StopReason::ToolUse => { + stream_finished::Reason::Done(stream_finished::Done {}) + } + _ => stream_finished::Reason::Other(stream_finished::Other {}), + }; + } + StreamEvent::Metadata(metadata) => { + if let Some(usage) = metadata.usage() { + input_tokens = usage.input_tokens(); + output_tokens = usage.output_tokens(); + } + } + _ => {} + }, + Ok(None) => { + log::info!("[bedrock] Stream ended normally"); + break; + } + Err(e) => { + log::error!("[bedrock] Stream error: {e}"); + yield Err(Arc::new(AIApiError::Stream { + stream_type: "bedrock_converse", + source: anyhow::anyhow!("Bedrock stream error: {}", e), + })); + return; + } + } + } + + log::info!("[bedrock] Stream finished: input_tokens={input_tokens}, output_tokens={output_tokens}"); + let finished_event = build_stream_finished(stop_reason, input_tokens, output_tokens); + yield Ok(finished_event); + }; + + Box::pin(stream) +} + +pub(super) fn build_stream_init(request_id: &str, conversation_id: &str) -> ResponseEvent { + ResponseEvent { + r#type: Some(api::response_event::Type::Init( + api::response_event::StreamInit { + conversation_id: conversation_id.to_string(), + request_id: request_id.to_string(), + run_id: String::new(), + }, + )), + } +} + +pub(super) fn build_stream_finished( + reason: stream_finished::Reason, + input_tokens: i32, + output_tokens: i32, +) -> ResponseEvent { + let total_tokens = (input_tokens + output_tokens) as u32; + + let mut byok_token_usage = std::collections::HashMap::new(); + if total_tokens > 0 { + #[allow(deprecated)] + byok_token_usage.insert( + "bedrock".to_string(), + stream_finished::ModelTokenUsage { + model_id: String::new(), + total_tokens, + token_usage_by_category: std::collections::HashMap::new(), + }, + ); + } + + #[allow(deprecated)] + let conversation_usage_metadata = Some(stream_finished::ConversationUsageMetadata { + context_window_usage: 0.0, + summarized: false, + credits_spent: 0.0, + token_usage: vec![], + tool_usage_metadata: None, + warp_token_usage: std::collections::HashMap::new(), + byok_token_usage, + }); + + ResponseEvent { + r#type: Some(api::response_event::Type::Finished( + api::response_event::StreamFinished { + reason: Some(reason), + token_usage: vec![], + should_refresh_model_config: false, + request_cost: None, + conversation_usage_metadata, + }, + )), + } +} + +fn build_add_agent_output_message( + task_id: &str, + message_id: &str, + initial_text: &str, +) -> ResponseEvent { + let message = api::Message { + id: message_id.to_string(), + task_id: task_id.to_string(), + request_id: String::new(), + timestamp: None, + server_message_data: String::new(), + citations: vec![], + message: Some(api::message::Message::AgentOutput( + api::message::AgentOutput { + text: initial_text.to_string(), + }, + )), + }; + + let action = ClientAction { + action: Some(api::client_action::Action::AddMessagesToTask( + api::client_action::AddMessagesToTask { + task_id: task_id.to_string(), + messages: vec![message], + }, + )), + }; + + ResponseEvent { + r#type: Some(api::response_event::Type::ClientActions( + api::response_event::ClientActions { + actions: vec![action], + }, + )), + } +} + +fn build_append_text(task_id: &str, message_id: &str, text_delta: &str) -> ResponseEvent { + let message = api::Message { + id: message_id.to_string(), + task_id: task_id.to_string(), + request_id: String::new(), + timestamp: None, + server_message_data: String::new(), + citations: vec![], + message: Some(api::message::Message::AgentOutput( + api::message::AgentOutput { + text: text_delta.to_string(), + }, + )), + }; + + let mask = prost_types::FieldMask { + paths: vec!["message.agent_output.text".to_string()], + }; + + let action = ClientAction { + action: Some(api::client_action::Action::AppendToMessageContent( + api::client_action::AppendToMessageContent { + task_id: task_id.to_string(), + message: Some(message), + mask: Some(mask), + }, + )), + }; + + ResponseEvent { + r#type: Some(api::response_event::Type::ClientActions( + api::response_event::ClientActions { + actions: vec![action], + }, + )), + } +} + +fn build_tool_call_message( + task_id: &str, + tool_use_id: &str, + tool_name: &str, + tool_input_json: &str, +) -> ResponseEvent { + let _tool_use_id = tool_use_id.to_string(); + let _tool_name = tool_name.to_string(); + let _tool_input_json = tool_input_json.to_string(); + + let message = api::Message { + id: Uuid::new_v4().to_string(), + task_id: task_id.to_string(), + request_id: String::new(), + timestamp: None, + server_message_data: String::new(), + citations: vec![], + message: Some(api::message::Message::AgentOutput( + api::message::AgentOutput { + text: format!("[Tool call: {} ({})]", _tool_name, _tool_use_id), + }, + )), + }; + + let action = ClientAction { + action: Some(api::client_action::Action::AddMessagesToTask( + api::client_action::AddMessagesToTask { + task_id: task_id.to_string(), + messages: vec![message], + }, + )), + }; + + ResponseEvent { + r#type: Some(api::response_event::Type::ClientActions( + api::response_event::ClientActions { + actions: vec![action], + }, + )), + } +} + + diff --git a/app/src/ai/bedrock/stream_tests.rs b/app/src/ai/bedrock/stream_tests.rs new file mode 100644 index 00000000..d6a8cc37 --- /dev/null +++ b/app/src/ai/bedrock/stream_tests.rs @@ -0,0 +1,69 @@ +use warp_multi_agent_api::{self as api, response_event::stream_finished}; + +use super::stream::*; + +#[test] +fn test_build_stream_init_has_valid_ids() { + let event = build_stream_init("req-123", "conv-456"); + + match event.r#type { + Some(api::response_event::Type::Init(init)) => { + assert_eq!(init.request_id, "req-123"); + assert_eq!(init.conversation_id, "conv-456"); + assert_eq!(init.run_id, ""); + } + other => panic!("Expected Init event, got {:?}", other), + } +} + +#[test] +fn test_build_stream_finished_done_reason() { + let reason = stream_finished::Reason::Done(stream_finished::Done {}); + let event = build_stream_finished(reason, 100, 50); + + match event.r#type { + Some(api::response_event::Type::Finished(finished)) => { + assert!(matches!( + finished.reason, + Some(stream_finished::Reason::Done(_)) + )); + assert!(!finished.should_refresh_model_config); + let metadata = finished.conversation_usage_metadata.unwrap(); + assert_eq!(metadata.byok_token_usage.get("bedrock").unwrap().total_tokens, 150); + } + other => panic!("Expected Finished event, got {:?}", other), + } +} + +#[test] +fn test_build_stream_finished_max_token_limit() { + let reason = + stream_finished::Reason::MaxTokenLimit(stream_finished::ReachedMaxTokenLimit {}); + let event = build_stream_finished(reason, 200, 100); + + match event.r#type { + Some(api::response_event::Type::Finished(finished)) => { + assert!(matches!( + finished.reason, + Some(stream_finished::Reason::MaxTokenLimit(_)) + )); + } + other => panic!("Expected Finished event, got {:?}", other), + } +} + +#[test] +fn test_build_stream_finished_other_reason() { + let reason = stream_finished::Reason::Other(stream_finished::Other {}); + let event = build_stream_finished(reason, 0, 0); + + match event.r#type { + Some(api::response_event::Type::Finished(finished)) => { + assert!(matches!( + finished.reason, + Some(stream_finished::Reason::Other(_)) + )); + } + other => panic!("Expected Finished event, got {:?}", other), + } +} diff --git a/app/src/ai/blocklist/block.rs b/app/src/ai/blocklist/block.rs index 272050ab..ff49e1a5 100644 --- a/app/src/ai/blocklist/block.rs +++ b/app/src/ai/blocklist/block.rs @@ -3607,8 +3607,8 @@ impl AIBlock { } let ai_settings = AISettings::as_ref(ctx); - let login_command = ai_settings.aws_bedrock_auth_refresh_command.value().clone(); - let auto_login_enabled = *ai_settings.aws_bedrock_auto_login.value(); + let login_command = ai_settings.bedrock_auth_refresh_command.value().clone(); + let auto_login_enabled = *ai_settings.bedrock_auto_login.value(); // If auto-login is enabled, run the login command automatically if auto_login_enabled { @@ -5646,7 +5646,7 @@ pub enum AIBlockAction { ToggleReferencesSection, ToggleAutoexecuteReadonlyCommandsSpeedbumpCheckbox, ToggleAutoreadFilesSpeedbumpCheckbox, - ToggleAwsBedrockAutoLogin, + ToggleBedrockAutoLogin, ToggleCodebaseSearchSpeedbump(Option), StartNewConversationButtonClicked { action_id: AIAgentActionId, @@ -6246,11 +6246,11 @@ impl TypedActionView for AIBlock { AIBlockAction::RunAwsLoginCommand => { ctx.emit(AIBlockEvent::RunAwsLoginCommand); } - AIBlockAction::ToggleAwsBedrockAutoLogin => { + AIBlockAction::ToggleBedrockAutoLogin => { AISettings::handle(ctx).update(ctx, |settings, ctx| { - let current = *settings.aws_bedrock_auto_login.value(); + let current = *settings.bedrock_auto_login.value(); let new_value = !current; - report_if_error!(settings.aws_bedrock_auto_login.set_value(new_value, ctx)); + report_if_error!(settings.bedrock_auto_login.set_value(new_value, ctx)); }); } AIBlockAction::ConfigureAwsLoginCommand => { diff --git a/app/src/ai/blocklist/controller/response_stream.rs b/app/src/ai/blocklist/controller/response_stream.rs index 7b29a233..f427b1d9 100644 --- a/app/src/ai/blocklist/controller/response_stream.rs +++ b/app/src/ai/blocklist/controller/response_stream.rs @@ -8,15 +8,20 @@ use warp_multi_agent_api::response_event; use warpui::{Entity, ModelContext, SingletonEntity}; use crate::{ - ai::agent::{ - api::{self, generate_multi_agent_output, ConvertToAPITypeError}, - conversation::AIConversationId, - AIIdentifiers, CancellationReason, + ai::{ + agent::{ + api::{self, generate_multi_agent_output, ConvertToAPITypeError}, + conversation::AIConversationId, + AIIdentifiers, CancellationReason, + }, + bedrock::{client::BedrockClientConfig, models::is_bedrock_model}, }, network::NetworkStatus, report_error, send_telemetry_from_ctx, server::server_api::ServerApiProvider, + settings::ai::AISettings, }; +use settings::Setting; #[derive(Debug, Clone, PartialEq, Eq, Hash)] pub struct ResponseStreamId(String); @@ -79,6 +84,30 @@ pub struct ResponseStream { } impl ResponseStream { + fn bedrock_config_if_applicable( + model_id: &str, + ctx: &ModelContext, + ) -> Option { + let settings = AISettings::as_ref(ctx); + if !*settings.bedrock_enabled.value() { + return None; + } + let configured_models = settings.bedrock_models.value().clone(); + if !is_bedrock_model(model_id, &configured_models) { + return None; + } + let auth_method = *settings.bedrock_auth_method.value(); + Some(BedrockClientConfig { + auth_method, + profile: settings.bedrock_profile.value().clone(), + region: settings.bedrock_region.value().clone(), + access_key_id: settings.bedrock_access_key_id.value().clone(), + secret_access_key: settings.bedrock_secret_access_key.value().clone(), + cross_region_inference: *settings.bedrock_cross_region_inference.value(), + fallback_to_warp: *settings.bedrock_fallback_to_warp.value(), + }) + } + pub fn new( params: api::RequestParams, ai_identifiers: AIIdentifiers, @@ -90,11 +119,12 @@ impl ResponseStream { let start_time = Local::now(); let request_id = Uuid::new_v4(); + let bedrock_config = Self::bedrock_config_if_applicable(params.model.as_str(), ctx); let params_clone = params.clone(); let _ = ctx.spawn( async move { - generate_multi_agent_output(server_api, params_clone, cancellation_rx).await + generate_multi_agent_output(server_api, bedrock_config, params_clone, cancellation_rx).await }, move |me, stream, ctx| { me.handle_response_stream_result(request_id, stream, ctx); @@ -155,9 +185,10 @@ impl ResponseStream { let request_id = Uuid::new_v4(); self.current_request_id = Some(request_id); let params = self.params.clone(); + let bedrock_config = Self::bedrock_config_if_applicable(params.model.as_str(), ctx); let server_api = ServerApiProvider::as_ref(ctx).get(); let _ = ctx.spawn( - async move { generate_multi_agent_output(server_api, params, cancellation_rx).await }, + async move { generate_multi_agent_output(server_api, bedrock_config, params, cancellation_rx).await }, move |me, stream, ctx| { me.handle_response_stream_result(request_id, stream, ctx); }, diff --git a/app/src/ai/blocklist/inline_action/aws_bedrock_credentials_error.rs b/app/src/ai/blocklist/inline_action/aws_bedrock_credentials_error.rs index bcd7205c..33362b3a 100644 --- a/app/src/ai/blocklist/inline_action/aws_bedrock_credentials_error.rs +++ b/app/src/ai/blocklist/inline_action/aws_bedrock_credentials_error.rs @@ -78,7 +78,7 @@ impl AwsBedrockCredentialsErrorView { // Subscribe to AISettings changes to update checkbox state ctx.subscribe_to_model(&AISettings::handle(ctx), |_me, _, event, ctx| { - if matches!(event, AISettingsChangedEvent::AwsBedrockAutoLogin { .. }) { + if matches!(event, AISettingsChangedEvent::BedrockAutoLogin { .. }) { ctx.notify(); } }); @@ -125,7 +125,7 @@ impl View for AwsBedrockCredentialsErrorView { .finish(); } - let auto_login_enabled = *AISettings::as_ref(app).aws_bedrock_auto_login.value(); + let auto_login_enabled = *AISettings::as_ref(app).bedrock_auto_login.value(); // Helper closures to create elements (since Box can't be cloned) let make_alert_icon = || { @@ -280,8 +280,8 @@ impl TypedActionView for AwsBedrockCredentialsErrorView { } AwsBedrockCredentialsErrorAction::ToggleAutoLogin => { AISettings::handle(ctx).update(ctx, |settings, ctx| { - let current = *settings.aws_bedrock_auto_login.value(); - report_if_error!(settings.aws_bedrock_auto_login.set_value(!current, ctx)); + let current = *settings.bedrock_auto_login.value(); + report_if_error!(settings.bedrock_auto_login.set_value(!current, ctx)); }); ctx.notify(); } diff --git a/app/src/ai/blocklist/passive_suggestions/maa.rs b/app/src/ai/blocklist/passive_suggestions/maa.rs index 4e402f21..2a8d3a59 100644 --- a/app/src/ai/blocklist/passive_suggestions/maa.rs +++ b/app/src/ai/blocklist/passive_suggestions/maa.rs @@ -174,7 +174,7 @@ impl PassiveSuggestionsModel { let stream_handle = ctx.spawn( async move { let stream_result = - generate_multi_agent_output(server_api, request_params, cancellation_rx).await; + generate_multi_agent_output(server_api, None, request_params, cancellation_rx).await; extract_suggestion_from_stream(stream_result).await }, move |me, result, ctx| { diff --git a/app/src/ai/execution_profiles/model_menu_items.rs b/app/src/ai/execution_profiles/model_menu_items.rs index 71e730dd..c187bd02 100644 --- a/app/src/ai/execution_profiles/model_menu_items.rs +++ b/app/src/ai/execution_profiles/model_menu_items.rs @@ -1,4 +1,4 @@ -use crate::ai::llms::{is_using_api_key_for_provider, DisableReason, LLMId, LLMInfo}; +use crate::ai::llms::{is_using_api_key_for_provider, DisableReason, LLMId, LLMInfo, LLMProvider}; use crate::menu::{MenuItem, MenuItemFields, MenuTooltipPosition}; use itertools::Itertools; use std::sync::Arc; @@ -80,6 +80,7 @@ fn make_item_fields( llm.menu_display_name() }; let is_using_api_key = is_using_api_key_for_provider(&llm.provider, app); + let is_bedrock = llm.provider == LLMProvider::Bedrock; let mut item = if let Some(position_id_fn) = position_id_fn { let position_id = position_id_fn(&llm.id); @@ -89,7 +90,11 @@ fn make_item_fields( Flex::row().with_cross_axis_alignment(CrossAxisAlignment::Center); let icon_container = Container::new( - ConstrainedBox::new(if is_using_api_key { + ConstrainedBox::new(if is_bedrock { + Icon::BedrockLogo + .to_warpui_icon(appearance.theme().foreground()) + .finish() + } else if is_using_api_key { Icon::Key .to_warpui_icon(appearance.theme().foreground()) .finish() diff --git a/app/src/ai/llms.rs b/app/src/ai/llms.rs index d2c22096..51c7d691 100644 --- a/app/src/ai/llms.rs +++ b/app/src/ai/llms.rs @@ -16,11 +16,17 @@ use crate::{ network::{NetworkStatus, NetworkStatusEvent, NetworkStatusKind}, report_error, server::server_api::ServerApiProvider, + settings::ai::{AISettings, AISettingsChangedEvent, BedrockModelConfig}, workspaces::user_workspaces::{UserWorkspaces, UserWorkspacesEvent}, }; +use settings::Setting; + use super::execution_profiles::profiles::AIExecutionProfilesModel; +#[cfg(not(target_family = "wasm"))] +use super::bedrock::models::get_effective_models; + pub use ai::LLMId; /// Checks if a user's' API key is being used for the given provider. @@ -36,6 +42,7 @@ pub fn is_using_api_key_for_provider(provider: &LLMProvider, app: &AppContext) - LLMProvider::OpenAI => api_keys.is_some_and(|keys| keys.openai.is_some()), LLMProvider::Anthropic => api_keys.is_some_and(|keys| keys.anthropic.is_some()), LLMProvider::Google => api_keys.is_some_and(|keys| keys.google.is_some()), + LLMProvider::Bedrock => true, _ => false, } } @@ -89,6 +96,7 @@ pub enum LLMProvider { Anthropic, Google, Xai, + Bedrock, Unknown, } @@ -99,6 +107,7 @@ impl LLMProvider { LLMProvider::OpenAI => Some(Icon::OpenAILogo), LLMProvider::Anthropic => Some(Icon::ClaudeLogo), LLMProvider::Google => Some(Icon::GeminiLogo), + LLMProvider::Bedrock => Some(Icon::BedrockLogo), LLMProvider::Xai => None, LLMProvider::Unknown => None, } @@ -498,11 +507,9 @@ struct AvailableLLMsUpdate { pub struct LLMPreferences { models_by_feature: ModelsByFeature, last_update: Option, - // Stores temporary model overrides for a given terminal view. - // NOTE: We only store an override if the model selected by the user is different - // from the base LLM for the active profile. This means that if the user selects the - // profile's default model and changes their profile, the model will update to that profile's default. base_llm_for_terminal_view: HashMap, + #[cfg(not(target_family = "wasm"))] + bedrock_models_fetched: bool, } impl LLMPreferences { @@ -534,12 +541,34 @@ impl LLMPreferences { } }); + #[cfg(not(target_family = "wasm"))] + ctx.subscribe_to_model(&AISettings::handle(ctx), |me, event, ctx| { + if matches!( + event, + AISettingsChangedEvent::BedrockEnabled { .. } + | AISettingsChangedEvent::BedrockModels { .. } + | AISettingsChangedEvent::BedrockCrossRegionInference { .. } + | AISettingsChangedEvent::BedrockRegion { .. } + ) { + if matches!(event, AISettingsChangedEvent::BedrockEnabled { .. }) { + let enabled = *AISettings::as_ref(ctx).bedrock_enabled.value(); + if enabled && !me.bedrock_models_fetched { + me.trigger_bedrock_discovery(ctx); + } + } + me.inject_bedrock_models(ctx); + ctx.emit(LLMPreferencesEvent::UpdatedAvailableLLMs); + } + }); + let base_llm_for_terminal_view = HashMap::new(); - let me = Self { + let mut me = Self { models_by_feature, last_update: None, base_llm_for_terminal_view, + #[cfg(not(target_family = "wasm"))] + bedrock_models_fetched: false, }; // In agent mode eval builds, eagerly kick off a fetch of the model list from the server @@ -549,9 +578,103 @@ impl LLMPreferences { #[cfg(feature = "agent_mode_evals")] me.refresh_available_models(ctx); + #[cfg(not(target_family = "wasm"))] + { + me.inject_bedrock_models(ctx); + if *AISettings::as_ref(ctx).bedrock_enabled.value() { + me.trigger_bedrock_discovery(ctx); + } + } + me } + #[cfg(not(target_family = "wasm"))] + fn trigger_bedrock_discovery(&mut self, ctx: &mut ModelContext) { + use crate::ai::bedrock::client::BedrockClientConfig; + use crate::ai::bedrock::discovery::discover_inference_profiles; + + self.bedrock_models_fetched = true; + + let settings = AISettings::as_ref(ctx); + let config = BedrockClientConfig { + auth_method: settings.bedrock_auth_method.value().clone(), + profile: settings.bedrock_profile.value().clone(), + region: settings.bedrock_region.value().clone(), + access_key_id: settings.bedrock_access_key_id.value().clone(), + secret_access_key: settings.bedrock_secret_access_key.value().clone(), + cross_region_inference: *settings.bedrock_cross_region_inference.value(), + fallback_to_warp: *settings.bedrock_fallback_to_warp.value(), + }; + + ctx.spawn( + async move { discover_inference_profiles(&config).await }, + |me, result, ctx| match result { + Ok(models) => { + AISettings::handle(ctx).update(ctx, |settings, ctx| { + let _ = settings.bedrock_models.set_value(models, ctx); + }); + me.inject_bedrock_models(ctx); + ctx.emit(LLMPreferencesEvent::UpdatedAvailableLLMs); + } + Err(e) => { + log::error!("Failed to discover Bedrock inference profiles: {e}"); + } + }, + ); + } + + #[cfg(not(target_family = "wasm"))] + fn inject_bedrock_models(&mut self, ctx: &AppContext) { + self.models_by_feature + .agent_mode + .choices + .retain(|m| m.provider != LLMProvider::Bedrock); + + let settings = AISettings::as_ref(ctx); + if !*settings.bedrock_enabled.value() { + return; + } + + let user_models: Vec = settings.bedrock_models.value().clone(); + let region = settings.bedrock_region.value().clone(); + let cross_region = *settings.bedrock_cross_region_inference.value(); + + let effective = get_effective_models(&user_models); + for model in effective { + let model_id = if cross_region && !region.is_empty() { + super::bedrock::models::apply_cross_region_prefix(&model.model_id, ®ion) + } else { + model.model_id.clone() + }; + + let llm_info = LLMInfo { + id: LLMId::from(model_id.as_str()), + display_name: model.display_name.clone(), + base_model_name: model.display_name.clone(), + reasoning_level: None, + usage_metadata: LLMUsageMetadata { + request_multiplier: 1, + credit_multiplier: None, + }, + description: Some("AWS Bedrock".to_string()), + disable_reason: None, + vision_supported: model.vision_supported, + spec: None, + provider: LLMProvider::Bedrock, + host_configs: HashMap::from([( + LLMModelHost::AwsBedrock, + RoutingHostConfig { + enabled: true, + model_routing_host: LLMModelHost::AwsBedrock, + }, + )]), + discount_percentage: None, + }; + self.models_by_feature.agent_mode.choices.push(llm_info); + } + } + /// Returns the `LLMInfo` for the base LLM to be used for an Agent Mode request. pub fn get_active_base_model<'a>( &'a self, @@ -924,6 +1047,9 @@ impl LLMPreferences { } } + #[cfg(not(target_family = "wasm"))] + self.inject_bedrock_models(ctx); + // Clear any model selections where the model is no longer supported. let profiles_model = AIExecutionProfilesModel::handle(ctx); profiles_model.update(ctx, |profiles, ctx| { diff --git a/app/src/ai/mod.rs b/app/src/ai/mod.rs index 5f1cb351..2d653639 100644 --- a/app/src/ai/mod.rs +++ b/app/src/ai/mod.rs @@ -15,6 +15,9 @@ pub mod artifacts; pub(crate) mod attachment_utils; #[cfg(not(target_family = "wasm"))] pub mod aws_credentials; +#[cfg(not(target_family = "wasm"))] +#[allow(dead_code)] +pub mod bedrock; pub(crate) mod block_context; pub(crate) mod blocklist; pub mod control_code_parser; diff --git a/app/src/settings/ai.rs b/app/src/settings/ai.rs index 5b35893b..88073c73 100644 --- a/app/src/settings/ai.rs +++ b/app/src/settings/ai.rs @@ -394,6 +394,65 @@ impl ThinkingDisplayMode { } } +/// Authentication method for AWS Bedrock. +#[derive( + Default, + Debug, + serde::Serialize, + serde::Deserialize, + PartialEq, + Copy, + Clone, + EnumIter, + schemars::JsonSchema, + settings_value::SettingsValue, +)] +#[schemars( + description = "Authentication method for AWS Bedrock.", + rename_all = "snake_case" +)] +pub enum BedrockAuthMethod { + #[default] + Profile, + StaticKeys, + Sso, +} + +settings::macros::implement_setting_for_enum!( + BedrockAuthMethod, + AISettings, + SupportedPlatforms::DESKTOP, + SyncToCloud::Globally(RespectUserSyncSetting::Yes), + private: false, + toml_path: "ai.bedrock.auth_method", + description: "Authentication method for AWS Bedrock.", +); + +impl BedrockAuthMethod { + pub fn display_name(&self) -> &'static str { + match self { + BedrockAuthMethod::Profile => "AWS Profile", + BedrockAuthMethod::StaticKeys => "Static Keys", + BedrockAuthMethod::Sso => "SSO", + } + } +} + +/// Configuration for a single Bedrock model. +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, schemars::JsonSchema)] +#[schemars(description = "Configuration for a single AWS Bedrock model.")] +pub struct BedrockModelConfig { + #[schemars(description = "The Bedrock model ID (e.g. anthropic.claude-sonnet-4-20250514-v1:0).")] + pub model_id: String, + #[schemars(description = "Display name shown in the model picker.")] + pub display_name: String, + #[serde(default)] + #[schemars(description = "Whether the model supports image/vision input.")] + pub vision_supported: bool, +} + +impl settings_value::SettingsValue for BedrockModelConfig {} + /// Tracks the state of the quota reset banner #[derive( Debug, @@ -1007,53 +1066,106 @@ define_settings_group!(AISettings, settings: [ sync_to_cloud: SyncToCloud::Globally(RespectUserSyncSetting::Yes), private: true, } - // Whether to use locally loaded AWS credentials for Bedrock-enabled requests. - aws_bedrock_credentials_enabled: AwsBedrockCredentialsEnabled { + // Whether direct Bedrock integration is enabled (client calls Bedrock API directly). + bedrock_enabled: BedrockEnabled { type: bool, default: false, supported_platforms: SupportedPlatforms::DESKTOP, sync_to_cloud: SyncToCloud::Globally(RespectUserSyncSetting::Yes), private: false, - toml_path: "cloud_platform.third_party_api_keys.aws_bedrock_credentials_enabled", - description: "Whether Warp should use your local AWS credentials for Bedrock-enabled requests.", + toml_path: "ai.bedrock.enabled", + description: "Whether to use AWS Bedrock directly for AI requests.", } - // Whether to automatically run the AWS login command when Bedrock credentials are expired. - // - // When true, the configured login command will be run automatically without asking. - // When false (default), a prompt will be shown asking for permission. - aws_bedrock_auto_login: AwsBedrockAutoLogin { - type: bool, - default: false, - supported_platforms: SupportedPlatforms::DESKTOP, - sync_to_cloud: SyncToCloud::Globally(RespectUserSyncSetting::Yes), - private: false, - toml_path: "cloud_platform.third_party_api_keys.aws_bedrock_auto_login", - description: "Whether to automatically run the AWS login command when Bedrock credentials expire.", - } - // Command to run to refresh AWS credentials when using Bedrock auto-login. - aws_bedrock_auth_refresh_command: AwsBedrockAuthRefreshCommand { - type: String, - default: "aws login".to_string(), - supported_platforms: SupportedPlatforms::DESKTOP, - sync_to_cloud: SyncToCloud::Globally(RespectUserSyncSetting::Yes), - private: false, - toml_path: "cloud_platform.third_party_api_keys.aws_bedrock_auth_refresh_command", - description: "The command to run to refresh AWS credentials for Bedrock.", - } - // AWS profile name to use when loading credentials from the local AWS credential/config chain. - aws_bedrock_profile: AwsBedrockProfile { + // Authentication method for Bedrock: "profile", "static_keys", or "sso". + bedrock_auth_method: BedrockAuthMethod, + // AWS profile name to use when auth_method is Profile or SSO. + bedrock_profile: BedrockProfile { type: String, default: "default".to_string(), supported_platforms: SupportedPlatforms::DESKTOP, sync_to_cloud: SyncToCloud::Globally(RespectUserSyncSetting::Yes), private: false, - toml_path: "cloud_platform.third_party_api_keys.aws_bedrock_profile", + toml_path: "ai.bedrock.profile", description: "The AWS profile name to use for Bedrock credentials.", } - // Whether the AWS Bedrock login banner has been permanently dismissed. - // - // Not a user-visible setting - we model it as a setting so we can track state. - aws_bedrock_login_banner_dismissed: AwsBedrockLoginBannerDismissed { + // AWS region for Bedrock API calls. Empty string means auto-detect from profile/config. + bedrock_region: BedrockRegion { + type: String, + default: String::new(), + supported_platforms: SupportedPlatforms::DESKTOP, + sync_to_cloud: SyncToCloud::Globally(RespectUserSyncSetting::Yes), + private: false, + toml_path: "ai.bedrock.region", + description: "AWS region for Bedrock API calls. Leave empty to auto-detect from profile.", + } + // Whether to automatically add cross-region inference prefixes to model IDs. + bedrock_cross_region_inference: BedrockCrossRegionInference { + type: bool, + default: true, + supported_platforms: SupportedPlatforms::DESKTOP, + sync_to_cloud: SyncToCloud::Globally(RespectUserSyncSetting::Yes), + private: false, + toml_path: "ai.bedrock.cross_region_inference", + description: "Whether to automatically add cross-region inference prefixes to model IDs.", + } + // Whether to fall back to routing through the Warp server when Bedrock credentials fail. + bedrock_fallback_to_warp: BedrockFallbackToWarp { + type: bool, + default: true, + supported_platforms: SupportedPlatforms::DESKTOP, + sync_to_cloud: SyncToCloud::Globally(RespectUserSyncSetting::Yes), + private: false, + toml_path: "ai.bedrock.fallback_to_warp", + description: "Whether to fall back to the Warp server when Bedrock credentials are invalid.", + } + // Custom Bedrock model configurations. + bedrock_models: BedrockModels { + type: Vec, + default: Vec::new(), + supported_platforms: SupportedPlatforms::DESKTOP, + sync_to_cloud: SyncToCloud::Globally(RespectUserSyncSetting::Yes), + private: false, + toml_path: "ai.bedrock.models", + description: "Custom AWS Bedrock model configurations.", + } + // Whether to automatically run the login command when Bedrock credentials expire. + bedrock_auto_login: BedrockAutoLogin { + type: bool, + default: false, + supported_platforms: SupportedPlatforms::DESKTOP, + sync_to_cloud: SyncToCloud::Globally(RespectUserSyncSetting::Yes), + private: false, + toml_path: "ai.bedrock.auto_login", + description: "Whether to automatically run the login command when Bedrock credentials expire.", + } + // Command to run to refresh AWS credentials (e.g. "aws sso login"). + bedrock_auth_refresh_command: BedrockAuthRefreshCommand { + type: String, + default: "aws sso login".to_string(), + supported_platforms: SupportedPlatforms::DESKTOP, + sync_to_cloud: SyncToCloud::Globally(RespectUserSyncSetting::Yes), + private: false, + toml_path: "ai.bedrock.auth_refresh_command", + description: "The command to run to refresh AWS credentials for Bedrock.", + } + // AWS access key ID for static key authentication (stored in OS keychain). + bedrock_access_key_id: BedrockAccessKeyId { + type: String, + default: String::new(), + supported_platforms: SupportedPlatforms::DESKTOP, + sync_to_cloud: SyncToCloud::Never, + private: true, + } + // AWS secret access key for static key authentication (stored in OS keychain). + bedrock_secret_access_key: BedrockSecretAccessKey { + type: String, + default: String::new(), + supported_platforms: SupportedPlatforms::DESKTOP, + sync_to_cloud: SyncToCloud::Never, + private: true, + } + // Whether the Bedrock login banner has been permanently dismissed. + bedrock_login_banner_dismissed: BedrockLoginBannerDismissed { type: bool, default: false, supported_platforms: SupportedPlatforms::DESKTOP, diff --git a/app/src/settings_view/ai_page.rs b/app/src/settings_view/ai_page.rs index ddb118e1..a8c18634 100644 --- a/app/src/settings_view/ai_page.rs +++ b/app/src/settings_view/ai_page.rs @@ -1,5 +1,3 @@ -#[cfg(not(target_family = "wasm"))] -use crate::ai::aws_credentials::refresh_aws_credentials; use crate::ai::blocklist::agent_view::agent_input_footer::editor::{ AgentToolbarEditorMode, AgentToolbarInlineEditor, }; @@ -25,8 +23,9 @@ use crate::settings::InputSettings; use crate::settings::{ AIAutoDetectionEnabled, AICommandDenylist, AISettingsChangedEvent, AgentModeCodingPermissionsType, AgentModeCommandExecutionDenylist, - AgentModeCommandExecutionPredicate, AgentModeQuerySuggestionsEnabled, AwsBedrockAutoLogin, - AwsBedrockCredentialsEnabled, CanUseWarpCreditsWithByok, CodeSettings, CodebaseContextEnabled, + AgentModeCommandExecutionPredicate, AgentModeQuerySuggestionsEnabled, BedrockAutoLogin, + BedrockAuthMethod, BedrockCrossRegionInference, BedrockEnabled, BedrockFallbackToWarp, + CanUseWarpCreditsWithByok, CodeSettings, CodebaseContextEnabled, FileBasedMcpEnabled, GitOperationsAutogenEnabled, IncludeAgentCommandsInHistory, IntelligentAutosuggestionsEnabled, MemoryEnabled, NLDInTerminalEnabled, NaturalLanguageAutosuggestionsEnabled, OrchestrationEnabled, RuleSuggestionsEnabled, @@ -53,7 +52,7 @@ use warp_core::context_flag::ContextFlag; use warp_core::features::FeatureFlag; use warp_core::ui::theme::color::internal_colors; use warpui::elements::{ - Border, ChildView, ConstrainedBox, CornerRadius, CrossAxisAlignment, Expanded, Fill, + ChildView, ConstrainedBox, CornerRadius, CrossAxisAlignment, Fill, HyperlinkLens, MainAxisAlignment, MainAxisSize, MouseStateHandle, Radius, Shrinkable, Text, }; use warpui::fonts::{Properties, Weight}; @@ -99,6 +98,8 @@ pub enum AISubpage { Knowledge, /// Third-party CLI agent settings. ThirdPartyCLIAgents, + /// AWS Bedrock direct provider configuration. + Bedrock, } impl AISubpage { @@ -108,6 +109,7 @@ impl AISubpage { SettingsSection::AgentProfiles => Some(Self::Profiles), SettingsSection::Knowledge => Some(Self::Knowledge), SettingsSection::ThirdPartyCLIAgents => Some(Self::ThirdPartyCLIAgents), + SettingsSection::Bedrock => Some(Self::Bedrock), // AgentMCPServers renders the standalone MCPServers page, not an AI subpage. _ => None, } @@ -1468,7 +1470,6 @@ impl AISettingsPageView { } widgets.push(Box::new(CLIAgentWidget::default())); widgets.push(Box::new(ApiKeysWidget::new(ctx))); - widgets.push(Box::new(AwsBedrockWidget::new(ctx))); widgets.push(Box::new(AgentAttributionWidget::default())); widgets.push(Box::new(OtherAIWidget::default())); if FeatureFlag::AgentModeComputerUse.is_enabled() { @@ -1508,7 +1509,6 @@ impl AISettingsPageView { widgets.push(Box::new(VoiceWidget::default())); } widgets.push(Box::new(ApiKeysWidget::new(ctx))); - widgets.push(Box::new(AwsBedrockWidget::new(ctx))); widgets.push(Box::new(AgentAttributionWidget::default())); widgets.push(Box::new(OtherAIWidget::default())); if FeatureFlag::AgentModeComputerUse.is_enabled() { @@ -1529,6 +1529,9 @@ impl AISettingsPageView { Some(AISubpage::ThirdPartyCLIAgents) => { widgets.push(Box::new(CLIAgentWidget::default())); } + Some(AISubpage::Bedrock) => { + widgets.push(Box::new(BedrockSettingsWidget::new(ctx))); + } } // Subpage widgets render their own subheader-sized titles internally, @@ -2107,9 +2110,13 @@ pub enum AISettingsPageAction { RemoveFromMCPDenylist(uuid::Uuid), CreateProfile, SignupAnonymousUser, - ToggleAwsBedrockAutoLogin, - ToggleAwsBedrockCredentialsEnabled, - RefreshAwsBedrockCredentials, + ToggleBedrockAutoLogin, + ToggleBedrockEnabled, + RefreshAwsBedrock, + SetBedrockAuthMethod(BedrockAuthMethod), + SetBedrockProfile(String), + ToggleBedrockCrossRegionInference, + ToggleBedrockFallbackToWarp, ToggleCloudAgentComputerUse, ToggleFileBasedMcp, ToggleIncludeAgentCommandsInHistory, @@ -2753,24 +2760,81 @@ impl TypedActionView for AISettingsPageView { AISettingsPageAction::SignupAnonymousUser => { ctx.emit(AISettingsPageEvent::SignupAnonymousUser); } - AISettingsPageAction::ToggleAwsBedrockAutoLogin => { + AISettingsPageAction::ToggleBedrockAutoLogin => { AISettings::handle(ctx).update(ctx, |settings, ctx| { - report_if_error!(settings.aws_bedrock_auto_login.toggle_and_save_value(ctx)); + report_if_error!(settings.bedrock_auto_login.toggle_and_save_value(ctx)); }); ctx.notify(); } - AISettingsPageAction::ToggleAwsBedrockCredentialsEnabled => { + AISettingsPageAction::ToggleBedrockEnabled => { AISettings::handle(ctx).update(ctx, |settings, ctx| { report_if_error!(settings - .aws_bedrock_credentials_enabled + .bedrock_enabled .toggle_and_save_value(ctx)); }); ctx.notify(); } - AISettingsPageAction::RefreshAwsBedrockCredentials => { + AISettingsPageAction::RefreshAwsBedrock => { #[cfg(not(target_family = "wasm"))] - ApiKeyManager::handle(ctx).update(ctx, |manager, ctx| { - drop(refresh_aws_credentials(manager, ctx)); + { + use crate::ai::bedrock::client::BedrockClientConfig; + use crate::ai::bedrock::discovery::discover_inference_profiles; + use settings::Setting; + + let ai_settings = AISettings::as_ref(ctx); + let config = BedrockClientConfig { + auth_method: ai_settings.bedrock_auth_method.value().clone(), + profile: ai_settings.bedrock_profile.value().clone(), + region: ai_settings.bedrock_region.value().clone(), + access_key_id: ai_settings.bedrock_access_key_id.value().clone(), + secret_access_key: ai_settings.bedrock_secret_access_key.value().clone(), + cross_region_inference: *ai_settings.bedrock_cross_region_inference.value(), + fallback_to_warp: *ai_settings.bedrock_fallback_to_warp.value(), + }; + + ctx.spawn( + async move { discover_inference_profiles(&config).await }, + |_me, result, ctx| match result { + Ok(models) => { + AISettings::handle(ctx).update(ctx, |settings, ctx| { + let _ = settings.bedrock_models.set_value(models, ctx); + }); + } + Err(e) => { + log::error!( + "Failed to discover Bedrock inference profiles: {e}" + ); + } + }, + ); + } + ctx.notify(); + } + AISettingsPageAction::SetBedrockAuthMethod(method) => { + AISettings::handle(ctx).update(ctx, |settings, ctx| { + report_if_error!(settings.bedrock_auth_method.set_value(*method, ctx)); + }); + ctx.notify(); + } + AISettingsPageAction::SetBedrockProfile(profile) => { + AISettings::handle(ctx).update(ctx, |settings, ctx| { + report_if_error!(settings.bedrock_profile.set_value(profile.clone(), ctx)); + }); + ctx.notify(); + } + AISettingsPageAction::ToggleBedrockCrossRegionInference => { + AISettings::handle(ctx).update(ctx, |settings, ctx| { + report_if_error!(settings + .bedrock_cross_region_inference + .toggle_and_save_value(ctx)); + }); + ctx.notify(); + } + AISettingsPageAction::ToggleBedrockFallbackToWarp => { + AISettings::handle(ctx).update(ctx, |settings, ctx| { + report_if_error!(settings + .bedrock_fallback_to_warp + .toggle_and_save_value(ctx)); }); ctx.notify(); } @@ -6301,25 +6365,85 @@ impl SettingsWidget for ApiKeysWidget { } } -struct AwsBedrockWidget { - aws_auth_refresh_command_editor: ViewHandle, - aws_auth_refresh_profile_editor: ViewHandle, - credentials_enabled_toggle: SwitchStateHandle, +struct BedrockSettingsWidget { + enabled_toggle: SwitchStateHandle, + cross_region_toggle: SwitchStateHandle, + fallback_toggle: SwitchStateHandle, auto_login_toggle: SwitchStateHandle, - refresh_credentials_button: ViewHandle, + auth_method_dropdown: ViewHandle>, + profile_dropdown: ViewHandle>, + region_editor: ViewHandle, + auth_refresh_command_editor: ViewHandle, + access_key_editor: ViewHandle, + secret_key_editor: ViewHandle, + refresh_button: ViewHandle, } -impl AwsBedrockWidget { +impl BedrockSettingsWidget { fn new(ctx: &mut ViewContext<::View>) -> Self { let ai_settings = AISettings::as_ref(ctx); - let is_any_ai_enabled = ai_settings.is_any_ai_enabled(ctx); + let is_enabled = *ai_settings.bedrock_enabled.value(); - let aws_auth_refresh_command = ai_settings.aws_bedrock_auth_refresh_command.value().clone(); - let aws_auth_refresh_profile = ai_settings.aws_bedrock_profile.value().clone(); - let is_usage_enabled = is_any_ai_enabled - && UserWorkspaces::as_ref(ctx).is_aws_bedrock_credentials_enabled(ctx); + let region_val = ai_settings.bedrock_region.value().clone(); + let auth_cmd_val = ai_settings.bedrock_auth_refresh_command.value().clone(); + let access_key_val = ai_settings.bedrock_access_key_id.value().clone(); + let secret_key_val = ai_settings.bedrock_secret_access_key.value().clone(); - let aws_auth_refresh_command_editor = ctx.add_typed_action_view(move |ctx| { + let auth_method_dropdown = ctx.add_typed_action_view(|ctx| { + let mut dropdown = Dropdown::new(ctx); + let methods = [ + BedrockAuthMethod::Profile, + BedrockAuthMethod::StaticKeys, + BedrockAuthMethod::Sso, + ]; + let current = AISettings::as_ref(ctx).bedrock_auth_method.value().clone(); + let selected_index = methods + .iter() + .position(|m| *m == current) + .unwrap_or(0); + dropdown.add_items( + methods + .into_iter() + .map(|m| { + DropdownItem::new( + m.display_name(), + AISettingsPageAction::SetBedrockAuthMethod(m), + ) + }) + .collect(), + ctx, + ); + dropdown.set_selected_by_index(selected_index, ctx); + dropdown + }); + + let profile_dropdown = ctx.add_typed_action_view(|ctx| { + use crate::ai::bedrock::discovery::list_aws_profiles; + + let mut dropdown = Dropdown::new(ctx); + let profiles = list_aws_profiles(); + let current_profile = AISettings::as_ref(ctx).bedrock_profile.value().clone(); + let items: Vec<_> = profiles + .iter() + .map(|p| { + DropdownItem::new( + p.as_str(), + AISettingsPageAction::SetBedrockProfile(p.clone()), + ) + }) + .collect(); + let selected_index = profiles + .iter() + .position(|p| *p == current_profile) + .unwrap_or(0); + dropdown.add_items(items, ctx); + if !profiles.is_empty() { + dropdown.set_selected_by_index(selected_index, ctx); + } + dropdown + }); + + let region_editor = ctx.add_typed_action_view(move |ctx| { let appearance = Appearance::as_ref(ctx); let options = SingleLineEditorOptions { is_password: false, @@ -6336,38 +6460,20 @@ impl AwsBedrockWidget { ..Default::default() }; let mut editor = EditorView::single_line(options, ctx); - editor.set_placeholder_text("aws login", ctx); - editor.set_buffer_text(&aws_auth_refresh_command, ctx); + editor.set_placeholder_text("auto-detect from profile", ctx); + editor.set_buffer_text(®ion_val, ctx); editor }); - AISettingsPageView::update_editor_interaction_state( - aws_auth_refresh_command_editor.clone(), - is_usage_enabled, - ctx, - ); - ctx.subscribe_to_view(&aws_auth_refresh_command_editor, |_, editor, event, ctx| { + ctx.subscribe_to_view(®ion_editor, |_, editor, event, ctx| { if matches!(event, EditorEvent::Blurred | EditorEvent::Enter) { - let buffer_text = editor.as_ref(ctx).buffer_text(ctx); - let should_reset = buffer_text.trim().is_empty(); - let value = if should_reset { - "aws login".to_string() - } else { - buffer_text - }; + let value = editor.as_ref(ctx).buffer_text(ctx); AISettings::handle(ctx).update(ctx, |settings, ctx| { - let _ = settings - .aws_bedrock_auth_refresh_command - .set_value(value, ctx); + let _ = settings.bedrock_region.set_value(value, ctx); }); - if should_reset { - editor.update(ctx, |editor, ctx| { - editor.set_buffer_text("aws login", ctx); - }); - } } }); - let aws_auth_refresh_profile_editor = ctx.add_typed_action_view(move |ctx| { + let auth_refresh_command_editor = ctx.add_typed_action_view(move |ctx| { let appearance = Appearance::as_ref(ctx); let options = SingleLineEditorOptions { is_password: false, @@ -6384,330 +6490,201 @@ impl AwsBedrockWidget { ..Default::default() }; let mut editor = EditorView::single_line(options, ctx); - editor.set_placeholder_text("default", ctx); - editor.set_buffer_text(&aws_auth_refresh_profile, ctx); + editor.set_placeholder_text("aws sso login", ctx); + editor.set_buffer_text(&auth_cmd_val, ctx); editor }); - AISettingsPageView::update_editor_interaction_state( - aws_auth_refresh_profile_editor.clone(), - is_usage_enabled, - ctx, - ); - ctx.subscribe_to_view(&aws_auth_refresh_profile_editor, |_, editor, event, ctx| { + ctx.subscribe_to_view(&auth_refresh_command_editor, |_, editor, event, ctx| { if matches!(event, EditorEvent::Blurred | EditorEvent::Enter) { let buffer_text = editor.as_ref(ctx).buffer_text(ctx); - let should_reset = buffer_text.trim().is_empty(); - let value = if should_reset { - "default".to_string() + let value = if buffer_text.trim().is_empty() { + "aws sso login".to_string() } else { buffer_text }; AISettings::handle(ctx).update(ctx, |settings, ctx| { - let _ = settings.aws_bedrock_profile.set_value(value, ctx); + let _ = settings.bedrock_auth_refresh_command.set_value(value, ctx); }); - if should_reset { - editor.update(ctx, |editor, ctx| { - editor.set_buffer_text("default", ctx); - }); - } } }); - let refresh_credentials_button = ctx.add_typed_action_view(|_| { - ActionButton::new("Refresh", SecondaryTheme) + let access_key_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("AKIA...", ctx); + editor.set_buffer_text(&access_key_val, ctx); + editor + }); + ctx.subscribe_to_view(&access_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.bedrock_access_key_id.set_value(value, ctx); + }); + } + }); + + let secret_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("wJalr...", ctx); + editor.set_buffer_text(&secret_key_val, ctx); + editor + }); + ctx.subscribe_to_view(&secret_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.bedrock_secret_access_key.set_value(value, ctx); + }); + } + }); + + let refresh_button = ctx.add_typed_action_view(|_| { + ActionButton::new("Refresh AWS Bedrock", SecondaryTheme) .with_icon(Icon::RefreshCw04) .with_size(ButtonSize::Small) .on_click(|ctx| { - ctx.dispatch_typed_action(AISettingsPageAction::RefreshAwsBedrockCredentials); + ctx.dispatch_typed_action(AISettingsPageAction::RefreshAwsBedrock); }) }); - refresh_credentials_button.update(ctx, |button, ctx| { - button.set_disabled(!is_usage_enabled, ctx); + refresh_button.update(ctx, |button, ctx| { + button.set_disabled(!is_enabled, ctx); }); - // Keep enablement in sync with the Global AI toggle. - let aws_auth_refresh_command_editor_clone = aws_auth_refresh_command_editor.clone(); - let aws_auth_refresh_profile_editor_clone = aws_auth_refresh_profile_editor.clone(); - let refresh_credentials_button_clone = refresh_credentials_button.clone(); + let profile_dropdown_clone = profile_dropdown.clone(); + let region_editor_clone = region_editor.clone(); + let auth_refresh_command_editor_clone = auth_refresh_command_editor.clone(); + let access_key_editor_clone = access_key_editor.clone(); + let secret_key_editor_clone = secret_key_editor.clone(); + let refresh_button_clone = refresh_button.clone(); ctx.subscribe_to_model(&AISettings::handle(ctx), move |_, _, event, ctx| { - if matches!( - event, - AISettingsChangedEvent::IsAnyAIEnabled { .. } - | AISettingsChangedEvent::AwsBedrockCredentialsEnabled { .. } - ) { - let is_any_ai_enabled = AISettings::as_ref(ctx).is_any_ai_enabled(ctx); - let is_usage_enabled = is_any_ai_enabled - && UserWorkspaces::as_ref(ctx).is_aws_bedrock_credentials_enabled(ctx); - + if matches!(event, AISettingsChangedEvent::BedrockEnabled { .. }) { + let is_enabled = *AISettings::as_ref(ctx).bedrock_enabled.value(); + profile_dropdown_clone.update(ctx, |dropdown, ctx| { + if is_enabled { + dropdown.set_enabled(ctx); + } else { + dropdown.set_disabled(ctx); + } + }); AISettingsPageView::update_editor_interaction_state( - aws_auth_refresh_command_editor_clone.clone(), - is_usage_enabled, + region_editor_clone.clone(), + is_enabled, ctx, ); AISettingsPageView::update_editor_interaction_state( - aws_auth_refresh_profile_editor_clone.clone(), - is_usage_enabled, + auth_refresh_command_editor_clone.clone(), + is_enabled, ctx, ); - refresh_credentials_button_clone.update(ctx, |button, ctx| { - button.set_disabled(!is_usage_enabled, ctx); + AISettingsPageView::update_editor_interaction_state( + access_key_editor_clone.clone(), + is_enabled, + ctx, + ); + AISettingsPageView::update_editor_interaction_state( + secret_key_editor_clone.clone(), + is_enabled, + ctx, + ); + refresh_button_clone.update(ctx, |button, ctx| { + button.set_disabled(!is_enabled, ctx); }); - ctx.notify(); } }); - let aws_auth_refresh_command_editor_clone = aws_auth_refresh_command_editor.clone(); - let aws_auth_refresh_profile_editor_clone = aws_auth_refresh_profile_editor.clone(); - let refresh_credentials_button_clone = refresh_credentials_button.clone(); - ctx.subscribe_to_model( - &UserWorkspaces::handle(ctx), - move |_, workspace, event, ctx| { - if let UserWorkspacesEvent::TeamsChanged = event { - let is_any_ai_enabled = AISettings::as_ref(ctx).is_any_ai_enabled(ctx); - let is_usage_enabled = is_any_ai_enabled - && workspace - .as_ref(ctx) - .is_aws_bedrock_credentials_enabled(ctx); - - AISettingsPageView::update_editor_interaction_state( - aws_auth_refresh_command_editor_clone.clone(), - is_usage_enabled, - ctx, - ); - AISettingsPageView::update_editor_interaction_state( - aws_auth_refresh_profile_editor_clone.clone(), - is_usage_enabled, - ctx, - ); - refresh_credentials_button_clone.update(ctx, |button, ctx| { - button.set_disabled(!is_usage_enabled, ctx); - }); - - ctx.notify(); - } - }, - ); - Self { - aws_auth_refresh_command_editor, - aws_auth_refresh_profile_editor, - credentials_enabled_toggle: SwitchStateHandle::default(), + enabled_toggle: SwitchStateHandle::default(), + cross_region_toggle: SwitchStateHandle::default(), + fallback_toggle: SwitchStateHandle::default(), auto_login_toggle: SwitchStateHandle::default(), - refresh_credentials_button, + auth_method_dropdown, + profile_dropdown, + region_editor, + auth_refresh_command_editor, + access_key_editor, + secret_key_editor, + refresh_button, } } - fn render_aws_bedrock_section( - &self, + fn render_input( appearance: &Appearance, + label: &'static str, + editor: ViewHandle, + is_enabled: bool, app: &AppContext, - is_bedrock_available: bool, ) -> Box { - let ai_settings = AISettings::as_ref(app); - let user_workspaces = UserWorkspaces::as_ref(app); - let is_any_ai_enabled = ai_settings.is_any_ai_enabled(app); - let is_section_enabled = is_any_ai_enabled && is_bedrock_available; - let is_admin_enforced = matches!( - user_workspaces.aws_bedrock_host_enablement_setting(), - crate::workspaces::workspace::HostEnablementSetting::Enforce - ); - let is_toggleable = - is_section_enabled && user_workspaces.is_aws_bedrock_credentials_toggleable(); - let are_credentials_enabled = user_workspaces.is_aws_bedrock_credentials_enabled(app); - let is_usage_enabled = is_section_enabled && are_credentials_enabled; - let toggle_description = if is_admin_enforced { - "Warp loads and sends local AWS CLI credentials for Bedrock-supported models. This setting is managed by your organization.".to_string() - } else { - "Warp loads and sends local AWS CLI credentials for Bedrock-supported models." - .to_string() + let padding = Some(Coords { + top: 10., + bottom: 10., + left: 16., + right: 16., + }); + let editor_style = UiComponentStyles { + padding, + background: Some(appearance.theme().surface_2().into()), + ..Default::default() }; - let mut column = Flex::column().with_spacing(16.).with_child( - Flex::column() - .with_child(render_ai_setting_toggle::( - "Use AWS Bedrock credentials", - AISettingsPageAction::ToggleAwsBedrockCredentialsEnabled, - are_credentials_enabled, - is_toggleable, - self.credentials_enabled_toggle.clone(), - &RefCell::new(HashMap::new()), - app, - )) - .with_child(render_ai_setting_description( - toggle_description, - is_section_enabled, - app, - )) - .finish(), - ); - - /// Helper function to render the UI for an input field. - fn render_input( - appearance: &Appearance, - label: &'static str, - editor: ViewHandle, - is_enabled: bool, - app: &AppContext, - ) -> Box { - let padding = Some(Coords { - top: 10., - bottom: 10., - left: 16., - right: 16., - }); - let editor_style = UiComponentStyles { - padding, - background: Some(appearance.theme().surface_2().into()), - ..Default::default() - }; - - let label = Text::new_inline(label, appearance.ui_font_family(), CONTENT_FONT_SIZE) - .with_color(styles::header_font_color(is_enabled, app).into()) - .finish(); - - let input = appearance - .ui_builder() - .text_input(editor) - .with_style(editor_style) - .build() - .finish(); - - Flex::column() - .with_spacing(8.) - .with_child(label) - .with_child(input) - .finish() - } - - fn render_credential_status_card( - refresh_button: &ViewHandle, - appearance: &Appearance, - are_credentials_enabled: bool, - app: &AppContext, - ) -> Box { - let (title_color, detail_color) = ( - styles::header_font_color(are_credentials_enabled, app), - styles::description_font_color(are_credentials_enabled, app), - ); - let (title_text, detail_text, icon) = ApiKeyManager::as_ref(app) - .aws_credentials_state() - .user_facing_components(); - - let icon = Container::new( - ConstrainedBox::new(icon.to_warpui_icon(title_color).finish()) - .with_width(16.) - .with_height(16.) - .finish(), - ) - .with_horizontal_padding(4.) + let label = Text::new_inline(label, appearance.ui_font_family(), CONTENT_FONT_SIZE) + .with_color(styles::header_font_color(is_enabled, app).into()) .finish(); - let text_column = Flex::column() - .with_cross_axis_alignment(CrossAxisAlignment::Start) - .with_spacing(4.) - .with_child( - Text::new_inline(title_text, appearance.ui_font_family(), CONTENT_FONT_SIZE) - .with_style(Properties::default().weight(Weight::Semibold)) - .with_color(title_color.into()) - .finish(), - ) - .with_child( - Text::new(detail_text, appearance.ui_font_family(), CONTENT_FONT_SIZE) - .with_color(detail_color.into()) - .soft_wrap(true) - .finish(), - ); + let input = appearance + .ui_builder() + .text_input(editor) + .with_style(editor_style) + .build() + .finish(); - Container::new( - Flex::row() - .with_main_axis_size(MainAxisSize::Max) - .with_cross_axis_alignment(CrossAxisAlignment::Center) - .with_spacing(12.) - .with_child( - Expanded::new( - 1., - Flex::row() - .with_cross_axis_alignment(CrossAxisAlignment::Center) - .with_spacing(12.) - .with_child(icon) - .with_child(Expanded::new(1., text_column.finish()).finish()) - .finish(), - ) - .finish(), - ) - .with_child(ChildView::new(refresh_button).finish()) - .finish(), - ) - .with_uniform_padding(12.) - .with_background(appearance.theme().surface_2()) - .with_border(Border::all(1.).with_border_fill(appearance.theme().outline())) - .with_corner_radius(CornerRadius::with_all(Radius::Pixels(6.))) + Flex::column() + .with_spacing(8.) + .with_child(label) + .with_child(input) .finish() - } - - column.add_child( - Container::new(render_credential_status_card( - &self.refresh_credentials_button, - appearance, - are_credentials_enabled, - app, - )) - .with_margin_top(-styles::DESCRIPTION_MARGIN_BOTTOM) - .finish(), - ); - column.add_child(render_input( - appearance, - "Login Command", - self.aws_auth_refresh_command_editor.clone(), - is_usage_enabled, - app, - )); - column.add_child(render_input( - appearance, - "AWS Profile", - self.aws_auth_refresh_profile_editor.clone(), - is_usage_enabled, - app, - )); - - let auto_login_enabled = *AISettings::as_ref(app).aws_bedrock_auto_login.value(); - - let toggle = render_ai_setting_toggle::( - "Automatically run login command", - AISettingsPageAction::ToggleAwsBedrockAutoLogin, - auto_login_enabled, - is_usage_enabled, - self.auto_login_toggle.clone(), - &RefCell::new(HashMap::new()), - app, - ); - let description = render_ai_setting_description( - "When enabled, the login command will run automatically when AWS Bedrock credentials expire.", - is_usage_enabled, - app, - ); - column.add_child( - Flex::column() - .with_child(toggle) - .with_child(description) - .finish(), - ); - - column.finish() } } -impl SettingsWidget for AwsBedrockWidget { +impl SettingsWidget for BedrockSettingsWidget { type View = AISettingsPageView; fn search_terms(&self) -> &str { - "aws bedrock amazon credentials login profile" + "aws bedrock amazon credentials login profile region sso static keys" } - fn should_render(&self, app: &AppContext) -> bool { - // Only show if admin has enabled AWS Bedrock for the workspace - UserWorkspaces::as_ref(app).is_aws_bedrock_available_from_workspace() + fn should_render(&self, _app: &AppContext) -> bool { + true } fn render( @@ -6717,26 +6694,167 @@ impl SettingsWidget for AwsBedrockWidget { app: &AppContext, ) -> Box { let ai_settings = AISettings::as_ref(app); - let is_any_ai_enabled = ai_settings.is_any_ai_enabled(app); - let is_bedrock_available = - UserWorkspaces::as_ref(app).is_aws_bedrock_available_from_workspace(); + let is_enabled = *ai_settings.bedrock_enabled.value(); + let auth_method = ai_settings.bedrock_auth_method.value().clone(); + let cross_region = *ai_settings.bedrock_cross_region_inference.value(); + let fallback = *ai_settings.bedrock_fallback_to_warp.value(); + let auto_login = *ai_settings.bedrock_auto_login.value(); - let column = Flex::column() - .with_child(render_separator(appearance)) - .with_child( - build_sub_header( - appearance, - "AWS Bedrock", - Some(styles::header_font_color(is_any_ai_enabled, app)), - ) - .with_padding_bottom(HEADER_PADDING) + let mut column = Flex::column().with_spacing(16.); + + column.add_child(render_ai_setting_toggle::( + "Enable AWS Bedrock", + AISettingsPageAction::ToggleBedrockEnabled, + is_enabled, + true, + self.enabled_toggle.clone(), + &RefCell::new(HashMap::new()), + app, + )); + column.add_child(render_ai_setting_description( + "Route AI requests directly through AWS Bedrock using your own credentials.", + true, + app, + )); + + column.add_child(render_separator(appearance)); + + let auth_label = Text::new_inline( + "Authentication Method", + appearance.ui_font_family(), + CONTENT_FONT_SIZE, + ) + .with_color(styles::header_font_color(is_enabled, app).into()) + .finish(); + column.add_child( + Flex::column() + .with_spacing(8.) + .with_child(auth_label) + .with_child(ChildView::new(&self.auth_method_dropdown).finish()) .finish(), - ) - .with_child(self.render_aws_bedrock_section(appearance, app, is_bedrock_available)); + ); - Container::new(column.finish()) - .with_margin_bottom(HEADER_PADDING) - .finish() + match auth_method { + BedrockAuthMethod::Profile | BedrockAuthMethod::Sso => { + let profile_label = Text::new_inline( + "AWS Profile", + appearance.ui_font_family(), + CONTENT_FONT_SIZE, + ) + .with_color(styles::header_font_color(is_enabled, app).into()) + .finish(); + column.add_child( + Flex::column() + .with_spacing(8.) + .with_child(profile_label) + .with_child(ChildView::new(&self.profile_dropdown).finish()) + .finish(), + ); + if auth_method == BedrockAuthMethod::Sso { + column.add_child(Self::render_input( + appearance, + "Login Command", + self.auth_refresh_command_editor.clone(), + is_enabled, + app, + )); + + column.add_child( + Flex::column() + .with_child(render_ai_setting_toggle::( + "Auto-run login on expiry", + AISettingsPageAction::ToggleBedrockAutoLogin, + auto_login, + is_enabled, + self.auto_login_toggle.clone(), + &RefCell::new(HashMap::new()), + app, + )) + .with_child(render_ai_setting_description( + "Automatically run the login command when credentials expire.", + is_enabled, + app, + )) + .finish(), + ); + } + } + BedrockAuthMethod::StaticKeys => { + column.add_child(Self::render_input( + appearance, + "Access Key ID", + self.access_key_editor.clone(), + is_enabled, + app, + )); + column.add_child(Self::render_input( + appearance, + "Secret Access Key", + self.secret_key_editor.clone(), + is_enabled, + app, + )); + } + } + + column.add_child(render_separator(appearance)); + + column.add_child(Self::render_input( + appearance, + "Region", + self.region_editor.clone(), + is_enabled, + app, + )); + column.add_child(render_ai_setting_description( + "Leave empty to auto-detect from your AWS profile/config.", + is_enabled, + app, + )); + + column.add_child( + Flex::column() + .with_child(render_ai_setting_toggle::( + "Cross-region inference", + AISettingsPageAction::ToggleBedrockCrossRegionInference, + cross_region, + is_enabled, + self.cross_region_toggle.clone(), + &RefCell::new(HashMap::new()), + app, + )) + .with_child(render_ai_setting_description( + "Automatically add geographic prefixes to model IDs for higher availability.", + is_enabled, + app, + )) + .finish(), + ); + + column.add_child( + Flex::column() + .with_child(render_ai_setting_toggle::( + "Fallback to Warp server", + AISettingsPageAction::ToggleBedrockFallbackToWarp, + fallback, + is_enabled, + self.fallback_toggle.clone(), + &RefCell::new(HashMap::new()), + app, + )) + .with_child(render_ai_setting_description( + "When enabled, requests will route through the Warp server if Bedrock credentials are invalid.", + is_enabled, + app, + )) + .finish(), + ); + + column.add_child(render_separator(appearance)); + + column.add_child(ChildView::new(&self.refresh_button).finish()); + + column.finish() } } diff --git a/app/src/settings_view/mod.rs b/app/src/settings_view/mod.rs index 15cdbdd8..8ffae34b 100644 --- a/app/src/settings_view/mod.rs +++ b/app/src/settings_view/mod.rs @@ -211,6 +211,7 @@ pub enum SettingsSection { AgentMCPServers, Knowledge, ThirdPartyCLIAgents, + Bedrock, /// Internal backing-page identifier for CodeSettingsPageView. Multiple subpages /// (CodeIndexing, EditorAndCodeReview) share this single backing page, /// so this variant is needed as the key in `settings_pages`. @@ -240,6 +241,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::CodeIndexing => write!(f, "Indexing and projects"), SettingsSection::EditorAndCodeReview => write!(f, "Editor and Code Review"), SettingsSection::CloudEnvironments => write!(f, "Environments"), @@ -264,6 +266,7 @@ impl SettingsSection { | Self::AgentMCPServers | Self::Knowledge | Self::ThirdPartyCLIAgents + | Self::Bedrock ) } @@ -301,6 +304,7 @@ impl SettingsSection { Self::AgentMCPServers, Self::Knowledge, Self::ThirdPartyCLIAgents, + Self::Bedrock, ] } @@ -341,6 +345,7 @@ 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), "Indexing and projects" | "CodeIndexing" => Ok(Self::CodeIndexing), "Editor and Code Review" | "EditorAndCodeReview" => Ok(Self::EditorAndCodeReview), "CloudEnvironments" => Ok(Self::CloudEnvironments), diff --git a/app/src/terminal/input/slash_commands/mod.rs b/app/src/terminal/input/slash_commands/mod.rs index a30f34f9..7bf1255e 100644 --- a/app/src/terminal/input/slash_commands/mod.rs +++ b/app/src/terminal/input/slash_commands/mod.rs @@ -318,16 +318,16 @@ impl Input { // Handle the slash command action based on its kind match command.name { - add_mcp if command.name == commands::ADD_MCP.name => { + _add_mcp if command.name == commands::ADD_MCP.name => { ctx.dispatch_typed_action(&TerminalAction::OpenAddMCPPane); } - add_prompt if command.name == commands::ADD_PROMPT.name => { + _add_prompt if command.name == commands::ADD_PROMPT.name => { ctx.dispatch_typed_action(&TerminalAction::OpenAddPromptPane); } - add_rule if command.name == commands::ADD_RULE.name => { + _add_rule if command.name == commands::ADD_RULE.name => { ctx.dispatch_typed_action(&TerminalAction::OpenAddRulePane); } - agent_or_new + _agent_or_new if command.name == commands::NEW.name || command.name == commands::AGENT.name => { if !self @@ -382,7 +382,7 @@ impl Input { origin: AgentViewEntryOrigin::SlashCommand { trigger }, }); } - cloud_agent if command.name == commands::CLOUD_AGENT.name => { + _cloud_agent if command.name == commands::CLOUD_AGENT.name => { let prompt = argument.and_then(|argument| { let trimmed = argument.trim(); if trimmed.is_empty() { @@ -396,17 +396,17 @@ impl Input { initial_prompt: prompt, }); } - create_docker_sandbox if command.name == commands::CREATE_DOCKER_SANDBOX.name => { + _create_docker_sandbox if command.name == commands::CREATE_DOCKER_SANDBOX.name => { ctx.emit(Event::CreateDockerSandbox); } - conversations if command.name == commands::CONVERSATIONS.name => { + _conversations if command.name == commands::CONVERSATIONS.name => { if FeatureFlag::AgentView.is_enabled() { self.open_conversation_menu(ctx); } else { ctx.dispatch_typed_action(&TerminalAction::OpenConversationsPalette); } } - rename_tab if command.name == commands::RENAME_TAB.name => { + _rename_tab if command.name == commands::RENAME_TAB.name => { let Some(name) = argument .map(|name| name.trim()) .filter(|name| !name.is_empty()) @@ -420,7 +420,7 @@ impl Input { ctx.dispatch_typed_action(&WorkspaceAction::SetActiveTabName(name.to_owned())); } - create_env if command.name == commands::CREATE_ENVIRONMENT.name => { + _create_env if command.name == commands::CREATE_ENVIRONMENT.name => { // If the user included args after the slash command, treat them as repo paths/URLs. let repos = argument .map(|arg| { @@ -433,7 +433,7 @@ impl Input { ctx.emit(Event::TriggerEnvironmentSetup { repos }); } - create_project if command.name == commands::CREATE_NEW_PROJECT.name => { + _create_project if command.name == commands::CREATE_NEW_PROJECT.name => { if argument.is_none_or(|args| args.is_empty()) { show_error_toast( "Please describe the project you want to create after /create-new-project" @@ -446,7 +446,7 @@ impl Input { let args = argument.expect("args are Some()"); self.initiate_create_new_project(args.to_owned(), ctx); } - edit if command.name == commands::EDIT.name => { + _edit if command.name == commands::EDIT.name => { #[cfg(feature = "local_fs")] match argument { Some(args) if !args.is_empty() => { @@ -539,7 +539,7 @@ impl Input { return true; } } - export_to_clipboard if command.name == commands::EXPORT_TO_CLIPBOARD.name => { + _export_to_clipboard if command.name == commands::EXPORT_TO_CLIPBOARD.name => { let history = BlocklistAIHistoryModel::handle(ctx); let Some(conversation) = history .as_ref(ctx) @@ -564,7 +564,7 @@ impl Input { toast_stack.add_ephemeral_toast(toast, window_id, ctx); }); } - export_to_file if command.name == commands::EXPORT_TO_FILE.name => { + _export_to_file if command.name == commands::EXPORT_TO_FILE.name => { #[cfg(not(target_family = "wasm"))] { self.export_conversation_to_file( @@ -581,76 +581,76 @@ impl Input { return true; } } - index if command.name == commands::INDEX.name => { + _index if command.name == commands::INDEX.name => { ctx.dispatch_typed_action(&TerminalAction::IndexProjectSpeedbump); } - init if command.name == commands::INIT.name => { + _init if command.name == commands::INIT.name => { ctx.dispatch_typed_action(&TerminalAction::InitProject); } - changelog if command.name == commands::CHANGELOG.name => { + _changelog if command.name == commands::CHANGELOG.name => { if !FeatureFlag::Changelog.is_enabled() { return false; } ctx.dispatch_typed_action(&WorkspaceAction::ViewLatestChangelog); } - feedback if command.name == commands::FEEDBACK.name => { + _feedback if command.name == commands::FEEDBACK.name => { ctx.dispatch_typed_action(&WorkspaceAction::SendFeedback); } - open_code_review if command.name == commands::OPEN_CODE_REVIEW.name => { + _open_code_review if command.name == commands::OPEN_CODE_REVIEW.name => { ctx.dispatch_typed_action(&TerminalAction::ToggleCodeReviewPane { entrypoint: CodeReviewPaneEntrypoint::SlashCommand, }); } - open_mcp_servers if command.name == commands::OPEN_MCP_SERVERS.name => { + _open_mcp_servers if command.name == commands::OPEN_MCP_SERVERS.name => { ctx.dispatch_typed_action(&TerminalAction::OpenViewMCPPane); } - open_settings_file if command.name == commands::OPEN_SETTINGS_FILE.name => { + _open_settings_file if command.name == commands::OPEN_SETTINGS_FILE.name => { if !FeatureFlag::SettingsFile.is_enabled() || !cfg!(feature = "local_fs") { return false; } ctx.dispatch_typed_action(&WorkspaceAction::OpenSettingsFile); } - open_project_rules if command.name == commands::OPEN_PROJECT_RULES.name => { + _open_project_rules if command.name == commands::OPEN_PROJECT_RULES.name => { ctx.dispatch_typed_action(&TerminalAction::OpenProjectRulesPane); } - open_rules if command.name == commands::OPEN_RULES.name => { + _open_rules if command.name == commands::OPEN_RULES.name => { ctx.dispatch_typed_action(&TerminalAction::OpenRulesPane); } - edit_skill if command.name == commands::EDIT_SKILL.name => { + _edit_skill if command.name == commands::EDIT_SKILL.name => { if !FeatureFlag::ListSkills.is_enabled() { return false; } // Open the skill selector menu - user will select a skill from the inline menu self.open_skill_selector(ctx); } - invoke_skill if command.name == commands::INVOKE_SKILL.name => { + _invoke_skill if command.name == commands::INVOKE_SKILL.name => { if !FeatureFlag::ListSkills.is_enabled() { return false; } // Open the skill selector menu for invocation - skill command will be inserted into buffer self.open_invoke_skill_selector(ctx); } - models if command.name == commands::MODEL.name => { + _models if command.name == commands::MODEL.name => { self.open_model_selector(ctx); } - profiles if command.name == commands::PROFILE.name => { + _profiles if command.name == commands::PROFILE.name => { if !FeatureFlag::InlineProfileSelector.is_enabled() { return false; } self.open_profile_selector(ctx); } - prompts if command.name == commands::PROMPTS.name => { + _prompts if command.name == commands::PROMPTS.name => { if FeatureFlag::AgentView.is_enabled() { self.open_prompts_menu(ctx); } else { return false; } } - rewind if command.name == commands::REWIND.name => { + _rewind if command.name == commands::REWIND.name => { self.open_rewind_menu(ctx); } - pr_comments if command.name == commands::PR_COMMENTS.name => { + _pr_comments if command.name == commands::PR_COMMENTS.name => { if !FeatureFlag::PRCommentsSlashCommand.is_enabled() { return false; } @@ -671,10 +671,10 @@ impl Input { ) }); } - usage if command.name == commands::USAGE.name => { + _usage if command.name == commands::USAGE.name => { ctx.dispatch_typed_action(&TerminalAction::OpenBillingAndUsagePane); } - remote_control if command.name == commands::REMOTE_CONTROL.name => { + _remote_control if command.name == commands::REMOTE_CONTROL.name => { if !FeatureFlag::CreatingSharedSessions.is_enabled() || !FeatureFlag::HOARemoteControl.is_enabled() { @@ -691,7 +691,7 @@ impl Input { } ctx.emit(Event::StartRemoteControl); } - cost if command.name == commands::COST.name => { + _cost if command.name == commands::COST.name => { let history = BlocklistAIHistoryModel::handle(ctx); let conversation = history .as_ref(ctx) @@ -715,7 +715,7 @@ impl Input { ctx.dispatch_typed_action(&TerminalAction::ToggleUsageFooter); } } - fork if command.name == commands::FORK.name => { + _fork if command.name == commands::FORK.name => { let Some(conversation_id) = self .ai_context_model .as_ref(ctx) @@ -740,11 +740,11 @@ impl Input { destination, }); } - fork_from if command.name == commands::FORK_FROM.name => { + _fork_from if command.name == commands::FORK_FROM.name => { self.open_user_query_menu(UserQueryMenuAction::ForkFrom, ctx); return true; } - fork_and_compact if command.name == commands::FORK_AND_COMPACT.name => { + _fork_and_compact if command.name == commands::FORK_AND_COMPACT.name => { let Some(conversation_id) = self .ai_context_model .as_ref(ctx) @@ -772,7 +772,7 @@ impl Input { destination, }); } - compact_and if command.name == commands::COMPACT_AND.name => { + _compact_and if command.name == commands::COMPACT_AND.name => { if self .ai_context_model .as_ref(ctx) @@ -791,7 +791,7 @@ impl Input { initial_prompt: argument.cloned(), }); } - queue if command.name == commands::QUEUE.name => { + _queue if command.name == commands::QUEUE.name => { let Some(conversation_id) = self .ai_context_model .as_ref(ctx) @@ -820,13 +820,13 @@ impl Input { self.submit_queued_prompt(prompt, ctx); } } - open_repo if command.name == commands::OPEN_REPO.name => { + _open_repo if command.name == commands::OPEN_REPO.name => { if !FeatureFlag::InlineRepoMenu.is_enabled() { return false; } self.open_repos_menu(ctx); } - command_that_just_sends_ai_request_with_prefix + _command_that_just_sends_ai_request_with_prefix if command.name == commands::COMPACT.name || command.name == commands::PLAN.name || command.name == commands::ORCHESTRATE.name => diff --git a/app/src/terminal/view.rs b/app/src/terminal/view.rs index b8b619fd..6e8bf8c2 100644 --- a/app/src/terminal/view.rs +++ b/app/src/terminal/view.rs @@ -3950,8 +3950,8 @@ impl TerminalView { } ctx.subscribe_to_model(&AISettings::handle(ctx), |me, _, ai_settings_event, ctx| { - if let AISettingsChangedEvent::AwsBedrockCredentialsEnabled { .. } = ai_settings_event { - if !UserWorkspaces::as_ref(ctx).is_aws_bedrock_credentials_enabled(ctx) { + if let AISettingsChangedEvent::BedrockEnabled { .. } = ai_settings_event { + if !UserWorkspaces::as_ref(ctx).is_bedrock_enabled(ctx) { me.remove_aws_bedrock_login_banner(ctx); } } @@ -9318,7 +9318,7 @@ impl TerminalView { AwsBedrockLoginBannerAction::DontShowAgain => { AISettings::handle(ctx).update(ctx, |ai_settings, ctx| { report_if_error!(ai_settings - .aws_bedrock_login_banner_dismissed + .bedrock_login_banner_dismissed .set_value(true, ctx)); }); } @@ -9338,7 +9338,7 @@ impl TerminalView { /// user interaction (e.g. "do you want to override X profile? y/n" is common) fn run_aws_login_command(&mut self, ctx: &mut ViewContext) { let login_command = AISettings::as_ref(ctx) - .aws_bedrock_auth_refresh_command + .bedrock_auth_refresh_command .value() .clone(); @@ -9377,7 +9377,7 @@ impl TerminalView { } // Check if AWS Bedrock is available in the workspace - if !UserWorkspaces::as_ref(ctx).is_aws_bedrock_credentials_enabled(ctx) { + if !UserWorkspaces::as_ref(ctx).is_bedrock_enabled(ctx) { return; } diff --git a/app/src/terminal/view/inline_banner/session_state.rs b/app/src/terminal/view/inline_banner/session_state.rs index a63e0187..784f0bbe 100644 --- a/app/src/terminal/view/inline_banner/session_state.rs +++ b/app/src/terminal/view/inline_banner/session_state.rs @@ -17,14 +17,14 @@ impl ByoLlmAuthBannerSessionState { pub fn new(ctx: &mut ModelContext) -> Self { // Initialize from the persisted permanent dismissal setting let dismissed = *AISettings::as_ref(ctx) - .aws_bedrock_login_banner_dismissed + .bedrock_login_banner_dismissed .value(); // Subscribe to changes in the permanent dismissal setting ctx.subscribe_to_model(&AISettings::handle(ctx), |state, event, ctx| { - if let AISettingsChangedEvent::AwsBedrockLoginBannerDismissed { .. } = event { + if let AISettingsChangedEvent::BedrockLoginBannerDismissed { .. } = event { let permanently_dismissed = *AISettings::as_ref(ctx) - .aws_bedrock_login_banner_dismissed + .bedrock_login_banner_dismissed .value(); if permanently_dismissed && !state.dismissed { state.dismissed = true; diff --git a/app/src/workspace/action.rs b/app/src/workspace/action.rs index 3bfd52ed..b8fe92ef 100644 --- a/app/src/workspace/action.rs +++ b/app/src/workspace/action.rs @@ -541,7 +541,7 @@ pub enum WorkspaceAction { ResetBuildPlanMigrationModalState, /// Reset the AWS Bedrock login banner dismissed state (for debugging). #[cfg(debug_assertions)] - DebugResetAwsBedrockLoginBannerDismissed, + DebugResetBedrockLoginBannerDismissed, /// Open the Oz Launch Modal (for debugging) #[cfg(debug_assertions)] OpenOzLaunchModal, @@ -941,7 +941,7 @@ impl WorkspaceAction { #[cfg(debug_assertions)] OpenBuildPlanMigrationModal | ResetBuildPlanMigrationModalState - | DebugResetAwsBedrockLoginBannerDismissed + | DebugResetBedrockLoginBannerDismissed | OpenOzLaunchModal | ResetOzLaunchModalState | OpenOpenWarpLaunchModal diff --git a/app/src/workspace/mod.rs b/app/src/workspace/mod.rs index a4432dfb..1326c78d 100644 --- a/app/src/workspace/mod.rs +++ b/app/src/workspace/mod.rs @@ -198,9 +198,9 @@ pub fn init(app: &mut AppContext) { ) .with_context_predicate(id!("Workspace")), EditableBinding::new( - "workspace:debug_reset_aws_bedrock_login_banner_dismissed", + "workspace:debug_reset_bedrock_login_banner_dismissed", "[Debug] Un-dismiss AWS login banner", - WorkspaceAction::DebugResetAwsBedrockLoginBannerDismissed, + WorkspaceAction::DebugResetBedrockLoginBannerDismissed, ) .with_context_predicate(id!("Workspace")), EditableBinding::new( diff --git a/app/src/workspace/view.rs b/app/src/workspace/view.rs index dbf0e189..86f74f90 100644 --- a/app/src/workspace/view.rs +++ b/app/src/workspace/view.rs @@ -21233,11 +21233,11 @@ impl TypedActionView for Workspace { log::info!("Build plan migration modal dismissed state has been reset"); } #[cfg(debug_assertions)] - DebugResetAwsBedrockLoginBannerDismissed => { + DebugResetBedrockLoginBannerDismissed => { // Reset the AWS Bedrock login banner dismissed state for debugging AISettings::handle(ctx).update(ctx, |ai_settings, ctx| { if let Err(e) = ai_settings - .aws_bedrock_login_banner_dismissed + .bedrock_login_banner_dismissed .set_value(false, ctx) { log::warn!( diff --git a/app/src/workspaces/user_workspaces.rs b/app/src/workspaces/user_workspaces.rs index 8b17db6e..a81cb73d 100644 --- a/app/src/workspaces/user_workspaces.rs +++ b/app/src/workspaces/user_workspaces.rs @@ -511,7 +511,7 @@ impl UserWorkspaces { ) } - pub fn is_aws_bedrock_credentials_enabled(&self, app: &AppContext) -> bool { + pub fn is_bedrock_enabled(&self, app: &AppContext) -> bool { // i.e. did the admin go and toggle on aws bedrock in the admin panel? if !self.is_aws_bedrock_available_from_workspace() { return false; @@ -520,7 +520,7 @@ impl UserWorkspaces { match self.aws_bedrock_host_enablement_setting() { HostEnablementSetting::Enforce => true, HostEnablementSetting::RespectUserSetting => *AISettings::as_ref(app) - .aws_bedrock_credentials_enabled + .bedrock_enabled .value(), } } diff --git a/app/src/workspaces/user_workspaces_tests.rs b/app/src/workspaces/user_workspaces_tests.rs index ba56ba1f..dc64f8c7 100644 --- a/app/src/workspaces/user_workspaces_tests.rs +++ b/app/src/workspaces/user_workspaces_tests.rs @@ -254,7 +254,7 @@ fn test_aws_bedrock_credentials_default_off_when_admin_respects_user_setting() { app.read(|ctx| { assert!( - !UserWorkspaces::as_ref(ctx).is_aws_bedrock_credentials_enabled(ctx), + !UserWorkspaces::as_ref(ctx).is_bedrock_enabled(ctx), "respect-user-setting should default the local Bedrock credentials toggle to off" ); assert!( @@ -303,13 +303,13 @@ fn test_aws_bedrock_credentials_respect_user_setting() { AISettings::handle(&app).update(&mut app, |settings, ctx| { let _ = settings - .aws_bedrock_credentials_enabled + .bedrock_enabled .set_value(false, ctx); }); app.read(|ctx| { assert!( - !UserWorkspaces::as_ref(ctx).is_aws_bedrock_credentials_enabled(ctx), + !UserWorkspaces::as_ref(ctx).is_bedrock_enabled(ctx), "respect-user-setting should honor the local Bedrock credentials toggle" ); assert!( @@ -358,13 +358,13 @@ fn test_aws_bedrock_credentials_enforced_by_admin() { AISettings::handle(&app).update(&mut app, |settings, ctx| { let _ = settings - .aws_bedrock_credentials_enabled + .bedrock_enabled .set_value(false, ctx); }); app.read(|ctx| { assert!( - UserWorkspaces::as_ref(ctx).is_aws_bedrock_credentials_enabled(ctx), + UserWorkspaces::as_ref(ctx).is_bedrock_enabled(ctx), "enforced Bedrock host policy should ignore the local Bedrock credentials toggle" ); assert!( diff --git a/crates/warp_core/src/ui/icons.rs b/crates/warp_core/src/ui/icons.rs index 87192ba0..dd30df0e 100644 --- a/crates/warp_core/src/ui/icons.rs +++ b/crates/warp_core/src/ui/icons.rs @@ -278,6 +278,7 @@ pub enum Icon { CopilotLogo, PiLogo, AuggieLogo, + BedrockLogo, CursorLogo, NLD, Oz, @@ -579,6 +580,7 @@ impl From for &'static str { Icon::ClaudeLogo => "bundled/svg/claude.svg", Icon::GeminiLogo => "bundled/svg/gemini_cli.svg", Icon::OpenAILogo => "bundled/svg/openai.svg", + Icon::BedrockLogo => "bundled/svg/bedrock.svg", Icon::AmpLogo => "bundled/svg/amp.svg", Icon::DroidLogo => "bundled/svg/droid.svg", Icon::OpenCodeLogo => "bundled/svg/opencode.svg",