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:
@@ -41,25 +41,53 @@ Environment variables:
|
||||
- `Error_<timestamp>.txt` snapshot files written to the repository root on request/stream failures (includes the serialized Bedrock context window, tool definitions, protobuf request debug payload, captured Bedrock diagnostic lines, and log tails)
|
||||
- Per-event Bedrock diagnostic logs written to `bedrock-diagnostics.log` in the active Warp log directory
|
||||
|
||||
### Bedrock Translator Architecture
|
||||
The Bedrock integration uses a **translator service pattern** where Warp proto types flow in, get converted to Bedrock SDK types, and responses are translated back:
|
||||
### AI Provider Architecture
|
||||
|
||||
Galaxy supports multiple AI backends via a **provider dispatch pattern**. Provider selection
|
||||
is controlled by settings (`ai.openai.enabled` takes priority over `ai.bedrock.enabled`).
|
||||
|
||||
```
|
||||
Warp UI (proto) → translator.rs → request_translator.rs → Bedrock API
|
||||
Warp UI (proto) ← response_translator.rs ← Bedrock stream
|
||||
Provider dispatch: response_stream.rs → resolve_provider_config() → ProviderConfig enum
|
||||
↓ Bedrock ↓ OpenAI
|
||||
bedrock/translator.rs openai/translator.rs
|
||||
```
|
||||
|
||||
Key files in `app/src/ai/bedrock/`:
|
||||
**Shared types** in `app/src/ai/provider/`:
|
||||
- `types.rs` — `ConversationMessage`, `MessageRole`, `MessageContent`, `ContentPart`, `ToolDefinition`
|
||||
- `mod.rs` — `ProviderConfig` enum (Bedrock | OpenAI | None)
|
||||
|
||||
**Bedrock provider** in `app/src/ai/bedrock/`:
|
||||
- `translator.rs` — Orchestrator: takes `api::Request` + config, returns `ResponseStream`
|
||||
- `request_translator.rs` — Converts Warp proto → Bedrock SDK types (messages, system prompt, tools, sanitization)
|
||||
- `response_translator.rs` — Converts Bedrock stream events → Warp proto `ResponseEvent`s
|
||||
- `convert.rs` — Shared types (`ConversationMessage`, `ToolDefinition`) and Bedrock SDK type builders
|
||||
- `convert.rs` — Re-exports shared types + Bedrock SDK type builders
|
||||
- `client.rs` — AWS SDK client construction and `converse_stream` call
|
||||
- `models.rs` — Model registry and cross-region inference prefix logic
|
||||
- `discovery.rs` — AWS profile listing and model discovery (STS identity check + ListFoundationModels)
|
||||
- `diagnostic.rs` — Debug logging (enabled via `GALAXY_BEDROCK_DIAGNOSTICS=1`)
|
||||
- `external_config.rs` — Fallback config from Claude Code/OpenCode settings
|
||||
|
||||
**OpenAI/LiteLLM provider** in `app/src/ai/openai/`:
|
||||
- `translator.rs` — Orchestrator: same pattern as Bedrock, targets OpenAI chat completions API
|
||||
- `client.rs` — `reqwest`-based HTTP client for `POST /v1/chat/completions` with streaming
|
||||
- `convert.rs` — `ConversationMessage` → OpenAI JSON format (system/user/assistant/tool roles, function calling)
|
||||
- `request_translator.rs` — OpenAI-specific message sanitization (lighter than Bedrock's strict alternation rules)
|
||||
- `response_translator.rs` — SSE stream parser → Warp proto `ResponseEvent`s
|
||||
|
||||
**Provider settings** (in settings TOML):
|
||||
- `ai.bedrock.enabled` — Use AWS Bedrock directly (default: true)
|
||||
- `ai.openai.enabled` — Use OpenAI-compatible endpoint, e.g. LiteLLM (default: false, takes priority)
|
||||
- `ai.openai.base_url` — Endpoint URL (default: `http://localhost:4000/v1`)
|
||||
- `ai.openai.api_key` — Optional API key (stored in keychain)
|
||||
- `ai.openai.model` — Model name override sent to the endpoint
|
||||
- `ai.openai.models` — Array of `OpenAIModelConfig` objects (model_id, display_name, vision_supported, context_size, provider)
|
||||
|
||||
**OpenAI/LiteLLM model discovery**:
|
||||
- Models can be auto-fetched from the `/models` endpoint via the Settings > OpenAI / LiteLLM page
|
||||
- Fetched models include context window sizes from `max_model_len` / `context_window` / `max_input_tokens` fields
|
||||
- Models injected into `LLMPreferences` use `LLMProvider::LiteLLM` and show the OpenAI icon in the picker
|
||||
- Provider is inferred from model ID (claude→anthropic, gpt→openai, gemini→google)
|
||||
|
||||
Key invariants:
|
||||
- Known tools are in `KNOWN_TOOLS` constant in `response_translator.rs`
|
||||
- Tool definitions are built via `tool_definition_for_name()` in `convert_request.rs`; includes `recall_tool_history` for retrieving past tool results
|
||||
|
||||
@@ -169,7 +169,6 @@ fn add_features(target_family: &str, target_os: &str) {
|
||||
if env::var("PROFILE").ok().is_some_and(|val| val == "debug") {
|
||||
println!("cargo:rustc-cfg=feature=\"agent_mode_debug\"");
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
fn build_and_link_sentry() {
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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>,
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
@@ -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,
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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,
|
||||
}
|
||||
@@ -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,
|
||||
}
|
||||
+107
-2
@@ -462,6 +462,50 @@ fn default_context_size() -> u32 {
|
||||
|
||||
impl settings_value::SettingsValue for BedrockModelConfig {}
|
||||
|
||||
/// Configuration for a single OpenAI-compatible (LiteLLM) model.
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, schemars::JsonSchema)]
|
||||
#[schemars(description = "Configuration for a single OpenAI-compatible model (e.g. via LiteLLM).")]
|
||||
pub struct OpenAIModelConfig {
|
||||
#[schemars(description = "The model ID to send in the API request (e.g. claude-sonnet-4-20250514).")]
|
||||
pub model_id: String,
|
||||
#[schemars(description = "Display name shown in the model picker.")]
|
||||
pub display_name: String,
|
||||
#[serde(default)]
|
||||
#[schemars(description = "Whether the model supports image/vision input.")]
|
||||
pub vision_supported: bool,
|
||||
#[serde(default = "default_context_size")]
|
||||
#[schemars(description = "Maximum context window size in tokens.")]
|
||||
pub context_size: u32,
|
||||
#[serde(default)]
|
||||
#[schemars(description = "Optional provider hint (e.g. anthropic, openai, google) for icon display.")]
|
||||
pub provider: Option<String>,
|
||||
}
|
||||
|
||||
impl settings_value::SettingsValue for OpenAIModelConfig {}
|
||||
|
||||
/// Configuration for a single OpenAI-compatible provider endpoint.
|
||||
///
|
||||
/// Multiple providers can be configured simultaneously (e.g. LiteLLM for cloud models,
|
||||
/// Ollama for local models, etc.). Each provider has its own endpoint, credentials, and model list.
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, schemars::JsonSchema)]
|
||||
#[schemars(
|
||||
description = "Configuration for an OpenAI-compatible provider endpoint (e.g. LiteLLM, Ollama, vLLM)."
|
||||
)]
|
||||
pub struct OpenAIProviderConfig {
|
||||
#[schemars(description = "Display name for this provider (shown in model picker).")]
|
||||
pub name: String,
|
||||
#[schemars(description = "Base URL for the OpenAI-compatible API endpoint.")]
|
||||
pub base_url: String,
|
||||
#[serde(default)]
|
||||
#[schemars(description = "API key for this endpoint (optional if the proxy handles auth).")]
|
||||
pub api_key: Option<String>,
|
||||
#[serde(default)]
|
||||
#[schemars(description = "Models available from this provider.")]
|
||||
pub models: Vec<OpenAIModelConfig>,
|
||||
}
|
||||
|
||||
impl settings_value::SettingsValue for OpenAIProviderConfig {}
|
||||
|
||||
/// Tracks the state of the quota reset banner
|
||||
#[derive(
|
||||
Debug,
|
||||
@@ -1183,6 +1227,67 @@ define_settings_group!(AISettings, settings: [
|
||||
sync_to_cloud: SyncToCloud::Globally(RespectUserSyncSetting::Yes),
|
||||
private: true,
|
||||
}
|
||||
// Whether the OpenAI-compatible (LiteLLM) provider is enabled.
|
||||
openai_enabled: OpenAIEnabled {
|
||||
type: bool,
|
||||
default: false,
|
||||
supported_platforms: SupportedPlatforms::DESKTOP,
|
||||
sync_to_cloud: SyncToCloud::Globally(RespectUserSyncSetting::Yes),
|
||||
private: false,
|
||||
toml_path: "ai.openai.enabled",
|
||||
description: "Whether to use an OpenAI-compatible endpoint (e.g. LiteLLM) for AI requests.",
|
||||
}
|
||||
// Base URL for the OpenAI-compatible API endpoint.
|
||||
openai_base_url: OpenAIBaseUrl {
|
||||
type: String,
|
||||
default: "http://localhost:4000/v1".to_string(),
|
||||
supported_platforms: SupportedPlatforms::DESKTOP,
|
||||
sync_to_cloud: SyncToCloud::Globally(RespectUserSyncSetting::Yes),
|
||||
private: false,
|
||||
toml_path: "ai.openai.base_url",
|
||||
description: "Base URL for the OpenAI-compatible API endpoint (e.g. LiteLLM proxy).",
|
||||
}
|
||||
// API key for the OpenAI-compatible endpoint (optional if proxy handles auth).
|
||||
openai_api_key: OpenAIApiKey {
|
||||
type: String,
|
||||
default: String::new(),
|
||||
supported_platforms: SupportedPlatforms::DESKTOP,
|
||||
sync_to_cloud: SyncToCloud::Never,
|
||||
private: false,
|
||||
toml_path: "ai.openai.api_key",
|
||||
description: "API key for the OpenAI-compatible endpoint (optional if proxy handles auth).",
|
||||
}
|
||||
// Model name to send to the OpenAI-compatible endpoint. Empty = use selected model ID.
|
||||
openai_model: OpenAIModel {
|
||||
type: String,
|
||||
default: String::new(),
|
||||
supported_platforms: SupportedPlatforms::DESKTOP,
|
||||
sync_to_cloud: SyncToCloud::Globally(RespectUserSyncSetting::Yes),
|
||||
private: false,
|
||||
toml_path: "ai.openai.model",
|
||||
description: "Model name to send to the OpenAI-compatible endpoint. Leave empty to use the selected model ID.",
|
||||
}
|
||||
// Custom OpenAI-compatible model configurations (fetched from LiteLLM or manually configured).
|
||||
openai_models: OpenAIModels {
|
||||
type: Vec<OpenAIModelConfig>,
|
||||
default: Vec::new(),
|
||||
supported_platforms: SupportedPlatforms::DESKTOP,
|
||||
sync_to_cloud: SyncToCloud::Globally(RespectUserSyncSetting::Yes),
|
||||
private: false,
|
||||
toml_path: "ai.openai.models",
|
||||
description: "Custom OpenAI-compatible model configurations (e.g. from LiteLLM).",
|
||||
}
|
||||
// Multiple OpenAI-compatible provider endpoints (LiteLLM, Ollama, vLLM, etc.).
|
||||
// Each provider has its own name, base_url, api_key, and model list.
|
||||
openai_providers: OpenAIProviders {
|
||||
type: Vec<OpenAIProviderConfig>,
|
||||
default: Vec::new(),
|
||||
supported_platforms: SupportedPlatforms::DESKTOP,
|
||||
sync_to_cloud: SyncToCloud::Globally(RespectUserSyncSetting::Yes),
|
||||
private: false,
|
||||
toml_path: "ai.providers",
|
||||
description: "Multiple OpenAI-compatible provider endpoints (e.g. LiteLLM, Ollama, local models).",
|
||||
}
|
||||
// Whether or not the user wants agent mode requests to use their saved rules.
|
||||
memory_enabled: MemoryEnabled {
|
||||
type: bool,
|
||||
@@ -1728,8 +1833,8 @@ impl AISettings {
|
||||
*self.file_based_mcp_enabled
|
||||
}
|
||||
|
||||
pub fn is_orchestration_enabled(&self, app: &galaxyui::AppContext) -> bool {
|
||||
FeatureFlag::Orchestration.is_enabled() && self.is_any_ai_enabled(app)
|
||||
pub fn is_orchestration_enabled(&self, _app: &galaxyui::AppContext) -> bool {
|
||||
false
|
||||
}
|
||||
|
||||
/// Determines whether a quota reset banner should be displayed to the user.
|
||||
|
||||
@@ -26,7 +26,7 @@ use crate::settings::{
|
||||
AgentModeCommandExecutionPredicate, AgentModeQuerySuggestionsEnabled, BedrockAuthMethod,
|
||||
BedrockAutoLogin, BedrockEnabled, CodeSettings, CodebaseContextEnabled, FileBasedMcpEnabled,
|
||||
GitOperationsAutogenEnabled, IncludeAgentCommandsInHistory, IntelligentAutosuggestionsEnabled,
|
||||
MemoryEnabled, NLDInTerminalEnabled, NaturalLanguageAutosuggestionsEnabled,
|
||||
MemoryEnabled, NLDInTerminalEnabled, NaturalLanguageAutosuggestionsEnabled, OpenAIEnabled,
|
||||
RuleSuggestionsEnabled, SharedBlockTitleGenerationEnabled, ShouldRenderCLIAgentToolbar,
|
||||
ShouldRenderUseAgentToolbarForUserCommands, ShowAgentTips, ShowConversationHistory,
|
||||
ShowHintText, ThinkingDisplayMode, VoiceInputEnabled, WarpDriveContextEnabled,
|
||||
@@ -97,6 +97,8 @@ pub enum AISubpage {
|
||||
ThirdPartyCLIAgents,
|
||||
/// AWS Bedrock direct provider configuration.
|
||||
Bedrock,
|
||||
/// OpenAI-compatible (LiteLLM) provider configuration.
|
||||
OpenAI,
|
||||
}
|
||||
|
||||
impl AISubpage {
|
||||
@@ -107,6 +109,7 @@ impl AISubpage {
|
||||
SettingsSection::Knowledge => Some(Self::Knowledge),
|
||||
SettingsSection::ThirdPartyCLIAgents => Some(Self::ThirdPartyCLIAgents),
|
||||
SettingsSection::Bedrock => Some(Self::Bedrock),
|
||||
SettingsSection::OpenAI => Some(Self::OpenAI),
|
||||
// AgentMCPServers renders the standalone MCPServers page, not an AI subpage.
|
||||
_ => None,
|
||||
}
|
||||
@@ -1393,6 +1396,130 @@ impl AISettingsPageView {
|
||||
}
|
||||
}
|
||||
|
||||
/// Fetches models from the LiteLLM endpoint and updates settings.
|
||||
fn fetch_litellm_models(&mut self, ctx: &mut ViewContext<Self>) {
|
||||
let settings = AISettings::as_ref(ctx);
|
||||
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)
|
||||
}
|
||||
};
|
||||
|
||||
let _ = ctx.spawn(
|
||||
async move {
|
||||
use crate::settings::ai::OpenAIModelConfig;
|
||||
|
||||
let url = format!("{}/models", base_url.trim_end_matches('/'));
|
||||
let client = reqwest::Client::new();
|
||||
let mut request = client.get(&url);
|
||||
if let Some(ref key) = api_key {
|
||||
request = request.header("Authorization", format!("Bearer {key}"));
|
||||
}
|
||||
|
||||
let response = match request.send().await {
|
||||
Ok(r) => r,
|
||||
Err(e) => {
|
||||
log::error!("[litellm] Failed to fetch models: {e}");
|
||||
return Vec::new();
|
||||
}
|
||||
};
|
||||
|
||||
if !response.status().is_success() {
|
||||
log::error!(
|
||||
"[litellm] Model fetch returned HTTP {}",
|
||||
response.status()
|
||||
);
|
||||
return Vec::new();
|
||||
}
|
||||
|
||||
let body: serde_json::Value = match response.json().await {
|
||||
Ok(v) => v,
|
||||
Err(e) => {
|
||||
log::error!("[litellm] Failed to parse models response: {e}");
|
||||
return Vec::new();
|
||||
}
|
||||
};
|
||||
|
||||
// LiteLLM /models endpoint returns OpenAI-compatible format:
|
||||
// { "data": [{ "id": "model-name", "max_model_len": N, ... }] }
|
||||
let models: Vec<OpenAIModelConfig> = body["data"]
|
||||
.as_array()
|
||||
.unwrap_or(&vec![])
|
||||
.iter()
|
||||
.filter_map(|m| {
|
||||
let id = m["id"].as_str()?;
|
||||
// Try multiple context window fields used by different proxies
|
||||
let context_size = m["max_model_len"]
|
||||
.as_u64()
|
||||
.or_else(|| m["context_window"].as_u64())
|
||||
.or_else(|| m["max_input_tokens"].as_u64())
|
||||
.unwrap_or(200_000) as u32;
|
||||
|
||||
// Derive display name from model ID
|
||||
let display_name = id
|
||||
.split('/')
|
||||
.last()
|
||||
.unwrap_or(id)
|
||||
.replace('-', " ")
|
||||
.replace('_', " ");
|
||||
// Capitalize first letter of each word
|
||||
let display_name = display_name
|
||||
.split_whitespace()
|
||||
.map(|word| {
|
||||
let mut chars = word.chars();
|
||||
match chars.next() {
|
||||
None => String::new(),
|
||||
Some(c) => {
|
||||
c.to_uppercase().to_string() + chars.as_str()
|
||||
}
|
||||
}
|
||||
})
|
||||
.collect::<Vec<_>>()
|
||||
.join(" ");
|
||||
|
||||
// Infer provider from model ID prefix
|
||||
let provider = if id.contains("claude") || id.contains("anthropic") {
|
||||
Some("anthropic".to_string())
|
||||
} else if id.contains("gpt") || id.contains("o1") || id.contains("o3") {
|
||||
Some("openai".to_string())
|
||||
} else if id.contains("gemini") {
|
||||
Some("google".to_string())
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
Some(OpenAIModelConfig {
|
||||
model_id: id.to_string(),
|
||||
display_name,
|
||||
vision_supported: m["supports_vision"].as_bool().unwrap_or(false),
|
||||
context_size,
|
||||
provider,
|
||||
})
|
||||
})
|
||||
.collect();
|
||||
|
||||
log::info!(
|
||||
"[litellm] Fetched {} model(s) from {}",
|
||||
models.len(),
|
||||
url
|
||||
);
|
||||
models
|
||||
},
|
||||
|_view, models, ctx| {
|
||||
if !models.is_empty() {
|
||||
AISettings::handle(ctx).update(ctx, |settings, ctx| {
|
||||
let _ = settings.openai_models.set_value(models, ctx);
|
||||
});
|
||||
}
|
||||
ctx.notify();
|
||||
},
|
||||
);
|
||||
}
|
||||
|
||||
fn build_page(
|
||||
subpage: Option<AISubpage>,
|
||||
ctx: &mut ViewContext<Self>,
|
||||
@@ -1505,6 +1632,12 @@ impl AISettingsPageView {
|
||||
let title: Option<&str> = None;
|
||||
return (PageType::new_uncategorized(widgets, title), None);
|
||||
}
|
||||
Some(AISubpage::OpenAI) => {
|
||||
let widget = OpenAISettingsWidget::new(ctx);
|
||||
widgets.push(Box::new(widget));
|
||||
let title: Option<&str> = None;
|
||||
return (PageType::new_uncategorized(widgets, title), None);
|
||||
}
|
||||
}
|
||||
|
||||
// Subpage widgets render their own subheader-sized titles internally,
|
||||
@@ -2088,6 +2221,8 @@ pub enum AISettingsPageAction {
|
||||
SetBedrockAuthMethod(BedrockAuthMethod),
|
||||
SetBedrockProfile(String),
|
||||
ToggleBedrockCrossRegionInference,
|
||||
ToggleOpenAIEnabled,
|
||||
FetchOpenAIModels,
|
||||
ToggleFileBasedMcp,
|
||||
ToggleIncludeAgentCommandsInHistory,
|
||||
ToggleAgentAttribution,
|
||||
@@ -2756,6 +2891,16 @@ impl TypedActionView for AISettingsPageView {
|
||||
});
|
||||
ctx.notify();
|
||||
}
|
||||
AISettingsPageAction::ToggleOpenAIEnabled => {
|
||||
AISettings::handle(ctx).update(ctx, |settings, ctx| {
|
||||
report_if_error!(settings.openai_enabled.toggle_and_save_value(ctx));
|
||||
});
|
||||
ctx.notify();
|
||||
}
|
||||
AISettingsPageAction::FetchOpenAIModels => {
|
||||
// Trigger a fetch of models from the LiteLLM endpoint
|
||||
self.fetch_litellm_models(ctx);
|
||||
}
|
||||
AISettingsPageAction::ToggleFileBasedMcp => {
|
||||
AISettings::handle(ctx).update(ctx, |settings, ctx| {
|
||||
report_if_error!(settings.file_based_mcp_enabled.toggle_and_save_value(ctx));
|
||||
@@ -3731,7 +3876,6 @@ impl SettingsWidget for ActiveAIWidget {
|
||||
.finish(),
|
||||
);
|
||||
|
||||
|
||||
if self.is_next_command_toggleable(app) {
|
||||
column.add_child(self.render_next_command_section(view, app));
|
||||
}
|
||||
@@ -6231,7 +6375,11 @@ impl SettingsWidget for BedrockSettingsWidget {
|
||||
let description = format!(
|
||||
"{} model{} configured via settings.toml.",
|
||||
configured_models.len(),
|
||||
if configured_models.len() == 1 { "" } else { "s" }
|
||||
if configured_models.len() == 1 {
|
||||
""
|
||||
} else {
|
||||
"s"
|
||||
}
|
||||
);
|
||||
column.add_child(render_ai_setting_description(description, is_enabled, app));
|
||||
} else {
|
||||
@@ -6246,6 +6394,274 @@ impl SettingsWidget for BedrockSettingsWidget {
|
||||
}
|
||||
}
|
||||
|
||||
struct OpenAISettingsWidget {
|
||||
enabled_toggle: SwitchStateHandle,
|
||||
base_url_editor: ViewHandle<EditorView>,
|
||||
api_key_editor: ViewHandle<EditorView>,
|
||||
fetch_button: MouseStateHandle,
|
||||
}
|
||||
|
||||
impl OpenAISettingsWidget {
|
||||
fn new(ctx: &mut ViewContext<<Self as SettingsWidget>::View>) -> Self {
|
||||
let ai_settings = AISettings::as_ref(ctx);
|
||||
|
||||
let base_url_val = ai_settings.openai_base_url.value().clone();
|
||||
let api_key_val = ai_settings.openai_api_key.value().clone();
|
||||
|
||||
let base_url_editor = ctx.add_typed_action_view(move |ctx| {
|
||||
let appearance = Appearance::as_ref(ctx);
|
||||
let options = SingleLineEditorOptions {
|
||||
is_password: false,
|
||||
text: TextOptions {
|
||||
font_size_override: Some(appearance.ui_font_size()),
|
||||
font_family_override: Some(appearance.monospace_font_family()),
|
||||
text_colors_override: Some(TextColors {
|
||||
default_color: appearance.theme().active_ui_text_color(),
|
||||
disabled_color: appearance.theme().disabled_ui_text_color(),
|
||||
hint_color: appearance.theme().disabled_ui_text_color(),
|
||||
}),
|
||||
..Default::default()
|
||||
},
|
||||
..Default::default()
|
||||
};
|
||||
let mut editor = EditorView::single_line(options, ctx);
|
||||
editor.set_placeholder_text("http://localhost:4000/v1", ctx);
|
||||
editor.set_buffer_text(&base_url_val, ctx);
|
||||
editor
|
||||
});
|
||||
ctx.subscribe_to_view(&base_url_editor, |_, editor, event, ctx| {
|
||||
if matches!(event, EditorEvent::Blurred | EditorEvent::Enter) {
|
||||
let value = editor.as_ref(ctx).buffer_text(ctx);
|
||||
AISettings::handle(ctx).update(ctx, |settings, ctx| {
|
||||
let _ = settings.openai_base_url.set_value(value, ctx);
|
||||
});
|
||||
}
|
||||
});
|
||||
|
||||
let api_key_editor = ctx.add_typed_action_view(move |ctx| {
|
||||
let appearance = Appearance::as_ref(ctx);
|
||||
let options = SingleLineEditorOptions {
|
||||
is_password: true,
|
||||
text: TextOptions {
|
||||
font_size_override: Some(appearance.ui_font_size()),
|
||||
font_family_override: Some(appearance.monospace_font_family()),
|
||||
text_colors_override: Some(TextColors {
|
||||
default_color: appearance.theme().active_ui_text_color(),
|
||||
disabled_color: appearance.theme().disabled_ui_text_color(),
|
||||
hint_color: appearance.theme().disabled_ui_text_color(),
|
||||
}),
|
||||
..Default::default()
|
||||
},
|
||||
..Default::default()
|
||||
};
|
||||
let mut editor = EditorView::single_line(options, ctx);
|
||||
editor.set_placeholder_text("sk-... (optional)", ctx);
|
||||
editor.set_buffer_text(&api_key_val, ctx);
|
||||
editor
|
||||
});
|
||||
ctx.subscribe_to_view(&api_key_editor, |_, editor, event, ctx| {
|
||||
if matches!(event, EditorEvent::Blurred | EditorEvent::Enter) {
|
||||
let value = editor.as_ref(ctx).buffer_text(ctx);
|
||||
AISettings::handle(ctx).update(ctx, |settings, ctx| {
|
||||
let _ = settings.openai_api_key.set_value(value, ctx);
|
||||
});
|
||||
}
|
||||
});
|
||||
|
||||
let base_url_editor_clone = base_url_editor.clone();
|
||||
let api_key_editor_clone = api_key_editor.clone();
|
||||
ctx.subscribe_to_model(&AISettings::handle(ctx), move |_, _, event, ctx| {
|
||||
if matches!(event, AISettingsChangedEvent::OpenAIEnabled { .. }) {
|
||||
let is_enabled = *AISettings::as_ref(ctx).openai_enabled.value();
|
||||
AISettingsPageView::update_editor_interaction_state(
|
||||
base_url_editor_clone.clone(),
|
||||
is_enabled,
|
||||
ctx,
|
||||
);
|
||||
AISettingsPageView::update_editor_interaction_state(
|
||||
api_key_editor_clone.clone(),
|
||||
is_enabled,
|
||||
ctx,
|
||||
);
|
||||
ctx.notify();
|
||||
}
|
||||
});
|
||||
|
||||
Self {
|
||||
enabled_toggle: SwitchStateHandle::default(),
|
||||
base_url_editor,
|
||||
api_key_editor,
|
||||
fetch_button: MouseStateHandle::default(),
|
||||
}
|
||||
}
|
||||
|
||||
fn render_input(
|
||||
appearance: &Appearance,
|
||||
label: &'static str,
|
||||
editor: ViewHandle<EditorView>,
|
||||
is_enabled: bool,
|
||||
app: &AppContext,
|
||||
) -> Box<dyn Element> {
|
||||
let padding = Some(Coords {
|
||||
top: 10.,
|
||||
bottom: 10.,
|
||||
left: 16.,
|
||||
right: 16.,
|
||||
});
|
||||
let editor_style = UiComponentStyles {
|
||||
padding,
|
||||
background: Some(appearance.theme().surface_2().into()),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
let label = Text::new_inline(label, appearance.ui_font_family(), CONTENT_FONT_SIZE)
|
||||
.with_color(styles::header_font_color(is_enabled, app).into())
|
||||
.finish();
|
||||
|
||||
let input = appearance
|
||||
.ui_builder()
|
||||
.text_input(editor)
|
||||
.with_style(editor_style)
|
||||
.build()
|
||||
.finish();
|
||||
|
||||
Flex::column()
|
||||
.with_spacing(8.)
|
||||
.with_child(label)
|
||||
.with_child(input)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
impl SettingsWidget for OpenAISettingsWidget {
|
||||
type View = AISettingsPageView;
|
||||
|
||||
fn search_terms(&self) -> &str {
|
||||
"openai litellm custom provider endpoint api key models"
|
||||
}
|
||||
|
||||
fn should_render(&self, _app: &AppContext) -> bool {
|
||||
true
|
||||
}
|
||||
|
||||
fn render(
|
||||
&self,
|
||||
_view: &Self::View,
|
||||
appearance: &Appearance,
|
||||
app: &AppContext,
|
||||
) -> Box<dyn Element> {
|
||||
let ai_settings = AISettings::as_ref(app);
|
||||
let is_enabled = *ai_settings.openai_enabled.value();
|
||||
|
||||
let mut column = Flex::column().with_spacing(16.);
|
||||
|
||||
column.add_child(render_ai_setting_toggle::<OpenAIEnabled>(
|
||||
"Enable OpenAI-Compatible Provider",
|
||||
AISettingsPageAction::ToggleOpenAIEnabled,
|
||||
is_enabled,
|
||||
true,
|
||||
self.enabled_toggle.clone(),
|
||||
&RefCell::new(HashMap::new()),
|
||||
app,
|
||||
));
|
||||
column.add_child(render_ai_setting_description(
|
||||
"Route AI requests through an OpenAI-compatible endpoint (e.g. LiteLLM proxy).",
|
||||
true,
|
||||
app,
|
||||
));
|
||||
|
||||
column.add_child(render_separator(appearance));
|
||||
|
||||
column.add_child(Self::render_input(
|
||||
appearance,
|
||||
"Base URL",
|
||||
self.base_url_editor.clone(),
|
||||
is_enabled,
|
||||
app,
|
||||
));
|
||||
column.add_child(render_ai_setting_description(
|
||||
"The OpenAI-compatible API base URL (e.g. http://localhost:4000/v1).",
|
||||
is_enabled,
|
||||
app,
|
||||
));
|
||||
|
||||
column.add_child(Self::render_input(
|
||||
appearance,
|
||||
"API Key",
|
||||
self.api_key_editor.clone(),
|
||||
is_enabled,
|
||||
app,
|
||||
));
|
||||
column.add_child(render_ai_setting_description(
|
||||
"Optional. Leave empty if the proxy handles authentication.",
|
||||
is_enabled,
|
||||
app,
|
||||
));
|
||||
|
||||
column.add_child(render_separator(appearance));
|
||||
|
||||
// Fetch models button
|
||||
let fetch_button = appearance
|
||||
.ui_builder()
|
||||
.button(ButtonVariant::Secondary, self.fetch_button.clone())
|
||||
.with_text_label("Fetch Models from Endpoint".to_owned())
|
||||
.build()
|
||||
.on_click(move |ctx, _, _| {
|
||||
ctx.dispatch_typed_action(AISettingsPageAction::FetchOpenAIModels);
|
||||
})
|
||||
.finish();
|
||||
column.add_child(fetch_button);
|
||||
column.add_child(render_ai_setting_description(
|
||||
"Queries the /models endpoint and populates the model list with available models and their context window sizes.",
|
||||
is_enabled,
|
||||
app,
|
||||
));
|
||||
|
||||
column.add_child(render_separator(appearance));
|
||||
|
||||
// Show configured models count
|
||||
let configured_models: Vec<_> = ai_settings.openai_models.value().clone();
|
||||
if !configured_models.is_empty() {
|
||||
let description = format!(
|
||||
"{} model{} configured via settings.toml.",
|
||||
configured_models.len(),
|
||||
if configured_models.len() == 1 {
|
||||
""
|
||||
} else {
|
||||
"s"
|
||||
}
|
||||
);
|
||||
column.add_child(render_ai_setting_description(description, is_enabled, app));
|
||||
|
||||
// Show first few model names
|
||||
let preview: String = configured_models
|
||||
.iter()
|
||||
.take(5)
|
||||
.map(|m| m.display_name.as_str())
|
||||
.collect::<Vec<_>>()
|
||||
.join(", ");
|
||||
let suffix = if configured_models.len() > 5 {
|
||||
format!(" (+{} more)", configured_models.len() - 5)
|
||||
} else {
|
||||
String::new()
|
||||
};
|
||||
column.add_child(render_ai_setting_description(
|
||||
format!("Models: {preview}{suffix}"),
|
||||
is_enabled,
|
||||
app,
|
||||
));
|
||||
} else {
|
||||
column.add_child(render_ai_setting_description(
|
||||
"No models configured. Use 'Fetch Models' or add them to ~/.galaxy/settings.toml under [ai.openai].",
|
||||
is_enabled,
|
||||
app,
|
||||
));
|
||||
}
|
||||
|
||||
column.finish()
|
||||
}
|
||||
}
|
||||
|
||||
mod styles {
|
||||
use galaxy_core::ui::{appearance::Appearance, theme::Fill};
|
||||
use galaxyui::{AppContext, SingletonEntity};
|
||||
|
||||
@@ -202,6 +202,7 @@ pub enum SettingsSection {
|
||||
Knowledge,
|
||||
ThirdPartyCLIAgents,
|
||||
Bedrock,
|
||||
OpenAI,
|
||||
/// Internal backing-page identifier for CodeSettingsPageView. Multiple subpages
|
||||
/// (CodeIndexing, EditorAndCodeReview) share this single backing page,
|
||||
/// so this variant is needed as the key in `settings_pages`.
|
||||
@@ -239,6 +240,7 @@ impl Display for SettingsSection {
|
||||
SettingsSection::Knowledge => write!(f, "Knowledge"),
|
||||
SettingsSection::ThirdPartyCLIAgents => write!(f, "Third party CLI agents"),
|
||||
SettingsSection::Bedrock => write!(f, "AWS Bedrock"),
|
||||
SettingsSection::OpenAI => write!(f, "OpenAI / LiteLLM"),
|
||||
SettingsSection::Warpify => write!(f, "Wormhole"),
|
||||
SettingsSection::CodeIndexing => write!(f, "Indexing and projects"),
|
||||
SettingsSection::EditorAndCodeReview => write!(f, "Editor and Code Review"),
|
||||
@@ -263,6 +265,7 @@ impl SettingsSection {
|
||||
| Self::Knowledge
|
||||
| Self::ThirdPartyCLIAgents
|
||||
| Self::Bedrock
|
||||
| Self::OpenAI
|
||||
)
|
||||
}
|
||||
|
||||
@@ -294,6 +297,7 @@ impl SettingsSection {
|
||||
Self::Knowledge,
|
||||
Self::ThirdPartyCLIAgents,
|
||||
Self::Bedrock,
|
||||
Self::OpenAI,
|
||||
]
|
||||
}
|
||||
|
||||
|
||||
@@ -31,9 +31,7 @@ use crate::{
|
||||
model_menu_items::{available_model_menu_items, has_reasoning_variants, is_auto},
|
||||
profiles::{AIExecutionProfilesModel, AIExecutionProfilesModelEvent, ClientProfileId},
|
||||
},
|
||||
llms::{
|
||||
dedupe_model_display_names, LLMId, LLMInfo, LLMPreferences, LLMPreferencesEvent,
|
||||
},
|
||||
llms::{dedupe_model_display_names, LLMId, LLMInfo, LLMPreferences, LLMPreferencesEvent},
|
||||
},
|
||||
appearance::Appearance,
|
||||
cloud_object::model::generic_string_model::StringModel,
|
||||
@@ -1476,7 +1474,6 @@ impl ProfileModelSelector {
|
||||
container.finish()
|
||||
}
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
impl TypedActionView for ProfileModelSelector {
|
||||
@@ -1614,7 +1611,6 @@ impl View for ProfileModelSelector {
|
||||
let model_menu = ChildView::new(&self.model_dropdown).finish();
|
||||
let positioning = self.get_menu_positioning(app, false);
|
||||
stack.add_positioned_overlay_child(model_menu, positioning);
|
||||
|
||||
}
|
||||
|
||||
let is_udi_enabled =
|
||||
|
||||
@@ -55,7 +55,8 @@ impl InputClassifier for HeuristicClassifier {
|
||||
return InputType::Shell;
|
||||
}
|
||||
|
||||
let result = self.classify_input(input, context)
|
||||
let result = self
|
||||
.classify_input(input, context)
|
||||
.await
|
||||
.map(|result| result.to_input_type())
|
||||
.unwrap_or(context.current_input_type);
|
||||
|
||||
@@ -2,8 +2,8 @@ mod generation;
|
||||
mod model_loader;
|
||||
|
||||
use std::path::PathBuf;
|
||||
use std::sync::atomic::{AtomicBool, Ordering};
|
||||
use std::sync::Arc;
|
||||
use std::sync::atomic::{AtomicBool, Ordering};
|
||||
|
||||
use anyhow::{Context as _, Result};
|
||||
use candle_core::{DType, Tensor};
|
||||
|
||||
@@ -91,7 +91,9 @@ impl InputClassificationTask {
|
||||
temperature: 0.0,
|
||||
top_p: 1.0,
|
||||
};
|
||||
let output = engine.generate_cancellable(&prompt, &config, cancel).await?;
|
||||
let output = engine
|
||||
.generate_cancellable(&prompt, &config, cancel)
|
||||
.await?;
|
||||
Ok(Self::parse_output(&output))
|
||||
}
|
||||
}
|
||||
|
||||
@@ -179,7 +179,11 @@ async fn scenario_tab_naming_development_contexts() {
|
||||
},
|
||||
TabCase {
|
||||
label: "Python data science",
|
||||
commands: vec!["jupyter notebook", "pip install pandas", "python analysis.py"],
|
||||
commands: vec![
|
||||
"jupyter notebook",
|
||||
"pip install pandas",
|
||||
"python analysis.py",
|
||||
],
|
||||
cwd: "/home/user/research/data-pipeline",
|
||||
},
|
||||
TabCase {
|
||||
|
||||
@@ -270,7 +270,10 @@ async fn test_classify_multiple_sequential_calls() {
|
||||
("git status", InputCategory::Shell),
|
||||
("what does this error mean?", InputCategory::AgentPrompt),
|
||||
("npm install express", InputCategory::Shell),
|
||||
("refactor this to use async/await", InputCategory::AgentPrompt),
|
||||
(
|
||||
"refactor this to use async/await",
|
||||
InputCategory::AgentPrompt,
|
||||
),
|
||||
];
|
||||
|
||||
for (input, expected_category) in inputs {
|
||||
@@ -286,7 +289,10 @@ async fn test_classify_multiple_sequential_calls() {
|
||||
.await
|
||||
.expect("classification failed");
|
||||
|
||||
println!("\"{input}\" -> {:?} (confidence: {:.2})", result.category, result.confidence);
|
||||
println!(
|
||||
"\"{input}\" -> {:?} (confidence: {:.2})",
|
||||
result.category, result.confidence
|
||||
);
|
||||
assert_eq!(
|
||||
result.category, expected_category,
|
||||
"expected {expected_category:?} for \"{input}\", got {:?}",
|
||||
|
||||
Reference in New Issue
Block a user