feat: introduce Rig agent runtime migration
This commit is contained in:
@@ -25,7 +25,7 @@ pub async fn generate_multi_agent_output(
|
||||
let supported_cli_agent_tools =
|
||||
supported_tools_override.unwrap_or_else(|| get_supported_cli_agent_tools(¶ms));
|
||||
let mut logging_metadata = HashMap::new();
|
||||
if let Some(metadata) = params.metadata {
|
||||
if let Some(ref metadata) = params.metadata {
|
||||
logging_metadata.insert(
|
||||
"is_autodetected_user_query".to_owned(),
|
||||
prost_types::Value {
|
||||
@@ -56,6 +56,12 @@ pub async fn generate_multi_agent_output(
|
||||
redaction::redact_inputs(&mut params.input);
|
||||
}
|
||||
|
||||
let rig_params = matches!(
|
||||
&provider_config,
|
||||
ProviderConfig::OpenAI(config) if config.use_rig
|
||||
)
|
||||
.then(|| params.clone());
|
||||
|
||||
let mut request = api::Request {
|
||||
task_context: Some(api::request::TaskContext {
|
||||
tasks: params.tasks,
|
||||
@@ -138,6 +144,14 @@ pub async fn generate_multi_agent_output(
|
||||
};
|
||||
|
||||
match provider_config {
|
||||
ProviderConfig::OpenAI(config) if config.use_rig => {
|
||||
Ok(crate::ai::runtime::rig_openai_response_stream(
|
||||
config,
|
||||
rig_params.expect("Rig request parameters should be retained for a Rig model"),
|
||||
&mut request,
|
||||
cancellation_rx,
|
||||
))
|
||||
}
|
||||
ProviderConfig::OpenAI(config) => {
|
||||
let translator_request = openai_translator::TranslatorRequest {
|
||||
config,
|
||||
|
||||
@@ -24,7 +24,7 @@ use crate::ai::acp::{
|
||||
resolve_acp_permissions, validate_acp_dispatch, validate_acp_launch_identity, AcpRuntimeModel,
|
||||
AcpSessionHandleSlot, AcpSessionMetadata, AcpSteeringRequest, GalaxyMcpTarget,
|
||||
};
|
||||
use crate::ai::agent::api::{self, generate_multi_agent_output, ConvertToAPITypeError};
|
||||
use crate::ai::agent::api::{self, ConvertToAPITypeError};
|
||||
use crate::ai::agent::conversation::AIConversationId;
|
||||
#[cfg(not(target_family = "wasm"))]
|
||||
use crate::ai::agent::AIAgentInput;
|
||||
@@ -35,6 +35,7 @@ use crate::ai::blocklist::BlocklistAIPermissions;
|
||||
use crate::ai::llms::{LLMId, LLMPreferences};
|
||||
use crate::ai::openai::client::OpenAIClientConfig;
|
||||
use crate::ai::provider::ProviderConfig;
|
||||
use crate::ai::runtime::ProviderRuntime;
|
||||
use crate::network::NetworkStatus;
|
||||
#[cfg(not(target_family = "wasm"))]
|
||||
use crate::pane_group::PaneGroup;
|
||||
@@ -233,6 +234,8 @@ impl ResponseStream {
|
||||
model: Some(model_id.to_string()),
|
||||
max_input_tokens: client_config.max_input_tokens,
|
||||
max_output_tokens: client_config.max_output_tokens,
|
||||
use_rig: client_config.use_rig,
|
||||
supports_system_messages: client_config.supports_system_messages,
|
||||
});
|
||||
}
|
||||
}
|
||||
@@ -400,15 +403,16 @@ impl ResponseStream {
|
||||
cancellation_rx: oneshot::Receiver<()>,
|
||||
ctx: &mut ModelContext<Self>,
|
||||
) {
|
||||
let _ =
|
||||
ctx.spawn(
|
||||
async move {
|
||||
generate_multi_agent_output(provider_config, params, cancellation_rx).await
|
||||
},
|
||||
move |me, stream, ctx| {
|
||||
me.handle_response_stream_result(request_id, stream, ctx);
|
||||
},
|
||||
);
|
||||
let _ = ctx.spawn(
|
||||
async move {
|
||||
ProviderRuntime::new(provider_config)
|
||||
.start_turn(params, cancellation_rx)
|
||||
.await
|
||||
},
|
||||
move |me, stream, ctx| {
|
||||
me.handle_response_stream_result(request_id, stream, ctx);
|
||||
},
|
||||
);
|
||||
}
|
||||
|
||||
pub fn new(
|
||||
@@ -604,15 +608,16 @@ impl ResponseStream {
|
||||
self.current_request_id = Some(request_id);
|
||||
let params = self.params.clone();
|
||||
let provider_config = Self::resolve_provider_config(params.model.as_str(), ctx);
|
||||
let _ =
|
||||
ctx.spawn(
|
||||
async move {
|
||||
generate_multi_agent_output(provider_config, params, cancellation_rx).await
|
||||
},
|
||||
move |me, stream, ctx| {
|
||||
me.handle_response_stream_result(request_id, stream, ctx);
|
||||
},
|
||||
);
|
||||
let _ = ctx.spawn(
|
||||
async move {
|
||||
ProviderRuntime::new(provider_config)
|
||||
.start_turn(params, cancellation_rx)
|
||||
.await
|
||||
},
|
||||
move |me, stream, ctx| {
|
||||
me.handle_response_stream_result(request_id, stream, ctx);
|
||||
},
|
||||
);
|
||||
}
|
||||
|
||||
fn should_fallback_to_coding_model(
|
||||
|
||||
@@ -168,6 +168,8 @@ impl CrosscheckReviewer {
|
||||
model: Some(model_id.to_string()),
|
||||
max_input_tokens: client_config.max_input_tokens,
|
||||
max_output_tokens: Some(REVIEWER_MAX_OUTPUT_TOKENS),
|
||||
use_rig: client_config.use_rig,
|
||||
supports_system_messages: client_config.supports_system_messages,
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
@@ -999,6 +999,8 @@ impl LLMPreferences {
|
||||
model: None, // filled per-request from model_id
|
||||
max_input_tokens: Some(openai_model_context_size(model)),
|
||||
max_output_tokens: model.max_output_tokens,
|
||||
use_rig: model.use_rig,
|
||||
supports_system_messages: model.supports_system_messages(),
|
||||
};
|
||||
self.openai_provider_routing
|
||||
.insert(model.model_id.clone(), client_config);
|
||||
@@ -2115,6 +2117,8 @@ async fn fetch_from_litellm_model_info(
|
||||
max_input_tokens,
|
||||
max_output_tokens,
|
||||
provider,
|
||||
use_rig: false,
|
||||
supports_system_messages: model_info["supports_system_messages"].as_bool(),
|
||||
})
|
||||
})
|
||||
.collect();
|
||||
@@ -2236,6 +2240,8 @@ async fn fetch_from_openai_models(
|
||||
max_input_tokens,
|
||||
max_output_tokens,
|
||||
provider,
|
||||
use_rig: false,
|
||||
supports_system_messages: m["supports_system_messages"].as_bool(),
|
||||
})
|
||||
})
|
||||
.collect();
|
||||
|
||||
@@ -55,6 +55,7 @@ pub(crate) mod remote_agent_context;
|
||||
pub(crate) mod remote_context_files;
|
||||
pub mod request_usage_model;
|
||||
pub(crate) mod restored_conversations;
|
||||
pub(crate) mod runtime;
|
||||
pub(crate) mod skills;
|
||||
pub(crate) mod voice;
|
||||
pub use agent_tips::*;
|
||||
|
||||
@@ -11,6 +11,8 @@ pub struct OpenAIClientConfig {
|
||||
pub model: Option<String>,
|
||||
pub max_input_tokens: Option<u32>,
|
||||
pub max_output_tokens: Option<u32>,
|
||||
pub use_rig: bool,
|
||||
pub supports_system_messages: bool,
|
||||
}
|
||||
|
||||
pub struct OpenAIClient {
|
||||
|
||||
@@ -33,14 +33,14 @@ pub struct OpenAIStreamContext {
|
||||
pub tool_result_archive: Vec<ConversationMessage>,
|
||||
}
|
||||
|
||||
struct StreamUsage {
|
||||
input_tokens: i32,
|
||||
output_tokens: i32,
|
||||
cache_read_tokens: i32,
|
||||
cache_write_tokens: i32,
|
||||
cost_in_cents: f32,
|
||||
model_id: String,
|
||||
max_context_tokens: Option<u32>,
|
||||
pub(crate) struct StreamUsage {
|
||||
pub(crate) input_tokens: i32,
|
||||
pub(crate) output_tokens: i32,
|
||||
pub(crate) cache_read_tokens: i32,
|
||||
pub(crate) cache_write_tokens: i32,
|
||||
pub(crate) cost_in_cents: f32,
|
||||
pub(crate) model_id: String,
|
||||
pub(crate) max_context_tokens: Option<u32>,
|
||||
}
|
||||
|
||||
pub fn openai_stream_to_response_events(
|
||||
@@ -533,7 +533,10 @@ fn build_tool_call_message(
|
||||
)
|
||||
}
|
||||
|
||||
fn build_stream_finished(reason: stream_finished::Reason, usage: StreamUsage) -> ResponseEvent {
|
||||
pub(crate) fn build_stream_finished(
|
||||
reason: stream_finished::Reason,
|
||||
usage: StreamUsage,
|
||||
) -> ResponseEvent {
|
||||
let StreamUsage {
|
||||
input_tokens,
|
||||
output_tokens,
|
||||
|
||||
@@ -24,27 +24,32 @@ pub struct TranslatorRequest {
|
||||
pub global_rules: Vec<(String, String)>,
|
||||
}
|
||||
|
||||
pub async fn execute(
|
||||
params: TranslatorRequest,
|
||||
request: &mut api::Request,
|
||||
) -> Result<ResponseStream, OpenAIError> {
|
||||
let client = OpenAIClient::from_config(params.config.clone());
|
||||
pub(crate) struct PreparedTurn {
|
||||
pub(crate) task_id: String,
|
||||
pub(crate) needs_create_task: bool,
|
||||
pub(crate) user_query: Option<String>,
|
||||
pub(crate) messages: Vec<ConversationMessage>,
|
||||
pub(crate) system_prompt: Option<String>,
|
||||
pub(crate) tools: Vec<crate::ai::provider::types::ToolDefinition>,
|
||||
pub(crate) model_id: String,
|
||||
pub(crate) persistent_message_count: usize,
|
||||
}
|
||||
|
||||
let task_id = params.root_task_id.unwrap_or_else(|| {
|
||||
pub(crate) fn prepare_turn(params: &TranslatorRequest, request: &mut api::Request) -> PreparedTurn {
|
||||
let task_id = params.root_task_id.clone().unwrap_or_else(|| {
|
||||
request
|
||||
.task_context
|
||||
.as_ref()
|
||||
.and_then(|tc| tc.tasks.first())
|
||||
.map(|t| t.id.clone())
|
||||
.map(|task| task.id.clone())
|
||||
.unwrap_or_else(|| uuid::Uuid::new_v4().to_string())
|
||||
});
|
||||
|
||||
let needs_create_task = request
|
||||
.task_context
|
||||
.as_ref()
|
||||
.map(|tc| tc.tasks.is_empty())
|
||||
.map(|task_context| task_context.tasks.is_empty())
|
||||
.unwrap_or(true);
|
||||
|
||||
let model_id = if params.model_id.is_empty() || params.model_id == "auto" {
|
||||
params
|
||||
.config
|
||||
@@ -52,7 +57,6 @@ pub async fn execute(
|
||||
.clone()
|
||||
.unwrap_or_else(|| "anthropic/claude-sonnet-4-6".to_string())
|
||||
} else {
|
||||
// If a model override is configured in settings, use it
|
||||
params
|
||||
.config
|
||||
.model
|
||||
@@ -60,25 +64,17 @@ pub async fn execute(
|
||||
.unwrap_or_else(|| params.model_id.clone())
|
||||
};
|
||||
|
||||
log::info!(
|
||||
"[openai] Translator: model={model_id}, task_id={task_id}, needs_create_task={needs_create_task}"
|
||||
);
|
||||
|
||||
request_translator::inject_input_messages_into_task(request);
|
||||
|
||||
let new_input_messages = request_translator::extract_new_input_messages(request);
|
||||
let new_input_count = new_input_messages.len();
|
||||
|
||||
let persistent_message_count = params.message_history.len() + new_input_messages.len();
|
||||
let mut messages = Vec::new();
|
||||
|
||||
// Prepend progressive summary as first message pair if present
|
||||
if let Some(ref summary) = params.progressive_summary {
|
||||
if let Some(summary) = ¶ms.progressive_summary {
|
||||
messages.push(ConversationMessage {
|
||||
role: MessageRole::User,
|
||||
content: MessageContent::Text(format!(
|
||||
"<conversation-history-summary>\n{}\n</conversation-history-summary>\n\n\
|
||||
The above summarizes earlier conversation history. The detailed messages below are the most recent exchanges.",
|
||||
summary
|
||||
"<conversation-history-summary>\n{summary}\n</conversation-history-summary>\n\n\
|
||||
The above summarizes earlier conversation history. The detailed messages below are the most recent exchanges."
|
||||
)),
|
||||
});
|
||||
messages.push(ConversationMessage {
|
||||
@@ -90,26 +86,44 @@ pub async fn execute(
|
||||
});
|
||||
}
|
||||
|
||||
let history_len = params.message_history.len();
|
||||
messages.extend(params.message_history);
|
||||
|
||||
if !new_input_messages.is_empty() {
|
||||
log::info!(
|
||||
"[openai] Appending {} new input messages to history of {}",
|
||||
new_input_messages.len(),
|
||||
history_len
|
||||
);
|
||||
messages.extend(new_input_messages);
|
||||
}
|
||||
|
||||
messages.extend(params.message_history.clone());
|
||||
messages.extend(new_input_messages);
|
||||
for message in &mut messages {
|
||||
message.truncate_tool_results_for_provider_request();
|
||||
}
|
||||
|
||||
sanitize_messages_for_openai(&mut messages);
|
||||
|
||||
let system_prompt = request_translator::extract_system_prompt(request, ¶ms.global_rules);
|
||||
let tools = request_translator::extract_tools(request);
|
||||
PreparedTurn {
|
||||
task_id,
|
||||
needs_create_task,
|
||||
user_query: request_translator::extract_user_query_text(request),
|
||||
messages,
|
||||
system_prompt: request_translator::extract_system_prompt(request, ¶ms.global_rules),
|
||||
tools: request_translator::extract_tools(request),
|
||||
model_id,
|
||||
persistent_message_count,
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn execute(
|
||||
params: TranslatorRequest,
|
||||
request: &mut api::Request,
|
||||
) -> Result<ResponseStream, OpenAIError> {
|
||||
let client = OpenAIClient::from_config(params.config.clone());
|
||||
let PreparedTurn {
|
||||
task_id,
|
||||
needs_create_task,
|
||||
user_query,
|
||||
mut messages,
|
||||
system_prompt,
|
||||
tools,
|
||||
model_id,
|
||||
persistent_message_count,
|
||||
} = prepare_turn(¶ms, request);
|
||||
|
||||
log::info!(
|
||||
"[openai] Translator: model={model_id}, task_id={task_id}, needs_create_task={needs_create_task}"
|
||||
);
|
||||
|
||||
log::info!(
|
||||
"[openai] Sending {} messages, system_prompt={}, tools={}",
|
||||
@@ -118,8 +132,6 @@ pub async fn execute(
|
||||
tools.len()
|
||||
);
|
||||
|
||||
let user_query_text = request_translator::extract_user_query_text(request);
|
||||
|
||||
let max_output_tokens = params
|
||||
.config
|
||||
.max_output_tokens
|
||||
@@ -139,9 +151,8 @@ pub async fn execute(
|
||||
|
||||
// Store the message history for the controller
|
||||
if let Ok(mut sent) = params.messages_sent.lock() {
|
||||
let persistent_count = history_len + new_input_count;
|
||||
if persistent_count > 0 && messages.len() >= persistent_count {
|
||||
*sent = messages.split_off(messages.len() - persistent_count);
|
||||
if persistent_message_count > 0 && messages.len() >= persistent_message_count {
|
||||
*sent = messages.split_off(messages.len() - persistent_message_count);
|
||||
} else {
|
||||
*sent = messages;
|
||||
}
|
||||
@@ -152,7 +163,7 @@ pub async fn execute(
|
||||
OpenAIStreamContext {
|
||||
task_id,
|
||||
needs_create_task,
|
||||
user_query: user_query_text,
|
||||
user_query,
|
||||
messages_sent: params.messages_sent.clone(),
|
||||
model_id,
|
||||
max_context_tokens: params.config.max_input_tokens,
|
||||
|
||||
@@ -1,131 +1,6 @@
|
||||
use serde_json::Value as JsonValue;
|
||||
|
||||
pub const MAX_TOOL_RESULT_CHARS_FOR_PROVIDER_REQUEST: usize = 64_000;
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub struct ConversationMessage {
|
||||
pub role: MessageRole,
|
||||
pub content: MessageContent,
|
||||
}
|
||||
|
||||
impl ConversationMessage {
|
||||
pub fn truncate_tool_results_for_provider_request(&mut self) {
|
||||
truncate_tool_results_in_content(&mut self.content);
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq)]
|
||||
pub enum MessageRole {
|
||||
User,
|
||||
Assistant,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub enum MessageContent {
|
||||
Text(String),
|
||||
ToolUse {
|
||||
tool_use_id: String,
|
||||
name: String,
|
||||
input: JsonValue,
|
||||
},
|
||||
ToolResult {
|
||||
tool_use_id: String,
|
||||
content: String,
|
||||
is_error: bool,
|
||||
},
|
||||
MultiPart(Vec<ContentPart>),
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub enum ContentPart {
|
||||
Text(String),
|
||||
Image {
|
||||
data: Vec<u8>,
|
||||
mime_type: String,
|
||||
},
|
||||
ToolUse {
|
||||
tool_use_id: String,
|
||||
name: String,
|
||||
input: JsonValue,
|
||||
},
|
||||
ToolResult {
|
||||
tool_use_id: String,
|
||||
content: String,
|
||||
is_error: bool,
|
||||
},
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub struct ToolDefinition {
|
||||
pub name: String,
|
||||
pub description: String,
|
||||
pub input_schema: JsonValue,
|
||||
}
|
||||
|
||||
fn truncate_tool_results_in_content(content: &mut MessageContent) {
|
||||
match content {
|
||||
MessageContent::Text(_) | MessageContent::ToolUse { .. } => {}
|
||||
MessageContent::ToolResult { content, .. } => truncate_tool_result_text(content),
|
||||
MessageContent::MultiPart(parts) => {
|
||||
for part in parts {
|
||||
if let ContentPart::ToolResult { content, .. } = part {
|
||||
truncate_tool_result_text(content);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn truncate_tool_result_text(content: &mut String) {
|
||||
let char_count = content.chars().count();
|
||||
if char_count <= MAX_TOOL_RESULT_CHARS_FOR_PROVIDER_REQUEST {
|
||||
return;
|
||||
}
|
||||
|
||||
let omitted_chars = char_count.saturating_sub(MAX_TOOL_RESULT_CHARS_FOR_PROVIDER_REQUEST);
|
||||
let marker = format!("\n... [tool result truncated; omitted {omitted_chars} chars] ...\n");
|
||||
let marker_chars = marker.chars().count();
|
||||
let retained_chars = MAX_TOOL_RESULT_CHARS_FOR_PROVIDER_REQUEST.saturating_sub(marker_chars);
|
||||
let head_chars = retained_chars / 2;
|
||||
let tail_chars = retained_chars.saturating_sub(head_chars);
|
||||
let head: String = content.chars().take(head_chars).collect();
|
||||
let tail: String = content
|
||||
.chars()
|
||||
.rev()
|
||||
.take(tail_chars)
|
||||
.collect::<String>()
|
||||
.chars()
|
||||
.rev()
|
||||
.collect();
|
||||
*content = format!("{head}{marker}{tail}");
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn truncates_large_tool_results_for_provider_request() {
|
||||
let prefix = "start:";
|
||||
let suffix = ":end";
|
||||
let middle = "x".repeat(MAX_TOOL_RESULT_CHARS_FOR_PROVIDER_REQUEST + 1_000);
|
||||
let mut message = ConversationMessage {
|
||||
role: MessageRole::User,
|
||||
content: MessageContent::ToolResult {
|
||||
tool_use_id: "toolu_1".to_string(),
|
||||
content: format!("{prefix}{middle}{suffix}"),
|
||||
is_error: false,
|
||||
},
|
||||
};
|
||||
|
||||
message.truncate_tool_results_for_provider_request();
|
||||
|
||||
let MessageContent::ToolResult { content, .. } = message.content else {
|
||||
panic!("expected tool result");
|
||||
};
|
||||
assert!(content.len() <= MAX_TOOL_RESULT_CHARS_FOR_PROVIDER_REQUEST + 128);
|
||||
assert!(content.starts_with(prefix));
|
||||
assert!(content.ends_with(suffix));
|
||||
assert!(content.contains("tool result truncated"));
|
||||
}
|
||||
}
|
||||
// Keep this module as a compatibility import path while provider-neutral message
|
||||
// types move out of the application crate.
|
||||
pub use galaxy_agent_core::{
|
||||
ContentPart, ConversationMessage, MessageContent, MessageRole, ToolDefinition,
|
||||
MAX_TOOL_RESULT_CHARS_FOR_PROVIDER_REQUEST,
|
||||
};
|
||||
|
||||
@@ -0,0 +1,5 @@
|
||||
mod provider;
|
||||
mod rig;
|
||||
|
||||
pub(crate) use provider::ProviderRuntime;
|
||||
pub(crate) use rig::rig_openai_response_stream;
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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) = ¤t_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) = ¤t_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;
|
||||
@@ -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"]
|
||||
);
|
||||
}
|
||||
+1
-1
@@ -14,7 +14,7 @@ fn main() -> Result<()> {
|
||||
ChannelConfig {
|
||||
app_id: AppId::new("com", "samsung", "Galaxy"),
|
||||
logfile_name: "galaxy.log".into(),
|
||||
server_config: WarpServerConfig::production(),
|
||||
server_config: WarpServerConfig::disabled(),
|
||||
oz_config: OzConfig::production(),
|
||||
telemetry_config: None,
|
||||
autoupdate_config: None,
|
||||
|
||||
+46
-3
@@ -874,10 +874,27 @@ pub struct OpenAIModelConfig {
|
||||
description = "Optional provider hint (e.g. anthropic, openai, google) for icon display."
|
||||
)]
|
||||
pub provider: Option<String>,
|
||||
#[serde(default)]
|
||||
#[schemars(
|
||||
description = "Route this model through Galaxy's Rig runtime. This is an opt-in migration path."
|
||||
)]
|
||||
pub use_rig: bool,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
#[schemars(
|
||||
description = "Whether this endpoint accepts system-role messages. Set false for ChatGPT-backed LiteLLM models that reject them."
|
||||
)]
|
||||
pub supports_system_messages: Option<bool>,
|
||||
}
|
||||
|
||||
impl settings_value::SettingsValue for OpenAIModelConfig {}
|
||||
|
||||
impl OpenAIModelConfig {
|
||||
pub fn supports_system_messages(&self) -> bool {
|
||||
self.supports_system_messages
|
||||
.unwrap_or_else(|| !self.model_id.starts_with("codex-gpt-"))
|
||||
}
|
||||
}
|
||||
|
||||
/// Configuration for a single OpenAI-compatible provider endpoint.
|
||||
///
|
||||
/// Multiple providers can be configured simultaneously (e.g. LiteLLM for cloud models,
|
||||
@@ -901,6 +918,30 @@ pub struct OpenAIProviderConfig {
|
||||
|
||||
impl settings_value::SettingsValue for OpenAIProviderConfig {}
|
||||
|
||||
const INITIAL_LITELLM_BASE_URL: &str = "https://ai.ryserve.net/v1";
|
||||
const INITIAL_RIG_MODEL_ID: &str = "codex-gpt-5.6-sol-xhigh";
|
||||
|
||||
fn default_openai_providers() -> Vec<OpenAIProviderConfig> {
|
||||
vec![OpenAIProviderConfig {
|
||||
name: "LiteLLM (ai.ryserve.net)".to_string(),
|
||||
base_url: INITIAL_LITELLM_BASE_URL.to_string(),
|
||||
// Credentials are deliberately never committed. Set this locally in
|
||||
// ~/.galaxy/settings.toml before sending a request.
|
||||
api_key: None,
|
||||
models: vec![OpenAIModelConfig {
|
||||
model_id: INITIAL_RIG_MODEL_ID.to_string(),
|
||||
display_name: "Codex GPT-5.6 SOL (xhigh)".to_string(),
|
||||
vision_supported: false,
|
||||
context_size: default_context_size(),
|
||||
max_input_tokens: None,
|
||||
max_output_tokens: None,
|
||||
provider: Some("openai".to_string()),
|
||||
use_rig: true,
|
||||
supports_system_messages: Some(false),
|
||||
}],
|
||||
}]
|
||||
}
|
||||
|
||||
/// Cached metadata and runtime session options for an ACP agent.
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, schemars::JsonSchema)]
|
||||
pub struct AcpAgentSettings {
|
||||
@@ -1447,7 +1488,7 @@ define_settings_group!(AISettings, settings: [
|
||||
// Whether the OpenAI-compatible (LiteLLM) provider is enabled.
|
||||
openai_enabled: OpenAIEnabled {
|
||||
type: bool,
|
||||
default: false,
|
||||
default: true,
|
||||
supported_platforms: SupportedPlatforms::DESKTOP,
|
||||
sync_to_cloud: SyncToCloud::Globally(RespectUserSyncSetting::Yes),
|
||||
private: false,
|
||||
@@ -1498,9 +1539,11 @@ define_settings_group!(AISettings, settings: [
|
||||
// Each provider has its own name, base_url, api_key, and model list.
|
||||
openai_providers: OpenAIProviders {
|
||||
type: Vec<OpenAIProviderConfig>,
|
||||
default: Vec::new(),
|
||||
default: default_openai_providers(),
|
||||
supported_platforms: SupportedPlatforms::DESKTOP,
|
||||
sync_to_cloud: SyncToCloud::Globally(RespectUserSyncSetting::Yes),
|
||||
// Provider entries may contain API keys, so the complete setting must
|
||||
// remain local even when preference sync is enabled.
|
||||
sync_to_cloud: SyncToCloud::Never,
|
||||
private: false,
|
||||
toml_path: "ai.providers",
|
||||
description: "Multiple OpenAI-compatible provider endpoints (e.g. LiteLLM, Ollama, local models).",
|
||||
|
||||
@@ -345,6 +345,37 @@ fn test_toolbar_command_map_roundtrip() {
|
||||
assert_eq!(original, restored);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn initial_litellm_provider_maps_codex_model_to_rig_without_a_committed_key() {
|
||||
let providers = default_openai_providers();
|
||||
|
||||
assert_eq!(providers.len(), 1);
|
||||
let provider = &providers[0];
|
||||
assert_eq!(provider.base_url, INITIAL_LITELLM_BASE_URL);
|
||||
assert_eq!(provider.api_key, None);
|
||||
assert_eq!(provider.models.len(), 1);
|
||||
let model = &provider.models[0];
|
||||
assert_eq!(model.model_id, INITIAL_RIG_MODEL_ID);
|
||||
assert_eq!(model.use_rig, true);
|
||||
assert_eq!(model.supports_system_messages, Some(false));
|
||||
assert_eq!(model.supports_system_messages(), false);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn codex_litellm_model_infers_missing_system_message_capability() {
|
||||
let mut model = default_openai_providers().remove(0).models.remove(0);
|
||||
model.supports_system_messages = None;
|
||||
|
||||
assert_eq!(model.supports_system_messages(), false);
|
||||
|
||||
model.model_id = "gpt-4o".to_string();
|
||||
assert_eq!(model.supports_system_messages(), true);
|
||||
|
||||
model.model_id = INITIAL_RIG_MODEL_ID.to_string();
|
||||
model.supports_system_messages = Some(true);
|
||||
assert_eq!(model.supports_system_messages(), true);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_toolbar_command_map_matched_agent() {
|
||||
App::test((), |mut app| async move {
|
||||
|
||||
Reference in New Issue
Block a user