From 629dde96ff369ec4f95701257b2025906d6aa5d9 Mon Sep 17 00:00:00 2001 From: Shardul Vaidya Date: Mon, 7 Apr 2025 17:02:30 -0400 Subject: [PATCH 01/10] one tokio handle doesn't do it --- crates/language_models/src/provider/bedrock.rs | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/crates/language_models/src/provider/bedrock.rs b/crates/language_models/src/provider/bedrock.rs index aa603534e0f647..730d23018987c9 100644 --- a/crates/language_models/src/provider/bedrock.rs +++ b/crates/language_models/src/provider/bedrock.rs @@ -480,6 +480,7 @@ impl BedrockModel { fn stream_completion( &self, request: bedrock::Request, + handle: tokio::runtime::Handle, cx: &AsyncApp, ) -> Result< BoxFuture<'static, BoxStream<'static, Result>>, @@ -488,10 +489,9 @@ impl BedrockModel { .get_or_init_client(cx) .cloned() .context("Bedrock client not initialized")?; - let owned_handle = self.handler.clone(); Ok(async move { - let request = bedrock::stream_completion(runtime_client, request, owned_handle); + let request = bedrock::stream_completion(runtime_client, request, handle); request.await.unwrap_or_else(|e| { futures::stream::once(async move { Err(BedrockError::ClientError(e)) }).boxed() }) @@ -579,7 +579,7 @@ impl LanguageModel for BedrockModel { let owned_handle = self.handler.clone(); - let request = self.stream_completion(request, cx); + let request = self.stream_completion(request, owned_handle.clone(), cx); let future = self.request_limiter.stream(async move { let response = request.map_err(|err| anyhow!(err))?.await; Ok(map_to_language_model_completion_events( From 1ab71998586e5b2afa2953d1986911429fdb2321 Mon Sep 17 00:00:00 2001 From: Shardul Vaidya Date: Tue, 8 Apr 2025 21:44:24 -0400 Subject: [PATCH 02/10] Stream completion has been simplified, but still same bug --- crates/bedrock/src/bedrock.rs | 109 +++---- .../language_models/src/provider/bedrock.rs | 289 +++++++----------- 2 files changed, 153 insertions(+), 245 deletions(-) diff --git a/crates/bedrock/src/bedrock.rs b/crates/bedrock/src/bedrock.rs index 5b34388a9afc49..7378ff8acbaff3 100644 --- a/crates/bedrock/src/bedrock.rs +++ b/crates/bedrock/src/bedrock.rs @@ -1,9 +1,6 @@ mod models; -use std::collections::HashMap; -use std::pin::Pin; - -use anyhow::{Error, Result, anyhow}; +use anyhow::{Context, Error, Result, anyhow}; use aws_sdk_bedrockruntime as bedrock; pub use aws_sdk_bedrockruntime as bedrock_client; pub use aws_sdk_bedrockruntime::types::{ @@ -24,6 +21,7 @@ pub use bedrock::types::{ use futures::stream::{self, BoxStream, Stream}; use serde::{Deserialize, Serialize}; use serde_json::{Number, Value}; +use std::collections::HashMap; use thiserror::Error; pub use crate::models::*; @@ -33,68 +31,59 @@ pub async fn stream_completion( request: Request, handle: tokio::runtime::Handle, ) -> Result>, Error> { - handle - .spawn(async move { - let mut response = bedrock::Client::converse_stream(&client) - .model_id(request.model.clone()) - .set_messages(request.messages.into()); + let mut response = bedrock::Client::converse_stream(&client) + .model_id(request.model.clone()) + .set_messages(request.messages.into()); - if let Some(Thinking::Enabled { - budget_tokens: Some(budget_tokens), - }) = request.thinking - { - response = - response.additional_model_request_fields(Document::Object(HashMap::from([( - "thinking".to_string(), - Document::from(HashMap::from([ - ("type".to_string(), Document::String("enabled".to_string())), - ( - "budget_tokens".to_string(), - Document::Number(AwsNumber::PosInt(budget_tokens)), - ), - ])), - )]))); - } + if let Some(Thinking::Enabled { + budget_tokens: Some(budget_tokens), + }) = request.thinking + { + let thinking_config = HashMap::from([ + ("type".to_string(), Document::String("enabled".to_string())), + ( + "budget_tokens".to_string(), + Document::Number(AwsNumber::PosInt(budget_tokens)), + ), + ]); + response = response.additional_model_request_fields(Document::Object(HashMap::from([( + "thinking".to_string(), + Document::from(thinking_config), + )]))); + } - if request.tools.is_some() && !request.tools.as_ref().unwrap().tools.is_empty() { - response = response.set_tool_config(request.tools); - } + if request + .tools + .as_ref() + .map_or(false, |t| !t.tools.is_empty()) + { + response = response.set_tool_config(request.tools); + } - let response = response.send().await; + // Send the request and create the stream + let output = handle + .spawn(response.send()) + .await? + .context("Failed to send API request to Bedrock"); - match response { - Ok(output) => { - let stream: Pin< - Box< - dyn Stream> - + Send, - >, - > = Box::pin(stream::unfold(output.stream, |mut stream| async move { - match stream.recv().await { - Ok(Some(output)) => Some(({ Ok(output) }, stream)), - Ok(None) => None, - Err(err) => { - Some(( - // TODO: Figure out how we can capture Throttling Exceptions - Err(BedrockError::ClientError(anyhow!( - "{:?}", - aws_sdk_bedrockruntime::error::DisplayErrorContext(err) - ))), - stream, - )) - } - } - })); - Ok(stream) - } - Err(err) => Err(anyhow!( - "{:?}", - aws_sdk_bedrockruntime::error::DisplayErrorContext(err) + let stream = Box::pin(stream::unfold( + output?.stream, + move |mut stream| async move { + match stream.recv().await { + Ok(Some(output)) => Some((Ok(output), stream)), + Ok(None) => None, + Err(err) => Some(( + Err(BedrockError::ClientError(anyhow!( + "{:?}", + aws_sdk_bedrockruntime::error::DisplayErrorContext(err) + ))), + stream, )), } - }) - .await - .map_err(|err| anyhow!("failed to spawn task: {err:?}"))? + }, + )); + + Ok(stream) } pub fn aws_document_to_value(document: &Document) -> Value { diff --git a/crates/language_models/src/provider/bedrock.rs b/crates/language_models/src/provider/bedrock.rs index 730d23018987c9..47c10712383a29 100644 --- a/crates/language_models/src/provider/bedrock.rs +++ b/crates/language_models/src/provider/bedrock.rs @@ -1,7 +1,3 @@ -use std::pin::Pin; -use std::str::FromStr; -use std::sync::Arc; - use crate::ui::InstructionListItem; use anyhow::{Context as _, Result, anyhow}; use aws_config::stalled_stream_protection::StalledStreamProtectionConfig; @@ -41,6 +37,10 @@ use serde::{Deserialize, Serialize}; use serde_json::Value; use settings::{Settings, SettingsStore}; use smol::lock::OnceCell; +use std::pin::Pin; +use std::str::FromStr; +use std::sync::Arc; +use std::time::{SystemTime, UNIX_EPOCH}; use strum::{EnumIter, IntoEnumIterator, IntoStaticStr}; use theme::ThemeSettings; use tokio::runtime::Handle; @@ -582,10 +582,7 @@ impl LanguageModel for BedrockModel { let request = self.stream_completion(request, owned_handle.clone(), cx); let future = self.request_limiter.stream(async move { let response = request.map_err(|err| anyhow!(err))?.await; - Ok(map_to_language_model_completion_events( - response, - owned_handle, - )) + Ok(map_to_language_model_completion_events(response)) }); async move { Ok(future.await?.boxed()) }.boxed() } @@ -781,7 +778,6 @@ pub fn get_bedrock_tokens( pub fn map_to_language_model_completion_events( events: Pin>>>, - handle: Handle, ) -> impl Stream> { struct RawToolUse { id: String, @@ -794,184 +790,107 @@ pub fn map_to_language_model_completion_events( tool_uses_by_index: HashMap, } - futures::stream::unfold( - State { - events, - tool_uses_by_index: HashMap::default(), - }, - move |mut state: State| { - let inner_handle = handle.clone(); - async move { - inner_handle - .spawn(async { - while let Some(event) = state.events.next().await { - match event { - Ok(event) => match event { - ConverseStreamOutput::ContentBlockDelta(cb_delta) => { - match cb_delta.delta { - Some(ContentBlockDelta::Text(text_out)) => { - let completion_event = - LanguageModelCompletionEvent::Text(text_out); - return Some((Some(Ok(completion_event)), state)); - } - - Some(ContentBlockDelta::ToolUse(text_out)) => { - if let Some(tool_use) = state - .tool_uses_by_index - .get_mut(&cb_delta.content_block_index) - { - tool_use.input_json.push_str(text_out.input()); - } - } - - Some(ContentBlockDelta::ReasoningContent(thinking)) => { - match thinking { - ReasoningContentBlockDelta::RedactedContent( - redacted, - ) => { - let thinking_event = - LanguageModelCompletionEvent::Thinking( - String::from_utf8( - redacted.into_inner(), - ) - .unwrap_or("REDACTED".to_string()), - ); - - return Some(( - Some(Ok(thinking_event)), - state, - )); - } - ReasoningContentBlockDelta::Signature(_sig) => { - } - ReasoningContentBlockDelta::Text(thoughts) => { - let thinking_event = - LanguageModelCompletionEvent::Thinking( - thoughts.to_string(), - ); - - return Some(( - Some(Ok(thinking_event)), - state, - )); - } - _ => {} - } - } - _ => {} - } - } - ConverseStreamOutput::ContentBlockStart(cb_start) => { - if let Some(ContentBlockStart::ToolUse(text_out)) = - cb_start.start - { - let tool_use = RawToolUse { - id: text_out.tool_use_id, - name: text_out.name, - input_json: String::new(), - }; - - state - .tool_uses_by_index - .insert(cb_start.content_block_index, tool_use); - } - } - ConverseStreamOutput::ContentBlockStop(cb_stop) => { - if let Some(tool_use) = state - .tool_uses_by_index - .remove(&cb_stop.content_block_index) - { - let tool_use_event = LanguageModelToolUse { - id: tool_use.id.into(), - name: tool_use.name.into(), - input: if tool_use.input_json.is_empty() { - Value::Null - } else { - serde_json::Value::from_str( - &tool_use.input_json, - ) - .map_err(|err| anyhow!(err)) - .unwrap() - }, - }; - - return Some(( - Some(Ok(LanguageModelCompletionEvent::ToolUse( - tool_use_event, - ))), - state, - )); - } - } - - ConverseStreamOutput::Metadata(cb_meta) => { - if let Some(metadata) = cb_meta.usage { - let completion_event = - LanguageModelCompletionEvent::UsageUpdate( - TokenUsage { - input_tokens: metadata.input_tokens as u32, - output_tokens: metadata.output_tokens - as u32, - cache_creation_input_tokens: default(), - cache_read_input_tokens: default(), - }, - ); - return Some((Some(Ok(completion_event)), state)); - } - } - ConverseStreamOutput::MessageStop(message_stop) => { - let reason = match message_stop.stop_reason { - StopReason::ContentFiltered => { - LanguageModelCompletionEvent::Stop( - language_model::StopReason::EndTurn, - ) - } - StopReason::EndTurn => { - LanguageModelCompletionEvent::Stop( - language_model::StopReason::EndTurn, - ) - } - StopReason::GuardrailIntervened => { - LanguageModelCompletionEvent::Stop( - language_model::StopReason::EndTurn, - ) - } - StopReason::MaxTokens => { - LanguageModelCompletionEvent::Stop( - language_model::StopReason::EndTurn, - ) - } - StopReason::StopSequence => { - LanguageModelCompletionEvent::Stop( - language_model::StopReason::EndTurn, - ) - } - StopReason::ToolUse => { - LanguageModelCompletionEvent::Stop( - language_model::StopReason::ToolUse, - ) - } - _ => LanguageModelCompletionEvent::Stop( - language_model::StopReason::EndTurn, - ), - }; - return Some((Some(Ok(reason)), state)); - } - _ => {} - }, - - Err(err) => return Some((Some(Err(anyhow!(err))), state)), + let initial_state = State { + events, + tool_uses_by_index: HashMap::default(), + }; + + futures::stream::unfold(initial_state, |mut state| async move { + match state.events.next().await { + Some(event_result) => match event_result { + Ok(event) => { + let result = match event { + ConverseStreamOutput::ContentBlockDelta(cb_delta) => match cb_delta.delta { + Some(ContentBlockDelta::Text(text)) => { + dbg!(format!( + "Converted Chunk: {}, {text}", + SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap() + .as_millis() + )); + Some(Ok(LanguageModelCompletionEvent::Text(text))) } + Some(ContentBlockDelta::ToolUse(tool_output)) => { + if let Some(tool_use) = state + .tool_uses_by_index + .get_mut(&cb_delta.content_block_index) + { + tool_use.input_json.push_str(tool_output.input()); + } + None + } + Some(ContentBlockDelta::ReasoningContent(thinking)) => match thinking { + ReasoningContentBlockDelta::Text(thoughts) => Some(Ok( + LanguageModelCompletionEvent::Thinking(thoughts.to_string()), + )), + ReasoningContentBlockDelta::RedactedContent(redacted) => { + let content = String::from_utf8(redacted.into_inner()) + .unwrap_or("REDACTED".to_string()); + Some(Ok(LanguageModelCompletionEvent::Thinking(content))) + } + _ => None, + }, + _ => None, + }, + ConverseStreamOutput::ContentBlockStart(cb_start) => { + if let Some(ContentBlockStart::ToolUse(tool_start)) = cb_start.start { + state.tool_uses_by_index.insert( + cb_start.content_block_index, + RawToolUse { + id: tool_start.tool_use_id, + name: tool_start.name, + input_json: String::new(), + }, + ); + } + None } - None - }) - .await - .log_err() - .flatten() - } - }, - ) - .filter_map(|event| async move { event }) + ConverseStreamOutput::ContentBlockStop(cb_stop) => state + .tool_uses_by_index + .remove(&cb_stop.content_block_index) + .map(|tool_use| { + let input = if tool_use.input_json.is_empty() { + Value::Null + } else { + serde_json::Value::from_str(&tool_use.input_json) + .unwrap_or(Value::Null) + }; + + Ok(LanguageModelCompletionEvent::ToolUse( + LanguageModelToolUse { + id: tool_use.id.into(), + name: tool_use.name.into(), + input, + }, + )) + }), + ConverseStreamOutput::Metadata(cb_meta) => cb_meta.usage.map(|metadata| { + Ok(LanguageModelCompletionEvent::UsageUpdate(TokenUsage { + input_tokens: metadata.input_tokens as u32, + output_tokens: metadata.output_tokens as u32, + cache_creation_input_tokens: default(), + cache_read_input_tokens: default(), + })) + }), + ConverseStreamOutput::MessageStop(message_stop) => { + let stop_reason = match message_stop.stop_reason { + StopReason::ToolUse => language_model::StopReason::ToolUse, + _ => language_model::StopReason::EndTurn, + }; + Some(Ok(LanguageModelCompletionEvent::Stop(stop_reason))) + } + _ => None, + }; + + Some((result, state)) + } + Err(err) => Some((Some(Err(anyhow!(err))), state)), + }, + None => None, + } + }) + .filter_map(|result| async move { result }) } struct ConfigurationView { From 6e1a9a037ab8a8f38572b75ca29c97e3fcccdc57 Mon Sep 17 00:00:00 2001 From: Shardul Vaidya Date: Wed, 9 Apr 2025 11:48:54 -0400 Subject: [PATCH 03/10] Added dbg timers to see how the system behaves --- crates/bedrock/src/bedrock.rs | 15 ++++++++++++--- crates/language_models/src/provider/bedrock.rs | 10 +++++----- 2 files changed, 17 insertions(+), 8 deletions(-) diff --git a/crates/bedrock/src/bedrock.rs b/crates/bedrock/src/bedrock.rs index 7378ff8acbaff3..8b7d093852779c 100644 --- a/crates/bedrock/src/bedrock.rs +++ b/crates/bedrock/src/bedrock.rs @@ -1,5 +1,6 @@ mod models; +use std::alloc::System; use anyhow::{Context, Error, Result, anyhow}; use aws_sdk_bedrockruntime as bedrock; pub use aws_sdk_bedrockruntime as bedrock_client; @@ -22,6 +23,7 @@ use futures::stream::{self, BoxStream, Stream}; use serde::{Deserialize, Serialize}; use serde_json::{Number, Value}; use std::collections::HashMap; +use std::time::{SystemTime, UNIX_EPOCH}; use thiserror::Error; pub use crate::models::*; @@ -59,10 +61,13 @@ pub async fn stream_completion( { response = response.set_tool_config(request.tools); } - // Send the request and create the stream let output = handle - .spawn(response.send()) + .spawn({ + let sent = SystemTime::now().duration_since(UNIX_EPOCH).unwrap().as_millis(); + dbg!(format!("Request sent: {sent}")); + response.send() + }) .await? .context("Failed to send API request to Bedrock"); @@ -70,7 +75,11 @@ pub async fn stream_completion( output?.stream, move |mut stream| async move { match stream.recv().await { - Ok(Some(output)) => Some((Ok(output), stream)), + Ok(Some(output)) => { + let rcvd = SystemTime::now().duration_since(UNIX_EPOCH).unwrap().as_millis(); + dbg!(format!("Received at: {rcvd}")); + Some((Ok(output), stream)) + }, Ok(None) => None, Err(err) => Some(( Err(BedrockError::ClientError(anyhow!( diff --git a/crates/language_models/src/provider/bedrock.rs b/crates/language_models/src/provider/bedrock.rs index 47c10712383a29..7650b0884128fa 100644 --- a/crates/language_models/src/provider/bedrock.rs +++ b/crates/language_models/src/provider/bedrock.rs @@ -802,12 +802,12 @@ pub fn map_to_language_model_completion_events( let result = match event { ConverseStreamOutput::ContentBlockDelta(cb_delta) => match cb_delta.delta { Some(ContentBlockDelta::Text(text)) => { + let rcvd = SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap() + .as_millis(); dbg!(format!( - "Converted Chunk: {}, {text}", - SystemTime::now() - .duration_since(UNIX_EPOCH) - .unwrap() - .as_millis() + "Converted Chunk: {rcvd}, {text}", )); Some(Ok(LanguageModelCompletionEvent::Text(text))) } From 4d648e2ad1e97e1838146f5666b9f60a97a63de7 Mon Sep 17 00:00:00 2001 From: Shardul Vaidya Date: Wed, 9 Apr 2025 13:52:56 -0400 Subject: [PATCH 04/10] Fixed streaming issues --- crates/bedrock/src/bedrock.rs | 18 ++------ .../language_models/src/provider/bedrock.rs | 42 +++++++------------ 2 files changed, 19 insertions(+), 41 deletions(-) diff --git a/crates/bedrock/src/bedrock.rs b/crates/bedrock/src/bedrock.rs index 8b7d093852779c..50d6319577c690 100644 --- a/crates/bedrock/src/bedrock.rs +++ b/crates/bedrock/src/bedrock.rs @@ -1,6 +1,5 @@ mod models; -use std::alloc::System; use anyhow::{Context, Error, Result, anyhow}; use aws_sdk_bedrockruntime as bedrock; pub use aws_sdk_bedrockruntime as bedrock_client; @@ -19,11 +18,10 @@ pub use bedrock::types::{ ToolResultContentBlock as BedrockToolResultContentBlock, ToolResultStatus as BedrockToolResultStatus, ToolUseBlock as BedrockToolUseBlock, }; -use futures::stream::{self, BoxStream, Stream}; +use futures::stream::{self, BoxStream}; use serde::{Deserialize, Serialize}; use serde_json::{Number, Value}; use std::collections::HashMap; -use std::time::{SystemTime, UNIX_EPOCH}; use thiserror::Error; pub use crate::models::*; @@ -31,7 +29,6 @@ pub use crate::models::*; pub async fn stream_completion( client: bedrock::Client, request: Request, - handle: tokio::runtime::Handle, ) -> Result>, Error> { let mut response = bedrock::Client::converse_stream(&client) .model_id(request.model.clone()) @@ -61,23 +58,14 @@ pub async fn stream_completion( { response = response.set_tool_config(request.tools); } - // Send the request and create the stream - let output = handle - .spawn({ - let sent = SystemTime::now().duration_since(UNIX_EPOCH).unwrap().as_millis(); - dbg!(format!("Request sent: {sent}")); - response.send() - }) - .await? - .context("Failed to send API request to Bedrock"); + + let output = response.send().await.context("Failed to send API request to Bedrock"); let stream = Box::pin(stream::unfold( output?.stream, move |mut stream| async move { match stream.recv().await { Ok(Some(output)) => { - let rcvd = SystemTime::now().duration_since(UNIX_EPOCH).unwrap().as_millis(); - dbg!(format!("Received at: {rcvd}")); Some((Ok(output), stream)) }, Ok(None) => None, diff --git a/crates/language_models/src/provider/bedrock.rs b/crates/language_models/src/provider/bedrock.rs index 7650b0884128fa..51ca254813659b 100644 --- a/crates/language_models/src/provider/bedrock.rs +++ b/crates/language_models/src/provider/bedrock.rs @@ -40,10 +40,8 @@ use smol::lock::OnceCell; use std::pin::Pin; use std::str::FromStr; use std::sync::Arc; -use std::time::{SystemTime, UNIX_EPOCH}; use strum::{EnumIter, IntoEnumIterator, IntoStaticStr}; use theme::ThemeSettings; -use tokio::runtime::Handle; use ui::{Icon, IconName, List, Tooltip, prelude::*}; use util::{ResultExt, default}; @@ -480,23 +478,24 @@ impl BedrockModel { fn stream_completion( &self, request: bedrock::Request, - handle: tokio::runtime::Handle, cx: &AsyncApp, - ) -> Result< - BoxFuture<'static, BoxStream<'static, Result>>, - > { - let runtime_client = self - .get_or_init_client(cx) + ) -> BoxFuture<'static, Result>>> { + let Ok(runtime_client) = self + .get_or_init_client(&cx) .cloned() - .context("Bedrock client not initialized")?; + .context("Bedrock client not initialized") + else { + return futures::future::ready(Err(anyhow!("App state dropped"))).boxed(); + }; - Ok(async move { - let request = bedrock::stream_completion(runtime_client, request, handle); - request.await.unwrap_or_else(|e| { - futures::stream::once(async move { Err(BedrockError::ClientError(e)) }).boxed() - }) + match Tokio::spawn(cx, bedrock::stream_completion(runtime_client, request)) { + Ok(res) => { + async {res.await.map_err(|err| anyhow!(err))?}.boxed() + } + Err(err) => { + futures::future::ready(Err(anyhow!(err))).boxed() + } } - .boxed()) } } @@ -577,11 +576,9 @@ impl LanguageModel for BedrockModel { Err(err) => return futures::future::ready(Err(err)).boxed(), }; - let owned_handle = self.handler.clone(); - - let request = self.stream_completion(request, owned_handle.clone(), cx); + let request = self.stream_completion(request, cx); let future = self.request_limiter.stream(async move { - let response = request.map_err(|err| anyhow!(err))?.await; + let response = request.await.map_err(|e| anyhow!(e))?; Ok(map_to_language_model_completion_events(response)) }); async move { Ok(future.await?.boxed()) }.boxed() @@ -802,13 +799,6 @@ pub fn map_to_language_model_completion_events( let result = match event { ConverseStreamOutput::ContentBlockDelta(cb_delta) => match cb_delta.delta { Some(ContentBlockDelta::Text(text)) => { - let rcvd = SystemTime::now() - .duration_since(UNIX_EPOCH) - .unwrap() - .as_millis(); - dbg!(format!( - "Converted Chunk: {rcvd}, {text}", - )); Some(Ok(LanguageModelCompletionEvent::Text(text))) } Some(ContentBlockDelta::ToolUse(tool_output)) => { From 5be21a907f91486fc30b392e563dfa8cedb90da1 Mon Sep 17 00:00:00 2001 From: Shardul Vaidya Date: Wed, 9 Apr 2025 14:54:19 -0400 Subject: [PATCH 05/10] Machete --- crates/bedrock/Cargo.toml | 2 -- 1 file changed, 2 deletions(-) diff --git a/crates/bedrock/Cargo.toml b/crates/bedrock/Cargo.toml index 84fd58460185ea..f8f6fa46017309 100644 --- a/crates/bedrock/Cargo.toml +++ b/crates/bedrock/Cargo.toml @@ -25,5 +25,3 @@ serde.workspace = true serde_json.workspace = true strum.workspace = true thiserror.workspace = true -tokio = { workspace = true, features = ["rt", "rt-multi-thread"] } -workspace-hack.workspace = true From 6f274fc607f03bf78110633f6d8f8c86a0fe1eba Mon Sep 17 00:00:00 2001 From: Shardul Vaidya Date: Wed, 9 Apr 2025 14:55:39 -0400 Subject: [PATCH 06/10] Maybe don't delete workspace-hack? --- crates/bedrock/Cargo.toml | 1 + 1 file changed, 1 insertion(+) diff --git a/crates/bedrock/Cargo.toml b/crates/bedrock/Cargo.toml index f8f6fa46017309..3000af50bb71be 100644 --- a/crates/bedrock/Cargo.toml +++ b/crates/bedrock/Cargo.toml @@ -25,3 +25,4 @@ serde.workspace = true serde_json.workspace = true strum.workspace = true thiserror.workspace = true +workspace-hack.workspace = true From 7dfc173d4cf325e76934174adf68210c129a2be3 Mon Sep 17 00:00:00 2001 From: Shardul Vaidya Date: Wed, 9 Apr 2025 14:56:21 -0400 Subject: [PATCH 07/10] Fmt --- crates/bedrock/src/bedrock.rs | 9 +++++---- crates/language_models/src/provider/bedrock.rs | 13 ++++++------- 2 files changed, 11 insertions(+), 11 deletions(-) diff --git a/crates/bedrock/src/bedrock.rs b/crates/bedrock/src/bedrock.rs index 50d6319577c690..c55760cfad0dc3 100644 --- a/crates/bedrock/src/bedrock.rs +++ b/crates/bedrock/src/bedrock.rs @@ -59,15 +59,16 @@ pub async fn stream_completion( response = response.set_tool_config(request.tools); } - let output = response.send().await.context("Failed to send API request to Bedrock"); + let output = response + .send() + .await + .context("Failed to send API request to Bedrock"); let stream = Box::pin(stream::unfold( output?.stream, move |mut stream| async move { match stream.recv().await { - Ok(Some(output)) => { - Some((Ok(output), stream)) - }, + Ok(Some(output)) => Some((Ok(output), stream)), Ok(None) => None, Err(err) => Some(( Err(BedrockError::ClientError(anyhow!( diff --git a/crates/language_models/src/provider/bedrock.rs b/crates/language_models/src/provider/bedrock.rs index 51ca254813659b..ea5e1ed349464a 100644 --- a/crates/language_models/src/provider/bedrock.rs +++ b/crates/language_models/src/provider/bedrock.rs @@ -479,7 +479,10 @@ impl BedrockModel { &self, request: bedrock::Request, cx: &AsyncApp, - ) -> BoxFuture<'static, Result>>> { + ) -> BoxFuture< + 'static, + Result>>, + > { let Ok(runtime_client) = self .get_or_init_client(&cx) .cloned() @@ -489,12 +492,8 @@ impl BedrockModel { }; match Tokio::spawn(cx, bedrock::stream_completion(runtime_client, request)) { - Ok(res) => { - async {res.await.map_err(|err| anyhow!(err))?}.boxed() - } - Err(err) => { - futures::future::ready(Err(anyhow!(err))).boxed() - } + Ok(res) => async { res.await.map_err(|err| anyhow!(err))? }.boxed(), + Err(err) => futures::future::ready(Err(anyhow!(err))).boxed(), } } } From e75ea2a0de4e57516b7ee23e89c71982c05cd818 Mon Sep 17 00:00:00 2001 From: Shardul Vaidya Date: Thu, 1 May 2025 11:36:27 -0400 Subject: [PATCH 08/10] Caught up to main --- Cargo.lock | 1 - .../language_models/src/provider/bedrock.rs | 228 ++++-------------- 2 files changed, 52 insertions(+), 177 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index c1b165309ec99d..369133b091b65d 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1904,7 +1904,6 @@ dependencies = [ "serde_json", "strum 0.27.1", "thiserror 2.0.12", - "tokio", "workspace-hack", ] diff --git a/crates/language_models/src/provider/bedrock.rs b/crates/language_models/src/provider/bedrock.rs index ac28146886eb23..9362cad2cc1b85 100644 --- a/crates/language_models/src/provider/bedrock.rs +++ b/crates/language_models/src/provider/bedrock.rs @@ -43,9 +43,6 @@ use serde::{Deserialize, Serialize}; use serde_json::Value; use settings::{Settings, SettingsStore}; use smol::lock::OnceCell; -use std::pin::Pin; -use std::str::FromStr; -use std::sync::Arc; use strum::{EnumIter, IntoEnumIterator, IntoStaticStr}; use theme::ThemeSettings; use ui::{Icon, IconName, List, Tooltip, prelude::*}; @@ -801,7 +798,7 @@ pub fn get_bedrock_tokens( pub fn map_to_language_model_completion_events( events: Pin>>>, -) -> impl Stream> { +) -> impl Stream> { struct RawToolUse { id: String, name: String, @@ -837,13 +834,25 @@ pub fn map_to_language_model_completion_events( None } Some(ContentBlockDelta::ReasoningContent(thinking)) => match thinking { - ReasoningContentBlockDelta::Text(thoughts) => Some(Ok( - LanguageModelCompletionEvent::Thinking(thoughts.to_string()), - )), + ReasoningContentBlockDelta::Text(thoughts) => { + Some(Ok(LanguageModelCompletionEvent::Thinking { + text: thoughts.clone(), + signature: None, + })) + } + ReasoningContentBlockDelta::Signature(sig) => { + Some(Ok(LanguageModelCompletionEvent::Thinking { + text: "".into(), + signature: Some(sig), + })) + } ReasoningContentBlockDelta::RedactedContent(redacted) => { let content = String::from_utf8(redacted.into_inner()) .unwrap_or("REDACTED".to_string()); - Some(Ok(LanguageModelCompletionEvent::Thinking(content))) + Some(Ok(LanguageModelCompletionEvent::Thinking { + text: content, + signature: None, + })) } _ => None, }, @@ -872,176 +881,43 @@ pub fn map_to_language_model_completion_events( serde_json::Value::from_str(&tool_use.input_json) .unwrap_or(Value::Null) }; - Some(ContentBlockDelta::ToolUse(text_out)) => { - if let Some(tool_use) = state - .tool_uses_by_index - .get_mut(&cb_delta.content_block_index) - { - tool_use.input_json.push_str(text_out.input()); - } - } - - Some(ContentBlockDelta::ReasoningContent(thinking)) => { - match thinking { - ReasoningContentBlockDelta::RedactedContent( - redacted, - ) => { - let thinking_event = - LanguageModelCompletionEvent::Thinking { - text: String::from_utf8( - redacted.into_inner(), - ) - .unwrap_or("REDACTED".to_string()), - signature: None, - }; - - return Some(( - Some(Ok(thinking_event)), - state, - )); - } - ReasoningContentBlockDelta::Signature( - signature, - ) => { - return Some(( - Some(Ok(LanguageModelCompletionEvent::Thinking { - text: "".to_string(), - signature: Some(signature) - })), - state, - )); - } - ReasoningContentBlockDelta::Text(thoughts) => { - let thinking_event = - LanguageModelCompletionEvent::Thinking { - text: thoughts.to_string(), - signature: None - }; - - return Some(( - Some(Ok(thinking_event)), - state, - )); - } - _ => {} - } - } - _ => {} - } - } - ConverseStreamOutput::ContentBlockStart(cb_start) => { - if let Some(ContentBlockStart::ToolUse(text_out)) = - cb_start.start - { - let tool_use = RawToolUse { - id: text_out.tool_use_id, - name: text_out.name, - input_json: String::new(), - }; - - state - .tool_uses_by_index - .insert(cb_start.content_block_index, tool_use); - } - } - ConverseStreamOutput::ContentBlockStop(cb_stop) => { - if let Some(tool_use) = state - .tool_uses_by_index - .remove(&cb_stop.content_block_index) - { - let tool_use_event = LanguageModelToolUse { - id: tool_use.id.into(), - name: tool_use.name.into(), - is_input_complete: true, - raw_input: tool_use.input_json.clone(), - input: if tool_use.input_json.is_empty() { - Value::Null - } else { - serde_json::Value::from_str( - &tool_use.input_json, - ) - .map_err(|err| anyhow!(err)) - .unwrap() - }, - }; - - return Some(( - Some(Ok(LanguageModelCompletionEvent::ToolUse( - tool_use_event, - ))), - state, - )); - } - } - - ConverseStreamOutput::Metadata(cb_meta) => { - if let Some(metadata) = cb_meta.usage { - let completion_event = - LanguageModelCompletionEvent::UsageUpdate( - TokenUsage { - input_tokens: metadata.input_tokens as u32, - output_tokens: metadata.output_tokens - as u32, - cache_creation_input_tokens: default(), - cache_read_input_tokens: default(), - }, - ); - return Some((Some(Ok(completion_event)), state)); - } - } - ConverseStreamOutput::MessageStop(message_stop) => { - let reason = match message_stop.stop_reason { - StopReason::ContentFiltered => { - LanguageModelCompletionEvent::Stop( - language_model::StopReason::EndTurn, - ) - } - StopReason::EndTurn => { - LanguageModelCompletionEvent::Stop( - language_model::StopReason::EndTurn, - ) - } - StopReason::GuardrailIntervened => { - LanguageModelCompletionEvent::Stop( - language_model::StopReason::EndTurn, - ) - } - StopReason::MaxTokens => { - LanguageModelCompletionEvent::Stop( - language_model::StopReason::EndTurn, - ) - } - StopReason::StopSequence => { - LanguageModelCompletionEvent::Stop( - language_model::StopReason::EndTurn, - ) - } - StopReason::ToolUse => { - LanguageModelCompletionEvent::Stop( - language_model::StopReason::ToolUse, - ) - } - _ => LanguageModelCompletionEvent::Stop( - language_model::StopReason::EndTurn, - ), - }; - return Some((Some(Ok(reason)), state)); - } - _ => {} - }, - Err(err) => return Some((Some(Err(anyhow!(err).into())), state)), - } + Ok(LanguageModelCompletionEvent::ToolUse( + LanguageModelToolUse { + id: tool_use.id.into(), + name: tool_use.name.into(), + is_input_complete: true, + raw_input: tool_use.input_json.clone(), + input, + }, + )) + }), + ConverseStreamOutput::Metadata(cb_meta) => cb_meta.usage.map(|metadata| { + Ok(LanguageModelCompletionEvent::UsageUpdate(TokenUsage { + input_tokens: metadata.input_tokens as u32, + output_tokens: metadata.output_tokens as u32, + cache_creation_input_tokens: default(), + cache_read_input_tokens: default(), + })) + }), + ConverseStreamOutput::MessageStop(message_stop) => { + let stop_reason = match message_stop.stop_reason { + StopReason::ToolUse => language_model::StopReason::ToolUse, + _ => language_model::StopReason::EndTurn, + }; + Some(Ok(LanguageModelCompletionEvent::Stop(stop_reason))) } - None - }) - .await - .log_err() - .flatten() - } - }, - ) - .filter_map(|event| async move { event }) + _ => None, + }; + + Some((result, state)) + } + Err(err) => Some((Some(Err(LanguageModelCompletionError::Other(anyhow!(err)))), state)), + }, + None => None, + } + }) + .filter_map(|result| async move { result }) } struct ConfigurationView { From 6ada85b20fd672f8ba75754603038ef70709659e Mon Sep 17 00:00:00 2001 From: Shardul Vaidya Date: Thu, 1 May 2025 11:44:30 -0400 Subject: [PATCH 09/10] Fmt --- crates/language_models/src/provider/bedrock.rs | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/crates/language_models/src/provider/bedrock.rs b/crates/language_models/src/provider/bedrock.rs index 9362cad2cc1b85..7fff1040434c14 100644 --- a/crates/language_models/src/provider/bedrock.rs +++ b/crates/language_models/src/provider/bedrock.rs @@ -912,7 +912,10 @@ pub fn map_to_language_model_completion_events( Some((result, state)) } - Err(err) => Some((Some(Err(LanguageModelCompletionError::Other(anyhow!(err)))), state)), + Err(err) => Some(( + Some(Err(LanguageModelCompletionError::Other(anyhow!(err)))), + state, + )), }, None => None, } From 13de1e44e8940533f848484763395febcb90d5e2 Mon Sep 17 00:00:00 2001 From: Shardul Vaidya Date: Wed, 25 Jun 2025 11:41:19 -0400 Subject: [PATCH 10/10] merge --- crates/language_models/src/provider/bedrock.rs | 12 +++++++++--- 1 file changed, 9 insertions(+), 3 deletions(-) diff --git a/crates/language_models/src/provider/bedrock.rs b/crates/language_models/src/provider/bedrock.rs index c3d14516a4f60d..9ae1be28a232f7 100644 --- a/crates/language_models/src/provider/bedrock.rs +++ b/crates/language_models/src/provider/bedrock.rs @@ -571,7 +571,7 @@ impl LanguageModel for BedrockModel { let request = self.stream_completion(request, cx); let future = self.request_limiter.stream(async move { - let response = request.map_err(|err| anyhow!(err))?.await; + let response = request.await.map_err(|err| anyhow!(err))?; let events = map_to_language_model_completion_events(response); if deny_tool_calls { @@ -974,8 +974,14 @@ pub fn map_to_language_model_completion_events( Ok(LanguageModelCompletionEvent::UsageUpdate(TokenUsage { input_tokens: metadata.input_tokens as u64, output_tokens: metadata.output_tokens as u64, - cache_creation_input_tokens: default(), - cache_read_input_tokens: default(), + cache_creation_input_tokens: metadata + .cache_write_input_tokens + .unwrap_or_default() + as u64, + cache_read_input_tokens: metadata + .cache_read_input_tokens + .unwrap_or_default() + as u64, })) }), ConverseStreamOutput::MessageStop(message_stop) => {