Add unified models UI and Rig Bedrock runtime
This commit is contained in:
@@ -8,13 +8,17 @@ license.workspace = true
|
||||
[dependencies]
|
||||
async-stream.workspace = true
|
||||
async-trait.workspace = true
|
||||
aws-sdk-bedrockruntime.workspace = true
|
||||
base64.workspace = true
|
||||
futures.workspace = true
|
||||
galaxy_agent_core.workspace = true
|
||||
rig-core.workspace = true
|
||||
rig-bedrock.workspace = true
|
||||
serde_json.workspace = true
|
||||
uuid.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
aws-smithy-http-client.workspace = true
|
||||
bytes.workspace = true
|
||||
rig-core = { workspace = true, features = ["test-utils"] }
|
||||
tokio = { workspace = true, features = ["macros", "rt"] }
|
||||
|
||||
@@ -0,0 +1,162 @@
|
||||
use async_trait::async_trait;
|
||||
use aws_sdk_bedrockruntime::Client as AwsBedrockClient;
|
||||
use galaxy_agent_core::{
|
||||
AgentError, AgentErrorKind, AgentEventStream, AgentRuntime, RuntimeCapabilities,
|
||||
RuntimeDescriptor, RuntimeKind, TurnControl, TurnRequest,
|
||||
};
|
||||
use rig_bedrock::client::Client as RigBedrockClient;
|
||||
use rig_bedrock::completion::CompletionModel;
|
||||
use rig_core::client::CompletionClient;
|
||||
use rig_core::completion::CompletionRequest;
|
||||
|
||||
use crate::request::build_completion_request;
|
||||
use crate::stream::start_model_turn;
|
||||
|
||||
const INFERENCE_PROFILE_PREFIXES: &[&str] = &["us.", "eu.", "apac.", "jp.", "au.", "global."];
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub struct BedrockRigConfig {
|
||||
pub model: String,
|
||||
pub region: String,
|
||||
pub cross_region_inference: bool,
|
||||
pub prompt_caching: bool,
|
||||
pub max_output_tokens: Option<u64>,
|
||||
}
|
||||
|
||||
/// A Rig Bedrock client built from Galaxy's already-resolved AWS SDK client.
|
||||
///
|
||||
/// Credential/profile/SSO resolution remains in Galaxy's explicit Bedrock
|
||||
/// configuration boundary. Rig receives the resulting SDK client and owns the
|
||||
/// Converse request/stream conversion from that point onward.
|
||||
#[derive(Clone)]
|
||||
pub struct BedrockRuntime {
|
||||
client: RigBedrockClient,
|
||||
config: BedrockRigConfig,
|
||||
resolved_model: String,
|
||||
descriptor: RuntimeDescriptor,
|
||||
}
|
||||
|
||||
impl BedrockRuntime {
|
||||
pub fn from_aws_client(
|
||||
client: AwsBedrockClient,
|
||||
config: BedrockRigConfig,
|
||||
) -> Result<Self, AgentError> {
|
||||
let resolved_model =
|
||||
resolve_bedrock_model_id(&config.model, &config.region, config.cross_region_inference)?;
|
||||
let descriptor = RuntimeDescriptor {
|
||||
id: format!("rig-bedrock:{resolved_model}"),
|
||||
display_name: format!("Rig / Bedrock / {resolved_model}"),
|
||||
kind: RuntimeKind::Provider,
|
||||
capabilities: RuntimeCapabilities {
|
||||
model_selection: true,
|
||||
session_resume: false,
|
||||
steering: false,
|
||||
tool_permissions: false,
|
||||
},
|
||||
};
|
||||
|
||||
Ok(Self {
|
||||
client: RigBedrockClient::from(client),
|
||||
config,
|
||||
resolved_model,
|
||||
descriptor,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn resolved_model(&self) -> &str {
|
||||
&self.resolved_model
|
||||
}
|
||||
|
||||
pub fn completion_model(&self) -> CompletionModel {
|
||||
let model = self.client.completion_model(&self.resolved_model);
|
||||
if self.config.prompt_caching {
|
||||
model.with_prompt_caching()
|
||||
} else {
|
||||
model
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl AgentRuntime for BedrockRuntime {
|
||||
fn descriptor(&self) -> &RuntimeDescriptor {
|
||||
&self.descriptor
|
||||
}
|
||||
|
||||
async fn start_turn(
|
||||
&self,
|
||||
request: TurnRequest,
|
||||
control: TurnControl,
|
||||
) -> Result<AgentEventStream, AgentError> {
|
||||
let max_output_tokens = request.max_output_tokens.or(self.config.max_output_tokens);
|
||||
let mut completion_request =
|
||||
build_bedrock_completion_request(request, self.config.max_output_tokens)?;
|
||||
// Context markers and inference-profile expansion are Galaxy model
|
||||
// configuration, not identifiers that Rig should send unchanged.
|
||||
completion_request.model = Some(self.resolved_model.clone());
|
||||
start_model_turn(
|
||||
self.completion_model(),
|
||||
completion_request,
|
||||
control,
|
||||
max_output_tokens,
|
||||
)
|
||||
.await
|
||||
}
|
||||
}
|
||||
|
||||
pub fn build_bedrock_completion_request(
|
||||
request: TurnRequest,
|
||||
configured_max_output_tokens: Option<u64>,
|
||||
) -> Result<CompletionRequest, AgentError> {
|
||||
build_completion_request(request, configured_max_output_tokens, true, true, None)
|
||||
}
|
||||
|
||||
pub fn resolve_bedrock_model_id(
|
||||
configured_model: &str,
|
||||
region: &str,
|
||||
cross_region_inference: bool,
|
||||
) -> Result<String, AgentError> {
|
||||
let model = strip_context_marker(configured_model.trim());
|
||||
if model.is_empty() {
|
||||
return Err(AgentError::new(
|
||||
AgentErrorKind::Configuration,
|
||||
"Bedrock model ID is empty",
|
||||
));
|
||||
}
|
||||
|
||||
if !cross_region_inference
|
||||
|| model.starts_with("arn:")
|
||||
|| INFERENCE_PROFILE_PREFIXES
|
||||
.iter()
|
||||
.any(|prefix| model.starts_with(prefix))
|
||||
{
|
||||
return Ok(model.to_string());
|
||||
}
|
||||
|
||||
let prefix = inference_profile_prefix(region);
|
||||
Ok(prefix
|
||||
.map(|prefix| format!("{prefix}.{model}"))
|
||||
.unwrap_or_else(|| model.to_string()))
|
||||
}
|
||||
|
||||
fn strip_context_marker(model: &str) -> &str {
|
||||
model
|
||||
.get(..model.len().saturating_sub(4))
|
||||
.filter(|_| model.ends_with("[1m]") || model.ends_with("[1M]"))
|
||||
.unwrap_or(model)
|
||||
}
|
||||
|
||||
fn inference_profile_prefix(region: &str) -> Option<&'static str> {
|
||||
match region {
|
||||
region if region.starts_with("us-") || region.starts_with("ca-") => Some("us"),
|
||||
region if region.starts_with("eu-") || region == "il-central-1" => Some("eu"),
|
||||
"ap-northeast-1" | "ap-northeast-3" => Some("jp"),
|
||||
"ap-southeast-2" | "ap-southeast-4" | "ap-southeast-6" => Some("au"),
|
||||
region if region.starts_with("ap-") => Some("apac"),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
#[path = "bedrock_tests.rs"]
|
||||
mod tests;
|
||||
@@ -0,0 +1,279 @@
|
||||
use aws_sdk_bedrockruntime::config::Region;
|
||||
use aws_smithy_http_client::test_util::NeverClient;
|
||||
use futures::StreamExt;
|
||||
use galaxy_agent_core::{
|
||||
AgentEvent, AgentRuntime, ContentPart, ConversationMessage, MessageContent, MessageRole,
|
||||
StopReason, ToolDefinition, TurnCommand, TurnRequest, Usage,
|
||||
};
|
||||
use rig_bedrock::streaming::{BedrockStreamingResponse, BedrockUsage};
|
||||
use rig_core::completion::{AssistantContent, CompletionError, GetTokenUsage, Message};
|
||||
use rig_core::message::{DocumentSourceKind, ToolResultContent, UserContent};
|
||||
|
||||
use super::*;
|
||||
use crate::stream::{completion_error_stop_reason, map_usage};
|
||||
|
||||
#[test]
|
||||
fn resolves_context_marker_and_us_inference_profile() {
|
||||
assert_eq!(
|
||||
resolve_bedrock_model_id("anthropic.claude-sonnet-4-6[1m]", "us-east-1", true,).unwrap(),
|
||||
"us.anthropic.claude-sonnet-4-6"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolves_each_supported_inference_geography() {
|
||||
for (region, expected_prefix) in [
|
||||
("eu-west-1", "eu"),
|
||||
("il-central-1", "eu"),
|
||||
("ap-northeast-1", "jp"),
|
||||
("ap-southeast-2", "au"),
|
||||
("ap-southeast-1", "apac"),
|
||||
("ca-central-1", "us"),
|
||||
] {
|
||||
assert_eq!(
|
||||
resolve_bedrock_model_id("anthropic.claude-test", region, true).unwrap(),
|
||||
format!("{expected_prefix}.anthropic.claude-test")
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn preserves_arns_existing_profiles_and_unknown_regions() {
|
||||
let arn = "arn:aws:bedrock:us-east-1:123:application-inference-profile/example";
|
||||
assert_eq!(
|
||||
resolve_bedrock_model_id(arn, "us-east-1", true).unwrap(),
|
||||
arn
|
||||
);
|
||||
assert_eq!(
|
||||
resolve_bedrock_model_id("global.anthropic.claude-test", "us-east-1", true).unwrap(),
|
||||
"global.anthropic.claude-test"
|
||||
);
|
||||
assert_eq!(
|
||||
resolve_bedrock_model_id("anthropic.claude-test", "me-south-1", true).unwrap(),
|
||||
"anthropic.claude-test"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn prefixes_amazon_models_instead_of_mistaking_provider_for_geography() {
|
||||
assert_eq!(
|
||||
resolve_bedrock_model_id("amazon.nova-pro-v1:0", "us-east-1", true).unwrap(),
|
||||
"us.amazon.nova-pro-v1:0"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_an_empty_model_id() {
|
||||
let error = resolve_bedrock_model_id(" ", "us-east-1", false).unwrap_err();
|
||||
assert_eq!(error.kind, AgentErrorKind::Configuration);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn normalizes_bedrock_usage_and_max_token_stop() {
|
||||
let response = BedrockStreamingResponse {
|
||||
usage: Some(BedrockUsage {
|
||||
input_tokens: 100,
|
||||
output_tokens: 25,
|
||||
total_tokens: 125,
|
||||
cache_read_input_tokens: Some(40),
|
||||
cache_write_input_tokens: Some(10),
|
||||
}),
|
||||
};
|
||||
assert_eq!(
|
||||
map_usage(response.token_usage()),
|
||||
Usage {
|
||||
input_tokens: 100,
|
||||
output_tokens: 25,
|
||||
cached_input_tokens: 40,
|
||||
cache_creation_input_tokens: 10,
|
||||
}
|
||||
);
|
||||
assert_eq!(
|
||||
completion_error_stop_reason(&CompletionError::ProviderError(
|
||||
"Exceeded max tokens".to_string(),
|
||||
)),
|
||||
Some(StopReason::MaxTokens)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn constructs_rig_client_from_galaxys_resolved_aws_client_without_network() {
|
||||
let sdk_config = aws_sdk_bedrockruntime::Config::builder()
|
||||
.behavior_version_latest()
|
||||
.region(Region::new("us-east-1"))
|
||||
.http_client(NeverClient::new())
|
||||
.build();
|
||||
let aws_client = AwsBedrockClient::from_conf(sdk_config);
|
||||
let client = BedrockRuntime::from_aws_client(
|
||||
aws_client,
|
||||
BedrockRigConfig {
|
||||
model: "anthropic.claude-test[1M]".to_string(),
|
||||
region: "us-east-1".to_string(),
|
||||
cross_region_inference: true,
|
||||
prompt_caching: true,
|
||||
max_output_tokens: Some(8_192),
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(client.resolved_model(), "us.anthropic.claude-test");
|
||||
let completion_model = client.completion_model();
|
||||
assert_eq!(completion_model.model, client.resolved_model());
|
||||
assert!(completion_model.prompt_caching);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn cancellation_before_bedrock_stream_start_never_contacts_aws() {
|
||||
let never_client = NeverClient::new();
|
||||
let sdk_config = aws_sdk_bedrockruntime::Config::builder()
|
||||
.behavior_version_latest()
|
||||
.region(Region::new("us-east-1"))
|
||||
.http_client(never_client.clone())
|
||||
.build();
|
||||
let runtime = BedrockRuntime::from_aws_client(
|
||||
AwsBedrockClient::from_conf(sdk_config),
|
||||
BedrockRigConfig {
|
||||
model: "anthropic.claude-test".to_string(),
|
||||
region: "us-east-1".to_string(),
|
||||
cross_region_inference: false,
|
||||
prompt_caching: false,
|
||||
max_output_tokens: None,
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
let (sender, control) = galaxy_agent_core::turn_control();
|
||||
sender.send(TurnCommand::Cancel).await.unwrap();
|
||||
|
||||
let events = runtime
|
||||
.start_turn(
|
||||
TurnRequest::new(
|
||||
"anthropic.claude-test",
|
||||
vec![ConversationMessage {
|
||||
role: MessageRole::User,
|
||||
content: MessageContent::Text("Hello".to_string()),
|
||||
}],
|
||||
),
|
||||
control,
|
||||
)
|
||||
.await
|
||||
.unwrap()
|
||||
.collect::<Vec<_>>()
|
||||
.await
|
||||
.into_iter()
|
||||
.collect::<Result<Vec<_>, _>>()
|
||||
.unwrap();
|
||||
|
||||
assert!(matches!(events[0], AgentEvent::TurnStarted { .. }));
|
||||
assert_eq!(
|
||||
events[1],
|
||||
AgentEvent::TurnStopped {
|
||||
reason: StopReason::Cancelled,
|
||||
}
|
||||
);
|
||||
assert_eq!(never_client.num_calls(), 0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn bedrock_request_preserves_system_image_reasoning_tool_and_token_semantics() {
|
||||
let mut request = TurnRequest::new(
|
||||
"anthropic.claude-test",
|
||||
vec![
|
||||
ConversationMessage {
|
||||
role: MessageRole::User,
|
||||
content: MessageContent::MultiPart(vec![
|
||||
ContentPart::Text("Describe the image".to_string()),
|
||||
ContentPart::Image {
|
||||
data: vec![1, 2, 3, 4],
|
||||
mime_type: "image/png".to_string(),
|
||||
},
|
||||
]),
|
||||
},
|
||||
ConversationMessage {
|
||||
role: MessageRole::Assistant,
|
||||
content: MessageContent::MultiPart(vec![
|
||||
ContentPart::Reasoning {
|
||||
text: "I should inspect the manifest.".to_string(),
|
||||
signature: Some("signed-reasoning".to_string()),
|
||||
},
|
||||
ContentPart::ToolUse {
|
||||
tool_use_id: "call-1".to_string(),
|
||||
name: "read_files".to_string(),
|
||||
input: serde_json::json!({"files": ["Cargo.toml"]}),
|
||||
},
|
||||
]),
|
||||
},
|
||||
ConversationMessage {
|
||||
role: MessageRole::User,
|
||||
content: MessageContent::ToolResult {
|
||||
tool_use_id: "call-1".to_string(),
|
||||
content: "permission denied".to_string(),
|
||||
is_error: true,
|
||||
},
|
||||
},
|
||||
],
|
||||
);
|
||||
request.system_prompt = Some("Use Galaxy tools safely".to_string());
|
||||
request.max_output_tokens = Some(4_096);
|
||||
request.tools.push(ToolDefinition {
|
||||
name: "read_files".to_string(),
|
||||
description: "Read project files".to_string(),
|
||||
input_schema: serde_json::json!({
|
||||
"type": "object",
|
||||
"properties": {"files": {"type": "array"}}
|
||||
}),
|
||||
});
|
||||
|
||||
let converted = build_bedrock_completion_request(request, Some(8_192)).unwrap();
|
||||
assert!(converted.additional_params.is_none());
|
||||
assert_eq!(converted.max_tokens, Some(4_096));
|
||||
assert_eq!(converted.tools.len(), 1);
|
||||
assert_eq!(converted.tools[0].name, "read_files");
|
||||
|
||||
let messages = converted.chat_history.iter().collect::<Vec<_>>();
|
||||
let [system, user, assistant, result] = messages.as_slice() else {
|
||||
panic!("expected system, user, assistant, and tool-result messages");
|
||||
};
|
||||
assert!(matches!(
|
||||
system,
|
||||
Message::System { content } if content == "Use Galaxy tools safely"
|
||||
));
|
||||
|
||||
let Message::User { content } = user else {
|
||||
panic!("expected a user image message");
|
||||
};
|
||||
let user_content = content.iter().collect::<Vec<_>>();
|
||||
assert!(matches!(
|
||||
user_content.as_slice(),
|
||||
[UserContent::Text(text), UserContent::Image(image)]
|
||||
if text.text == "Describe the image"
|
||||
&& matches!(&image.data, DocumentSourceKind::Base64(data) if data == "AQIDBA==")
|
||||
));
|
||||
|
||||
let Message::Assistant { content, .. } = assistant else {
|
||||
panic!("expected an assistant tool call");
|
||||
};
|
||||
let assistant_content = content.iter().collect::<Vec<_>>();
|
||||
let [
|
||||
AssistantContent::Reasoning(reasoning),
|
||||
AssistantContent::ToolCall(call),
|
||||
] = assistant_content.as_slice()
|
||||
else {
|
||||
panic!("expected signed reasoning followed by a tool call");
|
||||
};
|
||||
assert_eq!(reasoning.display_text(), "I should inspect the manifest.");
|
||||
assert_eq!(reasoning.first_signature(), Some("signed-reasoning"));
|
||||
assert_eq!(call.id, "call-1");
|
||||
assert_eq!(call.function.name, "read_files");
|
||||
|
||||
let Message::User { content } = result else {
|
||||
panic!("expected a user tool result");
|
||||
};
|
||||
let Some(UserContent::ToolResult(result)) = content.iter().next() else {
|
||||
panic!("expected tool result content");
|
||||
};
|
||||
assert_eq!(result.id, "call-1");
|
||||
assert!(matches!(
|
||||
result.content.iter().next(),
|
||||
Some(ToolResultContent::Text(text)) if text.text == "[ERROR] permission denied"
|
||||
));
|
||||
}
|
||||
@@ -1,5 +1,9 @@
|
||||
//! Rig-backed implementations of Galaxy's provider-neutral agent runtime.
|
||||
|
||||
mod bedrock;
|
||||
mod openai_compatible;
|
||||
mod request;
|
||||
mod stream;
|
||||
|
||||
pub use bedrock::*;
|
||||
pub use openai_compatible::*;
|
||||
|
||||
@@ -1,22 +1,14 @@
|
||||
use async_trait::async_trait;
|
||||
use futures::{FutureExt, StreamExt};
|
||||
use galaxy_agent_core::{
|
||||
AgentError, AgentErrorKind, AgentEvent, AgentEventStream, AgentRuntime, ContentPart,
|
||||
ConversationMessage, MessageContent, MessageRole, RuntimeCapabilities, RuntimeDescriptor,
|
||||
RuntimeKind, StopReason, ToolCall, TurnCommand, TurnControl, TurnRequest, Usage,
|
||||
AgentError, AgentErrorKind, AgentEventStream, AgentRuntime, RuntimeCapabilities,
|
||||
RuntimeDescriptor, RuntimeKind, TurnControl, TurnRequest,
|
||||
};
|
||||
use rig_core::OneOrMany;
|
||||
use rig_core::client::CompletionClient;
|
||||
use rig_core::completion::{
|
||||
AssistantContent, CompletionError, CompletionModel, CompletionRequest, GetTokenUsage, Message,
|
||||
ToolDefinition,
|
||||
};
|
||||
use rig_core::message::{
|
||||
DocumentSourceKind, Image, ImageMediaType, MimeType, ToolResultContent, UserContent,
|
||||
};
|
||||
use rig_core::completion::{CompletionModel, CompletionRequest};
|
||||
use rig_core::providers::openai;
|
||||
use rig_core::streaming::StreamedAssistantContent;
|
||||
use uuid::Uuid;
|
||||
|
||||
use crate::request::build_completion_request as build_provider_completion_request;
|
||||
use crate::stream::start_model_turn as start_provider_model_turn;
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub struct OpenAICompatibleRuntimeConfig {
|
||||
@@ -91,143 +83,13 @@ where
|
||||
M: CompletionModel + Send + Sync + 'static,
|
||||
M::StreamingResponse: Send + Sync + 'static,
|
||||
{
|
||||
let runtime_request_id = Uuid::new_v4().to_string();
|
||||
let max_output_tokens = request.max_output_tokens.or(configured_max_output_tokens);
|
||||
let completion_request = build_completion_request(
|
||||
request,
|
||||
configured_max_output_tokens,
|
||||
supports_system_messages,
|
||||
)?;
|
||||
let stream_future = model.stream(completion_request).fuse();
|
||||
let initial_control = control.clone();
|
||||
let control_future = initial_control.receive().fuse();
|
||||
futures::pin_mut!(stream_future, control_future);
|
||||
|
||||
let mut rig_stream = futures::select_biased! {
|
||||
command = control_future => match command {
|
||||
Ok(TurnCommand::Cancel) => {
|
||||
return Ok(stopped_before_stream(runtime_request_id));
|
||||
}
|
||||
Ok(TurnCommand::Steer { .. }) | Err(_) => {
|
||||
stream_future.await.map_err(map_completion_error)?
|
||||
}
|
||||
},
|
||||
result = stream_future => result.map_err(map_completion_error)?,
|
||||
};
|
||||
|
||||
let events = async_stream::stream! {
|
||||
yield Ok(AgentEvent::TurnStarted {
|
||||
runtime_request_id,
|
||||
});
|
||||
|
||||
let mut control_open = true;
|
||||
let mut last_output_tokens = 0;
|
||||
loop {
|
||||
let next_item = rig_stream.next().fuse();
|
||||
let next_command = if control_open {
|
||||
futures::future::Either::Left(control.receive())
|
||||
} else {
|
||||
futures::future::Either::Right(futures::future::pending())
|
||||
}
|
||||
.fuse();
|
||||
futures::pin_mut!(next_item, next_command);
|
||||
|
||||
futures::select_biased! {
|
||||
command = next_command => {
|
||||
match command {
|
||||
Ok(TurnCommand::Cancel) => {
|
||||
rig_stream.cancel();
|
||||
yield Ok(AgentEvent::TurnStopped {
|
||||
reason: StopReason::Cancelled,
|
||||
});
|
||||
return;
|
||||
}
|
||||
Ok(TurnCommand::Steer { .. }) => {
|
||||
// Steering is not advertised by this runtime yet.
|
||||
}
|
||||
Err(_) => control_open = false,
|
||||
}
|
||||
}
|
||||
item = next_item => {
|
||||
let Some(item) = item else {
|
||||
yield Ok(AgentEvent::TurnStopped {
|
||||
reason: if max_output_tokens.is_some_and(|max| {
|
||||
last_output_tokens >= max
|
||||
}) {
|
||||
StopReason::MaxTokens
|
||||
} else {
|
||||
StopReason::Completed
|
||||
},
|
||||
});
|
||||
return;
|
||||
};
|
||||
|
||||
match item {
|
||||
Ok(StreamedAssistantContent::Text(text)) => {
|
||||
if !text.text.is_empty() {
|
||||
yield Ok(AgentEvent::TextDelta { text: text.text });
|
||||
}
|
||||
}
|
||||
Ok(StreamedAssistantContent::Reasoning(reasoning)) => {
|
||||
let text = reasoning.display_text();
|
||||
if !text.is_empty() {
|
||||
yield Ok(AgentEvent::ReasoningDelta { text });
|
||||
}
|
||||
}
|
||||
Ok(StreamedAssistantContent::ReasoningDelta { reasoning, .. }) => {
|
||||
if !reasoning.is_empty() {
|
||||
yield Ok(AgentEvent::ReasoningDelta { text: reasoning });
|
||||
}
|
||||
}
|
||||
Ok(StreamedAssistantContent::ToolCall { tool_call, .. }) => {
|
||||
yield Ok(AgentEvent::Tool {
|
||||
event: galaxy_agent_core::ToolEvent::Proposed {
|
||||
call: ToolCall {
|
||||
id: tool_call.id,
|
||||
name: tool_call.function.name,
|
||||
arguments: tool_call.function.arguments,
|
||||
},
|
||||
},
|
||||
});
|
||||
}
|
||||
Ok(StreamedAssistantContent::ToolCallDelta { .. }) => {
|
||||
// Rig emits a complete ToolCall after its deltas, which
|
||||
// is the canonical event Galaxy consumes.
|
||||
}
|
||||
Ok(StreamedAssistantContent::Final(response)) => {
|
||||
let mapped_usage = map_usage(response.token_usage());
|
||||
last_output_tokens = mapped_usage.output_tokens;
|
||||
yield Ok(AgentEvent::UsageUpdated {
|
||||
usage: mapped_usage,
|
||||
});
|
||||
}
|
||||
Ok(StreamedAssistantContent::Unknown(value)) => {
|
||||
yield Err(AgentError::new(
|
||||
AgentErrorKind::Protocol,
|
||||
format!("Rig returned an unsupported provider event: {value}"),
|
||||
));
|
||||
return;
|
||||
}
|
||||
Err(error) => {
|
||||
yield Err(map_completion_error(error));
|
||||
return;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
Ok(Box::pin(events))
|
||||
}
|
||||
|
||||
fn stopped_before_stream(runtime_request_id: String) -> AgentEventStream {
|
||||
Box::pin(futures::stream::iter([
|
||||
Ok(AgentEvent::TurnStarted { runtime_request_id }),
|
||||
Ok(AgentEvent::TurnStopped {
|
||||
reason: StopReason::Cancelled,
|
||||
}),
|
||||
]))
|
||||
start_provider_model_turn(model, completion_request, control, max_output_tokens).await
|
||||
}
|
||||
|
||||
fn build_completion_request(
|
||||
@@ -235,208 +97,17 @@ fn build_completion_request(
|
||||
configured_max_output_tokens: Option<u64>,
|
||||
supports_system_messages: bool,
|
||||
) -> Result<CompletionRequest, AgentError> {
|
||||
let mut messages = Vec::new();
|
||||
if let Some(system_prompt) = request.system_prompt {
|
||||
if supports_system_messages {
|
||||
messages.push(Message::System {
|
||||
content: system_prompt,
|
||||
});
|
||||
} else {
|
||||
messages.push(Message::User {
|
||||
content: OneOrMany::one(UserContent::text(system_prompt)),
|
||||
});
|
||||
}
|
||||
}
|
||||
for message in request.messages {
|
||||
messages.push(convert_message(message)?);
|
||||
}
|
||||
|
||||
let chat_history = OneOrMany::many(messages).map_err(|_| {
|
||||
AgentError::new(
|
||||
AgentErrorKind::InvalidRequest,
|
||||
"a Rig turn requires at least one conversation message",
|
||||
)
|
||||
})?;
|
||||
|
||||
Ok(CompletionRequest {
|
||||
model: Some(request.model.as_str().to_string()),
|
||||
preamble: None,
|
||||
chat_history,
|
||||
documents: Vec::new(),
|
||||
tools: request
|
||||
.tools
|
||||
.into_iter()
|
||||
.map(|tool| ToolDefinition {
|
||||
name: tool.name,
|
||||
description: tool.description,
|
||||
parameters: tool.input_schema,
|
||||
})
|
||||
.collect(),
|
||||
temperature: None,
|
||||
max_tokens: request.max_output_tokens.or(configured_max_output_tokens),
|
||||
tool_choice: None,
|
||||
additional_params: Some(serde_json::json!({
|
||||
build_provider_completion_request(
|
||||
request,
|
||||
configured_max_output_tokens,
|
||||
supports_system_messages,
|
||||
false,
|
||||
Some(serde_json::json!({
|
||||
"stream_options": { "include_usage": true }
|
||||
})),
|
||||
output_schema: None,
|
||||
})
|
||||
}
|
||||
|
||||
fn convert_message(message: ConversationMessage) -> Result<Message, AgentError> {
|
||||
match message.role {
|
||||
MessageRole::User => Ok(Message::User {
|
||||
content: user_content(message.content)?,
|
||||
}),
|
||||
MessageRole::Assistant => Ok(Message::Assistant {
|
||||
id: None,
|
||||
content: assistant_content(message.content)?,
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
fn user_content(content: MessageContent) -> Result<OneOrMany<UserContent>, AgentError> {
|
||||
let parts = match content {
|
||||
MessageContent::Text(text) => vec![UserContent::text(text)],
|
||||
MessageContent::ToolResult {
|
||||
tool_use_id,
|
||||
content,
|
||||
is_error,
|
||||
} => vec![UserContent::tool_result(
|
||||
tool_use_id,
|
||||
OneOrMany::one(ToolResultContent::text(tool_result_text(content, is_error))),
|
||||
)],
|
||||
MessageContent::MultiPart(parts) => parts
|
||||
.into_iter()
|
||||
.map(convert_user_part)
|
||||
.collect::<Result<Vec<_>, _>>()?,
|
||||
MessageContent::ToolUse { .. } => {
|
||||
return Err(invalid_role("tool use", "user"));
|
||||
}
|
||||
};
|
||||
one_or_many(parts, "user")
|
||||
}
|
||||
|
||||
fn assistant_content(content: MessageContent) -> Result<OneOrMany<AssistantContent>, AgentError> {
|
||||
let parts = match content {
|
||||
MessageContent::Text(text) => vec![AssistantContent::text(text)],
|
||||
MessageContent::ToolUse {
|
||||
tool_use_id,
|
||||
name,
|
||||
input,
|
||||
} => vec![AssistantContent::tool_call(tool_use_id, name, input)],
|
||||
MessageContent::MultiPart(parts) => parts
|
||||
.into_iter()
|
||||
.map(convert_assistant_part)
|
||||
.collect::<Result<Vec<_>, _>>()?,
|
||||
MessageContent::ToolResult { .. } => {
|
||||
return Err(invalid_role("tool result", "assistant"));
|
||||
}
|
||||
};
|
||||
one_or_many(parts, "assistant")
|
||||
}
|
||||
|
||||
fn convert_user_part(part: ContentPart) -> Result<UserContent, AgentError> {
|
||||
match part {
|
||||
ContentPart::Text(text) => Ok(UserContent::text(text)),
|
||||
ContentPart::Image { data, mime_type } => Ok(UserContent::image_raw(
|
||||
data,
|
||||
ImageMediaType::from_mime_type(&mime_type),
|
||||
None,
|
||||
)),
|
||||
ContentPart::ToolResult {
|
||||
tool_use_id,
|
||||
content,
|
||||
is_error,
|
||||
} => Ok(UserContent::tool_result(
|
||||
tool_use_id,
|
||||
OneOrMany::one(ToolResultContent::text(tool_result_text(content, is_error))),
|
||||
)),
|
||||
ContentPart::ToolUse { .. } => Err(invalid_role("tool use", "user")),
|
||||
}
|
||||
}
|
||||
|
||||
fn convert_assistant_part(part: ContentPart) -> Result<AssistantContent, AgentError> {
|
||||
match part {
|
||||
ContentPart::Text(text) => Ok(AssistantContent::text(text)),
|
||||
ContentPart::Image { data, mime_type } => Ok(AssistantContent::Image(Image {
|
||||
data: DocumentSourceKind::Raw(data),
|
||||
media_type: ImageMediaType::from_mime_type(&mime_type),
|
||||
detail: None,
|
||||
additional_params: None,
|
||||
})),
|
||||
ContentPart::ToolUse {
|
||||
tool_use_id,
|
||||
name,
|
||||
input,
|
||||
} => Ok(AssistantContent::tool_call(tool_use_id, name, input)),
|
||||
ContentPart::ToolResult { .. } => Err(invalid_role("tool result", "assistant")),
|
||||
}
|
||||
}
|
||||
|
||||
fn tool_result_text(content: String, is_error: bool) -> String {
|
||||
if is_error {
|
||||
format!("[ERROR] {content}")
|
||||
} else {
|
||||
content
|
||||
}
|
||||
}
|
||||
|
||||
fn one_or_many<T: Clone>(parts: Vec<T>, role: &str) -> Result<OneOrMany<T>, AgentError> {
|
||||
OneOrMany::many(parts).map_err(|_| {
|
||||
AgentError::new(
|
||||
AgentErrorKind::InvalidRequest,
|
||||
format!("{role} message has no content"),
|
||||
)
|
||||
})
|
||||
}
|
||||
|
||||
fn invalid_role(content: &str, role: &str) -> AgentError {
|
||||
AgentError::new(
|
||||
AgentErrorKind::InvalidRequest,
|
||||
format!("{content} content cannot appear in a {role} message"),
|
||||
)
|
||||
}
|
||||
|
||||
fn map_usage(usage: rig_core::completion::Usage) -> Usage {
|
||||
Usage {
|
||||
input_tokens: usage.input_tokens,
|
||||
output_tokens: usage.output_tokens,
|
||||
cached_input_tokens: usage.cached_input_tokens,
|
||||
cache_creation_input_tokens: usage.cache_creation_input_tokens,
|
||||
}
|
||||
}
|
||||
|
||||
fn map_completion_error(error: CompletionError) -> AgentError {
|
||||
let status = error
|
||||
.provider_response_status()
|
||||
.map(|status| status.as_u16());
|
||||
let kind = match status {
|
||||
Some(401 | 403) => AgentErrorKind::Authentication,
|
||||
Some(429) => AgentErrorKind::RateLimited,
|
||||
Some(400 | 404 | 413 | 422) => AgentErrorKind::InvalidRequest,
|
||||
Some(500..=599) => AgentErrorKind::Provider,
|
||||
Some(_) => AgentErrorKind::Provider,
|
||||
None => match &error {
|
||||
CompletionError::HttpError(_)
|
||||
| CompletionError::UrlError(_)
|
||||
| CompletionError::RequestError(_) => AgentErrorKind::Transport,
|
||||
CompletionError::JsonError(_) | CompletionError::ResponseError(_) => {
|
||||
AgentErrorKind::Protocol
|
||||
}
|
||||
CompletionError::ProviderError(_) | CompletionError::ProviderResponse(_) => {
|
||||
AgentErrorKind::Provider
|
||||
}
|
||||
_ => AgentErrorKind::Provider,
|
||||
},
|
||||
};
|
||||
let mut mapped = AgentError::new(kind, error.to_string());
|
||||
mapped.recoverable = matches!(
|
||||
kind,
|
||||
AgentErrorKind::RateLimited | AgentErrorKind::Transport
|
||||
);
|
||||
mapped
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
#[path = "openai_compatible_tests.rs"]
|
||||
mod tests;
|
||||
|
||||
@@ -1,8 +1,11 @@
|
||||
use futures::StreamExt;
|
||||
use galaxy_agent_core::{
|
||||
AgentEvent, AgentRuntime, ConversationMessage, MessageContent, MessageRole, ToolEvent,
|
||||
AgentEvent, AgentRuntime, ConversationMessage, MessageContent, MessageRole, StopReason,
|
||||
ToolEvent, TurnCommand, Usage,
|
||||
};
|
||||
use rig_core::client::CompletionClient;
|
||||
use rig_core::completion::{AssistantContent, Message};
|
||||
use rig_core::message::{ToolResultContent, UserContent};
|
||||
use rig_core::providers::openai;
|
||||
use rig_core::test_utils::MockStreamingClient;
|
||||
|
||||
|
||||
@@ -0,0 +1,211 @@
|
||||
use base64::Engine;
|
||||
use base64::prelude::BASE64_STANDARD;
|
||||
use galaxy_agent_core::{
|
||||
AgentError, AgentErrorKind, ContentPart, ConversationMessage, MessageContent, MessageRole,
|
||||
TurnRequest,
|
||||
};
|
||||
use rig_core::OneOrMany;
|
||||
use rig_core::completion::{AssistantContent, CompletionRequest, Message, ToolDefinition};
|
||||
use rig_core::message::{
|
||||
DocumentSourceKind, Image, ImageMediaType, MimeType, Reasoning, ToolResultContent, UserContent,
|
||||
};
|
||||
|
||||
pub(crate) fn build_completion_request(
|
||||
request: TurnRequest,
|
||||
configured_max_output_tokens: Option<u64>,
|
||||
supports_system_messages: bool,
|
||||
encode_images_as_base64: bool,
|
||||
additional_params: Option<serde_json::Value>,
|
||||
) -> Result<CompletionRequest, AgentError> {
|
||||
let mut messages = Vec::new();
|
||||
if let Some(system_prompt) = request.system_prompt {
|
||||
if supports_system_messages {
|
||||
messages.push(Message::System {
|
||||
content: system_prompt,
|
||||
});
|
||||
} else {
|
||||
messages.push(Message::User {
|
||||
content: OneOrMany::one(UserContent::text(system_prompt)),
|
||||
});
|
||||
}
|
||||
}
|
||||
for message in request.messages {
|
||||
messages.push(convert_message(message, encode_images_as_base64)?);
|
||||
}
|
||||
|
||||
let chat_history = OneOrMany::many(messages).map_err(|_| {
|
||||
AgentError::new(
|
||||
AgentErrorKind::InvalidRequest,
|
||||
"a Rig turn requires at least one conversation message",
|
||||
)
|
||||
})?;
|
||||
|
||||
Ok(CompletionRequest {
|
||||
model: Some(request.model.as_str().to_string()),
|
||||
preamble: None,
|
||||
chat_history,
|
||||
documents: Vec::new(),
|
||||
tools: request
|
||||
.tools
|
||||
.into_iter()
|
||||
.map(|tool| ToolDefinition {
|
||||
name: tool.name,
|
||||
description: tool.description,
|
||||
parameters: tool.input_schema,
|
||||
})
|
||||
.collect(),
|
||||
temperature: None,
|
||||
max_tokens: request.max_output_tokens.or(configured_max_output_tokens),
|
||||
tool_choice: None,
|
||||
additional_params,
|
||||
output_schema: None,
|
||||
})
|
||||
}
|
||||
|
||||
fn convert_message(
|
||||
message: ConversationMessage,
|
||||
encode_images_as_base64: bool,
|
||||
) -> Result<Message, AgentError> {
|
||||
match message.role {
|
||||
MessageRole::User => Ok(Message::User {
|
||||
content: user_content(message.content, encode_images_as_base64)?,
|
||||
}),
|
||||
MessageRole::Assistant => Ok(Message::Assistant {
|
||||
id: None,
|
||||
content: assistant_content(message.content, encode_images_as_base64)?,
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
fn user_content(
|
||||
content: MessageContent,
|
||||
encode_images_as_base64: bool,
|
||||
) -> Result<OneOrMany<UserContent>, AgentError> {
|
||||
let parts = match content {
|
||||
MessageContent::Text(text) => vec![UserContent::text(text)],
|
||||
MessageContent::ToolResult {
|
||||
tool_use_id,
|
||||
content,
|
||||
is_error,
|
||||
} => vec![UserContent::tool_result(
|
||||
tool_use_id,
|
||||
OneOrMany::one(ToolResultContent::text(tool_result_text(content, is_error))),
|
||||
)],
|
||||
MessageContent::MultiPart(parts) => parts
|
||||
.into_iter()
|
||||
.map(|part| convert_user_part(part, encode_images_as_base64))
|
||||
.collect::<Result<Vec<_>, _>>()?,
|
||||
MessageContent::ToolUse { .. } => {
|
||||
return Err(invalid_role("tool use", "user"));
|
||||
}
|
||||
};
|
||||
one_or_many(parts, "user")
|
||||
}
|
||||
|
||||
fn assistant_content(
|
||||
content: MessageContent,
|
||||
encode_images_as_base64: bool,
|
||||
) -> Result<OneOrMany<AssistantContent>, AgentError> {
|
||||
let parts = match content {
|
||||
MessageContent::Text(text) => vec![AssistantContent::text(text)],
|
||||
MessageContent::ToolUse {
|
||||
tool_use_id,
|
||||
name,
|
||||
input,
|
||||
} => vec![AssistantContent::tool_call(tool_use_id, name, input)],
|
||||
MessageContent::MultiPart(parts) => parts
|
||||
.into_iter()
|
||||
.map(|part| convert_assistant_part(part, encode_images_as_base64))
|
||||
.collect::<Result<Vec<_>, _>>()?,
|
||||
MessageContent::ToolResult { .. } => {
|
||||
return Err(invalid_role("tool result", "assistant"));
|
||||
}
|
||||
};
|
||||
one_or_many(parts, "assistant")
|
||||
}
|
||||
|
||||
fn convert_user_part(
|
||||
part: ContentPart,
|
||||
encode_images_as_base64: bool,
|
||||
) -> Result<UserContent, AgentError> {
|
||||
match part {
|
||||
ContentPart::Text(text) => Ok(UserContent::text(text)),
|
||||
ContentPart::Reasoning { .. } => Err(invalid_role("reasoning", "user")),
|
||||
ContentPart::Image { data, mime_type } => {
|
||||
let media_type = ImageMediaType::from_mime_type(&mime_type);
|
||||
if encode_images_as_base64 {
|
||||
Ok(UserContent::image_base64(
|
||||
BASE64_STANDARD.encode(data),
|
||||
media_type,
|
||||
None,
|
||||
))
|
||||
} else {
|
||||
Ok(UserContent::image_raw(data, media_type, None))
|
||||
}
|
||||
}
|
||||
ContentPart::ToolResult {
|
||||
tool_use_id,
|
||||
content,
|
||||
is_error,
|
||||
} => Ok(UserContent::tool_result(
|
||||
tool_use_id,
|
||||
OneOrMany::one(ToolResultContent::text(tool_result_text(content, is_error))),
|
||||
)),
|
||||
ContentPart::ToolUse { .. } => Err(invalid_role("tool use", "user")),
|
||||
}
|
||||
}
|
||||
|
||||
fn convert_assistant_part(
|
||||
part: ContentPart,
|
||||
encode_images_as_base64: bool,
|
||||
) -> Result<AssistantContent, AgentError> {
|
||||
match part {
|
||||
ContentPart::Text(text) => Ok(AssistantContent::text(text)),
|
||||
ContentPart::Reasoning { text, signature } => Ok(AssistantContent::Reasoning(
|
||||
Reasoning::new_with_signature(&text, signature),
|
||||
)),
|
||||
ContentPart::Image { data, mime_type } => Ok(AssistantContent::Image(Image {
|
||||
data: if encode_images_as_base64 {
|
||||
DocumentSourceKind::Base64(BASE64_STANDARD.encode(data))
|
||||
} else {
|
||||
DocumentSourceKind::Raw(data)
|
||||
},
|
||||
media_type: ImageMediaType::from_mime_type(&mime_type),
|
||||
detail: None,
|
||||
additional_params: None,
|
||||
})),
|
||||
ContentPart::ToolUse {
|
||||
tool_use_id,
|
||||
name,
|
||||
input,
|
||||
} => Ok(AssistantContent::tool_call(tool_use_id, name, input)),
|
||||
ContentPart::ToolResult { .. } => Err(invalid_role("tool result", "assistant")),
|
||||
}
|
||||
}
|
||||
|
||||
fn tool_result_text(content: String, is_error: bool) -> String {
|
||||
// Rig core does not yet carry Bedrock's optional ToolResultStatus. Keep
|
||||
// Galaxy's structured error state in the domain model and make the error
|
||||
// semantic explicit in the provider-visible result text.
|
||||
if is_error {
|
||||
format!("[ERROR] {content}")
|
||||
} else {
|
||||
content
|
||||
}
|
||||
}
|
||||
|
||||
fn one_or_many<T: Clone>(parts: Vec<T>, role: &str) -> Result<OneOrMany<T>, AgentError> {
|
||||
OneOrMany::many(parts).map_err(|_| {
|
||||
AgentError::new(
|
||||
AgentErrorKind::InvalidRequest,
|
||||
format!("{role} message has no content"),
|
||||
)
|
||||
})
|
||||
}
|
||||
|
||||
fn invalid_role(content: &str, role: &str) -> AgentError {
|
||||
AgentError::new(
|
||||
AgentErrorKind::InvalidRequest,
|
||||
format!("{content} content cannot appear in a {role} message"),
|
||||
)
|
||||
}
|
||||
@@ -0,0 +1,208 @@
|
||||
use futures::{FutureExt, StreamExt};
|
||||
use galaxy_agent_core::{
|
||||
AgentError, AgentErrorKind, AgentEvent, AgentEventStream, StopReason, ToolCall, TurnCommand,
|
||||
TurnControl, Usage,
|
||||
};
|
||||
use rig_core::completion::{CompletionError, CompletionModel, CompletionRequest, GetTokenUsage};
|
||||
use rig_core::streaming::StreamedAssistantContent;
|
||||
use uuid::Uuid;
|
||||
|
||||
pub(crate) async fn start_model_turn<M>(
|
||||
model: M,
|
||||
completion_request: CompletionRequest,
|
||||
control: TurnControl,
|
||||
max_output_tokens: Option<u64>,
|
||||
) -> Result<AgentEventStream, AgentError>
|
||||
where
|
||||
M: CompletionModel + Send + Sync + 'static,
|
||||
M::StreamingResponse: Send + Sync + 'static,
|
||||
{
|
||||
let runtime_request_id = Uuid::new_v4().to_string();
|
||||
let stream_future = model.stream(completion_request).fuse();
|
||||
let initial_control = control.clone();
|
||||
let control_future = initial_control.receive().fuse();
|
||||
futures::pin_mut!(stream_future, control_future);
|
||||
|
||||
let mut rig_stream = futures::select_biased! {
|
||||
command = control_future => match command {
|
||||
Ok(TurnCommand::Cancel) => {
|
||||
return Ok(stopped_before_stream(runtime_request_id));
|
||||
}
|
||||
Ok(TurnCommand::Steer { .. }) | Err(_) => {
|
||||
stream_future.await.map_err(map_completion_error)?
|
||||
}
|
||||
},
|
||||
result = stream_future => result.map_err(map_completion_error)?,
|
||||
};
|
||||
|
||||
let events = async_stream::stream! {
|
||||
yield Ok(AgentEvent::TurnStarted {
|
||||
runtime_request_id,
|
||||
});
|
||||
|
||||
let mut control_open = true;
|
||||
let mut last_output_tokens = 0;
|
||||
loop {
|
||||
let next_item = rig_stream.next().fuse();
|
||||
let next_command = if control_open {
|
||||
futures::future::Either::Left(control.receive())
|
||||
} else {
|
||||
futures::future::Either::Right(futures::future::pending())
|
||||
}
|
||||
.fuse();
|
||||
futures::pin_mut!(next_item, next_command);
|
||||
|
||||
futures::select_biased! {
|
||||
command = next_command => {
|
||||
match command {
|
||||
Ok(TurnCommand::Cancel) => {
|
||||
rig_stream.cancel();
|
||||
yield Ok(AgentEvent::TurnStopped {
|
||||
reason: StopReason::Cancelled,
|
||||
});
|
||||
return;
|
||||
}
|
||||
Ok(TurnCommand::Steer { .. }) => {
|
||||
// Steering is not advertised by provider runtimes yet.
|
||||
}
|
||||
Err(_) => control_open = false,
|
||||
}
|
||||
}
|
||||
item = next_item => {
|
||||
let Some(item) = item else {
|
||||
yield Ok(AgentEvent::TurnStopped {
|
||||
reason: if max_output_tokens.is_some_and(|max| {
|
||||
last_output_tokens >= max
|
||||
}) {
|
||||
StopReason::MaxTokens
|
||||
} else {
|
||||
StopReason::Completed
|
||||
},
|
||||
});
|
||||
return;
|
||||
};
|
||||
|
||||
match item {
|
||||
Ok(StreamedAssistantContent::Text(text)) => {
|
||||
if !text.text.is_empty() {
|
||||
yield Ok(AgentEvent::TextDelta { text: text.text });
|
||||
}
|
||||
}
|
||||
Ok(StreamedAssistantContent::Reasoning(reasoning)) => {
|
||||
let text = reasoning.display_text();
|
||||
yield Ok(AgentEvent::ReasoningCompleted {
|
||||
text,
|
||||
signature: reasoning.first_signature().map(str::to_string),
|
||||
});
|
||||
}
|
||||
Ok(StreamedAssistantContent::ReasoningDelta { reasoning, .. }) => {
|
||||
if !reasoning.is_empty() {
|
||||
yield Ok(AgentEvent::ReasoningDelta { text: reasoning });
|
||||
}
|
||||
}
|
||||
Ok(StreamedAssistantContent::ToolCall { tool_call, .. }) => {
|
||||
yield Ok(AgentEvent::Tool {
|
||||
event: galaxy_agent_core::ToolEvent::Proposed {
|
||||
call: ToolCall {
|
||||
id: tool_call.id,
|
||||
name: tool_call.function.name,
|
||||
arguments: tool_call.function.arguments,
|
||||
},
|
||||
},
|
||||
});
|
||||
}
|
||||
Ok(StreamedAssistantContent::ToolCallDelta { .. }) => {
|
||||
// Rig emits a complete ToolCall after its deltas, which
|
||||
// is the canonical event Galaxy consumes.
|
||||
}
|
||||
Ok(StreamedAssistantContent::Final(response)) => {
|
||||
let mapped_usage = map_usage(response.token_usage());
|
||||
last_output_tokens = mapped_usage.output_tokens;
|
||||
yield Ok(AgentEvent::UsageUpdated {
|
||||
usage: mapped_usage,
|
||||
});
|
||||
}
|
||||
Ok(StreamedAssistantContent::Unknown(value)) => {
|
||||
yield Err(AgentError::new(
|
||||
AgentErrorKind::Protocol,
|
||||
format!("Rig returned an unsupported provider event: {value}"),
|
||||
));
|
||||
return;
|
||||
}
|
||||
Err(error) => {
|
||||
if let Some(reason) = completion_error_stop_reason(&error) {
|
||||
yield Ok(AgentEvent::TurnStopped { reason });
|
||||
return;
|
||||
}
|
||||
yield Err(map_completion_error(error));
|
||||
return;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
Ok(Box::pin(events))
|
||||
}
|
||||
|
||||
fn stopped_before_stream(runtime_request_id: String) -> AgentEventStream {
|
||||
Box::pin(futures::stream::iter([
|
||||
Ok(AgentEvent::TurnStarted { runtime_request_id }),
|
||||
Ok(AgentEvent::TurnStopped {
|
||||
reason: StopReason::Cancelled,
|
||||
}),
|
||||
]))
|
||||
}
|
||||
|
||||
pub(crate) fn map_usage(usage: rig_core::completion::Usage) -> Usage {
|
||||
Usage {
|
||||
input_tokens: usage.input_tokens,
|
||||
output_tokens: usage.output_tokens,
|
||||
cached_input_tokens: usage.cached_input_tokens,
|
||||
cache_creation_input_tokens: usage.cache_creation_input_tokens,
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn completion_error_stop_reason(error: &CompletionError) -> Option<StopReason> {
|
||||
match error {
|
||||
// rig-bedrock 0.40 currently surfaces Bedrock's MaxTokens stop as a
|
||||
// provider error. Normalize it here so the UI sees the same semantic
|
||||
// stop reason as every other Rig-backed provider.
|
||||
CompletionError::ProviderError(message) if message == "Exceeded max tokens" => {
|
||||
Some(StopReason::MaxTokens)
|
||||
}
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn map_completion_error(error: CompletionError) -> AgentError {
|
||||
let status = error
|
||||
.provider_response_status()
|
||||
.map(|status| status.as_u16());
|
||||
let kind = match status {
|
||||
Some(401 | 403) => AgentErrorKind::Authentication,
|
||||
Some(429) => AgentErrorKind::RateLimited,
|
||||
Some(400 | 404 | 413 | 422) => AgentErrorKind::InvalidRequest,
|
||||
Some(500..=599) => AgentErrorKind::Provider,
|
||||
Some(_) => AgentErrorKind::Provider,
|
||||
None => match &error {
|
||||
CompletionError::HttpError(_)
|
||||
| CompletionError::UrlError(_)
|
||||
| CompletionError::RequestError(_) => AgentErrorKind::Transport,
|
||||
CompletionError::JsonError(_) | CompletionError::ResponseError(_) => {
|
||||
AgentErrorKind::Protocol
|
||||
}
|
||||
CompletionError::ProviderError(_) | CompletionError::ProviderResponse(_) => {
|
||||
AgentErrorKind::Provider
|
||||
}
|
||||
_ => AgentErrorKind::Provider,
|
||||
},
|
||||
};
|
||||
let mut mapped = AgentError::new(kind, error.to_string());
|
||||
mapped.recoverable = matches!(
|
||||
kind,
|
||||
AgentErrorKind::RateLimited | AgentErrorKind::Transport
|
||||
);
|
||||
mapped
|
||||
}
|
||||
Reference in New Issue
Block a user