feat: introduce Rig agent runtime migration

This commit is contained in:
2026-08-04 02:15:18 -05:00
parent d9cf0d8ae3
commit 4c7270db8d
39 changed files with 2551 additions and 211 deletions
@@ -0,0 +1,432 @@
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,
};
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::providers::openai;
use rig_core::streaming::StreamedAssistantContent;
use uuid::Uuid;
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct OpenAICompatibleRuntimeConfig {
pub base_url: String,
pub api_key: Option<String>,
pub model: String,
pub max_output_tokens: Option<u64>,
pub supports_system_messages: bool,
}
#[derive(Clone, Debug)]
pub struct OpenAICompatibleRuntime {
config: OpenAICompatibleRuntimeConfig,
descriptor: RuntimeDescriptor,
}
impl OpenAICompatibleRuntime {
pub fn new(config: OpenAICompatibleRuntimeConfig) -> Self {
let descriptor = RuntimeDescriptor {
id: format!("rig-openai-compatible:{}", config.model),
display_name: format!("Rig / {}", config.model),
kind: RuntimeKind::Provider,
capabilities: RuntimeCapabilities {
model_selection: true,
session_resume: false,
steering: false,
tool_permissions: false,
},
};
Self { config, descriptor }
}
}
#[async_trait]
impl AgentRuntime for OpenAICompatibleRuntime {
fn descriptor(&self) -> &RuntimeDescriptor {
&self.descriptor
}
async fn start_turn(
&self,
request: TurnRequest,
control: TurnControl,
) -> Result<AgentEventStream, AgentError> {
let client = openai::CompletionsClient::builder()
// Rig 0.40 requires an API-key builder value. An empty key preserves
// compatibility with unauthenticated local OpenAI-compatible servers.
.api_key(self.config.api_key.as_deref().unwrap_or_default())
.base_url(&self.config.base_url)
.build()
.map_err(|error| AgentError::new(AgentErrorKind::Configuration, error.to_string()))?;
let model = client.completion_model(&self.config.model);
start_model_turn(
model,
request,
control,
self.config.max_output_tokens,
self.config.supports_system_messages,
)
.await
}
}
async fn start_model_turn<M>(
model: M,
request: TurnRequest,
control: TurnControl,
configured_max_output_tokens: Option<u64>,
supports_system_messages: bool,
) -> Result<AgentEventStream, AgentError>
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::ToolProposed {
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,
}),
]))
}
fn build_completion_request(
request: TurnRequest,
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!({
"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,
..
} => vec![UserContent::tool_result(
tool_use_id,
OneOrMany::one(ToolResultContent::text(content)),
)],
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,
..
} => Ok(UserContent::tool_result(
tool_use_id,
OneOrMany::one(ToolResultContent::text(content)),
)),
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 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;