Compare commits

...
Author SHA1 Message Date
Nick Pismenkov a320a64086 fix linter 2026-03-06 18:43:26 -08:00
Nick Pismenkov 13813cbb18 Merge branch 'main' into feat/images 2026-03-06 18:41:54 -08:00
Nick Pismenkov 833738bc85 review fixes 2026-03-06 18:41:06 -08:00
Nick Pismenkov f4d290f5ed feat: Support processing images by IronClaw 2026-03-06 17:55:41 -08:00
29 changed files with 1639 additions and 67 deletions
+13
View File
@@ -0,0 +1,13 @@
{
"permissions": {
"allow": [
"Bash(cargo check:*)",
"Bash(cargo clippy:*)",
"Bash(cargo test:*)",
"Bash(cargo fmt:*)",
"Bash(grep:*)",
"Bash(env:*)",
"Skill(ship)"
]
}
}
+32
View File
@@ -17,6 +17,16 @@ use crate::error::Error;
use crate::llm::{ChatMessage, Reasoning, ReasoningContext, RespondResult};
use crate::tools::redact_params;
/// Represents image generation sentinel data in tool output.
#[derive(serde::Deserialize)]
struct ImageGeneratedSentinel<'a> {
#[serde(rename = "type")]
ty: &'a str,
data: &'a str,
media_type: &'a str,
path: &'a str,
}
/// Result of the agentic loop execution.
pub(super) enum AgenticLoopResult {
/// Completed with a response.
@@ -640,6 +650,28 @@ impl Agent {
&message.metadata,
)
.await;
// Check for image_generated sentinel and emit SSE event
if let Ok(sentinel) =
serde_json::from_str::<ImageGeneratedSentinel>(output)
&& sentinel.ty == "image_generated"
{
let data_url = format!(
"data:{};base64,{}",
sentinel.media_type, sentinel.data
);
let _ = self
.channels
.send_status(
&message.channel,
StatusUpdate::ImageGenerated {
data_url,
path: sentinel.path.to_string(),
},
&message.metadata,
)
.await;
}
}
// Record result in thread
+29 -2
View File
@@ -16,7 +16,7 @@ use chrono::{DateTime, Utc};
use serde::{Deserialize, Serialize};
use uuid::Uuid;
use crate::llm::{ChatMessage, ToolCall};
use crate::llm::{ChatMessage, ImageAttachment, ToolCall};
/// A session containing one or more threads.
#[derive(Debug, Clone, Serialize, Deserialize)]
@@ -250,6 +250,22 @@ impl Thread {
&mut self.turns[turn_number]
}
/// Start a new turn with user input and image attachments.
pub fn start_turn_with_images(
&mut self,
user_input: impl Into<String>,
images: Vec<ImageAttachment>,
) -> &mut Turn {
let turn_number = self.turns.len();
let mut turn = Turn::new(turn_number, user_input);
turn.images = images;
self.turns.push(turn);
self.state = ThreadState::Processing;
self.updated_at = Utc::now();
// turn_number was len() before push, so it's a valid index after push
&mut self.turns[turn_number]
}
/// Complete the current turn with a response.
pub fn complete_turn(&mut self, response: impl Into<String>) {
if let Some(turn) = self.turns.last_mut() {
@@ -320,7 +336,14 @@ impl Thread {
pub fn messages(&self) -> Vec<ChatMessage> {
let mut messages = Vec::new();
for turn in &self.turns {
messages.push(ChatMessage::user(&turn.user_input));
if turn.images.is_empty() {
messages.push(ChatMessage::user(&turn.user_input));
} else {
messages.push(ChatMessage::user_with_images(
&turn.user_input,
turn.images.clone(),
));
}
if let Some(ref response) = turn.response {
messages.push(ChatMessage::assistant(response));
}
@@ -407,6 +430,9 @@ pub struct Turn {
pub completed_at: Option<DateTime<Utc>>,
/// Error message (if failed).
pub error: Option<String>,
/// Images attached to this turn's user input.
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub images: Vec<ImageAttachment>,
}
impl Turn {
@@ -421,6 +447,7 @@ impl Turn {
started_at: Utc::now(),
completed_at: None,
error: None,
images: Vec::new(),
}
}
+5 -1
View File
@@ -264,7 +264,11 @@ impl Agent {
.threads
.get_mut(&thread_id)
.ok_or_else(|| Error::from(crate::error::JobError::NotFound { id: thread_id }))?;
thread.start_turn(content);
if message.images.is_empty() {
thread.start_turn(content);
} else {
thread.start_turn_with_images(content, message.images.clone());
}
thread.messages()
};
+40
View File
@@ -376,6 +376,46 @@ impl AppBuilder {
}
let ws = Arc::new(ws);
tools.register_memory_tools(Arc::clone(&ws));
// Register image tools if image generation models are available
match llm.list_models().await {
Ok(models) => {
if let Some(image_model) =
crate::llm::image_models::suggest_image_model(&models)
{
tools.register_image_tools(self.config.llm.nearai.clone(), Arc::clone(&ws));
tracing::info!(
"Image generation tools registered (model: {})",
image_model
);
} else {
tracing::debug!(
"No image generation models detected in available models: {:?}",
models
);
}
// Register vision analysis tool if vision models are available
if let Some(vision_model) =
crate::llm::vision_models::suggest_vision_model(&models)
{
tools.register_vision_tools(Arc::clone(&ws));
tracing::info!(
"Image analysis tool registered (vision model: {})",
vision_model
);
} else {
tracing::debug!("No vision-capable models detected in available models");
}
}
Err(e) => {
tracing::warn!(
"Failed to list available models for image tool registration: {}",
e
);
}
}
Some(ws)
} else {
None
+12
View File
@@ -9,6 +9,7 @@ use futures::Stream;
use uuid::Uuid;
use crate::error::ChannelError;
use crate::llm::ImageAttachment;
/// A message received from an external channel.
#[derive(Debug, Clone)]
@@ -29,6 +30,8 @@ pub struct IncomingMessage {
pub received_at: DateTime<Utc>,
/// Channel-specific metadata.
pub metadata: serde_json::Value,
/// Images attached to this message.
pub images: Vec<ImageAttachment>,
}
impl IncomingMessage {
@@ -47,6 +50,7 @@ impl IncomingMessage {
thread_id: None,
received_at: Utc::now(),
metadata: serde_json::Value::Null,
images: Vec::new(),
}
}
@@ -67,6 +71,12 @@ impl IncomingMessage {
self.user_name = Some(name.into());
self
}
/// Attach image attachments.
pub fn with_images(mut self, images: Vec<ImageAttachment>) -> Self {
self.images = images;
self
}
}
/// Stream of incoming messages.
@@ -163,6 +173,8 @@ pub enum StatusUpdate {
success: bool,
message: String,
},
/// An image was generated or edited by a tool.
ImageGenerated { data_url: String, path: String },
}
impl StatusUpdate {
+3
View File
@@ -585,6 +585,9 @@ impl Channel for ReplChannel {
eprintln!("\x1b[31m {extension_name}: {message}\x1b[0m");
}
}
StatusUpdate::ImageGenerated { path, .. } => {
eprintln!(" \x1b[36m[image]\x1b[0m {path}");
}
}
Ok(())
}
+5
View File
@@ -2591,6 +2591,11 @@ fn status_to_wit(status: &StatusUpdate, metadata: &serde_json::Value) -> wit_cha
),
metadata_json,
},
StatusUpdate::ImageGenerated { path, .. } => wit_channel::StatusUpdate {
status: wit_channel::StatusType::Status,
message: format!("Image generated: {}", path),
metadata_json,
},
}
}
+13
View File
@@ -369,6 +369,19 @@ impl Channel for GatewayChannel {
success,
message,
},
StatusUpdate::ImageGenerated { data_url, path } => {
tracing::debug!(
path = %path,
data_url_len = data_url.len(),
thread_id = ?thread_id,
"Broadcasting ImageGenerated SSE event"
);
SseEvent::ImageGenerated {
data_url,
path,
thread_id,
}
}
};
self.state.sse.broadcast(event);
+1
View File
@@ -247,6 +247,7 @@ pub fn convert_messages(messages: &[OpenAiMessage]) -> Result<Vec<ChatMessage>,
tool_call_id: None,
name: m.name.clone(),
tool_calls: None,
images: Vec::new(),
}),
}
})
+31 -12
View File
@@ -43,6 +43,7 @@ use crate::channels::web::types::*;
use crate::channels::web::util::{build_turns_from_db_messages, truncate_preview};
use crate::db::Database;
use crate::extensions::ExtensionManager;
use crate::llm::ImageAttachment;
use crate::orchestrator::job_manager::ContainerJobManager;
use crate::tools::ToolRegistry;
use crate::workspace::Workspace;
@@ -626,6 +627,17 @@ async fn chat_send_handler(
msg = msg.with_metadata(serde_json::json!({"thread_id": thread_id}));
}
// Convert image data to ImageAttachment
let images: Vec<ImageAttachment> = req
.images
.into_iter()
.map(|img| ImageAttachment {
media_type: img.media_type,
data: img.data,
})
.collect();
msg = msg.with_images(images);
let msg_id = msg.id;
tracing::debug!(
"[chat_send_handler] Created message id={}, content={:?}",
@@ -951,18 +963,25 @@ async fn chat_history_handler(
tool_calls: t
.tool_calls
.iter()
.map(|tc| ToolCallInfo {
name: tc.name.clone(),
has_result: tc.result.is_some(),
has_error: tc.error.is_some(),
result_preview: tc.result.as_ref().map(|r| {
let s = match r {
serde_json::Value::String(s) => s.clone(),
other => other.to_string(),
};
truncate_preview(&s, 500)
}),
error: tc.error.clone(),
.map(|tc| {
// Image tools need full results (large base64 data), don't truncate
let limit = match tc.name.as_str() {
"image_generate" | "image_edit" | "image_analyze" => usize::MAX,
_ => 500,
};
ToolCallInfo {
name: tc.name.clone(),
has_result: tc.result.is_some(),
has_error: tc.error.is_some(),
result_preview: tc.result.as_ref().map(|r| {
let s = match r {
serde_json::Value::String(s) => s.clone(),
other => other.to_string(),
};
truncate_preview(&s, limit)
}),
error: tc.error.clone(),
}
})
.collect(),
})
+5
View File
@@ -55,6 +55,10 @@ impl SseManager {
/// Broadcast an event to all connected clients.
pub fn broadcast(&self, event: SseEvent) {
// Log image events for debugging
if matches!(&event, SseEvent::ImageGenerated { .. }) {
tracing::debug!("Broadcasting image_generated SSE event to all connected clients");
}
// Ignore send errors (no receivers is fine)
let _ = self.tx.send(event);
}
@@ -143,6 +147,7 @@ impl SseManager {
SseEvent::JobResult { .. } => "job_result",
SseEvent::Heartbeat => "heartbeat",
SseEvent::ExtensionStatus { .. } => "extension_status",
SseEvent::ImageGenerated { .. } => "image_generated",
};
Ok(Event::default().event(event_type).data(data))
});
+148 -6
View File
@@ -41,6 +41,9 @@ const SLASH_COMMANDS = [
let _slashSelected = -1;
let _slashMatches = [];
// --- Image Attachments ---
let stagedImages = []; // Array of { media_type, data, previewUrl }
// --- Tool Activity State ---
let _activeGroup = null;
let _activeToolCards = {};
@@ -113,6 +116,78 @@ document.getElementById('token-input').addEventListener('keydown', (e) => {
}
})();
// --- Image Attachment Handlers ---
// Handle file picker selection
document.getElementById('image-input').addEventListener('change', (e) => {
const files = e.target.files;
if (files) {
for (let file of files) {
if (file.type.startsWith('image/')) {
const reader = new FileReader();
reader.onload = (evt) => {
const base64Data = evt.target.result.split(',')[1]; // Remove data URL prefix
stagedImages.push({
media_type: file.type,
data: base64Data,
previewUrl: evt.target.result,
});
renderImagePreviews();
};
reader.readAsDataURL(file);
}
}
}
// Reset file input so the same file can be selected again
e.target.value = '';
});
// Handle paste event
document.getElementById('chat-input').addEventListener('paste', (e) => {
const items = e.clipboardData.items;
for (let item of items) {
if (item.type.startsWith('image/')) {
e.preventDefault();
const file = item.getAsFile();
const reader = new FileReader();
reader.onload = (evt) => {
const base64Data = evt.target.result.split(',')[1];
stagedImages.push({
media_type: item.type,
data: base64Data,
previewUrl: evt.target.result,
});
renderImagePreviews();
};
reader.readAsDataURL(file);
}
}
});
function renderImagePreviews() {
const strip = document.getElementById('image-preview-strip');
if (stagedImages.length === 0) {
strip.style.display = 'none';
return;
}
strip.style.display = 'flex';
strip.innerHTML = '';
stagedImages.forEach((img, idx) => {
const container = document.createElement('div');
container.className = 'image-preview';
container.innerHTML = `
<img src="${img.previewUrl}" alt="Preview">
<button class="image-preview-remove" onclick="removeImage(${idx})" title="Remove">×</button>
`;
strip.appendChild(container);
});
}
function removeImage(idx) {
stagedImages.splice(idx, 1);
renderImagePreviews();
}
// --- API helper ---
function apiFetch(path, options) {
@@ -315,6 +390,17 @@ function connectSSE() {
setToolCardOutput(data.name, data.preview);
});
eventSource.addEventListener('image_generated', (e) => {
const data = JSON.parse(e.data);
console.log('Received image_generated event:', { thread_id: data.thread_id, path: data.path, data_url_len: data.data_url ? data.data_url.length : 0 });
if (!isCurrentThread(data.thread_id)) {
console.log('Image event ignored: not current thread', { currentThreadId, eventThreadId: data.thread_id });
return;
}
console.log('Adding generated image to chat');
addGeneratedImage(data.data_url, data.path);
});
eventSource.addEventListener('stream_chunk', (e) => {
const data = JSON.parse(e.data);
if (!isCurrentThread(data.thread_id)) return;
@@ -430,19 +516,28 @@ function sendMessage() {
return;
}
const content = input.value.trim();
if (!content) return;
if (!content && stagedImages.length === 0) return;
addMessage('user', content);
input.value = '';
autoResizeTextarea(input);
input.focus();
const images = stagedImages.map(img => ({
media_type: img.media_type,
data: img.data,
}));
apiFetch('/api/chat/send', {
method: 'POST',
body: { content, thread_id: currentThreadId || undefined },
body: { content, thread_id: currentThreadId || undefined, images },
}).catch((err) => {
addMessage('system', 'Failed to send: ' + err.message);
});
// Clear staged images after sending
stagedImages = [];
renderImagePreviews();
}
function enableChatInput() {
@@ -858,6 +953,30 @@ function finalizeActivityGroup() {
_activeToolCards = {};
}
function addGeneratedImage(dataUrl, path) {
const container = document.getElementById('chat-messages');
console.log('addGeneratedImage called', { dataUrl_len: dataUrl ? dataUrl.length : 0, path });
const card = document.createElement('div');
card.className = 'generated-image-card';
const img = document.createElement('img');
img.src = dataUrl;
img.alt = 'Generated image';
img.className = 'generated-image';
img.onerror = () => console.error('Failed to load image from data URL:', dataUrl.substring(0, 100));
img.onload = () => console.log('Image loaded successfully from data URL');
const pathLabel = document.createElement('div');
pathLabel.className = 'generated-image-path';
pathLabel.textContent = 'Saved to: ' + path;
card.appendChild(img);
card.appendChild(pathLabel);
container.appendChild(card);
console.log('Image card appended to DOM');
container.scrollTop = container.scrollHeight;
}
function showApproval(data) {
const container = document.getElementById('chat-messages');
const card = document.createElement('div');
@@ -1223,10 +1342,33 @@ function createToolCallsSummaryElement(toolCalls) {
item.appendChild(nameSpan);
if (tc.result_preview) {
const preview = document.createElement('div');
preview.className = 'tool-call-preview';
preview.textContent = tc.result_preview;
item.appendChild(preview);
// Check if this is an image result
try {
const parsed = JSON.parse(tc.result_preview);
if (parsed.type === 'image_generated' && parsed.data && parsed.media_type) {
const dataUrl = `data:${parsed.media_type};base64,${parsed.data}`;
const imgDiv = document.createElement('div');
imgDiv.className = 'generated-image-card';
const img = document.createElement('img');
img.src = dataUrl;
img.alt = 'Generated image';
img.className = 'generated-image';
imgDiv.appendChild(img);
item.appendChild(imgDiv);
} else {
// Regular text result
const preview = document.createElement('div');
preview.className = 'tool-call-preview';
preview.textContent = tc.result_preview;
item.appendChild(preview);
}
} catch {
// Not JSON, display as text
const preview = document.createElement('div');
preview.className = 'tool-call-preview';
preview.textContent = tc.result_preview;
item.appendChild(preview);
}
}
if (tc.error) {
const errDiv = document.createElement('div');
+3
View File
@@ -129,7 +129,10 @@
<div class="chat-container">
<div class="chat-messages" id="chat-messages"></div>
<div id="slash-autocomplete" class="slash-autocomplete" style="display:none"></div>
<div class="image-preview-strip" id="image-preview-strip" style="display:none;"></div>
<div class="chat-input">
<input type="file" id="image-input" accept="image/*" multiple style="display:none">
<button id="attach-btn" class="attach-btn" title="Attach image" onclick="document.getElementById('image-input').click()">📎</button>
<textarea id="chat-input" placeholder="Message or / for commands..." rows="1"></textarea>
<button id="send-btn" onclick="sendMessage()">Send</button>
</div>
+98
View File
@@ -1093,6 +1093,37 @@ body {
font-style: italic;
}
/* Generated image card */
.generated-image-card {
align-self: flex-start;
width: 50%;
background: var(--bg-secondary);
border: 1px solid var(--border);
border-radius: var(--radius-lg);
overflow: hidden;
margin: 8px 0;
box-shadow: var(--shadow);
display: flex;
flex-direction: column;
flex-shrink: 0;
}
.generated-image {
display: block;
width: 100%;
border-radius: var(--radius-lg);
object-fit: contain;
}
.generated-image-path {
padding: 8px 12px;
font-size: 12px;
color: var(--text-secondary);
background: var(--bg-tertiary);
border-top: 1px solid var(--border);
word-break: break-all;
}
/* Tool calls summary (persisted between user/assistant messages) */
.tool-calls-summary {
background: var(--bg-secondary);
@@ -1325,6 +1356,73 @@ body {
cursor: not-allowed;
}
.attach-btn {
padding: 8px 12px;
background: transparent;
color: var(--text-secondary);
border: 1px solid var(--border);
border-radius: var(--radius);
cursor: pointer;
font-size: 16px;
transition: all 0.2s;
}
.attach-btn:hover {
background: var(--bg);
color: var(--text);
border-color: var(--accent);
}
.image-preview-strip {
display: flex;
padding: 12px 16px 0 16px;
gap: 12px;
background: var(--bg-secondary);
overflow-x: auto;
border-top: 1px solid var(--border);
}
.image-preview {
position: relative;
width: 80px;
height: 80px;
flex-shrink: 0;
border-radius: var(--radius);
overflow: hidden;
background: var(--bg);
border: 1px solid var(--border);
}
.image-preview img {
width: 100%;
height: 100%;
object-fit: cover;
}
.image-preview-remove {
position: absolute;
top: -1px;
right: -1px;
width: 24px;
height: 24px;
padding: 0;
background: rgba(0, 0, 0, 0.6);
color: white;
border: none;
border-radius: 0;
font-size: 18px;
font-weight: bold;
cursor: pointer;
display: flex;
align-items: center;
justify-content: center;
transition: background 0.2s;
}
.image-preview-remove:hover {
background: rgba(0, 0, 0, 0.8);
}
/* Memory Tab */
.memory-container {
flex: 1;
+34 -2
View File
@@ -5,10 +5,18 @@ use uuid::Uuid;
// --- Chat ---
#[derive(Debug, Clone, Deserialize, Serialize)]
pub struct ImageData {
pub media_type: String,
pub data: String, // base64-encoded
}
#[derive(Debug, Deserialize)]
pub struct SendMessageRequest {
pub content: String,
pub thread_id: Option<String>,
#[serde(default)]
pub images: Vec<ImageData>,
}
#[derive(Debug, Serialize)]
@@ -225,6 +233,17 @@ pub enum SseEvent {
#[serde(skip_serializing_if = "Option::is_none")]
message: Option<String>,
},
/// An image was generated or edited.
#[serde(rename = "image_generated")]
ImageGenerated {
/// Base64 data URL: "data:image/png;base64,..."
data_url: String,
/// Workspace path where the image is saved.
path: String,
#[serde(skip_serializing_if = "Option::is_none")]
thread_id: Option<String>,
},
}
// --- Memory ---
@@ -606,6 +625,8 @@ pub enum WsClientMessage {
Message {
content: String,
thread_id: Option<String>,
#[serde(default)]
images: Vec<ImageData>,
},
/// Approve or deny a pending tool execution.
#[serde(rename = "approval")]
@@ -673,6 +694,7 @@ impl WsServerMessage {
SseEvent::JobStatus { .. } => "job_status",
SseEvent::JobResult { .. } => "job_result",
SseEvent::ExtensionStatus { .. } => "extension_status",
SseEvent::ImageGenerated { .. } => "image_generated",
};
let data = serde_json::to_value(event).unwrap_or(serde_json::Value::Null);
WsServerMessage::Event {
@@ -791,9 +813,14 @@ mod tests {
let json = r#"{"type":"message","content":"hello","thread_id":"t1"}"#;
let msg: WsClientMessage = serde_json::from_str(json).unwrap();
match msg {
WsClientMessage::Message { content, thread_id } => {
WsClientMessage::Message {
content,
thread_id,
images,
} => {
assert_eq!(content, "hello");
assert_eq!(thread_id.as_deref(), Some("t1"));
assert!(images.is_empty());
}
_ => panic!("Expected Message variant"),
}
@@ -804,9 +831,14 @@ mod tests {
let json = r#"{"type":"message","content":"hi"}"#;
let msg: WsClientMessage = serde_json::from_str(json).unwrap();
match msg {
WsClientMessage::Message { content, thread_id } => {
WsClientMessage::Message {
content,
thread_id,
images,
} => {
assert_eq!(content, "hi");
assert!(thread_id.is_none());
assert!(images.is_empty());
}
_ => panic!("Expected Message variant"),
}
+18 -1
View File
@@ -22,6 +22,7 @@ use crate::agent::submission::Submission;
use crate::channels::IncomingMessage;
use crate::channels::web::server::GatewayState;
use crate::channels::web::types::{WsClientMessage, WsServerMessage};
use crate::llm::ImageAttachment;
/// Tracks active WebSocket connections.
pub struct WsConnectionTracker {
@@ -156,12 +157,26 @@ async fn handle_client_message(
direct_tx: &mpsc::Sender<WsServerMessage>,
) {
match msg {
WsClientMessage::Message { content, thread_id } => {
WsClientMessage::Message {
content,
thread_id,
images,
} => {
let mut incoming = IncomingMessage::new("gateway", user_id, &content);
if let Some(ref tid) = thread_id {
incoming = incoming.with_thread(tid);
}
// Convert image data to ImageAttachment
let image_attachments: Vec<ImageAttachment> = images
.into_iter()
.map(|img| ImageAttachment {
media_type: img.media_type,
data: img.data,
})
.collect();
incoming = incoming.with_images(image_attachments);
let tx_guard = state.msg_tx.read().await;
if let Some(ref tx) = *tx_guard {
if tx.send(incoming).await.is_err() {
@@ -349,6 +364,7 @@ mod tests {
WsClientMessage::Message {
content: "hello agent".to_string(),
thread_id: Some("t1".to_string()),
images: vec![],
},
&state,
"user1",
@@ -373,6 +389,7 @@ mod tests {
WsClientMessage::Message {
content: "hello".to_string(),
thread_id: None,
images: vec![],
},
&state,
"user1",
+139
View File
@@ -0,0 +1,139 @@
//! Detection of image generation models across inference providers.
/// Check if a model name indicates image generation capability.
///
/// Detects models like:
/// - FLUX (Black Forest Labs): `flux`, `flux.2`, `flux-pro`, etc.
/// - DALL-E (OpenAI): `dall-e-2`, `dall-e-3`, etc.
/// - Stable Diffusion: `stable-diffusion`, `sdxl`, etc.
/// - Imagen (Google): `imagen`, `imagen-2`, etc.
/// - Other generation models
pub fn is_image_generation_model(model: &str) -> bool {
let model_lower = model.to_lowercase();
// FLUX models
if model_lower.contains("flux") {
return true;
}
// DALL-E models
if model_lower.contains("dall-e") || model_lower.contains("dalle") {
return true;
}
// Stable Diffusion models
if model_lower.contains("stable-diffusion")
|| model_lower.contains("sdxl")
|| model_lower.contains("stability")
{
return true;
}
// Imagen models
if model_lower.contains("imagen") {
return true;
}
// Midjourney (if exposed via API)
if model_lower.contains("midjourney") {
return true;
}
// Replicate FLUX via API
if model_lower.contains("black-forest-labs") || model_lower.contains("lucataco") {
return true;
}
false
}
/// Check if any model in a list is an image generation model.
pub fn has_image_generation_model(models: &[String]) -> bool {
models.iter().any(|m| is_image_generation_model(m))
}
/// Suggest the best image generation model from available models.
///
/// Priority: FLUX > DALL-E > others
pub fn suggest_image_model(models: &[String]) -> Option<String> {
// Prefer FLUX
if let Some(flux) = models.iter().find(|m| m.to_lowercase().contains("flux")) {
return Some(flux.clone());
}
// Then DALL-E
if let Some(dalle) = models
.iter()
.find(|m| m.to_lowercase().contains("dall-e") || m.to_lowercase().contains("dalle"))
{
return Some(dalle.clone());
}
// Then any other image model
models
.iter()
.find(|m| is_image_generation_model(m))
.cloned()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_flux_detection() {
assert!(is_image_generation_model(
"black-forest-labs/FLUX.2-klein-4B"
));
assert!(is_image_generation_model("flux"));
assert!(is_image_generation_model("flux-pro"));
}
#[test]
fn test_dalle_detection() {
assert!(is_image_generation_model("dall-e-3"));
assert!(is_image_generation_model("dall-e-2"));
assert!(is_image_generation_model("dalle-3"));
}
#[test]
fn test_stable_diffusion_detection() {
assert!(is_image_generation_model("stable-diffusion-3"));
assert!(is_image_generation_model("sdxl"));
}
#[test]
fn test_imagen_detection() {
assert!(is_image_generation_model("imagen"));
assert!(is_image_generation_model("imagen-3"));
}
#[test]
fn test_non_image_models() {
assert!(!is_image_generation_model("claude-3-5-sonnet"));
assert!(!is_image_generation_model("gpt-4"));
assert!(!is_image_generation_model("gemini-pro"));
}
#[test]
fn test_suggest_image_model() {
let models = vec![
"gpt-4".to_string(),
"black-forest-labs/FLUX.2-klein-4B".to_string(),
"dall-e-3".to_string(),
];
// Should prefer FLUX
assert_eq!(
suggest_image_model(&models),
Some("black-forest-labs/FLUX.2-klein-4B".to_string())
);
}
#[test]
fn test_suggest_dalle_when_no_flux() {
let models = vec!["gpt-4".to_string(), "dall-e-3".to_string()];
assert_eq!(suggest_image_model(&models), Some("dall-e-3".to_string()));
}
}
+5 -2
View File
@@ -10,6 +10,7 @@
pub mod circuit_breaker;
pub mod costs;
pub mod failover;
pub mod image_models;
mod nearai_chat;
mod provider;
mod reasoning;
@@ -20,13 +21,15 @@ pub mod retry;
mod rig_adapter;
pub mod session;
pub mod smart_routing;
pub mod vision_models;
pub use circuit_breaker::{CircuitBreakerConfig, CircuitBreakerProvider};
pub use failover::{CooldownConfig, FailoverProvider};
pub use nearai_chat::{ModelInfo, NearAiChatProvider};
pub use provider::{
ChatMessage, CompletionRequest, CompletionResponse, FinishReason, LlmProvider, ModelMetadata,
Role, ToolCall, ToolCompletionRequest, ToolCompletionResponse, ToolDefinition, ToolResult,
ChatMessage, CompletionRequest, CompletionResponse, FinishReason, ImageAttachment, LlmProvider,
ModelMetadata, Role, ToolCall, ToolCompletionRequest, ToolCompletionResponse, ToolDefinition,
ToolResult,
};
pub use reasoning::{
ActionPlan, Reasoning, ReasoningContext, RespondOutput, RespondResult, SILENT_REPLY_TOKEN,
+70 -33
View File
@@ -671,7 +671,7 @@ struct ChatCompletionRequest {
struct ChatCompletionMessage {
role: String,
#[serde(skip_serializing_if = "Option::is_none")]
content: Option<String>,
content: Option<serde_json::Value>,
#[serde(skip_serializing_if = "Option::is_none")]
tool_call_id: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
@@ -839,10 +839,15 @@ fn flatten_tool_messages(messages: Vec<ChatCompletionMessage>) -> Vec<ChatComple
if let (true, Some(calls)) = (msg.role == "assistant", &msg.tool_calls) {
// Convert assistant tool_calls into descriptive text
let mut parts: Vec<String> = Vec::new();
if let Some(ref text) = msg.content
&& !text.is_empty()
{
parts.push(text.clone());
if let Some(content) = &msg.content {
// Extract string from JSON value
let text = match content {
serde_json::Value::String(s) => s.as_str(),
_ => "",
};
if !text.is_empty() {
parts.push(text.to_string());
}
}
for tc in calls {
parts.push(format!(
@@ -852,7 +857,7 @@ fn flatten_tool_messages(messages: Vec<ChatCompletionMessage>) -> Vec<ChatComple
}
ChatCompletionMessage {
role: "assistant".to_string(),
content: Some(parts.join("\n")),
content: Some(serde_json::json!(parts.join("\n"))),
tool_call_id: None,
name: None,
@@ -861,10 +866,16 @@ fn flatten_tool_messages(messages: Vec<ChatCompletionMessage>) -> Vec<ChatComple
} else if msg.role == "tool" {
// Convert tool result into a user message
let tool_name = msg.name.as_deref().unwrap_or("unknown");
let result = msg.content.as_deref().unwrap_or("");
let result = match &msg.content {
Some(serde_json::Value::String(s)) => s.as_str(),
_ => "",
};
ChatCompletionMessage {
role: "user".to_string(),
content: Some(format!("[Tool `{}` returned: {}]", tool_name, result)),
content: Some(serde_json::json!(format!(
"[Tool `{}` returned: {}]",
tool_name, result
))),
tool_call_id: None,
name: None,
@@ -902,8 +913,23 @@ impl From<ChatMessage> for ChatCompletionMessage {
let content = if role == "assistant" && tool_calls.is_some() && msg.content.is_empty() {
None
} else if !msg.images.is_empty() && role == "user" {
// User message with images: create a content array with text and image parts
let mut parts = vec![serde_json::json!({
"type": "text",
"text": msg.content
})];
for img in msg.images {
parts.push(serde_json::json!({
"type": "image_url",
"image_url": {
"url": format!("data:{};base64,{}", img.media_type, img.data)
}
}));
}
Some(serde_json::Value::Array(parts))
} else {
Some(msg.content)
Some(serde_json::json!(msg.content))
};
Self {
@@ -1068,7 +1094,7 @@ mod tests {
let msg = ChatMessage::user("Hello");
let chat_msg: ChatCompletionMessage = msg.into();
assert_eq!(chat_msg.role, "user");
assert_eq!(chat_msg.content, Some("Hello".to_string()));
assert_eq!(chat_msg.content, Some(serde_json::json!("Hello")));
}
#[test]
@@ -1142,14 +1168,14 @@ mod tests {
let messages = vec![
ChatCompletionMessage {
role: "system".to_string(),
content: Some("You are helpful.".to_string()),
content: Some(serde_json::json!("You are helpful.")),
tool_call_id: None,
name: None,
tool_calls: None,
},
ChatCompletionMessage {
role: "user".to_string(),
content: Some("Hello".to_string()),
content: Some(serde_json::json!("Hello")),
tool_call_id: None,
name: None,
tool_calls: None,
@@ -1166,7 +1192,7 @@ mod tests {
let messages = vec![
ChatCompletionMessage {
role: "user".to_string(),
content: Some("test".to_string()),
content: Some(serde_json::json!("test")),
tool_call_id: None,
name: None,
tool_calls: None,
@@ -1187,7 +1213,7 @@ mod tests {
},
ChatCompletionMessage {
role: "tool".to_string(),
content: Some("hi".to_string()),
content: Some(serde_json::json!("hi")),
tool_call_id: Some("call_1".to_string()),
name: Some("echo".to_string()),
tool_calls: None,
@@ -1200,24 +1226,28 @@ mod tests {
// Assistant tool_calls → plain assistant text
assert_eq!(result[1].role, "assistant");
assert!(result[1].tool_calls.is_none());
assert!(
result[1]
.content
.as_ref()
.unwrap()
.contains("[Called tool `echo`")
);
if let Some(content) = &result[1].content {
if let serde_json::Value::String(s) = content {
assert!(s.contains("[Called tool `echo`"));
} else {
panic!("Content should be a string");
}
} else {
panic!("Content should be present");
}
// Tool result → user message
assert_eq!(result[2].role, "user");
assert!(result[2].tool_call_id.is_none());
assert!(
result[2]
.content
.as_ref()
.unwrap()
.contains("[Tool `echo` returned: hi]")
);
if let Some(content) = &result[2].content {
if let serde_json::Value::String(s) = content {
assert!(s.contains("[Tool `echo` returned: hi]"));
} else {
panic!("Content should be a string");
}
} else {
panic!("Content should be present");
}
}
#[test]
@@ -1225,7 +1255,7 @@ mod tests {
let messages = vec![
ChatCompletionMessage {
role: "assistant".to_string(),
content: Some("Let me check that.".to_string()),
content: Some(serde_json::json!("Let me check that.")),
tool_call_id: None,
name: None,
tool_calls: Some(vec![ChatCompletionToolCall {
@@ -1239,7 +1269,7 @@ mod tests {
},
ChatCompletionMessage {
role: "tool".to_string(),
content: Some("found it".to_string()),
content: Some(serde_json::json!("found it")),
tool_call_id: Some("call_1".to_string()),
name: Some("search".to_string()),
tool_calls: None,
@@ -1247,9 +1277,16 @@ mod tests {
];
let result = flatten_tool_messages(messages);
let text = result[0].content.as_ref().unwrap();
assert!(text.starts_with("Let me check that."));
assert!(text.contains("[Called tool `search`"));
if let Some(content) = result[0].content.as_ref() {
if let serde_json::Value::String(text) = content {
assert!(text.starts_with("Let me check that."));
assert!(text.contains("[Called tool `search`"));
} else {
panic!("Content should be a string");
}
} else {
panic!("Content should be present");
}
}
#[test]
+29
View File
@@ -16,6 +16,15 @@ pub enum Role {
Tool,
}
/// An image attachment for user messages.
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ImageAttachment {
/// MIME type (e.g., "image/jpeg", "image/png", "image/gif", "image/webp")
pub media_type: String,
/// Base64-encoded image data (without data URL prefix)
pub data: String,
}
/// A message in a conversation.
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ChatMessage {
@@ -31,6 +40,9 @@ pub struct ChatMessage {
/// to appear on the assistant message preceding tool result messages).
#[serde(skip_serializing_if = "Option::is_none")]
pub tool_calls: Option<Vec<ToolCall>>,
/// Images attached to user messages.
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub images: Vec<ImageAttachment>,
}
impl ChatMessage {
@@ -42,6 +54,7 @@ impl ChatMessage {
tool_call_id: None,
name: None,
tool_calls: None,
images: Vec::new(),
}
}
@@ -53,6 +66,19 @@ impl ChatMessage {
tool_call_id: None,
name: None,
tool_calls: None,
images: Vec::new(),
}
}
/// Create a user message with image attachments.
pub fn user_with_images(content: impl Into<String>, images: Vec<ImageAttachment>) -> Self {
Self {
role: Role::User,
content: content.into(),
tool_call_id: None,
name: None,
tool_calls: None,
images,
}
}
@@ -64,6 +90,7 @@ impl ChatMessage {
tool_call_id: None,
name: None,
tool_calls: None,
images: Vec::new(),
}
}
@@ -82,6 +109,7 @@ impl ChatMessage {
} else {
Some(tool_calls)
},
images: Vec::new(),
}
}
@@ -97,6 +125,7 @@ impl ChatMessage {
tool_call_id: Some(tool_call_id.into()),
name: Some(name.into()),
tool_calls: None,
images: Vec::new(),
}
}
}
+31 -3
View File
@@ -10,8 +10,8 @@ use rig::completion::{
ToolDefinition as RigToolDefinition, Usage as RigUsage,
};
use rig::message::{
Message as RigMessage, ToolChoice as RigToolChoice, ToolFunction, ToolResult as RigToolResult,
ToolResultContent, UserContent,
DocumentSourceKind, Image, ImageMediaType, Message as RigMessage, ToolChoice as RigToolChoice,
ToolFunction, ToolResult as RigToolResult, ToolResultContent, UserContent,
};
use rust_decimal::Decimal;
use serde::Serialize;
@@ -230,7 +230,33 @@ fn convert_messages(messages: &[ChatMessage]) -> (Option<String>, Vec<RigMessage
}
}
crate::llm::Role::User => {
history.push(RigMessage::user(&msg.content));
if msg.images.is_empty() {
history.push(RigMessage::user(&msg.content));
} else {
// User message with images: create multi-part content
let mut parts: Vec<UserContent> = vec![UserContent::text(&msg.content)];
for img in &msg.images {
let media_type = match img.media_type.to_lowercase().as_str() {
"image/jpeg" => ImageMediaType::JPEG,
"image/png" => ImageMediaType::PNG,
"image/gif" => ImageMediaType::GIF,
"image/webp" => ImageMediaType::WEBP,
_ => ImageMediaType::JPEG,
};
parts.push(UserContent::Image(Image {
data: DocumentSourceKind::Base64(img.data.clone()),
media_type: Some(media_type),
detail: None,
additional_params: Default::default(),
}));
}
if let Ok(many) = OneOrMany::many(parts) {
history.push(RigMessage::User { content: many });
} else {
// Fallback to text only
history.push(RigMessage::user(&msg.content));
}
}
}
crate::llm::Role::Assistant => {
if let Some(ref tool_calls) = msg.tool_calls {
@@ -635,6 +661,7 @@ mod tests {
tool_call_id: None,
name: Some("search".to_string()),
tool_calls: None,
images: vec![],
}];
let (_preamble, history) = convert_messages(&messages);
match &history[0] {
@@ -784,6 +811,7 @@ mod tests {
tool_call_id: None,
name: Some("search".to_string()),
tool_calls: None,
images: vec![],
};
let messages = vec![assistant_msg, tool_result_msg];
let (_preamble, history) = convert_messages(&messages);
+160
View File
@@ -0,0 +1,160 @@
//! Detection of vision-capable models across inference providers.
/// Check if a model name indicates vision capability.
///
/// Detects models like:
/// - Claude (Anthropic): `claude-opus`, `claude-sonnet`, etc.
/// - GPT (OpenAI): `gpt-4-vision`, `gpt-4-turbo`, `gpt-4o`, etc.
/// - Gemini (Google): `gemini-pro-vision`, `gemini-2.0-flash`, etc.
/// - Llama (Meta): `llama-2-vision`, etc.
/// - Other vision-capable models
pub fn is_vision_model(model: &str) -> bool {
let model_lower = model.to_lowercase();
// Claude models (Anthropic)
if model_lower.contains("claude") {
return true;
}
// GPT-4 models with vision support
if (model_lower.contains("gpt-4")
|| model_lower.contains("gpt-4o")
|| model_lower.contains("gpt-4-turbo")
|| model_lower.contains("gpt-4-vision"))
&& !model_lower.contains("gpt-4-mini")
{
return true;
}
// Gemini models
if model_lower.contains("gemini") {
return true;
}
// Llava and other vision models
if model_lower.contains("llava")
|| model_lower.contains("vision")
|| model_lower.contains("multimodal")
{
return true;
}
false
}
/// Check if any model in a list is a vision-capable model.
pub fn has_vision_model(models: &[String]) -> bool {
models.iter().any(|m| is_vision_model(m))
}
/// Suggest the best vision model from available models.
///
/// Priority: Claude > GPT-4 > Gemini > others
pub fn suggest_vision_model(models: &[String]) -> Option<String> {
// Prefer Claude
if let Some(claude) = models.iter().find(|m| m.to_lowercase().contains("claude")) {
return Some(claude.clone());
}
// Then GPT-4
if let Some(gpt4) = models
.iter()
.find(|m| m.to_lowercase().contains("gpt-4") && !m.to_lowercase().contains("gpt-4-mini"))
{
return Some(gpt4.clone());
}
// Then Gemini
if let Some(gemini) = models.iter().find(|m| m.to_lowercase().contains("gemini")) {
return Some(gemini.clone());
}
// Then any other vision model
models.iter().find(|m| is_vision_model(m)).cloned()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_claude_detection() {
assert!(is_vision_model("claude-opus-4-20250514"));
assert!(is_vision_model("claude-sonnet-4-20250514"));
assert!(is_vision_model("claude-haiku-3-5-sonnet"));
}
#[test]
fn test_gpt4_detection() {
assert!(is_vision_model("gpt-4-turbo"));
assert!(is_vision_model("gpt-4o"));
assert!(is_vision_model("gpt-4-vision"));
assert!(is_vision_model("gpt-4-32k"));
}
#[test]
fn test_gpt4_mini_not_vision() {
assert!(!is_vision_model("gpt-4-mini"));
}
#[test]
fn test_gemini_detection() {
assert!(is_vision_model("gemini-pro-vision"));
assert!(is_vision_model("gemini-2.0-flash"));
assert!(is_vision_model("gemini-1.5-pro"));
}
#[test]
fn test_llava_detection() {
assert!(is_vision_model("llava-1.6"));
assert!(is_vision_model("llava-v1-7b"));
}
#[test]
fn test_multimodal_detection() {
assert!(is_vision_model("my-multimodal-model"));
assert!(is_vision_model("custom-vision-model"));
}
#[test]
fn test_non_vision_models() {
assert!(!is_vision_model("text-davinci-3"));
assert!(!is_vision_model("llama-2-7b"));
assert!(!is_vision_model("mistral-7b"));
}
#[test]
fn test_suggest_vision_model() {
let models = vec![
"gpt-4-turbo".to_string(),
"claude-opus-4-20250514".to_string(),
"gemini-2.0-flash".to_string(),
];
// Should prefer Claude
assert_eq!(
suggest_vision_model(&models),
Some("claude-opus-4-20250514".to_string())
);
}
#[test]
fn test_suggest_gpt4_when_no_claude() {
let models = vec!["gpt-4-turbo".to_string(), "gemini-2.0-flash".to_string()];
assert_eq!(
suggest_vision_model(&models),
Some("gpt-4-turbo".to_string())
);
}
#[test]
fn test_suggest_gemini_when_no_claude_or_gpt4() {
let models = vec!["gemini-2.0-flash".to_string(), "text-davinci-3".to_string()];
assert_eq!(
suggest_vision_model(&models),
Some("gemini-2.0-flash".to_string())
);
}
}
+236
View File
@@ -0,0 +1,236 @@
//! Image analysis tool for vision-capable LLMs.
//!
//! Reads images from the workspace and prepares them for vision analysis.
//! The LLM can then analyze the image content based on the user's query.
use std::sync::Arc;
use async_trait::async_trait;
use serde_json::{Value, json};
use crate::context::JobContext;
use crate::tools::tool::{ApprovalRequirement, Tool, ToolError, ToolOutput};
use crate::workspace::Workspace;
/// Tool for analyzing images using a vision-capable LLM.
pub struct ImageAnalyzeTool {
workspace: Arc<Workspace>,
}
impl ImageAnalyzeTool {
/// Create a new image analysis tool.
pub fn new(workspace: Arc<Workspace>) -> Self {
Self { workspace }
}
/// Infer media type from file extension.
fn infer_media_type(path: &str) -> &'static str {
let lower_path = path.to_lowercase();
if lower_path.ends_with(".png") || lower_path.ends_with(".b64") {
"image/png"
} else if lower_path.ends_with(".jpg") || lower_path.ends_with(".jpeg") {
"image/jpeg"
} else if lower_path.ends_with(".gif") {
"image/gif"
} else if lower_path.ends_with(".webp") {
"image/webp"
} else {
"image/png" // Default to PNG
}
}
}
#[async_trait]
impl Tool for ImageAnalyzeTool {
fn name(&self) -> &str {
"image_analyze"
}
fn description(&self) -> &str {
"Analyze an image using the LLM's vision capabilities. Provide the workspace path to the image and a question or prompt about what you want to know about the image."
}
fn parameters_schema(&self) -> Value {
json!({
"type": "object",
"properties": {
"path": {
"type": "string",
"description": "Workspace path to the image (e.g., 'images/generated/abc123.b64')"
},
"query": {
"type": "string",
"description": "What do you want to know about the image? (e.g., 'describe the objects in this image', 'is there text in this image?')"
}
},
"required": ["path", "query"]
})
}
async fn execute(&self, params: Value, _ctx: &JobContext) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
// Parse parameters
let path = params
.get("path")
.and_then(|v| v.as_str())
.ok_or_else(|| {
ToolError::InvalidParameters("Missing or invalid 'path' parameter".to_string())
})?
.to_string();
let query = params
.get("query")
.and_then(|v| v.as_str())
.ok_or_else(|| {
ToolError::InvalidParameters("Missing or invalid 'query' parameter".to_string())
})?
.to_string();
if query.is_empty() {
return Err(ToolError::InvalidParameters(
"Query cannot be empty".to_string(),
));
}
// Read image from workspace
let doc = self.workspace.read(&path).await.map_err(|e| {
ToolError::ExecutionFailed(format!("Failed to read image from workspace: {}", e))
})?;
// Infer media type from path
let media_type = Self::infer_media_type(&path).to_string();
// Return the image data and query so the agent can include the image in its vision analysis
Ok(ToolOutput::success(
json!({
"type": "image_analysis_ready",
"path": path,
"query": query,
"data": doc.content,
"media_type": media_type,
"instruction": format!("The user wants you to analyze this image with the following query: {}", query)
}),
start.elapsed(),
))
}
fn requires_approval(&self, _params: &Value) -> ApprovalRequirement {
// Image analysis is read-only, no approval needed
ApprovalRequirement::Never
}
fn sensitive_params(&self) -> &[&str] {
&[]
}
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[test]
fn test_infer_media_type_png() {
assert_eq!(
ImageAnalyzeTool::infer_media_type("images/test.png"),
"image/png"
);
assert_eq!(
ImageAnalyzeTool::infer_media_type("images/test.b64"),
"image/png"
);
}
#[test]
fn test_infer_media_type_jpeg() {
assert_eq!(
ImageAnalyzeTool::infer_media_type("images/test.jpg"),
"image/jpeg"
);
assert_eq!(
ImageAnalyzeTool::infer_media_type("images/test.jpeg"),
"image/jpeg"
);
}
#[test]
fn test_infer_media_type_gif() {
assert_eq!(
ImageAnalyzeTool::infer_media_type("images/test.gif"),
"image/gif"
);
}
#[test]
fn test_infer_media_type_webp() {
assert_eq!(
ImageAnalyzeTool::infer_media_type("images/test.webp"),
"image/webp"
);
}
#[test]
fn test_infer_media_type_default() {
assert_eq!(
ImageAnalyzeTool::infer_media_type("images/test.unknown"),
"image/png"
);
}
#[test]
fn test_parameters_schema_required_fields() {
let schema = json!({
"type": "object",
"properties": {
"path": {
"type": "string",
"description": "Workspace path to the image (e.g., 'images/generated/abc123.b64')"
},
"query": {
"type": "string",
"description": "What do you want to know about the image? (e.g., 'describe the objects in this image', 'is there text in this image?')"
}
},
"required": ["path", "query"]
});
assert_eq!(schema["type"], "object");
assert!(schema["properties"]["path"].is_object());
assert!(schema["properties"]["query"].is_object());
assert_eq!(schema["required"], json!(["path", "query"]));
}
#[test]
fn test_infer_media_type_uppercase_extension_defaults() {
// Uppercase extensions are now case-insensitively matched
assert_eq!(
ImageAnalyzeTool::infer_media_type("images/test.PNG"),
"image/png"
);
assert_eq!(
ImageAnalyzeTool::infer_media_type("images/test.JPG"),
"image/jpeg"
);
}
#[test]
fn test_infer_media_type_nested_path() {
assert_eq!(
ImageAnalyzeTool::infer_media_type("images/generated/2024-03-06/deep/nested/image.png"),
"image/png"
);
}
#[test]
fn test_infer_media_type_multiple_dots() {
assert_eq!(
ImageAnalyzeTool::infer_media_type("images/my.test.image.png"),
"image/png"
);
assert_eq!(
ImageAnalyzeTool::infer_media_type("images/file.backup.jpg"),
"image/jpeg"
);
}
}
+231
View File
@@ -0,0 +1,231 @@
//! Image editing tool using NEAR AI cloud-api (FLUX model).
use std::sync::Arc;
use std::time::Duration;
use async_trait::async_trait;
use base64::Engine;
use secrecy::ExposeSecret;
use serde_json::{Value, json};
use uuid::Uuid;
use crate::config::NearAiConfig;
use crate::context::JobContext;
use crate::tools::tool::{ApprovalRequirement, Tool, ToolError, ToolOutput, ToolRateLimitConfig};
use crate::workspace::Workspace;
/// Tool for editing existing images using NEAR AI cloud-api (FLUX).
pub struct ImageEditTool {
config: NearAiConfig,
client: reqwest::Client,
workspace: Arc<Workspace>,
}
impl ImageEditTool {
/// Create a new image editing tool.
pub fn new(config: NearAiConfig, workspace: Arc<Workspace>) -> Self {
Self {
config,
client: reqwest::Client::new(),
workspace,
}
}
}
#[async_trait]
impl Tool for ImageEditTool {
fn name(&self) -> &str {
"image_edit"
}
fn description(&self) -> &str {
"Edit an existing image using NEAR AI cloud-api (FLUX) by providing the workspace path and a description of changes. \
Returns the edited image saved to the workspace."
}
fn parameters_schema(&self) -> Value {
json!({
"type": "object",
"properties": {
"path": {
"type": "string",
"description": "Workspace path to the source image (e.g., 'images/generated/abc123.b64')"
},
"prompt": {
"type": "string",
"description": "Description of the edits to apply (max 4000 characters)"
},
"size": {
"type": "string",
"enum": ["1024x1024", "1792x1024", "1024x1792"],
"description": "Image dimensions. Default: 1024x1024"
}
},
"required": ["path", "prompt"]
})
}
async fn execute(&self, params: Value, _ctx: &JobContext) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
// Parse parameters
let path = params
.get("path")
.and_then(|v| v.as_str())
.ok_or_else(|| {
ToolError::InvalidParameters("Missing or invalid 'path' parameter".to_string())
})?
.to_string();
let prompt = params
.get("prompt")
.and_then(|v| v.as_str())
.ok_or_else(|| {
ToolError::InvalidParameters("Missing or invalid 'prompt' parameter".to_string())
})?
.to_string();
if prompt.is_empty() {
return Err(ToolError::InvalidParameters(
"Prompt cannot be empty".to_string(),
));
}
if prompt.len() > 4000 {
return Err(ToolError::InvalidParameters(format!(
"Prompt exceeds 4000 character limit (got {})",
prompt.len()
)));
}
let size = params
.get("size")
.and_then(|v| v.as_str())
.unwrap_or("1024x1024");
// Read base64 image data from workspace
let doc = self.workspace.read(&path).await.map_err(|e| {
ToolError::ExecutionFailed(format!("Failed to read image from workspace: {}", e))
})?;
// Decode base64 to bytes
let image_bytes = base64::engine::general_purpose::STANDARD
.decode(&doc.content)
.map_err(|e| {
ToolError::ExecutionFailed(format!("Failed to decode base64 image data: {}", e))
})?;
// Build multipart form
let form = reqwest::multipart::Form::new()
.text("model", "black-forest-labs/FLUX.2-klein-4B")
.part(
"image",
reqwest::multipart::Part::bytes(image_bytes).file_name("image.png"),
)
.text("prompt", prompt.clone())
.text("n", "1")
.text("size", size.to_string())
.text("response_format", "b64_json");
// Call NEAR AI cloud-api edit endpoint
let endpoint = format!(
"{}/v1/images/edits",
self.config.base_url.trim_end_matches('/')
);
let auth_header = if let Some(api_key) = &self.config.api_key {
format!("Bearer {}", api_key.expose_secret())
} else {
"Bearer ".to_string()
};
let response = self
.client
.post(&endpoint)
.header("Authorization", auth_header)
.multipart(form)
.timeout(Duration::from_secs(120))
.send()
.await
.map_err(|e| ToolError::ExternalService(format!("NEAR AI image edit failed: {}", e)))?;
if !response.status().is_success() {
let status = response.status();
let error_text = response
.text()
.await
.unwrap_or_else(|_| "Unknown error".to_string());
return Err(ToolError::ExternalService(format!(
"NEAR AI image edit error ({}): {}",
status, error_text
)));
}
let response_json: Value = response.json().await.map_err(|e| {
ToolError::ExternalService(format!("Failed to parse NEAR AI response: {}", e))
})?;
// Extract base64 edited image data
let edited_base64 = response_json
.get("data")
.and_then(|d| d.get(0))
.and_then(|item| item.get("b64_json"))
.and_then(|v| v.as_str())
.ok_or_else(|| {
ToolError::ExternalService(
"Invalid NEAR AI response structure: missing base64 data".to_string(),
)
})?
.to_string();
// Generate unique filename for edited image
let edit_id = Uuid::new_v4().to_string();
let filename = format!("images/generated/{}_edit.png", edit_id);
// Store edited image to workspace
let edit_path = format!("images/generated/{}_edit.b64", edit_id);
self.workspace
.write(&edit_path, &edited_base64)
.await
.map_err(|e| {
ToolError::ExecutionFailed(format!(
"Failed to save edited image to workspace: {}",
e
))
})?;
// Return sentinel JSON for agent_loop to detect and emit SSE event
Ok(ToolOutput::success(
json!({
"type": "image_generated",
"path": edit_path,
"data": edited_base64,
"media_type": "image/png",
"prompt": prompt,
"size": size,
"filename": filename,
"source_path": path
}),
start.elapsed(),
))
}
fn requires_approval(&self, _params: &Value) -> ApprovalRequirement {
// Image editing is read-only on external state
ApprovalRequirement::Never
}
fn rate_limit_config(&self) -> Option<ToolRateLimitConfig> {
// DALL-E is expensive; rate limit aggressively
Some(ToolRateLimitConfig::new(6, 30))
}
fn sensitive_params(&self) -> &[&str] {
&[]
}
fn execution_timeout(&self) -> std::time::Duration {
// Image editing can take 2+ minutes on the NEAR AI cloud-api
std::time::Duration::from_secs(180)
}
}
+203
View File
@@ -0,0 +1,203 @@
//! Image generation tool using NEAR AI cloud-api (FLUX model).
use std::sync::Arc;
use std::time::Duration;
use async_trait::async_trait;
use secrecy::ExposeSecret;
use serde_json::{Value, json};
use uuid::Uuid;
use crate::config::NearAiConfig;
use crate::context::JobContext;
use crate::tools::tool::{ApprovalRequirement, Tool, ToolError, ToolOutput, ToolRateLimitConfig};
use crate::workspace::Workspace;
/// Tool for generating images from text prompts using NEAR AI cloud-api (FLUX).
pub struct ImageGenerateTool {
config: NearAiConfig,
client: reqwest::Client,
workspace: Arc<Workspace>,
}
impl ImageGenerateTool {
/// Create a new image generation tool.
pub fn new(config: NearAiConfig, workspace: Arc<Workspace>) -> Self {
Self {
config,
client: reqwest::Client::new(),
workspace,
}
}
}
#[async_trait]
impl Tool for ImageGenerateTool {
fn name(&self) -> &str {
"image_generate"
}
fn description(&self) -> &str {
"Generate an image from a text prompt using NEAR AI cloud-api (FLUX.2-klein-4B). \
Returns the generated image saved to the workspace."
}
fn parameters_schema(&self) -> Value {
json!({
"type": "object",
"properties": {
"prompt": {
"type": "string",
"description": "Detailed text description of the image to generate (max 4000 characters)"
},
"size": {
"type": "string",
"enum": ["1024x1024", "1792x1024", "1024x1792"],
"description": "Image dimensions. Default: 1024x1024"
}
},
"required": ["prompt"]
})
}
async fn execute(&self, params: Value, _ctx: &JobContext) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
// Parse parameters
let prompt = params
.get("prompt")
.and_then(|v| v.as_str())
.ok_or_else(|| {
ToolError::InvalidParameters("Missing or invalid 'prompt' parameter".to_string())
})?
.to_string();
if prompt.is_empty() {
return Err(ToolError::InvalidParameters(
"Prompt cannot be empty".to_string(),
));
}
if prompt.len() > 4000 {
return Err(ToolError::InvalidParameters(format!(
"Prompt exceeds 4000 character limit (got {})",
prompt.len()
)));
}
let size = params
.get("size")
.and_then(|v| v.as_str())
.unwrap_or("1024x1024");
// Call NEAR AI cloud-api for image generation (FLUX model)
let request_body = json!({
"model": "black-forest-labs/FLUX.2-klein-4B",
"prompt": prompt,
"n": 1,
"size": size,
"response_format": "b64_json"
});
let endpoint = format!(
"{}/v1/images/generations",
self.config.base_url.trim_end_matches('/')
);
let auth_header = if let Some(api_key) = &self.config.api_key {
format!("Bearer {}", api_key.expose_secret())
} else {
// Fallback: use default NEAR AI cloud-api without explicit key
// (expects auth via environment or other mechanism)
"Bearer ".to_string()
};
let response = self
.client
.post(&endpoint)
.header("Authorization", auth_header)
.json(&request_body)
.timeout(Duration::from_secs(120))
.send()
.await
.map_err(|e| {
ToolError::ExternalService(format!("NEAR AI image generation failed: {}", e))
})?;
if !response.status().is_success() {
let status = response.status();
let error_text = response
.text()
.await
.unwrap_or_else(|_| "Unknown error".to_string());
return Err(ToolError::ExternalService(format!(
"NEAR AI image generation error ({}): {}",
status, error_text
)));
}
let response_json: Value = response.json().await.map_err(|e| {
ToolError::ExternalService(format!("Failed to parse NEAR AI response: {}", e))
})?;
// Extract base64 image data
let base64_data = response_json
.get("data")
.and_then(|d| d.get(0))
.and_then(|item| item.get("b64_json"))
.and_then(|v| v.as_str())
.ok_or_else(|| {
ToolError::ExternalService(
"Invalid NEAR AI response structure: missing base64 data".to_string(),
)
})?
.to_string();
// Generate unique filename
let image_id = Uuid::new_v4().to_string();
let filename = format!("images/generated/{}.png", image_id);
// Store the image file (with extension) containing the base64 data
let image_path = format!("images/generated/{}.b64", image_id);
self.workspace
.write(&image_path, &base64_data)
.await
.map_err(|e| {
ToolError::ExecutionFailed(format!("Failed to save image to workspace: {}", e))
})?;
// Return sentinel JSON for agent_loop to detect and emit SSE event
Ok(ToolOutput::success(
json!({
"type": "image_generated",
"path": image_path,
"data": base64_data,
"media_type": "image/png",
"prompt": prompt,
"size": size,
"filename": filename
}),
start.elapsed(),
))
}
fn requires_approval(&self, _params: &Value) -> ApprovalRequirement {
// Image generation from a prompt is read-only on external state
// so no approval needed
ApprovalRequirement::Never
}
fn rate_limit_config(&self) -> Option<ToolRateLimitConfig> {
// DALL-E is expensive; rate limit aggressively
Some(ToolRateLimitConfig::new(6, 30))
}
fn sensitive_params(&self) -> &[&str] {
&[]
}
fn execution_timeout(&self) -> std::time::Duration {
// Image generation can take 2+ minutes on the NEAR AI cloud-api
std::time::Duration::from_secs(180)
}
}
+6
View File
@@ -4,6 +4,9 @@ mod echo;
pub mod extension_tools;
mod file;
mod http;
mod image_analyze;
mod image_edit;
mod image_gen;
mod job;
mod json;
mod memory;
@@ -23,6 +26,9 @@ pub use extension_tools::{
};
pub use file::{ApplyPatchTool, ListDirTool, ReadFileTool, WriteFileTool};
pub use http::HttpTool;
pub use image_analyze::ImageAnalyzeTool;
pub use image_edit::ImageEditTool;
pub use image_gen::ImageGenerateTool;
pub use job::{
CancelJobTool, CreateJobTool, JobEventsTool, JobPromptTool, JobStatusTool, ListJobsTool,
PromptQueue, SchedulerSlot,
+36 -5
View File
@@ -17,11 +17,11 @@ use crate::skills::registry::SkillRegistry;
use crate::tools::builder::{BuildSoftwareTool, BuilderConfig, LlmSoftwareBuilder};
use crate::tools::builtin::{
ApplyPatchTool, CancelJobTool, CreateJobTool, EchoTool, ExtensionInfoTool, HttpTool,
JobEventsTool, JobPromptTool, JobStatusTool, JsonTool, ListDirTool, ListJobsTool,
MemoryReadTool, MemorySearchTool, MemoryTreeTool, MemoryWriteTool, PromptQueue, ReadFileTool,
ShellTool, SkillInstallTool, SkillListTool, SkillRemoveTool, SkillSearchTool, TimeTool,
ToolActivateTool, ToolAuthTool, ToolInstallTool, ToolListTool, ToolRemoveTool, ToolSearchTool,
WriteFileTool,
ImageEditTool, ImageGenerateTool, JobEventsTool, JobPromptTool, JobStatusTool, JsonTool,
ListDirTool, ListJobsTool, MemoryReadTool, MemorySearchTool, MemoryTreeTool, MemoryWriteTool,
PromptQueue, ReadFileTool, ShellTool, SkillInstallTool, SkillListTool, SkillRemoveTool,
SkillSearchTool, TimeTool, ToolActivateTool, ToolAuthTool, ToolInstallTool, ToolListTool,
ToolRemoveTool, ToolSearchTool, WriteFileTool,
};
use crate::tools::rate_limiter::RateLimiter;
use crate::tools::tool::{Tool, ToolDomain};
@@ -71,6 +71,9 @@ const PROTECTED_TOOL_NAMES: &[&str] = &[
"message",
"web_fetch",
"restart",
"image_generate",
"image_edit",
"image_analyze",
];
/// Registry of available tools.
@@ -302,6 +305,34 @@ impl ToolRegistry {
tracing::info!("Registered 4 memory tools");
}
/// Register image generation tools with NEAR AI config and workspace.
///
/// Image tools require NEAR AI cloud-api access and workspace for storing generated images.
pub fn register_image_tools(
&self,
config: crate::config::NearAiConfig,
workspace: Arc<Workspace>,
) {
self.register_sync(Arc::new(ImageGenerateTool::new(
config.clone(),
Arc::clone(&workspace),
)));
self.register_sync(Arc::new(ImageEditTool::new(config, workspace)));
tracing::info!("Registered 2 image tools (NEAR AI FLUX)");
}
/// Register image analysis tool with workspace access.
///
/// Vision tool allows analyzing images using the LLM's vision capabilities.
pub fn register_vision_tools(&self, workspace: Arc<Workspace>) {
self.register_sync(Arc::new(crate::tools::builtin::ImageAnalyzeTool::new(
workspace,
)));
tracing::info!("Registered 1 vision tool (image analysis)");
}
/// Register job management tools.
///
/// Job tools allow the LLM to create, list, check status, and cancel jobs.
+3
View File
@@ -203,6 +203,7 @@ mod tests {
thread_id: None,
received_at: Utc::now(),
metadata: serde_json::json!({}),
images: vec![],
};
let fired = engine.check_event_triggers(&matching_msg).await;
assert!(
@@ -223,6 +224,7 @@ mod tests {
thread_id: None,
received_at: Utc::now(),
metadata: serde_json::json!({}),
images: vec![],
};
let fired_neg = engine.check_event_triggers(&non_matching_msg).await;
assert_eq!(fired_neg, 0, "Expected 0 routines fired on non-match");
@@ -286,6 +288,7 @@ mod tests {
thread_id: None,
received_at: Utc::now(),
metadata: serde_json::json!({}),
images: vec![],
};
let fired1 = engine.check_event_triggers(&msg).await;
assert!(fired1 >= 1, "First fire should work");