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
+5
View File
@@ -0,0 +1,5 @@
mod provider;
mod rig;
pub(crate) use provider::ProviderRuntime;
pub(crate) use rig::rig_openai_response_stream;
+27
View File
@@ -0,0 +1,27 @@
use futures::channel::oneshot;
use crate::ai::agent::api::{self, ConvertToAPITypeError};
use crate::ai::provider::ProviderConfig;
/// Application-facing provider runtime dispatcher.
///
/// OpenAI-compatible models can opt into the provider-neutral Rig runtime;
/// other models continue through their current translators while migration is
/// in progress. Both paths preserve the existing UI response stream contract.
pub(crate) struct ProviderRuntime {
provider_config: ProviderConfig,
}
impl ProviderRuntime {
pub(crate) fn new(provider_config: ProviderConfig) -> Self {
Self { provider_config }
}
pub(crate) async fn start_turn(
self,
params: api::RequestParams,
cancellation_rx: oneshot::Receiver<()>,
) -> Result<api::ResponseStream, ConvertToAPITypeError> {
api::generate_multi_agent_output(self.provider_config, params, cancellation_rx).await
}
}
+319
View File
@@ -0,0 +1,319 @@
use std::sync::Arc;
use futures::channel::oneshot;
use futures::{FutureExt, StreamExt};
use galaxy_agent_core::{
turn_control, AgentError, AgentEvent, AgentRuntime, MessageContent, MessageRole, StopReason,
TurnCommand, TurnRequest, Usage,
};
use galaxy_agent_rig::{OpenAICompatibleRuntime, OpenAICompatibleRuntimeConfig};
use uuid::Uuid;
use warp_multi_agent_api::response_event::stream_finished;
use warp_multi_agent_api::{self as api, ClientAction, ResponseEvent};
use crate::ai::agent::api::{Event, RequestParams, ResponseStream};
use crate::ai::bedrock::response_translator::{
build_add_agent_output_message, build_append_text, build_create_task, build_stream_init,
build_user_query_message,
};
use crate::ai::openai::client::OpenAIClientConfig;
use crate::ai::openai::response_translator::{build_stream_finished, StreamUsage};
use crate::ai::openai::translator::{prepare_turn, PreparedTurn, TranslatorRequest};
use crate::ai::provider::types::ConversationMessage;
use crate::server::server_api::AIApiError;
pub(crate) fn rig_openai_response_stream(
config: OpenAIClientConfig,
params: RequestParams,
request: &mut api::Request,
cancellation_rx: oneshot::Receiver<()>,
) -> ResponseStream {
let translator_request = TranslatorRequest {
config: config.clone(),
model_id: params.model.as_str().to_string(),
root_task_id: params.root_task_id,
message_history: params.bedrock_message_history,
tool_result_archive: params.bedrock_tool_result_archive,
progressive_summary: params.bedrock_progressive_summary,
messages_sent: params.bedrock_messages_sent,
global_rules: params.global_rules,
};
let PreparedTurn {
task_id,
needs_create_task,
user_query,
messages,
system_prompt,
tools: _,
model_id,
persistent_message_count,
} = prepare_turn(&translator_request, request);
store_messages_sent(
&translator_request.messages_sent,
&messages,
persistent_message_count,
);
let conversation_id = request
.metadata
.as_ref()
.map(|metadata| metadata.conversation_id.clone())
.filter(|id| !id.is_empty());
let mut turn_request = TurnRequest::new(model_id.clone(), messages);
turn_request.conversation_id = conversation_id.clone();
turn_request.system_prompt = system_prompt;
// Phase 2 deliberately validates the model streaming seam. Galaxy tool
// execution moves behind AgentRuntime in Phase 3; exposing the legacy tool
// list here would split ownership across both systems.
turn_request.tools = Vec::new();
turn_request.max_output_tokens = config.max_output_tokens.map(u64::from);
let runtime = OpenAICompatibleRuntime::new(OpenAICompatibleRuntimeConfig {
base_url: config.base_url,
api_key: config.api_key,
model: model_id.clone(),
max_output_tokens: config.max_output_tokens.map(u64::from),
supports_system_messages: config.supports_system_messages,
});
let messages_sent = translator_request.messages_sent;
let max_context_tokens = config.max_input_tokens;
let stream = async_stream::stream! {
let (control_sender, control) = turn_control();
let start_future = runtime.start_turn(turn_request, control).fuse();
let cancel_future = cancellation_rx.fuse();
futures::pin_mut!(start_future, cancel_future);
let mut agent_events = futures::select_biased! {
_ = cancel_future => {
let _ = control_sender.try_send(TurnCommand::Cancel);
match start_future.await {
Ok(stream) => stream,
Err(error) => {
yield Err(agent_error(error));
return;
}
}
}
result = start_future => match result {
Ok(stream) => stream,
Err(error) => {
yield Err(agent_error(error));
return;
}
},
};
let request_id = Uuid::new_v4().to_string();
let conversation_id = conversation_id.unwrap_or_else(|| Uuid::new_v4().to_string());
let mut initialized = false;
let mut current_text_message_id: Option<String> = None;
let mut current_reasoning_message_id: Option<String> = None;
let mut full_text = String::new();
let mut usage = Usage::default();
loop {
let next_event = agent_events.next().fuse();
futures::pin_mut!(next_event);
futures::select_biased! {
_ = cancel_future => {
let _ = control_sender.try_send(TurnCommand::Cancel);
}
event = next_event => {
let Some(event) = event else {
yield Err(Arc::new(AIApiError::UnexpectedEof));
return;
};
let event = match event {
Ok(event) => event,
Err(error) => {
yield Err(agent_error(error));
return;
}
};
match event {
AgentEvent::TurnStarted { .. } => {
initialized = true;
yield Ok(build_stream_init(&request_id, &conversation_id));
if needs_create_task {
yield Ok(build_create_task(&task_id));
}
if let Some(user_query) = &user_query {
yield Ok(build_user_query_message(&task_id, user_query));
}
}
AgentEvent::TextDelta { text } => {
full_text.push_str(&text);
if let Some(message_id) = &current_text_message_id {
yield Ok(build_append_text(&task_id, message_id, &text));
} else {
let message_id = Uuid::new_v4().to_string();
yield Ok(build_add_agent_output_message(&task_id, &message_id, &text));
current_text_message_id = Some(message_id);
}
}
AgentEvent::ReasoningDelta { text } => {
if let Some(message_id) = &current_reasoning_message_id {
yield Ok(build_append_reasoning(&task_id, message_id, &text));
} else {
let message_id = Uuid::new_v4().to_string();
yield Ok(build_add_reasoning(&task_id, &message_id, &text));
current_reasoning_message_id = Some(message_id);
}
}
AgentEvent::UsageUpdated { usage: updated } => usage = updated,
AgentEvent::TurnStopped { reason } => {
if !initialized {
yield Ok(build_stream_init(&request_id, &conversation_id));
}
store_assistant_text(&messages_sent, full_text);
yield Ok(build_stream_finished(
map_stop_reason(reason),
StreamUsage {
input_tokens: saturating_i32(usage.input_tokens),
output_tokens: saturating_i32(usage.output_tokens),
cache_read_tokens: saturating_i32(usage.cached_input_tokens),
cache_write_tokens: saturating_i32(
usage.cache_creation_input_tokens,
),
cost_in_cents: 0.0,
model_id,
max_context_tokens,
},
));
return;
}
AgentEvent::ToolProposed { .. }
| AgentEvent::PermissionRequested { .. }
| AgentEvent::ToolStarted { .. }
| AgentEvent::ToolCompleted { .. } => {
yield Err(agent_error(AgentError::new(
galaxy_agent_core::AgentErrorKind::Protocol,
"the Phase 2 Rig runtime emitted a tool event while tools are disabled",
)));
return;
}
}
}
}
}
};
Box::pin(stream)
}
fn store_messages_sent(
messages_sent: &std::sync::Arc<std::sync::Mutex<Vec<ConversationMessage>>>,
messages: &[ConversationMessage],
persistent_message_count: usize,
) {
let Ok(mut sent) = messages_sent.lock() else {
return;
};
if persistent_message_count > 0 && messages.len() >= persistent_message_count {
*sent = messages[messages.len() - persistent_message_count..].to_vec();
} else {
*sent = messages.to_vec();
}
}
fn store_assistant_text(
messages_sent: &std::sync::Arc<std::sync::Mutex<Vec<ConversationMessage>>>,
text: String,
) {
if text.is_empty() {
return;
}
if let Ok(mut sent) = messages_sent.lock() {
sent.push(ConversationMessage {
role: MessageRole::Assistant,
content: MessageContent::Text(text),
});
}
}
fn build_add_reasoning(task_id: &str, message_id: &str, text: &str) -> ResponseEvent {
reasoning_action(task_id, message_id, text, false)
}
fn build_append_reasoning(task_id: &str, message_id: &str, text: &str) -> ResponseEvent {
reasoning_action(task_id, message_id, text, true)
}
fn reasoning_action(task_id: &str, message_id: &str, text: &str, append: bool) -> ResponseEvent {
let message = api::Message {
id: message_id.to_string(),
task_id: task_id.to_string(),
request_id: String::new(),
timestamp: None,
server_message_data: String::new(),
citations: Vec::new(),
fetched_memories: Vec::new(),
message: Some(api::message::Message::AgentReasoning(
api::message::AgentReasoning {
reasoning: text.to_string(),
finished_duration: None,
},
)),
};
let action = if append {
api::client_action::Action::AppendToMessageContent(
api::client_action::AppendToMessageContent {
task_id: task_id.to_string(),
message: Some(message),
mask: Some(prost_types::FieldMask {
paths: vec!["agent_reasoning.reasoning".to_string()],
}),
},
)
} else {
api::client_action::Action::AddMessagesToTask(api::client_action::AddMessagesToTask {
task_id: task_id.to_string(),
messages: vec![message],
})
};
ResponseEvent {
r#type: Some(api::response_event::Type::ClientActions(
api::response_event::ClientActions {
actions: vec![ClientAction {
action: Some(action),
}],
},
)),
}
}
fn map_stop_reason(reason: StopReason) -> stream_finished::Reason {
match reason {
StopReason::Completed => stream_finished::Reason::Done(stream_finished::Done {}),
StopReason::MaxTokens => {
stream_finished::Reason::MaxTokenLimit(stream_finished::ReachedMaxTokenLimit {})
}
StopReason::ContextWindowExceeded => stream_finished::Reason::ContextWindowExceeded(
stream_finished::ContextWindowExceeded {},
),
StopReason::Cancelled
| StopReason::Refusal
| StopReason::ToolLoopLimit
| StopReason::Other(_) => stream_finished::Reason::Other(stream_finished::Other {}),
}
}
fn saturating_i32(value: u64) -> i32 {
i32::try_from(value).unwrap_or(i32::MAX)
}
fn agent_error(error: AgentError) -> Arc<AIApiError> {
Arc::new(
AIApiError::Stream {
stream_type: "rig_openai_compatible",
source: anyhow::anyhow!(error),
}
.into_quota_limit_if_provider_budget_exhausted(),
)
}
#[cfg(test)]
#[path = "rig_tests.rs"]
mod tests;
+59
View File
@@ -0,0 +1,59 @@
use galaxy_agent_core::StopReason;
use warp_multi_agent_api::response_event::stream_finished;
use super::{build_add_reasoning, build_append_reasoning, map_stop_reason, saturating_i32};
#[test]
fn stop_reasons_map_to_the_existing_ui_contract() {
assert!(matches!(
map_stop_reason(StopReason::Completed),
stream_finished::Reason::Done(_)
));
assert!(matches!(
map_stop_reason(StopReason::MaxTokens),
stream_finished::Reason::MaxTokenLimit(_)
));
assert!(matches!(
map_stop_reason(StopReason::Cancelled),
stream_finished::Reason::Other(_)
));
}
#[test]
fn token_counts_saturate_at_the_proto_limit() {
assert_eq!(saturating_i32(u64::MAX), i32::MAX);
}
#[test]
fn reasoning_events_match_the_existing_ui_message_contract() {
let add = build_add_reasoning("task", "message", "think");
let append = build_append_reasoning("task", "message", " more");
let Some(warp_multi_agent_api::response_event::Type::ClientActions(add)) = add.r#type else {
panic!("expected client actions");
};
let Some(warp_multi_agent_api::client_action::Action::AddMessagesToTask(add)) =
&add.actions[0].action
else {
panic!("expected add-message action");
};
assert!(matches!(
add.messages[0].message.as_ref(),
Some(warp_multi_agent_api::message::Message::AgentReasoning(reasoning))
if reasoning.reasoning == "think"
));
let Some(warp_multi_agent_api::response_event::Type::ClientActions(append)) = append.r#type
else {
panic!("expected client actions");
};
let Some(warp_multi_agent_api::client_action::Action::AppendToMessageContent(append)) =
&append.actions[0].action
else {
panic!("expected append-message action");
};
assert_eq!(
append.mask.as_ref().unwrap().paths,
["agent_reasoning.reasoning"]
);
}