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:
@@ -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())
|
||||
}
|
||||
}
|
||||
@@ -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,
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -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");
|
||||
}
|
||||
@@ -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;
|
||||
@@ -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);
|
||||
}
|
||||
@@ -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__")
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
Reference in New Issue
Block a user