Rebrand to Galaxy, major improvements to Bedrock support, still needs some TLC though
This commit is contained in:
@@ -1,3 +1,5 @@
|
||||
use std::sync::Arc;
|
||||
|
||||
use anyhow::Result;
|
||||
use aws_config::BehaviorVersion;
|
||||
use aws_sdk_bedrockruntime::config::Region;
|
||||
@@ -6,6 +8,7 @@ use aws_sdk_bedrockruntime::Client as BedrockRuntimeClient;
|
||||
use crate::settings::ai::BedrockAuthMethod;
|
||||
|
||||
use super::convert::{build_converse_request, ConversationMessage, ToolDefinition};
|
||||
use super::diagnostic::BedrockDiagnosticLogger;
|
||||
use super::models::apply_cross_region_prefix;
|
||||
use super::stream::bedrock_stream_to_response_events;
|
||||
use crate::ai::agent::api::ResponseStream;
|
||||
@@ -102,12 +105,14 @@ impl BedrockClient {
|
||||
&self,
|
||||
model_id: &str,
|
||||
task_id: &str,
|
||||
needs_create_task: bool,
|
||||
messages: Vec<ConversationMessage>,
|
||||
system_prompt: Option<String>,
|
||||
tools: Vec<ToolDefinition>,
|
||||
max_tokens: i32,
|
||||
temperature: Option<f32>,
|
||||
cross_region_inference: bool,
|
||||
diagnostic_logger: Option<Arc<BedrockDiagnosticLogger>>,
|
||||
) -> Result<ResponseStream, BedrockError> {
|
||||
let effective_model_id = if cross_region_inference {
|
||||
apply_cross_region_prefix(model_id, &self.region)
|
||||
@@ -122,8 +127,26 @@ impl BedrockClient {
|
||||
tools.len()
|
||||
);
|
||||
|
||||
let converted =
|
||||
build_converse_request(messages, system_prompt, tools, max_tokens, temperature, None, None);
|
||||
let converted = build_converse_request(
|
||||
messages.clone(),
|
||||
system_prompt.clone(),
|
||||
tools.clone(),
|
||||
max_tokens,
|
||||
temperature,
|
||||
None,
|
||||
None,
|
||||
);
|
||||
|
||||
if let Some(ref logger) = diagnostic_logger {
|
||||
logger.log_bedrock_input(
|
||||
&messages,
|
||||
&system_prompt,
|
||||
&tools,
|
||||
max_tokens,
|
||||
temperature,
|
||||
cross_region_inference,
|
||||
);
|
||||
}
|
||||
|
||||
let mut request = self
|
||||
.runtime_client
|
||||
@@ -147,6 +170,9 @@ impl BedrockClient {
|
||||
} else {
|
||||
display_msg
|
||||
};
|
||||
if let Some(ref logger) = diagnostic_logger {
|
||||
logger.log_result_fail(&msg);
|
||||
}
|
||||
if msg.contains("AccessDenied") || msg.contains("access denied") {
|
||||
BedrockError::AccessDenied(msg)
|
||||
} else if msg.contains("ThrottlingException") || msg.contains("throttl") {
|
||||
@@ -161,7 +187,12 @@ impl BedrockClient {
|
||||
})?;
|
||||
|
||||
log::info!("[bedrock] Stream connected successfully");
|
||||
Ok(Box::pin(bedrock_stream_to_response_events(output, task_id.to_string())))
|
||||
Ok(Box::pin(bedrock_stream_to_response_events(
|
||||
output,
|
||||
task_id.to_string(),
|
||||
needs_create_task,
|
||||
diagnostic_logger,
|
||||
)))
|
||||
}
|
||||
|
||||
pub fn runtime_client(&self) -> &BedrockRuntimeClient {
|
||||
|
||||
Reference in New Issue
Block a user