Add unified models UI and Rig Bedrock runtime

This commit is contained in:
2026-08-04 17:25:19 -05:00
parent a3c68e9c30
commit b0ad07f6f2
41 changed files with 2122 additions and 564 deletions
+4
View File
@@ -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"] }
+162
View File
@@ -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"
));
}
+4
View File
@@ -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::*;
+13 -342
View File
@@ -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;
+211
View File
@@ -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"),
)
}
+208
View File
@@ -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
}