Add OpenAI/LiteLLM provider support with settings UI

- Add openai/ provider module with translator, client, convert, request/response translators
- Add shared provider/ types (ConversationMessage, MessageRole, ProviderConfig enum)
- Wire OpenAI-compatible provider dispatch alongside Bedrock in response_stream.rs
- Add ai.openai.* settings (enabled, base_url, api_key, model, models)
- Add OpenAI/LiteLLM settings page with model fetch, picker, and config UI
- Extend model menu items and llms.rs to surface LiteLLM models
- Update WARP.md with OpenAI provider architecture docs
This commit is contained in:
Ryan Ward
2026-06-17 14:14:40 -05:00
parent 59cfd0e2f5
commit 5ea378a38d
32 changed files with 2442 additions and 137 deletions
+64 -32
View File
@@ -5,13 +5,14 @@ use futures_util::StreamExt;
use galaxy_core::features::FeatureFlag;
use warp_multi_agent_api as api;
use crate::ai::bedrock::client::BedrockClientConfig;
use crate::ai::bedrock::translator::{self, TranslatorRequest};
use crate::ai::openai::translator as openai_translator;
use crate::ai::provider::ProviderConfig;
use super::{convert_to::convert_input, ConvertToAPITypeError, RequestParams, ResponseStream};
pub async fn generate_multi_agent_output(
bedrock_config: Option<BedrockClientConfig>,
provider_config: ProviderConfig,
mut params: RequestParams,
cancellation_rx: futures::channel::oneshot::Receiver<()>,
) -> Result<ResponseStream, ConvertToAPITypeError> {
@@ -128,19 +129,6 @@ pub async fn generate_multi_agent_output(
mcp_context: params.mcp_context.map(Into::into),
};
let Some(config) = bedrock_config else {
log::error!("[bedrock] No Bedrock config available. Cannot process request.");
let err = Arc::new(crate::server::server_api::AIApiError::Stream {
stream_type: "bedrock_converse",
source: anyhow::anyhow!(
"No AI backend available. Please configure Bedrock credentials in Settings > AI."
),
});
let (tx, rx) = async_channel::unbounded();
let _ = tx.send(Err(err)).await;
return Ok(Box::pin(rx));
};
let model_id = request
.settings
.as_ref()
@@ -148,26 +136,70 @@ pub async fn generate_multi_agent_output(
.map(|mc| mc.base.clone())
.unwrap_or_default();
let translator_request = TranslatorRequest {
config,
model_id,
root_task_id: params.root_task_id.clone(),
bedrock_message_history: params.bedrock_message_history.clone(),
bedrock_tool_result_archive: params.bedrock_tool_result_archive.clone(),
bedrock_progressive_summary: params.bedrock_progressive_summary.clone(),
bedrock_messages_sent: params.bedrock_messages_sent.clone(),
};
match provider_config {
ProviderConfig::Bedrock(config) => {
let translator_request = TranslatorRequest {
config,
model_id,
root_task_id: params.root_task_id.clone(),
bedrock_message_history: params.bedrock_message_history.clone(),
bedrock_tool_result_archive: params.bedrock_tool_result_archive.clone(),
bedrock_progressive_summary: params.bedrock_progressive_summary.clone(),
bedrock_messages_sent: params.bedrock_messages_sent.clone(),
};
match translator::execute(translator_request, &mut request).await {
Ok(stream) => {
let output_stream = stream.take_until(cancellation_rx);
Ok(Box::pin(output_stream))
match translator::execute(translator_request, &mut request).await {
Ok(stream) => {
let output_stream = stream.take_until(cancellation_rx);
Ok(Box::pin(output_stream))
}
Err(e) => {
log::error!("[bedrock] Translator error: {e}");
let err = Arc::new(crate::server::server_api::AIApiError::Stream {
stream_type: "bedrock_converse",
source: anyhow::anyhow!("{e}"),
});
let (tx, rx) = async_channel::unbounded();
let _ = tx.send(Err(err)).await;
Ok(Box::pin(rx))
}
}
}
Err(e) => {
log::error!("[bedrock] Translator error: {e}");
ProviderConfig::OpenAI(config) => {
let translator_request = openai_translator::TranslatorRequest {
config,
model_id,
root_task_id: params.root_task_id.clone(),
message_history: params.bedrock_message_history.clone(),
tool_result_archive: params.bedrock_tool_result_archive.clone(),
progressive_summary: params.bedrock_progressive_summary.clone(),
messages_sent: params.bedrock_messages_sent.clone(),
};
match openai_translator::execute(translator_request, &mut request).await {
Ok(stream) => {
let output_stream = stream.take_until(cancellation_rx);
Ok(Box::pin(output_stream))
}
Err(e) => {
log::error!("[openai] Translator error: {e}");
let err = Arc::new(crate::server::server_api::AIApiError::Stream {
stream_type: "openai_chat_completions",
source: anyhow::anyhow!("{e}"),
});
let (tx, rx) = async_channel::unbounded();
let _ = tx.send(Err(err)).await;
Ok(Box::pin(rx))
}
}
}
ProviderConfig::None => {
log::error!("No AI provider configured. Cannot process request.");
let err = Arc::new(crate::server::server_api::AIApiError::Stream {
stream_type: "bedrock_converse",
source: anyhow::anyhow!("{e}"),
stream_type: "provider_dispatch",
source: anyhow::anyhow!(
"No AI backend available. Please configure a provider in Settings > AI."
),
});
let (tx, rx) = async_channel::unbounded();
let _ = tx.send(Err(err)).await;
+5 -50
View File
@@ -11,6 +11,11 @@ use serde_json::Value as JsonValue;
use super::external_config::ExternalBedrockConfig;
// Re-export shared provider types so existing imports from bedrock::convert continue to work.
pub use crate::ai::provider::types::{
ContentPart, ConversationMessage, MessageContent, MessageRole, ToolDefinition,
};
#[derive(Clone, Debug)]
pub struct CachingConfig {
pub enabled: bool,
@@ -42,56 +47,6 @@ pub struct ConvertedRequest {
pub tool_config: Option<ToolConfiguration>,
}
#[derive(Clone, Debug)]
pub struct ConversationMessage {
pub role: MessageRole,
pub content: MessageContent,
}
#[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),
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,
}
pub fn build_converse_request(
messages: Vec<ConversationMessage>,
system_prompt: Option<String>,
+1
View File
@@ -92,6 +92,7 @@ fn test_get_effective_models_custom_overrides() {
model_id: "custom.model-v1:0".to_string(),
display_name: "Custom Model".to_string(),
vision_supported: false,
context_size: 200_000,
}];
let models = get_effective_models(&custom);
assert_eq!(models.len(), 1);
+6 -1
View File
@@ -992,7 +992,12 @@ pub fn extract_tools(request: &api::Request) -> Vec<ToolDefinition> {
// Filter out suggest_next_prompt — its action executor waits on a oneshot
// channel for UI interaction that never fires in the Bedrock path, causing
// the conversation to stay InProgress forever.
tools.retain(|t| t.name != "suggest_next_prompt");
// Filter out start_agent/send_message_to_agent — sub-agents are disabled.
tools.retain(|t| {
t.name != "suggest_next_prompt"
&& t.name != "start_agent"
&& t.name != "send_message_to_agent"
});
tools
}
+4 -4
View File
@@ -560,7 +560,7 @@ pub fn bedrock_stream_to_response_events(
Box::pin(stream)
}
pub(crate) fn build_create_task(task_id: &str) -> ResponseEvent {
pub fn build_create_task(task_id: &str) -> ResponseEvent {
let task = api::Task {
id: task_id.to_string(),
description: String::new(),
@@ -617,7 +617,7 @@ fn build_user_query_message(task_id: &str, query_text: &str) -> ResponseEvent {
}
}
pub(super) fn build_stream_init(request_id: &str, conversation_id: &str) -> ResponseEvent {
pub fn build_stream_init(request_id: &str, conversation_id: &str) -> ResponseEvent {
ResponseEvent {
r#type: Some(api::response_event::Type::Init(
api::response_event::StreamInit {
@@ -629,7 +629,7 @@ pub(super) fn build_stream_init(request_id: &str, conversation_id: &str) -> Resp
}
}
pub(super) fn build_stream_finished(
pub fn build_stream_finished(
reason: stream_finished::Reason,
input_tokens: i32,
output_tokens: i32,
@@ -820,7 +820,7 @@ fn build_append_text(task_id: &str, message_id: &str, text_delta: &str) -> Respo
}
}
fn build_tool_call_message(
pub fn build_tool_call_message(
task_id: &str,
tool_use_id: &str,
tool_name: &str,
@@ -19,7 +19,15 @@ fn test_build_stream_init_has_valid_ids() {
#[test]
fn test_build_stream_finished_done_reason() {
let reason = stream_finished::Reason::Done(stream_finished::Done {});
let event = build_stream_finished(reason, 100, 50, 20, 10, "anthropic.claude-sonnet-4-6", false);
let event = build_stream_finished(
reason,
100,
50,
20,
10,
"anthropic.claude-sonnet-4-6",
false,
);
match event.r#type {
Some(api::response_event::Type::Finished(finished)) => {
@@ -15,6 +15,8 @@ use crate::{
AIIdentifiers, CancellationReason,
},
bedrock::client::BedrockClientConfig,
openai::client::OpenAIClientConfig,
provider::ProviderConfig,
},
network::NetworkStatus,
report_error, send_telemetry_from_ctx,
@@ -83,26 +85,54 @@ pub struct ResponseStream {
}
impl ResponseStream {
fn bedrock_config_if_applicable(
_model_id: &str,
ctx: &ModelContext<Self>,
) -> Option<BedrockClientConfig> {
fn resolve_provider_config(model_id: &str, ctx: &ModelContext<Self>) -> ProviderConfig {
let settings = AISettings::as_ref(ctx);
if !*settings.bedrock_enabled.value() {
return None;
// Check if OpenAI/LiteLLM provider is enabled
if *settings.openai_enabled.value() {
let base_url = settings.openai_base_url.value().clone();
let api_key = {
let key = settings.openai_api_key.value().clone();
if key.is_empty() {
None
} else {
Some(key)
}
};
// Use the model override from settings if set, otherwise use the selected model ID.
// This allows LiteLLM models to pass through their actual model_id to the proxy.
let model = {
let m = settings.openai_model.value().clone();
if m.is_empty() {
Some(model_id.to_string())
} else {
Some(m)
}
};
return ProviderConfig::OpenAI(OpenAIClientConfig {
base_url,
api_key,
model,
});
}
let auth_method = *settings.bedrock_auth_method.value();
Some(
BedrockClientConfig {
auth_method,
profile: settings.bedrock_profile.value().clone(),
region: settings.bedrock_region.value().clone(),
access_key_id: settings.bedrock_access_key_id.value().clone(),
secret_access_key: settings.bedrock_secret_access_key.value().clone(),
cross_region_inference: *settings.bedrock_cross_region_inference.value(),
}
.with_external_fallbacks(),
)
// Fall back to Bedrock
if *settings.bedrock_enabled.value() {
let auth_method = *settings.bedrock_auth_method.value();
return ProviderConfig::Bedrock(
BedrockClientConfig {
auth_method,
profile: settings.bedrock_profile.value().clone(),
region: settings.bedrock_region.value().clone(),
access_key_id: settings.bedrock_access_key_id.value().clone(),
secret_access_key: settings.bedrock_secret_access_key.value().clone(),
cross_region_inference: *settings.bedrock_cross_region_inference.value(),
}
.with_external_fallbacks(),
);
}
ProviderConfig::None
}
pub fn new(
@@ -115,11 +145,11 @@ impl ResponseStream {
let start_time = Local::now();
let request_id = Uuid::new_v4();
let bedrock_config = Self::bedrock_config_if_applicable(params.model.as_str(), ctx);
let provider_config = Self::resolve_provider_config(params.model.as_str(), ctx);
let params_clone = params.clone();
let _ = ctx.spawn(
async move {
generate_multi_agent_output(bedrock_config, params_clone, cancellation_rx).await
generate_multi_agent_output(provider_config, params_clone, cancellation_rx).await
},
move |me, stream, ctx| {
me.handle_response_stream_result(request_id, stream, ctx);
@@ -192,11 +222,11 @@ impl ResponseStream {
let request_id = Uuid::new_v4();
self.current_request_id = Some(request_id);
let params = self.params.clone();
let bedrock_config = Self::bedrock_config_if_applicable(params.model.as_str(), ctx);
let provider_config = Self::resolve_provider_config(params.model.as_str(), ctx);
let _ =
ctx.spawn(
async move {
generate_multi_agent_output(bedrock_config, params, cancellation_rx).await
generate_multi_agent_output(provider_config, params, cancellation_rx).await
},
move |me, stream, ctx| {
me.handle_response_stream_result(request_id, stream, ctx);
@@ -432,7 +432,6 @@ impl PassiveSuggestionsModel {
);
}
}
}
impl Entity for PassiveSuggestionsModel {
@@ -81,6 +81,7 @@ fn make_item_fields<A: Action + Clone>(
};
let is_using_api_key = is_using_api_key_for_provider(&llm.provider, app);
let is_bedrock = llm.provider == LLMProvider::Bedrock;
let is_litellm = llm.provider == LLMProvider::LiteLLM;
let mut item = if let Some(position_id_fn) = position_id_fn {
let position_id = position_id_fn(&llm.id);
@@ -94,6 +95,10 @@ fn make_item_fields<A: Action + Clone>(
Icon::BedrockLogo
.to_galaxyui_icon(appearance.theme().foreground())
.finish()
} else if is_litellm {
Icon::OpenAILogo
.to_galaxyui_icon(appearance.theme().foreground())
.finish()
} else if is_using_api_key {
Icon::Key
.to_galaxyui_icon(appearance.theme().foreground())
+90 -3
View File
@@ -16,7 +16,7 @@ use crate::{
network::{NetworkStatus, NetworkStatusEvent, NetworkStatusKind},
report_error,
server::server_api::ServerApiProvider,
settings::ai::{AISettings, AISettingsChangedEvent, BedrockModelConfig},
settings::ai::{AISettings, AISettingsChangedEvent, BedrockModelConfig, OpenAIModelConfig},
workspaces::user_workspaces::{UserWorkspaces, UserWorkspacesEvent},
};
@@ -43,6 +43,7 @@ pub fn is_using_api_key_for_provider(provider: &LLMProvider, app: &AppContext) -
LLMProvider::Anthropic => api_keys.is_some_and(|keys| keys.anthropic.is_some()),
LLMProvider::Google => api_keys.is_some_and(|keys| keys.google.is_some()),
LLMProvider::Bedrock => true,
LLMProvider::LiteLLM => true,
_ => false,
}
}
@@ -97,6 +98,8 @@ pub enum LLMProvider {
Google,
Xai,
Bedrock,
/// Models served through an OpenAI-compatible proxy (e.g. LiteLLM).
LiteLLM,
Unknown,
}
@@ -108,6 +111,7 @@ impl LLMProvider {
LLMProvider::Anthropic => Some(Icon::ClaudeLogo),
LLMProvider::Google => Some(Icon::GeminiLogo),
LLMProvider::Bedrock => Some(Icon::BedrockLogo),
LLMProvider::LiteLLM => Some(Icon::OpenAILogo),
LLMProvider::Xai => None,
LLMProvider::Unknown => None,
}
@@ -551,6 +555,15 @@ impl LLMPreferences {
me.inject_bedrock_models(ctx);
ctx.emit(LLMPreferencesEvent::UpdatedAvailableLLMs);
}
if matches!(
event,
AISettingsChangedEvent::OpenAIEnabled { .. }
| AISettingsChangedEvent::OpenAIModels { .. }
| AISettingsChangedEvent::OpenAIBaseUrl { .. }
) {
me.inject_openai_models(ctx);
ctx.emit(LLMPreferencesEvent::UpdatedAvailableLLMs);
}
});
let base_llm_for_terminal_view = HashMap::new();
@@ -572,6 +585,7 @@ impl LLMPreferences {
{
Self::ensure_default_models_in_settings(ctx);
me.inject_bedrock_models(ctx);
me.inject_openai_models(ctx);
}
me
@@ -582,8 +596,7 @@ impl LLMPreferences {
use crate::ai::bedrock::models::DEFAULT_BEDROCK_MODELS;
let settings = AISettings::as_ref(ctx);
let mut current_models: Vec<BedrockModelConfig> =
settings.bedrock_models.value().clone();
let mut current_models: Vec<BedrockModelConfig> = settings.bedrock_models.value().clone();
let existing_ids: std::collections::HashSet<String> =
current_models.iter().map(|m| m.model_id.clone()).collect();
@@ -786,6 +799,80 @@ impl LLMPreferences {
}
}
/// Injects models from the OpenAI-compatible (LiteLLM) provider into the available model lists.
#[cfg(not(target_family = "wasm"))]
fn inject_openai_models(&mut self, ctx: &AppContext) {
// Remove any previously injected LiteLLM models
self.models_by_feature
.agent_mode
.choices
.retain(|m| m.provider != LLMProvider::LiteLLM);
self.models_by_feature
.coding
.choices
.retain(|m| m.provider != LLMProvider::LiteLLM);
if let Some(ref mut cli) = self.models_by_feature.cli_agent {
cli.choices.retain(|m| m.provider != LLMProvider::LiteLLM);
}
let settings = AISettings::as_ref(ctx);
if !*settings.openai_enabled.value() {
return;
}
let user_models: Vec<OpenAIModelConfig> = settings.openai_models.value().clone();
if user_models.is_empty() {
return;
}
let base_url = settings.openai_base_url.value().clone();
let description_label = if base_url.contains("localhost") || base_url.contains("127.0.0.1")
{
"LiteLLM (local)".to_string()
} else {
"LiteLLM".to_string()
};
for model in &user_models {
let llm_info = LLMInfo {
id: LLMId::from(model.model_id.as_str()),
display_name: model.display_name.clone(),
base_model_name: model.display_name.clone(),
reasoning_level: None,
usage_metadata: LLMUsageMetadata {
request_multiplier: 1,
credit_multiplier: None,
},
description: Some(description_label.clone()),
disable_reason: None,
vision_supported: model.vision_supported,
spec: None,
provider: LLMProvider::LiteLLM,
host_configs: HashMap::from([(
LLMModelHost::DirectApi,
RoutingHostConfig {
enabled: true,
model_routing_host: LLMModelHost::DirectApi,
},
)]),
discount_percentage: None,
};
self.models_by_feature
.agent_mode
.choices
.push(llm_info.clone());
self.models_by_feature.coding.choices.push(llm_info.clone());
if let Some(ref mut cli) = self.models_by_feature.cli_agent {
cli.choices.push(llm_info);
}
}
log::info!(
"[openai/litellm] Injected {} model(s) into available choices",
user_models.len()
);
}
/// Returns the `LLMInfo` for the base LLM to be used for an Agent Mode request.
pub fn get_active_base_model<'a>(
&'a self,
+4
View File
@@ -29,10 +29,14 @@ pub(crate) mod get_relevant_files;
pub(crate) mod harness_display;
pub(crate) mod llms;
pub mod onboarding;
#[cfg(not(target_family = "wasm"))]
pub mod openai;
pub(crate) mod persisted_workspace;
pub(crate) mod predict;
#[allow(dead_code)]
pub mod prompt_builder;
#[cfg(not(target_family = "wasm"))]
pub mod provider;
pub mod request_usage_model;
pub(crate) mod restored_conversations;
pub(crate) mod skills;
+92
View File
@@ -0,0 +1,92 @@
use std::fmt;
use bytes::Bytes;
use futures::Stream;
use reqwest::header::{HeaderMap, HeaderValue, AUTHORIZATION, CONTENT_TYPE};
#[derive(Clone, Debug)]
pub struct OpenAIClientConfig {
pub base_url: String,
pub api_key: Option<String>,
pub model: Option<String>,
}
pub struct OpenAIClient {
http: reqwest::Client,
base_url: String,
api_key: Option<String>,
}
#[derive(Debug)]
pub enum OpenAIError {
ConnectionFailed(String),
AuthenticationFailed(String),
RateLimited(String),
BadRequest(String),
ServerError(String),
#[allow(dead_code)]
StreamError(String),
}
impl fmt::Display for OpenAIError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::ConnectionFailed(msg) => write!(f, "Connection failed: {msg}"),
Self::AuthenticationFailed(msg) => write!(f, "Authentication failed: {msg}"),
Self::RateLimited(msg) => write!(f, "Rate limited: {msg}"),
Self::BadRequest(msg) => write!(f, "Bad request: {msg}"),
Self::ServerError(msg) => write!(f, "Server error: {msg}"),
Self::StreamError(msg) => write!(f, "Stream error: {msg}"),
}
}
}
impl OpenAIClient {
pub fn from_config(config: OpenAIClientConfig) -> Self {
let http = reqwest::Client::new();
Self {
http,
base_url: config.base_url,
api_key: config.api_key,
}
}
pub async fn chat_completions_stream(
&self,
request_body: serde_json::Value,
) -> Result<impl Stream<Item = Result<Bytes, reqwest::Error>>, OpenAIError> {
let url = format!("{}/chat/completions", self.base_url.trim_end_matches('/'));
let mut headers = HeaderMap::new();
headers.insert(CONTENT_TYPE, HeaderValue::from_static("application/json"));
if let Some(ref key) = self.api_key {
headers.insert(
AUTHORIZATION,
HeaderValue::from_str(&format!("Bearer {key}"))
.map_err(|e| OpenAIError::BadRequest(format!("Invalid API key header: {e}")))?,
);
}
let response = self
.http
.post(&url)
.headers(headers)
.json(&request_body)
.send()
.await
.map_err(|e| OpenAIError::ConnectionFailed(e.to_string()))?;
let status = response.status();
if !status.is_success() {
let body = response.text().await.unwrap_or_default();
return Err(match status.as_u16() {
401 => OpenAIError::AuthenticationFailed(body),
429 => OpenAIError::RateLimited(body),
400 => OpenAIError::BadRequest(body),
_ => OpenAIError::ServerError(format!("HTTP {status}: {body}")),
});
}
Ok(response.bytes_stream())
}
}
+229
View File
@@ -0,0 +1,229 @@
use serde_json::{json, Value as JsonValue};
use crate::ai::provider::types::{
ContentPart, ConversationMessage, MessageContent, MessageRole, ToolDefinition,
};
pub fn build_openai_request(
messages: Vec<ConversationMessage>,
system_prompt: Option<String>,
tools: Vec<ToolDefinition>,
max_tokens: i32,
temperature: Option<f32>,
model: &str,
) -> JsonValue {
let mut openai_messages: Vec<JsonValue> = Vec::new();
if let Some(prompt) = system_prompt {
if !prompt.is_empty() {
openai_messages.push(json!({
"role": "system",
"content": prompt,
}));
}
}
for msg in messages {
match convert_message(msg) {
ConvertedMessages::Single(m) => openai_messages.push(m),
ConvertedMessages::Multiple(ms) => openai_messages.extend(ms),
}
}
let mut request = json!({
"model": model,
"messages": openai_messages,
"max_tokens": max_tokens,
"stream": true,
"stream_options": { "include_usage": true },
});
if let Some(temp) = temperature {
request["temperature"] = json!(temp);
}
if !tools.is_empty() {
let tool_defs: Vec<JsonValue> = tools.into_iter().map(convert_tool_definition).collect();
request["tools"] = json!(tool_defs);
}
request
}
enum ConvertedMessages {
Single(JsonValue),
Multiple(Vec<JsonValue>),
}
fn convert_message(msg: ConversationMessage) -> ConvertedMessages {
match msg.role {
MessageRole::User => convert_user_message(msg.content),
MessageRole::Assistant => convert_assistant_message(msg.content),
}
}
fn convert_user_message(content: MessageContent) -> ConvertedMessages {
match content {
MessageContent::Text(text) => ConvertedMessages::Single(json!({
"role": "user",
"content": text,
})),
MessageContent::ToolResult {
tool_use_id,
content,
is_error,
} => {
let mut msg = json!({
"role": "tool",
"tool_call_id": tool_use_id,
"content": content,
});
if is_error {
msg["content"] = json!(format!("[ERROR] {content}"));
}
ConvertedMessages::Single(msg)
}
MessageContent::ToolUse { .. } => {
// User messages shouldn't contain tool_use, but handle gracefully
ConvertedMessages::Single(json!({
"role": "user",
"content": "[unexpected tool_use in user message]",
}))
}
MessageContent::MultiPart(parts) => {
let mut messages = Vec::new();
let mut text_parts: Vec<String> = Vec::new();
for part in parts {
match part {
ContentPart::Text(text) => text_parts.push(text),
ContentPart::ToolResult {
tool_use_id,
content,
is_error,
} => {
// Flush any accumulated text as a user message first
if !text_parts.is_empty() {
messages.push(json!({
"role": "user",
"content": text_parts.join("\n"),
}));
text_parts.clear();
}
let result_content = if is_error {
format!("[ERROR] {content}")
} else {
content
};
messages.push(json!({
"role": "tool",
"tool_call_id": tool_use_id,
"content": result_content,
}));
}
ContentPart::ToolUse { .. } => {
text_parts.push("[unexpected tool_use in user message]".to_string());
}
}
}
if !text_parts.is_empty() {
messages.push(json!({
"role": "user",
"content": text_parts.join("\n"),
}));
}
if messages.len() == 1 {
ConvertedMessages::Single(messages.into_iter().next().unwrap())
} else {
ConvertedMessages::Multiple(messages)
}
}
}
}
fn convert_assistant_message(content: MessageContent) -> ConvertedMessages {
match content {
MessageContent::Text(text) => ConvertedMessages::Single(json!({
"role": "assistant",
"content": text,
})),
MessageContent::ToolUse {
tool_use_id,
name,
input,
} => ConvertedMessages::Single(json!({
"role": "assistant",
"content": null,
"tool_calls": [{
"id": tool_use_id,
"type": "function",
"function": {
"name": name,
"arguments": input.to_string(),
}
}]
})),
MessageContent::ToolResult { .. } => {
// Assistant messages shouldn't contain tool_result
ConvertedMessages::Single(json!({
"role": "assistant",
"content": "[unexpected tool_result in assistant message]",
}))
}
MessageContent::MultiPart(parts) => {
let mut text_content = String::new();
let mut tool_calls: Vec<JsonValue> = Vec::new();
for part in parts {
match part {
ContentPart::Text(text) => {
if !text_content.is_empty() {
text_content.push('\n');
}
text_content.push_str(&text);
}
ContentPart::ToolUse {
tool_use_id,
name,
input,
} => {
tool_calls.push(json!({
"id": tool_use_id,
"type": "function",
"function": {
"name": name,
"arguments": input.to_string(),
}
}));
}
ContentPart::ToolResult { .. } => {}
}
}
let mut msg = json!({ "role": "assistant" });
if !text_content.is_empty() {
msg["content"] = json!(text_content);
} else {
msg["content"] = JsonValue::Null;
}
if !tool_calls.is_empty() {
msg["tool_calls"] = json!(tool_calls);
}
ConvertedMessages::Single(msg)
}
}
}
fn convert_tool_definition(tool: ToolDefinition) -> JsonValue {
json!({
"type": "function",
"function": {
"name": tool.name,
"description": tool.description,
"parameters": tool.input_schema,
}
})
}
+272
View File
@@ -0,0 +1,272 @@
use serde_json::json;
use crate::ai::openai::convert::build_openai_request;
use crate::ai::provider::types::{
ContentPart, ConversationMessage, MessageContent, MessageRole, ToolDefinition,
};
#[test]
fn test_simple_text_message_conversion() {
let messages = vec![ConversationMessage {
role: MessageRole::User,
content: MessageContent::Text("Hello world".to_string()),
}];
let request = build_openai_request(messages, None, vec![], 1024, None, "test-model");
let msgs = request["messages"].as_array().unwrap();
assert_eq!(msgs.len(), 1);
assert_eq!(msgs[0]["role"], "user");
assert_eq!(msgs[0]["content"], "Hello world");
assert_eq!(request["model"], "test-model");
assert_eq!(request["max_tokens"], 1024);
assert_eq!(request["stream"], true);
}
#[test]
fn test_system_prompt_placement() {
let messages = vec![ConversationMessage {
role: MessageRole::User,
content: MessageContent::Text("Hi".to_string()),
}];
let request = build_openai_request(
messages,
Some("You are a helpful assistant.".to_string()),
vec![],
1024,
None,
"test-model",
);
let msgs = request["messages"].as_array().unwrap();
assert_eq!(msgs.len(), 2);
assert_eq!(msgs[0]["role"], "system");
assert_eq!(msgs[0]["content"], "You are a helpful assistant.");
assert_eq!(msgs[1]["role"], "user");
}
#[test]
fn test_assistant_tool_use_conversion() {
let messages = vec![ConversationMessage {
role: MessageRole::Assistant,
content: MessageContent::ToolUse {
tool_use_id: "call_123".to_string(),
name: "run_shell_command".to_string(),
input: json!({"command": "ls -la"}),
},
}];
let request = build_openai_request(messages, None, vec![], 1024, None, "test-model");
let msgs = request["messages"].as_array().unwrap();
assert_eq!(msgs.len(), 1);
assert_eq!(msgs[0]["role"], "assistant");
assert!(msgs[0]["content"].is_null());
let tool_calls = msgs[0]["tool_calls"].as_array().unwrap();
assert_eq!(tool_calls.len(), 1);
assert_eq!(tool_calls[0]["id"], "call_123");
assert_eq!(tool_calls[0]["type"], "function");
assert_eq!(tool_calls[0]["function"]["name"], "run_shell_command");
assert_eq!(
tool_calls[0]["function"]["arguments"],
json!({"command": "ls -la"}).to_string()
);
}
#[test]
fn test_tool_result_conversion() {
let messages = vec![ConversationMessage {
role: MessageRole::User,
content: MessageContent::ToolResult {
tool_use_id: "call_123".to_string(),
content: "file1.txt\nfile2.txt".to_string(),
is_error: false,
},
}];
let request = build_openai_request(messages, None, vec![], 1024, None, "test-model");
let msgs = request["messages"].as_array().unwrap();
assert_eq!(msgs.len(), 1);
assert_eq!(msgs[0]["role"], "tool");
assert_eq!(msgs[0]["tool_call_id"], "call_123");
assert_eq!(msgs[0]["content"], "file1.txt\nfile2.txt");
}
#[test]
fn test_tool_result_error_conversion() {
let messages = vec![ConversationMessage {
role: MessageRole::User,
content: MessageContent::ToolResult {
tool_use_id: "call_456".to_string(),
content: "command not found".to_string(),
is_error: true,
},
}];
let request = build_openai_request(messages, None, vec![], 1024, None, "test-model");
let msgs = request["messages"].as_array().unwrap();
assert_eq!(msgs[0]["role"], "tool");
assert_eq!(msgs[0]["content"], "[ERROR] command not found");
}
#[test]
fn test_multipart_assistant_message() {
let messages = vec![ConversationMessage {
role: MessageRole::Assistant,
content: MessageContent::MultiPart(vec![
ContentPart::Text("I'll run that command for you.".to_string()),
ContentPart::ToolUse {
tool_use_id: "call_abc".to_string(),
name: "run_shell_command".to_string(),
input: json!({"command": "pwd"}),
},
]),
}];
let request = build_openai_request(messages, None, vec![], 1024, None, "test-model");
let msgs = request["messages"].as_array().unwrap();
assert_eq!(msgs.len(), 1);
assert_eq!(msgs[0]["role"], "assistant");
assert_eq!(msgs[0]["content"], "I'll run that command for you.");
let tool_calls = msgs[0]["tool_calls"].as_array().unwrap();
assert_eq!(tool_calls.len(), 1);
assert_eq!(tool_calls[0]["id"], "call_abc");
assert_eq!(tool_calls[0]["function"]["name"], "run_shell_command");
}
#[test]
fn test_multipart_user_message_with_tool_results() {
let messages = vec![ConversationMessage {
role: MessageRole::User,
content: MessageContent::MultiPart(vec![
ContentPart::ToolResult {
tool_use_id: "call_1".to_string(),
content: "result 1".to_string(),
is_error: false,
},
ContentPart::ToolResult {
tool_use_id: "call_2".to_string(),
content: "result 2".to_string(),
is_error: false,
},
]),
}];
let request = build_openai_request(messages, None, vec![], 1024, None, "test-model");
let msgs = request["messages"].as_array().unwrap();
assert_eq!(msgs.len(), 2);
assert_eq!(msgs[0]["role"], "tool");
assert_eq!(msgs[0]["tool_call_id"], "call_1");
assert_eq!(msgs[1]["role"], "tool");
assert_eq!(msgs[1]["tool_call_id"], "call_2");
}
#[test]
fn test_tool_definitions_conversion() {
let tools = vec![
ToolDefinition {
name: "run_shell_command".to_string(),
description: "Runs a shell command".to_string(),
input_schema: json!({
"type": "object",
"properties": {
"command": {"type": "string"}
},
"required": ["command"]
}),
},
ToolDefinition {
name: "read_files".to_string(),
description: "Reads files from disk".to_string(),
input_schema: json!({
"type": "object",
"properties": {
"files": {"type": "array", "items": {"type": "string"}}
}
}),
},
];
let request = build_openai_request(vec![], None, tools, 1024, None, "test-model");
let tool_defs = request["tools"].as_array().unwrap();
assert_eq!(tool_defs.len(), 2);
assert_eq!(tool_defs[0]["type"], "function");
assert_eq!(tool_defs[0]["function"]["name"], "run_shell_command");
assert_eq!(
tool_defs[0]["function"]["description"],
"Runs a shell command"
);
assert_eq!(tool_defs[0]["function"]["parameters"]["type"], "object");
assert_eq!(tool_defs[1]["function"]["name"], "read_files");
}
#[test]
fn test_temperature_handling() {
let messages = vec![ConversationMessage {
role: MessageRole::User,
content: MessageContent::Text("test".to_string()),
}];
// Temperature absent when None
let request = build_openai_request(messages.clone(), None, vec![], 1024, None, "test-model");
assert!(request.get("temperature").is_none());
// Temperature present when Some
let request = build_openai_request(messages, None, vec![], 1024, Some(0.7), "test-model");
assert!(request.get("temperature").is_some());
}
#[test]
fn test_full_conversation_roundtrip() {
let messages = vec![
ConversationMessage {
role: MessageRole::User,
content: MessageContent::Text("List files".to_string()),
},
ConversationMessage {
role: MessageRole::Assistant,
content: MessageContent::ToolUse {
tool_use_id: "call_1".to_string(),
name: "run_shell_command".to_string(),
input: json!({"command": "ls"}),
},
},
ConversationMessage {
role: MessageRole::User,
content: MessageContent::ToolResult {
tool_use_id: "call_1".to_string(),
content: "file1.rs\nfile2.rs".to_string(),
is_error: false,
},
},
ConversationMessage {
role: MessageRole::Assistant,
content: MessageContent::Text("Here are the files: file1.rs and file2.rs".to_string()),
},
];
let request = build_openai_request(
messages,
Some("You are Galaxy AI.".to_string()),
vec![],
4096,
None,
"claude-sonnet",
);
let msgs = request["messages"].as_array().unwrap();
assert_eq!(msgs.len(), 5); // system + 4 conversation messages
assert_eq!(msgs[0]["role"], "system");
assert_eq!(msgs[1]["role"], "user");
assert_eq!(msgs[2]["role"], "assistant");
assert_eq!(msgs[3]["role"], "tool");
assert_eq!(msgs[4]["role"], "assistant");
}
+13
View File
@@ -0,0 +1,13 @@
pub mod client;
pub mod convert;
pub mod request_translator;
pub mod response_translator;
pub mod translator;
#[cfg(test)]
#[path = "convert_tests.rs"]
mod convert_tests;
#[cfg(test)]
#[path = "request_translator_tests.rs"]
mod request_translator_tests;
+140
View File
@@ -0,0 +1,140 @@
use crate::ai::provider::types::{ContentPart, ConversationMessage, MessageContent, MessageRole};
/// Sanitizes messages for OpenAI API compatibility.
///
/// OpenAI is more lenient than Bedrock — it doesn't require strict user/assistant
/// alternation and allows system messages anywhere. The main constraints are:
/// - Tool results must reference a valid tool_call_id from a preceding assistant message
/// - Tool calls in assistant messages must eventually have matching tool results
pub fn sanitize_messages_for_openai(messages: &mut Vec<ConversationMessage>) {
remove_orphaned_tool_results(messages);
synthesize_missing_tool_results(messages);
}
/// Removes tool_result messages that reference tool_use_ids not found in any
/// preceding assistant message.
fn remove_orphaned_tool_results(messages: &mut Vec<ConversationMessage>) {
let mut known_tool_use_ids: std::collections::HashSet<String> =
std::collections::HashSet::new();
// First pass: collect all tool_use_ids from assistant messages
for msg in messages.iter() {
if msg.role != MessageRole::Assistant {
continue;
}
collect_tool_use_ids(&msg.content, &mut known_tool_use_ids);
}
// Second pass: remove tool_results that reference unknown IDs
messages.retain(|msg| {
if msg.role != MessageRole::User {
return true;
}
match &msg.content {
MessageContent::ToolResult { tool_use_id, .. } => {
known_tool_use_ids.contains(tool_use_id)
}
MessageContent::MultiPart(parts) => {
// Keep the message if it has at least one non-orphaned part
parts.iter().any(|part| match part {
ContentPart::ToolResult { tool_use_id, .. } => {
known_tool_use_ids.contains(tool_use_id)
}
_ => true,
})
}
_ => true,
}
});
}
/// For any assistant tool_use that doesn't have a matching tool_result in a
/// subsequent user message, synthesize an error result.
fn synthesize_missing_tool_results(messages: &mut Vec<ConversationMessage>) {
let mut pending_tool_use_ids: Vec<(String, usize)> = Vec::new();
let mut answered_ids: std::collections::HashSet<String> = std::collections::HashSet::new();
// Collect all tool_use IDs and all answered IDs
for (i, msg) in messages.iter().enumerate() {
match msg.role {
MessageRole::Assistant => {
collect_tool_use_ids_with_index(&msg.content, i, &mut pending_tool_use_ids);
}
MessageRole::User => {
collect_tool_result_ids(&msg.content, &mut answered_ids);
}
}
}
// Find unanswered tool_uses and synthesize results
let mut synthetic_results: Vec<ConversationMessage> = Vec::new();
for (tool_use_id, _) in pending_tool_use_ids {
if !answered_ids.contains(&tool_use_id) {
synthetic_results.push(ConversationMessage {
role: MessageRole::User,
content: MessageContent::ToolResult {
tool_use_id,
content: "Tool call result unavailable (conversation was interrupted)."
.to_string(),
is_error: true,
},
});
}
}
if !synthetic_results.is_empty() {
messages.extend(synthetic_results);
}
}
fn collect_tool_use_ids(content: &MessageContent, ids: &mut std::collections::HashSet<String>) {
match content {
MessageContent::ToolUse { tool_use_id, .. } => {
ids.insert(tool_use_id.clone());
}
MessageContent::MultiPart(parts) => {
for part in parts {
if let ContentPart::ToolUse { tool_use_id, .. } = part {
ids.insert(tool_use_id.clone());
}
}
}
_ => {}
}
}
fn collect_tool_use_ids_with_index(
content: &MessageContent,
index: usize,
ids: &mut Vec<(String, usize)>,
) {
match content {
MessageContent::ToolUse { tool_use_id, .. } => {
ids.push((tool_use_id.clone(), index));
}
MessageContent::MultiPart(parts) => {
for part in parts {
if let ContentPart::ToolUse { tool_use_id, .. } = part {
ids.push((tool_use_id.clone(), index));
}
}
}
_ => {}
}
}
fn collect_tool_result_ids(content: &MessageContent, ids: &mut std::collections::HashSet<String>) {
match content {
MessageContent::ToolResult { tool_use_id, .. } => {
ids.insert(tool_use_id.clone());
}
MessageContent::MultiPart(parts) => {
for part in parts {
if let ContentPart::ToolResult { tool_use_id, .. } = part {
ids.insert(tool_use_id.clone());
}
}
}
_ => {}
}
}
@@ -0,0 +1,148 @@
use serde_json::json;
use crate::ai::openai::request_translator::sanitize_messages_for_openai;
use crate::ai::provider::types::{ContentPart, ConversationMessage, MessageContent, MessageRole};
#[test]
fn test_removes_orphaned_tool_results() {
let mut messages = vec![
ConversationMessage {
role: MessageRole::User,
content: MessageContent::Text("Hello".to_string()),
},
// This tool result references a tool_use that doesn't exist
ConversationMessage {
role: MessageRole::User,
content: MessageContent::ToolResult {
tool_use_id: "nonexistent_id".to_string(),
content: "some result".to_string(),
is_error: false,
},
},
];
sanitize_messages_for_openai(&mut messages);
assert_eq!(messages.len(), 1);
matches!(&messages[0].content, MessageContent::Text(_));
}
#[test]
fn test_keeps_valid_tool_results() {
let mut messages = vec![
ConversationMessage {
role: MessageRole::Assistant,
content: MessageContent::ToolUse {
tool_use_id: "valid_id".to_string(),
name: "run_shell_command".to_string(),
input: json!({"command": "ls"}),
},
},
ConversationMessage {
role: MessageRole::User,
content: MessageContent::ToolResult {
tool_use_id: "valid_id".to_string(),
content: "file1.txt".to_string(),
is_error: false,
},
},
];
sanitize_messages_for_openai(&mut messages);
assert_eq!(messages.len(), 2);
}
#[test]
fn test_synthesizes_missing_tool_results() {
let mut messages = vec![
ConversationMessage {
role: MessageRole::Assistant,
content: MessageContent::ToolUse {
tool_use_id: "unanswered_id".to_string(),
name: "read_files".to_string(),
input: json!({"files": ["test.rs"]}),
},
},
// No corresponding tool result!
];
sanitize_messages_for_openai(&mut messages);
// Should have synthesized a tool result
assert_eq!(messages.len(), 2);
match &messages[1].content {
MessageContent::ToolResult {
tool_use_id,
is_error,
..
} => {
assert_eq!(tool_use_id, "unanswered_id");
assert!(*is_error);
}
_ => panic!("Expected ToolResult"),
}
}
#[test]
fn test_does_not_require_user_assistant_alternation() {
let mut messages = vec![
ConversationMessage {
role: MessageRole::User,
content: MessageContent::Text("First message".to_string()),
},
ConversationMessage {
role: MessageRole::User,
content: MessageContent::Text("Second message".to_string()),
},
ConversationMessage {
role: MessageRole::Assistant,
content: MessageContent::Text("Response".to_string()),
},
];
sanitize_messages_for_openai(&mut messages);
// Both user messages should remain — OpenAI allows consecutive same-role
assert_eq!(messages.len(), 3);
}
#[test]
fn test_does_not_require_starting_with_user() {
let mut messages = vec![ConversationMessage {
role: MessageRole::Assistant,
content: MessageContent::Text("I start the conversation".to_string()),
}];
sanitize_messages_for_openai(&mut messages);
// Should NOT prepend a user message (unlike Bedrock)
assert_eq!(messages.len(), 1);
assert_eq!(messages[0].role, MessageRole::Assistant);
}
#[test]
fn test_multipart_tool_uses_all_get_results() {
let mut messages = vec![ConversationMessage {
role: MessageRole::Assistant,
content: MessageContent::MultiPart(vec![
ContentPart::ToolUse {
tool_use_id: "id_1".to_string(),
name: "grep".to_string(),
input: json!({"queries": ["test"]}),
},
ContentPart::ToolUse {
tool_use_id: "id_2".to_string(),
name: "file_glob".to_string(),
input: json!({"patterns": ["*.rs"]}),
},
]),
}];
sanitize_messages_for_openai(&mut messages);
// Should synthesize results for both unanswered tool calls
assert_eq!(messages.len(), 3);
assert_eq!(messages[1].role, MessageRole::User);
assert_eq!(messages[2].role, MessageRole::User);
}
+515
View File
@@ -0,0 +1,515 @@
use std::sync::{Arc, Mutex};
use bytes::Bytes;
use futures::stream::BoxStream;
use futures::Stream;
use serde_json::Value as JsonValue;
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;
use crate::ai::bedrock::response_translator::{
build_create_task, build_stream_init, context_window_for_model,
};
use crate::ai::provider::types::{ContentPart, ConversationMessage, MessageContent, MessageRole};
use crate::server::server_api::AIApiError;
struct ToolCallAccumulator {
#[allow(dead_code)]
index: usize,
id: String,
name: String,
arguments: String,
}
pub fn openai_stream_to_response_events(
byte_stream: impl Stream<Item = Result<Bytes, reqwest::Error>> + Send + 'static,
task_id: String,
needs_create_task: bool,
user_query: Option<String>,
messages_sent: Arc<Mutex<Vec<ConversationMessage>>>,
model_id: String,
_tool_result_archive: Vec<ConversationMessage>,
) -> BoxStream<'static, Event> {
use futures::StreamExt;
let request_id = Uuid::new_v4().to_string();
let conversation_id = Uuid::new_v4().to_string();
let stream = async_stream::stream! {
log::info!("[openai] Stream started: task_id={task_id}, request_id={request_id}");
let init_event = build_stream_init(&request_id, &conversation_id);
yield Ok(init_event);
if needs_create_task {
let create_task_event = build_create_task(&task_id);
yield Ok(create_task_event);
}
if let Some(ref query_text) = user_query {
let user_query_msg = build_user_query_message(&task_id, query_text);
yield Ok(user_query_msg);
}
let mut current_text_message_id: Option<String> = None;
let mut full_text = String::new();
let mut tool_calls: Vec<ToolCallAccumulator> = Vec::new();
let mut input_tokens: i32 = 0;
let mut output_tokens: i32 = 0;
let mut stop_reason = stream_finished::Reason::Done(api::response_event::stream_finished::Done {});
let mut line_buffer = String::new();
futures::pin_mut!(byte_stream);
while let Some(chunk_result) = byte_stream.next().await {
let chunk = match chunk_result {
Ok(bytes) => bytes,
Err(e) => {
log::error!("[openai] Stream chunk error: {e}");
yield Err(Arc::new(AIApiError::Stream {
stream_type: "openai_chat_completions",
source: anyhow::anyhow!("{e}"),
}));
return;
}
};
let chunk_str = String::from_utf8_lossy(&chunk);
line_buffer.push_str(&chunk_str);
// Process complete SSE lines
while let Some(line_end) = line_buffer.find('\n') {
let line = line_buffer[..line_end].trim_end_matches('\r').to_string();
line_buffer = line_buffer[line_end + 1..].to_string();
if line.is_empty() {
continue;
}
if line == "data: [DONE]" {
log::info!("[openai] Stream complete: [DONE]");
break;
}
if let Some(data) = line.strip_prefix("data: ") {
let parsed: JsonValue = match serde_json::from_str(data) {
Ok(v) => v,
Err(e) => {
log::warn!("[openai] Failed to parse SSE data: {e}");
continue;
}
};
// Extract usage from the chunk (may appear in any chunk or final one)
if let Some(usage) = parsed.get("usage") {
if let Some(prompt) = usage.get("prompt_tokens").and_then(|v| v.as_i64()) {
input_tokens = prompt as i32;
}
if let Some(completion) = usage.get("completion_tokens").and_then(|v| v.as_i64()) {
output_tokens = completion as i32;
}
}
// Process choices
let choices = match parsed.get("choices").and_then(|v| v.as_array()) {
Some(c) => c,
None => continue,
};
for choice in choices {
// Check finish_reason
if let Some(reason) = choice.get("finish_reason").and_then(|v| v.as_str()) {
match reason {
"stop" => {
stop_reason = stream_finished::Reason::Done(
api::response_event::stream_finished::Done {},
);
}
"tool_calls" => {
stop_reason = stream_finished::Reason::Done(
api::response_event::stream_finished::Done {},
);
}
"length" => {
stop_reason = stream_finished::Reason::MaxTokenLimit(
stream_finished::ReachedMaxTokenLimit {},
);
}
_ => {}
}
}
let delta = match choice.get("delta") {
Some(d) => d,
None => continue,
};
// Handle text content
if let Some(content) = delta.get("content").and_then(|v| v.as_str()) {
if !content.is_empty() {
full_text.push_str(content);
if let Some(ref msg_id) = current_text_message_id {
let event = build_append_text(&task_id, msg_id, content);
yield Ok(event);
} else {
let msg_id = Uuid::new_v4().to_string();
let event = build_add_agent_output_message(&task_id, &msg_id, content);
current_text_message_id = Some(msg_id);
yield Ok(event);
}
}
}
// Handle tool calls
if let Some(tc_array) = delta.get("tool_calls").and_then(|v| v.as_array()) {
for tc in tc_array {
let index = tc.get("index").and_then(|v| v.as_u64()).unwrap_or(0) as usize;
// Extend tool_calls vector if needed
while tool_calls.len() <= index {
tool_calls.push(ToolCallAccumulator {
index: tool_calls.len(),
id: String::new(),
name: String::new(),
arguments: String::new(),
});
}
if let Some(id) = tc.get("id").and_then(|v| v.as_str()) {
tool_calls[index].id = id.to_string();
}
if let Some(function) = tc.get("function") {
if let Some(name) = function.get("name").and_then(|v| v.as_str()) {
tool_calls[index].name = name.to_string();
}
if let Some(args) = function.get("arguments").and_then(|v| v.as_str()) {
tool_calls[index].arguments.push_str(args);
}
}
}
}
}
}
}
}
// Emit tool call messages for completed tool calls
let mut assistant_parts: Vec<ContentPart> = Vec::new();
if !full_text.is_empty() {
assistant_parts.push(ContentPart::Text(full_text.clone()));
}
for tc in &tool_calls {
if tc.id.is_empty() || tc.name.is_empty() {
continue;
}
let event = build_tool_call_message(&task_id, &tc.id, &tc.name, &tc.arguments);
yield Ok(event);
let input: JsonValue = serde_json::from_str(&tc.arguments).unwrap_or(serde_json::json!({}));
assistant_parts.push(ContentPart::ToolUse {
tool_use_id: tc.id.clone(),
name: tc.name.clone(),
input,
});
}
// Store the complete assistant message in messages_sent
if !assistant_parts.is_empty() {
let assistant_msg = if assistant_parts.len() == 1 {
match assistant_parts.into_iter().next().unwrap() {
ContentPart::Text(text) => ConversationMessage {
role: MessageRole::Assistant,
content: MessageContent::Text(text),
},
ContentPart::ToolUse { tool_use_id, name, input } => ConversationMessage {
role: MessageRole::Assistant,
content: MessageContent::ToolUse { tool_use_id, name, input },
},
_ => unreachable!(),
}
} else {
ConversationMessage {
role: MessageRole::Assistant,
content: MessageContent::MultiPart(assistant_parts),
}
};
if let Ok(mut sent) = messages_sent.lock() {
sent.push(assistant_msg);
}
}
// Emit hallucinated tool error results (tools the model called that aren't known)
for tc in &tool_calls {
if tc.id.is_empty() || tc.name.is_empty() {
continue;
}
if !is_known_tool(&tc.name) {
log::warn!("[openai] Model called unknown tool: {}", tc.name);
let error_result = ConversationMessage {
role: MessageRole::User,
content: MessageContent::ToolResult {
tool_use_id: tc.id.clone(),
content: format!(
"Error: '{}' is not a valid tool. Please use one of the available tools.",
tc.name
),
is_error: true,
},
};
if let Ok(mut sent) = messages_sent.lock() {
sent.push(error_result);
}
}
}
let cost = estimate_cost_cents(input_tokens as u32, output_tokens as u32, &model_id);
let finished_event = build_stream_finished(stop_reason, input_tokens, output_tokens, cost, &model_id);
yield Ok(finished_event);
log::info!("[openai] Stream finished: input_tokens={input_tokens}, output_tokens={output_tokens}");
};
Box::pin(stream)
}
fn build_user_query_message(task_id: &str, query_text: &str) -> ResponseEvent {
let message = api::Message {
id: Uuid::new_v4().to_string(),
task_id: task_id.to_string(),
request_id: String::new(),
timestamp: None,
server_message_data: String::new(),
citations: vec![],
message: Some(api::message::Message::UserQuery(api::message::UserQuery {
query: query_text.to_string(),
..Default::default()
})),
};
let action = ClientAction {
action: Some(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![action],
},
)),
}
}
fn build_add_agent_output_message(
task_id: &str,
message_id: &str,
initial_text: &str,
) -> 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![],
message: Some(api::message::Message::AgentOutput(
api::message::AgentOutput {
text: initial_text.to_string(),
},
)),
};
let action = ClientAction {
action: Some(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![action],
},
)),
}
}
fn build_append_text(task_id: &str, message_id: &str, text_delta: &str) -> 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![],
message: Some(api::message::Message::AgentOutput(
api::message::AgentOutput {
text: text_delta.to_string(),
},
)),
};
let mask = prost_types::FieldMask {
paths: vec!["agent_output.text".to_string()],
};
let action = ClientAction {
action: Some(api::client_action::Action::AppendToMessageContent(
api::client_action::AppendToMessageContent {
task_id: task_id.to_string(),
message: Some(message),
mask: Some(mask),
},
)),
};
ResponseEvent {
r#type: Some(api::response_event::Type::ClientActions(
api::response_event::ClientActions {
actions: vec![action],
},
)),
}
}
fn build_tool_call_message(
task_id: &str,
tool_use_id: &str,
tool_name: &str,
tool_input_json: &str,
) -> ResponseEvent {
// Reuse the Bedrock tool call message builder since the proto output is identical
crate::ai::bedrock::response_translator::build_tool_call_message(
task_id,
tool_use_id,
tool_name,
tool_input_json,
)
}
fn build_stream_finished(
reason: stream_finished::Reason,
input_tokens: i32,
output_tokens: i32,
cost_in_cents: f32,
model_id: &str,
) -> ResponseEvent {
let total_tokens = (input_tokens + output_tokens) as u32;
let mut byok_token_usage = std::collections::HashMap::new();
if total_tokens > 0 {
#[allow(deprecated)]
byok_token_usage.insert(
"openai".to_string(),
stream_finished::ModelTokenUsage {
model_id: String::new(),
total_tokens,
token_usage_by_category: std::collections::HashMap::new(),
},
);
}
let token_usage = vec![stream_finished::TokenUsage {
model_id: "openai".to_string(),
total_input: input_tokens as u32,
output: output_tokens as u32,
input_cache_read: 0,
input_cache_write: 0,
cost_in_cents,
}];
let max_context_tokens = context_window_for_model(model_id);
let context_usage = if max_context_tokens > 0 {
input_tokens as f32 / max_context_tokens as f32
} else {
0.0
};
#[allow(deprecated)]
let conversation_usage_metadata = Some(stream_finished::ConversationUsageMetadata {
context_window_usage: context_usage,
summarized: false,
credits_spent: 0.0,
token_usage: vec![],
tool_usage_metadata: None,
warp_token_usage: std::collections::HashMap::new(),
byok_token_usage,
});
ResponseEvent {
r#type: Some(api::response_event::Type::Finished(
api::response_event::StreamFinished {
reason: Some(reason),
token_usage,
should_refresh_model_config: false,
request_cost: None,
conversation_usage_metadata,
},
)),
}
}
/// LiteLLM proxies to various backends — estimate cost based on model name.
/// These are rough estimates; actual billing comes from LiteLLM.
fn estimate_cost_cents(input_tokens: u32, output_tokens: u32, model_id: &str) -> f32 {
let lower = model_id.to_lowercase();
let (input_rate, output_rate) = if lower.contains("opus") {
(15.0, 75.0)
} else if lower.contains("haiku") {
(0.80, 4.0)
} else if lower.contains("sonnet") {
(3.0, 15.0)
} else if lower.contains("gpt-4o") {
(2.50, 10.0)
} else if lower.contains("gpt-4") {
(30.0, 60.0)
} else if lower.contains("gpt-3.5") {
(0.50, 1.50)
} else {
(3.0, 15.0) // Default to Sonnet-tier pricing
};
let input_cost = input_tokens as f64 * input_rate * 100.0 / 1_000_000.0;
let output_cost = output_tokens as f64 * output_rate * 100.0 / 1_000_000.0;
(input_cost + output_cost) as f32
}
const KNOWN_TOOLS: &[&str] = &[
"run_shell_command",
"read_files",
"apply_file_diffs",
"grep",
"file_glob",
"search_codebase",
"write_to_long_running_shell_command",
"read_shell_command_output",
"read_mcp_resource",
"read_documents",
"create_documents",
"edit_documents",
"start_agent",
"send_message_to_agent",
"ask_user_question",
"suggest_next_prompt",
"read_skill",
"fetch_conversation",
"recall_tool_history",
];
fn is_known_tool(name: &str) -> bool {
KNOWN_TOOLS.contains(&name) || name.starts_with("mcp__")
}
+148
View File
@@ -0,0 +1,148 @@
use std::sync::{Arc, Mutex};
use warp_multi_agent_api as api;
use crate::ai::agent::api::ResponseStream;
use crate::ai::bedrock::request_translator;
use crate::ai::provider::types::{ConversationMessage, MessageContent, MessageRole};
use super::client::{OpenAIClient, OpenAIClientConfig, OpenAIError};
use super::convert::build_openai_request;
use super::request_translator::sanitize_messages_for_openai;
use super::response_translator::openai_stream_to_response_events;
pub struct TranslatorRequest {
pub config: OpenAIClientConfig,
pub model_id: String,
pub root_task_id: Option<String>,
pub message_history: Vec<ConversationMessage>,
pub tool_result_archive: Vec<ConversationMessage>,
pub progressive_summary: Option<String>,
pub messages_sent: Arc<Mutex<Vec<ConversationMessage>>>,
}
pub async fn execute(
params: TranslatorRequest,
request: &mut api::Request,
) -> Result<ResponseStream, OpenAIError> {
let client = OpenAIClient::from_config(params.config.clone());
let task_id = params.root_task_id.unwrap_or_else(|| {
request
.task_context
.as_ref()
.and_then(|tc| tc.tasks.first())
.map(|t| t.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())
.unwrap_or(true);
let model_id = if params.model_id.is_empty() || params.model_id == "auto" {
params
.config
.model
.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
.clone()
.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 mut messages = Vec::new();
// Prepend progressive summary as first message pair if present
if let Some(ref summary) = params.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
)),
});
messages.push(ConversationMessage {
role: MessageRole::Assistant,
content: MessageContent::Text(
"Understood, I have the prior context. Continuing with the recent conversation."
.to_string(),
),
});
}
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);
}
sanitize_messages_for_openai(&mut messages);
let system_prompt = request_translator::extract_system_prompt(request);
let tools = request_translator::extract_tools(request);
log::info!(
"[openai] Sending {} messages, system_prompt={}, tools={}",
messages.len(),
system_prompt.is_some(),
tools.len()
);
let user_query_text = request_translator::extract_user_query_text(request);
let request_body = build_openai_request(
messages.clone(),
system_prompt,
tools,
64000,
None,
&model_id,
);
let byte_stream = client.chat_completions_stream(request_body).await?;
// 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);
} else {
*sent = messages;
}
}
let stream = openai_stream_to_response_events(
byte_stream,
task_id,
needs_create_task,
user_query_text,
params.messages_sent.clone(),
model_id,
params.tool_result_archive,
);
Ok(stream)
}
+10
View File
@@ -0,0 +1,10 @@
pub mod types;
use crate::ai::bedrock::client::BedrockClientConfig;
use crate::ai::openai::client::OpenAIClientConfig;
pub enum ProviderConfig {
Bedrock(BedrockClientConfig),
OpenAI(OpenAIClientConfig),
None,
}
+51
View File
@@ -0,0 +1,51 @@
use serde_json::Value as JsonValue;
#[derive(Clone, Debug)]
pub struct ConversationMessage {
pub role: MessageRole,
pub content: MessageContent,
}
#[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),
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,
}