Skip to repository content3784 lines · 147.3 KB · rust
tenant.openagents/omega
No repository description is available.
OpenAgents Git authority 2026-07-28T04:46:25.928Z Public web read
NIP-34 coordinate
30617:7649603503856e5148d571eac2766b288a8ff1e9e35d380337a1d2b0015b4f92:omegaMaintainersHidden in public view
References2 branches · 1 tag
Read-only clone
git clone https://openagents.com/git/tenant.openagents/omega.gitBrowse files
bedrock.rs
1use std::pin::Pin;
2use std::sync::Arc;
3
4use anyhow::{Context as _, Result, anyhow};
5use async_lock::OnceCell;
6use aws_config::stalled_stream_protection::StalledStreamProtectionConfig;
7use aws_config::{BehaviorVersion, Region};
8use aws_credential_types::provider::{ProvideCredentials, SharedCredentialsProvider};
9use aws_credential_types::{Credentials, Token};
10use aws_http_client::AwsHttpClient;
11use aws_sigv4::http_request::{SignableBody, SignableRequest, SigningSettings, sign};
12use aws_sigv4::sign::v4;
13use bedrock::BedrockSystemContentBlock;
14use bedrock::bedrock_client::Client as BedrockClient;
15use bedrock::bedrock_client::config::timeout::TimeoutConfig;
16use bedrock::bedrock_client::types::{
17 CachePointBlock, CachePointType, ContentBlockDelta, ContentBlockStart, ConverseStreamOutput,
18 ReasoningContentBlockDelta, StopReason,
19};
20use bedrock::{
21 BedrockAnyToolChoice, BedrockAutoToolChoice, BedrockBlob, BedrockError, BedrockImageBlock,
22 BedrockImageFormat, BedrockImageSource, BedrockInnerContent, BedrockMessage, BedrockModelMode,
23 BedrockStreamingResponse, BedrockThinkingBlock, BedrockThinkingTextBlock, BedrockTool,
24 BedrockToolChoice, BedrockToolConfig, BedrockToolInputSchema, BedrockToolResultBlock,
25 BedrockToolResultContentBlock, BedrockToolResultStatus, BedrockToolSpec, BedrockToolUseBlock,
26 ConverseModel, MantleModel, MantleProtocol, value_to_aws_document,
27};
28use collections::{BTreeMap, HashMap};
29use credentials_provider::CredentialsProvider;
30use futures::{
31 AsyncBufReadExt, AsyncReadExt, FutureExt, Stream, StreamExt, future::BoxFuture, io::BufReader,
32 stream::BoxStream,
33};
34use gpui::{
35 App, AsyncApp, Context, Entity, FocusHandle, Subscription, Task, TaskExt, Window, actions,
36};
37use gpui_tokio::Tokio;
38use http_client::{
39 AsyncBody, CustomHeaders, HttpClient, Method, Request as HttpRequest, RequestBuilderExt,
40 http::{HeaderValue, header::AUTHORIZATION},
41};
42use language_model::{
43 AuthenticateError, EnvVar, IconOrSvg, InlineDescription, LanguageModel,
44 LanguageModelCompletionError, LanguageModelCompletionEvent, LanguageModelEffortLevel,
45 LanguageModelId, LanguageModelName, LanguageModelProvider, LanguageModelProviderId,
46 LanguageModelProviderName, LanguageModelProviderState, LanguageModelRequest,
47 LanguageModelToolChoice, LanguageModelToolResultContent, LanguageModelToolSchemaFormat,
48 LanguageModelToolUse, MessageContent, ProviderSettingsView, RateLimiter, Role,
49 SubPageProviderSettings, TokenUsage, env_var,
50};
51use open_ai::responses::Request as OpenAiResponseRequest;
52use open_ai::responses::{ResponseOutputItem, StreamEvent as OpenAiResponseStreamEvent};
53use schemars::JsonSchema;
54use serde::{Deserialize, Serialize};
55use serde_json::Value;
56use settings::{
57 BedrockAvailableModel as AvailableModel, BedrockMantleAvailableModel as MantleAvailableModel,
58 Settings, SettingsStore,
59};
60use std::sync::LazyLock;
61use std::time::SystemTime;
62use strum::{EnumIter, IntoEnumIterator, IntoStaticStr};
63use ui::{ButtonLink, ConfiguredApiCard, Divider, List, ListBulletItem, prelude::*};
64use ui_input::InputField;
65use util::ResultExt;
66
67use crate::AllLanguageModelSettings;
68use crate::provider::open_ai::{
69 ChatCompletionMaxTokensParameter, OpenAiEventMapper, OpenAiResponseEventMapper, into_open_ai,
70 into_open_ai_response,
71};
72use language_model::util::{fix_streamed_json, parse_tool_arguments};
73use open_ai::{ReasoningEffort, RequestError, ResponseStreamEvent};
74
75actions!(bedrock, [Tab, TabPrev]);
76
77const PROVIDER_ID: LanguageModelProviderId = LanguageModelProviderId::new("amazon-bedrock");
78const PROVIDER_NAME: LanguageModelProviderName = LanguageModelProviderName::new("Amazon Bedrock");
79pub(crate) const RESERVED_HEADER_NAMES: &[&str] = &[
80 "host",
81 "x-amz-date",
82 "x-amz-security-token",
83 "x-amz-content-sha256",
84 "amz-sdk-invocation-id",
85 "amz-sdk-request",
86];
87
88/// Credentials stored in the keychain for static authentication.
89/// Region is handled separately since it's orthogonal to auth method.
90#[derive(Default, Clone, Deserialize, Serialize, PartialEq, Debug)]
91pub struct BedrockCredentials {
92 pub access_key_id: String,
93 pub secret_access_key: String,
94 pub session_token: Option<String>,
95 pub bearer_token: Option<String>,
96}
97
98/// Resolved authentication configuration for Bedrock.
99/// Settings take priority over UX-provided credentials.
100#[derive(Clone, Debug, PartialEq)]
101pub enum BedrockAuth {
102 /// Use default AWS credential provider chain (IMDSv2, PodIdentity, env vars, etc.)
103 Automatic,
104 /// Use AWS named profile from ~/.aws/credentials or ~/.aws/config
105 NamedProfile { profile_name: String },
106 /// Use AWS SSO profile
107 SingleSignOn { profile_name: String },
108 /// Use IAM credentials (access key + secret + optional session token)
109 IamCredentials {
110 access_key_id: String,
111 secret_access_key: String,
112 session_token: Option<String>,
113 },
114 /// Use Bedrock API Key (bearer token authentication)
115 ApiKey { api_key: String },
116}
117
118impl BedrockCredentials {
119 /// Convert stored credentials to the appropriate auth variant.
120 /// Prefers API key if present, otherwise uses IAM credentials.
121 fn into_auth(self) -> Option<BedrockAuth> {
122 if let Some(api_key) = self.bearer_token.filter(|t| !t.is_empty()) {
123 Some(BedrockAuth::ApiKey { api_key })
124 } else if !self.access_key_id.is_empty() && !self.secret_access_key.is_empty() {
125 Some(BedrockAuth::IamCredentials {
126 access_key_id: self.access_key_id,
127 secret_access_key: self.secret_access_key,
128 session_token: self.session_token.filter(|t| !t.is_empty()),
129 })
130 } else {
131 None
132 }
133 }
134}
135
136#[derive(Default, Clone, Debug, PartialEq)]
137pub struct AmazonBedrockSettings {
138 pub available_models: Vec<AvailableModel>,
139 pub mantle_available_models: Vec<MantleAvailableModel>,
140 pub custom_headers: CustomHeaders,
141 pub region: Option<String>,
142 pub endpoint: Option<String>,
143 pub profile_name: Option<String>,
144 pub role_arn: Option<String>,
145 pub authentication_method: Option<BedrockAuthMethod>,
146 pub allow_global: Option<bool>,
147 pub guardrail_identifier: Option<String>,
148 pub guardrail_version: Option<String>,
149}
150
151#[derive(Clone, Debug, PartialEq, Serialize, Deserialize, EnumIter, IntoStaticStr, JsonSchema)]
152pub enum BedrockAuthMethod {
153 #[serde(rename = "named_profile")]
154 NamedProfile,
155 #[serde(rename = "sso")]
156 SingleSignOn,
157 #[serde(rename = "api_key")]
158 ApiKey,
159 /// IMDSv2, PodIdentity, env vars, etc.
160 #[serde(rename = "default")]
161 Automatic,
162}
163
164impl From<settings::BedrockAuthMethodContent> for BedrockAuthMethod {
165 fn from(value: settings::BedrockAuthMethodContent) -> Self {
166 match value {
167 settings::BedrockAuthMethodContent::SingleSignOn => BedrockAuthMethod::SingleSignOn,
168 settings::BedrockAuthMethodContent::Automatic => BedrockAuthMethod::Automatic,
169 settings::BedrockAuthMethodContent::NamedProfile => BedrockAuthMethod::NamedProfile,
170 settings::BedrockAuthMethodContent::ApiKey => BedrockAuthMethod::ApiKey,
171 }
172 }
173}
174
175fn mantle_protocol_from_settings(value: settings::BedrockMantleProtocolContent) -> MantleProtocol {
176 match value {
177 settings::BedrockMantleProtocolContent::ChatCompletions => MantleProtocol::ChatCompletions,
178 settings::BedrockMantleProtocolContent::Responses => MantleProtocol::Responses,
179 }
180}
181
182#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize, JsonSchema)]
183#[serde(tag = "type", rename_all = "lowercase")]
184pub enum ModelMode {
185 #[default]
186 Default,
187 Thinking {
188 /// The maximum number of tokens to use for reasoning. Must be lower than the model's `max_output_tokens`.
189 budget_tokens: Option<u64>,
190 },
191 AdaptiveThinking {
192 effort: bedrock::BedrockAdaptiveThinkingEffort,
193 },
194}
195
196impl From<ModelMode> for BedrockModelMode {
197 fn from(value: ModelMode) -> Self {
198 match value {
199 ModelMode::Default => BedrockModelMode::Default,
200 ModelMode::Thinking { budget_tokens } => BedrockModelMode::Thinking { budget_tokens },
201 ModelMode::AdaptiveThinking { effort } => BedrockModelMode::AdaptiveThinking { effort },
202 }
203 }
204}
205
206impl From<BedrockModelMode> for ModelMode {
207 fn from(value: BedrockModelMode) -> Self {
208 match value {
209 BedrockModelMode::Default => ModelMode::Default,
210 BedrockModelMode::Thinking { budget_tokens } => ModelMode::Thinking { budget_tokens },
211 BedrockModelMode::AdaptiveThinking { effort } => ModelMode::AdaptiveThinking { effort },
212 }
213 }
214}
215
216/// The URL of the base AWS service.
217///
218/// Right now we're just using this as the key to store the AWS credentials
219/// under in the keychain.
220const AMAZON_AWS_URL: &str = "https://amazonaws.com";
221
222// These environment variables all use a `ZED_` prefix because we don't want to overwrite the user's AWS credentials.
223static ZED_BEDROCK_ACCESS_KEY_ID_VAR: LazyLock<EnvVar> = env_var!("ZED_ACCESS_KEY_ID");
224static ZED_BEDROCK_SECRET_ACCESS_KEY_VAR: LazyLock<EnvVar> = env_var!("ZED_SECRET_ACCESS_KEY");
225static ZED_BEDROCK_SESSION_TOKEN_VAR: LazyLock<EnvVar> = env_var!("ZED_SESSION_TOKEN");
226static ZED_AWS_PROFILE_VAR: LazyLock<EnvVar> = env_var!("ZED_AWS_PROFILE");
227static ZED_BEDROCK_REGION_VAR: LazyLock<EnvVar> = env_var!("ZED_AWS_REGION");
228static ZED_AWS_ENDPOINT_VAR: LazyLock<EnvVar> = env_var!("ZED_AWS_ENDPOINT");
229static ZED_BEDROCK_BEARER_TOKEN_VAR: LazyLock<EnvVar> = env_var!("ZED_BEDROCK_BEARER_TOKEN");
230
231/// AWS Regions where the `bedrock-mantle` endpoint is available.
232/// See <https://docs.aws.amazon.com/bedrock/latest/userguide/bedrock-mantle.html#regions>.
233const MANTLE_SUPPORTED_REGIONS: &[&str] = &[
234 "us-east-2",
235 "us-east-1",
236 "us-west-2",
237 "ap-southeast-3",
238 "ap-south-1",
239 "ap-southeast-2",
240 "ap-northeast-1",
241 "eu-central-1",
242 "eu-west-1",
243 "eu-west-2",
244 "eu-south-1",
245 "eu-north-1",
246 "sa-east-1",
247 "us-gov-west-1",
248];
249
250fn mantle_endpoint_url(region: &str) -> String {
251 format!("https://bedrock-mantle.{region}.api.aws/openai/v1")
252}
253
254enum MantleAuth {
255 ApiKey { api_key: String },
256 SigV4 { credentials: Credentials },
257}
258
259impl MantleAuth {
260 fn apply(&self, request: &mut HttpRequest<AsyncBody>, body: &[u8], region: &str) -> Result<()> {
261 match self {
262 MantleAuth::ApiKey { api_key } => {
263 let value = HeaderValue::from_str(&format!("Bearer {}", api_key.trim()))
264 .context("building Mantle bearer token authorization header")?;
265 request.headers_mut().insert(AUTHORIZATION, value);
266 }
267 MantleAuth::SigV4 { credentials } => {
268 sign_mantle_request_sigv4(request, body, credentials, region)?;
269 }
270 }
271
272 Ok(())
273 }
274}
275
276fn sign_mantle_request_sigv4(
277 request: &mut HttpRequest<AsyncBody>,
278 body: &[u8],
279 credentials: &Credentials,
280 region: &str,
281) -> Result<()> {
282 sign_mantle_request_sigv4_at(request, body, credentials, region, SystemTime::now())
283}
284
285fn sign_mantle_request_sigv4_at(
286 request: &mut HttpRequest<AsyncBody>,
287 body: &[u8],
288 credentials: &Credentials,
289 region: &str,
290 time: SystemTime,
291) -> Result<()> {
292 if !request
293 .headers()
294 .contains_key(http_client::http::header::HOST)
295 && let Some(authority) = request.uri().authority()
296 {
297 let host = HeaderValue::from_str(authority.as_str())
298 .context("invalid host header derived from Mantle request URI")?;
299 request
300 .headers_mut()
301 .insert(http_client::http::header::HOST, host);
302 }
303
304 let identity = credentials.clone().into();
305 let signing_params: aws_sigv4::http_request::SigningParams = v4::SigningParams::builder()
306 .identity(&identity)
307 .region(region)
308 .name("bedrock-mantle")
309 .time(time)
310 .settings(SigningSettings::default())
311 .build()
312 .context("building Mantle SigV4 signing params")?
313 .into();
314
315 let method = request.method().as_str();
316 let uri = request.uri().to_string();
317 let headers = request
318 .headers()
319 .iter()
320 .map(|(name, value)| {
321 value
322 .to_str()
323 .map(|value| (name.as_str(), value))
324 .with_context(|| format!("header {name} is not valid UTF-8 and cannot be signed"))
325 })
326 .collect::<Result<Vec<_>>>()?;
327
328 let signable_request =
329 SignableRequest::new(method, uri, headers.into_iter(), SignableBody::Bytes(body))
330 .context("constructing Mantle SigV4 request")?;
331
332 let (instructions, _signature) = sign(signable_request, &signing_params)
333 .context("signing Mantle request with SigV4")?
334 .into_parts();
335 instructions.apply_to_request_http1x(request);
336
337 Ok(())
338}
339
340pub struct State {
341 /// The resolved authentication method. Settings take priority over UX credentials.
342 auth: Option<BedrockAuth>,
343 /// Raw settings from settings.json
344 settings: Option<AmazonBedrockSettings>,
345 /// Whether credentials came from environment variables (only relevant for static credentials)
346 credentials_from_env: bool,
347 credentials_provider: Arc<dyn CredentialsProvider>,
348 _subscription: Subscription,
349}
350
351impl State {
352 fn reset_auth(&self, cx: &mut Context<Self>) -> Task<Result<()>> {
353 let credentials_provider = self.credentials_provider.clone();
354 cx.spawn(async move |this, cx| {
355 credentials_provider
356 .delete_credentials(AMAZON_AWS_URL, cx)
357 .await
358 .log_err();
359 this.update(cx, |this, cx| {
360 this.auth = None;
361 this.credentials_from_env = false;
362 cx.notify();
363 })
364 })
365 }
366
367 fn set_static_credentials(
368 &mut self,
369 credentials: BedrockCredentials,
370 cx: &mut Context<Self>,
371 ) -> Task<Result<()>> {
372 let auth = credentials.clone().into_auth();
373 let credentials_provider = self.credentials_provider.clone();
374 cx.spawn(async move |this, cx| {
375 credentials_provider
376 .write_credentials(
377 AMAZON_AWS_URL,
378 "Bearer",
379 &serde_json::to_vec(&credentials)?,
380 cx,
381 )
382 .await?;
383 this.update(cx, |this, cx| {
384 this.auth = auth;
385 this.credentials_from_env = false;
386 cx.notify();
387 })
388 })
389 }
390
391 fn is_authenticated(&self) -> bool {
392 self.auth.is_some()
393 }
394
395 /// Resolve authentication. Settings take priority over UX-provided credentials.
396 fn authenticate(&self, cx: &mut Context<Self>) -> Task<Result<(), AuthenticateError>> {
397 if self.is_authenticated() {
398 return Task::ready(Ok(()));
399 }
400
401 // Step 1: Check if settings specify an auth method (enterprise control)
402 if let Some(settings) = &self.settings {
403 if let Some(method) = &settings.authentication_method {
404 let profile_name = settings
405 .profile_name
406 .clone()
407 .unwrap_or_else(|| "default".to_string());
408
409 let auth = match method {
410 BedrockAuthMethod::Automatic => BedrockAuth::Automatic,
411 BedrockAuthMethod::NamedProfile => BedrockAuth::NamedProfile { profile_name },
412 BedrockAuthMethod::SingleSignOn => BedrockAuth::SingleSignOn { profile_name },
413 BedrockAuthMethod::ApiKey => {
414 // ApiKey method means "use static credentials from keychain/env"
415 // Fall through to load them below
416 return self.load_static_credentials(cx);
417 }
418 };
419
420 return cx.spawn(async move |this, cx| {
421 this.update(cx, |this, cx| {
422 this.auth = Some(auth);
423 this.credentials_from_env = false;
424 cx.notify();
425 })?;
426 Ok(())
427 });
428 }
429 }
430
431 // Step 2: No settings auth method - try to load static credentials
432 self.load_static_credentials(cx)
433 }
434
435 /// Load static credentials from environment variables or keychain.
436 fn load_static_credentials(
437 &self,
438 cx: &mut Context<Self>,
439 ) -> Task<Result<(), AuthenticateError>> {
440 let credentials_provider = self.credentials_provider.clone();
441 cx.spawn(async move |this, cx| {
442 // Try environment variables first
443 let (auth, from_env) = if let Some(bearer_token) = &ZED_BEDROCK_BEARER_TOKEN_VAR.value {
444 if !bearer_token.is_empty() {
445 (
446 Some(BedrockAuth::ApiKey {
447 api_key: bearer_token.to_string(),
448 }),
449 true,
450 )
451 } else {
452 (None, false)
453 }
454 } else if let Some(access_key_id) = &ZED_BEDROCK_ACCESS_KEY_ID_VAR.value {
455 if let Some(secret_access_key) = &ZED_BEDROCK_SECRET_ACCESS_KEY_VAR.value {
456 if !access_key_id.is_empty() && !secret_access_key.is_empty() {
457 let session_token = ZED_BEDROCK_SESSION_TOKEN_VAR
458 .value
459 .as_deref()
460 .filter(|s| !s.is_empty())
461 .map(|s| s.to_string());
462 (
463 Some(BedrockAuth::IamCredentials {
464 access_key_id: access_key_id.to_string(),
465 secret_access_key: secret_access_key.to_string(),
466 session_token,
467 }),
468 true,
469 )
470 } else {
471 (None, false)
472 }
473 } else {
474 (None, false)
475 }
476 } else {
477 (None, false)
478 };
479
480 // If we got auth from env vars, use it
481 if let Some(auth) = auth {
482 this.update(cx, |this, cx| {
483 this.auth = Some(auth);
484 this.credentials_from_env = from_env;
485 cx.notify();
486 })?;
487 return Ok(());
488 }
489
490 // Try keychain
491 let (_, credentials_bytes) = credentials_provider
492 .read_credentials(AMAZON_AWS_URL, cx)
493 .await?
494 .ok_or(AuthenticateError::CredentialsNotFound)?;
495
496 let credentials_str = String::from_utf8(credentials_bytes)
497 .with_context(|| format!("invalid {PROVIDER_NAME} credentials"))?;
498
499 let credentials: BedrockCredentials =
500 serde_json::from_str(&credentials_str).context("failed to parse credentials")?;
501
502 let auth = credentials
503 .into_auth()
504 .ok_or(AuthenticateError::CredentialsNotFound)?;
505
506 this.update(cx, |this, cx| {
507 this.auth = Some(auth);
508 this.credentials_from_env = false;
509 cx.notify();
510 })?;
511
512 Ok(())
513 })
514 }
515
516 /// Get the resolved region. Checks env var, then settings, then defaults to us-east-1.
517 fn get_region(&self) -> String {
518 // Priority: env var > settings > default
519 if let Some(region) = ZED_BEDROCK_REGION_VAR.value.as_deref() {
520 if !region.is_empty() {
521 return region.to_string();
522 }
523 }
524
525 self.settings
526 .as_ref()
527 .and_then(|s| s.region.clone())
528 .unwrap_or_else(|| "us-east-1".to_string())
529 }
530
531 fn get_allow_global(&self) -> bool {
532 self.settings
533 .as_ref()
534 .and_then(|s| s.allow_global)
535 .unwrap_or(false)
536 }
537
538 fn get_guardrail_config(&self) -> (Option<String>, Option<String>) {
539 self.settings.as_ref().map_or((None, None), |s| {
540 (s.guardrail_identifier.clone(), s.guardrail_version.clone())
541 })
542 }
543}
544
545pub struct BedrockLanguageModelProvider {
546 http_client: AwsHttpClient,
547 plain_http_client: Arc<dyn HttpClient>,
548 handle: tokio::runtime::Handle,
549 state: Entity<State>,
550}
551
552impl BedrockLanguageModelProvider {
553 pub fn new(
554 http_client: Arc<dyn HttpClient>,
555 credentials_provider: Arc<dyn CredentialsProvider>,
556 cx: &mut App,
557 ) -> Self {
558 let state = cx.new(|cx| State {
559 auth: None,
560 settings: Some(AllLanguageModelSettings::get_global(cx).bedrock.clone()),
561 credentials_from_env: false,
562 credentials_provider,
563 _subscription: cx.observe_global::<SettingsStore>(|_, cx| {
564 cx.notify();
565 }),
566 });
567
568 Self {
569 http_client: AwsHttpClient::new(http_client.clone()),
570 plain_http_client: http_client,
571 handle: Tokio::handle(cx),
572 state,
573 }
574 }
575
576 fn create_language_model(&self, model: bedrock::ConverseModel) -> Arc<dyn LanguageModel> {
577 Arc::new(BedrockModel {
578 id: LanguageModelId::from(model.id().to_string()),
579 model,
580 http_client: self.http_client.clone(),
581 handle: self.handle.clone(),
582 state: self.state.clone(),
583 client: OnceCell::new(),
584 request_limiter: RateLimiter::new(4),
585 })
586 }
587
588 fn create_mantle_language_model(&self, model: bedrock::MantleModel) -> Arc<dyn LanguageModel> {
589 Arc::new(BedrockMantleModel {
590 id: LanguageModelId::from(model.id().to_string()),
591 model,
592 http_client: self.plain_http_client.clone(),
593 state: self.state.clone(),
594 credentials_provider: Arc::new(OnceCell::new()),
595 request_limiter: RateLimiter::new(4),
596 })
597 }
598}
599
600impl LanguageModelProvider for BedrockLanguageModelProvider {
601 fn id(&self) -> LanguageModelProviderId {
602 PROVIDER_ID
603 }
604
605 fn name(&self) -> LanguageModelProviderName {
606 PROVIDER_NAME
607 }
608
609 fn icon(&self) -> IconOrSvg {
610 IconOrSvg::Icon(IconName::AiBedrock)
611 }
612
613 fn default_model(&self, _cx: &App) -> Option<Arc<dyn LanguageModel>> {
614 Some(self.create_language_model(bedrock::ConverseModel::default()))
615 }
616
617 fn default_fast_model(&self, cx: &App) -> Option<Arc<dyn LanguageModel>> {
618 let region = self.state.read(cx).get_region();
619 Some(self.create_language_model(bedrock::ConverseModel::default_fast(region.as_str())))
620 }
621
622 fn provided_models(&self, cx: &App) -> Vec<Arc<dyn LanguageModel>> {
623 let bedrock_settings = &AllLanguageModelSettings::get_global(cx).bedrock;
624 let mut models = BTreeMap::default();
625
626 for model in bedrock::ConverseModel::iter() {
627 if !matches!(model, bedrock::ConverseModel::Custom { .. }) {
628 models.insert(model.id().to_string(), model);
629 }
630 }
631
632 // Override with available models from settings
633 for model in bedrock_settings.available_models.iter() {
634 models.insert(
635 model.name.clone(),
636 bedrock::ConverseModel::Custom {
637 name: model.name.clone(),
638 display_name: model.display_name.clone(),
639 max_tokens: model.max_tokens,
640 max_output_tokens: model.max_output_tokens,
641 default_temperature: model.default_temperature,
642 cache_configuration: model.cache_configuration.as_ref().map(|config| {
643 bedrock::BedrockModelCacheConfiguration {
644 max_cache_anchors: config.max_cache_anchors,
645 min_total_token: config.min_total_token,
646 }
647 }),
648 },
649 );
650 }
651
652 let mut models: Vec<Arc<dyn LanguageModel>> = models
653 .into_values()
654 .map(|model| self.create_language_model(model))
655 .collect();
656
657 let mut mantle_models = BTreeMap::default();
658
659 for model in bedrock::MantleModel::iter() {
660 if !matches!(model, bedrock::MantleModel::Custom { .. }) {
661 mantle_models.insert(model.id().to_string(), model);
662 }
663 }
664
665 // Override with available Mantle models from settings
666 for model in bedrock_settings.mantle_available_models.iter() {
667 mantle_models.insert(
668 model.name.clone(),
669 bedrock::MantleModel::Custom {
670 name: model.name.clone(),
671 display_name: model.display_name.clone(),
672 max_tokens: model.max_tokens,
673 max_output_tokens: model.max_output_tokens,
674 protocol: mantle_protocol_from_settings(model.protocol),
675 supports_tools: model.supports_tools.unwrap_or(false),
676 supports_images: model.supports_images.unwrap_or(false),
677 supports_thinking: model.supports_thinking.unwrap_or(false),
678 },
679 );
680 }
681
682 models.extend(
683 mantle_models
684 .into_values()
685 .map(|model| self.create_mantle_language_model(model)),
686 );
687
688 models
689 }
690
691 fn is_authenticated(&self, cx: &App) -> bool {
692 self.state.read(cx).is_authenticated()
693 }
694
695 fn authenticate(&self, cx: &mut App) -> Task<Result<(), AuthenticateError>> {
696 self.state.update(cx, |state, cx| state.authenticate(cx))
697 }
698
699 fn settings_view(&self, _cx: &mut App) -> Option<ProviderSettingsView> {
700 let state = self.state.clone();
701 Some(ProviderSettingsView::SubPage(
702 SubPageProviderSettings::new(move |window, cx| {
703 cx.new(|cx| ConfigurationView::new(state.clone(), window, cx))
704 .into()
705 })
706 .description(InlineDescription::Text(
707 "To use Omega Agent with Bedrock, set a custom authentication strategy in your settings or use static credentials. Mantle-only models (e.g. GPT-5.5, GPT-5.4, Grok 4.3) additionally require IAM permissions for the `bedrock-mantle` endpoint.".into(),
708 )),
709 ))
710 }
711}
712
713impl LanguageModelProviderState for BedrockLanguageModelProvider {
714 type ObservableEntity = State;
715
716 fn observable_entity(&self) -> Option<Entity<Self::ObservableEntity>> {
717 Some(self.state.clone())
718 }
719}
720
721struct BedrockModel {
722 id: LanguageModelId,
723 model: ConverseModel,
724 http_client: AwsHttpClient,
725 handle: tokio::runtime::Handle,
726 client: OnceCell<BedrockClient>,
727 state: Entity<State>,
728 request_limiter: RateLimiter,
729}
730
731impl BedrockModel {
732 fn get_or_init_client(&self, cx: &AsyncApp) -> anyhow::Result<&BedrockClient> {
733 self.client
734 .get_or_try_init_blocking(|| {
735 let (auth, endpoint, region) = cx.read_entity(&self.state, |state, _cx| {
736 let endpoint = state.settings.as_ref().and_then(|s| s.endpoint.clone());
737 let region = state.get_region();
738 (state.auth.clone(), endpoint, region)
739 });
740
741 let mut config_builder = aws_config::defaults(BehaviorVersion::latest())
742 .stalled_stream_protection(StalledStreamProtectionConfig::disabled())
743 .http_client(self.http_client.clone())
744 .region(Region::new(region))
745 .timeout_config(TimeoutConfig::disabled());
746
747 if let Some(endpoint_url) = endpoint
748 && !endpoint_url.is_empty()
749 {
750 config_builder = config_builder.endpoint_url(endpoint_url);
751 }
752
753 match auth {
754 Some(BedrockAuth::Automatic) | None => {
755 // Use default AWS credential provider chain
756 }
757 Some(BedrockAuth::NamedProfile { profile_name })
758 | Some(BedrockAuth::SingleSignOn { profile_name }) => {
759 if !profile_name.is_empty() {
760 config_builder = config_builder.profile_name(profile_name);
761 }
762 }
763 Some(BedrockAuth::IamCredentials {
764 access_key_id,
765 secret_access_key,
766 session_token,
767 }) => {
768 let aws_creds = Credentials::new(
769 access_key_id,
770 secret_access_key,
771 session_token,
772 None,
773 "zed-bedrock-provider",
774 );
775 config_builder = config_builder.credentials_provider(aws_creds);
776 }
777 Some(BedrockAuth::ApiKey { api_key }) => {
778 config_builder = config_builder
779 .auth_scheme_preference(["httpBearerAuth".into()]) // https://github.com/smithy-lang/smithy-rs/pull/4241
780 .token_provider(Token::new(api_key, None));
781 }
782 }
783
784 let config = self.handle.block_on(config_builder.load());
785
786 anyhow::Ok(BedrockClient::new(&config))
787 })
788 .context("initializing Bedrock client")?;
789
790 self.client.get().context("Bedrock client not initialized")
791 }
792
793 fn stream_completion(
794 &self,
795 request: bedrock::Request,
796 cx: &AsyncApp,
797 ) -> BoxFuture<
798 'static,
799 Result<BoxStream<'static, Result<BedrockStreamingResponse, anyhow::Error>>, BedrockError>,
800 > {
801 let Ok(runtime_client) = self
802 .get_or_init_client(cx)
803 .cloned()
804 .context("Bedrock client not initialized")
805 else {
806 return futures::future::ready(Err(BedrockError::Other(anyhow!("App state dropped"))))
807 .boxed();
808 };
809 let extra_headers = self.state.read_with(cx, |_, cx| {
810 AllLanguageModelSettings::get_global(cx)
811 .bedrock
812 .custom_headers
813 .clone()
814 });
815
816 let task = Tokio::spawn(
817 cx,
818 bedrock::stream_completion(runtime_client, request, extra_headers),
819 );
820 async move { task.await.map_err(|e| BedrockError::Other(e.into()))? }.boxed()
821 }
822}
823
824impl LanguageModel for BedrockModel {
825 fn id(&self) -> LanguageModelId {
826 self.id.clone()
827 }
828
829 fn name(&self) -> LanguageModelName {
830 LanguageModelName::from(self.model.display_name().to_string())
831 }
832
833 fn provider_id(&self) -> LanguageModelProviderId {
834 PROVIDER_ID
835 }
836
837 fn provider_name(&self) -> LanguageModelProviderName {
838 PROVIDER_NAME
839 }
840
841 fn supports_tools(&self) -> bool {
842 self.model.supports_tool_use()
843 }
844
845 fn supports_images(&self) -> bool {
846 self.model.supports_images()
847 }
848
849 fn supports_thinking(&self) -> bool {
850 self.model.supports_thinking()
851 }
852
853 fn refusal_fallback_model_id(&self) -> Option<&'static str> {
854 if self
855 .model
856 .id()
857 .starts_with(anthropic::FABLE_MODEL_ID_PREFIX)
858 {
859 Some(anthropic::FABLE_FALLBACK_MODEL_ID)
860 } else {
861 None
862 }
863 }
864
865 fn supported_effort_levels(&self) -> Vec<language_model::LanguageModelEffortLevel> {
866 if self.model.supports_adaptive_thinking() {
867 vec![
868 language_model::LanguageModelEffortLevel {
869 name: "Low".into(),
870 value: "low".into(),
871 is_default: false,
872 },
873 language_model::LanguageModelEffortLevel {
874 name: "Medium".into(),
875 value: "medium".into(),
876 is_default: false,
877 },
878 language_model::LanguageModelEffortLevel {
879 name: "High".into(),
880 value: "high".into(),
881 is_default: true,
882 },
883 language_model::LanguageModelEffortLevel {
884 name: "XHigh".into(),
885 value: "xhigh".into(),
886 is_default: false,
887 },
888 language_model::LanguageModelEffortLevel {
889 name: "Max".into(),
890 value: "max".into(),
891 is_default: false,
892 },
893 ]
894 .into_iter()
895 .filter(|effort_level| {
896 effort_level.value != "xhigh" || self.model.supports_xhigh_adaptive_thinking()
897 })
898 .collect()
899 } else {
900 Vec::new()
901 }
902 }
903
904 fn supports_tool_choice(&self, choice: LanguageModelToolChoice) -> bool {
905 match choice {
906 LanguageModelToolChoice::Auto | LanguageModelToolChoice::Any => {
907 self.model.supports_tool_use()
908 }
909 // Add support for None - we'll filter tool calls at response
910 LanguageModelToolChoice::None => self.model.supports_tool_use(),
911 }
912 }
913
914 fn supports_streaming_tools(&self) -> bool {
915 true
916 }
917
918 fn telemetry_id(&self) -> String {
919 format!("bedrock/{}", self.model.id())
920 }
921
922 fn max_token_count(&self) -> u64 {
923 self.model.max_token_count()
924 }
925
926 fn max_output_tokens(&self) -> Option<u64> {
927 Some(self.model.max_output_tokens())
928 }
929
930 fn stream_completion(
931 &self,
932 request: LanguageModelRequest,
933 cx: &AsyncApp,
934 ) -> BoxFuture<
935 'static,
936 Result<
937 BoxStream<'static, Result<LanguageModelCompletionEvent, LanguageModelCompletionError>>,
938 LanguageModelCompletionError,
939 >,
940 > {
941 if request.contains_custom_tool_input() {
942 return async move {
943 Err(anyhow::anyhow!("Bedrock does not support custom tools").into())
944 }
945 .boxed();
946 }
947
948 let (region, allow_global, guardrail_identifier, guardrail_version) =
949 cx.read_entity(&self.state, |state, _cx| {
950 let (gid, gv) = state.get_guardrail_config();
951 (state.get_region(), state.get_allow_global(), gid, gv)
952 });
953
954 let model_id = match self.model.cross_region_inference_id(®ion, allow_global) {
955 Ok(s) => s,
956 Err(e) => {
957 return async move { Err(e.into()) }.boxed();
958 }
959 };
960
961 let deny_tool_calls = request.tool_choice == Some(LanguageModelToolChoice::None);
962
963 let request = match into_bedrock(
964 request,
965 model_id,
966 self.model.default_temperature(),
967 self.model.max_output_tokens(),
968 self.model.thinking_mode(),
969 self.model.supports_caching(),
970 self.model.supports_tool_use(),
971 guardrail_identifier,
972 guardrail_version,
973 ) {
974 Ok(request) => request,
975 Err(err) => return futures::future::ready(Err(err.into())).boxed(),
976 };
977
978 let request = self.stream_completion(request, cx);
979 let display_name = self.model.display_name().to_string();
980 let future = self.request_limiter.stream(async move {
981 let response = request.await.map_err(|err| match err {
982 BedrockError::Validation(ref msg) => {
983 if msg.contains("model identifier is invalid") {
984 LanguageModelCompletionError::Other(anyhow!(
985 "{display_name} is not available in {region}. \
986 Try switching to a region where this model is supported."
987 ))
988 } else {
989 LanguageModelCompletionError::BadRequestFormat {
990 provider: PROVIDER_NAME,
991 message: msg.clone(),
992 }
993 }
994 }
995 BedrockError::RateLimited => LanguageModelCompletionError::RateLimitExceeded {
996 provider: PROVIDER_NAME,
997 retry_after: None,
998 },
999 BedrockError::ServiceUnavailable => {
1000 LanguageModelCompletionError::ServerOverloaded {
1001 provider: PROVIDER_NAME,
1002 retry_after: None,
1003 }
1004 }
1005 BedrockError::AccessDenied(msg) => LanguageModelCompletionError::PermissionError {
1006 provider: PROVIDER_NAME,
1007 message: msg,
1008 },
1009 BedrockError::InternalServer(msg) => {
1010 LanguageModelCompletionError::ApiInternalServerError {
1011 provider: PROVIDER_NAME,
1012 message: msg,
1013 }
1014 }
1015 other => LanguageModelCompletionError::Other(anyhow!(other)),
1016 })?;
1017 let events = map_to_language_model_completion_events(response);
1018
1019 if deny_tool_calls {
1020 Ok(deny_tool_use_events(events).boxed())
1021 } else {
1022 Ok(events.boxed())
1023 }
1024 });
1025
1026 async move { Ok(future.await?.boxed()) }.boxed()
1027 }
1028}
1029
1030const MANTLE_SELECTABLE_REASONING_EFFORTS: &[ReasoningEffort] = &[
1031 ReasoningEffort::Low,
1032 ReasoningEffort::Medium,
1033 ReasoningEffort::High,
1034 ReasoningEffort::XHigh,
1035];
1036
1037fn mantle_default_reasoning_effort(model: &MantleModel) -> Option<ReasoningEffort> {
1038 model.supports_thinking().then_some(ReasoningEffort::Medium)
1039}
1040
1041fn mantle_selected_reasoning_effort(
1042 request: &LanguageModelRequest,
1043 model: &MantleModel,
1044) -> Option<ReasoningEffort> {
1045 if !model.supports_thinking() {
1046 return None;
1047 }
1048
1049 if request.thinking_allowed {
1050 request
1051 .thinking_effort
1052 .as_deref()
1053 .and_then(|effort| effort.parse::<ReasoningEffort>().ok())
1054 .filter(|effort| *effort != ReasoningEffort::None)
1055 .or_else(|| mantle_default_reasoning_effort(model))
1056 } else {
1057 Some(ReasoningEffort::None)
1058 }
1059}
1060
1061fn mantle_supported_effort_levels(model: &MantleModel) -> Vec<LanguageModelEffortLevel> {
1062 let Some(default_effort) = mantle_default_reasoning_effort(model) else {
1063 return Vec::new();
1064 };
1065
1066 MANTLE_SELECTABLE_REASONING_EFFORTS
1067 .iter()
1068 .copied()
1069 .map(|effort| LanguageModelEffortLevel {
1070 name: effort.label().into(),
1071 value: effort.value().into(),
1072 is_default: effort == default_effort,
1073 })
1074 .collect()
1075}
1076
1077/// Special-cases Mantle authorization failures with a message that points at
1078/// the separate `bedrock-mantle` IAM policy namespace instead of regular
1079/// `bedrock-runtime` permissions.
1080fn map_mantle_error(model: &MantleModel, error: RequestError) -> LanguageModelCompletionError {
1081 if let RequestError::HttpResponseError { status_code, .. } = &error
1082 && *status_code == http_client::http::StatusCode::FORBIDDEN
1083 {
1084 return LanguageModelCompletionError::PermissionError {
1085 provider: PROVIDER_NAME,
1086 message: format!(
1087 "Bedrock Mantle denied this request for {}. Mantle-only models require IAM \
1088 permissions for the `bedrock-mantle` endpoint (for example via the \
1089 `AmazonBedrockMantleInferenceAccess` managed policy) in addition to whatever \
1090 permissions your existing Bedrock credentials already have.",
1091 model.display_name()
1092 ),
1093 };
1094 }
1095 error.into()
1096}
1097
1098/// Resolves an AWS credentials provider for profile/SSO/automatic auth.
1099/// Cached in `cell` since building it may read config files from disk;
1100/// credentials themselves are still re-resolved on every call. Async so this
1101/// never blocks the foreground thread (unlike `BedrockModel::get_or_init_client`).
1102async fn resolve_mantle_credentials_provider(
1103 cell: &OnceCell<SharedCredentialsProvider>,
1104 profile_name: Option<String>,
1105 region: String,
1106) -> Result<SharedCredentialsProvider> {
1107 let provider = cell
1108 .get_or_try_init(move || async move {
1109 let mut config_builder =
1110 aws_config::defaults(BehaviorVersion::latest()).region(Region::new(region));
1111
1112 if let Some(profile_name) = profile_name.filter(|name| !name.is_empty()) {
1113 config_builder = config_builder.profile_name(profile_name);
1114 }
1115
1116 let config = config_builder.load().await;
1117 config
1118 .credentials_provider()
1119 .context("no AWS credentials provider is configured")
1120 })
1121 .await
1122 .context("resolving AWS credentials for Bedrock Mantle")?;
1123 Ok(provider.clone())
1124}
1125
1126/// Resolves provider settings into concrete Mantle request auth. A configured
1127/// Bedrock API key is sent as bearer auth; every AWS-credential-based method
1128/// signs the Mantle HTTP request directly with SigV4.
1129async fn resolve_mantle_auth(
1130 credentials_provider: Arc<OnceCell<SharedCredentialsProvider>>,
1131 auth: Option<BedrockAuth>,
1132 region: String,
1133) -> Result<MantleAuth> {
1134 match auth {
1135 Some(BedrockAuth::ApiKey { api_key }) => Ok(MantleAuth::ApiKey { api_key }),
1136 Some(BedrockAuth::IamCredentials {
1137 access_key_id,
1138 secret_access_key,
1139 session_token,
1140 }) => Ok(MantleAuth::SigV4 {
1141 credentials: Credentials::new(
1142 access_key_id,
1143 secret_access_key,
1144 session_token,
1145 None,
1146 "zed-bedrock-provider",
1147 ),
1148 }),
1149 Some(BedrockAuth::NamedProfile { profile_name })
1150 | Some(BedrockAuth::SingleSignOn { profile_name }) => {
1151 let provider = resolve_mantle_credentials_provider(
1152 &credentials_provider,
1153 Some(profile_name),
1154 region.clone(),
1155 )
1156 .await?;
1157 let credentials = provider
1158 .provide_credentials()
1159 .await
1160 .context("failed to resolve AWS credentials")?;
1161 Ok(MantleAuth::SigV4 { credentials })
1162 }
1163 Some(BedrockAuth::Automatic) | None => {
1164 let provider =
1165 resolve_mantle_credentials_provider(&credentials_provider, None, region.clone())
1166 .await?;
1167 let credentials = provider
1168 .provide_credentials()
1169 .await
1170 .context("failed to resolve AWS credentials")?;
1171 Ok(MantleAuth::SigV4 { credentials })
1172 }
1173 }
1174}
1175
1176#[derive(Deserialize)]
1177#[serde(untagged)]
1178enum MantleChatStreamResult {
1179 Ok(ResponseStreamEvent),
1180 Err { error: MantleChatStreamError },
1181}
1182
1183#[derive(Deserialize)]
1184struct MantleChatStreamError {
1185 message: String,
1186}
1187
1188fn parse_mantle_chat_stream_line(line: &str) -> Result<ResponseStreamEvent> {
1189 match serde_json::from_str(line) {
1190 Ok(MantleChatStreamResult::Ok(response)) => Ok(response),
1191 Ok(MantleChatStreamResult::Err { error }) => Err(anyhow!(error.message)),
1192 Err(error) => {
1193 log::error!(
1194 "Failed to parse Mantle chat completion stream event: `{}`\nResponse: `{}`",
1195 error,
1196 line,
1197 );
1198 Err(anyhow!(error))
1199 }
1200 }
1201}
1202
1203fn parse_mantle_response_stream_line(line: &str) -> Result<open_ai::responses::StreamEvent> {
1204 serde_json::from_str(line).map_err(|error| {
1205 log::error!(
1206 "Failed to parse Mantle responses stream event: `{}`\nResponse: `{}`",
1207 error,
1208 line,
1209 );
1210 anyhow!(error)
1211 })
1212}
1213
1214async fn stream_mantle_sse<Request, Event>(
1215 client: &dyn HttpClient,
1216 provider_name: &str,
1217 url: &str,
1218 region: &str,
1219 auth: &MantleAuth,
1220 request: Request,
1221 extra_headers: &CustomHeaders,
1222 parse_stream_line: fn(&str) -> Result<Event>,
1223) -> std::result::Result<BoxStream<'static, Result<Event>>, RequestError>
1224where
1225 Request: Serialize,
1226 Event: Send + 'static,
1227{
1228 let body = serde_json::to_vec(&request).map_err(|error| RequestError::Other(error.into()))?;
1229 let mut request = HttpRequest::builder()
1230 .method(Method::POST)
1231 .uri(url)
1232 .header("Content-Type", "application/json")
1233 .extra_headers(extra_headers)
1234 .body(AsyncBody::from(body.clone()))
1235 .map_err(|error| RequestError::Other(error.into()))?;
1236
1237 auth.apply(&mut request, &body, region)
1238 .map_err(RequestError::Other)?;
1239
1240 let mut response = client.send(request).await?;
1241 if response.status().is_success() {
1242 let reader = BufReader::new(response.into_body());
1243 Ok(reader
1244 .lines()
1245 .filter_map(move |line| async move {
1246 match line {
1247 Ok(line) => {
1248 let line = line
1249 .strip_prefix("data: ")
1250 .or_else(|| line.strip_prefix("data:"))?;
1251 if line == "[DONE]" || line.is_empty() {
1252 None
1253 } else {
1254 Some(parse_stream_line(line))
1255 }
1256 }
1257 Err(error) => Some(Err(anyhow!(error))),
1258 }
1259 })
1260 .boxed())
1261 } else {
1262 let mut body = String::new();
1263 response
1264 .body_mut()
1265 .read_to_string(&mut body)
1266 .await
1267 .map_err(|error| RequestError::Other(error.into()))?;
1268
1269 Err(RequestError::HttpResponseError {
1270 provider: provider_name.to_owned(),
1271 status_code: response.status(),
1272 body,
1273 headers: response.headers().clone(),
1274 })
1275 }
1276}
1277
1278fn strip_unsupported_mantle_response_fields(request: &mut OpenAiResponseRequest) {
1279 request.context_management = None;
1280}
1281
1282#[derive(Debug, PartialEq, Eq, Clone, Copy)]
1283enum MantleMessageSnapshotKind {
1284 Undetermined,
1285 Cumulative,
1286 Independent,
1287}
1288
1289struct MantleMessageSnapshot {
1290 item_id: String,
1291 text: String,
1292 emitted_text_length: usize,
1293 kind: MantleMessageSnapshotKind,
1294 phase: Option<String>,
1295 output_index: usize,
1296 content_index: Option<usize>,
1297}
1298
1299/// The most recently completed message item, used to detect Mantle's cumulative
1300/// snapshot replays.
1301#[derive(Clone)]
1302struct MantlePreviousMessage {
1303 text: String,
1304 phase: Option<String>,
1305}
1306
1307/// Adapts Mantle's cumulative Responses message snapshots to incremental events.
1308///
1309/// Mantle can emit a new message item containing all text produced so far, then
1310/// replay the corresponding output deltas. The shared OpenAI mapper assumes each
1311/// delta is new text, so forwarding those events directly duplicates the reply.
1312///
1313/// Only adjacent, same-phase, strict-prefix extensions are collapsed. Equal,
1314/// shrinking, divergent, or different-phase messages stay visible. Any non-message
1315/// output item (reasoning, tool call, etc.) is a hard boundary: it is forwarded
1316/// unchanged and clears the snapshot used for collapsing, so a message after a
1317/// tool call is never merged with the one before it.
1318struct MantleResponseEventMapper {
1319 open_ai_mapper: OpenAiResponseEventMapper,
1320 current_message: Option<MantleMessageSnapshot>,
1321 previous_message: Option<MantlePreviousMessage>,
1322 pending_message_events: Vec<Result<LanguageModelCompletionEvent, LanguageModelCompletionError>>,
1323}
1324
1325impl MantleResponseEventMapper {
1326 fn new() -> Self {
1327 Self {
1328 open_ai_mapper: OpenAiResponseEventMapper::new(PROVIDER_ID),
1329 current_message: None,
1330 previous_message: None,
1331 pending_message_events: Vec::new(),
1332 }
1333 }
1334
1335 fn map_stream(
1336 mut self,
1337 events: BoxStream<'static, Result<OpenAiResponseStreamEvent>>,
1338 ) -> BoxStream<'static, Result<LanguageModelCompletionEvent, LanguageModelCompletionError>>
1339 {
1340 events
1341 .flat_map(move |event| {
1342 futures::stream::iter(match event {
1343 Ok(event) => self.map_event(event),
1344 Err(error) => {
1345 let mut events = self.finish_current_message();
1346 events.push(Err(LanguageModelCompletionError::from(error)));
1347 events
1348 }
1349 })
1350 })
1351 .boxed()
1352 }
1353
1354 fn map_event(
1355 &mut self,
1356 event: OpenAiResponseStreamEvent,
1357 ) -> Vec<Result<LanguageModelCompletionEvent, LanguageModelCompletionError>> {
1358 match event {
1359 OpenAiResponseStreamEvent::OutputItemAdded {
1360 output_index,
1361 sequence_number,
1362 item,
1363 } => match &item {
1364 ResponseOutputItem::Message(message) => {
1365 let mut events = self.finish_current_message();
1366 let item_id = message.id.clone();
1367 let phase = message.phase.clone();
1368 if let Some(item_id) = item_id {
1369 self.current_message = Some(MantleMessageSnapshot {
1370 item_id,
1371 text: String::new(),
1372 emitted_text_length: 0,
1373 kind: MantleMessageSnapshotKind::Undetermined,
1374 phase,
1375 output_index,
1376 content_index: None,
1377 });
1378
1379 self.map_message_added(OpenAiResponseStreamEvent::OutputItemAdded {
1380 output_index,
1381 sequence_number,
1382 item,
1383 });
1384 } else {
1385 self.flush_pending_message(&mut events, false);
1386 events.extend(self.open_ai_mapper.map_event(
1387 OpenAiResponseStreamEvent::OutputItemAdded {
1388 output_index,
1389 sequence_number,
1390 item,
1391 },
1392 ));
1393 }
1394 events
1395 }
1396 _ => {
1397 let mut events = self.finish_current_message();
1398 self.previous_message = None;
1399 events.extend(self.open_ai_mapper.map_event(
1400 OpenAiResponseStreamEvent::OutputItemAdded {
1401 output_index,
1402 sequence_number,
1403 item,
1404 },
1405 ));
1406 events
1407 }
1408 },
1409 OpenAiResponseStreamEvent::OutputTextDelta {
1410 item_id,
1411 output_index,
1412 content_index,
1413 delta,
1414 } => self.map_text_delta(item_id, output_index, content_index, delta),
1415 OpenAiResponseStreamEvent::OutputTextDone {
1416 item_id,
1417 output_index,
1418 content_index,
1419 text,
1420 } => {
1421 let suffix = self
1422 .current_message
1423 .as_ref()
1424 .filter(|message| message.item_id == item_id)
1425 .and_then(|message| text.strip_prefix(&message.text))
1426 .map(str::to_string);
1427 if let Some(suffix) = suffix {
1428 if suffix.is_empty() {
1429 return Vec::new();
1430 }
1431 return self.map_text_delta(item_id, output_index, content_index, suffix);
1432 }
1433
1434 let mut events = self.finish_current_message();
1435 events.extend(self.open_ai_mapper.map_event(
1436 OpenAiResponseStreamEvent::OutputTextDone {
1437 item_id,
1438 output_index,
1439 content_index,
1440 text,
1441 },
1442 ));
1443 events
1444 }
1445 OpenAiResponseStreamEvent::ContentPartAdded { ref item_id, .. }
1446 | OpenAiResponseStreamEvent::ContentPartDone { ref item_id, .. }
1447 | OpenAiResponseStreamEvent::RefusalDelta { ref item_id, .. }
1448 | OpenAiResponseStreamEvent::RefusalDone { ref item_id, .. }
1449 if self
1450 .current_message
1451 .as_ref()
1452 .is_some_and(|message| &message.item_id == item_id) =>
1453 {
1454 self.open_ai_mapper.map_event(event)
1455 }
1456 OpenAiResponseStreamEvent::OutputItemDone {
1457 output_index,
1458 sequence_number,
1459 item,
1460 } => {
1461 if let ResponseOutputItem::Message(message) = &item
1462 && let Some(item_id) = message.id.as_ref()
1463 && let Some(current_message) = self.current_message.as_mut()
1464 && ¤t_message.item_id == item_id
1465 {
1466 current_message.phase = message.phase.clone();
1467 }
1468
1469 let is_non_message = !matches!(item, ResponseOutputItem::Message(_));
1470 let was_merged = matches!(
1471 self.current_message.as_ref(),
1472 Some(current_message)
1473 if current_message.kind == MantleMessageSnapshotKind::Cumulative
1474 );
1475 let mut events = self.finish_current_message();
1476 if is_non_message {
1477 self.previous_message = None;
1478 }
1479 let done_events =
1480 self.open_ai_mapper
1481 .map_event(OpenAiResponseStreamEvent::OutputItemDone {
1482 output_index,
1483 sequence_number,
1484 item,
1485 });
1486 if was_merged {
1487 // The message was merged into the previous one, so the phase/
1488 // reasoning metadata this `OutputItemDone` would otherwise emit
1489 // again is already reflected by the metadata emitted earlier.
1490 events.extend(done_events.into_iter().filter(|event| {
1491 !matches!(event, Ok(LanguageModelCompletionEvent::ReasoningDetails(_)))
1492 }));
1493 } else {
1494 events.extend(done_events);
1495 }
1496 events
1497 }
1498 event => {
1499 let mut events = self.finish_current_message();
1500 events.extend(self.open_ai_mapper.map_event(event));
1501 events
1502 }
1503 }
1504 }
1505
1506 fn map_message_added(&mut self, event: OpenAiResponseStreamEvent) {
1507 debug_assert!(
1508 self.pending_message_events.is_empty(),
1509 "pending message events must be flushed before a new message is added"
1510 );
1511 self.pending_message_events = self.open_ai_mapper.map_event(event);
1512 }
1513
1514 fn map_text_delta(
1515 &mut self,
1516 item_id: String,
1517 output_index: usize,
1518 content_index: Option<usize>,
1519 delta: String,
1520 ) -> Vec<Result<LanguageModelCompletionEvent, LanguageModelCompletionError>> {
1521 if !self
1522 .current_message
1523 .as_ref()
1524 .is_some_and(|current_message| current_message.item_id == item_id)
1525 {
1526 let mut events = self.finish_current_message();
1527 events.extend(self.open_ai_mapper.map_event(
1528 OpenAiResponseStreamEvent::OutputTextDelta {
1529 item_id,
1530 output_index,
1531 content_index,
1532 delta,
1533 },
1534 ));
1535 return events;
1536 }
1537
1538 let Some(current_message) = self.current_message.as_mut() else {
1539 return Vec::new();
1540 };
1541 current_message.text.push_str(&delta);
1542 current_message.output_index = output_index;
1543 current_message.content_index = content_index;
1544 let kind = current_message.kind;
1545
1546 match kind {
1547 MantleMessageSnapshotKind::Undetermined => {
1548 let snapshot = current_message.text.clone();
1549 let phase = current_message.phase.clone();
1550 let previous = self.previous_message.clone();
1551
1552 if let Some(previous) = previous
1553 && previous.phase == phase
1554 {
1555 // Strict same-phase extension: previous text is a proper prefix of
1556 // the new snapshot. Collapse the replayed prefix and emit only the
1557 // suffix.
1558 if snapshot.starts_with(&previous.text) && snapshot.len() > previous.text.len()
1559 {
1560 if let Some(current_message) = self.current_message.as_mut() {
1561 current_message.kind = MantleMessageSnapshotKind::Cumulative;
1562 current_message.emitted_text_length = previous.text.len();
1563 }
1564 let mut events = Vec::new();
1565 self.flush_pending_message(&mut events, true);
1566 events.extend(self.map_snapshot_suffix(output_index, content_index));
1567 return events;
1568 }
1569
1570 // The snapshot is still a prefix of (or equal to) the previous
1571 // message. A genuine cumulative replay keeps growing past the
1572 // previous text, so defer the decision and withhold emission
1573 // until it either exceeds the previous text or diverges.
1574 if previous.text.starts_with(&snapshot) {
1575 return Vec::new();
1576 }
1577 }
1578
1579 // Diverged from the previous message, had a different phase, or there
1580 // was no previous message: this is an independent message.
1581 if let Some(current_message) = self.current_message.as_mut() {
1582 current_message.kind = MantleMessageSnapshotKind::Independent;
1583 }
1584 let mut events = Vec::new();
1585 self.flush_pending_message(&mut events, false);
1586 events.extend(self.map_text_event(item_id, output_index, content_index, snapshot));
1587 events
1588 }
1589 MantleMessageSnapshotKind::Cumulative => {
1590 self.map_snapshot_suffix(output_index, content_index)
1591 }
1592 MantleMessageSnapshotKind::Independent => {
1593 self.map_text_event(item_id, output_index, content_index, delta)
1594 }
1595 }
1596 }
1597
1598 fn map_snapshot_suffix(
1599 &mut self,
1600 output_index: usize,
1601 content_index: Option<usize>,
1602 ) -> Vec<Result<LanguageModelCompletionEvent, LanguageModelCompletionError>> {
1603 let (item_id, suffix) = {
1604 let Some(current_message) = self.current_message.as_mut() else {
1605 return Vec::new();
1606 };
1607
1608 if current_message.text.len() <= current_message.emitted_text_length {
1609 return Vec::new();
1610 }
1611
1612 let suffix = current_message.text[current_message.emitted_text_length..].to_string();
1613 current_message.emitted_text_length = current_message.text.len();
1614 (current_message.item_id.clone(), suffix)
1615 };
1616
1617 self.map_text_event(item_id, output_index, content_index, suffix)
1618 }
1619
1620 fn map_text_event(
1621 &mut self,
1622 item_id: String,
1623 output_index: usize,
1624 content_index: Option<usize>,
1625 delta: String,
1626 ) -> Vec<Result<LanguageModelCompletionEvent, LanguageModelCompletionError>> {
1627 if delta.is_empty() {
1628 return Vec::new();
1629 }
1630
1631 self.open_ai_mapper
1632 .map_event(OpenAiResponseStreamEvent::OutputTextDelta {
1633 item_id,
1634 output_index,
1635 content_index,
1636 delta,
1637 })
1638 }
1639
1640 fn finish_current_message(
1641 &mut self,
1642 ) -> Vec<Result<LanguageModelCompletionEvent, LanguageModelCompletionError>> {
1643 let mut events = Vec::new();
1644 self.finish_current_message_into(&mut events);
1645 events
1646 }
1647
1648 fn finish_current_message_into(
1649 &mut self,
1650 events: &mut Vec<Result<LanguageModelCompletionEvent, LanguageModelCompletionError>>,
1651 ) {
1652 let Some(current_message) = self.current_message.take() else {
1653 self.flush_pending_message(events, false);
1654 return;
1655 };
1656
1657 match current_message.kind {
1658 MantleMessageSnapshotKind::Cumulative => {
1659 self.flush_pending_message(events, true);
1660 }
1661 MantleMessageSnapshotKind::Independent => {
1662 self.flush_pending_message(events, false);
1663 }
1664 MantleMessageSnapshotKind::Undetermined => {
1665 // The message never exceeded or diverged from the previous one, so
1666 // it could not be confirmed as a cumulative replay. Treat it as an
1667 // independent message: emit the withheld StartMessage and any text
1668 // accumulated while deferring the decision.
1669 self.flush_pending_message(events, false);
1670 if !current_message.text.is_empty() {
1671 events.extend(self.map_text_event(
1672 current_message.item_id.clone(),
1673 current_message.output_index,
1674 current_message.content_index,
1675 current_message.text.clone(),
1676 ));
1677 }
1678 }
1679 }
1680
1681 if !current_message.text.is_empty() {
1682 self.previous_message = Some(MantlePreviousMessage {
1683 text: current_message.text,
1684 phase: current_message.phase,
1685 });
1686 }
1687 }
1688
1689 fn flush_pending_message(
1690 &mut self,
1691 events: &mut Vec<Result<LanguageModelCompletionEvent, LanguageModelCompletionError>>,
1692 suppress_start_message: bool,
1693 ) {
1694 for event in self.pending_message_events.drain(..) {
1695 // When merging into the previous message, the phase/reasoning metadata
1696 // can't have changed since it was last emitted (a merge only happens for
1697 // adjacent, same-phase messages with no reasoning item in between), so
1698 // suppress both the StartMessage and the now-redundant ReasoningDetails.
1699 if suppress_start_message
1700 && matches!(
1701 event,
1702 Ok(LanguageModelCompletionEvent::StartMessage { .. })
1703 | Ok(LanguageModelCompletionEvent::ReasoningDetails(_))
1704 )
1705 {
1706 continue;
1707 }
1708 events.push(event);
1709 }
1710 }
1711}
1712
1713struct BedrockMantleModel {
1714 id: LanguageModelId,
1715 model: MantleModel,
1716 http_client: Arc<dyn HttpClient>,
1717 state: Entity<State>,
1718 credentials_provider: Arc<OnceCell<SharedCredentialsProvider>>,
1719 request_limiter: RateLimiter,
1720}
1721
1722impl BedrockMantleModel {
1723 fn stream_mantle_request<Request, Event>(
1724 &self,
1725 request: Request,
1726 cx: &AsyncApp,
1727 endpoint: &'static str,
1728 parse_stream_line: fn(&str) -> Result<Event>,
1729 ) -> BoxFuture<'static, Result<BoxStream<'static, Result<Event>>, LanguageModelCompletionError>>
1730 where
1731 Request: Serialize + Send + 'static,
1732 Event: Send + 'static,
1733 {
1734 let http_client = self.http_client.clone();
1735 let model = self.model.clone();
1736 let credentials_provider = self.credentials_provider.clone();
1737 let (auth, region) = cx.read_entity(&self.state, |state, _cx| {
1738 (state.auth.clone(), state.get_region())
1739 });
1740 let url = format!("{}/{}", mantle_endpoint_url(®ion), endpoint);
1741 let extra_headers = cx.read_entity(&self.state, |_, cx| {
1742 AllLanguageModelSettings::get_global(cx)
1743 .bedrock
1744 .custom_headers
1745 .clone()
1746 });
1747 let provider_name = PROVIDER_NAME.0.to_string();
1748 let auth_task = Tokio::spawn_result(
1749 cx,
1750 resolve_mantle_auth(credentials_provider, auth, region.clone()),
1751 );
1752
1753 let future = self.request_limiter.stream(async move {
1754 let auth = auth_task
1755 .await
1756 .map_err(LanguageModelCompletionError::Other)?;
1757 stream_mantle_sse(
1758 http_client.as_ref(),
1759 &provider_name,
1760 &url,
1761 ®ion,
1762 &auth,
1763 request,
1764 &extra_headers,
1765 parse_stream_line,
1766 )
1767 .await
1768 .map_err(|err| map_mantle_error(&model, err))
1769 });
1770
1771 async move { Ok(future.await?.boxed()) }.boxed()
1772 }
1773
1774 fn stream_completion(
1775 &self,
1776 request: open_ai::Request,
1777 cx: &AsyncApp,
1778 ) -> BoxFuture<
1779 'static,
1780 Result<BoxStream<'static, Result<ResponseStreamEvent>>, LanguageModelCompletionError>,
1781 > {
1782 self.stream_mantle_request(
1783 request,
1784 cx,
1785 "chat/completions",
1786 parse_mantle_chat_stream_line,
1787 )
1788 }
1789
1790 fn stream_response(
1791 &self,
1792 request: OpenAiResponseRequest,
1793 cx: &AsyncApp,
1794 ) -> BoxFuture<
1795 'static,
1796 Result<
1797 BoxStream<'static, Result<open_ai::responses::StreamEvent>>,
1798 LanguageModelCompletionError,
1799 >,
1800 > {
1801 let mut request = request;
1802 strip_unsupported_mantle_response_fields(&mut request);
1803 self.stream_mantle_request(request, cx, "responses", parse_mantle_response_stream_line)
1804 }
1805}
1806
1807impl LanguageModel for BedrockMantleModel {
1808 fn id(&self) -> LanguageModelId {
1809 self.id.clone()
1810 }
1811
1812 fn name(&self) -> LanguageModelName {
1813 LanguageModelName::from(self.model.display_name().to_string())
1814 }
1815
1816 fn provider_id(&self) -> LanguageModelProviderId {
1817 PROVIDER_ID
1818 }
1819
1820 fn provider_name(&self) -> LanguageModelProviderName {
1821 PROVIDER_NAME
1822 }
1823
1824 fn supports_tools(&self) -> bool {
1825 self.model.supports_tools()
1826 }
1827
1828 fn tool_input_format(&self) -> LanguageModelToolSchemaFormat {
1829 LanguageModelToolSchemaFormat::JsonSchemaSubset
1830 }
1831
1832 fn supports_images(&self) -> bool {
1833 self.model.supports_images()
1834 }
1835
1836 fn supports_tool_choice(&self, choice: LanguageModelToolChoice) -> bool {
1837 match choice {
1838 LanguageModelToolChoice::Auto | LanguageModelToolChoice::Any => {
1839 self.model.supports_tools()
1840 }
1841 LanguageModelToolChoice::None => true,
1842 }
1843 }
1844
1845 fn supports_streaming_tools(&self) -> bool {
1846 true
1847 }
1848
1849 fn supports_thinking(&self) -> bool {
1850 self.model.supports_thinking()
1851 }
1852
1853 fn supported_effort_levels(&self) -> Vec<LanguageModelEffortLevel> {
1854 mantle_supported_effort_levels(&self.model)
1855 }
1856
1857 fn supports_split_token_display(&self) -> bool {
1858 true
1859 }
1860
1861 fn telemetry_id(&self) -> String {
1862 format!("bedrock-mantle/{}", self.model.id())
1863 }
1864
1865 fn max_token_count(&self) -> u64 {
1866 self.model.max_token_count()
1867 }
1868
1869 fn max_output_tokens(&self) -> Option<u64> {
1870 Some(self.model.max_output_tokens())
1871 }
1872
1873 fn stream_completion(
1874 &self,
1875 request: LanguageModelRequest,
1876 cx: &AsyncApp,
1877 ) -> BoxFuture<
1878 'static,
1879 Result<
1880 BoxStream<'static, Result<LanguageModelCompletionEvent, LanguageModelCompletionError>>,
1881 LanguageModelCompletionError,
1882 >,
1883 > {
1884 let region = cx.read_entity(&self.state, |state, _cx| state.get_region());
1885
1886 if !MANTLE_SUPPORTED_REGIONS.contains(®ion.as_str()) {
1887 let display_name = self.model.display_name().to_string();
1888 let supported = MANTLE_SUPPORTED_REGIONS.join(", ");
1889 return futures::future::ready(Err(LanguageModelCompletionError::Other(anyhow!(
1890 "{display_name} is not available in {region} because Bedrock Mantle isn't offered \
1891 there. Try switching to one of the following regions: {supported}."
1892 ))))
1893 .boxed();
1894 }
1895
1896 let model_id = self.model.request_id().to_string();
1897 let max_output_tokens = Some(self.model.max_output_tokens());
1898
1899 match self.model.protocol() {
1900 MantleProtocol::Responses => {
1901 let request = match into_open_ai_response(
1902 request,
1903 &model_id,
1904 self.model.supports_tools(),
1905 false,
1906 max_output_tokens,
1907 mantle_default_reasoning_effort(&self.model),
1908 self.model.supports_thinking(),
1909 &PROVIDER_ID,
1910 ) {
1911 Ok(request) => request,
1912 Err(error) => return async move { Err(error.into()) }.boxed(),
1913 };
1914 let completions = self.stream_response(request, cx);
1915 async move {
1916 let mapper = MantleResponseEventMapper::new();
1917 Ok(mapper.map_stream(completions.await?).boxed())
1918 }
1919 .boxed()
1920 }
1921 MantleProtocol::ChatCompletions => {
1922 let reasoning_effort = mantle_selected_reasoning_effort(&request, &self.model);
1923 let request = match into_open_ai(
1924 request,
1925 &model_id,
1926 self.model.supports_tools(),
1927 false,
1928 max_output_tokens,
1929 ChatCompletionMaxTokensParameter::MaxCompletionTokens,
1930 reasoning_effort,
1931 false,
1932 ) {
1933 Ok(request) => request,
1934 Err(error) => return async move { Err(error.into()) }.boxed(),
1935 };
1936 let completions = self.stream_completion(request, cx);
1937 async move {
1938 let mapper = OpenAiEventMapper::new();
1939 Ok(mapper.map_stream(completions.await?).boxed())
1940 }
1941 .boxed()
1942 }
1943 }
1944 }
1945}
1946
1947fn deny_tool_use_events(
1948 events: impl Stream<Item = Result<LanguageModelCompletionEvent, LanguageModelCompletionError>>,
1949) -> impl Stream<Item = Result<LanguageModelCompletionEvent, LanguageModelCompletionError>> {
1950 events.map(|event| {
1951 match event {
1952 Ok(LanguageModelCompletionEvent::ToolUse(tool_use)) => {
1953 // Convert tool use to an error message if model decided to call it
1954 Ok(LanguageModelCompletionEvent::Text(format!(
1955 "\n\n[Error: Tool calls are disabled in this context. Attempted to call '{}']",
1956 tool_use.name
1957 )))
1958 }
1959 other => other,
1960 }
1961 })
1962}
1963
1964pub fn into_bedrock(
1965 request: LanguageModelRequest,
1966 model: String,
1967 default_temperature: f32,
1968 max_output_tokens: u64,
1969 thinking_mode: BedrockModelMode,
1970 supports_caching: bool,
1971 supports_tool_use: bool,
1972 guardrail_identifier: Option<String>,
1973 guardrail_version: Option<String>,
1974) -> Result<bedrock::Request> {
1975 if request.contains_custom_tool_input() {
1976 anyhow::bail!("Bedrock does not support custom tools");
1977 }
1978
1979 let mut new_messages: Vec<BedrockMessage> = Vec::new();
1980 let mut system_message = String::new();
1981
1982 // Track whether messages contain tool content - Bedrock requires toolConfig
1983 // when tool blocks are present, so we may need to add a dummy tool
1984 let mut messages_contain_tool_content = false;
1985
1986 for message in request.messages {
1987 if message.contents_empty() {
1988 continue;
1989 }
1990
1991 match message.role {
1992 Role::User | Role::Assistant => {
1993 let mut bedrock_message_content: Vec<BedrockInnerContent> = message
1994 .content
1995 .into_iter()
1996 .filter_map(|content| match content {
1997 MessageContent::Text(text) => {
1998 if !text.is_empty() {
1999 Some(BedrockInnerContent::Text(text))
2000 } else {
2001 None
2002 }
2003 }
2004 MessageContent::Compaction(_) => None,
2005 MessageContent::Thinking { text, signature } => {
2006 if model.contains(ConverseModel::DeepSeekR1.request_id()) {
2007 // DeepSeekR1 doesn't support thinking blocks
2008 // And the AWS API demands that you strip them
2009 return None;
2010 }
2011 if signature.is_none() {
2012 // Thinking blocks without a signature are invalid
2013 // (e.g. from cancellation mid-think) and must be
2014 // stripped to avoid API errors.
2015 return None;
2016 }
2017 let thinking = BedrockThinkingTextBlock::builder()
2018 .text(text)
2019 .set_signature(signature)
2020 .build()
2021 .context("failed to build reasoning block")
2022 .log_err()?;
2023
2024 Some(BedrockInnerContent::ReasoningContent(
2025 BedrockThinkingBlock::ReasoningText(thinking),
2026 ))
2027 }
2028 MessageContent::RedactedThinking(blob) => {
2029 if model.contains(ConverseModel::DeepSeekR1.request_id()) {
2030 // DeepSeekR1 doesn't support thinking blocks
2031 // And the AWS API demands that you strip them
2032 return None;
2033 }
2034 let redacted =
2035 BedrockThinkingBlock::RedactedContent(BedrockBlob::new(blob));
2036
2037 Some(BedrockInnerContent::ReasoningContent(redacted))
2038 }
2039 MessageContent::ToolUse(tool_use) => {
2040 messages_contain_tool_content = true;
2041 let input =
2042 if let language_model::LanguageModelToolUseInput::Json(input) =
2043 &tool_use.input
2044 {
2045 if input.is_null() {
2046 // Bedrock API requires valid JsonValue, not null, for tool use input
2047 value_to_aws_document(&serde_json::json!({}))
2048 } else {
2049 value_to_aws_document(input)
2050 }
2051 } else {
2052 value_to_aws_document(&serde_json::json!({}))
2053 };
2054 BedrockToolUseBlock::builder()
2055 .name(tool_use.name.to_string())
2056 .tool_use_id(tool_use.id.to_string())
2057 .input(input)
2058 .build()
2059 .context("failed to build Bedrock tool use block")
2060 .log_err()
2061 .map(BedrockInnerContent::ToolUse)
2062 }
2063 MessageContent::ToolResult(tool_result) => {
2064 messages_contain_tool_content = true;
2065 let mut builder = BedrockToolResultBlock::builder()
2066 .tool_use_id(tool_result.tool_use_id.to_string());
2067 for part in tool_result.content {
2068 let block = match part {
2069 LanguageModelToolResultContent::Text(text) => {
2070 BedrockToolResultContentBlock::Text(text.to_string())
2071 }
2072 LanguageModelToolResultContent::Image(image) => {
2073 use base64::Engine;
2074
2075 match base64::engine::general_purpose::STANDARD
2076 .decode(image.source.as_bytes())
2077 {
2078 Ok(image_bytes) => {
2079 match BedrockImageBlock::builder()
2080 .format(BedrockImageFormat::Png)
2081 .source(BedrockImageSource::Bytes(
2082 BedrockBlob::new(image_bytes),
2083 ))
2084 .build()
2085 {
2086 Ok(image_block) => {
2087 BedrockToolResultContentBlock::Image(
2088 image_block,
2089 )
2090 }
2091 Err(err) => {
2092 BedrockToolResultContentBlock::Text(
2093 format!(
2094 "[Failed to build image block: {}]",
2095 err
2096 ),
2097 )
2098 }
2099 }
2100 }
2101 Err(err) => {
2102 BedrockToolResultContentBlock::Text(format!(
2103 "[Failed to decode tool result image: {}]",
2104 err
2105 ))
2106 }
2107 }
2108 }
2109 };
2110 builder = builder.content(block);
2111 }
2112 builder
2113 .status({
2114 if tool_result.is_error {
2115 BedrockToolResultStatus::Error
2116 } else {
2117 BedrockToolResultStatus::Success
2118 }
2119 })
2120 .build()
2121 .context("failed to build Bedrock tool result block")
2122 .log_err()
2123 .map(BedrockInnerContent::ToolResult)
2124 }
2125 MessageContent::Image(image) => {
2126 use base64::Engine;
2127
2128 let image_bytes = base64::engine::general_purpose::STANDARD
2129 .decode(image.source.as_bytes())
2130 .context("failed to decode base64 image data")
2131 .log_err()?;
2132
2133 BedrockImageBlock::builder()
2134 .format(BedrockImageFormat::Png)
2135 .source(BedrockImageSource::Bytes(BedrockBlob::new(image_bytes)))
2136 .build()
2137 .context("failed to build Bedrock image block")
2138 .log_err()
2139 .map(BedrockInnerContent::Image)
2140 }
2141 })
2142 .collect();
2143 if message.cache && supports_caching && !bedrock_message_content.is_empty() {
2144 bedrock_message_content.push(BedrockInnerContent::CachePoint(
2145 CachePointBlock::builder()
2146 .r#type(CachePointType::Default)
2147 .build()
2148 .context("failed to build cache point block")?,
2149 ));
2150 }
2151 let bedrock_role = match message.role {
2152 Role::User => bedrock::BedrockRole::User,
2153 Role::Assistant => bedrock::BedrockRole::Assistant,
2154 Role::System => unreachable!("System role should never occur here"),
2155 };
2156 if bedrock_message_content.is_empty() {
2157 continue;
2158 }
2159
2160 if let Some(last_message) = new_messages.last_mut()
2161 && last_message.role == bedrock_role
2162 {
2163 last_message.content.extend(bedrock_message_content);
2164 continue;
2165 }
2166 new_messages.push(
2167 BedrockMessage::builder()
2168 .role(bedrock_role)
2169 .set_content(Some(bedrock_message_content))
2170 .build()
2171 .context("failed to build Bedrock message")?,
2172 );
2173 }
2174 Role::System => {
2175 if !system_message.is_empty() {
2176 system_message.push_str("\n\n");
2177 }
2178 system_message.push_str(&message.string_contents());
2179 }
2180 }
2181 }
2182
2183 let mut tool_spec: Vec<BedrockTool> = if supports_tool_use {
2184 request
2185 .tools
2186 .iter()
2187 .map(|tool| {
2188 let language_model::LanguageModelRequestToolInput::Function {
2189 input_schema, ..
2190 } = &tool.input
2191 else {
2192 anyhow::bail!("Bedrock does not support custom tools");
2193 };
2194 Ok(BedrockTool::ToolSpec(
2195 BedrockToolSpec::builder()
2196 .name(tool.name.clone())
2197 .description(tool.description.clone())
2198 .input_schema(BedrockToolInputSchema::Json(value_to_aws_document(
2199 input_schema,
2200 )))
2201 .build()
2202 .context("failed to build Bedrock tool spec")?,
2203 ))
2204 })
2205 .collect::<Result<_>>()?
2206 } else {
2207 Vec::new()
2208 };
2209
2210 // Bedrock requires toolConfig when messages contain tool use/result blocks.
2211 // If no tools are defined but messages contain tool content (e.g., when
2212 // summarising a conversation that used tools), add a dummy tool to satisfy
2213 // the API requirement.
2214 if supports_tool_use && tool_spec.is_empty() && messages_contain_tool_content {
2215 tool_spec.push(BedrockTool::ToolSpec(
2216 BedrockToolSpec::builder()
2217 .name("_placeholder")
2218 .description("Placeholder tool to satisfy Bedrock API requirements when conversation history contains tool usage")
2219 .input_schema(BedrockToolInputSchema::Json(value_to_aws_document(
2220 &serde_json::json!({"type": "object", "properties": {}}),
2221 )))
2222 .build()
2223 .context("failed to build placeholder tool spec")?,
2224 ));
2225 }
2226
2227 if !tool_spec.is_empty() && supports_caching {
2228 tool_spec.push(BedrockTool::CachePoint(
2229 CachePointBlock::builder()
2230 .r#type(CachePointType::Default)
2231 .build()
2232 .context("failed to build cache point block")?,
2233 ));
2234 }
2235
2236 let tool_choice = match request.tool_choice {
2237 Some(LanguageModelToolChoice::Auto) | None => {
2238 BedrockToolChoice::Auto(BedrockAutoToolChoice::builder().build())
2239 }
2240 Some(LanguageModelToolChoice::Any) => {
2241 BedrockToolChoice::Any(BedrockAnyToolChoice::builder().build())
2242 }
2243 Some(LanguageModelToolChoice::None) => {
2244 // For None, we still use Auto but will filter out tool calls in the response
2245 BedrockToolChoice::Auto(BedrockAutoToolChoice::builder().build())
2246 }
2247 };
2248 let tool_config = if tool_spec.is_empty() {
2249 None
2250 } else {
2251 Some(
2252 BedrockToolConfig::builder()
2253 .set_tools(Some(tool_spec))
2254 .tool_choice(tool_choice)
2255 .build()?,
2256 )
2257 };
2258
2259 let mut system_blocks: Vec<BedrockSystemContentBlock> = Vec::new();
2260 if !system_message.is_empty() {
2261 system_blocks.push(BedrockSystemContentBlock::Text(system_message));
2262 if supports_caching {
2263 system_blocks.push(BedrockSystemContentBlock::CachePoint(
2264 CachePointBlock::builder()
2265 .r#type(CachePointType::Default)
2266 .build()
2267 .context("failed to build system cache point block")?,
2268 ));
2269 }
2270 }
2271
2272 let thinking = if request.thinking_allowed {
2273 match thinking_mode {
2274 BedrockModelMode::Thinking { budget_tokens } => {
2275 Some(bedrock::Thinking::Enabled { budget_tokens })
2276 }
2277 BedrockModelMode::AdaptiveThinking {
2278 effort: default_effort,
2279 } => {
2280 let effort = request
2281 .thinking_effort
2282 .as_deref()
2283 .and_then(|e| match e {
2284 "low" => Some(bedrock::BedrockAdaptiveThinkingEffort::Low),
2285 "medium" => Some(bedrock::BedrockAdaptiveThinkingEffort::Medium),
2286 "high" => Some(bedrock::BedrockAdaptiveThinkingEffort::High),
2287 "xhigh" => Some(bedrock::BedrockAdaptiveThinkingEffort::XHigh),
2288 "max" => Some(bedrock::BedrockAdaptiveThinkingEffort::Max),
2289 _ => None,
2290 })
2291 .unwrap_or(default_effort);
2292 Some(bedrock::Thinking::Adaptive { effort })
2293 }
2294 BedrockModelMode::Default => None,
2295 }
2296 } else if model.contains(ConverseModel::ClaudeOpus5.request_id()) {
2297 // On Claude Opus 5, omitting the `thinking` field no longer means
2298 // "off": the model runs adaptive thinking by default, so features
2299 // that suppress thinking (e.g. inline assist) must opt out
2300 // explicitly. Earlier Claude models treat omission as "off" and must
2301 // keep omitting the field. No effort accompanies the opt-out because
2302 // `disabled` combined with effort `xhigh`/`max` is a 400.
2303 // <https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-anthropic-claude-opus-5.html>
2304 Some(bedrock::Thinking::Disabled)
2305 } else {
2306 None
2307 };
2308
2309 Ok(bedrock::Request {
2310 model,
2311 messages: new_messages,
2312 max_tokens: max_output_tokens,
2313 system: system_blocks,
2314 tools: tool_config,
2315 thinking,
2316 metadata: None,
2317 stop_sequences: Vec::new(),
2318 temperature: request.temperature.or(Some(default_temperature)),
2319 top_k: None,
2320 top_p: None,
2321 guardrail_identifier,
2322 guardrail_version,
2323 })
2324}
2325
2326pub fn map_to_language_model_completion_events(
2327 events: Pin<Box<dyn Send + Stream<Item = Result<BedrockStreamingResponse, anyhow::Error>>>>,
2328) -> impl Stream<Item = Result<LanguageModelCompletionEvent, LanguageModelCompletionError>> {
2329 struct RawToolUse {
2330 id: String,
2331 name: String,
2332 input_json: String,
2333 }
2334
2335 struct State {
2336 events: Pin<Box<dyn Send + Stream<Item = Result<BedrockStreamingResponse, anyhow::Error>>>>,
2337 tool_uses_by_index: HashMap<i32, RawToolUse>,
2338 emitted_tool_use: bool,
2339 }
2340
2341 let initial_state = State {
2342 events,
2343 tool_uses_by_index: HashMap::default(),
2344 emitted_tool_use: false,
2345 };
2346
2347 futures::stream::unfold(initial_state, |mut state| async move {
2348 match state.events.next().await {
2349 Some(event_result) => match event_result {
2350 Ok(event) => {
2351 let result = match event {
2352 ConverseStreamOutput::ContentBlockDelta(cb_delta) => match cb_delta.delta {
2353 Some(ContentBlockDelta::Text(text)) => {
2354 Some(Ok(LanguageModelCompletionEvent::Text(text)))
2355 }
2356 Some(ContentBlockDelta::ToolUse(tool_output)) => {
2357 if let Some(tool_use) = state
2358 .tool_uses_by_index
2359 .get_mut(&cb_delta.content_block_index)
2360 {
2361 tool_use.input_json.push_str(tool_output.input());
2362 if let Ok(input) = serde_json::from_str::<serde_json::Value>(
2363 &fix_streamed_json(&tool_use.input_json),
2364 ) {
2365 Some(Ok(LanguageModelCompletionEvent::ToolUse(
2366 LanguageModelToolUse {
2367 id: tool_use.id.clone().into(),
2368 name: tool_use.name.clone().into(),
2369 is_input_complete: false,
2370 raw_input: tool_use.input_json.clone(),
2371 input:
2372 language_model::LanguageModelToolUseInput::Json(
2373 input,
2374 ),
2375 thought_signature: None,
2376 },
2377 )))
2378 } else {
2379 None
2380 }
2381 } else {
2382 None
2383 }
2384 }
2385 Some(ContentBlockDelta::ReasoningContent(thinking)) => match thinking {
2386 ReasoningContentBlockDelta::Text(thoughts) => {
2387 Some(Ok(LanguageModelCompletionEvent::Thinking {
2388 text: thoughts,
2389 signature: None,
2390 }))
2391 }
2392 ReasoningContentBlockDelta::Signature(sig) => {
2393 Some(Ok(LanguageModelCompletionEvent::Thinking {
2394 text: "".into(),
2395 signature: Some(sig),
2396 }))
2397 }
2398 ReasoningContentBlockDelta::RedactedContent(redacted) => {
2399 let content = String::from_utf8(redacted.into_inner())
2400 .unwrap_or("REDACTED".to_string());
2401 Some(Ok(LanguageModelCompletionEvent::Thinking {
2402 text: content,
2403 signature: None,
2404 }))
2405 }
2406 _ => None,
2407 },
2408 _ => None,
2409 },
2410 ConverseStreamOutput::ContentBlockStart(cb_start) => {
2411 if let Some(ContentBlockStart::ToolUse(tool_start)) = cb_start.start {
2412 state.tool_uses_by_index.insert(
2413 cb_start.content_block_index,
2414 RawToolUse {
2415 id: tool_start.tool_use_id,
2416 name: tool_start.name,
2417 input_json: String::new(),
2418 },
2419 );
2420 }
2421 None
2422 }
2423 ConverseStreamOutput::MessageStart(_) => None,
2424 ConverseStreamOutput::ContentBlockStop(cb_stop) => state
2425 .tool_uses_by_index
2426 .remove(&cb_stop.content_block_index)
2427 .map(|tool_use| {
2428 state.emitted_tool_use = true;
2429
2430 let input = parse_tool_arguments(&tool_use.input_json)
2431 .unwrap_or_else(|_| Value::Object(Default::default()));
2432
2433 Ok(LanguageModelCompletionEvent::ToolUse(
2434 LanguageModelToolUse {
2435 id: tool_use.id.into(),
2436 name: tool_use.name.into(),
2437 is_input_complete: true,
2438 raw_input: tool_use.input_json,
2439 input: language_model::LanguageModelToolUseInput::Json(
2440 input,
2441 ),
2442 thought_signature: None,
2443 },
2444 ))
2445 }),
2446 ConverseStreamOutput::Metadata(cb_meta) => cb_meta.usage.map(|metadata| {
2447 Ok(LanguageModelCompletionEvent::UsageUpdate(TokenUsage {
2448 input_tokens: metadata.input_tokens as u64,
2449 output_tokens: metadata.output_tokens as u64,
2450 cache_creation_input_tokens: metadata
2451 .cache_write_input_tokens
2452 .unwrap_or_default()
2453 as u64,
2454 cache_read_input_tokens: metadata
2455 .cache_read_input_tokens
2456 .unwrap_or_default()
2457 as u64,
2458 }))
2459 }),
2460 ConverseStreamOutput::MessageStop(message_stop) => {
2461 let stop_reason = if state.emitted_tool_use {
2462 // Some models (e.g. Kimi) send EndTurn even when
2463 // they've made tool calls. Trust the content over
2464 // the stop reason.
2465 language_model::StopReason::ToolUse
2466 } else {
2467 match message_stop.stop_reason {
2468 StopReason::ToolUse => language_model::StopReason::ToolUse,
2469 _ => language_model::StopReason::EndTurn,
2470 }
2471 };
2472 Some(Ok(LanguageModelCompletionEvent::Stop(stop_reason)))
2473 }
2474 _ => None,
2475 };
2476
2477 Some((result, state))
2478 }
2479 Err(err) => Some((
2480 Some(Err(LanguageModelCompletionError::Other(anyhow!(err)))),
2481 state,
2482 )),
2483 },
2484 None => None,
2485 }
2486 })
2487 .filter_map(|result| async move { result })
2488}
2489
2490struct ConfigurationView {
2491 access_key_id_editor: Entity<InputField>,
2492 secret_access_key_editor: Entity<InputField>,
2493 session_token_editor: Entity<InputField>,
2494 bearer_token_editor: Entity<InputField>,
2495 state: Entity<State>,
2496 load_credentials_task: Option<Task<()>>,
2497 focus_handle: FocusHandle,
2498}
2499
2500impl ConfigurationView {
2501 const PLACEHOLDER_ACCESS_KEY_ID_TEXT: &'static str = "XXXXXXXXXXXXXXXX";
2502 const PLACEHOLDER_SECRET_ACCESS_KEY_TEXT: &'static str =
2503 "XXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXX";
2504 const PLACEHOLDER_SESSION_TOKEN_TEXT: &'static str = "XXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXX";
2505 const PLACEHOLDER_BEARER_TOKEN_TEXT: &'static str = "XXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXX";
2506
2507 fn new(state: Entity<State>, window: &mut Window, cx: &mut Context<Self>) -> Self {
2508 let focus_handle = cx.focus_handle();
2509
2510 cx.observe(&state, |_, _, cx| {
2511 cx.notify();
2512 })
2513 .detach();
2514
2515 let access_key_id_editor = cx.new(|cx| {
2516 InputField::new(window, cx, Self::PLACEHOLDER_ACCESS_KEY_ID_TEXT)
2517 .label("Access Key ID")
2518 .tab_index(0)
2519 .tab_stop(true)
2520 });
2521
2522 let secret_access_key_editor = cx.new(|cx| {
2523 InputField::new(window, cx, Self::PLACEHOLDER_SECRET_ACCESS_KEY_TEXT)
2524 .label("Secret Access Key")
2525 .tab_index(1)
2526 .tab_stop(true)
2527 });
2528
2529 let session_token_editor = cx.new(|cx| {
2530 InputField::new(window, cx, Self::PLACEHOLDER_SESSION_TOKEN_TEXT)
2531 .label("Session Token (Optional)")
2532 .tab_index(2)
2533 .tab_stop(true)
2534 });
2535
2536 let bearer_token_editor = cx.new(|cx| {
2537 InputField::new(window, cx, Self::PLACEHOLDER_BEARER_TOKEN_TEXT)
2538 .label("Bedrock API Key")
2539 .tab_index(3)
2540 .tab_stop(true)
2541 });
2542
2543 let load_credentials_task = Some(cx.spawn({
2544 let state = state.clone();
2545 async move |this, cx| {
2546 if let Some(task) = Some(state.update(cx, |state, cx| state.authenticate(cx))) {
2547 // We don't log an error, because "not signed in" is also an error.
2548 let _ = task.await;
2549 }
2550 this.update(cx, |this, cx| {
2551 this.load_credentials_task = None;
2552 cx.notify();
2553 })
2554 .log_err();
2555 }
2556 }));
2557
2558 Self {
2559 access_key_id_editor,
2560 secret_access_key_editor,
2561 session_token_editor,
2562 bearer_token_editor,
2563 state,
2564 load_credentials_task,
2565 focus_handle,
2566 }
2567 }
2568
2569 fn save_credentials(
2570 &mut self,
2571 _: &menu::Confirm,
2572 _window: &mut Window,
2573 cx: &mut Context<Self>,
2574 ) {
2575 let access_key_id = self
2576 .access_key_id_editor
2577 .read(cx)
2578 .text(cx)
2579 .trim()
2580 .to_string();
2581 let secret_access_key = self
2582 .secret_access_key_editor
2583 .read(cx)
2584 .text(cx)
2585 .trim()
2586 .to_string();
2587 let session_token = self
2588 .session_token_editor
2589 .read(cx)
2590 .text(cx)
2591 .trim()
2592 .to_string();
2593 let session_token = if session_token.is_empty() {
2594 None
2595 } else {
2596 Some(session_token)
2597 };
2598 let bearer_token = self
2599 .bearer_token_editor
2600 .read(cx)
2601 .text(cx)
2602 .trim()
2603 .to_string();
2604 let bearer_token = if bearer_token.is_empty() {
2605 None
2606 } else {
2607 Some(bearer_token)
2608 };
2609
2610 let state = self.state.clone();
2611 cx.spawn(async move |_, cx| {
2612 state
2613 .update(cx, |state, cx| {
2614 let credentials = BedrockCredentials {
2615 access_key_id,
2616 secret_access_key,
2617 session_token,
2618 bearer_token,
2619 };
2620
2621 state.set_static_credentials(credentials, cx)
2622 })
2623 .await
2624 })
2625 .detach_and_log_err(cx);
2626 }
2627
2628 fn reset_credentials(&mut self, window: &mut Window, cx: &mut Context<Self>) {
2629 self.access_key_id_editor
2630 .update(cx, |editor, cx| editor.set_text("", window, cx));
2631 self.secret_access_key_editor
2632 .update(cx, |editor, cx| editor.set_text("", window, cx));
2633 self.session_token_editor
2634 .update(cx, |editor, cx| editor.set_text("", window, cx));
2635 self.bearer_token_editor
2636 .update(cx, |editor, cx| editor.set_text("", window, cx));
2637
2638 let state = self.state.clone();
2639 cx.spawn(async move |_, cx| state.update(cx, |state, cx| state.reset_auth(cx)).await)
2640 .detach_and_log_err(cx);
2641 }
2642
2643 fn on_tab(&mut self, _: &menu::SelectNext, window: &mut Window, cx: &mut Context<Self>) {
2644 window.focus_next(cx);
2645 }
2646
2647 fn on_tab_prev(
2648 &mut self,
2649 _: &menu::SelectPrevious,
2650 window: &mut Window,
2651 cx: &mut Context<Self>,
2652 ) {
2653 window.focus_prev(cx);
2654 }
2655}
2656
2657impl Render for ConfigurationView {
2658 fn render(&mut self, _window: &mut Window, cx: &mut Context<Self>) -> impl IntoElement {
2659 let state = self.state.read(cx);
2660 let env_var_set = state.credentials_from_env;
2661 let auth = state.auth.clone();
2662 let settings_auth_method = state
2663 .settings
2664 .as_ref()
2665 .and_then(|s| s.authentication_method.clone());
2666
2667 if self.load_credentials_task.is_some() {
2668 return div().child(Label::new("Loading credentials...")).into_any();
2669 }
2670
2671 let configured_label = match &auth {
2672 Some(BedrockAuth::Automatic) => {
2673 "Using automatic credentials (AWS default chain)".into()
2674 }
2675 Some(BedrockAuth::NamedProfile { profile_name }) => {
2676 format!("Using AWS profile: {profile_name}")
2677 }
2678 Some(BedrockAuth::SingleSignOn { profile_name }) => {
2679 format!("Using AWS SSO profile: {profile_name}")
2680 }
2681 Some(BedrockAuth::IamCredentials { .. }) if env_var_set => {
2682 format!(
2683 "Using IAM credentials from {} and {} environment variables",
2684 ZED_BEDROCK_ACCESS_KEY_ID_VAR.name, ZED_BEDROCK_SECRET_ACCESS_KEY_VAR.name
2685 )
2686 }
2687 Some(BedrockAuth::IamCredentials { .. }) => "Using IAM credentials".into(),
2688 Some(BedrockAuth::ApiKey { .. }) if env_var_set => {
2689 format!(
2690 "Using Bedrock API Key from {} environment variable",
2691 ZED_BEDROCK_BEARER_TOKEN_VAR.name
2692 )
2693 }
2694 Some(BedrockAuth::ApiKey { .. }) => "Using Bedrock API Key".into(),
2695 None => "Not authenticated".into(),
2696 };
2697
2698 // Determine if credentials can be reset
2699 // Settings-derived auth (non-ApiKey) cannot be reset from UI
2700 let is_settings_derived = matches!(
2701 settings_auth_method,
2702 Some(BedrockAuthMethod::Automatic)
2703 | Some(BedrockAuthMethod::NamedProfile)
2704 | Some(BedrockAuthMethod::SingleSignOn)
2705 );
2706
2707 let tooltip_label = if env_var_set {
2708 Some(format!(
2709 "To reset your credentials, unset the {}, {}, and {} or {} environment variables.",
2710 ZED_BEDROCK_ACCESS_KEY_ID_VAR.name,
2711 ZED_BEDROCK_SECRET_ACCESS_KEY_VAR.name,
2712 ZED_BEDROCK_SESSION_TOKEN_VAR.name,
2713 ZED_BEDROCK_BEARER_TOKEN_VAR.name
2714 ))
2715 } else if is_settings_derived {
2716 Some(
2717 "Authentication method is configured in settings. Edit settings.json to change."
2718 .to_string(),
2719 )
2720 } else {
2721 None
2722 };
2723
2724 let credentials_control = if self.state.read(cx).is_authenticated() {
2725 ConfiguredApiCard::new("bedrock-reset", configured_label)
2726 .disabled(env_var_set || is_settings_derived)
2727 .on_click(cx.listener(|this, _, window, cx| this.reset_credentials(window, cx)))
2728 .when_some(tooltip_label, |this, label| this.tooltip_label(label))
2729 .into_any_element()
2730 } else {
2731 self.render_static_credentials_ui().into_any_element()
2732 };
2733
2734 v_flex()
2735 .min_w_0()
2736 .w_full()
2737 .track_focus(&self.focus_handle)
2738 .on_action(cx.listener(Self::on_tab))
2739 .on_action(cx.listener(Self::on_tab_prev))
2740 .on_action(cx.listener(ConfigurationView::save_credentials))
2741 .gap_1()
2742 .child(Headline::new("Amazon Bedrock").size(HeadlineSize::Small))
2743 .child(
2744 Label::new(
2745 "To use Omega Agent with Bedrock, you can set a custom authentication strategy through your settings file or use static credentials.",
2746 )
2747 .color(Color::Muted),
2748 )
2749 .child(
2750 Label::new("But first, to access models on AWS, you need to:")
2751 .mt_1()
2752 .color(Color::Muted),
2753 )
2754 .child(
2755 List::new()
2756 .child(
2757 ListBulletItem::new("")
2758 .child(
2759 Label::new(
2760 "Grant permissions to the strategy you'll use according to the:",
2761 )
2762 .color(Color::Muted),
2763 )
2764 .child(ButtonLink::new(
2765 "Prerequisites",
2766 "https://docs.aws.amazon.com/bedrock/latest/userguide/inference-prereq.html",
2767 )),
2768 )
2769 .child(
2770 ListBulletItem::new("")
2771 .child(
2772 Label::new("Select the models you would like access to:")
2773 .color(Color::Muted),
2774 )
2775 .child(ButtonLink::new(
2776 "Bedrock Model Catalog",
2777 "https://us-east-1.console.aws.amazon.com/bedrock/home?region=us-east-1#/model-catalog",
2778 )),
2779 ),
2780 )
2781 .child(credentials_control)
2782 .into_any()
2783 }
2784}
2785
2786impl ConfigurationView {
2787 fn render_static_credentials_ui(&self) -> impl IntoElement {
2788 let list_item = List::new()
2789 .child(
2790 ListBulletItem::new("")
2791 .child(
2792 Label::new(
2793 "For access keys: Create an IAM user in the AWS console with programmatic access",
2794 )
2795 .color(Color::Muted),
2796 )
2797 .child(ButtonLink::new(
2798 "IAM Console",
2799 "https://us-east-1.console.aws.amazon.com/iam/home?region=us-east-1#/users",
2800 )),
2801 )
2802 .child(
2803 ListBulletItem::new("")
2804 .child(
2805 Label::new("For Bedrock API Keys: Generate an API key from the")
2806 .color(Color::Muted),
2807 )
2808 .child(ButtonLink::new(
2809 "Bedrock Console",
2810 "https://docs.aws.amazon.com/bedrock/latest/userguide/api-keys-use.html",
2811 )),
2812 )
2813 .child(
2814 ListBulletItem::new("")
2815 .child(
2816 Label::new("Attach the necessary Bedrock permissions to")
2817 .color(Color::Muted),
2818 )
2819 .child(ButtonLink::new(
2820 "this user",
2821 "https://docs.aws.amazon.com/bedrock/latest/userguide/inference-prereq.html",
2822 )),
2823 )
2824 .child(
2825 ListBulletItem::new(
2826 "Enter either access keys OR a Bedrock API Key below (not both)",
2827 )
2828 .label_color(Color::Muted),
2829 );
2830
2831 v_flex()
2832 .my_2()
2833 .tab_group()
2834 .gap_1p5()
2835 .child(Divider::horizontal())
2836 .child(Label::new("Static Credentials").mt_2())
2837 .child(
2838 Label::new(
2839 "This method uses your AWS access key ID and secret access key, or a Bedrock API Key.",
2840 )
2841 .color(Color::Muted),
2842 )
2843 .child(list_item)
2844 .child(
2845 v_flex()
2846 .gap_1()
2847 .child(self.access_key_id_editor.clone())
2848 .child(self.secret_access_key_editor.clone())
2849 .child(self.session_token_editor.clone()),
2850 )
2851 .child(
2852 Label::new(format!(
2853 "You can also set the {}, {} and {} environment variables (or {} for Bedrock API Key authentication) and restart Omega.",
2854 ZED_BEDROCK_ACCESS_KEY_ID_VAR.name,
2855 ZED_BEDROCK_SECRET_ACCESS_KEY_VAR.name,
2856 ZED_BEDROCK_REGION_VAR.name,
2857 ZED_BEDROCK_BEARER_TOKEN_VAR.name
2858 ))
2859 .size(LabelSize::Small)
2860 .color(Color::Muted),
2861 )
2862 .child(
2863 Label::new(format!(
2864 "Optionally, if your environment uses AWS CLI profiles, you can set {}; if it requires a custom endpoint, you can set {}; and if it requires a Session Token, you can set {}.",
2865 ZED_AWS_PROFILE_VAR.name,
2866 ZED_AWS_ENDPOINT_VAR.name,
2867 ZED_BEDROCK_SESSION_TOKEN_VAR.name
2868 ))
2869 .size(LabelSize::Small)
2870 .color(Color::Muted)
2871 .mt_1()
2872 .mb_2p5(),
2873 )
2874 .child(Divider::horizontal())
2875 .child(Label::new("Using the API key").mt_2().mb_1())
2876 .child(self.bearer_token_editor.clone())
2877 .child(
2878 Label::new(format!(
2879 "Region is configured via {} environment variable or settings.json (defaults to us-east-1).",
2880 ZED_BEDROCK_REGION_VAR.name
2881 ))
2882 .size(LabelSize::Small)
2883 .color(Color::Muted)
2884 )
2885 }
2886}
2887
2888#[cfg(test)]
2889mod tests {
2890 use super::*;
2891 use language_model::LanguageModelRequestMessage;
2892 use open_ai::responses::{
2893 ResponseFunctionToolCall, ResponseOutputMessage, ResponseReasoningItem,
2894 };
2895
2896 fn into_bedrock_request(messages: Vec<LanguageModelRequestMessage>) -> bedrock::Request {
2897 into_bedrock(
2898 LanguageModelRequest {
2899 messages,
2900 ..Default::default()
2901 },
2902 "claude-sonnet-4-5".to_string(),
2903 1.0,
2904 4096,
2905 BedrockModelMode::Default,
2906 true,
2907 true,
2908 None,
2909 None,
2910 )
2911 .unwrap()
2912 }
2913
2914 #[test]
2915 fn test_thinking_disallowed_sends_explicit_opt_out_only_on_opus_5() {
2916 // Claude Opus 5 runs adaptive thinking by default when the `thinking`
2917 // field is omitted, so suppressing thinking requires an explicit
2918 // `disabled` opt-out. Earlier Claude models treat omission as "off".
2919 for (model, expects_explicit_opt_out) in [
2920 ("us.anthropic.claude-opus-5", true),
2921 ("global.anthropic.claude-opus-5", true),
2922 ("us.anthropic.claude-opus-4-8", false),
2923 ] {
2924 let request = into_bedrock(
2925 LanguageModelRequest {
2926 messages: vec![LanguageModelRequestMessage {
2927 role: Role::User,
2928 content: vec![MessageContent::Text("Hi".into())],
2929 cache: false,
2930 reasoning_details: None,
2931 }],
2932 thinking_allowed: false,
2933 ..Default::default()
2934 },
2935 model.to_string(),
2936 1.0,
2937 128_000,
2938 BedrockModelMode::AdaptiveThinking {
2939 effort: bedrock::BedrockAdaptiveThinkingEffort::High,
2940 },
2941 true,
2942 true,
2943 None,
2944 None,
2945 )
2946 .unwrap();
2947
2948 if expects_explicit_opt_out {
2949 assert!(
2950 matches!(request.thinking, Some(bedrock::Thinking::Disabled)),
2951 "{model} should send an explicit thinking opt-out"
2952 );
2953 } else {
2954 assert!(
2955 request.thinking.is_none(),
2956 "{model} should omit the thinking field entirely"
2957 );
2958 }
2959 }
2960 }
2961
2962 #[test]
2963 fn test_cache_marked_message_that_filters_to_empty_is_dropped() {
2964 let request = into_bedrock_request(vec![
2965 LanguageModelRequestMessage {
2966 role: Role::User,
2967 content: vec![MessageContent::Text("What's the weather?".into())],
2968 cache: false,
2969 reasoning_details: None,
2970 },
2971 LanguageModelRequestMessage {
2972 role: Role::Assistant,
2973 content: vec![MessageContent::Thinking {
2974 text: "Let me think about this...".into(),
2975 signature: None,
2976 }],
2977 cache: true,
2978 reasoning_details: None,
2979 },
2980 LanguageModelRequestMessage {
2981 role: Role::User,
2982 content: vec![MessageContent::Text("Summarize this conversation.".into())],
2983 cache: false,
2984 reasoning_details: None,
2985 },
2986 ]);
2987
2988 for message in &request.messages {
2989 assert!(
2990 message
2991 .content()
2992 .iter()
2993 .any(|block| !matches!(block, BedrockInnerContent::CachePoint(_))),
2994 "message must not consist solely of cache points: {:?}",
2995 message
2996 );
2997 }
2998 assert!(
2999 request
3000 .messages
3001 .iter()
3002 .all(|message| *message.role() == bedrock::BedrockRole::User),
3003 "the assistant message stripped to empty content should be dropped entirely"
3004 );
3005 }
3006
3007 #[test]
3008 fn test_cache_marked_message_with_content_gets_cache_point() {
3009 let request = into_bedrock_request(vec![LanguageModelRequestMessage {
3010 role: Role::User,
3011 content: vec![MessageContent::Text("What's the weather?".into())],
3012 cache: true,
3013 reasoning_details: None,
3014 }]);
3015
3016 assert_eq!(request.messages.len(), 1);
3017 assert!(
3018 matches!(
3019 request.messages[0].content().last(),
3020 Some(BedrockInnerContent::CachePoint(_))
3021 ),
3022 "a cache-marked message with content should end with a cache point"
3023 );
3024 }
3025
3026 #[test]
3027 fn test_sign_mantle_request_sigv4_uses_mantle_service() {
3028 let credentials = Credentials::new(
3029 "AKIDEXAMPLE",
3030 "wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY",
3031 None,
3032 None,
3033 "test",
3034 );
3035 let body = br#"{"model":"openai.gpt-5.5"}"#;
3036 let mut request = HttpRequest::builder()
3037 .method(Method::POST)
3038 .uri("https://bedrock-mantle.us-east-1.api.aws/openai/v1/responses")
3039 .header("Content-Type", "application/json")
3040 .body(AsyncBody::from(body.to_vec()))
3041 .unwrap();
3042 let time = std::time::UNIX_EPOCH + std::time::Duration::from_secs(1_700_000_000);
3043
3044 sign_mantle_request_sigv4_at(&mut request, body, &credentials, "us-east-1", time).unwrap();
3045
3046 assert_eq!(
3047 request
3048 .headers()
3049 .get(http_client::http::header::HOST)
3050 .and_then(|value| value.to_str().ok()),
3051 Some("bedrock-mantle.us-east-1.api.aws")
3052 );
3053 assert_eq!(
3054 request
3055 .headers()
3056 .get("x-amz-date")
3057 .and_then(|value| value.to_str().ok()),
3058 Some("20231114T221320Z")
3059 );
3060 let authorization = request
3061 .headers()
3062 .get(AUTHORIZATION)
3063 .and_then(|value| value.to_str().ok())
3064 .unwrap();
3065 assert!(authorization.starts_with("AWS4-HMAC-SHA256 "));
3066 assert!(
3067 authorization
3068 .contains("Credential=AKIDEXAMPLE/20231114/us-east-1/bedrock-mantle/aws4_request")
3069 );
3070 assert!(authorization.contains("SignedHeaders=content-type;host"));
3071 assert!(authorization.contains("Signature="));
3072 }
3073
3074 #[test]
3075 fn test_mantle_endpoint_url_uses_openai_path_prefix() {
3076 assert_eq!(
3077 mantle_endpoint_url("us-east-1"),
3078 "https://bedrock-mantle.us-east-1.api.aws/openai/v1"
3079 );
3080 assert_eq!(
3081 mantle_endpoint_url("us-west-2"),
3082 "https://bedrock-mantle.us-west-2.api.aws/openai/v1"
3083 );
3084 }
3085
3086 #[test]
3087 fn test_mantle_protocol_from_settings() {
3088 assert_eq!(
3089 mantle_protocol_from_settings(settings::BedrockMantleProtocolContent::ChatCompletions),
3090 MantleProtocol::ChatCompletions
3091 );
3092 assert_eq!(
3093 mantle_protocol_from_settings(settings::BedrockMantleProtocolContent::Responses),
3094 MantleProtocol::Responses
3095 );
3096 }
3097
3098 #[test]
3099 fn test_mantle_supported_regions_matches_docs() {
3100 assert!(MANTLE_SUPPORTED_REGIONS.contains(&"us-east-1"));
3101 assert!(MANTLE_SUPPORTED_REGIONS.contains(&"eu-west-1"));
3102 assert!(!MANTLE_SUPPORTED_REGIONS.contains(&"ap-southeast-1"));
3103 }
3104
3105 fn mantle_message_item(id: &str) -> ResponseOutputItem {
3106 mantle_message_item_with_phase(id, None)
3107 }
3108
3109 fn mantle_message_item_with_phase(id: &str, phase: Option<&str>) -> ResponseOutputItem {
3110 ResponseOutputItem::Message(ResponseOutputMessage {
3111 id: Some(id.to_string()),
3112 content: Vec::new(),
3113 role: Some("assistant".to_string()),
3114 status: Some("in_progress".to_string()),
3115 phase: phase.map(str::to_string),
3116 })
3117 }
3118
3119 fn mantle_function_call_item(id: &str, name: &str, call_id: &str) -> ResponseOutputItem {
3120 ResponseOutputItem::FunctionCall(ResponseFunctionToolCall {
3121 id: Some(id.to_string()),
3122 arguments: String::new(),
3123 call_id: Some(call_id.to_string()),
3124 name: Some(name.to_string()),
3125 status: Some("in_progress".to_string()),
3126 })
3127 }
3128
3129 fn mantle_reasoning_item(id: &str) -> ResponseOutputItem {
3130 ResponseOutputItem::Reasoning(ResponseReasoningItem {
3131 id: Some(id.to_string()),
3132 summary: Vec::new(),
3133 content: Vec::new(),
3134 encrypted_content: None,
3135 status: Some("in_progress".to_string()),
3136 })
3137 }
3138
3139 fn text_delta(item_id: &str, output_index: usize, delta: &str) -> OpenAiResponseStreamEvent {
3140 OpenAiResponseStreamEvent::OutputTextDelta {
3141 item_id: item_id.to_string(),
3142 output_index,
3143 content_index: Some(0),
3144 delta: delta.to_string(),
3145 }
3146 }
3147
3148 fn item_added(output_index: usize, item: ResponseOutputItem) -> OpenAiResponseStreamEvent {
3149 OpenAiResponseStreamEvent::OutputItemAdded {
3150 output_index,
3151 sequence_number: None,
3152 item,
3153 }
3154 }
3155
3156 fn item_done(output_index: usize, item: ResponseOutputItem) -> OpenAiResponseStreamEvent {
3157 OpenAiResponseStreamEvent::OutputItemDone {
3158 output_index,
3159 sequence_number: None,
3160 item,
3161 }
3162 }
3163
3164 fn map_mantle_response_events(
3165 events: Vec<OpenAiResponseStreamEvent>,
3166 ) -> Vec<LanguageModelCompletionEvent> {
3167 let mut mapper = MantleResponseEventMapper::new();
3168 events
3169 .into_iter()
3170 .flat_map(|event| mapper.map_event(event))
3171 .map(|event| event.expect("Mantle response event should map successfully"))
3172 .collect()
3173 }
3174
3175 // Keeps only StartMessage and Text events so phase/reasoning tests can assert on
3176 // message boundaries without depending on incidental ReasoningDetails metadata.
3177 fn start_messages_and_texts(
3178 events: Vec<LanguageModelCompletionEvent>,
3179 ) -> Vec<LanguageModelCompletionEvent> {
3180 events
3181 .into_iter()
3182 .filter(|event| {
3183 matches!(
3184 event,
3185 LanguageModelCompletionEvent::StartMessage { .. }
3186 | LanguageModelCompletionEvent::Text(_)
3187 )
3188 })
3189 .collect()
3190 }
3191
3192 #[test]
3193 fn mantle_response_mapper_coalesces_cumulative_message_snapshots() {
3194 let first_message = mantle_message_item("msg_1");
3195 let second_message = mantle_message_item("msg_2");
3196 let events = map_mantle_response_events(vec![
3197 OpenAiResponseStreamEvent::OutputItemAdded {
3198 output_index: 0,
3199 sequence_number: None,
3200 item: first_message.clone(),
3201 },
3202 OpenAiResponseStreamEvent::OutputTextDelta {
3203 item_id: "msg_1".to_string(),
3204 output_index: 0,
3205 content_index: Some(0),
3206 delta: "Plan: rename".to_string(),
3207 },
3208 OpenAiResponseStreamEvent::OutputItemDone {
3209 output_index: 0,
3210 sequence_number: None,
3211 item: first_message,
3212 },
3213 OpenAiResponseStreamEvent::OutputItemAdded {
3214 output_index: 1,
3215 sequence_number: None,
3216 item: second_message.clone(),
3217 },
3218 OpenAiResponseStreamEvent::OutputTextDelta {
3219 item_id: "msg_2".to_string(),
3220 output_index: 1,
3221 content_index: Some(0),
3222 delta: "Plan: rename".to_string(),
3223 },
3224 OpenAiResponseStreamEvent::OutputTextDelta {
3225 item_id: "msg_2".to_string(),
3226 output_index: 1,
3227 content_index: Some(0),
3228 delta: " and regenerate".to_string(),
3229 },
3230 OpenAiResponseStreamEvent::OutputItemDone {
3231 output_index: 1,
3232 sequence_number: None,
3233 item: second_message,
3234 },
3235 OpenAiResponseStreamEvent::Completed {
3236 response: open_ai::responses::ResponseSummary::default(),
3237 },
3238 ]);
3239
3240 assert_eq!(
3241 events,
3242 vec![
3243 LanguageModelCompletionEvent::StartMessage {
3244 message_id: "msg_1".to_string(),
3245 },
3246 LanguageModelCompletionEvent::Text("Plan: rename".to_string()),
3247 LanguageModelCompletionEvent::Text(" and regenerate".to_string()),
3248 LanguageModelCompletionEvent::Stop(language_model::StopReason::EndTurn),
3249 ]
3250 );
3251 }
3252
3253 #[test]
3254 fn mantle_response_mapper_preserves_independent_messages() {
3255 let first_message = mantle_message_item("msg_1");
3256 let second_message = mantle_message_item("msg_2");
3257 let events = map_mantle_response_events(vec![
3258 OpenAiResponseStreamEvent::OutputItemAdded {
3259 output_index: 0,
3260 sequence_number: None,
3261 item: first_message.clone(),
3262 },
3263 OpenAiResponseStreamEvent::OutputTextDelta {
3264 item_id: "msg_1".to_string(),
3265 output_index: 0,
3266 content_index: Some(0),
3267 delta: "Plan".to_string(),
3268 },
3269 OpenAiResponseStreamEvent::OutputItemDone {
3270 output_index: 0,
3271 sequence_number: None,
3272 item: first_message,
3273 },
3274 OpenAiResponseStreamEvent::OutputItemAdded {
3275 output_index: 1,
3276 sequence_number: None,
3277 item: second_message.clone(),
3278 },
3279 OpenAiResponseStreamEvent::OutputTextDelta {
3280 item_id: "msg_2".to_string(),
3281 output_index: 1,
3282 content_index: Some(0),
3283 delta: "Final answer".to_string(),
3284 },
3285 OpenAiResponseStreamEvent::OutputItemDone {
3286 output_index: 1,
3287 sequence_number: None,
3288 item: second_message,
3289 },
3290 ]);
3291
3292 assert_eq!(
3293 events,
3294 vec![
3295 LanguageModelCompletionEvent::StartMessage {
3296 message_id: "msg_1".to_string(),
3297 },
3298 LanguageModelCompletionEvent::Text("Plan".to_string()),
3299 LanguageModelCompletionEvent::StartMessage {
3300 message_id: "msg_2".to_string(),
3301 },
3302 LanguageModelCompletionEvent::Text("Final answer".to_string()),
3303 ]
3304 );
3305 }
3306
3307 #[test]
3308 fn mantle_response_mapper_coalesces_chained_cumulative_snapshots() {
3309 // Each merge must thread `previous_message`/`emitted_text_length`
3310 // through to the next.
3311 let first_message = mantle_message_item("msg_1");
3312 let second_message = mantle_message_item("msg_2");
3313 let third_message = mantle_message_item("msg_3");
3314 let events = map_mantle_response_events(vec![
3315 item_added(0, first_message.clone()),
3316 text_delta("msg_1", 0, "Plan"),
3317 item_done(0, first_message),
3318 item_added(1, second_message.clone()),
3319 text_delta("msg_2", 1, "Plan and go"),
3320 item_done(1, second_message),
3321 item_added(2, third_message.clone()),
3322 text_delta("msg_3", 2, "Plan and go further"),
3323 item_done(2, third_message),
3324 ]);
3325
3326 assert_eq!(
3327 events,
3328 vec![
3329 LanguageModelCompletionEvent::StartMessage {
3330 message_id: "msg_1".to_string(),
3331 },
3332 LanguageModelCompletionEvent::Text("Plan".to_string()),
3333 LanguageModelCompletionEvent::Text(" and go".to_string()),
3334 LanguageModelCompletionEvent::Text(" further".to_string()),
3335 ]
3336 );
3337 }
3338
3339 #[test]
3340 fn mantle_response_mapper_keeps_message_state_across_content_events() {
3341 let first_message = mantle_message_item("msg_1");
3342 let second_message = mantle_message_item("msg_2");
3343 let events = map_mantle_response_events(vec![
3344 item_added(0, first_message.clone()),
3345 text_delta("msg_1", 0, "Plan"),
3346 item_done(0, first_message),
3347 item_added(1, second_message.clone()),
3348 OpenAiResponseStreamEvent::ContentPartAdded {
3349 item_id: "msg_2".to_string(),
3350 output_index: 1,
3351 content_index: 0,
3352 part: serde_json::json!({"type": "output_text", "text": ""}),
3353 },
3354 text_delta("msg_2", 1, "Plan"),
3355 text_delta("msg_2", 1, " and continue"),
3356 OpenAiResponseStreamEvent::OutputTextDone {
3357 item_id: "msg_2".to_string(),
3358 output_index: 1,
3359 content_index: Some(0),
3360 text: "Plan and continue".to_string(),
3361 },
3362 OpenAiResponseStreamEvent::ContentPartDone {
3363 item_id: "msg_2".to_string(),
3364 output_index: 1,
3365 content_index: 0,
3366 part: serde_json::json!({"type": "output_text", "text": "Plan and continue"}),
3367 },
3368 item_done(1, second_message),
3369 ]);
3370
3371 assert_eq!(
3372 events,
3373 vec![
3374 LanguageModelCompletionEvent::StartMessage {
3375 message_id: "msg_1".to_string(),
3376 },
3377 LanguageModelCompletionEvent::Text("Plan".to_string()),
3378 LanguageModelCompletionEvent::Text(" and continue".to_string()),
3379 ]
3380 );
3381 }
3382
3383 #[test]
3384 fn mantle_response_mapper_resets_on_non_message_item_done() {
3385 let first_message = mantle_message_item("msg_1");
3386 let second_message = mantle_message_item("msg_2");
3387 let tool_call = mantle_function_call_item("call_1", "get_weather", "server_call_1");
3388 let events = map_mantle_response_events(vec![
3389 item_added(0, first_message.clone()),
3390 text_delta("msg_1", 0, "Plan"),
3391 item_done(0, first_message),
3392 item_done(1, tool_call),
3393 item_added(2, second_message.clone()),
3394 text_delta("msg_2", 2, "Plan continued"),
3395 item_done(2, second_message),
3396 ]);
3397
3398 assert_eq!(
3399 start_messages_and_texts(events),
3400 vec![
3401 LanguageModelCompletionEvent::StartMessage {
3402 message_id: "msg_1".to_string(),
3403 },
3404 LanguageModelCompletionEvent::Text("Plan".to_string()),
3405 LanguageModelCompletionEvent::StartMessage {
3406 message_id: "msg_2".to_string(),
3407 },
3408 LanguageModelCompletionEvent::Text("Plan continued".to_string()),
3409 ]
3410 );
3411 }
3412
3413 #[test]
3414 fn mantle_response_mapper_keeps_post_tool_message_separate_when_extending_prior_text() {
3415 // A tool call sits between two messages. Even though the second message's
3416 // text happens to extend the first message's text, the tool call is a hard
3417 // boundary and the second message must remain a separate, visible message.
3418 let first_message = mantle_message_item("msg_1");
3419 let tool_call = mantle_function_call_item("call_1", "get_weather", "server_call_1");
3420 let second_message = mantle_message_item("msg_2");
3421 let events = map_mantle_response_events(vec![
3422 item_added(0, first_message.clone()),
3423 text_delta("msg_1", 0, "Plan"),
3424 item_done(0, first_message),
3425 item_added(1, tool_call.clone()),
3426 OpenAiResponseStreamEvent::FunctionCallArgumentsDone {
3427 item_id: "call_1".to_string(),
3428 output_index: 1,
3429 arguments: "{\"city\":\"Boston\"}".to_string(),
3430 sequence_number: None,
3431 },
3432 item_done(1, tool_call),
3433 item_added(2, second_message.clone()),
3434 text_delta("msg_2", 2, "Plan continued"),
3435 item_done(2, second_message),
3436 ]);
3437
3438 assert_eq!(
3439 events,
3440 vec![
3441 LanguageModelCompletionEvent::StartMessage {
3442 message_id: "msg_1".to_string(),
3443 },
3444 LanguageModelCompletionEvent::Text("Plan".to_string()),
3445 LanguageModelCompletionEvent::ToolUse(LanguageModelToolUse {
3446 id: language_model::LanguageModelToolUseId::from("server_call_1"),
3447 name: Arc::<str>::from("get_weather"),
3448 is_input_complete: true,
3449 input: language_model::LanguageModelToolUseInput::Json(
3450 serde_json::json!({"city": "Boston"}),
3451 ),
3452 raw_input: "{\"city\":\"Boston\"}".to_string(),
3453 thought_signature: None,
3454 }),
3455 LanguageModelCompletionEvent::StartMessage {
3456 message_id: "msg_2".to_string(),
3457 },
3458 LanguageModelCompletionEvent::Text("Plan continued".to_string()),
3459 ]
3460 );
3461 }
3462
3463 #[test]
3464 fn mantle_response_mapper_treats_reasoning_as_collapse_boundary() {
3465 // A reasoning item between two messages breaks adjacency, so the second
3466 // message is emitted independently even though it repeats the first's text.
3467 let first_message = mantle_message_item("msg_1");
3468 let reasoning = mantle_reasoning_item("rsn_1");
3469 let second_message = mantle_message_item("msg_2");
3470 let events = map_mantle_response_events(vec![
3471 item_added(0, first_message.clone()),
3472 text_delta("msg_1", 0, "Plan"),
3473 item_done(0, first_message),
3474 item_added(1, reasoning.clone()),
3475 item_done(1, reasoning),
3476 item_added(2, second_message.clone()),
3477 text_delta("msg_2", 2, "Plan"),
3478 item_done(2, second_message),
3479 ]);
3480
3481 assert_eq!(
3482 start_messages_and_texts(events),
3483 vec![
3484 LanguageModelCompletionEvent::StartMessage {
3485 message_id: "msg_1".to_string(),
3486 },
3487 LanguageModelCompletionEvent::Text("Plan".to_string()),
3488 LanguageModelCompletionEvent::StartMessage {
3489 message_id: "msg_2".to_string(),
3490 },
3491 LanguageModelCompletionEvent::Text("Plan".to_string()),
3492 ]
3493 );
3494 }
3495
3496 #[test]
3497 fn mantle_response_mapper_preserves_equal_and_shrinking_messages() {
3498 // Equal and strict-prefix (shrinking) adjacent messages must not be dropped:
3499 // they are emitted as independent messages rather than collapsed away.
3500 let first_message = mantle_message_item("msg_1");
3501 let second_message = mantle_message_item("msg_2");
3502 let third_message = mantle_message_item("msg_3");
3503 let events = map_mantle_response_events(vec![
3504 item_added(0, first_message.clone()),
3505 text_delta("msg_1", 0, "Hello world"),
3506 item_done(0, first_message),
3507 // Equal to the previous message.
3508 item_added(1, second_message.clone()),
3509 text_delta("msg_2", 1, "Hello world"),
3510 item_done(1, second_message),
3511 // Strict prefix of (shorter than) the previous message.
3512 item_added(2, third_message.clone()),
3513 text_delta("msg_3", 2, "Hello"),
3514 item_done(2, third_message),
3515 ]);
3516
3517 assert_eq!(
3518 events,
3519 vec![
3520 LanguageModelCompletionEvent::StartMessage {
3521 message_id: "msg_1".to_string(),
3522 },
3523 LanguageModelCompletionEvent::Text("Hello world".to_string()),
3524 LanguageModelCompletionEvent::StartMessage {
3525 message_id: "msg_2".to_string(),
3526 },
3527 LanguageModelCompletionEvent::Text("Hello world".to_string()),
3528 LanguageModelCompletionEvent::StartMessage {
3529 message_id: "msg_3".to_string(),
3530 },
3531 LanguageModelCompletionEvent::Text("Hello".to_string()),
3532 ]
3533 );
3534 }
3535
3536 #[test]
3537 fn mantle_response_mapper_keeps_different_phase_messages_separate() {
3538 // Two messages share a text prefix but have different phases; only same-phase
3539 // strict extensions collapse, so the second message stays visible.
3540 let first_message = mantle_message_item_with_phase("msg_1", Some("commentary"));
3541 let second_message = mantle_message_item_with_phase("msg_2", Some("final_answer"));
3542 let events = map_mantle_response_events(vec![
3543 item_added(0, first_message.clone()),
3544 text_delta("msg_1", 0, "Plan: rename"),
3545 item_done(0, first_message),
3546 item_added(1, second_message.clone()),
3547 text_delta("msg_2", 1, "Plan: rename and regenerate"),
3548 item_done(1, second_message),
3549 ]);
3550
3551 assert_eq!(
3552 start_messages_and_texts(events),
3553 vec![
3554 LanguageModelCompletionEvent::StartMessage {
3555 message_id: "msg_1".to_string(),
3556 },
3557 LanguageModelCompletionEvent::Text("Plan: rename".to_string()),
3558 LanguageModelCompletionEvent::StartMessage {
3559 message_id: "msg_2".to_string(),
3560 },
3561 LanguageModelCompletionEvent::Text("Plan: rename and regenerate".to_string()),
3562 ]
3563 );
3564 }
3565
3566 #[test]
3567 fn mantle_response_mapper_coalesces_same_phase_strict_extension() {
3568 // Same phase, strict extension: the replayed prefix is collapsed and only the
3569 // new suffix is emitted, so no second StartMessage appears.
3570 let first_message = mantle_message_item_with_phase("msg_1", Some("final_answer"));
3571 let second_message = mantle_message_item_with_phase("msg_2", Some("final_answer"));
3572 let events = map_mantle_response_events(vec![
3573 item_added(0, first_message.clone()),
3574 text_delta("msg_1", 0, "Plan: rename"),
3575 item_done(0, first_message),
3576 item_added(1, second_message.clone()),
3577 text_delta("msg_2", 1, "Plan: rename"),
3578 text_delta("msg_2", 1, " and regenerate"),
3579 item_done(1, second_message),
3580 ]);
3581
3582 assert_eq!(
3583 start_messages_and_texts(events),
3584 vec![
3585 LanguageModelCompletionEvent::StartMessage {
3586 message_id: "msg_1".to_string(),
3587 },
3588 LanguageModelCompletionEvent::Text("Plan: rename".to_string()),
3589 LanguageModelCompletionEvent::Text(" and regenerate".to_string()),
3590 ]
3591 );
3592 }
3593
3594 #[test]
3595 fn mantle_response_mapper_does_not_duplicate_reasoning_details_on_merge() {
3596 // Merging a second message into the first must not emit any additional
3597 // phase/reasoning metadata events beyond what the first message alone would
3598 // have produced, since the metadata can't have changed by merging.
3599 fn reasoning_details_count(events: &[LanguageModelCompletionEvent]) -> usize {
3600 events
3601 .iter()
3602 .filter(|event| matches!(event, LanguageModelCompletionEvent::ReasoningDetails(_)))
3603 .count()
3604 }
3605
3606 let first_message = mantle_message_item_with_phase("msg_1", Some("final_answer"));
3607 let baseline_events = map_mantle_response_events(vec![
3608 item_added(0, first_message.clone()),
3609 text_delta("msg_1", 0, "Plan: rename"),
3610 item_done(0, first_message.clone()),
3611 ]);
3612
3613 let second_message = mantle_message_item_with_phase("msg_2", Some("final_answer"));
3614 let merged_events = map_mantle_response_events(vec![
3615 item_added(0, first_message.clone()),
3616 text_delta("msg_1", 0, "Plan: rename"),
3617 item_done(0, first_message),
3618 item_added(1, second_message.clone()),
3619 text_delta("msg_2", 1, "Plan: rename"),
3620 text_delta("msg_2", 1, " and regenerate"),
3621 item_done(1, second_message),
3622 ]);
3623
3624 assert_eq!(
3625 reasoning_details_count(&merged_events),
3626 reasoning_details_count(&baseline_events),
3627 );
3628 }
3629
3630 #[test]
3631 fn test_builtin_mantle_models_support_thinking() {
3632 assert!(MantleModel::Gpt5_6Sol.supports_thinking());
3633 assert!(MantleModel::Gpt5_6Terra.supports_thinking());
3634 assert!(MantleModel::Gpt5_6Luna.supports_thinking());
3635 assert!(MantleModel::Gpt5_5.supports_thinking());
3636 assert!(MantleModel::Gpt5_4.supports_thinking());
3637 assert!(MantleModel::Grok4_3.supports_thinking());
3638 assert_eq!(
3639 mantle_default_reasoning_effort(&MantleModel::Gpt5_6Sol),
3640 Some(ReasoningEffort::Medium)
3641 );
3642 assert_eq!(
3643 mantle_default_reasoning_effort(&MantleModel::Gpt5_6Terra),
3644 Some(ReasoningEffort::Medium)
3645 );
3646 assert_eq!(
3647 mantle_default_reasoning_effort(&MantleModel::Gpt5_6Luna),
3648 Some(ReasoningEffort::Medium)
3649 );
3650 assert_eq!(
3651 mantle_default_reasoning_effort(&MantleModel::Gpt5_5),
3652 Some(ReasoningEffort::Medium)
3653 );
3654 assert_eq!(
3655 mantle_default_reasoning_effort(&MantleModel::Grok4_3),
3656 Some(ReasoningEffort::Medium)
3657 );
3658 }
3659
3660 #[test]
3661 fn test_mantle_supported_effort_levels_hide_none() {
3662 let effort_levels = mantle_supported_effort_levels(&MantleModel::Gpt5_5);
3663 let values = effort_levels
3664 .iter()
3665 .map(|level| level.value.as_ref())
3666 .collect::<Vec<_>>();
3667
3668 assert_eq!(values, ["low", "medium", "high", "xhigh"]);
3669 assert_eq!(
3670 effort_levels
3671 .iter()
3672 .find(|level| level.is_default)
3673 .map(|level| level.value.as_ref()),
3674 Some("medium")
3675 );
3676 }
3677
3678 #[test]
3679 fn test_custom_mantle_model_can_disable_thinking() {
3680 let model = MantleModel::Custom {
3681 name: "custom-mantle-model".to_string(),
3682 display_name: None,
3683 max_tokens: 128_000,
3684 max_output_tokens: None,
3685 protocol: MantleProtocol::Responses,
3686 supports_tools: true,
3687 supports_images: false,
3688 supports_thinking: false,
3689 };
3690
3691 assert!(!model.supports_thinking());
3692 assert_eq!(mantle_default_reasoning_effort(&model), None);
3693 assert!(mantle_supported_effort_levels(&model).is_empty());
3694 assert_eq!(
3695 mantle_selected_reasoning_effort(
3696 &LanguageModelRequest {
3697 thinking_effort: Some("high".to_string()),
3698 ..Default::default()
3699 },
3700 &model,
3701 ),
3702 None
3703 );
3704 }
3705
3706 #[test]
3707 fn test_disabled_mantle_thinking_serializes_none() {
3708 let request = into_open_ai_response(
3709 LanguageModelRequest {
3710 thinking_allowed: false,
3711 ..Default::default()
3712 },
3713 MantleModel::Grok4_3.request_id(),
3714 true,
3715 false,
3716 Some(MantleModel::Grok4_3.max_output_tokens()),
3717 mantle_default_reasoning_effort(&MantleModel::Grok4_3),
3718 MantleModel::Grok4_3.supports_thinking(),
3719 &PROVIDER_ID,
3720 )
3721 .unwrap();
3722
3723 assert_eq!(
3724 serde_json::to_value(&request).unwrap()["reasoning"],
3725 serde_json::json!({ "effort": "none" })
3726 );
3727 }
3728
3729 #[test]
3730 fn test_mantle_reasoning_passes_known_efforts_through() {
3731 for effort in ["low", "medium", "high", "xhigh", "minimal", "max"] {
3732 assert_eq!(
3733 mantle_selected_reasoning_effort(
3734 &LanguageModelRequest {
3735 thinking_allowed: true,
3736 thinking_effort: Some(effort.to_string()),
3737 ..Default::default()
3738 },
3739 &MantleModel::Gpt5_5,
3740 )
3741 .map(|effort| effort.value()),
3742 Some(effort)
3743 );
3744 }
3745
3746 assert_eq!(
3747 mantle_selected_reasoning_effort(
3748 &LanguageModelRequest {
3749 thinking_allowed: true,
3750 thinking_effort: Some("none".to_string()),
3751 ..Default::default()
3752 },
3753 &MantleModel::Gpt5_5,
3754 ),
3755 Some(ReasoningEffort::Medium)
3756 );
3757 }
3758
3759 #[test]
3760 fn test_strip_unsupported_mantle_response_fields_removes_context_management() {
3761 let mut request = into_open_ai_response(
3762 LanguageModelRequest {
3763 compact_at_tokens: Some(10_000),
3764 ..Default::default()
3765 },
3766 "openai.gpt-5.5",
3767 true,
3768 false,
3769 Some(128_000),
3770 Some(ReasoningEffort::Medium),
3771 false,
3772 &PROVIDER_ID,
3773 )
3774 .unwrap();
3775
3776 assert!(request.context_management.is_some());
3777 strip_unsupported_mantle_response_fields(&mut request);
3778 assert!(request.context_management.is_none());
3779
3780 let request = serde_json::to_value(&request).unwrap();
3781 assert!(request.get("context_management").is_none());
3782 }
3783}
3784